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": "large-v3", "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_cuda_warmup_fails(self) -> None: """cuDNN 부재처럼 '장치는 보이지만 추론이 죽는' 경우를 잡는다.""" factory = _FakeTranscriberFactory(failing_device="cuda") transcriber = MODULE.build_transcriber( "small", "auto", factory=factory, cuda_available=True ) self.assertEqual(transcriber.device, "cpu") self.assertEqual(transcriber.compute_type, "int8") self.assertEqual( factory.built, [("small", "cuda", "float16"), ("small", "cpu", "int8")] ) def test_auto_keeps_cuda_when_warmup_succeeds(self) -> None: factory = _FakeTranscriberFactory() transcriber = MODULE.build_transcriber( "small", "auto", factory=factory, cuda_available=True ) self.assertEqual(transcriber.device, "cuda") self.assertEqual(len(factory.built), 1) def test_explicit_cuda_never_downgrades_silently(self) -> None: factory = _FakeTranscriberFactory(failing_device="cuda") with self.assertRaises(MODULE.TranscriptionError): MODULE.build_transcriber( "small", "cuda", factory=factory, cuda_available=True ) self.assertEqual(factory.built, [("small", "cuda", "float16")]) def test_cpu_warmup_failure_is_not_retried(self) -> None: factory = _FakeTranscriberFactory(failing_device="cpu") with self.assertRaises(MODULE.TranscriptionError): MODULE.build_transcriber( "small", "auto", factory=factory, cuda_available=False ) self.assertEqual(factory.built, [("small", "cpu", "int8")]) 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"]) if __name__ == "__main__": unittest.main()