diff --git a/CHANGELOG.md b/CHANGELOG.md index bacd9b5..c14abe9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,48 @@ the GitHub Release body, so a release with no entry here fails. Versioning follows [docs/versioning.md](docs/versioning.md). +## [0.5.0] - 2026-09-04 + +### Added + +- `agent_core.components.cycle` — the product-neutral write → audit → feedback → + re-write loop, extracted from ApodexHarness. `WriteAuditCycle` owns max-rounds, + the wall-clock budget with between-round estimation, the missing-artifact + output check, writer-exception and missing-output capture as synthetic audits, + per-round persistence of the audit trail, and failure-isolated round-observer + dispatch. Ships the `AuditReport` / `WriterOutput` / `AuditFinding` / + `CycleOutput` contracts with JSON round-trip, `DefaultFeedbackRenderer`, + `SessionBackedWriter` / `SessionBackedAuditor`, `ScoreThresholdAuditor`, + `MajorityVoteAuditor`, `select_best_attempt`, `select_by_answer_consensus`, and + the `BestSoFarObserver` / `PlateauAbortObserver` / `MetricsObserver` built-ins. +- `agent_core.components.verifier` — the `Verifier` and `Generator` protocols, + the `Verdict` model with nesting, `GroundTruth` / `VerifierContext` with + automatic oracle-field stripping at runtime, the six composers (Pipeline, + Ensemble, Fallback, Cascade, Parallel, ConsensusVerifier), and the + `AuditReport` ↔ `Verdict` bridge so a `Verifier` can stand in for a + `CycleAuditor`. +- `agent_core.components.observers.conclude_phase_observer.ConcludePhaseObserver` + — nudges a session toward emitting its structured output once it crosses a + ratio of its turn budget, instead of exploring until `max_turns` kills it and + returning empty content. `cycle.builders` depends on it. + +Scoring vocabularies stay free-form strings by design, and concrete writers and +auditors stay in the product. See +[docs/cycle-verifier-boundary.md](docs/cycle-verifier-boundary.md). + +### Fixed + +- `SessionBackedWriter._ensure_session` declared `-> str` while returning the + `str | None` attribute it caches into. Behavior is unchanged — the bus does + return an id — but the annotation no longer overstates what was proven. + +### Consumer action + +None required. This release only adds modules; nothing existing changed shape. +A product adopting these should replace its own copies with import aliases +rather than keeping both, since divergence is the failure this package exists to +prevent. + ## [0.4.0] - 2026-09-03 ### Added diff --git a/agent_core/components/cycle/__init__.py b/agent_core/components/cycle/__init__.py new file mode 100644 index 0000000..95bdb0a --- /dev/null +++ b/agent_core/components/cycle/__init__.py @@ -0,0 +1,102 @@ +"""WriteAuditCycle — iterate-until-pass artifact production. + +Generic write → audit → feedback → re-write loop over arbitrary +caller-supplied :class:`Writer` and :class:`Auditor` implementations. +The cycle owns the loop mechanics (max-rounds, wall-clock budget, +output existence check, exception capture, per-round persistence, +observer hooks); callers own the role-specific prompts, tools, and +rendering. + +Quick start:: + + from agent_core.components.cycle import ( + WriteAuditCycle, SessionBackedWriter, SessionBackedAuditor, + DefaultFeedbackRenderer, + ) + + writer = SessionBackedWriter( + bus=bus, task_id=task_id, + role_id="paper_writer", name="paper_writer", + system_prompt=PAPER_WRITER_PROMPT, + initial_prompt="Write a 3-section research paper on X.", + tools=[file_editor, read_text, web_search], + max_turns=80, llm_timeout=300, + work_dir=work_dir, + output_check=lambda wd: missing_in( + wd, ["sections/intro.md", "sections/method.md"], + ), + ) + auditor = SessionBackedAuditor( + bus=bus, task_id=task_id, + role_id="paper_auditor", + system_prompt=PAPER_AUDITOR_PROMPT, # must define AuditReport JSON schema + # tools defaults to [] — Codex-style hard tool restriction. + tools=[read_text, grep], # whitelist read-only inspection + max_turns=40, llm_timeout=300, + conclude_ratio=0.8, + ) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=work_dir, + output_check=writer.output_check, + max_rounds=10, + ) + output: CycleOutput = await cycle.run() + +See ``internal-docs/WRITE_AUDIT_CYCLE_REQUIREMENTS.md`` for the full spec +(FR1–FR11, NFR1–NFR7) and +``internal-docs/architecture/guide-write-audit-cycle.md`` for the user guide. +""" +from agent_core.components.cycle.builders import ( + MajorityVoteAuditor, + ScoreThresholdAuditor, + SessionBackedAuditor, + SessionBackedWriter, + parse_score_response, + select_best_attempt, + select_by_answer_consensus, +) +from agent_core.components.cycle.default_renderer import DefaultFeedbackRenderer +from agent_core.components.cycle.observers import ( + BaseRoundObserver, + CycleContext, + RoundIntervention, + RoundObserver, +) +from agent_core.components.cycle.observers_builtin import ( + BestSoFarObserver, + MetricsObserver, + PlateauAbortObserver, +) +from agent_core.components.cycle.protocols import FeedbackRenderer +from agent_core.components.cycle.types import ( + AuditFinding, + AuditReport, + CycleOutput, + WriterOutput, +) +from agent_core.components.cycle.write_audit_cycle import WriteAuditCycle + +__all__ = [ + "AuditFinding", + "AuditReport", + "BaseRoundObserver", + "BestSoFarObserver", + "CycleContext", + "CycleOutput", + "DefaultFeedbackRenderer", + "FeedbackRenderer", + "MajorityVoteAuditor", + "MetricsObserver", + "PlateauAbortObserver", + "RoundIntervention", + "RoundObserver", + "ScoreThresholdAuditor", + "SessionBackedAuditor", + "SessionBackedWriter", + "WriteAuditCycle", + "WriterOutput", + "parse_score_response", + "select_best_attempt", + "select_by_answer_consensus", +] diff --git a/agent_core/components/cycle/builders.py b/agent_core/components/cycle/builders.py new file mode 100644 index 0000000..51e799f --- /dev/null +++ b/agent_core/components/cycle/builders.py @@ -0,0 +1,1395 @@ +"""Session-backed default implementations of Writer / Auditor. + +These wrap :class:`AgentBus` so callers don't have to repeat the +boilerplate of creating a session, submitting a task, collecting the +result, and parsing structured output. Sophisticated callers +(implementing :class:`Writer` / :class:`Auditor` directly) bypass these. + +Two ready-made auditor flavours ship here: + +- :class:`SessionBackedAuditor` — generic AuditReport-emitting auditor + that expects the LLM to produce a JSON object matching the full + :class:`~agent_core.components.cycle.types.AuditReport` schema (verdict + + findings + summary + …). +- :class:`ScoreThresholdAuditor` — score-based grader. The LLM emits a + short ``{"score"|"points": int, "feedback"|"explanation": str}`` JSON + blob and the auditor maps the score against a caller-supplied + ``pass_threshold`` to produce the verdict. Useful for rubric-based + workflows (Olympiad-style proof grading, code review with rubric, + generic score-out-of-N evaluators) where the auditor *is* a grader, + not a critic that authors structured findings itself. + +Invariants this module enforces: + +- ``SessionBackedWriter`` creates **exactly one persistent session** on + the first call to :meth:`write` and reuses it for every subsequent + round (FR9). Conversation history accumulates across rounds. +- ``SessionBackedAuditor`` / ``ScoreThresholdAuditor`` create a + **fresh session per round** with + ``name=name_template.format(round=round_num)`` so prior-round + reasoning does not contaminate (FR9). Auditor sessions default to + ``tools_override=[]`` (empty tool set) — Codex-style hard restriction + (PRD §2.5). Caller whitelists explicitly via ``tools=[...]``. +- :class:`ConcludePhaseObserver` is auto-injected into every auditor + session so the LLM is nudged toward emitting structured output before + turn exhaustion (FR7). +- ``llm_timeout`` is propagated through to + ``AgentBus.create_session(llm_timeout=...)`` (FR8). +""" +from __future__ import annotations + +import asyncio +import collections +import json +import logging +import statistics +from collections.abc import Awaitable, Callable +from pathlib import Path +from typing import Any, cast + +from agent_core.components.cycle.protocols import CycleAuditor +from agent_core.components.cycle.types import ( + AuditFinding, + AuditReport, + WriterOutput, +) +from agent_core.components.observers.conclude_phase_observer import ConcludePhaseObserver +from agent_core.tool import Tool + +logger = logging.getLogger(__name__) + + +# Default revision instruction appended after the rendered feedback on +# rounds > 0. Generic; override via ``revision_instruction``. +_DEFAULT_REVISION_INSTRUCTION = ( + "Revise the artifact addressing every finding above. Maintain " + "everything that was already correct. Do not start over from " + "scratch unless the feedback explicitly requires it." +) + +# Default audit prompt envelope; the auditor's ``system_prompt`` is +# expected to contain the AuditReport JSON schema. This template only +# wraps the writer's output and a final reminder. Override via +# ``audit_prompt_template``. +_DEFAULT_AUDIT_PROMPT_TEMPLATE = ( + "The writer produced this artifact for round {round}.\n\n" + "---\n" + "{writer_content}\n" + "---\n\n" + "Files reported produced this round: {files}\n\n" + "Audit per your instructions and emit your AuditReport JSON now." +) + + +# ── SessionBackedWriter ───────────────────────────────────────────────────── + + +class SessionBackedWriter: + """:class:`Writer` backed by a persistent ``AgentBus`` session. + + The writer's ``initial_prompt`` is submitted on round 0. On + subsequent rounds, the rendered ``feedback_md`` is submitted as the + next task on the same session, with the ``revision_instruction`` + appended. Because the session is persistent, the LLM sees the entire + conversation (initial task + every prior round's feedback + every + prior round's writer output) when deciding what to revise. + + Parameters + ---------- + bus + The shared :class:`AgentBus` instance. + task_id + Parent task id. Sub-agent sessions live under this scope. + role_id + Stable role identifier registered in :class:`AgentRegistry` (or + a free-form string when ``llm_override`` + ``tools`` are passed + explicitly). + name + Session name. Combined with ``task_id`` to form + ``session_id`` (idempotent reuse — see ``AgentBus.create_session``). + system_prompt + Writer's persistent system prompt. Defines what the writer is / + does and how it should respond to feedback. + initial_prompt + The task description submitted on round 0. The "what to write". + tools + Tools the writer may call. Default ``None`` → resolved by + ResourceManager from ``role_id``. + max_turns + Per-task turn budget. Default 80. + llm_timeout + Per-LLM-call timeout (seconds). Default 300. Propagated to + ``AgentBus.create_session``. + work_dir + Filesystem root for artifacts. Stored on the writer so the + cycle can borrow ``output_check``. + output_check + ``(work_dir) -> list[str]`` returning missing required artifacts. + Stored on the writer so the cycle can read it as + ``writer.output_check``. The writer itself does not call this. + llm_override + Optional alternate LLM. Default ``None`` → resolved by + ResourceManager from ``role_id``. + observers_factory + Optional ``() -> list[observer]`` returning observers to attach + to every round's task. Each round calls this fresh, so observers + with internal state get a clean instance per round. + revision_instruction + Text appended after ``feedback_md`` on rounds > 0. Default is + a generic "revise addressing every finding" message. + timeout + ``collect`` timeout per round. Default 600s. + """ + + def __init__( + self, + *, + bus: Any, + task_id: str, + role_id: str, + name: str, + system_prompt: str, + initial_prompt: str, + work_dir: Path | str, + output_check: Callable[[Path], list[str]], + tools: list[Tool] | None = None, + max_turns: int = 80, + llm_timeout: int = 300, + llm_override: Any = None, + observers_factory: Callable[[], list[Any]] | None = None, + revision_instruction: str = _DEFAULT_REVISION_INSTRUCTION, + timeout: float = 600.0, + ) -> None: + self._bus = bus + self._task_id = task_id + self.role_id = role_id + self._name = name + self._system_prompt = system_prompt + self._initial_prompt = initial_prompt + self._tools = tools + self._max_turns = max_turns + self._llm_timeout = llm_timeout + self._llm_override = llm_override + self._observers_factory = observers_factory + self._revision_instruction = revision_instruction + self._timeout = timeout + + self.work_dir = Path(work_dir) + self.output_check = output_check + + self._session_id: str | None = None + + async def _ensure_session(self) -> str: + if self._session_id is None: + # Bind through a declared local first. ``self._bus`` is duck-typed, + # so the await yields ``Any``; assigning it straight to the + # ``str | None`` attribute left the declared ``-> str`` return + # unprovable, which is what the type checker objected to. + session_id: str = await self._bus.create_session( + task_id=self._task_id, + name=self._name, + role_id=self.role_id, + system_prompt=self._system_prompt, + tools_override=self._tools, + llm_override=self._llm_override, + max_turns=self._max_turns, + llm_timeout=self._llm_timeout, + ) + self._session_id = session_id + return self._session_id + + async def generate( + self, + prev_audit: Any, + round_num: int, + feedback_md: str, + ) -> WriterOutput: + del prev_audit # the Protocol passes it; the renderer already used it + sid = await self._ensure_session() + + if round_num == 0: + prompt = self._initial_prompt + else: + prompt = ( + f"{feedback_md}\n\n{self._revision_instruction}" + if feedback_md + else self._revision_instruction + ) + + observers = ( + self._observers_factory() if self._observers_factory else [] + ) + job_id = await self._bus.submit_task_to_session( + sid, prompt, observers=observers, + ) + cr = await self._bus.collect([job_id], timeout=self._timeout) + + sub_result = _extract_one_result(cr, label=f"writer round {round_num}") + return WriterOutput( + content=sub_result.final_content, + files=list(sub_result.metadata.get("files", []) or []), + message_count=int(sub_result.metadata.get("message_count", 0) or 0), + metadata=dict(sub_result.metadata), + loop_result=sub_result, + ) + + +# ── SessionBackedAuditor ──────────────────────────────────────────────────── + + +class SessionBackedAuditor: + """:class:`Auditor` backed by a fresh ``AgentBus`` session per round. + + Defaults to ``tools_override=[]`` — auditors run with no tools at + all. The auditor inspects the writer's output as text inside its + prompt; it does not rummage. Caller whitelists tools explicitly via + ``tools=[read_text, grep, ...]`` when read-only inspection is + required. + + Auto-injects :class:`ConcludePhaseObserver` so the auditor reliably + emits its structured output before turn exhaustion (FR7). + + Parameters + ---------- + bus + The shared :class:`AgentBus`. + task_id + Parent task id. + role_id + Stable role identifier. + name_template + Format string with one ``{round}`` placeholder for fresh per-round + session names. Default ``"auditor_R{round}"``. + system_prompt + Auditor's system prompt. **Should** include the AuditReport + JSON schema so the auditor knows what to emit. + tools + Whitelisted tool subset. Default ``[]`` (Codex-style hard + restriction). Pass ``None`` to fall back to ResourceManager — + only do this when the auditor genuinely needs the role's full + tool set. + max_turns + Per-task turn budget. Default 40 (auditors should be fast). + llm_timeout + Per-LLM-call timeout. Default 300. + conclude_ratio + ConcludePhaseObserver threshold. Default 0.8. + llm_override + Optional alternate LLM (e.g. a stronger reasoning model just for + review — Codex's ``review_model`` pattern). + observers_factory + Optional extra observers per round. + audit_prompt_template + Prompt envelope wrapping the writer's output. Override when the + default phrasing doesn't suit. + timeout + ``collect`` timeout per round. + """ + + def __init__( + self, + *, + bus: Any, + task_id: str, + role_id: str, + system_prompt: str, + name_template: str = "auditor_R{round}", + tools: list[Tool] | None = None, + max_turns: int = 40, + llm_timeout: int = 300, + conclude_ratio: float = 0.8, + llm_override: Any = None, + observers_factory: Callable[[], list[Any]] | None = None, + audit_prompt_template: str = _DEFAULT_AUDIT_PROMPT_TEMPLATE, + timeout: float = 600.0, + ) -> None: + self._bus = bus + self._task_id = task_id + self.role_id = role_id + self._system_prompt = system_prompt + self._name_template = name_template + # Default to empty tool set — Codex-style hard restriction. + # An explicit ``None`` is the (rare) opt-in to ResourceManager + # fallback. Sentinel _DEFAULT to disambiguate. + self._tools: list[Tool] = list(tools) if tools is not None else [] + self._max_turns = max_turns + self._llm_timeout = llm_timeout + self._conclude_ratio = conclude_ratio + self._llm_override = llm_override + self._observers_factory = observers_factory + self._audit_prompt_template = audit_prompt_template + self._timeout = timeout + + def _build_observers(self) -> list[Any]: + observers: list[Any] = [ + ConcludePhaseObserver(conclude_ratio=self._conclude_ratio), + ] + if self._observers_factory: + extra = self._observers_factory() + if extra: + observers.extend(extra) + return observers + + def _build_prompt( + self, writer_output: WriterOutput, round_num: int, + ) -> str: + files_text = ( + ", ".join(writer_output.files) + if writer_output.files + else "(none reported)" + ) + return self._audit_prompt_template.format( + round=round_num, + writer_content=writer_output.content, + files=files_text, + ) + + async def verify( + self, + writer_output: WriterOutput, + round_num: int, + ) -> AuditReport: + name = self._name_template.format(round=round_num) + sid = await self._bus.create_session( + task_id=self._task_id, + name=name, + role_id=self.role_id, + system_prompt=self._system_prompt, + tools_override=self._tools, # default [] — empty tool set + llm_override=self._llm_override, + max_turns=self._max_turns, + llm_timeout=self._llm_timeout, + ) + observers = self._build_observers() + prompt = self._build_prompt(writer_output, round_num) + job_id = await self._bus.submit_task_to_session( + sid, prompt, observers=observers, + ) + cr = await self._bus.collect([job_id], timeout=self._timeout) + + try: + sub_result = _extract_one_result( + cr, label=f"auditor round {round_num}", + ) + except _CycleRuntimeError as exc: + return _runtime_error_audit(str(exc)) + + return _parse_audit_report( + sub_result.final_content, + metadata={ + "auditor_session": sid, + "auditor_role": self.role_id, + "round": round_num, + }, + ) + + +# ── Helpers ───────────────────────────────────────────────────────────────── + + +class _CycleRuntimeError(RuntimeError): + """Internal — raised when a session task neither completes nor fails.""" + + +def _extract_one_result(cr: Any, *, label: str) -> Any: + if cr.completed: + return cr.completed[0] + if cr.failed: + return cr.failed[0] + raise _CycleRuntimeError( + f"{label}: AgentBus.collect returned no result (timeout / pending)", + ) + + +def _runtime_error_audit(message: str) -> AuditReport: + return AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="auditor_runtime_error", + severity="error", + short_message=message[:200], + detailed_message=message, + suggested_action=( + "Inspect AgentBus / session state. The cycle treats " + "this as iterate so it can retry on the next round." + ), + ), + ], + summary=f"Auditor session did not return a result: {message}", + confidence=1.0, + metadata={"synthetic": True, "reason": "auditor_runtime_error"}, + ) + + +def _parse_audit_report( + raw: str, *, metadata: dict[str, Any] | None = None, +) -> AuditReport: + """Extract AuditReport JSON from the auditor's free-form response. + + Looks for a fenced ``json`` block first, then for the largest + ``{...}`` substring. On failure, returns a structured + ``auditor_parse_failure`` audit so the cycle can iterate rather + than hang. The auditor's full raw text is preserved in + ``raw_text`` for debugging. + """ + data, parse_error = _try_parse_json(raw) + if data is None: + return AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="auditor_parse_failure", + severity="error", + short_message=( + "Auditor did not emit parseable AuditReport JSON" + ), + detailed_message=parse_error or "no json found in response", + suggested_action=( + "Inspect auditor's system prompt — it must " + "instruct the LLM to emit structured AuditReport " + "JSON." + ), + ), + ], + summary=( + "The auditor's response did not contain parseable JSON. " + "Treating verdict as iterate so the cycle can retry." + ), + raw_text=raw, + confidence=1.0, + metadata={ + "synthetic": True, + "reason": "auditor_parse_failure", + **(metadata or {}), + }, + ) + + try: + report = AuditReport.from_dict(data) + report.raw_text = raw + if metadata: + report.metadata = {**report.metadata, **metadata} + return report + except (KeyError, TypeError, ValueError) as exc: + return AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="auditor_schema_mismatch", + severity="error", + short_message=( + f"Auditor JSON did not match AuditReport schema: " + f"{exc!s}"[:200] + ), + detailed_message=str(exc), + ), + ], + summary=( + "The auditor's JSON parsed but did not match the " + "AuditReport schema. Treating as iterate." + ), + raw_text=raw, + confidence=1.0, + metadata={ + "synthetic": True, + "reason": "auditor_schema_mismatch", + **(metadata or {}), + }, + ) + + +def _try_parse_json(raw: str) -> tuple[dict[str, Any] | None, str | None]: + """Best-effort JSON extraction. Returns (data, error_message).""" + if not raw: + return None, "empty response" + + # Prefer a fenced ```json ... ``` block. + fenced = _extract_fenced_json(raw) + if fenced is not None: + try: + data = json.loads(fenced) + except json.JSONDecodeError as exc: + return None, f"fenced block did not parse: {exc!s}" + if isinstance(data, dict): + return cast("dict[str, Any]", data), None + return None, "fenced JSON was not an object" + + # Fall back to the widest {...} substring. + start = raw.find("{") + end = raw.rfind("}") + if start < 0 or end <= start: + return None, "no JSON object found" + try: + data = json.loads(raw[start : end + 1]) + except json.JSONDecodeError as exc: + return None, f"bare JSON did not parse: {exc!s}" + if isinstance(data, dict): + return cast("dict[str, Any]", data), None + return None, "JSON was not an object" + + +def _extract_fenced_json(raw: str) -> str | None: + marker = "```json" + idx = raw.find(marker) + if idx < 0: + return None + rest = raw[idx + len(marker):] + end = rest.find("```") + if end < 0: + return None + return rest[:end].strip() + + +# ── ScoreThresholdAuditor ─────────────────────────────────────────────────── + + +# Default user-prompt envelope for the grader. Kept generic — domain +# context (problem statement, ground-truth, rubric) is the caller's job +# to compose into either ``system_prompt`` or ``audit_prompt_template``. +_DEFAULT_GRADE_PROMPT_TEMPLATE = ( + "Grade the following candidate solution for round {round}.\n\n" + "---\n" + "{writer_content}\n" + "---\n\n" + "Files reported produced this round: {files}\n\n" + "Emit your grade as a single JSON object inside a ```json fenced " + "block. Schema:\n" + " {{\"score\": , \"feedback\": }}\n" + "No prose outside the JSON block." +) + + +class ScoreThresholdAuditor: + """:class:`Auditor` that maps a grader LLM's score to a verdict. + + Designed for the common "rubric grader" pattern where the auditor + LLM emits a tiny JSON blob — ``{"score": 7, "feedback": "..."}`` or + ``{"points": 6, "explanation": "..."}`` — and the cycle decides + whether to keep iterating based on whether the score crosses + ``pass_threshold``. + + Compared with :class:`SessionBackedAuditor`, this auditor: + + - Asks the LLM for a much simpler JSON shape (no ``findings`` + authoring required from the model). + - Maps ``score >= pass_threshold`` → ``pass_verdict`` (default + ``"success"``); else → ``fail_verdict`` (default ``"iterate"``). + - Always packs the parsed ``score`` into ``AuditReport.metadata`` + (under key ``"score"``) and into a single + ``category="grade"`` :class:`AuditFinding` so the writer's next + round sees the score in its rendered feedback. + - Sets ``confidence = (score - score_min) / (score_max - score_min)`` + (clamped to [0, 1]) when ``score_range`` is non-empty. + + Common patterns: + + - **Iterate-until-pass** (default): score ≥ threshold ends the cycle + with ``success=True``. Cheaper — early-exits on a good attempt. + - **Always-run-k** (parity with fixed-budget benchmark agents like + MiroVerifier's IMO-GVR): pass ``pass_verdict="iterate"`` so a high + score does *not* terminate the cycle. The cycle runs the full + ``max_rounds`` and the workflow picks the best attempt by score + from ``CycleOutput.history`` (use :func:`select_best_attempt`). + + Parameters + ---------- + bus + Shared :class:`AgentBus`. + task_id + Parent task id. + role_id + Stable role identifier. + system_prompt + Grader's system prompt — should describe the rubric and + instruct the LLM to emit the score JSON. The default + :data:`audit_prompt_template` already restates the JSON schema + on every round; you may rely on either or both. + name_template + Format string with one ``{round}`` placeholder. Default + ``"grader_R{round}"``. + pass_threshold + Minimum score that counts as a pass. Default ``6`` (matches the + IMO-GVR rubric's "essentially correct" cutoff). Inclusive. + pass_verdict + Verdict emitted when ``score >= pass_threshold``. Default + ``"success"``. Set to ``"iterate"`` for "always-run-k" mode. + fail_verdict + Verdict emitted when ``score < pass_threshold`` or the score is + unparseable. Default ``"iterate"``. + score_keys + Tuple of JSON keys to try, in order, when extracting the score. + Default ``("score", "points")`` — matches both common + conventions (raw "score" plus IMO-GVR's "points"). + feedback_keys + Tuple of JSON keys to try when extracting the feedback string. + Default ``("feedback", "explanation")``. + score_range + ``(min, max)`` for confidence calculation. Default ``(0, 7)`` + (IMO rubric). Pass ``None`` to disable confidence derivation. + audit_prompt_template + Prompt envelope wrapping the writer's output. ``{round}``, + ``{writer_content}``, ``{files}`` are substituted. Override when + you need to inject domain context (problem, ground-truth, + rubric) directly per-round. + tools, max_turns, llm_timeout, conclude_ratio, llm_override, + observers_factory, timeout + Same semantics as :class:`SessionBackedAuditor`. + """ + + def __init__( + self, + *, + bus: Any, + task_id: str, + role_id: str, + system_prompt: str, + name_template: str = "grader_R{round}", + pass_threshold: int | float = 6, + pass_verdict: str = "success", + fail_verdict: str = "iterate", + score_keys: tuple[str, ...] = ("score", "points"), + feedback_keys: tuple[str, ...] = ("feedback", "explanation"), + score_range: tuple[float, float] | None = (0, 7), + audit_prompt_template: str = _DEFAULT_GRADE_PROMPT_TEMPLATE, + tools: list[Tool] | None = None, + max_turns: int = 8, + llm_timeout: int = 300, + conclude_ratio: float = 0.8, + llm_override: Any = None, + observers_factory: Callable[[], list[Any]] | None = None, + timeout: float = 600.0, + ) -> None: + if not score_keys: + raise ValueError("score_keys must be non-empty") + if not feedback_keys: + raise ValueError("feedback_keys must be non-empty") + if score_range is not None and score_range[1] <= score_range[0]: + raise ValueError( + f"score_range max must exceed min, got {score_range!r}", + ) + + self._bus = bus + self._task_id = task_id + self.role_id = role_id + self._system_prompt = system_prompt + self._name_template = name_template + self._pass_threshold = pass_threshold + self._pass_verdict = pass_verdict + self._fail_verdict = fail_verdict + self._score_keys = tuple(score_keys) + self._feedback_keys = tuple(feedback_keys) + self._score_range = score_range + self._audit_prompt_template = audit_prompt_template + # Same Codex-style default as SessionBackedAuditor: empty tool + # set unless the caller explicitly opts into ResourceManager + # fallback by passing ``tools=None`` after construction. + self._tools: list[Tool] = list(tools) if tools is not None else [] + self._max_turns = max_turns + self._llm_timeout = llm_timeout + self._conclude_ratio = conclude_ratio + self._llm_override = llm_override + self._observers_factory = observers_factory + self._timeout = timeout + + def _build_observers(self) -> list[Any]: + observers: list[Any] = [ + ConcludePhaseObserver(conclude_ratio=self._conclude_ratio), + ] + if self._observers_factory: + extra = self._observers_factory() + if extra: + observers.extend(extra) + return observers + + def _build_prompt( + self, writer_output: WriterOutput, round_num: int, + ) -> str: + files_text = ( + ", ".join(writer_output.files) + if writer_output.files + else "(none reported)" + ) + return self._audit_prompt_template.format( + round=round_num, + writer_content=writer_output.content, + files=files_text, + ) + + async def verify( + self, + writer_output: WriterOutput, + round_num: int, + ) -> AuditReport: + name = self._name_template.format(round=round_num) + sid = await self._bus.create_session( + task_id=self._task_id, + name=name, + role_id=self.role_id, + system_prompt=self._system_prompt, + tools_override=self._tools, + llm_override=self._llm_override, + max_turns=self._max_turns, + llm_timeout=self._llm_timeout, + ) + observers = self._build_observers() + prompt = self._build_prompt(writer_output, round_num) + job_id = await self._bus.submit_task_to_session( + sid, prompt, observers=observers, + ) + cr = await self._bus.collect([job_id], timeout=self._timeout) + + try: + sub_result = _extract_one_result( + cr, label=f"grader round {round_num}", + ) + except _CycleRuntimeError as exc: + return _runtime_error_audit(str(exc)) + + return self._score_to_audit_report( + sub_result.final_content, + metadata={ + "auditor_session": sid, + "auditor_role": self.role_id, + "round": round_num, + }, + ) + + def _score_to_audit_report( + self, + raw: str, + *, + metadata: dict[str, Any] | None = None, + ) -> AuditReport: + return parse_score_response( + raw, + pass_threshold=self._pass_threshold, + pass_verdict=self._pass_verdict, + fail_verdict=self._fail_verdict, + score_keys=self._score_keys, + feedback_keys=self._feedback_keys, + score_range=self._score_range, + metadata=metadata, + ) + + +def parse_score_response( + raw: str, + *, + pass_threshold: int | float = 6, + pass_verdict: str = "success", + fail_verdict: str = "iterate", + score_keys: tuple[str, ...] = ("score", "points"), + feedback_keys: tuple[str, ...] = ("feedback", "explanation"), + score_range: tuple[float, float] | None = (0, 7), + metadata: dict[str, Any] | None = None, +) -> AuditReport: + """Build an :class:`AuditReport` from a grader LLM's raw response. + + Public helper extracting :class:`ScoreThresholdAuditor`'s parsing + + verdict-mapping logic. Useful when a workflow drives the grader + via raw LLM calls (no :class:`AgentBus` session) but still wants + the same structured output and ``WriteAuditCycle`` integration — + the IMO-proof workflow at ``workflows/imo_proof/`` is the canonical + consumer. + + Behaviour: + + - JSON parses + score key present + value coerces to a number → + verdict derived from ``pass_threshold``; one ``category="grade"`` + finding carries the score; ``metadata['score']`` populated. + - JSON parses but no score key (or non-numeric) → + ``grader_score_missing`` finding, verdict = ``fail_verdict``. + - JSON did not parse → + ``grader_parse_failure`` finding, verdict = ``fail_verdict``. + + The grader's full raw text is always preserved in + :attr:`AuditReport.raw_text` for debugging / replay. Caller-supplied + ``metadata`` is merged into :attr:`AuditReport.metadata` (with the + parser's own keys — ``score``, ``feedback``, ``passed``, + ``pass_threshold``, ``grader_data`` — taking precedence). + """ + meta_base: dict[str, Any] = dict(metadata or {}) + data, parse_error = _try_parse_json(raw) + if data is None: + return AuditReport( + verdict=fail_verdict, + findings=[ + AuditFinding( + category="grader_parse_failure", + severity="error", + short_message=( + "Grader did not emit parseable JSON" + ), + detailed_message=parse_error or "no JSON found", + suggested_action=( + "Inspect grader system prompt — it must " + "instruct the LLM to emit a JSON object " + "with the score and feedback keys." + ), + ), + ], + summary=( + "The grader's response did not contain parseable " + f"JSON. Treating verdict as {fail_verdict!r}." + ), + raw_text=raw, + confidence=0.0, + metadata={ + **meta_base, + "synthetic": True, + "reason": "grader_parse_failure", + }, + ) + + score = _extract_first_number(data, score_keys) + feedback = _extract_first_str(data, feedback_keys) or "" + + if score is None: + return AuditReport( + verdict=fail_verdict, + findings=[ + AuditFinding( + category="grader_score_missing", + severity="error", + short_message=( + "Grader JSON did not contain a numeric " + f"score under {score_keys!r}" + ), + detailed_message=( + f"Parsed JSON: {data!r}. Expected one of " + f"{score_keys!r} to be a number." + ), + ), + ], + summary=feedback or ( + "Grader returned JSON without a numeric score." + ), + raw_text=raw, + confidence=0.0, + metadata={ + **meta_base, + "synthetic": True, + "reason": "grader_score_missing", + "grader_data": data, + }, + ) + + passed = score >= pass_threshold + verdict = pass_verdict if passed else fail_verdict + confidence = _confidence_from_score(score, score_range) + + finding = AuditFinding( + category="grade", + severity="info" if passed else "warning", + short_message=( + f"Grader score: {score} " + f"(threshold {pass_threshold}, " + f"{'pass' if passed else 'fail'})" + ), + detailed_message=feedback, + suggested_action=( + "" if passed else + "Revise the artifact to address the grader's feedback " + "and re-submit on the next round." + ), + metadata={"score": score, "passed": passed}, + ) + + summary_lines = [ + f"Grader score: {score} (threshold " + f"{pass_threshold}, " + f"{'pass' if passed else 'fail'})", + ] + if feedback: + summary_lines.append("") + summary_lines.append(feedback) + summary = "\n".join(summary_lines) + + return AuditReport( + verdict=verdict, + findings=[finding], + summary=summary, + confidence=confidence, + raw_text=raw, + metadata={ + **meta_base, + "score": score, + "feedback": feedback, + "passed": passed, + "pass_threshold": pass_threshold, + "grader_data": data, + }, + ) + + +# ── Score / selection helpers ─────────────────────────────────────────────── + + +def _extract_first_number( + data: dict[str, Any], keys: tuple[str, ...], +) -> float | None: + """Return ``float(data[k])`` for the first ``k`` in ``keys`` whose + value is numeric (``int`` / ``float``) or a numeric string. Returns + ``None`` if no key yields a number.""" + for k in keys: + if k not in data: + continue + v = data[k] + if isinstance(v, bool): + # bool is a subclass of int in Python; reject explicitly. + continue + if isinstance(v, (int, float)): + return float(v) + if isinstance(v, str): + try: + return float(v.strip()) + except ValueError: + continue + return None + + +def _extract_first_str( + data: dict[str, Any], keys: tuple[str, ...], +) -> str | None: + """Return the first ``k`` in ``keys`` whose value is a non-empty + string. Returns ``None`` if no key yields a string.""" + for k in keys: + v = data.get(k) + if isinstance(v, str) and v.strip(): + return v + return None + + +def _confidence_from_score( + score: float, + score_range: tuple[float, float] | None, +) -> float: + """Linearly map ``score`` into ``[0, 1]`` over ``score_range``; + clamp to the unit interval. Returns ``0.0`` when ``score_range`` is + ``None`` (caller opted out).""" + if score_range is None: + return 0.0 + lo, hi = score_range + if hi <= lo: + return 0.0 + raw = (score - lo) / (hi - lo) + if raw < 0.0: + return 0.0 + if raw > 1.0: + return 1.0 + return raw + + +def select_best_attempt( + history: list[tuple[WriterOutput, AuditReport]], + *, + score_key: str = "score", + default_score: float = float("-inf"), +) -> tuple[int, WriterOutput, AuditReport]: + """Pick the highest-scoring attempt from a cycle's history. + + Companion to :class:`ScoreThresholdAuditor` for "always-run-k" + workflows: with ``pass_verdict="iterate"`` the cycle runs the full + ``max_rounds`` and the caller post-selects the best attempt from + ``CycleOutput.history``. + + Selection rule mirrors MiroVerifier IMO-GVR's ``_best_by_score``: + highest score wins; ties go to the *latest* attempt (later + revisions are at least as good as earlier ones). + + Parameters + ---------- + history + ``(WriterOutput, AuditReport)`` pairs as produced by + :class:`~agent_core.components.cycle.WriteAuditCycle`. Empty history + raises :class:`ValueError`. + score_key + Key under :attr:`AuditReport.metadata` carrying the score. + Default ``"score"`` (what :class:`ScoreThresholdAuditor` + writes). + default_score + Score used when an audit's metadata lacks ``score_key``. Default + ``-inf`` so unscored audits never beat scored ones. + + Returns + ------- + tuple + ``(zero_indexed_attempt_number, writer_output, audit_report)`` + for the selected attempt. + """ + if not history: + raise ValueError("history is empty; cannot select") + + def _key(item: tuple[int, tuple[WriterOutput, AuditReport]]) -> tuple[float, int]: + idx, (_w, audit) = item + raw = audit.metadata.get(score_key, default_score) + try: + score = float(raw) + except (TypeError, ValueError): + score = default_score + return (score, idx) + + best_idx, (best_w, best_a) = max(enumerate(history), key=_key) + return best_idx, best_w, best_a + + +# ── select_by_answer_consensus ────────────────────────────────────────────── + + +async def select_by_answer_consensus( + history: list[tuple[WriterOutput, AuditReport]], + *, + equivalence: Callable[[str, str], Awaitable[bool]], + answer_key: str = "extracted_answer", + score_key: str = "score", + default_score: float = float("-inf"), +) -> tuple[int, WriterOutput, AuditReport]: + """Pick the best attempt via answer-equivalence bucketing. + + Mirrors MiroVerifier@fac1b9e:agents/imo_gvr.py:_select answer mode + (the "bucket form" the user requested). Used by imo_proof's + answer mode where the writer emits a final symbolic answer per + round and we want consensus across the K sequential GVR attempts. + + Algorithm: + + 1. Filter ``history`` to *valid* attempts — those where + ``audit.metadata[answer_key]`` is a non-empty string. + 2. If none valid, fall back to :func:`select_best_attempt` over + the full history (parity with MiroVerifier's ``_select``). + 3. Bucket valid attempts: each placed in the first existing + bucket whose head answer it's equivalent to (per + ``equivalence``); else opens a new bucket. ``equivalence`` + raising is treated as "not equivalent" (defensive — the + upstream rule-based grader can throw on weird LaTeX). + 4. Sort buckets by ``(sum_of_scores, bucket_size)`` descending. + 5. Within the winning bucket, run :func:`select_best_attempt` + (highest score; ties → latest). + + Note this is *async* (unlike :func:`select_best_attempt`) because + realistic equivalence checks may dispatch to LLM rule-based + graders (`is_equivalent` was async in the original). + + Parameters + ---------- + history + ``(WriterOutput, AuditReport)`` pairs from the cycle. Empty + history raises :class:`ValueError`. + equivalence + Async callable ``(a, b) -> bool`` deciding if two answer + strings represent the same answer. + answer_key + Metadata key holding the extracted answer string. Default + ``"extracted_answer"`` (what + :class:`workflows.imo_proof.auditor.IMOProofGrader` writes + when ``mode="answer"``). + score_key + Metadata key for the per-attempt score. Default ``"score"``. + default_score + Score used when an audit lacks ``score_key`` (or it's + non-numeric). Default ``-inf`` — unscored attempts never win + a bucket selection. + + Returns + ------- + tuple + ``(zero_indexed_attempt_number, writer_output, audit_report)``. + """ + if not history: + raise ValueError("history is empty; cannot select") + + def _score_of(audit: AuditReport) -> float: + raw = audit.metadata.get(score_key, default_score) + try: + return float(raw) + except (TypeError, ValueError): + return default_score + + valid: list[tuple[int, str]] = [] + for i, (_w, audit) in enumerate(history): + ans = audit.metadata.get(answer_key) + if isinstance(ans, str) and ans.strip(): + valid.append((i, ans)) + + if not valid: + return select_best_attempt( + history, score_key=score_key, default_score=default_score + ) + + # Bucket by equivalence — each bucket holds (idx, answer) pairs; + # the head's answer represents the bucket. + buckets: list[list[tuple[int, str]]] = [] + for idx, ans in valid: + placed = False + for bucket in buckets: + head_ans = bucket[0][1] + if head_ans.strip() == ans.strip(): + bucket.append((idx, ans)) + placed = True + break + try: + eq = await equivalence(ans, head_ans) + except Exception: + eq = False + if eq: + bucket.append((idx, ans)) + placed = True + break + if not placed: + buckets.append([(idx, ans)]) + + # Sort buckets by (sum_score, bucket_size) descending + def _bucket_key(bucket: list[tuple[int, str]]) -> tuple[float, int]: + score_sum = sum(_score_of(history[idx][1]) for idx, _ in bucket) + return (score_sum, len(bucket)) + + buckets.sort(key=_bucket_key, reverse=True) + winner = buckets[0] + + # Within winning bucket: best-by-score (ties → latest, per + # select_best_attempt parity). + winner_history = [history[idx] for idx, _ in winner] + rel_idx, best_w, best_a = select_best_attempt( + winner_history, score_key=score_key, default_score=default_score + ) + abs_idx = winner[rel_idx][0] + return abs_idx, best_w, best_a + + +# ── MajorityVoteAuditor ───────────────────────────────────────────────────── + + +# Aggregator over the per-judge numeric scores. Default: median (robust +# to one outlier judge in n=3). Use ``statistics.mean`` for averaging, +# ``max`` / ``min`` for optimistic / pessimistic, or any custom callable. +ScoreAggregator = Callable[[list[float]], float] + + +class MajorityVoteAuditor: + """Aggregate ``n`` independent calls of a base :class:`Auditor`. + + Wraps any caller-supplied auditor (typically a + :class:`ScoreThresholdAuditor` or a stateless raw-LLM grader) and + runs it ``n_votes`` times **in parallel** per round, then collapses + the results into a single :class:`AuditReport`: + + - ``verdict`` — majority over the per-judge verdicts. Ties are + broken in favour of the first verdict observed (i.e. + ``Counter.most_common(1)`` semantics). + - ``metadata['score']`` — ``score_aggregator(per_judge_scores)``; + default :func:`statistics.median`. + - ``confidence`` — fraction of judges that returned a parseable + AuditReport (i.e. ``n_succeeded / n_votes``). + - ``findings`` — one synthetic ``category="grade"`` finding + carrying the aggregated score plus the per-judge score list. + - ``summary`` — concatenation of every judge's free-form feedback. + - ``raw_text`` — every judge's raw response, separated by a clear + delimiter for debugging / replay. + - ``metadata['individual_scores']`` / ``['individual_verdicts']`` + preserved for downstream selection / analysis. + + Reduces single-judge grader variance — particularly useful for + rubric-based scoring where one model run can plausibly land on + score=6 vs. score=7. With ``n_votes=3`` and median, you need two + out of three judges to agree on the lower score for the verdict + to flip — much more stable. + + Note: ``MajorityVoteAuditor`` and + :class:`agent_core.components.verifier.ConsensusVerifier` + are **not equivalent** despite both being "majority over N parallel + runs" — they vote on different fields: + + - This class votes on ``AuditReport.verdict`` label (pass/fail). + - ``ConsensusVerifier`` votes on ``Verdict.metadata[answer_key]`` + (extracted candidate answer). + + Both legitimately coexist; pick the one that matches the topology. + (Earlier docstrings suggested this collapses into + ``cycle_auditor_from_verifier(ConsensusVerifier(...))`` — that was + wrong; see ``internal-docs/designs/2026-04-30-verifier-first-class-lite.md`` + §0.1.) + + The base auditor is responsible for its own per-call freshness + (e.g. :class:`SessionBackedAuditor` / :class:`ScoreThresholdAuditor` + create a fresh session per call already; for stateless raw-LLM + auditors, repeated calls naturally produce independent samples + when the LLM has non-zero temperature). + + Parameters + ---------- + base + Any :class:`Auditor` implementation. Called ``n_votes`` times + per :meth:`audit` invocation. + n_votes + Number of independent base-auditor calls per round. Default + ``3`` (cheapest non-trivial majority). Must be ≥ 1. + score_aggregator + Function reducing the list of per-judge numeric scores into a + single number. Default :func:`statistics.median`. Callers may + pass ``statistics.mean`` for averaging or a custom callable. + role_id + Override the auditor's role identifier. Default copies + ``base.role_id``. + + Notes + ----- + Per-judge exceptions are caught: if a base call raises (or + returns a non-AuditReport), it is excluded from the aggregate + and contributes to a lowered ``confidence``. If **every** base + call fails, the wrapper emits a synthetic + ``category="majority_vote_failed"`` finding and verdict + ``"iterate"``. + """ + + role_id: str = "majority_vote_auditor" + + def __init__( + self, + *, + base: CycleAuditor, + n_votes: int = 3, + score_aggregator: ScoreAggregator = statistics.median, + role_id: str | None = None, + ) -> None: + if n_votes < 1: + raise ValueError( + f"n_votes must be >= 1, got {n_votes!r}", + ) + self._base = base + self._n = n_votes + self._aggregator = score_aggregator + self.role_id = role_id or base.role_id + + async def verify( + self, + writer_output: WriterOutput, + round_num: int, + ) -> AuditReport: + results: list[Any] = await asyncio.gather( + *[ + self._base.verify(writer_output, round_num) + for _ in range(self._n) + ], + return_exceptions=True, + ) + + valid: list[AuditReport] = [ + r for r in results if isinstance(r, AuditReport) + ] + errors: list[BaseException] = [ + r for r in results if isinstance(r, BaseException) + ] + + if not valid: + error_msgs = "; ".join( + f"{type(e).__name__}: {str(e)[:120]}" for e in errors[:5] + ) + return AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="majority_vote_failed", + severity="error", + short_message=( + f"All {self._n} base audits failed" + ), + detailed_message=error_msgs, + suggested_action=( + "Inspect the base auditor — every parallel " + "call raised. The cycle treats this as " + "iterate so it can retry on the next round." + ), + ), + ], + summary=( + f"Majority vote failed: all {self._n} base " + f"auditors raised exceptions." + ), + confidence=0.0, + metadata={ + "synthetic": True, + "reason": "majority_vote_failed", + "n_votes": self._n, + "n_succeeded": 0, + }, + ) + + verdict_counts = collections.Counter(r.verdict for r in valid) + majority_verdict, _ = verdict_counts.most_common(1)[0] + + per_judge_scores: list[float] = [] + for r in valid: + score_meta = r.metadata.get("score") + if isinstance(score_meta, (int, float)) and not isinstance( + score_meta, bool, + ): + per_judge_scores.append(float(score_meta)) + + agg_score: float | None + if per_judge_scores: + try: + agg_score = float(self._aggregator(per_judge_scores)) + except Exception: + logger.exception( + "MajorityVoteAuditor: score_aggregator raised; " + "falling back to median", + ) + agg_score = float(statistics.median(per_judge_scores)) + else: + agg_score = None + + feedbacks = [ + (r.metadata.get("feedback") or r.summary or "").strip() + for r in valid + ] + summary_lines = [ + f"Majority vote of {len(valid)}/{self._n} judges:", + f" verdict={majority_verdict!r} score={agg_score} " + f"(per-judge scores={per_judge_scores})", + "", + ] + for i, fb in enumerate(feedbacks, start=1): + summary_lines.append(f"=== Judge {i} ===") + summary_lines.append(fb if fb else "(no feedback)") + summary_lines.append("") + summary = "\n".join(summary_lines).rstrip() + + finding = AuditFinding( + category="grade", + severity="info" if ( + agg_score is not None and agg_score >= 6 + ) else "warning", + short_message=( + f"Majority vote score: {agg_score} " + f"({len(valid)}/{self._n} judges, " + f"verdict={majority_verdict})" + ), + detailed_message=summary, + metadata={ + "score": agg_score, + "individual_scores": per_judge_scores, + "individual_verdicts": [r.verdict for r in valid], + }, + ) + + raw_text = "\n\n=== JUDGE BREAK ===\n\n".join( + r.raw_text for r in valid if r.raw_text + ) + + return AuditReport( + verdict=majority_verdict, + findings=[finding], + summary=summary, + confidence=len(valid) / self._n, + raw_text=raw_text, + metadata={ + "score": agg_score, + "feedback": summary, + "individual_scores": per_judge_scores, + "individual_verdicts": [r.verdict for r in valid], + "n_votes": self._n, + "n_succeeded": len(valid), + "n_failed": len(errors), + }, + ) diff --git a/agent_core/components/cycle/default_renderer.py b/agent_core/components/cycle/default_renderer.py new file mode 100644 index 0000000..eab4da2 --- /dev/null +++ b/agent_core/components/cycle/default_renderer.py @@ -0,0 +1,130 @@ +"""Default FeedbackRenderer — turns an AuditReport into the next-round prompt. + +The renderer is deliberately verbose: it emits both a Markdown table +(human-readable) and a JSON detail block (machine-readable) in the same +prompt segment, plus the auditor's free-form ``summary`` paragraph +verbatim above. This lets downstream LLMs read either form depending on +their tool / parsing capability, and gives humans inspecting the trace a +readable artifact. + +Design notes: + +- Empty findings + non-terminal verdict produces an explicit + "re-attempt" message rather than an empty string (FR6). Returning "" + would make the writer think the audit was successful and the writer's + prior round was fine, when in fact the auditor failed to articulate + what went wrong. +- ``raw_text`` is **not** included in the output. It is for debugging + only — including it would spam the writer's context with the + auditor's chain-of-thought / probing. +- The output is plain text (Markdown). The cycle injects it as a + ``HumanMessage`` content; we do not nest it in any envelope here. +""" +from __future__ import annotations + +import json +from collections.abc import Iterable + +from agent_core.components.cycle.types import AuditFinding, AuditReport + +_HEADER = "## Audit feedback for the previous round\n" + +_REATTEMPT_MESSAGE = ( + "**Audit verdict**: `{verdict}`. The auditor did not list specific " + "findings, which means the artifact is not yet acceptable but the " + "auditor could not pinpoint a single issue. Re-attempt the task in " + "full — review the requirements, do not assume the prior draft is " + "salvageable." +) + +_TABLE_HEADER = ( + "| # | Severity | Category | Where | Short message | Suggested action |\n" + "|---|---|---|---|---|---|\n" +) + + +class DefaultFeedbackRenderer: + """Hybrid Markdown + JSON renderer (FR6). + + Output structure (top-to-bottom): + + 1. Section header. + 2. Verdict line. + 3. Auditor's free-form ``summary`` paragraph (if present). + 4. **If findings empty + verdict not in success set** — an explicit + "re-attempt" instruction. + 5. **Else** — Markdown summary table (one row per finding) followed + by a fenced JSON block with the full structured findings. + + The implementation does not know which verdicts are "success" — the + renderer is told via the ``success_verdicts`` constructor argument + so it can choose between (4) and (5) without hardcoding domain + vocabulary. + """ + + def __init__( + self, + *, + success_verdicts: Iterable[str] = ("success",), + ) -> None: + self._success = frozenset(success_verdicts) + + def render(self, audit: AuditReport) -> str: + parts: list[str] = [_HEADER] + parts.append(f"**Verdict**: `{audit.verdict}`") + if audit.confidence: + parts.append(f" (auditor confidence: {audit.confidence:.2f})") + parts.append("\n\n") + + if audit.summary.strip(): + parts.append("**Auditor summary**:\n\n") + parts.append(audit.summary.strip()) + parts.append("\n\n") + + is_terminal_success = audit.verdict in self._success + if not audit.findings and not is_terminal_success: + parts.append( + _REATTEMPT_MESSAGE.format(verdict=audit.verdict), + ) + parts.append("\n") + return "".join(parts) + + if audit.findings: + parts.append("**Findings**:\n\n") + parts.append(_TABLE_HEADER) + for i, f in enumerate(audit.findings, start=1): + parts.append(_render_table_row(i, f)) + parts.append("\n") + + parts.append("**Findings (structured)**:\n\n") + parts.append("```json\n") + parts.append( + json.dumps( + [f.to_dict() for f in audit.findings], + indent=2, + ensure_ascii=False, + ), + ) + parts.append("\n```\n") + + return "".join(parts) + + +def _render_table_row(idx: int, f: AuditFinding) -> str: + where = "" + if f.file: + where = f.file + if f.line is not None: + where = f"{f.file}:{f.line}" + return ( + f"| {idx} | {_md_cell(f.severity)} | {_md_cell(f.category)} " + f"| {_md_cell(where)} | {_md_cell(f.short_message)} " + f"| {_md_cell(f.suggested_action)} |\n" + ) + + +def _md_cell(text: str) -> str: + """Escape pipe characters and newlines so a single cell stays one row.""" + if not text: + return "" + return text.replace("|", "\\|").replace("\n", " ") diff --git a/agent_core/components/cycle/observers.py b/agent_core/components/cycle/observers.py new file mode 100644 index 0000000..afe065c --- /dev/null +++ b/agent_core/components/cycle/observers.py @@ -0,0 +1,156 @@ +"""RoundObserver — pluggable hooks for :class:`WriteAuditCycle`. + +Mirrors :class:`agent_core.loop_types.LoopObserver` but at cycle +granularity (rounds, not LLM turns). Lets callers inject behaviour +between rounds without subclassing the cycle: + +- best-so-far snapshotting / external upload +- mid-cycle metrics emission +- score-plateau early termination +- ATIF / external trace export +- domain-specific intervention (e.g. abort if 3 rounds saw the same + exception) + +The cycle never trusts an observer to do the right thing — every +hook call is wrapped in try/except, exceptions are logged and +swallowed so observer bugs cannot kill a long-running cycle. + +Hook ordering per round:: + + on_round_start(round_num, prev_audit) + ↓ + writer.generate() # may raise → synth audit; on_writer_done skipped + ↓ + on_writer_done(round_num, writer_out) → may abort + ↓ + output_check + auditor.verify() (or synth audit) + ↓ + on_audit_done(round_num, writer_out, audit) → may abort + +Outside the round loop: + + on_cycle_start(ctx) — once, before round 0 + on_cycle_end(output) — once, on every exit path (including + wall-clock / observer-abort / max-rounds) +""" +from __future__ import annotations + +from dataclasses import dataclass, field +from enum import StrEnum +from pathlib import Path +from typing import Protocol, runtime_checkable + +from agent_core.components.cycle.types import AuditReport, CycleOutput, WriterOutput + + +class RoundIntervention(StrEnum): + """Verdict an observer can return from a per-round hook. + + ``CONTINUE`` (the default for any hook returning ``None``) lets + the cycle proceed normally. ``ABORT`` causes the cycle to + terminate immediately with ``CycleOutput.success=False`` and + ``reason="observer_aborted"``. The current round's writer output + + audit (if any) are preserved in ``final_writer_output`` / + ``final_audit`` and the partial history is returned as usual. + """ + + CONTINUE = "continue" + ABORT = "abort" + + +@dataclass +class CycleContext: + """Read-only handle to cycle state, passed to ``on_cycle_start``. + + Observers should treat ``history_so_far`` as a *live view*: the + cycle keeps appending to the same list as rounds complete, so + later hooks can inspect prior rounds without the cycle re-passing + them. Observers should not mutate the list. + """ + + work_dir: Path + max_rounds: int + max_wall_seconds: float | None + history_so_far: list[tuple[WriterOutput, AuditReport]] = field( + default_factory=list[tuple[WriterOutput, AuditReport]], + ) + + +@runtime_checkable +class RoundObserver(Protocol): + """Pluggable per-round hook for :class:`WriteAuditCycle`. + + All methods are async. All methods are optional in spirit — the + cycle calls every hook on every observer, so implementations + that don't care about a particular event simply ``return None`` + or ``pass``. (Use :class:`BaseRoundObserver` for a default-noop + base class.) + + Methods returning :class:`RoundIntervention` may abort the cycle + by returning :attr:`RoundIntervention.ABORT`. Other return values + (including ``None``) are treated as :attr:`RoundIntervention.CONTINUE`. + """ + + async def on_cycle_start(self, ctx: CycleContext) -> None: + ... + + async def on_round_start( + self, + round_num: int, + prev_audit: AuditReport | None, + ) -> None: + ... + + async def on_writer_done( + self, + round_num: int, + writer_out: WriterOutput, + ) -> RoundIntervention | None: + ... + + async def on_audit_done( + self, + round_num: int, + writer_out: WriterOutput, + audit: AuditReport, + ) -> RoundIntervention | None: + ... + + async def on_cycle_end(self, output: CycleOutput) -> None: + ... + + +class BaseRoundObserver: + """Default-noop base class for :class:`RoundObserver` impls. + + Subclass and override only the hooks you care about. Implements + every Protocol method as ``return None``. + """ + + async def on_cycle_start(self, ctx: CycleContext) -> None: + return None + + async def on_round_start( + self, + round_num: int, + prev_audit: AuditReport | None, + ) -> None: + return None + + async def on_writer_done( + self, + round_num: int, + writer_out: WriterOutput, + ) -> RoundIntervention | None: + return None + + async def on_audit_done( + self, + round_num: int, + writer_out: WriterOutput, + audit: AuditReport, + ) -> RoundIntervention | None: + return None + + async def on_cycle_end(self, output: CycleOutput) -> None: + return None diff --git a/agent_core/components/cycle/observers_builtin.py b/agent_core/components/cycle/observers_builtin.py new file mode 100644 index 0000000..231562d --- /dev/null +++ b/agent_core/components/cycle/observers_builtin.py @@ -0,0 +1,232 @@ +"""Built-in :class:`RoundObserver` implementations. + +These are the most common cycle-extension patterns, shipped so callers +don't have to re-derive them. All are pure additions on top of +:class:`agent_core.components.cycle.observers.BaseRoundObserver` — feel free to +read them as templates for your own observers. +""" +from __future__ import annotations + +import asyncio +import inspect +import logging +from collections.abc import Awaitable, Callable +from typing import Any + +from agent_core.components.cycle.builders import select_best_attempt +from agent_core.components.cycle.observers import ( + BaseRoundObserver, + RoundIntervention, +) +from agent_core.components.cycle.types import AuditReport, WriterOutput + +logger = logging.getLogger(__name__) + + +# ── BestSoFarObserver ─────────────────────────────────────────────────────── + + +# A callback receiving (best_index_zero_based, best_writer_out, best_audit). +# May be sync or async. Errors are caught + logged; the cycle continues. +BestSoFarCallback = Callable[ + [int, WriterOutput, AuditReport], + Awaitable[None] | None, +] + + +class BestSoFarObserver(BaseRoundObserver): + """Invoke ``callback(best_idx, best_writer_out, best_audit)`` + after every audit, with the highest-scoring attempt seen so far. + + Score selection uses :func:`agent_core.components.cycle.select_best_attempt` + (highest ``audit.metadata['score']``, ties → latest). The + observer maintains its own append-only history; it does not read + cycle internals. + + Designed for the "still-running cycle should always have a usable + best-so-far artefact persisted somewhere" pattern — most directly + the MiroVerifier IMO-GVR ``_upload_content`` per-attempt upload + that protects against agent timeouts. Make ``callback`` upload + the proof to your sandbox / object store / DB. + + Synthetic audits (writer exceptions, output_missing) typically + have no ``metadata['score']`` and are treated as worse than any + scored attempt, so they don't poison the best-so-far selection. + """ + + def __init__(self, callback: BestSoFarCallback) -> None: + self._callback = callback + self._history: list[tuple[WriterOutput, AuditReport]] = [] + + async def on_audit_done( + self, + round_num: int, + writer_out: WriterOutput, + audit: AuditReport, + ) -> RoundIntervention | None: + self._history.append((writer_out, audit)) + try: + best_idx, best_w, best_a = select_best_attempt(self._history) + except ValueError: + return None + try: + ret = self._callback(best_idx, best_w, best_a) + if inspect.isawaitable(ret): + await ret + except Exception: + logger.exception( + "BestSoFarObserver callback raised on round %d", round_num, + ) + return None + + +# ── PlateauAbortObserver ──────────────────────────────────────────────────── + + +class PlateauAbortObserver(BaseRoundObserver): + """Abort the cycle when scores have plateaued. + + Triggers :class:`RoundIntervention.ABORT` when the best score in + the most recent ``plateau_rounds`` rounds is no greater than the + best score in any earlier round. Useful for cost control on + always-run-k workflows where the LLM has clearly converged and + further rounds won't help. + + Reads ``audit.metadata['score']``; rounds without a score + (synthetic audits, grader errors) are skipped — they neither + trigger plateau nor reset the counter. + + Parameters + ---------- + plateau_rounds + Number of consecutive trailing rounds that must show no + improvement over historical best to trigger abort. Default + ``3``. Must be ≥ 1. + min_rounds + Skip plateau check until at least this many scored rounds + have happened. Default ``plateau_rounds`` itself — that is, + the earliest possible abort is after exactly + ``plateau_rounds`` consecutive scored rounds with no + improvement. + """ + + def __init__( + self, + plateau_rounds: int = 3, + *, + min_rounds: int | None = None, + ) -> None: + if plateau_rounds < 1: + raise ValueError( + f"plateau_rounds must be >= 1, got {plateau_rounds!r}", + ) + self._plateau = plateau_rounds + self._min_rounds = ( + min_rounds if min_rounds is not None else plateau_rounds + ) + self._scores: list[float] = [] + + async def on_audit_done( + self, + round_num: int, + writer_out: WriterOutput, + audit: AuditReport, + ) -> RoundIntervention | None: + del round_num, writer_out + score = audit.metadata.get("score") + if not isinstance(score, (int, float)): + return None + self._scores.append(float(score)) + + if len(self._scores) < self._min_rounds: + return None + if len(self._scores) < self._plateau + 1: + # Need at least one earlier round to compare against. + return None + recent = self._scores[-self._plateau:] + earlier = self._scores[: -self._plateau] + if max(recent) <= max(earlier): + logger.info( + "PlateauAbortObserver: aborting; recent_max=%s earlier_max=%s", + max(recent), max(earlier), + ) + return RoundIntervention.ABORT + return None + + +# ── MetricsObserver (lightweight logging) ─────────────────────────────────── + + +MetricsCallback = Callable[[dict[str, Any]], Awaitable[None] | None] + + +class MetricsObserver(BaseRoundObserver): + """Emit a metrics dict per round. + + Each ``on_audit_done`` fires the callback with:: + + { + "round": , + "duration_seconds": , + "score": , + "verdict": , + "writer_chars": , + "audit_findings": , + } + + The default callback logs at INFO level. Pass a custom callback + to push to Prometheus / a metrics client / a JSON line file. + Sync or async callbacks both work; exceptions are logged + swallowed. + """ + + def __init__( + self, + callback: MetricsCallback | None = None, + ) -> None: + self._callback = callback + self._round_starts: dict[int, float] = {} + + async def on_round_start( + self, + round_num: int, + prev_audit: AuditReport | None, + ) -> None: + del prev_audit + loop = asyncio.get_event_loop() + self._round_starts[round_num] = loop.time() + + async def on_audit_done( + self, + round_num: int, + writer_out: WriterOutput, + audit: AuditReport, + ) -> RoundIntervention | None: + loop = asyncio.get_event_loop() + duration = loop.time() - self._round_starts.pop( + round_num, loop.time(), + ) + score_meta = audit.metadata.get("score") + metrics = { + "round": round_num, + "duration_seconds": duration, + "score": ( + float(score_meta) + if isinstance(score_meta, (int, float)) + else None + ), + "verdict": audit.verdict, + "writer_chars": len(writer_out.content), + "audit_findings": len(audit.findings), + } + if self._callback is None: + logger.info("cycle round metrics: %s", metrics) + else: + try: + ret = self._callback(metrics) + if inspect.isawaitable(ret): + await ret + except Exception: + logger.exception( + "MetricsObserver callback raised on round %d", round_num, + ) + return None diff --git a/agent_core/components/cycle/protocols.py b/agent_core/components/cycle/protocols.py new file mode 100644 index 0000000..4d8e7ba --- /dev/null +++ b/agent_core/components/cycle/protocols.py @@ -0,0 +1,65 @@ +"""Caller-side protocols for the WriteAuditCycle. + +The cycle's generator extension lives in +``agent_core.components.verifier.Generator``. The verifier extension +surface is the cycle-domain ``CycleAuditor`` Protocol defined here — +distinct from the unified ``components.verifier.Verifier`` Protocol +because cycle auditors operate on ``WriterOutput`` + ``round_num`` +and return ``AuditReport`` rather than the unified +``(subject, ctx) → Verdict`` shape. To plug a unified ``Verifier`` +into a cycle, wrap it via +:func:`agent_core.components.verifier.cycle_auditor_from_verifier`. + +The framework provides default session-backed implementations +(``SessionBackedWriter`` / ``SessionBackedAuditor`` / +``DefaultFeedbackRenderer``). Sophisticated callers can implement +``CycleAuditor`` directly — for example, an "auditor team" that fans +out to N sub-verifiers and aggregates their verdicts into one +:class:`AuditReport`. +""" +from __future__ import annotations + +from typing import Protocol, runtime_checkable + +from agent_core.components.cycle.types import AuditReport, WriterOutput + + +class CycleAuditor(Protocol): + """Cycle-domain verifier contract used by ``WriteAuditCycle``. + + Distinct from the unified ``components.verifier.Verifier`` Protocol: + cycle auditors receive the writer's output + round number and emit + an ``AuditReport`` (with verdict label, structured findings, raw + text). Use :func:`cycle_auditor_from_verifier` to bridge from a + unified ``Verifier``. + """ + + role_id: str + + async def verify( + self, + writer_output: WriterOutput, + round_num: int, + ) -> AuditReport: ... + + +@runtime_checkable +class FeedbackRenderer(Protocol): + """Turns an AuditReport into the prompt segment for the next generator round. + + The default renderer (:class:`DefaultFeedbackRenderer`) emits a + Markdown summary table + JSON detail block + the verifier's + free-form ``summary`` paragraph above. Callers can supply their own + renderer when they need a different format. + """ + + def render(self, audit: AuditReport) -> str: + """Render the audit as a feedback string. + + Implementations must not return empty strings or generic + placeholders like ``"please address some issues"`` — when the + audit has zero findings but a non-terminal verdict, return an + explicit ``"no specific findings, full re-attempt required"`` + or equivalent (FR6). + """ + ... diff --git a/agent_core/components/cycle/types.py b/agent_core/components/cycle/types.py new file mode 100644 index 0000000..81d3220 --- /dev/null +++ b/agent_core/components/cycle/types.py @@ -0,0 +1,274 @@ +"""Structured data model for the WriteAuditCycle. + +These types form the cycle's stable contract: + +- :class:`AuditFinding` — one structured diagnostic produced by an auditor. +- :class:`AuditReport` — full audit verdict + structured findings + free-form prose. +- :class:`WriterOutput` — what a writer round produced. +- :class:`CycleOutput` — final cycle result with full history. + +All four round-trip losslessly through :meth:`to_json` / :meth:`from_json`, +which is how the per-round audit trail is persisted on disk (FR4). + +Hybrid structured + free-form output (FR3): :class:`AuditReport` carries +both a list of structured ``findings`` (machine-readable, parsed by the +orchestrator to decide loop continuation) and a free-form ``summary`` +paragraph (human/LLM-readable, consumed by the writer to actually +understand the critique). LLMs are asked to emit JSON inside a fenced +block plus prose outside, which is more reliable than encoding prose +inside JSON. +""" +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass, field +from typing import Any, Literal + +Severity = Literal["error", "warning", "info"] + + +# ── AuditFinding ──────────────────────────────────────────────────────────── + + +@dataclass +class AuditFinding: + """One structured diagnostic produced by an auditor. + + ``category`` is free-form but conventionally drawn from a stable + vocabulary (e.g. ``"missing_section"``, ``"citation_invalid"``, + ``"writer_exception"``, ``"output_missing"``, ``"auditor_parse_failure"``). + Free-form is preferred over a closed enum because vocabularies vary + by domain (paper, code, dataset, plan). + + ``target_role`` is an optional hint about which writer-side role + should act on this finding when the cycle has multiple writers + (e.g. method-section writer vs. results-section writer). The cycle + itself does not route on this field; it is informational for the + feedback renderer. + + ``metadata`` is an extensible bag for caller-specific payload. + """ + + category: str + severity: Severity + short_message: str + detailed_message: str = "" + file: str | None = None + line: int | None = None + snippet: str = "" + suggested_action: str = "" + target_role: str | None = None + metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> AuditFinding: + return cls( + category=data["category"], + severity=data["severity"], + short_message=data["short_message"], + detailed_message=data.get("detailed_message", ""), + file=data.get("file"), + line=data.get("line"), + snippet=data.get("snippet", ""), + suggested_action=data.get("suggested_action", ""), + target_role=data.get("target_role"), + metadata=dict(data.get("metadata") or {}), + ) + + +# ── AuditReport ───────────────────────────────────────────────────────────── + + +@dataclass +class AuditReport: + """Full audit verdict for one cycle round. + + ``verdict`` is free-form. The cycle compares it against + ``terminal_verdicts`` (default ``{"success", "abandon"}``) to decide + whether to terminate. Common values: ``"success"`` (artifact passes), + ``"iterate"`` (writer should revise), ``"abandon"`` (unrecoverable). + + ``findings`` is the structured part — what the orchestrator reads. + ``summary`` is the free-form part — what the writer reads. + + ``confidence`` is the auditor's self-reported confidence in its own + verdict, in ``[0.0, 1.0]``. Optional but recommended; renderers may + surface it to downstream readers. + + ``raw_text`` is the auditor's full unparsed output, kept for + debugging and replay. The default renderer does not include it in + the next-round prompt — only the structured findings + summary. + """ + + verdict: str + findings: list[AuditFinding] = field(default_factory=list[AuditFinding]) + summary: str = "" + confidence: float = 0.0 + raw_text: str = "" + metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + + def to_dict(self) -> dict[str, Any]: + return { + "verdict": self.verdict, + "findings": [f.to_dict() for f in self.findings], + "summary": self.summary, + "confidence": self.confidence, + "raw_text": self.raw_text, + "metadata": dict(self.metadata), + } + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> AuditReport: + raw_findings: list[Any] = data.get("findings") or [] + return cls( + verdict=data["verdict"], + findings=[AuditFinding.from_dict(f) for f in raw_findings], + summary=data.get("summary", ""), + confidence=float(data.get("confidence", 0.0)), + raw_text=data.get("raw_text", ""), + metadata=dict(data.get("metadata") or {}), + ) + + def to_json(self, *, indent: int | None = 2) -> str: + return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False) + + @classmethod + def from_json(cls, text: str) -> AuditReport: + return cls.from_dict(json.loads(text)) + + +# ── WriterOutput ──────────────────────────────────────────────────────────── + + +@dataclass +class WriterOutput: + """What a writer round produced. + + ``content`` is the writer's final text (the LLM's last AI message). + ``files`` lists artifact paths the writer touched, relative to + ``work_dir`` — populated by the writer implementation, not by the + cycle. ``message_count`` is the number of messages in the writer's + session at the end of this round (informational; useful for trimmer + diagnostics). + + ``loop_result`` is the underlying ``AgentLoopResult`` when the + writer was run via ``run_agent_loop``. Optional — protocols-based + writers that don't use the agent loop may leave it ``None``. + """ + + content: str + files: list[str] = field(default_factory=list[str]) + message_count: int = 0 + metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + loop_result: Any = None + + def to_dict(self) -> dict[str, Any]: + # ``loop_result`` is intentionally dropped — it carries + # langchain message objects which are not JSON-serialisable. + # The audit trail is for the AuditReport; writer output is + # reconstructable from the writer's session messages. + return { + "content": self.content, + "files": list(self.files), + "message_count": self.message_count, + "metadata": dict(self.metadata), + } + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> WriterOutput: + return cls( + content=data["content"], + files=list(data.get("files") or []), + message_count=int(data.get("message_count", 0)), + metadata=dict(data.get("metadata") or {}), + loop_result=None, + ) + + +# ── CycleOutput ───────────────────────────────────────────────────────────── + + +@dataclass +class CycleOutput: + """Final cycle result. + + ``success`` is ``True`` when the final audit's verdict is in the + cycle's "success" set (typically ``{"success"}``); ``False`` for any + abandon / max-rounds-exhausted termination. + + ``rounds_used`` counts every writer attempt (including those whose + output_check failed and the auditor was skipped). + + ``history`` is the complete trail: one + ``(WriterOutput, AuditReport)`` pair per round, including synthetic + audits constructed by the cycle when the writer raised or + output_check failed. + + ``reason`` is a short string explaining termination, e.g. + ``"verdict=success"``, ``"verdict=abandon"``, + ``"max_rounds_exhausted"``. + """ + + success: bool + rounds_used: int + final_writer_output: WriterOutput | None + final_audit: AuditReport | None + history: list[tuple[WriterOutput, AuditReport]] = field( + default_factory=list[tuple[WriterOutput, AuditReport]], + ) + reason: str = "" + + def to_dict(self) -> dict[str, Any]: + return { + "success": self.success, + "rounds_used": self.rounds_used, + "final_writer_output": ( + self.final_writer_output.to_dict() + if self.final_writer_output is not None + else None + ), + "final_audit": ( + self.final_audit.to_dict() + if self.final_audit is not None + else None + ), + "history": [ + {"writer": w.to_dict(), "audit": a.to_dict()} + for w, a in self.history + ], + "reason": self.reason, + } + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> CycleOutput: + fwo_raw: Any = data.get("final_writer_output") + fa_raw: Any = data.get("final_audit") + raw_history: list[Any] = data.get("history") or [] + return cls( + success=bool(data["success"]), + rounds_used=int(data["rounds_used"]), + final_writer_output=( + WriterOutput.from_dict(fwo_raw) if fwo_raw else None + ), + final_audit=( + AuditReport.from_dict(fa_raw) if fa_raw else None + ), + history=[ + ( + WriterOutput.from_dict(item["writer"]), + AuditReport.from_dict(item["audit"]), + ) + for item in raw_history + ], + reason=data.get("reason", ""), + ) + + def to_json(self, *, indent: int | None = 2) -> str: + return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False) + + @classmethod + def from_json(cls, text: str) -> CycleOutput: + return cls.from_dict(json.loads(text)) diff --git a/agent_core/components/cycle/write_audit_cycle.py b/agent_core/components/cycle/write_audit_cycle.py new file mode 100644 index 0000000..68f6128 --- /dev/null +++ b/agent_core/components/cycle/write_audit_cycle.py @@ -0,0 +1,484 @@ +"""WriteAuditCycle — iterate-until-pass loop over Writer + Auditor. + +Pseudocode (FR2):: + + for round_num in 0 .. max_rounds: + feedback_md = renderer.render(prev_audit) if prev_audit else "" + try: + writer_out = await writer.generate(prev_audit, round_num, feedback_md) + except Exception: + prev_audit = synth_audit_from_exception(...) + persist(prev_audit); continue + missing = output_check(work_dir) + if missing: + prev_audit = synth_audit_from_missing(missing) + persist(prev_audit); continue + audit = await auditor.verify(writer_out, round_num) + persist(audit) + if audit.verdict in terminal_set: + return CycleOutput(success=verdict_is_success, ...) + prev_audit = audit + return CycleOutput(success=False, reason="max_rounds_exhausted", ...) + +Invariants: + +- Every round persists exactly one ``audit_round_N.json`` to + ``work_dir`` — including synthetic audits constructed by the cycle + for output_check failures and writer exceptions (FR4). +- Writer exceptions are **never** propagated to the caller. They become + ``writer_exception`` findings and the loop continues (FR10). The only + terminations are: terminal verdict, max_rounds exhaustion, or + cancellation from outside. +- Auditor is **not invoked** when ``output_check`` reports missing + artifacts — auditing vacuous output is a known false-positive source + (FR5). The cycle synthesizes a structured ``output_missing`` audit + itself. +""" +from __future__ import annotations + +import logging +import time +import traceback +from collections import deque +from collections.abc import Callable, Iterable +from pathlib import Path + +from agent_core.components.cycle.default_renderer import DefaultFeedbackRenderer +from agent_core.components.cycle.observers import ( + CycleContext, + RoundIntervention, + RoundObserver, +) +from agent_core.components.cycle.protocols import ( + CycleAuditor, + FeedbackRenderer, +) +from agent_core.components.cycle.types import ( + AuditFinding, + AuditReport, + CycleOutput, + WriterOutput, +) +from agent_core.components.verifier.protocols import Generator + +logger = logging.getLogger(__name__) + + +# Default verdict set that *terminates* the loop. ``success`` ends with +# ``CycleOutput.success=True``; ``abandon`` ends with ``False``. Callers +# override via ``terminal_verdicts`` + ``success_verdicts`` constructor +# arguments. +_DEFAULT_TERMINAL: frozenset[str] = frozenset({"success", "abandon"}) +_DEFAULT_SUCCESS: frozenset[str] = frozenset({"success"}) + + +class WriteAuditCycle: + """Generic iterate-until-pass orchestrator. + + Parameters + ---------- + writer + Caller-supplied :class:`Writer` implementation. + auditor + Caller-supplied :class:`Auditor` implementation. + work_dir + Filesystem root for artifacts + per-round audit JSON. Created if + absent. + output_check + ``(work_dir) -> list[str]`` returning the list of *missing* + required artifacts. Empty list means "all required artifacts + present, proceed to audit". Required. + feedback_renderer + Defaults to :class:`DefaultFeedbackRenderer`. + max_rounds + Hard upper bound on writer attempts. Default 10. + terminal_verdicts + Set of verdict strings that terminate the loop. Default + ``{"success", "abandon"}``. Free-form strings — see PRD §10.2. + success_verdicts + Subset of ``terminal_verdicts`` that count as ``success=True``. + Default ``{"success"}``. + audit_path_template + Filename pattern for per-round persistence. Default + ``"audit_round_{round}.json"``. + max_wall_seconds + Optional wall-clock budget (seconds). When set, the cycle + checks remaining time **between rounds** (never mid-round — + protects partial work). If the next-round estimate + (``wall_clock_safety_ratio × max(recent durations)``) exceeds + the remaining budget, the cycle exits cleanly with + ``reason="wall_clock_exhausted"``. Default ``None`` (no + budget). Mirrors MiroVerifier IMO-GVR's deadline behaviour. + wall_clock_safety_ratio + Multiplier on the maximum of recent round durations when + estimating whether the next round will fit in the remaining + budget. Default ``1.1``. Must be ≥ 1.0. + observers + Optional list of :class:`RoundObserver` instances to receive + per-round hooks (start/end of cycle, after writer, after + audit). See :mod:`agent_core.components.cycle.observers`. Observer + exceptions are caught + logged so a buggy observer can never + kill the cycle. Observers may return + :attr:`RoundIntervention.ABORT` from ``on_writer_done`` / + ``on_audit_done`` to terminate the cycle early with + ``reason="observer_aborted"``. + """ + + def __init__( + self, + *, + writer: Generator, + auditor: CycleAuditor, + work_dir: Path | str, + output_check: Callable[[Path], list[str]], + feedback_renderer: FeedbackRenderer | None = None, + max_rounds: int = 10, + terminal_verdicts: Iterable[str] = _DEFAULT_TERMINAL, + success_verdicts: Iterable[str] = _DEFAULT_SUCCESS, + audit_path_template: str = "audit_round_{round}.json", + max_wall_seconds: float | None = None, + wall_clock_safety_ratio: float = 1.1, + observers: list[RoundObserver] | None = None, + ) -> None: + if max_rounds < 1: + raise ValueError( + f"max_rounds must be >= 1, got {max_rounds!r}", + ) + if max_wall_seconds is not None and max_wall_seconds <= 0: + raise ValueError( + f"max_wall_seconds must be > 0 when set, " + f"got {max_wall_seconds!r}", + ) + if wall_clock_safety_ratio < 1.0: + raise ValueError( + f"wall_clock_safety_ratio must be >= 1.0, " + f"got {wall_clock_safety_ratio!r}", + ) + self.writer = writer + self.auditor = auditor + self.work_dir = Path(work_dir) + self.output_check = output_check + self.renderer = feedback_renderer or DefaultFeedbackRenderer( + success_verdicts=success_verdicts, + ) + self.max_rounds = max_rounds + self.terminal_verdicts = frozenset(terminal_verdicts) + self.success_verdicts = frozenset(success_verdicts) + self.audit_path_template = audit_path_template + self.max_wall_seconds = max_wall_seconds + self.wall_clock_safety_ratio = wall_clock_safety_ratio + self._observers: list[RoundObserver] = list(observers or []) + if not self.success_verdicts.issubset(self.terminal_verdicts): + raise ValueError( + "success_verdicts must be a subset of terminal_verdicts; " + f"got success={set(self.success_verdicts)!r} " + f"terminal={set(self.terminal_verdicts)!r}", + ) + + async def run(self) -> CycleOutput: + """Run the cycle until terminal verdict, max_rounds exhaustion, + wall-clock budget exhaustion, or observer-driven abort. + + Never raises — writer exceptions are captured as + ``writer_exception`` findings (FR10) and observer exceptions + are caught + logged. + """ + self.work_dir.mkdir(parents=True, exist_ok=True) + prev_audit: AuditReport | None = None + history: list[tuple[WriterOutput, AuditReport]] = [] + last_writer_output: WriterOutput | None = None + + cycle_start = time.monotonic() + round_durations: deque[float] = deque(maxlen=3) + + ctx = CycleContext( + work_dir=self.work_dir, + max_rounds=self.max_rounds, + max_wall_seconds=self.max_wall_seconds, + history_so_far=history, # live view + ) + await self._fire("on_cycle_start", ctx) + + for round_num in range(self.max_rounds): + # Wall-clock check (between rounds; never mid-round). + if ( + self.max_wall_seconds is not None + and round_durations + ): + elapsed = time.monotonic() - cycle_start + remaining = self.max_wall_seconds - elapsed + est_next = ( + max(round_durations) * self.wall_clock_safety_ratio + ) + if remaining < est_next: + logger.info( + "Wall-clock budget exhausted: remaining=%.2fs " + "est_next=%.2fs (rounds_used=%d)", + remaining, est_next, round_num, + ) + return await self._finalize( + success=False, + rounds_used=round_num, + final_writer_output=last_writer_output, + final_audit=prev_audit, + history=history, + reason="wall_clock_exhausted", + ) + + round_t0 = time.monotonic() + await self._fire("on_round_start", round_num, prev_audit) + + feedback_md = ( + self.renderer.render(prev_audit) if prev_audit else "" + ) + + try: + writer_out = await self.writer.generate( + prev_audit, round_num, feedback_md, + ) + except Exception as exc: + logger.warning( + "Writer raised on round %d: %s", round_num, exc, + ) + synth = _synth_writer_exception_audit(exc) + self._persist(synth, round_num) + empty_out = WriterOutput(content="") + history.append((empty_out, synth)) + round_durations.append(time.monotonic() - round_t0) + if ( + await self._fire_intervention( + "on_audit_done", round_num, empty_out, synth, + ) + == RoundIntervention.ABORT + ): + return await self._finalize( + success=False, + rounds_used=round_num + 1, + final_writer_output=empty_out, + final_audit=synth, + history=history, + reason="observer_aborted", + ) + prev_audit = synth + last_writer_output = empty_out + continue + + if ( + await self._fire_intervention( + "on_writer_done", round_num, writer_out, + ) + == RoundIntervention.ABORT + ): + return await self._finalize( + success=False, + rounds_used=round_num + 1, + final_writer_output=writer_out, + final_audit=prev_audit, + history=history, + reason="observer_aborted", + ) + + missing = self.output_check(self.work_dir) + if missing: + logger.info( + "Round %d: output_check missing=%s — auditor skipped", + round_num, missing, + ) + synth = _synth_output_missing_audit(missing) + self._persist(synth, round_num) + history.append((writer_out, synth)) + round_durations.append(time.monotonic() - round_t0) + if ( + await self._fire_intervention( + "on_audit_done", round_num, writer_out, synth, + ) + == RoundIntervention.ABORT + ): + return await self._finalize( + success=False, + rounds_used=round_num + 1, + final_writer_output=writer_out, + final_audit=synth, + history=history, + reason="observer_aborted", + ) + prev_audit = synth + last_writer_output = writer_out + continue + + audit = await self.auditor.verify(writer_out, round_num) + self._persist(audit, round_num) + history.append((writer_out, audit)) + round_durations.append(time.monotonic() - round_t0) + last_writer_output = writer_out + + if ( + await self._fire_intervention( + "on_audit_done", round_num, writer_out, audit, + ) + == RoundIntervention.ABORT + ): + return await self._finalize( + success=False, + rounds_used=round_num + 1, + final_writer_output=writer_out, + final_audit=audit, + history=history, + reason="observer_aborted", + ) + + if audit.verdict in self.terminal_verdicts: + success = audit.verdict in self.success_verdicts + return await self._finalize( + success=success, + rounds_used=round_num + 1, + final_writer_output=writer_out, + final_audit=audit, + history=history, + reason=f"verdict={audit.verdict}", + ) + + prev_audit = audit + + # Loop fell through — max_rounds exhausted with no terminal + # verdict. last_writer_output / prev_audit reflect the final + # round's outcome. + return await self._finalize( + success=False, + rounds_used=self.max_rounds, + final_writer_output=last_writer_output, + final_audit=prev_audit, + history=history, + reason="max_rounds_exhausted", + ) + + async def _finalize( + self, + *, + success: bool, + rounds_used: int, + final_writer_output: WriterOutput | None, + final_audit: AuditReport | None, + history: list[tuple[WriterOutput, AuditReport]], + reason: str, + ) -> CycleOutput: + """Build the CycleOutput, fire on_cycle_end, return.""" + output = CycleOutput( + success=success, + rounds_used=rounds_used, + final_writer_output=final_writer_output, + final_audit=final_audit, + history=history, + reason=reason, + ) + await self._fire("on_cycle_end", output) + return output + + async def _fire(self, hook_name: str, *args: object) -> None: + """Fire a no-return hook on every observer; isolate errors.""" + for obs in self._observers: + try: + await getattr(obs, hook_name)(*args) + except Exception: + logger.exception( + "Observer %r raised in %s", obs, hook_name, + ) + + async def _fire_intervention( + self, hook_name: str, *args: object, + ) -> RoundIntervention: + """Fire an intervention hook; ``ABORT`` from any observer wins. + + Errors from individual observers are isolated; the hook is + still called on every other observer (so abort decisions are + not silently masked by an unrelated bug). + """ + result = RoundIntervention.CONTINUE + for obs in self._observers: + try: + ret = await getattr(obs, hook_name)(*args) + if ret == RoundIntervention.ABORT: + result = RoundIntervention.ABORT + except Exception: + logger.exception( + "Observer %r raised in %s", obs, hook_name, + ) + return result + + def _persist(self, audit: AuditReport, round_num: int) -> None: + """Write ``audit_round_{round}.json`` under work_dir. + + Failures are logged but never raised — losing a trail entry is + annoying but must not abort the cycle. + """ + path = self.work_dir / self.audit_path_template.format( + round=round_num, + ) + try: + path.write_text(audit.to_json(), encoding="utf-8") + except Exception: + logger.exception( + "Failed to persist audit trail to %s", path, + ) + + +# ── Synthetic audits constructed by the cycle ─────────────────────────────── + + +def _synth_output_missing_audit(missing: list[str]) -> AuditReport: + """Build a synthetic AuditReport for a failed output_check (FR5).""" + findings = [ + AuditFinding( + category="output_missing", + severity="error", + short_message=f"Required artifact missing: {path}", + detailed_message=( + f"The cycle's output_check reported {path!r} as missing " + f"after the writer round. Auditor was not invoked. " + f"Re-attempt the round and ensure this artifact is " + f"produced." + ), + file=path, + suggested_action=f"Produce the artifact at {path!r}.", + ) + for path in missing + ] + return AuditReport( + verdict="iterate", + findings=findings, + summary=( + f"Required output artifacts are missing " + f"({len(missing)}). The auditor was not invoked because " + f"auditing missing artifacts produces unreliable verdicts." + ), + confidence=1.0, + metadata={"synthetic": True, "reason": "output_missing"}, + ) + + +def _synth_writer_exception_audit(exc: Exception) -> AuditReport: + """Build a synthetic AuditReport for a writer exception (FR10).""" + tb = traceback.format_exc() + return AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="writer_exception", + severity="error", + short_message=( + f"{type(exc).__name__}: {str(exc)[:200]}" + ), + detailed_message=tb, + suggested_action=( + "Inspect the writer-side error and re-attempt " + "the round. The cycle will continue regardless." + ), + ), + ], + summary=( + f"The writer raised {type(exc).__name__} on this round. " + f"The cycle continues to the next round." + ), + confidence=1.0, + metadata={"synthetic": True, "reason": "writer_exception"}, + ) diff --git a/agent_core/components/observers/conclude_phase_observer.py b/agent_core/components/observers/conclude_phase_observer.py new file mode 100644 index 0000000..12aa962 --- /dev/null +++ b/agent_core/components/observers/conclude_phase_observer.py @@ -0,0 +1,98 @@ +"""Nudges an agent session toward emitting its final structured output. + +The well-known failure mode this guards against: an agent (typically an +auditor or report writer) keeps probing tool calls until ``max_turns`` is +hit, then exits with an empty final content because it never bothered to +write the structured output its task requires. This is especially common +when the auditor has a non-trivial tool budget — it explores until the +loop kills it. + +The observer fires once when the session crosses ``conclude_ratio`` of +its ``max_turns`` budget (default 0.8, i.e. last 20% of turns). The +injected message instructs the LLM to stop further investigation and +emit its final structured output now. Subsequent turns do not re-fire, +so the message is not repeated. + +Independent of any specific cycle / orchestration: any session whose +terminal output is mandatory benefits — auditors, report writers, +summarizers, finalize-then-stop solvers. + +Usage:: + + observer = ConcludePhaseObserver(conclude_ratio=0.8) + await bus.submit_task_to_session( + session_id, prompt, observers=[observer, ...], + ) +""" +from __future__ import annotations + +import logging + +from agent_core.loop_types import ( + BaseObserver, + Intervention, + TurnContext, +) + +logger = logging.getLogger(__name__) + + +_DEFAULT_MESSAGE = ( + "[conclude phase] You are nearing the turn budget for this task. " + "Stop further investigation and emit your final structured output " + "now (the JSON block your task requires). Use only the information " + "you already have. Further tool calls will be cut off." +) + + +class ConcludePhaseObserver(BaseObserver): + """Critical observer that nudges the agent to conclude before turn exhaustion. + + Parameters + ---------- + conclude_ratio + Fraction of ``max_turns`` at which the nudge fires. Must be in + ``(0.0, 1.0]``. Default 0.8 means "fire once in the last 20% of + the turn budget". + message + Override the injected message text. Default is a generic + instruction to emit the final structured output. Override when + the session's required output format is specific (e.g. "emit your + AuditReport JSON now"). + """ + + critical = True + + def __init__( + self, + conclude_ratio: float = 0.8, + message: str | None = None, + ) -> None: + if not 0.0 < conclude_ratio <= 1.0: + raise ValueError( + "conclude_ratio must be in (0.0, 1.0], got " + f"{conclude_ratio!r}", + ) + self._ratio = conclude_ratio + self._message = message if message is not None else _DEFAULT_MESSAGE + self._fired = False + + async def on_turn_end(self, ctx: TurnContext) -> Intervention | None: + if self._fired: + return None + if ctx.max_turns <= 0: + return None + progress = ctx.turn / ctx.max_turns + if progress < self._ratio: + return None + self._fired = True + logger.info( + "ConcludePhaseObserver firing at turn=%d/%d (progress=%.2f, " + "ratio=%.2f, role=%s)", + ctx.turn, + ctx.max_turns, + progress, + self._ratio, + ctx.role_id, + ) + return Intervention(inject_messages=[self._message]) diff --git a/agent_core/components/verifier/__init__.py b/agent_core/components/verifier/__init__.py new file mode 100644 index 0000000..55dd0a1 --- /dev/null +++ b/agent_core/components/verifier/__init__.py @@ -0,0 +1,53 @@ +"""Unified verifier protocol + composers. + +``(subject, ctx) → Verdict`` covers LLM-rubric, search-grounded, +multi-phase, debate, code-exec, and cross-trajectory consensus verify +forms. Six composers (Pipeline / Ensemble / Fallback / Cascade / +Parallel / ConsensusVerifier) each implement the Verifier protocol so +they nest freely. + +This package ships only the contract and the composition primitives. +Concrete verifiers belong to the consuming product, beside the workflow +whose subject they judge: the scoring rubric, the domain vocabulary, and +the oracle are all product decisions. See ``docs/cycle-verifier-boundary.md``. +""" + +from agent_core.components.verifier._compat import ( + audit_report_from_verdict, + cycle_auditor_from_verifier, + verdict_from_audit_report, +) +from agent_core.components.verifier.composers import ( + Cascade, + ConsensusVerifier, + Ensemble, + Fallback, + Parallel, + Pipeline, +) +from agent_core.components.verifier.protocols import ( + Finding, + Generator, + GroundTruth, + Verdict, + Verifier, + VerifierContext, +) + +__all__ = [ + "Cascade", + "ConsensusVerifier", + "Ensemble", + "Fallback", + "Finding", + "Generator", + "GroundTruth", + "Parallel", + "Pipeline", + "Verdict", + "Verifier", + "VerifierContext", + "audit_report_from_verdict", + "cycle_auditor_from_verifier", + "verdict_from_audit_report", +] diff --git a/agent_core/components/verifier/_compat.py b/agent_core/components/verifier/_compat.py new file mode 100644 index 0000000..6f60992 --- /dev/null +++ b/agent_core/components/verifier/_compat.py @@ -0,0 +1,154 @@ +"""Bridge between cycle-domain ``AuditReport`` and unified ``Verdict``. + +The cycle GVR engine retains its richer ``AuditReport`` shape (verdict +label, structured findings with audit-domain fields, summary, raw text). +Verifier composers use the unified ``Verdict`` shape. These helpers +convert between the two at the boundary — workflow authors only call +them when stitching a verifier composer into a cycle engine. +""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import TYPE_CHECKING, cast, get_args + +from agent_core.components.cycle.protocols import CycleAuditor +from agent_core.components.cycle.types import Severity, WriterOutput +from agent_core.components.verifier.protocols import ( + Finding, + Verdict, + Verifier, + VerifierContext, +) + +if TYPE_CHECKING: + from agent_core.components.cycle.types import AuditReport + +# Must match ``cycle/write_audit_cycle.py`` ``_DEFAULT_SUCCESS`` so that a +# round-tripped Verdict whose ``passed=True`` becomes a verdict label the +# cycle treats as terminal. +_TERMINAL_OK_LABELS: frozenset[str] = frozenset({"success"}) +# Non-terminal label so a failed Verdict re-enters the cycle loop instead +# of being treated as terminal-abandon. +_DEFAULT_FAIL_LABEL = "iterate" +_VALID_SEVERITIES: frozenset[str] = frozenset(get_args(Severity)) + + +def verdict_from_audit_report(report: AuditReport) -> Verdict: + """Convert a cycle ``AuditReport`` into a unified ``Verdict``.""" + return Verdict( + passed=report.verdict in _TERMINAL_OK_LABELS, + score=report.confidence, + findings=[ + Finding( + severity=str(f.severity), + message=f.short_message, + location=f.file, + ) + for f in report.findings + ], + reasoning=report.summary, + metadata={ + **report.metadata, + "_cycle_verdict": report.verdict, + "_cycle_raw_text": report.raw_text, + }, + ) + + +def audit_report_from_verdict( + verdict: Verdict, + *, + default_fail_label: str = _DEFAULT_FAIL_LABEL, +) -> AuditReport: + """Convert a ``Verdict`` back into a cycle ``AuditReport``.""" + from agent_core.components.cycle.types import AuditFinding, AuditReport + + label = verdict.metadata.get("_cycle_verdict") + if label is None: + label = "success" if verdict.passed else default_fail_label + raw = verdict.metadata.get("_cycle_raw_text", "") + leftover_metadata = { + k: v + for k, v in verdict.metadata.items() + if k not in {"_cycle_verdict", "_cycle_raw_text"} + } + return AuditReport( + verdict=label, + findings=[ + AuditFinding( + category=str(f.severity), + severity=_coerce_severity(f.severity), + short_message=f.message, + file=f.location, + ) + for f in verdict.findings + ], + summary=verdict.reasoning, + confidence=verdict.score or 0.0, + raw_text=raw, + metadata=leftover_metadata, + ) + + +def _coerce_severity(value: str) -> Severity: + if value in _VALID_SEVERITIES: + return cast(Severity, value) + return "info" + + +def _default_ctx_factory(round_num: int) -> VerifierContext: + return VerifierContext( + is_runtime=True, + metadata={"round_num": round_num}, + ) + + +class _CycleAuditorAdapter: + """Wrap a unified ``Verifier`` to satisfy the cycle ``CycleAuditor``. + + Translates the cycle's ``(WriterOutput, round_num) → AuditReport`` + call shape into ``(subject, ctx) → Verdict`` and converts the + returned :class:`Verdict` back via :func:`audit_report_from_verdict`. + """ + + def __init__( + self, + verifier: Verifier, + ctx_factory: Callable[[int], VerifierContext], + role_id: str, + ) -> None: + self._inner = verifier + self._ctx_factory = ctx_factory + self.role_id = role_id + + async def verify( + self, + writer_output: WriterOutput, + round_num: int, + ) -> AuditReport: + ctx = self._ctx_factory(round_num) + verdict = await self._inner.verify(writer_output, ctx) + return audit_report_from_verdict(verdict) + + +def cycle_auditor_from_verifier( + verifier: Verifier, + *, + ctx_factory: Callable[[int], VerifierContext] | None = None, + role_id: str | None = None, +) -> CycleAuditor: + """Adapt a unified :class:`Verifier` to the cycle ``CycleAuditor`` shape. + + The returned object structurally satisfies + :class:`agent_core.components.cycle.protocols.CycleAuditor` and can + be passed directly as ``WriteAuditCycle(auditor=...)``. + + ``ctx_factory`` builds the :class:`VerifierContext` for each round + (defaults to ``is_runtime=True`` with ``round_num`` in metadata). + """ + return _CycleAuditorAdapter( + verifier, + ctx_factory or _default_ctx_factory, + role_id or getattr(verifier, "role_id", "verifier"), + ) diff --git a/agent_core/components/verifier/composers.py b/agent_core/components/verifier/composers.py new file mode 100644 index 0000000..f9cd05d --- /dev/null +++ b/agent_core/components/verifier/composers.py @@ -0,0 +1,319 @@ +"""Composer verifiers — each implements the Verifier protocol so they nest. + +Six composition patterns: + +- ``Pipeline`` short-circuits on first ``passed=False`` +- ``Ensemble`` runs N verifiers in parallel and aggregates scores/passes +- ``Fallback`` primary fail → backups try in order +- ``Cascade`` cheap verifier with high confidence skips the expensive one +- ``Parallel`` runs N verifiers in parallel without aggregation, + all results stashed in ``sub_verdicts`` +- ``ConsensusVerifier`` cross-trajectory consensus — majority vote on a + metadata key (e.g. extracted answer) +""" + +from __future__ import annotations + +import asyncio +import collections +import statistics +from typing import Any + +from agent_core.components.verifier.protocols import ( + Finding, + Verdict, + Verifier, + VerifierContext, +) + +_AGGREGATORS = {"majority", "median", "min", "mean", "max"} + + +class Pipeline: + """Run verifiers sequentially; return on first ``passed=False``.""" + + def __init__(self, *verifiers: Verifier, role_id: str = "pipeline") -> None: + if not verifiers: + raise ValueError("Pipeline requires at least one verifier") + self._verifiers = verifiers + self.role_id = role_id + + async def verify(self, subject: Any, ctx: VerifierContext) -> Verdict: + sub: list[Verdict] = [] + for v in self._verifiers: + r = await v.verify(subject, ctx) + sub.append(r) + if not r.passed: + return Verdict( + passed=False, + sub_verdicts=sub, + reasoning=f"short-circuited at {v.role_id}: {r.reasoning}", + ) + return Verdict(passed=True, sub_verdicts=sub) + + +class Ensemble: + """Run N verifiers in parallel and aggregate.""" + + def __init__( + self, + *verifiers: Verifier, + aggregator: str = "majority", + role_id: str = "ensemble", + ) -> None: + if not verifiers: + raise ValueError("Ensemble requires at least one verifier") + if aggregator not in _AGGREGATORS: + raise ValueError( + f"unknown aggregator {aggregator!r}; expected one of {_AGGREGATORS}" + ) + self._verifiers = verifiers + self._aggregator = aggregator + self.role_id = role_id + + async def verify(self, subject: Any, ctx: VerifierContext) -> Verdict: + results = await asyncio.gather( + *(v.verify(subject, ctx) for v in self._verifiers) + ) + passes = [r.passed for r in results] + scores = [r.score for r in results if r.score is not None] + + if self._aggregator == "majority": + passed = passes.count(True) > len(passes) / 2 + else: + passed = all(passes) + + score: float | None + if not scores: + score = None + elif self._aggregator == "median": + score = statistics.median(scores) + elif self._aggregator == "min": + score = min(scores) + elif self._aggregator == "max": + score = max(scores) + elif self._aggregator == "mean": + score = statistics.fmean(scores) + else: + # ``majority`` aggregates pass/fail; use median for score. + score = statistics.median(scores) + + return Verdict( + score=score, + passed=passed, + sub_verdicts=list(results), + reasoning=f"ensemble({self._aggregator}, n={len(results)})", + ) + + +class Fallback: + """Try primary; if it fails, try backups in order.""" + + def __init__( + self, + primary: Verifier, + *backups: Verifier, + role_id: str = "fallback", + ) -> None: + self._primary = primary + self._backups = backups + self.role_id = role_id + + async def verify(self, subject: Any, ctx: VerifierContext) -> Verdict: + sub: list[Verdict] = [] + first = await self._primary.verify(subject, ctx) + sub.append(first) + if first.passed: + return Verdict(passed=True, sub_verdicts=sub, score=first.score) + + for backup in self._backups: + r = await backup.verify(subject, ctx) + sub.append(r) + if r.passed: + return Verdict( + passed=True, + sub_verdicts=sub, + score=r.score, + reasoning=f"primary failed; recovered via {backup.role_id}", + ) + + return Verdict( + passed=False, + sub_verdicts=sub, + reasoning="primary and all backups failed", + ) + + +class Cascade: + """Cheap-then-expensive: skip expensive when cheap is confident.""" + + def __init__( + self, + cheap: Verifier, + expensive: Verifier, + *, + confidence_threshold: float = 0.9, + role_id: str = "cascade", + ) -> None: + if not 0.0 <= confidence_threshold <= 1.0: + raise ValueError("confidence_threshold must be in [0, 1]") + self._cheap = cheap + self._expensive = expensive + self._threshold = confidence_threshold + self.role_id = role_id + + async def verify(self, subject: Any, ctx: VerifierContext) -> Verdict: + cheap_v = await self._cheap.verify(subject, ctx) + if cheap_v.score is not None and cheap_v.score >= self._threshold: + return Verdict( + passed=cheap_v.passed, + score=cheap_v.score, + sub_verdicts=[cheap_v], + reasoning=f"cheap verifier confident (≥{self._threshold})", + ) + expensive_v = await self._expensive.verify(subject, ctx) + return Verdict( + passed=expensive_v.passed, + score=expensive_v.score, + sub_verdicts=[cheap_v, expensive_v], + reasoning=f"escalated to {self._expensive.role_id}", + ) + + +class Parallel: + """Run N verifiers in parallel; collect all results without aggregating.""" + + def __init__(self, *verifiers: Verifier, role_id: str = "parallel") -> None: + if not verifiers: + raise ValueError("Parallel requires at least one verifier") + self._verifiers = verifiers + self.role_id = role_id + + async def verify(self, subject: Any, ctx: VerifierContext) -> Verdict: + results = await asyncio.gather( + *(v.verify(subject, ctx) for v in self._verifiers) + ) + return Verdict( + passed=all(r.passed for r in results), + sub_verdicts=list(results), + reasoning=f"parallel(n={len(results)})", + ) + + +class ConsensusVerifier: + """Run N verifiers in parallel and aggregate by majority answer. + + Cross-trajectory consensus primitive: each sub-verifier surfaces a + candidate answer in ``Verdict.metadata[answer_key]``; the wrapper + counts answers and passes when a strict majority of valid voters + agree (``top_count * 2 > n_valid``). + + Resilient to sub-verifier exceptions — failures are excluded from + the vote and reported via ``metadata.n_succeeded`` plus a warning + Finding. Returns ``passed=False`` if every sub-verifier failed or + no sub-verdict carried the answer key. + + Note: ``ConsensusVerifier`` and + :class:`agent_core.components.cycle.builders.MajorityVoteAuditor` + are **not equivalent** despite both being "majority over N parallel + runs": + + - ``ConsensusVerifier`` votes on ``Verdict.metadata[answer_key]`` + (extracted candidate answer) — N sub-verifiers run on the **same** + subject and agree on what the answer is. + - ``MajorityVoteAuditor`` votes on ``AuditReport.verdict`` label + (pass/fail/iterate) — N grader calls run on the **same** writer + output and agree on the verdict. + + Different vote fields, different intents. Both legitimately coexist + and pick the one that matches your topology. (Earlier docstrings + suggested one collapses into the other — that was wrong; see + ``internal-docs/designs/2026-04-30-verifier-first-class-lite.md`` + §0.1.) + """ + + def __init__( + self, + *verifiers: Verifier, + answer_key: str = "answer", + role_id: str = "consensus", + ) -> None: + if not verifiers: + raise ValueError("ConsensusVerifier requires at least one verifier") + self._verifiers = verifiers + self._answer_key = answer_key + self.role_id = role_id + + async def verify(self, subject: Any, ctx: VerifierContext) -> Verdict: + n_total = len(self._verifiers) + results = await asyncio.gather( + *(v.verify(subject, ctx) for v in self._verifiers), + return_exceptions=True, + ) + valid: list[Verdict] = [r for r in results if isinstance(r, Verdict)] + n_valid = len(valid) + + findings: list[Finding] = [] + if n_valid < n_total: + findings.append( + Finding( + severity="warning", + message=f"{n_total - n_valid}/{n_total} sub-verifiers failed", + ) + ) + + if not valid: + return Verdict( + passed=False, + sub_verdicts=[], + findings=findings, + reasoning="all sub-verifiers failed", + metadata={"n_total": n_total, "n_succeeded": 0}, + ) + + answers = [v.metadata.get(self._answer_key) for v in valid] + non_none = [a for a in answers if a is not None] + + if not non_none: + findings.append( + Finding( + severity="warning", + message=( + f"no sub-verdict carried metadata[{self._answer_key!r}]" + ), + ) + ) + return Verdict( + passed=False, + sub_verdicts=valid, + findings=findings, + reasoning=( + f"consensus failed: no answers in " + f"metadata[{self._answer_key!r}]" + ), + metadata={ + "n_total": n_total, + "n_succeeded": n_valid, + }, + ) + + counter = collections.Counter(non_none) + consensus_answer, top_count = counter.most_common(1)[0] + agreement = top_count / n_valid + passed = top_count * 2 > n_valid + + return Verdict( + passed=passed, + sub_verdicts=valid, + findings=findings, + reasoning=( + f"consensus({n_valid}/{n_total} valid, " + f"agreement={agreement:.0%}): {consensus_answer!r}" + ), + metadata={ + "n_total": n_total, + "n_succeeded": n_valid, + "consensus_answer": consensus_answer, + "agreement": agreement, + }, + ) diff --git a/agent_core/components/verifier/protocols.py b/agent_core/components/verifier/protocols.py new file mode 100644 index 0000000..140c901 --- /dev/null +++ b/agent_core/components/verifier/protocols.py @@ -0,0 +1,120 @@ +"""Verifier protocol + supporting dataclasses. + +Pure typing — no kernel / state / workflow imports. The ``is_runtime`` +field on ``VerifierContext`` provides framework-level information +isolation: oracle fields on ``GroundTruth`` (reference answer, formal +spec, test cases) are stripped automatically when ``is_runtime=True``, +so the verifier physically cannot see them during business runtime. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Protocol, runtime_checkable + + +@dataclass +class GroundTruth: + """Reference material a verifier may consult. + + Public fields are visible in both runtime and eval; underscore + fields are oracle-only and only surface when ``is_runtime=False``. + """ + + rubric: str | None = None + metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + _reference: str | None = None + _formal_spec: str | None = None + _test_cases: list[Any] | None = None + + +@dataclass +class VerifierContext: + """Context threaded through ``Verifier.verify`` calls.""" + + is_runtime: bool + _ground_truth: GroundTruth | None = None + call_llm: Any | None = None + call_search: Any | None = None + metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + + @property + def ground_truth(self) -> GroundTruth | None: + """Return ground truth with oracle fields stripped at runtime.""" + if self._ground_truth is None: + return None + if self.is_runtime: + return GroundTruth( + rubric=self._ground_truth.rubric, + metadata=dict(self._ground_truth.metadata), + ) + return self._ground_truth + + +@dataclass +class Finding: + severity: str + message: str + location: str | None = None + + +def _no_sub_verdicts() -> list[Verdict]: + """Empty default for :attr:`Verdict.sub_verdicts`. + + ``field(default_factory=list[Verdict])`` -- the spelling used for every + other collection field here -- cannot work inside ``Verdict``'s own body, + because ``default_factory`` is evaluated at class-creation time when the + name does not exist yet. ``from __future__ import annotations`` keeps this + function's return annotation lazy, so the type is still declared. + """ + return [] + + +@dataclass +class Verdict: + """Unified verifier return value. + + ``sub_verdicts`` carries nested results so composers can express + arbitrarily deep trees without bespoke types. + """ + + score: float | None = None + passed: bool = False + findings: list[Finding] = field(default_factory=list[Finding]) + reasoning: str = "" + sub_verdicts: list[Verdict] = field(default_factory=_no_sub_verdicts) + metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + + +@runtime_checkable +class Verifier(Protocol): + """Unified verifier contract: ``(subject, ctx) → Verdict``.""" + + role_id: str + + async def verify( + self, + subject: Any, + ctx: VerifierContext, + ) -> Verdict: ... + + +@runtime_checkable +class Generator(Protocol): + """Produces an artifact for one round. + + Mirrors ``cycle.Writer`` semantics: stateful across rounds (typically + a persistent agent session). The ``prev_verdict`` argument is loosely + typed because cycle GVR engines pass an ``AuditReport`` while pure + Verifier flows pass a ``Verdict``; both are handled uniformly until + the cycle engine migrates to ``Verdict`` natively. + """ + + role_id: str + + async def generate( + self, + prev_verdict: Any, + round_num: int, + feedback_md: str, + ) -> Any: ... diff --git a/docs/cycle-verifier-boundary.md b/docs/cycle-verifier-boundary.md new file mode 100644 index 0000000..464a668 --- /dev/null +++ b/docs/cycle-verifier-boundary.md @@ -0,0 +1,65 @@ +# Write-audit cycle and verifier boundary + +AgentCore owns the product-neutral iterate-until-pass machinery: + +- the write → audit → feedback → re-write loop and its termination rules; +- max-rounds, wall-clock budget with between-round estimation, and the + missing-artifact output check; +- writer-exception and missing-output capture as synthetic audits, so a crash + becomes a finding the next round can read rather than an aborted cycle; +- per-round persistence of the audit trail (`audit_round_{N}.json`); +- the round-observer protocol and its failure-isolated dispatch; +- the `AuditReport` / `WriterOutput` / `AuditFinding` / `CycleOutput` data + contracts and their JSON round-trip; +- the `Verifier` / `Generator` protocols, the `Verdict` model, and the six + composers (Pipeline, Ensemble, Fallback, Cascade, Parallel, + ConsensusVerifier), each of which implements `Verifier` so they nest; +- the `AuditReport` ↔ `Verdict` bridge, so a `Verifier` can be used wherever a + `CycleAuditor` is expected. + +## What the product owns + +Concrete writers and auditors. That is the whole point of the split: the cycle +knows how many rounds are left and whether an artifact exists, and knows +nothing about what makes an artifact good. + +- **Scoring vocabulary.** `AuditFinding.category` is a free-form `str` on + purpose, not an enum — the useful vocabulary differs between a paper, a + patch, a dataset and a plan, and freezing one here would force every product + to translate into someone else's taxonomy. +- **Verdict words.** `AuditReport.verdict` is likewise free-form. The cycle + compares it against the caller-supplied `terminal_verdicts` set, which + defaults to `{"success", "abandon"}` only so that the common case needs no + configuration. +- **Prompts, tools, and role definitions** for the writer and auditor sessions. +- **The oracle.** `GroundTruth` carries underscore-prefixed oracle fields + (`_reference`, `_formal_spec`, `_test_cases`) that `VerifierContext` strips + automatically when `is_runtime=True`. Core enforces the isolation; the + product decides what the reference answer is and when it may be seen. +- **Feedback rendering**, beyond the default Markdown renderer. Core ships + `DefaultFeedbackRenderer` and guarantees it never returns an empty string, + because an empty feedback body silently turns a revise round into a re-roll. + +## Injection seams + +`SessionBackedWriter` and `SessionBackedAuditor` take `bus: Any` and call +`create_session` / `submit_task_to_session` / `collect` structurally. They work +against `agent_core.components.agent_bus`, but the annotation is deliberately +loose so a product can substitute its own session runtime without implementing +a Protocol it does not otherwise need. + +`output_check` is a `Callable[[Path], list[str]]` returning the missing paths — +a plain callable rather than a policy object, because "what must exist when the +writer is done" is a per-task fact, not a per-product one. + +Observers receive a live view of `history_so_far`. Every observer hook is +dispatched inside `try/except`: an observer that raises is logged and skipped, +never allowed to abort a cycle that would otherwise have produced an artifact. + +## Not in this package + +Claim-level verification engines, evidence calibration, debate arbitration, +benchmark judges, and cross-trajectory consensus over a specific answer shape +all stay in the product. They are verification *policies* over +product-specific subjects; only the contract they implement and the primitives +that compose them belong here. diff --git a/pyproject.toml b/pyproject.toml index 3d6f8fa..8c7c273 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "apodex-agent-core" -version = "0.4.0" +version = "0.5.0" description = "Shared, product-neutral runtime primitives for Apodex agents" readme = "README.md" license = "Apache-2.0" diff --git a/tests/test_cycle_answer_consensus_shared.py b/tests/test_cycle_answer_consensus_shared.py new file mode 100644 index 0000000..0ad1b5a --- /dev/null +++ b/tests/test_cycle_answer_consensus_shared.py @@ -0,0 +1,182 @@ +"""TDD-first tests for ``select_by_answer_consensus``. + +Mirrors MiroVerifier@fac1b9e:agents/imo_gvr.py:_select answer mode: +attempts grouped by answer-equivalence buckets; pick the bucket with +highest summed score (tie → larger bucket); within winning bucket, +``select_best_attempt`` rule. +""" +from __future__ import annotations + +from collections.abc import Awaitable, Callable + +import pytest + +from agent_core.components.cycle import AuditReport, WriterOutput +from agent_core.components.cycle.builders import select_by_answer_consensus + + +def _attempt( + *, + content: str = "writer output", + score: float | None, + extracted: str, + extra: dict | None = None, +) -> tuple[WriterOutput, AuditReport]: + metadata = {"score": score, "extracted_answer": extracted} + if extra: + metadata.update(extra) + audit = AuditReport( + verdict="iterate", + findings=[], + summary="", + confidence=1.0, + metadata=metadata, + ) + return (WriterOutput(content=content), audit) + + +# --- equivalence helpers used in tests --------------------------------- + +async def _exact(a: str, b: str) -> bool: + return a.strip() == b.strip() + + +def _make_equivalence_set(*groups: list[str]) -> Callable[[str, str], Awaitable[bool]]: + """Return an equivalence function that maps each input to its group.""" + bucket_of: dict[str, int] = {} + for i, group in enumerate(groups): + for s in group: + bucket_of[s] = i + + async def eq(a: str, b: str) -> bool: + return bucket_of.get(a, -1) == bucket_of.get(b, -2) + + return eq + + +# --- tests -------------------------------------------------------------- + + +class TestSimpleCases: + @pytest.mark.asyncio + async def test_empty_history_raises(self) -> None: + with pytest.raises(ValueError): + await select_by_answer_consensus([], equivalence=_exact) + + @pytest.mark.asyncio + async def test_single_attempt_returns_it(self) -> None: + history = [_attempt(score=5, extracted="42")] + idx, _w, a = await select_by_answer_consensus(history, equivalence=_exact) + assert idx == 0 + assert a.metadata["extracted_answer"] == "42" + + +class TestBucketing: + @pytest.mark.asyncio + async def test_same_answer_bucketed_together(self) -> None: + history = [ + _attempt(score=3, extracted="42"), + _attempt(score=5, extracted="42"), + ] + idx, _w, a = await select_by_answer_consensus(history, equivalence=_exact) + # Both in same bucket, best-by-score → idx 1 + assert idx == 1 + assert a.metadata["score"] == 5 + + @pytest.mark.asyncio + async def test_different_answers_separate_buckets_higher_summed_score_wins( + self, + ) -> None: + history = [ + _attempt(score=7, extracted="x+1"), + _attempt(score=3, extracted="y-1"), + _attempt(score=4, extracted="y-1"), + ] + # bucket A "x+1": sum=7 + # bucket B "y-1": sum=7 + # tie on sum → larger bucket wins → B (size=2) + idx, _w, a = await select_by_answer_consensus(history, equivalence=_exact) + assert a.metadata["extracted_answer"] == "y-1" + # within bucket B: best-by-score → idx 2 (score 4) + assert idx == 2 + + @pytest.mark.asyncio + async def test_higher_summed_score_wins_over_larger_bucket(self) -> None: + history = [ + _attempt(score=7, extracted="x+1"), + _attempt(score=3, extracted="y-1"), + _attempt(score=3, extracted="y-1"), + ] + # bucket A: sum=7, size=1 + # bucket B: sum=6, size=2 + # higher summed score wins → A + idx, _w, a = await select_by_answer_consensus(history, equivalence=_exact) + assert idx == 0 + assert a.metadata["extracted_answer"] == "x+1" + + @pytest.mark.asyncio + async def test_equivalent_answers_via_callable_bucket_together(self) -> None: + # "1/2" and "\frac{1}{2}" treated as equivalent by custom callable + eq = _make_equivalence_set(["1/2", "\\frac{1}{2}"], ["x+1"]) + history = [ + _attempt(score=5, extracted="1/2"), + _attempt(score=4, extracted="\\frac{1}{2}"), + _attempt(score=7, extracted="x+1"), + ] + # bucket A "1/2 ≡ \frac{1}{2}": sum=9, size=2 + # bucket B "x+1": sum=7, size=1 + # higher sum → A + idx, _w, a = await select_by_answer_consensus(history, equivalence=eq) + assert a.metadata["extracted_answer"] == "1/2" # best-by-score in bucket A + assert idx == 0 + + +class TestFallback: + @pytest.mark.asyncio + async def test_no_extracted_answers_falls_back_to_best_by_score(self) -> None: + history = [ + _attempt(score=3, extracted=""), + _attempt(score=7, extracted=""), + _attempt(score=5, extracted=""), + ] + idx, _w, a = await select_by_answer_consensus(history, equivalence=_exact) + assert idx == 1 + assert a.metadata["score"] == 7 + + @pytest.mark.asyncio + async def test_some_invalid_some_valid_only_valid_bucketed(self) -> None: + history = [ + _attempt(score=7, extracted=""), + _attempt(score=3, extracted="x+1"), + ] + # only valid attempt is idx=1 → bucket {1}, returns it + idx, _w, _a = await select_by_answer_consensus(history, equivalence=_exact) + assert idx == 1 + + +class TestEquivalenceExceptions: + @pytest.mark.asyncio + async def test_equivalence_raising_treated_as_not_equivalent(self) -> None: + async def flaky(a: str, b: str) -> bool: + raise RuntimeError("boom") + + history = [ + _attempt(score=5, extracted="x"), + _attempt(score=3, extracted="y"), + ] + # equivalence always raises → every attempt becomes its own + # bucket → highest score bucket wins (idx 0) + idx, _w, _a = await select_by_answer_consensus(history, equivalence=flaky) + assert idx == 0 + + +class TestTieBreaking: + @pytest.mark.asyncio + async def test_within_bucket_ties_go_to_latest(self) -> None: + history = [ + _attempt(score=5, extracted="42"), + _attempt(score=5, extracted="42"), + ] + idx, _w, _a = await select_by_answer_consensus(history, equivalence=_exact) + # Same bucket, same score → latest wins (parity with select_best_attempt) + assert idx == 1 diff --git a/tests/test_cycle_builders_shared.py b/tests/test_cycle_builders_shared.py new file mode 100644 index 0000000..fecbc3d --- /dev/null +++ b/tests/test_cycle_builders_shared.py @@ -0,0 +1,779 @@ +"""Tests for SessionBackedWriter / SessionBackedAuditor. + +Mocked-AgentBus tests covering: + +- writer creates exactly ONE session reused across rounds (FR9) +- auditor creates a FRESH session per round (FR9) +- ConcludePhaseObserver auto-injected into auditor (FR7) +- llm_timeout propagated to create_session (FR8) +- tools_override defaults to [] for auditor; explicit list becomes whitelist +- llm_override propagated for both writer and auditor +- audit JSON parsing (fenced + bare + parse-failure fallback) +""" +from __future__ import annotations + +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from agent_core.components.agent_bus.models import ( + CollectResult, + SubAgentResult, +) +from agent_core.components.cycle.builders import ( + MajorityVoteAuditor, + ScoreThresholdAuditor, + SessionBackedAuditor, + SessionBackedWriter, + _confidence_from_score, + _extract_first_number, + _extract_first_str, + _parse_audit_report, + select_best_attempt, +) +from agent_core.components.cycle.types import AuditFinding, AuditReport, WriterOutput +from agent_core.components.observers.conclude_phase_observer import ConcludePhaseObserver + + +def _ok_result(content: str = "ok", metadata: dict | None = None) -> SubAgentResult: + return SubAgentResult( + question="x", + role_id="any", + final_content=content, + success=True, + metadata=metadata or {}, + ) + + +def _make_bus(*, content: str = "ok", session_id: str = "task-1::w") -> Any: + bus = AsyncMock() + bus.create_session = AsyncMock(return_value=session_id) + bus.submit_task_to_session = AsyncMock(return_value="job-1") + bus.collect = AsyncMock( + return_value=CollectResult(completed=[_ok_result(content)]), + ) + return bus + + +# ── SessionBackedWriter ───────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_writer_creates_one_session_reused_across_rounds( + tmp_path: Path, +) -> None: + bus = _make_bus() + writer = SessionBackedWriter( + bus=bus, task_id="t1", role_id="paper_writer", name="paper_w", + system_prompt="you write papers", initial_prompt="write a paper", + work_dir=tmp_path, output_check=lambda _: [], + llm_timeout=420, + ) + + await writer.generate(None, 0, "") + await writer.generate(None, 1, "feedback round 1") + await writer.generate(None, 2, "feedback round 2") + + # Exactly ONE create_session call across all rounds + assert bus.create_session.await_count == 1 + # Three submissions (one per round) + assert bus.submit_task_to_session.await_count == 3 + + +@pytest.mark.asyncio +async def test_writer_round_0_submits_initial_prompt(tmp_path: Path) -> None: + bus = _make_bus() + writer = SessionBackedWriter( + bus=bus, task_id="t", role_id="r", name="w", + system_prompt="sp", initial_prompt="write X", + work_dir=tmp_path, output_check=lambda _: [], + ) + await writer.generate(None, 0, "") + args, _kwargs = bus.submit_task_to_session.await_args + submitted_prompt = args[1] + assert "write X" in submitted_prompt + + +@pytest.mark.asyncio +async def test_writer_round_n_submits_feedback_plus_revision( + tmp_path: Path, +) -> None: + bus = _make_bus() + writer = SessionBackedWriter( + bus=bus, task_id="t", role_id="r", name="w", + system_prompt="sp", initial_prompt="write X", + work_dir=tmp_path, output_check=lambda _: [], + revision_instruction="REVISE_NOW", + ) + await writer.generate(None, 1, "FEEDBACK_HERE") + args, _ = bus.submit_task_to_session.await_args + submitted_prompt = args[1] + assert "FEEDBACK_HERE" in submitted_prompt + assert "REVISE_NOW" in submitted_prompt + + +@pytest.mark.asyncio +async def test_writer_propagates_llm_timeout_and_override( + tmp_path: Path, +) -> None: + bus = _make_bus() + fake_llm = object() + writer = SessionBackedWriter( + bus=bus, task_id="t", role_id="r", name="w", + system_prompt="sp", initial_prompt="x", + work_dir=tmp_path, output_check=lambda _: [], + llm_timeout=999, llm_override=fake_llm, + ) + await writer.generate(None, 0, "") + _, kwargs = bus.create_session.await_args + assert kwargs["llm_timeout"] == 999 + assert kwargs["llm_override"] is fake_llm + + +@pytest.mark.asyncio +async def test_writer_returns_writer_output_with_content_and_metadata( + tmp_path: Path, +) -> None: + bus = AsyncMock() + bus.create_session = AsyncMock(return_value="t::w") + bus.submit_task_to_session = AsyncMock(return_value="j1") + bus.collect = AsyncMock( + return_value=CollectResult(completed=[ + _ok_result( + "draft v1", + metadata={"files": ["intro.md"], "message_count": 12}, + ), + ]), + ) + writer = SessionBackedWriter( + bus=bus, task_id="t", role_id="r", name="w", + system_prompt="sp", initial_prompt="x", + work_dir=tmp_path, output_check=lambda _: [], + ) + out = await writer.generate(None, 0, "") + assert out.content == "draft v1" + assert out.files == ["intro.md"] + assert out.message_count == 12 + + +# ── SessionBackedAuditor ──────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_auditor_creates_fresh_session_per_round( + tmp_path: Path, +) -> None: + del tmp_path # not used here but shared fixture style + # Each call to create_session returns a per-round id + sid_returns = iter(["t::a_R0", "t::a_R1", "t::a_R2"]) + bus = AsyncMock() + bus.create_session = AsyncMock(side_effect=lambda **_: next(sid_returns)) + bus.submit_task_to_session = AsyncMock(return_value="job") + bus.collect = AsyncMock( + return_value=CollectResult(completed=[ + _ok_result('{"verdict": "iterate", "findings": []}'), + ]), + ) + auditor = SessionBackedAuditor( + bus=bus, task_id="t", role_id="auditor", + system_prompt="emit JSON", + ) + + for r in range(3): + await auditor.verify(WriterOutput(content="x"), r) + + assert bus.create_session.await_count == 3 + # Fresh names per round + names = [ + call.kwargs["name"] for call in bus.create_session.await_args_list + ] + assert names == ["auditor_R0", "auditor_R1", "auditor_R2"] + + +@pytest.mark.asyncio +async def test_auditor_injects_conclude_phase_observer() -> None: + bus = _make_bus(content='{"verdict": "iterate", "findings": []}') + auditor = SessionBackedAuditor( + bus=bus, task_id="t", role_id="r", + system_prompt="sp", + conclude_ratio=0.7, + ) + await auditor.verify(WriterOutput(content="x"), 0) + _, kwargs = bus.submit_task_to_session.await_args + obs_list = kwargs["observers"] + assert len(obs_list) >= 1 + has_conclude = any( + isinstance(o, ConcludePhaseObserver) for o in obs_list + ) + assert has_conclude + + +@pytest.mark.asyncio +async def test_auditor_default_tools_is_empty_list() -> None: + bus = _make_bus(content='{"verdict": "success", "findings": []}') + auditor = SessionBackedAuditor( + bus=bus, task_id="t", role_id="r", system_prompt="sp", + ) + await auditor.verify(WriterOutput(content="x"), 0) + _, kwargs = bus.create_session.await_args + # Default is empty tool set — Codex-style hard restriction + assert kwargs["tools_override"] == [] + + +@pytest.mark.asyncio +async def test_auditor_explicit_tools_become_whitelist() -> None: + bus = _make_bus(content='{"verdict": "success", "findings": []}') + fake_tool = object() + auditor = SessionBackedAuditor( + bus=bus, task_id="t", role_id="r", system_prompt="sp", + tools=[fake_tool], # type: ignore[list-item] + ) + await auditor.verify(WriterOutput(content="x"), 0) + _, kwargs = bus.create_session.await_args + assert kwargs["tools_override"] == [fake_tool] + + +@pytest.mark.asyncio +async def test_auditor_propagates_llm_timeout_and_override() -> None: + bus = _make_bus(content='{"verdict": "success", "findings": []}') + fake_llm = object() + auditor = SessionBackedAuditor( + bus=bus, task_id="t", role_id="r", system_prompt="sp", + llm_timeout=777, llm_override=fake_llm, + ) + await auditor.verify(WriterOutput(content="x"), 0) + _, kwargs = bus.create_session.await_args + assert kwargs["llm_timeout"] == 777 + assert kwargs["llm_override"] is fake_llm + + +@pytest.mark.asyncio +async def test_auditor_handles_runtime_failure_as_iterate() -> None: + """If AgentBus returns no completed/failed results, auditor synthesizes.""" + bus = AsyncMock() + bus.create_session = AsyncMock(return_value="t::a") + bus.submit_task_to_session = AsyncMock(return_value="job") + bus.collect = AsyncMock(return_value=CollectResult()) + auditor = SessionBackedAuditor( + bus=bus, task_id="t", role_id="r", system_prompt="sp", + ) + report = await auditor.verify(WriterOutput(content="x"), 0) + assert report.verdict == "iterate" + assert any( + f.category == "auditor_runtime_error" for f in report.findings + ) + + +# ── _parse_audit_report ───────────────────────────────────────────────────── + + +def test_parse_fenced_json_block() -> None: + raw = """Some prose first. + +```json +{ + "verdict": "iterate", + "findings": [ + {"category": "x", "severity": "error", "short_message": "m"} + ], + "summary": "needs work", + "confidence": 0.9 +} +``` + +Trailing prose. +""" + r = _parse_audit_report(raw) + assert r.verdict == "iterate" + assert len(r.findings) == 1 + assert r.findings[0].category == "x" + assert r.summary == "needs work" + assert r.raw_text == raw # preserved for debugging + + +def test_parse_bare_json_object() -> None: + raw = 'prefix prose {"verdict": "success", "findings": []} suffix' + r = _parse_audit_report(raw) + assert r.verdict == "success" + assert r.findings == [] + + +def test_parse_invalid_json_synthesizes_parse_failure() -> None: + r = _parse_audit_report("totally non-JSON output here") + assert r.verdict == "iterate" + assert any( + f.category == "auditor_parse_failure" for f in r.findings + ) + assert r.metadata.get("reason") == "auditor_parse_failure" + + +def test_parse_empty_response_synthesizes_parse_failure() -> None: + r = _parse_audit_report("") + assert r.verdict == "iterate" + assert any( + f.category == "auditor_parse_failure" for f in r.findings + ) + + +def test_parse_schema_mismatch_synthesizes_schema_error() -> None: + """JSON parses but is missing the required 'verdict' key.""" + r = _parse_audit_report('{"not_a_verdict": "x"}') + assert r.verdict == "iterate" + assert any( + f.category == "auditor_schema_mismatch" for f in r.findings + ) + + +# ── ScoreThresholdAuditor ─────────────────────────────────────────────────── + + +def _make_grader_bus(content: str) -> Any: + bus = AsyncMock() + bus.create_session = AsyncMock(return_value="t::g") + bus.submit_task_to_session = AsyncMock(return_value="job-g") + bus.collect = AsyncMock( + return_value=CollectResult(completed=[_ok_result(content)]), + ) + return bus + + +@pytest.mark.asyncio +async def test_score_threshold_pass_emits_success_verdict() -> None: + bus = _make_grader_bus( + '{"score": 7, "feedback": "looks great"}', + ) + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade per rubric", + pass_threshold=6, + ) + report = await auditor.verify(WriterOutput(content="proof v1"), 0) + assert report.verdict == "success" + assert report.metadata["score"] == 7.0 + assert report.metadata["passed"] is True + assert report.confidence == pytest.approx(1.0) + assert report.findings[0].category == "grade" + assert report.findings[0].severity == "info" + assert "looks great" in report.summary + + +@pytest.mark.asyncio +async def test_score_threshold_fail_emits_iterate_verdict() -> None: + bus = _make_grader_bus( + '{"score": 1, "feedback": "missing case analysis"}', + ) + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade per rubric", + pass_threshold=6, + ) + report = await auditor.verify(WriterOutput(content="proof"), 0) + assert report.verdict == "iterate" + assert report.metadata["score"] == 1.0 + assert report.metadata["passed"] is False + assert report.findings[0].severity == "warning" + assert "missing case analysis" in report.summary + + +@pytest.mark.asyncio +async def test_score_threshold_alternate_keys_points_explanation() -> None: + """IMO-GVR's grader emits {points, explanation} not {score, feedback}.""" + bus = _make_grader_bus( + '{"points": 6, "explanation": "minor gap"}', + ) + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade", + pass_threshold=6, + ) + report = await auditor.verify(WriterOutput(content="x"), 0) + assert report.verdict == "success" + assert report.metadata["score"] == 6.0 + assert report.metadata["feedback"] == "minor gap" + + +@pytest.mark.asyncio +async def test_score_threshold_always_iterate_for_run_k_mode() -> None: + """pass_verdict='iterate' lets the cycle run all max_rounds even on + a perfect score — parity with MiroVerifier IMO-GVR.""" + bus = _make_grader_bus('{"score": 7, "feedback": "perfect"}') + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade", + pass_threshold=6, + pass_verdict="iterate", # never terminate via verdict + ) + report = await auditor.verify(WriterOutput(content="x"), 0) + assert report.verdict == "iterate" + assert report.metadata["score"] == 7.0 + assert report.metadata["passed"] is True + + +@pytest.mark.asyncio +async def test_score_threshold_parse_failure_synthesizes_finding() -> None: + bus = _make_grader_bus("totally non-JSON garbage") + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade", + ) + report = await auditor.verify(WriterOutput(content="x"), 0) + assert report.verdict == "iterate" + assert any( + f.category == "grader_parse_failure" for f in report.findings + ) + assert report.metadata.get("reason") == "grader_parse_failure" + + +@pytest.mark.asyncio +async def test_score_threshold_score_missing_synthesizes_finding() -> None: + """JSON parses but has no recognised score key.""" + bus = _make_grader_bus('{"verdict": "ok"}') + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade", + ) + report = await auditor.verify(WriterOutput(content="x"), 0) + assert report.verdict == "iterate" + assert any( + f.category == "grader_score_missing" for f in report.findings + ) + + +@pytest.mark.asyncio +async def test_score_threshold_string_score_is_coerced() -> None: + """Some models emit '"score": "7"' — should still parse.""" + bus = _make_grader_bus('{"score": "7", "feedback": "ok"}') + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade", + pass_threshold=6, + ) + report = await auditor.verify(WriterOutput(content="x"), 0) + assert report.verdict == "success" + assert report.metadata["score"] == 7.0 + + +@pytest.mark.asyncio +async def test_score_threshold_creates_fresh_session_per_round() -> None: + """Same per-round freshness invariant as SessionBackedAuditor.""" + sid_returns = iter(["t::g_R0", "t::g_R1", "t::g_R2"]) + bus = AsyncMock() + bus.create_session = AsyncMock(side_effect=lambda **_: next(sid_returns)) + bus.submit_task_to_session = AsyncMock(return_value="job") + bus.collect = AsyncMock( + return_value=CollectResult(completed=[ + _ok_result('{"score": 5, "feedback": "fix x"}'), + ]), + ) + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", + system_prompt="grade", + ) + for r in range(3): + await auditor.verify(WriterOutput(content="x"), r) + assert bus.create_session.await_count == 3 + names = [ + call.kwargs["name"] for call in bus.create_session.await_args_list + ] + assert names == ["grader_R0", "grader_R1", "grader_R2"] + + +@pytest.mark.asyncio +async def test_score_threshold_default_tools_is_empty_list() -> None: + bus = _make_grader_bus('{"score": 7, "feedback": "ok"}') + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", system_prompt="grade", + ) + await auditor.verify(WriterOutput(content="x"), 0) + _, kwargs = bus.create_session.await_args + assert kwargs["tools_override"] == [] + + +@pytest.mark.asyncio +async def test_score_threshold_injects_conclude_observer() -> None: + bus = _make_grader_bus('{"score": 7, "feedback": "ok"}') + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", system_prompt="grade", + conclude_ratio=0.6, + ) + await auditor.verify(WriterOutput(content="x"), 0) + _, kwargs = bus.submit_task_to_session.await_args + obs_list = kwargs["observers"] + assert any(isinstance(o, ConcludePhaseObserver) for o in obs_list) + + +@pytest.mark.asyncio +async def test_score_threshold_handles_runtime_failure_as_fail_verdict() -> None: + """If AgentBus.collect returns nothing, audit synthesises a runtime + error report (verdict 'iterate' to keep the cycle alive).""" + bus = AsyncMock() + bus.create_session = AsyncMock(return_value="t::g") + bus.submit_task_to_session = AsyncMock(return_value="job") + bus.collect = AsyncMock(return_value=CollectResult()) + auditor = ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="grader", system_prompt="grade", + ) + report = await auditor.verify(WriterOutput(content="x"), 0) + assert report.verdict == "iterate" + assert any( + f.category == "auditor_runtime_error" for f in report.findings + ) + + +def test_score_threshold_rejects_invalid_score_range() -> None: + bus = AsyncMock() + with pytest.raises(ValueError, match="score_range"): + ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="g", system_prompt="x", + score_range=(7, 0), # backwards + ) + + +def test_score_threshold_rejects_empty_score_keys() -> None: + bus = AsyncMock() + with pytest.raises(ValueError, match="score_keys"): + ScoreThresholdAuditor( + bus=bus, task_id="t", role_id="g", system_prompt="x", + score_keys=(), + ) + + +# ── Score / selection helpers ─────────────────────────────────────────────── + + +def test_extract_first_number_prefers_first_key() -> None: + assert _extract_first_number( + {"score": 7, "points": 3}, ("score", "points"), + ) == 7.0 + assert _extract_first_number( + {"points": 3}, ("score", "points"), + ) == 3.0 + + +def test_extract_first_number_skips_bool_and_non_numeric() -> None: + # bool is a subclass of int; ScoreThresholdAuditor must not treat True as 1. + assert _extract_first_number({"score": True}, ("score",)) is None + assert _extract_first_number({"score": "abc"}, ("score",)) is None + assert _extract_first_number({}, ("score",)) is None + + +def test_extract_first_str_prefers_non_empty() -> None: + assert _extract_first_str( + {"feedback": "", "explanation": "real"}, + ("feedback", "explanation"), + ) == "real" + assert _extract_first_str({}, ("feedback",)) is None + # whitespace-only counts as empty + assert _extract_first_str( + {"feedback": " "}, ("feedback",), + ) is None + + +def test_confidence_from_score_clamps() -> None: + assert _confidence_from_score(7, (0, 7)) == 1.0 + assert _confidence_from_score(0, (0, 7)) == 0.0 + assert _confidence_from_score(3.5, (0, 7)) == pytest.approx(0.5) + # Out-of-range scores clamp. + assert _confidence_from_score(8, (0, 7)) == 1.0 + assert _confidence_from_score(-1, (0, 7)) == 0.0 + # None range disables. + assert _confidence_from_score(7, None) == 0.0 + + +def _hist_entry(content: str, score: float | None) -> tuple[WriterOutput, AuditReport]: + metadata = {"score": score} if score is not None else {} + return ( + WriterOutput(content=content), + AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="grade", severity="info", + short_message=f"score={score}", + ), + ], + metadata=metadata, + ), + ) + + +def test_select_best_attempt_picks_highest_score() -> None: + history = [ + _hist_entry("v1", 1.0), + _hist_entry("v2", 7.0), + _hist_entry("v3", 6.0), + ] + idx, w, a = select_best_attempt(history) + assert idx == 1 + assert w.content == "v2" + assert a.metadata["score"] == 7.0 + + +def test_select_best_attempt_ties_go_to_latest() -> None: + """IMO-GVR rule: when scores tie, the *latest* attempt wins + (revisions are at least as good as earlier ones).""" + history = [ + _hist_entry("v1", 7.0), + _hist_entry("v2", 7.0), + _hist_entry("v3", 7.0), + ] + idx, w, _a = select_best_attempt(history) + assert idx == 2 + assert w.content == "v3" + + +def test_select_best_attempt_unscored_audit_treated_as_minimum() -> None: + """Unscored audits should never beat scored ones.""" + history = [ + _hist_entry("v1", 1.0), + _hist_entry("v2", None), # no score in metadata + ] + idx, w, _a = select_best_attempt(history) + assert idx == 0 + assert w.content == "v1" + + +def test_select_best_attempt_empty_history_raises() -> None: + with pytest.raises(ValueError, match="empty"): + select_best_attempt([]) + + +# ── MajorityVoteAuditor ──────────────────────────────────────────────────── + + +class _FakeAuditor: + """Returns scripted AuditReports in order. Each call pops one.""" + + role_id = "fake_judge" + + def __init__(self, reports: list) -> None: + self._reports = list(reports) + self.calls = 0 + + async def verify(self, writer_output, round_num): + self.calls += 1 + item = self._reports.pop(0) + if isinstance(item, Exception): + raise item + return item + + +def _grade_report(score: float, verdict: str = "iterate") -> AuditReport: + return AuditReport( + verdict=verdict, + findings=[ + AuditFinding( + category="grade", severity="info", + short_message=f"score={score}", + metadata={"score": score}, + ), + ], + summary=f"feedback for score {score}", + confidence=score / 7.0, + metadata={"score": score, "feedback": f"feedback for score {score}"}, + ) + + +@pytest.mark.asyncio +async def test_majority_vote_aggregates_median_score() -> None: + base = _FakeAuditor([ + _grade_report(7.0), + _grade_report(6.0), + _grade_report(7.0), + ]) + voter = MajorityVoteAuditor(base=base, n_votes=3) + report = await voter.verify(WriterOutput(content="x"), 0) + assert report.metadata["score"] == 7.0 # median([7,6,7]) = 7 + assert report.metadata["individual_scores"] == [7.0, 6.0, 7.0] + assert report.confidence == 1.0 + assert report.metadata["n_votes"] == 3 + assert report.metadata["n_succeeded"] == 3 + + +@pytest.mark.asyncio +async def test_majority_vote_majority_verdict() -> None: + base = _FakeAuditor([ + _grade_report(7.0, verdict="success"), + _grade_report(6.0, verdict="success"), + _grade_report(1.0, verdict="iterate"), + ]) + voter = MajorityVoteAuditor(base=base, n_votes=3) + report = await voter.verify(WriterOutput(content="x"), 0) + assert report.verdict == "success" + assert report.metadata["individual_verdicts"] == [ + "success", "success", "iterate", + ] + + +@pytest.mark.asyncio +async def test_majority_vote_partial_failures_lower_confidence() -> None: + base = _FakeAuditor([ + _grade_report(7.0), + RuntimeError("judge 2 boom"), + _grade_report(6.0), + ]) + voter = MajorityVoteAuditor(base=base, n_votes=3) + report = await voter.verify(WriterOutput(content="x"), 0) + assert report.confidence == pytest.approx(2 / 3) + assert report.metadata["n_succeeded"] == 2 + assert report.metadata["n_failed"] == 1 + assert report.metadata["score"] == 6.5 # median([7, 6]) + + +@pytest.mark.asyncio +async def test_majority_vote_all_failures_synthesizes_report() -> None: + base = _FakeAuditor([ + RuntimeError("a"), + RuntimeError("b"), + RuntimeError("c"), + ]) + voter = MajorityVoteAuditor(base=base, n_votes=3) + report = await voter.verify(WriterOutput(content="x"), 0) + assert report.verdict == "iterate" + assert any( + f.category == "majority_vote_failed" for f in report.findings + ) + assert report.metadata["n_succeeded"] == 0 + assert report.confidence == 0.0 + + +@pytest.mark.asyncio +async def test_majority_vote_calls_base_n_times_in_parallel() -> None: + base = _FakeAuditor([ + _grade_report(7.0), _grade_report(7.0), _grade_report(7.0), + _grade_report(6.0), _grade_report(6.0), _grade_report(6.0), + ]) + voter = MajorityVoteAuditor(base=base, n_votes=3) + await voter.verify(WriterOutput(content="x"), 0) + assert base.calls == 3 + await voter.verify(WriterOutput(content="y"), 1) + assert base.calls == 6 + + +@pytest.mark.asyncio +async def test_majority_vote_custom_aggregator() -> None: + """statistics.mean instead of median.""" + import statistics as _s + base = _FakeAuditor([ + _grade_report(7.0), _grade_report(6.0), _grade_report(2.0), + ]) + voter = MajorityVoteAuditor( + base=base, n_votes=3, score_aggregator=_s.mean, + ) + report = await voter.verify(WriterOutput(content="x"), 0) + assert report.metadata["score"] == pytest.approx(5.0) + + +@pytest.mark.asyncio +async def test_majority_vote_inherits_role_id() -> None: + base = _FakeAuditor([_grade_report(7.0)]) + voter = MajorityVoteAuditor(base=base, n_votes=1) + assert voter.role_id == "fake_judge" + + +def test_majority_vote_invalid_n_rejected() -> None: + base = _FakeAuditor([]) + with pytest.raises(ValueError, match="n_votes"): + MajorityVoteAuditor(base=base, n_votes=0) diff --git a/tests/test_cycle_default_renderer_shared.py b/tests/test_cycle_default_renderer_shared.py new file mode 100644 index 0000000..3a813f3 --- /dev/null +++ b/tests/test_cycle_default_renderer_shared.py @@ -0,0 +1,124 @@ +"""Tests for DefaultFeedbackRenderer.""" +from __future__ import annotations + +import json +import re + +from agent_core.components.cycle.default_renderer import DefaultFeedbackRenderer +from agent_core.components.cycle.types import AuditFinding, AuditReport + + +def test_multi_finding_emits_table_and_json_block() -> None: + audit = AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="missing_section", + severity="error", + short_message="add intro", + suggested_action="write 200-word introduction", + ), + AuditFinding( + category="citation_invalid", + severity="warning", + short_message="cite missing", + file="paper.md", line=42, + suggested_action="add citation", + ), + ], + summary="Two issues blocking acceptance.", + ) + out = DefaultFeedbackRenderer().render(audit) + # Table header is present + assert "| # | Severity | Category" in out + # Both findings show up as rows + assert "missing_section" in out + assert "citation_invalid" in out + assert "paper.md:42" in out + # JSON block is present and valid + m = re.search(r"```json\n(.+?)\n```", out, re.DOTALL) + assert m, "expected fenced json block" + parsed = json.loads(m.group(1)) + assert isinstance(parsed, list) + assert len(parsed) == 2 + assert parsed[0]["category"] == "missing_section" + # Free-form summary is rendered verbatim + assert "Two issues blocking acceptance." in out + + +def test_zero_findings_non_terminal_emits_reattempt_message() -> None: + audit = AuditReport(verdict="iterate", findings=[], summary="") + out = DefaultFeedbackRenderer().render(audit) + assert "re-attempt" in out.lower() or "re-attempt" in out + assert "iterate" in out + # No empty string — the whole point of FR6 + assert out.strip() != "" + # Must NOT include a generic placeholder + assert "please address some issues" not in out.lower() + + +def test_zero_findings_success_verdict_skips_reattempt() -> None: + """Success verdict + no findings is a clean pass — no re-attempt nudge.""" + audit = AuditReport(verdict="success", findings=[], summary="LGTM.") + out = DefaultFeedbackRenderer().render(audit) + assert "re-attempt" not in out.lower() + assert "LGTM." in out + + +def test_single_finding_table_well_formed() -> None: + audit = AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="x", severity="error", short_message="m", + ), + ], + ) + out = DefaultFeedbackRenderer().render(audit) + lines = out.splitlines() + # Header + alignment separator + one data row, in order + header_idx = next( + (i for i, line in enumerate(lines) if line.startswith("| # |")), + -1, + ) + assert header_idx >= 0, f"missing header in:\n{out}" + assert lines[header_idx + 1].startswith("|---") + assert lines[header_idx + 2].startswith("| 1 |") + assert "| error |" in lines[header_idx + 2] + assert "| x |" in lines[header_idx + 2] + + +def test_pipe_in_message_escaped() -> None: + audit = AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="x", severity="error", + short_message="contains | pipe", + ), + ], + ) + out = DefaultFeedbackRenderer().render(audit) + # Table row must escape the pipe so the row stays well-formed + assert "contains \\| pipe" in out + + +def test_custom_success_verdicts() -> None: + """Caller can override what counts as a 'success' verdict.""" + audit = AuditReport(verdict="merge", findings=[], summary="") + out = DefaultFeedbackRenderer( + success_verdicts={"merge", "ship"}, + ).render(audit) + # 'merge' is success here → no re-attempt message + assert "re-attempt" not in out.lower() + + +def test_summary_only_no_findings_terminal() -> None: + audit = AuditReport( + verdict="success", findings=[], + summary="Looks great.", confidence=0.92, + ) + out = DefaultFeedbackRenderer().render(audit) + assert "Looks great." in out + # Confidence surfaced when non-zero + assert "0.92" in out diff --git a/tests/test_cycle_integration_smoke_shared.py b/tests/test_cycle_integration_smoke_shared.py new file mode 100644 index 0000000..53fb9ca --- /dev/null +++ b/tests/test_cycle_integration_smoke_shared.py @@ -0,0 +1,204 @@ +"""End-to-end integration smoke for WriteAuditCycle. + +Wires SessionBackedWriter + SessionBackedAuditor + WriteAuditCycle +against a deterministic mocked AgentBus. No real LLM in the loop. + +Validates the three §7.2 acceptance scenarios: +1. Round-0 success terminates with success=True +2. Iterate-then-success produces a complete two-round history +3. max_rounds_exhausted produces the expected reason + full history +""" +from __future__ import annotations + +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from agent_core.components.agent_bus.models import ( + CollectResult, + SubAgentResult, +) +from agent_core.components.cycle import ( + SessionBackedAuditor, + SessionBackedWriter, + WriteAuditCycle, +) + + +def _ok(content: str) -> SubAgentResult: + return SubAgentResult( + question="x", role_id="r", final_content=content, success=True, + ) + + +def _build_bus(content_script: list[str]) -> Any: + """Bus that returns scripted final_content per submission, in order. + + Each submit_task_to_session triggers the next collect to return the + next scripted content. We use a counter via list.pop(0). + """ + pending = list(content_script) + + async def _collect(_job_ids: list[str], **_kwargs: Any) -> CollectResult: + if not pending: + raise AssertionError("collect called more times than scripted") + return CollectResult(completed=[_ok(pending.pop(0))]) + + bus = AsyncMock() + sid_counter = iter( + f"task-1::session-{i}" for i in range(100) + ) + bus.create_session = AsyncMock(side_effect=lambda **_: next(sid_counter)) + bus.submit_task_to_session = AsyncMock(return_value="job-x") + bus.collect = _collect + return bus + + +def _make_writer(bus: Any, tmp: Path) -> SessionBackedWriter: + return SessionBackedWriter( + bus=bus, task_id="task-1", role_id="writer", name="writer", + system_prompt="you write", initial_prompt="write the artifact", + work_dir=tmp, + output_check=lambda _: [], # always pass + ) + + +def _make_auditor(bus: Any) -> SessionBackedAuditor: + return SessionBackedAuditor( + bus=bus, task_id="task-1", role_id="auditor", + system_prompt="emit AuditReport JSON", + max_turns=10, + ) + + +# ── Scenario 1: round-0 success ───────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_smoke_round_0_success(tmp_path: Path) -> None: + """Writer produces, output_check passes, auditor verdict=success → done.""" + bus = _build_bus(content_script=[ + # Round 0 — writer + "draft v1", + # Round 0 — auditor (returns JSON verdict=success) + '```json\n{"verdict": "success", "findings": [], "summary": "LGTM"}\n```', + ]) + writer = _make_writer(bus, tmp_path) + auditor = _make_auditor(bus) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=writer.output_check, + max_rounds=5, + ) + out = await cycle.run() + assert out.success is True + assert out.rounds_used == 1 + assert out.reason == "verdict=success" + assert len(out.history) == 1 + # Persistence happened + assert (tmp_path / "audit_round_0.json").exists() + + +# ── Scenario 2: iterate-then-success ──────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_smoke_iterate_then_success(tmp_path: Path) -> None: + bus = _build_bus(content_script=[ + # Round 0 — writer + "draft v1", + # Round 0 — auditor: iterate + ( + '```json\n{"verdict": "iterate", "findings": [' + '{"category": "missing_section", "severity": "error",' + ' "short_message": "add intro"}], "summary": "fix intro"}\n```' + ), + # Round 1 — writer (revised) + "draft v2 with intro", + # Round 1 — auditor: success + '```json\n{"verdict": "success", "findings": [], "summary": "LGTM"}\n```', + ]) + writer = _make_writer(bus, tmp_path) + auditor = _make_auditor(bus) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=writer.output_check, + max_rounds=5, + ) + out = await cycle.run() + assert out.success is True + assert out.rounds_used == 2 + assert len(out.history) == 2 + assert out.history[0][1].verdict == "iterate" + assert out.history[1][1].verdict == "success" + # Both round files persisted + assert (tmp_path / "audit_round_0.json").exists() + assert (tmp_path / "audit_round_1.json").exists() + # Writer reused the same persistent session: only ONE create_session + # call from the writer (auditor creates fresh per round → 2 of those). + # Total create_session calls = 1 (writer) + 2 (auditor) = 3. + assert bus.create_session.await_count == 3 + + +# ── Scenario 3: max_rounds exhausted ──────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_smoke_max_rounds_exhausted(tmp_path: Path) -> None: + bus = _build_bus(content_script=[ + # Round 0 + "draft v1", + '```json\n{"verdict": "iterate", "findings": [], "summary": "redo"}\n```', + # Round 1 + "draft v2", + '```json\n{"verdict": "iterate", "findings": [], "summary": "redo"}\n```', + # Round 2 + "draft v3", + '```json\n{"verdict": "iterate", "findings": [], "summary": "redo"}\n```', + ]) + writer = _make_writer(bus, tmp_path) + auditor = _make_auditor(bus) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=writer.output_check, + max_rounds=3, + ) + out = await cycle.run() + assert out.success is False + assert out.reason == "max_rounds_exhausted" + assert out.rounds_used == 3 + assert len(out.history) == 3 + + +# ── Scenario 4: output_check failure short-circuits the auditor ───────────── + + +@pytest.mark.asyncio +async def test_smoke_output_check_failure_short_circuits_auditor( + tmp_path: Path, +) -> None: + """Real wiring: when output_check reports missing, no auditor session is created.""" + # Auditor is never reached, so script only has the writer's outputs. + bus = _build_bus(content_script=["draft v1", "draft v2"]) + writer = _make_writer(bus, tmp_path) + auditor = _make_auditor(bus) + # Output_check reports missing every round + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, + output_check=lambda _: ["sections/intro.md"], + max_rounds=2, + ) + out = await cycle.run() + assert out.success is False + assert out.reason == "max_rounds_exhausted" + # Synthetic audits per round — auditor session never created + assert bus.create_session.await_count == 1 # only the writer + # Both rounds' synthetic audits are persisted + assert (tmp_path / "audit_round_0.json").exists() + assert (tmp_path / "audit_round_1.json").exists() + # Findings name the missing artifact + final_findings = out.final_audit.findings if out.final_audit else [] + assert any(f.file == "sections/intro.md" for f in final_findings) diff --git a/tests/test_cycle_observers_shared.py b/tests/test_cycle_observers_shared.py new file mode 100644 index 0000000..0e59976 --- /dev/null +++ b/tests/test_cycle_observers_shared.py @@ -0,0 +1,514 @@ +"""Tests for RoundObserver hooks + builtin observers + cycle wall-clock. + +Covers: + +- D1 wall-clock budget: between-rounds check, reason="wall_clock_exhausted", + safety ratio honoured, no max_wall_seconds = no check +- D2 RoundObserver: every hook fires on every observer; observer + exception is isolated; ABORT intervention from any hook terminates + cycle with reason="observer_aborted"; on_cycle_end always fires +- BestSoFarObserver: callback invoked per round with the best-so-far + selection; sync + async callbacks both work; callback errors don't + crash cycle +- PlateauAbortObserver: triggers ABORT after no-improvement window; + rounds without scores skipped; minimum-rounds gate +- MetricsObserver: emits per-round metrics dict; default logger + callback works +""" +from __future__ import annotations + +import asyncio +import logging +from pathlib import Path + +import pytest + +from agent_core.components.cycle import ( + AuditFinding, + AuditReport, + BaseRoundObserver, + BestSoFarObserver, + CycleContext, + MetricsObserver, + PlateauAbortObserver, + RoundIntervention, + WriteAuditCycle, + WriterOutput, +) + +# ── Test fixtures ─────────────────────────────────────────────────────────── + + +class _ScriptedWriter: + role_id = "writer" + + def __init__( + self, + outputs: list[str], + sleep_between: float = 0.0, + ) -> None: + self._outputs = list(outputs) + self._sleep = sleep_between + self.calls: list[int] = [] + + async def generate(self, prev_audit, round_num, feedback_md): + del prev_audit, feedback_md + self.calls.append(round_num) + if self._sleep: + await asyncio.sleep(self._sleep) + if not self._outputs: + raise AssertionError("scripted writer exhausted") + return WriterOutput(content=self._outputs.pop(0)) + + +class _ScriptedAuditor: + role_id = "auditor" + + def __init__( + self, + verdicts: list[str], + scores: list[float | None] | None = None, + ) -> None: + self._verdicts = list(verdicts) + self._scores = list(scores) if scores else [None] * len(verdicts) + self.calls: list[int] = [] + + async def verify(self, writer_output, round_num): + self.calls.append(round_num) + verdict = self._verdicts.pop(0) + score = self._scores.pop(0) if self._scores else None + meta: dict = {} + if score is not None: + meta["score"] = score + return AuditReport( + verdict=verdict, + findings=[ + AuditFinding( + category="grade", severity="info", + short_message=f"score={score}", + ), + ], + metadata=meta, + ) + + +# ── D1: Wall-clock budget ─────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_wall_clock_budget_exits_between_rounds( + tmp_path: Path, +) -> None: + """Each round sleeps 0.1s; budget is 0.25s. Should run round 0 + successfully (records duration ~0.1s), check at round 1 boundary + (elapsed=0.1, remaining=0.15, est_next=0.11), pass; then check + again at round 2 boundary (elapsed=0.2, remaining=0.05, + est_next=0.11), fail → exit.""" + writer = _ScriptedWriter(["v0", "v1", "v2"], sleep_between=0.1) + auditor = _ScriptedAuditor(["iterate", "iterate", "iterate"]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, + output_check=lambda _: [], + max_rounds=10, + max_wall_seconds=0.25, + ) + output = await cycle.run() + assert output.reason == "wall_clock_exhausted" + assert output.success is False + assert 1 <= output.rounds_used <= 3 # ran 1-2 rounds before exhausting + + +@pytest.mark.asyncio +async def test_wall_clock_safety_ratio_honoured(tmp_path: Path) -> None: + """With ratio=2.0 and round_duration ~0.05s, budget 0.15s should + only allow round 0 (after, est_next=0.10, remaining=0.10, equal → + fail).""" + writer = _ScriptedWriter(["v0", "v1", "v2"], sleep_between=0.05) + auditor = _ScriptedAuditor(["iterate", "iterate", "iterate"]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, + output_check=lambda _: [], + max_rounds=10, + max_wall_seconds=0.15, + wall_clock_safety_ratio=2.0, + ) + output = await cycle.run() + assert output.reason == "wall_clock_exhausted" + + +@pytest.mark.asyncio +async def test_no_wall_clock_means_no_check(tmp_path: Path) -> None: + """Default: cycle runs to max_rounds regardless of duration.""" + writer = _ScriptedWriter(["v0", "v1"], sleep_between=0.05) + auditor = _ScriptedAuditor(["iterate", "iterate"]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, + output_check=lambda _: [], + max_rounds=2, + ) + output = await cycle.run() + assert output.reason == "max_rounds_exhausted" + assert output.rounds_used == 2 + + +def test_wall_clock_invalid_value_rejected(tmp_path: Path) -> None: + writer = _ScriptedWriter(["x"]) + auditor = _ScriptedAuditor(["iterate"]) + with pytest.raises(ValueError, match="max_wall_seconds"): + WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_wall_seconds=0, + ) + + +def test_wall_clock_safety_ratio_below_one_rejected(tmp_path: Path) -> None: + writer = _ScriptedWriter(["x"]) + auditor = _ScriptedAuditor(["iterate"]) + with pytest.raises(ValueError, match="wall_clock_safety_ratio"): + WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], wall_clock_safety_ratio=0.5, + ) + + +# ── D2: RoundObserver hook firing ────────────────────────────────────────── + + +class _RecordingObserver(BaseRoundObserver): + """Records every hook call for inspection.""" + + def __init__(self) -> None: + self.events: list[tuple] = [] + + async def on_cycle_start(self, ctx): + self.events.append(("cycle_start", ctx.max_rounds)) + + async def on_round_start(self, round_num, prev_audit): + self.events.append(( + "round_start", round_num, + prev_audit.verdict if prev_audit else None, + )) + + async def on_writer_done(self, round_num, writer_out): + self.events.append(("writer_done", round_num, writer_out.content)) + return None + + async def on_audit_done(self, round_num, writer_out, audit): + self.events.append(("audit_done", round_num, audit.verdict)) + return None + + async def on_cycle_end(self, output): + self.events.append(("cycle_end", output.reason)) + + +@pytest.mark.asyncio +async def test_all_hooks_fire_in_order(tmp_path: Path) -> None: + obs = _RecordingObserver() + writer = _ScriptedWriter(["v0", "v1"]) + auditor = _ScriptedAuditor(["iterate", "success"]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=3, observers=[obs], + ) + output = await cycle.run() + assert output.reason == "verdict=success" + + expected = [ + ("cycle_start", 3), + ("round_start", 0, None), + ("writer_done", 0, "v0"), + ("audit_done", 0, "iterate"), + ("round_start", 1, "iterate"), + ("writer_done", 1, "v1"), + ("audit_done", 1, "success"), + ("cycle_end", "verdict=success"), + ] + assert obs.events == expected + + +@pytest.mark.asyncio +async def test_observer_exception_isolated(tmp_path: Path) -> None: + """Observer raising in any hook is logged but cycle continues.""" + + class _BoomObserver(BaseRoundObserver): + async def on_round_start(self, round_num, prev_audit): + raise RuntimeError("boom") + + obs = _BoomObserver() + rec = _RecordingObserver() + writer = _ScriptedWriter(["v0"]) + auditor = _ScriptedAuditor(["success"]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=1, observers=[obs, rec], + ) + output = await cycle.run() + assert output.success is True + # Recording observer still saw all events + assert any(e[0] == "audit_done" for e in rec.events) + assert any(e[0] == "cycle_end" for e in rec.events) + + +@pytest.mark.asyncio +async def test_observer_abort_from_writer_done(tmp_path: Path) -> None: + class _AbortAfterWriter(BaseRoundObserver): + async def on_writer_done(self, round_num, writer_out): + return RoundIntervention.ABORT + + writer = _ScriptedWriter(["v0", "v1"]) + auditor = _ScriptedAuditor(["success", "success"]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=5, + observers=[_AbortAfterWriter()], + ) + output = await cycle.run() + assert output.reason == "observer_aborted" + assert output.rounds_used == 1 + assert output.final_writer_output.content == "v0" + # Auditor should NOT have been called (abort fired before audit) + assert auditor.calls == [] + + +@pytest.mark.asyncio +async def test_observer_abort_from_audit_done(tmp_path: Path) -> None: + class _AbortAfterAudit(BaseRoundObserver): + async def on_audit_done(self, round_num, writer_out, audit): + if round_num >= 1: + return RoundIntervention.ABORT + return None + + writer = _ScriptedWriter(["v0", "v1", "v2"]) + auditor = _ScriptedAuditor(["iterate", "iterate", "iterate"]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=10, + observers=[_AbortAfterAudit()], + ) + output = await cycle.run() + assert output.reason == "observer_aborted" + assert output.rounds_used == 2 # ran round 0 + round 1, aborted after + + +@pytest.mark.asyncio +async def test_on_cycle_end_fires_on_every_exit(tmp_path: Path) -> None: + """Ensures on_cycle_end is called for verdict-success, max-rounds, + wall-clock, and observer-abort exits.""" + end_reasons: list[str] = [] + + class _EndCapture(BaseRoundObserver): + async def on_cycle_end(self, output): + end_reasons.append(output.reason) + + # Case 1: terminal verdict + cycle = WriteAuditCycle( + writer=_ScriptedWriter(["v0"]), + auditor=_ScriptedAuditor(["success"]), + work_dir=tmp_path / "a", + output_check=lambda _: [], max_rounds=1, + observers=[_EndCapture()], + ) + await cycle.run() + + # Case 2: max_rounds + cycle = WriteAuditCycle( + writer=_ScriptedWriter(["v0"]), + auditor=_ScriptedAuditor(["iterate"]), + work_dir=tmp_path / "b", + output_check=lambda _: [], max_rounds=1, + observers=[_EndCapture()], + ) + await cycle.run() + + assert end_reasons == ["verdict=success", "max_rounds_exhausted"] + + +@pytest.mark.asyncio +async def test_cycle_context_passed_to_on_cycle_start(tmp_path: Path) -> None: + captured: list[CycleContext] = [] + + class _CtxCapture(BaseRoundObserver): + async def on_cycle_start(self, ctx): + captured.append(ctx) + + cycle = WriteAuditCycle( + writer=_ScriptedWriter(["v0"]), + auditor=_ScriptedAuditor(["success"]), + work_dir=tmp_path, + output_check=lambda _: [], max_rounds=5, + max_wall_seconds=99.0, + observers=[_CtxCapture()], + ) + await cycle.run() + assert len(captured) == 1 + ctx = captured[0] + assert ctx.max_rounds == 5 + assert ctx.max_wall_seconds == 99.0 + assert ctx.work_dir == tmp_path + + +# ── BestSoFarObserver ────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_best_so_far_observer_invokes_callback_per_round( + tmp_path: Path, +) -> None: + snapshots: list[tuple[int, str, float]] = [] + + def _cb(idx, w, a): + snapshots.append((idx, w.content, a.metadata.get("score"))) + + obs = BestSoFarObserver(callback=_cb) + writer = _ScriptedWriter(["v0", "v1", "v2"]) + auditor = _ScriptedAuditor( + ["iterate", "iterate", "iterate"], + scores=[3.0, 7.0, 5.0], + ) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=3, observers=[obs], + ) + await cycle.run() + + assert snapshots == [ + (0, "v0", 3.0), + (1, "v1", 7.0), + (1, "v1", 7.0), # round 2 (score 5) didn't beat round 1 (score 7) + ] + + +@pytest.mark.asyncio +async def test_best_so_far_observer_supports_async_callback( + tmp_path: Path, +) -> None: + received: list[float] = [] + + async def _async_cb(idx, w, a): + await asyncio.sleep(0) + received.append(a.metadata.get("score")) + + obs = BestSoFarObserver(callback=_async_cb) + writer = _ScriptedWriter(["v0", "v1"]) + auditor = _ScriptedAuditor( + ["iterate", "iterate"], scores=[1.0, 7.0], + ) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=2, observers=[obs], + ) + await cycle.run() + assert received == [1.0, 7.0] + + +@pytest.mark.asyncio +async def test_best_so_far_callback_error_does_not_crash_cycle( + tmp_path: Path, +) -> None: + def _bad_cb(idx, w, a): + raise RuntimeError("kaboom") + + obs = BestSoFarObserver(callback=_bad_cb) + writer = _ScriptedWriter(["v0", "v1"]) + auditor = _ScriptedAuditor(["iterate", "success"], scores=[1.0, 7.0]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=5, observers=[obs], + ) + output = await cycle.run() + assert output.success is True + + +# ── PlateauAbortObserver ─────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_plateau_abort_after_no_improvement(tmp_path: Path) -> None: + """Scores: 3, 5, 5, 5 → no improvement over historical max=5 in + last 3 → abort at round 4 (after 4 rounds).""" + obs = PlateauAbortObserver(plateau_rounds=3) + writer = _ScriptedWriter(["v0", "v1", "v2", "v3", "v4"]) + auditor = _ScriptedAuditor( + ["iterate"] * 5, scores=[3.0, 5.0, 5.0, 5.0, 5.0], + ) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=10, observers=[obs], + ) + output = await cycle.run() + assert output.reason == "observer_aborted" + # rounds 0,1,2,3 ran (4 rounds). After round 3 the plateau check + # had history [3,5,5,5] — recent=[5,5,5], earlier=[3] → max(recent)=5 + # > max(earlier)=3 → no abort. After round 4 (score 5): + # history=[3,5,5,5,5], recent=[5,5,5], earlier=[3,5] → max(recent)=5 + # = max(earlier)=5 → abort. + assert output.rounds_used == 5 + + +@pytest.mark.asyncio +async def test_plateau_skips_rounds_without_score(tmp_path: Path) -> None: + """Synthetic audits (score=None) should not advance plateau counter.""" + obs = PlateauAbortObserver(plateau_rounds=2) + writer = _ScriptedWriter(["v0", "v1", "v2"]) + auditor = _ScriptedAuditor( + ["iterate", "iterate", "success"], scores=[5.0, None, 7.0], + ) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=3, observers=[obs], + ) + output = await cycle.run() + # Score history seen by plateau: [5.0, 7.0] → improvement → no abort + assert output.reason == "verdict=success" + + +def test_plateau_invalid_param_rejected() -> None: + with pytest.raises(ValueError, match="plateau_rounds"): + PlateauAbortObserver(plateau_rounds=0) + + +# ── MetricsObserver ──────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_metrics_observer_emits_per_round(tmp_path: Path) -> None: + captured: list[dict] = [] + + obs = MetricsObserver(callback=lambda m: captured.append(m)) + writer = _ScriptedWriter(["abc", "defgh"]) + auditor = _ScriptedAuditor( + ["iterate", "success"], scores=[3.0, 7.0], + ) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=2, observers=[obs], + ) + await cycle.run() + + assert len(captured) == 2 + assert captured[0]["round"] == 0 + assert captured[0]["score"] == 3.0 + assert captured[0]["verdict"] == "iterate" + assert captured[0]["writer_chars"] == 3 + assert captured[1]["round"] == 1 + assert captured[1]["score"] == 7.0 + assert captured[1]["writer_chars"] == 5 + + +@pytest.mark.asyncio +async def test_metrics_observer_default_logs( + tmp_path: Path, caplog: pytest.LogCaptureFixture, +) -> None: + obs = MetricsObserver() # no callback → logs at INFO + writer = _ScriptedWriter(["x"]) + auditor = _ScriptedAuditor(["success"], scores=[7.0]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, work_dir=tmp_path, + output_check=lambda _: [], max_rounds=1, observers=[obs], + ) + with caplog.at_level(logging.INFO, logger="agent_core.components.cycle.observers_builtin"): + await cycle.run() + assert any("cycle round metrics" in r.message for r in caplog.records) diff --git a/tests/test_cycle_protocols_shared.py b/tests/test_cycle_protocols_shared.py new file mode 100644 index 0000000..83b1986 --- /dev/null +++ b/tests/test_cycle_protocols_shared.py @@ -0,0 +1,66 @@ +"""Tests for cycle protocols + verifier protocol contract within cycle scope.""" +from __future__ import annotations + +from agent_core.components.cycle.protocols import FeedbackRenderer +from agent_core.components.cycle.types import AuditReport, WriterOutput +from agent_core.components.verifier import Generator, Verifier + + +class _FakeGenerator: + role_id = "fake_writer" + + async def generate( + self, + prev_audit: AuditReport | None, + round_num: int, + feedback_md: str, + ) -> WriterOutput: + return WriterOutput(content=f"round {round_num}") + + +class _FakeVerifier: + role_id = "fake_auditor" + + async def verify( + self, + writer_output: WriterOutput, + round_num: int, + ) -> AuditReport: + return AuditReport(verdict="success") + + +class _FakeRenderer: + def render(self, audit: AuditReport) -> str: + return f"verdict={audit.verdict}" + + +class _MissingMethod: + role_id = "broken" + + +def test_generator_protocol_passes_isinstance() -> None: + assert isinstance(_FakeGenerator(), Generator) + + +def test_verifier_protocol_passes_isinstance() -> None: + assert isinstance(_FakeVerifier(), Verifier) + + +def test_feedback_renderer_protocol_passes_isinstance() -> None: + assert isinstance(_FakeRenderer(), FeedbackRenderer) + + +def test_missing_method_fails_isinstance() -> None: + """runtime_checkable protocols verify method presence.""" + assert not isinstance(_MissingMethod(), Generator) + assert not isinstance(_MissingMethod(), Verifier) + + +def test_renderer_does_not_require_role_id() -> None: + """FeedbackRenderer is method-only — no role_id attribute.""" + + class _PureRenderer: + def render(self, audit: AuditReport) -> str: + return "x" + + assert isinstance(_PureRenderer(), FeedbackRenderer) diff --git a/tests/test_cycle_types_shared.py b/tests/test_cycle_types_shared.py new file mode 100644 index 0000000..9483abf --- /dev/null +++ b/tests/test_cycle_types_shared.py @@ -0,0 +1,177 @@ +"""Tests for agent_core/components/cycle/types.py — round-trip and defaults.""" +from __future__ import annotations + +from agent_core.components.cycle.types import ( + AuditFinding, + AuditReport, + CycleOutput, + WriterOutput, +) + +# ── AuditFinding ──────────────────────────────────────────────────────────── + + +def test_finding_minimal_construction() -> None: + f = AuditFinding( + category="missing_section", + severity="error", + short_message="Missing intro", + ) + assert f.detailed_message == "" + assert f.file is None + assert f.line is None + assert f.metadata == {} + + +def test_finding_round_trip_minimal() -> None: + f = AuditFinding( + category="x", severity="warning", short_message="msg", + ) + restored = AuditFinding.from_dict(f.to_dict()) + assert restored == f + + +def test_finding_round_trip_full() -> None: + f = AuditFinding( + category="citation_invalid", + severity="error", + short_message="cite missing", + detailed_message="full detail", + file="paper.md", + line=42, + snippet="...as Smith et al. show...", + suggested_action="add citation [SMITH2024]", + target_role="method_writer", + metadata={"source_id": "abc-123"}, + ) + restored = AuditFinding.from_dict(f.to_dict()) + assert restored == f + + +# ── AuditReport ───────────────────────────────────────────────────────────── + + +def test_report_default_construction() -> None: + r = AuditReport(verdict="success") + assert r.findings == [] + assert r.summary == "" + assert r.confidence == 0.0 + assert r.raw_text == "" + + +def test_report_round_trip_with_findings() -> None: + r = AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="a", severity="error", short_message="m1", + ), + AuditFinding( + category="b", severity="warning", short_message="m2", + file="f.py", line=10, + ), + ], + summary="Two issues found in this round.", + confidence=0.85, + raw_text="raw auditor output", + metadata={"audit_round": 2}, + ) + restored = AuditReport.from_dict(r.to_dict()) + assert restored == r + + +def test_report_json_round_trip() -> None: + r = AuditReport( + verdict="success", + findings=[ + AuditFinding( + category="info_only", + severity="info", + short_message="LGTM", + ), + ], + summary="No blockers.", + confidence=0.95, + ) + text = r.to_json() + restored = AuditReport.from_json(text) + assert restored == r + + +def test_report_from_dict_tolerates_missing_optional_fields() -> None: + """Older serialized forms missing newer fields should still load.""" + minimal = {"verdict": "iterate", "findings": []} + r = AuditReport.from_dict(minimal) + assert r.verdict == "iterate" + assert r.findings == [] + assert r.summary == "" + assert r.confidence == 0.0 + + +# ── WriterOutput ──────────────────────────────────────────────────────────── + + +def test_writer_output_round_trip_drops_loop_result() -> None: + """loop_result is intentionally not serialized.""" + wo = WriterOutput( + content="draft text", + files=["sections/intro.md", "sections/method.md"], + message_count=42, + metadata={"turns_used": 12}, + loop_result={"unserializable": object()}, + ) + restored = WriterOutput.from_dict(wo.to_dict()) + assert restored.content == wo.content + assert restored.files == wo.files + assert restored.message_count == wo.message_count + assert restored.metadata == wo.metadata + assert restored.loop_result is None # dropped on round-trip + + +# ── CycleOutput ───────────────────────────────────────────────────────────── + + +def test_cycle_output_empty_history() -> None: + co = CycleOutput( + success=False, + rounds_used=0, + final_writer_output=None, + final_audit=None, + reason="never_started", + ) + restored = CycleOutput.from_dict(co.to_dict()) + assert restored == co + + +def test_cycle_output_full_round_trip() -> None: + wo1 = WriterOutput(content="draft v1", files=["a.md"], message_count=10) + a1 = AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="missing_section", + severity="error", + short_message="add intro", + ), + ], + summary="Need an intro.", + ) + wo2 = WriterOutput(content="draft v2", files=["a.md"], message_count=20) + a2 = AuditReport(verdict="success", summary="LGTM.", confidence=0.9) + co = CycleOutput( + success=True, + rounds_used=2, + final_writer_output=wo2, + final_audit=a2, + history=[(wo1, a1), (wo2, a2)], + reason="verdict=success", + ) + text = co.to_json() + restored = CycleOutput.from_json(text) + assert restored.success == co.success + assert restored.rounds_used == co.rounds_used + assert restored.final_audit == a2 + assert len(restored.history) == 2 + assert restored.history[0][1] == a1 + assert restored.history[1][1] == a2 + assert restored.reason == "verdict=success" diff --git a/tests/test_verifier_compat_shared.py b/tests/test_verifier_compat_shared.py new file mode 100644 index 0000000..f2a4bae --- /dev/null +++ b/tests/test_verifier_compat_shared.py @@ -0,0 +1,196 @@ +"""Round-trip Verdict ↔ AuditReport via _compat helpers.""" + +from __future__ import annotations + +import inspect + +import pytest + +from agent_core.components.cycle.types import ( + AuditFinding, + AuditReport, + WriterOutput, +) +from agent_core.components.verifier import ( + Finding, + Verdict, + VerifierContext, + cycle_auditor_from_verifier, +) +from agent_core.components.verifier._compat import ( + audit_report_from_verdict, + verdict_from_audit_report, +) + + +def test_audit_report_to_verdict_passed(): + report = AuditReport( + verdict="success", + findings=[ + AuditFinding( + category="missing_section", + severity="warning", + short_message="missing intro", + file="paper.md", + ) + ], + summary="Looks fine overall.", + confidence=0.92, + raw_text="", + metadata={"reviewer": "auditor-3"}, + ) + v = verdict_from_audit_report(report) + assert v.passed is True + assert v.score == 0.92 + assert v.reasoning == "Looks fine overall." + assert len(v.findings) == 1 + assert v.findings[0].severity == "warning" + assert v.findings[0].message == "missing intro" + assert v.findings[0].location == "paper.md" + assert v.metadata["_cycle_verdict"] == "success" + assert v.metadata["_cycle_raw_text"] == "" + assert v.metadata["reviewer"] == "auditor-3" + + +def test_audit_report_to_verdict_failed(): + report = AuditReport(verdict="iterate", confidence=0.4) + v = verdict_from_audit_report(report) + assert v.passed is False + + +def test_verdict_to_audit_report_roundtrip_preserves_label(): + original = AuditReport( + verdict="abandon", + findings=[ + AuditFinding( + category="parse_failure", + severity="error", + short_message="bad json", + ) + ], + summary="Can't proceed.", + confidence=0.1, + raw_text="", + metadata={"k": "v"}, + ) + bridged = audit_report_from_verdict(verdict_from_audit_report(original)) + assert bridged.verdict == "abandon" + assert bridged.summary == "Can't proceed." + assert bridged.confidence == 0.1 + assert bridged.raw_text == "" + assert bridged.metadata["k"] == "v" + assert "_cycle_verdict" not in bridged.metadata + assert len(bridged.findings) == 1 + assert bridged.findings[0].severity == "error" + + +def test_verdict_to_audit_report_uses_default_fail_label(): + v = Verdict(passed=False, reasoning="failure") + report = audit_report_from_verdict(v, default_fail_label="iterate") + assert report.verdict == "iterate" + + +def test_verdict_to_audit_report_passed_uses_success_label(): + v = Verdict(passed=True) + report = audit_report_from_verdict(v) + assert report.verdict == "success" + + +def test_finding_with_unknown_severity_coerces_to_info(): + v = Verdict( + passed=False, + findings=[Finding(severity="critical", message="boom")], + ) + report = audit_report_from_verdict(v) + assert report.findings[0].severity == "info" + assert report.findings[0].category == "critical" + + +# ── cycle_auditor_from_verifier adapter ──────────────────────────────── + + +class _RecordingVerifier: + role_id = "recording" + + def __init__(self, verdict: Verdict) -> None: + self._verdict = verdict + self.last_subject: object | None = None + self.last_ctx: VerifierContext | None = None + + async def verify(self, subject, ctx): + self.last_subject = subject + self.last_ctx = ctx + return self._verdict + + +@pytest.mark.asyncio +async def test_adapter_threads_writer_output_and_round_to_verifier(): + inner = _RecordingVerifier(Verdict(passed=True, score=0.8)) + auditor = cycle_auditor_from_verifier(inner) + writer_out = WriterOutput(content="hello") + report = await auditor.verify(writer_out, round_num=2) + + assert isinstance(report, AuditReport) + assert inner.last_subject is writer_out + assert inner.last_ctx is not None + assert inner.last_ctx.is_runtime is True + assert inner.last_ctx.metadata["round_num"] == 2 + + +@pytest.mark.asyncio +async def test_adapter_converts_passed_verdict_to_success(): + inner = _RecordingVerifier( + Verdict(passed=True, score=0.9, reasoning="all good"), + ) + auditor = cycle_auditor_from_verifier(inner) + report = await auditor.verify(WriterOutput(content="x"), round_num=0) + assert report.verdict == "success" + assert report.confidence == 0.9 + assert report.summary == "all good" + + +@pytest.mark.asyncio +async def test_adapter_converts_failed_verdict_to_iterate(): + inner = _RecordingVerifier( + Verdict( + passed=False, + findings=[Finding(severity="error", message="boom")], + ), + ) + auditor = cycle_auditor_from_verifier(inner) + report = await auditor.verify(WriterOutput(content="x"), round_num=0) + assert report.verdict == "iterate" + assert len(report.findings) == 1 + assert report.findings[0].severity == "error" + + +@pytest.mark.asyncio +async def test_adapter_custom_ctx_factory(): + inner = _RecordingVerifier(Verdict(passed=True)) + + def factory(round_num: int) -> VerifierContext: + return VerifierContext( + is_runtime=False, + metadata={"round_num": round_num, "tag": "eval"}, + ) + + auditor = cycle_auditor_from_verifier(inner, ctx_factory=factory) + await auditor.verify(WriterOutput(content="x"), round_num=5) + + assert inner.last_ctx is not None + assert inner.last_ctx.is_runtime is False + assert inner.last_ctx.metadata["tag"] == "eval" + + +def test_adapter_satisfies_cycle_auditor_shape(): + inner = _RecordingVerifier(Verdict(passed=True)) + auditor = cycle_auditor_from_verifier(inner) + assert hasattr(auditor, "role_id") and isinstance(auditor.role_id, str) + assert inspect.iscoroutinefunction(auditor.verify) + assert auditor.role_id == "recording" + + +def test_adapter_role_id_override(): + inner = _RecordingVerifier(Verdict(passed=True)) + auditor = cycle_auditor_from_verifier(inner, role_id="override") + assert auditor.role_id == "override" diff --git a/tests/test_verifier_composers_shared.py b/tests/test_verifier_composers_shared.py new file mode 100644 index 0000000..1d374a7 --- /dev/null +++ b/tests/test_verifier_composers_shared.py @@ -0,0 +1,387 @@ +"""Composer semantics + nesting.""" + +from __future__ import annotations + +import pytest + +from agent_core.components.verifier import ( + Cascade, + ConsensusVerifier, + Ensemble, + Fallback, + Parallel, + Pipeline, + Verdict, + VerifierContext, +) + + +class _Stub: + """Returns a fixed verdict; records call count.""" + + def __init__( + self, + *, + passed: bool = True, + score: float | None = None, + role_id: str = "stub", + ) -> None: + self._passed = passed + self._score = score + self.role_id = role_id + self.calls = 0 + + async def verify(self, subject, ctx): + self.calls += 1 + return Verdict( + passed=self._passed, + score=self._score, + reasoning=f"{self.role_id} returned passed={self._passed}", + ) + + +@pytest.fixture +def ctx(): + return VerifierContext(is_runtime=True) + + +# ── Pipeline ─────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_pipeline_all_pass(ctx): + a, b, c = _Stub(passed=True), _Stub(passed=True), _Stub(passed=True) + p = Pipeline(a, b, c) + out = await p.verify("x", ctx) + assert out.passed is True + assert len(out.sub_verdicts) == 3 + assert a.calls == b.calls == c.calls == 1 + + +@pytest.mark.asyncio +async def test_pipeline_short_circuits_on_first_fail(ctx): + a = _Stub(passed=True, role_id="a") + b = _Stub(passed=False, role_id="b") + c = _Stub(passed=True, role_id="c") + p = Pipeline(a, b, c) + out = await p.verify("x", ctx) + assert out.passed is False + assert len(out.sub_verdicts) == 2 + assert c.calls == 0 + assert "short-circuited at b" in out.reasoning + + +def test_pipeline_rejects_empty(): + with pytest.raises(ValueError): + Pipeline() + + +# ── Ensemble ─────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_ensemble_majority_pass(ctx): + e = Ensemble( + _Stub(passed=True), + _Stub(passed=True), + _Stub(passed=False), + aggregator="majority", + ) + out = await e.verify("x", ctx) + assert out.passed is True + assert len(out.sub_verdicts) == 3 + + +@pytest.mark.asyncio +async def test_ensemble_majority_fail(ctx): + e = Ensemble( + _Stub(passed=False), + _Stub(passed=False), + _Stub(passed=True), + aggregator="majority", + ) + out = await e.verify("x", ctx) + assert out.passed is False + + +@pytest.mark.asyncio +async def test_ensemble_score_aggregation(ctx): + e = Ensemble( + _Stub(passed=True, score=0.6), + _Stub(passed=True, score=0.8), + _Stub(passed=True, score=1.0), + aggregator="median", + ) + out = await e.verify("x", ctx) + assert out.score == 0.8 + + +@pytest.mark.asyncio +async def test_ensemble_score_min(ctx): + e = Ensemble( + _Stub(passed=True, score=0.5), + _Stub(passed=True, score=0.9), + aggregator="min", + ) + out = await e.verify("x", ctx) + assert out.score == 0.5 + + +def test_ensemble_rejects_unknown_aggregator(): + with pytest.raises(ValueError): + Ensemble(_Stub(), aggregator="totally-bogus") + + +# ── Fallback ─────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_fallback_primary_passes(ctx): + primary = _Stub(passed=True, role_id="primary") + backup = _Stub(passed=True, role_id="backup") + f = Fallback(primary, backup) + out = await f.verify("x", ctx) + assert out.passed is True + assert backup.calls == 0 + + +@pytest.mark.asyncio +async def test_fallback_primary_fails_backup_recovers(ctx): + primary = _Stub(passed=False, role_id="p") + backup_a = _Stub(passed=False, role_id="ba") + backup_b = _Stub(passed=True, role_id="bb") + f = Fallback(primary, backup_a, backup_b) + out = await f.verify("x", ctx) + assert out.passed is True + assert "bb" in out.reasoning + assert primary.calls == backup_a.calls == backup_b.calls == 1 + + +@pytest.mark.asyncio +async def test_fallback_all_fail(ctx): + f = Fallback( + _Stub(passed=False), + _Stub(passed=False), + _Stub(passed=False), + ) + out = await f.verify("x", ctx) + assert out.passed is False + + +# ── Cascade ──────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_cascade_cheap_confident_skips_expensive(ctx): + cheap = _Stub(passed=True, score=0.95, role_id="cheap") + expensive = _Stub(passed=True, score=0.99, role_id="exp") + c = Cascade(cheap, expensive, confidence_threshold=0.9) + out = await c.verify("x", ctx) + assert out.passed is True + assert expensive.calls == 0 + assert "confident" in out.reasoning + + +@pytest.mark.asyncio +async def test_cascade_low_confidence_escalates(ctx): + cheap = _Stub(passed=True, score=0.5, role_id="cheap") + expensive = _Stub(passed=False, score=0.4, role_id="exp") + c = Cascade(cheap, expensive, confidence_threshold=0.9) + out = await c.verify("x", ctx) + assert out.passed is False + assert expensive.calls == 1 + assert len(out.sub_verdicts) == 2 + + +@pytest.mark.asyncio +async def test_cascade_no_cheap_score_escalates(ctx): + cheap = _Stub(passed=True, score=None, role_id="cheap") + expensive = _Stub(passed=True, score=0.7, role_id="exp") + c = Cascade(cheap, expensive, confidence_threshold=0.9) + await c.verify("x", ctx) + assert expensive.calls == 1 + + +def test_cascade_threshold_validation(): + with pytest.raises(ValueError): + Cascade(_Stub(), _Stub(), confidence_threshold=1.5) + + +# ── Parallel ─────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_parallel_collects_all(ctx): + stubs = [_Stub(passed=True, role_id=f"s{i}") for i in range(5)] + p = Parallel(*stubs) + out = await p.verify("x", ctx) + assert out.passed is True + assert len(out.sub_verdicts) == 5 + + +@pytest.mark.asyncio +async def test_parallel_fails_if_any_fails(ctx): + p = Parallel( + _Stub(passed=True), + _Stub(passed=True), + _Stub(passed=False), + ) + out = await p.verify("x", ctx) + assert out.passed is False + assert len(out.sub_verdicts) == 3 + + +# ── Nesting ──────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_pipeline_of_ensemble(ctx): + """Pipeline can short-circuit on a composer's verdict.""" + inner_ok = Ensemble(_Stub(passed=True), _Stub(passed=True)) + inner_fail = Ensemble(_Stub(passed=False), _Stub(passed=False)) + final = _Stub(passed=True) + p = Pipeline(inner_ok, inner_fail, final) + out = await p.verify("x", ctx) + assert out.passed is False + assert len(out.sub_verdicts) == 2 + assert final.calls == 0 + + +@pytest.mark.asyncio +async def test_fallback_of_parallel_recovers(ctx): + failing = Parallel(_Stub(passed=False), _Stub(passed=True)) + succeeding = Parallel(_Stub(passed=True), _Stub(passed=True)) + f = Fallback(failing, succeeding) + out = await f.verify("x", ctx) + assert out.passed is True + assert len(out.sub_verdicts) == 2 + + +# ── ConsensusVerifier ────────────────────────────────────────────────── + + +class _AnswerStub: + """Verifier that surfaces a fixed answer in metadata.""" + + def __init__( + self, + *, + answer: str | None, + passed: bool = True, + role_id: str = "answer-stub", + ) -> None: + self._answer = answer + self._passed = passed + self.role_id = role_id + + async def verify(self, subject, ctx): + meta: dict[str, str] = {} + if self._answer is not None: + meta["answer"] = self._answer + return Verdict(passed=self._passed, metadata=meta) + + +class _RaisingStub: + role_id = "raising" + + async def verify(self, subject, ctx): + raise RuntimeError("simulated grader failure") + + +@pytest.mark.asyncio +async def test_consensus_strict_majority(ctx): + c = ConsensusVerifier( + _AnswerStub(answer="42"), + _AnswerStub(answer="42"), + _AnswerStub(answer="7"), + ) + out = await c.verify("subject", ctx) + assert out.passed is True + assert out.metadata["consensus_answer"] == "42" + assert out.metadata["n_succeeded"] == 3 + assert out.metadata["agreement"] == pytest.approx(2 / 3) + + +@pytest.mark.asyncio +async def test_consensus_no_majority_two_two_split(ctx): + c = ConsensusVerifier( + _AnswerStub(answer="a"), + _AnswerStub(answer="a"), + _AnswerStub(answer="b"), + _AnswerStub(answer="b"), + ) + out = await c.verify("subject", ctx) + assert out.passed is False # 2-2 split is not strict majority + # Counter.most_common breaks ties first-seen + assert out.metadata["consensus_answer"] == "a" + assert out.metadata["agreement"] == 0.5 + + +@pytest.mark.asyncio +async def test_consensus_resilient_to_single_failure(ctx): + c = ConsensusVerifier( + _AnswerStub(answer="42"), + _AnswerStub(answer="42"), + _RaisingStub(), + ) + out = await c.verify("subject", ctx) + assert out.passed is True # 2/2 valid agree → strict majority + assert out.metadata["n_total"] == 3 + assert out.metadata["n_succeeded"] == 2 + assert out.metadata["consensus_answer"] == "42" + assert any( + f.severity == "warning" and "sub-verifiers failed" in f.message + for f in out.findings + ) + + +@pytest.mark.asyncio +async def test_consensus_all_failed(ctx): + c = ConsensusVerifier(_RaisingStub(), _RaisingStub()) + out = await c.verify("subject", ctx) + assert out.passed is False + assert out.metadata["n_succeeded"] == 0 + assert out.sub_verdicts == [] + + +@pytest.mark.asyncio +async def test_consensus_missing_answer_key(ctx): + c = ConsensusVerifier( + _AnswerStub(answer=None, passed=True), + _AnswerStub(answer=None, passed=True), + ) + out = await c.verify("subject", ctx) + assert out.passed is False + assert "consensus_answer" not in out.metadata + assert any( + "no sub-verdict carried metadata" in f.message for f in out.findings + ) + + +@pytest.mark.asyncio +async def test_consensus_custom_answer_key(ctx): + class _ExtractedAnswerStub: + role_id = "extracted" + + def __init__(self, value: str) -> None: + self._v = value + + async def verify(self, subject, ctx): + return Verdict( + passed=True, metadata={"extracted_answer": self._v} + ) + + c = ConsensusVerifier( + _ExtractedAnswerStub("π/4"), + _ExtractedAnswerStub("π/4"), + _ExtractedAnswerStub("0"), + answer_key="extracted_answer", + ) + out = await c.verify("subject", ctx) + assert out.passed is True + assert out.metadata["consensus_answer"] == "π/4" + + +def test_consensus_rejects_empty(): + with pytest.raises(ValueError): + ConsensusVerifier() diff --git a/tests/test_verifier_protocols_shared.py b/tests/test_verifier_protocols_shared.py new file mode 100644 index 0000000..a1189c2 --- /dev/null +++ b/tests/test_verifier_protocols_shared.py @@ -0,0 +1,86 @@ +"""Verifier protocol contracts + ground-truth oracle isolation.""" + +from __future__ import annotations + +import pytest + +from agent_core.components.verifier import ( + Finding, + GroundTruth, + Verdict, + Verifier, + VerifierContext, +) + + +def test_runtime_strips_oracle_fields(): + gt = GroundTruth( + rubric="grade as IMO proof", + _reference="the official answer", + _formal_spec="lean code", + _test_cases=[1, 2, 3], + ) + ctx = VerifierContext(is_runtime=True, _ground_truth=gt) + runtime_view = ctx.ground_truth + assert runtime_view is not None + assert runtime_view.rubric == "grade as IMO proof" + assert runtime_view._reference is None + assert runtime_view._formal_spec is None + assert runtime_view._test_cases is None + + +def test_eval_keeps_oracle_fields(): + gt = GroundTruth( + rubric="grade as IMO proof", + _reference="the official answer", + _formal_spec="lean code", + ) + ctx = VerifierContext(is_runtime=False, _ground_truth=gt) + eval_view = ctx.ground_truth + assert eval_view is not None + assert eval_view._reference == "the official answer" + assert eval_view._formal_spec == "lean code" + + +def test_no_ground_truth_returns_none(): + ctx = VerifierContext(is_runtime=True, _ground_truth=None) + assert ctx.ground_truth is None + + +def test_runtime_metadata_isolation(): + gt = GroundTruth(rubric="r", metadata={"k": "v"}) + ctx = VerifierContext(is_runtime=True, _ground_truth=gt) + runtime_view = ctx.ground_truth + runtime_view.metadata["mutate"] = "after" + assert "mutate" not in gt.metadata + + +def test_verdict_defaults(): + v = Verdict() + assert v.score is None + assert v.passed is False + assert v.findings == [] + assert v.sub_verdicts == [] + assert v.metadata == {} + + +def test_finding_required_fields(): + f = Finding(severity="error", message="boom") + assert f.severity == "error" + assert f.message == "boom" + assert f.location is None + + +class _StubVerifier: + role_id = "stub" + + async def verify(self, subject, ctx): + return Verdict(passed=True) + + +@pytest.mark.asyncio +async def test_verifier_protocol_runtime_check(): + v = _StubVerifier() + assert isinstance(v, Verifier) + out = await v.verify("x", VerifierContext(is_runtime=True)) + assert out.passed is True diff --git a/tests/test_write_audit_cycle_shared.py b/tests/test_write_audit_cycle_shared.py new file mode 100644 index 0000000..2a57702 --- /dev/null +++ b/tests/test_write_audit_cycle_shared.py @@ -0,0 +1,367 @@ +"""Tests for WriteAuditCycle — six §7.1 cases against fake Writer/Auditor.""" +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +from agent_core.components.cycle.types import ( + AuditFinding, + AuditReport, + WriterOutput, +) +from agent_core.components.cycle.write_audit_cycle import WriteAuditCycle + + +class _ScriptedWriter: + """Returns scripted WriterOutputs in order. Re-uses last on overflow.""" + + role_id = "scripted_writer" + + def __init__(self, outputs: list[WriterOutput]) -> None: + self.outputs = outputs + self.calls: list[tuple[AuditReport | None, int, str]] = [] + + async def generate( + self, + prev_audit: AuditReport | None, + round_num: int, + feedback_md: str, + ) -> WriterOutput: + self.calls.append((prev_audit, round_num, feedback_md)) + idx = min(round_num, len(self.outputs) - 1) + return self.outputs[idx] + + +class _RaisingWriter: + role_id = "raising_writer" + + def __init__(self) -> None: + self.calls = 0 + + async def generate(self, *args, **kwargs): + self.calls += 1 + raise RuntimeError("kaboom") + + +class _ScriptedAuditor: + """Returns scripted AuditReports in order. Re-uses last on overflow.""" + + role_id = "scripted_auditor" + + def __init__(self, reports: list[AuditReport]) -> None: + self.reports = reports + self.calls: list[tuple[WriterOutput, int]] = [] + + async def verify( + self, writer_output: WriterOutput, round_num: int, + ) -> AuditReport: + self.calls.append((writer_output, round_num)) + idx = min(round_num, len(self.reports) - 1) + return self.reports[idx] + + +def _no_missing(_wd: Path) -> list[str]: + return [] + + +def _always_missing(missing: list[str]): + def _check(_wd: Path) -> list[str]: + return list(missing) + return _check + + +# ── Case 1: terminal-success on round 0 ───────────────────────────────────── + + +@pytest.mark.asyncio +async def test_success_on_round_0(tmp_path: Path) -> None: + writer = _ScriptedWriter([WriterOutput(content="draft v1")]) + auditor = _ScriptedAuditor([ + AuditReport(verdict="success", summary="LGTM"), + ]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + ) + out = await cycle.run() + assert out.success is True + assert out.rounds_used == 1 + assert out.reason == "verdict=success" + assert len(out.history) == 1 + assert len(writer.calls) == 1 + # Round 0 has no prior feedback + assert writer.calls[0] == (None, 0, "") + assert (tmp_path / "audit_round_0.json").exists() + + +# ── Case 2: iterate then success on a later round ─────────────────────────── + + +@pytest.mark.asyncio +async def test_iterate_then_success(tmp_path: Path) -> None: + writer = _ScriptedWriter([ + WriterOutput(content="v1"), + WriterOutput(content="v2"), + ]) + auditor = _ScriptedAuditor([ + AuditReport( + verdict="iterate", + findings=[ + AuditFinding( + category="missing_section", + severity="error", + short_message="add intro", + ), + ], + summary="Need intro.", + ), + AuditReport(verdict="success", summary="LGTM"), + ]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + ) + out = await cycle.run() + assert out.success is True + assert out.rounds_used == 2 + assert len(out.history) == 2 + # Round 1's feedback_md must be non-empty (contains the iterate audit's findings) + _, _, fb = writer.calls[1] + assert "add intro" in fb + assert "missing_section" in fb + # Both round files persisted + assert (tmp_path / "audit_round_0.json").exists() + assert (tmp_path / "audit_round_1.json").exists() + + +# ── Case 3: max_rounds exhausted ──────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_max_rounds_exhausted(tmp_path: Path) -> None: + writer = _ScriptedWriter([WriterOutput(content="x")]) + auditor = _ScriptedAuditor([AuditReport(verdict="iterate")]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + max_rounds=3, + ) + out = await cycle.run() + assert out.success is False + assert out.rounds_used == 3 + assert out.reason == "max_rounds_exhausted" + assert len(out.history) == 3 + assert writer.calls[0][2] == "" # round 0 has no feedback + assert writer.calls[1][2] != "" # round 1 sees prior audit + assert writer.calls[2][2] != "" + + +# ── Case 4: output_check failure short-circuits auditor ───────────────────── + + +@pytest.mark.asyncio +async def test_output_check_failure_short_circuits_auditor( + tmp_path: Path, +) -> None: + writer = _ScriptedWriter([WriterOutput(content="empty draft")]) + auditor = _ScriptedAuditor([ + AuditReport(verdict="success", summary="should not be called"), + ]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, + output_check=_always_missing(["sections/intro.md", "data/refs.bib"]), + max_rounds=2, + ) + out = await cycle.run() + # Auditor was never invoked + assert auditor.calls == [] + # Synthetic audit names every missing artifact + assert out.final_audit is not None + cats = {f.category for f in out.final_audit.findings} + assert cats == {"output_missing"} + files = {f.file for f in out.final_audit.findings} + assert files == {"sections/intro.md", "data/refs.bib"} + # max_rounds_exhausted because the writer keeps failing output_check + assert out.reason == "max_rounds_exhausted" + # Persistence still fires for synthetic audits + persisted = json.loads((tmp_path / "audit_round_0.json").read_text()) + assert persisted["metadata"]["synthetic"] is True + assert persisted["metadata"]["reason"] == "output_missing" + + +# ── Case 5: writer exception captured as writer_exception finding ─────────── + + +@pytest.mark.asyncio +async def test_writer_exception_captured_and_continues( + tmp_path: Path, +) -> None: + writer = _RaisingWriter() + auditor = _ScriptedAuditor([AuditReport(verdict="success")]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + max_rounds=2, + ) + out = await cycle.run() + # The cycle did NOT raise + assert writer.calls == 2 + # Auditor never invoked because output_check / writer never produced + assert auditor.calls == [] + # Both rounds have a synthetic writer_exception audit + assert out.success is False + assert out.reason == "max_rounds_exhausted" + for _, audit in out.history: + assert audit.metadata.get("reason") == "writer_exception" + assert any( + f.category == "writer_exception" for f in audit.findings + ) + # Exception type is surfaced in the short message + assert any("RuntimeError" in f.short_message for f in audit.findings) + # Persistence happened for both rounds + assert (tmp_path / "audit_round_0.json").exists() + assert (tmp_path / "audit_round_1.json").exists() + + +# ── Case 6: persistence round-trips to_json / from_json ───────────────────── + + +@pytest.mark.asyncio +async def test_persisted_audit_round_trips(tmp_path: Path) -> None: + writer = _ScriptedWriter([WriterOutput(content="v1")]) + auditor = _ScriptedAuditor([ + AuditReport( + verdict="success", + findings=[ + AuditFinding( + category="info_only", + severity="info", + short_message="all good", + metadata={"score": 0.95}, + ), + ], + summary="Looks great.", + confidence=0.95, + metadata={"reviewer_id": "auditor_R0"}, + ), + ]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + ) + await cycle.run() + text = (tmp_path / "audit_round_0.json").read_text() + restored = AuditReport.from_json(text) + assert restored.verdict == "success" + assert restored.findings[0].category == "info_only" + assert restored.findings[0].metadata == {"score": 0.95} + assert restored.summary == "Looks great." + assert restored.confidence == 0.95 + assert restored.metadata == {"reviewer_id": "auditor_R0"} + + +# ── Extra: invariants ─────────────────────────────────────────────────────── + + +def test_max_rounds_must_be_positive() -> None: + with pytest.raises(ValueError, match="max_rounds"): + WriteAuditCycle( + writer=_ScriptedWriter([]), + auditor=_ScriptedAuditor([]), + work_dir=Path("/tmp"), + output_check=_no_missing, + max_rounds=0, + ) + + +def test_success_must_subset_terminal() -> None: + with pytest.raises(ValueError, match="success_verdicts"): + WriteAuditCycle( + writer=_ScriptedWriter([]), + auditor=_ScriptedAuditor([]), + work_dir=Path("/tmp"), + output_check=_no_missing, + terminal_verdicts={"abandon"}, + success_verdicts={"success"}, # not in terminal + ) + + +@pytest.mark.asyncio +async def test_abandon_terminates_with_failure(tmp_path: Path) -> None: + writer = _ScriptedWriter([WriterOutput(content="x")]) + auditor = _ScriptedAuditor([AuditReport(verdict="abandon")]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + ) + out = await cycle.run() + assert out.success is False + assert out.reason == "verdict=abandon" + assert out.rounds_used == 1 + + +@pytest.mark.asyncio +async def test_custom_terminal_verdicts(tmp_path: Path) -> None: + """Caller can use domain-specific verdict vocabulary.""" + writer = _ScriptedWriter([WriterOutput(content="x")]) + auditor = _ScriptedAuditor([AuditReport(verdict="merge")]) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + terminal_verdicts={"merge", "discard"}, + success_verdicts={"merge"}, + ) + out = await cycle.run() + assert out.success is True + assert out.reason == "verdict=merge" + + +# ── Day 5: cycle accepts unified Verifier via adapter ─────────────────────── + + +@pytest.mark.asyncio +async def test_cycle_runs_with_adapter_wrapped_verifier(tmp_path: Path) -> None: + """End-to-end: components.verifier.Verifier → cycle via adapter.""" + from agent_core.components.verifier import ( + Verdict, + cycle_auditor_from_verifier, + ) + + class _UnifiedVerifier: + role_id = "unified" + + def __init__(self, verdicts: list[Verdict]) -> None: + self.verdicts = verdicts + self.calls: list[int] = [] + + async def verify(self, subject, ctx) -> Verdict: + self.calls.append(ctx.metadata["round_num"]) + idx = min(len(self.calls) - 1, len(self.verdicts) - 1) + return self.verdicts[idx] + + writer = _ScriptedWriter([ + WriterOutput(content="draft v1"), + WriterOutput(content="draft v2"), + ]) + inner = _UnifiedVerifier([ + Verdict(passed=False, score=0.3, reasoning="needs work"), + Verdict(passed=True, score=0.95, reasoning="ship it"), + ]) + auditor = cycle_auditor_from_verifier(inner) + cycle = WriteAuditCycle( + writer=writer, auditor=auditor, + work_dir=tmp_path, output_check=_no_missing, + max_rounds=3, + ) + out = await cycle.run() + + assert out.success is True + assert out.reason == "verdict=success" + assert out.rounds_used == 2 + assert inner.calls == [0, 1] + assert out.final_audit is not None + assert out.final_audit.confidence == 0.95 + assert out.final_audit.summary == "ship it" diff --git a/uv.lock b/uv.lock index 95b8de1..28499df 100644 --- a/uv.lock +++ b/uv.lock @@ -50,7 +50,7 @@ wheels = [ [[package]] name = "apodex-agent-core" -version = "0.4.0" +version = "0.5.0" source = { editable = "." } dependencies = [ { name = "anthropic", extra = ["bedrock"] },