From b9511c071888eb2353b4589648c2b4ebca281a4d Mon Sep 17 00:00:00 2001 From: PJ Date: Fri, 17 Apr 2026 22:50:54 +0700 Subject: [PATCH] feat(agent): socket server with PAUSE/STATE/RESUME flow Accept waits for an SDK HELLO then hands back a Conn. Conn.Snapshot sends a PAUSE, blocks on the matching STATE (id-correlated), and leaves the SDK paused until Release sends RESUME. Conn.Close sends GOODBYE best-effort. v0.1 supports one client at a time; transport is left to the caller so tests can use TCP loopback while production wires via adb reverse to localabstract:uatu-agent. --- internal/agent/server.go | 138 +++++++++++++++++ internal/agent/server_test.go | 277 ++++++++++++++++++++++++++++++++++ 2 files changed, 415 insertions(+) create mode 100644 internal/agent/server.go create mode 100644 internal/agent/server_test.go diff --git a/internal/agent/server.go b/internal/agent/server.go new file mode 100644 index 0000000..64c9d97 --- /dev/null +++ b/internal/agent/server.go @@ -0,0 +1,138 @@ +package agent + +import ( + "context" + "errors" + "fmt" + "net" + "time" +) + +type Server struct { + listener net.Listener +} + +func NewServer(listener net.Listener) *Server { + return &Server{listener: listener} +} + +func (s *Server) Addr() net.Addr { return s.listener.Addr() } + +// Accept waits for the next SDK client and performs the HELLO handshake. +// Only one Conn may be active at a time; subsequent Accepts block until the +// current connection closes. +func (s *Server) Accept(ctx context.Context) (*Conn, error) { + cancelCloser := closeListenerOnCancel(ctx, s.listener) + defer cancelCloser() + + rawConn, err := s.listener.Accept() + if err != nil { + if ctx.Err() != nil { + return nil, ctx.Err() + } + return nil, fmt.Errorf("accept: %w", err) + } + hello, err := readWithDeadline(ctx, rawConn) + if err != nil { + rawConn.Close() + return nil, fmt.Errorf("read hello: %w", err) + } + if hello.Type != MessageTypeHello { + rawConn.Close() + return nil, fmt.Errorf("expected HELLO, got %q", hello.Type) + } + return &Conn{rawConn: rawConn, hello: hello}, nil +} + +func (s *Server) Close() error { return s.listener.Close() } + +type Conn struct { + rawConn net.Conn + hello Message + nextID uint64 +} + +func (c *Conn) Hello() Message { return c.hello } + +func (c *Conn) RemoteAddr() net.Addr { return c.rawConn.RemoteAddr() } + +// Snapshot sends PAUSE with a fresh id and blocks until the SDK returns the +// matching STATE. The SDK's main thread stays paused until Release is called. +func (c *Conn) Snapshot(ctx context.Context) (Message, error) { + c.nextID++ + id := c.nextID + + if err := writeWithDeadline(ctx, c.rawConn, Pause(id)); err != nil { + return Message{}, fmt.Errorf("send pause: %w", err) + } + message, err := readWithDeadline(ctx, c.rawConn) + if err != nil { + return Message{}, fmt.Errorf("read state: %w", err) + } + if message.Type != MessageTypeState { + return Message{}, fmt.Errorf("expected STATE, got %q", message.Type) + } + if message.ID != id { + return Message{}, fmt.Errorf("state id mismatch: sent %d, got %d", id, message.ID) + } + return message, nil +} + +// Release sends RESUME, freeing the SDK's paused main thread. +func (c *Conn) Release(ctx context.Context) error { + return writeWithDeadline(ctx, c.rawConn, Resume(c.nextID)) +} + +// Close sends GOODBYE (best effort) and closes the underlying connection. +func (c *Conn) Close() error { + _ = writeWithDeadline(context.Background(), c.rawConn, Goodbye("shutdown")) + return c.rawConn.Close() +} + +func readWithDeadline(ctx context.Context, conn net.Conn) (Message, error) { + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetReadDeadline(deadline) + defer conn.SetReadDeadline(time.Time{}) + } + done := make(chan struct{}) + defer close(done) + go func() { + select { + case <-ctx.Done(): + _ = conn.SetReadDeadline(time.Unix(1, 0)) + case <-done: + } + }() + message, err := ReadMessage(conn) + if err != nil && ctx.Err() != nil { + return Message{}, ctx.Err() + } + return message, err +} + +func writeWithDeadline(ctx context.Context, conn net.Conn, message Message) error { + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetWriteDeadline(deadline) + defer conn.SetWriteDeadline(time.Time{}) + } + err := WriteMessage(conn, message) + if err != nil && ctx.Err() != nil { + return ctx.Err() + } + return err +} + +func closeListenerOnCancel(ctx context.Context, listener net.Listener) (cancel func()) { + done := make(chan struct{}) + go func() { + select { + case <-ctx.Done(): + _ = listener.Close() + case <-done: + } + }() + return func() { close(done) } +} + +// ErrClosed is returned when a Conn method is called after Close. +var ErrClosed = errors.New("agent: connection closed") diff --git a/internal/agent/server_test.go b/internal/agent/server_test.go new file mode 100644 index 0000000..7f825f6 --- /dev/null +++ b/internal/agent/server_test.go @@ -0,0 +1,277 @@ +package agent + +import ( + "context" + "encoding/json" + "net" + "strings" + "sync" + "testing" + "time" +) + +// fakeSDK drives the client side of an agent connection the way the real SDK +// would: HELLO on connect, then respond to PAUSE with STATE, honor RESUME, +// and close on GOODBYE. +type fakeSDK struct { + conn net.Conn + snapshotFunc func(id uint64) map[string]json.RawMessage +} + +func (f *fakeSDK) sendHello(version, platform, appPackage string) error { + return WriteMessage(f.conn, Hello(version, platform, appPackage)) +} + +func (f *fakeSDK) serveOne() error { + message, err := ReadMessage(f.conn) + if err != nil { + return err + } + switch message.Type { + case MessageTypePause: + snapshots := f.snapshotFunc(message.ID) + return WriteMessage(f.conn, State(message.ID, snapshots)) + case MessageTypeResume: + return nil + case MessageTypeGoodbye: + return nil + default: + return nil + } +} + +func newLoopbackServer(t *testing.T) *Server { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { listener.Close() }) + return NewServer(listener) +} + +func TestServer_AcceptHandshake(t *testing.T) { + server := newLoopbackServer(t) + + connectErr := make(chan error, 1) + go func() { + client, err := net.Dial("tcp", server.Addr().String()) + if err != nil { + connectErr <- err + return + } + sdk := &fakeSDK{conn: client} + connectErr <- sdk.sendHello("0.0.1", "android", "com.example") + }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + conn, err := server.Accept(ctx) + if err != nil { + t.Fatalf("Accept: %v", err) + } + defer conn.Close() + + if got := conn.Hello(); got.Type != MessageTypeHello || got.Version != "0.0.1" || got.AppPackage != "com.example" { + t.Errorf("unexpected hello: %+v", got) + } + if err := <-connectErr; err != nil { + t.Fatalf("client side: %v", err) + } +} + +func TestServer_SnapshotAndRelease(t *testing.T) { + server := newLoopbackServer(t) + + var wg sync.WaitGroup + wg.Go(func() { + client, err := net.Dial("tcp", server.Addr().String()) + if err != nil { + t.Errorf("dial: %v", err) + return + } + sdk := &fakeSDK{ + conn: client, + snapshotFunc: func(id uint64) map[string]json.RawMessage { + return map[string]json.RawMessage{ + "screen": json.RawMessage(`"home"`), + "ledger.balance": json.RawMessage(`1500`), + } + }, + } + if err := sdk.sendHello("0.0.1", "android", "com.x"); err != nil { + t.Errorf("hello: %v", err) + return + } + for range 2 { + if err := sdk.serveOne(); err != nil { + t.Errorf("pause: %v", err) + return + } + if err := sdk.serveOne(); err != nil { + t.Errorf("resume: %v", err) + return + } + } + }) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + conn, err := server.Accept(ctx) + if err != nil { + t.Fatalf("Accept: %v", err) + } + defer conn.Close() + + for expected := uint64(1); expected <= 2; expected++ { + state, err := conn.Snapshot(ctx) + if err != nil { + t.Fatalf("Snapshot #%d: %v", expected, err) + } + if state.ID != expected { + t.Errorf("snapshot #%d: id=%d", expected, state.ID) + } + if string(state.Snapshots["screen"]) != `"home"` { + t.Errorf("snapshot #%d: screen=%s", expected, state.Snapshots["screen"]) + } + if err := conn.Release(ctx); err != nil { + t.Fatalf("Release #%d: %v", expected, err) + } + } + wg.Wait() +} + +func TestServer_AcceptRequiresHello(t *testing.T) { + server := newLoopbackServer(t) + + go func() { + client, err := net.Dial("tcp", server.Addr().String()) + if err != nil { + return + } + defer client.Close() + // Send a PAUSE instead of HELLO — server should reject. + _ = WriteMessage(client, Pause(1)) + }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + _, err := server.Accept(ctx) + if err == nil || !strings.Contains(err.Error(), "expected HELLO") { + t.Fatalf("expected HELLO-required error, got %v", err) + } +} + +func TestServer_AcceptCancelsOnContext(t *testing.T) { + server := newLoopbackServer(t) + + ctx, cancel := context.WithCancel(context.Background()) + acceptErr := make(chan error, 1) + go func() { _, err := server.Accept(ctx); acceptErr <- err }() + + time.Sleep(50 * time.Millisecond) + cancel() + + select { + case err := <-acceptErr: + if err == nil { + t.Errorf("expected error after cancel, got nil") + } + case <-time.After(2 * time.Second): + t.Errorf("accept did not return after cancel") + } +} + +func TestConn_SnapshotRejectsIDMismatch(t *testing.T) { + server := newLoopbackServer(t) + + go func() { + client, _ := net.Dial("tcp", server.Addr().String()) + defer client.Close() + _ = WriteMessage(client, Hello("0.0.1", "android", "com.x")) + // Read the PAUSE but respond with a wrong id. + msg, _ := ReadMessage(client) + _ = WriteMessage(client, State(msg.ID+99, map[string]json.RawMessage{})) + }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + conn, err := server.Accept(ctx) + if err != nil { + t.Fatalf("Accept: %v", err) + } + defer conn.Close() + + _, err = conn.Snapshot(ctx) + if err == nil || !strings.Contains(err.Error(), "id mismatch") { + t.Errorf("expected id-mismatch error, got %v", err) + } +} + +func TestConn_CloseSendsGoodbye(t *testing.T) { + server := newLoopbackServer(t) + + received := make(chan Message, 1) + go func() { + client, _ := net.Dial("tcp", server.Addr().String()) + defer client.Close() + _ = WriteMessage(client, Hello("0.0.1", "android", "com.x")) + // Drain until GOODBYE. + for { + msg, err := ReadMessage(client) + if err != nil { + return + } + if msg.Type == MessageTypeGoodbye { + received <- msg + return + } + } + }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + conn, err := server.Accept(ctx) + if err != nil { + t.Fatal(err) + } + if err := conn.Close(); err != nil { + t.Fatal(err) + } + + select { + case msg := <-received: + if msg.Reason != "shutdown" { + t.Errorf("expected reason=shutdown, got %q", msg.Reason) + } + case <-time.After(time.Second): + t.Error("client did not receive GOODBYE") + } +} + +func TestConn_SnapshotTimesOutIfSDKSilent(t *testing.T) { + server := newLoopbackServer(t) + + go func() { + client, _ := net.Dial("tcp", server.Addr().String()) + defer client.Close() + _ = WriteMessage(client, Hello("0.0.1", "android", "com.x")) + // Never respond to PAUSE. + time.Sleep(2 * time.Second) + }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + conn, err := server.Accept(ctx) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + fastCtx, fastCancel := context.WithTimeout(ctx, 200*time.Millisecond) + defer fastCancel() + _, err = conn.Snapshot(fastCtx) + if err == nil { + t.Errorf("expected timeout error, got nil") + } +}