"""Focused RBAC, IDOR, and session-read audit regression tests.""" from __future__ import annotations import unittest from typing import Any from unittest.mock import AsyncMock, patch from fastapi import HTTPException from . import session_persistence from .deps import Principal, Role from .routes import sessions from .services import persona as persona_service, state_machine from .store import InProcSession, TurnRecord, store def _principal( *, user_id: str, role: Role = Role.LEARNER, ) -> Principal: return Principal( user_id=user_id, role=role, cohort_ids=["cohort-a"] if role == Role.TEACHER else [], email=f"{role.value}-{user_id[-4:]}@example.test", display_name=f"{role.value.title()} {user_id[-4:]}", ) def _session( *, session_id: str, learner_id: str, ) -> InProcSession: card = persona_service.P1 return InProcSession( session_id=session_id, case_id=f"case-{session_id[-12:]}", learner_id=learner_id, persona_code=card.code, theory_mode="humanistic", persona=card, state=state_machine.SessionState( resistance=card.base_resistance(), ideation_stage=card.ideation_baseline(), ), ) def _turn( *, seq: int, speaker: str, text: str, visible_to: tuple[str, ...] = ("client", "counselor", "evaluator"), ) -> TurnRecord: return TurnRecord( turn_seq=seq, speaker=speaker, stage="rapport", text=text, text_masked=text, created_at=1_800_000_000.0 + seq, visible_to=visible_to, ) class _Acquire: def __init__(self, conn: "_FakeConn") -> None: self.conn = conn async def __aenter__(self) -> "_FakeConn": return self.conn async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None: return None class _FakeConn: def __init__( self, *, fetchrow_results: list[Any] | None = None, fetch_results: list[list[Any]] | None = None, ) -> None: self.fetchrow_results = list(fetchrow_results or []) self.fetch_results = list(fetch_results or []) self.executed: list[tuple[str, tuple[Any, ...]]] = [] async def fetchrow(self, *args: Any, **kwargs: Any) -> Any: if not self.fetchrow_results: raise AssertionError("unexpected fetchrow") return self.fetchrow_results.pop(0) async def fetch(self, *args: Any, **kwargs: Any) -> list[Any]: if not self.fetch_results: raise AssertionError("unexpected fetch") return self.fetch_results.pop(0) async def execute(self, query: str, *args: Any) -> str: self.executed.append((query, args)) return "INSERT 0 1" def _audit_calls(conn: _FakeConn) -> list[tuple[str, tuple[Any, ...]]]: return [ call for call in conn.executed if "INSERT INTO audit.audit_log" in call[0] ] class LearnerSessionIdorTest(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self) -> None: store._sessions.clear() async def asyncTearDown(self) -> None: store._sessions.clear() async def test_get_session_detail_rejects_other_learner_session_id(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000101", ) intruder = _principal( user_id="00000000-0000-0000-0000-000000000202", ) sess = _session( session_id="00000000-0000-0000-0000-00000000a101", learner_id=owner.user_id, ) sess.turns.append( _turn( seq=1, speaker="counselor", text="owner-visible turn must not grant access", visible_to=("counselor", "evaluator"), ) ) store.put(sess) with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=None), ), patch.object(sessions.turn_runtime, "runtime_fallback_allowed", return_value=True), patch.object( sessions, "_review_ready", AsyncMock(side_effect=AssertionError("review lookup should not run")), ) as review_ready, ): with self.assertRaises(HTTPException) as caught: await sessions.get_session_detail(sess.session_id, intruder) self.assertEqual(caught.exception.status_code, 403) self.assertIn("does not belong", caught.exception.detail) review_ready.assert_not_awaited() async def test_super_admin_primary_learner_can_review_other_learner_session(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000103", ) super_admin = Principal( user_id="00000000-0000-0000-0000-000000000303", role=Role.LEARNER, admin_access=True, super_admin=True, email="yunchan@twentyoz.kr", display_name="Yun Chan", ) sess = _session( session_id="00000000-0000-0000-0000-00000000a103", learner_id=owner.user_id, ) store.put(sess) with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=None), ) as load_session, patch.object(sessions.turn_runtime, "runtime_fallback_allowed", return_value=True), ): loaded, review_principal = await sessions._load_review_session_or_404( sess.session_id, super_admin, ) self.assertIs(loaded, sess) self.assertEqual(review_principal.role, Role.ADMIN) self.assertEqual(load_session.await_count, 2) async def test_super_admin_primary_learner_review_loaders_use_admin_principal(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000105", ) super_admin = Principal( user_id="00000000-0000-0000-0000-000000000305", role=Role.LEARNER, admin_access=True, super_admin=True, email="yunchan@twentyoz.kr", display_name="Yun Chan", ) sess = _session( session_id="00000000-0000-0000-0000-00000000a105", learner_id=owner.user_id, ) sess.ended = True sess.ended_at = 1_800_000_120.0 sess.turns.extend( [ _turn(seq=1, speaker="counselor", text="admin-visible learner turn"), _turn(seq=2, speaker="client", text="admin-visible client turn"), ] ) store.put(sess) with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=None), ), patch.object(sessions.turn_runtime, "runtime_fallback_allowed", return_value=True), patch.object( sessions.session_persistence, "load_session_evaluation", AsyncMock(return_value=(None, False)), ) as load_session_evaluation, patch.object( sessions.session_persistence, "load_case_worksheet", AsyncMock(return_value=(None, False)), ) as load_case_worksheet, patch.object( sessions.session_persistence, "load_session_review_status", AsyncMock( return_value=( { "status": "viewed", "reviewer_id": "00000000-0000-0000-0000-000000000999", }, True, ) ), ) as load_session_review_status, ): response = await sessions.get_session_review(sess.session_id, super_admin) load_session_evaluation.assert_awaited_once() load_case_worksheet.assert_awaited_once() load_session_review_status.assert_awaited_once() evaluation_principal = load_session_evaluation.await_args.args[1] worksheet_principal = load_case_worksheet.await_args.args[1] review_status_principal = load_session_review_status.await_args.args[1] self.assertEqual(load_session_evaluation.await_args.args[0], sess.session_id) self.assertEqual(load_case_worksheet.await_args.args[0], sess.session_id) self.assertEqual(load_session_review_status.await_args.args[0], sess.session_id) self.assertIs(evaluation_principal, worksheet_principal) self.assertIs(evaluation_principal, review_status_principal) self.assertEqual(evaluation_principal.user_id, super_admin.user_id) self.assertEqual(evaluation_principal.role, Role.ADMIN) self.assertTrue(evaluation_principal.super_admin) self.assertEqual(response.session_id, sess.session_id) self.assertIsNotNone(response.teacherReview) self.assertEqual(response.teacherReview.status, "viewed") async def test_admin_access_flag_alone_does_not_review_other_learner_session(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000104", ) delegated_admin = Principal( user_id="00000000-0000-0000-0000-000000000304", role=Role.LEARNER, admin_access=True, super_admin=False, email="delegate@example.test", display_name="Delegated Admin", ) sess = _session( session_id="00000000-0000-0000-0000-00000000a104", learner_id=owner.user_id, ) store.put(sess) with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=None), ), patch.object(sessions.turn_runtime, "runtime_fallback_allowed", return_value=True), ): with self.assertRaises(HTTPException) as caught: await sessions._load_review_session_or_404(sess.session_id, delegated_admin) self.assertEqual(caught.exception.status_code, 403) self.assertIn("does not belong", caught.exception.detail) async def test_session_detail_filters_evaluator_only_turns(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000111", ) sess = _session( session_id="00000000-0000-0000-0000-00000000e111", learner_id=owner.user_id, ) sess.turns.extend( [ _turn(seq=1, speaker="counselor", text="learner normal turn"), _turn(seq=2, speaker="client", text="client normal turn"), _turn( seq=3, speaker="counselor", text="SECRET_EVALUATOR_ONLY_DETAIL", visible_to=("evaluator",), ), ] ) store.put(sess) with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=None), ), patch.object(sessions.turn_runtime, "runtime_fallback_allowed", return_value=True), patch.object(sessions, "_review_ready", AsyncMock(return_value=False)), ): response = await sessions.get_session_detail(sess.session_id, owner) self.assertEqual([turn.text for turn in response.turns], ["learner normal turn", "client normal turn"]) self.assertEqual([turn.speaker for turn in response.turns], ["learner", "client"]) async def test_session_review_filters_evaluator_only_turns_and_payload(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000112", ) sess = _session( session_id="00000000-0000-0000-0000-00000000e112", learner_id=owner.user_id, ) sess.ended = True sess.ended_at = 1_800_000_120.0 sess.turns.extend( [ _turn(seq=1, speaker="counselor", text="visible learner review turn"), _turn(seq=2, speaker="client", text="visible client review turn"), _turn( seq=3, speaker="client", text="SECRET_EVALUATOR_ONLY_REVIEW", visible_to=("evaluator",), ), ] ) store.put(sess) evaluation_record = { "status": "ready", "payload": { "strengths": ["SECRET_EVALUATOR_ONLY_REVIEW"], "improvements": ["SECRET_EVALUATOR_ONLY_REVIEW"], "supervisor_rationale": "SECRET_EVALUATOR_ONLY_REVIEW", "alternative_utterances": ["SECRET_EVALUATOR_ONLY_REVIEW"], }, } with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=None), ), patch.object(sessions.turn_runtime, "runtime_fallback_allowed", return_value=True), patch.object( sessions.session_persistence, "load_session_evaluation", AsyncMock(return_value=(evaluation_record, True)), ), ): response = await sessions.get_session_review(sess.session_id, owner) self.assertEqual([turn.text for turn in response.turns], ["visible learner review turn", "visible client review turn"]) self.assertFalse(response.reviewReady) rendered = " ".join( [ response.summary, response.clientFeedback or "", response.nextLine or "", *[turn.text for turn in response.turns], *[point.body for point in response.goodMoments], *[point.body for point in response.growthPoints], ] ) self.assertNotIn("SECRET_EVALUATOR_ONLY_REVIEW", rendered) async def test_teacher_can_read_session_review_but_not_save_learner_worksheet(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000114", ) teacher = _principal( user_id="00000000-0000-0000-0000-000000000901", role=Role.TEACHER, ) sess = _session( session_id="00000000-0000-0000-0000-00000000e114", learner_id=owner.user_id, ) sess.ended = True sess.ended_at = 1_800_000_120.0 sess.turns.extend( [ _turn(seq=1, speaker="counselor", text="teacher review visible learner turn"), _turn(seq=2, speaker="client", text="teacher review visible client turn"), ] ) with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=sess), ) as load_session, patch.object( sessions.session_persistence, "load_session_evaluation", AsyncMock(return_value=(None, False)), ), patch.object( sessions.session_persistence, "load_case_worksheet", AsyncMock(return_value=(None, False)), ), ): response = await sessions.get_session_review(sess.session_id, teacher) load_session.assert_awaited_once_with( sess.session_id, teacher, allow_ended=True, include_turn_evaluation=True, ) self.assertEqual(response.session_id, sess.session_id) self.assertEqual( [turn.text for turn in response.turns], ["teacher review visible learner turn", "teacher review visible client turn"], ) with self.assertRaises(HTTPException) as caught: await sessions.save_session_review_worksheet( sess.session_id, sessions.ReviewCaseWorksheetSaveRequest(sections=[], limitations=[]), teacher, ) self.assertEqual(caught.exception.status_code, 403) async def test_submit_turn_sends_only_client_visible_history_to_engine(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000113", ) sess = _session( session_id="00000000-0000-0000-0000-00000000e113", learner_id=owner.user_id, ) sess.turns.extend( [ _turn(seq=1, speaker="counselor", text="client-visible history"), _turn( seq=2, speaker="client", text="SECRET_EVALUATOR_ONLY_ENGINE_CONTEXT", visible_to=("evaluator",), ), ] ) store.put(sess) captured_recent_turns: list[dict[str, str]] | None = None async def successful_turn(ctx, engine, **kwargs): nonlocal captured_recent_turns captured_recent_turns = list(ctx.memory.recent_turns) assert ctx.state_after is not None return sessions.orchestrator.TurnResult( turn_seq=ctx.state_after.turn_seq, stage=ctx.state_after.stage.value, effective_openness=ctx.state_after.effective_openness, client_reply="client reply", safety_flagged=False, state_after=ctx.state_after, ) with ( patch.object( sessions.session_persistence, "load_session", AsyncMock(return_value=None), ), patch.object(sessions.turn_runtime, "runtime_fallback_allowed", return_value=True), patch.object(sessions.orchestrator, "run_turn_generate", successful_turn), ): await sessions.submit_turn( sess.session_id, sessions.TurnRequest(text="new learner turn"), owner, ) self.assertEqual(captured_recent_turns, [{"speaker": "counselor", "text": "client-visible history"}]) self.assertNotIn( "SECRET_EVALUATOR_ONLY_ENGINE_CONTEXT", " ".join(turn["text"] for turn in captured_recent_turns or []), ) async def test_learner_session_list_filters_other_runtime_sessions(self) -> None: owner = _principal( user_id="00000000-0000-0000-0000-000000000303", ) other = _principal( user_id="00000000-0000-0000-0000-000000000404", ) owned_session = _session( session_id="00000000-0000-0000-0000-00000000b303", learner_id=owner.user_id, ) other_session = _session( session_id="00000000-0000-0000-0000-00000000b404", learner_id=other.user_id, ) store.put(owned_session) store.put(other_session) with ( patch.object( sessions.session_persistence, "list_sessions", AsyncMock(return_value=([], False)), ), patch.object( sessions.session_persistence, "list_session_archives", AsyncMock(return_value=({}, False)), ), patch.object(sessions, "require_runtime_fallback_allowed", return_value=None), patch.object(sessions, "_review_ready", AsyncMock(return_value=False)), ): response = await sessions.list_learner_sessions(owner) self.assertEqual(response.source, "runtime") self.assertEqual([item.session_id for item in response.sessions], [owned_session.session_id]) class TeacherAdminAuditTest(unittest.IsolatedAsyncioTestCase): async def test_teacher_list_sessions_inserts_read_audit_log(self) -> None: teacher = _principal( user_id="00000000-0000-0000-0000-000000000505", role=Role.TEACHER, ) session_id = "00000000-0000-0000-0000-00000000c505" row = {"id": session_id} sess = _session( session_id=session_id, learner_id="00000000-0000-0000-0000-000000000606", ) conn = _FakeConn( fetchrow_results=[None], fetch_results=[[row], []], ) acquire_calls: list[dict[str, Any]] = [] def fake_acquire(**kwargs: Any) -> _Acquire: acquire_calls.append(kwargs) return _Acquire(conn) with ( patch.object(session_persistence, "get_pool", return_value=object()), patch.object(session_persistence, "acquire", fake_acquire), patch.object(session_persistence, "_session_from_rows", return_value=sess), ): found, durable = await session_persistence.list_sessions(teacher) self.assertTrue(durable) self.assertEqual(found, [sess]) self.assertEqual(acquire_calls[0]["role"], "teacher") audit = _audit_calls(conn) self.assertEqual(len(audit), 1) _, args = audit[0] self.assertEqual(args[0], teacher.user_id) self.assertEqual(args[1], "read_session") self.assertEqual(args[2], "session_list") self.assertEqual(args[3], "sessions") self.assertEqual(args[4]["access"], "list_sessions") self.assertEqual(args[4]["role"], "teacher") self.assertEqual(args[4]["result_count"], 1) async def test_admin_load_session_inserts_read_audit_log(self) -> None: admin = _principal( user_id="00000000-0000-0000-0000-000000000707", role=Role.ADMIN, ) session_id = "00000000-0000-0000-0000-00000000d707" learner_id = "00000000-0000-0000-0000-000000000808" sess = _session(session_id=session_id, learner_id=learner_id) conn = _FakeConn( fetchrow_results=[ {"id": session_id, "ended_at": None}, None, ], fetch_results=[[]], ) acquire_calls: list[dict[str, Any]] = [] def fake_acquire(**kwargs: Any) -> _Acquire: acquire_calls.append(kwargs) return _Acquire(conn) with ( patch.object(session_persistence, "get_pool", return_value=object()), patch.object(session_persistence, "acquire", fake_acquire), patch.object(session_persistence, "_session_from_rows", return_value=sess), ): found = await session_persistence.load_session( session_id, admin, allow_ended=True, ) self.assertEqual(found, sess) self.assertEqual(acquire_calls[0]["role"], "admin") audit = _audit_calls(conn) self.assertEqual(len(audit), 1) _, args = audit[0] self.assertEqual(args[0], admin.user_id) self.assertEqual(args[1], "read_session") self.assertEqual(args[2], "session") self.assertEqual(args[3], session_id) self.assertEqual(args[4]["access"], "load_session") self.assertEqual(args[4]["role"], "admin") self.assertEqual(args[4]["learner_id"], learner_id) if __name__ == "__main__": unittest.main()