"""승인된 신규 페르소나의 세션 시작 계약 회귀 테스트.""" from __future__ import annotations import unittest from dataclasses import replace from datetime import datetime, timezone from types import SimpleNamespace from typing import Any from unittest.mock import AsyncMock, patch from . import persona_repository, session_persistence from .deps import Principal, Role from .routes import sessions from .services import memory, persona as persona_service, state_machine from .store import store class _Transaction: async def __aenter__(self) -> None: return None async def __aexit__(self, exc_type: object, exc: object, tb: object) -> bool: return False class _Acquire: def __init__(self, conn: "_ContractConnection") -> None: self.conn = conn async def __aenter__(self) -> "_ContractConnection": return self.conn async def __aexit__(self, exc_type: object, exc: object, tb: object) -> bool: return False class _ContractConnection: """저장소·세션 SQL 계약을 보존하는 최소 상태형 DB 대역.""" def __init__(self) -> None: self.personas: list[dict[str, Any]] = [] self.session_insert_args: tuple[Any, ...] | None = None self.case_id = "00000000-0000-0000-0000-00000000ca5e" self.session_id = "00000000-0000-0000-0000-000000005355" def transaction(self) -> _Transaction: return _Transaction() async def fetchval(self, query: str, *args: Any) -> int: if "MAX(version)" not in query: raise AssertionError(f"unexpected fetchval: {query}") code = str(args[0]).upper() versions = [ int(row["version"]) for row in self.personas if str(row["code"]).upper() == code ] return (max(versions) if versions else 0) + 1 async def fetch(self, query: str, *args: Any) -> list[dict[str, Any]]: if "FROM app.persona_card" not in query or "status = 'approved'" not in query: raise AssertionError(f"unexpected fetch: {query}") latest: dict[str, dict[str, Any]] = {} for row in self.personas: if row["status"] != "approved": continue code = str(row["code"]).upper() if code not in latest or int(row["version"]) > int(latest[code]["version"]): latest[code] = row return [latest[code] for code in sorted(latest)] async def fetchrow(self, query: str, *args: Any) -> dict[str, Any] | None: now = datetime(2026, 8, 27, tzinfo=timezone.utc) if "INSERT INTO app.persona_card" in query: row = { "persona_id": str(args[0]), "code": str(args[1]).upper(), "version": int(args[2]), "status": str(args[3]), "display_name": args[4], "difficulty": args[5], "theory_target": list(args[6]), "demographics": dict(args[7]), "presenting": dict(args[8]), "history": dict(args[9]), "big5": dict(args[10]), "resistance": dict(args[11]), "speech_style": dict(args[12]), "affect_baseline": dict(args[13]), "ccd": dict(args[14]), "dsm5_dimensional": dict(args[15]), "source_provenance": args[16], "is_synthetic": bool(args[17]), "triggers": dict(args[18]), "created_by": args[19], "approved_by": None, "created_at": now, "approved_at": None, } self.personas.append(row) return row if "UPDATE app.persona_card" in query and "status IN ('draft', 'review')" in query: persona_id = str(args[0]) for row in self.personas: if row["persona_id"] != persona_id or row["status"] not in {"draft", "review"}: continue row["status"] = str(args[1]) row["approved_by"] = args[2] row["approved_at"] = now if row["status"] == "approved" else None return row return None if "FROM app.persona_card" in query and "status = 'approved'" in query: code = str(args[0]).upper() matches = [ row for row in self.personas if row["status"] == "approved" and str(row["code"]).upper() == code ] matches.sort(key=lambda row: int(row["version"]), reverse=True) return matches[0] if matches else None if "INSERT INTO app.case_profile" in query: return {"case_id": self.case_id, "last_session_no": 0} if "UPDATE app.case_profile" in query: return {"last_session_no": int(args[1])} if "INSERT INTO app.sessions" in query: self.session_insert_args = args return { "id": self.session_id, "runtime_case_id": args[0], "case_id": args[1], "learner_id": args[2], "persona_code": args[5], "session_no": args[8], "theory_mode": args[9], "started_at": now, "ended_at": None, "prev_rapport_credit": args[10], } raise AssertionError(f"unexpected fetchrow: {query}") async def execute(self, query: str, *args: Any) -> str: if "audit.audit_log" in query or "app.session_state" in query: return "INSERT 0 1" raise AssertionError(f"unexpected execute: {query}") def _principal(role: Role, user_id: str) -> Principal: return Principal( user_id=user_id, role=role, cohort_ids=[], email=f"{role.value}@example.test", display_name=role.value.title(), consent_at=1.0 if role == Role.LEARNER else None, profile_completed_at=1.0 if role == Role.LEARNER else None, ) class PersonaSessionContractTest(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self) -> None: store._sessions.clear() sessions._RECALL_CACHE.clear() async def asyncTearDown(self) -> None: store._sessions.clear() sessions._RECALL_CACHE.clear() def test_ideation_baseline_clamps_database_card_to_runtime_safety_contract(self) -> None: below_schema_minimum = replace( persona_service.P2, affect_baseline={ **persona_service.P2.affect_baseline, "suicide_ideation_stage": 0, }, ) above_runtime_cap = replace( persona_service.P2, affect_baseline={ **persona_service.P2.affect_baseline, "suicide_ideation_stage": 5, }, ) self.assertEqual(below_schema_minimum.ideation_baseline(), 1) self.assertEqual(below_schema_minimum.openness_params().ideation_baseline, 1) self.assertEqual(above_runtime_cap.ideation_baseline(), 3) self.assertEqual(above_runtime_cap.openness_params().ideation_baseline, 3) async def test_created_reviewed_and_approved_persona_starts_pinned_session(self) -> None: conn = _ContractConnection() teacher = _principal(Role.TEACHER, "00000000-0000-0000-0000-000000000901") learner = _principal(Role.LEARNER, "00000000-0000-0000-0000-000000000902") card = replace( persona_service.P1, code="P8", display_name="신규 내담자(가명)", presenting={ "complaint": "새로 등록한 주호소", "surface": "신뢰가 쌓인 뒤 구체화됨", }, source_provenance="contract test", ) def acquire(**_: Any) -> _Acquire: return _Acquire(conn) def close_background(coro: Any) -> SimpleNamespace: coro.close() return SimpleNamespace() with ( patch.object(persona_repository, "get_pool", return_value=object()), patch.object(persona_repository, "acquire", acquire), patch.object(session_persistence, "get_pool", return_value=object()), patch.object(session_persistence, "acquire", acquire), patch.object( sessions, "_build_seed_recall", AsyncMock(return_value=memory.RecallContext()), ), patch.object(sessions.asyncio, "create_task", close_background), ): created = await persona_repository.create_persona_draft( card=card, author_id=teacher.user_id, role="teacher", submit_for_review=True, ) approved = await persona_repository.update_persona_review_status( persona_id=created.persona_id, action="approve", reviewer_id=teacher.user_id, role="teacher", ) catalog = await persona_repository.list_catalog_personas() response = await sessions.start_session( sessions.SessionStartRequest(persona_code="P8"), learner, ) self.assertEqual(created.status, "review") self.assertIsNotNone(approved) assert approved is not None self.assertEqual(approved.status, "approved") self.assertEqual([entry.card.code for entry in catalog], ["P8"]) self.assertEqual(response.session_id, conn.session_id) self.assertEqual(response.persona_id, created.persona_id) self.assertEqual(response.persona_version, created.version) self.assertFalse(response.degraded) self.assertIsNotNone(conn.session_insert_args) assert conn.session_insert_args is not None self.assertEqual(str(conn.session_insert_args[3]), created.persona_id) self.assertEqual(conn.session_insert_args[4], created.version) self.assertEqual(conn.session_insert_args[5], "P8") async def test_create_session_raises_typed_error_when_durable_write_fails(self) -> None: with ( patch.object(session_persistence, "get_pool", side_effect=RuntimeError("db down")), patch.object(session_persistence, "runtime_fallback_allowed", return_value=False), ): with self.assertRaises(session_persistence.SessionCreationPersistenceError) as caught: await session_persistence.create_session( learner_id="00000000-0000-0000-0000-000000000902", card=persona_service.P1, theory_mode="humanistic", state=state_machine.SessionState(), ) self.assertEqual(str(caught.exception), "session_persistence_unavailable") self.assertIsInstance(caught.exception.__cause__, RuntimeError) async def test_start_route_exposes_stable_persistence_failure(self) -> None: learner = _principal(Role.LEARNER, "00000000-0000-0000-0000-000000000902") catalog_persona = persona_repository.CatalogPersona( card=replace(persona_service.P1, code="P8"), persona_id="00000000-0000-0000-0000-000000000908", version=1, source="database", ) with ( patch.object(sessions, "get_catalog_persona", AsyncMock(return_value=catalog_persona)), patch.object( sessions.session_persistence, "get_case_context", AsyncMock(return_value=None), ), patch.object( sessions, "_build_seed_recall", AsyncMock(return_value=memory.RecallContext()), ), patch.object( sessions.session_persistence, "create_session", AsyncMock( side_effect=session_persistence.SessionCreationPersistenceError( "session_persistence_unavailable" ) ), ), ): with self.assertRaises(sessions.HTTPException) as caught: await sessions.start_session( sessions.SessionStartRequest(persona_code="P8"), learner, ) self.assertEqual(caught.exception.status_code, 503) self.assertEqual(caught.exception.detail, "session_persistence_unavailable") if __name__ == "__main__": unittest.main()