Team Ai
Apppublic

WaveCut/Bonsai-Chat-WebGPU

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