362 lines
12 KiB
Python
362 lines
12 KiB
Python
"""Explicit transcription model management (Hex-inspired design).
|
|
|
|
Adapts the model-management concept from the Hex macOS app to the fcp-mcp-server
|
|
stack (Python + MCP), keeping our conventions: rational, allowlisted model
|
|
names, lazy optional imports, graceful degradation, and I/O confined to the
|
|
model cache directory.
|
|
|
|
This module is the skeleton of the manager. The catalog and cache primitives
|
|
are real; the MCP handlers and wiring into ``transcribe()`` come in a later
|
|
phase (see docs/TRANSCRIPTION-MODELS.md).
|
|
|
|
Model weights are the same Systran/faster-whisper artifacts that
|
|
``transcribe.py`` already loads, so install status here matches what
|
|
``WhisperModel(model_size, ...)`` would download on its own.
|
|
"""
|
|
|
|
import json
|
|
import logging
|
|
import re
|
|
import threading
|
|
from pathlib import Path
|
|
from typing import Callable, List, Optional
|
|
|
|
from .transcribe import ALLOWED_MODELS
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# faster-whisper (CTranslate2) caches HF snapshots under ~/.cache/huggingface/hub
|
|
# as models--Systran--faster-whisper-<size>. This is the on-disk truth for
|
|
# "is downloaded".
|
|
_HF_REPO = "Systran/faster-whisper"
|
|
|
|
# Default cache root — matches what faster-whisper (HF hub) uses on its own,
|
|
# so a default install lines up with anything already on disk.
|
|
_DEFAULT_MODELS_DIR = Path.home() / ".cache" / "huggingface" / "hub"
|
|
|
|
# Config file for the persisted model selection + models dir.
|
|
_CONFIG_DIR = Path.home() / ".fcp-mcp-server"
|
|
_CONFIG_FILE = _CONFIG_DIR / "config.json"
|
|
|
|
# Path-traversal guard: only allow [A-Za-z0-9._-] in an internal model name
|
|
# (already enforced by ALLOWED_MODELS, but the cache-dir helper is belt + braces).
|
|
_SAFE_NAME_RE = re.compile(r"^[\w.-]+$")
|
|
|
|
# Progress callback signature used across the module.
|
|
ProgressCallback = Callable[[float], None]
|
|
|
|
|
|
def get_models_dir() -> Path:
|
|
"""The configured models root (falls back to the default HF hub cache)."""
|
|
configured = _load_config().get("models_dir")
|
|
if configured:
|
|
return Path(configured)
|
|
return _DEFAULT_MODELS_DIR
|
|
|
|
|
|
def save_models_dir(path: str) -> str:
|
|
"""Persist the models root directory. Returns the stored value."""
|
|
if not path.strip():
|
|
raise ValueError("models_dir cannot be empty")
|
|
_CONFIG_DIR.mkdir(parents=True, exist_ok=True)
|
|
data = _load_config()
|
|
data["models_dir"] = str(Path(path).expanduser())
|
|
_write_config(data)
|
|
return data["models_dir"]
|
|
|
|
|
|
def _load_config() -> dict:
|
|
"""Read the full config JSON (never raises; returns {} on error)."""
|
|
try:
|
|
return json.loads(_CONFIG_FILE.read_text(encoding="utf-8"))
|
|
except (OSError, ValueError):
|
|
return {}
|
|
|
|
|
|
def _write_config(data: dict) -> None:
|
|
_CONFIG_DIR.mkdir(parents=True, exist_ok=True)
|
|
_CONFIG_FILE.write_text(json.dumps(data, indent=2), encoding="utf-8")
|
|
|
|
|
|
def _hf_snapshot_dir(model_size: str) -> Path:
|
|
"""Cache folder for a given model size's HF snapshot, under the models root.
|
|
|
|
``Systran/faster-whisper`` -> ``<root>/models--Systran--faster-whisper-<size>``.
|
|
"""
|
|
name = "-".join(_HF_REPO.replace("/", "--").split("-")) + f"-{model_size}"
|
|
return get_models_dir() / f"models--{name}"
|
|
|
|
|
|
def model_cache_dir(model_size: str) -> Path:
|
|
"""The on-disk cache directory for ``model_size`` (as downloaded on disk)."""
|
|
if not _SAFE_NAME_RE.match(model_size):
|
|
raise ValueError(f"Unsafe model name: {model_size!r}")
|
|
return _hf_snapshot_dir(model_size)
|
|
|
|
|
|
def _load_bundled_catalog() -> Optional[list[dict]]:
|
|
"""Read the curated catalog from the bundled ``models.json`` (lazy, cached)."""
|
|
global _catalog
|
|
if _catalog is not None:
|
|
return _catalog
|
|
path = Path(__file__).with_name("models.json")
|
|
try:
|
|
with path.open(encoding="utf-8") as fh:
|
|
_catalog = json.load(fh)
|
|
except (OSError, ValueError):
|
|
logger.warning("Failed to load bundled models.json (%s)", path)
|
|
_catalog = []
|
|
return _catalog
|
|
|
|
|
|
_catalog: Optional[list[dict]] = None
|
|
|
|
|
|
def load_catalog() -> list[dict]:
|
|
"""All curated models (or ``[]`` if the bundled file is unreadable)."""
|
|
return list(_load_bundled_catalog() or [])
|
|
|
|
|
|
def get_catalog_model(internal_name: str) -> Optional[dict]:
|
|
"""The catalog entry whose ``internal_name`` equals ``internal_name``."""
|
|
for entry in load_catalog():
|
|
if entry.get("internal_name") == internal_name:
|
|
return entry
|
|
return None
|
|
|
|
|
|
def is_model_downloaded(model_size: str) -> bool:
|
|
"""True when the model's cache snapshot exists and isn't an empty dir.
|
|
|
|
Never raises on I/O; reports ``False`` for missing/unreadable dirs so
|
|
callers can offer a download instead of crashing.
|
|
"""
|
|
try:
|
|
d = model_cache_dir(model_size)
|
|
except ValueError:
|
|
return False
|
|
if not d.is_dir():
|
|
return False
|
|
try:
|
|
return any(d.iterdir())
|
|
except OSError:
|
|
return False
|
|
|
|
|
|
def list_installed_models() -> List[str]:
|
|
"""Model sizes present in the cache, filtered to the known allowlist."""
|
|
root = get_models_dir()
|
|
if not root.is_dir():
|
|
return []
|
|
installed: List[str] = []
|
|
for entry in root.iterdir():
|
|
if not entry.is_dir():
|
|
continue
|
|
for size in ALLOWED_MODELS:
|
|
if entry.name == _hf_snapshot_dir(size).name and is_model_downloaded(size):
|
|
installed.append(size)
|
|
break
|
|
# Stable, de-duped order following the allowlist.
|
|
return [s for s in ALLOWED_MODELS if s in installed]
|
|
|
|
|
|
def download_model(
|
|
model_size: str,
|
|
*,
|
|
progress_cb: Optional[ProgressCallback] = None,
|
|
cancel_event: Optional[threading.Event] = None,
|
|
) -> Optional[Path]:
|
|
"""Download a model snapshot to the cache.
|
|
|
|
Validates ``model_size`` against ``ALLOWED_MODELS`` and returns ``None``
|
|
(with a logged install hint) when ``huggingface_hub`` is unavailable —
|
|
the same graceful-degradation contract as ``transcribe()``.
|
|
|
|
``cancel_event`` (a ``threading.Event``) aborts the download on the next
|
|
progress tick if set; the partially-downloaded snapshot is removed.
|
|
"""
|
|
if model_size not in ALLOWED_MODELS:
|
|
raise ValueError(
|
|
f"model_size must be one of {', '.join(ALLOWED_MODELS)}, got {model_size!r}"
|
|
)
|
|
try:
|
|
from huggingface_hub import snapshot_download
|
|
except ImportError:
|
|
logger.info(
|
|
"huggingface_hub not installed; install the [transcribe] extra "
|
|
"to enable model download"
|
|
)
|
|
return None
|
|
|
|
target = model_cache_dir(model_size)
|
|
|
|
# tqdm hook that both reports progress and honors cancellation.
|
|
from tqdm import tqdm
|
|
|
|
class _CancelableTqdm(tqdm):
|
|
def update(self, n=1):
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
raise _DownloadCancelledError(model_size)
|
|
super().update(n)
|
|
|
|
def _on_progress(current: int, total: int) -> None:
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
raise _DownloadCancelledError(model_size)
|
|
if progress_cb is not None and total > 0:
|
|
progress_cb(current / total)
|
|
|
|
try:
|
|
snapshot_download(
|
|
repo_id=_HF_REPO,
|
|
revision=model_size,
|
|
# write into the same snapshot folder faster-whisper expects,
|
|
# under the configured models root.
|
|
cache_dir=str(get_models_dir()),
|
|
local_dir=str(target),
|
|
tqdm_class=_CancelableTqdm,
|
|
local_dir_use_symlinks=False,
|
|
)
|
|
# Local snapshot already landed in `target`; simpler than a temp+move.
|
|
except _DownloadCancelledError:
|
|
logger.info("Download cancelled for model %s", model_size)
|
|
_remove_dir(target)
|
|
return None
|
|
except Exception:
|
|
logger.warning("Failed to download model %s", model_size)
|
|
_remove_dir(target)
|
|
return None
|
|
return target
|
|
|
|
|
|
class _DownloadCancelledError(Exception):
|
|
"""Internal signal raised to abort a model download."""
|
|
|
|
|
|
def _remove_dir(path: Path) -> None:
|
|
import shutil
|
|
|
|
try:
|
|
if path.is_dir():
|
|
shutil.rmtree(path, ignore_errors=True)
|
|
except OSError:
|
|
logger.warning("Failed to clean up partial download %s", path)
|
|
|
|
|
|
def delete_model(model_size: str) -> bool:
|
|
"""Remove a model's cache snapshot from disk. Returns True if something was removed."""
|
|
try:
|
|
d = model_cache_dir(model_size)
|
|
except ValueError:
|
|
return False
|
|
if not d.is_dir():
|
|
return False
|
|
try:
|
|
import shutil
|
|
|
|
shutil.rmtree(d, ignore_errors=True)
|
|
except OSError:
|
|
logger.warning("Failed to delete model cache %s", d)
|
|
return False
|
|
return not d.exists()
|
|
|
|
|
|
def _load_selection() -> str:
|
|
"""The persisted selected model size (or ``""`` when absent/invalid)."""
|
|
return str(_load_config().get("selected_model", ""))
|
|
|
|
|
|
def save_selected_model(model_size: str) -> str:
|
|
"""Persist the selected model size. Validates against ``ALLOWED_MODELS``.
|
|
|
|
Returns the persisted value so callers can confirm it round-tripped.
|
|
"""
|
|
if model_size not in ALLOWED_MODELS:
|
|
raise ValueError(
|
|
f"model_size must be one of {', '.join(ALLOWED_MODELS)}, got {model_size!r}"
|
|
)
|
|
data = _load_config()
|
|
data["selected_model"] = model_size
|
|
_write_config(data)
|
|
return model_size
|
|
|
|
|
|
def load_selected_model() -> str:
|
|
"""The effective selected model, with a safe fallback.
|
|
|
|
- Returns the persisted selection when it's in ``ALLOWED_MODELS``.
|
|
- If the persisted model isn't on disk but another is installed, returns that
|
|
installed one (never clears the user's persisted value on a false-negative
|
|
availability scan, mirroring Hex's rule).
|
|
- Otherwise returns ``""`` so callers can fall back to ``"base"``.
|
|
"""
|
|
selected = _load_selection()
|
|
if selected and selected in ALLOWED_MODELS and is_model_downloaded(selected):
|
|
return selected
|
|
installed = list_installed_models()
|
|
if installed:
|
|
return installed[0]
|
|
return ""
|
|
|
|
|
|
def load_hf_token() -> str:
|
|
"""The persisted HuggingFace token for diarization (or ``""``)."""
|
|
return str(_load_config().get("hf_token", ""))
|
|
|
|
|
|
def save_hf_token(token: str) -> str:
|
|
"""Persist the HuggingFace token used for speaker diarization."""
|
|
data = _load_config()
|
|
data["hf_token"] = str(token or "").strip()
|
|
_write_config(data)
|
|
return data["hf_token"]
|
|
|
|
|
|
# Transcription languages, matching the codes used across the app UI.
|
|
# ``""`` / ``"auto"`` means "detect automatically".
|
|
ALLOWED_LANGUAGES = {
|
|
"auto",
|
|
"pt",
|
|
"en",
|
|
"es",
|
|
"fr",
|
|
"de",
|
|
"it",
|
|
"nl",
|
|
"ja",
|
|
"ko",
|
|
"zh",
|
|
}
|
|
|
|
|
|
def load_transcript_language() -> str:
|
|
"""The persisted transcription language (``""``/``"auto"`` = auto-detect)."""
|
|
lang = str(_load_config().get("language", ""))
|
|
return lang if lang in ALLOWED_LANGUAGES else "auto"
|
|
|
|
|
|
def save_transcript_language(lang: str) -> str:
|
|
"""Persist the transcription language. Returns the stored value."""
|
|
val = str(lang or "auto").strip().lower()
|
|
if val not in ALLOWED_LANGUAGES:
|
|
raise ValueError(
|
|
f"language must be one of {', '.join(sorted(ALLOWED_LANGUAGES))}, got {lang!r}"
|
|
)
|
|
data = _load_config()
|
|
data["language"] = val
|
|
_write_config(data)
|
|
return val
|
|
|
|
|
|
def load_num_speakers() -> str:
|
|
"""The persisted expected participant count (``""`` = auto-detect)."""
|
|
return str(_load_config().get("num_speakers", ""))
|
|
|
|
|
|
def save_num_speakers(num: str) -> str:
|
|
"""Persist the expected participant count (empty string = auto)."""
|
|
val = str(num or "").strip()
|
|
data = _load_config()
|
|
data["num_speakers"] = val
|
|
_write_config(data)
|
|
return val
|