From 9a7d2a21809ff06913f9e45be74178e85db1bfed Mon Sep 17 00:00:00 2001 From: Neil <4138956+nwparker@users.noreply.github.com> Date: Fri, 29 May 2026 01:57:20 -0700 Subject: [PATCH] fix speech model download cancellation (#3043) --- src/main/speech/model-manager.test.ts | 60 ++++++++++++++++++++++++++- src/main/speech/model-manager.ts | 26 ++++++++++-- 2 files changed, 80 insertions(+), 6 deletions(-) diff --git a/src/main/speech/model-manager.test.ts b/src/main/speech/model-manager.test.ts index 5e43cc528..d799b768f 100644 --- a/src/main/speech/model-manager.test.ts +++ b/src/main/speech/model-manager.test.ts @@ -6,7 +6,8 @@ import { beforeEach, describe, expect, it, vi } from 'vitest' import { SPEECH_MODEL_CATALOG } from './model-catalog' import { ModelManager } from './model-manager' -const { spawnMock } = vi.hoisted(() => ({ +const { httpsGetMock, spawnMock } = vi.hoisted(() => ({ + httpsGetMock: vi.fn(), spawnMock: vi.fn() })) @@ -21,6 +22,11 @@ vi.mock('child_process', async () => { return { ...(actual as Record), spawn: spawnMock } }) +vi.mock('https', async () => { + const actual = await vi.importActual('https') + return { ...(actual as Record), get: httpsGetMock } +}) + type ModelManagerInternals = { verifyArchiveSha256: (archivePath: string, expectedSha256: string) => Promise downloadFile: ( @@ -28,7 +34,8 @@ type ModelManagerInternals = { dest: string, expectedSize: number, modelId: string, - isAborted: () => boolean + isAborted: () => boolean, + signal?: AbortSignal ) => Promise extractArchive: ( archivePath: string, @@ -40,6 +47,7 @@ type ModelManagerInternals = { describe('ModelManager', () => { beforeEach(() => { + httpsGetMock.mockReset() spawnMock.mockReset() }) @@ -85,6 +93,54 @@ describe('ModelManager', () => { } }) + it('aborts an in-flight model download request when cancelled', async () => { + const dir = mkdtempSync(join(tmpdir(), 'orca-model-manager-')) + try { + const manifest = SPEECH_MODEL_CATALOG[0] + const errorHandlers: ((err: Error) => void)[] = [] + const request = { + destroy: vi.fn((err?: Error) => { + queueMicrotask(() => { + for (const handler of errorHandlers) { + handler(err ?? new Error('destroyed')) + } + }) + return request + }), + on: vi.fn((event: string, cb: (err: Error) => void) => { + if (event === 'error') { + errorHandlers.push(cb) + } + return request + }) + } + httpsGetMock.mockImplementation( + ( + _url: URL, + options: { signal?: AbortSignal } | ((response: unknown) => void), + _cb?: (response: unknown) => void + ) => { + if (typeof options !== 'function') { + options.signal?.addEventListener('abort', () => request.destroy(new Error('Aborted')), { + once: true + }) + } + return request + } + ) + const manager = new ModelManager(dir) + + const download = manager.downloadModel(manifest.id) + manager.cancelDownload(manifest.id) + await expect(download).resolves.toBeUndefined() + + expect(request.destroy).toHaveBeenCalledWith(expect.any(Error)) + expect((await manager.getModelState(manifest.id)).status).toBe('not-downloaded') + } finally { + rmSync(dir, { recursive: true, force: true }) + } + }) + it('clears extraction abort polling when the child does not close', async () => { vi.useFakeTimers() const dir = mkdtempSync(join(tmpdir(), 'orca-model-manager-')) diff --git a/src/main/speech/model-manager.ts b/src/main/speech/model-manager.ts index e31fad215..5ddd739d7 100644 --- a/src/main/speech/model-manager.ts +++ b/src/main/speech/model-manager.ts @@ -113,10 +113,14 @@ export class ModelManager { const archivePath = join(this.modelsDir, `${modelId}.tar.bz2`) let aborted = false + const abortController = new AbortController() const handle: DownloadHandle = { abort: () => { aborted = true + // Why: a stalled HTTPS request may never deliver another data chunk; + // cancellation must tear down the request immediately. + abortController.abort() } } this.activeDownloads.set(modelId, handle) @@ -127,7 +131,8 @@ export class ModelManager { archivePath, manifest.sizeBytes, modelId, - () => aborted + () => aborted, + abortController.signal ) if (aborted) { @@ -223,9 +228,15 @@ export class ModelManager { expectedSize: number, modelId: string, isAborted: () => boolean, + signal?: AbortSignal, redirectCount = 0 ): Promise { return new Promise((resolve, reject) => { + if (signal?.aborted) { + reject(new Error('Aborted')) + return + } + let parsedUrl: URL try { parsedUrl = new URL(url) @@ -239,7 +250,8 @@ export class ModelManager { return } - const request = httpsGet(parsedUrl, (response: IncomingMessage) => { + let request: ReturnType + const onResponse = (response: IncomingMessage): void => { if ( response.statusCode === 301 || response.statusCode === 302 || @@ -278,6 +290,7 @@ export class ModelManager { expectedSize, modelId, isAborted, + signal, redirectCount + 1 ) .then(resolve) @@ -298,8 +311,9 @@ export class ModelManager { response.on('data', (chunk: Buffer) => { if (isAborted()) { + request.destroy(new Error('Aborted')) response.destroy() - fileStream.close() + fileStream.destroy() return } downloaded += chunk.length @@ -316,7 +330,11 @@ export class ModelManager { } }) .catch(reject) - }) + } + + request = signal + ? httpsGet(parsedUrl, { signal }, onResponse) + : httpsGet(parsedUrl, onResponse) request.on('error', reject) })