"""Deterministic WebSocket contract tests for the voice gateway route.""" from __future__ import annotations import json import unittest from types import SimpleNamespace from unittest.mock import AsyncMock, patch from fastapi import HTTPException from .deps import Principal, Role from .persona_repository import PersonaVoiceMap from .routes import voice as voice_routes from .services.voice import TTSChunk, VoicePreset from .store import TurnRecord SESSION_ID = "voice-ws-contract-session" VOICE_PRESET = VoicePreset(preset="neutral", openai_voice="sage") def _principal( role: Role = Role.LEARNER, *, consent_at: float | None = 1.0, profile_completed_at: float | None = 1.0, ) -> Principal: return Principal( user_id="00000000-0000-0000-0000-000000000201", role=role, cohort_ids=[], email=f"voice-ws-{role.value}@hs.ac.kr", display_name="Voice WS Contract", consent_at=consent_at if role == Role.LEARNER else None, profile_completed_at=profile_completed_at if role == Role.LEARNER else None, ) def _control(payload: dict[str, object]) -> dict[str, object]: return {"text": json.dumps(payload)} def _binary(data: bytes) -> dict[str, object]: return {"bytes": data} class FakeWebSocket: def __init__(self, incoming: list[dict[str, object]] | None = None) -> None: self._incoming = list(incoming or []) self.accepted = False self.client_state = voice_routes.WebSocketState.CONNECTING self.cookies: dict[str, str] = {} self.query_params: dict[str, str] = {} self.sent_json: list[dict[str, object]] = [] self.sent_text: list[str] = [] self.sent_bytes: list[bytes] = [] self.close_codes: list[int] = [] async def accept(self) -> None: self.accepted = True self.client_state = voice_routes.WebSocketState.CONNECTED async def receive(self) -> dict[str, object]: if self._incoming: return self._incoming.pop(0) return {"type": "websocket.disconnect"} async def send_text(self, data: str) -> None: self.sent_text.append(data) self.sent_json.append(json.loads(data)) async def send_bytes(self, data: bytes) -> None: self.sent_bytes.append(data) async def close(self, code: int = 1000) -> None: self.close_codes.append(code) self.client_state = voice_routes.WebSocketState.DISCONNECTED class VoiceWebSocketContractTest(unittest.IsolatedAsyncioTestCase): def _bind_result(self) -> tuple[str, VoicePreset, None, dict[str, object]]: return ( SESSION_ID, VOICE_PRESET, None, {"degraded": False, "persona_catalog_source": "session"}, ) def test_provider_events_get_internal_taxonomy_without_raw_payload(self) -> None: events = voice_routes._safe_provider_events( [ { "type": "SIGH", "confidence": 0.81, "text": "raw transcript must drop", }, { "kind": "voice_activity", "start_ms": 10, "raw_text": "drop", }, { "label": "vendor custom marker", "score": 0.44, }, ] ) self.assertEqual( events, [ { "type": "SIGH", "confidence": 0.81, "event_type": "sigh", "category": "paralinguistic", }, { "kind": "voice_activity", "start_ms": 10, "event_type": "voice_activity", "category": "speech_activity", }, { "label": "vendor custom marker", "score": 0.44, "event_type": "vendor_custom_marker", "category": "unknown", }, ], ) def test_text_tts_uses_only_the_persisted_client_visible_reply(self) -> None: session = SimpleNamespace( turns=[ TurnRecord( turn_seq=2, speaker="counselor", stage="초기", text="raw learner text", text_masked="masked learner text", ), TurnRecord( turn_seq=2, speaker="client", stage="초기", text="raw client reply", text_masked="마스킹된 내담자 응답", ), TurnRecord( turn_seq=3, speaker="client", stage="초기", text="hidden evaluator reply", text_masked="hidden evaluator reply", visible_to=("evaluator",), ), ] ) self.assertEqual( voice_routes._client_turn_text_for_speech(session, 2), "마스킹된 내담자 응답", ) self.assertIsNone(voice_routes._client_turn_text_for_speech(session, 3)) self.assertIsNone(voice_routes._client_turn_text_for_speech(session, 99)) def test_text_tts_maps_logical_turn_to_db_transcript_sequence(self) -> None: session = SimpleNamespace( turns=[ TurnRecord( turn_seq=1, speaker="counselor", stage="초기", text="첫 질문", text_masked="첫 질문", ), TurnRecord( turn_seq=2, speaker="client", stage="초기", text="첫 응답", text_masked="첫 응답", ), TurnRecord( turn_seq=3, speaker="counselor", stage="초기", text="둘째 질문", text_masked="둘째 질문", ), TurnRecord( turn_seq=4, speaker="client", stage="초기", text="둘째 응답", text_masked="둘째 응답", ), ] ) self.assertEqual( voice_routes._client_turn_text_for_speech(session, 1), "첫 응답" ) self.assertEqual( voice_routes._client_turn_text_for_speech(session, 2), "둘째 응답" ) self.assertIsNone(voice_routes._client_turn_text_for_speech(session, 3)) async def test_text_turn_speech_returns_openai_audio_for_owned_persisted_turn( self, ) -> None: session = SimpleNamespace( persona=SimpleNamespace(code="P1"), turns=[ TurnRecord( turn_seq=4, speaker="client", stage="초기", text="내담자 응답", text_masked="내담자 응답", ) ], ) synthesized: list[tuple[str, VoicePreset]] = [] async def synthesize(text: str, voice: VoicePreset): synthesized.append((text, voice)) yield TTSChunk(audio=b"mp3-a") yield TTSChunk(audio=b"mp3-b") with ( patch.object( voice_routes, "_practice_access_error", AsyncMock(return_value=None), ), patch.object( voice_routes.turn_runtime, "load_owned_session", AsyncMock(return_value=(session, None)), ), patch.object( voice_routes, "_resolve_session_voice", AsyncMock(return_value=VOICE_PRESET), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes.voice_service, "tts_provider", return_value="openai", ), patch.object( voice_routes.voice_service, "synthesize_stream", new=synthesize, ), ): response = await voice_routes.voice_speech( voice_routes.VoiceSpeechRequest(session_id=SESSION_ID, turn_seq=4), _principal(), ) self.assertEqual(response.status_code, 200) self.assertEqual(response.media_type, "audio/mpeg") self.assertEqual(response.body, b"mp3-amp3-b") self.assertEqual(response.headers["cache-control"], "no-store") self.assertEqual(response.headers["x-vignette-tts-provider"], "openai") self.assertEqual(synthesized, [("내담자 응답", VOICE_PRESET)]) async def test_audio_start_binary_chunks_audio_end_ping_close_contract( self, ) -> None: websocket = FakeWebSocket( [ _control({"type": "audio_start", "format": "webm"}), _binary(b"chunk-one"), _binary(b"chunk-two"), _control({"type": "ping"}), _control( { "type": "audio_end", "format": "webm", "silence_ms": "450", "barge_in": "true", "provider_events": [ { "type": "sigh", "confidence": 0.82, "text": "raw transcript must not persist", }, {"kind": "noise", "label": "x" * 120}, "invalid", ], } ), _control({"type": "close"}), ] ) handle_utterance = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal()), ), patch.object( voice_routes, "_bind_session", AsyncMock(return_value=self._bind_result()), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes, "_handle_utterance", handle_utterance, ), patch.object( voice_routes, "_run_turn_and_speak", AsyncMock(), ), patch.object( voice_routes.time, "monotonic", side_effect=[10.0, 12.0], ), ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] self.assertTrue(websocket.accepted) self.assertEqual(websocket.close_codes, [1000]) self.assertEqual( [ (message.get("type"), message.get("state")) for message in websocket.sent_json ], [("ready", "idle"), ("state", "listening"), ("pong", None)], ) handle_utterance.assert_awaited_once() context = handle_utterance.await_args.args[1] utterance = handle_utterance.await_args.args[2] self.assertEqual(context.session_id, SESSION_ID) self.assertEqual(context.principal.user_id, _principal().user_id) self.assertEqual(context.voice_preset, VOICE_PRESET) self.assertEqual(utterance.audio, b"chunk-onechunk-two") self.assertEqual(utterance.fmt, "webm") self.assertIsNone(utterance.sample_rate) self.assertIsNone(utterance.channels) self.assertIsNone(utterance.sample_width) self.assertEqual(utterance.audio_started_at, 10.0) self.assertEqual(utterance.audio_ended_at, 12.0) self.assertEqual(utterance.prosody.silence_ms, 450) self.assertIs(utterance.prosody.barge_in, True) self.assertEqual( utterance.prosody.provider_events, [ { "type": "sigh", "confidence": 0.82, "event_type": "sigh", "category": "paralinguistic", }, { "kind": "noise", "label": "x" * 80, "event_type": "background_noise", "category": "audio_quality", }, ], ) def test_pcm_upload_is_wrapped_as_wav_before_stt(self) -> None: audio, fmt = voice_routes._normalize_audio_upload( b"\x00\x00\xff\x7f", fmt="pcm", sample_rate=16000, channels=1, sample_width=2, ) self.assertEqual(fmt, "wav") self.assertTrue(audio.startswith(b"RIFF")) self.assertEqual(audio[8:12], b"WAVE") self.assertEqual(audio[12:16], b"fmt ") self.assertEqual(int.from_bytes(audio[24:28], "little"), 16000) self.assertEqual(int.from_bytes(audio[22:24], "little"), 1) self.assertEqual(audio[36:40], b"data") self.assertEqual(int.from_bytes(audio[40:44], "little"), 4) self.assertEqual(audio[44:], b"\x00\x00\xff\x7f") async def test_pcm_control_metadata_flows_to_handle_utterance(self) -> None: websocket = FakeWebSocket( [ _control( { "type": "audio_start", "format": "pcm", "sample_rate": 16000, "channels": 1, "sample_width": 2, } ), _binary(b"\x00\x00\xff\x7f"), _control({"type": "audio_end", "format": "pcm"}), _control({"type": "close"}), ] ) handle_utterance = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal()), ), patch.object( voice_routes, "_bind_session", AsyncMock(return_value=self._bind_result()), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes, "_handle_utterance", handle_utterance, ), ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] handle_utterance.assert_awaited_once() utterance = handle_utterance.await_args.args[2] self.assertEqual(utterance.fmt, "pcm") self.assertEqual(utterance.sample_rate, 16000) self.assertEqual(utterance.channels, 1) self.assertEqual(utterance.sample_width, 2) async def test_text_turn_strips_text_runs_turn_and_ping_close_still_work( self, ) -> None: websocket = FakeWebSocket( [ _control({"type": "text_turn", "text": " I need help practicing. "}), _control({"type": "ping"}), _control({"type": "close"}), ] ) run_turn = AsyncMock() handle_utterance = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal()), ), patch.object( voice_routes, "_bind_session", AsyncMock(return_value=self._bind_result()), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes, "_handle_utterance", handle_utterance, ), patch.object( voice_routes, "_run_turn_and_speak", run_turn, ), ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] self.assertEqual(websocket.close_codes, [1000]) self.assertEqual( [ (message.get("type"), message.get("state")) for message in websocket.sent_json ], [("ready", "idle"), ("pong", None)], ) handle_utterance.assert_not_awaited() run_turn.assert_awaited_once() context = run_turn.await_args.args[1] turn = run_turn.await_args.args[2] self.assertEqual(context.session_id, SESSION_ID) self.assertEqual(context.principal.user_id, _principal().user_id) self.assertEqual(context.voice_preset, VOICE_PRESET) self.assertEqual(turn.learner_text, "I need help practicing.") async def test_turn_persistence_failure_uses_structured_voice_error(self) -> None: websocket = FakeWebSocket( [ _control({"type": "text_turn", "text": "I need help practicing."}), _control({"type": "close"}), ] ) run_turn = AsyncMock( side_effect=HTTPException( status_code=503, detail=( "voice session turn append persistence unavailable; " "runtime fallback is disabled in prod" ), ) ) with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal()), ), patch.object( voice_routes, "_bind_session", AsyncMock(return_value=self._bind_result()), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes, "_run_turn_and_speak", run_turn, ), ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] self.assertEqual(websocket.close_codes, [1000]) self.assertEqual(websocket.sent_json[0].get("type"), "ready") self.assertEqual(websocket.sent_json[0].get("state"), "idle") self.assertEqual( websocket.sent_json[-2], { "type": "error", "code": "turn_persistence_unavailable", "detail": "voice turn persistence unavailable; retry the utterance", }, ) self.assertEqual(websocket.sent_json[-1], {"type": "state", "state": "idle"}) async def test_stt_result_waits_for_final_transcript_before_running_turn( self, ) -> None: websocket = FakeWebSocket( [ _control( { "type": "stt_result", "text": "I am still talking", "final": False, "silence_ms": 2500, } ), _control({"type": "close"}), ] ) run_turn = AsyncMock() handle_utterance = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal()), ), patch.object( voice_routes, "_bind_session", AsyncMock(return_value=self._bind_result()), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes, "_handle_utterance", handle_utterance, ), patch.object( voice_routes, "_run_turn_and_speak", run_turn, ), ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] handle_utterance.assert_not_awaited() run_turn.assert_not_awaited() self.assertIn( { "type": "eot", "ready": False, "reason": "final_transcript_pending", "silence_ms": 2500, "threshold_ms": 1200, }, websocket.sent_json, ) self.assertEqual( websocket.sent_json[-1], {"type": "state", "state": "listening"} ) async def test_stt_result_runs_turn_only_after_eot_ready(self) -> None: websocket = FakeWebSocket( [ _control({"type": "audio_start", "format": "pcm"}), _control( { "type": "stt_result", "text": " I am done now. ", "final": True, "silence_ms": 1300, "barge_in": False, "provider_events": [ { "type": "speech_final", "confidence": 0.91, "text": "raw transcript must not persist", } ], } ), _control({"type": "close"}), ] ) run_turn = AsyncMock() handle_utterance = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal()), ), patch.object( voice_routes, "_bind_session", AsyncMock(return_value=self._bind_result()), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes, "_handle_utterance", handle_utterance, ), patch.object( voice_routes, "_run_turn_and_speak", run_turn, ), patch.object( voice_routes.time, "monotonic", side_effect=[10.0, 12.0], ), ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] handle_utterance.assert_not_awaited() run_turn.assert_awaited_once() self.assertIn( { "type": "eot", "ready": True, "reason": "ready", "silence_ms": 1300, "threshold_ms": 1200, }, websocket.sent_json, ) self.assertIn( { "type": "transcript", "text": "I am done now.", "final": True, "speaker": "counselor", }, websocket.sent_json, ) self.assertIn({"type": "state", "state": "thinking"}, websocket.sent_json) turn = run_turn.await_args.args[2] self.assertEqual(turn.learner_text, "I am done now.") self.assertEqual(turn.prosody.duration_s, 2.0) self.assertEqual(turn.prosody.silence_ms, 1300) self.assertIs(turn.prosody.barge_in, False) self.assertEqual( turn.prosody.provider_events, [ { "type": "speech_final", "confidence": 0.91, "event_type": "speech_final", "category": "speech_activity", } ], ) async def test_oversize_binary_audio_reports_error_and_drops_utterance( self, ) -> None: websocket = FakeWebSocket( [ _control({"type": "audio_start", "format": "webm"}), _binary(b"12345"), _control({"type": "close"}), ] ) handle_utterance = AsyncMock() run_turn = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal()), ), patch.object( voice_routes, "_bind_session", AsyncMock(return_value=self._bind_result()), ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ), patch.object( voice_routes, "_handle_utterance", handle_utterance, ), patch.object( voice_routes, "_run_turn_and_speak", run_turn, ), patch.object( voice_routes, "_MAX_AUDIO_BYTES", 4, ), ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] self.assertEqual(websocket.close_codes, [1000]) self.assertEqual( websocket.sent_json[-1], { "type": "error", "detail": "audio too large; please send a shorter utterance", }, ) handle_utterance.assert_not_awaited() run_turn.assert_not_awaited() async def test_unauthenticated_client_closes_before_session_or_voice_checks( self, ) -> None: websocket = FakeWebSocket() bind_session = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=None), ), patch.object( voice_routes, "_bind_session", bind_session, ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ) as is_available, ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] self.assertTrue(websocket.accepted) self.assertEqual(websocket.close_codes, [voice_routes.WS_CLOSE_UNAUTHORIZED]) self.assertEqual( websocket.sent_json, [{"type": "error", "detail": "not authenticated"}] ) bind_session.assert_not_awaited() is_available.assert_not_called() async def test_non_learner_client_closes_before_session_or_voice_checks( self, ) -> None: websocket = FakeWebSocket() bind_session = AsyncMock() with ( patch.object( voice_routes, "_principal_from_websocket", AsyncMock(return_value=_principal(Role.TEACHER)), ), patch.object( voice_routes, "_bind_session", bind_session, ), patch.object( voice_routes.voice_service, "is_available", return_value=True, ) as is_available, ): await voice_routes.voice_ws(websocket) # type: ignore[arg-type] self.assertTrue(websocket.accepted) self.assertEqual(websocket.close_codes, [voice_routes.WS_CLOSE_UNAUTHORIZED]) self.assertEqual( websocket.sent_json, [{"type": "error", "detail": "only learners can use voice"}], ) bind_session.assert_not_awaited() is_available.assert_not_called() async def test_bind_session_uses_db_voice_map_for_existing_session(self) -> None: websocket = FakeWebSocket() websocket.query_params = {"session_id": SESSION_ID} sess = SimpleNamespace(persona=SimpleNamespace(code="P2")) voice_map = PersonaVoiceMap( provider="openai", voice_id="voice-p2-custom", base_params={ "preset": "calm-adult-male", "openai_voice": "onyx", "rate": 1.08, "instructions": "Low, guarded delivery.", }, ) get_voice_map = AsyncMock(return_value=voice_map) with ( patch.object( voice_routes, "_load_voice_session", AsyncMock(return_value=(sess, None)), ), patch.object( voice_routes, "get_session_voice_map", get_voice_map, ), ): session_id, voice, err, meta = await voice_routes._bind_session( websocket, _principal() ) self.assertEqual(session_id, SESSION_ID) self.assertIsNone(err) self.assertEqual(meta["persona_catalog_source"], "session") self.assertEqual( voice, VoicePreset( preset="calm-adult-male", openai_voice="onyx", rate=1.08, instructions="Low, guarded delivery.", ), ) get_voice_map.assert_awaited_once_with(SESSION_ID) async def test_bind_session_rejects_existing_session_without_onboarding( self, ) -> None: websocket = FakeWebSocket() websocket.query_params = {"session_id": SESSION_ID} load_voice_session = AsyncMock() with ( patch.object( voice_routes, "user_onboarding_complete", AsyncMock(return_value=False), ), patch.object( voice_routes, "user_has_consent", AsyncMock(return_value=True), ), patch.object( voice_routes, "_load_voice_session", load_voice_session, ), ): session_id, voice, err, meta = await voice_routes._bind_session( websocket, _principal(profile_completed_at=None), ) self.assertIsNone(session_id) self.assertIsNone(voice) self.assertEqual(err, "onboarding_required") self.assertEqual(meta, {}) load_voice_session.assert_not_awaited() async def test_bind_session_rejects_existing_session_without_consent(self) -> None: websocket = FakeWebSocket() websocket.query_params = {"session_id": SESSION_ID} load_voice_session = AsyncMock() with ( patch.object( voice_routes, "user_has_consent", AsyncMock(return_value=False), ), patch.object( voice_routes, "_load_voice_session", load_voice_session, ), ): session_id, voice, err, meta = await voice_routes._bind_session( websocket, _principal(consent_at=None), ) self.assertIsNone(session_id) self.assertIsNone(voice) self.assertEqual(err, "consent_required") self.assertEqual(meta, {}) load_voice_session.assert_not_awaited() async def test_bind_session_explicit_preset_overrides_db_voice_map(self) -> None: websocket = FakeWebSocket() websocket.query_params = {"session_id": SESSION_ID, "preset": "soft-young-fem"} sess = SimpleNamespace(persona=SimpleNamespace(code="P2")) get_voice_map = AsyncMock() with ( patch.object( voice_routes, "_load_voice_session", AsyncMock(return_value=(sess, None)), ), patch.object( voice_routes, "get_session_voice_map", get_voice_map, ), ): session_id, voice, err, _ = await voice_routes._bind_session( websocket, _principal() ) self.assertEqual(session_id, SESSION_ID) self.assertIsNone(err) self.assertEqual(voice.preset, "soft-young-fem") self.assertEqual(voice.openai_voice, "coral") get_voice_map.assert_not_awaited() async def test_catalog_voice_map_is_used_for_dev_persona_binding_helper( self, ) -> None: voice_map = PersonaVoiceMap( provider="openai", voice_id="verse", base_params={"preset": "soft-young-fem", "rate": 0.9}, ) get_voice_map = AsyncMock(return_value=voice_map) with patch.object(voice_routes, "get_persona_voice_map", get_voice_map): voice = await voice_routes._resolve_catalog_voice( persona_id="00000000-0000-0000-0000-000000000301", version=7, persona_code="P1", explicit_preset=None, ) self.assertEqual(voice.preset, "soft-young-fem") self.assertEqual(voice.openai_voice, "verse") self.assertAlmostEqual(voice.rate, 0.9) get_voice_map.assert_awaited_once_with( persona_id="00000000-0000-0000-0000-000000000301", version=7, ) if __name__ == "__main__": unittest.main()