Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/close-speech-stream-adapters.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
'@livekit/agents': patch
---

Close temporary speech stream adapters when their owning pipeline or fallback lifecycle ends.
38 changes: 38 additions & 0 deletions agents/src/stt/fallback_adapter.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import { APIConnectionError, APIError } from '../_exceptions.js';
import { initializeLogger } from '../log.js';
import type { APIConnectOptions } from '../types.js';
import { type AudioBuffer, delay } from '../utils.js';
import { VAD, VADStream } from '../vad.js';
import { FallbackAdapter } from './fallback_adapter.js';
import { STT, type SpeechEvent, SpeechEventType, SpeechStream } from './stt.js';
import { FakeSTT, RecognizeSentinel, emptyAudioFrame } from './testing/fake_stt.js';
Expand Down Expand Up @@ -82,6 +83,30 @@ class RetryTimelineStream extends SpeechStream {
}
}

class FakeVAD extends VAD {
label = 'fake-vad';

constructor() {
super({ updateInterval: 100 });
}

stream(): VADStream {
return new (class extends VADStream {})(this);
}
}

class NonStreamingSTT extends FakeSTT {
closeCount = 0;

constructor() {
super({ capabilities: { streaming: false, interimResults: false } });
}

override async close(): Promise<void> {
this.closeCount++;
}
}

describe('FallbackAdapter', () => {
beforeAll(() => {
initializeLogger({ pretty: false });
Expand Down Expand Up @@ -281,6 +306,19 @@ describe('FallbackAdapter', () => {

expect(received).toHaveLength(0);
});

it('close closes automatically created stream adapters', async () => {
const stt = new NonStreamingSTT();
const baseline = stt.listenerCount('metrics_collected');
const adapter = new FallbackAdapter({ sttInstances: [stt], vad: new FakeVAD() });

expect(stt.listenerCount('metrics_collected')).toBe(baseline + 1);

await adapter.close();

expect(stt.listenerCount('metrics_collected')).toBe(baseline);
expect(stt.closeCount).toBe(0);
});
});

describe('FallbackSpeechStream (streaming path)', () => {
Expand Down
13 changes: 10 additions & 3 deletions agents/src/stt/fallback_adapter.ts
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ export class FallbackAdapter extends STT {
private _status: STTStatus[] = [];
private _logger = log();
private _metricsForwarders = new Map<STT, (m: STTMetrics) => void>();
private _ownedStreamAdapters: StreamAdapter[] = [];
// Last child that produced output or returned a recognize result. Surfaced
// via the dynamic label/model/provider getters so OTel attributes like
// `gen_ai.request.model` on `user_turn` (refreshed on every STT event by
Expand All @@ -124,9 +125,13 @@ export class FallbackAdapter extends STT {
);
}

const wrapped = opts.sttInstances.map((s) =>
s.capabilities.streaming ? s : new StreamAdapter(s, opts.vad!),
);
const ownedStreamAdapters: StreamAdapter[] = [];
const wrapped = opts.sttInstances.map((s) => {
if (s.capabilities.streaming) return s;
const adapter = new StreamAdapter(s, opts.vad!);
ownedStreamAdapters.push(adapter);
return adapter;
});

// Pick the primary's granularity only if every instance supports aligned
// transcripts — otherwise consumers can't rely on a consistent format.
Expand All @@ -145,6 +150,7 @@ export class FallbackAdapter extends STT {
});

this.sttInstances = wrapped;
this._ownedStreamAdapters = ownedStreamAdapters;
this.attemptTimeoutMs = opts.attemptTimeoutMs ?? 10_000;
this.maxRetryPerSTT = opts.maxRetryPerSTT ?? 1;
this.retryIntervalMs = opts.retryIntervalMs ?? 5_000;
Expand Down Expand Up @@ -327,6 +333,7 @@ export class FallbackAdapter extends STT {
if (m) s.off('metrics_collected' as keyof STTCallbacks, m);
}
this._metricsForwarders.clear();
await Promise.all(this._ownedStreamAdapters.map((adapter) => adapter.close()));
}
}

Expand Down
24 changes: 17 additions & 7 deletions agents/src/stt/stream_adapter.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4,19 +4,28 @@
import type { AudioFrame } from '@livekit/rtc-node';
import { ThrowsPromise } from '@livekit/throws-transformer/throws';
import { log } from '../log.js';
import type { STTMetrics } from '../metrics/base.js';
import type { APIConnectOptions } from '../types.js';
import { isStreamClosedError } from '../utils.js';
import type { VAD, VADStream } from '../vad.js';
import { VADEventType } from '../vad.js';
import type { ConversationItemAddedEvent } from '../voice/events.js';
import type { SpeechEvent } from './stt.js';
import type { STTError, SpeechEvent } from './stt.js';
import { STT, SpeechEventType, SpeechStream } from './stt.js';

export class StreamAdapter extends STT {
#stt: STT;
#vad: VAD;
label: string;

#forwardMetrics = (metrics: STTMetrics) => {
this.emit('metrics_collected', metrics);
};

#forwardError = (error: STTError) => {
this.emit('error', error);
};

