#!/usr/bin/env python3 """Fail-closed, metadata-only readiness probe for the public voice sidecars. The probe never sends audio or synthesis text. It proves the local Whisper WebSocket protocol by validating its first ``ready`` frame, and proves MeloTTS through its loopback ``/health`` document. Only bounded, non-PII metadata is printed. """ from __future__ import annotations import argparse import asyncio import json import sys from collections.abc import Iterable, Mapping from typing import Any from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit from urllib.request import ProxyHandler, Request, build_opener DEFAULT_STT_URL = "ws://127.0.0.1:9882/v1/listen" DEFAULT_TTS_URL = "http://127.0.0.1:9883" DEFAULT_STT_PROVIDER = "local_whisper" DEFAULT_STT_MODEL = "small" DEFAULT_STT_LANGUAGE = "ko" DEFAULT_STT_DEVICE = "cpu" DEFAULT_TTS_PROVIDER = "melotts" DEFAULT_TTS_MODEL = "melotts-korean" DEFAULT_TTS_LANGUAGE = "KR" MAX_METADATA_BYTES = 64 * 1024 LOOPBACK_HOSTS = frozenset({"127.0.0.1", "localhost", "::1"}) class ProbeError(RuntimeError): """A stable, non-sensitive readiness failure code.""" def _loopback_parts(url: str, *, schemes: set[str]) -> Any: parts = urlsplit(url) if ( parts.scheme not in schemes or parts.hostname not in LOOPBACK_HOSTS or parts.username is not None or parts.password is not None or parts.fragment ): raise ProbeError("loopback_url_required") return parts def build_stt_probe_url( base_url: str, *, model: str, language: str, ) -> str: parts = _loopback_parts(base_url, schemes={"ws"}) query = dict(parse_qsl(parts.query, keep_blank_values=True)) query.update( { "model": model, "language": language, "sample_rate": "16000", "channels": "1", } ) return urlunsplit( (parts.scheme, parts.netloc, parts.path, urlencode(query), "") ) def validate_stt_ready( payload: object, *, provider: str, model: str, language: str, device: str, ) -> dict[str, object]: if not isinstance(payload, Mapping): raise ProbeError("stt_ready_not_object") expected = { "type": "ready", "provider": provider, "model": model, "language": language, "device": device, "sample_rate": 16_000, } for key, value in expected.items(): if payload.get(key) != value: raise ProbeError(f"stt_{key}_mismatch") expected_compute_type = "int8" if device == "cpu" else "float16" if payload.get("compute_type") != expected_compute_type: raise ProbeError("stt_compute_type_mismatch") return { "provider": provider, "model": model, "language": language, "device": device, "compute_type": expected_compute_type, "sample_rate": 16_000, } def validate_tts_health( payload: object, *, provider: str, model: str, language: str, ) -> dict[str, object]: if not isinstance(payload, Mapping): raise ProbeError("tts_health_not_object") expected = { "status": "ok", "provider": provider, "model": model, "language": language, "license": "MIT", "reference_policy": "pretrained-multispeaker-no-external-reference", } for key, value in expected.items(): if payload.get(key) != value: raise ProbeError(f"tts_{key}_mismatch") sample_rate = payload.get("sample_rate") speakers = payload.get("speakers") if not isinstance(sample_rate, int) or sample_rate <= 0: raise ProbeError("tts_sample_rate_invalid") if not isinstance(speakers, list) or not speakers: raise ProbeError("tts_speakers_missing") return { "provider": provider, "model": model, "language": language, "license": "MIT", "sample_rate": sample_rate, } async def probe_stt( url: str, *, provider: str, model: str, language: str, device: str, timeout_seconds: float, ) -> dict[str, object]: from websockets.asyncio.client import connect probe_url = build_stt_probe_url(url, model=model, language=language) try: async with asyncio.timeout(timeout_seconds): async with connect( probe_url, max_size=MAX_METADATA_BYTES, max_queue=1, ping_interval=None, ) as socket: frame = await socket.recv() except ProbeError: raise except Exception as exc: raise ProbeError("stt_connection_failed") from exc if not isinstance(frame, str) or len(frame.encode("utf-8")) > MAX_METADATA_BYTES: raise ProbeError("stt_ready_frame_invalid") try: payload = json.loads(frame) except ValueError as exc: raise ProbeError("stt_ready_json_invalid") from exc return validate_stt_ready( payload, provider=provider, model=model, language=language, device=device, ) def probe_tts( url: str, *, provider: str, model: str, language: str, timeout_seconds: float, ) -> dict[str, object]: parts = _loopback_parts(url, schemes={"http"}) health_path = parts.path.rstrip("/") + "/health" health_url = urlunsplit((parts.scheme, parts.netloc, health_path, "", "")) request = Request( health_url, method="GET", headers={"Accept": "application/json", "User-Agent": "vignette-readiness/1"}, ) try: with build_opener(ProxyHandler({})).open( request, timeout=timeout_seconds ) as response: content_type = response.headers.get_content_type() body = response.read(MAX_METADATA_BYTES + 1) status = response.status except Exception as exc: raise ProbeError("tts_connection_failed") from exc if status != 200: raise ProbeError("tts_health_status_invalid") if content_type != "application/json" or len(body) > MAX_METADATA_BYTES: raise ProbeError("tts_health_response_invalid") try: payload = json.loads(body) except (UnicodeDecodeError, ValueError) as exc: raise ProbeError("tts_health_json_invalid") from exc return validate_tts_health( payload, provider=provider, model=model, language=language, ) def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--component", choices=("all", "stt", "tts"), default="all") parser.add_argument("--stt-url", default=DEFAULT_STT_URL) parser.add_argument("--stt-provider", default=DEFAULT_STT_PROVIDER) parser.add_argument("--stt-model", default=DEFAULT_STT_MODEL) parser.add_argument("--stt-language", default=DEFAULT_STT_LANGUAGE) parser.add_argument("--stt-device", choices=("cpu", "cuda"), default=DEFAULT_STT_DEVICE) parser.add_argument("--tts-url", default=DEFAULT_TTS_URL) parser.add_argument("--tts-provider", default=DEFAULT_TTS_PROVIDER) parser.add_argument("--tts-model", default=DEFAULT_TTS_MODEL) parser.add_argument("--tts-language", default=DEFAULT_TTS_LANGUAGE) parser.add_argument("--timeout-seconds", type=float, default=5.0) return parser def main(argv: Iterable[str] | None = None) -> int: args = build_parser().parse_args(list(argv) if argv is not None else None) if not 0.1 <= args.timeout_seconds <= 30: print('{"ok":false,"error":"timeout_invalid"}', file=sys.stderr) return 2 result: dict[str, object] = { "schema_version": "vignette.voice-sidecar-readiness.v1", "ok": True, } try: if args.component in {"all", "stt"}: result["stt"] = asyncio.run( probe_stt( args.stt_url, provider=args.stt_provider, model=args.stt_model, language=args.stt_language, device=args.stt_device, timeout_seconds=args.timeout_seconds, ) ) if args.component in {"all", "tts"}: result["tts"] = probe_tts( args.tts_url, provider=args.tts_provider, model=args.tts_model, language=args.tts_language, timeout_seconds=args.timeout_seconds, ) except ProbeError as exc: print( json.dumps( {"ok": False, "component": args.component, "error": str(exc)}, separators=(",", ":"), ), file=sys.stderr, ) return 1 print(json.dumps(result, separators=(",", ":"), sort_keys=True)) return 0 if __name__ == "__main__": raise SystemExit(main())