WaveCut/Bonsai-Chat-WebGPU
5
1import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';2import { BrowserEngineClient } from './client';3import type { EngineWorkerMessage } from './protocol';4 5class FakeWorker {6 static instances: FakeWorker[] = [];7 8 onmessage: ((event: MessageEvent<EngineWorkerMessage>) => void) | null = null;9 onerror: ((event: ErrorEvent) => void) | null = null;10 onmessageerror: (() => void) | null = null;11 readonly postMessage = vi.fn<(message: unknown) => void>();12 readonly terminate = vi.fn<() => void>();13 14 constructor(..._args: unknown[]) {15 FakeWorker.instances.push(this);16 }17 18 crash(message: string): void {19 this.onerror?.({ message } as ErrorEvent);20 }21 22 send(message: EngineWorkerMessage): void {23 this.onmessage?.({ data: message } as MessageEvent<EngineWorkerMessage>);24 }25}26 27function latestWorker(): FakeWorker {28 const worker = FakeWorker.instances.at(-1);29 if (!worker) throw new Error('Expected a fake worker instance.');30 return worker;31}32 33describe('BrowserEngineClient worker lifecycle', () => {34 beforeEach(() => {35 FakeWorker.instances = [];36 vi.stubGlobal('Worker', FakeWorker);37 });38 39 afterEach(() => {40 vi.unstubAllGlobals();41 });42 43 it('terminates and recreates the worker when model loading is aborted', async () => {44 const client = new BrowserEngineClient();45 const firstWorker = latestWorker();46 const controller = new AbortController();47 const loading = client.loadModel({48 manifestUrl: '/manifest/models.json',49 modelId: '1_7b',50 backend: 'webgpu',51 }, {52 requestId: 'load-request',53 signal: controller.signal,54 });55 56 firstWorker.send({57 type: 'event',58 requestId: 'load-request',59 event: 'progress',60 phase: 'load',61 loadedBytes: 0,62 totalBytes: 100,63 shardIndex: null,64 shardCount: 1,65 shardPath: null,66 stageProgress: 0,67 nativeStage: null,68 residentBytes: null,69 shards: [],70 });71 72 controller.abort();73 74 await expect(loading).rejects.toMatchObject({ name: 'AbortError' });75 expect(firstWorker.terminate).toHaveBeenCalledOnce();76 expect(FakeWorker.instances).toHaveLength(2);77 78 const replacement = latestWorker();79 firstWorker.crash('late event from terminated worker');80 expect(replacement.terminate).not.toHaveBeenCalled();81 const estimate = client.storageEstimate();82 const request = replacement.postMessage.mock.calls[0]?.[0] as { requestId?: string };83 expect(request.requestId).toBeTruthy();84 replacement.send({85 type: 'response',86 requestId: request.requestId!,87 method: 'storageEstimate',88 ok: true,89 result: { usageBytes: 100, quotaBytes: 1_000, persisted: false },90 });91 92 await expect(estimate).resolves.toEqual({93 usageBytes: 100,94 quotaBytes: 1_000,95 persisted: false,96 });97 client.close();98 });99 100 it('also hard-restarts when the observable warmup is aborted', async () => {101 const client = new BrowserEngineClient();102 const worker = latestWorker();103 const controller = new AbortController();104 const loading = client.loadModel({105 manifestUrl: '/manifest/models.json',106 modelId: '1_7b',107 backend: 'webgpu',108 }, {109 requestId: 'warmup-request',110 signal: controller.signal,111 });112 113 worker.send({114 type: 'event',115 requestId: 'warmup-request',116 event: 'progress',117 phase: 'warmup',118 loadedBytes: 100,119 totalBytes: 100,120 shardIndex: null,121 shardCount: 1,122 shardPath: null,123 stageProgress: 0.5,124 nativeStage: 'webgpu_prefill_kernels',125 residentBytes: 100,126 shards: [],127 });128 controller.abort();129 130 await expect(loading).rejects.toMatchObject({ name: 'AbortError' });131 expect(worker.terminate).toHaveBeenCalledOnce();132 expect(FakeWorker.instances).toHaveLength(2);133 client.close();134 });135 136 it('hard-restarts only after soft download-abort cleanup times out', async () => {137 vi.useFakeTimers();138 try {139 const client = new BrowserEngineClient();140 const worker = latestWorker();141 const controller = new AbortController();142 const loading = client.loadModel({143 manifestUrl: '/manifest/models.json',144 modelId: '1_7b',145 backend: 'webgpu',146 }, {147 requestId: 'stuck-download',148 signal: controller.signal,149 });150 const loadingFailure = expect(loading).rejects.toMatchObject({ name: 'AbortError' });151 worker.send({152 type: 'event',153 requestId: 'stuck-download',154 event: 'progress',155 phase: 'download',156 loadedBytes: 1,157 totalBytes: 100,158 shardIndex: 0,159 shardCount: 1,160 shardPath: 'model.gguf',161 stageProgress: null,162 nativeStage: null,163 residentBytes: null,164 shards: [{165 index: 0,166 path: 'model.gguf',167 loadedBytes: 1,168 totalBytes: 100,169 state: 'downloading',170 }],171 });172 173 controller.abort();174 await vi.advanceTimersByTimeAsync(29_999);175 expect(worker.terminate).not.toHaveBeenCalled();176 await vi.advanceTimersByTimeAsync(1);177 178 await loadingFailure;179 expect(worker.terminate).toHaveBeenCalledOnce();180 expect(FakeWorker.instances).toHaveLength(2);181 client.close();182 } finally {183 vi.useRealTimers();184 }185 });186 187 it('lets the runtime clean partial shards when download is aborted', async () => {188 const client = new BrowserEngineClient();189 const worker = latestWorker();190 const controller = new AbortController();191 const loading = client.loadModel({192 manifestUrl: '/manifest/models.json',193 modelId: '1_7b',194 backend: 'webgpu',195 }, {196 requestId: 'download-request',197 signal: controller.signal,198 });199 worker.send({200 type: 'event',201 requestId: 'download-request',202 event: 'progress',203 phase: 'download',204 loadedBytes: 50,205 totalBytes: 100,206 shardIndex: 0,207 shardCount: 1,208 shardPath: 'model.gguf',209 stageProgress: null,210 nativeStage: null,211 residentBytes: null,212 shards: [{213 index: 0,214 path: 'model.gguf',215 loadedBytes: 50,216 totalBytes: 100,217 state: 'downloading',218 }],219 });220 221 controller.abort();222 223 expect(worker.terminate).not.toHaveBeenCalled();224 expect(FakeWorker.instances).toHaveLength(1);225 const abortRequest = worker.postMessage.mock.calls226 .map((call) => call[0] as { requestId: string; method: string })227 .find((request) => request.method === 'abort');228 expect(abortRequest).toBeTruthy();229 worker.send({230 type: 'response',231 requestId: 'download-request',232 method: 'loadModel',233 ok: false,234 error: {235 code: 'SHARD_DOWNLOAD_ABORTED',236 message: 'Partial shard removed.',237 details: { retryFromByteZero: true },238 },239 });240 if (!abortRequest) throw new Error('Expected abort request.');241 worker.send({242 type: 'response',243 requestId: abortRequest.requestId,244 method: 'abort',245 ok: true,246 result: { targetRequestId: 'download-request', aborted: true },247 });248 249 await expect(loading).rejects.toMatchObject({ code: 'SHARD_DOWNLOAD_ABORTED' });250 expect(worker.terminate).not.toHaveBeenCalled();251 client.close();252 });253 254 it('keeps the loaded model when generation responds to a soft abort', async () => {255 const client = new BrowserEngineClient();256 const worker = latestWorker();257 const controller = new AbortController();258 const generation = client.generate({259 messages: [{ role: 'user', content: 'hello' }],260 }, {261 requestId: 'soft-generation',262 signal: controller.signal,263 });264 const generationFailure = expect(generation).rejects.toMatchObject({ code: 'ABORTED' });265 266 controller.abort();267 268 const abortRequest = worker.postMessage.mock.calls269 .map((call) => call[0] as { requestId: string; method: string })270 .find((request) => request.method === 'abort');271 expect(abortRequest).toBeTruthy();272 worker.send({273 type: 'response',274 requestId: 'soft-generation',275 method: 'generate',276 ok: false,277 error: { code: 'ABORTED', message: 'Generation stopped.' },278 });279 if (!abortRequest) throw new Error('Expected abort request.');280 worker.send({281 type: 'response',282 requestId: abortRequest.requestId,283 method: 'abort',284 ok: true,285 result: { targetRequestId: 'soft-generation', aborted: true },286 });287 288 await generationFailure;289 expect(worker.terminate).not.toHaveBeenCalled();290 expect(FakeWorker.instances).toHaveLength(1);291 client.close();292 });293 294 it('routes and softly aborts diagnostic sequence scoring', async () => {295 const client = new BrowserEngineClient();296 const worker = latestWorker();297 const controller = new AbortController();298 const scoring = client.scoreSequence({299 promptTokenIds: Array.from({ length: 38 }, (_, index) => index + 1),300 referenceTokenIds: Array.from({ length: 1_024 }, (_, index) => index + 1_000),301 topK: 5,302 }, {303 requestId: 'soft-score-sequence',304 signal: controller.signal,305 });306 const scoringFailure = expect(scoring).rejects.toMatchObject({ code: 'ABORTED' });307 308 expect(worker.postMessage).toHaveBeenCalledWith(expect.objectContaining({309 requestId: 'soft-score-sequence',310 method: 'scoreSequence',311 }));312 controller.abort();313 const abortRequest = worker.postMessage.mock.calls314 .map((call) => call[0] as { requestId: string; method: string })315 .find((request) => request.method === 'abort');316 expect(abortRequest).toBeTruthy();317 worker.send({318 type: 'response',319 requestId: 'soft-score-sequence',320 method: 'scoreSequence',321 ok: false,322 error: { code: 'ABORTED', message: 'Sequence scoring stopped.' },323 });324 if (!abortRequest) throw new Error('Expected abort request.');325 worker.send({326 type: 'response',327 requestId: abortRequest.requestId,328 method: 'abort',329 ok: true,330 result: { targetRequestId: 'soft-score-sequence', aborted: true },331 });332 333 await scoringFailure;334 expect(worker.terminate).not.toHaveBeenCalled();335 client.close();336 });337 338 it('keeps the worker and loaded model while native generation reaches an abort checkpoint', async () => {339 vi.useFakeTimers();340 try {341 const client = new BrowserEngineClient();342 const worker = latestWorker();343 const controller = new AbortController();344 const generation = client.generate({345 messages: [{ role: 'user', content: 'hello' }],346 }, {347 requestId: 'hung-generation',348 signal: controller.signal,349 });350 const generationFailure = expect(generation).rejects.toMatchObject({ code: 'ABORTED' });351 352 controller.abort();353 await vi.advanceTimersByTimeAsync(60_000);354 expect(worker.terminate).not.toHaveBeenCalled();355 expect(FakeWorker.instances).toHaveLength(1);356 357 const abortRequest = worker.postMessage.mock.calls358 .map((call) => call[0] as { requestId: string; method: string })359 .find((request) => request.method === 'abort');360 expect(abortRequest).toBeTruthy();361 worker.send({362 type: 'response',363 requestId: 'hung-generation',364 method: 'generate',365 ok: false,366 error: { code: 'ABORTED', message: 'Generation stopped at the next native checkpoint.' },367 });368 if (!abortRequest) throw new Error('Expected abort request.');369 worker.send({370 type: 'response',371 requestId: abortRequest.requestId,372 method: 'abort',373 ok: true,374 result: { targetRequestId: 'hung-generation', aborted: true },375 });376 await generationFailure;377 expect(worker.terminate).not.toHaveBeenCalled();378 expect(FakeWorker.instances).toHaveLength(1);379 client.close();380 } finally {381 vi.useRealTimers();382 }383 });384 385 it('rejects pending work on a crash and recreates the worker for the next request', async () => {386 const client = new BrowserEngineClient();387 const firstWorker = latestWorker();388 const pending = client.storageEstimate();389 390 firstWorker.crash('synthetic worker crash');391 392 await expect(pending).rejects.toMatchObject({393 code: 'ENGINE_WORKER_FAILED',394 message: expect.stringContaining('synthetic worker crash'),395 details: { recoverable: true, nextAction: 'reload-model' },396 });397 expect(firstWorker.terminate).toHaveBeenCalledOnce();398 expect(FakeWorker.instances).toHaveLength(1);399 400 const recovered = client.storageEstimate();401 expect(FakeWorker.instances).toHaveLength(2);402 const replacement = latestWorker();403 const request = replacement.postMessage.mock.calls[0]?.[0] as { requestId?: string };404 replacement.send({405 type: 'response',406 requestId: request.requestId!,407 method: 'storageEstimate',408 ok: true,409 result: { usageBytes: null, quotaBytes: null, persisted: false },410 });411 412 await expect(recovered).resolves.toEqual({413 usageBytes: null,414 quotaBytes: null,415 persisted: false,416 });417 client.close();418 });419 420 it('surfaces WebGPU device loss and replaces the invalid worker before the next action', async () => {421 const client = new BrowserEngineClient();422 const firstWorker = latestWorker();423 const generation = client.generate({424 messages: [{ role: 'user', content: 'hello' }],425 }, { requestId: 'generation' });426 427 firstWorker.send({428 type: 'response',429 requestId: 'generation',430 method: 'generate',431 ok: false,432 error: {433 code: 'WEBGPU_DEVICE_LOST',434 message: 'The WebGPU device was lost during generate.',435 details: { recoverable: true, nextAction: 'reload-model' },436 },437 });438 439 await expect(generation).rejects.toMatchObject({440 name: 'EngineClientError',441 code: 'WEBGPU_DEVICE_LOST',442 details: { recoverable: true, nextAction: 'reload-model' },443 });444 expect(firstWorker.terminate).toHaveBeenCalledOnce();445 expect(FakeWorker.instances).toHaveLength(2);446 447 const replacement = latestWorker();448 const reload = client.loadModel({449 manifestUrl: '/manifest/models.json',450 modelId: '1_7b',451 backend: 'webgpu',452 }, { requestId: 'reload' });453 expect(replacement.postMessage).toHaveBeenCalledWith(expect.objectContaining({454 requestId: 'reload',455 method: 'loadModel',456 }));457 replacement.send({458 type: 'response',459 requestId: 'reload',460 method: 'loadModel',461 ok: false,462 error: { code: 'MODEL_GATE_REJECTED', message: 'synthetic stop' },463 });464 await expect(reload).rejects.toMatchObject({ code: 'MODEL_GATE_REJECTED' });465 client.close();466 });467 468 it('types a crash during generation as model-invalidating before recreating the worker', async () => {469 const client = new BrowserEngineClient();470 const firstWorker = latestWorker();471 const generation = client.generate({472 messages: [{ role: 'user', content: 'hello' }],473 });474 475 firstWorker.crash('GPU worker process exited');476 477 await expect(generation).rejects.toMatchObject({478 code: 'ENGINE_WORKER_FAILED',479 details: { recoverable: true, nextAction: 'reload-model' },480 });481 expect(firstWorker.terminate).toHaveBeenCalledOnce();482 483 const reload = client.loadModel({484 manifestUrl: '/manifest/models.json',485 modelId: '1_7b',486 backend: 'webgpu',487 }, { requestId: 'crash-reload' });488 const replacement = latestWorker();489 expect(replacement.postMessage).toHaveBeenCalledWith(expect.objectContaining({490 requestId: 'crash-reload',491 method: 'loadModel',492 }));493 replacement.send({494 type: 'response',495 requestId: 'crash-reload',496 method: 'loadModel',497 ok: false,498 error: { code: 'MODEL_GATE_REJECTED', message: 'synthetic stop' },499 });500 await expect(reload).rejects.toMatchObject({ code: 'MODEL_GATE_REJECTED' });501 client.close();502 });503 504 it('routes live tool-call progress independently from visible text tokens', async () => {505 const client = new BrowserEngineClient();506 const worker = latestWorker();507 const onToken = vi.fn();508 const onToolCallProgress = vi.fn();509 const generation = client.generate({510 messages: [{ role: 'user', content: 'Create an app.' }],511 }, {512 requestId: 'tool-progress',513 onToken,514 onToolCallProgress,515 });516 517 worker.send({518 type: 'event',519 requestId: 'tool-progress',520 event: 'tool-call',521 index: 0,522 id: 'call_artifact',523 name: 'html_artifact',524 argumentCharacters: 4_096,525 argumentDelta: '<main>',526 });527 528 expect(onToolCallProgress).toHaveBeenCalledWith(expect.objectContaining({529 name: 'html_artifact',530 argumentCharacters: 4_096,531 argumentDelta: '<main>',532 }));533 expect(onToken).not.toHaveBeenCalled();534 535 worker.send({536 type: 'response',537 requestId: 'tool-progress',538 method: 'generate',539 ok: true,540 result: {541 text: '',542 reasoningText: '',543 tokenIds: null,544 tokenTrace: null,545 finishReason: 'tool_calls',546 toolCalls: [{547 id: 'call_artifact',548 type: 'function',549 function: { name: 'html_artifact', arguments: '{"html":"<main />"}' },550 }],551 usage: null,552 timings: null,553 },554 });555 556 await expect(generation).resolves.toMatchObject({ finishReason: 'tool_calls' });557 client.close();558 });559 560 it('never recreates a worker after close', async () => {561 const client = new BrowserEngineClient();562 const worker = latestWorker();563 const pending = client.storageEstimate();564 565 client.close();566 client.restart();567 568 await expect(pending).rejects.toThrow('closed');569 await expect(client.storageEstimate()).rejects.toThrow('closed');570 expect(worker.terminate).toHaveBeenCalledOnce();571 expect(FakeWorker.instances).toHaveLength(1);572 });573});574 