144 lines
5.2 KiB
Python
144 lines
5.2 KiB
Python
from __future__ import annotations
|
|
|
|
import importlib.util
|
|
import sys
|
|
import unittest
|
|
from pathlib import Path
|
|
from urllib.parse import parse_qs, urlsplit
|
|
|
|
|
|
SCRIPT_PATH = Path(__file__).with_name("probe-public-voice-sidecars.py")
|
|
SPEC = importlib.util.spec_from_file_location("public_voice_sidecar_probe", SCRIPT_PATH)
|
|
assert SPEC is not None and SPEC.loader is not None
|
|
MODULE = importlib.util.module_from_spec(SPEC)
|
|
sys.modules[SPEC.name] = MODULE
|
|
SPEC.loader.exec_module(MODULE)
|
|
|
|
|
|
class SttProbeContractTest(unittest.TestCase):
|
|
def _payload(self, **overrides: object) -> dict[str, object]:
|
|
payload: dict[str, object] = {
|
|
"type": "ready",
|
|
"provider": "local_whisper",
|
|
"model": "small",
|
|
"language": "ko",
|
|
"device": "cpu",
|
|
"compute_type": "int8",
|
|
"sample_rate": 16_000,
|
|
}
|
|
payload.update(overrides)
|
|
return payload
|
|
|
|
def test_exact_public_metadata_is_accepted(self) -> None:
|
|
result = MODULE.validate_stt_ready(
|
|
self._payload(),
|
|
provider="local_whisper",
|
|
model="small",
|
|
language="ko",
|
|
device="cpu",
|
|
)
|
|
self.assertEqual(result["compute_type"], "int8")
|
|
|
|
def test_provider_model_device_and_compute_type_mismatch_fail_closed(self) -> None:
|
|
for field, value in (
|
|
("provider", "deepgram"),
|
|
("model", "large-v3"),
|
|
("device", "cuda"),
|
|
("compute_type", "float16"),
|
|
):
|
|
with self.subTest(field=field):
|
|
with self.assertRaises(MODULE.ProbeError):
|
|
MODULE.validate_stt_ready(
|
|
self._payload(**{field: value}),
|
|
provider="local_whisper",
|
|
model="small",
|
|
language="ko",
|
|
device="cpu",
|
|
)
|
|
|
|
def test_probe_url_replaces_query_with_exact_stream_contract(self) -> None:
|
|
url = MODULE.build_stt_probe_url(
|
|
"ws://127.0.0.1:9882/v1/listen?model=base&channels=2",
|
|
model="small",
|
|
language="ko",
|
|
)
|
|
query = parse_qs(urlsplit(url).query)
|
|
self.assertEqual(query["model"], ["small"])
|
|
self.assertEqual(query["language"], ["ko"])
|
|
self.assertEqual(query["sample_rate"], ["16000"])
|
|
self.assertEqual(query["channels"], ["1"])
|
|
|
|
def test_non_loopback_urls_and_embedded_credentials_are_rejected(self) -> None:
|
|
for url in (
|
|
"wss://example.com/v1/listen",
|
|
"ws://user:secret@127.0.0.1:9882/v1/listen",
|
|
):
|
|
with self.subTest(url=url):
|
|
with self.assertRaises(MODULE.ProbeError):
|
|
MODULE.build_stt_probe_url(url, model="small", language="ko")
|
|
|
|
|
|
class TtsProbeContractTest(unittest.TestCase):
|
|
def _payload(self, **overrides: object) -> dict[str, object]:
|
|
payload: dict[str, object] = {
|
|
"status": "ok",
|
|
"provider": "melotts",
|
|
"model": "melotts-korean",
|
|
"language": "KR",
|
|
"license": "MIT",
|
|
"reference_policy": "pretrained-multispeaker-no-external-reference",
|
|
"sample_rate": 44_100,
|
|
"speakers": ["KR"],
|
|
}
|
|
payload.update(overrides)
|
|
return payload
|
|
|
|
def test_exact_public_metadata_is_accepted(self) -> None:
|
|
result = MODULE.validate_tts_health(
|
|
self._payload(),
|
|
provider="melotts",
|
|
model="melotts-korean",
|
|
language="KR",
|
|
)
|
|
self.assertEqual(result["license"], "MIT")
|
|
|
|
def test_provider_model_license_and_reference_mismatch_fail_closed(self) -> None:
|
|
for field, value in (
|
|
("provider", "openai"),
|
|
("model", "other"),
|
|
("license", "unknown"),
|
|
("reference_policy", "external-reference"),
|
|
):
|
|
with self.subTest(field=field):
|
|
with self.assertRaises(MODULE.ProbeError):
|
|
MODULE.validate_tts_health(
|
|
self._payload(**{field: value}),
|
|
provider="melotts",
|
|
model="melotts-korean",
|
|
language="KR",
|
|
)
|
|
|
|
def test_empty_speaker_catalog_and_invalid_sample_rate_fail_closed(self) -> None:
|
|
for override in ({"speakers": []}, {"sample_rate": 0}):
|
|
with self.subTest(override=override):
|
|
with self.assertRaises(MODULE.ProbeError):
|
|
MODULE.validate_tts_health(
|
|
self._payload(**override),
|
|
provider="melotts",
|
|
model="melotts-korean",
|
|
language="KR",
|
|
)
|
|
|
|
|
|
class CliContractTest(unittest.TestCase):
|
|
def test_defaults_are_the_public_local_stack(self) -> None:
|
|
args = MODULE.build_parser().parse_args([])
|
|
self.assertEqual(args.stt_provider, "local_whisper")
|
|
self.assertEqual(args.stt_model, "small")
|
|
self.assertEqual(args.stt_device, "cpu")
|
|
self.assertEqual(args.tts_provider, "melotts")
|
|
self.assertEqual(args.tts_model, "melotts-korean")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|