vignette/apps/api/app/test_persona_session_contract.py

314 lines
12 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 "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()