fix(security): pin curated Hugging Face revisions

This commit is contained in:
debpalash
2026-08-09 21:07:01 +00:00
parent 4945dec04e
commit f6f2bc5dcd
13 changed files with 332 additions and 19 deletions
+20 -7
View File
@@ -20,6 +20,7 @@ from pydantic import BaseModel
from core import prefs
from core.failure import is_hf_connectivity_error
from services.hf_revisions import revision_for
from utils import hf_progress
from utils import download_aggregator
# Weight-floor scan (MM2-07 / #352) lives in ``models.py`` — the lowest module in
@@ -170,7 +171,7 @@ def _repo_cancelled(repo_id: str) -> bool:
return repo_id in _cancelled
def _segmented_snapshot(repo_id: str, *, endpoint: "str | None") -> str:
def _segmented_snapshot(repo_id: str, *, endpoint: "str | None", revision: str) -> str:
"""Fetch every file of a repo via the segmented downloader into the HF
cache, mirroring hf_hub_download's blob+snapshot+refs layout so the result
is indistinguishable from snapshot_download (FDL-09) — keeping /models
@@ -187,10 +188,10 @@ def _segmented_snapshot(repo_id: str, *, endpoint: "str | None") -> str:
token = _resolve_token()
api = HfApi(endpoint=endpoint, token=token)
info = api.repo_info(repo_id, repo_type="model")
info = api.repo_info(repo_id, repo_type="model", revision=revision)
commit = info.sha
files = [s.rfilename for s in (info.siblings or [])]
if not commit or not files:
if commit != revision or not files:
raise RuntimeError("repo_info returned no commit/siblings")
repo_dir = os.path.join(_C.HF_HUB_CACHE, repo_folder_name(repo_id=repo_id, repo_type="model"))
@@ -407,6 +408,7 @@ async def install_model(req: InstallModelRequest):
# parallel-files worker count, and honour an optional mirror endpoint.
dl_kwargs: dict = {
"repo_id": req.repo_id,
"revision": revision_for(req.repo_id),
"max_workers": _download_max_workers(),
}
_tqdm_cls = hf_progress.tracked_tqdm_class()
@@ -447,11 +449,15 @@ async def install_model(req: InstallModelRequest):
# bytes that will actually download — BEFORE any byte flows. Seeds
# the overall aggregator so its bar/ETA are correct from the first
# event. Degrades gracefully (totals=None) on older/gated repos.
_preflight_kwargs = {"repo_id": req.repo_id, "dry_run": True}
_preflight_kwargs = {
"repo_id": req.repo_id,
"revision": dl_kwargs["revision"],
"dry_run": True,
}
if _endpoint:
_preflight_kwargs["endpoint"] = _endpoint
try:
_plan = snapshot_download(**_preflight_kwargs)
_plan = snapshot_download(**_preflight_kwargs) # nosec B615 -- immutable revision_for pin
_summary = compute_plan(_plan)
# Disk-space guard (before a single byte flows): the preflight
# gives an exact "to download" size, so reject an install that
@@ -515,7 +521,11 @@ async def install_model(req: InstallModelRequest):
_snapshot_path = None
if _attempt == 1 and _segmented_enabled() and not _xet_active():
try:
_snapshot_path = _segmented_snapshot(req.repo_id, endpoint=_endpoint)
_snapshot_path = _segmented_snapshot(
req.repo_id,
endpoint=_endpoint,
revision=dl_kwargs["revision"],
)
except _InstallCancelled:
raise
except Exception as _seg_err:
@@ -525,8 +535,11 @@ async def install_model(req: InstallModelRequest):
)
_snapshot_path = None
if _snapshot_path is None:
_snapshot_path = snapshot_download(**dl_kwargs)
_snapshot_path = snapshot_download(**dl_kwargs) # nosec B615 -- immutable revision_for pin
_validate_snapshot_has_weights(req.repo_id, _snapshot_path)
from huggingface_hub.constants import HF_HUB_CACHE
from services.hf_revisions import remember_revision
remember_revision(req.repo_id, dl_kwargs["revision"], HF_HUB_CACHE)
break
except Exception as net_err:
# #1224: a truncated body ("peer closed connection without
+29 -3
View File
@@ -36,6 +36,7 @@ from __future__ import annotations
import base64
import json
import os
import re
import struct
import sys
import traceback
@@ -50,8 +51,30 @@ MAX_FRAME_BYTES = 64 * 1024 * 1024
#: the loaded model at synthesize time; this is the handshake default).
MOSS_SAMPLE_RATE = 24000
#: HF repo id for the weights, overridable for air-gapped / mirror installs.
#: Reviewed HF repo and immutable revision for the default remote-code model.
_DEFAULT_REPO = "OpenMOSS-Team/MOSS-TTS-v1.5"
_DEFAULT_REVISION = "cdd3b911b1585e3f2dbc7775ef10f9926f58850a"
_SHA = re.compile(r"[0-9a-f]{40}\Z")
def _model_source() -> tuple[str, str]:
"""Return a pinned model source; custom remote code requires two opt-ins."""
repo = os.environ.get("OMNIVOICE_MOSS_TTS_V15_MODEL", _DEFAULT_REPO)
if repo == _DEFAULT_REPO:
return repo, _DEFAULT_REVISION
unsafe = os.environ.get("OMNIVOICE_MOSS_TTS_V15_TRUST_REMOTE_CODE", "").lower()
revision = os.environ.get("OMNIVOICE_MOSS_TTS_V15_REVISION", "")
if unsafe not in {"1", "true", "yes", "on"}:
raise RuntimeError(
"A custom MOSS model contains executable remote code. Set "
"OMNIVOICE_MOSS_TTS_V15_TRUST_REMOTE_CODE=1 only after auditing it."
)
if not _SHA.fullmatch(revision):
raise RuntimeError(
"A custom MOSS model requires OMNIVOICE_MOSS_TTS_V15_REVISION="
"<40-character commit SHA>; branches and tags are mutable."
)
return repo, revision
#: ISO-639-1 → MOSS language name. MOSS's ``build_user_message`` takes a
#: language *name* ("French"), not a code. Unknown codes are omitted so the
@@ -130,14 +153,16 @@ def _load_model(stdout):
import torch
from transformers import AutoModel, AutoProcessor
repo = os.environ.get("OMNIVOICE_MOSS_TTS_V15_MODEL", _DEFAULT_REPO)
repo, revision = _model_source()
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.bfloat16 if device == "cuda" else torch.float32
# "sdpa" works on CUDA + CPU and needs no extra dep. flash_attention_2
# (Ampere+ CUDA, optional flash-attn) is opt-in via env.
attn = os.environ.get("OMNIVOICE_MOSS_TTS_V15_ATTN", "sdpa")
processor = AutoProcessor.from_pretrained(repo, trust_remote_code=True)
processor = AutoProcessor.from_pretrained(
repo, revision=revision, trust_remote_code=True,
)
# The audio tokenizer is a separate sub-module that must be moved to the
# device independently (easy to miss — see upstream README).
processor.audio_tokenizer = processor.audio_tokenizer.to(device)
@@ -146,6 +171,7 @@ def _load_model(stdout):
model = AutoModel.from_pretrained(
repo,
revision=revision,
trust_remote_code=True,
attn_implementation=attn,
torch_dtype=dtype,
+9 -3
View File
@@ -209,6 +209,12 @@ def repair_repo_cache(repo_id: str, cache_dir: str | None = None) -> dict:
)
return summary
# Resolve the exact revision before deleting even a broken pointer. An
# unreviewed repository can never turn cache repair into a mutable-
# branch network fetch.
from services.hf_revisions import installed_revision
revision = installed_revision(repo_id, cache_root)
def _remove(paths: list[str]) -> int:
n = 0
for path in paths:
@@ -232,13 +238,13 @@ def repair_repo_cache(repo_id: str, cache_dir: str | None = None) -> dict:
from huggingface_hub import snapshot_download
dl_kwargs: dict = {"repo_id": repo_id}
dl_kwargs: dict = {"repo_id": repo_id, "revision": revision}
if cache_dir:
dl_kwargs["cache_dir"] = cache_dir
endpoint = os.environ.get("HF_ENDPOINT")
if endpoint:
dl_kwargs["endpoint"] = endpoint
snapshot_download(**dl_kwargs)
snapshot_download(**dl_kwargs) # nosec B615 -- installed immutable revision
summary["restored"] = True
# Verify-after-repair: hub's memoized symlink probe can claim support
@@ -267,7 +273,7 @@ def repair_repo_cache(repo_id: str, cache_dir: str | None = None) -> dict:
)
return summary
summary["removed"] += _remove(still_broken)
snapshot_download(**dl_kwargs)
snapshot_download(**dl_kwargs) # nosec B615 -- installed immutable revision
remaining = find_dangling_entries(repo_dir)
if remaining:
summary["error"] = (
+90
View File
@@ -0,0 +1,90 @@
"""Immutable Hugging Face revisions for VoiceStudio's curated repositories.
Branch names are mutable supply-chain inputs. Every repo the product offers is
resolved here to a reviewed commit SHA; download, preflight, and repair paths
must call :func:`revision_for` instead of following ``main``.
"""
from __future__ import annotations
import os
import re
from pathlib import Path
_SHA = re.compile(r"[0-9a-f]{40}\Z")
CURATED_REVISIONS: dict[str, str] = {
"k2-fsa/OmniVoice": "c5fdb5ccb189668d56333f77ba2629f4cd7535f4",
"Systran/faster-whisper-large-v3": "edaa852ec7e145841d8ffdb056a99866b5f0a478",
"mlx-community/whisper-large-v3-mlx": "49e6aa286ad60c14352c404340ded53710378a11",
"mlx-community/whisper-large-v3-turbo": "a4aaeec0636e6fef84abdcbe3544cb2bf7e9f6fb",
"openai/whisper-large-v3": "06f233fe06e710322aca913c1bc4249a0d71fce1",
"mlx-community/whisper-tiny-mlx": "6caf9c55601caafbe6508a8b0d216bdf4783c4e8",
"deepdml/faster-whisper-large-v3-turbo-ct2": "4df90f75321148c3a29a9e2351b7ddf8f5b115a8",
"Systran/faster-distil-whisper-large-v3": "c3058b475261292e64a0412df1d2681c06260fab",
"Systran/faster-whisper-medium": "08e178d48790749d25932bbc082711ddcfdfbc4f",
"Systran/faster-whisper-small": "536b0662742c02347bc0e980a01041f333bce120",
"Systran/faster-whisper-base": "ebe41f70d5b6dfa9166e2c581c45c9c0cfc57b66",
"nvidia/parakeet-tdt-0.6b-v3": "541d1f99c6b0c3cd0b11a95167540bb8edefd82b",
"nvidia/parakeet-tdt-0.6b-v2": "ae9ad07059c7c739ffaf932226a8fe64ae2620b0",
"mlx-community/parakeet-tdt-0.6b-v3": "ed2b7e8c15f9aaa0b5772e2efb986255eaef7e15",
"UsefulSensors/moonshine-base": "7a73d8d55ac0ba2ef3ae761593f6784b51f96dcf",
"UsefulSensors/moonshine-tiny": "390624ed33d594443aa4aa221f5b9f283b545b5a",
"csukuangfj/sherpa-onnx-nemo-parakeet-tdt-0.6b-v3-int8": "2bda32ec70b097a55adaa07d9a7173915b43cc78",
"csukuangfj/sherpa-onnx-nemo-parakeet-tdt-0.6b-v2-int8": "1ab9323565ddb038682214b292f588070a538ce2",
"csukuangfj/sherpa-onnx-streaming-zipformer-bilingual-zh-en-2023-02-20": "98590b7ed6443e77b714204da2757d75e1a642f4",
"csukuangfj/sherpa-onnx-streaming-paraformer-bilingual-zh-en": "8e40c43232a1c5c66c82111efc5820d3accca11b",
"csukuangfj/sherpa-onnx-streaming-zipformer-en-20M-2023-02-17": "d42f2d9f7ca24806fb667456a18a9f1b60f70d16",
"csukuangfj/sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23": "204ad334e2e683fd295359930cc16fc0432a23ac",
"csukuangfj/sherpa-onnx-whisper-tiny": "65176e2deb88badc814a94058666cadccc29b61c",
"pyannote/speaker-diarization-3.1": "84fd25912480287da0247647c3d2b4853cb3ee5d",
"OpenMOSS-Team/MOSS-TTS-Nano-100M": "44502f80dbf9743528fa921cc544d662c685ebec",
"KittenML/kitten-tts-mini-0.8": "c02725660cea441db4c383af69f1f26f5cd00947",
"mlx-community/Kokoro-82M-bf16": "a71e4d38b236d968966a2002c4c895dbd12b1c3c",
"mlx-community/csm-1b-8bit": "fcf0cc857eade3615a60f30722cf5197d4f88406",
"mlx-community/Qwen3-TTS-12Hz-1.7B-VoiceDesign-4bit": "5c390979e4b93af5f2932f90742ca99c7dd04687",
"mlx-community/Dia-1.6B": "de4fa8c178ca5cc4e9d884b55b03fcfaa0995162",
"mlx-community/Llama-OuteTTS-1.0-1B-4bit": "3ac2cff406f7de16a3216c60d0108571a916acc0",
"mlx-community/Chatterbox-TTS-4bit": "a3c8ded2d711d6395410d645b3a97c79fd563a13",
"mlx-community/MeloTTS-English-v3-MLX": "837d15fd72bc35a15033234ce5ea242367ca1960",
"OpenMOSS-Team/MOSS-TTS-v1.5": "cdd3b911b1585e3f2dbc7775ef10f9926f58850a",
}
def revision_for(repo_id: str) -> str:
"""Return the immutable revision for a curated repo, or raise."""
try:
return CURATED_REVISIONS[repo_id]
except KeyError as exc:
raise ValueError(f"No reviewed revision is pinned for {repo_id!r}") from exc
def _repo_dir(repo_id: str, cache_dir: str) -> Path:
return Path(cache_dir) / ("models--" + repo_id.replace("/", "--"))
def remember_revision(repo_id: str, revision: str, cache_dir: str) -> None:
"""Persist the exact installed revision for later in-place repair."""
if not _SHA.fullmatch(revision):
raise ValueError("Hugging Face revision must be a 40-character commit SHA")
repo_dir = _repo_dir(repo_id, cache_dir)
repo_dir.mkdir(parents=True, exist_ok=True)
marker = repo_dir / "voicestudio-revision"
temporary = marker.with_suffix(f".tmp-{os.getpid()}")
temporary.write_text(revision + "\n", encoding="ascii")
os.replace(temporary, marker)
def installed_revision(repo_id: str, cache_dir: str) -> str:
"""Return VoiceStudio's recorded revision, falling back to the curated pin."""
repo_dir = _repo_dir(repo_id, cache_dir)
# New installs write the first marker. ``refs/main`` preserves the commit
# resolved by older VoiceStudio/huggingface_hub installs, so upgrades repair
# the bytes the user actually installed rather than silently changing them.
for marker in (repo_dir / "voicestudio-revision", repo_dir / "refs" / "main"):
try:
revision = marker.read_text(encoding="ascii").strip()
except OSError:
continue
if _SHA.fullmatch(revision):
return revision
return revision_for(repo_id)
+11 -3
View File
@@ -1626,7 +1626,15 @@ def _repair_model_cache(checkpoint: str, *, force: bool = False) -> bool:
logger.warning("Cannot import snapshot_download to repair cache: %s", imp_err)
_last_repair_error = f"{type(imp_err).__name__}: {imp_err}"
return False
dl_kwargs: dict = {"repo_id": checkpoint}
try:
from services.hf_cache_repair import hf_cache_home
from services.hf_revisions import installed_revision
revision = installed_revision(checkpoint, hf_cache_home())
except (OSError, ValueError) as revision_err:
_last_repair_error = str(revision_err)
logger.warning("Refusing unpinned model repair for %s: %s", checkpoint, revision_err)
return False
dl_kwargs: dict = {"repo_id": checkpoint, "revision": revision}
# Explicit endpoint (HF_ENDPOINT / pref) wins; otherwise the automatic
# endpoint selection's cached pick applies (services.endpoint_race).
try:
@@ -1647,12 +1655,12 @@ def _repair_model_cache(checkpoint: str, *, force: bool = False) -> bool:
"""One snapshot_download, tolerating an hf_hub that rejects the optional
symlink knob. Lets real failures (network, gated repo, disk) propagate."""
try:
snapshot_download(**dl_kwargs)
snapshot_download(**dl_kwargs) # nosec B615 -- installed immutable revision
except TypeError:
# Older/newer huggingface_hub may not accept local_dir_use_symlinks
# on a cache-only call — retry without the optional knob.
dl_kwargs.pop("local_dir_use_symlinks", None)
snapshot_download(**dl_kwargs)
snapshot_download(**dl_kwargs) # nosec B615 -- installed immutable revision
# Bounded retries (#739): an incomplete cache *is* an interrupted download, so
# a single transient blip mid-repair shouldn't drop the user back to a manual
+8 -1
View File
@@ -272,11 +272,14 @@ def _resolve_model_dir(spec: SherpaModelSpec, *, download: bool = True) -> str:
so we never pull the bundled fp32 weights or test wavs.
"""
from huggingface_hub import snapshot_download
from services.hf_revisions import revision_for
wanted = list(spec.files.values())
revision = revision_for(spec.repo_id)
try:
return snapshot_download(
repo_id=spec.repo_id,
revision=revision,
local_files_only=True,
allow_patterns=wanted,
)
@@ -284,7 +287,11 @@ def _resolve_model_dir(spec: SherpaModelSpec, *, download: bool = True) -> str:
if not download:
raise
logger.info("sherpa dictation: downloading %s on first use", spec.repo_id)
return snapshot_download(repo_id=spec.repo_id, allow_patterns=wanted)
return snapshot_download(
repo_id=spec.repo_id,
revision=revision,
allow_patterns=wanted,
)
def is_installed(spec: SherpaModelSpec) -> bool:
+3 -1
View File
@@ -104,7 +104,9 @@ default voice. `duration` (seconds) maps to MOSS's `tokens` argument at
| Variable | Default | Purpose |
|----------|---------|---------|
| `OMNIVOICE_MOSS_TTS_V15_DIR` | — | Path to the MOSS-TTS clone (required). |
| `OMNIVOICE_MOSS_TTS_V15_MODEL` | `OpenMOSS-Team/MOSS-TTS-v1.5` | HF repo id override (mirror / air-gapped). |
| `OMNIVOICE_MOSS_TTS_V15_MODEL` | `OpenMOSS-Team/MOSS-TTS-v1.5` | Advanced HF repo override. Custom repositories execute their modelling code and are rejected unless both safeguards below are set. Configure `HF_ENDPOINT` for a mirror instead. |
| `OMNIVOICE_MOSS_TTS_V15_REVISION` | — | Immutable 40-character commit SHA required with a custom model repository. |
| `OMNIVOICE_MOSS_TTS_V15_TRUST_REMOTE_CODE` | — | Set to `1` only after auditing a custom repository's Python code. The reviewed built-in repository needs no opt-in. |
| `OMNIVOICE_MOSS_TTS_V15_ATTN` | `sdpa` | Attention impl; set `flash_attention_2` on Ampere+ CUDA with `flash-attn` installed. |
## Common errors
+11 -1
View File
@@ -49,6 +49,12 @@ def _no_ambient_offline_mode(monkeypatch):
explicitly via monkeypatch.setenv."""
monkeypatch.delenv("HF_HUB_OFFLINE", raising=False)
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising=False)
from services import hf_revisions
monkeypatch.setitem(
hf_revisions.CURATED_REVISIONS,
"test/checkpoint",
"a" * 40,
)
def _mk_repo_cache(tmp_path, repo_id: str = "test/checkpoint"):
@@ -155,7 +161,11 @@ def test_repair_removes_only_broken_and_redownloads(tmp_path, monkeypatch):
assert summary["error"] == ""
# The broken entry is gone; snapshot_download was asked to restore it.
assert not os.path.lexists(snap / "model.safetensors")
assert calls == [{"repo_id": "test/checkpoint", "cache_dir": str(cache)}]
assert calls == [{
"repo_id": "test/checkpoint",
"revision": "a" * 40,
"cache_dir": str(cache),
}]
# Healthy entries and blobs are untouched.
assert (snap / "config.json").read_bytes() == b'{"ok": true}'
assert (snap / "tokenizer.json").read_bytes() == b'{"tok": 1}'
+52
View File
@@ -0,0 +1,52 @@
"""Curated Hugging Face inputs are immutable and repair preserves installs."""
from pathlib import Path
import yaml
from services import hf_revisions
def test_every_catalog_repo_has_an_immutable_revision():
catalog = yaml.safe_load(Path("backend/config/models.yaml").read_text(encoding="utf-8"))
missing = {
model["repo_id"]
for model in catalog["models"]
if model["repo_id"] not in hf_revisions.CURATED_REVISIONS
}
assert missing == set()
assert all(len(revision) == 40 for revision in hf_revisions.CURATED_REVISIONS.values())
assert all(int(revision, 16) >= 0 for revision in hf_revisions.CURATED_REVISIONS.values())
def test_installed_revision_round_trips_for_repair(tmp_path):
repo_id = "k2-fsa/OmniVoice"
installed = "f" * 40
hf_revisions.remember_revision(repo_id, installed, str(tmp_path))
assert hf_revisions.installed_revision(repo_id, str(tmp_path)) == installed
def test_missing_or_invalid_marker_falls_back_to_reviewed_pin(tmp_path):
repo_id = "k2-fsa/OmniVoice"
assert hf_revisions.installed_revision(repo_id, str(tmp_path)) == hf_revisions.revision_for(repo_id)
marker = tmp_path / "models--k2-fsa--OmniVoice" / "voicestudio-revision"
marker.parent.mkdir(parents=True)
marker.write_text("main\n", encoding="ascii")
assert hf_revisions.installed_revision(repo_id, str(tmp_path)) == hf_revisions.revision_for(repo_id)
def test_existing_hub_main_ref_is_preserved_for_upgrade_repair(tmp_path):
repo_id = "k2-fsa/OmniVoice"
existing = "e" * 40
ref = tmp_path / "models--k2-fsa--OmniVoice" / "refs" / "main"
ref.parent.mkdir(parents=True)
ref.write_text(existing + "\n", encoding="ascii")
assert hf_revisions.installed_revision(repo_id, str(tmp_path)) == existing
def test_unknown_repo_cannot_start_a_network_repair(tmp_path):
try:
hf_revisions.installed_revision("attacker/unreviewed", str(tmp_path))
except ValueError as exc:
assert "No reviewed revision" in str(exc)
else: # pragma: no cover - assertion message is clearer than pytest.raises here
raise AssertionError("unreviewed repository unexpectedly received a revision")
+44
View File
@@ -90,7 +90,10 @@ def test_install_worker_emits_install_error_and_skips_download(models_mod, monke
# call must never happen (the guard returns first). Fail loudly if it does.
import huggingface_hub
calls = []
def _fake_snapshot(**kwargs):
calls.append(kwargs)
if kwargs.get("dry_run"):
return [] # preflight plan input (compute_plan is stubbed anyway)
raise AssertionError("snapshot_download called for a real download despite disk-full reject")
@@ -116,5 +119,46 @@ def test_install_worker_emits_install_error_and_skips_download(models_mod, monke
errs = [e for e in events if e.get("phase") == "install_error"]
assert errs, f"expected an install_error event, got phases: {[e.get('phase') for e in events]}"
assert "disk space" in errs[0]["error"].lower()
from services.hf_revisions import revision_for
assert calls[0]["revision"] == revision_for(repo_id)
# And the resolving heartbeat must not have leaked — no infinite 'resolving'
# stream after the bail (the worker set the stop event before returning).
def test_install_preflight_download_and_repair_marker_share_revision(models_mod, monkeypatch):
download = importlib.import_module("api.routers.setup.download")
import asyncio
import huggingface_hub
from services import hf_revisions
repo_id = download.KNOWN_MODELS[0]["repo_id"]
expected = hf_revisions.revision_for(repo_id)
calls = []
remembered = []
def fake_snapshot(**kwargs):
calls.append(kwargs)
return [] if kwargs.get("dry_run") else "/cache/snapshots/" + expected
monkeypatch.setattr(huggingface_hub, "snapshot_download", fake_snapshot)
monkeypatch.setattr(download, "compute_plan", lambda _plan: {
"total_bytes": 1, "cached_bytes": 0, "to_download_bytes": 1,
"n_files": 1, "n_cached": 0,
})
monkeypatch.setattr(download, "disk_space_error", lambda *_a, **_k: None)
monkeypatch.setattr(download, "_segmented_enabled", lambda: False)
monkeypatch.setattr(download, "_validate_snapshot_has_weights", lambda *_a: None)
monkeypatch.setattr(hf_revisions, "remember_revision", lambda *args: remembered.append(args))
async def run_install():
await download.install_model(download.InstallModelRequest(repo_id=repo_id))
pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
if pending:
await asyncio.gather(*pending)
asyncio.run(run_install())
assert [call["revision"] for call in calls] == [expected, expected]
assert calls[0]["dry_run"] is True
assert "dry_run" not in calls[1]
assert remembered and remembered[0][0:2] == (repo_id, expected)
+3
View File
@@ -23,6 +23,7 @@ def model_manager(monkeypatch):
sys.modules.pop(mod_name, None)
import services.model_manager as mm
from services import hf_revisions
monkeypatch.setattr(mm, "_torch", None)
monkeypatch.setattr(mm, "_OmniVoice", None)
@@ -33,6 +34,7 @@ def model_manager(monkeypatch):
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising=False)
monkeypatch.setattr(mm, "_lazy_torch", lambda: SimpleNamespace(float16="float16"))
monkeypatch.setattr(mm, "get_best_device", lambda: "cpu")
monkeypatch.setitem(hf_revisions.CURATED_REVISIONS, "test/checkpoint", "a" * 40)
return mm
@@ -202,6 +204,7 @@ def test_repair_invokes_snapshot_download(model_manager, monkeypatch):
assert model_manager._repair_model_cache("test/checkpoint") is True
assert calls and calls[0]["repo_id"] == "test/checkpoint"
assert calls[0]["revision"] == "a" * 40
def test_repair_returns_false_when_download_fails(model_manager, monkeypatch):
+33
View File
@@ -63,6 +63,39 @@ def test_sidecar_script_ships():
assert MOSS_TTS_V15_SIDECAR_SCRIPT.is_file()
def test_default_model_source_is_pinned(monkeypatch):
from engines.moss_tts_v15 import main
monkeypatch.delenv("OMNIVOICE_MOSS_TTS_V15_MODEL", raising=False)
assert main._model_source() == (main._DEFAULT_REPO, main._DEFAULT_REVISION)
def test_custom_remote_code_is_rejected_without_explicit_opt_in(monkeypatch):
from engines.moss_tts_v15 import main
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_MODEL", "someone/custom-model")
monkeypatch.delenv("OMNIVOICE_MOSS_TTS_V15_TRUST_REMOTE_CODE", raising=False)
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_REVISION", "a" * 40)
with pytest.raises(RuntimeError, match="after auditing"):
main._model_source()
def test_custom_remote_code_requires_immutable_revision(monkeypatch):
from engines.moss_tts_v15 import main
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_MODEL", "someone/custom-model")
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_TRUST_REMOTE_CODE", "1")
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_REVISION", "main")
with pytest.raises(RuntimeError, match="40-character commit SHA"):
main._model_source()
def test_audited_custom_remote_code_requires_both_opt_ins(monkeypatch):
from engines.moss_tts_v15 import main
revision = "b" * 40
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_MODEL", "someone/custom-model")
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_TRUST_REMOTE_CODE", "true")
monkeypatch.setenv("OMNIVOICE_MOSS_TTS_V15_REVISION", revision)
assert main._model_source() == ("someone/custom-model", revision)
# ── hardware honesty (cross-platform rule) ─────────────────────────────────
+19
View File
@@ -174,6 +174,25 @@ def test_get_spec_accepts_repo_id():
assert not sd.is_sherpa_model(None)
def test_model_resolution_pins_offline_probe_and_download(monkeypatch):
from services import hf_revisions, sherpa_dictation as sd
import huggingface_hub
spec = sd.get_spec("sherpa-whisper-tiny")
calls = []
def fake_snapshot(**kwargs):
calls.append(kwargs)
if kwargs.get("local_files_only"):
raise FileNotFoundError("not cached")
return "/cache/pinned"
monkeypatch.setattr(huggingface_hub, "snapshot_download", fake_snapshot)
assert sd._resolve_model_dir(spec) == "/cache/pinned"
assert len(calls) == 2
assert all(call["revision"] == hf_revisions.revision_for(spec.repo_id) for call in calls)
# ── The 4 recognizer kinds construct + transcribe ───────────────────────────