mirror of
https://github.com/moeru-ai/airi.git
synced 2026-08-14 08:52:42 +00:00
fix(stage-ui): fail fast on Whisper worker errors (#1803)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user