Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 59 additions & 4 deletions apps/extension/src/lib/__tests__/connection-controller.test.ts
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
import { MIN_COMPATIBLE_PROTOCOL } from "../../transport/handshake";
import type { ConnectionStateHandler, Transport } from "../../transport/transport";
import type { ConnectionState, HandshakeResult } from "../../transport/types";
import type { ConnectionStateHandler, FrameHandler, Transport } from "../../transport/transport";
import type { ConnectionState, HandshakeResult, ProtocolFrame } from "../../transport/types";
import { __testing__, ConnectionController } from "../connection-controller";

vi.mock("../instance-id", () => ({
Expand Down Expand Up @@ -97,6 +97,7 @@ describe("computeConnectedState (protocol-based compat)", () => {
function makeMockTransport(initialState: ConnectionState = "disconnected") {
let state = initialState;
const stateHandlers = new Set<ConnectionStateHandler>();
const messageHandlers = new Set<FrameHandler>();
const transport = {
get state() {
return state;
Expand All @@ -110,7 +111,10 @@ function makeMockTransport(initialState: ConnectionState = "disconnected") {
for (const h of stateHandlers) h("disconnected");
}),
send: vi.fn(),
onMessage: vi.fn(() => ({ dispose: () => {} })),
onMessage: vi.fn((handler: FrameHandler) => {
messageHandlers.add(handler);
return { dispose: () => messageHandlers.delete(handler) };
}),
onConnectionStateChange: vi.fn((handler: ConnectionStateHandler) => {
stateHandlers.add(handler);
return { dispose: () => stateHandlers.delete(handler) };
Expand All @@ -119,6 +123,9 @@ function makeMockTransport(initialState: ConnectionState = "disconnected") {
state = next;
for (const h of stateHandlers) h(next);
},
emitMessage(frame: ProtocolFrame) {
for (const handler of messageHandlers) handler(frame);
},
};
return transport as typeof transport & Transport;
}
Expand All @@ -128,6 +135,10 @@ describe("ConnectionController connectionEnabled", () => {
vi.clearAllMocks();
});

afterEach(() => {
vi.useRealTimers();
});

it("does not connect on attach when connection is disabled", async () => {
const controller = new ConnectionController();
const transport = makeMockTransport();
Expand Down Expand Up @@ -201,4 +212,48 @@ describe("ConnectionController connectionEnabled", () => {

expect(controller.snapshot().state).toBe("disconnected");
});

it("disconnects and retries when handshake fails while the socket is still open", async () => {
const controller = new ConnectionController();
const transport = makeMockTransport();
await controller.attach(transport, { name: "Chrome", version: "120" }, true);
const request = transport.send.mock.calls[0]?.[0] as { id: string };

transport.emitMessage({
id: request.id,
error: { code: "protocol_error", message: "bad handshake" },
});
// The failure path disconnects once; that state change then drives
// recoverFromDisconnect(), which issues one more idempotent disconnect
// (cancelling the transport's auto-reconnect) before teardown + reconnect.
await vi.waitFor(() => expect(transport.disconnect).toHaveBeenCalledTimes(2));

// Reconnection is owned by the recovery path, and a fresh handshake is
// attempted on the new connection.
await vi.waitFor(() => expect(transport.connect).toHaveBeenCalledTimes(2));
await vi.waitFor(() => expect(transport.send.mock.calls.length).toBeGreaterThan(1));
});

it("binds each handshake to the connection that initiated it", async () => {
const controller = new ConnectionController();
const transport = makeMockTransport();
await controller.attach(transport, { name: "Chrome", version: "120" }, true);
const first = transport.send.mock.calls[0]?.[0] as { id: string };

// Socket drops; the recovery path reconnects on its own and starts a new
// handshake bound to the new connection.
transport.emitState("disconnected");
await vi.waitFor(() => expect(transport.connect).toHaveBeenCalledTimes(2));
const second = transport.send.mock.calls[transport.send.mock.calls.length - 1]?.[0] as {
id: string;
};
expect(second.id).not.toBe(first.id);

transport.emitMessage({ id: first.id, result: handshake("1.0", "1.0") });
await Promise.resolve();
expect(controller.snapshot().state).not.toBe("connected");

transport.emitMessage({ id: second.id, result: handshake("1.0", "1.0") });
await vi.waitFor(() => expect(controller.snapshot().state).toBe("connected"));
});
});
59 changes: 48 additions & 11 deletions apps/extension/src/lib/connection-controller.ts
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,8 @@ export class ConnectionController {
private lastError: string | null = null;
private connectionEnabled = true;
private listeners = new Set<Listener>();
private handshakeInFlight = false;
private connectionGeneration = 0;
private handshakeAbort: AbortController | null = null;
private lifecycleHooks: ConnectionLifecycleHooks = {};
private disconnectRecovery: Promise<void> | null = null;

Expand Down Expand Up @@ -87,6 +88,9 @@ export class ConnectionController {

transport.onConnectionStateChange((s) => {
if (s === "disconnected") {
// Cancel any in-flight handshake (PR #19) before local teardown and
// reconnect recovery (#17).
this.cancelHandshake();
this.handshake = null;
if (this.connectionEnabled) {
this.setState("disconnected");
Expand All @@ -96,7 +100,7 @@ export class ConnectionController {
}
if (!this.connectionEnabled) return;
if (s === "connected") {
void this.runHandshake(browser);
this.startHandshake(browser);
return;
}
this.setState(s);
Expand Down Expand Up @@ -142,18 +146,33 @@ export class ConnectionController {
this.fire();
}

private async runHandshake(browser: { name: string; version: string }): Promise<void> {
private startHandshake(browser: { name: string; version: string }): void {
this.cancelHandshake();
const generation = ++this.connectionGeneration;
const abort = new AbortController();
this.handshakeAbort = abort;
void this.runHandshake(browser, generation, abort.signal);
}

private async runHandshake(
browser: { name: string; version: string },
generation: number,
signal: AbortSignal,
): Promise<void> {
if (!this.transport) return;
if (!this.connectionEnabled) return;
if (this.handshakeInFlight) return;
this.handshakeInFlight = true;
this.setState("connecting");
try {
const outcome = await performHandshake(this.transport, {
instanceId: this.instanceId,
browser,
label: this.label,
});
const outcome = await performHandshake(
this.transport,
{
instanceId: this.instanceId,
browser,
label: this.label,
},
{ signal },
);
if (generation !== this.connectionGeneration || signal.aborted) return;
this.handshake = outcome.result;
const verdict = computeConnectedState(outcome.result);
if (verdict.kind === "rejected") {
Expand All @@ -166,15 +185,21 @@ export class ConnectionController {
this.lastError = null;
this.setState(verdict.kind);
} catch (err) {
if (generation !== this.connectionGeneration || signal.aborted || isAbortError(err)) return;
this.handshake = null;
this.lastError = err instanceof Error ? err.message : String(err);
// Disconnecting here emits a transport state change, which drives
// recoverFromDisconnect() — reconnection (with session teardown) is
// owned by that path; no separate retry timer is needed.
await this.transport.disconnect().catch(() => {});
this.setState("disconnected");
} finally {
this.handshakeInFlight = false;
if (generation === this.connectionGeneration) this.handshakeAbort = null;
}
}

private async applyDisabledState(): Promise<void> {
this.cancelHandshake();
this.handshake = null;
this.lastError = null;
await this.lifecycleHooks.beforeDisconnect?.();
Expand Down Expand Up @@ -214,6 +239,12 @@ export class ConnectionController {
return this.disconnectRecovery;
}

private cancelHandshake(): void {
this.connectionGeneration += 1;
this.handshakeAbort?.abort();
this.handshakeAbort = null;
}

private setState(next: ConnectionState): void {
if (this.currentState === next) return;
this.currentState = next;
Expand All @@ -232,6 +263,12 @@ export class ConnectionController {
}
}

function isAbortError(err: unknown): boolean {
return (
typeof err === "object" && err !== null && (err as { name?: string }).name === "AbortError"
);
}

/**
* Verdict of the symmetric post-handshake compat check.
*/
Expand Down
19 changes: 19 additions & 0 deletions apps/extension/src/transport/__tests__/handshake.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -198,6 +198,25 @@ describe("performHandshake", () => {
result: { server: "browser-skill-daemon", protocol_version: "1.0" },
});
});

it("stops waiting immediately when the connection generation is aborted", async () => {
const { transport } = deferredFakeTransport();
const controller = new AbortController();
const pending = performHandshake(
transport,
{
instanceId: "x",
browser: { name: "chrome", version: "131" },
label: "",
rpcId: "hs-abort",
},
{ signal: controller.signal, timeoutMs: 60_000 },
);

controller.abort();

await expect(pending).rejects.toMatchObject({ name: "AbortError" });
});
});

describe("detectBrowserMeta", () => {
Expand Down
32 changes: 32 additions & 0 deletions apps/extension/src/transport/__tests__/ws-transport.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,10 @@ class FakeSocket {
this.emit("close", { code, reason: "server-gone" });
}

emitClose(code = 1006): void {
this.emit("close", { code, reason: "delayed-close" });
}

private emit(type: string, ev: unknown): void {
for (const l of this.listeners[type] ?? []) {
// biome-ignore lint/suspicious/noExplicitAny: minimal fake mirrors WebSocket Event shape
Expand Down Expand Up @@ -202,6 +206,34 @@ describe("WSTransport", () => {
expect(FakeSocket.instances.length).toBe(beforeCount + 1);
});

it("cancels pending backoff and ignores delayed events from a replaced socket", async () => {
const t = new WSTransport({
url: "ws://127.0.0.1:52800",
webSocketFactory: (url) => new FakeSocket(url) as unknown as WebSocket,
});
const initial = t.connect();
const first = lastSocket();
first.open();
await initial;

first.serverClose();
const reconnect = t.connect();
expect(FakeSocket.instances).toHaveLength(2);
const second = lastSocket();
second.open();
await reconnect;

// A real browser can deliver the old close callback after the replacement
// socket is already open. It must not clear or disconnect the new socket.
first.emitClose();
t.send({ id: "current", method: "system.ping" });
expect(second.sent).toEqual([JSON.stringify({ id: "current", method: "system.ping" })]);
expect(t.state).toBe("connected");

await vi.advanceTimersByTimeAsync(5_000);
expect(FakeSocket.instances).toHaveLength(2);
});

it("ignores malformed inbound messages (non-JSON) without throwing", async () => {
const handler = vi.fn();
const t = new WSTransport({
Expand Down
15 changes: 14 additions & 1 deletion apps/extension/src/transport/handshake.ts
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ function ridToString(): string {
export function performHandshake(
transport: Transport,
input: HandshakeInput,
options: { timeoutMs?: number } = {},
options: { timeoutMs?: number; signal?: AbortSignal } = {},
): Promise<HandshakeOutcome> {
const id = input.rpcId ?? ridToString();
const params: HandshakeParams = {
Expand All @@ -75,6 +75,12 @@ export function performHandshake(
const timeoutMs = options.timeoutMs ?? 10_000;

return new Promise<HandshakeOutcome>((resolve, reject) => {
const onAbort = () => {
cleanup();
const err = new Error("[handshake] aborted because the connection changed");
err.name = "AbortError";
reject(err);
};
const timer = setTimeout(() => {
cleanup();
reject(new Error("[handshake] timed out waiting for daemon response"));
Expand All @@ -97,7 +103,14 @@ export function performHandshake(
function cleanup() {
clearTimeout(timer);
sub.dispose();
options.signal?.removeEventListener("abort", onAbort);
}

if (options.signal?.aborted) {
onAbort();
return;
}
options.signal?.addEventListener("abort", onAbort, { once: true });

try {
transport.send(req);
Expand Down
Loading