Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 22 additions & 2 deletions packages/cli/src/commands/models.test.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
import { CliRuntimeError, consumeCommandResult } from "../utils/commandResult.js";

const cachedRuntimePath =
"/cache/optional/sherpa-onnx-node@1.13.8/node_modules/sherpa-onnx-node/index.js";
const sherpa = {
SHERPA_RUNTIME_DIR: "/cache/optional/sherpa-onnx-node@1.13.8",
PARAKEET_MODEL_DIR: "/cache/parakeet/parakeet-tdt-0.6b-v3-int8",
Expand Down Expand Up @@ -44,7 +46,10 @@ describe("models install parakeet --json", () => {
consumeCommandResult();
vi.spyOn(console, "log").mockImplementation(() => {});
sherpa.sherpaUnsupportedReason.mockReturnValue(null);
sherpa.installSherpaRuntime.mockResolvedValue(true);
sherpa.installSherpaRuntime.mockResolvedValue({
installed: true,
runtimePath: cachedRuntimePath,
});
sherpa.ensureParakeetModel.mockResolvedValue(false);
});
afterEach(() => vi.restoreAllMocks());
Expand All @@ -58,17 +63,32 @@ describe("models install parakeet --json", () => {
model: "parakeet-tdt-0.6b-v3",
changed: true,
runtimeDir: sherpa.SHERPA_RUNTIME_DIR,
runtimePath: cachedRuntimePath,
modelDir: sherpa.PARAKEET_MODEL_DIR,
},
});
expect(dispose).toHaveBeenCalled();
});

it("reports changed: false when the runtime loads and the model verifies", async () => {
sherpa.installSherpaRuntime.mockResolvedValue(false);
sherpa.installSherpaRuntime.mockResolvedValue({
installed: false,
runtimePath: cachedRuntimePath,
});
expect((await install()).out).toMatchObject({ ok: true, changed: false });
});

it("reports a selected beside copy while retaining the cache destination", async () => {
const runtimePath = "/bundle/node_modules/sherpa-onnx-node/index.js";
sherpa.installSherpaRuntime.mockResolvedValue({ installed: false, runtimePath });
expect((await install()).out).toMatchObject({
ok: true,
changed: false,
runtimeDir: sherpa.SHERPA_RUNTIME_DIR,
runtimePath,
});
});

