import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { connectSocket } from "./ws"; import type { ChatEvent } from "./reducer"; /** A minimal scriptable WebSocket double. */ class FakeWebSocket { static OPEN = 1; static instances: FakeWebSocket[] = []; url: string; readyState = 0; sent: string[] = []; onopen: (() => void) | null = null; onmessage: ((event: { data: string }) => void) | null = null; onclose: (() => void) | null = null; constructor(url: string) { this.url = url; FakeWebSocket.instances.push(this); } send(data: string) { this.sent.push(data); } close() { this.readyState = 3; this.onclose?.(); } /* test helpers */ serverOpens() { this.readyState = FakeWebSocket.OPEN; this.onopen?.(); } serverSends(payload: unknown) { this.onmessage?.({ data: JSON.stringify(payload) }); } serverDrops() { this.readyState = 3; this.onclose?.(); } } describe("connectSocket", () => { let events: ChatEvent[]; const dispatch = (event: ChatEvent) => events.push(event); beforeEach(() => { events = []; FakeWebSocket.instances = []; vi.useFakeTimers(); }); afterEach(() => { vi.useRealTimers(); }); const connect = () => connectSocket({ url: "ws://test/ws", dispatch, webSocketImpl: FakeWebSocket as unknown as typeof WebSocket, initialRetryMs: 100, }); it("dispatches parsed server messages", () => { connect(); const ws = FakeWebSocket.instances[0]; ws.serverOpens(); ws.serverSends({ type: "hello", sessionId: "s1" }); expect(events).toContainEqual({ type: "server", message: { type: "hello", sessionId: "s1" }, }); }); it("sends only while the socket is open", () => { const socket = connect(); const ws = FakeWebSocket.instances[0]; socket.send({ type: "prompt", text: "too early" }); expect(ws.sent).toHaveLength(0); ws.serverOpens(); socket.send({ type: "prompt", text: "hello" }); expect(ws.sent).toEqual([ JSON.stringify({ type: "prompt", text: "hello" }), ]); }); it("reconnects with backoff after a drop", () => { connect(); FakeWebSocket.instances[0].serverOpens(); FakeWebSocket.instances[0].serverDrops(); expect(events).toContainEqual({ type: "connection", status: "offline" }); expect(FakeWebSocket.instances).toHaveLength(1); vi.advanceTimersByTime(100); expect(FakeWebSocket.instances).toHaveLength(2); // second drop: the delay doubles FakeWebSocket.instances[1].serverDrops(); vi.advanceTimersByTime(100); expect(FakeWebSocket.instances).toHaveLength(2); vi.advanceTimersByTime(100); expect(FakeWebSocket.instances).toHaveLength(3); }); it("does not reconnect after close()", () => { const socket = connect(); FakeWebSocket.instances[0].serverOpens(); socket.close(); vi.advanceTimersByTime(10_000); expect(FakeWebSocket.instances).toHaveLength(1); }); it("survives malformed frames", () => { connect(); const ws = FakeWebSocket.instances[0]; ws.serverOpens(); ws.onmessage?.({ data: "{not json" }); ws.serverSends({ type: "turn_started" }); expect(events).toContainEqual({ type: "server", message: { type: "turn_started" }, }); }); });