WaveCut/Bonsai-Chat-WebGPU
5
1import { sha256 as hashSha256 } from '@noble/hashes/sha2.js';2import { bytesToHex } from '@noble/hashes/utils.js';3import { describe, expect, it } from 'vitest';4import type { Wllama } from '../../vendor/wllama-bonsai/esm/index.js';5import { EngineRuntimeError } from './errors';6import type { ManifestModelV2 } from './manifest';7import { isShardDownloadFailureDetails, type EngineEvent } from './protocol';8import {9 BrowserEngineRuntime,10 resolveLoadTuning,11 resolveRuntimeBatchShape,12 selectWasmFlavor,13} from './runtime';14 15type DownloadOptions = {16 signal?: AbortSignal;17 headers?: Record<string, string>;18 progressCallback?: (progress: { loaded: number; total: number }) => void;19};20 21class FakeCacheManager {22 readonly blobs = new Map<string, Blob>();23 readonly metadata = new Map<string, Record<string, unknown>>();24 readonly downloads: Array<{ url: string; headers?: Record<string, string> }> = [];25 readonly deletes: string[] = [];26 readonly payloads = new Map<string, Blob>();27 failNetworkFor: string | null = null;28 readonly waitForAbortFor = new Set<string>();29 private readonly downloadWaiters: Array<{ count: number; resolve: () => void }> = [];30 31 waitForDownloads(count: number): Promise<void> {32 if (this.downloads.length >= count) return Promise.resolve();33 return new Promise<void>((resolve) => this.downloadWaiters.push({ count, resolve }));34 }35 36 private resolveDownloadWaiters(): void {37 for (let index = this.downloadWaiters.length - 1; index >= 0; index -= 1) {38 const waiter = this.downloadWaiters[index];39 if (waiter && this.downloads.length >= waiter.count) {40 this.downloadWaiters.splice(index, 1);41 waiter.resolve();42 }43 }44 }45 46 async download(url: string, options: DownloadOptions = {}): Promise<void> {47 this.downloads.push({ url, ...(options.headers ? { headers: options.headers } : {}) });48 this.resolveDownloadWaiters();49 const payload = this.payloads.get(url) ?? new Blob(['partial']);50 this.blobs.set(url, payload);51 options.progressCallback?.({ loaded: payload.size, total: payload.size });52 if (this.failNetworkFor === url) {53 throw new TypeError('synthetic network failure');54 }55 if (this.waitForAbortFor.has(url)) {56 await new Promise<never>((_resolve, reject) => {57 const rejectAborted = () => reject(new DOMException('synthetic abort', 'AbortError'));58 if (options.signal?.aborted) rejectAborted();59 else options.signal?.addEventListener('abort', rejectAborted, { once: true });60 });61 }62 }63 64 async open(url: string): Promise<Blob | null> {65 return this.blobs.get(url) ?? null;66 }67 68 async delete(url: string): Promise<void> {69 this.deletes.push(url);70 this.blobs.delete(url);71 }72 73 async getNameFromURL(url: string): Promise<string> {74 return `cache:${url}`;75 }76 77 async getMetadata(key: string): Promise<Record<string, unknown> | null> {78 return this.metadata.get(key) ?? null;79 }80 81 async writeMetadata(key: string, metadata: Record<string, unknown>): Promise<void> {82 this.metadata.set(key, metadata);83 }84}85 86type InspectCachedShards = (87 requestId: string,88 wllama: Wllama,89 model: ManifestModelV2,90 urls: readonly string[],91 signal: AbortSignal,92 sink: (event: EngineEvent) => void,93) => Promise<{ blobs: Array<Blob | null>; cachedBytes: number }>;94 95type DownloadShards = (96 requestId: string,97 wllama: Wllama,98 model: ManifestModelV2,99 urls: readonly string[],100 cached: { blobs: Array<Blob | null>; cachedBytes: number },101 signal: AbortSignal,102 sink: (event: EngineEvent) => void,103) => Promise<Blob[]>;104 105async function sha256(blob: Blob): Promise<string> {106 return bytesToHex(hashSha256(new Uint8Array(await blob.arrayBuffer())));107}108 109async function fixture(first: Blob, second: Blob): Promise<{110 model: ManifestModelV2;111 urls: [string, string];112}> {113 const result = await fixtureShards([first, second]);114 const firstUrl = result.urls[0];115 const secondUrl = result.urls[1];116 if (!firstUrl || !secondUrl) throw new Error('Two-shard fixture is incomplete.');117 return { model: result.model, urls: [firstUrl, secondUrl] };118}119 120async function fixtureShards(payloads: readonly Blob[]): Promise<{121 model: ManifestModelV2;122 urls: string[];123}> {124 const urls = payloads.map((_, index) => `https://models.test/shard-${index}.gguf`);125 const files = await Promise.all(payloads.map(async (payload, index) => ({126 path: `shard-${index}.gguf`,127 bytes: payload.size,128 sha256: await sha256(payload),129 })));130 return {131 urls,132 model: {133 id: '1_7b',134 displayName: 'Test Bonsai',135 architecture: 'test',136 source: { repo: 'test/model', revision: 'a'.repeat(40), file: 'model.gguf', bytes: 1, sha256: '0'.repeat(64) },137 files,138 downloadBytes: payloads.reduce((sum, payload) => sum + payload.size, 0),139 contextLength: 2048,140 defaultContext: 1024,141 cpuFallback: true,142 largestTensorBytes: 1,143 requiredLimits: {},144 chatTemplate: { bytes: 1, sha256: '0'.repeat(64), markers: { think: false, toolCall: false, toolResponse: false } },145 hybridDimensions: {},146 nextNTensorCount: 0,147 runtimePolicy: { flashAttention: false, tokenEmbeddingOnWebGPU: true, requireSingleWebGPUGraph: true },148 },149 };150}151 152function downloader(runtime: BrowserEngineRuntime): DownloadShards {153 return (runtime as unknown as { downloadShards: DownloadShards }).downloadShards.bind(runtime);154}155 156function cacheInspector(runtime: BrowserEngineRuntime): InspectCachedShards {157 return (runtime as unknown as {158 inspectCachedShards: InspectCachedShards;159 }).inspectCachedShards.bind(runtime);160}161 162describe('BrowserEngineRuntime per-shard retry', () => {163 it('accepts a size-matched pinned cache entry without rescanning the Blob', async () => {164 const expected = new Blob(['cached-model']);165 const cachedBlob = new Blob(['cached-model']);166 Object.defineProperty(cachedBlob, 'stream', {167 value: () => {168 throw new Error('The runtime attempted to rescan the cached Blob.');169 },170 });171 const { model, urls } = await fixtureShards([expected]);172 const url = urls[0];173 const file = model.files[0];174 if (!url || !file) throw new Error('One-shard fixture is incomplete.');175 const cache = new FakeCacheManager();176 cache.blobs.set(url, cachedBlob);177 cache.metadata.set(`cache:${url}`, { sha256: file.sha256 });178 179 const cached = await cacheInspector(new BrowserEngineRuntime())(180 'request-cached-no-scan',181 { cacheManager: cache } as unknown as Wllama,182 model,183 urls,184 new AbortController().signal,185 () => undefined,186 );187 188 expect(cached.blobs).toEqual([cachedBlob]);189 expect(cached.cachedBytes).toBe(cachedBlob.size);190 });191 192 it('hands a correctly sized download to the native loader without a second full-file scan', async () => {193 const expected = new Blob(['model-bytes']);194 const downloaded = new Blob(['model-bytes']);195 Object.defineProperty(downloaded, 'stream', {196 value: () => {197 throw new Error('The runtime attempted a duplicate full-file scan.');198 },199 });200 const { model, urls } = await fixtureShards([expected]);201 const url = urls[0];202 if (!url) throw new Error('One-shard fixture is incomplete.');203 const cache = new FakeCacheManager();204 cache.payloads.set(url, downloaded);205 206 const blobs = await downloader(new BrowserEngineRuntime())(207 'request-no-second-scan',208 { cacheManager: cache } as unknown as Wllama,209 model,210 urls,211 { blobs: [null], cachedBytes: 0 },212 new AbortController().signal,213 () => undefined,214 );215 216 expect(blobs).toEqual([downloaded]);217 });218 219 it('deletes a network-failed partial and retries only that shard with a full GET', async () => {220 const verified = new Blob(['verified']);221 const retryPayload = new Blob(['retry-ok']);222 const { model, urls } = await fixture(verified, retryPayload);223 const cache = new FakeCacheManager();224 cache.payloads.set(urls[1], retryPayload);225 cache.failNetworkFor = urls[1];226 const cached = { blobs: [verified, null], cachedBytes: verified.size };227 const download = downloader(new BrowserEngineRuntime());228 229 const failed = download(230 'request-1',231 { cacheManager: cache } as unknown as Wllama,232 model,233 urls,234 cached,235 new AbortController().signal,236 () => undefined,237 );238 239 await expect(failed).rejects.toSatisfy((error: unknown) => {240 if (!(error instanceof EngineRuntimeError) || !isShardDownloadFailureDetails(error.details)) return false;241 expect(error.code).toBe('SHARD_DOWNLOAD_FAILED');242 expect(error.details).toMatchObject({243 failure: 'network',244 shardIndex: 1,245 shardCount: 2,246 shardPath: 'shard-1.gguf',247 retryFromByteZero: true,248 partialDeleted: true,249 });250 return true;251 });252 expect(cache.deletes).toEqual([urls[1]]);253 expect(cache.blobs.has(urls[1])).toBe(false);254 expect(cached.blobs).toEqual([verified, null]);255 256 cache.failNetworkFor = null;257 const blobs = await download(258 'request-2',259 { cacheManager: cache } as unknown as Wllama,260 model,261 urls,262 cached,263 new AbortController().signal,264 () => undefined,265 );266 267 expect(blobs).toEqual([verified, retryPayload]);268 expect(cache.downloads.map(({ url }) => url)).toEqual([urls[1], urls[1]]);269 expect(cache.downloads.every(({ headers }) => headers?.Range === undefined)).toBe(true);270 });271 272 it('deletes a shard with an unexpected byte length before native loading', async () => {273 const expected = new Blob(['good']);274 const corrupted = new Blob(['corrupted']);275 const { model, urls } = await fixture(new Blob(['cached']), expected);276 const cache = new FakeCacheManager();277 cache.payloads.set(urls[1], corrupted);278 const cached = { blobs: [new Blob(['cached']), null], cachedBytes: 6 };279 280 await expect(downloader(new BrowserEngineRuntime())(281 'request-size',282 { cacheManager: cache } as unknown as Wllama,283 model,284 urls,285 cached,286 new AbortController().signal,287 () => undefined,288 )).rejects.toSatisfy((error: unknown) => (289 error instanceof EngineRuntimeError290 && error.code === 'SHARD_SIZE_MISMATCH'291 && isShardDownloadFailureDetails(error.details)292 && error.details.failure === 'verification'293 && error.details.retryFromByteZero294 ));295 expect(cache.deletes).toEqual([urls[1]]);296 expect(cache.blobs.has(urls[1])).toBe(false);297 });298 299 it('stops dequeue after a network failure, lets in-flight siblings finish, and retries only 0 + 3', async () => {300 const payloads = [301 new Blob(['retry-zero']),302 new Blob(['sibling-one']),303 new Blob(['sibling-two']),304 new Blob(['queued-three']),305 ];306 const { model, urls } = await fixtureShards(payloads);307 const cache = new FakeCacheManager();308 urls.forEach((url, index) => cache.payloads.set(url, payloads[index] as Blob));309 cache.failNetworkFor = urls[0] ?? null;310 const cached = { blobs: payloads.map(() => null), cachedBytes: 0 };311 const progress: number[] = [];312 const download = downloader(new BrowserEngineRuntime());313 314 await expect(download(315 'request-parallel-failure',316 { cacheManager: cache } as unknown as Wllama,317 model,318 urls,319 cached,320 new AbortController().signal,321 (event) => {322 if (event.event === 'progress') progress.push(event.loadedBytes);323 },324 )).rejects.toMatchObject({ code: 'SHARD_DOWNLOAD_FAILED' });325 326 expect(cache.downloads.map(({ url }) => url)).toEqual(urls.slice(0, 3));327 expect(cache.downloads.some(({ url }) => url === urls[3])).toBe(false);328 expect(cache.deletes).toEqual([urls[0]]);329 expect(cached.blobs).toEqual([null, payloads[1], payloads[2], null]);330 expect(progress[progress.length - 1]).toBe(payloads[1]!.size + payloads[2]!.size);331 332 cache.failNetworkFor = null;333 const retryStart = cache.downloads.length;334 const blobs = await download(335 'request-parallel-retry',336 { cacheManager: cache } as unknown as Wllama,337 model,338 urls,339 cached,340 new AbortController().signal,341 () => undefined,342 );343 344 expect(blobs).toEqual(payloads);345 expect(cache.downloads.slice(retryStart).map(({ url }) => url)).toEqual([urls[0], urls[3]]);346 expect(cache.downloads.slice(retryStart).every(({ headers }) => headers?.Range === undefined)).toBe(true);347 });348 349 it('aborts all in-flight shards, cleans each partial, and never dequeues the fourth shard', async () => {350 const expected = [351 new Blob(['complete-zero']),352 new Blob(['complete-one']),353 new Blob(['complete-two']),354 new Blob(['complete-three']),355 ];356 const partials = [new Blob(['p0']), new Blob(['p1']), new Blob(['p2'])];357 const { model, urls } = await fixtureShards(expected);358 const cache = new FakeCacheManager();359 urls.slice(0, 3).forEach((url, index) => {360 cache.payloads.set(url, partials[index] as Blob);361 cache.waitForAbortFor.add(url);362 });363 const cached = { blobs: expected.map(() => null), cachedBytes: 0 };364 const progress: number[] = [];365 const controller = new AbortController();366 const failure = downloader(new BrowserEngineRuntime())(367 'request-abort',368 { cacheManager: cache } as unknown as Wllama,369 model,370 urls,371 cached,372 controller.signal,373 (event) => {374 if (event.event === 'progress') progress.push(event.loadedBytes);375 },376 );377 378 await cache.waitForDownloads(3);379 controller.abort();380 381 await expect(failure).rejects.toSatisfy((error: unknown) => (382 error instanceof EngineRuntimeError383 && error.code === 'SHARD_DOWNLOAD_ABORTED'384 && isShardDownloadFailureDetails(error.details)385 && error.details.failure === 'abort'386 && error.details.shardIndex < 3387 && error.details.shardCount === 4388 && error.details.shardPath === `shard-${error.details.shardIndex}.gguf`389 && error.details.partialDeleted390 ));391 expect(cache.downloads.map(({ url }) => url)).toEqual(urls.slice(0, 3));392 expect(cache.downloads.some(({ url }) => url === urls[3])).toBe(false);393 expect([...cache.deletes].sort()).toEqual([...urls.slice(0, 3)].sort());394 expect(urls.slice(0, 3).every((url) => !cache.blobs.has(url))).toBe(true);395 expect(cached.blobs).toEqual([null, null, null, null]);396 expect(progress[progress.length - 1]).toBe(0);397 });398});399 400describe('resolveLoadTuning', () => {401 const benchmarkDefaults = {402 flashMode: 'off',403 kvCacheType: 'f16',404 wasmFlavor: 'auto',405 } as const;406 407 function tuningError(408 input: unknown,409 backend: 'auto' | 'webgpu' | 'wasm' = 'webgpu',410 ): EngineRuntimeError {411 try {412 resolveLoadTuning(input, backend);413 } catch (error) {414 if (error instanceof EngineRuntimeError) return error;415 throw error;416 }417 throw new Error('Expected benchmark tuning to be rejected.');418 }419 420 it('uses immutable release defaults when benchmark tuning is omitted', () => {421 expect(resolveLoadTuning(undefined, 'auto')).toEqual({422 scope: 'release-defaults',423 nBatch: null,424 nUbatch: null,425 flashMode: 'off',426 kvCacheType: 'f16',427 wasmFlavor: 'auto',428 });429 expect(resolveLoadTuning(undefined, 'wasm')).toEqual({430 scope: 'release-defaults',431 nBatch: null,432 nUbatch: null,433 flashMode: 'off',434 kvCacheType: 'f16',435 wasmFlavor: 'auto',436 });437 });438 439 it('accepts an explicit benchmark batch shape without changing release defaults', () => {440 expect(resolveLoadTuning({441 ...benchmarkDefaults,442 nBatch: 512,443 nUbatch: 128,444 }, 'auto')).toEqual({445 scope: 'benchmark',446 nBatch: 512,447 nUbatch: 128,448 flashMode: 'off',449 kvCacheType: 'f16',450 wasmFlavor: 'auto',451 });452 });453 454 it('rejects a microbatch larger than the requested batch', () => {455 expect(tuningError({456 ...benchmarkDefaults,457 nBatch: 128,458 nUbatch: 256,459 })).toMatchObject({460 code: 'INVALID_BENCHMARK_TUNING',461 message: 'Benchmark n_ubatch must not exceed n_batch.',462 });463 });464 465 it('accepts Flash Attention and quantized KV only on the explicit WebGPU backend', () => {466 expect(resolveLoadTuning({467 ...benchmarkDefaults,468 flashMode: 'auto',469 }, 'webgpu')).toMatchObject({470 scope: 'benchmark',471 flashMode: 'auto',472 kvCacheType: 'f16',473 wasmFlavor: 'auto',474 });475 for (const kvCacheType of ['q8_0', 'q4_0'] as const) {476 expect(resolveLoadTuning({477 ...benchmarkDefaults,478 flashMode: 'auto',479 kvCacheType,480 }, 'webgpu')).toMatchObject({481 scope: 'benchmark',482 flashMode: 'auto',483 kvCacheType,484 wasmFlavor: 'auto',485 });486 }487 488 expect(tuningError({489 ...benchmarkDefaults,490 flashMode: 'auto',491 }, 'auto')).toMatchObject({492 code: 'INVALID_BENCHMARK_TUNING',493 });494 expect(tuningError({495 ...benchmarkDefaults,496 flashMode: 'auto',497 kvCacheType: 'q8_0',498 }, 'wasm')).toMatchObject({499 code: 'INVALID_BENCHMARK_TUNING',500 });501 });502 503 it('rejects malformed and incomplete benchmark tuning objects', () => {504 for (const input of [null, [], 'off', 42]) {505 expect(tuningError(input)).toMatchObject({ code: 'INVALID_BENCHMARK_TUNING' });506 }507 for (const input of [508 { kvCacheType: 'f16' },509 { flashMode: 'off' },510 { flashMode: 'off', kvCacheType: 'f16' },511 { ...benchmarkDefaults, flashMode: 'enabled' },512 { ...benchmarkDefaults, kvCacheType: 'q5_0' },513 { ...benchmarkDefaults, wasmFlavor: 'future' },514 ]) {515 expect(tuningError(input)).toMatchObject({ code: 'INVALID_BENCHMARK_TUNING' });516 }517 });518 519 it('rejects quantized V cache unless Flash Attention is in auto mode', () => {520 for (const kvCacheType of ['q8_0', 'q4_0'] as const) {521 expect(tuningError({ ...benchmarkDefaults, kvCacheType })).toMatchObject({522 code: 'INVALID_BENCHMARK_TUNING',523 message: expect.stringMatching(/requires Flash Attention/i),524 });525 }526 });527 528 it('rejects invalid and overflowing signed glue batch values', () => {529 for (const nBatch of [0, -1, 1.5, Number.NaN, Number.POSITIVE_INFINITY, 2_147_483_648]) {530 expect(tuningError({ ...benchmarkDefaults, nBatch })).toMatchObject({531 code: 'INVALID_BATCH_SIZE',532 });533 }534 for (const nUbatch of [0, -1, 1.5, Number.NaN, Number.POSITIVE_INFINITY, Number.MAX_SAFE_INTEGER]) {535 expect(tuningError({ ...benchmarkDefaults, nUbatch })).toMatchObject({536 code: 'INVALID_UBATCH_SIZE',537 });538 }539 });540});541 542describe('resolveRuntimeBatchShape', () => {543 const releaseDefaults = resolveLoadTuning(undefined, 'auto');544 545 it('caps only the compat release path to short Safari-safe command buffers', () => {546 expect(resolveRuntimeBatchShape(releaseDefaults, 'compat')).toEqual({547 nBatch: 32,548 nUbatch: 16,549 });550 expect(resolveRuntimeBatchShape(releaseDefaults, 'jspi')).toEqual({551 nBatch: undefined,552 nUbatch: undefined,553 });554 });555 556 it('preserves explicit tuning and never defaults microbatch above batch', () => {557 expect(resolveRuntimeBatchShape({ nBatch: 96, nUbatch: 24 }, 'compat')).toEqual({558 nBatch: 96,559 nUbatch: 24,560 });561 expect(resolveRuntimeBatchShape({ nBatch: 8, nUbatch: null }, 'compat')).toEqual({562 nBatch: 8,563 nUbatch: 8,564 });565 });566});567 568describe('selectWasmFlavor', () => {569 it('uses the browser-native flavor for auto', () => {570 expect(selectWasmFlavor('auto', false)).toBe('jspi');571 expect(selectWasmFlavor('auto', true)).toBe('compat');572 });573 574 it('forces compat even when JSPI is available', () => {575 expect(selectWasmFlavor('compat', false)).toBe('compat');576 expect(selectWasmFlavor('compat', true)).toBe('compat');577 });578 579 it('selects JSPI when the browser supports it', () => {580 expect(selectWasmFlavor('jspi', false)).toBe('jspi');581 });582 583 it('fails loud when JSPI is requested but unavailable', () => {584 expect(() => selectWasmFlavor('jspi', true)).toThrowError(585 expect.objectContaining({586 code: 'BENCHMARK_WASM_FLAVOR_UNAVAILABLE',587 message: expect.stringMatching(/requires the compatibility runtime/i),588 }),589 );590 });591});592 