vignette/apps/api/app/test_session_continuity_guard.py

326 lines
12 KiB
Python

"""같은 내담자 케이스의 회기 순차성 회귀 테스트."""
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)