412 lines
16 KiB
Python
412 lines
16 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import importlib.util
|
|
import math
|
|
import struct
|
|
import sys
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
|
|
SCRIPT_PATH = Path(__file__).with_name("local-whisper-stt-server.py")
|
|
SPEC = importlib.util.spec_from_file_location("local_whisper_stt_server", 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)
|
|
|
|
|
|
SAMPLE_RATE = 16_000
|
|
|
|
|
|
def pcm_ms(milliseconds: float, *, amplitude: int) -> bytes:
|
|
"""Mono linear16 tone (or silence when amplitude is 0)."""
|
|
|
|
count = int(SAMPLE_RATE * milliseconds / 1000.0)
|
|
if amplitude <= 0:
|
|
return b"\x00\x00" * count
|
|
samples = [
|
|
int(amplitude * math.sin(2 * math.pi * 220 * index / SAMPLE_RATE))
|
|
for index in range(count)
|
|
]
|
|
return struct.pack(f"<{count}h", *samples)
|
|
|
|
|
|
SPEECH = lambda ms: pcm_ms(ms, amplitude=9000) # noqa: E731
|
|
SILENCE = lambda ms: pcm_ms(ms, amplitude=0) # noqa: E731
|
|
|
|
|
|
class FakeTranscriber:
|
|
def __init__(self, text: str = "안녕하세요") -> None:
|
|
self.text = text
|
|
self.calls: list[tuple[int, bool]] = []
|
|
|
|
def transcribe(self, pcm: bytes, *, sample_rate: int, final: bool):
|
|
self.calls.append((len(pcm), final))
|
|
words = (
|
|
(MODULE.Word(word=self.text, start=0.0, end=0.5, confidence=0.9),)
|
|
if final
|
|
else ()
|
|
)
|
|
return MODULE.Transcription(text=self.text, confidence=0.87, words=words)
|
|
|
|
|
|
def default_options(**overrides):
|
|
options = {
|
|
"model": MODULE.DEFAULT_MODEL,
|
|
"language": "ko",
|
|
"sample_rate": SAMPLE_RATE,
|
|
"channels": 1,
|
|
"endpointing_ms": MODULE.DEFAULT_ENDPOINTING_MS,
|
|
"utterance_end_ms": MODULE.DEFAULT_UTTERANCE_END_MS,
|
|
"interim_interval_ms": MODULE.DEFAULT_INTERIM_INTERVAL_MS,
|
|
"silence_rms": MODULE.DEFAULT_SILENCE_RMS,
|
|
}
|
|
options.update(overrides)
|
|
return options
|
|
|
|
|
|
def run(coro):
|
|
return asyncio.run(coro)
|
|
|
|
|
|
class FrameRmsTest(unittest.TestCase):
|
|
def test_silence_and_speech_are_separated(self) -> None:
|
|
self.assertEqual(MODULE.frame_rms(SILENCE(20)), 0)
|
|
self.assertGreater(MODULE.frame_rms(SPEECH(20)), MODULE.DEFAULT_SILENCE_RMS)
|
|
|
|
def test_empty_and_odd_length_input_is_silence(self) -> None:
|
|
self.assertEqual(MODULE.frame_rms(b""), 0)
|
|
self.assertEqual(MODULE.frame_rms(b"\x01"), 0)
|
|
|
|
|
|
class UtteranceBufferTest(unittest.TestCase):
|
|
def test_duration_tracks_appended_bytes(self) -> None:
|
|
buffer = MODULE.UtteranceBuffer()
|
|
buffer.append(SPEECH(500))
|
|
self.assertAlmostEqual(buffer.duration_ms(), 500.0, delta=1.0)
|
|
|
|
def test_silence_alone_never_finalizes(self) -> None:
|
|
buffer = MODULE.UtteranceBuffer()
|
|
buffer.append(SILENCE(3_000))
|
|
self.assertFalse(buffer.should_finalize())
|
|
self.assertFalse(buffer.should_emit_interim())
|
|
self.assertTrue(buffer.is_stalled())
|
|
|
|
def test_speech_then_endpointing_silence_finalizes(self) -> None:
|
|
buffer = MODULE.UtteranceBuffer(endpointing_ms=300)
|
|
buffer.append(SPEECH(800))
|
|
self.assertFalse(buffer.should_finalize())
|
|
buffer.append(SILENCE(200))
|
|
self.assertFalse(buffer.should_finalize())
|
|
buffer.append(SILENCE(200))
|
|
self.assertTrue(buffer.should_finalize())
|
|
|
|
def test_speech_after_silence_resets_the_silence_run(self) -> None:
|
|
buffer = MODULE.UtteranceBuffer(endpointing_ms=300)
|
|
buffer.append(SPEECH(400))
|
|
buffer.append(SILENCE(200))
|
|
buffer.append(SPEECH(100))
|
|
buffer.append(SILENCE(200))
|
|
self.assertFalse(buffer.should_finalize())
|
|
|
|
def test_interim_waits_for_the_configured_interval(self) -> None:
|
|
buffer = MODULE.UtteranceBuffer(interim_interval_ms=700)
|
|
buffer.append(SPEECH(300))
|
|
self.assertFalse(buffer.should_emit_interim())
|
|
buffer.append(SPEECH(500))
|
|
self.assertTrue(buffer.should_emit_interim())
|
|
buffer.mark_interim_emitted()
|
|
self.assertFalse(buffer.should_emit_interim())
|
|
|
|
def test_very_long_utterance_finalizes_even_without_silence(self) -> None:
|
|
buffer = MODULE.UtteranceBuffer()
|
|
for _ in range(int(MODULE.MAX_UTTERANCE_SECONDS) + 1):
|
|
buffer.append(SPEECH(1_000))
|
|
self.assertTrue(buffer.should_finalize())
|
|
|
|
def test_reset_clears_audio_and_state(self) -> None:
|
|
buffer = MODULE.UtteranceBuffer()
|
|
buffer.append(SPEECH(500))
|
|
buffer.reset()
|
|
self.assertEqual(buffer.pcm, b"")
|
|
self.assertFalse(buffer.saw_speech)
|
|
self.assertEqual(buffer.duration_ms(), 0.0)
|
|
|
|
|
|
class StreamOptionsTest(unittest.TestCase):
|
|
def test_defaults_apply_without_a_query(self) -> None:
|
|
options = MODULE.parse_stream_options("/v1/listen")
|
|
self.assertEqual(options["model"], MODULE.DEFAULT_MODEL)
|
|
self.assertEqual(options["language"], MODULE.DEFAULT_LANGUAGE)
|
|
self.assertEqual(options["sample_rate"], MODULE.DEFAULT_SAMPLE_RATE)
|
|
|
|
def test_values_are_clamped_into_range(self) -> None:
|
|
options = MODULE.parse_stream_options(
|
|
"/v1/listen?sample_rate=999999&endpointing=1&utterance_end_ms=99999"
|
|
)
|
|
self.assertEqual(options["sample_rate"], 48_000)
|
|
self.assertEqual(options["endpointing_ms"], 10)
|
|
self.assertEqual(options["utterance_end_ms"], 10_000)
|
|
|
|
def test_unparsable_numbers_fall_back_to_defaults(self) -> None:
|
|
options = MODULE.parse_stream_options("/v1/listen?sample_rate=abc")
|
|
self.assertEqual(options["sample_rate"], MODULE.DEFAULT_SAMPLE_RATE)
|
|
|
|
def test_model_allowlist_and_language_shape_fail_closed(self) -> None:
|
|
with self.assertRaises(ValueError):
|
|
MODULE.parse_stream_options("/v1/listen?model=../etc/passwd")
|
|
with self.assertRaises(ValueError):
|
|
MODULE.parse_stream_options("/v1/listen?language=ko;DROP")
|
|
|
|
|
|
class TranscriptPayloadTest(unittest.TestCase):
|
|
def test_interim_payload_omits_word_timestamps(self) -> None:
|
|
result = MODULE.Transcription(
|
|
text="안녕", confidence=0.5, words=(MODULE.Word("안녕", 0.0, 0.4),)
|
|
)
|
|
payload = MODULE.transcript_payload(
|
|
result, is_final=False, speech_final=False, duration_seconds=1.0
|
|
)
|
|
self.assertEqual(payload["words"], [])
|
|
self.assertFalse(payload["is_final"])
|
|
|
|
def test_final_payload_carries_word_timestamps(self) -> None:
|
|
result = MODULE.Transcription(
|
|
text="안녕", confidence=None, words=(MODULE.Word("안녕", 0.0, 0.4, 0.8),)
|
|
)
|
|
payload = MODULE.transcript_payload(
|
|
result, is_final=True, speech_final=True, duration_seconds=1.234
|
|
)
|
|
self.assertEqual(payload["words"][0]["word"], "안녕")
|
|
self.assertEqual(payload["words"][0]["confidence"], 0.8)
|
|
self.assertIsNone(payload["confidence"])
|
|
self.assertEqual(payload["duration"], 1.234)
|
|
|
|
|
|
class StreamSessionTest(unittest.TestCase):
|
|
def _session(self, transcriber, sent, **overrides):
|
|
async def send(payload):
|
|
sent.append(payload)
|
|
|
|
return MODULE.StreamSession(
|
|
options=default_options(**overrides),
|
|
transcriber=transcriber,
|
|
send=send,
|
|
)
|
|
|
|
def test_silence_only_stream_emits_nothing(self) -> None:
|
|
sent: list[dict] = []
|
|
transcriber = FakeTranscriber()
|
|
session = self._session(transcriber, sent)
|
|
run(session.feed(SILENCE(2_000)))
|
|
run(session.finalize(speech_final=True))
|
|
self.assertEqual(sent, [])
|
|
self.assertEqual(transcriber.calls, [])
|
|
|
|
def test_speech_emits_interim_then_final(self) -> None:
|
|
sent: list[dict] = []
|
|
session = self._session(FakeTranscriber(), sent)
|
|
run(session.feed(SPEECH(800)))
|
|
self.assertEqual(len(sent), 1)
|
|
self.assertFalse(sent[0]["is_final"])
|
|
run(session.feed(SILENCE(400)))
|
|
self.assertEqual(len(sent), 2)
|
|
self.assertTrue(sent[1]["is_final"])
|
|
self.assertTrue(sent[1]["speech_final"])
|
|
self.assertTrue(sent[1]["words"])
|
|
|
|
def test_final_appends_to_the_transcript_and_resets_the_buffer(self) -> None:
|
|
sent: list[dict] = []
|
|
session = self._session(FakeTranscriber("첫 문장"), sent)
|
|
run(session.feed(SPEECH(800)))
|
|
run(session.feed(SILENCE(400)))
|
|
self.assertEqual(session.final_segments, ["첫 문장"])
|
|
self.assertEqual(session.buffer.pcm, b"")
|
|
self.assertFalse(session.buffer.saw_speech)
|
|
|
|
def test_two_utterances_produce_two_finals(self) -> None:
|
|
sent: list[dict] = []
|
|
session = self._session(FakeTranscriber(), sent)
|
|
for _ in range(2):
|
|
run(session.feed(SPEECH(800)))
|
|
run(session.feed(SILENCE(400)))
|
|
finals = [item for item in sent if item["is_final"]]
|
|
self.assertEqual(len(finals), 2)
|
|
self.assertEqual(len(session.final_segments), 2)
|
|
|
|
def test_explicit_finalize_flushes_without_trailing_silence(self) -> None:
|
|
sent: list[dict] = []
|
|
session = self._session(FakeTranscriber(), sent)
|
|
run(session.feed(SPEECH(400)))
|
|
run(session.finalize(speech_final=True))
|
|
self.assertTrue(sent[-1]["is_final"])
|
|
|
|
def test_blank_model_output_is_not_emitted(self) -> None:
|
|
sent: list[dict] = []
|
|
session = self._session(FakeTranscriber(text=" "), sent)
|
|
run(session.feed(SPEECH(800)))
|
|
run(session.feed(SILENCE(400)))
|
|
self.assertEqual(sent, [])
|
|
self.assertEqual(session.final_segments, [])
|
|
|
|
def test_oversized_frame_fails_closed(self) -> None:
|
|
sent: list[dict] = []
|
|
session = self._session(FakeTranscriber(), sent)
|
|
with self.assertRaises(MODULE.TranscriptionError):
|
|
run(session.feed(b"\x00" * (MODULE.MAX_FRAME_BYTES + 1)))
|
|
|
|
def test_final_transcription_requests_word_timestamps(self) -> None:
|
|
sent: list[dict] = []
|
|
transcriber = FakeTranscriber()
|
|
session = self._session(transcriber, sent)
|
|
run(session.feed(SPEECH(800)))
|
|
run(session.feed(SILENCE(400)))
|
|
self.assertIn(False, [final for _, final in transcriber.calls])
|
|
self.assertIn(True, [final for _, final in transcriber.calls])
|
|
|
|
|
|
class DeviceResolutionTest(unittest.TestCase):
|
|
def test_cpu_request_uses_int8(self) -> None:
|
|
self.assertEqual(MODULE.resolve_device("cpu"), ("cpu", "int8"))
|
|
|
|
def test_invalid_device_fails_closed(self) -> None:
|
|
with self.assertRaises(ValueError):
|
|
MODULE.resolve_device("tpu")
|
|
|
|
def test_auto_returns_a_supported_pair(self) -> None:
|
|
device, compute = MODULE.resolve_device("auto")
|
|
self.assertIn(device, {"cuda", "cpu"})
|
|
self.assertIn(compute, {"float16", "int8"})
|
|
|
|
def test_auto_prefers_cuda_when_present(self) -> None:
|
|
self.assertEqual(
|
|
MODULE.resolve_device("auto", cuda_available=True), ("cuda", "float16")
|
|
)
|
|
|
|
def test_explicit_cuda_without_a_device_fails(self) -> None:
|
|
with self.assertRaises(ValueError):
|
|
MODULE.resolve_device("cuda", cuda_available=False)
|
|
|
|
|
|
class WarmupPcmTest(unittest.TestCase):
|
|
def test_warmup_buffer_is_bytes_of_the_expected_length(self) -> None:
|
|
pcm = MODULE.warmup_pcm(16_000, milliseconds=500)
|
|
self.assertIsInstance(pcm, bytes)
|
|
self.assertEqual(len(pcm), 16_000 // 2 * 2)
|
|
self.assertEqual(MODULE.frame_rms(pcm), 0)
|
|
|
|
def test_warmup_buffer_is_never_empty(self) -> None:
|
|
self.assertGreaterEqual(len(MODULE.warmup_pcm(8_000, milliseconds=0)), 2)
|
|
|
|
|
|
class _FakeTranscriberFactory:
|
|
"""Records constructions and can fail warmup on a chosen device."""
|
|
|
|
def __init__(self, *, failing_device: str | None = None) -> None:
|
|
self.failing_device = failing_device
|
|
self.built: list[tuple[str, str, str]] = []
|
|
|
|
def __call__(self, model: str, device: str, compute_type: str):
|
|
self.built.append((model, device, compute_type))
|
|
failing = self.failing_device
|
|
|
|
class _Transcriber:
|
|
def __init__(self) -> None:
|
|
self.device = device
|
|
self.compute_type = compute_type
|
|
self.model_name = model
|
|
|
|
def warmup(self) -> None:
|
|
if failing is not None and device == failing:
|
|
raise RuntimeError("Could not locate cudnn_ops64_9.dll")
|
|
|
|
return _Transcriber()
|
|
|
|
|
|
class TranscriberBuildTest(unittest.TestCase):
|
|
def test_auto_falls_back_to_cpu_when_the_cuda_probe_fails(self) -> None:
|
|
"""cuDNN 부재는 네이티브 크래시라 자식 프로세스 프로브로만 잡힌다."""
|
|
|
|
factory = _FakeTranscriberFactory()
|
|
probes: list[tuple[str, str]] = []
|
|
|
|
def prober(model: str, device: str) -> bool:
|
|
probes.append((model, device))
|
|
return False
|
|
|
|
transcriber = MODULE.build_transcriber(
|
|
"small", "auto", factory=factory, cuda_available=True, prober=prober
|
|
)
|
|
self.assertEqual(probes, [("small", "cuda")])
|
|
self.assertEqual(transcriber.device, "cpu")
|
|
self.assertEqual(transcriber.compute_type, "int8")
|
|
# CUDA 로는 아예 모델을 올리지 않는다. 올렸다면 그 자리에서 죽는다.
|
|
self.assertEqual(factory.built, [("small", "cpu", "int8")])
|
|
|
|
def test_auto_keeps_cuda_when_the_probe_succeeds(self) -> None:
|
|
factory = _FakeTranscriberFactory()
|
|
transcriber = MODULE.build_transcriber(
|
|
"small",
|
|
"auto",
|
|
factory=factory,
|
|
cuda_available=True,
|
|
prober=lambda model, device: True,
|
|
)
|
|
self.assertEqual(transcriber.device, "cuda")
|
|
self.assertEqual(factory.built, [("small", "cuda", "float16")])
|
|
|
|
def test_explicit_cuda_never_downgrades_silently(self) -> None:
|
|
factory = _FakeTranscriberFactory()
|
|
with self.assertRaises(MODULE.TranscriptionError):
|
|
MODULE.build_transcriber(
|
|
"small",
|
|
"cuda",
|
|
factory=factory,
|
|
cuda_available=True,
|
|
prober=lambda model, device: False,
|
|
)
|
|
self.assertEqual(factory.built, [])
|
|
|
|
def test_cpu_path_is_not_probed(self) -> None:
|
|
factory = _FakeTranscriberFactory()
|
|
probes: list[tuple[str, str]] = []
|
|
MODULE.build_transcriber(
|
|
"small",
|
|
"cpu",
|
|
factory=factory,
|
|
cuda_available=True,
|
|
prober=lambda model, device: probes.append((model, device)) or True,
|
|
)
|
|
self.assertEqual(probes, [])
|
|
self.assertEqual(factory.built, [("small", "cpu", "int8")])
|
|
|
|
def test_python_level_warmup_failure_still_fails_closed(self) -> None:
|
|
factory = _FakeTranscriberFactory(failing_device="cpu")
|
|
with self.assertRaises(MODULE.TranscriptionError):
|
|
MODULE.build_transcriber(
|
|
"small",
|
|
"auto",
|
|
factory=factory,
|
|
cuda_available=False,
|
|
prober=lambda model, device: True,
|
|
)
|
|
|
|
|
|
class CliTest(unittest.TestCase):
|
|
def test_server_refuses_to_start_without_enable(self) -> None:
|
|
self.assertEqual(MODULE.main([]), 2)
|
|
|
|
def test_host_is_loopback_only(self) -> None:
|
|
parser = MODULE.build_parser()
|
|
with self.assertRaises(SystemExit):
|
|
parser.parse_args(["--host", "0.0.0.0"])
|
|
|
|
def test_default_model_matches_the_cpu_public_runtime_contract(self) -> None:
|
|
parser = MODULE.build_parser()
|
|
self.assertEqual(parser.parse_args([]).model, "small")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|