Team Ai
Apppublic

WaveCut/Bonsai-Chat-WebGPU

sourceHugging Faceupdated 3mo agoView on Hugging Face
5likes
client.test.ts574 linesDownload Raw Back to engine
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