"""같은 내담자 케이스의 회기 순차성 회귀 테스트.""" from __future__ import annotations import unittest from types import SimpleNamespace from unittest.mock import AsyncMock, patch from fastapi import HTTPException from . import 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, exc, tb) -> bool: return False class _Acquire: def __init__(self, conn: "_ActiveSessionConnection") -> None: self.conn = conn async def __aenter__(self) -> "_ActiveSessionConnection": return self.conn async def __aexit__(self, exc_type, exc, tb) -> bool: return False class _ActiveSessionConnection: """새 row INSERT 전 active-case 확인 순서를 검증하는 최소 DB 대역.""" active_session_id = "00000000-0000-0000-0000-000000000777" def __init__(self) -> None: self.active_query = "" self.insert_attempted = False def transaction(self) -> _Transaction: return _Transaction() async def fetchrow(self, query: str, *args: object) -> dict[str, object] | None: if "UPDATE app.case_profile" in query: return {"last_session_no": 2} if "FROM app.sessions" in query and "ended_at IS NULL" in query: self.active_query = query return {"id": self.active_session_id} if "INSERT INTO app.sessions" in query: self.insert_attempted = True raise AssertionError("active session guard must run before INSERT") raise AssertionError(f"unexpected query: {query}") async def execute(self, query: str, *args: object) -> str: if "pg_advisory_xact_lock" in query: return "SELECT 1" raise AssertionError(f"unexpected execute: {query}") def _principal() -> Principal: return Principal( user_id="00000000-0000-0000-0000-000000000101", role=Role.LEARNER, cohort_ids=[], email="continuity-guard@example.test", display_name="연속성 검증 학습자", consent_at=1.0, profile_completed_at=1.0, ) def _catalog_persona() -> SimpleNamespace: return SimpleNamespace( card=persona_service.P1, persona_id="00000000-0000-0000-0000-0000000000a1", version=1, degraded=False, ) class SessionContinuityGuardTest(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() async def test_persistence_rejects_a_second_active_session_before_insert(self) -> None: conn = _ActiveSessionConnection() principal = _principal() with ( patch.object(session_persistence, "get_pool", return_value=object()), patch.object(session_persistence, "acquire", return_value=_Acquire(conn)), ): with self.assertRaises(session_persistence.ActiveSessionExistsError) as raised: await session_persistence.create_session( learner_id=principal.user_id, card=persona_service.P1, theory_mode="humanistic", state=state_machine.SessionState(), session_no=2, persona_id="00000000-0000-0000-0000-0000000000a1", persona_version=1, case_id="00000000-0000-0000-0000-00000000ca5e", ) self.assertEqual(raised.exception.session_id, conn.active_session_id) self.assertIn("FOR UPDATE", conn.active_query) self.assertIn("persona_id = $2::uuid", conn.active_query) self.assertNotIn("case_id = $1::uuid", conn.active_query) self.assertFalse(conn.insert_attempted) async def test_persistence_builds_start_state_after_case_lock_and_active_check( self, ) -> None: class ReadyConnection: def __init__(self) -> None: self.events: list[str] = [] def transaction(self) -> _Transaction: return _Transaction() async def fetchrow( self, query: str, *args: object, ) -> dict[str, object] | None: if "FROM app.case_profile" in query and "FOR UPDATE" in query: self.events.append("case_lock") return {"case_id": args[0], "last_session_no": 1} if "UPDATE app.case_profile" in query: self.events.append("case_counter") return {"last_session_no": 2} if "FROM app.sessions" in query and "ended_at IS NULL" in query: self.events.append("active_check") return None if "INSERT INTO app.sessions" in query: self.events.append("session_insert") return { "id": "00000000-0000-0000-0000-000000000302", "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": session_persistence.datetime.fromtimestamp( 1_000.0, tz=session_persistence.timezone.utc, ), "ended_at": None, "prev_rapport_credit": args[10], } raise AssertionError(f"unexpected query: {query}") async def execute(self, query: str, *args: object) -> str: if "pg_advisory_xact_lock" in query: self.events.append("scope_lock") return "SELECT 1" if "INSERT INTO app.session_state" not in query: raise AssertionError(f"unexpected execute: {query}") self.events.append("state_insert") return "INSERT 0 1" conn = ReadyConnection() principal = _principal() recalled_state = state_machine.SessionState(rapport_credit=0.37) factory_calls: list[tuple[str, int]] = [] async def locked_state_factory( case_id: str, session_no: int, ) -> state_machine.SessionState: self.assertEqual( conn.events, ["scope_lock", "active_check", "case_lock", "case_counter"], ) factory_calls.append((case_id, session_no)) return recalled_state with ( patch.object(session_persistence, "get_pool", return_value=object()), patch.object(session_persistence, "acquire", return_value=_Acquire(conn)), ): created = await session_persistence.create_session( learner_id=principal.user_id, card=persona_service.P1, theory_mode="humanistic", state=state_machine.SessionState(), session_no=2, persona_id="00000000-0000-0000-0000-0000000000a1", persona_version=1, case_id="00000000-0000-0000-0000-00000000ca5e", locked_state_factory=locked_state_factory, ) self.assertEqual( factory_calls, [("00000000-0000-0000-0000-00000000ca5e", 2)], ) self.assertEqual( conn.events, [ "scope_lock", "active_check", "case_lock", "case_counter", "session_insert", "state_insert", ], ) self.assertIsNotNone(created) assert created is not None self.assertIs(created.state, recalled_state) self.assertEqual(created.prev_rapport_credit, 0.37) async def test_start_returns_resumable_conflict_for_persisted_active_session(self) -> None: principal = _principal() catalog_persona = _catalog_persona() active_session_id = "00000000-0000-0000-0000-000000000777" with ( patch.object( sessions, "get_catalog_persona", AsyncMock(return_value=catalog_persona), ), patch.object( sessions.session_persistence, "get_case_context", AsyncMock( return_value=session_persistence.CaseContext( case_id="00000000-0000-0000-0000-00000000ca5e", last_session_no=1, ) ), ), patch.object( sessions, "_build_seed_recall", AsyncMock(return_value=memory.RecallContext()), ), patch.object( sessions.session_persistence, "create_session", AsyncMock( side_effect=session_persistence.ActiveSessionExistsError( active_session_id ) ), ), ): with self.assertRaises(HTTPException) as raised: await sessions.start_session( sessions.SessionStartRequest(persona_code=persona_service.P1.code), principal, ) self.assertEqual(raised.exception.status_code, 409) self.assertEqual( raised.exception.detail, { "code": "active_session_exists", "session_id": active_session_id, }, ) async def test_runtime_fallback_rejects_a_second_active_session(self) -> None: principal = _principal() catalog_persona = _catalog_persona() existing = store.create( learner_id=principal.user_id, persona=persona_service.P1, theory_mode="humanistic", state=state_machine.SessionState(), persona_id=catalog_persona.persona_id, persona_version=catalog_persona.version, ) 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(return_value=None), ), patch.object(sessions, "require_runtime_fallback_allowed"), ): with self.assertRaises(HTTPException) as raised: await sessions.start_session( sessions.SessionStartRequest(persona_code=persona_service.P1.code), principal, ) self.assertEqual(raised.exception.status_code, 409) self.assertEqual( raised.exception.detail, { "code": "active_session_exists", "session_id": existing.session_id, }, ) self.assertEqual(len(store.list()), 1)