mirror of
https://github.com/priyanshujain/sanderling.git
synced 2026-10-02 19:17:10 +00:00
chore: delete internal/agent package
This commit is contained in:
1 parent
6c32fb0e1d
commit
b88535cace
4 files changed
-772
No files matched your search
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user