282 lines
8.7 KiB
Python
282 lines
8.7 KiB
Python
#!/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())
|