"""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-. 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`` -> ``/models--Systran--faster-whisper-``. """ 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