326 lines
12 KiB
Python
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)
|