"""새 사례/이어가기 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()