fix(pipelines-audio): no longer concurrently tts
This commit is contained in:
@@ -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([])
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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'],
|
||||
},
|
||||
})
|
||||
Reference in New Issue
Block a user