Files
gart/code/fcpxml/model_manager.py
T

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