fix(stage-ui): fail fast on Whisper worker errors (#1803)

This commit is contained in:
juyua9
2026-07-20 20:52:59 +08:00
committed by GitHub
parent 72b4884dbf
commit 8a3f916641
2 changed files with 153 additions and 2 deletions
@@ -0,0 +1,121 @@
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
class MockWorker {
static instances: MockWorker[] = []
listeners = new Map<string, Set<(event: any) => void>>()
addEventListener = vi.fn((type: string, listener: (event: any) => void) => {
if (!this.listeners.has(type))
this.listeners.set(type, new Set())
this.listeners.get(type)!.add(listener)
})
removeEventListener = vi.fn((type: string, listener: (event: any) => void) => {
this.listeners.get(type)?.delete(listener)
})
postMessage = vi.fn()
terminate = vi.fn()
constructor() {
MockWorker.instances.push(this)
}
dispatch(type: string, event: any): void {
for (const listener of this.listeners.get(type) ?? [])
listener(event)
}
}
vi.stubGlobal('Worker', MockWorker)
vi.mock('../../../composables/use-inference-status', () => ({
removeInferenceStatus: vi.fn(),
updateInferenceStatus: vi.fn(),
}))
const enqueueMock = vi.fn((_id: string, _p: number, loader: () => Promise<unknown>) => loader())
const recordDeviceLoss = vi.fn()
vi.mock('../coordinator', () => ({
getGPUCoordinator: () => ({
recordDeviceLoss,
release: vi.fn(),
requestAllocation: vi.fn(() => ({ estimatedBytes: 0, modelId: 'whisper' })),
}),
getLoadQueue: () => ({
enqueue: enqueueMock,
}),
MODEL_VRAM_ESTIMATES: {},
}))
vi.mock('@proj-airi/stage-shared', () => ({
defaultPerfTracer: {
withMeasure: vi.fn((_category: string, _name: string, fn: () => unknown) => fn()),
},
}))
describe('whisper adapter worker failure handling', () => {
beforeEach(() => {
vi.useFakeTimers()
MockWorker.instances.length = 0
enqueueMock.mockClear()
enqueueMock.mockImplementation((_id: string, _p: number, loader: () => Promise<unknown>) => loader())
recordDeviceLoss.mockClear()
})
afterEach(() => {
vi.useRealTimers()
vi.restoreAllMocks()
})
it('rejects an in-flight model load as soon as the worker errors', async () => {
const { createWhisperAdapter } = await import('./whisper')
const adapter = createWhisperAdapter(new URL('whisper-worker.ts', import.meta.url))
const loading = adapter.load()
await vi.waitFor(() => expect(enqueueMock).toHaveBeenCalled())
const worker = MockWorker.instances.at(-1)!
expect(worker.postMessage).toHaveBeenCalledWith(expect.objectContaining({ type: 'load-model' }))
worker.dispatch('error', { error: new Error('Whisper worker crashed while loading') })
await expect(loading).rejects.toThrow('Whisper worker crashed while loading')
expect(adapter.state).toBe('error')
})
it('rejects an in-flight transcription as soon as the worker errors', async () => {
const { createWhisperAdapter } = await import('./whisper')
const adapter = createWhisperAdapter(new URL('whisper-worker.ts', import.meta.url))
const loading = adapter.load()
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).toBeDefined()
worker.dispatch('message', {
data: {
device: 'webgpu',
modelId: 'whisper',
requestId: loadRequest!.requestId,
type: 'model-ready',
},
})
await loading
expect(adapter.state).toBe('ready')
const transcribing = adapter.transcribe({ audio: 'data:audio/wav;base64,test', language: 'en' })
await vi.waitFor(() => {
expect(worker.postMessage).toHaveBeenCalledWith(expect.objectContaining({ type: 'run-inference' }))
})
worker.dispatch('error', { error: new Error('Whisper worker crashed during transcription') })
await expect(transcribing).rejects.toThrow('Whisper worker crashed during transcription')
expect(adapter.state).toBe('error')
})
})
@@ -94,6 +94,11 @@ export interface WhisperAdapter {
const LOAD_TIMEOUT = TIMEOUTS.WHISPER_LOAD
const TRANSCRIBE_TIMEOUT = TIMEOUTS.WHISPER_TRANSCRIBE
interface PendingWaiter {
cleanup: () => void
reject: (error: Error) => void
}
// ---------------------------------------------------------------------------
// Factory
// ---------------------------------------------------------------------------
@@ -106,6 +111,7 @@ export function createWhisperAdapter(workerUrl: string | URL): WhisperAdapter {
let messageListener: ((event: MessageEvent) => void) | null = null
let errorListener: ((event: ErrorEvent) => void) | null = null
const messageHandlers = new Set<(event: WhisperEvent) => void>()
const pendingWaiters = new Map<string, PendingWaiter>()
// NOTICE: Device-loss resilience state. See kokoro.ts for rationale.
let lastManifest: { device: string } | null = null
@@ -113,16 +119,37 @@ export function createWhisperAdapter(workerUrl: string | URL): WhisperAdapter {
const operationMutex = new Mutex()
function toWorkerError(event: ErrorEvent | Error, fallback = 'Whisper worker failed'): Error {
if (event instanceof Error)
return event
if (event.error instanceof Error)
return event.error
return new Error(event.message || fallback)
}
function rejectPendingWaiters(error: Error): void {
for (const [requestId, waiter] of pendingWaiters) {
waiter.cleanup()
pendingWaiters.delete(requestId)
waiter.reject(error)
}
}
function handleWorkerError(event: ErrorEvent | Error): void {
state = 'error'
operationMutex.cancel()
const code = classifyError(event instanceof Error ? event : (event as ErrorEvent).error ?? event)
const error = toWorkerError(event)
rejectPendingWaiters(error)
const code = classifyError(error)
if (code === 'DEVICE_LOST') {
deviceLossCount++
getGPUCoordinator().recordDeviceLoss({
modelId: MODEL_NAMES.WHISPER,
reason: classifyDeviceLossReason(event instanceof Error ? event : (event as ErrorEvent).error ?? event),
reason: classifyDeviceLossReason(error),
occurredAt: Date.now(),
})
}
@@ -218,6 +245,7 @@ export function createWhisperAdapter(workerUrl: string | URL): WhisperAdapter {
function cleanup(): void {
if (timeoutId !== undefined)
clearTimeout(timeoutId)
pendingWaiters.delete(requestId)
w.removeEventListener('message', handler)
if (abortListener && signal)
signal.removeEventListener('abort', abortListener)
@@ -245,6 +273,7 @@ export function createWhisperAdapter(workerUrl: string | URL): WhisperAdapter {
}
w.addEventListener('message', handler)
pendingWaiters.set(requestId, { cleanup, reject })
timeoutId = setTimeout(() => {
cleanup()
@@ -387,6 +416,7 @@ export function createWhisperAdapter(workerUrl: string | URL): WhisperAdapter {
function terminateAdapter(): void {
operationMutex.cancel()
rejectPendingWaiters(new InferenceAbortError('Whisper adapter terminated.'))
destroyWorker()
if (allocationToken) {
removeInferenceStatus(MODEL_NAMES.WHISPER)