From 79e984be43afea560e03a7a7c9c55456a5ed2046 Mon Sep 17 00:00:00 2001 From: PJ Date: Sun, 7 Jun 2026 16:54:33 +0530 Subject: [PATCH] fix(ioscompanion): reconnect after interrupted runner calls instead of restarting --- .../driver/ioscompanion/transport/runner.go | 44 ++++++++++-- .../ioscompanion/transport/runner_test.go | 72 ++++++++++++++++++- 2 files changed, 107 insertions(+), 9 deletions(-) diff --git a/internal/driver/ioscompanion/transport/runner.go b/internal/driver/ioscompanion/transport/runner.go index be64387..a62915e 100644 --- a/internal/driver/ioscompanion/transport/runner.go +++ b/internal/driver/ioscompanion/transport/runner.go @@ -23,11 +23,17 @@ var deadlineImmediate = time.Unix(1, 0) type runnerCompanion struct { uniqueDeviceIdentifier string bundleID string + address string mutex sync.Mutex conn net.Conn reader *bufio.Reader nextID int + + // dirty marks the connection desynced: a call was interrupted before its + // response was read, so the next call reconnects to the still-running + // server instead of misreading the stale response. + dirty bool } // DialRunner opens one persistent TCP connection to the simulator runner at @@ -42,6 +48,7 @@ func DialRunner(address, uniqueDeviceIdentifier, bundleID string) (Companion, er return &runnerCompanion{ uniqueDeviceIdentifier: uniqueDeviceIdentifier, bundleID: bundleID, + address: address, conn: conn, reader: bufio.NewReader(conn), }, nil @@ -69,6 +76,12 @@ func (c *runnerCompanion) call(ctx context.Context, method string, params map[st c.mutex.Lock() defer c.mutex.Unlock() + if c.dirty { + if err := c.reconnect(); err != nil { + return nil, fmt.Errorf("runner transport: %w: reconnect: %v", ErrCompanionUnavailable, err) + } + } + if params == nil { params = map[string]any{} } @@ -95,12 +108,12 @@ func (c *runnerCompanion) call(ctx context.Context, method string, params map[st } payload = append(payload, '\n') if _, err := c.conn.Write(payload); err != nil { - return nil, wrapTransport(ctx, "write", method, err) + return nil, c.wrapTransport(ctx, "write", method, err) } line, err := c.reader.ReadBytes('\n') if err != nil { - return nil, wrapTransport(ctx, "read", method, err) + return nil, c.wrapTransport(ctx, "read", method, err) } var response runnerResponse @@ -116,16 +129,33 @@ func (c *runnerCompanion) call(ctx context.Context, method string, params map[st return response.Result, nil } -// wrapTransport classifies a read/write failure. A cancelled or expired context -// is reported as such so callers see why the call was interrupted; either way -// the error wraps ErrCompanionUnavailable. -func wrapTransport(ctx context.Context, stage, method string, err error) error { +// wrapTransport classifies a read/write failure. A caller-imposed cancel or +// deadline is the caller's slowness budget, not a connection loss, so it does +// not carry the unavailable sentinel: a child restart would not make the call +// faster. Either way the connection is desynced and reconnects on the next +// call. +func (c *runnerCompanion) wrapTransport(ctx context.Context, stage, method string, err error) error { + c.dirty = true if ctxErr := ctx.Err(); ctxErr != nil { - return fmt.Errorf("runner transport: %w: %s %s interrupted: %v", ErrCompanionUnavailable, stage, method, ctxErr) + return fmt.Errorf("runner %s interrupted (%s): %w", method, stage, ctxErr) } return fmt.Errorf("runner transport: %w: %s %s: %v", ErrCompanionUnavailable, stage, method, err) } +// reconnect replaces the desynced connection with a fresh one to the same +// still-running server. +func (c *runnerCompanion) reconnect() error { + _ = c.conn.Close() + conn, err := net.Dial("tcp", c.address) + if err != nil { + return err + } + c.conn = conn + c.reader = bufio.NewReader(conn) + c.dirty = false + return nil +} + func (c *runnerCompanion) AccessibilityInfo(ctx context.Context) (string, error) { result, err := c.call(ctx, "snapshot", map[string]any{"bundleId": c.bundleID}) if err != nil { diff --git a/internal/driver/ioscompanion/transport/runner_test.go b/internal/driver/ioscompanion/transport/runner_test.go index 5788c5f..f53ef04 100644 --- a/internal/driver/ioscompanion/transport/runner_test.go +++ b/internal/driver/ioscompanion/transport/runner_test.go @@ -428,14 +428,82 @@ func TestContextCancellationUnblocksCall(t *testing.T) { if callErr == nil { t.Fatal("expected cancellation error") } - if !errors.Is(callErr, ErrCompanionUnavailable) { - t.Fatalf("cancellation error must wrap the sentinel: %v", callErr) + // A caller-imposed cancel is the caller's budget, not a connection + // loss: it must NOT wrap the sentinel, or a slow call would trigger a + // pointless child restart. + if !errors.Is(callErr, context.Canceled) { + t.Fatalf("cancellation error must carry the context error: %v", callErr) + } + if errors.Is(callErr, ErrCompanionUnavailable) { + t.Fatalf("cancellation error must not wrap the sentinel: %v", callErr) } case <-time.After(2 * time.Second): t.Fatal("cancelled call did not return within 2s") } } +func TestInterruptedCallReconnectsOnNextCall(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { listener.Close() }) + + // The server holds "describe" hostage and answers everything else, on + // every connection it accepts. A late reply to the interrupted request + // must never be misread by the following call. + accepted := make(chan net.Conn, 4) + go func() { + for { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + accepted <- conn + go func(c net.Conn) { + reader := bufio.NewReader(c) + for { + line, readErr := reader.ReadBytes('\n') + if readErr != nil { + return + } + var request runnerRequest + if json.Unmarshal(line, &request) != nil { + return + } + if request.Method == "describe" { + continue + } + response := `{"id":` + strconv.Itoa(request.ID) + `,"result":{"ok":true}}` + "\n" + if _, writeErr := c.Write([]byte(response)); writeErr != nil { + return + } + } + }(conn) + } + }() + + companion, err := DialRunner(listener.Addr().String(), "UDID", "com.example.app") + if err != nil { + t.Fatalf("DialRunner: %v", err) + } + t.Cleanup(func() { companion.Close() }) + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + if _, err := companion.Describe(ctx); err == nil { + t.Fatal("expected the held call to time out") + } + + // The next call must transparently reconnect and succeed. + if err := companion.Terminate(context.Background(), "com.example.app"); err != nil { + t.Fatalf("call after interrupt: %v", err) + } + if len(accepted) != 2 { + t.Fatalf("accepted %d connections, want 2 (reconnect)", len(accepted)) + } +} + func TestResponseIDMismatchIsSentinel(t *testing.T) { listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil {