"""G1 동맹 펄스 서비스의 비대칭·잠금·실패 회귀 검사.""" from __future__ import annotations import asyncio import unittest from contextlib import asynccontextmanager from typing import Any from unittest.mock import AsyncMock, patch from uuid import UUID, uuid4 from .contracts.engine_gateway import GenerateRequest, GenerateResponse from .contracts.measurement import ALLIANCE_DIMENSIONS from .deps import Principal, Role from .services import alliance_measurement as alliance def _turns() -> tuple[alliance.TranscriptTurn, ...]: return ( alliance.TranscriptTurn( turn_id=uuid4(), seq=1, speaker="counselor", text="오늘 함께 다루고 싶은 목표를 먼저 정해 볼까요?", ), alliance.TranscriptTurn( turn_id=uuid4(), seq=2, speaker="client", text="잠을 덜 미루는 방법을 찾고 싶어요.", ), ) def _assessment(score: float, evidence_index: int = 0) -> dict[str, Any]: return { dimension: { "score": score, "confidence": 0.8, "evidence_turn_indices": [evidence_index], "rationale": f"{dimension} 근거", } for dimension in ALLIANCE_DIMENSIONS } class RecordingEngine: engine_mode = "openai" live_client_provider = "claude_api" default_model = "gateway-default" def __init__( self, *, evaluator_error: Exception | None = None, evidence_index: int = 0, cleanup_error: Exception | None = None, ) -> None: self.requests: list[GenerateRequest] = [] self.closed_session_ids: list[str] = [] self.evaluator_error = evaluator_error self.evidence_index = evidence_index self.cleanup_error = cleanup_error async def generate(self, request: GenerateRequest) -> GenerateResponse: self.requests.append(request) if request.ai_role == "evaluator" and self.evaluator_error is not None: raise self.evaluator_error score = 0.75 if request.ai_role == "client" else 0.35 return GenerateResponse( text="", provider="claude_api" if request.ai_role == "client" else "openai", model="client-model" if request.ai_role == "client" else "observer-model", tokens_in=120, tokens_out=80, cost_usd=0.002, inference_geo="kr", structured=_assessment(score, self.evidence_index), ) async def close_session(self, session_id: str) -> bool: self.closed_session_ids.append(session_id) if self.cleanup_error is not None: raise self.cleanup_error return True class PersistenceConnection: def __init__(self, *, current_status: str = "awaiting_agents") -> None: self.current_status = current_status self.operations: list[tuple[str, tuple[Any, ...]]] = [] async def fetchval(self, query: str, *args: Any) -> Any: self.operations.append((query, args)) if "SELECT status FROM app.alliance_pulse" in query: return self.current_status return None async def execute(self, query: str, *args: Any) -> str: self.operations.append((query, args)) return "OK" def _acquire_for(conn: Any): @asynccontextmanager async def fake_acquire(**_: Any): yield conn return fake_acquire class AllianceAgentRunTest(unittest.IsolatedAsyncioTestCase): def test_prompt_anchors_keep_dimensions_independent(self) -> None: messages = alliance._messages( perspective="client_agent_report", checkpoint="post", turns=_turns(), ) system = messages[0].content self.assertIn("0.7~1.0=내담자의 명시적 수용·확인", system) self.assertIn("goal/task가 낮아도 독립적으로 높게", system) self.assertIn("내담자 후속 반응", system) self.assertEqual(alliance.PROMPT_BUNDLE_VERSION, "1.2.0") self.assertEqual(alliance.GENERATION_CONFIG["temperature"], 0.0) def test_provider_extra_fields_are_dropped_without_changing_scores(self) -> None: payload = _assessment(0.42) payload["goal"]["score_note"] = "설명 필드" payload["task"]["evidence_turn_indices_check"] = True payload["provider_comment"] = "schema 밖 설명" assessment, dropped = alliance._validated_assessment(payload) self.assertEqual(assessment.goal.score, 0.42) self.assertEqual(assessment.task.evidence_turn_indices, (0,)) self.assertEqual( dropped, ( "goal.score_note", "provider_comment", "task.evidence_turn_indices_check", ), ) class AlliancePulseIdempotencyTest(unittest.IsolatedAsyncioTestCase): async def _create_with_existing( self, *, stored_scores: dict[str, float], stored_evidence: tuple[UUID, ...], submitted_scores: alliance.AllianceScores, submitted_evidence: tuple[UUID, ...], ) -> alliance.LockedPulseResult: session_id = uuid4() learner_id = uuid4() pulse_id = uuid4() class Conn: async def fetchrow(self, query: str, *_args: Any) -> dict[str, Any] | None: if "count(t.id)::int AS turn_count" in query: return {"id": session_id, "ended_at": alliance._utc_now(), "turn_count": 4} if "INSERT INTO app.alliance_pulse" in query: return None if "JOIN app.self_assessment" in query: return { "pulse_id": pulse_id, "learner_id": learner_id, "scores": stored_scores, "evidence_turn_ids": list(stored_evidence), } raise AssertionError(f"unexpected fetchrow query: {query}") async def fetchval(self, query: str, *_args: Any) -> int: if "FROM app.turns" in query: return len(submitted_evidence) raise AssertionError(f"unexpected fetchval query: {query}") with patch.object(alliance, "acquire", _acquire_for(Conn())): return await alliance.create_locked_pulse( principal=Principal(str(learner_id), Role.LEARNER), session_id=session_id, checkpoint="post", scores=submitted_scores, evidence_turn_ids=submitted_evidence, ) async def test_same_locked_payload_returns_stable_idempotent_result(self) -> None: evidence = (uuid4(), uuid4()) scores = alliance.AllianceScores(goal=0.4, task=0.6, bond=0.8) result = await self._create_with_existing( stored_scores=scores.model_dump(), stored_evidence=tuple(reversed(evidence)), submitted_scores=scores, submitted_evidence=evidence, ) self.assertTrue(result.idempotent_replay) async def test_changed_locked_payload_is_conflict(self) -> None: evidence = (uuid4(),) with self.assertRaisesRegex( alliance.AlliancePulseConflictError, "different content", ): await self._create_with_existing( stored_scores={"goal": 0.4, "task": 0.6, "bond": 0.8}, stored_evidence=evidence, submitted_scores=alliance.AllianceScores(goal=0.41, task=0.6, bond=0.8), submitted_evidence=evidence, ) async def test_invalid_structured_attempt_is_audited_before_retry_success( self, ) -> None: class RetryEngine(RecordingEngine): async def generate(self, request: GenerateRequest) -> GenerateResponse: self.requests.append(request) structured = ( {"goal": _assessment(0.2)["goal"]} if len(self.requests) == 1 else _assessment(0.8) ) return GenerateResponse( text="", provider="claude_cli", model="test-model", structured=structured, ) engine = RetryEngine() run = await alliance.run_agent_assessment( pulse_id=uuid4(), session_id=uuid4(), checkpoint="post", perspective="independent_observer", turns=_turns(), engine=engine, # type: ignore[arg-type] ) self.assertIsNotNone(run.assessment) self.assertEqual(len(engine.requests), 2) self.assertEqual(len(run.prior_model_runs), 1) self.assertEqual(run.prior_model_runs[0].status, "error") self.assertTrue(run.prior_model_runs[0].error_code.startswith("agent_validation_")) self.assertEqual(run.model_run.status, "ready") self.assertEqual(run.model_run.metadata["prior_failed_attempts"], 1) self.assertEqual( run.prior_model_runs[0].input_evidence_hash, run.model_run.input_evidence_hash, ) self.assertNotEqual( run.prior_model_runs[0].model_run_id, run.model_run.model_run_id, ) self.assertEqual(len(set(engine.closed_session_ids)), 2) async def test_client_and_observer_are_independent_runs_with_evidence_provenance( self, ) -> None: pulse_id = uuid4() session_id = uuid4() turns = _turns() engine = RecordingEngine() client_run, observer_run = await asyncio.gather( alliance.run_agent_assessment( pulse_id=pulse_id, session_id=session_id, checkpoint="mid", perspective="client_agent_report", turns=turns, engine=engine, # type: ignore[arg-type] ), alliance.run_agent_assessment( pulse_id=pulse_id, session_id=session_id, checkpoint="mid", perspective="independent_observer", turns=turns, engine=engine, # type: ignore[arg-type] ), ) self.assertEqual(len(engine.requests), 2) request_by_role = {request.ai_role: request for request in engine.requests} self.assertNotEqual( request_by_role["client"].session_id, request_by_role["evaluator"].session_id, ) self.assertIn("client-report", request_by_role["client"].session_id or "") self.assertIn( "independent-observer", request_by_role["evaluator"].session_id or "" ) self.assertEqual( set(engine.closed_session_ids), { request_by_role["client"].session_id, request_by_role["evaluator"].session_id, }, ) self.assertNotEqual( client_run.model_run.model_run_id, observer_run.model_run.model_run_id ) self.assertNotEqual( client_run.model_run.prompt_bundle_hash, observer_run.model_run.prompt_bundle_hash, ) self.assertEqual( client_run.model_run.input_evidence_hash, observer_run.model_run.input_evidence_hash, ) self.assertEqual(client_run.assessment.goal.score, 0.75) # type: ignore[union-attr] self.assertEqual(observer_run.assessment.goal.score, 0.35) # type: ignore[union-attr] client_events = alliance._events_for_run( pulse_id=pulse_id, session_id=session_id, checkpoint="mid", turns=turns, run=client_run, ) self.assertEqual(len(client_events), 3) self.assertTrue( all( event.model_run_id == client_run.model_run.model_run_id for event in client_events ) ) self.assertTrue( all( event.evidence_turn_ids == (turns[0].turn_id,) for event in client_events ) ) async def test_prompt_bundle_hash_is_stable_while_input_hash_tracks_transcript( self, ) -> None: engine = RecordingEngine() first_turns = _turns() second_turns = ( first_turns[0], alliance.TranscriptTurn( turn_id=uuid4(), seq=2, speaker="client", text="이번에는 가족과의 갈등을 먼저 이야기하고 싶어요.", ), ) first = await alliance.run_agent_assessment( pulse_id=uuid4(), session_id=uuid4(), checkpoint="mid", perspective="independent_observer", turns=first_turns, engine=engine, # type: ignore[arg-type] ) second = await alliance.run_agent_assessment( pulse_id=uuid4(), session_id=uuid4(), checkpoint="mid", perspective="independent_observer", turns=second_turns, engine=engine, # type: ignore[arg-type] ) self.assertEqual( first.model_run.prompt_bundle_hash, second.model_run.prompt_bundle_hash, ) self.assertNotEqual( first.model_run.input_evidence_hash, second.model_run.input_evidence_hash, ) async def test_empty_transcript_degrades_without_calling_engine_or_inventing_scores( self, ) -> None: engine = RecordingEngine() run = await alliance._run_agent( pulse_id=uuid4(), session_id=uuid4(), checkpoint="pre", perspective="client_agent_report", turns=(), engine=engine, # type: ignore[arg-type] degradation_code="insufficient_transcript", ) self.assertEqual(engine.requests, []) self.assertEqual(run.model_run.status, "degraded") self.assertEqual(run.error_code, "insufficient_transcript") self.assertIsNone(run.assessment) events = alliance._events_for_run( pulse_id=uuid4(), session_id=run.model_run.session_id, # type: ignore[arg-type] checkpoint="pre", turns=(), run=run, ) self.assertTrue(all(event.status == "degraded" for event in events)) self.assertTrue(all(event.value is None for event in events)) self.assertTrue( all(event.error_code == "insufficient_transcript" for event in events) ) async def test_out_of_range_evidence_becomes_error_without_a_score(self) -> None: turns = _turns() engine = RecordingEngine(evidence_index=len(turns)) run = await alliance._run_agent( pulse_id=uuid4(), session_id=uuid4(), checkpoint="mid", perspective="independent_observer", turns=turns, engine=engine, # type: ignore[arg-type] ) self.assertEqual(run.model_run.status, "error") self.assertEqual(run.error_code, "evidence_out_of_range") self.assertIsNone(run.assessment) events = alliance._events_for_run( pulse_id=uuid4(), session_id=run.model_run.session_id, # type: ignore[arg-type] checkpoint="mid", turns=turns, run=run, ) self.assertTrue( all(event.status == "error" and event.value is None for event in events) ) async def test_unexpected_engine_and_cleanup_errors_still_return_error_provenance( self, ) -> None: engine = RecordingEngine( evaluator_error=RuntimeError("adapter exploded"), cleanup_error=RuntimeError("cleanup exploded"), ) run = await alliance._run_agent( pulse_id=uuid4(), session_id=uuid4(), checkpoint="mid", perspective="independent_observer", turns=_turns(), engine=engine, # type: ignore[arg-type] ) self.assertEqual(run.model_run.status, "error") self.assertEqual(run.error_code, "agent_runtimeerror") self.assertIsNone(run.assessment) self.assertEqual(len(engine.closed_session_ids), 1) class PulseInputTest(unittest.IsolatedAsyncioTestCase): async def test_masked_transcript_is_fail_closed_without_raw_text_fallback( self, ) -> None: pulse_id = uuid4() session_id = uuid4() class Conn: async def fetchrow(self, _query: str, _pulse_id: UUID) -> dict[str, Any]: return { "pulse_id": pulse_id, "session_id": session_id, "checkpoint": "mid", "status": "awaiting_agents", "self_assessment_id": uuid4(), "locked_at": alliance._utc_now(), } async def fetch( self, _query: str, _session_id: UUID ) -> list[dict[str, Any]]: return [ { "id": uuid4(), "seq": 1, "speaker": "counselor", "text": "주민번호가 포함된 원문", "text_masked": None, } ] with patch.object(alliance, "acquire", _acquire_for(Conn())): pulse_input = await alliance._load_pulse_input(pulse_id) self.assertEqual(pulse_input.turns, ()) self.assertEqual(pulse_input.degradation_code, "masked_transcript_unavailable") async def test_agent_run_refuses_pulse_without_locked_self_assessment(self) -> None: class Conn: async def fetchrow(self, _query: str, _pulse_id: UUID) -> dict[str, Any]: return { "pulse_id": _pulse_id, "session_id": uuid4(), "checkpoint": "mid", "status": "awaiting_agents", "self_assessment_id": None, "locked_at": None, } with patch.object(alliance, "acquire", _acquire_for(Conn())): with self.assertRaisesRegex( alliance.AlliancePulseStateError, "must be locked" ): await alliance._load_pulse_input(uuid4()) class AlliancePersistenceTest(unittest.IsolatedAsyncioTestCase): async def test_agent_results_are_persisted_before_atomic_reveal(self) -> None: pulse_id = uuid4() session_id = uuid4() turns = _turns() conn = PersistenceConnection() engine = RecordingEngine(evaluator_error=RuntimeError("observer unavailable")) pulse_input = alliance.PulseInput( session_id=session_id, checkpoint="mid", turns=turns, ) with ( patch.object( alliance, "_load_pulse_input", AsyncMock(return_value=pulse_input) ), patch.object(alliance, "acquire", _acquire_for(conn)), ): await alliance.run_alliance_agents(pulse_id, engine=engine) # type: ignore[arg-type] write_queries = [ query for query, _args in conn.operations if query.lstrip().startswith(("INSERT", "UPDATE")) ] self.assertEqual( sum("INSERT INTO audit.model_run" in query for query in write_queries), 2 ) self.assertEqual( sum( "INSERT INTO app.measurement_event" in query for query in write_queries ), 6, ) self.assertIn("UPDATE app.alliance_pulse", write_queries[-1]) update_args = next( args for query, args in reversed(conn.operations) if "UPDATE app.alliance_pulse" in query ) self.assertEqual( update_args, ( pulse_id, "degraded", "alliance_agent_partial_failure", ), ) async def test_no_transcript_persists_explicit_degraded_events_without_engine_calls( self, ) -> None: pulse_id = uuid4() conn = PersistenceConnection() engine = RecordingEngine() pulse_input = alliance.PulseInput( session_id=uuid4(), checkpoint="pre", turns=(), degradation_code="insufficient_transcript", ) with ( patch.object( alliance, "_load_pulse_input", AsyncMock(return_value=pulse_input) ), patch.object(alliance, "acquire", _acquire_for(conn)), ): await alliance.run_alliance_agents(pulse_id, engine=engine) # type: ignore[arg-type] self.assertEqual(engine.requests, []) event_operations = [ args for query, args in conn.operations if "INSERT INTO app.measurement_event" in query ] self.assertEqual(len(event_operations), 6) # INSERT parameter positions: value=$12, status=$16, error_code=$17. self.assertTrue(all(args[11] is None for args in event_operations)) self.assertTrue(all(args[15] == "degraded" for args in event_operations)) self.assertTrue( all(args[16] == "insufficient_transcript" for args in event_operations) ) update_args = next( args for query, args in reversed(conn.operations) if "UPDATE app.alliance_pulse" in query ) self.assertEqual(update_args, (pulse_id, "degraded", "insufficient_transcript")) async def test_background_processing_failure_is_persisted_on_the_pulse( self, ) -> None: pulse_id = uuid4() conn = PersistenceConnection() with ( patch.object( alliance, "_load_pulse_input", AsyncMock(side_effect=RuntimeError("db read exploded")), ), patch.object(alliance, "acquire", _acquire_for(conn)), ): await alliance.run_alliance_agents(pulse_id) query, args = next( (query, args) for query, args in conn.operations if "UPDATE app.alliance_pulse" in query ) self.assertIn("revealed_at = now()", query) self.assertEqual(args, (pulse_id, "alliance_processing_error")) async def test_list_query_hides_agent_and_supervisor_rows_until_reveal( self, ) -> None: session_id = uuid4() class Conn: def __init__(self) -> None: self.operations: list[tuple[str, tuple[Any, ...]]] = [] async def fetchval(self, query: str, *_args: Any) -> bool: self.operations.append((query, _args)) return True async def fetch(self, query: str, *_args: Any) -> list[Any]: self.operations.append((query, _args)) return [] conn = Conn() principal = Principal(str(uuid4()), Role.LEARNER, cohort_ids=["e2e-hanshin"]) with patch.object(alliance, "acquire", _acquire_for(conn)): result = await alliance.list_alliance_pulses( principal=principal, session_id=session_id, ) self.assertEqual(result, []) event_query, event_args = next( operation for operation in conn.operations if "FROM app.measurement_event me" in operation[0] ) self.assertIn("me.perspective = 'learner_self_report'", event_query) self.assertIn("me.pulse_id = ANY($2::uuid[])", event_query) self.assertNotIn("p.revealed_at IS NOT NULL", event_query) self.assertEqual(event_args, (session_id, [])) async def test_list_evidence_never_falls_back_to_raw_transcript(self) -> None: evidence_turn_id = uuid4() pulse_id = uuid4() class ReadConnection: def __init__(self) -> None: self.fetch_queries: list[str] = [] async def fetchval(self, _query: str, *_args: object) -> bool: return True async def fetch(self, query: str, *_args: object) -> list[dict[str, object]]: self.fetch_queries.append(query) if "FROM app.alliance_pulse" in query: return [] if "FROM app.measurement_event" in query: return [ { "measurement_id": uuid4(), "pulse_id": pulse_id, "dimension": "goal", "perspective": "independent_observer", "source_kind": "model_inferred", "value": 0.5, "confidence": 0.5, "status": "ready", "error_code": None, "evidence_turn_ids": [evidence_turn_id], "metadata": {}, "created_at": alliance._utc_now(), } ] return [] conn = ReadConnection() @asynccontextmanager async def fake_acquire(**_kwargs: object): yield conn with patch.object(alliance, "acquire", fake_acquire): await alliance.list_alliance_pulses( principal=Principal( str(uuid4()), Role.LEARNER, cohort_ids=["e2e-hanshin"], ), session_id=uuid4(), ) evidence_query = next( query for query in conn.fetch_queries if "FROM app.turns" in query ) self.assertIn("text_masked AS text", evidence_query) self.assertIn("NULLIF(btrim(text_masked), '') IS NOT NULL", evidence_query) self.assertNotIn("COALESCE", evidence_query) async def test_supervisor_cannot_rate_before_agent_reveal(self) -> None: class Conn: async def fetchrow(self, _query: str, *_args: Any) -> dict[str, Any]: return {"status": "awaiting_agents", "revealed_at": None} with patch.object(alliance, "acquire", _acquire_for(Conn())): with self.assertRaisesRegex( alliance.AlliancePulseStateError, "requires a revealed" ): await alliance.add_supervisor_rating( principal=Principal(str(uuid4()), Role.TEACHER), session_id=uuid4(), pulse_id=uuid4(), scores=alliance.AllianceScores(goal=0.5, task=0.5, bond=0.5), evidence_turn_ids=(uuid4(),), note="근거", ) class AllianceRecoveryTest(unittest.IsolatedAsyncioTestCase): async def asyncTearDown(self) -> None: tasks = tuple(alliance._background_tasks.values()) for task in tasks: if not task.done(): task.cancel() if tasks: await asyncio.gather(*tasks, return_exceptions=True) alliance._background_tasks.clear() async def test_scheduler_deduplicates_running_pulse_and_releases_key_on_completion( self, ) -> None: pulse_id = uuid4() started = asyncio.Event() release = asyncio.Event() calls: list[UUID] = [] async def controlled_runner(scheduled_pulse_id: UUID) -> None: calls.append(scheduled_pulse_id) started.set() await release.wait() alliance._background_tasks.clear() with patch.object(alliance, "run_alliance_agents", controlled_runner): self.assertTrue(alliance.schedule_alliance_agents(pulse_id)) await started.wait() self.assertFalse(alliance.schedule_alliance_agents(pulse_id)) self.assertEqual(calls, [pulse_id]) task = alliance._background_tasks[pulse_id] release.set() await task await asyncio.sleep(0) self.assertNotIn(pulse_id, alliance._background_tasks) self.assertTrue(alliance.schedule_alliance_agents(pulse_id)) await alliance._background_tasks[pulse_id] await asyncio.sleep(0) self.assertEqual(calls, [pulse_id, pulse_id]) async def test_recovery_reads_only_awaiting_rows_in_evaluator_context(self) -> None: pulse_ids = (uuid4(), uuid4(), uuid4()) acquire_contexts: list[dict[str, Any]] = [] class Conn: def __init__(self) -> None: self.query = "" async def fetch(self, query: str) -> list[dict[str, UUID]]: self.query = query return [{"pulse_id": pulse_id} for pulse_id in pulse_ids] conn = Conn() @asynccontextmanager async def recording_acquire(**kwargs: Any): acquire_contexts.append(kwargs) yield conn with ( patch.object(alliance, "acquire", recording_acquire), patch.object( alliance, "schedule_alliance_agents", side_effect=(True, False, True), ) as schedule, ): scheduled = await alliance.recover_pending_alliance_pulses() self.assertEqual(scheduled, 2) self.assertEqual( acquire_contexts, [{"ai_context": True, "ai_view": "evaluator"}], ) self.assertIn("WHERE status = 'awaiting_agents'", conn.query) self.assertEqual( [args.args[0] for args in schedule.call_args_list], list(pulse_ids), ) async def test_retry_does_not_write_if_pulse_became_terminal_during_inference( self, ) -> None: pulse_id = uuid4() conn = PersistenceConnection(current_status="ready") engine = RecordingEngine() pulse_input = alliance.PulseInput( session_id=uuid4(), checkpoint="mid", turns=_turns(), ) with ( patch.object( alliance, "_load_pulse_input", AsyncMock(return_value=pulse_input), ), patch.object(alliance, "acquire", _acquire_for(conn)), ): await alliance.run_alliance_agents( pulse_id, engine=engine, # type: ignore[arg-type] ) self.assertEqual(len(engine.requests), 2) self.assertFalse( any( query.lstrip().startswith(("INSERT", "UPDATE")) for query, _args in conn.operations ) ) async def test_failure_persistence_update_is_guarded_against_terminal_overwrite( self, ) -> None: pulse_id = uuid4() class TerminalConn: def __init__(self) -> None: self.status = "ready" self.query = "" async def execute(self, query: str, *_args: Any) -> str: self.query = query if self.status == "awaiting_agents": self.status = "error" return "UPDATE 0" conn = TerminalConn() with patch.object(alliance, "acquire", _acquire_for(conn)): await alliance._persist_pulse_failure( pulse_id, "alliance_processing_cancelled", ) self.assertEqual(conn.status, "ready") self.assertIn("AND status = 'awaiting_agents'", conn.query) if __name__ == "__main__": unittest.main()