it("reports a failed install as ok:false with exit 1", async () => {
sherpa.ensureParakeetModel.mockRejectedValue(new Error("tokens.txt did not match"));
expect(await install()).toEqual({
Expand Down
19 changes: 13 additions & 6 deletions packages/cli/src/commands/models.ts
Original file line number Diff line number Diff line change
Expand Up @@ -43,18 +43,18 @@ function downloadProgress(spin: Spinner) {
type Sherpa = typeof import("../whisper/sherpa.js");
type Spinner = ReturnType<typeof clack.spinner> | null;

/** Installs what is missing; true when anything changed. */
/** Installs missing pieces and carries the selected runtime provenance into the result. */
async function installMissing(sherpa: Sherpa, spin: Spinner, signal: AbortSignal) {
const runtimeInstalled = await sherpa.installSherpaRuntime({ signal });
const runtime = await sherpa.installSherpaRuntime({ signal });
spin?.message("Verifying the Parakeet model...");
const modelFetched = await sherpa.ensureParakeetModel({
signal,
onBytes: downloadProgress(spin),
});
return runtimeInstalled || modelFetched;
return { changed: runtime.installed || modelFetched, runtimePath: runtime.runtimePath };
}

/** A synchronous child's Ctrl-C shows as its error before the scope's listener runs. */
/** Cancellation can reach the child before the scope observes the signal. */
const wasCancelled = (err: unknown, signal: AbortSignal, sherpa: Sherpa) =>
signal.aborted || err instanceof sherpa.DecodeCancelled;

Expand All @@ -68,12 +68,19 @@ async function installParakeet(json: boolean): Promise<void> {
const cancellation = createRenderCancellationScope();
spin?.start("Checking the sherpa-onnx runtime (installing it from npm if it does not load)...");
try {
const changed = await installMissing(sherpa, spin, cancellation.signal);
const { changed, runtimePath } = await installMissing(sherpa, spin, cancellation.signal);
spin?.stop(c.success(changed ? "Parakeet installed" : "Parakeet is already installed"));
if (json) {
const { SHERPA_RUNTIME_DIR: runtimeDir, PARAKEET_MODEL_DIR: modelDir } = sherpa;
console.log(
JSON.stringify({ ok: true, model: PARAKEET_MODEL_LABEL, changed, runtimeDir, modelDir }),
JSON.stringify({
ok: true,
model: PARAKEET_MODEL_LABEL,
changed,
runtimeDir,
runtimePath,
modelDir,
}),
);
}
} catch (err) {
Expand Down
39 changes: 34 additions & 5 deletions packages/cli/src/commands/transcribe.test.ts
Original file line number Diff line number Diff line change
@@ -1,8 +1,17 @@
import { runCommand } from "citty";
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
import { chmodSync, existsSync, writeFileSync, readFileSync, mkdtempSync, rmSync } from "node:fs";
import {
chmodSync,
mkdirSync,
existsSync,
writeFileSync,
readFileSync,
mkdtempSync,
rmSync,
} from "node:fs";
import { join } from "node:path";
import { tmpdir } from "node:os";
import { pathToFileURL } from "node:url";
import { WhisperUnavailableError } from "../whisper/manager.js";
import { CliRuntimeError, consumeCommandResult } from "../utils/commandResult.js";

Expand All @@ -18,10 +27,24 @@ vi.mock("../whisper/transcribe.js", () => ({

// Engine selection: which runners look installed, and the sherpa decode child it spawns.
const runners = { sherpa: false, mlx: false };
vi.mock("../whisper/sherpa.js", async (importOriginal) => ({
...(await importOriginal<typeof import("../whisper/sherpa.js")>()),
sherpaParakeetInstalled: () => runners.sherpa,
}));
let runtimeDir: string;
vi.mock("../whisper/sherpa.js", async (importOriginal) => {
const sherpa = await importOriginal<typeof import("../whisper/sherpa.js")>();
return {
...sherpa,
sherpaParakeetInstalled: () => runners.sherpa,
transcribeWithSherpa: (
wavPath: string,
dir: string,
options: Parameters<typeof sherpa.transcribeWithSherpa>[2],
) =>
sherpa.transcribeWithSherpa(wavPath, dir, {
...options,
runtimeDir,
cliUrl: pathToFileURL(join(runtimeDir, "dist", "cli.js")).href,
}),
};
});
const mlxMock = vi.fn();
vi.mock("../whisper/parakeet.js", async (importOriginal) => ({
...(await importOriginal<typeof import("../whisper/parakeet.js")>()),
Expand Down Expand Up @@ -165,6 +188,12 @@ describe("transcribe command", () => {

describe("engine selection", () => {
beforeEach(() => {
runtimeDir = mkdtempSync(join(tmpdir(), "hf-transcribe-runtime-"));
dirs.push(runtimeDir);
const pkg = join(runtimeDir, "node_modules", "sherpa-onnx-node");
mkdirSync(pkg, { recursive: true });
writeFileSync(join(pkg, "package.json"), JSON.stringify({ main: "index.js" }));
writeFileSync(join(pkg, "index.js"), "module.exports = {};");
transcribeMock.mockImplementation(async (_in: string, dir: string) =>
fakeTranscript(dir, "whisper"),
);
Expand Down
25 changes: 23 additions & 2 deletions packages/cli/src/utils/optionalPackages.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -132,9 +132,9 @@ describe("loadOptionalPackage", () => {

describe("a copy installed beside the CLI", () => {
// node_modules/hyperframes/dist/cli.js with onnxruntime-node installed next to hyperframes.
function layout(version: string) {
function layout(version: string, name = "onnxruntime-node") {
const root = mkdtempSync(join(tmpdir(), "hf-beside-"));
const pkg = join(root, "node_modules", "onnxruntime-node");
const pkg = join(root, "node_modules", name);
mkdirSync(pkg, { recursive: true });
writeFileSync(join(pkg, "package.json"), JSON.stringify({ version, main: "index.js" }));
writeFileSync(join(pkg, "index.js"), `module.exports = { copy: "beside ${version}" };`);
Expand All @@ -153,6 +153,27 @@ describe("a copy installed beside the CLI", () => {
}
});

it.each(["1.13.8", "1.13.7"])("accepts only the Sherpa 1.13.8 pin, received %s", (version) => {
const { root, cliUrl } = layout(version, "sherpa-onnx-node");
try {
const result = loadBesideCli("sherpa-onnx-node", cliUrl);
if (version === "1.13.8") expect(result).toEqual({ copy: "beside 1.13.8" });
else expect(result).toBeNull();
} finally {
rmSync(root, { recursive: true, force: true });
}
});

it("keeps a missing optional-package entry eligible for cache fallback", () => {
const { root, cliUrl } = layout("1.21.1");
try {
rmSync(join(root, "node_modules", "onnxruntime-node", "index.js"));
expect(loadBesideCli("onnxruntime-node", cliUrl)).toBeNull();
} finally {
rmSync(root, { recursive: true, force: true });
}
});

it("ignores it when its manifest is unreadable, instead of throwing", () => {
const { root, cliUrl } = layout(OPTIONAL_PACKAGES["onnxruntime-node"]);
writeFileSync(join(root, "node_modules", "onnxruntime-node", "package.json"), "{ version: 1");
Expand Down
52 changes: 41 additions & 11 deletions packages/cli/src/utils/optionalPackages.ts
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,13 @@ export const OPTIONAL_PACKAGES = {
"@google/genai": "1.52.0",
} as const satisfies Record<OptionalPackage, string>;

export const PINNED_PACKAGES = {
...OPTIONAL_PACKAGES,
"sherpa-onnx-node": "1.13.8",
} as const;

type PinnedPackage = keyof typeof PINNED_PACKAGES;

export const CACHE_DIR = join(homedir(), ".cache", "hyperframes", "optional");

export interface OptionalPackageDeps {
Expand All @@ -53,7 +60,7 @@ export function installedOptionalPackageVersion(
cacheDir = CACHE_DIR,
cliUrl = import.meta.url,
): string | null {
if (pinnedCopyBesideCli(name, cliUrl)) return OPTIONAL_PACKAGES[name];
if (pinnedPackageBesideCli(name, cliUrl)?.resolution === "entry") return OPTIONAL_PACKAGES[name];
const dir = optionalPackageDir(name, cacheDir);
if (!isInstalled(dir, name)) return null;
return (JSON.parse(readFileSync(manifestPath(dir, name), "utf-8")) as { version: string })
Expand Down Expand Up @@ -113,28 +120,51 @@ export function isInstalled(dir: string, name: string): boolean {
return existsSync(manifestPath(dir, name));
}

export function loadInstalled(dir: string, name: string): unknown | null {
export function installedPackagePath(dir: string, name: string): string | null {
if (!isInstalled(dir, name)) return null;
return createRequire(join(dir, "package.json"))(name);
return createRequire(join(dir, "package.json")).resolve(name);
}

function pinnedCopyBesideCli(name: OptionalPackage, cliUrl: string): boolean {
function loadInstalled(dir: string, name: string): unknown | null {
const entry = installedPackagePath(dir, name);
return entry === null ? null : createRequire(join(dir, "package.json"))(entry);
}

export function pinnedPackageBesideCli(
name: PinnedPackage,
cliUrl = import.meta.url,
): { resolution: "entry" | "package"; path: string } | null {
const req = createRequire(cliUrl);
try {
const entry = realpathSync(req.resolve(name));
let entry: string;
try {
entry = realpathSync(req.resolve(name));
} catch {
const copy = (req.resolve.paths(name) ?? [])
.map((dir) => join(dir, name))
.find((dir) => existsSync(join(dir, "package.json")));
return copy && hasPinnedVersion(copy, name)
? { resolution: "package", path: realpathSync(copy) }
: null;
}
const copy = (req.resolve.paths(name) ?? [])
.map((dir) => join(dir, name))
.find((dir) => existsSync(dir) && entry.startsWith(realpathSync(dir) + sep));
if (!copy) return false;
const manifest = readFileSync(join(copy, "package.json"), "utf-8");
return (JSON.parse(manifest) as { version?: string }).version === OPTIONAL_PACKAGES[name];
if (!copy) return null;
return hasPinnedVersion(copy, name) ? { resolution: "entry", path: entry } : null;
} catch {
return false;
return null;
}
}

export function loadBesideCli(name: OptionalPackage, cliUrl = import.meta.url): unknown | null {
return pinnedCopyBesideCli(name, cliUrl) ? createRequire(cliUrl)(name) : null;
function hasPinnedVersion(copy: string, name: PinnedPackage): boolean {
const manifest = readFileSync(join(copy, "package.json"), "utf-8");
return (JSON.parse(manifest) as { version?: string }).version === PINNED_PACKAGES[name];
}

export function loadBesideCli(name: PinnedPackage, cliUrl = import.meta.url): unknown | null {
const copy = pinnedPackageBesideCli(name, cliUrl);
return copy?.resolution === "entry" ? createRequire(cliUrl)(copy.path) : null;
}

export function runNpm(args: string[], signal?: AbortSignal): Promise<void> {
Expand Down
Loading
Loading