회기 연속성과 멀티 케이스 계약을 영속화

This commit is contained in:
Yun Chan 2026-09-01 11:45:16 +09:00
parent be08c0b573
commit 72353ecd82
26 changed files with 2170 additions and 127 deletions

View file

@ -0,0 +1,357 @@
"""새 사례/이어가기 API의 learner 경계 계약 테스트."""
from __future__ import annotations
import unittest
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock, patch
from uuid import UUID
from fastapi import HTTPException
from .deps import Principal, Role
from .routes import sessions
from .services import persona as persona_service
from .store import InProcSession, store
LEARNER_ID = "00000000-0000-0000-0000-000000000741"
PERSONA_ID = "00000000-0000-0000-0000-000000000742"
CASE_ID = "00000000-0000-0000-0000-00000000ca5e"
ACTIVE_SESSION_ID = "00000000-0000-0000-0000-000000000743"
def _principal() -> Principal:
return Principal(
user_id=LEARNER_ID,
role=Role.LEARNER,
cohort_ids=[],
email="case-api-test@hs.ac.kr",
display_name="Case API Test",
consent_at=1.0,
profile_completed_at=1.0,
)
def _catalog_persona() -> SimpleNamespace:
return SimpleNamespace(
card=persona_service.P1,
persona_id=PERSONA_ID,
version=3,
degraded=False,
)
class _Acquire:
def __init__(self, conn: "_MemoryConnection") -> None:
self.conn = conn
async def __aenter__(self) -> "_MemoryConnection":
return self.conn
async def __aexit__(
self,
exc_type: object,
exc: object,
tb: object,
) -> bool:
return False
class _MemoryConnection:
"""foldout이 필요한 최소 learner-safe 행만 내는 DB 대역."""
def __init__(
self,
*,
case_row: dict[str, Any] | None,
summary_row: dict[str, Any] | None = None,
fact_rows: list[dict[str, Any]] | None = None,
) -> None:
self.case_row = case_row
self.summary_row = summary_row
self.fact_rows = fact_rows or []
self.calls: list[tuple[str, tuple[Any, ...]]] = []
async def fetchrow(self, query: str, *args: Any) -> dict[str, Any] | None:
self.calls.append((query, args))
if "FROM app.case_profile" in query:
return self.case_row
if "FROM app.session_summary AS ss" in query:
return self.summary_row
raise AssertionError(f"unexpected fetchrow: {query}")
async def fetch(self, query: str, *args: Any) -> list[dict[str, Any]]:
self.calls.append((query, args))
if "FROM app.pinned_fact AS pf" not in query:
raise AssertionError(f"unexpected fetch: {query}")
return self.fact_rows
class LearnerCaseListApiTest(unittest.IsolatedAsyncioTestCase):
async def test_case_list_projects_one_case_with_complete_case_local_stats(self) -> None:
principal = _principal()
catalog_persona = _catalog_persona()
active_started_at = datetime(2026, 8, 31, 9, 15, tzinfo=timezone.utc)
last_activity_at = datetime(2026, 8, 31, 10, 45, tzinfo=timezone.utc)
rows = [
{
"case_id": CASE_ID,
"last_session_no": 4,
"total_sessions": 4,
"completed_sessions": 3,
"total_turns": 18,
"total_duration_seconds": 5_400,
"active_session_id": ACTIVE_SESSION_ID,
"active_session_no": 4,
"active_started_at": active_started_at,
"last_activity_at": last_activity_at,
}
]
list_summaries = AsyncMock(return_value=rows)
with (
patch.object(
sessions,
"get_catalog_persona",
AsyncMock(return_value=catalog_persona),
),
patch.object(
sessions.session_persistence,
"list_case_summaries",
list_summaries,
),
):
response = await sessions.list_learner_cases(
persona_code=persona_service.P1.code,
principal=principal,
)
list_summaries.assert_awaited_once_with(
learner_id=principal.user_id,
persona_id=PERSONA_ID,
)
self.assertEqual(response.source, "database")
self.assertEqual(len(response.cases), 1)
case = response.cases[0]
self.assertEqual(case.case_id, CASE_ID)
self.assertEqual(case.persona_code, persona_service.P1.code)
self.assertEqual(case.persona_name, persona_service.P1.display_name)
self.assertEqual(case.last_session_no, 4)
self.assertEqual(case.progress.total_sessions, 4)
self.assertEqual(case.progress.completed_sessions, 3)
self.assertEqual(case.progress.total_turns, 18)
self.assertEqual(case.progress.total_duration_seconds, 5_400)
self.assertEqual(case.progress.active_session_id, ACTIVE_SESSION_ID)
self.assertEqual(case.progress.active_session_no, 4)
self.assertEqual(case.progress.active_started_at, active_started_at.isoformat())
self.assertEqual(case.progress.last_activity_at, last_activity_at.isoformat())
async def test_case_list_fails_closed_when_database_progress_is_unavailable(self) -> None:
principal = _principal()
with (
patch.object(
sessions,
"get_catalog_persona",
AsyncMock(return_value=_catalog_persona()),
),
patch.object(
sessions.session_persistence,
"list_case_summaries",
AsyncMock(
side_effect=sessions.session_persistence.CaseProgressUnavailableError(
"database unavailable"
)
),
),
):
with self.assertRaises(HTTPException) as raised:
await sessions.list_learner_cases(
persona_code=persona_service.P1.code,
principal=principal,
)
self.assertEqual(raised.exception.status_code, 503)
self.assertEqual(raised.exception.detail, "case_progress_unavailable")
class LearnerCaseMemoryPreviewApiTest(unittest.IsolatedAsyncioTestCase):
async def test_memory_preview_is_owner_scoped_and_never_selects_raw_session_data(self) -> None:
principal = _principal()
conn = _MemoryConnection(
case_row={"case_digest": "사례 요약 " + "" * 700},
summary_row={
"digest": "직전 회기 요약 " + "" * 700,
"open_threads": ["남은 주제 " + "" * 200, "다음 질문"],
},
fact_rows=[{"value": "기억 항목 " + "" * 200}],
)
with patch.object(
sessions.db,
"acquire",
return_value=_Acquire(conn),
) as acquire:
response = await sessions.get_learner_case_memory_preview(
case_id=UUID(CASE_ID),
principal=principal,
)
acquire.assert_called_once_with(role="learner", user_id=principal.user_id)
self.assertEqual(response.case_id, CASE_ID)
self.assertTrue(response.memory_available)
self.assertEqual(len(response.case_digest or ""), 600)
self.assertEqual(len(response.latest_session_digest or ""), 600)
self.assertEqual(len(response.open_threads[0]), 160)
self.assertTrue(response.open_threads[0].endswith(""))
self.assertEqual(response.open_threads[1], "다음 질문")
self.assertEqual(len(response.pinned_facts[0]), 160)
self.assertTrue(response.pinned_facts[0].endswith(""))
queried_sql = "\n".join(query for query, _ in conn.calls).lower()
self.assertIn("and learner_id = $2::uuid", queried_sql)
self.assertIn("and s.learner_id = $2::uuid", queried_sql)
self.assertIn("and cp.learner_id = $2::uuid", queried_sql)
self.assertNotIn("app.turns", queried_sql)
self.assertNotIn("transcript", queried_sql)
self.assertNotIn("evaluator", queried_sql)
self.assertNotIn("end_state", queried_sql)
self.assertNotIn(" text", queried_sql)
self.assertEqual({args[1] for _, args in conn.calls}, {principal.user_id})
response_payload = response.model_dump()
self.assertNotIn("transcript", response_payload)
self.assertNotIn("evaluator", response_payload)
self.assertNotIn("end_state", response_payload)
self.assertEqual(
set(response_payload),
{
"case_id",
"memory_available",
"case_digest",
"latest_session_digest",
"open_threads",
"pinned_facts",
},
)
async def test_memory_preview_returns_404_before_loading_unowned_case_details(self) -> None:
principal = _principal()
conn = _MemoryConnection(case_row=None)
with patch.object(
sessions.db,
"acquire",
return_value=_Acquire(conn),
):
with self.assertRaises(HTTPException) as raised:
await sessions.get_learner_case_memory_preview(
case_id=UUID(CASE_ID),
principal=principal,
)
self.assertEqual(raised.exception.status_code, 404)
self.assertEqual(raised.exception.detail, "case_not_found")
self.assertEqual(len(conn.calls), 1)
self.assertIn("FROM app.case_profile", conn.calls[0][0])
self.assertEqual(conn.calls[0][1], (CASE_ID, principal.user_id))
class SessionStartCaseModeApiTest(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self) -> None:
store._sessions.clear()
sessions._RECALL_CACHE.clear()
sessions._KB_CUES_CACHE.clear()
async def asyncTearDown(self) -> None:
store._sessions.clear()
sessions._RECALL_CACHE.clear()
sessions._KB_CUES_CACHE.clear()
async def test_fresh_start_rejects_selected_existing_case(self) -> None:
principal = _principal()
create_session = AsyncMock()
with (
patch.object(
sessions,
"get_catalog_persona",
AsyncMock(return_value=_catalog_persona()),
),
patch.object(
sessions.session_persistence,
"create_session",
create_session,
),
):
with self.assertRaises(HTTPException) as raised:
await sessions.start_session(
sessions.SessionStartRequest(
persona_code=persona_service.P1.code,
start_mode="fresh",
case_id=UUID(CASE_ID),
),
principal,
)
self.assertEqual(raised.exception.status_code, 422)
self.assertEqual(raised.exception.detail, "fresh_start_must_not_select_case")
create_session.assert_not_awaited()
async def test_continue_start_forwards_the_selected_case_to_durable_creation(self) -> None:
principal = _principal()
catalog_persona = _catalog_persona()
captured: dict[str, Any] = {}
async def fake_create_session(**kwargs: Any) -> InProcSession:
captured.update(kwargs)
return InProcSession(
session_id="selected-case-session",
case_id=str(kwargs["case_id"]),
learner_id=principal.user_id,
persona_code=persona_service.P1.code,
theory_mode=str(kwargs["theory_mode"]),
persona=persona_service.P1,
state=kwargs["state"],
persona_id=PERSONA_ID,
persona_version=3,
session_no=7,
prev_rapport_credit=float(kwargs["carry_rapport"]),
)
def close_background(coro: Any) -> SimpleNamespace:
coro.close()
return SimpleNamespace()
with (
patch.object(
sessions,
"get_catalog_persona",
AsyncMock(return_value=catalog_persona),
),
patch.object(
sessions.session_persistence,
"create_session",
AsyncMock(side_effect=fake_create_session),
),
patch.object(sessions.asyncio, "create_task", close_background),
):
response = await sessions.start_session(
sessions.SessionStartRequest(
persona_code=persona_service.P1.code,
start_mode="continue",
case_id=UUID(CASE_ID),
),
principal,
)
self.assertEqual(captured["case_id"], CASE_ID)
self.assertEqual(captured["start_mode"], "continue")
self.assertEqual(response.case_id, CASE_ID)
self.assertEqual(response.session_no, 7)
self.assertEqual(response.start_mode, "continue")
if __name__ == "__main__":
unittest.main()