vignette/apps/api/app/test_alliance_measurement.py
Yun Chan 16e791e044 G0~G8 성과·동맹 측정 OS 작업 일괄 고정
8월 7일까지 워킹트리에만 남아 있던 미커밋 작업을 커밋한다. 여러 사본
폴더(worktree·clone)에 흩어져 있던 중간 스냅샷을 정리하기 전에 원본을
git 이력으로 고정하는 것이 목적이다.

- contracts/routes/services: measurement, outcome_trajectory, rupture_repair,
  deliberate_practice, calibration_transfer, supervision_research,
  multimodal_alliance, continuous_improvement 계열 신규 모듈과 테스트
- infra/db/init: 07~16 마이그레이션(측정 기반~calibration transfer 실행)
- apps/web: 세션 리뷰 카드·관리 화면·E2E 스펙 추가
- docs/ops: G0~G8 라이브 통합·배포·롤백 증거 문서와 evidence JSON/PNG
- scripts: smoke·ledger·릴리스 에이전트·NAS 프리뷰 운영 스크립트

engine.public 로그 .bak과 apps/web/test-results 산출물은 커밋에서 제외했다.
2026-08-08 01:30:53 +09:00

882 lines
31 KiB
Python

"""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()