"""Regression tests for session turn persistence ordering.""" from __future__ import annotations import unittest from unittest.mock import AsyncMock, patch from .deps import Principal, Role from .engine_client import EngineError from .routes import sessions from .routes import voice as voice_routes from .services import orchestrator, persona as persona_service, state_machine from .services.voice import TTSChunk, TranscriptResult, VoicePreset from .store import InProcSession, store def _principal() -> Principal: return Principal( user_id="00000000-0000-0000-0000-000000000101", role=Role.LEARNER, cohort_ids=[], email="turn-test@hs.ac.kr", display_name="Turn Test", ) def _session(principal: Principal) -> InProcSession: card = persona_service.P1 sess = InProcSession( session_id="turn-persistence-session", case_id="turn-persistence-case", learner_id=principal.user_id, persona_code=card.code, theory_mode="humanistic", persona=card, state=state_machine.SessionState( resistance=card.base_resistance(), ideation_stage=card.ideation_baseline(), ), ) store.put(sess) return sess async def _consume_event_source(response: object) -> bytes: body = bytearray() iterator = getattr(response, "body_iterator") async for chunk in iterator: if isinstance(chunk, str): body.extend(chunk.encode("utf-8")) elif isinstance(chunk, (bytes, bytearray)): body.extend(chunk) else: body.extend(str(chunk).encode("utf-8")) return bytes(body) class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self) -> None: store._sessions.clear() async def asyncTearDown(self) -> None: store._sessions.clear() async def test_generate_turn_engine_failure_does_not_append_learner_turn(self) -> None: principal = _principal() sess = _session(principal) with patch.object( sessions.orchestrator, "run_turn_generate", AsyncMock(side_effect=EngineError("engine unavailable: test")), ): with self.assertRaises(sessions.HTTPException) as caught: await sessions.submit_turn( sess.session_id, sessions.TurnRequest(text="실패한 발화"), principal, ) self.assertEqual(caught.exception.status_code, 503) self.assertEqual(sess.turns, []) async def test_generate_turn_persists_client_engine_telemetry(self) -> None: principal = _principal() sess = _session(principal) async def successful_turn(ctx, engine, **kwargs): assert ctx.state_after is not None return orchestrator.TurnResult( turn_seq=ctx.state_after.turn_seq, stage=ctx.state_after.stage.value, effective_openness=ctx.state_after.effective_openness, client_reply="괜찮아요. 천천히 말해볼게요.", safety_flagged=False, state_after=ctx.state_after, llm_provider="claude_cli", model="gateway-default", tokens_in=17, tokens_out=23, cost_usd=0.012345, ) with patch.object(sessions.orchestrator, "run_turn_generate", successful_turn): response = await sessions.submit_turn( sess.session_id, sessions.TurnRequest(text="요즘 많이 힘들었겠어요."), principal, ) self.assertEqual(response.client_reply, "괜찮아요. 천천히 말해볼게요.") self.assertEqual(len(sess.turns), 2) learner_turn, client_turn = sess.turns self.assertIsNone(learner_turn.llm_provider) self.assertEqual(client_turn.llm_provider, "claude_cli") self.assertEqual(client_turn.model, "gateway-default") self.assertEqual(client_turn.tokens_in, 17) self.assertEqual(client_turn.tokens_out, 23) self.assertEqual(client_turn.cost_usd, 0.012345) async def test_stream_turn_persists_client_engine_telemetry(self) -> None: principal = _principal() sess = _session(principal) async def successful_stream(ctx, engine): assert ctx.state_after is not None yield orchestrator.StreamEvent("token", {"text": "괜찮아요."}) yield orchestrator.StreamEvent( "done", { "session_id": ctx.session_id, "stage": ctx.state_after.stage.value, "effective_openness": ctx.state_after.effective_openness, "turn_seq": ctx.state_after.turn_seq, "safety_flagged": False, "llm_provider": "claude_cli", "model": "gateway-default", "tokens_in": 31, "tokens_out": 37, "cost_usd": 0.023456, }, ) with patch.object(sessions.orchestrator, "run_turn_stream", successful_stream): response = await sessions.stream_turn( sess.session_id, sessions.TurnRequest(text="스트림 성공 발화"), principal, ) body = await _consume_event_source(response) self.assertIn(b"done", body) self.assertEqual(len(sess.turns), 2) client_turn = sess.turns[1] self.assertEqual(client_turn.llm_provider, "claude_cli") self.assertEqual(client_turn.model, "gateway-default") self.assertEqual(client_turn.tokens_in, 31) self.assertEqual(client_turn.tokens_out, 37) self.assertEqual(client_turn.cost_usd, 0.023456) async def test_run_turn_stream_parses_gateway_done_telemetry(self) -> None: class FakeStreamEngine: engine_mode = "claude_cli" default_model = None async def stream(self, req): yield "event: token" yield '{"ignored":"not data"}' yield 'data: {"text":"부분 응답"}' yield "event: done" yield ( 'data: {"provider":"claude_cli","model":"gateway-default",' '"tokens_in":5,"tokens_out":7,"cost_usd":0.034567}' ) principal = _principal() sess = _session(principal) ctx = orchestrator.prepare_turn( session_id=sess.session_id, case_id=sess.case_id, card=sess.persona, state=sess.state, learner_text="게이트웨이 스트림 테스트", recent_turns=[], ) events = [ event async for event in orchestrator.run_turn_stream(ctx, FakeStreamEngine()) # type: ignore[arg-type] ] self.assertEqual([event.event for event in events], ["token", "done"]) self.assertEqual(events[0].data["text"], "부분 응답") self.assertEqual(events[1].data["llm_provider"], "claude_cli") self.assertEqual(events[1].data["model"], "gateway-default") self.assertEqual(events[1].data["tokens_in"], 5) self.assertEqual(events[1].data["tokens_out"], 7) self.assertEqual(events[1].data["cost_usd"], 0.034567) async def test_run_turn_stream_treats_gateway_error_event_as_error(self) -> None: class FakeStreamEngine: engine_mode = "claude_cli" default_model = None async def stream(self, req): yield "event: token" yield 'data: {"text":"부분 응답"}' yield "event: error" yield 'data: {"detail":"engine unavailable: gateway"}' principal = _principal() sess = _session(principal) ctx = orchestrator.prepare_turn( session_id=sess.session_id, case_id=sess.case_id, card=sess.persona, state=sess.state, learner_text="게이트웨이 오류 테스트", recent_turns=[], ) events = [ event async for event in orchestrator.run_turn_stream(ctx, FakeStreamEngine()) # type: ignore[arg-type] ] self.assertEqual([event.event for event in events], ["token", "error"]) self.assertIn("engine unavailable", events[1].data["detail"]) async def test_stream_turn_engine_error_event_does_not_append_partial_turns(self) -> None: principal = _principal() sess = _session(principal) async def failing_stream(*args, **kwargs): yield orchestrator.StreamEvent("token", {"text": "부분 응답"}) yield orchestrator.StreamEvent("error", {"detail": "engine unavailable: stream"}) with patch.object(sessions.orchestrator, "run_turn_stream", failing_stream): response = await sessions.stream_turn( sess.session_id, sessions.TurnRequest(text="스트림 실패 발화"), principal, ) body = await _consume_event_source(response) self.assertIn(b"engine unavailable: stream", body) self.assertEqual(sess.turns, []) async def test_voice_turn_engine_failure_does_not_append_learner_turn(self) -> None: class FakeWebSocket: def __init__(self) -> None: self.messages: list[dict[str, object]] = [] self.client_state = voice_routes.WebSocketState.CONNECTED async def send_text(self, data: str) -> None: import json self.messages.append(json.loads(data)) principal = _principal() sess = _session(principal) websocket = FakeWebSocket() with patch.object( voice_routes.orchestrator, "run_turn_generate", AsyncMock(side_effect=EngineError("voice engine unavailable")), ): await voice_routes._run_turn_and_speak( websocket, # type: ignore[arg-type] session_id=sess.session_id, principal=principal, voice_preset=VoicePreset(preset="neutral", openai_voice="sage"), learner_text="음성 실패 발화", ) self.assertTrue( any( message.get("type") == "error" and "engine unavailable" in str(message.get("detail")) for message in websocket.messages ), websocket.messages, ) self.assertEqual(sess.turns, []) async def test_voice_audio_turn_persists_paralinguistic_metadata(self) -> None: class FakeWebSocket: def __init__(self) -> None: self.messages: list[dict[str, object]] = [] self.binary: list[bytes] = [] self.client_state = voice_routes.WebSocketState.CONNECTED async def send_text(self, data: str) -> None: import json self.messages.append(json.loads(data)) async def send_bytes(self, data: bytes) -> None: self.binary.append(data) async def successful_turn(ctx, engine, **kwargs): assert ctx.state_after is not None return orchestrator.TurnResult( turn_seq=ctx.state_after.turn_seq, stage=ctx.state_after.stage.value, effective_openness=ctx.state_after.effective_openness, client_reply="천천히 말해줘서 고마워요.", safety_flagged=False, state_after=ctx.state_after, llm_provider="claude_cli", model="gateway-default", tokens_in=11, tokens_out=13, cost_usd=0.0012, ) async def fake_synthesize_stream(text, voice_preset): yield TTSChunk(audio=b"tts-audio") principal = _principal() sess = _session(principal) websocket = FakeWebSocket() audio = b"\x00\x80" * 1600 with patch.object( voice_routes.voice_service, "transcribe", AsyncMock(return_value=TranscriptResult(text="오늘은 좀 힘들었어요.", duration=2.0)), ), patch.object( voice_routes.orchestrator, "run_turn_generate", successful_turn, ), patch.object( voice_routes.voice_service, "synthesize_stream", fake_synthesize_stream, ): await voice_routes._handle_utterance( websocket, # type: ignore[arg-type] session_id=sess.session_id, principal=principal, voice_preset=VoicePreset(preset="neutral", openai_voice="sage"), audio=audio, fmt="webm", silence_ms=1234, barge_in=True, ) self.assertEqual(len(sess.turns), 2) learner_turn, client_turn = sess.turns self.assertTrue(str(learner_turn.audio_ref).startswith("voice:webm:sha256:")) self.assertEqual(learner_turn.silence_ms, 1234) self.assertGreater(learner_turn.speech_rate or 0, 0) self.assertTrue(learner_turn.barge_in) self.assertIsNone(client_turn.audio_ref) self.assertEqual(client_turn.llm_provider, "claude_cli") self.assertTrue(any(message.get("type") == "tts_end" for message in websocket.messages)) if __name__ == "__main__": unittest.main()