"""Regression tests for session turn persistence ordering.""" from __future__ import annotations import json import unittest from types import SimpleNamespace from unittest.mock import AsyncMock, patch from . import turn_runtime from .deps import Principal, Role from .engine_client import EngineError from .routes import sessions from .routes import voice as voice_routes from .services import memory, orchestrator, persona as persona_service, state_machine from .services.voice import TTSChunk, TranscriptResult, VoicePreset from .store import InProcSession, TurnRecord, 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", consent_at=1.0, ) 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_generate_real_crisis_stops_before_engine_and_returns_109_resource(self) -> None: principal = _principal() sess = _session(principal) async def should_not_generate(*args, **kwargs): raise AssertionError("engine must not be called for learner_real crisis") with patch.object(sessions.engine_client, "generate", should_not_generate): response = await sessions.submit_turn( sess.session_id, sessions.TurnRequest(text="저 지금 죽고싶어요. 도와주세요."), principal, ) self.assertTrue(response.safety_flagged) self.assertTrue(response.conversation_stopped) self.assertEqual(response.crisis_kind, "learner_real") self.assertIsNotNone(response.crisis_resource) self.assertEqual(response.crisis_resource.number, "109") self.assertIsNone(response.client_reply) self.assertEqual(len(sess.turns), 1) self.assertEqual(sess.turns[0].speaker, "counselor") async def test_record_safety_event_writes_teacher_alert_payload(self) -> None: 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="저 지금 자살하고 싶어요. 도와주세요.", theory_mode=sess.theory_mode, ) result = orchestrator.TurnResult( turn_seq=ctx.state_after.turn_seq if ctx.state_after else 1, stage=ctx.state_after.stage.value if ctx.state_after else "라포", effective_openness=ctx.state_after.effective_openness if ctx.state_after else 0.0, client_reply=None, safety_flagged=True, state_after=ctx.state_after or sess.state, crisis_kind="learner_real", crisis_resource={"title": "자살예방상담전화 109", "number": "109"}, conversation_stopped=True, ) calls: list[tuple[str, tuple[object, ...]]] = [] class FakeConn: async def execute(self, query: str, *args: object) -> str: calls.append((query, args)) return "INSERT 0 1" class FakeAcquire: async def __aenter__(self) -> FakeConn: return FakeConn() async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None: return None with patch.object(turn_runtime.db, "acquire", return_value=FakeAcquire()): await turn_runtime.record_safety_event(sess, ctx, result) self.assertEqual(len(calls), 1) query, args = calls[0] self.assertIn("INSERT INTO app.safety_events", query) self.assertEqual(args[0], sess.session_id) self.assertEqual(args[1], "learner_real") self.assertGreaterEqual(args[2], 4) detail = json.loads(args[3]) self.assertTrue(detail["conversation_stopped"]) self.assertEqual(detail["crisis_resource"]["number"], "109") self.assertEqual(detail["alert_status"], "teacher_dashboard") async def test_stream_turn_persists_client_engine_telemetry(self) -> None: principal = _principal() sess = _session(principal) async def successful_stream(ctx, engine, **kwargs): 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_stream_real_crisis_stops_before_engine_and_persists_learner_only(self) -> None: principal = _principal() sess = _session(principal) def should_not_stream(*args, **kwargs): raise AssertionError("stream engine must not be called for learner_real crisis") with patch.object(sessions.engine_client, "stream", should_not_stream): response = await sessions.stream_turn( sess.session_id, sessions.TurnRequest(text="저 지금 자살하고 싶어요. 도와주세요."), principal, ) body = await _consume_event_source(response) rendered = body.decode("utf-8") self.assertIn("'event': 'safety'", rendered) self.assertIn("'event': 'done'", rendered) self.assertIn("109", rendered) self.assertIn("conversation_stopped", rendered) self.assertEqual(len(sess.turns), 1) self.assertEqual(sess.turns[0].speaker, "counselor") async def test_stream_turn_persists_fast_loop_evaluation_on_learner_turn(self) -> None: principal = _principal() sess = _session(principal) async def successful_stream(ctx, engine, **kwargs): 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": 9, "tokens_out": 10, "cost_usd": 0.001, }, ) async def fake_eval_hook(ctx, client_reply): return { "loop": "fast", "turn_seq": ctx.state_after.turn_seq, "stage": ctx.state_after.stage.value, "appropriateness": "pos", "appropriateness_note": f"응답 반영: {client_reply}", } with patch.object(sessions.orchestrator, "run_turn_stream", successful_stream), patch.object( sessions.evaluator, "make_eval_hook", return_value=fake_eval_hook, ): response = await sessions.stream_turn( sess.session_id, sessions.TurnRequest(text="스트림 평가 발화"), principal, ) await _consume_event_source(response) self.assertEqual(len(sess.turns), 2) learner_turn, client_turn = sess.turns self.assertEqual(learner_turn.speaker, "counselor") self.assertIsNotNone(learner_turn.evaluation) self.assertEqual(learner_turn.evaluation["appropriateness"], "pos") self.assertIn("조금 말해볼게요", learner_turn.evaluation["appropriateness_note"]) self.assertIsNone(client_turn.evaluation) async def test_start_session_uses_stable_case_context_and_seed_recall(self) -> None: principal = _principal() card = persona_service.P1 case_context = sessions.session_persistence.CaseContext( case_id="00000000-0000-0000-0000-00000000ca5e", last_session_no=1, ) catalog_persona = SimpleNamespace( card=card, persona_id="00000000-0000-0000-0000-0000000000a1", version=3, degraded=False, ) recall = memory.RecallContext( recall_summary="지난 회기에서 가족 이야기를 열어두었다.", carry={ "rapport_credit": 0.6, "resistance": card.base_resistance(), "ideation_stage": card.ideation_baseline(), }, ) async def fake_create_session(**kwargs): self.assertEqual(kwargs["case_id"], case_context.case_id) self.assertEqual(kwargs["session_no"], 2) self.assertGreater(kwargs["state"].rapport_credit, 0) return InProcSession( session_id="stable-case-session", case_id=kwargs["case_id"], learner_id=principal.user_id, persona_code=card.code, theory_mode=kwargs["theory_mode"], persona=card, state=kwargs["state"], session_no=kwargs["session_no"], prev_rapport_credit=kwargs["carry_rapport"], ) def close_background(coro): coro.close() return None with patch.object(sessions, "get_catalog_persona", AsyncMock(return_value=catalog_persona)), patch.object( sessions.session_persistence, "get_case_context", AsyncMock(return_value=case_context), ), patch.object( sessions, "_build_seed_recall", AsyncMock(return_value=recall), ), patch.object( sessions.session_persistence, "create_session", fake_create_session, ), patch.object(sessions.asyncio, "create_task", close_background): response = await sessions.start_session( sessions.SessionStartRequest(persona_code=card.code), principal, ) self.assertEqual(response.case_id, case_context.case_id) self.assertEqual(response.session_no, 2) self.assertEqual(response.recall_summary, recall.recall_summary) self.assertIs(sessions._RECALL_CACHE[response.session_id], recall) async def test_start_session_requires_learner_consent_before_catalog_lookup(self) -> None: principal = _principal() principal.consent_at = None with patch.object( sessions, "get_catalog_persona", AsyncMock(side_effect=AssertionError("consent gate must run before catalog lookup")), ) as get_persona: with self.assertRaises(sessions.HTTPException) as caught: await sessions.start_session( sessions.SessionStartRequest(persona_code=persona_service.P1.code), principal, ) self.assertEqual(caught.exception.status_code, 403) self.assertEqual(caught.exception.detail, "consent_required") get_persona.assert_not_awaited() 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)) async def test_review_exposes_voice_nonverbal_events_on_learner_turn(self) -> None: principal = _principal() sess = _session(principal) created_at = sess.created_at sess.turns.extend( [ TurnRecord( turn_seq=1, speaker="counselor", stage=sess.state.stage.value, text="learner voice turn", text_masked="learner voice turn", created_at=created_at + 1, audio_ref="voice:webm:sha256:test", silence_ms=1234, speech_rate=420.0, barge_in=True, ), TurnRecord( turn_seq=2, speaker="client", stage=sess.state.stage.value, text="client reply", text_masked="client reply", created_at=created_at + 2, audio_ref="voice:webm:sha256:client", silence_ms=2500, speech_rate=180.0, barge_in=True, ), ] ) response = await sessions.get_session_review(sess.session_id, principal) self.assertEqual(len(response.turns), 2) learner_turn, client_turn = response.turns self.assertEqual([event.kind for event in learner_turn.nonverbal], ["silence", "pace", "barge_in", "audio"]) self.assertEqual(learner_turn.nonverbal[0].label, "침묵") self.assertEqual(learner_turn.nonverbal[0].detail, "1.2초") self.assertEqual(learner_turn.nonverbal[1].detail, "분당 420자") self.assertEqual(client_turn.nonverbal, []) async def test_review_includes_case_formulation_worksheet_draft(self) -> None: principal = _principal() sess = _session(principal) created_at = sess.created_at sess.turns.extend( [ TurnRecord( turn_seq=1, speaker="counselor", stage=sess.state.stage.value, text="오늘은 어떤 목표로 이야기해보고 싶으세요?", text_masked="오늘은 어떤 목표로 이야기해보고 싶으세요?", created_at=created_at + 1, ), TurnRecord( turn_seq=2, speaker="client", stage=sess.state.stage.value, text="요즘 너무 불안하고 친구 관계 스트레스 때문에 잠을 잘 못 자요.", text_masked="요즘 너무 불안하고 친구 관계 스트레스 때문에 잠을 잘 못 자요.", created_at=created_at + 2, ), ] ) response = await sessions.get_session_review(sess.session_id, principal) worksheet = response.caseWorksheet self.assertEqual(worksheet.status, "draft_from_transcript") self.assertGreaterEqual(len(worksheet.sections), 5) exploration = worksheet.sections[0] self.assertEqual(exploration.key, "exploration_11") complaint = next(item for item in exploration.items if item.key == "presenting_complaint") self.assertEqual(complaint.confidence, "medium") self.assertEqual(complaint.evidence[0].turnId, "t2") self.assertIn("불안", complaint.value or "") risk = next(item for item in exploration.items if item.key == "risk") self.assertEqual(risk.confidence, "none") self.assertEqual(risk.evidence, []) self.assertIn("명시 근거", risk.emptyReason or "") if __name__ == "__main__": unittest.main()