diff --git a/internal/agent/protocol.go b/internal/agent/protocol.go deleted file mode 100644 index fe61fe7..0000000 --- a/internal/agent/protocol.go +++ /dev/null @@ -1,127 +0,0 @@ -package agent - -import ( - "encoding/binary" - "encoding/json" - "errors" - "fmt" - "io" -) - -type MessageType string - -const ( - MessageTypeHello MessageType = "HELLO" - MessageTypePause MessageType = "PAUSE" - MessageTypeResume MessageType = "RESUME" - MessageTypeState MessageType = "STATE" - MessageTypeExtractResult MessageType = "EXTRACT_RESULT" - MessageTypeGoodbye MessageType = "GOODBYE" -) - -const MaxFrameSize = 16 * 1024 * 1024 - -// ProtocolVersion is the wire-format version. Bump on any breaking change -// to the message schema or framing. Independent of the SDK release version. -const ProtocolVersion = 1 - -type Message struct { - Type MessageType `json:"type"` - ID uint64 `json:"id,omitempty"` - - ProtocolVersion int `json:"protocol_version,omitempty"` - Version string `json:"version,omitempty"` - Platform string `json:"platform,omitempty"` - AppPackage string `json:"app_package,omitempty"` - - Snapshots map[string]json.RawMessage `json:"snapshots,omitempty"` - Exceptions []Exception `json:"exceptions,omitempty"` - - Extractor string `json:"extractor,omitempty"` - Result json.RawMessage `json:"result,omitempty"` - Error string `json:"error,omitempty"` - - Reason string `json:"reason,omitempty"` -} - -// Exception mirrors an uncaught throwable captured by the SDK. -type Exception struct { - Class string `json:"class"` - Message string `json:"message,omitempty"` - StackTrace string `json:"stack_trace,omitempty"` - UnixMillis int64 `json:"unix_millis,omitempty"` -} - -func Hello(version, platform, appPackage string) Message { - return Message{ - Type: MessageTypeHello, - ProtocolVersion: ProtocolVersion, - Version: version, - Platform: platform, - AppPackage: appPackage, - } -} - -func Pause(id uint64) Message { return Message{Type: MessageTypePause, ID: id} } - -func Resume(id uint64) Message { return Message{Type: MessageTypeResume, ID: id} } - -func State(id uint64, snapshots map[string]json.RawMessage) Message { - return Message{Type: MessageTypeState, ID: id, Snapshots: snapshots} -} - -func ExtractResult(id uint64, extractor string, result json.RawMessage, extractorError string) Message { - return Message{ - Type: MessageTypeExtractResult, - ID: id, - Extractor: extractor, - Result: result, - Error: extractorError, - } -} - -func Goodbye(reason string) Message { - return Message{Type: MessageTypeGoodbye, Reason: reason} -} - -func WriteMessage(writer io.Writer, message Message) error { - payload, err := json.Marshal(message) - if err != nil { - return fmt.Errorf("marshal: %w", err) - } - if len(payload) > MaxFrameSize { - return fmt.Errorf("frame of %d bytes exceeds maximum %d", len(payload), MaxFrameSize) - } - var header [4]byte - binary.BigEndian.PutUint32(header[:], uint32(len(payload))) - if _, err := writer.Write(header[:]); err != nil { - return fmt.Errorf("write header: %w", err) - } - if _, err := writer.Write(payload); err != nil { - return fmt.Errorf("write payload: %w", err) - } - return nil -} - -func ReadMessage(reader io.Reader) (Message, error) { - var header [4]byte - if _, err := io.ReadFull(reader, header[:]); err != nil { - return Message{}, err - } - length := binary.BigEndian.Uint32(header[:]) - if length > MaxFrameSize { - return Message{}, fmt.Errorf("frame of %d bytes exceeds maximum %d", length, MaxFrameSize) - } - payload := make([]byte, length) - if _, err := io.ReadFull(reader, payload); err != nil { - return Message{}, fmt.Errorf("read payload: %w", err) - } - var message Message - if err := json.Unmarshal(payload, &message); err != nil { - return Message{}, fmt.Errorf("unmarshal: %w", err) - } - if message.Type == "" { - return Message{}, errors.New("missing type") - } - return message, nil -} diff --git a/internal/agent/protocol_test.go b/internal/agent/protocol_test.go deleted file mode 100644 index 19509ad..0000000 --- a/internal/agent/protocol_test.go +++ /dev/null @@ -1,153 +0,0 @@ -package agent - -import ( - "bytes" - "encoding/binary" - "encoding/json" - "errors" - "io" - "strings" - "testing" -) - -func roundTrip(t *testing.T, message Message) Message { - t.Helper() - var buffer bytes.Buffer - if err := WriteMessage(&buffer, message); err != nil { - t.Fatalf("WriteMessage: %v", err) - } - got, err := ReadMessage(&buffer) - if err != nil { - t.Fatalf("ReadMessage: %v", err) - } - return got -} - -func TestRoundTrip_Hello(t *testing.T) { - got := roundTrip(t, Hello("0.0.1", "android", "in.okcredit.merchant")) - if got.Type != MessageTypeHello || got.Version != "0.0.1" || got.Platform != "android" || got.AppPackage != "in.okcredit.merchant" { - t.Fatalf("hello round-trip failed: %+v", got) - } - if got.ProtocolVersion != ProtocolVersion { - t.Errorf("protocol_version: got %d, want %d", got.ProtocolVersion, ProtocolVersion) - } -} - -func TestRoundTrip_PauseResume(t *testing.T) { - for _, builder := range []func(uint64) Message{Pause, Resume} { - got := roundTrip(t, builder(42)) - if got.ID != 42 { - t.Errorf("id round-trip failed: %+v", got) - } - } -} - -func TestRoundTrip_State(t *testing.T) { - snapshots := map[string]json.RawMessage{ - "screen": json.RawMessage(`"customer_ledger"`), - "ledger.balance": json.RawMessage(`1500`), - "is_signed_in": json.RawMessage(`true`), - } - got := roundTrip(t, State(7, snapshots)) - if got.Type != MessageTypeState || got.ID != 7 { - t.Fatalf("state envelope wrong: %+v", got) - } - if string(got.Snapshots["screen"]) != `"customer_ledger"` { - t.Errorf("screen snapshot wrong: %s", got.Snapshots["screen"]) - } - if string(got.Snapshots["ledger.balance"]) != `1500` { - t.Errorf("balance snapshot wrong: %s", got.Snapshots["ledger.balance"]) - } -} - -func TestRoundTrip_ExtractResult(t *testing.T) { - got := roundTrip(t, ExtractResult(1, "ledger.balance", json.RawMessage(`2500`), "")) - if got.Extractor != "ledger.balance" || string(got.Result) != `2500` { - t.Fatalf("extract result round-trip failed: %+v", got) - } - - failed := roundTrip(t, ExtractResult(2, "ledger.balance", nil, "no active customer")) - if failed.Error != "no active customer" { - t.Errorf("extract error round-trip failed: %+v", failed) - } -} - -func TestRoundTrip_Goodbye(t *testing.T) { - got := roundTrip(t, Goodbye("app terminated")) - if got.Type != MessageTypeGoodbye || got.Reason != "app terminated" { - t.Fatalf("goodbye round-trip failed: %+v", got) - } -} - -func TestWriteMessage_FrameFormat(t *testing.T) { - var buffer bytes.Buffer - if err := WriteMessage(&buffer, Pause(99)); err != nil { - t.Fatal(err) - } - raw := buffer.Bytes() - if len(raw) < 4 { - t.Fatalf("frame too short: %d bytes", len(raw)) - } - length := binary.BigEndian.Uint32(raw[:4]) - if int(length) != len(raw)-4 { - t.Errorf("header length %d mismatches payload length %d", length, len(raw)-4) - } - if !strings.Contains(string(raw[4:]), `"type":"PAUSE"`) { - t.Errorf("payload does not contain PAUSE type: %s", raw[4:]) - } -} - -func TestReadMessage_ShortReaderReturnsEOF(t *testing.T) { - _, err := ReadMessage(bytes.NewReader(nil)) - if !errors.Is(err, io.EOF) { - t.Errorf("expected EOF on empty reader, got %v", err) - } -} - -func TestReadMessage_OversizedFrameRejected(t *testing.T) { - var header [4]byte - binary.BigEndian.PutUint32(header[:], uint32(MaxFrameSize+1)) - _, err := ReadMessage(bytes.NewReader(header[:])) - if err == nil || !strings.Contains(err.Error(), "exceeds maximum") { - t.Errorf("expected oversized-frame error, got %v", err) - } -} - -func TestReadMessage_MissingTypeRejected(t *testing.T) { - var buffer bytes.Buffer - payload := []byte(`{"id":1}`) - var header [4]byte - binary.BigEndian.PutUint32(header[:], uint32(len(payload))) - buffer.Write(header[:]) - buffer.Write(payload) - - _, err := ReadMessage(&buffer) - if err == nil || !strings.Contains(err.Error(), "missing type") { - t.Errorf("expected missing-type error, got %v", err) - } -} - -func TestWriteMessage_StreamsMultipleFrames(t *testing.T) { - var buffer bytes.Buffer - messages := []Message{ - Hello("v", "android", "com.x"), - Pause(1), - State(1, map[string]json.RawMessage{"x": json.RawMessage(`42`)}), - Resume(1), - Goodbye("done"), - } - for _, message := range messages { - if err := WriteMessage(&buffer, message); err != nil { - t.Fatal(err) - } - } - for index, want := range messages { - got, err := ReadMessage(&buffer) - if err != nil { - t.Fatalf("frame %d: %v", index, err) - } - if got.Type != want.Type { - t.Errorf("frame %d: got type %q, want %q", index, got.Type, want.Type) - } - } -} diff --git a/internal/agent/server.go b/internal/agent/server.go deleted file mode 100644 index 93ea20c..0000000 --- a/internal/agent/server.go +++ /dev/null @@ -1,145 +0,0 @@ -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) - } - if hello.ProtocolVersion != ProtocolVersion { - rawConn.Close() - return nil, fmt.Errorf("protocol version mismatch: host=%d sdk=%d", ProtocolVersion, hello.ProtocolVersion) - } - 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) - } - done := make(chan struct{}) - exited := make(chan struct{}) - go func() { - defer close(exited) - select { - case <-ctx.Done(): - _ = conn.SetReadDeadline(time.Unix(1, 0)) - case <-done: - } - }() - message, err := ReadMessage(conn) - close(done) - <-exited - _ = conn.SetReadDeadline(time.Time{}) - 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 deleted file mode 100644 index b71aa70..0000000 --- a/internal/agent/server_test.go +++ /dev/null @@ -1,347 +0,0 @@ -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_AcceptRejectsProtocolVersionMismatch(t *testing.T) { - server := newLoopbackServer(t) - - go func() { - client, err := net.Dial("tcp", server.Addr().String()) - if err != nil { - return - } - defer client.Close() - mismatched := Hello("0.0.1", "android", "com.x") - mismatched.ProtocolVersion = ProtocolVersion + 99 - _ = WriteMessage(client, mismatched) - }() - - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) - defer cancel() - _, err := server.Accept(ctx) - if err == nil || !strings.Contains(err.Error(), "protocol version mismatch") { - t.Fatalf("expected protocol-version-mismatch error, got %v", err) - } -} - -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 }() - - 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") - } -} - -// TestConn_SnapshotAfterAcceptContextCancel guards against a race in -// readWithDeadline where the watcher goroutine from Accept could clobber the -// conn's read deadline with a past time after Accept returned, causing the -// next read on the same conn (Snapshot) to time out instantly. -func TestConn_SnapshotAfterAcceptContextCancel(t *testing.T) { - for iteration := range 50 { - server := newLoopbackServer(t) - - clientDone := make(chan struct{}) - go func() { - defer close(clientDone) - client, err := net.Dial("tcp", server.Addr().String()) - if err != nil { - return - } - defer client.Close() - if err := WriteMessage(client, Hello("0.0.1", "android", "com.x")); err != nil { - return - } - msg, err := ReadMessage(client) - if err != nil { - return - } - _ = WriteMessage(client, State(msg.ID, map[string]json.RawMessage{"ok": json.RawMessage(`true`)})) - }() - - acceptCtx, acceptCancel := context.WithTimeout(context.Background(), time.Second) - conn, err := server.Accept(acceptCtx) - acceptCancel() - if err != nil { - t.Fatalf("iteration %d: Accept: %v", iteration, err) - } - - snapCtx, snapCancel := context.WithTimeout(context.Background(), 2*time.Second) - state, err := conn.Snapshot(snapCtx) - snapCancel() - if err != nil { - t.Fatalf("iteration %d: Snapshot: %v", iteration, err) - } - if string(state.Snapshots["ok"]) != `true` { - t.Errorf("iteration %d: unexpected snapshots: %v", iteration, state.Snapshots) - } - conn.Close() - <-clientDone - } -} - -func TestConn_SnapshotTimesOutIfSDKSilent(t *testing.T) { - server := newLoopbackServer(t) - - done := make(chan struct{}) - t.Cleanup(func() { close(done) }) - 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; stay alive until the test ends. - <-done - }() - - 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") - } -}