constructor(stt: STT, vad: VAD) {
super({
streaming: true,
Expand All @@ -28,13 +37,14 @@ export class StreamAdapter extends STT {
this.#vad = vad;
this.label = `stt.StreamAdapter<${this.#stt.label}>`;

this.#stt.on('metrics_collected', (metrics) => {
this.emit('metrics_collected', metrics);
});
this.#stt.on('metrics_collected', this.#forwardMetrics);
this.#stt.on('error', this.#forwardError);
}

this.#stt.on('error', (error) => {
this.emit('error', error);
});
async close(): Promise<void> {
this.#stt.off('metrics_collected', this.#forwardMetrics);
this.#stt.off('error', this.#forwardError);
await super.close();
}

_recognize(frame: AudioFrame, abortSignal?: AbortSignal): Promise<SpeechEvent> {
Expand Down
117 changes: 113 additions & 4 deletions agents/src/tts/fallback_adapter.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ class MockSynthesizeStream extends SynthesizeStream {
constructor(
private mockTts: MockTTS,
private shouldFail: boolean,
private blocked: boolean,
connOptions?: APIConnectOptions,
) {
super(mockTts, connOptions);
Expand All @@ -36,6 +37,13 @@ class MockSynthesizeStream extends SynthesizeStream {
}

protected async run(): Promise<void> {
if (this.blocked) {
if (this.abortSignal.aborted) return;
await new Promise<void>((resolve) => {
this.abortSignal.addEventListener('abort', () => resolve(), { once: true });
});
return;
}
if (this.shouldFail) {
if (this.mockTts.failAfterInput) {
// Simulate a provider that receives text but dies before emitting
Expand Down Expand Up @@ -78,11 +86,19 @@ class MockChunkedStream extends ChunkedStream {
private mockTts: MockTTS,
text: string,
private shouldFail: boolean,
private blocked: boolean,
connOptions?: APIConnectOptions,
) {
super(text, mockTts, connOptions);
}
protected async run(): Promise<void> {
if (this.blocked) {
if (this.abortSignal.aborted) return;
await new Promise<void>((resolve) => {
this.abortSignal.addEventListener('abort', () => resolve(), { once: true });
});
return;
}
if (this.shouldFail) {
throw new APIError('mock TTS failed immediately');
}
Expand All @@ -98,25 +114,50 @@ class MockChunkedStream extends ChunkedStream {
class MockTTS extends TTS {
label: string;
shouldFail = false;
blocked = false;
/** When failing, first consume a token (and mark started) before throwing. */
failAfterInput = false;
/** Simulated latency between receiving text and sending it to the provider. */
sendDelayMs = 0;
/** The started time the stream recorded when it "sent" text to the provider. */
lastMarkedTime?: number;

constructor(label: string, sampleRate: number = SAMPLE_RATE) {
super(sampleRate, 1, { streaming: true });
closeCount = 0;

constructor(label: string, sampleRate: number = SAMPLE_RATE, streaming = true) {
super(sampleRate, 1, { streaming });
this.label = label;
}

synthesize(text: string, connOptions?: APIConnectOptions): ChunkedStream {
return new MockChunkedStream(this, text, this.shouldFail, connOptions);
return new MockChunkedStream(this, text, this.shouldFail, this.blocked, connOptions);
}

stream(options?: { connOptions?: APIConnectOptions }): SynthesizeStream {
return new MockSynthesizeStream(this, this.shouldFail, options?.connOptions);
return new MockSynthesizeStream(this, this.shouldFail, this.blocked, options?.connOptions);
}

override async close(): Promise<void> {
this.closeCount++;
}
}

function textInput(text = 'hello test'): ReadableStream<string> {
return new ReadableStream({
start(controller) {
controller.enqueue(text);
controller.close();
},
});
}

async function consume(stream: SynthesizeStream): Promise<number> {
stream.updateInputStream(textInput());
let frames = 0;
for await (const event of stream) {
if (event !== SynthesizeStream.END_OF_STREAM) frames++;
}
return frames;
}

describe('TTS FallbackAdapter', () => {
Expand All @@ -126,6 +167,74 @@ describe('TTS FallbackAdapter', () => {
process.on('unhandledRejection', () => {});
});

it('closes temporary stream adapters after each request', async () => {
const nonStreaming = new MockTTS('non-streaming', SAMPLE_RATE, false);
const adapter = new FallbackAdapter({ ttsInstances: [nonStreaming] });
const baseline = nonStreaming.listenerCount('metrics_collected');

try {
for (let i = 0; i < 3; i++) {
expect(await consume(adapter.stream())).toBeGreaterThan(0);
expect(nonStreaming.listenerCount('metrics_collected')).toBe(baseline);
}
expect(nonStreaming.closeCount).toBe(0);
} finally {
await adapter.close();
}
});

it('closes a temporary stream adapter after failure and fallback', async () => {
const nonStreaming = new MockTTS('non-streaming', SAMPLE_RATE, false);
nonStreaming.shouldFail = true;
const adapter = new FallbackAdapter({
ttsInstances: [nonStreaming, new MockTTS('fallback')],
maxRetryPerTTS: 0,
recoveryDelayMs: 60_000,
});
const baseline = nonStreaming.listenerCount('metrics_collected');

try {
expect(await consume(adapter.stream())).toBeGreaterThan(0);
expect(nonStreaming.listenerCount('metrics_collected')).toBe(baseline);
expect(nonStreaming.closeCount).toBe(0);
} finally {
await adapter.close();
}
});

it('closes a temporary stream adapter after cancellation', async () => {
const nonStreaming = new MockTTS('non-streaming', SAMPLE_RATE, false);
nonStreaming.blocked = true;
const adapter = new FallbackAdapter({ ttsInstances: [nonStreaming] });
const baseline = nonStreaming.listenerCount('metrics_collected');
const stream = adapter.stream();
stream.updateInputStream(textInput());

try {
const deadline = Date.now() + 1_000;
while (
nonStreaming.listenerCount('metrics_collected') === baseline &&
Date.now() < deadline
) {
await new Promise((resolve) => setTimeout(resolve, 5));
}
expect(nonStreaming.listenerCount('metrics_collected')).toBeGreaterThan(baseline);

stream.close();
const closeDeadline = Date.now() + 1_000;
while (
nonStreaming.listenerCount('metrics_collected') > baseline &&
Date.now() < closeDeadline
) {
await new Promise((resolve) => setTimeout(resolve, 5));
}
expect(nonStreaming.listenerCount('metrics_collected')).toBe(baseline);
} finally {
stream.close();
await adapter.close();
}
});

it('should fall back to the next TTS when the primary stream fails before any pushText', async () => {
const primary = new MockTTS('primary');
primary.shouldFail = true;
Expand Down
16 changes: 13 additions & 3 deletions agents/src/tts/fallback_adapter.ts
Original file line number Diff line number Diff line change
Expand Up @@ -440,7 +440,6 @@ class FallbackSynthesizeStream extends SynthesizeStream {
})();

for (let i = 0; i < this.adapter.ttsInstances.length; i++) {
const tts = this.adapter.getStreamingInstance(i);
const originalTts = this.adapter.ttsInstances[i]!;
const status = this.adapter.status[i]!;
let lastRequestId: string = '';
Expand All @@ -450,7 +449,11 @@ class FallbackSynthesizeStream extends SynthesizeStream {
this.adapter.markUnAvailable(i);
continue;
}
const resampler = this.adapter.createResamplerForTTS(i);
const tts = this.adapter.getStreamingInstance(i);
let stream!: SynthesizeStream;
let resampler: AudioResampler | null = null;
const closeStream = () => stream?.close();
this.abortSignal.addEventListener('abort', closeStream, { once: true });

// ttfb measures the fallback adapter as a whole: anchor on the first
// time a sentence was handed to any underlying TTS — even one that
Expand All @@ -459,14 +462,16 @@ class FallbackSynthesizeStream extends SynthesizeStream {
let captureStartedTime: () => void = () => {};

try {
resampler = this.adapter.createResamplerForTTS(i);
this._logger.debug({ tts: originalTts.label }, 'attempting TTS stream');

const connOptions: APIConnectOptions = {
...this.connOptions,
maxRetry: this.adapter.maxRetryPerTTS,
};

const stream = tts.stream({ connOptions });
stream = tts.stream({ connOptions });
if (this.abortSignal.aborted) stream.close();
let bufferIndex = 0;
let streamOutputCompleted = false;

Expand Down Expand Up @@ -610,7 +615,12 @@ class FallbackSynthesizeStream extends SynthesizeStream {
// the stream may have received text and failed before emitting audio;
// its started time must still anchor the fallback's ttfb
captureStartedTime();
this.abortSignal.removeEventListener('abort', closeStream);
stream?.close();
resampler?.close();
if (tts !== originalTts) {
await tts.close();
}
}
}
await readInputLLMStream.catch(() => {});
Expand Down
Loading
Loading