diff --git a/packages/stage-ui/src/libs/inference/adapters/kokoro.test.ts b/packages/stage-ui/src/libs/inference/adapters/kokoro.test.ts index 57fcf5b9e..236f4f10a 100644 --- a/packages/stage-ui/src/libs/inference/adapters/kokoro.test.ts +++ b/packages/stage-ui/src/libs/inference/adapters/kokoro.test.ts @@ -170,6 +170,45 @@ describe('kokoro adapter - device loss resilience', () => { })) }) + it('should not restart the worker when generation is aborted', async () => { + const { createKokoroAdapter } = await import('./kokoro') + const adapter = createKokoroAdapter() + + const loading = adapter.loadModel('q4', 'webgpu') + await vi.waitFor(() => expect(enqueueMock).toHaveBeenCalled()) + + const worker = MockWorker.instances.at(-1)! + const loadRequest = worker.postMessage.mock.calls.find(([message]) => message.type === 'load-model')?.[0] + expect(loadRequest?.requestId).toBeDefined() + worker.dispatch('message', { + data: { + type: 'model-ready', + requestId: loadRequest.requestId, + device: 'webgpu', + metadata: { voices: { af_heart: { name: 'Heart' } } }, + }, + }) + await loading + expect(adapter.state).toBe('ready') + + const controller = new AbortController() + const generating = adapter.generate('hello', 'af_heart' as any, { signal: controller.signal }) + + await vi.waitFor(() => { + expect(worker.postMessage).toHaveBeenCalledWith(expect.objectContaining({ type: 'run-inference' })) + }) + + controller.abort('cancel speech') + + await expect(generating).rejects.toMatchObject({ name: 'AbortError' }) + expect(adapter.state).toBe('ready') + expect(worker.terminate).not.toHaveBeenCalled() + expect(worker.postMessage).toHaveBeenCalledWith(expect.objectContaining({ + type: 'cancel', + targetRequestId: expect.any(String), + })) + }) + it('should classify worker device-loss errors before restarting', async () => { const { createKokoroAdapter } = await import('./kokoro') const adapter = createKokoroAdapter() diff --git a/packages/stage-ui/src/libs/inference/adapters/kokoro.ts b/packages/stage-ui/src/libs/inference/adapters/kokoro.ts index b9aac0570..9961b5dc2 100644 --- a/packages/stage-ui/src/libs/inference/adapters/kokoro.ts +++ b/packages/stage-ui/src/libs/inference/adapters/kokoro.ts @@ -463,6 +463,15 @@ export function createKokoroAdapter(): KokoroAdapter { if (error === notReadyError) throw error + // Cancellation is a caller-controlled lifecycle outcome, not a worker + // failure. Keep the loaded model available and avoid restarting the + // worker after waitForWorkerMessage has already posted `cancel`. + if ((error as Error)?.name === 'AbortError') { + if (state === 'running') + state = 'ready' + throw error + } + handleWorkerError(error instanceof Error ? error : new Error(String(error))) throw error })