fix(stt): kept whisper setup downloads alive

Kept the STT subprocess referenced while download and stream requests are pending so setup cannot exit before the worker answers.

Propagated worker download errors to setup callers and verified completed downloads leave the expected cache files.

Fixes #3939
This commit is contained in:
roboomp
2026-07-01 00:07:14 +00:00
parent ebdc7280cb
commit efcd6cafd6
4 changed files with 191 additions and 30 deletions
+4
View File
@@ -2,6 +2,10 @@
## [Unreleased]
### Fixed
- Fixed `omp setup speech` returning before Whisper STT downloads finish by keeping the STT worker referenced while setup awaits it, and surfaced worker download errors instead of collapsing them to a silent false result. ([#3939](https://github.com/can1357/oh-my-pi/issues/3939))
## [16.2.10] - 2026-06-30
### Changed
+89 -27
View File
@@ -4,12 +4,12 @@ import {
createWorkerHandle,
createWorkerSubprocess,
logWorkerMessage,
type RefCountedWorkerHandle,
resolveWorkerSpawnCmd,
SMOKE_TEST_TIMEOUT_MS,
type SpawnedSubprocess,
smokeTestWorker,
spawnWorkerOrUnavailable,
type WorkerHandle,
} from "../subprocess/worker-client";
import { tinyWorkerEnv } from "../tiny/title-client";
import { safeSend } from "../utils/ipc";
@@ -18,7 +18,7 @@ import type { SttModelKey } from "./models";
type PendingRequest =
| { kind: "transcribe"; modelKey: SttModelKey; resolve: (text: string) => void; reject: (error: Error) => void }
| { kind: "download"; modelKey: SttModelKey; resolve: (ok: boolean) => void };
| { kind: "download"; modelKey: SttModelKey; resolve: (result: SttDownloadResult) => void };
export interface SttTranscribeOptions {
language?: string;
@@ -30,6 +30,11 @@ export interface SttDownloadOptions {
onProgress?: (event: SttProgressEvent) => void;
}
export interface SttDownloadResult {
ok: boolean;
error?: string;
}
/** Live streaming session handle returned by {@link SttClient.startStream}. */
export interface SttStreamHandle {
/** Feed 16 kHz mono float samples as the recorder produces them. */
@@ -79,30 +84,57 @@ export function createSttSubprocess(): SpawnedSubprocess<SttWorkerOutbound> {
function wrapSubprocess(
spawned: SpawnedSubprocess<SttWorkerOutbound>,
): WorkerHandle<SttWorkerInbound, SttWorkerOutbound> {
): RefCountedWorkerHandle<SttWorkerInbound, SttWorkerOutbound> {
const { proc } = spawned;
return createWorkerHandle<SttWorkerInbound, SttWorkerOutbound>(spawned, message => safeSend(proc, message, "stt"));
return {
...createWorkerHandle<SttWorkerInbound, SttWorkerOutbound>(spawned, message => safeSend(proc, message, "stt")),
ref() {
try {
proc.ref();
} catch {
// Already gone.
}
},
unref() {
try {
proc.unref();
} catch {
// Already gone.
}
},
};
}
function spawnSttWorker(): WorkerHandle<SttWorkerInbound, SttWorkerOutbound> {
function spawnInlineUnavailableWorker(error: unknown): RefCountedWorkerHandle<SttWorkerInbound, SttWorkerOutbound> {
return {
...createUnavailableWorker<SttWorkerInbound, SttWorkerOutbound>(error),
ref() {},
unref() {},
};
}
function spawnSttWorker(): RefCountedWorkerHandle<SttWorkerInbound, SttWorkerOutbound> {
return spawnWorkerOrUnavailable(
() => wrapSubprocess(createSttSubprocess()),
createUnavailableWorker<SttWorkerInbound, SttWorkerOutbound>,
spawnInlineUnavailableWorker,
"stt worker spawn failed; speech-to-text disabled",
);
}
export class SttClient {
#worker: WorkerHandle<SttWorkerInbound, SttWorkerOutbound> | null = null;
#worker: RefCountedWorkerHandle<SttWorkerInbound, SttWorkerOutbound> | null = null;
#unsubscribeMessage: (() => void) | null = null;
#unsubscribeError: (() => void) | null = null;
#pending = new Map<string, PendingRequest>();
#streams = new Map<string, StreamState>();
#progressListeners = new Set<(event: SttProgressEvent) => void>();
#nextRequestId = 0;
#spawnWorker: () => WorkerHandle<SttWorkerInbound, SttWorkerOutbound>;
#refed = false;
#spawnWorker: () => RefCountedWorkerHandle<SttWorkerInbound, SttWorkerOutbound>;
constructor(spawnWorker: () => WorkerHandle<SttWorkerInbound, SttWorkerOutbound> = spawnSttWorker) {
constructor(
spawnWorker: () => RefCountedWorkerHandle<SttWorkerInbound, SttWorkerOutbound> = spawnSttWorker,
) {
this.#spawnWorker = spawnWorker;
}
@@ -121,11 +153,11 @@ export class SttClient {
const worker = this.#ensureWorker();
const id = String(++this.#nextRequestId);
const { promise, resolve, reject } = Promise.withResolvers<string>();
this.#pending.set(id, { kind: "transcribe", modelKey, resolve, reject });
this.#addPending(id, { kind: "transcribe", modelKey, resolve, reject });
const abort = (): void => {
const pending = this.#pending.get(id);
if (pending?.kind !== "transcribe") return;
this.#pending.delete(id);
this.#deletePending(id);
pending.reject(new DOMException("The operation was aborted.", "AbortError"));
};
options.signal?.addEventListener("abort", abort, { once: true });
@@ -134,7 +166,7 @@ export class SttClient {
return await promise;
} finally {
options.signal?.removeEventListener("abort", abort);
this.#pending.delete(id);
this.#deletePending(id);
}
}
@@ -163,6 +195,7 @@ export class SttClient {
settled = true;
this.#streams.delete(id);
signal?.removeEventListener("abort", onAbort);
this.#syncWorkerRef();
apply();
};
this.#streams.set(id, {
@@ -173,6 +206,7 @@ export class SttClient {
reject,
finish,
});
this.#syncWorkerRef();
worker.send({ type: "stream_start", id, modelKey, language: options.language });
const handle: SttStreamHandle = {
pushAudio: audio => {
@@ -193,19 +227,19 @@ export class SttClient {
return handle;
}
async downloadModel(modelKey: SttModelKey, options: SttDownloadOptions = {}): Promise<boolean> {
if (options.signal?.aborted) return false;
async downloadModel(modelKey: SttModelKey, options: SttDownloadOptions = {}): Promise<SttDownloadResult> {
if (options.signal?.aborted) return { ok: false };
const unsubscribe = options.onProgress ? this.onProgress(options.onProgress) : undefined;
try {
const worker = this.#ensureWorker();
const id = String(++this.#nextRequestId);
const { promise, resolve } = Promise.withResolvers<boolean>();
this.#pending.set(id, { kind: "download", modelKey, resolve });
const { promise, resolve } = Promise.withResolvers<SttDownloadResult>();
this.#addPending(id, { kind: "download", modelKey, resolve });
const abort = (): void => {
const pending = this.#pending.get(id);
if (pending?.kind !== "download") return;
this.#pending.delete(id);
pending.resolve(false);
this.#deletePending(id);
pending.resolve({ ok: false });
};
options.signal?.addEventListener("abort", abort, { once: true });
try {
@@ -213,14 +247,15 @@ export class SttClient {
return await promise;
} finally {
options.signal?.removeEventListener("abort", abort);
this.#pending.delete(id);
this.#deletePending(id);
}
} catch (error) {
const message = error instanceof Error ? error.message : String(error);
logger.debug("stt: local model download failed", {
modelKey,
error: error instanceof Error ? error.message : String(error),
error: message,
});
return false;
return { ok: false, error: message };
} finally {
unsubscribe?.();
}
@@ -236,9 +271,10 @@ export class SttClient {
for (const pending of this.#pending.values()) {
this.#emitProgress({ modelKey: pending.modelKey, status: "error" });
if (pending.kind === "transcribe") pending.reject(new Error("stt worker terminated"));
else pending.resolve(false);
else pending.resolve({ ok: false });
}
this.#pending.clear();
this.#refed = false;
this.#failStreams(new Error("stt worker terminated"));
try {
await worker?.terminate();
@@ -247,7 +283,7 @@ export class SttClient {
}
}
#ensureWorker(): WorkerHandle<SttWorkerInbound, SttWorkerOutbound> {
#ensureWorker(): RefCountedWorkerHandle<SttWorkerInbound, SttWorkerOutbound> {
if (this.#worker) return this.#worker;
const worker = this.#spawnWorker();
this.#worker = worker;
@@ -256,6 +292,32 @@ export class SttClient {
return worker;
}
/** Register a pending request and keep the worker referenced while work is in flight. */
#addPending(id: string, request: PendingRequest): void {
this.#pending.set(id, request);
this.#syncWorkerRef();
}
/** Drop a pending request and unref the worker once no request or stream is active. */
#deletePending(id: string): void {
if (this.#pending.delete(id)) this.#syncWorkerRef();
}
/**
* STT workers start unreferenced so an idle warm model never blocks exit.
* Setup/download commands must keep the worker alive while awaiting IPC, or
* Bun can drain the event loop immediately after `Preparing Speech-to-Text`.
*/
#syncWorkerRef(): void {
const worker = this.#worker;
if (!worker) return;
const shouldRef = this.#pending.size > 0 || this.#streams.size > 0;
if (shouldRef === this.#refed) return;
this.#refed = shouldRef;
if (shouldRef) worker.ref();
else worker.unref();
}
#handleMessage(message: SttWorkerOutbound): void {
if (message.type === "log") {
logWorkerMessage(message);
@@ -287,19 +349,19 @@ export class SttClient {
}
return;
}
this.#pending.delete(message.id);
this.#deletePending(message.id);
if (message.type === "transcription") {
if (pending.kind === "transcribe") pending.resolve(message.text);
return;
}
if (message.type === "downloaded") {
if (pending.kind === "download") pending.resolve(true);
if (pending.kind === "download") pending.resolve({ ok: true });
return;
}
// message.type === "error"
this.#emitProgress({ modelKey: pending.modelKey, status: "error" });
if (pending.kind === "transcribe") pending.reject(new Error(message.error));
else pending.resolve(false);
else pending.resolve({ ok: false, error: message.error });
}
#emitProgress(event: SttProgressEvent): void {
@@ -318,7 +380,7 @@ export class SttClient {
for (const pending of this.#pending.values()) {
this.#emitProgress({ modelKey: pending.modelKey, status: "error" });
if (pending.kind === "transcribe") pending.reject(error);
else pending.resolve(false);
else pending.resolve({ ok: false, error: error.message });
}
this.#pending.clear();
this.#failStreams(error);
+8 -2
View File
@@ -90,7 +90,7 @@ export async function downloadSttModel(
): Promise<void> {
const spec = resolveSttModelSpec(key);
const files = new Map<string, { loaded: number; total: number }>();
const ok = await sttClient.downloadModel(spec.key, {
const result = await sttClient.downloadModel(spec.key, {
signal: options?.signal,
onProgress: event => {
if ((event.status === "progress" || event.status === "progress_total") && event.file) {
@@ -117,7 +117,13 @@ export async function downloadSttModel(
});
},
});
if (!ok) throw new Error(`Failed to download speech model (${spec.repo}). Check your network connection.`);
if (!result.ok) {
const detail = result.error ? `: ${result.error}` : ". Check your network connection.";
throw new Error(`Failed to download speech model (${spec.repo})${detail}`);
}
if (!(await isSttModelCached(spec.key))) {
throw new Error(`Speech model download finished without required files (${spec.repo}).`);
}
}
// ── Public API ─────────────────────────────────────────────────────
@@ -1,7 +1,8 @@
import { describe, expect, it } from "bun:test";
import { SttClient } from "@oh-my-pi/pi-coding-agent/stt/asr-client";
import type { SttWorkerInbound, SttWorkerOutbound } from "@oh-my-pi/pi-coding-agent/stt/asr-protocol";
import { TinyTitleClient } from "@oh-my-pi/pi-coding-agent/tiny/title-client";
import type { TinyTitleWorkerInbound, TinyTitleWorkerOutbound } from "@oh-my-pi/pi-coding-agent/tiny/title-protocol";
class FakeTinyWorker {
terminated = false;
refCalls = 0;
@@ -49,6 +50,49 @@ class FakeTinyWorker {
}
}
class FakeSttWorker {
terminated = false;
refCalls = 0;
unrefCalls = 0;
#messageHandlers = new Set<(message: SttWorkerOutbound) => void>();
#errorHandlers = new Set<(error: Error) => void>();
#onSend: (message: SttWorkerInbound, worker: FakeSttWorker) => void;
constructor(onSend: (message: SttWorkerInbound, worker: FakeSttWorker) => void) {
this.#onSend = onSend;
}
send(message: SttWorkerInbound): void {
this.#onSend(message, this);
}
onMessage(handler: (message: SttWorkerOutbound) => void): () => void {
this.#messageHandlers.add(handler);
return () => this.#messageHandlers.delete(handler);
}
onError(handler: (error: Error) => void): () => void {
this.#errorHandlers.add(handler);
return () => this.#errorHandlers.delete(handler);
}
async terminate(): Promise<void> {
this.terminated = true;
}
ref(): void {
this.refCalls += 1;
}
unref(): void {
this.unrefCalls += 1;
}
emit(message: SttWorkerOutbound): void {
for (const handler of this.#messageHandlers) handler(message);
}
}
describe("tiny title client prompt options", () => {
it("forwards a custom system prompt on local title requests", async () => {
let sent: TinyTitleWorkerInbound | undefined;
@@ -198,3 +242,48 @@ describe("issue #3291 — tiny-model downloads keep the worker referenced", () =
}
});
});
describe("issue #3939 — stt downloads keep the worker referenced", () => {
it("references the worker while a download request is pending", async () => {
let downloadRequestId = "";
const worker = new FakeSttWorker(message => {
if (message.type === "download") downloadRequestId = message.id;
});
const client = new SttClient(() => worker);
try {
const download = client.downloadModel("turbo");
expect(downloadRequestId).not.toBe("");
expect(worker.refCalls).toBe(1);
expect(worker.unrefCalls).toBe(0);
worker.emit({ type: "downloaded", id: downloadRequestId });
expect(await download).toEqual({ ok: true });
expect(worker.unrefCalls).toBe(1);
} finally {
await client.terminate();
}
});
it("surfaces worker download errors to setup callers", async () => {
let downloadRequestId = "";
const worker = new FakeSttWorker(message => {
if (message.type === "download") downloadRequestId = message.id;
});
const client = new SttClient(() => worker);
try {
const download = client.downloadModel("turbo");
expect(downloadRequestId).not.toBe("");
worker.emit({ type: "error", id: downloadRequestId, error: "Error: Hub returned 403" });
expect(await download).toEqual({ ok: false, error: "Error: Hub returned 403" });
expect(worker.unrefCalls).toBe(1);
} finally {
await client.terminate();
}
});
});