chore: adiciona .gitignore e commit.command
This commit is contained in:
@@ -0,0 +1,361 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user