diff --git a/.changeset/clean-fallback-teardown.md b/.changeset/clean-fallback-teardown.md new file mode 100644 index 0000000000..39a8c262d8 --- /dev/null +++ b/.changeset/clean-fallback-teardown.md @@ -0,0 +1,5 @@ +--- +'@livekit/agents': patch +--- + +Keep healthy STT providers available when fallback streams close with transcripts in flight. diff --git a/agents/etc/agents.api.md b/agents/etc/agents.api.md index ec6f879ddf..3a6fce63cb 100644 --- a/agents/etc/agents.api.md +++ b/agents/etc/agents.api.md @@ -7699,6 +7699,8 @@ abstract class SpeechStream implements AsyncIterableIterator { // (undocumented) detachInputStream(): void; endInput(): void; + // @internal + get _failed(): boolean; flush(): void; // (undocumented) protected static readonly FLUSH_SENTINEL: unique symbol; @@ -10149,13 +10151,13 @@ export const zipFunctionCallsAndOutputs: (event: FunctionToolsExecutedEvent) => // src/llm/tool_context.ts:746:3 - (ae-unresolved-link) The @link reference could not be resolved: The reference is ambiguous because "ToolFlag" has more than one declaration; you need to add a TSDoc member reference selector // src/metrics/base.ts:198:3 - (ae-forgotten-export) The symbol "RealtimeModelMetricsInputTokenDetails" needs to be exported by the entry point index.d.ts // src/metrics/base.ts:202:3 - (ae-forgotten-export) The symbol "RealtimeModelMetricsOutputTokenDetails" needs to be exported by the entry point index.d.ts -// src/stt/stt.ts:364:3 - (ae-unresolved-link) The @link reference could not be resolved: The package "@livekit/agents" does not have an export "STT" +// src/stt/stt.ts:365:3 - (ae-unresolved-link) The @link reference could not be resolved: The package "@livekit/agents" does not have an export "STT" // src/utils.ts:550:3 - (ae-unresolved-link) The @link reference could not be resolved: The package "@livekit/agents" does not have an export "cancelled" // src/voice/agent_session.ts:387:3 - (ae-unresolved-link) The @link reference could not be resolved: This type of declaration is not supported yet by the resolver // src/voice/agent_session.ts:1025:5 - (ae-forgotten-export) The symbol "RecordingOptions" needs to be exported by the entry point index.d.ts -// src/voice/agent_session.ts:1695:5 - (ae-forgotten-export) The symbol "STTError" needs to be exported by the entry point index.d.ts -// src/voice/agent_session.ts:1695:5 - (ae-forgotten-export) The symbol "TTSError" needs to be exported by the entry point index.d.ts -// src/voice/agent_session.ts:1695:5 - (ae-forgotten-export) The symbol "LLMError" needs to be exported by the entry point index.d.ts +// src/voice/agent_session.ts:1696:5 - (ae-forgotten-export) The symbol "STTError" needs to be exported by the entry point index.d.ts +// src/voice/agent_session.ts:1696:5 - (ae-forgotten-export) The symbol "TTSError" needs to be exported by the entry point index.d.ts +// src/voice/agent_session.ts:1696:5 - (ae-forgotten-export) The symbol "LLMError" needs to be exported by the entry point index.d.ts // src/voice/amd.ts:315:3 - (ae-unresolved-link) The @link reference could not be resolved: The reference is ambiguous because "waitForTrackPublication" has more than one declaration; you need to add a TSDoc member reference selector // src/voice/amd.ts:315:3 - (ae-unresolved-link) The @link reference could not be resolved: The package "@livekit/agents" does not have an export "gateListening" // src/voice/amd.ts:323:3 - (ae-unresolved-link) The @link reference could not be resolved: The package "@livekit/agents" does not have an export "aclose" diff --git a/agents/src/stt/fallback_adapter.test.ts b/agents/src/stt/fallback_adapter.test.ts index db616664bc..cc2fdd6c2e 100644 --- a/agents/src/stt/fallback_adapter.test.ts +++ b/agents/src/stt/fallback_adapter.test.ts @@ -314,6 +314,82 @@ describe('FallbackSpeechStream (streaming path)', () => { expect(adapter.status[1]?.available).toBe(true); }); + it('keeps the provider available when a late transcript arrives after stream close', async () => { + const primary = new FakeSTT({ + label: 'primary', + fakeTranscript: 'late transcript', + fakeTimeoutMs: 50, + }); + const adapter = new FallbackAdapter({ + sttInstances: [primary], + maxRetryPerSTT: 0, + }); + + const availabilityChanges: Array<{ stt: STT; available: boolean }> = []; + (adapter as unknown as EventEmitter).on( + 'stt_availability_changed', + (ev: { stt: STT; available: boolean }) => { + availabilityChanges.push(ev); + }, + ); + + const stream = adapter.stream(); + await primary.streamCh.next(); + stream.close(); + await delay(150); + + expect(availabilityChanges).toEqual([]); + expect(adapter.status[0]?.available).toBe(true); + + await adapter.close(); + }); + + it('closes recovery probes when a late transcript arrives after stream close', async () => { + const primary = new FakeSTT({ + label: 'primary', + fakeException: new APIError('primary down'), + }); + const fallback = new FakeSTT({ + label: 'fallback', + fakeTranscript: 'late transcript', + fakeTimeoutMs: 50, + }); + const adapter = new FallbackAdapter({ + sttInstances: [primary, fallback], + maxRetryPerSTT: 0, + }); + + (adapter as unknown as EventEmitter).on( + 'stt_availability_changed', + (ev: { stt: STT; available: boolean }) => { + if (ev.stt === primary && !ev.available) { + // Keep the recovery probe alive until the parent stream tears it down. + primary.updateOptions({ fakeException: null }); + } + }, + ); + + const stream = adapter.stream(); + stream.endInput(); + + await primary.streamCh.next(); // failed main stream + const recoveryProbe = (await primary.streamCh.next()).value; + await fallback.streamCh.next(); + + stream.close(); + await delay(150); + + let recoveryProbeClosed = false; + try { + recoveryProbe?.endInput(); + } catch { + recoveryProbeClosed = true; + } + await adapter.close(); + + expect(recoveryProbeClosed).toBe(true); + }); + it('stream switches to the secondary provider when the primary errors', async () => { const primary = new FakeSTT({ label: 'primary', diff --git a/agents/src/stt/fallback_adapter.ts b/agents/src/stt/fallback_adapter.ts index f733545bbf..49bd2fcb4b 100644 --- a/agents/src/stt/fallback_adapter.ts +++ b/agents/src/stt/fallback_adapter.ts @@ -26,6 +26,8 @@ interface STTStatus { available: boolean; recoveringRecognizeTask: Task | null; recoveringStreamTask: Task | null; + waitingStreams: Set; + recoveryClosed: boolean; } /** @@ -69,8 +71,10 @@ const DEFAULT_FALLBACK_API_CONNECT_OPTIONS: APIConnectOptions = { * * When the primary STT fails, the adapter switches to the next available * provider in the list for the active session. Failed providers are monitored - * by a parallel probe stream that receives the same live audio — when a probe - * yields a non-empty FINAL_TRANSCRIPT the provider is marked available again. + * by one probe stream per provider. The probe receives its owner's live audio; + * another waiting stream takes over if the owner closes. A non-empty + * FINAL_TRANSCRIPT marks the provider available again. + * If every provider is unavailable, normal streams retry them in priority order. * * Non-streaming STTs are automatically wrapped with {@link StreamAdapter} * provided a `vad` is passed in. @@ -153,6 +157,8 @@ export class FallbackAdapter extends STT { available: true, recoveringRecognizeTask: null, recoveringStreamTask: null, + waitingStreams: new Set(), + recoveryClosed: false, })); this.setupEventForwarding(); @@ -226,23 +232,35 @@ export class FallbackAdapter extends STT { private tryRecoverRecognize(stt: STT, frame: Parameters[0]): void { const idx = this.sttInstances.indexOf(stt); const status = this._status[idx]; - if (!status) return; + if (!status || status.recoveryClosed) return; if (status.recoveringRecognizeTask && !status.recoveringRecognizeTask.done) return; - status.recoveringRecognizeTask = Task.from(async (controller) => { + const task = Task.from(async (controller) => { try { await stt.recognize(frame, controller.signal); + if (controller.signal.aborted || status.recoveryClosed) return; status.available = true; - this._logger.info({ stt: stt.label }, `${stt.label} recovered`); + this._logger.info({ stt: stt.label }, 'STT recovered'); this.emitAvailabilityChanged(stt, true); } catch (e) { + if (controller.signal.aborted || status.recoveryClosed) return; if (e instanceof APIError) { - this._logger.warn({ stt: stt.label, err: e }, `${stt.label} recovery failed`); + this._logger.warn( + { stt: stt.label, errorType: e instanceof Error ? e.constructor.name : typeof e }, + 'STT recovery failed', + ); } else { - this._logger.debug({ stt: stt.label, err: e }, `${stt.label} recovery unexpected error`); + this._logger.debug( + { stt: stt.label, errorType: e instanceof Error ? e.constructor.name : typeof e }, + 'STT recovery failed', + ); } } }); + status.recoveringRecognizeTask = task; + task.addDoneCallback(() => { + if (status.recoveringRecognizeTask === task) status.recoveringRecognizeTask = null; + }); } // Skip the base class's `metrics_collected` emit: the active child's own @@ -276,17 +294,10 @@ export class FallbackAdapter extends STT { this._setActiveStt(stt); return result; } catch (e) { - if (e instanceof APIError) { - this._logger.warn( - { stt: stt.label, err: e }, - `${stt.label} failed, switching to next STT`, - ); - } else { - this._logger.warn( - { stt: stt.label, err: e }, - `${stt.label} unexpected error, switching to next STT`, - ); - } + this._logger.warn( + { stt: stt.label, errorType: e instanceof Error ? e.constructor.name : typeof e }, + 'STT failed, switching to next provider', + ); if (status.available) { status.available = false; this.emitAvailabilityChanged(stt, false); @@ -312,12 +323,12 @@ export class FallbackAdapter extends STT { override async close(): Promise { const tasks: Task[] = []; for (const status of this._status) { + status.recoveryClosed = true; + status.waitingStreams.clear(); if (status.recoveringRecognizeTask && !status.recoveringRecognizeTask.done) { tasks.push(status.recoveringRecognizeTask); } - if (status.recoveringStreamTask && !status.recoveringStreamTask.done) { - tasks.push(status.recoveringStreamTask); - } + if (status.recoveringStreamTask) tasks.push(status.recoveringStreamTask); } if (tasks.length > 0) { await cancelAndWait(tasks, 1000); @@ -333,7 +344,8 @@ export class FallbackAdapter extends STT { class FallbackSpeechStream extends SpeechStream { label = 'stt.FallbackSpeechStream'; private fallbackAdapter: FallbackAdapter; - private recoveringStreams: SpeechStream[] = []; + private recoveringStreams = new Map>(); + private inputEnded = false; private _logger = log(); constructor(adapter: FallbackAdapter, connOptions: APIConnectOptions) { @@ -357,65 +369,129 @@ class FallbackSpeechStream extends SpeechStream { if (!this.output.closed) this.output.close(); } - private tryRecoverStream(sttInstance: STT): void { + private tryRecoverStream(sttInstance: STT): boolean { + if (this.abortSignal.aborted) return false; const idx = this.fallbackAdapter.sttInstances.indexOf(sttInstance); const status = this.fallbackAdapter.status[idx]; - if (!status) return; - if (status.recoveringStreamTask && !status.recoveringStreamTask.done) return; - - const probe = sttInstance.stream({ - connOptions: { - maxRetry: 0, - timeoutMs: this.fallbackAdapter.attemptTimeoutMs, - retryIntervalMs: this.fallbackAdapter.retryIntervalMs, - }, - }); - this.recoveringStreams.push(probe); + if (!status || status.available || status.recoveryClosed) return false; + if (status.recoveringStreamTask && !status.recoveringStreamTask.done) { + status.waitingStreams.add(this); + return false; + } + status.waitingStreams.delete(this); + + let probe: SpeechStream; + try { + probe = sttInstance.stream({ + connOptions: { + maxRetry: 0, + timeoutMs: this.fallbackAdapter.attemptTimeoutMs, + retryIntervalMs: this.fallbackAdapter.retryIntervalMs, + }, + }); + } catch (error) { + this._logger.warn( + { + stt: sttInstance.label, + errorType: error instanceof Error ? error.constructor.name : typeof error, + }, + 'STT recovery failed', + ); + return false; + } + const closeProbe = () => { + try { + probe.close(); + } catch { + /* already closed */ + } + }; + if (this.abortSignal.aborted || status.recoveryClosed) { + closeProbe(); + return false; + } - // Absorb child 'error' events while the probe is active. JS EventEmitter - // crashes if 'error' fires with no listener; the probe's iterator ends - // naturally on failure, so we don't need to do anything with the payload. + // Absorb provider error events; each probe records its own terminal outcome. const errorSink: (e: STTError) => void = () => {}; sttInstance.on('error', errorSink); - status.recoveringStreamTask = Task.from(async (controller) => { + const task = Task.from(async (controller) => { + controller.signal.addEventListener('abort', closeProbe, { once: true }); try { - let gotTranscript = false; + let recovered = false; for await (const ev of probe) { - if (controller.signal.aborted) break; + if (controller.signal.aborted || this.abortSignal.aborted) break; if (ev.type === SpeechEventType.FINAL_TRANSCRIPT) { const text = ev.alternatives?.[0]?.text; if (!text) continue; - gotTranscript = true; + recovered = true; break; } } - if (!gotTranscript) return; - status.available = true; - this._logger.info({ stt: sttInstance.label }, `${sttInstance.label} recovered`); - this.fallbackAdapter.emitAvailabilityChanged(sttInstance, true); + if ( + !recovered || + controller.signal.aborted || + this.abortSignal.aborted || + status.recoveryClosed + ) + return; + if (!status.available) { + status.available = true; + this._logger.info({ stt: sttInstance.label }, 'STT recovered'); + this.fallbackAdapter.emitAvailabilityChanged(sttInstance, true); + } } catch (e) { + if (controller.signal.aborted || this.abortSignal.aborted) return; if (e instanceof APIError) { this._logger.warn( - { stt: sttInstance.label, err: e }, - `${sttInstance.label} recovery failed`, + { + stt: sttInstance.label, + errorType: e instanceof Error ? e.constructor.name : typeof e, + }, + 'STT recovery failed', ); } else { this._logger.debug( - { stt: sttInstance.label, err: e }, - `${sttInstance.label} recovery unexpected error`, + { + stt: sttInstance.label, + errorType: e instanceof Error ? e.constructor.name : typeof e, + }, + 'STT recovery failed', ); } } finally { + controller.signal.removeEventListener('abort', closeProbe); sttInstance.off('error', errorSink); - probe.close(); - const i = this.recoveringStreams.indexOf(probe); - if (i >= 0) this.recoveringStreams.splice(i, 1); + closeProbe(); } }); + this.recoveringStreams.set(probe, task); + status.recoveringStreamTask = task; + task.addDoneCallback(() => { + this.recoveringStreams.delete(probe); + if (status.recoveringStreamTask !== task) return; + status.recoveringStreamTask = null; + if (status.available || status.recoveryClosed) { + status.waitingStreams.clear(); + return; + } + for (const stream of status.waitingStreams) { + status.waitingStreams.delete(stream); + if (stream.tryRecoverStream(sttInstance)) break; + } + }); + if (this.inputEnded) { + try { + probe.endInput(); + } catch { + closeProbe(); + } + } + return true; } protected async run(): Promise { + if (this.abortSignal.aborted) return; const startTime = Date.now(); const allFailed = this.fallbackAdapter.status.every((s) => !s.available); if (allFailed) { @@ -429,16 +505,11 @@ class FallbackSpeechStream extends SpeechStream { // type to `never` based on its initial value. TS's control-flow analysis // for closures can't always see that outer code reassigns the var. const mainRef: { current: SpeechStream | null } = { current: null }; - // Tracks whether the forwarder has finished draining `this.input`. - // Children elected after this point never receive input, so we must - // end their input immediately on election (mirrors Python's check for - // forward_input_task.done() before starting a new one). - let forwarderFinished = false; // Forwarder runs as a Task so we can cancel+await it on terminal failure. const forwarderTask = Task.from(async (controller) => { for await (const item of this.input) { if (controller.signal.aborted || this.abortSignal.aborted) break; - for (const probe of [...this.recoveringStreams]) { + for (const probe of this.recoveringStreams.keys()) { try { if (typeof item === 'symbol') probe.flush(); else probe.pushFrame(item); @@ -452,138 +523,130 @@ class FallbackSpeechStream extends SpeechStream { if (typeof item === 'symbol') current.flush(); else current.pushFrame(item); } catch (e) { - this._logger.debug({ err: e }, 'error forwarding input to main stream'); + this._logger.debug( + { errorType: e instanceof Error ? e.constructor.name : typeof e }, + 'error forwarding input to main stream', + ); } } } - const endTarget = mainRef.current; - if (endTarget !== null) { + this.inputEnded = true; + for (const endTarget of [mainRef.current, ...this.recoveringStreams.keys()]) { try { - endTarget.endInput(); + endTarget?.endInput(); } catch { /* already ended */ } } - forwarderFinished = true; }); - for (let i = 0; i < this.fallbackAdapter.sttInstances.length; i++) { - const sttInstance = this.fallbackAdapter.sttInstances[i]!; - const status = this.fallbackAdapter.status[i]!; - if (!(status.available || allFailed)) { - this.tryRecoverStream(sttInstance); - continue; + const closeStreams = () => { + for (const status of this.fallbackAdapter.status) status.waitingStreams.delete(this); + for (const task of this.recoveringStreams.values()) { + task.cancel(); } + for (const stream of [mainRef.current, ...this.recoveringStreams.keys()]) { + try { + stream?.close(); + } catch { + // Continue closing the remaining streams if a provider throws. + } + } + }; - // Capture child errors: the base SpeechStream's mainTask emits an - // `error` event and then closes its output queue — consumers never - // see the throw via `for await`. Without this listener we can't - // distinguish a provider failure from a silent end-of-input. - let childErrored = false; - const errListener = (e: STTError) => { - if (!e.recoverable) childErrored = true; - }; - sttInstance.on('error', errListener); - - try { - const child = sttInstance.stream({ - connOptions: { - maxRetry: this.fallbackAdapter.maxRetryPerSTT, - timeoutMs: this.fallbackAdapter.attemptTimeoutMs, - retryIntervalMs: this.fallbackAdapter.retryIntervalMs, - }, - }); - // Keep child timestamps anchored to the parent stream's current retry attempt. - child.startTimeOffset = this.startTimeOffset + (Date.now() - startTime) / 1000; - mainRef.current = child; - // If the forwarder has already drained and exited (input EOF), it - // will never call endInput() on this child. End it here so the - // child's `for await (input)` loop can terminate cleanly instead - // of hanging forever. - if (forwarderFinished) { - try { - child.endInput(); - } catch { - /* already ended */ - } + this.abortSignal.addEventListener('abort', closeStreams, { once: true }); + try { + for (let i = 0; i < this.fallbackAdapter.sttInstances.length; i++) { + if (this.abortSignal.aborted) return; + const sttInstance = this.fallbackAdapter.sttInstances[i]!; + const status = this.fallbackAdapter.status[i]!; + if (!status.available && !allFailed) { + this.tryRecoverStream(sttInstance); + continue; } + // Absorb provider errors here; the child records its own terminal outcome. + const errListener = () => {}; + sttInstance.on('error', errListener); + try { - for await (const ev of child) { - this.fallbackAdapter._setActiveStt(sttInstance); - this.queue.put(ev); + const child = sttInstance.stream({ + connOptions: { + maxRetry: this.fallbackAdapter.maxRetryPerSTT, + timeoutMs: this.fallbackAdapter.attemptTimeoutMs, + retryIntervalMs: this.fallbackAdapter.retryIntervalMs, + }, + }); + // Keep child timestamps anchored to the parent stream's current retry attempt. + child.startTimeOffset = this.startTimeOffset + (Date.now() - startTime) / 1000; + mainRef.current = child; + try { + if (this.abortSignal.aborted) return; + // If the forwarder has already drained and exited (input EOF), it + // will never call endInput() on this child. End it here so the + // child's `for await (input)` loop can terminate cleanly instead + // of hanging forever. + if (this.inputEnded) { + try { + child.endInput(); + } catch { + /* already ended */ + } + } + + for await (const ev of child) { + // The parent can close while a child has a transcript in flight. + // Stop cleanly instead of treating the closed queue as a provider failure. + if (this.abortSignal.aborted || this.queue.closed) { + return; + } + this.fallbackAdapter._setActiveStt(sttInstance); + this.queue.put(ev); + } + } finally { + child.close(); } - } finally { - child.close(); - } - if (!childErrored) { - // Main stream ended cleanly (input EOF). - return; - } - if (status.available) { - status.available = false; - this.fallbackAdapter.emitAvailabilityChanged(sttInstance, false); - } - this._logger.warn( - { stt: sttInstance.label }, - `${sttInstance.label} failed, switching to next STT`, - ); - } catch (e) { - if (e instanceof APIError) { - this._logger.warn( - { stt: sttInstance.label, err: e }, - `${sttInstance.label} failed, switching to next STT`, - ); - } else { + if (this.abortSignal.aborted || (!child._failed && this.inputEnded)) { + return; + } + if (status.available) { + status.available = false; + this.fallbackAdapter.emitAvailabilityChanged(sttInstance, false); + } + this._logger.warn({ stt: sttInstance.label }, 'STT failed, switching to next provider'); + } catch (e) { + if (this.abortSignal.aborted) return; this._logger.warn( - { stt: sttInstance.label, err: e }, - `${sttInstance.label} unexpected error, switching to next STT`, + { + stt: sttInstance.label, + errorType: e instanceof Error ? e.constructor.name : typeof e, + }, + 'STT failed, switching to next provider', ); + if (status.available) { + status.available = false; + this.fallbackAdapter.emitAvailabilityChanged(sttInstance, false); + } + } finally { + sttInstance.off('error', errListener); + mainRef.current = null; } - if (status.available) { - status.available = false; - this.fallbackAdapter.emitAvailabilityChanged(sttInstance, false); - } - } finally { - sttInstance.off('error', errListener); - mainRef.current = null; - } - this.tryRecoverStream(sttInstance); - } - - // Terminal failure: drain + cancel the forwarder and every live probe - // task before throwing. - try { - this.input.close(); - } catch { - /* already closed */ - } - if (!forwarderTask.done) { - await cancelAndWait([forwarderTask], 1000); - } - const liveProbeTasks: Task[] = []; - for (let i = 0; i < this.fallbackAdapter.sttInstances.length; i++) { - const s = this.fallbackAdapter.status[i]; - if (s?.recoveringStreamTask && !s.recoveringStreamTask.done) { - liveProbeTasks.push(s.recoveringStreamTask); - } - } - if (liveProbeTasks.length > 0) { - await cancelAndWait(liveProbeTasks, 1000); - } - for (const probe of [...this.recoveringStreams]) { - try { - probe.close(); - } catch { - /* already closed */ + this.tryRecoverStream(sttInstance); } - } - const labels = this.fallbackAdapter.sttInstances.map((s) => s.label).join(', '); - throw new APIConnectionError({ - message: `all STTs failed (${labels}) after ${Date.now() - startTime}ms`, - }); + if (this.abortSignal.aborted) return; + const labels = this.fallbackAdapter.sttInstances.map((s) => s.label).join(', '); + throw new APIConnectionError({ + message: `all STTs failed (${labels}) after ${Date.now() - startTime}ms`, + }); + } finally { + this.abortSignal.removeEventListener('abort', closeStreams); + if (!this.input.closed) this.input.close(); + const tasks = [forwarderTask, ...this.recoveringStreams.values()]; + closeStreams(); + await cancelAndWait(tasks, 1000); + } } } diff --git a/agents/src/stt/fallback_adapter_lifecycle.test.ts b/agents/src/stt/fallback_adapter_lifecycle.test.ts new file mode 100644 index 0000000000..530e50a38e --- /dev/null +++ b/agents/src/stt/fallback_adapter_lifecycle.test.ts @@ -0,0 +1,647 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { AudioFrame } from '@livekit/rtc-node'; +import type { EventEmitter } from 'node:events'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { APIError } from '../_exceptions.js'; +import { asLanguageCode } from '../language.js'; +import { log } from '../log.js'; +import type { APIConnectOptions } from '../types.js'; +import { Future, delay } from '../utils.js'; +import { Agent, AgentTask } from '../voice/agent.js'; +import { AgentSession } from '../voice/agent_session.js'; +import { FallbackAdapter } from './fallback_adapter.js'; +import { STT, type SpeechEvent, SpeechEventType, SpeechStream } from './stt.js'; + +class ControlledSTT extends STT { + streams: ControlledStream[] = []; + + constructor(public label: string) { + super({ streaming: true, interimResults: true }); + } + + protected async _recognize(): Promise { + throw new Error('not used'); + } + + stream(options?: { connOptions?: APIConnectOptions }): ControlledStream { + const stream = new ControlledStream(this, undefined, options?.connOptions); + this.streams.push(stream); + return stream; + } +} + +/** Simulates a provider waiting for its final response after input EOF. */ +class ControlledStream extends SpeechStream { + label = 'controlled-stream'; + readonly started = new Future(); + private completion = new Future(); + + get isClosed(): boolean { + return this.closed; + } + + emitText(text: string): void { + if (this.closed) return; + this.queue.put({ + type: SpeechEventType.FINAL_TRANSCRIPT, + alternatives: [ + { text, language: asLanguageCode('en'), startTime: 0, endTime: 1, confidence: 1 }, + ], + }); + } + + fail(error: Error = new APIError('provider connection ended')): void { + this.completion.reject(error); + } + + finish(): void { + this.completion.resolve(); + } + + protected async run(): Promise { + this.started.resolve(); + if (this.abortSignal.aborted) return; + const onAbort = () => this.completion.resolve(); + this.abortSignal.addEventListener('abort', onAbort, { once: true }); + try { + await this.completion.await; + } finally { + this.abortSignal.removeEventListener('abort', onAbort); + } + } +} + +async function getStream(provider: ControlledSTT, index = 0): Promise { + await vi.waitFor(() => expect(provider.streams[index]).toBeDefined()); + const stream = provider.streams[index]!; + await stream.started.await; + return stream; +} + +describe('FallbackSpeechStream lifecycle', () => { + let primary: ControlledSTT; + let secondary: ControlledSTT; + let adapter: FallbackAdapter; + let streams: SpeechStream[]; + let availability: Array<{ label: string; available: boolean }>; + + beforeEach(() => { + primary = new ControlledSTT('primary'); + secondary = new ControlledSTT('secondary'); + adapter = new FallbackAdapter({ + sttInstances: [primary, secondary], + maxRetryPerSTT: 0, + }); + availability = []; + (adapter as unknown as EventEmitter).on( + 'stt_availability_changed', + ({ stt, available }: { stt: STT; available: boolean }) => { + availability.push({ label: stt.label, available }); + }, + ); + streams = []; + const createStream = adapter.stream.bind(adapter); + vi.spyOn(adapter, 'stream').mockImplementation((options) => { + const stream = createStream(options); + streams.push(stream); + return stream; + }); + }); + + afterEach(async () => { + for (const stream of streams) stream.close(); + for (const provider of [primary, secondary]) { + for (const stream of provider.streams) stream.close(); + } + await adapter.close(); + vi.restoreAllMocks(); + }); + + it('does not open a provider if closed before its run starts', async () => { + adapter.stream().close(); + await delay(0); + + expect(primary.streams).toHaveLength(0); + expect(secondary.streams).toHaveLength(0); + expect(availability).toEqual([]); + }); + + it('closes an idle child without waiting for another provider event', async () => { + const stream = adapter.stream(); + const child = await getStream(primary); + + stream.close(); + stream.close(); + + expect(child.isClosed).toBe(true); + await vi.waitFor(() => expect(primary.listenerCount('error')).toBe(0)); + expect(availability).toEqual([]); + }); + + it.each(['transcript', 'error'] as const)( + 'discards an in-flight %s when the parent closes', + async (event) => { + const stream = adapter.stream(); + const child = await getStream(primary); + if (event === 'transcript') child.emitText('late transcript'); + else child.fail(); + stream.close(); + + await vi.waitFor(() => expect(primary.listenerCount('error')).toBe(0)); + expect(await stream.next()).toEqual({ done: true, value: undefined }); + expect(availability).toEqual([]); + expect(adapter.status.map((status) => status.available)).toEqual([true, true]); + expect(primary.streams).toHaveLength(1); + expect(secondary.streams).toHaveLength(0); + + const replacement = adapter.stream(); + const replacementChild = await getStream(primary, 1); + replacementChild.emitText('replacement primary'); + expect((await replacement.next()).value?.alternatives?.[0]?.text).toBe('replacement primary'); + }, + ); + + it('ignores a synchronous provider failure during parent closure', async () => { + const createStream = vi.spyOn(primary, 'stream').mockImplementationOnce(() => { + stream.close(); + throw new APIError('provider closed during setup'); + }); + const stream = adapter.stream(); + + await vi.waitFor(() => expect(createStream).toHaveBeenCalled()); + expect(createStream).toHaveBeenCalledTimes(1); + expect(secondary.streams).toHaveLength(0); + expect(availability).toEqual([]); + }); + + it('does not start recovery or fallback if an availability listener closes the parent', async () => { + const stream = adapter.stream(); + const child = await getStream(primary); + (adapter as unknown as EventEmitter).on('stt_availability_changed', () => stream.close()); + + child.fail(); + + await vi.waitFor(() => expect(availability).toEqual([{ label: 'primary', available: false }])); + expect(primary.streams).toHaveLength(1); + expect(secondary.streams).toHaveLength(0); + }); + + it('closes an idle recovery probe and settles its task when the parent closes', async () => { + const stream = adapter.stream(); + const child = await getStream(primary); + child.fail(); + const probe = await getStream(primary, 1); + const fallback = await getStream(secondary); + const recoveryTask = adapter.status[0]!.recoveringStreamTask; + + stream.close(); + + expect(probe.isClosed).toBe(true); + expect(fallback.isClosed).toBe(true); + await recoveryTask!.result; + await vi.waitFor(() => expect(secondary.listenerCount('error')).toBe(0)); + expect(primary.listenerCount('error')).toBe(0); + expect(availability).toEqual([{ label: 'primary', available: false }]); + }); + + it("does not cancel another stream's recovery on clean EOF", async () => { + adapter.stream(); + const child = await getStream(primary); + child.fail(); + const probe = await getStream(primary, 1); + await getStream(secondary); + + const other = adapter.stream(); + const otherChild = await getStream(secondary, 1); + other.endInput(); + otherChild.finish(); + expect((await other.next()).done).toBe(true); + + probe.emitText('primary recovered'); + await vi.waitFor(() => expect(adapter.status[0]!.available).toBe(true)); + expect(availability).toEqual([ + { label: 'primary', available: false }, + { label: 'primary', available: true }, + ]); + }); + + it('ends a healthy stream after another child from the same provider fails', async () => { + adapter.stream(); + const failingChild = await getStream(primary); + const healthy = adapter.stream(); + const healthyChild = await getStream(primary, 1); + + failingChild.fail(); + await getStream(secondary); + healthy.endInput(); + healthyChild.finish(); + + let ended = false; + const completion = healthy.next().then((event) => { + ended = !!event.done; + }); + await vi.waitFor(() => expect(ended).toBe(true)); + await completion; + expect(secondary.streams).toHaveLength(1); + expect(availability).toEqual([{ label: 'primary', available: false }]); + }); + + it('keeps recovery alive when another stream with an active probe closes', async () => { + const owner = adapter.stream(); + (await getStream(primary)).fail(); + const oldProbe = await getStream(primary, 1); + await getStream(secondary); + + const replacement = adapter.stream(); + const replacementFallback = await getStream(secondary, 1); + expect(primary.streams).toHaveLength(2); + const oldPushFrame = vi.spyOn(oldProbe, 'pushFrame'); + const fallbackPushFrame = vi.spyOn(replacementFallback, 'pushFrame'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + replacement.pushFrame(frame); + await vi.waitFor(() => expect(fallbackPushFrame).toHaveBeenCalledWith(frame)); + expect(oldPushFrame).not.toHaveBeenCalled(); + owner.close(); + + const probe = await getStream(primary, 2); + expect(oldProbe.isClosed).toBe(true); + expect(probe.isClosed).toBe(false); + const pushFrame = vi.spyOn(probe, 'pushFrame'); + replacement.pushFrame(frame); + await vi.waitFor(() => expect(pushFrame).toHaveBeenCalledWith(frame)); + + probe.emitText('primary recovered in replacement'); + await vi.waitFor(() => expect(adapter.status[0]!.available).toBe(true)); + expect(availability).toEqual([ + { label: 'primary', available: false }, + { label: 'primary', available: true }, + ]); + }); + + it.each(['close', 'recover'] as const)('does not start a waiting probe on %s', async (action) => { + adapter.stream(); + (await getStream(primary)).fail(); + const firstProbe = await getStream(primary, 1); + await getStream(secondary); + adapter.stream(); + await getStream(secondary, 1); + expect(primary.streams).toHaveLength(2); + + if (action === 'close') { + await adapter.close(); + } else { + firstProbe.emitText('primary recovered'); + } + + await vi.waitFor(() => expect(primary.listenerCount('error')).toBe(0)); + expect(firstProbe.isClosed).toBe(true); + expect(primary.streams).toHaveLength(2); + expect(secondary.streams.every((stream) => !stream.isClosed)).toBe(true); + expect(adapter.status[0]!.recoveringStreamTask).toBeNull(); + expect(availability).toEqual([ + { label: 'primary', available: false }, + ...(action !== 'close' ? [{ label: 'primary', available: true }] : []), + ]); + }); + + it('skips a closed waiting stream when transferring recovery', async () => { + const owner = adapter.stream(); + (await getStream(primary)).fail(); + const oldProbe = await getStream(primary, 1); + await getStream(secondary); + const waiting = adapter.stream(); + await getStream(secondary, 1); + const replacement = adapter.stream(); + await getStream(secondary, 2); + expect(primary.streams).toHaveLength(2); + + waiting.close(); + owner.close(); + + const probe = await getStream(primary, 2); + expect(oldProbe.isClosed).toBe(true); + expect(primary.streams).toHaveLength(3); + const pushFrame = vi.spyOn(probe, 'pushFrame'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + replacement.pushFrame(frame); + await vi.waitFor(() => expect(pushFrame).toHaveBeenCalledWith(frame)); + }); + + it('continues transcription when a child exits before parent input EOF', async () => { + const stream = adapter.stream(); + (await getStream(primary)).finish(); + + const fallback = await getStream(secondary); + const pushFrame = vi.spyOn(fallback, 'pushFrame'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + stream.pushFrame(frame); + await vi.waitFor(() => expect(pushFrame).toHaveBeenCalledWith(frame)); + fallback.emitText('continued transcription'); + expect((await stream.next()).value?.alternatives?.[0]?.text).toBe('continued transcription'); + expect(availability).toEqual([{ label: 'primary', available: false }]); + }); + + it('keeps forwarding main audio when the recovery probe rejects input', async () => { + const stream = adapter.stream(); + (await getStream(primary)).fail(); + const probe = await getStream(primary, 1); + const fallback = await getStream(secondary); + vi.spyOn(probe, 'pushFrame').mockImplementation(() => { + throw new Error('probe input closed'); + }); + vi.spyOn(probe, 'flush').mockImplementation(() => { + throw new Error('probe input closed'); + }); + const pushFrame = vi.spyOn(fallback, 'pushFrame'); + const flush = vi.spyOn(fallback, 'flush'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + + stream.pushFrame(frame); + stream.flush(); + + await vi.waitFor(() => expect(flush).toHaveBeenCalledOnce()); + expect(pushFrame).toHaveBeenCalledWith(frame); + fallback.emitText('fallback still transcribes'); + expect((await stream.next()).value?.alternatives?.[0]?.text).toBe('fallback still transcribes'); + }); + + it('tries the next waiting stream if a replacement probe throws during setup', async () => { + const owner = adapter.stream(); + (await getStream(primary)).fail(); + await getStream(primary, 1); + await getStream(secondary); + adapter.stream(); + await getStream(secondary, 1); + const replacement = adapter.stream(); + await getStream(secondary, 2); + vi.spyOn(primary, 'stream').mockImplementationOnce(() => { + throw new Error('probe setup failed'); + }); + + owner.close(); + + const probe = await getStream(primary, 2); + const pushFrame = vi.spyOn(probe, 'pushFrame'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + replacement.pushFrame(frame); + await vi.waitFor(() => expect(pushFrame).toHaveBeenCalledWith(frame)); + }); + + it('retries the primary directly when every provider is unavailable', async () => { + for (const status of adapter.status) status.available = false; + const stream = adapter.stream(); + const child = await getStream(primary); + expect(secondary.streams).toHaveLength(0); + expect(adapter.status.every((status) => status.recoveringStreamTask === null)).toBe(true); + const pushFrame = vi.spyOn(child, 'pushFrame'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + + stream.pushFrame(frame); + await vi.waitFor(() => expect(pushFrame).toHaveBeenCalledWith(frame)); + child.emitText('direct retry succeeded'); + stream.endInput(); + child.finish(); + + expect((await stream.next()).value?.alternatives?.[0]?.text).toBe('direct retry succeeded'); + expect((await stream.next()).done).toBe(true); + expect(availability).toEqual([]); + expect(adapter.status.map((status) => status.available)).toEqual([false, false]); + }); + + it('retries the unavailable secondary after the unavailable primary fails', async () => { + for (const status of adapter.status) status.available = false; + const stream = adapter.stream(); + (await getStream(primary)).fail(); + const probe = await getStream(primary, 1); + const fallback = await getStream(secondary); + const pushFrame = vi.spyOn(fallback, 'pushFrame'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + stream.pushFrame(frame); + await vi.waitFor(() => expect(pushFrame).toHaveBeenCalledWith(frame)); + + probe.emitText('probe transcript'); + await vi.waitFor(() => expect(adapter.status[0]!.available).toBe(true)); + fallback.emitText('secondary transcript'); + stream.endInput(); + fallback.finish(); + + const texts: string[] = []; + for await (const event of stream) texts.push(event.alternatives![0].text); + expect(texts).toEqual(['secondary transcript']); + expect(availability).toEqual([{ label: 'primary', available: true }]); + expect(probe.isClosed).toBe(true); + }); + + it('retries normally alongside an existing probe without creating another probe', async () => { + const owner = adapter.stream(); + (await getStream(primary)).fail(); + const firstProbe = await getStream(primary, 1); + await getStream(secondary); + const recoveryTask = adapter.status[0]!.recoveringStreamTask; + adapter.status[1]!.available = false; + + const retry = adapter.stream(); + const retryChild = await getStream(primary, 2); + expect(adapter.status[0]!.recoveringStreamTask).toBe(recoveryTask); + expect(firstProbe.isClosed).toBe(false); + retryChild.fail(); + await getStream(secondary, 1); + expect(primary.streams).toHaveLength(3); + expect(adapter.status[0]!.recoveringStreamTask).toBe(recoveryTask); + + owner.close(); + + const replacementProbe = await getStream(primary, 3); + expect(firstProbe.isClosed).toBe(true); + const pushFrame = vi.spyOn(replacementProbe, 'pushFrame'); + const frame = new AudioFrame(new Int16Array(160), 16_000, 1, 160); + retry.pushFrame(frame); + await vi.waitFor(() => expect(pushFrame).toHaveBeenCalledWith(frame)); + }); + + it.each([APIError, Error])( + 'falls back after an input-ended child fails with %s', + async (ErrorType) => { + const stream = adapter.stream(); + const child = await getStream(primary); + stream.endInput(); + child.fail(new ErrorType('terminal failure')); + + const fallback = await getStream(secondary); + fallback.emitText('final fallback transcript'); + fallback.finish(); + + expect(child._failed).toBe(true); + expect((await stream.next()).value?.alternatives?.[0]?.text).toBe( + 'final fallback transcript', + ); + expect((await stream.next()).done).toBe(true); + }, + ); + + it('ignores a recognize recovery result after adapter shutdown', async () => { + const recovery = new Future(); + const result: SpeechEvent = { + type: SpeechEventType.FINAL_TRANSCRIPT, + alternatives: [ + { + text: 'recognized', + language: asLanguageCode('en'), + startTime: 0, + endTime: 1, + confidence: 1, + }, + ], + }; + vi.spyOn(primary, 'recognize') + .mockRejectedValueOnce(new APIError('primary failed')) + .mockImplementationOnce(() => recovery.await); + vi.spyOn(secondary, 'recognize').mockResolvedValue(result); + await adapter.recognize(new AudioFrame(new Int16Array(160), 16_000, 1, 160)); + const task = adapter.status[0]!.recoveringRecognizeTask!; + try { + await adapter.close(); + expect(task.done).toBe(false); + } finally { + recovery.resolve(result); + await task.result; + } + + expect(adapter.status[0]!.available).toBe(false); + expect(availability).toEqual([{ label: 'primary', available: false }]); + await vi.waitFor(() => expect(adapter.status[0]!.recoveringRecognizeTask).toBeNull()); + }); + + it('closes recovery probes when every unavailable provider fails its normal retry', async () => { + for (const status of adapter.status) status.available = false; + adapter.on('error', () => {}); + const stream = adapter.stream(); + (await getStream(primary)).fail(); + const primaryProbe = await getStream(primary, 1); + const fallback = await getStream(secondary); + + fallback.fail(); + + expect((await stream.next()).done).toBe(true); + await vi.waitFor(() => + expect(adapter.status.every((status) => status.recoveringStreamTask === null)).toBe(true), + ); + expect(primaryProbe.isClosed).toBe(true); + expect(primary.streams).toHaveLength(2); + expect(secondary.streams).toHaveLength(2); + expect(secondary.streams.every((child) => child.isClosed)).toBe(true); + expect(availability).toEqual([]); + }); + + it.each([ + ['APIError', APIError], + ['Error', Error], + ] as const)('logs safe metadata for a %s provider error', async (_name, ErrorType) => { + const error = new ErrorType('secret provider response'); + error.cause = new Error('secret credentials'); + const warn = vi.spyOn(log(), 'warn'); + vi.spyOn(primary, 'stream').mockImplementation(() => { + throw error; + }); + + adapter.stream(); + await getStream(secondary); + + expect(warn).toHaveBeenCalledWith( + { stt: 'primary', errorType: ErrorType.name }, + 'STT failed, switching to next provider', + ); + for (const [attributes] of warn.mock.calls) { + expect(attributes).not.toHaveProperty('err'); + expect(attributes).not.toHaveProperty('error'); + } + }); + + it('transfers the recovery probe through an AgentTask handoff', async () => { + const beginHandoff = new Future(); + const task = AgentTask.create({ instructions: 'question' }); + const parent = Agent.create({ + instructions: 'parent', + onEnter: async () => { + await beginHandoff.await; + await task.run(); + }, + }); + const session = new AgentSession({ + stt: adapter, + vad: null, + turnDetection: 'manual', + turnHandling: { interruption: { enabled: false } }, + }); + try { + await session.start({ agent: parent }); + (await getStream(primary)).fail(); + const oldProbe = await getStream(primary, 1); + await getStream(secondary); + + beginHandoff.resolve(); + await getStream(secondary, 1); + const probe = await getStream(primary, 2); + + expect(session.currentAgent).toBe(task); + expect(oldProbe.isClosed).toBe(true); + expect(primary.streams.filter((stream) => !stream.isClosed)).toEqual([probe]); + probe.emitText('primary recovered in task'); + await vi.waitFor(() => expect(adapter.status[0]!.available).toBe(true)); + expect(availability).toEqual([ + { label: 'primary', available: false }, + { label: 'primary', available: true }, + ]); + } finally { + beginHandoff.resolve(); + if (!task.done) task.complete(undefined); + await session.close(); + } + }); + + it('preserves recovery through an AgentTask handoff', async () => { + const beginHandoff = new Future(); + const task = AgentTask.create({ instructions: 'question' }); + const parent = Agent.create({ + instructions: 'parent', + onEnter: async () => { + await beginHandoff.await; + await task.run(); + }, + }); + const session = new AgentSession({ + stt: adapter, + vad: null, + turnDetection: 'manual', + turnHandling: { interruption: { enabled: false } }, + }); + try { + await session.start({ agent: parent }); + const oldChild = await getStream(primary); + beginHandoff.resolve(); + const taskChild = await getStream(primary, 1); + expect(session.currentAgent).toBe(task); + expect(adapter.stream).toHaveBeenCalledTimes(2); + + taskChild.fail(); + const probe = await getStream(primary, 2); + await getStream(secondary); + oldChild.emitText('late parent transcript'); + await delay(0); + probe.emitText('primary recovered in task'); + + await vi.waitFor(() => expect(adapter.status[0]!.available).toBe(true)); + expect(oldChild.isClosed).toBe(true); + expect(availability).toEqual([ + { label: 'primary', available: false }, + { label: 'primary', available: true }, + ]); + } finally { + beginHandoff.resolve(); + if (!task.done) task.complete(undefined); + await session.close(); + } + }); +}); diff --git a/agents/src/stt/stt.ts b/agents/src/stt/stt.ts index 6acbf9f007..04f1f0d6e7 100644 --- a/agents/src/stt/stt.ts +++ b/agents/src/stt/stt.ts @@ -315,6 +315,7 @@ export abstract class SpeechStream implements AsyncIterableIterator abstract label: string; protected closed = false; #stt: STT; + #failed = false; private deferredInputStream: DeferredReadableStream; private logger = log(); private _connOptions: APIConnectOptions; @@ -417,6 +418,7 @@ export abstract class SpeechStream implements AsyncIterableIterator } private emitError({ error, recoverable }: { error: Error; recoverable: boolean }) { + if (!recoverable) this.#failed = true; this.#stt.emit('error', { type: 'stt_error', timestamp: Date.now(), @@ -491,6 +493,11 @@ export abstract class SpeechStream implements AsyncIterableIterator return this.abortController.signal; } + /** Whether this stream ended with an unrecoverable error. @internal */ + get _failed(): boolean { + return this.#failed; + } + get startTimeOffset(): number { return this._startTimeOffset; }