diff --git a/src/gateway/server/ws-connection.test.ts b/src/gateway/server/ws-connection.test.ts index 92cc4a021cd7..2085b27570cb 100644 --- a/src/gateway/server/ws-connection.test.ts +++ b/src/gateway/server/ws-connection.test.ts @@ -15,12 +15,16 @@ const { attachGatewayWsMessageHandlerMock, attachWorkerWsMessageHandlerMock, broadcastPresenceSnapshotMock, + closeTalkRealtimeRelaySessionsForConnectionMock, + closeTalkTranscriptionRelaySessionsForConnectionMock, touchPresenceMock, upsertPresenceMock, } = vi.hoisted(() => ({ attachGatewayWsMessageHandlerMock: vi.fn(), attachWorkerWsMessageHandlerMock: vi.fn((_params: unknown) => vi.fn()), broadcastPresenceSnapshotMock: vi.fn(), + closeTalkRealtimeRelaySessionsForConnectionMock: vi.fn(), + closeTalkTranscriptionRelaySessionsForConnectionMock: vi.fn(), touchPresenceMock: vi.fn(), upsertPresenceMock: vi.fn(), })); @@ -38,6 +42,13 @@ vi.mock("../../infra/system-presence.js", () => ({ vi.mock("./presence-events.js", () => ({ broadcastPresenceSnapshot: broadcastPresenceSnapshotMock, })); +vi.mock("../talk-realtime-relay.js", () => ({ + closeTalkRealtimeRelaySessionsForConnection: closeTalkRealtimeRelaySessionsForConnectionMock, +})); +vi.mock("../talk-transcription-relay.js", () => ({ + closeTalkTranscriptionRelaySessionsForConnection: + closeTalkTranscriptionRelaySessionsForConnectionMock, +})); import { attachGatewayWsConnectionHandler } from "./ws-connection.js"; import { resolveSharedGatewaySessionGeneration } from "./ws-shared-generation.js"; @@ -91,6 +102,8 @@ describe("attachGatewayWsConnectionHandler", () => { attachGatewayWsMessageHandlerMock.mockReset(); attachWorkerWsMessageHandlerMock.mockClear(); broadcastPresenceSnapshotMock.mockReset(); + closeTalkRealtimeRelaySessionsForConnectionMock.mockReset(); + closeTalkTranscriptionRelaySessionsForConnectionMock.mockReset(); touchPresenceMock.mockReset(); upsertPresenceMock.mockReset(); }); @@ -262,6 +275,33 @@ describe("attachGatewayWsConnectionHandler", () => { expect(socket.ping).toHaveBeenCalledOnce(); }); + it("releases connection-owned Talk relays when a gateway connection closes", async () => { + const { passed, socket } = await connectTestWs(); + const handlerParams = passed as { + connId: string; + setClient: (client: unknown) => boolean; + }; + expect( + handlerParams.setClient({ + socket, + connect: { client: { id: "openclaw-control-ui", mode: "webchat" } }, + connId: handlerParams.connId, + usesSharedGatewayAuth: false, + }), + ).toBe(true); + + socket.emit("close", 1000, Buffer.from("done")); + + expect(closeTalkRealtimeRelaySessionsForConnectionMock).toHaveBeenCalledOnce(); + expect(closeTalkRealtimeRelaySessionsForConnectionMock).toHaveBeenCalledWith( + handlerParams.connId, + ); + expect(closeTalkTranscriptionRelaySessionsForConnectionMock).toHaveBeenCalledOnce(); + expect(closeTalkTranscriptionRelaySessionsForConnectionMock).toHaveBeenCalledWith( + handlerParams.connId, + ); + }); + it("continues protocol pings after pong and stops when the connection closes", async () => { vi.useFakeTimers(); const socket = Object.assign(createGatewayWsTestSocket({ ping: true }), { diff --git a/src/gateway/server/ws-connection.ts b/src/gateway/server/ws-connection.ts index 6357d1b7a487..201ac8fec0ee 100644 --- a/src/gateway/server/ws-connection.ts +++ b/src/gateway/server/ws-connection.ts @@ -28,6 +28,8 @@ import { } from "../server-constants.js"; import type { GatewayRequestContext, GatewayRequestHandlers } from "../server-methods/types.js"; import { formatError } from "../server-utils.js"; +import { closeTalkRealtimeRelaySessionsForConnection } from "../talk-realtime-relay.js"; +import { closeTalkTranscriptionRelaySessionsForConnection } from "../talk-transcription-relay.js"; import { formatForLog, logWs } from "../ws-log.js"; import { getHealthVersion, incrementPresenceVersion } from "./health-state.js"; import type { PreauthConnectionBudget } from "./preauth-connection-budget.js"; @@ -554,6 +556,8 @@ export function attachGatewayWsConnectionHandler(params: AttachGatewayWsConnecti } if (connectionKind === "gateway") { const context = buildRequestContext(); + closeTalkRealtimeRelaySessionsForConnection(connId); + closeTalkTranscriptionRelaySessionsForConnection(connId); context.unsubscribeAllSessionEvents(connId); // Detach (or, with a zero grace period, kill) any PTY shells this // connection owned; detached sessions stay reattachable via diff --git a/src/gateway/talk-realtime-relay-operations.ts b/src/gateway/talk-realtime-relay-operations.ts index 97b0d977df1a..8f791f39608b 100644 --- a/src/gateway/talk-realtime-relay-operations.ts +++ b/src/gateway/talk-realtime-relay-operations.ts @@ -6,6 +6,7 @@ import { import { registerClientVoiceConsultRun } from "../talk/client-voice-session.js"; import type { RealtimeVoiceToolResultOptions } from "../talk/provider-types.js"; import { abortChatRunById } from "./chat-abort.js"; +import { formatError } from "./server-utils.js"; import { cancelForcedConsults, submitForcedTalkRealtimeRelayToolResult, @@ -40,6 +41,7 @@ import { import { decodeTalkRelayAudioBase64 } from "./talk-relay-audio-base64.js"; import { closeExpiredTalkRelaySessions, + closeTalkRelaySessionsForConnection, requireActiveTalkRelaySession, } from "./talk-relay-session-lifecycle.js"; import { forgetUnifiedTalkSession } from "./talk-session-registry.js"; @@ -93,17 +95,36 @@ export function closeRelaySession(session: RelaySession, reason: "completed" | " forgetUnifiedTalkSession(session.id); clearTimeout(session.cleanupTimer); abortRelayAgentRuns(session, reason === "error" ? "relay-error" : "relay-closed"); - session.bridge.close(); - closeRelayVoiceSession(session); - broadcastToOwner(session.context, session.connId, { - relaySessionId: session.id, - type: "close", - reason, - talkEvent: session.harness.talk.emit({ - type: "session.closed", - payload: { reason }, - final: true, - }), + try { + session.bridge.close(); + } finally { + // Provider teardown may throw, but the relay must still reach its durable + // voice and owner-visible terminal state before that error is surfaced. + closeRelayVoiceSession(session); + broadcastToOwner(session.context, session.connId, { + relaySessionId: session.id, + type: "close", + reason, + talkEvent: session.harness.talk.emit({ + type: "session.closed", + payload: { reason }, + final: true, + }), + }); + } +} + +/** Releases every realtime relay session owned by a disconnected gateway connection. */ +export function closeTalkRealtimeRelaySessionsForConnection(connId: string): void { + closeTalkRelaySessionsForConnection({ + sessions: relaySessions.values(), + connId, + closeSession: (session) => closeRelaySession(session, "completed"), + onCloseError: (error, session) => { + session.context.logGateway.warn( + `failed to close realtime relay session after connection disconnect: ${formatError(error)}`, + ); + }, }); } diff --git a/src/gateway/talk-realtime-relay-session-create.ts b/src/gateway/talk-realtime-relay-session-create.ts index b148805c3a44..fc75eefc87b5 100644 --- a/src/gateway/talk-realtime-relay-session-create.ts +++ b/src/gateway/talk-realtime-relay-session-create.ts @@ -11,7 +11,10 @@ import { shouldAutoControlRealtimeVoiceAgentText, } from "../talk/agent-run-control.js"; import { resolveTalkSessionAgentId } from "../talk/agent-target.js"; -import { REALTIME_VOICE_AUDIO_FORMAT_PCM16_24KHZ } from "../talk/provider-types.js"; +import { + REALTIME_VOICE_AUDIO_FORMAT_PCM16_24KHZ, + type RealtimeVoiceCloseReason, +} from "../talk/provider-types.js"; import { createRealtimeVoiceSessionHarness } from "../talk/realtime-session-harness.js"; import type { TalkEventInput } from "../talk/talk-session-controller.js"; import { registerChatAbortController } from "./chat-abort.js"; @@ -102,13 +105,23 @@ export function createTalkRealtimeRelaySession( let currentOutputResponseId: string | undefined; let ready = false; let failureEmitted = false; + const constructionTerminal: { + current?: { kind: "error"; error: Error } | { kind: "close"; reason: RealtimeVoiceCloseReason }; + } = {}; const relayRef: { current?: RelaySession } = {}; + const getActiveRelay = (): RelaySession | undefined => { + const relay = relayRef.current; + return relay && relaySessions.get(relay.id) === relay ? relay : undefined; + }; let consultAgentRuntime: ReturnType["agent"] | undefined; const relaySessionKey = params.sessionKey?.trim(); const relayAgentId = relaySessionKey ? resolveTalkSessionAgentId(params.cfg ?? params.context.getRuntimeConfig(), relaySessionKey) : undefined; const runAgentConsult = async ({ prompt, signal }: { prompt: string; signal?: AbortSignal }) => { + if (!getActiveRelay()) { + throw new Error("Realtime gateway-relay session is closed"); + } const runtimeConfig = params.cfg ?? params.context.getRuntimeConfig(); const sessionKey = relaySessionKey; if (!sessionKey) { @@ -182,9 +195,13 @@ export function createTalkRealtimeRelaySession( tools: params.tools, markStrategy: "transport", audioSink: { - isOpen: () => Boolean(relayRef.current && relaySessions.has(relayRef.current.id)), + isOpen: () => Boolean(getActiveRelay()), sendAudio: (audio) => { - const turnId = relayRef.current ? ensureRelayTurn(relayRef.current) : undefined; + const relay = getActiveRelay(); + if (!relay) { + return; + } + const turnId = ensureRelayTurn(relay); emit( { relaySessionId, @@ -201,7 +218,11 @@ export function createTalkRealtimeRelaySession( ); }, clearAudio: (reason) => { - const turnId = relayRef.current ? ensureRelayTurn(relayRef.current) : undefined; + const relay = getActiveRelay(); + if (!relay) { + return; + } + const turnId = ensureRelayTurn(relay); emit( { relaySessionId, type: "clear", ...(reason ? { reason } : {}) }, { @@ -213,7 +234,11 @@ export function createTalkRealtimeRelaySession( ); }, sendMark: (markName) => { - const turnId = relayRef.current ? ensureRelayTurn(relayRef.current) : undefined; + const relay = getActiveRelay(); + if (!relay) { + return; + } + const turnId = ensureRelayTurn(relay); emit( { relaySessionId, type: "mark", markName }, { @@ -226,6 +251,9 @@ export function createTalkRealtimeRelaySession( }, }, onEvent: (event) => { + if (!getActiveRelay()) { + return; + } if (event.direction !== "server") { return; } @@ -260,9 +288,12 @@ export function createTalkRealtimeRelaySession( } }, onTranscript: (role, text, final) => { - const relay = relayRef.current; - const turnId = relay ? ensureRelayTurn(relay) : undefined; - if (final && relay) { + const relay = getActiveRelay(); + if (!relay) { + return; + } + const turnId = ensureRelayTurn(relay); + if (final) { enqueueRelayVoiceTranscript(relay, role, text); } const eventType = @@ -289,7 +320,6 @@ export function createTalkRealtimeRelaySession( return; } if ( - relay && pruneInactiveRelayAgentRuns(relay) > 0 && shouldAutoControlRealtimeVoiceAgentText(question) ) { @@ -301,11 +331,17 @@ export function createTalkRealtimeRelaySession( text: question, }) .then((result) => { + if (!getActiveRelay()) { + return; + } if (result.speak && !result.suppress && result.message.trim()) { bridge.sendUserMessage(buildRealtimeVoiceAgentControlSpeechMessage(result.message)); } }) .catch((error: unknown) => { + if (!getActiveRelay()) { + return; + } emit( { relaySessionId, type: "error", message: formatErrorMessage(error) }, { @@ -323,9 +359,12 @@ export function createTalkRealtimeRelaySession( } }, onToolCall: (toolCall) => { - const relay = relayRef.current; + const relay = getActiveRelay(); + if (!relay) { + return; + } let shouldSubmitWorkingResult = false; - if (relay && toolCall.name === REALTIME_VOICE_AGENT_CONSULT_TOOL_NAME) { + if (toolCall.name === REALTIME_VOICE_AGENT_CONSULT_TOOL_NAME) { const forcedConsult = relay.harness.forcedConsults.recordNativeConsult( toolCall.args, toolCall.callId, @@ -351,7 +390,7 @@ export function createTalkRealtimeRelaySession( } shouldSubmitWorkingResult = true; } - const turnId = relay ? ensureRelayTurn(relay) : undefined; + const turnId = ensureRelayTurn(relay); emit( { relaySessionId, @@ -369,15 +408,25 @@ export function createTalkRealtimeRelaySession( payload: { name: toolCall.name, args: toolCall.args }, }, ); - if (relay && shouldSubmitWorkingResult) { + if (shouldSubmitWorkingResult) { return submitRealtimeAgentConsultWorkingResponse(relay, toolCall.callId, turnId); } }, onReady: () => { + if (!getActiveRelay()) { + return; + } ready = true; emit({ relaySessionId, type: "ready" }, { type: "session.ready", payload: null }); }, onError: (error) => { + const active = getActiveRelay(); + if (!active) { + if (!relayRef.current) { + constructionTerminal.current ??= { kind: "error", error }; + } + return; + } const issue = realtimeRelayIssue({ message: formatErrorMessage(error), provider: params.provider.id, @@ -393,7 +442,10 @@ export function createTalkRealtimeRelaySession( }, onClose: (reason) => { const active = relaySessions.get(relaySessionId); - if (!active) { + if (!active || active !== relayRef.current) { + if (!relayRef.current) { + constructionTerminal.current ??= { kind: "close", reason }; + } return; } active.harness.close(); @@ -421,6 +473,21 @@ export function createTalkRealtimeRelaySession( ); }, }); + const earlyTerminal = constructionTerminal.current; + if (earlyTerminal) { + harness.close(); + try { + bridge.close(); + } catch (error) { + params.context.logGateway.warn( + `failed to close realtime relay bridge after provider terminated during creation: ${formatErrorMessage(error)}`, + ); + } + if (earlyTerminal.kind === "error") { + throw earlyTerminal.error; + } + throw new Error(`Realtime provider closed during session creation: ${earlyTerminal.reason}`); + } const initialSessionKey = params.sessionKey?.trim() || undefined; const relay: RelaySession = { id: relaySessionId, @@ -465,6 +532,10 @@ export function createTalkRealtimeRelaySession( relay.cleanupTimer.unref?.(); relaySessions.set(relaySessionId, relay); bridge.connect().catch((error: unknown) => { + const active = relaySessions.get(relaySessionId); + if (active !== relay) { + return; + } const issue = realtimeRelayIssue({ message: formatErrorMessage(error), provider: params.provider.id, @@ -477,10 +548,7 @@ export function createTalkRealtimeRelaySession( payload: issue, final: true, }); - const active = relaySessions.get(relaySessionId); - if (active) { - closeRelaySession(active, "error"); - } + closeRelaySession(active, "error"); }); return { diff --git a/src/gateway/talk-realtime-relay.test.ts b/src/gateway/talk-realtime-relay.test.ts index e213aec5205a..d715b63c274d 100644 --- a/src/gateway/talk-realtime-relay.test.ts +++ b/src/gateway/talk-realtime-relay.test.ts @@ -5,6 +5,7 @@ import path from "node:path"; * Tests talk realtime relay event forwarding and connection cleanup. */ import { afterEach, describe, expect, it, vi } from "vitest"; +import { useAutoCleanupTempDirTracker } from "../../test/helpers/temp-dir.js"; import { setActiveEmbeddedRun } from "../agents/embedded-agent-runner/runs.js"; import { testing as embeddedRunTesting } from "../agents/embedded-agent-runner/runs.test-support.js"; import { @@ -25,6 +26,7 @@ import { createChatRunState } from "./server-chat-state.js"; import { acknowledgeTalkRealtimeRelayMark, cancelTalkRealtimeRelayTurn, + closeTalkRealtimeRelaySessionsForConnection, createTalkRealtimeRelaySession as createTalkRealtimeRelaySessionRaw, ensureTalkRealtimeRelayVoiceSession, flushTalkRealtimeRelayVoiceWrites, @@ -36,6 +38,7 @@ import { } from "./talk-realtime-relay.js"; const activeRelaySessions = new Map(); +const tempDirs = useAutoCleanupTempDirTracker(afterEach); function createTalkRealtimeRelaySession( params: Parameters[0], @@ -92,6 +95,168 @@ describe("talk realtime gateway relay", () => { }; } + it("closes only realtime relays owned by the disconnected connection", async () => { + const envSnapshot = captureEnv(["OPENCLAW_STATE_DIR"]); + const tempDir = await fs.realpath(tempDirs.make("openclaw-relay-disconnect-")); + setTestEnvValue("OPENCLAW_STATE_DIR", tempDir); + const bridgeCloses: Array> = []; + const bridgeAudioSends: Array> = []; + const bridgeRequests: RealtimeVoiceBridgeCreateRequest[] = []; + const bridgeToolResults: Array> = []; + const provider = createIdleRelayProvider(); + provider.createBridge = (request) => { + const bridgeIndex = bridgeCloses.length; + const close = vi.fn(); + const sendAudio = vi.fn(); + const submitToolResult = vi.fn(); + bridgeRequests.push(request); + bridgeCloses.push(close); + bridgeAudioSends.push(sendAudio); + bridgeToolResults.push(submitToolResult); + return { + connect: vi.fn(async () => { + if (bridgeIndex === 1) { + throw new Error("late connect failure"); + } + }), + sendAudio, + setMediaTimestamp: vi.fn(), + handleBargeIn: vi.fn(), + submitToolResult, + acknowledgeMark: vi.fn(), + close, + isConnected: vi.fn(() => true), + }; + }; + try { + const logGateway = { warn: vi.fn() }; + const broadcastToConnIds = vi.fn(); + const context = { + broadcastToConnIds, + chatAbortControllers: new Map(), + getRuntimeConfig: () => ({}), + logGateway, + } as never; + const createSession = (connId: string) => + createTalkRealtimeRelaySession({ + context, + connId, + provider, + providerConfig: {}, + instructions: "brief", + tools: [], + }); + const firstOwned = createSession("conn-owner"); + const secondOwned = createSession("conn-owner"); + const unrelated = createSession("conn-other"); + ensureTalkRealtimeRelayVoiceSession({ + relaySessionId: firstOwned.relaySessionId, + connId: "conn-owner", + sessionKey: "agent:main:main", + }); + expect(clientVoiceSessionTesting.readRecord("main", firstOwned.relaySessionId)).toMatchObject( + { + status: "open", + }, + ); + bridgeCloses[0]?.mockImplementationOnce(() => { + throw new Error("provider close failed"); + }); + + expect(() => closeTalkRealtimeRelaySessionsForConnection("conn-owner")).not.toThrow(); + closeTalkRealtimeRelaySessionsForConnection("conn-owner"); + await Promise.resolve(); + await Promise.resolve(); + + expect(bridgeCloses[0]).toHaveBeenCalledOnce(); + expect(bridgeCloses[1]).toHaveBeenCalledOnce(); + expect(bridgeCloses[2]).not.toHaveBeenCalled(); + expect(logGateway.warn).toHaveBeenCalledWith( + "failed to close realtime relay session after connection disconnect: provider close failed", + ); + await vi.waitFor(() => + expect( + clientVoiceSessionTesting.readRecord("main", firstOwned.relaySessionId)?.status, + ).toBe("closed"), + ); + expect( + broadcastToConnIds.mock.calls.some( + ([event, payload]) => + event === "talk.event" && + payload.relaySessionId === firstOwned.relaySessionId && + payload.type === "close" && + payload.talkEvent?.type === "session.closed" && + payload.talkEvent.final === true, + ), + ).toBe(true); + expect( + broadcastToConnIds.mock.calls.some( + ([event, payload]) => + event === "talk.event" && + payload.relaySessionId === secondOwned.relaySessionId && + payload.type === "error", + ), + ).toBe(false); + const eventCountAfterClose = broadcastToConnIds.mock.calls.length; + const lateRequest = bridgeRequests[0]; + if (!lateRequest?.runAgentConsult) { + throw new Error("expected relay provider request to include the agent consult runner"); + } + lateRequest.onReady?.(); + lateRequest.onError?.(new Error("late provider error")); + lateRequest.onEvent?.({ direction: "server", type: "response.done" }); + lateRequest.onAudio(Buffer.from("late audio")); + lateRequest.onClearAudio("barge-in"); + lateRequest.onMark?.("late-mark"); + lateRequest.onTranscript?.("user", "late transcript", true); + lateRequest.onToolCall?.({ + itemId: "late-item", + callId: "late-call", + name: "openclaw_agent_consult", + args: { question: "late consult" }, + }); + lateRequest.onClose?.("error"); + await expect(lateRequest.runAgentConsult({ prompt: "late direct consult" })).rejects.toThrow( + "Realtime gateway-relay session is closed", + ); + await Promise.resolve(); + expect(broadcastToConnIds).toHaveBeenCalledTimes(eventCountAfterClose); + expect(bridgeToolResults[0]).not.toHaveBeenCalled(); + expect(() => + sendTalkRealtimeRelayAudio({ + relaySessionId: firstOwned.relaySessionId, + connId: "conn-owner", + audioBase64: "AQI=", + }), + ).toThrow("Unknown realtime relay session"); + expect(() => + sendTalkRealtimeRelayAudio({ + relaySessionId: secondOwned.relaySessionId, + connId: "conn-owner", + audioBase64: "AQI=", + }), + ).toThrow("Unknown realtime relay session"); + + sendTalkRealtimeRelayAudio({ + relaySessionId: unrelated.relaySessionId, + connId: "conn-other", + audioBase64: "AQI=", + }); + expect(bridgeAudioSends[2]).toHaveBeenCalledOnce(); + stopTalkRealtimeRelaySession({ + relaySessionId: unrelated.relaySessionId, + connId: "conn-other", + }); + closeTalkRealtimeRelaySessionsForConnection("conn-other"); + expect(bridgeCloses[2]).toHaveBeenCalledOnce(); + } finally { + clientVoiceSessionTesting.reset(); + closeOpenClawAgentDatabasesForTest(); + closeOpenClawStateDatabaseForTest(); + envSnapshot.restore(); + } + }); + it("injects the host agent runner only into gateway-relay bridge creation", () => { let bridgeRequest: RealtimeVoiceBridgeCreateRequest | undefined; const provider = createIdleRelayProvider(); @@ -122,6 +287,65 @@ describe("talk realtime gateway relay", () => { }); }); + it.each([ + { + name: "error before close", + terminate: (request: RealtimeVoiceBridgeCreateRequest) => { + request.onError?.(new Error("provider rejected session")); + request.onClose?.("error"); + }, + expectedError: "provider rejected session", + }, + { + name: "close before error", + terminate: (request: RealtimeVoiceBridgeCreateRequest) => { + request.onClose?.("completed"); + request.onError?.(new Error("late provider error")); + }, + expectedError: "Realtime provider closed during session creation: completed", + }, + ])( + "rejects a synchronous provider $name during bridge creation", + ({ terminate, expectedError }) => { + const connect = vi.fn(async () => undefined); + const sendAudio = vi.fn(); + const close = vi.fn(); + const bridge = { + ...createIdleRelayProvider().createBridge({} as never), + connect, + sendAudio, + close, + }; + const provider = createIdleRelayProvider(); + provider.createBridge = (request) => { + terminate(request); + return bridge; + }; + const broadcastToConnIds = vi.fn(); + + expect(() => + createTalkRealtimeRelaySession({ + context: { + broadcastToConnIds, + chatAbortControllers: new Map(), + getRuntimeConfig: () => ({}), + logGateway: { warn: vi.fn() }, + } as never, + connId: "conn-early-terminal", + provider, + providerConfig: {}, + instructions: "brief", + tools: [], + }), + ).toThrow(expectedError); + + expect(connect).not.toHaveBeenCalled(); + expect(sendAudio).not.toHaveBeenCalled(); + expect(close).toHaveBeenCalledOnce(); + expect(broadcastToConnIds).not.toHaveBeenCalled(); + }, + ); + it("appends finalized relay transcripts to the canonical agent session", async () => { const envSnapshot = captureEnv(["OPENCLAW_STATE_DIR"]); const tempDir = await fs.realpath( diff --git a/src/gateway/talk-realtime-relay.ts b/src/gateway/talk-realtime-relay.ts index 225a59f2b426..2ae9964ec36c 100644 --- a/src/gateway/talk-realtime-relay.ts +++ b/src/gateway/talk-realtime-relay.ts @@ -4,6 +4,7 @@ export { createTalkRealtimeRelaySession } from "./talk-realtime-relay-session-cr export { acknowledgeTalkRealtimeRelayMark, cancelTalkRealtimeRelayTurn, + closeTalkRealtimeRelaySessionsForConnection, ensureTalkRealtimeRelayVoiceSession, flushTalkRealtimeRelayVoiceWrites, registerTalkRealtimeRelayAgentRun, diff --git a/src/gateway/talk-relay-session-lifecycle.ts b/src/gateway/talk-relay-session-lifecycle.ts index aebb9fd68614..4d37aff98976 100644 --- a/src/gateway/talk-relay-session-lifecycle.ts +++ b/src/gateway/talk-relay-session-lifecycle.ts @@ -39,6 +39,27 @@ export function closeExpiredTalkRelaySessions(params: { + sessions: Iterable; + connId: string; + closeSession: CloseTalkRelaySession; + onCloseError: (error: unknown, session: TSession) => void; +}): void { + for (const session of params.sessions) { + if (session.connId !== params.connId) { + continue; + } + try { + params.closeSession(session); + } catch (error) { + params.onCloseError(error, session); + } + } +} + /** Returns the active session only when it belongs to the current connection. */ export function requireActiveTalkRelaySession(params: { sessions: ReadonlyMap; diff --git a/src/gateway/talk-transcription-relay.test.ts b/src/gateway/talk-transcription-relay.test.ts index 9eb57d7ba278..2d7456d6a323 100644 --- a/src/gateway/talk-transcription-relay.test.ts +++ b/src/gateway/talk-transcription-relay.test.ts @@ -4,8 +4,10 @@ import { afterEach, describe, expect, it, vi } from "vitest"; import type { RealtimeTranscriptionProviderPlugin } from "../plugins/types.js"; import type { RealtimeTranscriptionSessionCreateRequest } from "../realtime-transcription/provider-types.js"; +import { getUnifiedTalkSession, rememberUnifiedTalkSession } from "./talk-session-registry.js"; import { cancelTalkTranscriptionRelayTurn, + closeTalkTranscriptionRelaySessionsForConnection, createTalkTranscriptionRelaySession, sendTalkTranscriptionRelayAudio, stopTalkTranscriptionRelaySession, @@ -40,13 +42,15 @@ function createTranscriptionProvider( function createBroadcastContext() { const events: BroadcastEvent[] = []; + const logGateway = { warn: vi.fn() }; const context = { getRuntimeConfig: () => ({}), + logGateway, broadcastToConnIds: (event: string, payload: unknown, connIds: ReadonlySet) => { events.push({ event, payload, connIds: [...connIds] }); }, } as never; - return { context, events }; + return { context, events, logGateway }; } async function createStartedRelaySession( @@ -212,6 +216,101 @@ describe("talk transcription gateway relay", () => { }); }); + it("closes only transcription relays owned by the disconnected connection", async () => { + const firstOwned = createSttSessionMock(); + const secondOwned = createSttSessionMock(async () => { + throw new Error("late connect failure"); + }); + const unrelated = createSttSessionMock(); + const requests: RealtimeTranscriptionSessionCreateRequest[] = []; + const { context, events, logGateway } = createBroadcastContext(); + const createSession = (connId: string, sttSession: ReturnType) => + createTalkTranscriptionRelaySession({ + context, + connId, + provider: createTranscriptionProvider(sttSession, (request) => requests.push(request)), + providerConfig: {}, + }); + const firstSession = createSession("conn-owner", firstOwned); + const secondSession = createSession("conn-owner", secondOwned); + const unrelatedSession = createSession("conn-other", unrelated); + for (const session of [firstSession, secondSession]) { + rememberUnifiedTalkSession(session.transcriptionSessionId, { + kind: "transcription-relay", + connId: "conn-owner", + transcriptionSessionId: session.transcriptionSessionId, + }); + } + firstOwned.close.mockImplementationOnce(() => { + throw new Error("provider close failed"); + }); + + expect(() => closeTalkTranscriptionRelaySessionsForConnection("conn-owner")).not.toThrow(); + closeTalkTranscriptionRelaySessionsForConnection("conn-owner"); + await Promise.resolve(); + await Promise.resolve(); + + expect(firstOwned.close).toHaveBeenCalledOnce(); + expect(secondOwned.close).toHaveBeenCalledOnce(); + expect(unrelated.close).not.toHaveBeenCalled(); + expect(logGateway.warn).toHaveBeenCalledWith( + "failed to close transcription relay session after connection disconnect: provider close failed", + ); + for (const transcriptionSessionId of [ + firstSession.transcriptionSessionId, + secondSession.transcriptionSessionId, + ]) { + expect( + events.some( + (event) => + isRecord(event.payload) && + event.payload.transcriptionSessionId === transcriptionSessionId && + event.payload.type === "close" && + isRecord(event.payload.talkEvent) && + event.payload.talkEvent.type === "session.closed" && + event.payload.talkEvent.final === true, + ), + ).toBe(true); + expect(() => + sendTalkTranscriptionRelayAudio({ + transcriptionSessionId, + connId: "conn-owner", + audioBase64: "AQI=", + }), + ).toThrow("Unknown transcription Talk session"); + expect(() => getUnifiedTalkSession(transcriptionSessionId)).toThrow("Unknown Talk session"); + } + expect( + events.some( + (event) => + isRecord(event.payload) && + (event.payload.transcriptionSessionId === firstSession.transcriptionSessionId || + event.payload.transcriptionSessionId === secondSession.transcriptionSessionId) && + (event.payload.type === "ready" || event.payload.type === "error"), + ), + ).toBe(false); + const eventCountAfterClose = events.length; + requests[0]?.onSpeechStart?.(); + requests[0]?.onPartial?.("late partial"); + requests[0]?.onTranscript?.("late transcript"); + requests[0]?.onError?.(new Error("late provider error")); + await Promise.resolve(); + expect(events).toHaveLength(eventCountAfterClose); + + sendTalkTranscriptionRelayAudio({ + transcriptionSessionId: unrelatedSession.transcriptionSessionId, + connId: "conn-other", + audioBase64: "AQI=", + }); + expect(unrelated.sendAudio).toHaveBeenCalledOnce(); + stopTalkTranscriptionRelaySession({ + transcriptionSessionId: unrelatedSession.transcriptionSessionId, + connId: "conn-other", + }); + closeTalkTranscriptionRelaySessionsForConnection("conn-other"); + expect(unrelated.close).toHaveBeenCalledOnce(); + }); + it("rejects provider configs that do not match relay audio input", () => { const provider = createTranscriptionProvider(createSttSessionMock()); const { context } = createBroadcastContext(); diff --git a/src/gateway/talk-transcription-relay.ts b/src/gateway/talk-transcription-relay.ts index a5d0acb3ba18..2107f979891e 100644 --- a/src/gateway/talk-transcription-relay.ts +++ b/src/gateway/talk-transcription-relay.ts @@ -15,11 +15,14 @@ import { createTalkSessionController, } from "../talk/talk-session-controller.js"; import type { GatewayRequestContext } from "./server-methods/shared-types.js"; +import { formatError } from "./server-utils.js"; import { decodeTalkRelayAudioBase64 } from "./talk-relay-audio-base64.js"; import { closeExpiredTalkRelaySessions, + closeTalkRelaySessionsForConnection, requireActiveTalkRelaySession, } from "./talk-relay-session-lifecycle.js"; +import { forgetUnifiedTalkSession } from "./talk-session-registry.js"; /** * Gateway-owned relay for streaming speech-to-text providers used by Talk. @@ -181,17 +184,37 @@ function closeTranscriptionSession( } session.closed = true; transcriptionSessions.delete(session.id); + forgetUnifiedTalkSession(session.id); clearTimeout(session.cleanupTimer); - session.sttSession.close(); - broadcastToOwner(session.context, session.connId, { - transcriptionSessionId: session.id, - type: "close", - reason, - talkEvent: session.talk.emit({ - type: "session.closed", - payload: { reason }, - final: true, - }), + try { + session.sttSession.close(); + } finally { + // Provider teardown may throw, but the owner-visible terminal event must + // still complete so disconnect cleanup cannot leave ambiguous state. + broadcastToOwner(session.context, session.connId, { + transcriptionSessionId: session.id, + type: "close", + reason, + talkEvent: session.talk.emit({ + type: "session.closed", + payload: { reason }, + final: true, + }), + }); + } +} + +/** Releases every transcription relay owned by a disconnected gateway connection. */ +export function closeTalkTranscriptionRelaySessionsForConnection(connId: string): void { + closeTalkRelaySessionsForConnection({ + sessions: transcriptionSessions.values(), + connId, + closeSession: (session) => closeTranscriptionSession(session, "completed"), + onCloseError: (error, session) => { + session.context.logGateway.warn( + `failed to close transcription relay session after connection disconnect: ${formatError(error)}`, + ); + }, }); } @@ -251,18 +274,26 @@ export function createTalkTranscriptionRelaySession( }); }; const relayRef: { current?: TranscriptionRelaySession } = {}; - const ensureTurnId = (): string => { + const getActiveRelay = (): TranscriptionRelaySession | undefined => { const relay = relayRef.current; - return relay ? ensureTranscriptionTurn(relay) : "turn-1"; + return relay && transcriptionSessions.get(relay.id) === relay ? relay : undefined; }; const sttSession = params.provider.createSession({ cfg: params.context.getRuntimeConfig(), providerConfig: params.providerConfig, onSpeechStart: () => { - ensureTurnId(); + const relay = getActiveRelay(); + if (!relay) { + return; + } + ensureTranscriptionTurn(relay); }, onPartial: (text) => { - const turnId = ensureTurnId(); + const relay = getActiveRelay(); + if (!relay) { + return; + } + const turnId = ensureTranscriptionTurn(relay); emit( { transcriptionSessionId, type: "partial", text }, { @@ -273,7 +304,11 @@ export function createTalkTranscriptionRelaySession( ); }, onTranscript: (text) => { - const turnId = ensureTurnId(); + const relay = getActiveRelay(); + if (!relay) { + return; + } + const turnId = ensureTranscriptionTurn(relay); emit( { transcriptionSessionId, type: "transcript", text, final: true }, { @@ -283,21 +318,22 @@ export function createTalkTranscriptionRelaySession( final: true, }, ); - const relay = relayRef.current; - if (relay) { - const ended = relay.talk.endTurn({ turnId, payload: {} }); - if (ended.ok) { - broadcastToOwner(relay.context, relay.connId, { - transcriptionSessionId, - type: "transcript", - text: "", - final: true, - talkEvent: ended.event, - }); - } + const ended = relay.talk.endTurn({ turnId, payload: {} }); + if (ended.ok) { + broadcastToOwner(relay.context, relay.connId, { + transcriptionSessionId, + type: "transcript", + text: "", + final: true, + talkEvent: ended.event, + }); } }, onError: (error) => { + const relay = getActiveRelay(); + if (!relay) { + return; + } emit( { transcriptionSessionId, type: "error", message: error.message }, { @@ -306,10 +342,7 @@ export function createTalkTranscriptionRelaySession( final: true, }, ); - const relay = relayRef.current; - if (relay) { - closeTranscriptionSession(relay, "error"); - } + closeTranscriptionSession(relay, "error"); }, }); const relay: TranscriptionRelaySession = { @@ -334,9 +367,16 @@ export function createTalkTranscriptionRelaySession( sttSession .connect() .then(() => { + if (transcriptionSessions.get(transcriptionSessionId) !== relay) { + return; + } emit({ transcriptionSessionId, type: "ready" }, { type: "session.ready", payload: null }); }) .catch((error: unknown) => { + const active = transcriptionSessions.get(transcriptionSessionId); + if (active !== relay) { + return; + } emit( { transcriptionSessionId, @@ -349,10 +389,7 @@ export function createTalkTranscriptionRelaySession( final: true, }, ); - const active = transcriptionSessions.get(transcriptionSessionId); - if (active) { - closeTranscriptionSession(active, "error"); - } + closeTranscriptionSession(active, "error"); }); return {