diff --git a/docs/sdk.md b/docs/sdk.md index 5abb0a32..caeb5960 100644 --- a/docs/sdk.md +++ b/docs/sdk.md @@ -156,6 +156,16 @@ TLS handshakes share the existing dial timeout/cancellation budget. The agent connection then uses HTTP/1.1 `Upgrade: silkd` and remains a bidirectional stream; the guest protocol and port-forwarding frames do not change. +Data-plane calls share a handle's relay connection: after a call the SDK +keeps the connection for 30 seconds (`WithKeepAlive` tunes the window; 0 +dials per call) and the next call on that handle sends its request on it, so +a busy handle pays the dial, upgrade and TLS handshake once. A kept +connection counts as live for `idle_hibernate_seconds` until it closes, so +keep the window below that setting; `Close` and `Hibernate` drop it at once. +Long-lived streams (`Watch`, `OpenPty`, `DialPort`, an LSP session) take a +connection of their own, and a guest whose silkd predates the back-to-back +protocol gets one connection per call as before. + An explicit scheme in an owner, redirect, peer, or `Attach` address wins; a bare address inherits the entry client's scheme. Trust settings are shared across these connections. The SDK does not translate private addresses to diff --git a/e2e/cmd/rpcbench/main.go b/e2e/cmd/rpcbench/main.go index 05e95985..d155ff0b 100644 --- a/e2e/cmd/rpcbench/main.go +++ b/e2e/cmd/rpcbench/main.go @@ -1,8 +1,9 @@ // rpcbench measures the one-connection-per-RPC overhead on a live node and -// what a pre-dialed spare connection would buy (H-4's decision data): mode A -// dials+upgrades per RPC like the SDK does today; mode B keeps one dialed -// connection ahead, hiding the handshake behind the previous call. Run by -// hand against a claimed sandbox: +// what the alternatives buy: mode A dials+upgrades per RPC by hand; mode C +// is the SDK's own path, one kept connection serving RPCs back to back; mode +// B keeps one hand-dialed connection ahead, hiding the handshake behind the +// previous call (H-4's decision data). Run by hand against a claimed +// sandbox: // // rpcbench -addr -token -template -n 200 package main @@ -70,7 +71,20 @@ func run(addr, token, template string, n int) error { } a = append(a, time.Since(start)) } - report("A dial-per-RPC (today)", a) + report("A dial-per-RPC", a) + + if _, err := sb.Stat(ctx, "/"); err != nil { + return err + } + c := make([]time.Duration, 0, n) + for range n { + start := time.Now() + if _, err := sb.Stat(ctx, "/"); err != nil { + return err + } + c = append(c, time.Since(start)) + } + report("C SDK keep-alive", c) spare := make(chan net.Conn, 1) errs := make(chan error, 1) @@ -104,11 +118,11 @@ func run(addr, token, template string, n int) error { b = append(b, time.Since(start)) } cancel() - report("B pre-dialed spare ", b) + report("B pre-dialed spare", b) return nil } -// statRPC: the protocol is one RPC per connection. +// statRPC drives one hand-dialed connection through one RPC and closes it. func statRPC(conn net.Conn) error { defer func() { _ = conn.Close() }() sc := silkd.NewConn(conn) @@ -172,7 +186,7 @@ func dialAgent(ctx context.Context, addr, id, token string) (net.Conn, error) { func report(label string, samples []time.Duration) { slices.Sort(samples) pct := func(p float64) time.Duration { return samples[int(p*float64(len(samples)-1))] } - fmt.Printf("%s n=%d p50=%.2fms p90=%.2fms p99=%.2fms\n", + fmt.Printf("%-22s n=%d p50=%.2fms p90=%.2fms p99=%.2fms\n", label, len(samples), ms(pct(0.50)), ms(pct(0.90)), ms(pct(0.99))) } diff --git a/mcp/server_test.go b/mcp/server_test.go index 67521c89..fa80a61b 100644 --- a/mcp/server_test.go +++ b/mcp/server_test.go @@ -11,8 +11,19 @@ import ( "net/http/httptest" "strings" "testing" + + "github.com/cocoonstack/sandbox/protocol/wire" + "github.com/cocoonstack/sandbox/sdk/go/silkd/silkdtest" ) +var infoFrame = func() string { + frame, err := wire.EncodeResponse(&wire.InfoResp{Version: "test", Proto: wire.KeepAliveProto}) + if err != nil { + panic(err) + } + return string(frame) + "\n" +}() + func TestServeSpeaksMCP(t *testing.T) { srv := newTestServer(t, func(w http.ResponseWriter, r *http.Request) bool { if !strings.HasSuffix(r.URL.Path, "/checkpoint") { @@ -66,20 +77,9 @@ func TestServeSpeaksMCP(t *testing.T) { } func TestExecKeepsOutputWhenTheGuestDrops(t *testing.T) { - srv := newTestServer(t, func(w http.ResponseWriter, r *http.Request) bool { - if !strings.HasSuffix(r.URL.Path, "/agent") { - return false - } - conn, _, err := http.NewResponseController(w).Hijack() - if err != nil { - t.Errorf("hijack: %v", err) - return true - } - _, _ = io.WriteString(conn, "HTTP/1.1 101 Switching Protocols\r\nUpgrade: silkd\r\nConnection: Upgrade\r\n\r\n"+ - `{"type":"started","pid":7}`+"\n"+`{"type":"stdout","data":"cGFydGlhbA=="}`+"\n") - _ = conn.Close() - return true - }) + srv := newTestServer(t, agentRoute(t, func(string) (string, bool) { + return `{"type":"started","pid":7}` + "\n" + `{"type":"stdout","data":"cGFydGlhbA=="}` + "\n", true + })) replies := serveLines(t, srv, `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"create_sandbox","arguments":{}}}`, `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"exec","arguments":{"sandbox_id":"sb_1","command":"yes"}}}`, @@ -93,36 +93,12 @@ func TestExecKeepsOutputWhenTheGuestDrops(t *testing.T) { func TestReadFileStopsAtTheCap(t *testing.T) { chunk := `{"type":"data","data":"` + base64.StdEncoding.EncodeToString(make([]byte, 256<<10)) + `"}` + "\n" - srv := newTestServer(t, func(w http.ResponseWriter, r *http.Request) bool { - if !strings.HasSuffix(r.URL.Path, "/agent") { - return false - } - conn, _, err := http.NewResponseController(w).Hijack() - if err != nil { - t.Errorf("hijack: %v", err) - return true - } - defer conn.Close() - br := bufio.NewReader(conn) - if _, err = io.WriteString(conn, "HTTP/1.1 101 Switching Protocols\r\nUpgrade: silkd\r\nConnection: Upgrade\r\n\r\n"); err != nil { - return true - } - req, err := br.ReadString('\n') - if err != nil { - return true - } + srv := newTestServer(t, agentRoute(t, func(req string) (string, bool) { if strings.Contains(req, `"op":"fs_stat"`) { - _, _ = io.WriteString(conn, `{"type":"stat","info":{"kind":"file","size":0}}`+"\n") - return true + return `{"type":"stat","info":{"kind":"file","size":0}}` + "\n", false } - for range 8 { - if _, err := io.WriteString(conn, chunk); err != nil { - return true - } - } - _, _ = io.WriteString(conn, `{"type":"done"}`+"\n") - return true - }) + return strings.Repeat(chunk, 8) + `{"type":"done"}` + "\n", false + })) replies := serveLines(t, srv, `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"create_sandbox","arguments":{}}}`, `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"read_file","arguments":{"sandbox_id":"sb_1","path":"/dev/zero"}}}`, @@ -184,6 +160,35 @@ func newTestServer(t *testing.T, extra func(http.ResponseWriter, *http.Request) return srv } +func agentRoute(t *testing.T, reply func(req string) (frames string, drop bool)) func(http.ResponseWriter, *http.Request) bool { + t.Helper() + return func(w http.ResponseWriter, r *http.Request) bool { + if !strings.HasSuffix(r.URL.Path, "/agent") { + return false + } + conn, err := silkdtest.Upgrade(w) + if err != nil { + t.Errorf("upgrade: %v", err) + return true + } + defer conn.Close() + br := bufio.NewReader(conn) + for { + req, err := br.ReadString('\n') + if err != nil { + return true + } + frames, drop := infoFrame, false + if !strings.Contains(req, `"op":"info"`) { + frames, drop = reply(req) + } + if _, err := io.WriteString(conn, frames); err != nil || drop { + return true + } + } + } +} + func serveLines(t *testing.T, srv *server, lines ...string) map[int]map[string]any { t.Helper() var out bytes.Buffer diff --git a/protocol/wire/frame.go b/protocol/wire/frame.go index 9a7a2ecc..c3c42702 100644 --- a/protocol/wire/frame.go +++ b/protocol/wire/frame.go @@ -20,6 +20,8 @@ const ( // ProtoVersion is stamped into every request as "v"; silkd ignores // unknown fields, which is the forward-compatibility story. ProtoVersion = 1 + // KeepAliveProto is the InfoResp.Proto from which silkd serves RPCs back to back on one connection. + KeepAliveProto = 2 // MaxFrame mirrors silkd's frame cap. MaxFrame = 8 << 20 // BulkChunk mirrors silkd's BULK_CHUNK, the payload of one bulk data frame. diff --git a/sdk/go/client.go b/sdk/go/client.go index 1be27fb5..61515126 100644 --- a/sdk/go/client.go +++ b/sdk/go/client.go @@ -129,6 +129,7 @@ type Client struct { apiToken string hc *http.Client tlsConfig *tls.Config + keepAlive time.Duration } // Connect returns a client for a sandboxd node. @@ -142,7 +143,7 @@ func Connect(addr string, opts ...ClientOption) (*Client, error) { if err != nil { return nil, err } - c := &Client{addr: strings.TrimPrefix(u.String(), "http://"), scheme: u.Scheme, hc: &http.Client{}} + c := &Client{addr: strings.TrimPrefix(u.String(), "http://"), scheme: u.Scheme, hc: &http.Client{}, keepAlive: keepAliveIdle} for _, opt := range opts { opt(c) } diff --git a/sdk/go/keepalive.go b/sdk/go/keepalive.go new file mode 100644 index 00000000..d5dd1e5f --- /dev/null +++ b/sdk/go/keepalive.go @@ -0,0 +1,175 @@ +package sandbox + +import ( + "context" + "errors" + "slices" + "sync" + "sync/atomic" + "syscall" + "time" + + "github.com/cocoonstack/sandbox/protocol/wire" + "github.com/cocoonstack/sandbox/sdk/go/silkd" +) + +const ( + keepAliveIdle = 30 * time.Second + keepAliveConns = 8 +) + +// WithKeepAlive bounds how long a handle keeps an idle relay connection for +// its next call (default 30s; 0 dials per call). An open connection holds +// the sandbox's idle-hibernate clock, so keep it below the deployment's +// idle_hibernate_seconds. +func WithKeepAlive(idle time.Duration) ClientOption { + return func(c *Client) { c.keepAlive = idle } +} + +// agentPool parks a handle's idle relay connections between calls. +type agentPool struct { + mu sync.Mutex + idle []*agentConn +} + +// take returns a parked connection whose peer is still there, or nil. +func (p *agentPool) take() *agentConn { + for { + p.mu.Lock() + n := len(p.idle) + if n == 0 { + p.mu.Unlock() + return nil + } + c := p.idle[n-1] + p.idle = p.idle[:n-1] + p.mu.Unlock() + if c.timer.Stop() && c.quiet() { + return c + } + _ = c.Close() + } +} + +// park keeps c for the next call and closes it once idle has passed. +func (p *agentPool) park(c *agentConn, idle time.Duration) { + p.mu.Lock() + if len(p.idle) >= keepAliveConns { + p.mu.Unlock() + _ = c.Close() + return + } + if c.timer == nil { + c.timer = time.AfterFunc(idle, func() { p.evict(c) }) + } else { + c.timer.Reset(idle) + } + p.idle = append(p.idle, c) + p.mu.Unlock() +} + +func (p *agentPool) drain() { + p.mu.Lock() + idle := p.idle + p.idle = nil + p.mu.Unlock() + for _, c := range idle { + c.timer.Stop() + _ = c.Close() + } +} + +func (p *agentPool) evict(c *agentConn) { + p.mu.Lock() + if i := slices.Index(p.idle, c); i >= 0 { + p.idle = slices.Delete(p.idle, i, i+1) + } + p.mu.Unlock() + _ = c.Close() +} + +// agentConn is one relayed silkd connection, with the TCP socket's raw handle for the liveness peek. +type agentConn struct { + *silkd.Conn + sock syscall.RawConn + peeked bool + peekFn func(uintptr) bool + timer *time.Timer +} + +func newAgentConn(raw *upgradedConn) *agentConn { + c := &agentConn{Conn: silkd.NewConn(raw)} + if sc, ok := raw.tcp.(syscall.Conn); ok { + c.sock, _ = sc.SyscallConn() + } + c.peekFn = c.peek + return c +} + +// quiet reports whether the parked connection's peer has neither hung up nor spoken. +func (c *agentConn) quiet() bool { + return c.sock != nil && c.sock.Read(c.peekFn) == nil && c.peeked +} + +// sent sends req on c and hands it back, closing it on failure. +func (c *agentConn) sent(req wire.Request) (*agentConn, error) { + if err := c.Send(req); err != nil { + _ = c.Close() + return nil, err + } + return c, nil +} + +// probeProto pipelines info ahead of req and reads the daemon's answer. +func (c *agentConn) probeProto(ctx context.Context, req wire.Request) (uint32, error) { + if err := c.Send(wire.Info{}); err != nil { + return 0, err + } + if err := c.Send(req); err != nil { + return 0, err + } + stop := context.AfterFunc(ctx, func() { _ = c.Close() }) + defer stop() + info, err := expect[wire.InfoResp](ctx, c.Conn) + if err != nil { + return 0, err + } + return max(info.Proto, 1), nil +} + +// lease is one RPC's hold on a relay connection: close drops it, reuse parks it for the next call. +type lease struct { + s *Sandbox + conn *agentConn + stop func() bool + over atomic.Bool +} + +// done ends the RPC: a terminal frame, error frames included, parks the connection; any other failure drops it. +func (l *lease) done(err error) error { + if _, frame := errors.AsType[*wire.ErrorResp](err); err == nil || frame { + l.reuse() + } else { + l.close() + } + return err +} + +func (l *lease) reuse() { + if !l.over.CompareAndSwap(false, true) { + return + } + if !l.stop() || !canProbe || l.s.c.keepAlive <= 0 || l.s.proto.Load() < wire.KeepAliveProto { + _ = l.conn.Close() + return + } + l.s.pool.park(l.conn, l.s.c.keepAlive) +} + +func (l *lease) close() { + if !l.over.CompareAndSwap(false, true) { + return + } + l.stop() + _ = l.conn.Close() +} diff --git a/sdk/go/keepalive_test.go b/sdk/go/keepalive_test.go new file mode 100644 index 00000000..d4d855a6 --- /dev/null +++ b/sdk/go/keepalive_test.go @@ -0,0 +1,216 @@ +package sandbox + +import ( + "bufio" + "io" + "net" + "sync/atomic" + "testing" + "time" + + "github.com/cocoonstack/sandbox/protocol/wire" + "github.com/cocoonstack/sandbox/sdk/go/silkd" + "github.com/cocoonstack/sandbox/sdk/go/silkd/silkdtest" +) + +func TestCallsShareOneConnection(t *testing.T) { + var upgrades atomic.Int32 + fake := silkdtest.NewFake(t.TempDir()) + sb := testSandbox(t, newAgentServer(t, func(c net.Conn) { + upgrades.Add(1) + fake.ServeConn(c) + })) + for range 3 { + if _, err := sb.Stat(t.Context(), "/"); err != nil { + t.Fatalf("stat: %v", err) + } + } + if out, err := sb.Exec(t.Context(), "echo", "42"); err != nil || out != "42\n" { + t.Fatalf("exec: %q, %v", out, err) + } + if err := sb.WriteFile(t.Context(), "/f", []byte("hi"), nil); err != nil { + t.Fatalf("write: %v", err) + } + if _, err := sb.ReadFile(t.Context(), "/f"); err != nil { + t.Fatalf("read: %v", err) + } + if got := upgrades.Load(); got != 1 { + t.Errorf("upgrades = %d, want 1", got) + } +} + +func TestOldDaemonDialsPerCall(t *testing.T) { + var upgrades atomic.Int32 + sb := testSandbox(t, newAgentServer(t, func(c net.Conn) { + upgrades.Add(1) + silkdtest.ServeConnOnce(c) + })) + for range 3 { + if out, err := sb.Exec(t.Context(), "echo", "42"); err != nil || out != "42\n" { + t.Fatalf("exec: %q, %v", out, err) + } + } + if got := upgrades.Load(); got != 4 { + t.Errorf("upgrades = %d, want 4: the proto probe plus one per call", got) + } +} + +func TestOldDaemonConnectionIsNotParked(t *testing.T) { + if !canProbe { + t.Skip("no parked-connection probe on this platform") + } + client, server := net.Pipe() + t.Cleanup(func() { _ = server.Close() }) + sb := &Sandbox{c: &Client{keepAlive: time.Minute}} + t.Cleanup(sb.pool.drain) + sb.proto.Store(1) + l := lease{ + s: sb, + conn: &agentConn{Conn: silkd.NewConn(client)}, + stop: func() bool { return true }, + } + l.reuse() + sb.pool.mu.Lock() + parked := len(sb.pool.idle) + sb.pool.mu.Unlock() + if parked != 0 { + t.Errorf("parked connections = %d, want 0", parked) + } +} + +func TestKeepAliveOffDialsPerCall(t *testing.T) { + var upgrades atomic.Int32 + sb := testSandbox(t, newAgentServer(t, func(c net.Conn) { + upgrades.Add(1) + silkdtest.ServeConn(c) + }), WithKeepAlive(0)) + for range 3 { + if _, err := sb.Exec(t.Context(), "echo", "42"); err != nil { + t.Fatalf("exec: %v", err) + } + } + if got := upgrades.Load(); got != 3 { + t.Errorf("upgrades = %d, want 3", got) + } +} + +func TestIdleConnectionCloses(t *testing.T) { + served := make(chan struct{}, 1) + fake := silkdtest.NewFake(t.TempDir()) + sb := testSandbox(t, newAgentServer(t, func(c net.Conn) { + fake.ServeConn(c) + served <- struct{}{} + }), WithKeepAlive(50*time.Millisecond)) + if _, err := sb.Stat(t.Context(), "/"); err != nil { + t.Fatalf("stat: %v", err) + } + select { + case <-served: + case <-time.After(3 * time.Second): + t.Fatal("idle connection still open after the keep-alive window") + } +} + +func TestStalledStdinKeepsItsConnectionOut(t *testing.T) { + var upgrades atomic.Int32 + sb := testSandbox(t, newAgentServer(t, func(c net.Conn) { + upgrades.Add(1) + silkdtest.ServeConn(c) + })) + pr, pw := io.Pipe() + t.Cleanup(func() { _ = pw.Close() }) + if code, err := sb.Run(t.Context(), Cmd{Argv: []string{"echo", "hi"}, Stdin: pr}); err != nil || code != 0 { + t.Fatalf("run: %d, %v", code, err) + } + if _, err := sb.Exec(t.Context(), "echo", "42"); err != nil { + t.Fatalf("exec: %v", err) + } + if got := upgrades.Load(); got != 2 { + t.Errorf("upgrades = %d, want 2", got) + } +} + +func TestPeerHangUpIsNoticedBeforeReuse(t *testing.T) { + if !canProbe { + t.Skip("no parked-connection probe on this platform") + } + var upgrades atomic.Int32 + gone := make(chan struct{}, 2) + fake := silkdtest.NewFake(t.TempDir()) + sb := testSandbox(t, newAgentServer(t, func(c net.Conn) { + upgrades.Add(1) + fake.ServeConn(&hangUpAfterReply{Conn: c}) + gone <- struct{}{} + })) + sb.proto.Store(wire.KeepAliveProto) + if _, err := sb.Stat(t.Context(), "/"); err != nil { + t.Fatalf("stat: %v", err) + } + <-gone + if _, err := sb.Stat(t.Context(), "/"); err != nil { + t.Fatalf("stat after the peer hung up: %v", err) + } + if got := upgrades.Load(); got != 2 { + t.Errorf("upgrades = %d, want 2", got) + } +} + +func TestPeerQuiet(t *testing.T) { + if !canProbe { + t.Skip("no parked-connection probe on this platform") + } + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer l.Close() + accepted := make(chan net.Conn, 1) + go func() { + if c, acceptErr := l.Accept(); acceptErr == nil { + accepted <- c + } + }() + client, err := net.Dial("tcp", l.Addr().String()) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer client.Close() + server := <-accepted + probe := newAgentConn(&upgradedConn{Conn: client, tcp: client, r: bufio.NewReader(client)}) + if !probe.quiet() { + t.Error("a silent peer reported as gone") + } + if _, err := server.Write([]byte("x")); err != nil { + t.Fatalf("write: %v", err) + } + waitUntil(t, func() bool { return !probe.quiet() }, "unprompted byte not noticed") + var b [1]byte + if n, err := client.Read(b[:]); err != nil || n != 1 || b[0] != 'x' { + t.Errorf("read after the probe = %q, %v; want the peeked byte intact", b[:n], err) + } + _ = server.Close() + waitUntil(t, func() bool { return !probe.quiet() }, "hang-up not noticed") +} + +func waitUntil(t *testing.T, cond func() bool, msg string) { + t.Helper() + for start := time.Now(); time.Since(start) < 3*time.Second; time.Sleep(time.Millisecond) { + if cond() { + return + } + } + t.Fatal(msg) +} + +type hangUpAfterReply struct { + net.Conn + replied atomic.Bool +} + +func (c *hangUpAfterReply) Write(p []byte) (int, error) { + n, err := c.Conn.Write(p) + if c.replied.CompareAndSwap(false, true) { + _ = c.Close() + } + return n, err +} diff --git a/sdk/go/port.go b/sdk/go/port.go index 345cc07d..68e74ad9 100644 --- a/sdk/go/port.go +++ b/sdk/go/port.go @@ -132,16 +132,16 @@ func (s *Sandbox) ProxyPort(ctx context.Context, localAddr string, port uint16) // openStream drives the call → Ready → attach dance shared by every verb // that turns the connection into a raw byte stream (DialPort, Lsp.Request). func (s *Sandbox) openStream(ctx context.Context, req wire.Request) (*PortConn, error) { - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { return nil, err } if _, err = expect[wire.Ready](ctx, conn); err != nil { - done() + l.close() return nil, err } pr, pw := io.Pipe() - p := &PortConn{conn: conn, stop: done, out: pr} + p := &PortConn{conn: conn, stop: l.close, out: pr} go p.drain(ctx, pw) return p, nil } diff --git a/sdk/go/probe_other.go b/sdk/go/probe_other.go new file mode 100644 index 00000000..f4b99d92 --- /dev/null +++ b/sdk/go/probe_other.go @@ -0,0 +1,7 @@ +//go:build !unix + +package sandbox + +const canProbe = false + +func (*agentConn) peek(uintptr) bool { return true } diff --git a/sdk/go/probe_unix.go b/sdk/go/probe_unix.go new file mode 100644 index 00000000..f3892519 --- /dev/null +++ b/sdk/go/probe_unix.go @@ -0,0 +1,15 @@ +//go:build unix + +package sandbox + +import "syscall" + +const canProbe = true + +// peek looks at the socket's next byte without consuming it; peeked means there is none and the peer is still there. +func (c *agentConn) peek(fd uintptr) bool { + var b [1]byte + _, _, err := syscall.Recvfrom(int(fd), b[:], syscall.MSG_PEEK) + c.peeked = err == syscall.EAGAIN + return true +} diff --git a/sdk/go/proc.go b/sdk/go/proc.go index 433ac01f..b9ca12f9 100644 --- a/sdk/go/proc.go +++ b/sdk/go/proc.go @@ -63,10 +63,15 @@ func (s *Sandbox) Attach(ctx context.Context, pid uint32, stdout, stderr io.Writ } func (s *Sandbox) drainProc(ctx context.Context, req wire.Request, stdout, stderr io.Writer) (int32, bool, error) { - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { return 0, false, err } - defer done() - return pumpStdio(ctx, conn, stdout, stderr) + defer l.close() + code, exited, err := pumpStdio(ctx, conn, stdout, stderr) + // logs closes with done after the exit frame + if _, replay := req.(*wire.Logs); replay && exited && err == nil { + err = terminalErr(ctx, conn) + } + return code, exited, l.done(err) } diff --git a/sdk/go/proc_test.go b/sdk/go/proc_test.go index 1a984a86..3a12d2cf 100644 --- a/sdk/go/proc_test.go +++ b/sdk/go/proc_test.go @@ -24,7 +24,7 @@ func TestProcVerbs(t *testing.T) { `{"type":"exit","code":7}`, }, } - sb := testSandbox(t, newAgentServer(t, procServe(script))) + sb := legacySandbox(t, newAgentServer(t, procServe(script))) ctx := t.Context() pid, err := sb.Spawn(ctx, Cmd{Argv: []string{"sleep", "30"}}) @@ -68,7 +68,7 @@ func TestProcVerbsUnknownPid(t *testing.T) { script := map[string][]string{ "logs": {`{"type":"error","kind":"not_found","message":"no such pid"}`}, } - sb := testSandbox(t, newAgentServer(t, procServe(script))) + sb := legacySandbox(t, newAgentServer(t, procServe(script))) _, _, err := sb.Logs(t.Context(), 999, nil, nil) if err == nil || !strings.Contains(err.Error(), "not_found") { t.Fatalf("Logs unknown pid: %v, want typed not_found", err) diff --git a/sdk/go/pty.go b/sdk/go/pty.go index a068b2d4..98e9833a 100644 --- a/sdk/go/pty.go +++ b/sdk/go/pty.go @@ -54,15 +54,7 @@ func (p *Pty) Write(b []byte) (int, error) { // Resize adjusts the terminal window; it is a separate RPC keyed by pid. func (p *Pty) Resize(ctx context.Context, cols, rows uint16) error { - conn, done, err := p.sb.dial(ctx) - if err != nil { - return err - } - defer done() - if err := conn.Send(&wire.PtyResize{PID: p.PID, Cols: cols, Rows: rows}); err != nil { - return err - } - return terminalErr(ctx, conn) + return p.sb.doneRPC(ctx, &wire.PtyResize{PID: p.PID, Cols: cols, Rows: rows}) } // Close ends the pty session (silkd sees the disconnect and kills the shell). @@ -115,18 +107,18 @@ func (p *Pty) drain(ctx context.Context, pw *io.PipeWriter, stop func()) { // OpenPty starts a shell whose lifetime is governed by ctx or Close. func (s *Sandbox) OpenPty(ctx context.Context, opts PtyOpts) (*Pty, error) { req := &wire.PtyOpen{Cols: opts.Cols, Rows: opts.Rows, Cwd: opts.Cwd, Env: opts.Env, User: opts.User} - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { return nil, err } started, err := expect[wire.Started](ctx, conn) if err != nil { - done() + l.close() return nil, err } pr, pw := io.Pipe() - p := &Pty{PID: started.PID, sb: s, conn: conn, stop: done, out: pr} - go p.drain(ctx, pw, done) + p := &Pty{PID: started.PID, sb: s, conn: conn, stop: l.close, out: pr} + go p.drain(ctx, pw, l.close) return p, nil } diff --git a/sdk/go/pty_test.go b/sdk/go/pty_test.go index a5a950e6..66920caf 100644 --- a/sdk/go/pty_test.go +++ b/sdk/go/pty_test.go @@ -67,7 +67,7 @@ func TestPtyReleasesTheRelayWhenTheShellExits(t *testing.T) { _, _ = r.ReadByte() close(released) }) - pty, err := testSandbox(t, ts).OpenPty(t.Context(), PtyOpts{Cols: 80, Rows: 24}) + pty, err := legacySandbox(t, ts).OpenPty(t.Context(), PtyOpts{Cols: 80, Rows: 24}) if err != nil { t.Fatalf("OpenPty: %v", err) } diff --git a/sdk/go/sandbox.go b/sdk/go/sandbox.go index 44aec0e8..11c12ff4 100644 --- a/sdk/go/sandbox.go +++ b/sdk/go/sandbox.go @@ -7,6 +7,7 @@ import ( "io" "net/http" "strings" + "sync/atomic" "time" "github.com/cocoonstack/sandbox/protocol/wire" @@ -68,6 +69,8 @@ type Sandbox struct { c *Client token string owner string // data-plane address (owner node), from the claim + proto atomic.Uint32 + pool agentPool } // Owner returns the data-plane address of the node that owns the sandbox — @@ -102,20 +105,31 @@ func (s *Sandbox) Run(ctx context.Context, cmd Cmd) (int, error) { if len(cmd.Argv) == 0 { return 0, fmt.Errorf("empty argv") } - conn, done, err := s.call(ctx, &wire.Exec{Argv: cmd.Argv, Cwd: cmd.Cwd, Env: cmd.Env, User: cmd.User, Session: cmd.Session}) + conn, l, err := s.call(ctx, &wire.Exec{Argv: cmd.Argv, Cwd: cmd.Cwd, Env: cmd.Env, User: cmd.User, Session: cmd.Session}) if err != nil { return 0, err } - defer done() + defer l.close() + stdinDone := make(chan struct{}) if cmd.Stdin == nil { // a guest that already answered and closed fails this send; the frames it sent still come back below _ = conn.Send(wire.StdinClose{}) + close(stdinDone) } else { - go pumpStdin(conn, cmd.Stdin) + go func() { + defer close(stdinDone) + pumpStdin(conn, cmd.Stdin) + }() } code, exited, err := pumpStdio(ctx, conn, cmd.Stdout, cmd.Stderr) + // a pump still parked in Read would write into the next RPC, so its connection is not kept + select { + case <-stdinDone: + err = l.done(err) + default: + } if err != nil { return 0, err } @@ -175,12 +189,14 @@ func (s *Sandbox) Promote(ctx context.Context, template string) (*Template, erro // sessions, processes, and memory state intact. The TTL keeps running — a // hibernated sandbox is still reaped at its deadline. func (s *Sandbox) Hibernate(ctx context.Context) error { + s.pool.drain() return doNoContent(ctx, s.c, http.MethodPost, s.owner, "/v1/sandboxes/"+s.ID+"/hibernate", nil, s.token, "hibernate") } // Close releases the sandbox on its node; releasing one already gone is not // an error. It takes no ctx so it stays defer-friendly — bounded internally. func (s *Sandbox) Close() error { + s.pool.drain() ctx, cancel := context.WithTimeout(context.Background(), releaseTimeout) defer cancel() resp, err := s.c.roundTrip(ctx, http.MethodPost, s.owner, "/v1/sandboxes/"+s.ID+"/release", nil, s.token) @@ -194,29 +210,49 @@ func (s *Sandbox) Close() error { return apiError("release", resp) } -// dial opens one relayed silkd connection, closed by ctx cancellation or the returned cleanup; one connection carries one RPC. -func (s *Sandbox) dial(ctx context.Context) (*silkd.Conn, func(), error) { - raw, err := s.c.dialAgent(ctx, s.owner, s.ID, s.token) +// call sends req on a relay connection and hands back the lease that ends the RPC; ctx cancellation closes the connection under it. +func (s *Sandbox) call(ctx context.Context, req wire.Request) (*silkd.Conn, *lease, error) { + c, err := s.connect(ctx, req) if err != nil { return nil, nil, err } - conn := silkd.NewConn(raw) - stop := context.AfterFunc(ctx, func() { _ = conn.Close() }) - return conn, func() { stop(); _ = conn.Close() }, nil + return c.Conn, &lease{s: s, conn: c, stop: context.AfterFunc(ctx, func() { _ = c.Close() })}, nil } -// call dials one RPC connection and sends its leading request, cleaning up -// itself on a send failure; the returned cleanup must be deferred. -func (s *Sandbox) call(ctx context.Context, req wire.Request) (*silkd.Conn, func(), error) { - conn, done, err := s.dial(ctx) +// connect sends req on a parked connection or a fresh dial; a handle's first dial pipelines info ahead of req, and a daemon before proto 2 closes after answering it without reading req, hence the second dial. +func (s *Sandbox) connect(ctx context.Context, req wire.Request) (*agentConn, error) { + if c := s.pool.take(); c != nil { + return c.sent(req) + } + c, err := s.dial(ctx) if err != nil { - return nil, nil, err + return nil, err } - if err := conn.Send(req); err != nil { - done() - return nil, nil, err + if s.proto.Load() != 0 { + return c.sent(req) + } + proto, err := c.probeProto(ctx, req) + if err != nil { + _ = c.Close() + return nil, err + } + s.proto.Store(proto) + if proto >= wire.KeepAliveProto { + return c, nil + } + _ = c.Close() + if c, err = s.dial(ctx); err != nil { + return nil, err + } + return c.sent(req) +} + +func (s *Sandbox) dial(ctx context.Context) (*agentConn, error) { + raw, err := s.c.dialAgent(ctx, s.owner, s.ID, s.token) + if err != nil { + return nil, err } - return conn, done, nil + return newAgentConn(raw), nil } // pumpStdin chunks the reader into stdin frames; Send's own locking keeps diff --git a/sdk/go/sandbox_test.go b/sdk/go/sandbox_test.go index 02b2994d..1defa563 100644 --- a/sdk/go/sandbox_test.go +++ b/sdk/go/sandbox_test.go @@ -99,7 +99,7 @@ func TestUpgradeKeepsCoalescedBytes(t *testing.T) { t.Cleanup(ts.Close) var out strings.Builder - code, err := testSandbox(t, ts).Run(t.Context(), Cmd{Argv: []string{"echo", "42"}, Stdout: &out}) + code, err := legacySandbox(t, ts).Run(t.Context(), Cmd{Argv: []string{"echo", "42"}, Stdout: &out}) if err != nil { t.Fatalf("Run: %v", err) } @@ -144,13 +144,9 @@ func newAgentServer(t *testing.T, serve func(net.Conn)) *httptest.Server { w.WriteHeader(http.StatusUnauthorized) return } - conn, _, err := http.NewResponseController(w).Hijack() + conn, err := silkdtest.Upgrade(w) if err != nil { - t.Errorf("hijack: %v", err) - return - } - if _, err := io.WriteString(conn, upgrade101); err != nil { - _ = conn.Close() + t.Errorf("upgrade: %v", err) return } serve(conn) @@ -160,8 +156,17 @@ func newAgentServer(t *testing.T, serve func(net.Conn)) *httptest.Server { return ts } -func testSandbox(t *testing.T, ts *httptest.Server) *Sandbox { +func testSandbox(t *testing.T, ts *httptest.Server, opts ...ClientOption) *Sandbox { t.Helper() - c := testClient(t, ts) - return &Sandbox{ID: "sb_1", c: c, token: "tok", owner: c.addr} + c := testClient(t, ts, opts...) + sb := &Sandbox{ID: "sb_1", c: c, token: "tok", owner: c.addr} + t.Cleanup(sb.pool.drain) + return sb +} + +func legacySandbox(t *testing.T, ts *httptest.Server) *Sandbox { + t.Helper() + sb := testSandbox(t, ts) + sb.proto.Store(1) + return sb } diff --git a/sdk/go/silkd/conn_test.go b/sdk/go/silkd/conn_test.go index d934c814..1a5b7757 100644 --- a/sdk/go/silkd/conn_test.go +++ b/sdk/go/silkd/conn_test.go @@ -2,7 +2,6 @@ package silkd_test import ( "bytes" - "errors" "io" "net" "strings" @@ -47,8 +46,13 @@ func TestConnInfoRoundTrip(t *testing.T) { if !ok || info.Version != "silkdtest" { t.Errorf("got %#v, want silkdtest info", resp) } - if _, err := conn.Recv(); !errors.Is(err, io.EOF) { - t.Errorf("got %v after terminal frame, want EOF", err) + if err := conn.Send(wire.Info{}); err != nil { + t.Fatalf("send on the kept connection: %v", err) + } + if resp, err := conn.Recv(); err != nil { + t.Fatalf("recv on the kept connection: %v", err) + } else if _, ok := resp.(*wire.InfoResp); !ok { + t.Errorf("got %#v on the kept connection, want the next info", resp) } } diff --git a/sdk/go/silkd/silkdtest/fake.go b/sdk/go/silkd/silkdtest/fake.go index b7c194b1..52039e9d 100644 --- a/sdk/go/silkd/silkdtest/fake.go +++ b/sdk/go/silkd/silkdtest/fake.go @@ -35,15 +35,13 @@ func NewFake(root string) *Fake { return &Fake{Root: root, sessions: map[string]bool{}, branches: []string{"main"}, current: "main"} } -// ServeConn speaks one RPC on an already-open connection (e.g. after an HTTP -// hijack) and closes it. +// ServeConn serves RPCs on an already-open connection (e.g. after an HTTP +// hijack) until the peer closes it. func (f *Fake) ServeConn(conn net.Conn) { - defer func() { _ = conn.Close() }() - r := bufio.NewReader(conn) - req, err := recvRequest(r) - if err != nil { - return - } + serveConn(conn, bufio.NewReader(conn), wire.KeepAliveProto, f.serve) +} + +func (f *Fake) serve(conn net.Conn, r *bufio.Reader, req wire.Request) { switch req := req.(type) { case *wire.FsWrite: f.fsWrite(conn, r, req) diff --git a/sdk/go/silkd/silkdtest/silkdtest.go b/sdk/go/silkd/silkdtest/silkdtest.go index df29cf61..9451a976 100644 --- a/sdk/go/silkd/silkdtest/silkdtest.go +++ b/sdk/go/silkd/silkdtest/silkdtest.go @@ -6,19 +6,38 @@ import ( "fmt" "io" "net" + "net/http" "strings" "github.com/cocoonstack/sandbox/protocol/wire" ) -// Serve accepts connections until l closes, speaking one RPC per connection and closing it after the terminal frame, like silkd. +// Serve accepts connections until l closes, serving RPCs back to back on each like silkd. func Serve(l net.Listener) { acceptLoop(l, ServeConn) } -// ServeConn handles one RPC on an open connection, then closes it. +// ServeConn serves RPCs on an open connection until the peer closes it. func ServeConn(conn net.Conn) { - handle(conn, bufio.NewReader(conn)) + serveConn(conn, bufio.NewReader(conn), wire.KeepAliveProto, serveStateless) +} + +// ServeConnOnce serves one RPC and closes, like a daemon from before proto 2. +func ServeConnOnce(conn net.Conn) { + serveConn(conn, bufio.NewReader(conn), 1, serveStateless) +} + +// Upgrade hijacks an agent request and answers the 101, handing back the relay connection. +func Upgrade(w http.ResponseWriter) (net.Conn, error) { + conn, _, err := http.NewResponseController(w).Hijack() + if err != nil { + return nil, err + } + if _, err := io.WriteString(conn, "HTTP/1.1 101 Switching Protocols\r\nUpgrade: silkd\r\nConnection: Upgrade\r\n\r\n"); err != nil { + _ = conn.Close() + return nil, err + } + return conn, nil } // ListenHybrid serves the muxer handshake on a UDS: each connection must open with "CONNECT ", is answered "OK ", then speaks the Serve protocol. @@ -38,17 +57,34 @@ func ListenHybrid(sockPath string, port int) (io.Closer, error) { _ = c.Close() return } - handle(c, r) + serveConn(c, r, wire.KeepAliveProto, serveStateless) }) return l, nil } -func handle(conn net.Conn, r *bufio.Reader) { +// serveConn answers info with proto and dispatches the other requests until the peer closes, or after one when proto predates keep-alive; input frames outside an RPC are dropped as silkd drops them. +func serveConn(conn net.Conn, r *bufio.Reader, proto uint32, serve func(net.Conn, *bufio.Reader, wire.Request)) { defer func() { _ = conn.Close() }() - req, err := recvRequest(r) - if err != nil { - return + for { + req, err := recvRequest(r) + if err != nil { + return + } + switch req.(type) { + case *wire.Stdin, *wire.StdinClose, *wire.Data, *wire.DataEnd: + continue + case *wire.Info: + send(conn, &wire.InfoResp{Version: "silkdtest", Proto: proto}) + default: + serve(conn, r, req) + } + if proto < wire.KeepAliveProto { + return + } } +} + +func serveStateless(conn net.Conn, r *bufio.Reader, req wire.Request) { if !serveCommon(conn, r, req) { errFrame(conn, wire.KindUnimplemented, "silkdtest: "+req.Op()) } @@ -56,8 +92,6 @@ func handle(conn net.Conn, r *bufio.Reader) { func serveCommon(conn net.Conn, r *bufio.Reader, req wire.Request) bool { switch req := req.(type) { - case *wire.Info: - send(conn, &wire.InfoResp{Version: "silkdtest", Proto: wire.ProtoVersion}) case *wire.Exec: serveExec(conn, r, req) case *wire.PortForward: diff --git a/sdk/go/upgrade.go b/sdk/go/upgrade.go index c1f2241d..976f77c8 100644 --- a/sdk/go/upgrade.go +++ b/sdk/go/upgrade.go @@ -10,7 +10,7 @@ import ( "strings" ) -func (c *Client) dialAgent(ctx context.Context, addr, id, token string) (net.Conn, error) { +func (c *Client) dialAgent(ctx context.Context, addr, id, token string) (*upgradedConn, error) { if strings.ContainsAny(id, "\r\n\x00") || strings.ContainsAny(token, "\r\n\x00") { return nil, fmt.Errorf("agent upgrade: id or token contains a control character") } @@ -59,12 +59,13 @@ func (c *Client) dialAgent(ctx context.Context, addr, id, token string) (net.Con _ = raw.Close() return nil, err } - return &upgradedConn{Conn: conn, r: br}, nil + return &upgradedConn{Conn: conn, tcp: raw, r: br}, nil } type upgradedConn struct { net.Conn - r *bufio.Reader + tcp net.Conn // under any TLS, for the parked-connection probe + r *bufio.Reader } func (u *upgradedConn) Read(p []byte) (int, error) { return u.r.Read(p) } diff --git a/sdk/go/utils.go b/sdk/go/utils.go index 84de3509..9227c4df 100644 --- a/sdk/go/utils.go +++ b/sdk/go/utils.go @@ -19,35 +19,57 @@ type respPtr[T any] interface { // doneRPC sends a request that answers with Done or an error frame. func (s *Sandbox) doneRPC(ctx context.Context, req wire.Request) error { - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { return err } - defer done() - return terminalErr(ctx, conn) + defer l.close() + return l.done(terminalErr(ctx, conn)) } -// uploadRPC sends req, streams r as Data frames, and expects a terminal Done. +// uploadRPC streams r as Data frames after req and expects Done; an early terminal frame is the guest rejecting the upload and stops the stream. func (s *Sandbox) uploadRPC(ctx context.Context, req wire.Request, r io.Reader) error { - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { return err } - defer done() - if err := uploadStream(conn, r); err != nil { + defer l.close() + terminal := make(chan error, 1) + go func() { terminal <- terminalErr(ctx, conn) }() + buf := make([]byte, wire.BulkChunk) + for { + select { + case err := <-terminal: + return l.done(err) + default: + } + n, readErr := r.Read(buf) + if n > 0 { + if err := conn.Send(&wire.Data{Data: buf[:n]}); err != nil { + return err + } + } + if readErr == io.EOF { + break + } + if readErr != nil { + return readErr + } + } + if err := conn.Send(wire.DataEnd{}); err != nil { return err } - return terminalErr(ctx, conn) + return l.done(<-terminal) } // downloadRPC sends req and drains its Data stream into sink until Done. func (s *Sandbox) downloadRPC(ctx context.Context, req wire.Request, sink func([]byte) error) error { - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { return err } - defer done() - return drainData(ctx, conn, sink) + defer l.close() + return l.done(drainData(ctx, conn, sink)) } // pumpStdio copies stdout/stderr frames to the writers until the terminal frame: exit carries the code, done means the stream ended without one. @@ -83,12 +105,13 @@ func pumpStdio(ctx context.Context, conn *silkd.Conn, stdout, stderr io.Writer) // oneShotRPC sends req and returns its single typed reply frame. func oneShotRPC[T any, PT respPtr[T]](ctx context.Context, s *Sandbox, req wire.Request) (*T, error) { - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { return nil, err } - defer done() - return expect[T, PT](ctx, conn) + defer l.close() + v, err := expect[T, PT](ctx, conn) + return v, l.done(err) } // collectRPC sends req and gathers every streamed frame of type T until Done. @@ -107,12 +130,12 @@ func collectRPC[T any, PT respPtr[T]](ctx context.Context, s *Sandbox, req wire. func streamRPC[T any, PT respPtr[T]](ctx context.Context, s *Sandbox, req wire.Request) iter.Seq2[T, error] { return func(yield func(T, error) bool) { var zero T - conn, done, err := s.call(ctx, req) + conn, l, err := s.call(ctx, req) if err != nil { yield(zero, err) return } - defer done() + defer l.close() for { resp, err := recv(ctx, conn) if err != nil { @@ -127,7 +150,9 @@ func streamRPC[T any, PT respPtr[T]](ctx context.Context, s *Sandbox, req wire.R } switch r := resp.(type) { case *wire.Done: + l.reuse() case *wire.ErrorResp: + l.reuse() yield(zero, r) default: yield(zero, unexpected(resp)) @@ -137,26 +162,6 @@ func streamRPC[T any, PT respPtr[T]](ctx context.Context, s *Sandbox, req wire.R } } -// uploadStream chunks r into Data frames terminated by DataEnd; shared by the -// FsWrite payload and the FsPush tar stream. -func uploadStream(conn *silkd.Conn, r io.Reader) error { - buf := make([]byte, wire.BulkChunk) - for { - n, readErr := r.Read(buf) - if n > 0 { - if err := conn.Send(&wire.Data{Data: buf[:n]}); err != nil { - return err - } - } - if readErr == io.EOF { - return conn.Send(wire.DataEnd{}) - } - if readErr != nil { - return readErr - } - } -} - // drainData consumes Data frames into sink until Done; an error frame or an // unexpected frame is a Go error. func drainData(ctx context.Context, conn *silkd.Conn, sink func([]byte) error) error { diff --git a/sdk/go/watch.go b/sdk/go/watch.go index a0347325..8d5e1b50 100644 --- a/sdk/go/watch.go +++ b/sdk/go/watch.go @@ -82,17 +82,17 @@ func (w *Watcher) setErr(err error) { // returns once the sandbox acknowledges the watch is armed, so events caused // after Watch returns are guaranteed captured. func (s *Sandbox) Watch(ctx context.Context, path string, recursive bool) (*Watcher, error) { - conn, done, err := s.call(ctx, &wire.FsWatch{Path: path, Recursive: recursive}) + conn, l, err := s.call(ctx, &wire.FsWatch{Path: path, Recursive: recursive}) if err != nil { return nil, err } if _, err = expect[wire.Ready](ctx, conn); err != nil { - done() + l.close() return nil, err } w := &Watcher{ events: make(chan wire.Event, 16), - stop: done, + stop: l.close, closed: make(chan struct{}), } go w.drain(ctx, conn) diff --git a/sdk/go/watch_test.go b/sdk/go/watch_test.go index d43b4e8d..957be6d9 100644 --- a/sdk/go/watch_test.go +++ b/sdk/go/watch_test.go @@ -56,7 +56,7 @@ func TestWatchReleasesTheRelayWhenTheSandboxEnds(t *testing.T) { _, _ = r.ReadByte() close(released) }) - sb := testSandbox(t, ts) + sb := legacySandbox(t, ts) w, err := sb.Watch(t.Context(), "/work", true) if err != nil { t.Fatalf("Watch: %v", err)