fix(pipelines-audio): no longer concurrently tts

This commit is contained in:
Neko Ayaka
2026-03-31 15:51:24 +08:00
parent 660765b1aa
commit 4ba49a243a
5 changed files with 302 additions and 47 deletions
+3 -1
View File
@@ -30,7 +30,9 @@
],
"scripts": {
"build": "tsdown",
"typecheck": "tsc --noEmit"
"typecheck": "tsc --noEmit",
"test": "vitest",
"test:run": "vitest run --config vitest.config.ts"
},
"dependencies": {
"@moeru/eventa": "catalog:",
@@ -0,0 +1,190 @@
import type { PlaybackItem, TextSegment, TextToken, TtsRequest } from './types'
import { describe, expect, it, vi } from 'vitest'
import { createSpeechPipeline } from './speech-pipeline'
function delay(ms: number) {
return new Promise(resolve => setTimeout(resolve, ms))
}
function deferred<T>() {
let resolve!: (value: T | PromiseLike<T>) => void
let reject!: (reason?: unknown) => void
const promise = new Promise<T>((resolvePromise, rejectPromise) => {
resolve = resolvePromise
reject = rejectPromise
})
return {
promise,
resolve,
reject,
}
}
function createSegmenter(texts: string[]) {
return (_tokens: ReadableStream<TextToken>, meta: { streamId: string, intentId: string }) => {
let index = 0
return new ReadableStream<TextSegment>({
pull(controller) {
const text = texts[index]
if (text == null) {
controller.close()
return
}
controller.enqueue({
streamId: meta.streamId,
intentId: meta.intentId,
segmentId: `${meta.streamId}:${index}`,
text,
special: null,
reason: 'flush',
createdAt: Date.now(),
})
index += 1
},
})
}
}
function createPlaybackSpy() {
const scheduled: Array<PlaybackItem<string>> = []
return {
scheduled,
playback: {
schedule(item: PlaybackItem<string>) {
scheduled.push(item)
},
stopAll: vi.fn(),
stopByIntent: vi.fn(),
stopByOwner: vi.fn(),
onStart: vi.fn(),
onEnd: vi.fn(),
onInterrupt: vi.fn(),
onReject: vi.fn(),
},
}
}
describe('createSpeechPipeline', () => {
it('preserves playback order when TTS completes out of order', async () => {
const { scheduled, playback } = createPlaybackSpy()
const pipeline = createSpeechPipeline<string>({
ttsMaxConcurrent: 2,
segmenter: createSegmenter(['first', 'second', 'third']),
playback,
async tts(request) {
if (request.sequence === 0)
await delay(30)
else if (request.sequence === 1)
await delay(5)
return request.text
},
})
const intentFinished = new Promise<void>((resolve) => {
pipeline.on('onIntentEnd', () => resolve())
})
const intent = pipeline.openIntent()
intent.end()
await intentFinished
expect(scheduled.map(item => item.sequence)).toEqual([0, 1, 2])
expect(scheduled.map(item => item.text)).toEqual(['first', 'second', 'third'])
})
it('prefetches TTS requests up to the configured concurrency', async () => {
const { playback } = createPlaybackSpy()
const startedRequests: number[] = []
const pendingRequests = [
deferred<string>(),
deferred<string>(),
deferred<string>(),
]
let inFlight = 0
let maxInFlight = 0
const pipeline = createSpeechPipeline<string>({
ttsMaxConcurrent: 2,
segmenter: createSegmenter(['alpha', 'beta', 'gamma']),
playback,
async tts(request: TtsRequest) {
startedRequests.push(request.sequence)
inFlight += 1
maxInFlight = Math.max(maxInFlight, inFlight)
try {
return await pendingRequests[request.sequence]!.promise
}
finally {
inFlight -= 1
}
},
})
const intentFinished = new Promise<void>((resolve) => {
pipeline.on('onIntentEnd', () => resolve())
})
const intent = pipeline.openIntent()
intent.end()
await delay(0)
expect(startedRequests).toEqual([0, 1])
expect(maxInFlight).toBe(2)
pendingRequests[0]!.resolve('alpha')
pendingRequests[1]!.resolve('beta')
await delay(0)
expect(startedRequests).toEqual([0, 1, 2])
pendingRequests[2]!.resolve('gamma')
await intentFinished
expect(maxInFlight).toBe(2)
})
it('cancels in-flight TTS work without scheduling stale playback', async () => {
const { scheduled, playback } = createPlaybackSpy()
const abortedRequests: number[] = []
const pipeline = createSpeechPipeline<string>({
ttsMaxConcurrent: 2,
segmenter: createSegmenter(['left', 'right']),
playback,
tts(request, signal) {
return new Promise<string | null>((resolve) => {
signal.addEventListener('abort', () => {
abortedRequests.push(request.sequence)
resolve(null)
}, { once: true })
})
},
})
const intentCanceled = new Promise<void>((resolve) => {
pipeline.on('onIntentCancel', () => resolve())
})
const intent = pipeline.openIntent()
intent.end()
await delay(0)
intent.cancel('test-cancel')
await intentCanceled
expect(abortedRequests.sort()).toEqual([0, 1])
expect(scheduled).toEqual([])
})
})
+98 -46
View File
@@ -22,6 +22,12 @@ import { createPushStream } from './stream'
export interface SpeechPipelineOptions<TAudio> {
tts: (request: TtsRequest, signal: AbortSignal) => Promise<TAudio | null>
/**
* Maximum number of concurrent TTS generation tasks. Default is 4. Must be at least 1.
*
* @default 4
*/
ttsMaxConcurrent?: number
playback: {
schedule: (item: PlaybackItem<TAudio>) => void
stopAll: (reason: string) => void
@@ -58,6 +64,7 @@ export function createSpeechPipeline<TAudio>(options: SpeechPipelineOptions<TAud
const logger = options.logger ?? console
const priorityResolver = options.priority ?? createPriorityResolver()
const segmenter = options.segmenter ?? createTtsSegmentStream
const ttsMaxConcurrent = Math.max(1, options.ttsMaxConcurrent ?? 4)
const context = createContext()
const intents = new Map<string, IntentState>()
@@ -86,11 +93,91 @@ export function createSpeechPipeline<TAudio>(options: SpeechPipelineOptions<TAud
const tokenStream = intent.stream
const segmentStream = segmenter(tokenStream, { streamId: intent.streamId, intentId: intent.intentId })
const completedRequests = new Map<number, TtsResult<TAudio> | null>()
const inFlightTasks = new Set<Promise<void>>()
let nextRequestSequence = 0
let nextSequenceToSchedule = 0
function scheduleCompletedRequests() {
while (completedRequests.has(nextSequenceToSchedule)) {
const completedRequest = completedRequests.get(nextSequenceToSchedule) ?? null
completedRequests.delete(nextSequenceToSchedule)
if (completedRequest) {
options.playback.schedule({
id: createId('playback'),
streamId: completedRequest.streamId,
intentId: completedRequest.intentId,
segmentId: completedRequest.segmentId,
sequence: completedRequest.sequence,
ownerId: intent.ownerId,
priority: intent.priority,
text: completedRequest.text,
special: completedRequest.special,
audio: completedRequest.audio,
createdAt: Date.now(),
})
}
nextSequenceToSchedule += 1
}
}
function createTtsTask(request: TtsRequest) {
const task = (async () => {
let audio: TAudio | null = null
try {
audio = await options.tts(request, intent.controller.signal)
}
catch (err) {
logger.warn('TTS generation failed:', err)
if (intent.controller.signal.aborted)
return
}
if (intent.controller.signal.aborted) {
completedRequests.set(request.sequence, null)
scheduleCompletedRequests()
return
}
if (!audio) {
completedRequests.set(request.sequence, null)
scheduleCompletedRequests()
return
}
const ttsResult: TtsResult<TAudio> = {
streamId: request.streamId,
intentId: request.intentId,
segmentId: request.segmentId,
sequence: request.sequence,
text: request.text,
special: request.special,
audio,
createdAt: Date.now(),
}
context.emit(speechPipelineEventMap.onTtsResult, ttsResult)
completedRequests.set(request.sequence, ttsResult)
scheduleCompletedRequests()
})()
.finally(() => {
inFlightTasks.delete(task)
})
inFlightTasks.add(task)
return task
}
try {
const reader = segmentStream.getReader()
while (true) {
while (!intent.controller.signal.aborted && inFlightTasks.size >= ttsMaxConcurrent) {
await Promise.race(inFlightTasks)
}
const { value, done } = await reader.read()
if (done)
break
@@ -112,6 +199,7 @@ export function createSpeechPipeline<TAudio>(options: SpeechPipelineOptions<TAud
streamId: value.streamId,
intentId: value.intentId,
segmentId: value.segmentId,
sequence: nextRequestSequence++,
text: value.text,
special: value.special,
priority: intent.priority,
@@ -119,50 +207,11 @@ export function createSpeechPipeline<TAudio>(options: SpeechPipelineOptions<TAud
}
context.emit(speechPipelineEventMap.onTtsRequest, request)
let audio: TAudio | null = null
try {
audio = await options.tts(request, intent.controller.signal)
}
catch (err) {
logger.warn('TTS generation failed:', err)
if (intent.controller.signal.aborted)
break
continue
}
if (intent.controller.signal.aborted)
break
if (!audio)
continue
const ttsResult: TtsResult<TAudio> = {
streamId: request.streamId,
intentId: request.intentId,
segmentId: request.segmentId,
text: request.text,
special: request.special,
audio,
createdAt: Date.now(),
}
context.emit(speechPipelineEventMap.onTtsResult, ttsResult)
options.playback.schedule({
id: createId('playback'),
streamId: ttsResult.streamId,
intentId: ttsResult.intentId,
segmentId: ttsResult.segmentId,
ownerId: intent.ownerId,
priority: intent.priority,
text: ttsResult.text,
special: ttsResult.special,
audio: ttsResult.audio,
createdAt: Date.now(),
})
createTtsTask(request)
}
await Promise.allSettled(inFlightTasks)
scheduleCompletedRequests()
reader.releaseLock()
}
catch (err) {
@@ -177,11 +226,14 @@ export function createSpeechPipeline<TAudio>(options: SpeechPipelineOptions<TAud
}
intents.delete(intent.intentId)
activeIntent = null
if (activeIntent?.intentId === intent.intentId)
activeIntent = null
const next = pickNextIntent()
if (next)
void runIntent(next)
if (!activeIntent) {
const next = pickNextIntent()
if (next)
void runIntent(next)
}
}
}
+3
View File
@@ -27,6 +27,7 @@ export interface TtsRequest {
streamId: string
intentId: string
segmentId: string
sequence: number
text: string
special: string | null
priority: number
@@ -37,6 +38,7 @@ export interface TtsResult<TAudio> {
streamId: string
intentId: string
segmentId: string
sequence: number
text: string
special: string | null
audio: TAudio
@@ -48,6 +50,7 @@ export interface PlaybackItem<TAudio> {
streamId: string
intentId: string
segmentId: string
sequence: number
ownerId?: string
priority: number
text: string
@@ -0,0 +1,8 @@
import { defineConfig } from 'vitest/config'
export default defineConfig({
test: {
environment: 'node',
include: ['src/**/*.test.ts'],
},
})