diff --git a/ui/src/pages/chat/realtime-talk-google-live-lifecycle.ts b/ui/src/pages/chat/realtime-talk-google-live-lifecycle.ts new file mode 100644 index 000000000000..6a4cd8fa99e0 --- /dev/null +++ b/ui/src/pages/chat/realtime-talk-google-live-lifecycle.ts @@ -0,0 +1,148 @@ +import type { + RealtimeTalkJsonPcmWebSocketSessionResult, + RealtimeTalkTransportStartResult, +} from "./realtime-talk-shared.ts"; + +const GOOGLE_LIVE_WEBSOCKET_HOST = "generativelanguage.googleapis.com"; +const GOOGLE_LIVE_WEBSOCKET_PATH = + /^\/ws\/google\.ai\.generativelanguage\.v[0-9a-z]+\.GenerativeService\.BidiGenerateContent(?:Constrained)?$/; +export const GOOGLE_LIVE_SETUP_TIMEOUT_MS = 30_000; + +export function buildGoogleLiveUrl(session: RealtimeTalkJsonPcmWebSocketSessionResult): string { + let url: URL; + try { + url = new URL(session.websocketUrl); + } catch { + throw new Error("Invalid Google Live WebSocket URL"); + } + if (url.protocol !== "wss:") { + throw new Error("Google Live WebSocket URL must use wss://"); + } + if (url.hostname.toLowerCase() !== GOOGLE_LIVE_WEBSOCKET_HOST) { + throw new Error("Untrusted Google Live WebSocket host"); + } + if (url.username || url.password) { + throw new Error("Google Live WebSocket URL must not include credentials"); + } + if (!GOOGLE_LIVE_WEBSOCKET_PATH.test(url.pathname)) { + throw new Error("Untrusted Google Live WebSocket path"); + } + url.search = ""; + url.searchParams.set("access_token", session.clientSecret); + return url.toString(); +} + +type GoogleLiveConnectionState = + | "idle" + | "connecting" + | "ready" + | "active" + | "cancelled" + | "failed"; + +type StartupWaiter = { + resolve: (result: RealtimeTalkTransportStartResult) => void; + reject: (error: Error) => void; +}; + +export class GoogleLiveConnectionLifecycle { + private state: GoogleLiveConnectionState = "idle"; + private socket: WebSocket | null = null; + private waiter: StartupWaiter | null = null; + private error: Error | null = null; + + get currentState(): GoogleLiveConnectionState { + return this.state; + } + + get isActive(): boolean { + return this.state === "active"; + } + + get setupComplete(): boolean { + return this.state === "ready" || this.state === "active"; + } + + begin(socket: WebSocket): Promise { + this.state = "connecting"; + this.socket = socket; + this.error = null; + return new Promise((resolve, reject) => { + this.waiter = { resolve, reject }; + }); + } + + markReady(socket: WebSocket): boolean { + if (this.state !== "connecting" || this.socket !== socket) { + return false; + } + this.state = "ready"; + this.takeWaiter()?.resolve("ready"); + return true; + } + + finishStart(result: RealtimeTalkTransportStartResult): RealtimeTalkTransportStartResult { + if (this.error) { + throw this.error; + } + return this.state === "cancelled" ? "cancelled" : result; + } + + activate(): boolean { + if (this.state === "active") { + return false; + } + if (this.error) { + throw this.error; + } + if (this.state === "cancelled" || this.state === "idle") { + return false; + } + if (this.state !== "ready") { + throw new Error("Google Live transport activated before setup completed"); + } + this.state = "active"; + return true; + } + + failStartup(socket: WebSocket, error: Error): boolean { + if (this.socket !== socket || (this.state !== "connecting" && this.state !== "ready")) { + return false; + } + this.state = "failed"; + this.error = error; + this.takeWaiter()?.reject(error); + return true; + } + + cancel(): void { + if (this.state === "cancelled" || this.state === "failed") { + return; + } + this.state = "cancelled"; + this.takeWaiter()?.resolve("cancelled"); + } + + private takeWaiter(): StartupWaiter | null { + const waiter = this.waiter; + this.waiter = null; + return waiter; + } +} + +export function runRealtimeTalkCleanup(steps: Array<() => void>): void { + let firstError: Error | undefined; + for (const step of steps) { + try { + step(); + } catch (error) { + firstError ??= + error instanceof Error + ? error + : new Error("Realtime Talk cleanup failed", { cause: error }); + } + } + if (firstError) { + throw firstError; + } +} diff --git a/ui/src/pages/chat/realtime-talk-google-live-timeout.test.ts b/ui/src/pages/chat/realtime-talk-google-live-timeout.test.ts new file mode 100644 index 000000000000..1ed89bd74d44 --- /dev/null +++ b/ui/src/pages/chat/realtime-talk-google-live-timeout.test.ts @@ -0,0 +1,304 @@ +// @vitest-environment jsdom +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { GoogleLiveRealtimeTalkTransport } from "./realtime-talk-google-live.ts"; +import type { + RealtimeTalkCallbacks, + RealtimeTalkJsonPcmWebSocketSessionResult, +} from "./realtime-talk-shared.ts"; + +const SETUP_TIMEOUT_MS = 30_000; +const sockets: FakeGoogleLiveWebSocket[] = []; +const audioContexts: FakeAudioContext[] = []; +let stopInputTrack: ReturnType; + +class FakeGoogleLiveWebSocket extends EventTarget { + static OPEN = 1; + + readyState = FakeGoogleLiveWebSocket.OPEN; + binaryType: BinaryType = "blob"; + + constructor(readonly url: string) { + super(); + sockets.push(this); + } + + send(): void {} + + close(): void { + this.readyState = 3; + } + + emitOpen(): void { + this.dispatchEvent(new Event("open")); + } + + emitMessage(message: unknown): void { + this.dispatchEvent(new MessageEvent("message", { data: JSON.stringify(message) })); + } + + emitClose(): void { + this.dispatchEvent(new Event("close")); + } + + emitError(): void { + this.dispatchEvent(new Event("error")); + } +} + +class FakeAudioContext { + readonly currentTime = 0; + readonly destination = {}; + readonly sampleRate: number; + readonly close = vi.fn(async () => undefined); + + constructor(options?: { sampleRate?: number }) { + this.sampleRate = options?.sampleRate ?? 24_000; + audioContexts.push(this); + } + + createMediaStreamSource() { + return { connect() {}, disconnect() {} }; + } + + createScriptProcessor() { + return { connect() {}, disconnect() {}, onaudioprocess: null }; + } + + createGain() { + return { connect() {}, disconnect() {}, gain: { value: 1 } }; + } + + createAnalyser() { + return { + fftSize: 0, + smoothingTimeConstant: 0, + disconnect() {}, + getFloatTimeDomainData: (samples: Float32Array) => samples.fill(0.25), + }; + } +} + +function createSession(): RealtimeTalkJsonPcmWebSocketSessionResult { + return { + provider: "google", + transport: "provider-websocket", + protocol: "google-live-bidi", + clientSecret: ["auth_tokens", "browser-timeout-test"].join("/"), + websocketUrl: + "wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.v1alpha.GenerativeService.BidiGenerateContentConstrained", + audio: { + inputEncoding: "pcm16", + inputSampleRateHz: 16_000, + outputEncoding: "pcm16", + outputSampleRateHz: 24_000, + }, + }; +} + +function createTransport(callbacks: RealtimeTalkCallbacks = {}) { + return new GoogleLiveRealtimeTalkTransport(createSession(), { + callbacks, + client: { request: vi.fn(), addEventListener: vi.fn() } as never, + sessionKey: "main", + }); +} + +function latestSocket(): FakeGoogleLiveWebSocket { + const socket = sockets.at(-1); + if (!socket) { + throw new Error("missing Google Live WebSocket"); + } + return socket; +} + +async function beginTransport(transport: GoogleLiveRealtimeTalkTransport): Promise<{ + start: Promise<"ready" | "cancelled">; + socket: FakeGoogleLiveWebSocket; +}> { + const start = transport.start(); + await vi.advanceTimersByTimeAsync(0); + return { start, socket: latestSocket() }; +} + +describe("Google Live setup timeout", () => { + beforeEach(() => { + vi.useFakeTimers(); + sockets.length = 0; + audioContexts.length = 0; + stopInputTrack = vi.fn(); + vi.stubGlobal("WebSocket", FakeGoogleLiveWebSocket); + vi.stubGlobal("AudioContext", FakeAudioContext); + vi.stubGlobal("navigator", { + mediaDevices: { + getUserMedia: vi.fn(async () => ({ + getTracks: () => [{ stop: stopInputTrack }], + })), + }, + }); + }); + + afterEach(() => { + vi.useRealTimers(); + vi.unstubAllGlobals(); + }); + + it("releases browser resources when the WebSocket never opens", async () => { + const onStatus = vi.fn(); + const onTalkEvent = vi.fn(); + const transport = createTransport({ onStatus, onTalkEvent }); + + const { start, socket } = await beginTransport(transport); + socket.readyState = 0; + const rejected = expect(start).rejects.toThrow("Realtime connection timed out after 30000ms"); + await vi.advanceTimersByTimeAsync(SETUP_TIMEOUT_MS); + + await rejected; + expect(onStatus).not.toHaveBeenCalled(); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(audioContexts).toHaveLength(2); + expect(audioContexts.every((context) => context.close.mock.calls.length === 1)).toBe(true); + expect(socket.readyState).toBe(3); + expect(onTalkEvent).not.toHaveBeenCalled(); + expect(vi.getTimerCount()).toBe(0); + }); + + it("times out an open socket that never completes Google setup", async () => { + const onStatus = vi.fn(); + const transport = createTransport({ onStatus }); + + const { start, socket } = await beginTransport(transport); + socket.emitOpen(); + const rejected = expect(start).rejects.toThrow("Realtime connection timed out after 30000ms"); + await vi.advanceTimersByTimeAsync(SETUP_TIMEOUT_MS); + + await rejected; + expect(onStatus).not.toHaveBeenCalled(); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(audioContexts.every((context) => context.close.mock.calls.length === 1)).toBe(true); + }); + + it("does not publish provisional terminal callbacks when setup times out", async () => { + const onStatus = vi.fn(() => { + throw new Error("status callback must remain provisional"); + }); + const onTalkEvent = vi.fn(() => { + throw new Error("talk callback must remain provisional"); + }); + const transport = createTransport({ onStatus, onTalkEvent }); + + const { start, socket } = await beginTransport(transport); + socket.readyState = 0; + const rejected = expect(start).rejects.toThrow("Realtime connection timed out after 30000ms"); + await vi.advanceTimersByTimeAsync(SETUP_TIMEOUT_MS); + + await rejected; + expect(onStatus).not.toHaveBeenCalled(); + expect(onTalkEvent).not.toHaveBeenCalled(); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(audioContexts.every((context) => context.close.mock.calls.length === 1)).toBe(true); + expect(socket.readyState).toBe(3); + expect(vi.getTimerCount()).toBe(0); + }); + + it.each([ + ["close", "Realtime connection closed"], + ["error", "Realtime connection failed"], + ] as const)("rejects startup when the WebSocket emits %s", async (event, detail) => { + const onStatus = vi.fn(); + const transport = createTransport({ onStatus }); + const { start, socket } = await beginTransport(transport); + const rejected = expect(start).rejects.toThrow(detail); + + if (event === "close") { + socket.emitClose(); + } else { + socket.emitError(); + } + + await rejected; + expect(onStatus).not.toHaveBeenCalled(); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(socket.readyState).toBe(3); + expect(vi.getTimerCount()).toBe(0); + }); + + it("rejects when the socket closes after setup but before activation", async () => { + const onStatus = vi.fn(); + const transport = createTransport({ onStatus }); + const { start, socket } = await beginTransport(transport); + socket.emitOpen(); + socket.emitMessage({ setupComplete: {} }); + await Promise.resolve(); + const rejected = expect(start).rejects.toThrow("Realtime connection closed"); + + socket.emitClose(); + + await rejected; + expect(onStatus).not.toHaveBeenCalled(); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(socket.readyState).toBe(3); + expect(vi.getTimerCount()).toBe(0); + }); + + it("releases resources when a readiness callback throws during activation", async () => { + const onStatus = vi.fn(() => { + throw new Error("consumer failed"); + }); + const transport = createTransport({ onStatus }); + const { start, socket } = await beginTransport(transport); + socket.emitOpen(); + socket.emitMessage({ setupComplete: {} }); + await expect(start).resolves.toBe("ready"); + + expect(() => transport.activate()).toThrow("consumer failed"); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(socket.readyState).toBe(3); + expect(audioContexts.every((context) => context.close.mock.calls.length === 1)).toBe(true); + }); + + it("reclaims the meter when an input-level callback cancels activation", async () => { + let stopDuringActivation: () => void = () => undefined; + const onInputLevel = vi.fn(() => stopDuringActivation()); + const transport = createTransport({ onInputLevel }); + stopDuringActivation = () => transport.stop({ emitClosed: false }); + const { start, socket } = await beginTransport(transport); + socket.emitOpen(); + socket.emitMessage({ setupComplete: {} }); + await expect(start).resolves.toBe("ready"); + + expect(() => transport.activate()).toThrow("Google Live transport activation cancelled"); + expect(vi.getTimerCount()).toBe(0); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(socket.readyState).toBe(3); + expect(audioContexts.every((context) => context.close.mock.calls.length === 1)).toBe(true); + }); + + it("clears the deadline after Google setup completes", async () => { + const onStatus = vi.fn(); + const transport = createTransport({ onStatus }); + + const { start, socket } = await beginTransport(transport); + socket.emitOpen(); + socket.emitMessage({ setupComplete: {} }); + await expect(start).resolves.toBe("ready"); + expect(onStatus).not.toHaveBeenCalled(); + transport.activate(); + await vi.advanceTimersByTimeAsync(SETUP_TIMEOUT_MS); + + expect(onStatus).toHaveBeenCalledExactlyOnceWith("listening"); + transport.stop(); + }); + + it("clears the deadline when the transport stops", async () => { + const onStatus = vi.fn(); + const transport = createTransport({ onStatus }); + + const { start } = await beginTransport(transport); + transport.stop(); + await expect(start).resolves.toBe("cancelled"); + await vi.advanceTimersByTimeAsync(SETUP_TIMEOUT_MS); + + expect(onStatus).not.toHaveBeenCalled(); + expect(vi.getTimerCount()).toBe(0); + }); +}); diff --git a/ui/src/pages/chat/realtime-talk-google-live-video.test.ts b/ui/src/pages/chat/realtime-talk-google-live-video.test.ts index 563034745e7f..ab672650e708 100644 --- a/ui/src/pages/chat/realtime-talk-google-live-video.test.ts +++ b/ui/src/pages/chat/realtime-talk-google-live-video.test.ts @@ -84,6 +84,30 @@ function createTransport(callbacks: RealtimeTalkCallbacks, videoDeviceId?: strin ); } +async function beginTransport(transport: GoogleLiveRealtimeTalkTransport): Promise<{ + start: Promise<"ready" | "cancelled">; + ws: FakeGoogleLiveWebSocket; +}> { + const start = transport.start(); + await vi.advanceTimersByTimeAsync(0); + const ws = FakeGoogleLiveWebSocket.instance; + if (!ws) { + throw new Error("missing Google Live WebSocket"); + } + return { start, ws }; +} + +async function startTransport( + transport: GoogleLiveRealtimeTalkTransport, +): Promise { + const { start, ws } = await beginTransport(transport); + ws.emitOpen(); + ws.emitMessage({ setupComplete: {} }); + await expect(start).resolves.toBe("ready"); + transport.activate(); + return ws; +} + describe("Google Live Video Talk", () => { beforeEach(() => { vi.useFakeTimers(); @@ -142,16 +166,15 @@ describe("Google Live Video Talk", () => { const onVideoStream = vi.fn(); const transport = createTransport({ onStatus, onVideoStream }); - await transport.start(); + const { start, ws } = await beginTransport(transport); expect(getUserMedia).toHaveBeenCalledOnce(); expect(onVideoStream).not.toHaveBeenCalled(); await transport.setVideoEnabled(true); - const ws = FakeGoogleLiveWebSocket.instance; - if (!ws) { - throw new Error("missing Google Live WebSocket"); - } ws.emitOpen(); ws.emitMessage({ setupComplete: {} }); + await expect(start).resolves.toBe("ready"); + expect(ws.sent.some((message) => JSON.stringify(message).includes('"video"'))).toBe(false); + transport.activate(); await vi.advanceTimersByTimeAsync(0); expect(ws.sent).toContainEqual({ @@ -273,7 +296,7 @@ describe("Google Live Video Talk", () => { const onVideoStream = vi.fn(); const transport = createTransport({ onVideoStream }); - await transport.start(); + await startTransport(transport); await transport.setVideoEnabled(true); firstVideoTrack.dispatchEvent(new Event("ended")); @@ -304,7 +327,7 @@ describe("Google Live Video Talk", () => { vi.stubGlobal("navigator", { mediaDevices: { getUserMedia } }); const transport = createTransport({}); - await transport.start(); + await startTransport(transport); const enabling = transport.setVideoEnabled(true); await vi.waitFor(() => expect(getUserMedia).toHaveBeenCalledTimes(2)); transport.stop(); @@ -316,6 +339,45 @@ describe("Google Live Video Talk", () => { expect(FakeGoogleLiveWebSocket.instance?.readyState).toBe(3); }); + it("finishes active camera cleanup when the stream callback throws", async () => { + const audioStop = vi.fn(); + const videoStop = vi.fn(); + const audioTrack = { stop: audioStop } as unknown as MediaStreamTrack; + const videoTrack = Object.assign(new EventTarget(), { + stop: videoStop, + readyState: "live", + enabled: true, + muted: false, + }) as unknown as MediaStreamTrack; + const audio = { + getAudioTracks: () => [audioTrack], + getTracks: () => [audioTrack], + } as unknown as MediaStream; + const camera = { + getVideoTracks: () => [videoTrack], + getTracks: () => [videoTrack], + } as unknown as MediaStream; + vi.stubGlobal("navigator", { + mediaDevices: { + getUserMedia: vi.fn().mockResolvedValueOnce(audio).mockResolvedValueOnce(camera), + }, + }); + vi.spyOn(HTMLMediaElement.prototype, "play").mockResolvedValue(undefined); + const onVideoStream = vi.fn((stream: MediaStream | null) => { + if (!stream) { + throw new Error("stream callback failed"); + } + }); + const transport = createTransport({ onVideoStream }); + const ws = await startTransport(transport); + await transport.setVideoEnabled(true); + + expect(() => transport.stop()).toThrow("stream callback failed"); + expect(audioStop).toHaveBeenCalledOnce(); + expect(videoStop).toHaveBeenCalledOnce(); + expect(ws.readyState).toBe(3); + }); + it("switches an active camera and keeps video frame capture running", async () => { const audioTrack = { stop: vi.fn() } as unknown as MediaStreamTrack; const frontStop = vi.fn(); @@ -355,7 +417,7 @@ describe("Google Live Video Talk", () => { const onVideoStream = vi.fn(); const transport = createTransport({ onVideoStream }, "front"); - await transport.start(); + await startTransport(transport); await transport.setVideoEnabled(true); await transport.switchCamera("back"); @@ -370,4 +432,51 @@ describe("Google Live Video Talk", () => { transport.stop(); }); + + it("releases camera media when Google setup times out", async () => { + const audioStop = vi.fn(); + const videoStop = vi.fn(); + const audioTrack = { stop: audioStop } as unknown as MediaStreamTrack; + const videoTrack = Object.assign(new EventTarget(), { + stop: videoStop, + readyState: "live", + enabled: true, + muted: false, + }) as unknown as MediaStreamTrack; + const audio = { + getAudioTracks: () => [audioTrack], + getTracks: () => [audioTrack], + } as unknown as MediaStream; + const camera = { + getVideoTracks: () => [videoTrack], + getTracks: () => [videoTrack], + } as unknown as MediaStream; + const getUserMedia = vi.fn().mockResolvedValueOnce(audio).mockResolvedValueOnce(camera); + vi.stubGlobal("navigator", { mediaDevices: { getUserMedia } }); + const originalCreateElement = document.createElement.bind(document); + vi.spyOn(document, "createElement").mockImplementation((tagName: string) => { + const element = originalCreateElement(tagName); + if (element instanceof HTMLVideoElement) { + vi.spyOn(element, "play").mockResolvedValue(undefined); + } + return element; + }); + const onStatus = vi.fn(); + const onVideoStream = vi.fn(); + const transport = createTransport({ onStatus, onVideoStream }); + + const { start, ws } = await beginTransport(transport); + await transport.setVideoEnabled(true); + ws.emitOpen(); + const rejected = expect(start).rejects.toThrow("Realtime connection timed out after 30000ms"); + await vi.advanceTimersByTimeAsync(30_000); + + await rejected; + expect(onStatus).not.toHaveBeenCalled(); + expect(audioStop).toHaveBeenCalledOnce(); + expect(videoStop).toHaveBeenCalledOnce(); + expect(getUserMedia).toHaveBeenNthCalledWith(2, { video: true }); + expect(onVideoStream).not.toHaveBeenCalled(); + expect(ws.readyState).toBe(3); + }); }); diff --git a/ui/src/pages/chat/realtime-talk-google-live.test.ts b/ui/src/pages/chat/realtime-talk-google-live.test.ts index e59a073caf2e..f69e9c2cb7f5 100644 --- a/ui/src/pages/chat/realtime-talk-google-live.test.ts +++ b/ui/src/pages/chat/realtime-talk-google-live.test.ts @@ -223,6 +223,26 @@ function latestWebSocket(): MockGoogleLiveWebSocket { return ws; } +async function beginTransport(transport: GoogleLiveRealtimeTalkTransport): Promise<{ + start: Promise<"ready" | "cancelled">; + ws: MockGoogleLiveWebSocket; +}> { + const start = transport.start(); + await waitForFast(() => expect(wsInstances).toHaveLength(1)); + return { start, ws: latestWebSocket() }; +} + +async function startTransport( + transport: GoogleLiveRealtimeTalkTransport, +): Promise { + const { start, ws } = await beginTransport(transport); + ws.emitOpen(); + ws.emitMessage(encodeJsonFrame({ setupComplete: {} })); + await expect(start).resolves.toBe("ready"); + transport.activate(); + return ws; +} + function pumpMicrophone(samples: Float32Array): void { const processor = inputProcessors.at(-1); if (!processor) { @@ -275,12 +295,13 @@ describe("GoogleLiveRealtimeTalkTransport", () => { { callbacks: {}, client: createClient(), sessionKey: "main" }, ); - await transport.start(); + const { start } = await beginTransport(transport); expect(latestWebSocket().url).toBe( "wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.v1alpha.GenerativeService.BidiGenerateContentConstrained?access_token=auth_tokens%2Fbrowser-session", ); transport.stop(); + await expect(start).resolves.toBe("cancelled"); }); it.each([ @@ -301,7 +322,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { it("captures from the selected microphone with an exact constraint", async () => { const transport = createTransport({}, createClient(), "usb-mic"); - await transport.start(); + const { start } = await beginTransport(transport); expect(getUserMedia).toHaveBeenCalledWith({ audio: { @@ -312,13 +333,13 @@ describe("GoogleLiveRealtimeTalkTransport", () => { }, }); transport.stop(); + await expect(start).resolves.toBe("cancelled"); }); it("keeps the microphone processor inaudible locally", async () => { const transport = createTransport(); - await transport.start(); - latestWebSocket().emitOpen(); + await startTransport(transport); const processor = inputProcessors.at(-1); const sink = inputSinks.at(-1); @@ -359,12 +380,18 @@ describe("GoogleLiveRealtimeTalkTransport", () => { const onTalkEvent = vi.fn(); const transport = createTransport({ onStatus, onTalkEvent }); - await transport.start(); - const ws = latestWebSocket(); + const { start, ws } = await beginTransport(transport); + ws.emitOpen(); ws.emitMessage(encodeJsonFrame({ setupComplete: {} })); + await expect(start).resolves.toBe("ready"); expect(ws.binaryType).toBe("arraybuffer"); - await waitForFast(() => expect(onStatus).toHaveBeenCalledWith("listening")); + expect(onStatus).not.toHaveBeenCalled(); + expect(onTalkEvent).not.toHaveBeenCalled(); + transport.activate(); + transport.activate(); + expect(onStatus).toHaveBeenCalledWith("listening"); + expect(onStatus).toHaveBeenCalledOnce(); const readyEvent = requireFirstTalkEvent(onTalkEvent); expect(readyEvent.type).toBe("session.ready"); expect(readyEvent.sessionId).toBe("main:google:provider-websocket"); @@ -376,9 +403,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { const onTranscript = vi.fn(); const transport = createTransport({ onStatus, onTranscript }); - await expect(transport.start()).resolves.toBe("ready"); - const ws = latestWebSocket(); - ws.emitOpen(); + const ws = await startTransport(transport); pumpMicrophone(new Float32Array(4096)); ws.emitClose(); @@ -406,9 +431,8 @@ describe("GoogleLiveRealtimeTalkTransport", () => { const onTranscript = vi.fn(); const transport = createTransport({ onStatus, onTranscript }); - await expect(transport.start()).resolves.toBe("ready"); - const ws = latestWebSocket(); - ws.emitOpen(); + const ws = await startTransport(transport); + onStatus.mockClear(); ws.emitError(); ws.emitClose(); @@ -429,18 +453,25 @@ describe("GoogleLiveRealtimeTalkTransport", () => { it.each(["status", "talk event"] as const)( "releases socket resources when the terminal %s callback throws", async (callbackKind) => { - const throwingCallback = vi.fn(() => { - throw new Error("consumer failed"); - }); const transport = createTransport( callbackKind === "status" - ? { onStatus: throwingCallback } - : { onTalkEvent: throwingCallback }, + ? { + onStatus: vi.fn((status) => { + if (status === "error") { + throw new Error("consumer failed"); + } + }), + } + : { + onTalkEvent: vi.fn((event) => { + if (event.type === "session.closed") { + throw new Error("consumer failed"); + } + }), + }, ); - await expect(transport.start()).resolves.toBe("ready"); - const ws = latestWebSocket(); - ws.emitOpen(); + const ws = await startTransport(transport); expect(() => ws.emitError()).toThrow("consumer failed"); expect(stopInputTrack).toHaveBeenCalledOnce(); @@ -452,12 +483,30 @@ describe("GoogleLiveRealtimeTalkTransport", () => { }, ); + it("finishes cleanup when the input-level callback throws during activation", async () => { + const onInputLevel = vi.fn(() => { + throw new Error("meter callback failed"); + }); + const transport = createTransport({ onInputLevel }); + const { start, ws } = await beginTransport(transport); + ws.emitOpen(); + ws.emitMessage(encodeJsonFrame({ setupComplete: {} })); + await expect(start).resolves.toBe("ready"); + + expect(() => transport.activate()).toThrow("meter callback failed"); + expect(stopInputTrack).toHaveBeenCalledOnce(); + expect(inputProcessors).toHaveLength(0); + expect(ws.readyState).toBe(3); + for (const context of audioContexts) { + expect(context.close).toHaveBeenCalledOnce(); + } + }); + it("reports microphone activity and resets it when stopped", async () => { const onInputLevel = vi.fn(); const transport = createTransport({ onInputLevel }); - await transport.start(); - latestWebSocket().emitOpen(); + await startTransport(transport); pumpMicrophone(new Float32Array(4096)); pumpMicrophone(new Float32Array(4096).fill(0.25)); transport.stop(); @@ -470,8 +519,11 @@ describe("GoogleLiveRealtimeTalkTransport", () => { const onStatus = vi.fn(); const transport = createTransport({ onStatus }); - await transport.start(); - latestWebSocket().emitMessage(new Blob([JSON.stringify({ setupComplete: {} })])); + const { start, ws } = await beginTransport(transport); + ws.emitOpen(); + ws.emitMessage(new Blob([JSON.stringify({ setupComplete: {} })])); + await expect(start).resolves.toBe("ready"); + transport.activate(); await waitForFast(() => expect(onStatus).toHaveBeenCalledWith("listening")); }); @@ -479,8 +531,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { it("stops queued output when Google Live sends interruption", async () => { const onTalkEvent = vi.fn(); const transport = createTransport({ onTalkEvent }); - await transport.start(); - const ws = latestWebSocket(); + const ws = await startTransport(transport); ws.emitMessage( encodeJsonFrame({ @@ -508,8 +559,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { const onStatus = vi.fn(); const onTalkEvent = vi.fn(); const transport = createTransport({ onStatus, onTalkEvent }); - await transport.start(); - const ws = latestWebSocket(); + const ws = await startTransport(transport); ws.emitMessage( encodeJsonFrame({ @@ -557,8 +607,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { it("rejects an oversized first frame before decoding provider audio", async () => { const onStatus = vi.fn(); const transport = createTransport({ onStatus }); - await transport.start(); - const ws = latestWebSocket(); + const ws = await startTransport(transport); ws.emitMessage( encodeJsonFrame({ @@ -592,8 +641,9 @@ describe("GoogleLiveRealtimeTalkTransport", () => { const onTalkEvent = vi.fn(); const transport = createTransport({ onTalkEvent, onTranscript }); - await transport.start(); - latestWebSocket().emitMessage( + const ws = await startTransport(transport); + onTalkEvent.mockClear(); + ws.emitMessage( encodeJsonFrame({ serverContent: { inputTranscription: { text: "hello", finished: true }, @@ -638,8 +688,9 @@ describe("GoogleLiveRealtimeTalkTransport", () => { const onTranscript = vi.fn(() => transport.stop()); const transport = createTransport({ onTalkEvent, onTranscript }); - await transport.start(); - latestWebSocket().emitMessage( + const ws = await startTransport(transport); + onTalkEvent.mockClear(); + ws.emitMessage( encodeJsonFrame({ serverContent: { inputTranscription: { text: "overflow", finished: true }, @@ -657,22 +708,22 @@ describe("GoogleLiveRealtimeTalkTransport", () => { it("silently disposes a provisional Google Live transport", async () => { const onTalkEvent = vi.fn(); const transport = createTransport({ onTalkEvent }); - await transport.start(); - onTalkEvent.mockClear(); + const { start, ws } = await beginTransport(transport); transport.stop({ emitClosed: false }); + await expect(start).resolves.toBe("cancelled"); expect(onTalkEvent).not.toHaveBeenCalled(); - expect(latestWebSocket().readyState).toBe(3); + expect(ws.readyState).toBe(3); }); it("ignores late WebSocket events after stop", async () => { const onStatus = vi.fn(); const transport = createTransport({ onStatus }); - await transport.start(); - const ws = latestWebSocket(); + const { start, ws } = await beginTransport(transport); transport.stop(); + await expect(start).resolves.toBe("cancelled"); ws.emitOpen(); ws.emitMessage(new Blob([JSON.stringify({ setupComplete: {} })])); @@ -702,9 +753,10 @@ describe("GoogleLiveRealtimeTalkTransport", () => { }), } as unknown as RealtimeTalkTransportContext["client"]; const transport = createTransport({ onStatus }, client); - await transport.start(); + const ws = await startTransport(transport); + onStatus.mockClear(); - latestWebSocket().emitMessage( + ws.emitMessage( encodeJsonFrame({ toolCall: { functionCalls: [ @@ -744,8 +796,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { }), } as unknown as RealtimeTalkTransportContext["client"]; const transport = createTransport({}, client); - await transport.start(); - const ws = latestWebSocket(); + const ws = await startTransport(transport); ws.emitMessage( encodeJsonFrame({ @@ -802,8 +853,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { }); const transport = createTransport({ onStatus, onTalkEvent }, client); - await transport.start(); - const ws = latestWebSocket(); + const ws = await startTransport(transport); vi.spyOn(ws, "send").mockImplementation(() => { throw new Error("Google Live socket rejected the tool result"); }); @@ -868,9 +918,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { throw new Error(`unexpected request: ${method}`); }); const transport = createTransport({}, client); - await transport.start(); - const ws = latestWebSocket(); - ws.emitOpen(); + const ws = await startTransport(transport); ws.emitMessage( encodeJsonFrame({ serverContent: { @@ -941,9 +989,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { throw new Error(`unexpected request: ${method}`); }); const transport = createTransport({}, client); - await transport.start(); - const ws = latestWebSocket(); - ws.emitOpen(); + const ws = await startTransport(transport); ws.emitMessage( encodeJsonFrame({ serverContent: { @@ -1014,9 +1060,7 @@ describe("GoogleLiveRealtimeTalkTransport", () => { throw new Error(`unexpected request: ${method}`); }); const transport = createTransport({}, client); - await transport.start(); - const ws = latestWebSocket(); - ws.emitOpen(); + const ws = await startTransport(transport); ws.emitMessage( encodeJsonFrame({ serverContent: { diff --git a/ui/src/pages/chat/realtime-talk-google-live.ts b/ui/src/pages/chat/realtime-talk-google-live.ts index 3ba479910bb0..d507c45ec251 100644 --- a/ui/src/pages/chat/realtime-talk-google-live.ts +++ b/ui/src/pages/chat/realtime-talk-google-live.ts @@ -9,6 +9,12 @@ import { RealtimeTalkPcmOutputQueue, } from "./realtime-talk-audio.ts"; import { RealtimeTalkCameraController } from "./realtime-talk-camera-controller.ts"; +import { + buildGoogleLiveUrl, + GoogleLiveConnectionLifecycle, + GOOGLE_LIVE_SETUP_TIMEOUT_MS, + runRealtimeTalkCleanup, +} from "./realtime-talk-google-live-lifecycle.ts"; import { openRealtimeTalkCamera, openRealtimeTalkInput } from "./realtime-talk-input.ts"; import type { RealtimeTalkJsonPcmWebSocketSessionResult } from "./realtime-talk-shared.ts"; import { @@ -57,9 +63,6 @@ type PendingFunctionCall = { args: unknown; }; -const GOOGLE_LIVE_WEBSOCKET_HOST = "generativelanguage.googleapis.com"; -const GOOGLE_LIVE_WEBSOCKET_PATH = - /^\/ws\/google\.ai\.generativelanguage\.v[0-9a-z]+\.GenerativeService\.BidiGenerateContent(?:Constrained)?$/; const GOOGLE_LIVE_VIDEO_FRAME_INTERVAL_MS = 1_000; const GOOGLE_LIVE_VIDEO_MESSAGE_MAX_BYTES = 512 * 1024; @@ -81,32 +84,9 @@ function isGemini31LiveModel(model: string | undefined): boolean { return modelId.startsWith("gemini-3.1-") && modelId.includes("-live"); } -function buildGoogleLiveUrl(session: RealtimeTalkJsonPcmWebSocketSessionResult): string { - let url: URL; - try { - url = new URL(session.websocketUrl); - } catch { - throw new Error("Invalid Google Live WebSocket URL"); - } - if (url.protocol !== "wss:") { - throw new Error("Google Live WebSocket URL must use wss://"); - } - if (url.hostname.toLowerCase() !== GOOGLE_LIVE_WEBSOCKET_HOST) { - throw new Error("Untrusted Google Live WebSocket host"); - } - if (url.username || url.password) { - throw new Error("Google Live WebSocket URL must not include credentials"); - } - if (!GOOGLE_LIVE_WEBSOCKET_PATH.test(url.pathname)) { - throw new Error("Untrusted Google Live WebSocket path"); - } - url.search = ""; - url.searchParams.set("access_token", session.clientSecret); - return url.toString(); -} - export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { private ws: WebSocket | null = null; + private setupTimeout: ReturnType | null = null; private media: MediaStream | null = null; private inputContext: AudioContext | null = null; private outputContext: AudioContext | null = null; @@ -115,7 +95,8 @@ export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { private closed = false; private mediaSetupController: AbortController | null = null; private readonly camera: RealtimeTalkCameraController; - private setupComplete = false; + private readonly lifecycle = new GoogleLiveConnectionLifecycle(); + private cameraPublished = false; private videoFramesActive = false; private hasSentVideoFrame = false; private videoFrameTimer: ReturnType | null = null; @@ -134,9 +115,20 @@ export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { getDeviceId: () => this.ctx.videoDeviceId, setDeviceId: (deviceId) => (this.ctx.videoDeviceId = deviceId), isClosed: () => this.closed, - onStream: (stream) => this.ctx.callbacks.onVideoStream?.(stream), + onStream: (stream) => { + if (stream) { + if (!this.lifecycle.isActive) { + return; + } + this.cameraPublished = true; + this.ctx.callbacks.onVideoStream?.(stream); + } else if (this.cameraPublished) { + this.cameraPublished = false; + this.ctx.callbacks.onVideoStream?.(null); + } + }, onAcquired: () => { - if (this.setupComplete) { + if (this.lifecycle.isActive) { this.startVideoFrames(); } }, @@ -153,6 +145,7 @@ export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { } const wsUrl = buildGoogleLiveUrl(this.session); this.closed = false; + this.cameraPublished = false; this.mediaSetupController?.abort(); const mediaSetupController = new AbortController(); this.mediaSetupController = mediaSetupController; @@ -178,25 +171,79 @@ export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { this.media = media; this.inputContext = new AudioContext({ sampleRate: this.session.audio.inputSampleRateHz }); this.outputContext = new AudioContext({ sampleRate: this.session.audio.outputSampleRateHz }); - if (this.ctx.callbacks.onInputLevel) { - this.inputMeter = new RealtimeTalkMediaStreamMeter(this.ctx.callbacks.onInputLevel); - this.inputMeter.start(this.media, this.inputContext); - } - this.ws = new WebSocket(wsUrl); - this.ws.binaryType = "arraybuffer"; - this.ws.addEventListener("open", () => { - if (this.closed) { + const ws = new WebSocket(wsUrl); + this.ws = ws; + ws.binaryType = "arraybuffer"; + const startup = this.lifecycle.begin(ws); + this.setupTimeout = globalThis.setTimeout(() => { + if (this.closed || this.ws !== ws) { + return; + } + this.setupTimeout = null; + this.failConnection( + ws, + `Realtime connection timed out after ${GOOGLE_LIVE_SETUP_TIMEOUT_MS}ms`, + ); + }, GOOGLE_LIVE_SETUP_TIMEOUT_MS); + ws.addEventListener("open", () => { + if (this.closed || this.ws !== ws) { return; } this.send(this.session.initialMessage ?? { setup: {} }); + }); + ws.addEventListener("message", (event) => { + void this.handleMessage(ws, event.data); + }); + ws.addEventListener("close", () => { + this.failConnection(ws, "Realtime connection closed"); + }); + ws.addEventListener("error", () => { + this.failConnection(ws, "Realtime connection failed"); + }); + return this.lifecycle.finishStart(await startup); + } + + activate(): void { + if (this.closed || !this.lifecycle.activate()) { + return; + } + try { + this.ctx.callbacks.onStatus?.("listening"); + this.assertActivationCurrent(); + this.emitTalkEvent({ type: "session.ready" }); + this.assertActivationCurrent(); + if (this.ctx.callbacks.onInputLevel && this.media && this.inputContext) { + const inputMeter = new RealtimeTalkMediaStreamMeter(this.ctx.callbacks.onInputLevel); + this.inputMeter = inputMeter; + inputMeter.start(this.media, this.inputContext); + if (this.closed || !this.lifecycle.isActive || this.inputMeter !== inputMeter) { + // start() publishes synchronously before installing its interval. A + // reentrant stop must reclaim the interval that start() installs next. + inputMeter.stop(false); + } + this.assertActivationCurrent(); + } this.startMicrophonePump(); - }); - this.ws.addEventListener("message", (event) => { - void this.handleMessage(event.data); - }); - this.ws.addEventListener("close", () => this.failConnection("Realtime connection closed")); - this.ws.addEventListener("error", () => this.failConnection("Realtime connection failed")); - return "ready"; + if (this.camera.stream && !this.cameraPublished) { + this.cameraPublished = true; + this.ctx.callbacks.onVideoStream?.(this.camera.stream); + } + this.assertActivationCurrent(); + this.startVideoFrames(); + } catch (error) { + try { + this.stop({ emitClosed: false }); + } catch { + // Preserve the activation callback as the terminal cause after cleanup. + } + throw error; + } + } + + private assertActivationCurrent(): void { + if (this.closed || !this.lifecycle.isActive) { + throw new Error("Google Live transport activation cancelled"); + } } async setVideoEnabled(enabled: boolean): Promise { @@ -208,43 +255,73 @@ export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { } stop(options?: { emitClosed?: boolean }): void { - const emitClosed = !this.closed && options?.emitClosed !== false; + const emitClosed = !this.closed && this.lifecycle.isActive && options?.emitClosed !== false; this.closed = true; - try { - if (emitClosed) { - this.emitTalkEvent({ type: "session.closed", final: true }); - } - } finally { - this.releaseResources(); - } + this.lifecycle.cancel(); + runRealtimeTalkCleanup([ + () => { + if (emitClosed) { + this.emitTalkEvent({ type: "session.closed", final: true }); + } + }, + () => this.releaseResources(), + ]); } private releaseResources(): void { - this.mediaSetupController?.abort(); + const mediaSetupController = this.mediaSetupController; this.mediaSetupController = null; - this.setupComplete = false; - for (const controller of this.consultAbortControllers) { - controller.abort(); - } + this.clearSetupTimeout(); + const consultAbortControllers = [...this.consultAbortControllers]; this.consultAbortControllers.clear(); this.pendingCalls.clear(); - this.inputPump.stop(); - this.inputMeter?.stop(); + const inputMeter = this.inputMeter; this.inputMeter = null; - this.media?.getTracks().forEach((track) => track.stop()); + const media = this.media; this.media = null; - this.camera.release(); - this.stopOutput(); - void this.inputContext?.close(); + const inputContext = this.inputContext; this.inputContext = null; - void this.outputContext?.close(); + const outputContext = this.outputContext; this.outputContext = null; - this.ws?.close(); + const ws = this.ws; this.ws = null; + runRealtimeTalkCleanup([ + () => mediaSetupController?.abort(), + ...consultAbortControllers.map((controller) => () => controller.abort()), + () => this.inputPump.stop(), + () => inputMeter?.stop(), + ...(media?.getTracks() ?? []).map((track) => () => track.stop()), + () => this.camera.release(), + () => this.stopOutput(), + () => { + void inputContext?.close(); + }, + () => { + void outputContext?.close(); + }, + () => ws?.close(), + ]); } - private failConnection(detail: string): void { - if (this.closed) { + private clearSetupTimeout(): void { + if (this.setupTimeout !== null) { + globalThis.clearTimeout(this.setupTimeout); + this.setupTimeout = null; + } + } + + private failConnection(ws: WebSocket, detail: string): void { + // Error and close can arrive for the same socket, or after a replacement + // starts. Only the current lifecycle owner may report and release resources. + if (this.closed || this.ws !== ws) { + return; + } + if (this.lifecycle.failStartup(ws, new Error(detail))) { + try { + this.stop({ emitClosed: false }); + } catch { + // Startup rejection owns terminal precedence; cleanup still ran to completion. + } return; } try { @@ -283,8 +360,8 @@ export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { return false; } - private async handleMessage(data: unknown): Promise { - if (this.closed) { + private async handleMessage(ws: WebSocket, data: unknown): Promise { + if (this.closed || this.ws !== ws) { return; } let message: GoogleLiveMessage; @@ -293,14 +370,16 @@ export class GoogleLiveRealtimeTalkTransport implements RealtimeTalkTransport { } catch { return; } - if (this.closed) { + if (this.closed || this.ws !== ws) { return; } - if (message.setupComplete) { - this.setupComplete = true; - this.ctx.callbacks.onStatus?.("listening"); - this.emitTalkEvent({ type: "session.ready" }); - this.startVideoFrames(); + if (message.setupComplete && this.lifecycle.markReady(ws)) { + this.clearSetupTimeout(); + } + // The parent session adopts the candidate after start() resolves. Provider + // events remain provisional until activate() publishes that ownership. + if (!this.lifecycle.isActive) { + return; } const content = message.serverContent; if (content?.interrupted) {