322 lines
13 KiB
Python
322 lines
13 KiB
Python
"""승인된 신규 페르소나의 세션 시작 계약 회귀 테스트."""
|
|
|
|
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 "FROM app.case_profile" in query and "FOR UPDATE" in query:
|
|
return 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 "FROM app.sessions" in query and "ended_at IS NULL" in query:
|
|
return None
|
|
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
|
|
or "pg_advisory_xact_lock" 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()
|