357 lines
13 KiB
Python
357 lines
13 KiB
Python
"""새 사례/이어가기 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()
|