diff --git a/src/realtime-transcription/websocket-session.test.ts b/src/realtime-transcription/websocket-session.test.ts index 34fdd1b942bc..b6afe795bce5 100644 --- a/src/realtime-transcription/websocket-session.test.ts +++ b/src/realtime-transcription/websocket-session.test.ts @@ -1,11 +1,14 @@ // Realtime transcription websocket tests cover websocket session lifecycle. import { createServer } from "node:http"; import type { AddressInfo } from "node:net"; +import { setTimeout as delay } from "node:timers/promises"; import { expectDefined } from "@openclaw/normalization-core"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; -import type WebSocket from "ws"; -import { WebSocketServer } from "ws"; -import { createRealtimeTranscriptionWebSocketSession } from "./websocket-session.js"; +import WebSocket, { WebSocketServer } from "ws"; +import { + createRealtimeTranscriptionWebSocketSession, + type RealtimeTranscriptionWebSocketTransport, +} from "./websocket-session.js"; let cleanup: (() => Promise) | undefined; @@ -269,6 +272,240 @@ describe("createRealtimeTranscriptionWebSocketSession", () => { session.close(); }); + it("keeps replacement sockets owned when retired socket callbacks arrive late", async () => { + const connections: WebSocket[] = []; + const transports: RealtimeTranscriptionWebSocketTransport[] = []; + const onError = vi.fn(); + const onTranscript = vi.fn(); + const server = await createRealtimeServer({ + onConnection: (socket) => connections.push(socket), + }); + const session = createRealtimeTranscriptionWebSocketSession<{ text?: string }>({ + providerId: "test", + callbacks: { onError, onTranscript }, + url: server.url, + readyOnOpen: true, + closeTimeoutMs: 30, + reconnectDelayMs: 1, + onOpen: (transport) => transports.push(transport), + onMessage: (event, transport) => { + if (event.text) { + transport.callbacks.onTranscript?.(event.text); + } + }, + sendAudio: (audio, transport) => { + transport.sendBinary(audio); + }, + }); + + await session.connect(); + const retiredSocket = Reflect.get(session, "ws") as WebSocket; + session.close(); + await session.connect(); + await vi.waitFor(() => expect(connections).toHaveLength(2)); + + const retiredTransport = expectDefined(transports[0], "retired socket transport"); + expect(retiredTransport.isOpen()).toBe(false); + expect(retiredTransport.isReady()).toBe(false); + expect(retiredTransport.sendJson({ type: "retired" })).toBe(false); + retiredTransport.markReady(); + retiredTransport.failConnect(new Error("retired provider failed")); + retiredTransport.closeNow(); + + // ws delivers close/error/message asynchronously; replaying those actual + // socket callbacks proves a retired connection cannot poison its replacement. + retiredSocket.emit("message", Buffer.from(JSON.stringify({ text: "stale transcript" }))); + retiredSocket.emit("error", new Error("retired socket failed")); + retiredSocket.emit("close", 1000, Buffer.from("retired")); + + expect(onTranscript).not.toHaveBeenCalled(); + expect(onError).not.toHaveBeenCalled(); + expect(session.isConnected()).toBe(true); + + await delay(60); + expect(session.isConnected()).toBe(true); + expect(connections).toHaveLength(2); + + connections[1]?.send(JSON.stringify({ text: "current transcript" })); + await vi.waitFor(() => expect(onTranscript).toHaveBeenCalledWith("current transcript")); + session.close(); + }); + + it("discards superseded asynchronous connection preparation", async () => { + const connections: WebSocket[] = []; + const server = await createRealtimeServer({ + onConnection: (socket) => connections.push(socket), + }); + let resolveFirstUrl!: (url: string) => void; + const firstUrl = new Promise((resolve) => { + resolveFirstUrl = resolve; + }); + let connectionAttempt = 0; + const session = createRealtimeTranscriptionWebSocketSession({ + providerId: "test", + callbacks: {}, + url: async () => (++connectionAttempt === 1 ? await firstUrl : server.url), + readyOnOpen: true, + sendAudio: (audio, transport) => { + transport.sendBinary(audio); + }, + }); + + const supersededConnection = session.connect(); + session.close(); + await session.connect(); + resolveFirstUrl(server.url); + await supersededConnection; + + await delay(20); + expect(connections).toHaveLength(1); + expect(session.isConnected()).toBe(true); + session.close(); + }); + + it("cancels a retired reconnect delay before starting a replacement socket", async () => { + const connections: WebSocket[] = []; + const server = await createRealtimeServer({ + onConnection: (socket) => connections.push(socket), + }); + const session = createRealtimeTranscriptionWebSocketSession({ + providerId: "test", + callbacks: {}, + url: server.url, + readyOnOpen: true, + reconnectDelayMs: 50, + sendAudio: (audio, transport) => { + transport.sendBinary(audio); + }, + }); + + await session.connect(); + connections[0]?.close(1011, "retired connection"); + await vi.waitFor(() => expect(session.isConnected()).toBe(false)); + + session.close(); + await session.connect(); + await delay(100); + + expect(connections).toHaveLength(2); + expect(session.isConnected()).toBe(true); + session.close(); + }); + + it("reconnects after a healthy successor closes following a failed connection", async () => { + const connections: WebSocket[] = []; + const onError = vi.fn(); + const server = await createRealtimeServer({ + onConnection: (socket) => { + connections.push(socket); + socket.send( + JSON.stringify( + connections.length === 2 + ? { type: "error", message: "provider handshake rejected" } + : { type: "ready" }, + ), + ); + }, + }); + const session = createRealtimeTranscriptionWebSocketSession<{ + message?: string; + type?: string; + }>({ + providerId: "test", + callbacks: { onError }, + url: server.url, + reconnectDelayMs: 1, + onMessage: (event, transport) => { + if (event.type === "ready") { + transport.markReady(); + } else if (event.type === "error") { + transport.failConnect(new Error(event.message)); + } + }, + sendAudio: (audio, transport) => { + transport.sendBinary(audio); + }, + }); + + await session.connect(); + connections[0]?.close(1011, "first connection dropped"); + await vi.waitFor(() => expect(connections).toHaveLength(3)); + await vi.waitFor(() => expect(session.isConnected()).toBe(true)); + expect(onError).toHaveBeenCalledWith( + expect.objectContaining({ message: "provider handshake rejected" }), + ); + + connections[2]?.close(1011, "healthy connection dropped"); + await vi.waitFor(() => expect(connections).toHaveLength(4), { timeout: 500 }); + await vi.waitFor(() => expect(session.isConnected()).toBe(true)); + session.close(); + }); + + it("delivers graceful provider finals and sends the close signal only once", async () => { + const finalizedFrames: unknown[] = []; + const transcripts: string[] = []; + let providerSocket: WebSocket | undefined; + const server = await createRealtimeServer({ + onConnection: (socket) => { + providerSocket = socket; + }, + onText: (payload) => { + finalizedFrames.push(payload); + if (finalizedFrames.length === 1) { + providerSocket?.send(JSON.stringify({ text: "final provider transcript" })); + } + }, + }); + const session = createRealtimeTranscriptionWebSocketSession<{ text?: string }>({ + providerId: "test", + callbacks: { onTranscript: (text) => transcripts.push(text) }, + url: server.url, + readyOnOpen: true, + closeTimeoutMs: 100, + onClose: (transport) => { + transport.sendJson({ type: "finalize" }); + }, + onMessage: (event, transport) => { + if (event.text) { + transport.callbacks.onTranscript?.(event.text); + } + }, + sendAudio: (audio, transport) => { + transport.sendBinary(audio); + }, + }); + + await session.connect(); + session.close(); + session.close(); + + await vi.waitFor(() => expect(transcripts).toEqual(["final provider transcript"])); + await delay(20); + expect(finalizedFrames).toEqual([{ type: "finalize" }]); + }); + + it("terminates the captured socket when graceful provider shutdown expires", async () => { + const server = await createRealtimeServer(); + const session = createRealtimeTranscriptionWebSocketSession({ + providerId: "test", + callbacks: {}, + url: server.url, + readyOnOpen: true, + closeTimeoutMs: 20, + sendAudio: (audio, transport) => { + transport.sendBinary(audio); + }, + }); + + await session.connect(); + const ownedSocket = Reflect.get(session, "ws") as WebSocket; + const terminate = vi.spyOn(ownedSocket, "terminate"); + session.close(); + + await vi.waitFor(() => expect(terminate).toHaveBeenCalledOnce(), { timeout: 500 }); + expect(session.isConnected()).toBe(false); + }); + it("lets providers mark ready after a JSON handshake", async () => { const frames: unknown[] = []; const framesReady = createSignal(); diff --git a/src/realtime-transcription/websocket-session.ts b/src/realtime-transcription/websocket-session.ts index 45e6fa3dbb95..5a42737d608e 100644 --- a/src/realtime-transcription/websocket-session.ts +++ b/src/realtime-transcription/websocket-session.ts @@ -79,13 +79,12 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript private readySinceMs: number | undefined; private readonly reconnectSupervisor: RetrySupervisor; private reconnecting = false; - private suppressReconnect = false; private ws: WebSocket | null = null; + private connectionGeneration = 0; private readonly flowId = randomUUID(); private readonly options: RealtimeTranscriptionWebSocketSessionOptions; - private readonly transport: RealtimeTranscriptionWebSocketTransport; - private failConnect: ((error: Error) => void) | undefined; - private markReady: (() => void) | undefined; + private transport: RealtimeTranscriptionWebSocketTransport | undefined; + private cancelConnecting: (() => void) | undefined; constructor(options: RealtimeTranscriptionWebSocketSessionOptions) { this.options = options; @@ -98,35 +97,25 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript }, options.maxReconnectAttempts ?? DEFAULT_MAX_RECONNECT_ATTEMPTS, ); - this.transport = { - callbacks: options.callbacks, - closeNow: () => { - this.closed = true; - this.reconnectSupervisor.cancel(); - this.forceClose(); - }, - failConnect: (error) => this.failConnect?.(error), - isOpen: () => this.ws?.readyState === WebSocket.OPEN, - isReady: () => this.ready, - markReady: () => this.markReady?.(), - sendBinary: (payload) => this.sendBinary(payload), - sendJson: (payload) => this.sendJson(payload), - }; } async connect(): Promise { + const previousSocket = this.ws; + this.connectionGeneration += 1; + this.cancelConnecting?.(); + this.forceClose(previousSocket); this.closed = false; - this.suppressReconnect = false; this.readySinceMs = undefined; + this.reconnecting = false; this.reconnectSupervisor.reset(); - await this.doConnect(); + await this.doConnect(this.connectionGeneration); } sendAudio(audio: Buffer): void { if (this.closed || audio.byteLength === 0) { return; } - if (this.ws?.readyState === WebSocket.OPEN && this.ready) { + if (this.ws?.readyState === WebSocket.OPEN && this.ready && this.transport) { this.options.sendAudio(audio, this.transport); return; } @@ -136,22 +125,32 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript } close(): void { + if (this.closed) { + return; + } this.closed = true; + this.cancelConnecting?.(); this.connected = false; this.ready = false; this.readySinceMs = undefined; this.reconnectSupervisor.cancel(); this.clearQueuedAudio(); - if (!this.ws || this.ws.readyState !== WebSocket.OPEN) { - this.forceClose(); + const socket = this.ws; + const transport = this.transport; + if (!socket || socket.readyState !== WebSocket.OPEN || !transport) { + this.forceClose(socket); return; } try { - this.options.onClose?.(this.transport); + this.options.onClose?.(transport); } catch (error) { this.emitError(error); } - this.closeTimer = setTimeout(() => this.forceClose(), this.closeTimeoutMs); + if (this.ws === socket) { + // Keep the owning socket alive for provider final transcripts, but never + // let its shutdown deadline terminate a later connection generation. + this.closeTimer = setTimeout(() => this.forceClose(socket), this.closeTimeoutMs); + } } isConnected(): boolean { @@ -169,14 +168,22 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript return this.options.maxQueuedBytes ?? DEFAULT_MAX_QUEUED_BYTES; } - private async doConnect(): Promise { + private async doConnect(generation: number): Promise { await new Promise((resolve, reject) => { + if (generation !== this.connectionGeneration || this.closed) { + resolve(); + return; + } this.ready = false; const debugProxy = resolveDebugProxySettings(); const proxyAgent = createDebugProxyWebSocketAgent(debugProxy); let settled = false; let opened = false; let connectTimeout: ReturnType | undefined; + let socket: WebSocket | undefined; + + const ownsGeneration = () => generation === this.connectionGeneration; + const ownsSocket = () => ownsGeneration() && this.ws === socket; const normalizeError = (error: unknown) => error instanceof Error ? error : new Error(String(error)); @@ -186,6 +193,9 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript clearTimeout(connectTimeout); connectTimeout = undefined; } + if (this.cancelConnecting === finishClosedConnect) { + this.cancelConnecting = undefined; + } }; const finishClosedConnect = () => { @@ -201,11 +211,15 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript if (settled) { return; } + if (!ownsSocket()) { + finishClosedConnect(); + return; + } settled = true; clearConnectTimeout(); this.ready = true; this.readySinceMs = Date.now(); - this.flushQueuedAudio(); + this.flushQueuedAudio(transport); resolve(); }; @@ -213,16 +227,44 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript if (settled) { return; } + if (!ownsGeneration() || (socket && !ownsSocket())) { + finishClosedConnect(); + return; + } settled = true; clearConnectTimeout(); this.emitError(error); - this.suppressReconnect = true; - this.forceClose(); + this.forceClose(socket ?? this.ws); reject(error); }; + this.cancelConnecting = finishClosedConnect; - this.markReady = finishConnect; - this.failConnect = failConnect; + const transport: RealtimeTranscriptionWebSocketTransport = { + callbacks: this.options.callbacks, + closeNow: () => { + if (!ownsSocket()) { + return; + } + this.closed = true; + this.cancelConnecting?.(); + this.reconnectSupervisor.cancel(); + this.forceClose(socket); + }, + failConnect: (error) => { + if (ownsSocket()) { + failConnect(error); + } + }, + isOpen: () => ownsSocket() && socket?.readyState === WebSocket.OPEN, + isReady: () => ownsSocket() && this.ready, + markReady: () => { + if (ownsSocket()) { + finishConnect(); + } + }, + sendBinary: (payload) => this.send(payload, socket, generation), + sendJson: (payload) => this.send(JSON.stringify(payload), socket, generation), + }; connectTimeout = setTimeout(() => { failConnect( @@ -244,30 +286,35 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript if (settled) { return; } - if (this.closed) { + if (!ownsGeneration() || this.closed) { finishClosedConnect(); return; } this.currentUrl = connection.url; try { - this.ws = new WebSocket(this.currentUrl, { + socket = new WebSocket(this.currentUrl, { headers: connection.headers, maxPayload: REALTIME_TRANSCRIPTION_WS_MAX_PAYLOAD_BYTES, ...(proxyAgent ? { agent: proxyAgent } : {}), }); - this.ws.binaryType = "nodebuffer"; + socket.binaryType = "nodebuffer"; + this.ws = socket; + this.transport = transport; } catch (error) { failConnect(normalizeError(error)); return; } - this.ws.on("open", () => { + socket.on("open", () => { + if (!ownsSocket()) { + return; + } opened = true; this.connected = true; this.captureLocalOpen(); try { - this.options.onOpen?.(this.transport); + this.options.onOpen?.(transport); if (this.options.readyOnOpen) { finishConnect(); } @@ -276,7 +323,10 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript } }); - this.ws.on("message", (data) => { + socket.on("message", (data) => { + if (!ownsSocket()) { + return; + } const payload = data as Buffer; this.captureFrame("inbound", payload); try { @@ -284,13 +334,16 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript return; } const parseMessage = this.options.parseMessage ?? defaultParseMessage; - this.options.onMessage(parseMessage(payload) as Event, this.transport); + this.options.onMessage(parseMessage(payload) as Event, transport); } catch (error) { this.emitError(error); } }); - this.ws.on("error", (error) => { + socket.on("error", (error) => { + if (!ownsSocket()) { + return; + } const normalized = normalizeError(error); this.captureError(normalized); if (!opened || !settled) { @@ -300,7 +353,10 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript this.emitError(normalized); }); - this.ws.on("close", (code, reasonBuffer) => { + socket.on("close", (code, reasonBuffer) => { + if (!ownsSocket()) { + return; + } clearConnectTimeout(); this.captureClose(code, reasonBuffer); const readyForMs = this.readySinceMs === undefined ? 0 : Date.now() - this.readySinceMs; @@ -317,10 +373,6 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript if (this.closed) { return; } - if (this.suppressReconnect) { - this.suppressReconnect = false; - return; - } if (!opened || !settled) { failConnect( new Error( @@ -330,7 +382,7 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript ); return; } - void this.attemptReconnect(); + void this.attemptReconnect(generation); }); })(); }); @@ -349,8 +401,8 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript return { url, headers }; } - private async attemptReconnect(): Promise { - if (this.closed || this.reconnecting) { + private async attemptReconnect(generation: number): Promise { + if (generation !== this.connectionGeneration || this.closed || this.reconnecting) { return; } const retry = this.reconnectSupervisor.next(); @@ -366,16 +418,18 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript this.reconnecting = true; try { await sleepWithAbort(retry.delayMs, retry.signal); - if (!this.closed) { - await this.doConnect(); + if (generation === this.connectionGeneration && !this.closed) { + await this.doConnect(generation); } } catch { - if (!this.closed) { + if (generation === this.connectionGeneration && !this.closed) { this.reconnecting = false; - await this.attemptReconnect(); + await this.attemptReconnect(generation); } } finally { - this.reconnecting = false; + if (generation === this.connectionGeneration) { + this.reconnecting = false; + } } } @@ -398,11 +452,11 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript this.compactQueuedAudio(); } - private flushQueuedAudio(): void { + private flushQueuedAudio(transport: RealtimeTranscriptionWebSocketTransport): void { for (let index = this.queuedAudioHead; index < this.queuedAudio.length; index += 1) { const audio = this.queuedAudio[index]; if (audio) { - this.options.sendAudio(audio, this.transport); + this.options.sendAudio(audio, transport); } } this.clearQueuedAudio(); @@ -422,26 +476,28 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript this.queuedBytes = 0; } - private sendBinary(payload: Buffer): boolean { - if (this.ws?.readyState !== WebSocket.OPEN) { + private send( + payload: Buffer | string, + socket: WebSocket | undefined, + generation: number, + ): boolean { + if ( + !socket || + generation !== this.connectionGeneration || + this.ws !== socket || + socket.readyState !== WebSocket.OPEN + ) { return false; } this.captureFrame("outbound", payload); - this.ws.send(payload); + socket.send(payload); return true; } - private sendJson(payload: unknown): boolean { - if (this.ws?.readyState !== WebSocket.OPEN) { - return false; + private forceClose(socket: WebSocket | null | undefined = this.ws): void { + if (socket !== this.ws) { + return; } - const serialized = JSON.stringify(payload); - this.captureFrame("outbound", serialized); - this.ws.send(serialized); - return true; - } - - private forceClose(): void { if (this.closeTimer) { clearTimeout(this.closeTimer); this.closeTimer = undefined; @@ -449,10 +505,9 @@ class WebSocketRealtimeTranscriptionSession implements RealtimeTranscript this.connected = false; this.ready = false; this.readySinceMs = undefined; - if (this.ws) { - this.ws.close(1000, "Transcription session closed"); - this.ws = null; - } + this.ws = null; + this.transport = undefined; + socket?.terminate(); } private emitError(error: unknown): void {