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.
This commit is contained in:
pj committed 2026-04-17 22:50:54 +07:00
1 parent 2cc7617590
commit b9511c0718
2 files changed
+415

No files matched your search

+138
View File
@@ -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")
+277
View File
@@ -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")
}
}