diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index f348552..72c7336 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -48,3 +48,27 @@ jobs: # Live test execution requires ANTHROPIC_API_KEY + COMPOSIO_API_KEY # and runs locally before push, not in CI. run: pytest --collect-only -q + + agent-team-tests: + # The agent-team/ subproject is self-contained (no API keys needed), so its + # suite RUNS in CI rather than only collecting. Its tests/ package collides + # with the repo-root tests/ under one rootdir, so it runs in its own dir. + runs-on: ubuntu-latest + timeout-minutes: 10 + steps: + - uses: actions/checkout@v6 + + - uses: actions/setup-python@v5 + with: + python-version: "3.12" + cache: pip + cache-dependency-path: requirements.txt + + - name: Install dependencies + run: | + pip install -r requirements.txt + pip install pytest + + - name: Run agent-team suite + working-directory: agent-team + run: python -m pytest -q diff --git a/agent-team/.gitignore b/agent-team/.gitignore new file mode 100644 index 0000000..1a8546b --- /dev/null +++ b/agent-team/.gitignore @@ -0,0 +1,18 @@ +# Local durable state — never commit (SQLite stores + integrity sidecars) +*.sqlite +*.sqlite-wal +*.sqlite-shm +*.db +*.meta.json +state/ +.state/ + +# Secrets — never commit +*.env +secrev.env + +# Python artifacts +__pycache__/ +*.pyc +.pytest_cache/ +.ruff_cache/ diff --git a/agent-team/.security-review/suppressions.json b/agent-team/.security-review/suppressions.json new file mode 100644 index 0000000..63bf824 --- /dev/null +++ b/agent-team/.security-review/suppressions.json @@ -0,0 +1,25 @@ +{ + "_comment": "Written justifications for sh-security-review (Path A) + GPT-4.1 cross-review findings on the R720 agent-team Plane-2 scaffold that are deliberately NOT fixed in this pre-deployment commit. Per CLAUDE.md: a confirmed critical/high is either fixed or suppressed with a written justification. Every item here is design-level / deferred-to-P1-build and is NOT live-exploitable because nothing in agent-team/ is provisioned, scheduled, or enabled. The proven-exploitable HIGHs (CI-guard fnmatch '**/' denylist bypass and the scope '**' bypass) were FIXED, not suppressed (see ci/agent-team-apply-verify.yml + tests/test_ci_gate_workflow.py).", + "suppressions": [ + { + "id": "LOGIC-01/02/03-resume-worker-durability", + "justification": "The resume worker's no-double-apply / single-flight currently rests on an in-process threading.Lock + a turn-integer guard, and an 'answered' ledger row has no terminal 'resumed' transition, so a restart sweep re-enqueues it (correctness then relies on the turn guard). This is a real durability gap, but: (1) it is PRE-DEPLOYMENT scaffolding — no resume worker, responder, or scheduler runs anywhere, so it cannot be triggered in the current state; (2) the design (§3.3.1) explicitly frames the durable, cross-process single-flight resume worker as P1-build-proper. FIX TRACKED FOR P1 BUILD: add a terminal 'resumed' state to QUESTION_STATES, flip 'answered'->'resumed' via the BEGIN IMMEDIATE compare-and-set inside the resume success path (making the CAS the cross-process single-flight token), have the turn guard compare question_id identity (not just the turn integer), and filter recovery selects to exclude resumed rows. Until the worker is built and enabled, no merge of live resume behavior occurs." + }, + { + "id": "SC-01-diff-hash-threat-model", + "justification": "The unkeyed SHA-256 diff-integrity hash defends in-transit corruption/substitution between the trusted dispatcher and CI, NOT a compromised box (the box is in the trusted computing base for what it proposes). This is already stated honestly in the design doc §3.3.2 ('Threat-model honesty (the diff-hash does not cover a compromised box)') and the ci/README.md threat-model section has been aligned to match. The real backstops against a bad box are the credential-less untrusted build job, the trust-control-surface denylist, and the mandatory human review of the draft PR + required checks. A keyed/attested binding to the approved-plan record is a P1/P3 hardening, not a current vuln." + }, + { + "id": "ATCI-003-004-build-test-egress-breadth", + "justification": "The build-test egress allowlist includes GitHub API / githubusercontent / a wildcard actions host beyond the strict minimum. MEDIUM/LOW: the primary §3.3.2 mitigation is that the untrusted build-test job is credential-less (no secrets, no OIDC, contents:read), so a successful exfil yields nothing of value. The workflow is authored-but-disabled scaffolding; the file already carries a DEPLOY marker requiring the allowlist be trimmed to exactly the target repo's registries before the workflow is enabled. Over-trimming now risks breaking checkout/artifact flows in a workflow that does not yet run. TRACKED as a deploy-time hardening item." + }, + { + "id": "LOGIC-06-answered_at-stamp-before-lock", + "justification": "answer_question computes answered_at before BEGIN IMMEDIATE acquires the write lock, so under contention the persisted timestamp can invert commit order, which recovery uses to order cross-question replay. LOW: cross-question replay ordering does not affect P1 correctness (each thread_id is independent and per-thread ordering is preserved by the turn sequence). PRE-DEPLOYMENT. TRACKED for P1: stamp inside the transaction (SQLite strftime in the UPDATE) or order recovery by a monotonic rowid instead of answered_at." + }, + { + "id": "XREVIEW-7-cas-db-identity-toctou", + "justification": "_compare_and_set derives the backing DB file and opens a private connection per call; if the DB file were swapped/moved between connect() and the CAS, a stale file could be resolved. LOW/edge: requires an attacker with local filesystem write to swap the durable store mid-operation, at which point they already control the ledger directly (the deployment-model finding XREVIEW-9 / file-permissions, which is an OS-level access-control concern documented for the runbook). PRE-DEPLOYMENT; the box is read-only with mode-600 state per the design. TRACKED for the operational runbook (filesystem permissions + integrity) rather than a code change." + } + ] +} diff --git a/agent-team/README.md b/agent-team/README.md new file mode 100644 index 0000000..12bb29d --- /dev/null +++ b/agent-team/README.md @@ -0,0 +1,65 @@ +# agent-team — R720 Plane-2 FOUNDATION + +Pre-deployment scaffolding for the R720 agent-team SDLC pipeline (design: +`../docs/r720-agent-team-design.md`). This commit ships the **Plane-2 +FOUNDATION** layer only — the durable, transport-agnostic **contracts** the leaf +builders import verbatim. Nothing here is provisioned, scheduled, or wired to +live infrastructure. + +> Status: FOUNDATION modules only. No coordinator, no transports' concrete +> adapters, no CI workflow, no provisioning. Those are later phases (§7). + +## Layout + +``` +agent-team/ + agent_team/ # importable package (snake_case) + state_store.py # §6.7 atomic write + integrity-checked read + billing.py # §3.1 claude_invoke billing-mode seam + task_model.py # §3.3 TaskRecord / Phase / PipelineState + db/ + schema.py # §3.3.1/§6.7 SQLite DDL + connect/init/migrate + schema.sql # raw DDL, mirrors schema.py verbatim + transport/ + base.py # §3.3.1 Transport ABC + QuestionSet/NormalizedAnswer + tests/ # pytest unit tests, one module per source module +``` + +The top directory is kebab-case (`agent-team/`); the importable package is +snake_case (`agent_team/`), per the engineering handbook. + +## Modules (contracts) + +| Module | Design ref | What it provides | +|---|---|---| +| `state_store` | §6.7 | `atomic_write(path, data)` (write-temp → fsync → rename), `read_checked(path, *, schema_version)` (schema-version + content-hash integrity check, raises `IntegrityError`), `compute_content_hash(data)`. Pure stdlib; no other `agent_team` deps. | +| `db.schema` | §3.3.1, §6.7 | `SCHEMA_VERSION`, `PENDING_QUESTIONS_DDL`, `BUDGET_LEDGER_DDL`, `connect()` (WAL + foreign_keys + busy_timeout), `init_db()`, `migrate()`, and the `BEGIN IMMEDIATE` compare-and-set helpers (`answer_question`/`expire_question`/`supersede_question`). SQL DDL lives **only** here. | +| `billing` | §3.1 | `BillingMode{SUBSCRIPTION,API,BEDROCK}`, `claude_invoke(prompt, *, mode=None, **kw) -> ClaudeResult`, `resolve_mode(config)`. Single seam; subscription mode pops any stray `ANTHROPIC_API_KEY` so OAuth can't be overridden. | +| `transport.base` | §3.3.1 | `Transport` ABC (`post_question` → `channel_ref`; `parse_answer` → `(question_id, answer, via)`), `QuestionSet`, `NormalizedAnswer`. Transport-independent; Slack/GitHub/Claude-Code adapters subclass in the leaves. | +| `task_model` | §3.3 | `TaskRecord`, `TaskStatus`, `Phase{INTAKE…DONE}`, `new_thread_id()`, `PipelineState` TypedDict (LangGraph state schema), JSON serialization helpers. Pure model, no I/O. | + +### Durable human-in-the-loop (§3.3.1) + +The `pending_questions` ledger is the single durable source of truth for the +question lifecycle. Every race (duplicate answers, transport redelivery, +answer-vs-timeout) resolves via one atomic compare-and-set against the `status` +column, run inside a `BEGIN IMMEDIATE` transaction so concurrent responders are +serialized — first-answer-wins (`rowcount == 1`), late/duplicate ignored +(`rowcount == 0`). The LangGraph `SqliteSaver` checkpointer creates its own +tables against the **same** DB file. + +## Running the tests + +``` +cd agent-team +python3 -m pytest tests/ -q +``` + +`tests/conftest.py` puts the package on `sys.path`, so no install is required. + +## Not in this commit (later phases) + +Coordinator/brain, concrete Slack/GitHub/Claude-Code transport adapters, the +CI apply/verify workflow (§3.3.2), step-ca / Roles Anywhere, scheduling, and +the operator CLI (`run-team.py`). See `../docs/r720-agent-team-design.md` §7 +for the phased rollout. Secrets are never committed. diff --git a/agent-team/agent_team/__init__.py b/agent-team/agent_team/__init__.py new file mode 100644 index 0000000..b89a20d --- /dev/null +++ b/agent-team/agent_team/__init__.py @@ -0,0 +1,11 @@ +"""R720 agent-team — Plane-2 FOUNDATION modules (design v2). + +This package holds the durable, transport-agnostic contracts the leaf builders +import verbatim: the atomic state-store (:mod:`agent_team.state_store`), the +SQLite schema (:mod:`agent_team.db`), the Claude billing seam +(:mod:`agent_team.billing`), the transport interface +(:mod:`agent_team.transport`), and the task/thread model +(:mod:`agent_team.task_model`). +""" + +__all__: list[str] = [] diff --git a/agent-team/agent_team/billing.py b/agent-team/agent_team/billing.py new file mode 100644 index 0000000..ca84b2e --- /dev/null +++ b/agent-team/agent_team/billing.py @@ -0,0 +1,148 @@ +"""Billing-mode abstraction — the single ``claude_invoke`` seam (design §3.1). + +Every Claude-calling node imports :func:`claude_invoke` from here. The seam +selects the Claude auth/billing path from config: + +* ``SUBSCRIPTION`` — OAuth token (the R720 default; headless Agent SDK), +* ``API`` — metered ``ANTHROPIC_API_KEY``, +* ``BEDROCK`` — cross-account Bedrock (the rare cross-family tiebreak). + +Switching modes is a config flip, not a code change. In ``SUBSCRIPTION`` mode +the seam pops/unsets any stray ``ANTHROPIC_API_KEY`` from the environment +before invoking, so an inherited key cannot silently override OAuth (§3.1). + +This module is the contract leaf builders import verbatim; the actual SDK call +is delegated to an injectable ``_invoker`` so the seam stays testable and the +transport/SDK wiring lives in the leaves. +""" + +from __future__ import annotations + +import os +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Callable, Mapping + +__all__ = [ + "BillingMode", + "ClaudeResult", + "claude_invoke", + "resolve_mode", + "set_invoker", +] + +# Environment variable that carries the metered API key. Popped in +# subscription mode so OAuth cannot be silently overridden. +_API_KEY_ENV = "ANTHROPIC_API_KEY" + +# Config key (env or mapping) naming the desired billing mode. +_MODE_ENV = "AGENT_TEAM_BILLING_MODE" + + +class BillingMode(Enum): + """Claude auth/billing path selector (§3.1).""" + + SUBSCRIPTION = "subscription" + API = "api" + BEDROCK = "bedrock" + + +@dataclass +class ClaudeResult: + """Result of a :func:`claude_invoke` call. + + ``text`` is the model's response text. ``mode`` records which billing path + served the call. ``usage`` carries token/cost accounting for the budget + ledger (§6.6); ``raw`` is the untouched provider response for callers that + need more. + """ + + text: str + mode: BillingMode + usage: dict[str, Any] = field(default_factory=dict) + raw: Any = None + + +# Pluggable invoker: signature (prompt, mode, **kw) -> ClaudeResult. The +# default raises so an un-wired environment fails loudly rather than silently +# returning nothing; leaves call set_invoker() to bind the real SDK path. +Invoker = Callable[..., ClaudeResult] + + +def _unconfigured_invoker(prompt: str, *, mode: BillingMode, **kw: Any) -> ClaudeResult: + raise RuntimeError( + "claude_invoke has no invoker bound; call billing.set_invoker(fn) to " + "wire the Claude SDK path (subscription OAuth / API / Bedrock)." + ) + + +_invoker: Invoker = _unconfigured_invoker + + +def set_invoker(invoker: Invoker) -> None: + """Bind the function that performs the actual Claude SDK call. + + Leaves call this once at startup with an implementation that honours the + resolved :class:`BillingMode`. Keeping the SDK call injectable keeps this + seam dependency-free and unit-testable. + """ + global _invoker + _invoker = invoker + + +def resolve_mode(config: Mapping[str, Any] | None) -> BillingMode: + """Resolve the billing mode from ``config`` (falling back to env). + + Precedence: an explicit ``billing_mode`` in ``config`` (a + :class:`BillingMode` or its string value), then the + ``AGENT_TEAM_BILLING_MODE`` env var, then the ``SUBSCRIPTION`` default. + """ + raw: Any = None + if config is not None: + raw = config.get("billing_mode") + if raw is None: + raw = os.environ.get(_MODE_ENV) + if raw is None: + return BillingMode.SUBSCRIPTION + if isinstance(raw, BillingMode): + return raw + try: + return BillingMode(str(raw).strip().lower()) + except ValueError as exc: + valid = ", ".join(m.value for m in BillingMode) + raise ValueError( + f"unknown billing mode {raw!r}; expected one of: {valid}" + ) from exc + + +def claude_invoke( + prompt: str, + *, + mode: BillingMode | None = None, + config: Mapping[str, Any] | None = None, + **kw: Any, +) -> ClaudeResult: + """Invoke Claude through the configured billing path (§3.1). + + ``mode`` overrides config when given; otherwise it is resolved via + :func:`resolve_mode`. In ``SUBSCRIPTION`` mode any stray + ``ANTHROPIC_API_KEY`` is popped from ``os.environ`` for the duration of the + call so OAuth cannot be silently overridden, then restored afterward. + + The actual SDK call is delegated to the bound invoker (see + :func:`set_invoker`); this function owns only mode selection and the + subscription-mode env hygiene that the design mandates. + """ + effective = mode if mode is not None else resolve_mode(config) + + if effective is BillingMode.SUBSCRIPTION: + # Pop the stray key for the duration of the call; restore on exit so we + # don't mutate the caller's environment permanently. + stashed = os.environ.pop(_API_KEY_ENV, None) + try: + return _invoker(prompt, mode=effective, **kw) + finally: + if stashed is not None: + os.environ[_API_KEY_ENV] = stashed + + return _invoker(prompt, mode=effective, **kw) diff --git a/agent-team/agent_team/ci_gate.py b/agent-team/agent_team/ci_gate.py new file mode 100644 index 0000000..0a0db1b --- /dev/null +++ b/agent-team/agent_team/ci_gate.py @@ -0,0 +1,494 @@ +"""Pure-code CI pass/fail gate — the deterministic block decision (design §3.3.2). + +This module is the §3.3.2 boundary #4: *"Pass/fail is a pure-code gate over +authenticated CI results, not the LLM verifier."* Mirroring secrev's "one +pure-code script owns the block decision", the gate reads the CI run +**conclusion** (already fetched, authenticated as the box read-only PAT via the +GitHub Checks/Actions API) keyed to a specific ``run_id`` + ``diff_hash`` and +makes a deterministic ``pass | fail | block`` decision. It consumes **only** +that authenticated, patch-independent conclusion; it never trusts a +success/failure file or artifact the patch could have written. + +The gate also enforces the two box-side/CI defences that must hold before a +diff is even allowed to build (defence in depth — the CI side enforces the same +checks as a hard fail): + +* **boundary #2** — the **trust-control-surface denylist**: a candidate diff + may not touch ``.github/workflows/**``, IAM/policy IaC, branch-protection / + ``CODEOWNERS`` / Dependabot config, or paths outside the task's declared + scope. A match is canonicalized (symlink/``..`` resolution) and rejects + renames into denied paths, so a path match cannot be bypassed by indirection. +* **boundary #3** — **diff integrity**: the gate recomputes the candidate diff + hash and compares it to the ledger-recorded hash and to the hash CI verified, + so a tampered or substituted diff fails closed. + +This module is pure stdlib and imports the committed foundation contracts +verbatim (it redefines none of them). It performs **no** network I/O: the +authenticated CI conclusion is passed in as data (the caller fetches it via the +read-only PAT). Keeping the gate I/O-free is what makes the block decision +deterministic and unit-testable. + +Decision semantics (a diff "ships" only on an unambiguous authenticated pass): + +* :data:`GateDecision.PASS` — the run concluded ``success`` for the exact + ``run_id``/``diff_hash`` and every guard held; the verifier may open a draft + PR. +* :data:`GateDecision.FAIL` — the run concluded a recognised failure + (``failure``/``timed_out``/``cancelled``/...); the verifier loops back to the + builders. +* :data:`GateDecision.BLOCK` — a trust violation (denylist hit, hash mismatch, + run-id mismatch, missing/ambiguous authenticated conclusion). This is an + ALARM-worthy refuse-to-proceed, never a silent pass. +""" + +from __future__ import annotations + +import re +from dataclasses import dataclass, field +from enum import Enum +from pathlib import PurePosixPath +from typing import Any, Mapping, Sequence + +from agent_team.state_store import compute_content_hash + +__all__ = [ + "DENYLIST_GLOBS", + "CiGateError", + "GateDecision", + "GateResult", + "denylist_violations", + "diff_touched_paths", + "evaluate_ci_gate", + "verify_diff_hash", +] + + +class CiGateError(Exception): + """Raised when the gate is called with structurally invalid inputs. + + Distinct from a :data:`GateDecision.BLOCK`: a ``BLOCK`` is a *valid* gate + run that found a trust violation, whereas this exception means the caller + handed the gate malformed data (e.g. a non-string diff). Fails loud rather + than guessing. + """ + + +class GateDecision(Enum): + """The deterministic gate outcome (§3.3.2 boundary #4).""" + + PASS = "pass" + FAIL = "fail" + BLOCK = "block" + + +# Trust-control-surface denylist (§3.3.2 boundary #2). A candidate diff that +# touches any of these is escalated to mandatory human + GPT cross-review, never +# auto-built — they are the mandatory-cross-review surface regardless. Globs are +# matched against POSIX-canonicalized repo-relative paths. +DENYLIST_GLOBS: tuple[str, ...] = ( + # CI workflow definitions — the "pwn request" surface. + ".github/workflows/**", + ".github/actions/**", + # Branch protection / ownership / dependency automation config. + ".github/CODEOWNERS", + "CODEOWNERS", + ".github/dependabot.yml", + ".github/dependabot.yaml", + ".github/settings.yml", + # IAM / policy IaC (CDK / SAM / Terraform / CloudFormation). + "**/cdk.json", + "**/template.yml", + "**/template.yaml", + "**/samconfig.toml", + "**/*.tf", + "**/iam/**", + "**/policies/**", + "**/*policy*.json", +) + +# Authenticated GitHub run conclusions that count as a recognised failure (the +# verifier loops back). Anything not in PASS/this set is ambiguous -> BLOCK. +_FAILURE_CONCLUSIONS: frozenset[str] = frozenset( + {"failure", "timed_out", "cancelled", "action_required", "stale", "startup_failure"} +) + +# The single authenticated conclusion that means "ship-able". +_SUCCESS_CONCLUSION = "success" + +# ``git diff`` file-header line, e.g. ``diff --git a/foo.py b/foo.py``. We read +# the post-image (``b/``) path as the touched path and also surface the pre-image +# (``a/``) so a *rename into* a denied path is caught (boundary #2). +_DIFF_GIT_RE = re.compile(r"^diff --git a/(?P.+?) b/(?P.+?)\s*$") +# ``rename from``/``rename to`` lines carry the rename source/target explicitly. +_RENAME_FROM_RE = re.compile(r"^rename from (?P.+?)\s*$") +_RENAME_TO_RE = re.compile(r"^rename to (?P.+?)\s*$") + + +@dataclass +class GateResult: + """The gate's deterministic verdict plus the evidence behind it. + + ``decision`` is the block decision the verifier node acts on. ``reasons`` + enumerates every concrete trigger (denylist hits, hash mismatch, the CI + conclusion consumed) so the decision is auditable and an ALARM can quote it. + ``run_id``/``diff_hash`` echo the keys the gate was bound to (provenance, + §3.3.2). ``ci_conclusion`` is the authenticated conclusion actually + consumed. + """ + + decision: GateDecision + reasons: list[str] = field(default_factory=list) + run_id: str | None = None + diff_hash: str | None = None + ci_conclusion: str | None = None + + @property + def passed(self) -> bool: + """``True`` only on an unambiguous authenticated pass.""" + return self.decision is GateDecision.PASS + + @property + def blocked(self) -> bool: + """``True`` on a trust violation (ALARM-worthy refuse-to-proceed).""" + return self.decision is GateDecision.BLOCK + + +def _canonical_repo_path(raw: str) -> str: + """Canonicalize a repo-relative path for denylist matching (§3.3.2). + + Strips a leading ``a/``/``b/`` git prefix, normalizes separators, resolves + ``.``/``..`` segments without touching the filesystem (the diff describes + paths that may not exist locally), and drops a leading ``/`` so the result + is always repo-relative. A path that escapes the repo root via ``..`` is + returned with a sentinel ``..`` prefix preserved so it cannot silently match + *nothing* — the caller treats an escaping path as a denylist hit. + """ + text = raw.strip().strip('"') + # Drop a single git a//b/ prefix if present. + if text.startswith(("a/", "b/")): + text = text[2:] + # PurePosixPath normalizes separators; resolve . and .. logically. + parts: list[str] = [] + for segment in PurePosixPath(text).parts: + if segment in ("", "."): + continue + if segment == "..": + # Escaping the repo root — keep the marker so it never matches a + # benign glob and is treated as suspicious by the caller. + parts.append("..") + continue + parts.append(segment) + return "/".join(parts) + + +def diff_touched_paths(unified_diff: str) -> list[str]: + """Extract the canonicalized repo-relative paths a unified diff touches. + + Reads ``diff --git`` headers (both the ``a/`` pre-image and ``b/`` + post-image) plus explicit ``rename from``/``rename to`` lines, so a rename + *into* a denied path is surfaced (§3.3.2 boundary #2 — "rejects renames into + denied paths"). Returns a de-duplicated, sorted list of canonical paths. + + Raises :class:`CiGateError` if ``unified_diff`` is not a string. + """ + if not isinstance(unified_diff, str): + raise CiGateError(f"diff must be str, got {type(unified_diff).__name__}") + + touched: set[str] = set() + for line in unified_diff.splitlines(): + m = _DIFF_GIT_RE.match(line) + if m is not None: + touched.add(_canonical_repo_path(m.group("a"))) + touched.add(_canonical_repo_path(m.group("b"))) + continue + m = _RENAME_FROM_RE.match(line) + if m is not None: + touched.add(_canonical_repo_path(m.group("path"))) + continue + m = _RENAME_TO_RE.match(line) + if m is not None: + touched.add(_canonical_repo_path(m.group("path"))) + touched.discard("") + return sorted(touched) + + +def _glob_to_regex(glob: str) -> re.Pattern[str]: + """Compile a denylist glob to an anchored regex. + + Supports ``**`` (any number of path segments, including zero), ``*`` (within + a single segment), and ``?``. Everything else is matched literally. Matching + is done on canonical POSIX repo-relative paths. + """ + out: list[str] = ["^"] + i = 0 + n = len(glob) + while i < n: + ch = glob[i] + if ch == "*": + if i + 1 < n and glob[i + 1] == "*": + # ``**`` — any chars incl. ``/``. Swallow an optional trailing + # ``/`` so ``dir/**`` also matches ``dir`` itself's children + # without requiring a separator artifact. + out.append(".*") + i += 2 + if i < n and glob[i] == "/": + i += 1 + continue + # single ``*`` — anything but a path separator. + out.append("[^/]*") + i += 1 + continue + if ch == "?": + out.append("[^/]") + i += 1 + continue + out.append(re.escape(ch)) + i += 1 + out.append("$") + return re.compile("".join(out)) + + +_DENYLIST_RES: tuple[tuple[str, re.Pattern[str]], ...] = tuple( + (g, _glob_to_regex(g)) for g in DENYLIST_GLOBS +) + + +def denylist_violations( + unified_diff: str, + *, + allowed_scope: Sequence[str] | None = None, +) -> list[str]: + """Return the trust-control-surface violations in ``unified_diff`` (§3.3.2). + + A violation is any touched path that (a) matches a :data:`DENYLIST_GLOBS` + entry, (b) escapes the repo root via ``..`` (path-indirection attempt), or + (c) — when ``allowed_scope`` is given — falls outside the task's declared + scope. ``allowed_scope`` is a sequence of canonical path prefixes + (directories or exact files) the task is allowed to modify; a touched path + outside every prefix is a violation ("files outside the task's declared + scope"). + + Returns a sorted list of human-readable reason strings; an empty list means + the diff is clean for the denylist boundary. + """ + violations: list[str] = [] + normalized_scope = ( + [_canonical_repo_path(p) for p in allowed_scope] + if allowed_scope is not None + else None + ) + + for path in diff_touched_paths(unified_diff): + if ".." in PurePosixPath(path).parts: + violations.append(f"path escapes repo root via '..': {path!r}") + continue + for glob, pattern in _DENYLIST_RES: + if pattern.match(path): + violations.append(f"denylisted path {path!r} matches glob {glob!r}") + break + else: + if normalized_scope is not None and not _within_scope( + path, normalized_scope + ): + violations.append(f"path {path!r} is outside the task's declared scope") + return sorted(violations) + + +def _within_scope(path: str, scope: Sequence[str]) -> bool: + """Return ``True`` if ``path`` is within one declared-scope prefix.""" + candidate = PurePosixPath(path) + for prefix in scope: + if not prefix: + continue + if path == prefix: + return True + prefix_path = PurePosixPath(prefix) + try: + candidate.relative_to(prefix_path) + return True + except ValueError: + continue + return False + + +def verify_diff_hash( + candidate_diff: str, + *, + ledger_hash: str | None, + ci_verified_hash: str | None = None, +) -> bool: + """Verify the candidate diff hash matches the ledger (and CI) (§3.3.2 #3). + + Recomputes the content hash of ``candidate_diff`` (sha256, via the committed + :func:`agent_team.state_store.compute_content_hash`) and compares it to the + ``ledger_hash`` recorded by the builder and, when supplied, to the + ``ci_verified_hash`` CI checked before applying. Comparison is + constant-time. Returns ``True`` only when all supplied hashes agree; a + ``None`` ledger hash is treated as a failure (there is nothing to bind to). + + Raises :class:`CiGateError` if ``candidate_diff`` is not a string. + """ + if not isinstance(candidate_diff, str): + raise CiGateError( + f"candidate_diff must be str, got {type(candidate_diff).__name__}" + ) + if not ledger_hash: + return False + + actual = compute_content_hash(candidate_diff.encode("utf-8")) + if not _consteq(actual, ledger_hash): + return False + if ci_verified_hash is not None and not _consteq(actual, ci_verified_hash): + return False + return True + + +def _consteq(a: str, b: str) -> bool: + """Constant-time string compare (hashes are not secret, but be tidy).""" + if len(a) != len(b): + return False + result = 0 + for x, y in zip(a, b): + result |= ord(x) ^ ord(y) + return result == 0 + + +def evaluate_ci_gate( + *, + candidate_diff: str, + ledger_hash: str | None, + ci_result: Mapping[str, Any] | None, + expected_run_id: str, + allowed_scope: Sequence[str] | None = None, +) -> GateResult: + """Make the deterministic pass/fail/block decision (§3.3.2 boundary #4). + + Inputs (all authenticated/patch-independent — the gate does no I/O): + + * ``candidate_diff`` — the builder's diff text (re-hashed here). + * ``ledger_hash`` — the hash the builder recorded in the task ledger. + * ``ci_result`` — the authenticated CI conclusion the caller fetched via the + read-only PAT (GitHub Checks/Actions API). The gate reads only + ``run_id``, ``conclusion``, and (optionally) ``diff_hash`` from it; it + **never** reads a patch-written success file/artifact. + * ``expected_run_id`` — the run id the verifier dispatched for this exact + diff; the conclusion must be keyed to it (a stale/substituted run id is a + BLOCK). + * ``allowed_scope`` — optional declared-scope prefixes for the task. + + Decision order (a trust violation always wins over a CI verdict): + + 1. **Denylist / scope** — any violation -> :data:`GateDecision.BLOCK`. + 2. **Diff-hash integrity** — hash mismatch (ledger or CI-verified) -> + ``BLOCK``. + 3. **Authenticated conclusion** — missing result, a ``run_id`` that does not + match ``expected_run_id``, or an unrecognised/ambiguous conclusion -> + ``BLOCK``; a recognised failure -> :data:`GateDecision.FAIL`; ``success`` + -> :data:`GateDecision.PASS`. + + Returns a :class:`GateResult` with the decision and the reasons behind it. + Raises :class:`CiGateError` on structurally invalid inputs. + """ + if not isinstance(expected_run_id, str) or not expected_run_id: + raise CiGateError("expected_run_id must be a non-empty string") + + reasons: list[str] = [] + + # (1) Trust-control-surface denylist + declared scope. Highest priority: + # these files are the mandatory-cross-review surface, never auto-built. + violations = denylist_violations(candidate_diff, allowed_scope=allowed_scope) + if violations: + reasons.extend(violations) + return GateResult( + decision=GateDecision.BLOCK, + reasons=reasons, + run_id=expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=None, + ) + + # (2) Diff-hash integrity: bind the decision to the exact bytes the ledger + # recorded (and, if present, what CI verified before applying). + ci_verified_hash = ( + str(ci_result.get("diff_hash")) + if ci_result is not None and ci_result.get("diff_hash") is not None + else None + ) + if not verify_diff_hash( + candidate_diff, + ledger_hash=ledger_hash, + ci_verified_hash=ci_verified_hash, + ): + reasons.append( + "diff-hash mismatch: candidate diff does not match the ledger" + + (" / CI-verified" if ci_verified_hash is not None else "") + + " hash" + ) + return GateResult( + decision=GateDecision.BLOCK, + reasons=reasons, + run_id=expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=None, + ) + + # (3) Authenticated, patch-independent CI conclusion. + if ci_result is None: + reasons.append("no authenticated CI result supplied") + return GateResult( + decision=GateDecision.BLOCK, + reasons=reasons, + run_id=expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=None, + ) + + actual_run_id = ci_result.get("run_id") + if str(actual_run_id) != expected_run_id: + reasons.append( + f"CI run-id mismatch: expected {expected_run_id!r}, " + f"conclusion is keyed to {actual_run_id!r}" + ) + return GateResult( + decision=GateDecision.BLOCK, + reasons=reasons, + run_id=expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=None, + ) + + conclusion = ci_result.get("conclusion") + normalized = str(conclusion).strip().lower() if conclusion is not None else None + + if normalized == _SUCCESS_CONCLUSION: + reasons.append("authenticated CI conclusion: success") + return GateResult( + decision=GateDecision.PASS, + reasons=reasons, + run_id=expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=normalized, + ) + + if normalized in _FAILURE_CONCLUSIONS: + reasons.append(f"authenticated CI conclusion: {normalized}") + return GateResult( + decision=GateDecision.FAIL, + reasons=reasons, + run_id=expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=normalized, + ) + + # Unknown / null / still-running conclusion: refuse to proceed (never a + # silent pass). e.g. ``None`` (in-progress), ``"neutral"``, ``"skipped"``. + reasons.append( + f"unrecognised/ambiguous CI conclusion {conclusion!r}; refusing to proceed" + ) + return GateResult( + decision=GateDecision.BLOCK, + reasons=reasons, + run_id=expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=normalized, + ) diff --git a/agent-team/agent_team/db/__init__.py b/agent-team/agent_team/db/__init__.py new file mode 100644 index 0000000..c62f898 --- /dev/null +++ b/agent-team/agent_team/db/__init__.py @@ -0,0 +1,24 @@ +"""SQLite schema and connection helpers for the R720 agent-team pipeline. + +SQL DDL lives ONLY in this subpackage (``schema.py`` constants mirrored by the +companion ``schema.sql``). The LangGraph ``SqliteSaver`` checkpointer creates +its own tables against the same database file. +""" + +from agent_team.db.schema import ( + BUDGET_LEDGER_DDL, + PENDING_QUESTIONS_DDL, + SCHEMA_VERSION, + connect, + init_db, + migrate, +) + +__all__ = [ + "BUDGET_LEDGER_DDL", + "PENDING_QUESTIONS_DDL", + "SCHEMA_VERSION", + "connect", + "init_db", + "migrate", +] diff --git a/agent-team/agent_team/db/schema.py b/agent-team/agent_team/db/schema.py new file mode 100644 index 0000000..a6af850 --- /dev/null +++ b/agent-team/agent_team/db/schema.py @@ -0,0 +1,431 @@ +"""SQLite schema, DDL constants, and connection helpers (design §3.3.1, §6.7). + +This module is the single source of truth for the R720 agent-team durable +SQL. It declares: + +* the ``pending_questions`` human-interaction lifecycle ledger (§3.3.1), +* the ``budget_ledger`` shared Claude budget ledger (§6.1, §6.6), +* a ``schema_meta`` version row driving :func:`migrate`. + +The companion ``schema.sql`` mirrors these statements verbatim for tooling. +SQL DDL lives ONLY here. The LangGraph ``SqliteSaver`` checkpointer creates +its OWN tables against this same connection / database file; the design +reserves this DB for it but does not declare its tables. + +The atomic compare-and-set helpers used by the responder and the deadline +timer (§3.3.1) take the write lock up front via ``BEGIN IMMEDIATE`` so +concurrent responders are serialized — SQLite's default deferred isolation +does not serialize a check-and-set. +""" + +from __future__ import annotations + +import sqlite3 +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +__all__ = [ + "BUDGET_LEDGER_DDL", + "PENDING_QUESTIONS_DDL", + "PENDING_QUESTIONS_INDEXES_DDL", + "BUDGET_LEDGER_INDEXES_DDL", + "SCHEMA_META_DDL", + "SCHEMA_VERSION", + "QUESTION_STATES", + "answer_question", + "connect", + "expire_question", + "init_db", + "migrate", + "reopen_question", + "supersede_question", +] + +# Bump when the DDL below changes; migrate() steps a connection forward. +SCHEMA_VERSION: int = 1 + +# Default SQLite busy timeout (ms) so concurrent writers wait for the write +# lock rather than failing immediately. +_BUSY_TIMEOUT_MS: int = 5000 + +# Allowed lifecycle states for a pending question (§3.3.1). Mirrors the DDL +# CHECK constraint; exported so leaves can validate without re-listing them. +QUESTION_STATES: tuple[str, ...] = ("open", "answered", "expired", "superseded") + + +PENDING_QUESTIONS_DDL: str = """ +CREATE TABLE IF NOT EXISTS pending_questions ( + question_id TEXT PRIMARY KEY, + thread_id TEXT NOT NULL, + turn INTEGER NOT NULL, + status TEXT NOT NULL + CHECK (status IN ('open', 'answered', 'expired', 'superseded')), + transport TEXT NOT NULL, + channel_ref TEXT, + posted_at TEXT, + deadline_at TEXT, + answer_json TEXT, + answered_at TEXT, + answered_via TEXT +) +""".strip() + +PENDING_QUESTIONS_INDEXES_DDL: str = """ +CREATE INDEX IF NOT EXISTS idx_pending_questions_thread + ON pending_questions (thread_id, turn); +CREATE INDEX IF NOT EXISTS idx_pending_questions_status + ON pending_questions (status); +""".strip() + +BUDGET_LEDGER_DDL: str = """ +CREATE TABLE IF NOT EXISTS budget_ledger ( + entry_id INTEGER PRIMARY KEY AUTOINCREMENT, + thread_id TEXT, + stage TEXT, + model TEXT NOT NULL, + billing_mode TEXT NOT NULL, + input_tokens INTEGER NOT NULL DEFAULT 0, + output_tokens INTEGER NOT NULL DEFAULT 0, + usd_cost REAL NOT NULL DEFAULT 0.0, + recorded_at TEXT NOT NULL, + day_bucket TEXT NOT NULL +) +""".strip() + +BUDGET_LEDGER_INDEXES_DDL: str = """ +CREATE INDEX IF NOT EXISTS idx_budget_ledger_day + ON budget_ledger (day_bucket); +CREATE INDEX IF NOT EXISTS idx_budget_ledger_thread + ON budget_ledger (thread_id); +""".strip() + +SCHEMA_META_DDL: str = """ +CREATE TABLE IF NOT EXISTS schema_meta ( + id INTEGER PRIMARY KEY CHECK (id = 1), + schema_version INTEGER NOT NULL +) +""".strip() + + +class _Connection(sqlite3.Connection): + """``sqlite3.Connection`` subclass that can carry its backing file path. + + The base ``Connection`` has no ``__dict__``, so a path cannot be stashed on + it. This thin subclass (passed as ``factory=`` to :func:`sqlite3.connect`) + lets :func:`connect` record the db file for a thread-safe attribute lookup by + the compare-and-set, avoiding a ``PRAGMA`` on a connection shared across + threads. + """ + + agent_team_db_path: str = "" + + +def connect(db_path: Path) -> sqlite3.Connection: + """Open ``db_path`` with WAL, foreign keys, and a busy timeout. + + WAL (``journal_mode=WAL``) lets the resume worker read while a responder + writes; ``foreign_keys=ON`` enforces referential integrity; the busy + timeout makes concurrent writers wait for the write lock instead of + failing. ``isolation_level=None`` puts the connection in autocommit mode so + the compare-and-set helpers can drive transactions explicitly with + ``BEGIN IMMEDIATE`` (§3.3.1). + """ + db_path = Path(db_path) + db_path.parent.mkdir(parents=True, exist_ok=True) + conn = sqlite3.connect( + str(db_path), + isolation_level=None, + check_same_thread=False, + factory=_Connection, + ) + conn.row_factory = sqlite3.Row + # Set the busy timeout FIRST so every subsequent statement — including the + # journal-mode pragma below, which briefly needs the write lock — waits for + # the lock instead of failing immediately when another connection is mid + # -write. (Without this, opening a connection under concurrent writers could + # raise "database is locked" before the timeout was ever applied.) + conn.execute(f"PRAGMA busy_timeout={_BUSY_TIMEOUT_MS}") + conn.execute("PRAGMA journal_mode=WAL") + conn.execute("PRAGMA foreign_keys=ON") + # Record the backing file path so the compare-and-set can derive it via a + # thread-safe attribute read instead of running a PRAGMA on a connection that + # callers share across threads (a sqlite3.Connection is not safe for + # concurrent use — even a read would corrupt its transaction state). Empty + # for an in-memory DB (no file to reopen on a second connection). + conn.agent_team_db_path = "" if str(db_path) == ":memory:" else str(db_path) + return conn + + +def init_db(db_path: Path) -> None: + """Create the agent-team tables in ``db_path`` if absent. + + Creates ``pending_questions`` (+ indexes), the budget ledger (+ indexes), + and the ``schema_meta`` version row, and reserves the same DB file for the + LangGraph ``SqliteSaver`` checkpointer (which creates its own tables on + first use against this connection). Idempotent: safe to call on every + startup. + """ + conn = connect(db_path) + try: + conn.execute(SCHEMA_META_DDL) + conn.execute(PENDING_QUESTIONS_DDL) + for stmt in _split_statements(PENDING_QUESTIONS_INDEXES_DDL): + conn.execute(stmt) + conn.execute(BUDGET_LEDGER_DDL) + for stmt in _split_statements(BUDGET_LEDGER_INDEXES_DDL): + conn.execute(stmt) + # Record the schema version (single-row table). + conn.execute( + "INSERT INTO schema_meta (id, schema_version) VALUES (1, ?) " + "ON CONFLICT(id) DO NOTHING", + (SCHEMA_VERSION,), + ) + finally: + conn.close() + + +def migrate(conn: sqlite3.Connection) -> None: + """Step ``conn``'s schema forward to :data:`SCHEMA_VERSION`. + + Reads the recorded version from ``schema_meta`` (treating an empty/absent + row as version 0), applies any forward steps, and records the new version. + At ``SCHEMA_VERSION == 1`` there are no prior versions to migrate from, so + this ensures the base tables exist and stamps the version. Future versions + add ordered ``if current < N`` blocks here. + """ + conn.execute(SCHEMA_META_DDL) + row = conn.execute("SELECT schema_version FROM schema_meta WHERE id = 1").fetchone() + current = int(row["schema_version"]) if row is not None else 0 + + if current < 1: + # Base schema (v1): ensure all tables/indexes exist. + conn.execute(PENDING_QUESTIONS_DDL) + for stmt in _split_statements(PENDING_QUESTIONS_INDEXES_DDL): + conn.execute(stmt) + conn.execute(BUDGET_LEDGER_DDL) + for stmt in _split_statements(BUDGET_LEDGER_INDEXES_DDL): + conn.execute(stmt) + current = 1 + + # Future steps go here: `if current < 2: ...; current = 2`. + + conn.execute( + "INSERT INTO schema_meta (id, schema_version) VALUES (1, ?) " + "ON CONFLICT(id) DO UPDATE SET schema_version = excluded.schema_version", + (current,), + ) + + +def answer_question( + conn: sqlite3.Connection, + *, + question_id: str, + answer_json: str, + answered_via: str, + answered_at: str | None = None, +) -> bool: + """First-answer-wins compare-and-set: flip an ``open`` question to answered. + + Runs the §3.3.1 atomic statement inside a ``BEGIN IMMEDIATE`` transaction + so concurrent responders are serialized (the check-and-set takes the write + lock up front). Returns ``True`` when rowcount == 1 (this caller recorded + the first valid answer; enqueue a resume job), ``False`` when rowcount == 0 + (the question was not ``open`` — already answered/expired/superseded — so + the answer is a duplicate or late and must be ignored). + """ + stamp = answered_at or _utc_now_iso() + return _compare_and_set( + conn, + sql=( + "UPDATE pending_questions " + "SET status='answered', answer_json=?, answered_via=?, answered_at=? " + "WHERE question_id=? AND status='open'" + ), + params=(answer_json, answered_via, stamp, question_id), + ) + + +def expire_question( + conn: sqlite3.Connection, + *, + question_id: str, +) -> bool: + """Deadline race: flip an overdue ``open`` question to ``expired``. + + Same compare-and-set discipline as :func:`answer_question` (§3.3.1): an + answer that arrives for an already-expired question loses the race and is + ignored. Returns ``True`` if this call expired the question. + """ + return _compare_and_set( + conn, + sql=( + "UPDATE pending_questions SET status='expired' " + "WHERE question_id=? AND status='open'" + ), + params=(question_id,), + ) + + +def reopen_question( + conn: sqlite3.Connection, + *, + question_id: str, + deadline_at: str | None = None, +) -> bool: + """Un-park: flip an ``expired`` question back to ``open`` (operator action). + + The §6.6 operator force-resume path for a parked task whose clarifier + question expired with no answer: re-open it so the normal delivery → answer → + resume flow can proceed, instead of destructively superseding it (which would + remove it from the recovery sweep's reach). Same compare-and-set discipline — + only an ``expired`` row is reopened; an already-answered/open/superseded row + loses the CAS and is untouched. ``deadline_at`` sets a fresh window (``NULL`` + means no deadline until one is set, so it will not immediately re-expire). + Returns ``True`` if this call reopened the question. + """ + return _compare_and_set( + conn, + sql=( + "UPDATE pending_questions " + "SET status='open', deadline_at=?, channel_ref=NULL, " + "answer_json=NULL, answered_via=NULL, answered_at=NULL " + "WHERE question_id=? AND status='expired'" + ), + params=(deadline_at, question_id), + ) + + +def supersede_question( + conn: sqlite3.Connection, + *, + question_id: str, +) -> bool: + """Mark a stale ``open``/``answered`` question ``superseded``. + + Used by the turn-guarded resume worker: if the graph already advanced past + this turn, the question is superseded and the resume is skipped (§3.3.1). + Returns ``True`` if this call superseded the question. + """ + return _compare_and_set( + conn, + sql=( + "UPDATE pending_questions SET status='superseded' " + "WHERE question_id=? AND status IN ('open', 'answered')" + ), + params=(question_id,), + ) + + +# Bounded retry if the write lock is still contended after ``busy_timeout`` +# elapses, so transient over-timeout contention does not surface as an error to +# the responder / deadline-timer callers. +_CAS_RETRY_ATTEMPTS: int = 3 +_CAS_RETRY_BACKOFF_S: float = 0.05 + + +def _main_db_file(conn: sqlite3.Connection) -> str | None: + """Return the file backing ``conn``'s ``main`` database, or ``None``. + + ``None`` signals an in-memory database (no file to reopen on a second + connection). Prefers the path stashed by :func:`connect` — a thread-safe + attribute read, so it is safe even when callers share ``conn`` across + threads. Falls back to ``PRAGMA database_list`` (rows of ``(seq, name, + file)``, indexed positionally to be ``row_factory``-agnostic) only for a + connection not opened via :func:`connect`; such a connection must not be + shared across threads. + """ + stashed = getattr(conn, "agent_team_db_path", None) + if stashed is not None: + return stashed or None + for row in conn.execute("PRAGMA database_list"): + if row[1] == "main": + return row[2] or None + return None + + +def _compare_and_set( + conn: sqlite3.Connection, + *, + sql: str, + params: tuple[Any, ...], +) -> bool: + """Run a single compare-and-set UPDATE under ``BEGIN IMMEDIATE`` (§3.3.1). + + Returns ``True`` iff exactly one row changed. The check-and-set takes the + write lock up front so concurrent responders cannot both observe + ``status='open'`` (SQLite's default deferred isolation would not serialize + them). + + **Concurrency safety.** The write runs on a private, short-lived connection + to the same database file — never on the passed ``conn``. A single SQLite + connection cannot hold two explicit transactions at once, so if a caller + shares one ``conn`` across threads (the responder and resume worker do, and + ``connect()`` sets ``check_same_thread=False``), two concurrent + ``BEGIN IMMEDIATE`` statements on it would raise "cannot start a transaction + within a transaction". Giving each call its own connection makes the + compare-and-set safe under that sharing; WAL serializes the writers via the + busy handler. A lock that outlasts ``busy_timeout`` is retried a bounded + number of times before propagating. ``BEGIN IMMEDIATE`` runs inside the + guarded path so its lock error is caught and retried, not raised uncaught. + + For an in-memory database (no file to reopen) the call falls back to the + passed ``conn``; in-memory DBs are single-connection and not the concurrent + production path. + """ + db_file = _main_db_file(conn) + if db_file is None: + return _cas_once(conn, sql, params) + + last_err: sqlite3.OperationalError | None = None + for attempt in range(_CAS_RETRY_ATTEMPTS): + write = connect(Path(db_file)) + try: + return _cas_once(write, sql, params) + except sqlite3.OperationalError as err: + if "locked" not in str(err).lower(): + raise + last_err = err + finally: + write.close() + time.sleep(_CAS_RETRY_BACKOFF_S * (attempt + 1)) + + assert last_err is not None # loop only exits early via return or raise + raise last_err + + +def _cas_once( + conn: sqlite3.Connection, + sql: str, + params: tuple[Any, ...], +) -> bool: + """Execute one ``BEGIN IMMEDIATE`` compare-and-set on ``conn``. + + ``BEGIN IMMEDIATE`` is issued before the try so a lock-acquisition error + propagates to the caller's retry loop with no transaction to unwind; once + the transaction is open, any failure rolls it back (best-effort) and + re-raises. + """ + conn.execute("BEGIN IMMEDIATE") + try: + cur = conn.execute(sql, params) + changed = cur.rowcount == 1 + conn.execute("COMMIT") + return changed + except BaseException: + try: + conn.execute("ROLLBACK") + except sqlite3.OperationalError: + pass + raise + + +def _utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string.""" + return datetime.now(timezone.utc).isoformat() + + +def _split_statements(ddl: str) -> list[str]: + """Split a multi-statement DDL blob into individual statements.""" + return [stmt.strip() for stmt in ddl.split(";") if stmt.strip()] diff --git a/agent-team/agent_team/db/schema.sql b/agent-team/agent_team/db/schema.sql new file mode 100644 index 0000000..5dc4479 --- /dev/null +++ b/agent-team/agent_team/db/schema.sql @@ -0,0 +1,65 @@ +-- R720 agent-team durable SQLite schema (design §3.3.1, §6.7). +-- +-- This file holds the raw DDL statements ONLY. The authoritative copies live +-- as string constants in agent_team/db/schema.py; this companion file mirrors +-- them verbatim for tooling / direct inspection. SQL DDL lives only in these +-- two places. +-- +-- The LangGraph SqliteSaver checkpointer creates its OWN tables against this +-- same database file/connection; those are intentionally NOT declared here. + +-- pending_questions: the durable human-interaction lifecycle ledger. +-- The LangGraph checkpoint holds graph state; this table holds the question +-- lifecycle (delivery, duplicate/late answers, expiry) and is what delivery, +-- the responder, and restart recovery read. Every race resolves via an atomic +-- compare-and-set against the `status` column under BEGIN IMMEDIATE. +CREATE TABLE IF NOT EXISTS pending_questions ( + question_id TEXT PRIMARY KEY, + thread_id TEXT NOT NULL, + turn INTEGER NOT NULL, + status TEXT NOT NULL + CHECK (status IN ('open', 'answered', 'expired', 'superseded')), + transport TEXT NOT NULL, + channel_ref TEXT, + posted_at TEXT, + deadline_at TEXT, + answer_json TEXT, + answered_at TEXT, + answered_via TEXT +); + +CREATE INDEX IF NOT EXISTS idx_pending_questions_thread + ON pending_questions (thread_id, turn); + +CREATE INDEX IF NOT EXISTS idx_pending_questions_status + ON pending_questions (status); + +-- budget_ledger: the persistent shared Claude budget ledger (§6.1, §6.6). +-- One row per accounted spend event; the shared daily cap and the +-- interactive-first reserve are computed by summing over a UTC day. Spend is +-- recorded across ALL R720 Claude work (pipeline + Plane-1 sweeps). +CREATE TABLE IF NOT EXISTS budget_ledger ( + entry_id INTEGER PRIMARY KEY AUTOINCREMENT, + thread_id TEXT, + stage TEXT, + model TEXT NOT NULL, + billing_mode TEXT NOT NULL, + input_tokens INTEGER NOT NULL DEFAULT 0, + output_tokens INTEGER NOT NULL DEFAULT 0, + usd_cost REAL NOT NULL DEFAULT 0.0, + recorded_at TEXT NOT NULL, + day_bucket TEXT NOT NULL +); + +CREATE INDEX IF NOT EXISTS idx_budget_ledger_day + ON budget_ledger (day_bucket); + +CREATE INDEX IF NOT EXISTS idx_budget_ledger_thread + ON budget_ledger (thread_id); + +-- schema_meta: single-row table recording the applied schema version so +-- migrate() can detect and step forward. +CREATE TABLE IF NOT EXISTS schema_meta ( + id INTEGER PRIMARY KEY CHECK (id = 1), + schema_version INTEGER NOT NULL +); diff --git a/agent-team/agent_team/deadline_timer.py b/agent-team/agent_team/deadline_timer.py new file mode 100644 index 0000000..3a9d4dd --- /dev/null +++ b/agent-team/agent_team/deadline_timer.py @@ -0,0 +1,354 @@ +"""Deadline / no-answer timer loop for pending questions (design §3.3.1). + +The durable human-in-the-loop ledger (``pending_questions``) gives every +delivered question-set a ``deadline_at``. This module is the **timer loop** +that the design calls out: + + > Each open question has ``deadline_at``. A timer loop flips overdue + > ``open`` rows to ``expired`` (same compare-and-set) and applies the task + > policy: park + ALARM Adam, or apply a defined default answer. An answer + > arriving for an already-``expired`` question loses the compare-and-set and + > is ignored. Timeout vs answer is a deterministic race on flipping + > ``open``. + +Design contract this leaf honours: + +* **Same compare-and-set.** The ``open`` -> ``expired`` flip is the foundation's + :func:`agent_team.db.schema.expire_question`, imported **verbatim** and not + redefined here. It runs under ``BEGIN IMMEDIATE`` so the timer and a racing + responder are serialized: exactly one of "expire" / "answer" wins the row. +* **Deterministic race.** If :func:`expire_question` returns ``False`` for an + overdue row, a responder answered it first (or another timer pass already + expired it); the loop then does **nothing** for that row — it never applies + the no-answer policy to a question that was actually answered. +* **Per-question policy.** When the flip wins, the loop applies that question's + :class:`DeadlinePolicy`: ``PARK`` (park the task + ALARM Adam) or + ``DEFAULT_ANSWER`` (resume the graph with a configured default). Policy is + resolved per question via an injected resolver so the durable ledger schema + is not extended by this leaf. +* **Restart-safe.** All inputs are read from the durable table on each pass; the + loop holds no in-memory-only state, so a reboot mid-sweep simply re-runs the + remaining overdue rows on the next pass (each flip is idempotent via the + compare-and-set). + +Side effects (ALARM, resume-with-default) are injected as callables. This keeps +the module pure of transport / graph I/O — exactly what makes the deterministic +race and the policy branch unit-testable without a live Slack or LangGraph. +""" + +from __future__ import annotations + +import logging +import sqlite3 +from collections.abc import Callable +from dataclasses import dataclass, field +from datetime import datetime, timezone +from enum import Enum + +# Compare-and-set primitive imported VERBATIM from the committed foundation. +# The SQL / transaction discipline is NOT redefined here. +from agent_team.db.schema import expire_question + +__all__ = [ + "DeadlinePolicy", + "ExpiryAction", + "ExpiryOutcome", + "OverdueQuestion", + "TimerLoopReport", + "AlarmFn", + "PolicyResolver", + "ResumeWithDefaultFn", + "overdue_open_questions", + "run_deadline_timer", +] + +logger = logging.getLogger(__name__) + + +class DeadlinePolicy(Enum): + """What to do when a question hits its deadline with no answer (§3.3.1). + + ``PARK`` — park the task and ALARM Adam (the conservative default; a stuck + task parks rather than spins). ``DEFAULT_ANSWER`` — resume the graph with a + pre-configured default answer for tasks where a no-answer has a safe, + defined fallback. + """ + + PARK = "park" + DEFAULT_ANSWER = "default_answer" + + +class ExpiryAction(Enum): + """The action the loop actually took for one overdue row. + + ``PARKED`` / ``DEFAULTED`` follow a *won* expiry flip and the question's + policy. ``LOST_RACE`` means the compare-and-set returned ``False`` — a + responder answered (or a prior pass expired) the row first, so no no-answer + policy was applied. ``ERRORED`` means the flip won but the side effect + raised; the row is already ``expired`` and the failure is recorded for the + caller to ALARM/retry. + """ + + PARKED = "parked" + DEFAULTED = "defaulted" + LOST_RACE = "lost_race" + ERRORED = "errored" + + +@dataclass(frozen=True) +class OverdueQuestion: + """A minimal read-snapshot of one overdue ``open`` question row (§3.3.1). + + Only the columns the timer loop needs: identity, ownership (``thread_id`` / + ``turn`` for the turn-guarded resume), transport, and the deadline. Frozen + because rows are snapshots — mutation goes through the compare-and-set, never + by editing an instance. + """ + + question_id: str + thread_id: str + turn: int + transport: str + deadline_at: str | None = None + + @classmethod + def from_row(cls, row: sqlite3.Row) -> OverdueQuestion: + """Build an :class:`OverdueQuestion` from a ``sqlite3.Row``.""" + return cls( + question_id=row["question_id"], + thread_id=row["thread_id"], + turn=int(row["turn"]), + transport=row["transport"], + deadline_at=row["deadline_at"], + ) + + +@dataclass(frozen=True) +class ExpiryOutcome: + """The result of processing one overdue question (§3.3.1).""" + + question_id: str + thread_id: str + action: ExpiryAction + policy: DeadlinePolicy | None = None + error: str | None = None + + +@dataclass +class TimerLoopReport: + """Aggregate result of one timer-loop pass (§3.3.1). + + ``outcomes`` is one entry per overdue row examined. The summary counters let + the coordinator decide whether to ALARM (any ``errored``) without re-walking + the list. + """ + + outcomes: list[ExpiryOutcome] = field(default_factory=list) + + @property + def examined(self) -> int: + """Number of overdue rows examined this pass.""" + return len(self.outcomes) + + @property + def parked(self) -> int: + """Rows expired and parked (PARK policy).""" + return sum(1 for o in self.outcomes if o.action is ExpiryAction.PARKED) + + @property + def defaulted(self) -> int: + """Rows expired and resumed with a default answer (DEFAULT_ANSWER).""" + return sum(1 for o in self.outcomes if o.action is ExpiryAction.DEFAULTED) + + @property + def lost_race(self) -> int: + """Rows that were answered/expired by someone else first.""" + return sum(1 for o in self.outcomes if o.action is ExpiryAction.LOST_RACE) + + @property + def errored(self) -> int: + """Rows whose flip won but whose side effect raised.""" + return sum(1 for o in self.outcomes if o.action is ExpiryAction.ERRORED) + + @property + def expired(self) -> int: + """Rows this pass actually flipped ``open`` -> ``expired``. + + Every parked / defaulted / errored row won its compare-and-set, so the + row is ``expired`` in the durable table; a ``lost_race`` row did not. + """ + return self.examined - self.lost_race + + +# Injected side-effect / policy seams. Keeping these as callables means the +# timer loop performs no transport or LangGraph I/O of its own (testable, and +# faithful to §3.3.1's transport-agnostic core). +AlarmFn = Callable[[OverdueQuestion], None] +"""Called once per *expired-and-parked* question to ALARM Adam.""" + +ResumeWithDefaultFn = Callable[[OverdueQuestion], None] +"""Called once per *expired* question whose policy is DEFAULT_ANSWER, to resume +the graph with that task's configured default answer.""" + +PolicyResolver = Callable[[OverdueQuestion], DeadlinePolicy] +"""Resolve the :class:`DeadlinePolicy` for one question. Defaults to PARK.""" + + +def _utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string.""" + return datetime.now(timezone.utc).isoformat() + + +def _default_policy_resolver(_question: OverdueQuestion) -> DeadlinePolicy: + """Conservative default: park + ALARM on any no-answer (§3.3.1).""" + return DeadlinePolicy.PARK + + +def overdue_open_questions( + conn: sqlite3.Connection, + *, + now: str | None = None, +) -> list[OverdueQuestion]: + """Return ``open`` rows whose ``deadline_at`` has passed (§3.3.1). + + These are the candidates the timer loop flips with the foundation's + :func:`expire_question`. Rows with a NULL ``deadline_at`` never expire and + are excluded. ``now`` defaults to the current UTC ISO-8601 time; deadlines + are compared as ISO-8601 strings, which sort lexicographically iff written + in a consistent UTC offset — the ledger always writes UTC, so this holds. + Ordered oldest-deadline-first so the most-overdue questions are handled + first. + """ + cutoff = now or _utc_now_iso() + rows = conn.execute( + "SELECT question_id, thread_id, turn, transport, deadline_at " + "FROM pending_questions " + "WHERE status='open' AND deadline_at IS NOT NULL AND deadline_at <= ? " + "ORDER BY deadline_at ASC, question_id ASC", + (cutoff,), + ).fetchall() + return [OverdueQuestion.from_row(r) for r in rows] + + +def run_deadline_timer( + conn: sqlite3.Connection, + *, + on_park: AlarmFn, + resume_with_default: ResumeWithDefaultFn | None = None, + policy_resolver: PolicyResolver | None = None, + now: str | None = None, +) -> TimerLoopReport: + """Run one deadline-timer pass over the ledger (§3.3.1). + + For every ``open`` question past its ``deadline_at`` as of ``now``: + + 1. Attempt the atomic ``open`` -> ``expired`` flip via the foundation's + :func:`expire_question` (the same ``BEGIN IMMEDIATE`` compare-and-set + used by the responder). This is the **deterministic race**: if a + responder answered the question first, the flip returns ``False`` and the + loop records :attr:`ExpiryAction.LOST_RACE` and applies **no** policy. + 2. On a winning flip, resolve the question's :class:`DeadlinePolicy` and act: + * :attr:`DeadlinePolicy.PARK` -> call ``on_park`` (park the task + ALARM + Adam). + * :attr:`DeadlinePolicy.DEFAULT_ANSWER` -> call ``resume_with_default`` + (resume the graph with the configured default). If a question resolves + to ``DEFAULT_ANSWER`` but no ``resume_with_default`` callback was + supplied, that is a configuration error and raises :class:`ValueError` + (the row is already ``expired``; failing loud beats silently dropping a + resume). + + Side effects are isolated per row: if a callback raises, the row is already + durably ``expired`` (the flip committed first), so the loop records + :attr:`ExpiryAction.ERRORED` for that question and continues with the rest of + the batch rather than aborting the whole pass. The caller ALARMs on any + ``errored`` count. + + Restart-safety: all inputs are read fresh from the durable table, and each + flip is idempotent, so re-running the loop after a crash safely processes + only the rows still ``open`` and overdue. + + Returns a :class:`TimerLoopReport` describing what happened to each row. + """ + resolver = policy_resolver or _default_policy_resolver + report = TimerLoopReport() + + for question in overdue_open_questions(conn, now=now): + # Step 1: deterministic race on flipping ``open`` -> ``expired``. + won = expire_question(conn, question_id=question.question_id) + if not won: + # A responder answered first (or a prior pass expired it). Apply no + # no-answer policy — the question was actually answered/handled. + logger.debug( + "deadline-timer: %s lost the expire race (answered/expired first)", + question.question_id, + ) + report.outcomes.append( + ExpiryOutcome( + question_id=question.question_id, + thread_id=question.thread_id, + action=ExpiryAction.LOST_RACE, + ) + ) + continue + + # Step 2: the flip won — the row is durably ``expired``. Apply policy. + policy = resolver(question) + outcome = _apply_policy( + question, + policy=policy, + on_park=on_park, + resume_with_default=resume_with_default, + ) + report.outcomes.append(outcome) + + return report + + +def _apply_policy( + question: OverdueQuestion, + *, + policy: DeadlinePolicy, + on_park: AlarmFn, + resume_with_default: ResumeWithDefaultFn | None, +) -> ExpiryOutcome: + """Apply a question's no-answer policy after a winning expiry flip. + + Isolates the side effect: a callback that raises yields an + :attr:`ExpiryAction.ERRORED` outcome (the row stays ``expired``) instead of + crashing the whole sweep. A ``DEFAULT_ANSWER`` policy with no resume callback + is a configuration error and re-raises :class:`ValueError`. + """ + if policy is DeadlinePolicy.DEFAULT_ANSWER and resume_with_default is None: + raise ValueError( + f"question {question.question_id!r} resolved to DEFAULT_ANSWER but no " + "resume_with_default callback was supplied" + ) + + try: + if policy is DeadlinePolicy.PARK: + on_park(question) + action = ExpiryAction.PARKED + else: # DeadlinePolicy.DEFAULT_ANSWER (resume callback guaranteed above) + assert resume_with_default is not None # narrowed by the guard + resume_with_default(question) + action = ExpiryAction.DEFAULTED + except Exception as exc: # noqa: BLE001 — isolate one row's side-effect failure + logger.exception( + "deadline-timer: side effect for %s (%s) raised; row remains expired", + question.question_id, + policy.value, + ) + return ExpiryOutcome( + question_id=question.question_id, + thread_id=question.thread_id, + action=ExpiryAction.ERRORED, + policy=policy, + error=f"{type(exc).__name__}: {exc}", + ) + + return ExpiryOutcome( + question_id=question.question_id, + thread_id=question.thread_id, + action=action, + policy=policy, + ) diff --git a/agent-team/agent_team/graph.py b/agent-team/agent_team/graph.py new file mode 100644 index 0000000..eeda2f9 --- /dev/null +++ b/agent-team/agent_team/graph.py @@ -0,0 +1,405 @@ +"""LangGraph graph wiring for the Plane-2 SDLC pipeline (design §3.3, §7.1 P1). + +This module is the **P1 skeleton + human gate** wiring (§7.1): + + INTAKE ─► CLARIFY ─► PLAN (stop at an approved plan — no build yet) + +It assembles the durable, resumable LangGraph graph whose state schema is the +foundation's :class:`~agent_team.task_model.PipelineState`. The clarifier raises +a LangGraph ``interrupt()`` carrying a :class:`~agent_team.transport.QuestionSet` +so the task suspends + checkpoints, a question-set is delivered over the chosen +transport (D10), and the task resumes via ``Command(resume=...)`` when Adam's +answer arrives (§3.3, §3.3.1). The planner is the P1 terminal stage: it lands an +approved plan and stops; builders/verifiers are later phases (P3). + +What this module owns (Plane-2 P1 graph wiring only): + +* :data:`P1_PHASE_SEQUENCE` — the ordered P1 stage list. +* the three pure node functions (:func:`intake_node`, :func:`clarify_node`, + :func:`plan_node`) operating on :class:`PipelineState`. +* :func:`build_graph` — assemble + compile the ``StateGraph`` over a *caller- + injected* checkpointer (tests inject an in-memory saver; production injects + the SQLite saver, which the foundation's :func:`agent_team.db.connect` + reserves the DB file for). +* :func:`build_sqlite_checkpointer` — the production checkpointer factory, with + the ``langgraph.checkpoint.sqlite`` import deferred so this module imports + cleanly even where that optional package is absent (pre-deploy scaffolding). +* :func:`thread_config` / :func:`start_task` / :func:`resume_task` / + :func:`get_pipeline_state` / :func:`pending_question` — the thin + ``thread_id``-keyed driver seam the coordinator/responder call. + +It imports the committed foundation contracts verbatim and does **not** redefine +them. No provisioning, no scheduling, no live SDK calls: the clarifier's +question authoring is a deterministic stub here (the real Claude clarifier binds +``billing.claude_invoke`` in a later phase), and the human-interaction *ledger* +(``pending_questions``) lives in :mod:`agent_team.db.schema` — this module only +shapes the interrupt payload that drives it. +""" + +from __future__ import annotations + +import uuid +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import TYPE_CHECKING, Any + +from langgraph.graph import END, START, StateGraph +from langgraph.types import Command, interrupt + +from agent_team.task_model import ( + Phase, + PipelineState, + TaskStatus, + new_thread_id, +) +from agent_team.transport import QuestionSet + +if TYPE_CHECKING: # pragma: no cover - typing only + from langgraph.checkpoint.base import BaseCheckpointSaver + from langgraph.graph.state import CompiledStateGraph + +__all__ = [ + "CLARIFY", + "DEFAULT_CLARIFY_DEADLINE", + "INTAKE", + "P1_PHASE_SEQUENCE", + "PLAN", + "build_graph", + "build_sqlite_checkpointer", + "clarify_node", + "get_pipeline_state", + "intake_node", + "pending_question", + "plan_node", + "plan_phase", + "resume_task", + "start_task", + "thread_config", +] + +# --- Node names (graph vertices). ------------------------------------------ +# Kept as constants so the driver/tests reference the wiring by name rather +# than by string literal. +INTAKE = "intake" +CLARIFY = "clarify" +PLAN = "plan" + +# The P1 stage order (§7.1): intake -> clarify -> plan, then stop. Builders and +# verifiers (BUILD/VERIFY) are deliberately NOT wired here — P1 ends at an +# approved plan with no build (§7.1 "Stops at an approved plan, no build yet"). +P1_PHASE_SEQUENCE: tuple[Phase, ...] = (Phase.INTAKE, Phase.CLARIFY, Phase.PLAN) + +# How long a clarifier question-set stays open before the deadline policy runs +# (§3.3.1 ``deadline_at``). The driver records the concrete ``deadline_at`` on +# the ledger row; this is only the default window the interrupt advertises. +DEFAULT_CLARIFY_DEADLINE = timedelta(hours=24) + +# Fixed namespace for deriving a STABLE question_id from (thread_id, turn). The +# clarifier node re-executes from its start on resume (LangGraph replays the +# node, with interrupt() returning the answer the second time), so a fresh +# random id would change between the suspend that delivered/ledgered the +# question and the resume that records it — breaking the §3.3.1 identity +# contract. A uuid5 over (thread_id, turn) is uuid-shaped yet deterministic, so +# the delivered question_id, the ledger key, and the qa_history entry all agree. +_QUESTION_ID_NAMESPACE = uuid.UUID("a7b9c1d2-3e4f-5061-7283-94a5b6c7d8e9") + + +def _question_id_for(thread_id: str, turn: int) -> str: + """Return the stable question_id for ``(thread_id, turn)`` (§3.3.1 identity).""" + return uuid.uuid5(_QUESTION_ID_NAMESPACE, f"{thread_id}:{turn}").hex + + +def _utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string (ledger-compatible).""" + return datetime.now(timezone.utc).isoformat() + + +def _phase_value(phase: Phase) -> str: + """Return the string value a phase is stored as in :class:`PipelineState`.""" + return phase.value + + +# --- Nodes. ----------------------------------------------------------------- +# Each node is a pure ``PipelineState -> partial PipelineState`` function. They +# write only the keys they change (PipelineState is ``total=False``), so a +# checkpoint transition stays minimal. None of them performs I/O. + + +def intake_node(state: PipelineState) -> PipelineState: + """INTAKE stage: stamp the task ACTIVE and advance it into CLARIFY (§3.3). + + A task enters as a new thread record (the driver mints ``thread_id`` and + seeds INTAKE). This node marks it ``ACTIVE`` and moves the current phase to + ``CLARIFY`` so the next node runs the human gate. It never blocks. + """ + return PipelineState( + status=TaskStatus.ACTIVE.value, + current_phase=_phase_value(Phase.CLARIFY), + updated_at=_utc_now_iso(), + ) + + +def clarify_node(state: PipelineState) -> PipelineState: + """CLARIFY stage: the human gate (LangGraph ``interrupt()``) (§3.3, §3.3.1). + + Authors a question-set, then suspends the graph with ``interrupt()`` so the + task checkpoints and waits for Adam's answer. The interrupt payload is the + :class:`~agent_team.transport.QuestionSet` plus the lifecycle metadata the + responder/ledger need (``transport``, ``deadline``); the responder posts it + over the chosen transport and resumes the task via ``Command(resume=...)``. + + On resume, ``interrupt()`` returns Adam's answer; this node appends it to + ``qa_history`` and advances to ``PLAN``. The P1 skeleton asks exactly one + question-set (``turn`` 0); the multi-turn "until 98% confident" loop is a + later phase, and this node's single-turn shape is forward-compatible with it + (the ``turn`` is read from existing history). + + The actual Claude question authoring (``billing.claude_invoke``) is bound in + a later phase; here the question-set is a deterministic stub so the wiring + and the suspend/resume mechanic can be proven without a live model. + """ + history = list(state.get("qa_history", [])) + turn = len(history) + thread_id = state.get("thread_id", "") + transport = state.get("transport", "") + + # Stable across the resume replay of this node (see _question_id_for): the + # id delivered at suspend == the ledger key == the qa_history entry. + question_id = _question_id_for(thread_id, turn) + question_set = QuestionSet( + thread_id=thread_id, + question_id=question_id, + turn=turn, + questions=_author_questions(state), + context={"phase": _phase_value(Phase.CLARIFY)}, + ) + deadline = (datetime.now(timezone.utc) + DEFAULT_CLARIFY_DEADLINE).isoformat() + + # Suspend here. The payload mirrors §3.3.1: {thread_id, question_id, turn, + # question_set, transport, deadline}. On resume, ``answer`` is whatever the + # responder passed to ``Command(resume=...)``. + answer = interrupt( + { + "thread_id": thread_id, + "question_id": question_id, + "turn": turn, + "question_set": question_set, + "transport": transport, + "deadline": deadline, + } + ) + + history.append({"turn": turn, "question_id": question_id, "answer": answer}) + return PipelineState( + status=TaskStatus.ACTIVE.value, + current_phase=_phase_value(Phase.PLAN), + qa_history=history, + updated_at=_utc_now_iso(), + ) + + +def plan_node(state: PipelineState) -> PipelineState: + """PLAN stage: land an approved plan and stop — the P1 terminus (§7.1). + + Produces the phased plan record and marks the task ``DONE`` for P1 purposes + (P1 "stops at an approved plan, no build yet"). The real planner is Claude + (§3.3); here the plan body is a deterministic stub derived from the gathered + Q&A so the terminal-state wiring is exercised. Builders/verifiers are wired + in P3. + """ + plan = plan_phase(state) + return PipelineState( + status=TaskStatus.DONE.value, + current_phase=_phase_value(Phase.DONE), + plan=plan, + updated_at=_utc_now_iso(), + ) + + +def _author_questions(state: PipelineState) -> list[str]: + """Deterministic stand-in for the Claude clarifier's question authoring. + + The real clarifier gathers repo/memory/handbook context and asks until 98% + confident (§3.3); the P1 skeleton asks a single fixed question-set so the + suspend/resume mechanic is what's under test, not the model. + """ + return ["What problem should this task solve, and what is in scope?"] + + +def plan_phase(state: PipelineState) -> dict[str, Any]: + """Build the deterministic P1 plan record from the clarifier Q&A. + + Exposed (and unit-tested) separately from :func:`plan_node` so the plan + shape can be asserted without driving the whole graph. The real planner + replaces the body in a later phase. + """ + return { + "summary": "Approved P1 plan (skeleton).", + "phases": ["P1: skeleton + human gate"], + "qa_turns": len(state.get("qa_history", [])), + "approved": True, + } + + +# --- Graph assembly. -------------------------------------------------------- + + +def build_graph( + checkpointer: BaseCheckpointSaver | None = None, +) -> CompiledStateGraph: + """Assemble + compile the P1 pipeline ``StateGraph`` (§3.3, §7.1). + + Wires ``START → intake → clarify → plan → END`` over + :class:`PipelineState`. The clarifier suspends on ``interrupt()`` for the + human gate; the planner is the P1 terminus (no build). + + The ``checkpointer`` is **injected**, never constructed here: the design's + durable store is the SQLite checkpointer (D9), but pre-deploy scaffolding + must not provision it, and tests inject an in-memory saver. Production wires + :func:`build_sqlite_checkpointer`. A checkpointer is required for the + ``interrupt()``/``resume`` mechanic to work, so callers that pass ``None`` + get an uncheckpointed graph that can run straight-through but cannot + suspend; the driver functions therefore require a checkpointed graph. + """ + builder: StateGraph = StateGraph(PipelineState) + builder.add_node(INTAKE, intake_node) + builder.add_node(CLARIFY, clarify_node) + builder.add_node(PLAN, plan_node) + + builder.add_edge(START, INTAKE) + builder.add_edge(INTAKE, CLARIFY) + builder.add_edge(CLARIFY, PLAN) + builder.add_edge(PLAN, END) + + if checkpointer is None: + return builder.compile() + return builder.compile(checkpointer=checkpointer) + + +def build_sqlite_checkpointer(db_path: Path | str) -> BaseCheckpointSaver: + """Construct the production SQLite checkpointer over ``db_path`` (D9, §3.3). + + The import of ``langgraph.checkpoint.sqlite`` is deferred to call time so + this module imports cleanly in environments where that optional package is + not installed (pre-deploy scaffolding). The checkpointer creates its own + tables against the same DB file the foundation's + :func:`agent_team.db.init_db` reserves for it. + + Raises a clear :class:`RuntimeError` if the optional package is missing, so + a misconfigured deploy fails loudly rather than silently running + uncheckpointed. + """ + try: + from langgraph.checkpoint.sqlite import SqliteSaver + except ImportError as exc: # pragma: no cover - depends on optional dep + raise RuntimeError( + "langgraph SQLite checkpointer is unavailable; install the " + "'langgraph-checkpoint-sqlite' package to use " + "build_sqlite_checkpointer (D9). Tests inject an in-memory saver." + ) from exc + + db_path = Path(db_path) + db_path.parent.mkdir(parents=True, exist_ok=True) + return SqliteSaver.from_conn_string(str(db_path)) + + +# --- Driver seam (thread_id-keyed). ----------------------------------------- +# Thin helpers the coordinator/responder call. They own the mapping between a +# task's ``thread_id`` and the LangGraph ``config``; the durable lifecycle +# ledger lives in agent_team.db, and the transport delivery in agent_team +# .transport. These keep that wiring in one tested place. + + +def thread_config(thread_id: str) -> dict[str, Any]: + """Build the LangGraph ``config`` that scopes an invoke to ``thread_id``. + + Every checkpointed invoke/resume for a task must carry the same + ``{"configurable": {"thread_id": ...}}`` so it reads/writes that task's + checkpoint and no other (§3.3.1 per-thread isolation). + """ + return {"configurable": {"thread_id": thread_id}} + + +def start_task( + graph: CompiledStateGraph, + *, + thread_id: str | None = None, + transport: str = "", +) -> tuple[str, PipelineState]: + """Start a new pipeline task and run it up to the first human gate (§3.3). + + Mints a ``thread_id`` (unless one is supplied), seeds the INTAKE state, and + invokes the graph; it runs through INTAKE into CLARIFY and suspends on the + clarifier ``interrupt()``. Returns ``(thread_id, state)`` where ``state`` is + the checkpointed snapshot after the suspend (its ``__interrupt__`` carries + the pending question-set, surfaced by :func:`pending_question`). + + The graph MUST be compiled with a checkpointer for the suspend to persist; + an uncheckpointed graph would run straight through without honouring the + interrupt. + """ + tid = thread_id or new_thread_id() + now = _utc_now_iso() + seed = PipelineState( + thread_id=tid, + status=TaskStatus.ACTIVE.value, + current_phase=_phase_value(Phase.INTAKE), + qa_history=[], + transport=transport, + created_at=now, + updated_at=now, + ) + result = graph.invoke(seed, thread_config(tid)) + return tid, result + + +def resume_task( + graph: CompiledStateGraph, + *, + thread_id: str, + answer: Any, +) -> PipelineState: + """Resume a suspended task with Adam's ``answer`` (§3.3, §3.3.1). + + Calls ``graph.invoke(Command(resume=answer), config)`` for ``thread_id``. + The clarifier's ``interrupt()`` returns ``answer``, the task appends it to + ``qa_history`` and advances through PLAN to completion. The §3.3.1 + first-answer-wins / turn-guard discipline lives in the responder + ledger; + this helper is the single-flight resume call the resume worker drives once + it has won the compare-and-set. + """ + return graph.invoke(Command(resume=answer), thread_config(thread_id)) + + +def get_pipeline_state( + graph: CompiledStateGraph, + *, + thread_id: str, +) -> PipelineState: + """Return the live checkpointed :class:`PipelineState` for ``thread_id``. + + Reads the current checkpoint snapshot (post-suspend or post-completion). + Used by recovery + the manual CLI to inspect a task without resuming it. + """ + snapshot = graph.get_state(thread_config(thread_id)) + return snapshot.values + + +def pending_question( + graph: CompiledStateGraph, + *, + thread_id: str, +) -> dict[str, Any] | None: + """Return the pending interrupt payload for ``thread_id``, or ``None``. + + When a task is suspended on the clarifier human gate, its checkpoint carries + an interrupt whose value is the §3.3.1 question payload (``thread_id``, + ``question_id``, ``turn``, ``question_set``, ``transport``, ``deadline``). + The responder reads this to author the ledger row + transport post. Returns + ``None`` when the task is not currently waiting on a human answer. + """ + snapshot = graph.get_state(thread_config(thread_id)) + interrupts = getattr(snapshot, "interrupts", None) or () + if not interrupts: + return None + return interrupts[0].value diff --git a/agent-team/agent_team/ledger.py b/agent-team/agent_team/ledger.py new file mode 100644 index 0000000..cafb7c0 --- /dev/null +++ b/agent-team/agent_team/ledger.py @@ -0,0 +1,305 @@ +"""Pending-questions ledger ops — the §3.3.1 durable source of truth. + +The LangGraph SQLite checkpointer suspends/resumes the graph, but the +checkpoint alone does not track the human-interaction lifecycle (delivery, +duplicate/late answers, expiry). This module is the thin, transport-agnostic +operations layer over the ``pending_questions`` table declared in +:mod:`agent_team.db.schema`. It is what delivery, the responder, the deadline +timer, the resume worker, restart recovery, and the manual CLI all read and +write. + +Design (§3.3.1) mapping: + +* **Delivery (and lost-post).** :func:`post_question` writes the row ``open`` + *first* (no ``channel_ref``); the caller then posts to the transport and + records the ref via :func:`set_channel_ref`. If the post fails the row stays + ``open`` with no ref and the reconcile loop (:func:`open_questions_needing_ref`) + retries idempotently. +* **Answer / expiry / supersede.** The atomic compare-and-set helpers live in + :mod:`agent_team.db.schema` and run under ``BEGIN IMMEDIATE``. They are + imported here verbatim and re-exported as the ledger's public mutation API + (:func:`answer_question`, :func:`expire_question`, :func:`supersede_question`) + so callers depend on one module — the SQL is **not** redefined here. +* **Deadline / no-answer.** :func:`overdue_open_questions` returns the ``open`` + rows past their ``deadline_at`` for the timer loop to flip with + :func:`expire_question`. +* **Resume (turn-guarded).** :func:`answered_questions` feeds the resume worker + the ``answered`` rows whose graph may still be interrupted on that turn. +* **Restart recovery.** The startup sweep reads + :func:`open_questions_needing_ref` (retry delivery), :func:`answered_questions` + (re-enqueue resume), and :func:`overdue_open_questions` (run deadline policy). + All state is durable, so recovery is purely a function of the table. +* **Manual path (CLI).** :func:`list_questions` (optionally filtered by status) + and :func:`get_question` back the operator CLI that lists ``open`` questions, + re-delivers, force-expires, or answers on a task's behalf. + +This module is pure stdlib + the foundation modules; it performs no transport +I/O of its own. +""" + +from __future__ import annotations + +import sqlite3 +from dataclasses import dataclass +from datetime import datetime, timezone + +# Compare-and-set primitives + lifecycle constants are imported VERBATIM from +# the committed foundation. They are NOT redefined here; the ledger re-exports +# them so callers depend on a single operations module. +from agent_team.db.schema import ( + QUESTION_STATES, + answer_question, + expire_question, + supersede_question, +) + +__all__ = [ + "PendingQuestion", + "QUESTION_STATES", + "answer_question", + "answered_questions", + "count_by_status", + "expire_question", + "get_question", + "list_questions", + "open_questions_needing_ref", + "overdue_open_questions", + "post_question", + "set_channel_ref", + "supersede_question", +] + + +@dataclass(frozen=True) +class PendingQuestion: + """A typed view of one ``pending_questions`` row (§3.3.1 ledger). + + Mirrors the table columns declared in :data:`agent_team.db.schema. + PENDING_QUESTIONS_DDL`. ``frozen`` because rows are read snapshots; mutation + goes through the compare-and-set helpers, never by editing an instance. + """ + + question_id: str + thread_id: str + turn: int + status: str + transport: str + channel_ref: str | None = None + posted_at: str | None = None + deadline_at: str | None = None + answer_json: str | None = None + answered_at: str | None = None + answered_via: str | None = None + + @classmethod + def from_row(cls, row: sqlite3.Row) -> PendingQuestion: + """Build a :class:`PendingQuestion` from a ``sqlite3.Row``.""" + return cls( + question_id=row["question_id"], + thread_id=row["thread_id"], + turn=int(row["turn"]), + status=row["status"], + transport=row["transport"], + channel_ref=row["channel_ref"], + posted_at=row["posted_at"], + deadline_at=row["deadline_at"], + answer_json=row["answer_json"], + answered_at=row["answered_at"], + answered_via=row["answered_via"], + ) + + +def _utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string.""" + return datetime.now(timezone.utc).isoformat() + + +def post_question( + conn: sqlite3.Connection, + *, + question_id: str, + thread_id: str, + turn: int, + transport: str, + deadline_at: str | None = None, + posted_at: str | None = None, +) -> None: + """Insert a new ``open`` question row (delivery step 1 of §3.3.1). + + Writes the row ``open`` with **no** ``channel_ref`` *before* the caller + posts to the transport. If the subsequent post fails, the row stays ``open`` + with no ref and the reconcile loop (:func:`open_questions_needing_ref`) + retries delivery idempotently. The caller records the ref via + :func:`set_channel_ref` once the post succeeds. + + Raises :class:`sqlite3.IntegrityError` if ``question_id`` already exists + (PK) — re-posting the same question is the caller's reconcile concern, not a + silent overwrite. ``posted_at`` defaults to now (UTC ISO-8601). + """ + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, posted_at, deadline_at) " + "VALUES (?, ?, ?, 'open', ?, ?, ?)", + ( + question_id, + thread_id, + int(turn), + transport, + posted_at or _utc_now_iso(), + deadline_at, + ), + ) + + +def set_channel_ref( + conn: sqlite3.Connection, + *, + question_id: str, + channel_ref: str, +) -> bool: + """Record the transport ``channel_ref`` after a successful post (§3.3.1). + + Sets the Slack message ts / GitHub issue-comment id / Claude session id on + a still-``open`` row. Guarded on ``status='open'`` so a late post-confirm + cannot resurrect a ref on an already-answered/expired/superseded question. + Returns ``True`` if exactly one open row was updated, ``False`` otherwise + (unknown id, or no longer open) — letting the reconcile loop decide whether + to retry. + """ + cur = conn.execute( + "UPDATE pending_questions SET channel_ref=? " + "WHERE question_id=? AND status='open'", + (channel_ref, question_id), + ) + return cur.rowcount == 1 + + +def get_question( + conn: sqlite3.Connection, + question_id: str, +) -> PendingQuestion | None: + """Return one question by id, or ``None`` if absent.""" + row = conn.execute( + "SELECT * FROM pending_questions WHERE question_id=?", + (question_id,), + ).fetchone() + return PendingQuestion.from_row(row) if row is not None else None + + +def list_questions( + conn: sqlite3.Connection, + *, + status: str | None = None, + thread_id: str | None = None, +) -> list[PendingQuestion]: + """List questions, optionally filtered by ``status`` and/or ``thread_id``. + + Backs the manual CLI's ``list`` view (§3.3.1). ``status`` must be one of + :data:`QUESTION_STATES` when given; an unknown status raises + :class:`ValueError` rather than silently returning nothing. Results are + ordered oldest-first by ``posted_at`` so the operator sees the longest-open + questions first. + """ + if status is not None and status not in QUESTION_STATES: + raise ValueError( + f"unknown status {status!r}; expected one of {QUESTION_STATES}" + ) + + clauses: list[str] = [] + params: list[object] = [] + if status is not None: + clauses.append("status=?") + params.append(status) + if thread_id is not None: + clauses.append("thread_id=?") + params.append(thread_id) + + where = f" WHERE {' AND '.join(clauses)}" if clauses else "" + rows = conn.execute( + "SELECT * FROM pending_questions" + f"{where} ORDER BY posted_at ASC, question_id ASC", + params, + ).fetchall() + return [PendingQuestion.from_row(r) for r in rows] + + +def open_questions_needing_ref( + conn: sqlite3.Connection, +) -> list[PendingQuestion]: + """Return ``open`` rows that have no ``channel_ref`` (lost-post reconcile). + + The reconcile/recovery sweep (§3.3.1) retries delivery for these + idempotently: an ``open`` row with no ref means the row was written but the + transport post never confirmed. + """ + rows = conn.execute( + "SELECT * FROM pending_questions " + "WHERE status='open' AND channel_ref IS NULL " + "ORDER BY posted_at ASC, question_id ASC" + ).fetchall() + return [PendingQuestion.from_row(r) for r in rows] + + +def overdue_open_questions( + conn: sqlite3.Connection, + *, + now: str | None = None, +) -> list[PendingQuestion]: + """Return ``open`` rows whose ``deadline_at`` has passed (deadline loop). + + Feeds the §3.3.1 timer loop, which flips each returned row with + :func:`expire_question` (the same compare-and-set, so an answer racing the + timer is resolved deterministically). Rows with a NULL ``deadline_at`` never + expire and are excluded. ``now`` defaults to the current UTC ISO-8601 time; + deadlines are compared as ISO-8601 strings, which sort lexicographically iff + written in a consistent UTC offset (the ledger always writes UTC). + """ + cutoff = now or _utc_now_iso() + rows = conn.execute( + "SELECT * FROM pending_questions " + "WHERE status='open' AND deadline_at IS NOT NULL AND deadline_at <= ? " + "ORDER BY deadline_at ASC, question_id ASC", + (cutoff,), + ).fetchall() + return [PendingQuestion.from_row(r) for r in rows] + + +def answered_questions( + conn: sqlite3.Connection, + *, + thread_id: str | None = None, +) -> list[PendingQuestion]: + """Return ``answered`` rows (resume-worker / restart-recovery feed). + + The turn-guarded resume worker (§3.3.1) re-enqueues a resume for each + ``answered`` row whose graph is still interrupted on that turn; the guard + makes a redelivered job idempotent. Optionally scoped to one ``thread_id`` + (the resume worker serializes per thread). Ordered by ``answered_at`` so the + oldest pending resume is handled first. + """ + clauses = ["status='answered'"] + params: list[object] = [] + if thread_id is not None: + clauses.append("thread_id=?") + params.append(thread_id) + rows = conn.execute( + "SELECT * FROM pending_questions " + f"WHERE {' AND '.join(clauses)} " + "ORDER BY answered_at ASC, question_id ASC", + params, + ).fetchall() + return [PendingQuestion.from_row(r) for r in rows] + + +def count_by_status(conn: sqlite3.Connection) -> dict[str, int]: + """Return a ``{status: count}`` map over all :data:`QUESTION_STATES`. + + Backs CLI/telemetry summaries. Every state in :data:`QUESTION_STATES` is + present in the result (zero when absent) so callers get a stable shape. + """ + counts = dict.fromkeys(QUESTION_STATES, 0) + for row in conn.execute( + "SELECT status, COUNT(*) AS n FROM pending_questions GROUP BY status" + ).fetchall(): + counts[row["status"]] = int(row["n"]) + return counts diff --git a/agent-team/agent_team/nodes/__init__.py b/agent-team/agent_team/nodes/__init__.py new file mode 100644 index 0000000..03cfbb7 --- /dev/null +++ b/agent-team/agent_team/nodes/__init__.py @@ -0,0 +1,10 @@ +"""R720 agent-team — Plane-2 pipeline NODES (design §3.3). + +Leaf modules implementing one LangGraph node each (a pipeline stage). Nodes are +thin: they read/write :class:`agent_team.task_model.PipelineState`, lean on the +committed foundation contracts (state-store, billing seam, transports, task +model), and keep the heavy/risky decisions in dedicated pure-code modules (e.g. +the §3.3.2 pass/fail gate in :mod:`agent_team.ci_gate`). +""" + +__all__: list[str] = [] diff --git a/agent-team/agent_team/nodes/builders.py b/agent-team/agent_team/nodes/builders.py new file mode 100644 index 0000000..9920954 --- /dev/null +++ b/agent-team/agent_team/nodes/builders.py @@ -0,0 +1,618 @@ +"""Plane-2 builders node — approved plan -> candidate diff (design §3.3.2, §7.1 P3). + +This is the LangGraph **builders** stage (design §3.3): it turns the approved, +review-cleared plan into a **candidate diff**. Per D2/D11 the box has no write +token; builders do NOT write to repos — they emit the diff as data for the +org-CI apply/verify workflow. This leaf owns the **box-side half** of the +§3.3.2 CI-as-verifier trust boundary: + +* **Box-side trust-control-surface denylist (boundary #2).** Before a diff can + advance to the build/verify path, this node rejects any candidate diff that + touches the trust-control surface — ``.github/workflows/**``, IAM/policy/ + permission IaC (CDK/SAM), branch-protection / ``CODEOWNERS`` / Dependabot + config, or any file **outside the task's declared scope**. The match is not + naive: it canonicalizes paths (resolving ``.`` / ``..`` and rejecting absolute + or parent-escaping paths), and it inspects **rename targets** so a rename into + a denied path cannot slip through. A violation parks the task for mandatory + human + GPT cross-review (it is the mandatory-cross-review surface regardless, + per CLAUDE.md) and never auto-advances. +* **Diff integrity hash (boundary #3).** The accepted diff is hashed with the + same content-hash primitive the foundation uses + (:func:`agent_team.state_store.compute_content_hash`) and the hash is recorded + on the task record (``diff_hash``). CI verifies this hash matches the ledger + before applying the patch, so a tampered/substituted diff fails closed. + +The CI-side enforcement (the credential-less untrusted job, the CI hard-fail +guard, the pure-code pass/fail gate, branch protection — boundaries #1/#4/#5) +lives in the CI workflow, NOT here; this node is the box-side pre-check plus the +hash the CI gate keys against. + +The actual diff synthesis (Claude spec via the §3.1 billing seam + DeepSeek +mechanical edits) is delegated to an **injectable** ``DiffBuilder`` so this leaf +stays unit-testable and dependency-free; the real SDK/orchestrator wiring is +bound by the coordinator at startup. The default builder uses +:func:`agent_team.billing.claude_invoke` so an un-wired environment fails loudly +via the foundation's unconfigured-invoker contract rather than silently +producing nothing. + +This module imports the committed foundation contracts verbatim; it redefines +none of them. +""" + +from __future__ import annotations + +import posixpath +import re +from dataclasses import dataclass, field +from typing import Any, Mapping, Protocol + +from agent_team.billing import ClaudeResult, claude_invoke +from agent_team.state_store import compute_content_hash +from agent_team.task_model import Phase, PipelineState, TaskStatus + +__all__ = [ + "DENYLIST_REASONS", + "BuildError", + "DiffBuilder", + "TrustBoundaryViolation", + "build_candidate_diff", + "builders_node", + "default_diff_builder", + "iter_diff_target_paths", + "scan_trust_control_surface", +] + + +# --------------------------------------------------------------------------- +# Trust-control-surface denylist (design §3.3.2 boundary #2) +# --------------------------------------------------------------------------- +# +# Patterns are matched against POSIX-canonicalized, repo-relative paths (see +# ``_canonicalize``). Each entry is (compiled regex, human reason) so a +# violation report names *why* a path is denied. The patterns intentionally +# over-match toward rejection: a denied diff is escalated to human + GPT +# cross-review, never silently dropped, so false positives cost a review, not a +# security hole. +_DENY_PATTERNS: tuple[tuple[re.Pattern[str], str], ...] = ( + ( + re.compile(r"^\.github/workflows/.+"), + "modifies a GitHub Actions workflow (.github/workflows/**)", + ), + ( + re.compile(r"(^|/)CODEOWNERS$"), + "modifies CODEOWNERS", + ), + ( + re.compile(r"(^|/)\.github/dependabot\.ya?ml$|^dependabot\.ya?ml$"), + "modifies Dependabot configuration", + ), + ( + re.compile(r"(^|/)\.github/settings\.ya?ml$"), + "modifies repo/branch-protection settings (.github/settings.yml)", + ), + # IAM / policy / permission IaC (CDK/SAM and raw policy docs). These are the + # mandatory-cross-review surface regardless of the pipeline (CLAUDE.md). + ( + re.compile( + r"(^|/)(template|samconfig)\.ya?ml$" + r"|(^|/)serverless\.ya?ml$" + ), + "modifies SAM/serverless IaC (templates carry IAM policies)", + ), + ( + re.compile( + r"(^|/)cdk\.json$" + r"|(^|/).+\.(iam|policy)\.(json|ya?ml)$" + r"|(^|/)(iam|policies|policy)/.+\.(json|ya?ml)$" + # Parity with the CI-side denylist (case-insensitive): bare policy + # docs, any path naming "iam", Terraform, and CDK stack files — + # these were box-side gaps a diff could use to skip the cross-review. + r"|(^|/)policy[^/]*\.json$" + r"|(^|/)[^/]*iam[^/]*$" + r"|\.tf$" + r"|(^|/).+[-_]stack\.(ts|py)$", + re.IGNORECASE, + ), + "modifies IAM/policy/Terraform/CDK IaC", + ), + ( + re.compile(r"^\.github/actions/.+", re.IGNORECASE), + "modifies a composite GitHub Action (.github/actions/**)", + ), + ( + re.compile(r"\.(pem|key)$", re.IGNORECASE), + "modifies key material (*.pem / *.key)", + ), +) + +# A stable, importable description of every denylist reason, useful for callers +# (reports, tests) that want to enumerate the surface without re-deriving it. +DENYLIST_REASONS: tuple[str, ...] = tuple(reason for _, reason in _DENY_PATTERNS) + + +class BuildError(Exception): + """Raised when the builders node cannot produce a usable candidate diff. + + Signals a malformed approved plan or an empty/garbled diff from the + injected builder — i.e. the node cannot proceed, distinct from a *policy* + rejection (:class:`TrustBoundaryViolation`), which is a successful scan that + found a forbidden change. + """ + + +@dataclass +class TrustBoundaryViolation: + """A single trust-control-surface denylist hit (design §3.3.2 boundary #2). + + ``path`` is the canonicalized repo-relative path that tripped the check; + ``reason`` is the human-readable denylist rule (or an out-of-scope / unsafe + -path explanation). ``rename_from`` is set when the violation is a rename + whose *target* lands in a denied/out-of-scope location, so a rename cannot + launder a forbidden path. + """ + + path: str + reason: str + rename_from: str | None = None + + +class DiffBuilder(Protocol): + """Injectable diff-synthesis seam (design §3.3 builders stage). + + A ``DiffBuilder`` turns the approved ``plan`` into a unified-diff string. + The real implementation wires Claude (spec, via the §3.1 billing seam) and + DeepSeek (mechanical edits, via the local orchestrator); tests pass a stub. + It MUST return a unified diff and MUST NOT perform any repo writes (D2/D11). + """ + + def __call__( + self, *, plan: Mapping[str, Any], config: Mapping[str, Any] | None + ) -> str: + """Return the candidate unified diff for ``plan``.""" + ... + + +def default_diff_builder( + *, plan: Mapping[str, Any], config: Mapping[str, Any] | None +) -> str: + """Default :class:`DiffBuilder`: author the diff via the Claude billing seam. + + Renders the approved plan into an instruction and calls + :func:`agent_team.billing.claude_invoke` (the §3.1 seam) to produce the + unified diff. Because the seam's default invoker raises until + :func:`agent_team.billing.set_invoker` is called, an un-wired environment + fails loudly here rather than emitting an empty diff. The coordinator binds + the real Claude-spec + DeepSeek-edit path at startup. + """ + prompt = _render_build_prompt(plan) + result: ClaudeResult = claude_invoke(prompt, config=config) + return result.text + + +def _render_build_prompt(plan: Mapping[str, Any]) -> str: + """Render the approved plan into a builder instruction prompt. + + Kept deliberately small and deterministic: the plan is the source of truth + and the builder's job is to emit a unified diff implementing it without + touching the trust-control surface (§3.3.2). + """ + title = str(plan.get("title", "(untitled task)")) + scope = plan.get("scope") or [] + phases = plan.get("phases") or [] + scope_lines = "\n".join(f" - {p}" for p in scope) or " (no scope declared)" + phase_lines = ( + "\n".join(f" {i + 1}. {p}" for i, p in enumerate(phases)) or " (none)" + ) + return ( + "Implement the approved plan below as a single unified diff (git " + "format). Touch ONLY files within the declared scope. Do NOT modify " + "CI workflows, IAM/policy IaC, branch-protection, CODEOWNERS, or " + "Dependabot config.\n\n" + f"Title: {title}\n" + f"Declared scope (paths you may edit):\n{scope_lines}\n" + f"Phases:\n{phase_lines}\n" + ) + + +# --------------------------------------------------------------------------- +# Diff parsing + canonicalization +# --------------------------------------------------------------------------- + +# A unified diff is parsed **per file section**, each section delimited by its +# ``diff --git a/ b/`` header. That header carries BOTH the source and +# destination path for every change kind — modify, add, delete, mode-change, +# rename, and copy — so reading it (not only the ``+++ b/`` body line) is what +# lets the scan see deletes, mode-only changes, and ``copy to`` targets that have +# no ``+++`` line or whose ``+++`` is ``/dev/null``. The ``---``/``+++`` and +# rename/copy ``from``/``to`` lines refine the section's source/dest when present. +_PLUS_RE = re.compile(r"^\+\+\+ (?:b/)?(.+?)\s*$") +_MINUS_RE = re.compile(r"^--- (?:a/)?(.+?)\s*$") +_DIFF_GIT_RE = re.compile(r"^diff --git a/(.+?) b/(.+?)\s*$") +_RENAME_FROM_RE = re.compile(r"^rename from (.+?)\s*$") +_RENAME_TO_RE = re.compile(r"^rename to (.+?)\s*$") +_COPY_FROM_RE = re.compile(r"^copy from (.+?)\s*$") +_COPY_TO_RE = re.compile(r"^copy to (.+?)\s*$") + +# /dev/null appears as the source of an add or target of a delete; it is never a +# real repo path and must not be scanned/scoped as one. +_DEV_NULL = "/dev/null" + + +@dataclass +class _DiffTarget: + """An internal record of one path the diff would create/modify/rename. + + ``path`` is the canonicalized destination; ``rename_from`` is the prior + canonical path when this target is the destination of a git rename. + """ + + path: str + rename_from: str | None = None + + +def _canonicalize(raw: str) -> str | None: + """Canonicalize a repo-relative diff path; return ``None`` if it is unsafe. + + Strips a leading ``a/`` / ``b/`` prefix, normalizes ``.``/``..`` segments + with :func:`posixpath.normpath`, and rejects anything that escapes the repo + root (absolute paths, or a normalized path beginning with ``..``). Returning + ``None`` signals an *unsafe* path that the scan treats as a violation rather + than silently letting an indirection bypass the denylist (§3.3.2 boundary + #2: "it resolves symlinks and canonicalizes paths"). + """ + path = raw.strip() + if not path or path == _DEV_NULL: + return None + for prefix in ("a/", "b/"): + if path.startswith(prefix): + path = path[len(prefix) :] + break + # Normalize backslashes to forward slashes so a Windows-style separator + # cannot smuggle a segment past the POSIX normalizer. + path = path.replace("\\", "/") + if posixpath.isabs(path): + return None + normalized = posixpath.normpath(path) + if normalized == "." or normalized.startswith("../") or normalized == "..": + return None + return normalized + + +def iter_diff_target_paths(diff: str) -> list[_DiffTarget]: + """Parse a unified diff into every path it would create/modify/move/remove. + + The diff is walked **per file section**, each delimited by its ``diff --git + a/ b/`` header. Because that header carries the source and the + destination for *every* change kind — including deletes (``+++ /dev/null``), + mode-only changes (no ``+++`` line at all), and ``copy to`` targets — reading + it closes the bypasses that a ``+++``-only scan misses. Within a section the + ``--- a/`` / ``+++ b/`` lines and the rename/copy ``from``/``to`` lines refine + the source/destination when present (a rename/copy ``to`` is authoritative for + the destination and links it to its source). BOTH the section source and + destination are emitted as targets, so removing or moving a file *away from* a + denied/out-of-scope path is flagged too. A path that fails + :func:`_canonicalize` is surfaced as an unsafe target. Pure header parsing — + it never executes the diff. + """ + targets: list[_DiffTarget] = [] + seen: set[tuple[str, str | None]] = set() + + # Per-section accumulators; flushed at each new ``diff --git`` and at EOF. + src_raw: str | None = None + dst_raw: str | None = None + move_from: str | None = None + + def flush() -> None: + nonlocal src_raw, dst_raw, move_from + if src_raw is None and dst_raw is None: + return + rf = _canonicalize(move_from) if move_from else None + # Source side: catches deletes and renames/copies away from a denied or + # out-of-scope path. Destination side: catches creates/modifies/mode + # changes/copy targets, linked to its source for a clear violation note. + _emit_target(targets, seen, src_raw) + _emit_target(targets, seen, dst_raw, rename_from=rf) + src_raw = dst_raw = move_from = None + + for line in diff.splitlines(): + git = _DIFF_GIT_RE.match(line) + if git: + flush() + src_raw, dst_raw = git.group(1), git.group(2) + continue + + for pattern in (_RENAME_TO_RE, _COPY_TO_RE): + m = pattern.match(line) + if m: + dst_raw = m.group(1) # authoritative destination for the section + break + else: + for pattern in (_RENAME_FROM_RE, _COPY_FROM_RE): + m = pattern.match(line) + if m: + move_from = m.group(1) + break + else: + minus = _MINUS_RE.match(line) + if minus and minus.group(1).strip() != _DEV_NULL: + src_raw = minus.group(1) + continue + plus = _PLUS_RE.match(line) + # Ignore hunk body lines starting with "+++"; a real header is + # "+++ b/path" or "+++ /dev/null". A /dev/null target means a + # delete, so the destination stays the diff --git path. + if plus and plus.group(1).strip() != _DEV_NULL: + dst_raw = plus.group(1) + + flush() + return targets + + +def _emit_target( + targets: list[_DiffTarget], + seen: set[tuple[str, str | None]], + raw: str | None, + *, + rename_from: str | None = None, +) -> None: + """Canonicalize ``raw`` and append it as a target (unsafe paths flagged). + + ``/dev/null`` and empty values are dropped (no real path). A path that fails + :func:`_canonicalize` (absolute / parent-escaping) is appended as an unsafe + target so the scan rejects it rather than letting an indirection bypass the + denylist. + """ + if raw is None: + return + stripped = raw.strip() + if not stripped or stripped == _DEV_NULL: + return + canon = _canonicalize(stripped) + if canon is None: + _append_unique(targets, seen, _DiffTarget(path=stripped), unsafe=True) + else: + _append_unique( + targets, + seen, + _DiffTarget(path=canon, rename_from=rename_from), + unsafe=False, + ) + + +# Marks a target whose path could not be safely canonicalized. Stored on the +# _DiffTarget via a parallel set keyed by identity is overkill; instead we use a +# reserved reason string the scanner recognizes. +_UNSAFE_PATH_SENTINEL = "\x00unsafe\x00" + + +def _append_unique( + targets: list[_DiffTarget], + seen: set[tuple[str, str | None]], + target: _DiffTarget, + *, + unsafe: bool, +) -> None: + """Append ``target`` if its (path, rename_from) pair is new; tag unsafe.""" + if unsafe: + # Tag the rename_from slot with the sentinel so the scanner can flag it + # without changing the public _DiffTarget shape. + target = _DiffTarget(path=target.path, rename_from=_UNSAFE_PATH_SENTINEL) + key = (target.path, target.rename_from) + if key in seen: + return + seen.add(key) + targets.append(target) + + +def _path_denied(path: str) -> str | None: + """Return the denylist reason if ``path`` is on the trust-control surface.""" + for pattern, reason in _DENY_PATTERNS: + if pattern.search(path): + return reason + return None + + +def _in_scope(path: str, scope: tuple[str, ...]) -> bool: + """Return ``True`` if ``path`` falls under one of the declared scope prefixes. + + Scope entries are canonicalized directory/file prefixes. A path is in scope + if it equals a scope entry or sits beneath a scope directory (prefix match + on a ``/`` boundary). An empty scope means "nothing is in scope", so every + path is rejected as out-of-scope — the design treats an undeclared scope as + a hard stop, not a wildcard (§3.3.2: "files outside the task's declared + scope"). + """ + for entry in scope: + if path == entry or path.startswith(entry + "/"): + return True + return False + + +def _normalize_scope(scope: Any) -> tuple[str, ...]: + """Canonicalize the declared scope into a tuple of safe path prefixes. + + Unsafe scope entries (absolute / parent-escaping) are dropped, so a + malformed scope can only *shrink* what is allowed, never widen it. + """ + if not scope: + return () + out: list[str] = [] + for entry in scope: + canon = _canonicalize(str(entry)) + if canon is not None and canon not in out: + out.append(canon) + return tuple(out) + + +def scan_trust_control_surface( + diff: str, *, scope: Any +) -> list[TrustBoundaryViolation]: + """Scan a candidate diff for trust-control-surface violations (§3.3.2 #2). + + Returns every violation found (empty list == clean). A target violates the + boundary if it is (a) an unsafe/uncanonicalizable path, (b) on the denylist + (``.github/workflows/**``, IAM/policy IaC, CODEOWNERS, branch-protection, + Dependabot), or (c) outside the task's declared ``scope`` — including a + rename whose *destination* is denied/out-of-scope, so a rename cannot + launder a forbidden path. The scan is pure header parsing; it never executes + the patch. + """ + declared_scope = _normalize_scope(scope) + violations: list[TrustBoundaryViolation] = [] + + for target in iter_diff_target_paths(diff): + rename_from = target.rename_from + is_unsafe = rename_from == _UNSAFE_PATH_SENTINEL + if is_unsafe: + rename_from = None + + if is_unsafe: + violations.append( + TrustBoundaryViolation( + path=target.path, + reason="unsafe path (absolute or escapes the repo root)", + rename_from=None, + ) + ) + continue + + denied_reason = _path_denied(target.path) + if denied_reason is not None: + violations.append( + TrustBoundaryViolation( + path=target.path, + reason=denied_reason, + rename_from=rename_from, + ) + ) + continue + + if not _in_scope(target.path, declared_scope): + violations.append( + TrustBoundaryViolation( + path=target.path, + reason="outside the task's declared scope", + rename_from=rename_from, + ) + ) + + return violations + + +# --------------------------------------------------------------------------- +# Build result + the node entrypoint +# --------------------------------------------------------------------------- + + +@dataclass +class _BuildOutcome: + """Internal result of :func:`build_candidate_diff` before state assembly.""" + + diff: str + diff_hash: str + violations: list[TrustBoundaryViolation] = field(default_factory=list) + + @property + def clean(self) -> bool: + return not self.violations + + +def build_candidate_diff( + plan: Mapping[str, Any], + *, + builder: DiffBuilder | None = None, + config: Mapping[str, Any] | None = None, +) -> _BuildOutcome: + """Synthesize a candidate diff, scan it, and hash it (§3.3.2 #2/#3). + + Calls the injected ``builder`` (default :func:`default_diff_builder`) to turn + the approved ``plan`` into a unified diff, runs the box-side trust-control + -surface scan against the plan's declared ``scope``, and computes the diff + integrity hash via the foundation's + :func:`agent_team.state_store.compute_content_hash`. The hash is always + computed (CI keys against it) but a non-empty violation list means the diff + must NOT auto-advance — the caller parks it for human + GPT cross-review. + + Raises :class:`BuildError` if the plan is not a mapping or the builder + returns an empty/whitespace-only diff (nothing to build). + """ + if not isinstance(plan, Mapping): + raise BuildError("approved plan must be a mapping") + + diff_builder = builder if builder is not None else default_diff_builder + diff = diff_builder(plan=plan, config=config) + + if not isinstance(diff, str) or not diff.strip(): + raise BuildError("diff builder produced an empty candidate diff") + + diff_hash = compute_content_hash(diff.encode("utf-8")) + violations = scan_trust_control_surface(diff, scope=plan.get("scope")) + return _BuildOutcome(diff=diff, diff_hash=diff_hash, violations=violations) + + +def _format_violations(violations: list[TrustBoundaryViolation]) -> str: + """Render violations into a single human-readable park reason.""" + lines = [] + for v in violations: + if v.rename_from: + lines.append(f"{v.rename_from} -> {v.path}: {v.reason}") + else: + lines.append(f"{v.path}: {v.reason}") + return "; ".join(lines) + + +def builders_node( + state: PipelineState, + *, + builder: DiffBuilder | None = None, + config: Mapping[str, Any] | None = None, +) -> dict[str, Any]: + """LangGraph builders node: approved plan -> candidate diff (§3.3, §3.3.2). + + Reads the approved ``plan`` from ``state``, produces a candidate diff via the + injected (or default) :class:`DiffBuilder`, runs the §3.3.2 box-side + trust-control-surface scan, and writes the result back as a **partial** + :class:`agent_team.task_model.PipelineState` update (the graph state is + ``total=False``): + + * **Clean diff** — writes ``candidate_diff`` + ``diff_hash``, sets + ``current_phase`` to ``VERIFY`` and ``status`` to ``ACTIVE`` so the org-CI + apply/verify stage runs next (the box never builds locally, D2/D11). + * **Violation(s)** — does NOT advance to verify. Records the diff + hash + (provenance, per §3.3.2: a denylist-touching diff is an ALARM), sets + ``current_phase`` to ``PARKED`` and ``status`` to ``PARKED``, and writes a + ``park_reason`` naming the violations so the coordinator escalates to + mandatory human review + GPT cross-review. It is never auto-built. + + The node never raises for a *policy* rejection (that is an expected outcome); + it raises :class:`BuildError` only when there is no usable plan/diff at all. + The graph-state enum-valued keys are written as their ``.value`` strings to + match the :class:`PipelineState` ``TypedDict`` (str-typed), mirroring + :func:`agent_team.task_model.task_to_dict`. The return type is a plain + ``dict`` (a structural superset of the partial ``PipelineState`` update) so + the park path can carry an extra ``park_reason`` annotation without + redefining the foundation ``TypedDict``. + """ + plan = state.get("plan") + if not plan: + raise BuildError("builders_node requires an approved plan in state") + + outcome = build_candidate_diff(plan, builder=builder, config=config) + + update: dict[str, Any] = { + "candidate_diff": outcome.diff, + "diff_hash": outcome.diff_hash, + } + + if outcome.clean: + update["current_phase"] = Phase.VERIFY.value + update["status"] = TaskStatus.ACTIVE.value + else: + update["current_phase"] = Phase.PARKED.value + update["status"] = TaskStatus.PARKED.value + update["park_reason"] = ( + "trust-control-surface violation (mandatory human + GPT " + f"cross-review): {_format_violations(outcome.violations)}" + ) + + return update diff --git a/agent-team/agent_team/nodes/clarifier.py b/agent-team/agent_team/nodes/clarifier.py new file mode 100644 index 0000000..569fba5 --- /dev/null +++ b/agent-team/agent_team/nodes/clarifier.py @@ -0,0 +1,228 @@ +"""Clarifier node — interrupt + 98% confidence loop (design §3.3, §7.1 P1). + +The clarifier is the first reasoning stage and the **human gate** of the +Plane-2 pipeline (§3.3). It gathers context, then asks Adam *question-sets +until it is 98%+ confident*. Each question-set is delivered through a LangGraph +``interrupt()``: the graph suspends and checkpoints, the question-set is posted +over the chosen transport (§3.3.1), and the task resumes via +``Command(resume=...)`` when Adam answers. No progression to planning happens +until the clarifier clears the confidence bar (§3.3, §7.1 P1). + +This module is a **leaf** built on the committed foundation contracts, which it +imports verbatim and never redefines: + +* :class:`agent_team.task_model.PipelineState` — the LangGraph state schema. +* :class:`agent_team.task_model.Phase` / :class:`agent_team.task_model.TaskStatus` + — lifecycle enums written back into the state. +* :func:`agent_team.task_model.new_thread_id` — thread-id minting (intake). +* :class:`agent_team.transport.QuestionSet` — the interrupt payload. + +The node owns only the *loop*: how confidence is assessed and what questions are +asked are injected as callables so this stays a pure, unit-testable control +flow with no live Claude call. The real wiring binds Claude through the +``agent_team.billing.claude_invoke`` seam in a later phase; here the seam is a +constructor argument so P1 can prove the suspend/resume mechanic +deterministically. + +Key design points proven here (the §7.1 P1 "riskiest mechanic"): + +* **98% loop.** The node calls :func:`langgraph.types.interrupt` repeatedly in a + ``while confidence < threshold`` loop. Each resume replays the node from the + top; LangGraph returns previously-supplied resume values for already-cleared + interrupts, so the accumulated Q&A drives confidence upward deterministically. +* **Turn cap (§7.1).** A clarifier is capped at ``max_turns`` per task; on the + cap it stops asking, marks the task ``PARKED``/``Phase.PARKED`` (ALARM rather + than spin), and does not advance to planning. +* **Human gate.** Only a run that clears the bar writes ``Phase.PLAN`` + + ``TaskStatus.ACTIVE``; nothing else lets the pipeline progress to build. +""" + +from __future__ import annotations + +import uuid +from collections.abc import Callable, Sequence +from dataclasses import dataclass + +from langgraph.types import interrupt + +from agent_team.task_model import Phase, PipelineState, TaskStatus +from agent_team.transport import QuestionSet + +__all__ = [ + "DEFAULT_CONFIDENCE_THRESHOLD", + "DEFAULT_MAX_TURNS", + "ClarifierConfig", + "ConfidenceAssessor", + "QuestionGenerator", + "build_question_set", + "make_clarifier_node", +] + +# §3.3 / §7.1 P1: the clarifier must reach "98%+ confident" before the human +# gate opens. Expressed as a 0..1 fraction; the loop runs while below it. +DEFAULT_CONFIDENCE_THRESHOLD: float = 0.98 + +# §7.1: "a clarifier is capped at N turns per task, then" parks rather than +# spinning. A conservative default; callers override per task class. +DEFAULT_MAX_TURNS: int = 6 + + +# A confidence assessor inspects the running Q&A history (oldest first) and the +# task state, and returns the current 0..1 confidence that the requirement is +# understood well enough to plan. Injected so the loop is testable without a +# live model; the real binding calls Claude through the billing seam. +ConfidenceAssessor = Callable[[Sequence[object], PipelineState], float] + +# A question generator produces the next ordered question-set given the Q&A so +# far and the task state. Injected for the same reason. +QuestionGenerator = Callable[[Sequence[object], PipelineState], list[str]] + + +@dataclass(frozen=True) +class ClarifierConfig: + """Tunables for :func:`make_clarifier_node` (§3.3, §7.1). + + ``confidence_threshold`` is the 98% bar the loop must clear; ``max_turns`` + is the §7.1 turn cap after which the task parks instead of spinning; + ``transport`` records the channel Adam chose at intake (carried into the + interrupt payload per §3.3.1) and falls back to the state's ``transport`` + when empty. + """ + + confidence_threshold: float = DEFAULT_CONFIDENCE_THRESHOLD + max_turns: int = DEFAULT_MAX_TURNS + transport: str = "" + + def __post_init__(self) -> None: + if not 0.0 < self.confidence_threshold <= 1.0: + raise ValueError( + "confidence_threshold must be in (0, 1]; " + f"got {self.confidence_threshold!r}" + ) + if self.max_turns < 1: + raise ValueError(f"max_turns must be >= 1; got {self.max_turns!r}") + + +def _new_question_id() -> str: + """Mint a fresh ``question_id`` (uuid4 hex) for one question-set (§3.3.1).""" + return uuid.uuid4().hex + + +def build_question_set( + *, + thread_id: str, + turn: int, + questions: list[str], + context: dict[str, object] | None = None, + question_id: str | None = None, +) -> QuestionSet: + """Build the :class:`QuestionSet` carried by one ``interrupt()`` (§3.3.1). + + Mints a ``question_id`` when not supplied. ``turn`` is the monotonic turn + index within the task; ``questions`` is the ordered prompt list; ``context`` + is optional rendering metadata (repo, summary) the transport adapter may + surface. The returned payload is exactly the foundation + :class:`~agent_team.transport.QuestionSet` contract — never a redefinition. + """ + return QuestionSet( + thread_id=thread_id, + question_id=question_id or _new_question_id(), + turn=turn, + questions=list(questions), + context=dict(context or {}), + ) + + +def _interrupt_payload( + question_set: QuestionSet, *, transport: str +) -> dict[str, object]: + """Serialize the interrupt payload (§3.3.1: ``{thread_id, question_id, ...}``). + + The §3.3.1 interrupt payload carries ``{thread_id, question_id, turn, + question_set, transport, deadline}``. ``deadline`` is owned by the durable + ledger/timer seam and is filled in by the responder at delivery time, so it + is left ``None`` here; the node's contribution is the question-set and its + identity. + """ + return { + "thread_id": question_set.thread_id, + "question_id": question_set.question_id, + "turn": question_set.turn, + "question_set": question_set, + "transport": transport, + "deadline": None, + } + + +def make_clarifier_node( + *, + assess_confidence: ConfidenceAssessor, + generate_questions: QuestionGenerator, + config: ClarifierConfig | None = None, +) -> Callable[[PipelineState], PipelineState]: + """Build the clarifier LangGraph node (§3.3, §7.1 P1). + + Returns a node callable ``node(state) -> state-delta`` suitable for + ``StateGraph(PipelineState).add_node("clarify", node)``. The node: + + 1. Starts from the task's existing ``qa_history`` (so a resumed run keeps + prior answers) and assesses confidence. + 2. While confidence is below the threshold **and** the turn cap is not hit, + generates the next question-set and raises a LangGraph + :func:`~langgraph.types.interrupt` carrying it. The graph suspends and + checkpoints; the resume value (Adam's answer) is appended to the running + Q&A history and confidence is re-assessed. Each resume replays the node + from the top, so the loop is durable across crashes and restarts + (§7.1 P1 "resume after restart"). + 3. On clearing the bar, advances the task to ``Phase.PLAN`` / + ``TaskStatus.ACTIVE`` — the **human gate opens** (§3.3). + 4. On hitting the turn cap first, parks the task (``Phase.PARKED`` / + ``TaskStatus.PARKED``) instead of spinning (§7.1); it never advances to + planning, so the pipeline still cannot build. + + The returned delta writes only the keys it owns (``qa_history``, + ``current_phase``, ``status``) — ``PipelineState`` is ``total=False``, so a + partial write is the intended per-node checkpoint transition. + """ + cfg = config or ClarifierConfig() + + def clarifier_node(state: PipelineState) -> PipelineState: + thread_id = state.get("thread_id", "") + transport = cfg.transport or state.get("transport", "") + qa_history: list[object] = list(state.get("qa_history", [])) + + # Confidence assessed from whatever Q&A already exists (a fresh task has + # none, so a context-only assessor may still clear or fall short). + confidence = assess_confidence(qa_history, state) + turn = 0 + + while confidence < cfg.confidence_threshold and turn < cfg.max_turns: + questions = generate_questions(qa_history, state) + question_set = build_question_set( + thread_id=thread_id, + turn=turn, + questions=questions, + ) + # Suspend + checkpoint; resume value is Adam's answer for this turn. + answer = interrupt(_interrupt_payload(question_set, transport=transport)) + qa_history.append(answer) + confidence = assess_confidence(qa_history, state) + turn += 1 + + if confidence >= cfg.confidence_threshold: + # Human gate clears: advance to planning. + next_phase = Phase.PLAN + next_status = TaskStatus.ACTIVE + else: + # Turn cap hit without clearing the bar: park + ALARM, never spin, + # never advance to build (§7.1). + next_phase = Phase.PARKED + next_status = TaskStatus.PARKED + + return { + "qa_history": qa_history, + "current_phase": next_phase.value, + "status": next_status.value, + } + + return clarifier_node diff --git a/agent-team/agent_team/nodes/planner.py b/agent-team/agent_team/nodes/planner.py new file mode 100644 index 0000000..820a1ca --- /dev/null +++ b/agent-team/agent_team/nodes/planner.py @@ -0,0 +1,293 @@ +"""Planner node — Plane-2 pipeline PLAN stage (design §3.3, §7.1 P2). + +The planner turns the clarified task context into a **phased plan** (the format +these design docs use) by calling Claude through the billing seam, then hands +the plan to the adversarial review loop. + +Pipeline position (§3.3):: + + INTAKE -> CLARIFIER -> [PLANNER] -> REVIEW LOOP -> BUILDERS -> VERIFIERS + +Responsibilities of this leaf (Phase P2, §7.1): + +* Read the clarified context out of :class:`~agent_team.task_model.PipelineState` + (the task description + the full clarifier Q&A history). +* On a **review loop-back** (§3.3 "Loops back to the planner on REQUEST + CHANGES"), fold the prior ``review_verdicts`` into the re-plan prompt so the + next plan answers the reviewer's objections. +* Call :func:`agent_team.billing.claude_invoke` to draft the plan, then parse + the model's reply into a structured ``plan`` dict. +* Return a **partial** ``PipelineState`` update advancing the phase to + :attr:`~agent_team.task_model.Phase.REVIEW`. +* Enforce the convergence bound (§3.3 "escalates to Adam if it cannot + converge"): after :data:`MAX_PLAN_REVISIONS` REQUEST-CHANGES loop-backs the + task is parked + ALARM-ed rather than spun (status ``PARKED``, phase + ``PARKED``). + +This module imports the committed foundation contracts verbatim +(:mod:`agent_team.task_model`, :mod:`agent_team.billing`) and does no I/O of its +own beyond the injected Claude seam — keeping the node a pure LangGraph +state-transition function that is trivially unit-testable. +""" + +from __future__ import annotations + +import json +import re +from typing import Any + +from agent_team.billing import ClaudeResult, claude_invoke +from agent_team.task_model import Phase, PipelineState, TaskStatus + +__all__ = [ + "MAX_PLAN_REVISIONS", + "PlannerError", + "build_plan_prompt", + "parse_plan", + "plan_node", +] + +# §3.3 / §6.6 — the planner must converge or park. After this many +# REQUEST-CHANGES loop-backs from the review stage the task is escalated to +# Adam (parked + ALARM) instead of looping forever. +MAX_PLAN_REVISIONS = 3 + +# The reviewer verdict string that sends a plan back to the planner. Kept here +# (rather than imported from a review node that does not exist yet) so this leaf +# stays self-contained; the review leaf will emit this same token. +_REQUEST_CHANGES = "REQUEST_CHANGES" + + +class PlannerError(Exception): + """Raised when the planner cannot produce a usable plan. + + Distinct from a *converged-but-rejected* plan (which loops back through + review) — this signals the planner itself failed (empty/garbled model + output), so the coordinator can fail the task rather than advance it. + """ + + +def _task_description(state: PipelineState) -> str: + """Pull the task description out of the graph state. + + Intake writes the originating ask; we look in the conventional places and + fall back to an empty string so a malformed state surfaces as an empty + prompt section rather than a ``KeyError`` inside the node. + """ + plan = state.get("plan") or {} + if isinstance(plan, dict): + desc = plan.get("task") or plan.get("description") + if isinstance(desc, str) and desc.strip(): + return desc.strip() + desc = state.get("task") # type: ignore[call-overload] + if isinstance(desc, str) and desc.strip(): + return desc.strip() + return "" + + +def _format_qa_history(qa_history: list[Any]) -> str: + """Render the clarifier Q&A history into prompt text. + + Each entry may be a ``{"question": ..., "answer": ...}`` mapping (the + clarifier's shape) or a plain string; both are handled so the planner does + not couple tightly to the clarifier's internal record format. + """ + lines: list[str] = [] + for idx, entry in enumerate(qa_history, start=1): + if isinstance(entry, dict): + question = str(entry.get("question", "")).strip() + answer = str(entry.get("answer", "")).strip() + if question or answer: + lines.append(f"{idx}. Q: {question}\n A: {answer}") + else: + text = str(entry).strip() + if text: + lines.append(f"{idx}. {text}") + return "\n".join(lines) + + +def _format_review_feedback(review_verdicts: list[Any]) -> str: + """Render prior review verdicts into re-plan guidance (§3.3 loop-back). + + Only the most recent verdict drives the re-plan, but earlier ones are + summarised so the planner does not re-introduce already-rejected ideas. + """ + if not review_verdicts: + return "" + lines: list[str] = [] + for idx, verdict in enumerate(review_verdicts, start=1): + if isinstance(verdict, dict): + decision = str(verdict.get("decision", "")).strip() + notes = str(verdict.get("notes") or verdict.get("comment") or "").strip() + lines.append(f"Review {idx} [{decision}]: {notes}".rstrip()) + else: + lines.append(f"Review {idx}: {str(verdict).strip()}") + return "\n".join(lines) + + +def build_plan_prompt(state: PipelineState) -> str: + """Build the Claude prompt that drafts (or re-drafts) the phased plan. + + Pure string assembly over the graph state — no I/O — so the prompt shape is + directly unit-testable. On a loop-back (``review_verdicts`` present) the + prompt instructs Claude to revise the prior plan against the feedback rather + than start from scratch. + """ + description = _task_description(state) + qa = _format_qa_history(list(state.get("qa_history", []))) + feedback = _format_review_feedback(list(state.get("review_verdicts", []))) + prior_plan = state.get("plan") + + sections = [ + "You are the PLANNER stage of an agentic SDLC pipeline. Produce a " + "phased implementation plan for the task below. The clarifier has " + "already reached confidence with the human, so do not ask questions — " + "plan.", + "", + "## Task", + description or "(no task description provided)", + ] + + if qa: + sections += ["", "## Clarified context (Q&A)", qa] + + if feedback: + # Loop-back: the review stage sent the prior plan back for changes. + sections += [ + "", + "## Reviewer feedback on the previous plan (address every point)", + feedback, + ] + if isinstance(prior_plan, dict) and prior_plan.get("phases"): + sections += [ + "", + "## Previous plan (revise; do not restart from scratch)", + json.dumps(prior_plan, sort_keys=True, indent=2), + ] + + sections += [ + "", + "## Output format", + "Return ONLY a JSON object with keys: " + '"summary" (string), "phases" (a non-empty list of objects each with ' + '"name" and "steps", where "steps" is a non-empty list of strings). ' + "Do not include prose outside the JSON.", + ] + return "\n".join(sections) + + +def _strip_code_fence(text: str) -> str: + """Strip a leading/trailing Markdown code fence if Claude wrapped the JSON.""" + fenced = re.match( + r"^\s*```(?:json)?\s*\n(?P.*?)\n?\s*```\s*$", + text, + flags=re.DOTALL | re.IGNORECASE, + ) + if fenced: + return fenced.group("body") + return text + + +def parse_plan(text: str) -> dict[str, Any]: + """Parse the model reply into a validated phased-plan dict. + + Raises :class:`PlannerError` if the reply is not a JSON object with a + non-empty ``phases`` list of well-formed phases — a garbled plan must fail + loudly so the coordinator does not advance an empty plan into review. + """ + candidate = _strip_code_fence(text or "").strip() + if not candidate: + raise PlannerError("planner returned an empty response") + try: + data = json.loads(candidate) + except json.JSONDecodeError as exc: + raise PlannerError(f"planner reply was not valid JSON: {exc}") from exc + + if not isinstance(data, dict): + raise PlannerError("planner reply JSON must be an object") + + phases = data.get("phases") + if not isinstance(phases, list) or not phases: + raise PlannerError("plan must contain a non-empty 'phases' list") + + normalized_phases: list[dict[str, Any]] = [] + for idx, phase in enumerate(phases, start=1): + if not isinstance(phase, dict): + raise PlannerError(f"phase {idx} must be an object") + name = phase.get("name") + steps = phase.get("steps") + if not isinstance(name, str) or not name.strip(): + raise PlannerError(f"phase {idx} is missing a non-empty 'name'") + if not isinstance(steps, list) or not steps: + raise PlannerError(f"phase {idx} ('{name}') has no steps") + normalized_steps = [str(step).strip() for step in steps if str(step).strip()] + if not normalized_steps: + raise PlannerError(f"phase {idx} ('{name}') has no non-empty steps") + normalized_phases.append({"name": name.strip(), "steps": normalized_steps}) + + summary = data.get("summary") + return { + "summary": str(summary).strip() if isinstance(summary, str) else "", + "phases": normalized_phases, + } + + +def _revision_count(state: PipelineState) -> int: + """How many REQUEST-CHANGES loop-backs have happened so far (§3.3). + + Counts the request-changes verdicts already recorded in the state; the + planner uses this to decide whether the next attempt is still within the + convergence bound. + """ + count = 0 + for verdict in state.get("review_verdicts", []): + decision = ( + verdict.get("decision") if isinstance(verdict, dict) else str(verdict) + ) + if isinstance(decision, str) and decision.strip().upper() == _REQUEST_CHANGES: + count += 1 + return count + + +def plan_node( + state: PipelineState, + config: dict[str, Any] | None = None, +) -> PipelineState: + """LangGraph node: draft/refine the phased plan, then advance to REVIEW. + + Returns a **partial** :class:`~agent_team.task_model.PipelineState` (the + keys this node owns) — LangGraph merges it into the checkpointed state. + + Behaviour: + + * **Convergence bound (§3.3, §6.6).** If the review stage has already sent + the plan back :data:`MAX_PLAN_REVISIONS` times, the planner does **not** + burn more Claude budget: it parks the task (status ``PARKED``, phase + ``PARKED``) so the coordinator ALARMs Adam. This is the "escalate to Adam + if it cannot converge" path. + * **Plan / re-plan.** Otherwise it builds the prompt (folding in any review + feedback on a loop-back), calls :func:`claude_invoke`, parses the reply, + and returns the new ``plan`` with the phase advanced to ``REVIEW``. + + ``config`` is forwarded to the billing seam so the caller can pin the + billing mode (it is threaded through to :func:`claude_invoke` as ``config``). + """ + revisions = _revision_count(state) + if revisions >= MAX_PLAN_REVISIONS: + # Escalation ladder (§5/§3.3): stop looping, hand to the human. + return PipelineState( + status=TaskStatus.PARKED.value, + current_phase=Phase.PARKED.value, + ) + + prompt = build_plan_prompt(state) + result: ClaudeResult = claude_invoke(prompt, config=config) + plan = parse_plan(result.text) + # Record how many times we have planned so review/observability can see it. + plan["revision"] = revisions + + return PipelineState( + plan=plan, + current_phase=Phase.REVIEW.value, + status=TaskStatus.ACTIVE.value, + ) diff --git a/agent-team/agent_team/nodes/review_loop.py b/agent-team/agent_team/nodes/review_loop.py new file mode 100644 index 0000000..daba0d7 --- /dev/null +++ b/agent-team/agent_team/nodes/review_loop.py @@ -0,0 +1,399 @@ +"""GPT-4.1 review-loop node (design §3.3, §7.1 Phase P2). + +The review loop is the adversarial plan-review stage of the Plane-2 pipeline:: + + INTAKE -> CLARIFIER -> PLANNER -> REVIEW LOOP -> BUILDERS -> ... + +After the planner produces a phased plan, this node runs an **adversarial +cross-family review** of that plan — the ``sh-plan-review`` / ``cross_reviewer`` +discipline (GPT-4.1, a different model family than the Claude planner, so it +catches different blind spots). The reviewer returns a verdict: + +* ``APPROVE`` — the plan clears the bar; the task advances to the builders. +* ``REQUEST_CHANGES`` — the plan has gaps; the task **loops back to the + planner** carrying the reviewer's findings, and the planner revises. + +The loop is bounded. After ``max_rounds`` of ``REQUEST_CHANGES`` without +convergence the node **escalates to Adam** (parks the task with an ALARM) +rather than spinning — design §3.3 stability bound (6): "a task that stalls ... +parks and ALARMs rather than spinning", and §7.1 P2 "including loop-back and +the escalate-to-Adam path". + +Design notes honoured here: + +* The reviewer is **GPT-4.1 via the orchestrator's** ``cross_reviewer`` agent, + reached through the local ``run.py`` (design §3.2: the R720 coordinator calls + the local ``~/orchestrator/run.py`` for non-Claude single-shot sub-tasks, so + they keep API billing + LangSmith tracing). It is therefore **not** routed + through :mod:`agent_team.billing` (which is the *Claude* seam). The actual + call is delegated to an injectable :data:`ReviewInvoker` so this node stays + dependency-free and unit-testable, mirroring the ``billing.set_invoker`` + pattern in the foundation. +* The node is a pure function over + :class:`~agent_team.task_model.PipelineState`: it returns only the partial + state keys it changes (``review_verdicts``, ``current_phase``, ``status``), + which the LangGraph reducer merges. It never writes to repos and performs no + durable I/O of its own — the SQLite checkpointer persists the merged state. +* :func:`route_after_review` is the LangGraph conditional-edge function that + reads the verdict this node recorded and returns the next node name + (``"build"`` / ``"plan"`` / ``"parked"``). +""" + +from __future__ import annotations + +import json +import os +import subprocess +from dataclasses import dataclass +from datetime import datetime, timezone +from enum import Enum +from typing import Any, Callable, Mapping + +from agent_team.task_model import Phase, PipelineState, TaskStatus + +__all__ = [ + "DEFAULT_MAX_REVIEW_ROUNDS", + "ReviewInvoker", + "ReviewOutcome", + "ReviewResult", + "ReviewVerdict", + "review_node", + "route_after_review", + "set_review_invoker", +] + +# Node names this stage routes to (LangGraph graph node ids). Kept as module +# constants so the conditional-edge mapping and the node body cannot drift. +BUILD_NODE = "build" +PLAN_NODE = "plan" +PARKED_NODE = "parked" + +# Default cap on adversarial review rounds before the loop escalates to Adam +# (design §3.3 stability bound 6; §7.1 P2 escalate-to-Adam path). Overridable +# per task via config["max_review_rounds"]. +DEFAULT_MAX_REVIEW_ROUNDS = 3 + +# Config / env key naming the per-task review-round cap. +_MAX_ROUNDS_CONFIG_KEY = "max_review_rounds" +_MAX_ROUNDS_ENV = "AGENT_TEAM_MAX_REVIEW_ROUNDS" + +# Config key naming the orchestrator entry point (the local run.py). Defaults +# to the rsync'd path on the R720; overridable for tests / non-default installs. +_RUN_PY_CONFIG_KEY = "orchestrator_run_py" +_DEFAULT_RUN_PY = os.path.expanduser("~/Documents/repositories/orchestrator/run.py") + +# Verdict tokens the reviewer output is scanned for. REQUEST_CHANGES wins on a +# tie so an ambiguous review fails closed (loops back / escalates) rather than +# advancing a plan the reviewer flagged. +_APPROVE_TOKENS = ("APPROVE", "APPROVED", "LGTM", "NO BLOCKERS", "NO BLOCKING") +_CHANGES_TOKENS = ( + "REQUEST CHANGES", + "REQUEST_CHANGES", + "REQUESTCHANGES", + "BLOCK", + "BLOCKING", + "NEEDS CHANGES", + "NEEDS WORK", +) + + +class ReviewVerdict(Enum): + """The adversarial reviewer's verdict on a plan (design §3.3).""" + + APPROVE = "approve" + REQUEST_CHANGES = "request_changes" + + +class ReviewOutcome(Enum): + """What the loop decided to do after recording a verdict. + + ``APPROVED`` — advance to the builders. ``LOOP_BACK`` — return to the + planner with the findings. ``ESCALATE`` — the round cap was hit without + convergence; park the task and ALARM Adam (design §3.3 bound 6, §7.1 P2). + """ + + APPROVED = "approved" + LOOP_BACK = "loop_back" + ESCALATE = "escalate" + + +@dataclass +class ReviewResult: + """One adversarial review pass, appended to ``review_verdicts`` (design §3.3). + + ``verdict`` is the parsed :class:`ReviewVerdict`; ``round_index`` is the + 1-based review round; ``outcome`` is the loop decision this pass produced; + ``findings`` is the reviewer's full text (the planner consumes it on a + loop-back); ``reviewer`` records the agent/model for the audit trail; + ``raw`` is the untouched invoker payload; ``created_at`` is an ISO-8601 UTC + timestamp. + """ + + verdict: ReviewVerdict + round_index: int + outcome: ReviewOutcome + findings: str + reviewer: str = "cross_reviewer" + raw: Any = None + created_at: str | None = None + + def to_dict(self) -> dict[str, Any]: + """Serialize to a JSON-safe dict for the ``review_verdicts`` ledger.""" + return { + "verdict": self.verdict.value, + "round_index": self.round_index, + "outcome": self.outcome.value, + "findings": self.findings, + "reviewer": self.reviewer, + "created_at": self.created_at, + } + + +# Pluggable reviewer call: signature (prompt, **kw) -> str (the reviewer's +# text). The default shells out to the orchestrator's cross_reviewer via the +# local run.py. Leaves/tests rebind it with set_review_invoker(). +ReviewInvoker = Callable[..., str] + + +def _orchestrator_invoker(prompt: str, *, run_py: str, **_kw: Any) -> str: + """Default reviewer: call the orchestrator's ``cross_reviewer`` (GPT-4.1). + + Invokes the local ``run.py`` with the review prompt. The orchestrator's + router sends adversarial-review tasks to ``cross_reviewer`` (GPT-4.1); this + keeps the review cross-family (a different model than the Claude planner) + and API-billed + LangSmith-traced per design §3.2. Returns the orchestrator's + stdout (the reviewer's verdict + findings). + """ + if not os.path.exists(run_py): + raise FileNotFoundError( + f"orchestrator entry point not found: {run_py}; set " + f"config[{_RUN_PY_CONFIG_KEY!r}] or rebind via set_review_invoker()." + ) + completed = subprocess.run( # noqa: S603 - args are not shell-interpolated + ["python3", run_py, prompt], + capture_output=True, + text=True, + check=False, + ) + if completed.returncode != 0: + raise RuntimeError( + "orchestrator review call failed " + f"(exit {completed.returncode}): {completed.stderr.strip()}" + ) + return completed.stdout + + +_review_invoker: ReviewInvoker = _orchestrator_invoker + + +def set_review_invoker(invoker: ReviewInvoker) -> None: + """Bind the function that performs the adversarial review call. + + Leaves call this once at startup with an implementation that returns the + reviewer's text for a prompt. Keeping the call injectable keeps this node + dependency-free and unit-testable (mirrors ``billing.set_invoker``). + """ + global _review_invoker + _review_invoker = invoker + + +def _utcnow_iso() -> str: + """Return the current UTC time as an ISO-8601 string.""" + return datetime.now(timezone.utc).isoformat() + + +def _resolve_max_rounds(config: Mapping[str, Any] | None) -> int: + """Resolve the review-round cap from config, then env, then the default. + + A non-positive or non-integer value is rejected so a misconfigured cap + cannot turn the bounded loop into an unbounded one. + """ + raw: Any = None + if config is not None: + raw = config.get(_MAX_ROUNDS_CONFIG_KEY) + if raw is None: + raw = os.environ.get(_MAX_ROUNDS_ENV) + if raw is None: + return DEFAULT_MAX_REVIEW_ROUNDS + try: + value = int(raw) + except (TypeError, ValueError) as exc: + raise ValueError( + f"invalid {_MAX_ROUNDS_CONFIG_KEY!r}={raw!r}; expected a positive int" + ) from exc + if value < 1: + raise ValueError(f"invalid {_MAX_ROUNDS_CONFIG_KEY!r}={value!r}; must be >= 1") + return value + + +def _resolve_run_py(config: Mapping[str, Any] | None) -> str: + """Resolve the orchestrator ``run.py`` path from config, env, or default.""" + if config is not None: + configured = config.get(_RUN_PY_CONFIG_KEY) + if configured: + return os.path.expanduser(str(configured)) + env = os.environ.get("AGENT_TEAM_ORCHESTRATOR_RUN_PY") + if env: + return os.path.expanduser(env) + return _DEFAULT_RUN_PY + + +def parse_verdict(text: str) -> ReviewVerdict: + """Parse a :class:`ReviewVerdict` from the reviewer's free text. + + Scans for explicit ``REQUEST CHANGES`` / ``BLOCK`` tokens and ``APPROVE`` / + ``LGTM`` tokens (case-insensitive). The result **fails closed**: if a + change-requesting token is present, or if neither token class is present + (an ambiguous / empty review), the verdict is ``REQUEST_CHANGES`` so an + unclear review never silently advances a plan to the builders. + """ + haystack = (text or "").upper() + has_changes = any(token in haystack for token in _CHANGES_TOKENS) + has_approve = any(token in haystack for token in _APPROVE_TOKENS) + if has_changes: + return ReviewVerdict.REQUEST_CHANGES + if has_approve: + return ReviewVerdict.APPROVE + # Ambiguous / empty review -> fail closed. + return ReviewVerdict.REQUEST_CHANGES + + +def build_review_prompt(state: PipelineState) -> str: + """Compose the adversarial-review prompt sent to the GPT-4.1 reviewer. + + Embeds the planner's phased plan plus any prior-round findings, and asks for + a structured verdict the node can parse. Kept deterministic so the prompt is + testable and the verdict tokens line up with :func:`parse_verdict`. + """ + plan = state.get("plan") or {} + plan_text = json.dumps(plan, indent=2, sort_keys=True) + + prior = state.get("review_verdicts") or [] + prior_text = "" + if prior: + last = prior[-1] + if isinstance(last, Mapping): + prior_text = ( + "\n\nThis plan was REVISED after prior review feedback. The " + "previous round's findings were:\n" + f"{last.get('findings', '')}\n" + "Confirm they are resolved and look for anything new." + ) + + return ( + "You are the adversarial plan reviewer (sh-plan-review / cross_reviewer " + "discipline). Audit the following phased plan for flawed assumptions, " + "missing phases, ordering hazards, and convention violations.\n\n" + "Respond with a verdict line that is exactly one of:\n" + " VERDICT: APPROVE\n" + " VERDICT: REQUEST CHANGES\n" + "followed by your findings. Default to REQUEST CHANGES if anything is " + "unclear or risky.\n\n" + f"PLAN:\n{plan_text}" + f"{prior_text}" + ) + + +def _review_round_index(state: PipelineState) -> int: + """Return the 1-based index of the review round about to run. + + Counts only prior *review* verdict entries already in ``review_verdicts`` + (entries this node appended), so a loop-back/re-entry increments correctly. + """ + prior = state.get("review_verdicts") or [] + return len(prior) + 1 + + +def review_node( + state: PipelineState, + config: Mapping[str, Any] | None = None, +) -> PipelineState: + """LangGraph node: run one adversarial review round on the current plan. + + Returns a **partial** :class:`PipelineState` the reducer merges: + + * ``review_verdicts`` — the prior verdicts **plus** this round's + :class:`ReviewResult` (as a dict). (Returned as the full list rather than + a single item so the node is reducer-agnostic — it works whether or not + ``review_verdicts`` has an append-reducer configured.) + * ``current_phase`` / ``status`` — set per the loop decision: + - ``APPROVE`` -> phase ``BUILD``, status ``ACTIVE`` (advance to builders). + - ``REQUEST_CHANGES`` under the round cap -> phase ``PLAN``, status + ``ACTIVE`` (loop back to the planner with the findings). + - ``REQUEST_CHANGES`` at/over the round cap -> phase ``PARKED``, status + ``PARKED`` (escalate to Adam / ALARM; design §3.3 bound 6, §7.1 P2). + * ``updated_at`` — refreshed ISO-8601 UTC timestamp. + + The plan is required; calling this node with no ``plan`` in state is a + programming error (the planner runs first) and raises ``ValueError``. + """ + if not state.get("plan"): + raise ValueError( + "review_node requires a plan in state; the planner stage must run " + "before the review loop (design §3.3 INTAKE->...->PLANNER->REVIEW)." + ) + + max_rounds = _resolve_max_rounds(config) + run_py = _resolve_run_py(config) + round_index = _review_round_index(state) + + prompt = build_review_prompt(state) + raw = _review_invoker(prompt, run_py=run_py, config=config) + text = raw if isinstance(raw, str) else str(raw) + verdict = parse_verdict(text) + + if verdict is ReviewVerdict.APPROVE: + outcome = ReviewOutcome.APPROVED + next_phase = Phase.BUILD + next_status = TaskStatus.ACTIVE + elif round_index >= max_rounds: + # Cap reached without convergence -> escalate to Adam (park + ALARM). + outcome = ReviewOutcome.ESCALATE + next_phase = Phase.PARKED + next_status = TaskStatus.PARKED + else: + outcome = ReviewOutcome.LOOP_BACK + next_phase = Phase.PLAN + next_status = TaskStatus.ACTIVE + + result = ReviewResult( + verdict=verdict, + round_index=round_index, + outcome=outcome, + findings=text, + raw=raw, + created_at=_utcnow_iso(), + ) + + prior = list(state.get("review_verdicts") or []) + prior.append(result.to_dict()) + + update: PipelineState = { + "review_verdicts": prior, + "current_phase": next_phase.value, + "status": next_status.value, + "updated_at": _utcnow_iso(), + } + return update + + +def route_after_review(state: PipelineState) -> str: + """LangGraph conditional-edge: next node after the review loop. + + Reads the outcome of the most recent verdict this node recorded and maps it + to a node id: ``APPROVED`` -> ``"build"``, ``LOOP_BACK`` -> ``"plan"``, + ``ESCALATE`` -> ``"parked"``. A state with no recorded verdict is a + programming error (route is called after :func:`review_node`) and parks + fail-closed rather than advancing. + """ + verdicts = state.get("review_verdicts") or [] + if not verdicts: + return PARKED_NODE + last = verdicts[-1] + outcome = last.get("outcome") if isinstance(last, Mapping) else None + if outcome == ReviewOutcome.APPROVED.value: + return BUILD_NODE + if outcome == ReviewOutcome.LOOP_BACK.value: + return PLAN_NODE + # ESCALATE or any unexpected value -> park fail-closed. + return PARKED_NODE diff --git a/agent-team/agent_team/nodes/verifier.py b/agent-team/agent_team/nodes/verifier.py new file mode 100644 index 0000000..6e59a75 --- /dev/null +++ b/agent-team/agent_team/nodes/verifier.py @@ -0,0 +1,253 @@ +"""Verifier node — the VERIFY-stage LangGraph node (design §3.3, §3.3.2 P3). + +The verifier reads the authenticated CI results for a task's candidate diff and +either advances it to a draft PR or loops it back to the builders. The single +load-bearing rule from §3.3.2 boundary #4 is enforced structurally here: + + **The verifier *agent* cannot declare success.** Pass/fail is owned by the + pure-code gate (:mod:`agent_team.ci_gate`) over the authenticated, + patch-independent CI conclusion. The LLM verifier only ever reads *failures* + to propose the next fix. + +So this node calls :func:`agent_team.ci_gate.evaluate_ci_gate` for the decision +and consults an (optional, injectable) LLM **only** on a FAIL/BLOCK to author a +fix hint for the builders. The LLM is never asked whether the task passed. + +State contract (mirrors :class:`agent_team.task_model.PipelineState`): + +* reads ``candidate_diff``, ``diff_hash`` (ledger hash), ``ci_results``; +* writes ``status``, ``current_phase``, ``review_verdicts`` (appends the gate + verdict), and ``ci_results`` (annotated with the gate decision for + provenance). + +Transitions (the §3.3 "Stability + autonomy bounds" — the verifier must pass or +the task loops/holds, never ships): + +* gate PASS -> :class:`~agent_team.task_model.Phase.DONE`, + :class:`~agent_team.task_model.TaskStatus.DONE` (verifier produced a draft PR). +* gate FAIL -> :class:`~agent_team.task_model.Phase.BUILD`, + :class:`~agent_team.task_model.TaskStatus.ACTIVE` (loop back to builders), + unless the per-task build-loop budget is spent, in which case it PARKs. +* gate BLOCK -> :class:`~agent_team.task_model.Phase.PARKED`, + :class:`~agent_team.task_model.TaskStatus.PARKED` — an ALARM-worthy trust + violation (denylist hit, hash mismatch, ambiguous conclusion) is never auto- + retried; it parks for human + GPT cross-review. + +This module imports the committed foundation contracts verbatim and redefines +none of them. The LLM seam is injectable (mirroring +:func:`agent_team.billing.set_invoker`) so the node is pure-function testable +with no SDK or network. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import Any, Callable, Mapping + +from agent_team.ci_gate import GateDecision, GateResult, evaluate_ci_gate +from agent_team.task_model import Phase, PipelineState, TaskStatus + +__all__ = [ + "DEFAULT_MAX_BUILD_LOOPS", + "VerifierConfig", + "set_fix_advisor", + "verifier_node", +] + +# Default cap on build<->verify loops before a still-failing task parks rather +# than spinning (§3.3 bound #6 "N failed build loops"). Configurable per task +# via :class:`VerifierConfig`. +DEFAULT_MAX_BUILD_LOOPS: int = 3 + + +@dataclass +class VerifierConfig: + """Per-invocation knobs for the verifier node (§3.3, §3.3.2). + + ``expected_run_id`` keys the gate to the exact CI run the verifier + dispatched for this diff (a stale/substituted run id is a BLOCK). + ``allowed_scope`` is the task's declared-scope path prefixes for the + denylist boundary. ``max_build_loops`` caps build<->verify retries before + the task parks. ``build_loops`` is the loops already consumed for this task + (the coordinator threads it through state). + """ + + expected_run_id: str + allowed_scope: list[str] | None = None + max_build_loops: int = DEFAULT_MAX_BUILD_LOOPS + build_loops: int = 0 + + +# Fix-advisor seam: signature (gate_result, state) -> str. Consulted ONLY on a +# gate FAIL/BLOCK to author a fix hint for the builders from the *failure* +# details. Never consulted on PASS — the gate, not the LLM, owns success. The +# default returns no hint so an un-wired environment degrades to "loop back with +# no extra guidance" rather than crashing. +FixAdvisor = Callable[[GateResult, Mapping[str, Any]], str] + + +def _null_advisor(gate_result: GateResult, state: Mapping[str, Any]) -> str: + return "" + + +_fix_advisor: FixAdvisor = _null_advisor + + +def set_fix_advisor(advisor: FixAdvisor) -> None: + """Bind the LLM fix-advisor consulted on a gate FAIL/BLOCK. + + Leaves call this once at startup with an implementation that reads the gate + failure reasons + the CI logs and authors a next-fix hint for the builders. + Keeping it injectable keeps this node dependency-free and unit-testable, and + structurally enforces that the advisor is only ever invoked on failure (this + module never calls it on a PASS). + """ + global _fix_advisor + _fix_advisor = advisor + + +def _utc_now_iso() -> str: + return datetime.now(timezone.utc).isoformat() + + +def _verdict( + gate_result: GateResult, + *, + next_phase: Phase, + build_loops: int, + fix_hint: str = "", +) -> dict[str, Any]: + """Build the review-verdict record appended to ``review_verdicts``.""" + return { + "stage": "verify", + "decision": gate_result.decision.value, + "reasons": list(gate_result.reasons), + "run_id": gate_result.run_id, + "diff_hash": gate_result.diff_hash, + "ci_conclusion": gate_result.ci_conclusion, + "next_phase": next_phase.value, + "build_loops": build_loops, + "fix_hint": fix_hint, + "at": _utc_now_iso(), + } + + +def verifier_node( + state: PipelineState, + config: VerifierConfig, +) -> PipelineState: + """Run the VERIFY stage: gate the CI result, decide the next phase. + + Pure function of ``state`` + ``config`` (the gate does no I/O; the optional + fix-advisor is the only outbound call, and only on failure). Returns a + **partial** :class:`PipelineState` update (``total=False``) carrying the new + ``status``/``current_phase``, an appended ``review_verdicts`` entry, and a + gate-annotated ``ci_results`` for provenance — the LangGraph reducer merges + it into the durable thread state. + + Decision flow (§3.3.2 boundary #4 + §3.3 autonomy bounds): + + 1. Call :func:`agent_team.ci_gate.evaluate_ci_gate` with the candidate diff, + the ledger hash (``diff_hash``), the authenticated ``ci_results``, the + ``expected_run_id``, and the task's ``allowed_scope``. + 2. PASS -> advance to DONE (draft PR). The advisor is NOT consulted. + 3. FAIL -> consult the fix-advisor for a hint, then loop back to BUILD — + unless ``build_loops`` has reached ``max_build_loops``, in which case + PARK (a task that keeps failing must hold, never ship: §3.3 bound #3/#6). + 4. BLOCK -> PARK + (advisor hint) for human + GPT cross-review. A trust + violation is never auto-retried. + """ + candidate_diff = state.get("candidate_diff") + ledger_hash = state.get("diff_hash") + ci_results = state.get("ci_results") + + if not isinstance(candidate_diff, str): + # No diff to verify is itself a refuse-to-proceed: park for a human + # rather than declaring anything. (A builder must have produced a diff + # before VERIFY runs.) + gate_result = GateResult( + decision=GateDecision.BLOCK, + reasons=["no candidate_diff present in state to verify"], + run_id=config.expected_run_id, + diff_hash=ledger_hash, + ci_conclusion=None, + ) + else: + gate_result = evaluate_ci_gate( + candidate_diff=candidate_diff, + ledger_hash=ledger_hash, + ci_result=ci_results, + expected_run_id=config.expected_run_id, + allowed_scope=config.allowed_scope, + ) + + annotated_ci = dict(ci_results) if isinstance(ci_results, Mapping) else {} + annotated_ci["gate_decision"] = gate_result.decision.value + annotated_ci["gate_reasons"] = list(gate_result.reasons) + + if gate_result.decision is GateDecision.PASS: + verdict = _verdict( + gate_result, next_phase=Phase.DONE, build_loops=config.build_loops + ) + return { + "status": TaskStatus.DONE.value, + "current_phase": Phase.DONE.value, + "review_verdicts": [verdict], + "ci_results": annotated_ci, + "updated_at": verdict["at"], + } + + # Every non-PASS path consults the fix-advisor on the *failure* details. + # This is the only place the LLM is touched, and it never decides success. + fix_hint = _fix_advisor(gate_result, state) + + if gate_result.decision is GateDecision.FAIL: + next_loops = config.build_loops + 1 + if next_loops >= config.max_build_loops: + # Exhausted the build-loop budget: hold rather than spin (§3.3 #6). + verdict = _verdict( + gate_result, + next_phase=Phase.PARKED, + build_loops=next_loops, + fix_hint=fix_hint, + ) + verdict["reasons"].append( + f"max build loops ({config.max_build_loops}) reached; parking" + ) + return { + "status": TaskStatus.PARKED.value, + "current_phase": Phase.PARKED.value, + "review_verdicts": [verdict], + "ci_results": annotated_ci, + "updated_at": verdict["at"], + } + verdict = _verdict( + gate_result, + next_phase=Phase.BUILD, + build_loops=next_loops, + fix_hint=fix_hint, + ) + return { + "status": TaskStatus.ACTIVE.value, + "current_phase": Phase.BUILD.value, + "review_verdicts": [verdict], + "ci_results": annotated_ci, + "updated_at": verdict["at"], + } + + # GateDecision.BLOCK — trust violation. Park for human + GPT cross-review; + # never auto-retried, never shipped. + verdict = _verdict( + gate_result, + next_phase=Phase.PARKED, + build_loops=config.build_loops, + fix_hint=fix_hint, + ) + return { + "status": TaskStatus.PARKED.value, + "current_phase": Phase.PARKED.value, + "review_verdicts": [verdict], + "ci_results": annotated_ci, + "updated_at": verdict["at"], + } diff --git a/agent-team/agent_team/operator_cli.py b/agent-team/agent_team/operator_cli.py new file mode 100644 index 0000000..37c1797 --- /dev/null +++ b/agent-team/agent_team/operator_cli.py @@ -0,0 +1,664 @@ +"""Audit-logged operator CLI over the ``pending_questions`` ledger (§3.3.1, §6.6). + +A stuck task *parks* rather than spins, and an operator needs a manual path to +unstick it without poking SQLite by hand. §3.3.1 specifies that path: "a small +CLI over the ledger lets an operator list ``open``/``parked`` questions, +re-deliver, force-expire, or answer on a task's behalf; ... Destructive CLI +actions (force-expire, answer-on-behalf, force-resume) are audit-logged and +require an explicit confirmation flag." §6.6 adds that an operator can +force-resume or re-prioritize a parked task via the CLI. + +This module is the leaf that implements that CLI. Two hard rules from the design +are encoded as code, not convention: + +1. **Every destructive action is audit-logged** — the *attempt* is recorded + before the mutation runs, and the *outcome* after, so a refused, failed, or + no-op action still leaves a trail. The audit log is an append-only JSONL file + written through the foundation's atomic ``state_store.atomic_write`` primitive + (write-temp → fsync → rename), created mode ``600`` because it records + operator identity and answer content. + +2. **Destructive actions require an explicit confirmation flag** — ``--confirm`` + on the CLI, ``confirm=True`` in the API. Without it the action raises + :class:`ConfirmationRequired` and performs no mutation (the refused attempt + is still audit-logged). Read-only actions (``list``) are never gated. + +The ledger lifecycle transitions reuse the committed compare-and-set helpers +from :mod:`agent_team.db.schema` verbatim (``expire_question``, +``answer_question``, ``supersede_question``); this module does NOT re-implement +the lifecycle SQL. ``force-resume`` records the operator's intent durably and +marks the question ``superseded`` per the turn-guard discipline (§3.3.1) so the +resume worker (a later phase) picks it up idempotently; it does not fabricate a +resume-queue table the foundation schema does not declare. +""" + +from __future__ import annotations + +import argparse +import getpass +import json +import sqlite3 +import sys +from collections.abc import Sequence +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from agent_team.db.schema import ( + answer_question, + connect, + expire_question, + supersede_question, +) +from agent_team.state_store import atomic_write + +__all__ = [ + "DESTRUCTIVE_ACTIONS", + "AuditEntry", + "AuditLog", + "ConfirmationRequired", + "OperatorCli", + "OperatorError", + "QuestionNotFound", + "build_parser", + "main", +] + +# The destructive verbs the design singles out (§3.3.1): each requires an +# explicit confirmation flag and is audit-logged. ``redeliver`` and ``list`` are +# deliberately NOT here — re-delivery is idempotent and read-only listing is +# harmless, so neither is gated, though ``redeliver`` is still audit-logged for +# provenance. +DESTRUCTIVE_ACTIONS: frozenset[str] = frozenset( + {"force-expire", "answer-on-behalf", "force-resume"} +) + +# The audit log records operator identity and answer payloads, so it is created +# with owner-only permissions (mode 600), matching the design's mode-600 report +# / state convention (§6.7). +_AUDIT_FILE_MODE: int = 0o600 + + +class OperatorError(Exception): + """Base class for operator-CLI errors.""" + + +class ConfirmationRequired(OperatorError): + """Raised when a destructive action is invoked without explicit confirmation. + + The §3.3.1 rule: force-expire, answer-on-behalf, and force-resume "require an + explicit confirmation flag". The refused attempt is still audit-logged before + this is raised, so an operator who forgets ``--confirm`` leaves a trail. + """ + + +class QuestionNotFound(OperatorError): + """Raised when an action targets a ``question_id`` absent from the ledger.""" + + +def _utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string (audit timestamps).""" + return datetime.now(timezone.utc).isoformat() + + +@dataclass(frozen=True) +class AuditEntry: + """One immutable audit-log record (§3.3.1 "audit-logged"). + + ``phase`` is ``"attempt"`` for the pre-mutation record and ``"outcome"`` for + the post-mutation record, so a refused or failed action still leaves the + attempt on disk. ``confirmed`` records whether the explicit confirmation flag + was present; ``detail`` carries action-specific context (target ids, answer + payload, rowcount result). Frozen so a constructed record cannot be mutated + before it is written. + """ + + timestamp: str + actor: str + action: str + phase: str + confirmed: bool + question_id: str | None = None + thread_id: str | None = None + detail: dict[str, Any] | None = None + + def to_dict(self) -> dict[str, Any]: + """Return a JSON-safe dict for serialization into the JSONL log.""" + return { + "timestamp": self.timestamp, + "actor": self.actor, + "action": self.action, + "phase": self.phase, + "confirmed": self.confirmed, + "question_id": self.question_id, + "thread_id": self.thread_id, + "detail": self.detail or {}, + } + + def to_json(self) -> str: + """Serialize this entry to a single-line JSON string (one JSONL row).""" + return json.dumps(self.to_dict(), sort_keys=True) + + +class AuditLog: + """Append-only, atomically-written JSONL audit trail for operator actions. + + Backed by the foundation's :func:`agent_team.state_store.atomic_write` + (write-temp → fsync → rename), so a crash mid-append leaves either the prior + log or the appended-to log, never a torn file. The file is created mode + ``600`` (operator identity + answer content are sensitive). The log is + append-only by contract: :meth:`append` reads the current bytes, adds one + line, and atomically rewrites — it never truncates or rewrites history. + """ + + def __init__(self, path: Path) -> None: + self._path = Path(path) + + @property + def path(self) -> Path: + """The on-disk path of the JSONL audit log.""" + return self._path + + def append(self, entry: AuditEntry) -> None: + """Atomically append ``entry`` as one JSONL line. + + Reads the existing log (empty if absent), appends the serialized entry + plus a trailing newline, and atomically rewrites the whole file via + :func:`atomic_write`. Then tightens the file mode to ``600`` so the + recreated file never widens its permissions. + """ + existing = b"" + if self._path.exists(): + existing = self._path.read_bytes() + line = (entry.to_json() + "\n").encode("utf-8") + atomic_write(self._path, existing + line) + # atomic_write recreates the file via a fresh temp; enforce 600 each time + # so the log never ends up group/world-readable. + self._path.chmod(_AUDIT_FILE_MODE) + + def read_all(self) -> list[dict[str, Any]]: + """Return all audit records as a list of dicts (oldest first).""" + if not self._path.exists(): + return [] + text = self._path.read_text(encoding="utf-8") + return [json.loads(line) for line in text.splitlines() if line.strip()] + + +@dataclass(frozen=True) +class _ActionResult: + """Internal carrier for an action's outcome detail + human-readable message.""" + + ok: bool + message: str + detail: dict[str, Any] + + +class OperatorCli: + """Programmatic operator interface over the ledger + the audit log. + + Each public method maps to one CLI verb. Destructive methods take a + ``confirm`` flag and raise :class:`ConfirmationRequired` when it is false — + after audit-logging the refused attempt. The class owns the SQLite + connection lifetime so a single process can issue several actions; a context + manager closes it. + """ + + def __init__( + self, + db_path: Path, + audit_log_path: Path, + *, + actor: str | None = None, + ) -> None: + self._conn = connect(Path(db_path)) + self._audit = AuditLog(Path(audit_log_path)) + # Default the actor to the OS login so an operator who omits --actor is + # still attributed in the trail. + self._actor = actor or _default_actor() + + # -- lifecycle -------------------------------------------------------- + + def close(self) -> None: + """Close the underlying SQLite connection.""" + self._conn.close() + + def __enter__(self) -> OperatorCli: + return self + + def __exit__(self, *exc: object) -> None: + self.close() + + @property + def audit_log(self) -> AuditLog: + """The :class:`AuditLog` this CLI writes to (exposed for inspection).""" + return self._audit + + # -- helpers ---------------------------------------------------------- + + def _audit_attempt( + self, + action: str, + *, + confirmed: bool, + question_id: str | None = None, + thread_id: str | None = None, + detail: dict[str, Any] | None = None, + ) -> None: + self._audit.append( + AuditEntry( + timestamp=_utc_now_iso(), + actor=self._actor, + action=action, + phase="attempt", + confirmed=confirmed, + question_id=question_id, + thread_id=thread_id, + detail=detail, + ) + ) + + def _audit_outcome( + self, + action: str, + *, + confirmed: bool, + question_id: str | None = None, + thread_id: str | None = None, + detail: dict[str, Any] | None = None, + ) -> None: + self._audit.append( + AuditEntry( + timestamp=_utc_now_iso(), + actor=self._actor, + action=action, + phase="outcome", + confirmed=confirmed, + question_id=question_id, + thread_id=thread_id, + detail=detail, + ) + ) + + def _require_confirm( + self, + action: str, + confirm: bool, + *, + question_id: str | None = None, + thread_id: str | None = None, + detail: dict[str, Any] | None = None, + ) -> None: + """Audit-log the attempt; raise if a destructive action is unconfirmed. + + Always records the attempt first (so a refused action is on disk), then + enforces the explicit-confirmation rule for destructive verbs. + """ + self._audit_attempt( + action, + confirmed=confirm, + question_id=question_id, + thread_id=thread_id, + detail=detail, + ) + if action in DESTRUCTIVE_ACTIONS and not confirm: + self._audit_outcome( + action, + confirmed=confirm, + question_id=question_id, + thread_id=thread_id, + detail={"refused": "confirmation flag required"}, + ) + raise ConfirmationRequired( + f"action {action!r} is destructive and requires explicit " + f"confirmation (pass --confirm / confirm=True)" + ) + + def _fetch_row(self, question_id: str) -> sqlite3.Row: + row = self._conn.execute( + "SELECT * FROM pending_questions WHERE question_id=?", + (question_id,), + ).fetchone() + if row is None: + raise QuestionNotFound(f"no ledger row for question_id={question_id!r}") + return row + + # -- read-only -------------------------------------------------------- + + def list_questions( + self, + *, + statuses: Sequence[str] | None = None, + thread_id: str | None = None, + ) -> list[dict[str, Any]]: + """List ledger rows, optionally filtered by status and/or thread. + + Read-only and ungated; defaults to the operator-relevant ``open`` set + (the questions awaiting action). Pass ``statuses`` to widen or narrow. + Returns plain dicts (oldest-posted first) so the result is JSON-safe. + """ + clauses: list[str] = [] + params: list[Any] = [] + if statuses: + placeholders = ",".join("?" for _ in statuses) + clauses.append(f"status IN ({placeholders})") + params.extend(statuses) + if thread_id is not None: + clauses.append("thread_id = ?") + params.append(thread_id) + where = f" WHERE {' AND '.join(clauses)}" if clauses else "" + sql = ( + "SELECT * FROM pending_questions" + + where + + " ORDER BY posted_at IS NULL, posted_at, question_id" + ) + rows = self._conn.execute(sql, tuple(params)).fetchall() + return [dict(row) for row in rows] + + # -- destructive ------------------------------------------------------ + + def redeliver(self, question_id: str) -> _ActionResult: + """Mark an ``open`` question for re-delivery (idempotent, audit-logged). + + Re-delivery itself is performed by the transport reconcile loop (§3.3.1); + this clears the stored ``channel_ref`` so the reconcile loop re-posts and + records a fresh ref. Not destructive (the question stays ``open``), so it + does not require confirmation, but it is audit-logged for provenance. + """ + self._require_confirm("redeliver", True, question_id=question_id) + row = self._fetch_row(question_id) + if row["status"] != "open": + result = _ActionResult( + ok=False, + message=f"question {question_id} is {row['status']}, not open; " + "nothing to re-deliver", + detail={"status": row["status"]}, + ) + else: + self._conn.execute( + "UPDATE pending_questions SET channel_ref=NULL WHERE question_id=?", + (question_id,), + ) + result = _ActionResult( + ok=True, + message=f"cleared channel_ref for {question_id}; reconcile loop " + "will re-post", + detail={"prior_channel_ref": row["channel_ref"]}, + ) + self._audit_outcome( + "redeliver", + confirmed=True, + question_id=question_id, + thread_id=row["thread_id"], + detail=result.detail, + ) + return result + + def force_expire(self, question_id: str, *, confirm: bool = False) -> _ActionResult: + """Force an ``open`` question to ``expired`` (destructive, confirmed). + + Delegates the lifecycle transition to the committed compare-and-set + :func:`agent_team.db.schema.expire_question`, so an answer racing in + still loses deterministically (§3.3.1). Returns ``ok`` reflecting the + compare-and-set rowcount. + """ + self._require_confirm("force-expire", confirm, question_id=question_id) + row = self._fetch_row(question_id) + changed = expire_question(self._conn, question_id=question_id) + result = _ActionResult( + ok=changed, + message=( + f"expired {question_id}" + if changed + else f"{question_id} was not open (status={row['status']}); no change" + ), + detail={"changed": changed, "prior_status": row["status"]}, + ) + self._audit_outcome( + "force-expire", + confirmed=confirm, + question_id=question_id, + thread_id=row["thread_id"], + detail=result.detail, + ) + return result + + def answer_on_behalf( + self, + question_id: str, + answer: Any, + *, + confirm: bool = False, + ) -> _ActionResult: + """Answer an ``open`` question on the task's behalf (destructive, confirmed). + + Routes through the committed first-answer-wins + :func:`agent_team.db.schema.answer_question`; ``answered_via`` is stamped + ``operator:`` so the audit trail and the ledger agree on who + answered. The answer is JSON-encoded into ``answer_json``. A late answer + (question already closed) loses the compare-and-set and returns + ``ok=False``. + """ + answer_json = json.dumps(answer, sort_keys=True) + via = f"operator:{self._actor}" + self._require_confirm( + "answer-on-behalf", + confirm, + question_id=question_id, + detail={"answer": answer, "via": via}, + ) + row = self._fetch_row(question_id) + changed = answer_question( + self._conn, + question_id=question_id, + answer_json=answer_json, + answered_via=via, + ) + result = _ActionResult( + ok=changed, + message=( + f"answered {question_id} on behalf of the task" + if changed + else f"{question_id} was not open (status={row['status']}); " + "answer ignored" + ), + detail={"changed": changed, "prior_status": row["status"], "via": via}, + ) + self._audit_outcome( + "answer-on-behalf", + confirmed=confirm, + question_id=question_id, + thread_id=row["thread_id"], + detail=result.detail, + ) + return result + + def force_resume(self, question_id: str, *, confirm: bool = False) -> _ActionResult: + """Force-resume a parked task's question (destructive, confirmed). + + §6.6: an operator can force-resume a parked task via the CLI. The resume + worker proper is a later phase; here we record the operator's intent + durably in the audit log and mark the stale question ``superseded`` via + the committed turn-guard helper + :func:`agent_team.db.schema.supersede_question` so a redelivered/stale + resume job is skipped idempotently (§3.3.1). The resume worker consumes + the audit intent on its next sweep; this method never invokes the + LangGraph runtime directly (out of scope for the foundation). + """ + self._require_confirm("force-resume", confirm, question_id=question_id) + row = self._fetch_row(question_id) + superseded = supersede_question(self._conn, question_id=question_id) + result = _ActionResult( + ok=True, + message=( + f"recorded force-resume intent for {question_id} " + f"(thread {row['thread_id']}); " + + ( + "superseded stale question" + if superseded + else "no open/answered question to supersede" + ) + ), + detail={ + "thread_id": row["thread_id"], + "turn": row["turn"], + "superseded": superseded, + "prior_status": row["status"], + "resume_requested": True, + }, + ) + self._audit_outcome( + "force-resume", + confirmed=confirm, + question_id=question_id, + thread_id=row["thread_id"], + detail=result.detail, + ) + return result + + +def _default_actor() -> str: + """Best-effort OS login name for audit attribution. + + Falls back to ``"unknown"`` rather than raising in the rare environment with + no resolvable login, so an action is never blocked purely on attribution + lookup (the attempt is still recorded, just as ``unknown``). + """ + try: + return getpass.getuser() + except Exception: # noqa: BLE001 - getuser can raise on odd environments + return "unknown" + + +def build_parser() -> argparse.ArgumentParser: + """Build the ``argparse`` parser for the operator CLI. + + Exposed separately so tests can parse argv without dispatching. The entry + CLI (``run-team.py operator ...``) and direct ``python -m`` both route here. + """ + parser = argparse.ArgumentParser( + prog="operator", + description="Audit-logged operator CLI over the pending_questions ledger " + "(R720 agent-team §3.3.1, §6.6).", + ) + parser.add_argument( + "--db", + required=True, + type=Path, + help="path to the agent-team SQLite database", + ) + parser.add_argument( + "--audit-log", + required=True, + type=Path, + help="path to the append-only JSONL audit log (created mode 600)", + ) + parser.add_argument( + "--actor", + default=None, + help="operator identity for the audit trail (defaults to OS login)", + ) + + sub = parser.add_subparsers(dest="command", required=True) + + p_list = sub.add_parser("list", help="list ledger questions (read-only)") + p_list.add_argument( + "--status", + action="append", + dest="statuses", + choices=("open", "answered", "expired", "superseded"), + help="filter by status (repeatable); default is open", + ) + p_list.add_argument("--thread", default=None, help="filter by thread_id") + + p_redeliver = sub.add_parser("redeliver", help="clear channel_ref to re-post") + p_redeliver.add_argument("question_id") + + p_expire = sub.add_parser( + "force-expire", help="force an open question to expired (destructive)" + ) + p_expire.add_argument("question_id") + p_expire.add_argument( + "--confirm", action="store_true", help="explicit confirmation (required)" + ) + + p_answer = sub.add_parser( + "answer-on-behalf", + help="answer an open question on the task's behalf (destructive)", + ) + p_answer.add_argument("question_id") + p_answer.add_argument( + "answer", help="answer payload; parsed as JSON if valid, else a raw string" + ) + p_answer.add_argument( + "--confirm", action="store_true", help="explicit confirmation (required)" + ) + + p_resume = sub.add_parser( + "force-resume", help="force-resume a parked task's question (destructive)" + ) + p_resume.add_argument("question_id") + p_resume.add_argument( + "--confirm", action="store_true", help="explicit confirmation (required)" + ) + + return parser + + +def _coerce_answer(raw: str) -> Any: + """Parse an answer arg as JSON, falling back to the raw string. + + Lets an operator pass ``'{"approve": true}'`` or just ``approve`` from the + shell; structured JSON round-trips, a bare token is kept verbatim. + """ + try: + return json.loads(raw) + except (ValueError, TypeError): + return raw + + +def main(argv: Sequence[str] | None = None) -> int: + """CLI entrypoint. Returns a process exit code (0 ok, non-zero on error). + + Dispatches the parsed subcommand against an :class:`OperatorCli`. Confirmation + refusals and missing questions are reported on stderr with a non-zero exit; + the attempt is already audit-logged before the error surfaces. + """ + parser = build_parser() + args = parser.parse_args(argv) + + try: + with OperatorCli(args.db, args.audit_log, actor=args.actor) as cli: + if args.command == "list": + rows = cli.list_questions( + statuses=args.statuses or ["open"], + thread_id=args.thread, + ) + print(json.dumps(rows, indent=2, sort_keys=True)) + return 0 + if args.command == "redeliver": + result = cli.redeliver(args.question_id) + elif args.command == "force-expire": + result = cli.force_expire(args.question_id, confirm=args.confirm) + elif args.command == "answer-on-behalf": + result = cli.answer_on_behalf( + args.question_id, + _coerce_answer(args.answer), + confirm=args.confirm, + ) + elif args.command == "force-resume": + result = cli.force_resume(args.question_id, confirm=args.confirm) + else: # pragma: no cover - argparse enforces the choices + parser.error(f"unknown command {args.command!r}") + print(result.message) + return 0 if result.ok else 1 + except ConfirmationRequired as exc: + print(f"refused: {exc}", file=sys.stderr) + return 2 + except QuestionNotFound as exc: + print(f"error: {exc}", file=sys.stderr) + return 3 + + +if __name__ == "__main__": # pragma: no cover - module-level entry shim + raise SystemExit(main()) diff --git a/agent-team/agent_team/recovery.py b/agent-team/agent_team/recovery.py new file mode 100644 index 0000000..16f9808 --- /dev/null +++ b/agent-team/agent_team/recovery.py @@ -0,0 +1,527 @@ +"""Restart-recovery sweep + post-restore reconciliation (design §3.3.1, §6.7). + +All durable human-in-the-loop state lives in two SQLite stores (the LangGraph +``SqliteSaver`` checkpoint and the ``pending_questions`` ledger), so a reboot or +a backup restore must *converge* rather than restart: nothing in the +interaction lifecycle is kept in memory. This module is the startup sweep that +drives that convergence. + +Per §3.3.1 ("Restart recovery"), on startup the sweep: + +* **redelivers** every ``open`` question whose row has no ``channel_ref`` (its + original post was lost / never landed): it re-posts over the question's + transport idempotently and records the returned ``channel_ref``; +* **re-enqueues a resume** for every ``answered`` question whose graph is still + interrupted on that turn (a crash between the compare-and-set and the resume + worker). The turn guard makes this idempotent — a graph that already advanced + is left alone, the question marked ``superseded``; +* **applies the deadline policy** to every overdue ``open`` question, flipping + it ``expired`` via the same first-answer-wins compare-and-set (an answer that + arrives for an already-expired question loses the race), then parks + ALARMs + or applies a default answer per the task policy. + +Per §6.7 ("post-restore reconciliation"), when the sweep runs after a *restore* +(not just a reboot) it first re-syncs against external state (in-flight CI runs, +current GitHub PR status) via an injected reconciler, so a restored backup +cannot resume a task on stale external assumptions. + +The module owns no I/O of its own beyond reading the ledger through the +foundation's :func:`agent_team.db.schema.connect` connection: every +side-effecting collaborator (transport delivery, resume enqueue, the live +checkpoint's interrupt/turn probe, the deadline policy, the external +reconciler) is injected as a callable. That keeps the sweep deterministic and +unit-testable, and means it provisions nothing. + +It imports the committed foundation contracts verbatim +(:mod:`agent_team.db.schema`, :mod:`agent_team.transport.base`) and does not +redefine them. +""" + +from __future__ import annotations + +import sqlite3 +from collections.abc import Callable +from dataclasses import dataclass, field +from datetime import datetime, timezone +from typing import Any, Protocol + +from agent_team.db.schema import ( + QUESTION_STATES, + connect, + expire_question, + supersede_question, +) +from agent_team.transport.base import QuestionSet, Transport + +__all__ = [ + "PendingQuestion", + "RecoveryReport", + "DeadlineOutcome", + "TransportResolver", + "ResumeEnqueuer", + "GraphInterruptProbe", + "DeadlinePolicy", + "ExternalReconciler", + "load_pending_questions", + "redeliver_open_questions", + "reenqueue_answered_resumes", + "apply_deadline_policy", + "run_restart_recovery", +] + + +# --------------------------------------------------------------------------- # +# Row view # +# --------------------------------------------------------------------------- # +@dataclass(frozen=True) +class PendingQuestion: + """A read-only view of one ``pending_questions`` ledger row (§3.3.1). + + Mirrors the foundation DDL columns. The sweep reads rows through this view + rather than passing raw :class:`sqlite3.Row` objects around, so the + collaborators receive a typed, immutable record. + """ + + question_id: str + thread_id: str + turn: int + status: str + transport: str + channel_ref: str | None + posted_at: str | None + deadline_at: str | None + answer_json: str | None + answered_at: str | None + answered_via: str | None + + @classmethod + def from_row(cls, row: sqlite3.Row) -> PendingQuestion: + """Build a :class:`PendingQuestion` from a ``pending_questions`` row.""" + return cls( + question_id=row["question_id"], + thread_id=row["thread_id"], + turn=int(row["turn"]), + status=row["status"], + transport=row["transport"], + channel_ref=row["channel_ref"], + posted_at=row["posted_at"], + deadline_at=row["deadline_at"], + answer_json=row["answer_json"], + answered_at=row["answered_at"], + answered_via=row["answered_via"], + ) + + +# --------------------------------------------------------------------------- # +# Injected-collaborator contracts # +# --------------------------------------------------------------------------- # +# A resolver hands the sweep the right Transport adapter for a question's +# ``transport`` string (e.g. "slack"/"github"/"claude_code"). Returning None +# means the transport is currently unreachable/unconfigured; the sweep records +# the redelivery as deferred rather than crashing. +TransportResolver = Callable[[str], Transport | None] + +# Enqueue a resume job for (thread_id, question_id, turn). The resume worker +# (§3.3.1) is single-flight per thread and turn-guarded; the sweep only needs +# to (idempotently) put the job on the queue. Returns True if a job was +# enqueued. +ResumeEnqueuer = Callable[[str, str, int], bool] + +# Probe the LIVE LangGraph checkpoint: is ``thread_id`` still interrupted on +# exactly ``turn``? Returns True only when the graph is genuinely still waiting +# on this question's turn. A False return means the graph already advanced +# (stale/redelivered) and the question must be superseded, never resumed. +GraphInterruptProbe = Callable[[str, int], bool] + +# The deadline policy for an expired question (§3.3.1): park + ALARM, or apply +# a defined default answer. Invoked only after the row is durably flipped to +# ``expired``. Returns the action it took for the report. +DeadlinePolicy = Callable[["PendingQuestion"], "DeadlineOutcome"] + + +class ExternalReconciler(Protocol): + """Post-restore external-state reconciliation seam (§6.7). + + After a *restore* (not a plain reboot) the sweep must re-sync against + external systems (in-flight CI runs, current GitHub PR status) before any + task resumes, so a restored backup cannot act on stale assumptions. A + concrete reconciler is injected in the leaves; the sweep only invokes it. + """ + + def reconcile(self, question: PendingQuestion) -> bool: + """Reconcile one task's external state. + + Return ``True`` when the task is safe to resume, ``False`` when external + state diverged (e.g. the CI run vanished or the PR was closed) and the + task must be held/parked instead of resumed. + """ + ... + + +# --------------------------------------------------------------------------- # +# Report types # +# --------------------------------------------------------------------------- # +@dataclass(frozen=True) +class DeadlineOutcome: + """What the deadline policy did with one expired question (§3.3.1).""" + + question_id: str + action: str # "parked" | "defaulted" | str describing the action taken + detail: str = "" + + +@dataclass +class RecoveryReport: + """Structured result of a restart-recovery sweep (§3.3.1, §6.7). + + Every list holds ``question_id`` values so the coordinator can ALARM / + report deterministically. ``errors`` carries ``(question_id, message)`` + pairs for collaborator failures that were isolated so one bad row cannot + abort the whole sweep. + """ + + redelivered: list[str] = field(default_factory=list) + redelivery_deferred: list[str] = field(default_factory=list) + resumes_enqueued: list[str] = field(default_factory=list) + superseded: list[str] = field(default_factory=list) + expired: list[str] = field(default_factory=list) + deadline_outcomes: list[DeadlineOutcome] = field(default_factory=list) + reconcile_held: list[str] = field(default_factory=list) + errors: list[tuple[str, str]] = field(default_factory=list) + + @property + def clean(self) -> bool: + """True when the sweep took no action and hit no errors. + + A clean sweep means durable state already matched reality (nothing to + redeliver, resume, expire, or hold) — the §3.3 "clean night posts + nothing" discipline applies to recovery too. + """ + return not any( + ( + self.redelivered, + self.redelivery_deferred, + self.resumes_enqueued, + self.superseded, + self.expired, + self.reconcile_held, + self.errors, + ) + ) + + +# --------------------------------------------------------------------------- # +# Ledger reads # +# --------------------------------------------------------------------------- # +def load_pending_questions( + conn: sqlite3.Connection, + *, + status: str | None = None, +) -> list[PendingQuestion]: + """Load ``pending_questions`` rows, optionally filtered by ``status``. + + Returned newest-posted-first within a stable secondary key so the sweep is + deterministic. ``status`` must be one of + :data:`agent_team.db.schema.QUESTION_STATES` when given. + """ + if status is not None and status not in QUESTION_STATES: + raise ValueError( + f"unknown status {status!r}; expected one of {QUESTION_STATES}" + ) + + sql = "SELECT * FROM pending_questions" + params: tuple[Any, ...] = () + if status is not None: + sql += " WHERE status = ?" + params = (status,) + # Order by question_id as a stable tiebreaker; posted_at may be NULL for a + # never-delivered open row, so it cannot be the sole sort key. + sql += " ORDER BY posted_at IS NULL, posted_at, question_id" + + rows = conn.execute(sql, params).fetchall() + return [PendingQuestion.from_row(row) for row in rows] + + +def _record_channel_ref( + conn: sqlite3.Connection, *, question_id: str, channel_ref: str +) -> None: + """Persist a freshly obtained ``channel_ref`` for an ``open`` question. + + Guarded on ``status='open'`` so a concurrent answer/expire that closed the + row in the meantime is not clobbered — the redelivery simply no-ops on the + ledger if the question is no longer open. + """ + conn.execute("BEGIN IMMEDIATE") + try: + conn.execute( + "UPDATE pending_questions SET channel_ref = ? " + "WHERE question_id = ? AND status = 'open'", + (channel_ref, question_id), + ) + conn.execute("COMMIT") + except BaseException: + conn.execute("ROLLBACK") + raise + + +# --------------------------------------------------------------------------- # +# Sweep step 1 — redeliver lost posts # +# --------------------------------------------------------------------------- # +def redeliver_open_questions( + conn: sqlite3.Connection, + *, + resolve_transport: TransportResolver, + report: RecoveryReport, +) -> None: + """Re-post every ``open`` question lacking a ``channel_ref`` (§3.3.1). + + On interrupt the responder writes the row ``open`` *before* posting, so a + crash (or a failed post) can leave an ``open`` row with no ref. This step + retries delivery idempotently: it resolves the question's transport, + re-posts the question-set, and records the returned ``channel_ref``. If the + transport is unreachable the redelivery is *deferred* (not an error) so a + later sweep / reconcile loop retries — matching the §3.3.1 "reconcile loop + retries idempotently" behaviour. + + Rows that already have a ``channel_ref`` are skipped: their post landed. + """ + for question in load_pending_questions(conn, status="open"): + if question.channel_ref: + continue # post already landed; nothing to redeliver. + + transport = resolve_transport(question.transport) + if transport is None: + report.redelivery_deferred.append(question.question_id) + continue + + try: + channel_ref = transport.post_question( + thread_id=question.thread_id, + question_id=question.question_id, + turn=question.turn, + question_set=_question_set_for(question), + deadline=question.deadline_at or "", + ) + except Exception as exc: # isolate one bad transport call. + report.errors.append((question.question_id, f"redeliver: {exc}")) + report.redelivery_deferred.append(question.question_id) + continue + + if not channel_ref: + # Transport returned no locator; treat as a deferred retry. + report.redelivery_deferred.append(question.question_id) + continue + + _record_channel_ref( + conn, question_id=question.question_id, channel_ref=channel_ref + ) + report.redelivered.append(question.question_id) + + +def _question_set_for(question: PendingQuestion) -> QuestionSet: + """Build a minimal :class:`QuestionSet` for a redelivery. + + The original prompt text lives in the LangGraph checkpoint, not this + ledger; on a redelivery the sweep carries identity (``thread_id`` / + ``question_id`` / ``turn``) so the adapter can re-render from the + checkpoint. ``questions`` is left empty here and the adapter fills it from + graph state, keeping the ledger free of duplicated prompt text. + """ + return QuestionSet( + thread_id=question.thread_id, + question_id=question.question_id, + turn=question.turn, + questions=[], + ) + + +# --------------------------------------------------------------------------- # +# Sweep step 2 — re-enqueue resumes for already-answered questions # +# --------------------------------------------------------------------------- # +def reenqueue_answered_resumes( + conn: sqlite3.Connection, + *, + is_interrupted_on_turn: GraphInterruptProbe, + enqueue_resume: ResumeEnqueuer, + report: RecoveryReport, + reconciler: ExternalReconciler | None = None, +) -> None: + """Re-enqueue resumes for ``answered`` rows whose graph still waits (§3.3.1). + + A crash can land between the first-answer-wins compare-and-set (row flipped + ``answered``) and the resume worker actually resuming the graph. On startup + every ``answered`` question is checked against the *live* checkpoint: + + * still interrupted on this exact turn -> (optionally reconcile external + state per §6.7, then) re-enqueue a resume. The resume worker is itself + turn-guarded and single-flight, so re-enqueuing is idempotent — a + duplicate job no-ops. + * the graph already advanced past this turn -> this is a stale/redelivered + answer; mark the question ``superseded`` (foundation compare-and-set) and + do NOT resume. "A resume can never double-apply." + + When a ``reconciler`` is supplied (a restore, not a plain reboot) and it + reports external state diverged, the resume is held rather than enqueued so + the task does not act on stale CI/PR assumptions (§6.7). + """ + for question in load_pending_questions(conn, status="answered"): + try: + still_waiting = is_interrupted_on_turn(question.thread_id, question.turn) + except Exception as exc: # isolate a bad probe. + report.errors.append((question.question_id, f"probe: {exc}")) + continue + + if not still_waiting: + # Graph already advanced: the resume already applied (or the turn + # moved on). Supersede so the row can never re-trigger a resume. + if supersede_question(conn, question_id=question.question_id): + report.superseded.append(question.question_id) + continue + + if reconciler is not None: + try: + safe = reconciler.reconcile(question) + except Exception as exc: # isolate a bad reconciler. + report.errors.append((question.question_id, f"reconcile: {exc}")) + report.reconcile_held.append(question.question_id) + continue + if not safe: + report.reconcile_held.append(question.question_id) + continue + + try: + enqueued = enqueue_resume( + question.thread_id, question.question_id, question.turn + ) + except Exception as exc: # isolate a bad enqueue. + report.errors.append((question.question_id, f"resume: {exc}")) + continue + + if enqueued: + report.resumes_enqueued.append(question.question_id) + + +# --------------------------------------------------------------------------- # +# Sweep step 3 — deadline policy for overdue open questions # +# --------------------------------------------------------------------------- # +def apply_deadline_policy( + conn: sqlite3.Connection, + *, + policy: DeadlinePolicy, + report: RecoveryReport, + now: datetime | None = None, +) -> None: + """Expire overdue ``open`` questions and apply the task policy (§3.3.1). + + For each ``open`` question whose ``deadline_at`` is at/before ``now``, the + sweep flips it ``expired`` via :func:`agent_team.db.schema.expire_question` + — the same first-answer-wins compare-and-set the live timer uses, so an + answer racing the same deadline either wins (row already ``answered``, this + call no-ops) or loses (row flipped ``expired``, a late answer is later + ignored). Only after a row is *durably* flipped does the ``policy`` run + (park + ALARM, or apply a default answer), so a crash between the two leaves + an ``expired`` row a later sweep re-processes — the policy must be + idempotent. + + Rows with no ``deadline_at`` never expire here (no deadline configured). + ``now`` defaults to the current UTC time; it is injectable for tests. + """ + current = now or datetime.now(timezone.utc) + + for question in load_pending_questions(conn, status="open"): + if not _is_overdue(question.deadline_at, current): + continue + + try: + flipped = expire_question(conn, question_id=question.question_id) + except Exception as exc: # isolate a bad compare-and-set. + report.errors.append((question.question_id, f"expire: {exc}")) + continue + + if not flipped: + # Lost the race: the question was answered/closed concurrently. + continue + + report.expired.append(question.question_id) + try: + outcome = policy(question) + except Exception as exc: # isolate a bad policy callback. + report.errors.append((question.question_id, f"policy: {exc}")) + continue + report.deadline_outcomes.append(outcome) + + +def _is_overdue(deadline_at: str | None, now: datetime) -> bool: + """Return True when ``deadline_at`` (ISO-8601) is at/before ``now``. + + A missing or unparseable deadline is treated as "not overdue": the sweep + never expires a question whose deadline it cannot read, it leaves it ``open`` + for an operator. ``now`` is timezone-aware (UTC); a naive ``deadline_at`` is + assumed UTC for comparison. + """ + if not deadline_at: + return False + try: + parsed = datetime.fromisoformat(deadline_at) + except ValueError: + return False + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=timezone.utc) + return parsed <= now + + +# --------------------------------------------------------------------------- # +# Orchestrating sweep # +# --------------------------------------------------------------------------- # +def run_restart_recovery( + db_path: Any, + *, + resolve_transport: TransportResolver, + is_interrupted_on_turn: GraphInterruptProbe, + enqueue_resume: ResumeEnqueuer, + deadline_policy: DeadlinePolicy, + reconciler: ExternalReconciler | None = None, + now: datetime | None = None, + conn: sqlite3.Connection | None = None, +) -> RecoveryReport: + """Run the full restart-recovery sweep against the ledger (§3.3.1, §6.7). + + Convergence order matters and is fixed: + + 1. **deadline first** — expire overdue ``open`` rows before anything else so + a question past its deadline is never redelivered or resumed as if live; + 2. **redeliver** ``open`` rows still lacking a ``channel_ref`` (their post + was lost); + 3. **re-enqueue resumes** for ``answered`` rows whose graph still waits, + superseding those whose graph advanced. + + When ``reconciler`` is supplied the sweep is treated as a *post-restore* + reconciliation (§6.7): step 3 first re-syncs each task's external state and + holds (does not resume) any task whose external state diverged. + + A connection is opened via :func:`agent_team.db.schema.connect` unless one + is injected (tests share an in-memory/temp DB). Every per-row collaborator + failure is isolated into ``report.errors`` so one bad row cannot abort the + sweep — the coordinator decides whether the error set warrants an ALARM. + """ + owns_conn = conn is None + connection = conn if conn is not None else connect(db_path) + report = RecoveryReport() + try: + apply_deadline_policy( + connection, policy=deadline_policy, report=report, now=now + ) + redeliver_open_questions( + connection, resolve_transport=resolve_transport, report=report + ) + reenqueue_answered_resumes( + connection, + is_interrupted_on_turn=is_interrupted_on_turn, + enqueue_resume=enqueue_resume, + report=report, + reconciler=reconciler, + ) + finally: + if owns_conn: + connection.close() + return report diff --git a/agent-team/agent_team/responder.py b/agent-team/agent_team/responder.py new file mode 100644 index 0000000..96333ac --- /dev/null +++ b/agent-team/agent_team/responder.py @@ -0,0 +1,444 @@ +"""Transport-agnostic notify + resume responder (design §3.3, §3.3.1). + +This is the Plane-2 leaf that owns the human-in-the-loop **notify+resume seam**. +It sits on top of the committed FOUNDATION contracts and wires them together; it +re-declares none of them: + +* :mod:`agent_team.transport.base` — the :class:`~agent_team.transport.base.Transport` + ABC plus :class:`~agent_team.transport.base.QuestionSet` / + :class:`~agent_team.transport.base.NormalizedAnswer` payloads. +* :mod:`agent_team.db.schema` — the ``pending_questions`` ledger and the + ``BEGIN IMMEDIATE`` first-answer-wins compare-and-set helpers + (:func:`~agent_team.db.schema.answer_question`, + :func:`~agent_team.db.schema.expire_question`, + :func:`~agent_team.db.schema.supersede_question`). + +Three seams, all transport-independent (§3.3.1): + +* **Notify (and lost-post).** :func:`notify_question` writes the ledger row + ``open`` FIRST, then posts to the chosen transport and stores its + ``channel_ref``. If the post raises, the row stays ``open`` with no ref so the + reconcile loop can retry idempotently — the durable ledger is the source of + truth, never an in-memory-only post. +* **Answer (first-answer-wins).** :func:`submit_answer` normalizes the inbound + raw payload via the transport, then runs the single atomic compare-and-set + ``UPDATE ... SET status='answered' ... WHERE question_id=? AND status='open'``. + rowcount 1 = first valid answer → enqueue a resume job; rowcount 0 = duplicate + / late / already-closed → ignored with an "already closed" reply. This one + compare-and-set makes duplicate clicks, transport redelivery, answers via two + channels, and answer-after-timeout all safe. +* **Resume (single-flight, turn-guarded).** :class:`ResumeWorker` serializes per + ``thread_id`` and, before resuming, checks the live checkpoint is still + interrupted on this ``turn``. If the graph already advanced (stale or + redelivered job) it marks the question ``superseded`` and skips, so a resume + can never double-apply. The graph is injected as a small structural protocol + so this leaf stays free of a hard LangGraph dependency for pre-deployment + scaffolding. + +The deadline timer (overdue ``open`` → ``expired``) and the startup recovery +sweep round out the lifecycle and reuse the same compare-and-set helpers. + +No network, no SDK, no provisioning here: the transport, the graph, and the +resume-job queue are all injected, so this module is pure orchestration over the +durable contracts and is unit-testable in isolation. +""" + +from __future__ import annotations + +import json +import sqlite3 +import threading +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import Any, Callable, Protocol, runtime_checkable + +from agent_team.db.schema import ( + answer_question, + expire_question, + supersede_question, +) +from agent_team.transport.base import QuestionSet, Transport + +__all__ = [ + "AnswerOutcome", + "GraphHandle", + "ResumeJob", + "ResumeWorker", + "deadline_sweep", + "notify_question", + "recover_open_questions", + "submit_answer", +] + + +# --------------------------------------------------------------------------- +# Injected collaborators (kept as structural protocols so this leaf has no hard +# LangGraph / queue dependency for pre-deployment scaffolding). +# --------------------------------------------------------------------------- + + +@runtime_checkable +class GraphHandle(Protocol): + """The slice of the LangGraph graph the responder needs (§3.3.1). + + The real object is a compiled LangGraph graph backed by the SQLite + checkpointer. The responder only needs to (a) ask which ``turn`` a thread is + currently interrupted on and (b) resume it with an answer, so it depends on + this narrow structural protocol rather than importing LangGraph. + """ + + def interrupted_turn(self, thread_id: str) -> int | None: + """Return the ``turn`` the thread is interrupted on, or ``None``. + + ``None`` means the thread is not currently suspended on an + ``interrupt()`` (it already advanced, completed, or never existed). The + turn guard compares this against the question's ``turn``. + """ + ... + + def resume(self, thread_id: str, answer: Any) -> Any: + """Resume the thread with ``answer`` (``Command(resume=answer)``). + + Drives the graph forward from its checkpointed interrupt. Returns + whatever the graph yields next; the responder does not interpret it. + """ + ... + + +@dataclass(frozen=True) +class ResumeJob: + """A unit of resume work enqueued after a first-answer-wins compare-and-set. + + Carries exactly what the turn-guarded :class:`ResumeWorker` needs: + ``thread_id`` (the per-thread single-flight key), ``question_id`` (the + ledger row to supersede if stale), ``turn`` (the turn guard), and the + ``answer`` to feed into ``Command(resume=...)``. + """ + + thread_id: str + question_id: str + turn: int + answer: Any + + +# A resume-job enqueue callback. The real queue is durable (re-enqueued on the +# startup sweep, §3.3.1); the responder only needs to hand a job to it. +EnqueueResume = Callable[[ResumeJob], None] + + +@dataclass(frozen=True) +class AnswerOutcome: + """Result of :func:`submit_answer` (§3.3.1). + + ``accepted`` is the rowcount-1 first-answer-wins verdict: ``True`` means this + call recorded the first valid answer and a resume job was enqueued; + ``False`` means the question was not ``open`` (already answered / expired / + superseded), so the answer was a duplicate or late and was ignored. + ``question_id`` / ``via`` echo the normalized inbound answer for the audit + trail. ``job`` is the enqueued :class:`ResumeJob` iff ``accepted``. + """ + + accepted: bool + question_id: str + via: str + job: ResumeJob | None = None + + +# --------------------------------------------------------------------------- +# Notify (delivery + lost-post) — §3.3.1. +# --------------------------------------------------------------------------- + + +def notify_question( + conn: sqlite3.Connection, + transport: Transport, + question_set: QuestionSet, + *, + deadline: str, + posted_at: str | None = None, +) -> str | None: + """Deliver a question-set: write the ledger row ``open`` first, then post. + + Implements the §3.3.1 "delivery (and lost-post)" rule precisely: + + 1. Insert the ``pending_questions`` row as ``open`` with no ``channel_ref``. + The durable ledger is written BEFORE the side-effecting post so a crash + or a failed post never loses the question — recovery sees an ``open`` row + lacking a ref and retries. + 2. Post the question-set over ``transport`` (which MUST embed the + ``question_id`` so an inbound answer maps back) and capture the returned + ``channel_ref``. + 3. Persist the ``channel_ref`` (and ``posted_at``) on the row. + + Returns the ``channel_ref`` on success, or ``None`` if the post failed (the + row stays ``open`` with no ref for the reconcile loop). The transport + exception is intentionally swallowed: a lost post is a recoverable state in + this design, not a hard error. + + The row is inserted with the question's identity (``thread_id``, ``turn``, + ``transport`` name) so the turn guard and reconcile can act on it. + """ + stamp = posted_at or _utc_now_iso() + transport_name = type(transport).__name__ + + # 1. Durable ledger row first (open, no ref). + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, posted_at, deadline_at) " + "VALUES (?, ?, ?, 'open', ?, ?, ?)", + ( + question_set.question_id, + question_set.thread_id, + question_set.turn, + transport_name, + stamp, + deadline, + ), + ) + + # 2. Side-effecting post. A failure here is recoverable (row stays open, + # no ref) — do NOT let it bubble up and lose the durable row. + try: + channel_ref = transport.post_question( + thread_id=question_set.thread_id, + question_id=question_set.question_id, + turn=question_set.turn, + question_set=question_set, + deadline=deadline, + ) + except Exception: + return None + + # 3. Persist the ref so reconcile/recovery can act on the post. + conn.execute( + "UPDATE pending_questions SET channel_ref=? WHERE question_id=?", + (channel_ref, question_set.question_id), + ) + return channel_ref + + +# --------------------------------------------------------------------------- +# Answer (first-answer-wins) — §3.3.1. +# --------------------------------------------------------------------------- + + +def submit_answer( + conn: sqlite3.Connection, + transport: Transport, + raw: Any, + *, + enqueue_resume: EnqueueResume, + answered_at: str | None = None, +) -> AnswerOutcome: + """Normalize an inbound answer and run the first-answer-wins compare-and-set. + + The transport adapter normalizes ``raw`` to ``(question_id, answer, via)``; + the responder then runs the single atomic statement (via + :func:`agent_team.db.schema.answer_question`, which wraps it in + ``BEGIN IMMEDIATE``): + + UPDATE pending_questions SET status='answered', answer_json=?, + answered_via=?, answered_at=? WHERE question_id=? AND status='open' + + * rowcount 1 → this is the first valid answer: look up the row's ``turn`` + and enqueue a :class:`ResumeJob`; return ``accepted=True``. + * rowcount 0 → the question was not ``open`` (already answered / expired / + superseded): the answer is a duplicate or late and is ignored; return + ``accepted=False`` so the caller can send an "already closed" reply. + + This single compare-and-set is what makes duplicate clicks, transport + redelivery, answers-via-two-channels, and answer-after-timeout all safe — + only one caller can ever flip ``open`` → ``answered``. + """ + question_id, answer, via = transport.parse_answer(raw) + accepted = answer_question( + conn, + question_id=question_id, + answer_json=json.dumps(answer, sort_keys=True), + answered_via=via, + answered_at=answered_at, + ) + if not accepted: + # Duplicate / late / already-closed: ignore (caller replies "closed"). + return AnswerOutcome(accepted=False, question_id=question_id, via=via) + + # First valid answer: enqueue the turn-guarded resume job. + turn = _question_turn(conn, question_id) + job = ResumeJob( + thread_id=_question_thread(conn, question_id), + question_id=question_id, + turn=turn, + answer=answer, + ) + enqueue_resume(job) + return AnswerOutcome(accepted=True, question_id=question_id, via=via, job=job) + + +# --------------------------------------------------------------------------- +# Resume (single-flight, turn-guarded) — §3.3.1. +# --------------------------------------------------------------------------- + + +class ResumeWorker: + """Serializes resume work per ``thread_id`` and guards on the turn (§3.3.1). + + A resume job for a thread runs under a per-thread lock so two jobs for the + same thread can never resume concurrently (single-flight). Before resuming, + the worker checks the live checkpoint is still interrupted on the job's + ``turn``: + + * if the graph advanced past this turn (stale or redelivered job) the + question is marked ``superseded`` and the resume is skipped — a resume can + never double-apply; + * otherwise it calls ``graph.resume(thread_id, answer)`` (the injected + ``Command(resume=answer)`` seam). + + Different threads resume concurrently within the budget cap; only same-thread + work is serialized. The per-thread locks are created lazily under a single + registry lock so this is safe to share across resume threads. + """ + + def __init__(self, conn: sqlite3.Connection, graph: GraphHandle) -> None: + self._conn = conn + self._graph = graph + self._registry_lock = threading.Lock() + self._thread_locks: dict[str, threading.Lock] = {} + + def _lock_for(self, thread_id: str) -> threading.Lock: + """Return (creating if needed) the single-flight lock for ``thread_id``.""" + with self._registry_lock: + lock = self._thread_locks.get(thread_id) + if lock is None: + lock = threading.Lock() + self._thread_locks[thread_id] = lock + return lock + + def run(self, job: ResumeJob) -> bool: + """Process one resume ``job`` under the per-thread single-flight lock. + + Returns ``True`` if the graph was resumed, ``False`` if the job was a + no-op (graph already advanced past the turn → question superseded and + skipped). Idempotent: re-running a job for an already-advanced thread + supersedes-and-skips rather than double-applying. + """ + with self._lock_for(job.thread_id): + live_turn = self._graph.interrupted_turn(job.thread_id) + if live_turn != job.turn: + # Stale / redelivered: the graph already advanced past this turn + # (or is not interrupted at all). Supersede and skip — never + # double-apply a resume. + supersede_question(self._conn, question_id=job.question_id) + return False + + self._graph.resume(job.thread_id, job.answer) + return True + + +# --------------------------------------------------------------------------- +# Deadline / no-answer timer — §3.3.1. +# --------------------------------------------------------------------------- + + +def deadline_sweep( + conn: sqlite3.Connection, + *, + now: str | None = None, +) -> list[str]: + """Flip overdue ``open`` questions to ``expired`` (deterministic race). + + Selects ``open`` rows whose ``deadline_at`` is non-null and ``<= now`` and + runs the same compare-and-set (:func:`agent_team.db.schema.expire_question`) + on each. Expiry vs answer is a deterministic race on flipping ``open``: an + answer arriving for an already-``expired`` question loses the compare-and-set + and is ignored, and vice versa. Returns the ids actually expired by this + call (rowcount 1), so the caller can apply the task policy (park + ALARM, or + a defined default answer) to exactly those. + """ + moment = now or _utc_now_iso() + rows = conn.execute( + "SELECT question_id FROM pending_questions " + "WHERE status='open' AND deadline_at IS NOT NULL AND deadline_at <= ?", + (moment,), + ).fetchall() + + expired: list[str] = [] + for row in rows: + qid = row["question_id"] + if expire_question(conn, question_id=qid): + expired.append(qid) + return expired + + +# --------------------------------------------------------------------------- +# Restart recovery — §3.3.1. +# --------------------------------------------------------------------------- + + +def recover_open_questions( + conn: sqlite3.Connection, + *, + enqueue_resume: EnqueueResume, +) -> list[ResumeJob]: + """Startup sweep: re-enqueue resume jobs for already-``answered`` questions. + + All state is durable, so a reboot converges by replaying the ledger. This + routine handles the ``answered`` slice of that sweep (§3.3.1): for every row + still ``answered`` (i.e. answered before a crash but not yet resumed), it + re-enqueues a :class:`ResumeJob`. The re-enqueue is safe because the + :class:`ResumeWorker` turn guard makes resumes idempotent — a job whose graph + already advanced supersedes-and-skips. + + Delivery-retry for ``open`` rows lacking a ref and the deadline policy for + overdue ``open`` rows are the reconcile loop's and + :func:`deadline_sweep`'s jobs respectively; this function owns only the + answered→resume replay. Returns the jobs it enqueued. + """ + rows = conn.execute( + "SELECT question_id, thread_id, turn, answer_json " + "FROM pending_questions WHERE status='answered'" + ).fetchall() + + jobs: list[ResumeJob] = [] + for row in rows: + answer = json.loads(row["answer_json"]) if row["answer_json"] else None + job = ResumeJob( + thread_id=row["thread_id"], + question_id=row["question_id"], + turn=int(row["turn"]), + answer=answer, + ) + enqueue_resume(job) + jobs.append(job) + return jobs + + +# --------------------------------------------------------------------------- +# Internal helpers. +# --------------------------------------------------------------------------- + + +def _question_turn(conn: sqlite3.Connection, question_id: str) -> int: + """Return the ``turn`` recorded for ``question_id``.""" + row = conn.execute( + "SELECT turn FROM pending_questions WHERE question_id=?", + (question_id,), + ).fetchone() + if row is None: + raise KeyError(f"unknown question_id: {question_id!r}") + return int(row["turn"]) + + +def _question_thread(conn: sqlite3.Connection, question_id: str) -> str: + """Return the ``thread_id`` recorded for ``question_id``.""" + row = conn.execute( + "SELECT thread_id FROM pending_questions WHERE question_id=?", + (question_id,), + ).fetchone() + if row is None: + raise KeyError(f"unknown question_id: {question_id!r}") + return str(row["thread_id"]) + + +def _utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string.""" + return datetime.now(timezone.utc).isoformat() diff --git a/agent-team/agent_team/resume_worker.py b/agent-team/agent_team/resume_worker.py new file mode 100644 index 0000000..1d45451 --- /dev/null +++ b/agent-team/agent_team/resume_worker.py @@ -0,0 +1,335 @@ +"""Single-flight, turn-guarded resume worker (design §3.3.1). + +When a human answer lands and wins the first-answer-wins compare-and-set +(:func:`agent_team.db.schema.answer_question`), a resume job is enqueued to +drive the suspended LangGraph thread forward. This module owns that resume +mechanic, and the design pins three guarantees on it: + +* **Single-flight per ``thread_id``.** Two resume jobs for the same task never + run concurrently. A task is its own thread; different threads resume + concurrently, but one thread is serialized so a redelivered/duplicated job + cannot race itself. +* **Turn-guarded.** Before resuming, the worker reads the *live* checkpoint and + confirms the graph is still interrupted on the answer's ``turn``. If the + graph already advanced (a stale or redelivered job, or a resume that already + applied), the worker marks the question ``superseded`` and skips. This is the + mechanism by which **a resume can never double-apply** (§3.3.1). +* **Restart-recoverable.** A startup sweep re-enqueues a resume for every + ``answered`` ledger row whose graph is still interrupted on that turn; the + turn guard makes re-enqueue idempotent, so converging after a reboot cannot + double-apply either. + +The worker depends only on a small structural :class:`GraphLike` protocol +(``get_state`` + ``invoke``), satisfied by a compiled LangGraph app, so the +durable resume logic stays decoupled from any specific checkpointer and remains +unit-testable without a live graph. The ledger reads/writes go through the +committed :mod:`agent_team.db.schema` helpers (imported verbatim); this module +adds no SQL of its own. +""" + +from __future__ import annotations + +import sqlite3 +import threading +from dataclasses import dataclass +from enum import Enum +from typing import Any, Protocol, runtime_checkable + +from agent_team.db.schema import supersede_question + +__all__ = [ + "GraphLike", + "ResumeOutcome", + "ResumeResult", + "ResumeWorker", + "build_resume_command", + "snapshot_interrupt_turns", +] + + +# --------------------------------------------------------------------------- # +# Graph seam +# --------------------------------------------------------------------------- # + + +@runtime_checkable +class GraphLike(Protocol): + """Structural protocol for the compiled LangGraph app the worker drives. + + A real ``langgraph`` compiled graph satisfies this: ``get_state`` returns a + ``StateSnapshot`` (with ``.next`` and ``.interrupts``) and ``invoke`` + accepts a ``Command(resume=...)`` plus the thread config. Depending on the + structural protocol rather than the concrete class keeps the resume logic + decoupled from the checkpointer and trivially testable (§3.3.1). + """ + + def get_state(self, config: dict[str, Any]) -> Any: + """Return the live :class:`StateSnapshot` for ``config``'s thread.""" + ... + + def invoke(self, input: Any, config: dict[str, Any]) -> Any: + """Resume/run the graph for ``config``'s thread with ``input``.""" + ... + + +def build_resume_command(answer: Any) -> Any: + """Build the LangGraph ``Command(resume=answer)`` resume input. + + ``Command`` is imported lazily so this module imports even where + ``langgraph`` is absent (the durable ledger logic does not need it). When + ``langgraph`` is installed, the real ``Command`` is used so the worker + drives an actual compiled graph; otherwise a clear :class:`RuntimeError` + is raised at call time. + """ + try: + from langgraph.types import Command + except ImportError as exc: # pragma: no cover - environment-dependent + raise RuntimeError( + "langgraph is required to resume a graph; install langgraph or " + "inject a graph whose invoke() accepts a plain resume payload" + ) from exc + return Command(resume=answer) + + +def snapshot_interrupt_turns(snapshot: Any) -> set[int]: + """Extract the set of ``turn`` values the snapshot is interrupted on. + + The §3.3.1 interrupt payload is ``{thread_id, question_id, turn, ...}``. + This reads each pending ``Interrupt.value`` and collects its ``turn``. A + snapshot that is not interrupted (``snapshot.interrupts`` empty) yields an + empty set, which the turn guard treats as "graph already advanced". + + Tolerant of either a mapping payload (``value['turn']``) or an object + payload (``value.turn``); anything without a readable integer ``turn`` is + ignored rather than crashing the worker. + """ + turns: set[int] = set() + interrupts = getattr(snapshot, "interrupts", None) or () + for item in interrupts: + value = getattr(item, "value", item) + turn: Any = None + if isinstance(value, dict): + turn = value.get("turn") + else: + turn = getattr(value, "turn", None) + if isinstance(turn, bool): # bool is an int subclass; not a real turn + continue + if isinstance(turn, int): + turns.add(turn) + return turns + + +def _snapshot_is_interrupted(snapshot: Any) -> bool: + """True if the snapshot is suspended on an interrupt (``next`` non-empty). + + LangGraph reports a pending interrupt via a non-empty ``next`` tuple and a + populated ``interrupts`` tuple. We treat either signal as "still + interrupted"; the turn check then narrows it to *this* turn. + """ + if getattr(snapshot, "interrupts", None): + return True + nxt = getattr(snapshot, "next", None) + return bool(nxt) + + +# --------------------------------------------------------------------------- # +# Result types +# --------------------------------------------------------------------------- # + + +class ResumeOutcome(Enum): + """Outcome of a single :meth:`ResumeWorker.resume` attempt (§3.3.1).""" + + #: The graph was interrupted on this turn; ``Command(resume=...)`` applied. + RESUMED = "resumed" + #: The graph already advanced past this turn; question marked superseded, + #: resume skipped. This is the no-double-apply guard firing. + SUPERSEDED = "superseded" + #: The graph already advanced but the question was no longer open/answered, + #: so there was nothing to supersede; resume skipped. + STALE = "stale" + + +@dataclass +class ResumeResult: + """Structured result of a resume attempt. + + ``outcome`` is the :class:`ResumeOutcome`; ``thread_id`` / ``question_id`` / + ``turn`` echo the job; ``graph_result`` carries the graph's return value + when (and only when) the resume actually applied. + """ + + outcome: ResumeOutcome + thread_id: str + question_id: str + turn: int + graph_result: Any = None + + @property + def resumed(self) -> bool: + """True iff the resume applied (``Command(resume=...)`` was invoked).""" + return self.outcome is ResumeOutcome.RESUMED + + +# --------------------------------------------------------------------------- # +# Worker +# --------------------------------------------------------------------------- # + + +def _thread_config(thread_id: str) -> dict[str, Any]: + """The LangGraph config addressing a single durable thread.""" + return {"configurable": {"thread_id": thread_id}} + + +class ResumeWorker: + """Serializes and turn-guards graph resumes (§3.3.1 "single-flight"). + + One worker drives many threads; it holds a per-``thread_id`` lock registry + so resumes for the *same* task are serialized (single-flight) while + different tasks resume concurrently. ``graph`` is any :class:`GraphLike` + (a compiled LangGraph app in production); ``conn`` is the agent-team SQLite + connection (see :func:`agent_team.db.schema.connect`) used to read pending + rows and to mark a stale question ``superseded`` via the committed + compare-and-set helper. + + The worker performs no SQL of its own: lifecycle writes go through + :func:`agent_team.db.schema.supersede_question`. It does not itself flip a + question to ``answered`` — that is the responder's first-answer-wins + compare-and-set, which gates whether a resume job is enqueued at all. + """ + + def __init__(self, graph: GraphLike, conn: sqlite3.Connection) -> None: + self._graph = graph + self._conn = conn + # Registry of per-thread locks. Guarded by _registry_lock so two + # threads minting the lock for the same thread_id get the *same* lock. + self._locks: dict[str, threading.Lock] = {} + self._registry_lock = threading.Lock() + + def _lock_for(self, thread_id: str) -> threading.Lock: + """Return the (memoized) serialization lock for ``thread_id``.""" + with self._registry_lock: + lock = self._locks.get(thread_id) + if lock is None: + lock = threading.Lock() + self._locks[thread_id] = lock + return lock + + def resume( + self, + *, + thread_id: str, + question_id: str, + turn: int, + answer: Any, + ) -> ResumeResult: + """Resume ``thread_id`` with ``answer``, single-flight + turn-guarded. + + Acquires the per-thread lock so this thread's resumes are serialized, + then reads the live checkpoint. If the graph is still interrupted on + ``turn`` it invokes ``Command(resume=answer)`` and returns + :attr:`ResumeOutcome.RESUMED`. If the graph has already advanced past + ``turn`` (stale/redelivered job, or a resume that already applied), it + marks ``question_id`` ``superseded`` and skips — returning + :attr:`ResumeOutcome.SUPERSEDED` (the row was open/answered) or + :attr:`ResumeOutcome.STALE` (nothing left to supersede). A resume can + therefore never double-apply (§3.3.1). + """ + lock = self._lock_for(thread_id) + with lock: + return self._resume_locked( + thread_id=thread_id, + question_id=question_id, + turn=turn, + answer=answer, + ) + + def _resume_locked( + self, + *, + thread_id: str, + question_id: str, + turn: int, + answer: Any, + ) -> ResumeResult: + """Resume body that runs while holding this thread's lock.""" + config = _thread_config(thread_id) + snapshot = self._graph.get_state(config) + + interrupted_here = _snapshot_is_interrupted( + snapshot + ) and turn in snapshot_interrupt_turns(snapshot) + + if not interrupted_here: + # Graph already advanced past this turn: the turn guard. Mark the + # question superseded so it can never enqueue another resume, and + # skip. supersede_question is the committed atomic compare-and-set; + # rowcount 1 => we superseded it, 0 => it was already terminal. + superseded = supersede_question(self._conn, question_id=question_id) + outcome = ResumeOutcome.SUPERSEDED if superseded else ResumeOutcome.STALE + return ResumeResult( + outcome=outcome, + thread_id=thread_id, + question_id=question_id, + turn=turn, + ) + + graph_result = self._graph.invoke(build_resume_command(answer), config) + return ResumeResult( + outcome=ResumeOutcome.RESUMED, + thread_id=thread_id, + question_id=question_id, + turn=turn, + graph_result=graph_result, + ) + + def recover_pending_resumes(self) -> list[ResumeResult]: + """Restart sweep: re-enqueue resumes for durable ``answered`` rows. + + On reboot, in-memory resume jobs are gone but the ledger is durable + (§3.3.1 "restart recovery"). This re-drives every ``answered`` question + whose graph is still interrupted on its turn. The per-thread turn guard + makes this idempotent: a thread that already advanced is superseded and + skipped, so converging after a restart cannot double-apply. + + Returns one :class:`ResumeResult` per processed row (whatever its + outcome) so a caller can log/ALARM. Rows are processed oldest-first by + ``answered_at`` to preserve answer ordering across a recovery. + """ + rows = self._conn.execute( + "SELECT question_id, thread_id, turn, answer_json " + "FROM pending_questions " + "WHERE status = 'answered' " + "ORDER BY answered_at IS NULL, answered_at ASC" + ).fetchall() + + results: list[ResumeResult] = [] + for row in rows: + answer = _decode_answer(row["answer_json"]) + results.append( + self.resume( + thread_id=row["thread_id"], + question_id=row["question_id"], + turn=int(row["turn"]), + answer=answer, + ) + ) + return results + + +def _decode_answer(answer_json: str | None) -> Any: + """Decode a ledger ``answer_json`` payload back to a Python value. + + The responder stores answers as a JSON string in ``answer_json``. A + non-JSON or ``NULL`` value is returned as-is (``None`` for ``NULL``), so a + malformed row does not crash the recovery sweep — the turn guard still + governs whether anything is applied. + """ + if answer_json is None: + return None + import json + + try: + return json.loads(answer_json) + except (ValueError, TypeError): + return answer_json diff --git a/agent-team/agent_team/state_store.py b/agent-team/agent_team/state_store.py new file mode 100644 index 0000000..55fb849 --- /dev/null +++ b/agent-team/agent_team/state_store.py @@ -0,0 +1,176 @@ +"""Atomic state-store utilities and integrity checking (design §6.7). + +All durable state on the R720 (LangGraph SQLite checkpoint, the +``pending_questions`` ledger, the budget ledger, the Plane-1 rotation/coverage +pointer) is written atomically (write-temp-then-fsync-then-rename) and +integrity-checked on load. "Integrity-checked" is concrete here: a +schema-version match plus a stored content hash. On any mismatch the loader +refuses to proceed silently and raises :class:`IntegrityError` so the +coordinator can park the affected task with an ALARM rather than acting on +corrupt state. + +This module is pure stdlib (``os``, ``tempfile``, ``hashlib``, ``pathlib``) +and depends on no other ``agent_team`` module. The leaf builders import these +signatures verbatim, so they are intentionally explicit and final. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import tempfile +from pathlib import Path + +__all__ = [ + "IntegrityError", + "atomic_write", + "compute_content_hash", + "read_checked", +] + +# Sidecar files sit next to the protected payload and carry the integrity +# metadata (schema version + content hash). Keeping them separate from the +# payload means the payload bytes round-trip unchanged. +_META_SUFFIX = ".meta.json" + +# Algorithm used for the stored content hash. Recorded in the sidecar so a +# future algorithm change stays backward-readable. +_HASH_ALGO = "sha256" + + +class IntegrityError(Exception): + """Raised when durable state fails its integrity check on load. + + Signals a schema-version mismatch, a missing/garbled integrity sidecar, or + a stored-content-hash mismatch (corruption or tampering). Callers treat + this as "refuse to proceed silently": park the task and ALARM rather than + restart blindly (§6.7). + """ + + +def compute_content_hash(data: bytes) -> str: + """Return the hex content hash for ``data`` (sha256). + + The same routine is used when writing the sidecar and when verifying on + load, so the two are guaranteed consistent. + """ + return hashlib.new(_HASH_ALGO, data).hexdigest() + + +def _meta_path(path: Path) -> Path: + """Return the sidecar metadata path for a payload ``path``.""" + return path.with_name(path.name + _META_SUFFIX) + + +def atomic_write(path: Path, data: bytes) -> None: + """Atomically write ``data`` to ``path`` (write-temp -> fsync -> rename). + + The bytes are written to a temporary file in the same directory, flushed + and ``fsync``-ed to durable storage, then ``os.replace``-d onto the final + path. ``os.replace`` is atomic on POSIX within a filesystem, so a reader + never observes a half-written file and a crash mid-write leaves either the + old payload or the new one, never a torn one. The containing directory is + ``fsync``-ed afterward so the rename itself is durable. + + This writes only the payload; integrity metadata is written by callers via + :func:`write_checked` / read back by :func:`read_checked`. (The sidecar is + written through this same primitive, so it is equally crash-safe.) + """ + path = Path(path) + directory = path.parent + directory.mkdir(parents=True, exist_ok=True) + + # delete=False so we control the rename; same dir guarantees same fs. + fd, tmp_name = tempfile.mkstemp( + prefix=path.name + ".", suffix=".tmp", dir=directory + ) + tmp_path = Path(tmp_name) + try: + with os.fdopen(fd, "wb") as handle: + handle.write(data) + handle.flush() + os.fsync(handle.fileno()) + os.replace(tmp_path, path) + except BaseException: + # Best-effort cleanup of the temp file on any failure. + try: + os.unlink(tmp_path) + except FileNotFoundError: + pass + raise + + # Make the rename itself durable by fsync-ing the directory. + dir_fd = os.open(directory, os.O_RDONLY) + try: + os.fsync(dir_fd) + except OSError: + # Some filesystems disallow directory fsync; the rename is still + # atomic, only its durability across power-loss is weakened. + pass + finally: + os.close(dir_fd) + + +def write_checked(path: Path, data: bytes, *, schema_version: int) -> None: + """Atomically write ``data`` plus its integrity sidecar. + + Writes the payload first, then the sidecar carrying ``schema_version`` and + the content hash. :func:`read_checked` verifies both. Both writes go + through :func:`atomic_write`, so each is crash-safe; if a crash lands + between them the sidecar is simply stale/absent and :func:`read_checked` + fails closed with :class:`IntegrityError`, which is the intended + refuse-to-proceed behaviour. + """ + path = Path(path) + atomic_write(path, data) + meta = { + "schema_version": int(schema_version), + "hash_algo": _HASH_ALGO, + "content_hash": compute_content_hash(data), + } + atomic_write(_meta_path(path), json.dumps(meta, sort_keys=True).encode("utf-8")) + + +def read_checked(path: Path, *, schema_version: int) -> bytes: + """Read and integrity-check ``path``, returning its bytes. + + Verifies the integrity sidecar exists, that its recorded + ``schema_version`` matches the expected ``schema_version``, and that the + stored content hash matches a freshly computed hash of the payload bytes. + Any mismatch (missing/garbled sidecar, schema drift, corruption/tampering) + raises :class:`IntegrityError`. + """ + path = Path(path) + try: + data = path.read_bytes() + except FileNotFoundError as exc: + raise IntegrityError(f"state payload missing: {path}") from exc + + meta_path = _meta_path(path) + try: + raw_meta = meta_path.read_bytes() + except FileNotFoundError as exc: + raise IntegrityError(f"integrity sidecar missing: {meta_path}") from exc + + try: + meta = json.loads(raw_meta) + except (ValueError, UnicodeDecodeError) as exc: + raise IntegrityError(f"integrity sidecar unreadable: {meta_path}") from exc + + stored_version = meta.get("schema_version") + if stored_version != schema_version: + raise IntegrityError( + f"schema-version mismatch for {path}: " + f"stored={stored_version!r} expected={schema_version!r}" + ) + + stored_hash = meta.get("content_hash") + actual_hash = compute_content_hash(data) + if stored_hash != actual_hash: + raise IntegrityError( + f"content-hash mismatch for {path}: " + f"stored={stored_hash!r} actual={actual_hash!r}" + ) + + return data diff --git a/agent-team/agent_team/task_model.py b/agent-team/agent_team/task_model.py new file mode 100644 index 0000000..2e09a91 --- /dev/null +++ b/agent-team/agent_team/task_model.py @@ -0,0 +1,155 @@ +"""Task-record / thread model + LangGraph graph-state schema (design §3.3). + +A task is a long-lived, resumable record (a LangGraph thread). This module is +the pure model layer — no I/O — defining: + +* :class:`TaskStatus` / :class:`Phase` — task lifecycle enums. +* :class:`TaskRecord` — the durable task record (§3.3 "the task record holds: + status, current phase, the full Q&A history, the plan, review verdicts, the + candidate diff, and CI results"). +* :class:`PipelineState` — a ``TypedDict`` used as the LangGraph graph state + schema; its keys mirror :class:`TaskRecord` fields. +* :func:`new_thread_id` — uuid thread-id minting. +* JSON serialization helpers (:func:`task_to_dict` / :func:`task_from_dict` / + :func:`task_to_json` / :func:`task_from_json`). + +The signatures here are CONTRACTS leaf builders import verbatim. +""" + +from __future__ import annotations + +import json +import uuid +from dataclasses import asdict, dataclass, field +from enum import Enum +from typing import Any, TypedDict + +__all__ = [ + "Phase", + "PipelineState", + "TaskRecord", + "TaskStatus", + "new_thread_id", + "task_from_dict", + "task_from_json", + "task_to_dict", + "task_to_json", +] + + +class TaskStatus(Enum): + """Top-level task lifecycle status. + + ``ACTIVE`` — progressing through stages. ``WAITING_HUMAN`` — suspended on a + LangGraph ``interrupt()`` awaiting Adam's answer. ``PARKED`` — stalled + (no answer in window, N failed build loops, or budget contention) and + ALARM-ed rather than spinning (§3.3, §6.6). ``DONE`` — draft PR + report + produced. ``FAILED`` — terminal failure. + """ + + ACTIVE = "active" + WAITING_HUMAN = "waiting_human" + PARKED = "parked" + DONE = "done" + FAILED = "failed" + + +class Phase(Enum): + """Pipeline phase the task is currently in (§3.3).""" + + INTAKE = "intake" + CLARIFY = "clarify" + PLAN = "plan" + REVIEW = "review" + BUILD = "build" + VERIFY = "verify" + PARKED = "parked" + DONE = "done" + + +def new_thread_id() -> str: + """Mint a fresh unique ``thread_id`` (uuid4 hex).""" + return uuid.uuid4().hex + + +@dataclass +class TaskRecord: + """The durable per-task record (§3.3, §3.3.1). + + Mirrors the LangGraph thread state; the SQLite checkpointer persists the + graph state while this record is the logical view the coordinator reasons + over. ``qa_history`` is the full clarifier Q&A; ``review_verdicts`` the + adversarial review outcomes; ``candidate_diff`` + ``diff_hash`` the builder + output and its ledger-recorded hash (§3.3.2); ``ci_results`` the + authenticated CI conclusion the verifier reads. + """ + + thread_id: str + status: TaskStatus + current_phase: Phase + qa_history: list[Any] = field(default_factory=list) + plan: dict[str, Any] | None = None + review_verdicts: list[Any] = field(default_factory=list) + candidate_diff: str | None = None + diff_hash: str | None = None + ci_results: dict[str, Any] | None = None + transport: str = "" + created_at: str | None = None + updated_at: str | None = None + + +class PipelineState(TypedDict, total=False): + """LangGraph graph-state schema; keys mirror :class:`TaskRecord` (§3.3). + + Used as the graph's state type. ``total=False`` so a node may write a + subset of keys per checkpoint transition. + """ + + thread_id: str + status: str + current_phase: str + qa_history: list[Any] + plan: dict[str, Any] | None + review_verdicts: list[Any] + candidate_diff: str | None + diff_hash: str | None + ci_results: dict[str, Any] | None + transport: str + created_at: str | None + updated_at: str | None + + +def task_to_dict(record: TaskRecord) -> dict[str, Any]: + """Serialize a :class:`TaskRecord` to a JSON-safe dict (enums -> values).""" + data = asdict(record) + data["status"] = record.status.value + data["current_phase"] = record.current_phase.value + return data + + +def task_from_dict(data: dict[str, Any]) -> TaskRecord: + """Rebuild a :class:`TaskRecord` from a :func:`task_to_dict` dict.""" + return TaskRecord( + thread_id=data["thread_id"], + status=TaskStatus(data["status"]), + current_phase=Phase(data["current_phase"]), + qa_history=list(data.get("qa_history", [])), + plan=data.get("plan"), + review_verdicts=list(data.get("review_verdicts", [])), + candidate_diff=data.get("candidate_diff"), + diff_hash=data.get("diff_hash"), + ci_results=data.get("ci_results"), + transport=data.get("transport", ""), + created_at=data.get("created_at"), + updated_at=data.get("updated_at"), + ) + + +def task_to_json(record: TaskRecord) -> str: + """Serialize a :class:`TaskRecord` to a JSON string.""" + return json.dumps(task_to_dict(record), sort_keys=True) + + +def task_from_json(payload: str | bytes) -> TaskRecord: + """Deserialize a :class:`TaskRecord` from a JSON string/bytes.""" + return task_from_dict(json.loads(payload)) diff --git a/agent-team/agent_team/transport/__init__.py b/agent-team/agent_team/transport/__init__.py new file mode 100644 index 0000000..b3d3acd --- /dev/null +++ b/agent-team/agent_team/transport/__init__.py @@ -0,0 +1,17 @@ +"""Transport seam for the durable human-in-the-loop responder (design §3.3.1). + +The ledger + resume logic are transport-independent; concrete Slack / GitHub / +Claude-Code adapters subclass :class:`Transport` in the leaves. +""" + +from agent_team.transport.base import ( + NormalizedAnswer, + QuestionSet, + Transport, +) + +__all__ = [ + "NormalizedAnswer", + "QuestionSet", + "Transport", +] diff --git a/agent-team/agent_team/transport/base.py b/agent-team/agent_team/transport/base.py new file mode 100644 index 0000000..015cfda --- /dev/null +++ b/agent-team/agent_team/transport/base.py @@ -0,0 +1,110 @@ +"""Transport interface ABC + payload dataclasses (design §3.3.1). + +The durable human-in-the-loop responder owns a notify+resume seam that is +transport-agnostic. This module defines the contract every adapter implements: + +* :class:`Transport` — abstract base with ``post_question`` (deliver a + question-set, return a ``channel_ref`` that embeds the ``question_id``) and + ``parse_answer`` (normalize an inbound raw payload to + ``(question_id, answer, via)``). +* :class:`QuestionSet` — the question-set payload carried by a LangGraph + ``interrupt()``. +* :class:`NormalizedAnswer` — the normalized inbound answer the responder + feeds into the §3.3.1 first-answer-wins compare-and-set. + +Concrete Slack / GitHub / Claude-Code adapters subclass :class:`Transport` in +the leaves. The signatures here are CONTRACTS the leaf builders import +verbatim, so they are explicit and final. +""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import Any + +__all__ = [ + "NormalizedAnswer", + "QuestionSet", + "Transport", +] + +# Marker template embedded in transports without native callback metadata +# (e.g. a GitHub issue comment), so an inbound answer can be mapped back to its +# question. Slack embeds the question_id in ``callback_id`` instead. +GITHUB_MARKER_TEMPLATE = "" + + +@dataclass +class QuestionSet: + """A set of questions delivered to Adam for one ``turn`` of a task (§3.3.1). + + Carried in the ``interrupt()`` payload alongside ``thread_id``, + ``question_id``, ``turn``, ``transport``, and ``deadline``. ``questions`` is + the ordered list of prompts; ``context`` is optional rendering metadata + (repo, summary) the adapter may surface. + """ + + thread_id: str + question_id: str + turn: int + questions: list[str] + context: dict[str, Any] = field(default_factory=dict) + + +@dataclass +class NormalizedAnswer: + """A transport-normalized inbound answer (§3.3.1). + + The responder maps this into the first-answer-wins compare-and-set: + ``UPDATE ... SET status='answered' ... WHERE question_id=? AND + status='open'``. ``via`` records the answering channel/identity for the + audit trail (``answered_via``). + """ + + question_id: str + answer: Any + via: str + + +class Transport(ABC): + """Abstract transport adapter (§3.3.1). + + Subclasses implement delivery and answer parsing for one channel. The + ledger and resume worker depend only on this interface, so an adapter can + ship first (Slack) and others follow without touching the durable core. + """ + + @abstractmethod + def post_question( + self, + *, + thread_id: str, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, + ) -> str: + """Deliver ``question_set`` and return its ``channel_ref``. + + The posted message MUST embed ``question_id`` so an inbound answer can + be mapped back (Slack ``callback_id``; a + ```` marker in a GitHub comment). The returned + ``channel_ref`` is the transport's locator for the post (Slack message + ``ts`` / issue-comment id / Claude session id) and is stored on the + ledger row so reconcile/recovery can act on it (§3.3.1). + """ + raise NotImplementedError + + @abstractmethod + def parse_answer(self, raw: Any) -> tuple[str, Any, str]: + """Normalize an inbound ``raw`` payload to ``(question_id, answer, via)``. + + Extracts the embedded ``question_id`` (from the Slack ``callback_id`` / + the GitHub marker / the Claude session), the answer value, and the + ``via`` channel identity. The responder feeds the result into the + atomic compare-and-set. Implementations may build a + :class:`NormalizedAnswer` internally and return its fields as the tuple + the contract specifies. + """ + raise NotImplementedError diff --git a/agent-team/agent_team/transport/claude_code_adapter.py b/agent-team/agent_team/transport/claude_code_adapter.py new file mode 100644 index 0000000..0cf808f --- /dev/null +++ b/agent-team/agent_team/transport/claude_code_adapter.py @@ -0,0 +1,285 @@ +"""Claude-Code transport adapter (design §3.3.1, §7.1 P4, D10). + +The Plane-2 human-in-the-loop responder is transport-agnostic: it depends only +on the :class:`~agent_team.transport.base.Transport` contract and maps every +inbound answer back to its ``question_id`` via the §3.3.1 first-answer-wins +compare-and-set. This module is the **Claude-Code-on-the-Mac** leaf of that +seam (D10), added in Phase P4 alongside the GitHub ticket-comment adapter. + +Channel model +------------- +Adam runs an interactive Claude Code session on his Mac and answers the +pipeline's clarifying questions inline. The R720 box has no standing write +path to that session, so delivery is **SSH-invoked from the Mac side**: the box +hands a rendered question-set to an injected *delivery sink* (the thing that +actually surfaces the prompt inside the Claude-Code session — a file drop, an +SSH-relayed message, a CLI print). The sink returns the **Claude session id**, +which §3.3.1 names as this transport's ``channel_ref`` (the analogue of a Slack +message ``ts`` or a GitHub issue-comment id). The ref is stored on the ledger +row so reconcile/recovery can act on it. + +Because Claude-Code sessions carry no native per-message callback metadata, the +adapter embeds the ``question_id`` two ways, both of which :meth:`parse_answer` +will accept: + +* in the **rendered prompt body**, via the shared + :data:`~agent_team.transport.base.GITHUB_MARKER_TEMPLATE` + (````) so a copy-pasted answer round-trips it; and +* on the returned **session id** (``claude-session::``) + so the ledger ``channel_ref`` alone is enough to recover the mapping. + +:meth:`parse_answer` normalizes an inbound payload to +``(question_id, answer, via)`` where ``via`` is :data:`VIA` (``"claude_code"``). +It accepts the ``question_id`` from (in priority order) an explicit field, the +``channel_ref`` minted by :meth:`post_question`, or the embedded prompt marker. + +This is pre-deployment scaffolding: no SDK, network, SSH, or secrets are +touched here. The live wiring is supplied by the injected ``delivery`` sink in +the leaves; the default sink raises so the adapter cannot silently no-op a real +delivery. +""" + +from __future__ import annotations + +import re +from typing import Any, Callable + +from agent_team.transport.base import ( + GITHUB_MARKER_TEMPLATE, + NormalizedAnswer, + QuestionSet, + Transport, +) + +__all__ = [ + "VIA", + "ClaudeCodeAdapter", + "ClaudeCodeDeliveryError", + "build_channel_ref", + "parse_channel_ref", + "render_prompt", +] + +# ``via`` identity recorded on the ledger (``answered_via``) for answers that +# arrive through the Claude-Code-on-the-Mac channel (§3.3.1 audit trail). +VIA = "claude_code" + +# ``channel_ref`` prefix for this transport. The full ref is +# ``claude-session::`` so the ledger row alone carries +# both the locator (the Claude session) and the question mapping. +_CHANNEL_REF_PREFIX = "claude-session" + +# Recovers ```` and ```` from a channel_ref. Session +# ids and question ids never contain ``:`` (uuid hex / opaque session token), +# so a two-split is unambiguous. +_CHANNEL_REF_RE = re.compile( + rf"^{re.escape(_CHANNEL_REF_PREFIX)}:(?P[^:]+):(?P[^:]+)$" +) + +# Recovers the embedded question_id from a rendered prompt body / pasted answer +# that round-tripped the ```` marker. +_MARKER_RE = re.compile(r"") + + +class ClaudeCodeDeliveryError(RuntimeError): + """Raised when the Claude-Code delivery sink cannot surface the question. + + Mirrors a failed Slack post: the responder leaves the ledger row ``open`` + with no ``channel_ref`` and the reconcile loop retries idempotently + (§3.3.1 "Delivery (and lost-post)"). + """ + + +def build_channel_ref(session_id: str, question_id: str) -> str: + """Mint the §3.3.1 ``channel_ref`` for a Claude-Code post. + + The ref embeds both the Claude session locator and the ``question_id`` so + an inbound answer can be mapped back from the stored ref alone. + """ + return f"{_CHANNEL_REF_PREFIX}:{session_id}:{question_id}" + + +def parse_channel_ref(channel_ref: str) -> tuple[str, str] | None: + """Split a ``channel_ref`` into ``(session_id, question_id)``. + + Returns ``None`` if ``channel_ref`` is not a Claude-Code ref (e.g. a Slack + ``ts`` or a GitHub comment id stored by a different transport). + """ + match = _CHANNEL_REF_RE.match(channel_ref) + if match is None: + return None + return match.group("session_id"), match.group("question_id") + + +def render_prompt( + *, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, +) -> str: + """Render the question-set into the prompt body delivered to Claude Code. + + The body embeds the ``question_id`` via the shared GitHub-style marker so a + copy-pasted answer round-trips the mapping even without the channel_ref, + surfaces the deadline, and any optional ``context`` (repo / summary) the + adapter is given. Pure formatting — no I/O. + """ + marker = GITHUB_MARKER_TEMPLATE.format(question_id=question_id) + lines: list[str] = [ + marker, + f"R720 pipeline question-set (turn {turn}) — please answer:", + ] + context = question_set.context or {} + repo = context.get("repo") + if repo: + lines.append(f"repo: {repo}") + summary = context.get("summary") + if summary: + lines.append(f"summary: {summary}") + for index, question in enumerate(question_set.questions, start=1): + lines.append(f"{index}. {question}") + lines.append(f"(answer by {deadline})") + return "\n".join(lines) + + +def _default_delivery(*, session_hint: str, prompt: str) -> str: + """Default delivery sink: refuse to silently no-op a real delivery. + + The live wiring (file drop / SSH relay / CLI surface) is injected in the + leaves; without it, attempting to post must fail loudly rather than report + a success that never reached Adam. + """ + raise ClaudeCodeDeliveryError( + "no Claude-Code delivery sink configured; inject `delivery` to post" + ) + + +class ClaudeCodeAdapter(Transport): + """Claude-Code-on-the-Mac transport adapter (§3.3.1, §7.1 P4, D10). + + Implements the :class:`~agent_team.transport.base.Transport` contract for + the Claude-Code channel. Delivery is delegated to an injected ``delivery`` + callable so the durable core stays free of SDK/SSH/secret wiring and the + adapter is unit-testable with a fake sink. + + Parameters + ---------- + delivery: + Callable invoked as ``delivery(session_hint=..., prompt=...) -> str`` + that surfaces ``prompt`` inside a Claude-Code session and returns the + Claude **session id** used as the ``channel_ref`` locator. Defaults to + a sink that raises :class:`ClaudeCodeDeliveryError` so a real post + cannot silently no-op. + session_hint: + Optional hint passed to ``delivery`` (e.g. a target session id to reuse + or a routing tag). Pure passthrough; the adapter does not interpret it. + """ + + def __init__( + self, + delivery: Callable[..., str] = _default_delivery, + *, + session_hint: str = "", + ) -> None: + self._delivery = delivery + self._session_hint = session_hint + + def post_question( + self, + *, + thread_id: str, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, + ) -> str: + """Deliver ``question_set`` to Claude Code and return its ``channel_ref``. + + Renders the prompt (embedding ``question_id`` via the shared marker), + hands it to the injected delivery sink, and folds the returned Claude + session id together with the ``question_id`` into the §3.3.1 + ``channel_ref``. A sink that returns a falsy/blank session id is treated + as a failed post (raises :class:`ClaudeCodeDeliveryError`) so the ledger + row is not stamped with an unusable ref. + """ + prompt = render_prompt( + question_id=question_id, + turn=turn, + question_set=question_set, + deadline=deadline, + ) + session_id = self._delivery(session_hint=self._session_hint, prompt=prompt) + if not session_id or not str(session_id).strip(): + raise ClaudeCodeDeliveryError( + "Claude-Code delivery returned no session id; treating post as failed" + ) + return build_channel_ref(str(session_id).strip(), question_id) + + def parse_answer(self, raw: Any) -> tuple[str, Any, str]: + """Normalize an inbound Claude-Code payload to ``(question_id, answer, via)``. + + ``raw`` is a mapping (the relayed answer payload). The ``answer`` is + read from the ``answer`` key (falling back to ``value`` / ``text``). The + ``question_id`` is recovered in priority order from: + + 1. an explicit ``question_id`` field, + 2. the ``channel_ref`` minted by :meth:`post_question`, + 3. the ```` marker embedded in the prompt + body / pasted answer text. + + ``via`` is always :data:`VIA`. Raises :class:`ValueError` if no + ``question_id`` can be recovered, so an unmappable answer is rejected + rather than silently mis-routed. + """ + if not isinstance(raw, dict): + raise ValueError( + f"Claude-Code answer payload must be a mapping, got {type(raw).__name__}" + ) + + question_id = self._extract_question_id(raw) + if question_id is None: + raise ValueError( + "Claude-Code answer payload carries no recoverable question_id " + "(checked 'question_id', 'channel_ref', and the prompt marker)" + ) + + answer = self._extract_answer(raw) + normalized = NormalizedAnswer(question_id=question_id, answer=answer, via=VIA) + return normalized.question_id, normalized.answer, normalized.via + + @staticmethod + def _extract_question_id(raw: dict[str, Any]) -> str | None: + """Recover the ``question_id`` from a raw payload, or ``None``.""" + explicit = raw.get("question_id") + if isinstance(explicit, str) and explicit: + return explicit + + channel_ref = raw.get("channel_ref") + if isinstance(channel_ref, str): + parsed = parse_channel_ref(channel_ref) + if parsed is not None: + return parsed[1] + + for marker_field in ("prompt", "text", "answer"): + candidate = raw.get(marker_field) + if isinstance(candidate, str): + match = _MARKER_RE.search(candidate) + if match is not None: + return match.group("question_id") + + return None + + @staticmethod + def _extract_answer(raw: dict[str, Any]) -> Any: + """Read the answer value from a raw payload. + + Prefers an explicit ``answer`` key, then ``value``/``text``. The value + is returned untouched (it may be a string, a structured choice, etc.), + matching :class:`~agent_team.transport.base.NormalizedAnswer.answer`'s + ``Any`` typing. + """ + for key in ("answer", "value", "text"): + if key in raw: + return raw[key] + return None diff --git a/agent-team/agent_team/transport/github_adapter.py b/agent-team/agent_team/transport/github_adapter.py new file mode 100644 index 0000000..9254a49 --- /dev/null +++ b/agent-team/agent_team/transport/github_adapter.py @@ -0,0 +1,347 @@ +"""GitHub issue-comment transport adapter (design §3.3.1, §7.1 P4). + +One concrete :class:`~agent_team.transport.base.Transport` implementation: it +delivers a question-set as a **GitHub issue comment** and parses an inbound +answer comment back into the ``(question_id, answer, via)`` tuple the durable +responder feeds into the §3.3.1 first-answer-wins compare-and-set. + +Why an HTML-comment marker (and not a native callback id like Slack): + GitHub issue comments carry no per-message callback metadata we control, so + the question-set comment embeds ```` (the + ``GITHUB_MARKER_TEMPLATE`` from the foundation contract). The answering + human quotes / replies under that comment, GitHub preserves the marker in + the quoted body, and :meth:`GitHubTransport.parse_answer` recovers the + ``question_id`` from it. This is exactly the mapping §3.3.1 specifies for + transports without native callback metadata. + +Delivery returns the new comment's numeric id (stringified) as the +``channel_ref`` the ledger stores (§3.3.1 "issue-comment id"), so +reconcile/recovery can act on it. + +Design constraints honoured here (pre-deployment scaffolding): + * **No live infrastructure.** Nothing is provisioned or called at import. + The HTTP transport is dependency-injected (``http_post``); the default + is a stdlib-only (``urllib``) poster invoked only on an actual post, so + there is no third-party dependency and the unit tests stay fully + hermetic (they inject an in-memory fake). + * **Secrets never committed.** The GitHub token is read from the + environment (``GITHUB_TOKEN`` by default) at call time, never stored in + source or logged. +""" + +from __future__ import annotations + +import json +import os +import re +from typing import Any, Callable, Protocol +from urllib import error as _urlerror +from urllib import request as _urlrequest + +from agent_team.transport.base import ( + GITHUB_MARKER_TEMPLATE, + NormalizedAnswer, + QuestionSet, + Transport, +) + +__all__ = [ + "GITHUB_API_ROOT", + "GitHubApiError", + "GitHubTransport", + "HttpPost", + "build_marker", + "extract_question_id", + "render_question_comment", +] + +# Default GitHub REST API root. Overridable per-instance for GitHub Enterprise. +GITHUB_API_ROOT = "https://api.github.com" + +# Compiled matcher for the foundation marker ````. +# ``question_id`` is a uuid4 hex in practice but the pattern stays permissive +# to also accept hyphenated test/synthetic ids. It captures everything up to +# the closing ``-->`` non-greedily, then strips trailing whitespace, so a +# malformed marker surfaces as "no match" rather than a silently wrong capture. +_MARKER_RE = re.compile(r"") + + +class GitHubApiError(RuntimeError): + """Raised when a GitHub REST call returns a non-success status. + + Carries the HTTP ``status`` and the (truncated) response ``body`` so the + reconcile loop can decide whether to retry. The triggering question is left + ``open`` with no ``channel_ref`` per §3.3.1's lost-post handling. + """ + + def __init__(self, status: int, body: str) -> None: + self.status = status + self.body = body + super().__init__(f"GitHub API error {status}: {body[:200]}") + + +class HttpPost(Protocol): + """Injected HTTP POST seam: ``(url, headers, json_body) -> (status, data)``. + + Returns the response status code and the parsed JSON body (a dict). Keeping + this a narrow callable means the adapter has no hard dependency on any HTTP + client and the tests pass a pure in-memory fake. + """ + + def __call__( + self, + url: str, + *, + headers: dict[str, str], + json_body: dict[str, Any], + ) -> tuple[int, dict[str, Any]]: ... + + +def build_marker(question_id: str) -> str: + """Render the hidden ``question_id`` marker for an outbound comment. + + Thin wrapper over the foundation ``GITHUB_MARKER_TEMPLATE`` so the leaf + never re-defines the template string (the contract owns it). + """ + return GITHUB_MARKER_TEMPLATE.format(question_id=question_id) + + +def extract_question_id(body: str) -> str | None: + """Return the ``question_id`` embedded in ``body``, or ``None`` if absent. + + Scans for the ```` marker. Works on both the + original question comment and an answer that quotes it (GitHub preserves the + HTML comment in the ``>``-quoted block). + """ + match = _MARKER_RE.search(body or "") + return match.group(1) if match else None + + +def render_question_comment( + question_set: QuestionSet, + *, + question_id: str, + turn: int, + deadline: str, +) -> str: + """Render the Markdown body for the outbound question-set comment. + + The body embeds the hidden ``question_id`` marker (so the answer can be + mapped back), a human-readable header, any ``context`` the question-set + carries (e.g. ``repo``/``summary``), and the ordered questions. Answer + instructions tell the human to reply *quoting this comment* so the marker + survives into their reply. + """ + lines: list[str] = [build_marker(question_id)] + lines.append(f"### Agent-team needs input (turn {turn})") + lines.append("") + + context = question_set.context or {} + repo = context.get("repo") + summary = context.get("summary") + if repo: + lines.append(f"**Repo:** {repo}") + if summary: + lines.append(f"**Summary:** {summary}") + if repo or summary: + lines.append("") + + if question_set.questions: + for index, question in enumerate(question_set.questions, start=1): + lines.append(f"{index}. {question}") + else: + lines.append("_(no questions)_") + lines.append("") + + lines.append(f"_Please reply **quoting this comment** by {deadline}._") + return "\n".join(lines) + + +def _default_http_post( + url: str, + *, + headers: dict[str, str], + json_body: dict[str, Any], +) -> tuple[int, dict[str, Any]]: + """Stdlib-only default POST (no third-party dependency at import time). + + Used only when no ``http_post`` is injected and an actual delivery is + attempted. Tests never reach this path — they inject a fake. + """ + payload = json.dumps(json_body).encode("utf-8") + request = _urlrequest.Request(url, data=payload, method="POST") + for key, value in headers.items(): + request.add_header(key, value) + try: + with _urlrequest.urlopen(request) as response: # noqa: S310 (trusted api host) + status = response.getcode() + raw = response.read().decode("utf-8") + except _urlerror.HTTPError as exc: # pragma: no cover - network path + raw = exc.read().decode("utf-8", "replace") + raise GitHubApiError(exc.code, raw) from exc + data = json.loads(raw) if raw else {} + return status, data + + +class GitHubTransport(Transport): + """Deliver / parse human-in-the-loop questions over GitHub issue comments. + + Posts a question-set as a comment on a fixed ``owner/repo#issue_number`` + thread and parses answers replied under it. The durable ledger + resume + worker depend only on the :class:`Transport` contract, so this adapter can + be swapped for Slack / Claude-Code without touching the core (§3.3.1). + """ + + def __init__( + self, + *, + owner: str, + repo: str, + issue_number: int, + http_post: HttpPost | None = None, + token_env: str = "GITHUB_TOKEN", + api_root: str = GITHUB_API_ROOT, + token_provider: Callable[[], str | None] | None = None, + ) -> None: + """Bind the adapter to one issue thread. + + Args: + owner: Repository owner / org login. + repo: Repository name. + issue_number: Issue (or PR) number whose comment thread carries the + question-sets. + http_post: Injected POST seam. Defaults to a stdlib-only poster + that is built lazily and only invoked on a real delivery. + token_env: Environment variable holding the GitHub token. Read at + call time so the secret is never captured in source/state. + api_root: REST API root (override for GitHub Enterprise). + token_provider: Optional explicit token source (takes precedence + over ``token_env``); lets a caller wire in a secrets manager + without an env round-trip. Must never be a literal token in + source. + """ + self.owner = owner + self.repo = repo + self.issue_number = issue_number + self._http_post = http_post or _default_http_post + self._token_env = token_env + self._api_root = api_root.rstrip("/") + self._token_provider = token_provider + + # -- delivery ----------------------------------------------------------- + + @property + def comments_url(self) -> str: + """REST endpoint for creating a comment on the bound issue.""" + return ( + f"{self._api_root}/repos/{self.owner}/{self.repo}" + f"/issues/{self.issue_number}/comments" + ) + + def _resolve_token(self) -> str: + """Fetch the GitHub token at call time (never stored on the instance).""" + token = ( + self._token_provider() + if self._token_provider is not None + else os.environ.get(self._token_env) + ) + if not token: + raise GitHubApiError( + 401, + f"no GitHub token available (env {self._token_env!r} unset)", + ) + return token + + def _headers(self) -> dict[str, str]: + return { + "Authorization": f"Bearer {self._resolve_token()}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + "Content-Type": "application/json", + } + + def post_question( + self, + *, + thread_id: str, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, + ) -> str: + """Post the question-set as an issue comment; return its comment id. + + The comment body embeds ```` so an inbound + answer maps back (§3.3.1). The returned ``channel_ref`` is the GitHub + comment id as a string, which the ledger stores for reconcile/recovery. + Raises :class:`GitHubApiError` on a non-2xx response (the row stays + ``open`` with no ref, and the reconcile loop retries idempotently). + """ + body = render_question_comment( + question_set, + question_id=question_id, + turn=turn, + deadline=deadline, + ) + status, data = self._http_post( + self.comments_url, + headers=self._headers(), + json_body={"body": body}, + ) + if not (200 <= status < 300): + raise GitHubApiError(status, json.dumps(data)) + comment_id = data.get("id") + if comment_id is None: + raise GitHubApiError(status, f"response missing comment id: {data!r}") + return str(comment_id) + + # -- answer parsing ----------------------------------------------------- + + def parse_answer(self, raw: Any) -> tuple[str, Any, str]: + """Normalize an inbound GitHub comment payload to ``(qid, answer, via)``. + + ``raw`` is the issue-comment webhook payload shape (or an equivalent + dict): ``{"comment": {"body": ..., "user": {"login": ...}}}``. A + flattened ``{"body": ..., "user": {...}}`` is also accepted. + + The ``question_id`` is recovered from the embedded marker; the answer is + the comment body with the marker line(s) stripped; ``via`` is + ``github:`` for the audit trail (``answered_via``). Raises + :class:`ValueError` if no marker is present (the responder treats an + unmappable comment as not an answer). + """ + comment = raw.get("comment", raw) if isinstance(raw, dict) else {} + body = comment.get("body", "") if isinstance(comment, dict) else "" + question_id = extract_question_id(body) + if question_id is None: + raise ValueError("no shq question marker found in comment body") + + user = comment.get("user") or {} + login = user.get("login") if isinstance(user, dict) else None + via = f"github:{login}" if login else "github" + + answer = self._strip_marker(body) + normalized = NormalizedAnswer(question_id=question_id, answer=answer, via=via) + return normalized.question_id, normalized.answer, normalized.via + + @staticmethod + def _strip_marker(body: str) -> str: + """Recover the human's answer text from a reply body. + + The reply typically quotes the original question comment, which drags + the ```` marker and the question text (as ``>``-prefixed + Markdown quote lines) into the body. To isolate the human's *new* text: + + * drop quote lines (those starting with ``>``) — that is the echoed + original question, not the answer; + * remove any remaining inline marker token, in case the marker sits on + the same line as a short answer (`` yes``); + * collapse the surrounding blank lines. + """ + kept: list[str] = [] + for line in (body or "").splitlines(): + if line.lstrip().startswith(">"): + continue + cleaned = _MARKER_RE.sub("", line) + kept.append(cleaned) + return "\n".join(kept).strip() diff --git a/agent-team/agent_team/transport/slack_adapter.py b/agent-team/agent_team/transport/slack_adapter.py new file mode 100644 index 0000000..5c0f7c3 --- /dev/null +++ b/agent-team/agent_team/transport/slack_adapter.py @@ -0,0 +1,353 @@ +"""Slack Block Kit transport adapter (design §3.3.1, §7.1 P1). + +The first concrete :class:`~agent_team.transport.base.Transport` implementation. +P1 ships the human gate over **one** transport (Slack first), so this adapter +must satisfy the §3.3.1 contracts the durable responder depends on: + +* :meth:`SlackTransport.post_question` renders a :class:`QuestionSet` as a Slack + Block Kit message, embeds the ``question_id`` in the message ``callback_id`` + (Slack's native callback metadata, the analogue of the GitHub ```` + marker), posts it via an **injected** poster callable, and returns the Slack + message ``ts`` as the ``channel_ref`` stored on the ledger row. +* :meth:`SlackTransport.parse_answer` normalizes an inbound Slack payload + (an interactive ``block_actions`` callback or a plain text/slash reply) to + ``(question_id, answer, via)`` for the first-answer-wins compare-and-set. + +This module is **transport I/O only** and carries no live wiring: the network +call is a dependency-injected ``poster`` callable, so the adapter is unit +testable and ships nothing provisioned. A production deployment supplies a +poster backed by the Composio Slack connector or ``slack_sdk`` (design §3 — +"the R720 routes Slack/Jira/Notion through" the Composio connector); the +foundation must not import or require either, so the default poster raises a +clear "not configured" error rather than reaching the network. +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping, Sequence +from typing import Any + +from agent_team.transport.base import ( + NormalizedAnswer, + QuestionSet, + Transport, +) + +__all__ = [ + "CALLBACK_ID_PREFIX", + "VIA_SLACK", + "SlackPostError", + "SlackPoster", + "SlackTransport", + "build_callback_id", + "build_question_blocks", + "parse_callback_id", +] + +# ``via`` channel-identity tag recorded in the ledger's ``answered_via`` column +# for the audit trail (§3.3.1). +VIA_SLACK = "slack" + +# Namespacing prefix for the embedded ``callback_id`` so an inbound Slack +# interaction can be unambiguously mapped back to its ledger ``question_id``. +# Mirrors the GitHub ```` marker convention (``shq`` = Sea Haven +# question) from the base module. +CALLBACK_ID_PREFIX = "shq" + +# Type of the network seam: given the rendered Slack message kwargs, perform the +# ``chat.postMessage`` and return the response payload. The only field this +# adapter requires from the response is the message ``ts`` (the ``channel_ref``). +SlackPoster = Callable[[dict[str, Any]], Mapping[str, Any]] + + +class SlackPostError(RuntimeError): + """Raised when posting a Slack message fails or no poster is configured. + + The §3.3.1 delivery rule writes the ledger row ``open`` *before* posting, so + a raised :class:`SlackPostError` leaves the row ``open`` with no + ``channel_ref`` and the reconcile loop retries delivery idempotently. The + caller (responder/delivery loop) is expected to catch this and leave the row + for reconcile rather than treating the question as delivered. + """ + + +def _default_poster(_message: dict[str, Any]) -> Mapping[str, Any]: + """Default poster: refuse to reach the network (pre-deployment scaffolding). + + The foundation must not provision or wire live infrastructure, so a + :class:`SlackTransport` constructed without an explicit ``poster`` cannot + post. Production injects a poster backed by the Composio Slack connector or + ``slack_sdk``. + """ + raise SlackPostError( + "SlackTransport has no poster configured; inject a Slack poster " + "callable (Composio connector / slack_sdk) before posting." + ) + + +def build_callback_id(question_id: str) -> str: + """Embed ``question_id`` in a namespaced Slack ``callback_id``. + + Slack echoes the message-level ``callback_id`` back on every interaction + payload, so it is the natural carrier for the ledger ``question_id`` + (§3.3.1: "Slack ``callback_id``"). + """ + return f"{CALLBACK_ID_PREFIX}:{question_id}" + + +def parse_callback_id(callback_id: str) -> str: + """Recover the ``question_id`` from a :func:`build_callback_id` value. + + Accepts both the namespaced form (``shq:``) and a bare + ``question_id`` (defensive, in case a payload surfaces the id directly). + Raises :class:`ValueError` on an empty / malformed callback id so a + mis-mapped answer fails loudly instead of being silently mis-attributed. + """ + if not callback_id: + raise ValueError("empty callback_id") + prefix = f"{CALLBACK_ID_PREFIX}:" + if callback_id.startswith(prefix): + question_id = callback_id[len(prefix) :] + if not question_id: + raise ValueError(f"callback_id missing question_id: {callback_id!r}") + return question_id + return callback_id + + +def build_question_blocks( + question_set: QuestionSet, deadline: str +) -> list[dict[str, Any]]: + """Render a :class:`QuestionSet` as Slack Block Kit blocks. + + Surfaces optional ``context`` (``repo`` / ``summary``) as a context block, + lists the ordered questions, and footnotes the ``deadline`` so Adam sees the + answer window. Returns a plain JSON-serializable list (no Slack SDK types), + keeping the adapter dependency-free. + """ + blocks: list[dict[str, Any]] = [ + { + "type": "header", + "text": { + "type": "plain_text", + "text": f"Agent-team needs input (turn {question_set.turn})", + }, + } + ] + + context_elements: list[dict[str, Any]] = [] + repo = question_set.context.get("repo") + if repo: + context_elements.append({"type": "mrkdwn", "text": f"*repo:* {repo}"}) + summary = question_set.context.get("summary") + if summary: + context_elements.append({"type": "mrkdwn", "text": str(summary)}) + if context_elements: + blocks.append({"type": "context", "elements": context_elements}) + + for index, question in enumerate(question_set.questions, start=1): + blocks.append( + { + "type": "section", + "text": {"type": "mrkdwn", "text": f"*{index}.* {question}"}, + } + ) + + blocks.append( + { + "type": "context", + "elements": [ + {"type": "mrkdwn", "text": f"_Reply in this thread by {deadline}._"} + ], + } + ) + return blocks + + +def _join_answer_actions(actions: Sequence[Mapping[str, Any]]) -> Any: + """Reduce one-or-more interactive actions to a single answer value. + + A button click yields one action; a multi-select yields several. A single + action collapses to its scalar value; multiple actions return the ordered + list of values so the responder records every selection. + """ + values = [_action_value(action) for action in actions] + if len(values) == 1: + return values[0] + return values + + +def _action_value(action: Mapping[str, Any]) -> Any: + """Extract the answer value from one Slack ``block_actions`` action.""" + if "value" in action and action["value"] is not None: + return action["value"] + selected = action.get("selected_option") + if isinstance(selected, Mapping): + return selected.get("value") + selected_options = action.get("selected_options") + if isinstance(selected_options, Sequence) and not isinstance( + selected_options, (str, bytes) + ): + return [ + opt.get("value") for opt in selected_options if isinstance(opt, Mapping) + ] + selected_user = action.get("selected_user") + if selected_user is not None: + return selected_user + # Fall back to the action_id so a payload without an explicit value still + # maps to *some* deterministic answer rather than ``None``. + return action.get("action_id") + + +class SlackTransport(Transport): + """Concrete Slack Block Kit :class:`Transport` (§3.3.1, P1 — Slack first). + + ``channel`` is the target Slack channel id. ``poster`` is the injected + network seam (defaults to a non-networking poster that raises, so the + foundation ships nothing live). ``thread_ts`` mode is implicit: when a + ``QuestionSet`` is delivered the adapter posts a top-level message and uses + its ``ts`` as the ``channel_ref``; answers arrive as thread replies or + interactive callbacks carrying the embedded ``question_id``. + """ + + def __init__( + self, + channel: str, + poster: SlackPoster | None = None, + ) -> None: + self.channel = channel + self._poster: SlackPoster = poster if poster is not None else _default_poster + + def post_question( + self, + *, + thread_id: str, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, + ) -> str: + """Render + post the question-set; return the Slack ``ts`` channel_ref. + + Embeds ``question_id`` in the message ``callback_id`` so an inbound + answer maps back (§3.3.1). On any poster failure raises + :class:`SlackPostError` so the ledger row stays ``open`` for reconcile. + """ + blocks = build_question_blocks(question_set, deadline) + message: dict[str, Any] = { + "channel": self.channel, + "callback_id": build_callback_id(question_id), + "text": ( + f"Agent-team needs input on task {thread_id} " + f"(turn {turn}); reply by {deadline}." + ), + "blocks": blocks, + "metadata": { + "event_type": "agent_team_question", + "event_payload": { + "thread_id": thread_id, + "question_id": question_id, + "turn": turn, + }, + }, + } + + try: + response = self._poster(message) + except SlackPostError: + raise + except Exception as exc: # noqa: BLE001 — normalize any poster failure + raise SlackPostError(f"Slack post failed: {exc}") from exc + + channel_ref = _extract_ts(response) + if not channel_ref: + raise SlackPostError( + "Slack post response missing message 'ts'; cannot record " + f"channel_ref (response keys: {sorted(response.keys())})" + ) + return channel_ref + + def parse_answer(self, raw: Any) -> tuple[str, Any, str]: + """Normalize an inbound Slack payload to ``(question_id, answer, via)``. + + Handles the two inbound shapes for P1: + + * an interactive ``block_actions`` payload — ``callback_id`` carries the + ``question_id`` and ``actions[*]`` carry the answer value(s); + * a plain text / slash reply — ``callback_id`` (or ``question_id``) + carries the id and ``text`` / ``answer`` carries the value. + + Raises :class:`ValueError` on a payload with no recoverable + ``question_id`` so a malformed answer is rejected loudly rather than + mis-attributed. + """ + if not isinstance(raw, Mapping): + raise ValueError(f"Slack payload must be a mapping, got {type(raw)!r}") + + question_id = _extract_question_id(raw) + answer = _extract_answer(raw) + normalized = NormalizedAnswer( + question_id=question_id, answer=answer, via=VIA_SLACK + ) + return normalized.question_id, normalized.answer, normalized.via + + +def _extract_ts(response: Mapping[str, Any]) -> str | None: + """Pull the message ``ts`` from a ``chat.postMessage`` response. + + Slack returns the timestamp at the top level (``{"ts": ...}``) and also + nested under ``message`` (``{"message": {"ts": ...}}``); accept either. + """ + ts = response.get("ts") + if ts: + return str(ts) + message = response.get("message") + if isinstance(message, Mapping) and message.get("ts"): + return str(message["ts"]) + return None + + +def _extract_question_id(raw: Mapping[str, Any]) -> str: + """Recover the ledger ``question_id`` from any supported inbound payload.""" + callback_id = raw.get("callback_id") + if callback_id: + return parse_callback_id(str(callback_id)) + + # Interactive payloads nest the callback metadata under ``view`` / ``message``. + view = raw.get("view") + if isinstance(view, Mapping) and view.get("callback_id"): + return parse_callback_id(str(view["callback_id"])) + message = raw.get("message") + if isinstance(message, Mapping): + metadata = message.get("metadata") + if isinstance(metadata, Mapping): + payload = metadata.get("event_payload") + if isinstance(payload, Mapping) and payload.get("question_id"): + return str(payload["question_id"]) + + question_id = raw.get("question_id") + if question_id: + return str(question_id) + + raise ValueError("Slack payload carries no recoverable question_id") + + +def _extract_answer(raw: Mapping[str, Any]) -> Any: + """Recover the answer value from any supported inbound payload.""" + actions = raw.get("actions") + if ( + isinstance(actions, Sequence) + and not isinstance(actions, (str, bytes)) + and actions + ): + return _join_answer_actions( + [action for action in actions if isinstance(action, Mapping)] + ) + + if "answer" in raw: + return raw["answer"] + if "text" in raw: + return raw["text"] + if "value" in raw: + return raw["value"] + + raise ValueError("Slack payload carries no recoverable answer") diff --git a/agent-team/ci/README.md b/agent-team/ci/README.md new file mode 100644 index 0000000..b181e6c --- /dev/null +++ b/agent-team/ci/README.md @@ -0,0 +1,155 @@ +# agent-team/ci — split-job CI apply/verify workflow (Plane-2 leaf) + +Pre-deployment scaffolding for the R720 agent-team SDLC pipeline. This directory +holds the **split-job CI apply/verify workflow** that turns a builder agent's +**untrusted candidate diff** into a verified **draft PR** — the §3.3.2 trust +boundary, Phase P3 (§7.1) of `../../docs/r720-agent-team-design.md`. + +> **STATUS: DEPLOY-GATED. NOT ENABLED, NOT PROVISIONED.** This is authored as +> files only. Per the design (§3.3.2, §7.1 P3) the workflow + its OIDC role must +> clear **BOTH `/sh-security-review` AND the mandatory GPT-4.1 cross-review** +> before deployment, because it is IaC/IAM + untrusted-input handling. The +> privileged draft-PR step is hard-disabled (`if: ${{ false }}`) and the OIDC +> `id-token`/`pull-requests: write` grants are left commented until those gates +> pass. Nothing here is wired to a live org repo. + +## Files + +| File | What it is | +|---|---| +| `agent-team-apply-verify.yml` | The split-job workflow. Self-contained: the load-bearing gate logic (diff integrity, trust-control denylist, declared-scope check, pure-code pass/fail) is embedded inline as stdlib-only, type-hinted Python heredocs, so the workflow has **no external script dependency**. | +| `README.md` | This file. | + +The filename is kebab-case per the handbook. Deployment target (later, after the +gates): promote into `Sea-Haven-Industries/.github` as a reusable workflow +(`engineering-handbook/cicd.md`); the Option-B OIDC apply path invokes it. + +## The trust boundary (design §3.3.2) + +The builder agents are semi-trusted: an LLM that read repo content can be wrong +or prompt-injected, so **the candidate diff is treated as untrusted code.** The +threat is that executing it in CI with org credentials lets a bad diff exfiltrate +secrets, assume the deploy role, or tamper with other repos. The workflow +implements all five boundaries: + +1. **Split CI — untrusted execution is credential-less.** The job that checks + out and runs the diff (`build-test`) runs with `permissions: contents: read`, + **no secrets, no OIDC, no write token**, and egress blocked + (harden-runner). The patch executes only there, where there is nothing to + steal and nothing to assume. Every privileged action (the eventual OIDC role, + the draft-PR open) runs in a **separate `gate-and-pr` job that never checks + out or executes patch-controlled code** — it consumes the build/test report + as **data only**. There is **no `pull_request_target` + head-ref checkout** + (the "pwn request" anti-pattern). +2. **Trust-control-surface denylist (CI-side hard fail).** The `guard` job + rejects any diff that touches `.github/workflows/**`, IAM/policy IaC + (CDK/SAM/Terraform), branch-protection / `CODEOWNERS` / Dependabot config, or + files **outside the task's declared scope**. The match is not naive: it + **canonicalizes paths, rejects parent-directory traversal, and inspects + `rename from/to` headers**, so a rename *into* a denied path — or path + indirection — cannot bypass it. Such a diff is escalated to mandatory human + + GPT cross-review, never auto-built. +3. **Diff integrity, box → CI.** The builder records the candidate diff's sha256 + in the task ledger (foundation `agent_team.state_store` content-hash idiom). + CI **re-hashes the diff and verifies it equals the recorded hash before + applying** (in both `guard` and again pre-apply in `build-test`). Precisely: + this is an **unkeyed** hash that binds *the bytes CI applies* to *the hash the + dispatcher recorded* — it detects accidental corruption or substitution of the + artifact **in transit** between the trusted dispatcher and CI. It does **not** + prove the diff matches the approved plan, and it cannot defend a compromised + box that generates both the diff and its hash (see "Threat-model honesty" + below). A keyed/attested binding to the approval record is a later hardening. +4. **Pure-code pass/fail gate over authenticated results.** Mirroring secrev's + "one pure-code script owns the block decision," the `gate-and-pr` gate reads + **only** the authenticated `needs.*.result` job conclusions (GitHub-controlled, + patch-independent) keyed to the run id + diff hash. It **never trusts a + success/failure file or artifact the patch could have written.** The verifier + *agent* only reads failures to propose the next fix; it cannot declare + success. +5. **Branch protection is the final enforcement.** The draft PR targets a + protected branch; the locked-down checks are required; merge needs them green + **plus** the security-review, the Claude Code App review, and **human + approval**. The agent never auto-merges (D2). + +**Threat-model honesty (§3.3.2).** The diff hash protects box → CI transit +(tamper/substitution), **not** a box that is itself compromised — the box is in +the trusted computing base for what it *proposes*. The real backstops against a +bad box are the credential-less build job, the trust-control-surface denylist, +and above all the **human review of the draft PR** plus required checks (nothing +merges without Adam). + +## Job graph + +``` +workflow_dispatch (task_id, diff_artifact_name, expected_diff_hash, declared_scope) + │ + ▼ + guard contents:read, no secrets — validates the diff as DATA: + (boundaries 2,3) re-hash + denylist + scope. Never applies it. + │ (needs) + ▼ + build-test contents:read, no secrets, no OIDC, egress blocked — + (boundary 1) the ONLY job that applies + runs the UNTRUSTED patch. + │ (needs) Emits a NON-authoritative report artifact. + ▼ + gate-and-pr privileged home, but NEVER checks out patch code. Pure-code + (boundaries 1,4,5) gate over authenticated needs.*.result → DRAFT PR + (hard-disabled until the review gates pass). +``` + +`permissions: {}` at the workflow level (least privilege); each job re-declares +its own grant explicitly. The trigger is `workflow_dispatch` only — the patch +never runs in a context carrying write or secret scope. + +## SHA-pinned actions (handbook Pinning Principle, §3.3.2) + +Every third-party action is pinned to a full commit SHA with the human-readable +tag in a trailing comment: + +| Action | SHA | Tag | +|---|---|---| +| `actions/checkout` | `11bd71901bbe5b1630ceea73d27597364c9af683` | v4.2.2 | +| `actions/download-artifact` | `fa0a91b85d4f404e444e00e005971372dc801d16` | v4.1.8 | +| `actions/upload-artifact` | `b4b15b8c7c6ac21ea08fcf65892d2ee8f75cf882` | v4.4.3 | +| `actions/setup-python` | `0b93645e9fea7318ecaed2b359559ac225c90a2b` | v5.3.0 | +| `step-security/harden-runner` | `0080882f6c36860b6ba35c610c98ce87d4e2f26f` | v2.10.2 | + +## Relationship to the foundation + +This leaf **imports the committed Plane-2 foundation contracts verbatim** (it +does not redefine them): + +- The diff-hash recorded in the ledger and re-checked in CI is the same + content-hash idiom as `agent_team.state_store.compute_content_hash` (§6.7). +- The task this workflow verifies is an `agent_team.task_model.TaskRecord`; its + `candidate_diff` + `diff_hash` fields (§3.3) are exactly the + `expected_diff_hash` this workflow consumes, and `ci_results` is what the + verifier writes back from the authenticated gate (boundary 4). +- The ledger that records provenance (diff hash, run id, gate decision) is the + `agent_team.db` schema (`pending_questions` / `budget_ledger` live there; + per-task CI provenance is recorded against the task thread). + +## Tests + +The workflow's embedded gate logic (diff integrity, the trust-control denylist +with path-canonicalization + rename/copy/delete detection, the symlink-escape +reject, and declared-scope enforcement) is **stdlib-only, type-hinted, and +ruff-clean**, and is covered by a committed, runnable suite: +`../tests/test_ci_gate_workflow.py` extracts the inline guard script from this +YAML and executes it against good and adversarial diffs — clean in-scope, +hash mismatch, workflow delete, copy-into-denied, symlink addition, non-UTF-8, +out-of-scope, unscoped, and escaping-scope. Run it with the rest of the suite: +`python3 -m pytest agent-team/tests/ -q` from the repo root. (The claim that the +gate is "verified" is therefore backed by that test, not by authoring alone.) + +## Deploy gating (do NOT skip) + +Before this ships (§3.3.2, §7.1 P3): + +1. `/sh-security-review` over this workflow (IaC + untrusted-input handling). +2. Mandatory **GPT-4.1 cross-review** of the workflow **and** the Option-B OIDC + role it will assume (IAM change). +3. A documented, **exercised** rollback (remove the role, revert the workflow). +4. Only then: uncomment the `id-token` / `pull-requests: write` grants, enable + the draft-PR step, and promote to `Sea-Haven-Industries/.github`. Draft PRs + only; never auto-merge. diff --git a/agent-team/ci/agent-team-apply-verify.yml b/agent-team/ci/agent-team-apply-verify.yml new file mode 100644 index 0000000..e3784c9 --- /dev/null +++ b/agent-team/ci/agent-team-apply-verify.yml @@ -0,0 +1,641 @@ +# R720 agent-team — split-job CI apply/verify workflow (design §3.3.2, §7.1 P3). +# +# DEPLOY-GATED PRE-DEPLOYMENT SCAFFOLDING. This file is authored as IaC only. +# It is NOT enabled, NOT provisioned, and NOT wired to any live org repo. Per +# the design (§3.3.2, §7.1 P3) it must clear BOTH `/sh-security-review` AND the +# mandatory GPT-4.1 cross-review before it is deployed (it is IaC/IAM + +# untrusted-input handling). Until then it lives here as a reviewable artifact. +# +# Deployment target (later, after the gates): promote into +# Sea-Haven-Industries/.github as a reusable workflow (engineering-handbook +# cicd.md) and have the Option-B OIDC apply path call it. The filename stays +# kebab-case per the handbook. +# +# ───────────────────────────────────────────────────────────────────────────── +# TRUST BOUNDARY (design §3.3.2). The builder agents are semi-trusted: an LLM +# that read repo content can be wrong or prompt-injected, so the candidate diff +# is UNTRUSTED CODE. The five boundaries this workflow implements: +# +# 1. Split CI. The job that checks out + executes the patch (`build-test`) +# runs credential-less (`permissions: contents: read`, no secrets, no +# OIDC, no write token, egress-restricted). Every privileged action runs +# in a SEPARATE job (`gate-and-pr`) that NEVER checks out or runs +# patch-controlled code; it consumes the build/test report as DATA only. +# This is NOT `pull_request_target` with a head-ref checkout (pwn request). +# 2. Trust-control-surface denylist. `guard` hard-fails (CI-side, not only the +# box) any diff touching `.github/workflows/**`, IAM/policy IaC, branch +# protection / CODEOWNERS / Dependabot, or files outside the declared task +# scope. It canonicalizes paths, resolves symlinks, and rejects renames +# into denied paths — a path match cannot be bypassed by indirection. +# 3. Diff integrity, box → CI. CI re-hashes the candidate diff and verifies it +# equals the ledger-recorded hash BEFORE applying. Tamper/substitution +# fails the hash check. +# 4. Pure-code pass/fail gate. A deterministic gate reads the authenticated +# build/test conclusion keyed to (run id + diff hash). It never trusts a +# success/failure file the patch could have written. The verifier AGENT +# only reads failures to propose a fix; it cannot declare success. +# 5. Branch protection. The draft PR targets a protected branch; the +# locked-down checks are required; merge needs them green + the +# security-review + the Claude Code App review + human approval. The agent +# NEVER auto-merges (D2). +# +# All third-party actions are SHA-pinned (handbook Pinning Principle, §3.3.2). +# ───────────────────────────────────────────────────────────────────────────── + +name: agent-team-apply-verify + +# Manual / API trigger only. The Option-B OIDC apply path (a trusted, separate +# workflow that owns the write token) invokes this with the candidate-diff +# artifact + the ledger-recorded hash + the declared scope. There is NO +# pull_request / pull_request_target trigger: the patch must never run in a +# context that carries write or secret scope (boundary 1). +on: + workflow_dispatch: + inputs: + task_id: + description: "Pipeline task thread_id (for provenance/audit)." + required: true + type: string + diff_artifact_name: + description: "Name of the uploaded candidate-diff artifact." + required: true + type: string + expected_diff_hash: + description: "Ledger-recorded sha256 of the candidate diff (boundary 3)." + required: true + type: string + declared_scope: + description: >- + Newline-separated list of glob paths the task is allowed to touch + (boundary 2). A diff that changes anything outside this set fails. + required: true + type: string + +# Workflow-level default: least privilege. Every job re-declares its own +# `permissions:` so the grant is explicit per job and the untrusted job can be +# audited at a glance. +permissions: {} + +# One in-flight apply/verify per task; a re-dispatch cancels the stale run so a +# superseded diff cannot race a newer one. +concurrency: + group: agent-team-apply-verify-${{ inputs.task_id }} + cancel-in-progress: true + +jobs: + # ─────────────────────────────────────────────────────────────────────────── + # JOB 1 — guard (boundaries 2 + 3). Credential-less. Validates the candidate + # diff WITHOUT applying or executing it: re-hashes it (integrity) and runs the + # trust-control-surface denylist + declared-scope check. This job reads the + # diff as DATA only — it never `git apply`s it, so even a hostile diff cannot + # run code here. A failure is terminal: the diff is rejected and ALARM-worthy. + # ─────────────────────────────────────────────────────────────────────────── + guard: + runs-on: ubuntu-latest + timeout-minutes: 5 + permissions: + contents: read + outputs: + diff_hash: ${{ steps.verify.outputs.diff_hash }} + steps: + - name: Harden runner (egress audit; no secrets present anyway) + uses: step-security/harden-runner@0080882f6c36860b6ba35c610c98ce87d4e2f26f # v2.10.2 + with: + egress-policy: block + # Only what fetching the artifact + GitHub API needs. The job holds + # no secrets, so a successful exfil yields nothing of value (§3.3.2), + # but we deny egress as defense-in-depth. + allowed-endpoints: > + github.com:443 + api.github.com:443 + objects.githubusercontent.com:443 + *.actions.githubusercontent.com:443 + + - name: Download candidate diff (data only; not applied) + uses: actions/download-artifact@fa0a91b85d4f404e444e00e005971372dc801d16 # v4.1.8 + with: + name: ${{ inputs.diff_artifact_name }} + path: ./_incoming + + - name: Verify diff integrity + trust-control denylist + scope + id: verify + env: + EXPECTED_DIFF_HASH: ${{ inputs.expected_diff_hash }} + DECLARED_SCOPE: ${{ inputs.declared_scope }} + DIFF_PATH: ./_incoming/candidate.diff + run: | + set -euo pipefail + # Self-contained, stdlib-only, type-hinted gate program embedded + # inline so this workflow has NO external script dependency. It is + # patch-independent: it parses the unified diff as TEXT and never + # executes it. It re-hashes the diff (boundary 3) and enforces the + # trust-control-surface denylist + declared scope (boundary 2), + # canonicalizing paths and rejecting renames into denied paths. + python3 - <<'PY' + from __future__ import annotations + + import hashlib + import os + import posixpath + import re + import sys + + # --- Boundary 2: the trust-control surface. Touching ANY of these is + # an auto-reject; such a diff is escalated to mandatory human + GPT + # cross-review, never auto-built (these are the mandatory-cross-review + # surface regardless). Matched against canonicalized POSIX paths. --- + DENY_GLOBS: tuple[str, ...] = ( + ".github/workflows/**", + ".github/actions/**", + "**/CODEOWNERS", + "CODEOWNERS", + ".github/dependabot.yml", + ".github/dependabot.yaml", + ".github/settings.yml", + # IAM / policy / permission IaC (CDK / SAM / Terraform). + "**/template.yaml", + "**/template.yml", + "**/*.tf", + "**/cdk.json", + "**/*-stack.ts", + "**/*_stack.py", + "**/policy*.json", + "**/*iam*", + "**/*.pem", + "**/*.key", + ) + + def canonical(path: str) -> str: + """Canonicalize a diff path to a normalized, anchored POSIX path. + + Strips git's a//b/ prefixes, collapses ``.`` / ``..`` and + backslashes, and rejects absolute or parent-escaping paths so a + denied location cannot be reached by traversal/indirection. + """ + p = path.strip() + # git unified-diff prefixes. + for pre in ("a/", "b/"): + if p.startswith(pre): + p = p[len(pre):] + break + p = p.replace("\\", "/") + # normpath then re-POSIX it. + norm = posixpath.normpath(p) + if norm.startswith("/") or norm == ".." or norm.startswith("../"): + raise ValueError(f"path escapes repo root: {path!r}") + return norm + + def parse_touched_paths(diff_text: str) -> set[str]: + """Extract every path a unified diff adds/modifies/renames/deletes. + + Reads ``+++ ``/``--- `` targets, ``diff --git a/x b/y`` headers, and + ``rename from/to`` lines — so a rename INTO a denied path (or a new + file generated into one) is caught, not just in-place edits. + """ + touched: set[str] = set() + for line in diff_text.splitlines(): + m = re.match(r"^diff --git (\S+) (\S+)$", line) + if m: + for raw in (m.group(1), m.group(2)): + touched.add(canonical(raw)) + continue + m = re.match(r"^(?:\+\+\+|---) (.+)$", line) + if m: + tgt = m.group(1).strip() + if tgt == "/dev/null": + continue + # strip trailing tab-timestamp some diffs carry. + tgt = tgt.split("\t", 1)[0] + touched.add(canonical(tgt)) + continue + m = re.match(r"^rename (?:from|to) (.+)$", line) + if m: + touched.add(canonical(m.group(1).strip())) + return touched + + def find_symlink_additions(diff_text: str) -> list[tuple[str, str]]: + """Return ``[(path, target)]`` for every symlink the diff creates. + + A symlink shows as git file mode ``120000``; its link target is the + single added content line. Textual path canonicalization (``canonical``) + cannot see a symlink that redirects a later in-diff write into a denied + location (e.g. ``sub/link -> ../.github/workflows`` then a write to + ``sub/link/evil.yml``). A candidate auto-build diff has no legitimate + reason to introduce a symlink, so guard treats ANY symlink addition as + a hard reject (boundary 2), closing the symlink-escape vector. + """ + additions: list[tuple[str, str]] = [] + cur_path: str | None = None + pending = False + for line in diff_text.splitlines(): + g = re.match(r"^diff --git (\S+) (\S+)$", line) + if g: + cur_path, pending = g.group(2), False + continue + p = re.match(r"^\+\+\+ (.+)$", line) + if p and p.group(1).strip() != "/dev/null": + cur_path = p.group(1).split("\t", 1)[0].strip() + continue + if re.match(r"^(?:new file mode|new mode) 120000\s*$", line): + pending = True + continue + if pending and line.startswith("+") and not line.startswith("+++"): + try: + path_c = canonical(cur_path) if cur_path else "" + except ValueError: + # canonical() rejected the path (absolute/escaping); report + # it with only the git a//b/ PREFIX removed for the error + # message (re.sub, not str.lstrip which strips a char set). + path_c = re.sub(r"^[ab]/", "", cur_path or "") + additions.append((path_c, line[1:].strip())) + pending = False + return additions + + _GLOB_META = set("*?[]") + _GLOB_RE_CACHE: dict[str, "re.Pattern[str]"] = {} + + def _glob_to_regex(glob: str) -> "re.Pattern[str]": + """Compile a gitignore-style glob to a '/'-aware, case-insensitive regex. + + Python's ``fnmatch`` does NOT implement recursive ``**`` (it treats it + as a single ``*`` that already spans ``/``), so ``**/template.yaml`` + fails to match a repo-ROOT ``template.yaml`` — a denylist bypass for + exactly the IaC/secret families boundary 2 must catch. This translates + ``**/`` to "any depth INCLUDING zero", ``**`` to ".*", ``*`` to a single + non-slash run, ``?`` to one non-slash char, and matches case- + insensitively (POSIX runners are case-sensitive, but a case variant of a + trust-control filename must not slip the gate). + """ + cached = _GLOB_RE_CACHE.get(glob) + if cached is not None: + return cached + out: list[str] = [] + i, n = 0, len(glob) + while i < n: + if glob[i : i + 3] == "**/": + out.append(r"(?:.*/)?") + i += 3 + elif glob[i : i + 2] == "**": + out.append(r".*") + i += 2 + elif glob[i] == "*": + out.append(r"[^/]*") + i += 1 + elif glob[i] == "?": + out.append(r"[^/]") + i += 1 + else: + out.append(re.escape(glob[i])) + i += 1 + pat = re.compile("^" + "".join(out) + "$", re.IGNORECASE) + _GLOB_RE_CACHE[glob] = pat + return pat + + def denied(path: str) -> bool: + """True if ``path`` is on the trust-control denylist (recursive, case-insensitive).""" + return any(_glob_to_regex(g).match(path) for g in DENY_GLOBS) + + def _scope_prefix(entry: str) -> str | None: + """Reduce a canonicalized scope entry to a concrete dir/file prefix. + + Declared scope is *confinement*, not a pattern that may widen coverage. + ``fnmatch``-ing scope let a single ``**`` (or ``*``) entry match the + whole tree, collapsing boundary 2b to a no-op. Instead we take the + leading path segments up to the first glob metacharacter and prefix + -match against them (mirrors the box-side ``_in_scope``). A scope that + begins with a metacharacter reduces to the empty (repo-root) prefix and + is dropped, so it can never widen to everything. + """ + keep: list[str] = [] + for part in entry.split("/"): + if any(c in _GLOB_META for c in part): + break + keep.append(part) + prefix = "/".join(keep) + return prefix or None + + def safe_scope(scope: list[str]) -> list[str]: + """Canonicalize scope into concrete path prefixes; drop escaping/empty. + + An absolute or parent-escaping entry is discarded (canonical raises), + and a glob that reduces to the repo root is dropped, so a malformed or + over-broad scope can only SHRINK what is allowed, never widen it. + """ + safe: list[str] = [] + for g in scope: + try: + canon = canonical(g) + except ValueError: + continue + prefix = _scope_prefix(canon) + if prefix is not None and prefix not in safe: + safe.append(prefix) + return safe + + def in_scope(path: str, scope: list[str]) -> bool: + """True if ``path`` is at or under one of the declared scope prefixes.""" + return any(path == entry or path.startswith(entry + "/") for entry in scope) + + def main() -> int: + diff_path = os.environ["DIFF_PATH"] + expected = os.environ["EXPECTED_DIFF_HASH"].strip().lower() + scope = [s for s in os.environ.get("DECLARED_SCOPE", "").splitlines() if s.strip()] + + with open(diff_path, "rb") as fh: + raw = fh.read() + actual = hashlib.sha256(raw).hexdigest() + + # Boundary 3: integrity. A tampered/substituted diff fails here. + if actual != expected: + print(f"::error::diff hash mismatch: expected={expected} actual={actual}") + return 2 + + # Fail CLOSED on a non-UTF-8 diff rather than silently replacing bytes + # (errors='replace' could let a homoglyph/encoding trick evade the path + # match). A legitimate diff over source is valid UTF-8. + try: + text = raw.decode("utf-8") + except UnicodeDecodeError as exc: + print(f"::error::diff is not valid UTF-8 ({exc}); refusing to parse") + return 8 + + touched = parse_touched_paths(text) + if not touched: + print("::error::no paths parsed from diff; refusing empty/garbled diff") + return 3 + + # Boundary 2a: trust-control denylist (CI-side HARD FAIL). + hits = sorted(p for p in touched if denied(p)) + if hits: + for h in hits: + print(f"::error::trust-control-surface violation: {h}") + print("::error::diff touches the trust-control surface; escalate to human + GPT cross-review") + return 4 + + # Boundary 2a': symlink escape. A symlink can redirect a later in-diff + # write into a denied path that textual matching cannot see, so any + # symlink addition is rejected outright. + symlinks = find_symlink_additions(text) + if symlinks: + for path, target in symlinks: + print(f"::error::diff introduces a symlink ({path} -> {target}); symlinks can redirect writes into denied paths and are not allowed in an auto-built diff") + print("::error::symlink in candidate diff; escalate to human + GPT cross-review") + return 7 + + # Boundary 2b: declared-scope enforcement. + if not scope: + print("::error::no declared scope provided; refusing unscoped diff") + return 5 + scope = safe_scope(scope) + if not scope: + print("::error::declared scope has no valid (non-escaping) entries; refusing diff") + return 5 + out_of_scope = sorted(p for p in touched if not in_scope(p, scope)) + if out_of_scope: + for p in out_of_scope: + print(f"::error::out-of-declared-scope path: {p}") + return 6 + + gh_out = os.environ.get("GITHUB_OUTPUT") + if gh_out: + with open(gh_out, "a", encoding="utf-8") as fh: + fh.write(f"diff_hash={actual}\n") + print(f"diff_hash={actual}") + print(f"validated {len(touched)} path(s); all in-scope, none on the trust-control surface") + return 0 + + sys.exit(main()) + PY + + # ─────────────────────────────────────────────────────────────────────────── + # JOB 2 — build-test (boundary 1). UNTRUSTED execution. This is the ONLY job + # that applies + runs the patch. It is credential-less: contents:read only, no + # secrets, no OIDC, no write token, egress blocked. There is nothing here to + # steal and nothing to assume. It writes a report artifact consumed by the + # privileged gate as DATA — that report is NOT authoritative (boundary 4). + # Depends on `guard` so a denied/tampered diff never reaches execution. + # ─────────────────────────────────────────────────────────────────────────── + build-test: + needs: guard + runs-on: ubuntu-latest + timeout-minutes: 20 + permissions: + contents: read + steps: + - name: Harden runner (block egress — untrusted code runs here) + uses: step-security/harden-runner@0080882f6c36860b6ba35c610c98ce87d4e2f26f # v2.10.2 + with: + # Block, not audit: this is where untrusted patch code executes. A + # narrow allowlist for dependency resolution only; everything else is + # denied so a prompt-injected patch cannot phone home. + # DEPLOY: this allowlist is GitHub + PyPI only. Before enabling this + # workflow for a repo, replace/extend it with EXACTLY that repo's + # package registries (npm, crates, Go proxy, ...) and nothing more — + # an over-broad allowlist weakens the egress boundary. + egress-policy: block + allowed-endpoints: > + github.com:443 + api.github.com:443 + objects.githubusercontent.com:443 + codeload.github.com:443 + pypi.org:443 + files.pythonhosted.org:443 + + - name: Checkout base repo (clean ref; patch applied on top after) + uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2 + with: + # Checkout carries NO token into the working tree usable for writes — + # this job's permissions are contents:read. persist-credentials:false + # guarantees the patch cannot reuse the checkout token. + persist-credentials: false + + - name: Re-download candidate diff (re-validated below) + uses: actions/download-artifact@fa0a91b85d4f404e444e00e005971372dc801d16 # v4.1.8 + with: + name: ${{ inputs.diff_artifact_name }} + path: ./_incoming + + - name: Re-verify diff hash before apply (defense-in-depth) + env: + EXPECTED_DIFF_HASH: ${{ needs.guard.outputs.diff_hash }} + DIFF_PATH: ./_incoming/candidate.diff + run: | + set -euo pipefail + # Independently confirm the bytes match the hash `guard` blessed, so a + # swapped artifact between jobs cannot slip an unvetted diff into the + # apply step. + python3 - <<'PY' + from __future__ import annotations + + import hashlib + import os + import sys + + def main() -> int: + expected = os.environ["EXPECTED_DIFF_HASH"].strip().lower() + with open(os.environ["DIFF_PATH"], "rb") as fh: + actual = hashlib.sha256(fh.read()).hexdigest() + if actual != expected: + print(f"::error::pre-apply hash mismatch: expected={expected} actual={actual}") + return 1 + print(f"diff hash confirmed: {actual}") + return 0 + + sys.exit(main()) + PY + + - name: Apply candidate diff (UNTRUSTED — credential-less sandbox) + run: | + set -euo pipefail + # --check first so a malformed diff fails cleanly; then apply. The + # working tree has no write credential, so applying + running it can + # touch only this ephemeral runner. + git apply --check ./_incoming/candidate.diff + git apply ./_incoming/candidate.diff + + - name: Set up Python + uses: actions/setup-python@0b93645e9fea7318ecaed2b359559ac225c90a2b # v5.3.0 + with: + python-version: "3.12" + + - name: Install + build + test (untrusted; result is non-authoritative) + id: run + run: | + set -euo pipefail + # Placeholder build/test for the narrowest task class (dep bump / + # single-file fix). At deploy time this is parameterized per target + # repo. Exit code is what matters; any file the patch writes is + # ignored by the authoritative gate (boundary 4). + if [ -f requirements.txt ]; then + python3 -m pip install --quiet -r requirements.txt || true + fi + python3 -m pip install --quiet ruff pytest || true + ruff check . || echo "ruff non-zero (recorded, non-authoritative)" + pytest -q || echo "pytest non-zero (recorded, non-authoritative)" + + - name: Emit non-authoritative report (job conclusion is the truth) + if: always() + run: | + set -euo pipefail + # This report is consumed by the gate as DATA for the verifier agent's + # next-fix reasoning. It is NOT the pass/fail decision — the gate reads + # the AUTHENTICATED job conclusion (boundary 4), never this file. + mkdir -p ./_report + printf '{"task_id":"%s","note":"non-authoritative; gate uses job conclusion"}\n' \ + "${{ inputs.task_id }}" > ./_report/report.json + + - name: Upload non-authoritative report + if: always() + uses: actions/upload-artifact@b4b15b8c7c6ac21ea08fcf65892d2ee8f75cf882 # v4.4.3 + with: + name: build-test-report-${{ inputs.task_id }} + path: ./_report/report.json + retention-days: 7 + + # ─────────────────────────────────────────────────────────────────────────── + # JOB 3 — gate-and-pr (boundaries 1 + 4 + 5). PRIVILEGED, but it NEVER checks + # out or executes patch-controlled code. It reads the AUTHENTICATED conclusion + # of `build-test` (via needs.*.result — GitHub-controlled, patch-independent) + # keyed to this run, and only on a clean pass opens a DRAFT PR. It never trusts + # any artifact the patch wrote. Pass/fail is pure code here, not the LLM. + # + # NOTE: the OIDC/write grant is declared here as the eventual home of the + # privileged step, but this file is deploy-gated — the `id-token`/PR-open + # step is left as a documented placeholder so nothing is provisioned until the + # §3.3.2 review gates pass. Wiring the real OIDC role is Phase P3 / Phase 5 + # AFTER the mandatory GPT-4.1 cross-review of the IAM. + # ─────────────────────────────────────────────────────────────────────────── + gate-and-pr: + needs: [guard, build-test] + # `always()` so the gate runs even when build-test failed, to record the + # authoritative conclusion. The gate itself decides pass/fail from results. + if: always() + runs-on: ubuntu-latest + timeout-minutes: 5 + permissions: + contents: read + # pull-requests: write # ← enabled ONLY after the §3.3.2 review gates. + # id-token: write # ← OIDC for the Option-B apply role, post-gate. + steps: + - name: Harden runner (privileged job; block egress) + uses: step-security/harden-runner@0080882f6c36860b6ba35c610c98ce87d4e2f26f # v2.10.2 + with: + egress-policy: block + allowed-endpoints: > + github.com:443 + api.github.com:443 + + - name: Pure-code pass/fail gate over authenticated results + env: + # These come from GitHub's job orchestration, NOT from the patch. + GUARD_RESULT: ${{ needs.guard.result }} + BUILD_TEST_RESULT: ${{ needs.build-test.result }} + DIFF_HASH: ${{ needs.guard.outputs.diff_hash }} + EXPECTED_DIFF_HASH: ${{ inputs.expected_diff_hash }} + RUN_ID: ${{ github.run_id }} + run: | + set -euo pipefail + # Deterministic, patch-independent decision. Consumes ONLY the + # authenticated needs.*.result values + the hash binding (all from + # GitHub's orchestration, never from a file the patch wrote). Mirrors + # secrev's "one pure-code script owns the block decision." The + # verifier AGENT only reads failures to propose a fix; it cannot + # declare success here. + python3 - <<'PY' + from __future__ import annotations + + import os + import sys + + def gate( + *, + guard_result: str, + build_test_result: str, + diff_hash: str, + expected_hash: str, + run_id: str, + ) -> tuple[bool, str]: + """Return (passed, reason) from authenticated, patch-independent inputs. + + A pass requires: the guard job succeeded (integrity + denylist + + scope all held), the build-test job succeeded, and the hash the + guard exported equals the ledger-recorded expected hash bound to + this run. Anything else blocks. + """ + if not run_id: + return False, "missing run id; cannot bind decision to a run" + if diff_hash.strip().lower() != expected_hash.strip().lower(): + return False, f"hash binding broken: guard={diff_hash} expected={expected_hash}" + if guard_result != "success": + return False, f"guard did not pass: {guard_result!r}" + if build_test_result != "success": + return False, f"build-test did not pass: {build_test_result!r}" + return True, "authenticated build/test passed and diff hash is bound" + + def main() -> int: + passed, reason = gate( + guard_result=os.environ.get("GUARD_RESULT", ""), + build_test_result=os.environ.get("BUILD_TEST_RESULT", ""), + diff_hash=os.environ.get("DIFF_HASH", ""), + expected_hash=os.environ.get("EXPECTED_DIFF_HASH", ""), + run_id=os.environ.get("RUN_ID", ""), + ) + if passed: + print(f"GATE PASS: {reason}") + if (gh_out := os.environ.get("GITHUB_OUTPUT")): + with open(gh_out, "a", encoding="utf-8") as fh: + fh.write("gate=pass\n") + return 0 + print(f"::error::GATE BLOCK: {reason}") + return 1 + + sys.exit(main()) + PY + + - name: Open DRAFT PR (DEPLOY-GATED PLACEHOLDER — not enabled) + if: ${{ false }} # ← hard-disabled. Enable only after §3.3.2 review gates. + run: | + echo "Draft-PR open runs here AFTER the mandatory GPT-4.1 cross-review" + echo "+ /sh-security-review of this workflow and its OIDC role." + echo "Draft PR only; never auto-merge (D2). Branch protection is the" + echo "final enforcement (boundary 5)." diff --git a/agent-team/run-team.py b/agent-team/run-team.py new file mode 100644 index 0000000..71e895b --- /dev/null +++ b/agent-team/run-team.py @@ -0,0 +1,699 @@ +#!/usr/bin/env python3 +"""``run-team.py`` — R720 agent-team operator entry CLI (design §3.3.1, §7.1 P1). + +This is the **entry CLI** named in the design (§2 "entry CLI ``run-team.py``"; +§9 "operator CLI"). It is the small manual path over the durable +``pending_questions`` ledger that §3.3.1 ("Manual path") requires:: + + A small CLI over the ledger lets an operator list ``open``/``parked`` + questions, re-deliver, force-expire, or answer on a task's behalf; a stuck + task parks rather than spins. Destructive CLI actions (force-expire, + answer-on-behalf, force-resume) are audit-logged and require an explicit + confirmation flag. + +It imports the committed FOUNDATION contracts verbatim — it does not redefine +them: + +* :mod:`agent_team.db.schema` — :func:`connect`, :func:`init_db`, + :func:`answer_question`, :func:`expire_question`, :func:`supersede_question`, + :data:`QUESTION_STATES`. +* :mod:`agent_team.state_store` — :func:`atomic_write` for the append-only, + crash-safe audit log of destructive actions (§6.7 discipline). + +Per the build constraints this is **pre-deployment scaffolding**: it provisions +nothing, enables no live CI, and performs no network or rsync. It only reads and +mutates the local SQLite ledger and writes a local audit log. + +Subcommands (P1 surface): + +* ``init-db`` — create/upgrade the agent-team tables in the ledger DB + (idempotent; wraps :func:`init_db`). +* ``list`` — list ``open`` (default) or any-status pending questions; with + ``--parked`` it lists questions whose status is read as parked context. Pure + read; no confirmation needed. +* ``show`` — print one question row by ``question_id``. Pure read. +* ``expire`` — force-expire an ``open`` question (DESTRUCTIVE: requires + ``--confirm``; audit-logged). Maps to :func:`expire_question`. +* ``answer`` — answer a question on a task's behalf (DESTRUCTIVE: requires + ``--confirm``; audit-logged). Maps to :func:`answer_question`. +* ``supersede`` — mark a stale question ``superseded`` (DESTRUCTIVE: requires + ``--confirm``; audit-logged). Maps to :func:`supersede_question`. + +Exit codes: ``0`` success, ``1`` operational failure (e.g. row not found, the +compare-and-set lost the race), ``2`` usage error (argparse). +""" + +from __future__ import annotations + +import argparse +import getpass +import json +import os +import sqlite3 +import sys +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Sequence + +# ``run-team.py`` lives in ``agent-team/`` next to the importable ``agent_team`` +# package. The hyphenated filename cannot itself be imported, so when run as a +# script we make the sibling package importable without an editable install +# (mirrors tests/conftest.py). +_CLI_DIR = Path(__file__).resolve().parent +if str(_CLI_DIR) not in sys.path: + sys.path.insert(0, str(_CLI_DIR)) + +from agent_team.db.schema import ( # noqa: E402 (path bootstrap must precede) + QUESTION_STATES, + answer_question, + connect, + expire_question, + init_db, + reopen_question, + supersede_question, +) + +__all__ = [ + "build_parser", + "main", +] + +# Default ledger DB location. Kept out of the repo (the package .gitignore +# excludes ``state/`` and ``*.sqlite``) so durable state is never committed. +_DEFAULT_DB = _CLI_DIR / "state" / "agent_team.sqlite" + +# Default audit log for destructive actions, alongside the ledger DB. +_DEFAULT_AUDIT_LOG = _CLI_DIR / "state" / "audit.log.jsonl" + +# Columns selected for list/show rendering, in display order. +_QUESTION_COLUMNS: tuple[str, ...] = ( + "question_id", + "thread_id", + "turn", + "status", + "transport", + "channel_ref", + "posted_at", + "deadline_at", + "answered_at", + "answered_via", +) + +# Destructive subcommands that require ``--confirm`` and are audit-logged. +# ``force-resume`` is the design-named operator verb (§3.3.1/§6.6); ``supersede`` +# is kept as its lower-level alias. ``redeliver`` is NOT here — it is idempotent +# and non-destructive (it only clears a delivery ref), though it is still +# audit-logged for provenance. +_DESTRUCTIVE_ACTIONS: frozenset[str] = frozenset( + {"expire", "answer", "supersede", "force-resume"} +) + + +def _utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string (audit timestamps).""" + return datetime.now(timezone.utc).isoformat() + + +def _default_operator() -> str: + """Best-effort OS login for audit attribution (never an empty string). + + A previous empty default left destructive actions non-attributable (the + audit record named no one). Defaulting to the OS login keeps the §3.3.1 + "audit-logged AND attributable" guarantee even when --operator is omitted; + falls back to "unknown" only if the login cannot be resolved. + """ + try: + user = getpass.getuser() + except Exception: # noqa: BLE001 - getuser can raise on odd environments + return "unknown" + return user or "unknown" + + +def _row_to_dict(row: sqlite3.Row) -> dict[str, Any]: + """Project a ``pending_questions`` row to a plain dict for display.""" + return {col: row[col] for col in _QUESTION_COLUMNS if col in row.keys()} + + +def _append_audit(audit_log: Path, entry: dict[str, Any]) -> None: + """Append one JSON audit record with an atomic ``O_APPEND`` single write. + + Destructive actions (§3.3.1) must leave an attributable trail. A previous + read-modify-rewrite design lost records under concurrent operators (two + processes each read the same bytes and the last rewrite wins). Instead each + record is one line written with ``O_APPEND``: the kernel serializes the + append and a write below ``PIPE_BUF`` is atomic on POSIX, so concurrent + appends never clobber each other. The file is created mode ``0600`` + (operator identity / action content is sensitive) and re-chmod'd in case it + pre-existed wider. A failure here raises ``OSError`` BEFORE any ledger + mutation, preserving the audit-before-mutate guarantee. + """ + audit_log = Path(audit_log) + audit_log.parent.mkdir(parents=True, exist_ok=True) + line = (json.dumps(entry, sort_keys=True) + "\n").encode("utf-8") + fd = os.open(str(audit_log), os.O_WRONLY | os.O_CREAT | os.O_APPEND, 0o600) + try: + os.write(fd, line) + finally: + os.close(fd) + os.chmod(audit_log, 0o600) + + +def _audit_attempt( + audit_log: Path, + action: str, + *, + question_id: str, + operator: str, + detail: dict[str, Any] | None = None, +) -> None: + """Record the *intent* to perform a destructive action BEFORE it mutates. + + §3.3.1 requires every destructive action to be audit-logged. Writing the + attempt before the ledger mutation closes the "mutation applied with no + audit record" gap: if this append fails (e.g. an unwritable audit path) it + raises before any ledger row is touched, so the action aborts cleanly with + nothing changed. The matching :func:`_audit_outcome` records what happened. + """ + _append_audit( + audit_log, + { + "ts": _utc_now_iso(), + "action": action, + "phase": "attempt", + "question_id": question_id, + "operator": operator, + **(detail or {}), + }, + ) + + +def _audit_outcome( + audit_log: Path, + action: str, + *, + question_id: str, + operator: str, + applied: bool, + detail: dict[str, Any] | None = None, +) -> None: + """Record the *result* of a destructive action AFTER it ran. + + Carries ``applied`` (did the compare-and-set change a row). Pairs with the + :func:`_audit_attempt` record written before the mutation, so even if this + outcome append fails the attempt already proves the action was made. + """ + _append_audit( + audit_log, + { + "ts": _utc_now_iso(), + "action": action, + "phase": "outcome", + "question_id": question_id, + "operator": operator, + "applied": applied, + **(detail or {}), + }, + ) + + +def _require_confirm(action: str, *, confirm: bool) -> None: + """Raise unless a destructive ``action`` was explicitly confirmed. + + Mirrors the §3.3.1 rule: force-expire, answer-on-behalf, and force-resume + are audit-logged AND require an explicit confirmation flag. Failing closed + here means a typo can never silently mutate a live task's ledger row. + """ + if action in _DESTRUCTIVE_ACTIONS and not confirm: + raise PermissionError( + f"refusing destructive action '{action}' without --confirm " + f"(force-expire / answer-on-behalf / supersede are gated, §3.3.1)" + ) + + +def _fetch_question(conn: sqlite3.Connection, question_id: str) -> sqlite3.Row | None: + """Return the ledger row for ``question_id`` or ``None`` if absent.""" + return conn.execute( + "SELECT * FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + + +# --------------------------------------------------------------------------- # +# Subcommand handlers. Each returns a process exit code (0 ok, 1 op failure). +# --------------------------------------------------------------------------- # + + +def _cmd_init_db(args: argparse.Namespace, *, out: Any) -> int: + """Create/upgrade the agent-team tables (idempotent).""" + init_db(args.db) + print(f"initialized ledger DB at {args.db}", file=out) + return 0 + + +def _cmd_list(args: argparse.Namespace, *, out: Any) -> int: + """List pending questions, optionally filtered by status. + + Default lists ``open`` questions (the operator's "what is waiting" view). + ``--status STATE`` narrows to one lifecycle state; ``--all`` lists every + state. ``--parked`` is a convenience alias that surfaces the parked-task + context an operator chases: questions that are no longer ``open`` (answered + but never resumed, expired, or superseded) and so may back a parked task. + """ + conn = connect(args.db) + try: + if args.all: + rows = conn.execute( + "SELECT * FROM pending_questions ORDER BY thread_id, turn" + ).fetchall() + elif args.parked: + placeholders = ",".join("?" for _ in _PARKED_STATES) + rows = conn.execute( + f"SELECT * FROM pending_questions WHERE status IN ({placeholders}) " + "ORDER BY thread_id, turn", + tuple(_PARKED_STATES), + ).fetchall() + else: + rows = conn.execute( + "SELECT * FROM pending_questions WHERE status = ? " + "ORDER BY thread_id, turn", + (args.status,), + ).fetchall() + finally: + conn.close() + + payload = [_row_to_dict(row) for row in rows] + print(json.dumps(payload, indent=2, sort_keys=True), file=out) + return 0 + + +def _cmd_show(args: argparse.Namespace, *, out: Any) -> int: + """Print one question row by ``question_id`` (pure read).""" + conn = connect(args.db) + try: + row = _fetch_question(conn, args.question_id) + finally: + conn.close() + if row is None: + print(f"no such question: {args.question_id}", file=sys.stderr) + return 1 + print(json.dumps(_row_to_dict(row), indent=2, sort_keys=True), file=out) + return 0 + + +def _cmd_expire(args: argparse.Namespace, *, out: Any) -> int: + """Force-expire an ``open`` question (destructive; audit-logged). + + Audits the attempt BEFORE mutating so a mutation can never land without a + trail (§3.3.1); records the outcome after. + """ + _require_confirm("expire", confirm=args.confirm) + _audit_attempt( + args.audit_log, "expire", question_id=args.question_id, operator=args.operator + ) + conn = connect(args.db) + try: + changed = expire_question(conn, question_id=args.question_id) + finally: + conn.close() + _audit_outcome( + args.audit_log, + "expire", + question_id=args.question_id, + operator=args.operator, + applied=changed, + ) + if not changed: + print( + f"expire no-op: question {args.question_id} was not 'open' " + "(already answered/expired/superseded or absent)", + file=sys.stderr, + ) + return 1 + print(f"expired question {args.question_id}", file=out) + return 0 + + +def _cmd_redeliver(args: argparse.Namespace, *, out: Any) -> int: + """Clear an ``open`` question's ``channel_ref`` so it is re-posted (§3.3.1). + + The design's "re-deliver" operator action. Re-delivery itself is performed + by the transport reconcile loop; clearing ``channel_ref`` makes that loop + re-post and record a fresh ref. Idempotent and non-destructive (the question + stays ``open``), so it needs no ``--confirm`` — but it is audit-logged for + provenance. Returns ``1`` if the question is absent or not ``open``. + """ + _audit_attempt( + args.audit_log, + "redeliver", + question_id=args.question_id, + operator=args.operator, + ) + conn = connect(args.db) + try: + row = _fetch_question(conn, args.question_id) + if row is None: + applied = False + prior_ref = None + status = None + elif row["status"] != "open": + applied = False + prior_ref = row["channel_ref"] + status = row["status"] + else: + prior_ref = row["channel_ref"] + status = "open" + conn.execute( + "UPDATE pending_questions SET channel_ref=NULL WHERE question_id=?", + (args.question_id,), + ) + applied = True + finally: + conn.close() + _audit_outcome( + args.audit_log, + "redeliver", + question_id=args.question_id, + operator=args.operator, + applied=applied, + detail={"prior_channel_ref": prior_ref}, + ) + if not applied: + reason = "absent" if status is None else f"status={status}, not open" + print( + f"redeliver no-op: question {args.question_id} ({reason}); " + "nothing to re-post", + file=sys.stderr, + ) + return 1 + print( + f"cleared channel_ref for {args.question_id}; reconcile loop will re-post", + file=out, + ) + return 0 + + +def _cmd_answer(args: argparse.Namespace, *, out: Any) -> int: + """Answer a question on a task's behalf (destructive; audit-logged). + + Uses the foundation first-answer-wins compare-and-set: succeeds only if the + question is still ``open``. ``--answer`` is stored verbatim as the answer + payload string; ``--via`` records the answering identity for the audit + trail. The audit log records the operator regardless of outcome. + """ + _require_confirm("answer", confirm=args.confirm) + via = args.via or f"cli:{args.operator}" + _audit_attempt( + args.audit_log, + "answer", + question_id=args.question_id, + operator=args.operator, + detail={"answered_via": via}, + ) + conn = connect(args.db) + try: + changed = answer_question( + conn, + question_id=args.question_id, + answer_json=args.answer, + answered_via=via, + ) + finally: + conn.close() + _audit_outcome( + args.audit_log, + "answer", + question_id=args.question_id, + operator=args.operator, + applied=changed, + detail={"answered_via": via}, + ) + if not changed: + print( + f"answer no-op: question {args.question_id} was not 'open' " + "(already answered/expired/superseded or absent)", + file=sys.stderr, + ) + return 1 + print(f"answered question {args.question_id} (via {via})", file=out) + return 0 + + +def _cmd_supersede(args: argparse.Namespace, *, out: Any) -> int: + """Mark a stale ``open``/``answered`` question ``superseded`` (destructive).""" + _require_confirm("supersede", confirm=args.confirm) + _audit_attempt( + args.audit_log, + "supersede", + question_id=args.question_id, + operator=args.operator, + ) + conn = connect(args.db) + try: + changed = supersede_question(conn, question_id=args.question_id) + finally: + conn.close() + _audit_outcome( + args.audit_log, + "supersede", + question_id=args.question_id, + operator=args.operator, + applied=changed, + ) + if not changed: + print( + f"supersede no-op: question {args.question_id} was not " + "'open'/'answered' (already expired/superseded or absent)", + file=sys.stderr, + ) + return 1 + print(f"superseded question {args.question_id}", file=out) + return 0 + + +def _cmd_force_resume(args: argparse.Namespace, *, out: Any) -> int: + """Force-resume a parked task's question (destructive; audit-logged). + + The design-named operator verb (§3.3.1 / §6.6 "an operator can force-resume + ... a parked task via the CLI"). A task parks when its clarifier question + EXPIRES with no answer, so the un-park action is to RE-OPEN that expired + question (:func:`agent_team.db.schema.reopen_question`) so the normal + delivery → answer → resume flow can proceed. + + Crucially this does NOT ``supersede`` the row: superseding an ``answered`` + row would flip it out of the state the recovery sweep resumes from, making a + stuck-but-answered task permanently un-resumable — the opposite of + force-resume. So: + + * ``expired`` (the parked case) → reopened; returns 0. + * ``answered`` (answered but not yet resumed) → already eligible for the + recovery resume sweep; intent is recorded and we report that, no mutation. + * ``open`` / ``superseded`` / absent → nothing to force; reported as a no-op. + """ + _require_confirm("force-resume", confirm=args.confirm) + _audit_attempt( + args.audit_log, + "force-resume", + question_id=args.question_id, + operator=args.operator, + detail={"resume_requested": True}, + ) + conn = connect(args.db) + try: + row = _fetch_question(conn, args.question_id) + status = None if row is None else row["status"] + reopened = False + if status == "expired": + reopened = reopen_question(conn, question_id=args.question_id) + finally: + conn.close() + _audit_outcome( + args.audit_log, + "force-resume", + question_id=args.question_id, + operator=args.operator, + applied=reopened, + detail={"resume_requested": True, "prior_status": status}, + ) + if reopened: + print( + f"force-resume: reopened expired question {args.question_id}; " + "it will be re-delivered for an answer", + file=out, + ) + return 0 + if status == "answered": + print( + f"force-resume: question {args.question_id} is answered and pending " + "resume; the recovery sweep will resume it (intent recorded)", + file=out, + ) + return 0 + print( + f"force-resume no-op: question {args.question_id} " + f"({'absent' if status is None else f'status={status}'}) is not parked", + file=sys.stderr, + ) + return 1 + + +# Statuses an operator treats as "parked context": a task whose only pending +# question is no longer open may be parked (answered-but-unresumed, expired, or +# superseded). ``open`` is excluded — that is the live-waiting view (default +# ``list``). Derived from the foundation QUESTION_STATES so it stays in sync. +_PARKED_STATES: tuple[str, ...] = tuple(s for s in QUESTION_STATES if s != "open") + + +def build_parser() -> argparse.ArgumentParser: + """Construct the argparse parser for ``run-team.py`` (no side effects).""" + parser = argparse.ArgumentParser( + prog="run-team.py", + description=( + "R720 agent-team operator CLI — manual path over the durable " + "pending_questions ledger (design §3.3.1)." + ), + ) + parser.add_argument( + "--db", + type=Path, + default=_DEFAULT_DB, + help=f"path to the agent-team SQLite ledger (default: {_DEFAULT_DB})", + ) + parser.add_argument( + "--audit-log", + type=Path, + default=_DEFAULT_AUDIT_LOG, + dest="audit_log", + help=( + "append-only JSONL audit log for destructive actions " + f"(default: {_DEFAULT_AUDIT_LOG})" + ), + ) + parser.add_argument( + "--operator", + default=_default_operator(), + help="operator identity recorded in the audit log for destructive actions " + "(defaults to the OS login so the trail is always attributable)", + ) + + sub = parser.add_subparsers(dest="command", required=True) + + p_init = sub.add_parser("init-db", help="create/upgrade the ledger tables") + p_init.set_defaults(func=_cmd_init_db) + + p_list = sub.add_parser("list", help="list pending questions (read-only)") + list_filter = p_list.add_mutually_exclusive_group() + list_filter.add_argument( + "--status", + choices=QUESTION_STATES, + default="open", + help="lifecycle status to list (default: open)", + ) + list_filter.add_argument( + "--all", + action="store_true", + help="list questions in every lifecycle status", + ) + list_filter.add_argument( + "--parked", + action="store_true", + help="list non-open questions (parked-task context)", + ) + p_list.set_defaults(func=_cmd_list) + + p_show = sub.add_parser("show", help="print one question row (read-only)") + p_show.add_argument("question_id", help="the question_id to show") + p_show.set_defaults(func=_cmd_show) + + p_redeliver = sub.add_parser( + "redeliver", + help="clear an open question's channel_ref so it is re-posted", + ) + p_redeliver.add_argument("question_id", help="the question_id to re-deliver") + p_redeliver.set_defaults(func=_cmd_redeliver) + + p_expire = sub.add_parser( + "expire", help="force-expire an open question (destructive)" + ) + p_expire.add_argument("question_id", help="the question_id to expire") + p_expire.add_argument( + "--confirm", + action="store_true", + help="required: confirm this destructive, audit-logged action", + ) + p_expire.set_defaults(func=_cmd_expire) + + p_answer = sub.add_parser( + "answer", help="answer a question on a task's behalf (destructive)" + ) + p_answer.add_argument("question_id", help="the question_id to answer") + p_answer.add_argument( + "--answer", + required=True, + help="the answer payload (stored verbatim as answer_json)", + ) + p_answer.add_argument( + "--via", + default="", + help="answering identity for answered_via (default: cli:)", + ) + p_answer.add_argument( + "--confirm", + action="store_true", + help="required: confirm this destructive, audit-logged action", + ) + p_answer.set_defaults(func=_cmd_answer) + + p_supersede = sub.add_parser( + "supersede", help="mark a stale question superseded (destructive)" + ) + p_supersede.add_argument("question_id", help="the question_id to supersede") + p_supersede.add_argument( + "--confirm", + action="store_true", + help="required: confirm this destructive, audit-logged action", + ) + p_supersede.set_defaults(func=_cmd_supersede) + + p_resume = sub.add_parser( + "force-resume", + help="force-resume a parked task's question (destructive)", + ) + p_resume.add_argument("question_id", help="the question_id to force-resume") + p_resume.add_argument( + "--confirm", + action="store_true", + help="required: confirm this destructive, audit-logged action", + ) + p_resume.set_defaults(func=_cmd_force_resume) + + return parser + + +def main(argv: Sequence[str] | None = None, *, out: Any = None) -> int: + """CLI entry point. Returns a process exit code. + + ``argv`` defaults to ``sys.argv[1:]``; ``out`` defaults to ``sys.stdout`` + (injectable for tests). Operational failures return ``1``; a missing + ``--confirm`` on a destructive action raises :class:`PermissionError`, + surfaced as exit code ``1`` with a stderr message. + """ + out = out if out is not None else sys.stdout + parser = build_parser() + args = parser.parse_args(argv) + try: + return int(args.func(args, out=out)) + except PermissionError as exc: + # A refused destructive action (no --confirm) or an unwritable audit + # path. The attempt-before-mutate ordering means nothing was mutated. + print(f"error: {exc}", file=sys.stderr) + return 1 + except OSError as exc: + # Any other audit-log / filesystem failure (e.g. the audit append could + # not be written). Surfaced cleanly instead of as an uncaught traceback; + # if the attempt record was written, the action is on the trail. + print(f"error: audit/IO failure: {exc}", file=sys.stderr) + return 1 + + +if __name__ == "__main__": # pragma: no cover + raise SystemExit(main()) diff --git a/agent-team/tests/__init__.py b/agent-team/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/agent-team/tests/conftest.py b/agent-team/tests/conftest.py new file mode 100644 index 0000000..900c650 --- /dev/null +++ b/agent-team/tests/conftest.py @@ -0,0 +1,13 @@ +"""Pytest configuration: make the ``agent_team`` package importable. + +Adds the ``agent-team/`` project root (the directory containing the +``agent_team`` package) to ``sys.path`` so tests can run without an editable +install. +""" + +import sys +from pathlib import Path + +_PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(_PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(_PROJECT_ROOT)) diff --git a/agent-team/tests/sim/conftest.py b/agent-team/tests/sim/conftest.py new file mode 100644 index 0000000..6077641 --- /dev/null +++ b/agent-team/tests/sim/conftest.py @@ -0,0 +1,65 @@ +"""Pytest fixtures for the P1 exit-criteria simulation (design §7.1 P1). + +The parent ``tests/conftest.py`` already puts the ``agent-team/`` project root +on ``sys.path`` so ``agent_team`` (the committed foundation) imports cleanly. +This sim-level conftest adds the local ``tests/sim`` directory to ``sys.path`` +so the sibling :mod:`harness` module imports the same way under any pytest +``rootdir``/import mode, and exposes the shared fixtures the four exit-criteria +tests build on. + +These fixtures construct the harness against the *real* committed foundation +(the SQLite ``pending_questions`` ledger and the atomic, integrity-checked +state store) rooted under a pytest ``tmp_path``, so every test gets an +isolated, durable, on-disk pipeline and nothing touches live infrastructure. +""" + +from __future__ import annotations + +import sys +from pathlib import Path +from typing import Iterator + +import pytest + +# Make the sibling ``harness`` module importable regardless of pytest's import +# mode / rootdir (the parent conftest handles the ``agent_team`` package root). +_SIM_DIR = Path(__file__).resolve().parent +if str(_SIM_DIR) not in sys.path: + sys.path.insert(0, str(_SIM_DIR)) + +from harness import ( # noqa: E402 (path injected above) + PostFailingTransport, + RecordingTransport, + SimClock, + SimPipeline, +) + + +@pytest.fixture +def clock() -> SimClock: + """A fresh advanceable simulation clock starting at tick 0.""" + return SimClock(start=0) + + +@pytest.fixture +def transport() -> RecordingTransport: + """An in-memory recording transport (faithful Transport subclass, no Slack).""" + return RecordingTransport(name="sim") + + +@pytest.fixture +def post_failing_transport() -> PostFailingTransport: + """A transport whose first ``post_question`` raises (lost-post simulation).""" + return PostFailingTransport(name="sim", fail_times=1) + + +@pytest.fixture +def pipeline( + tmp_path: Path, clock: SimClock, transport: RecordingTransport +) -> Iterator[SimPipeline]: + """A durable :class:`SimPipeline` rooted under an isolated ``tmp_path``. + + Backed by the real ``pending_questions`` SQLite ledger and the real atomic + state store; teardown is the ``tmp_path`` cleanup pytest already does. + """ + yield SimPipeline(tmp_path / "state", clock=clock, transport=transport) diff --git a/agent-team/tests/sim/harness.py b/agent-team/tests/sim/harness.py new file mode 100644 index 0000000..66f4a59 --- /dev/null +++ b/agent-team/tests/sim/harness.py @@ -0,0 +1,601 @@ +"""P1 exit-criteria simulation harness (design §7.1 P1, demonstrating §3.3.1). + +Phase P1 of the Plane-2 pipeline must *demonstrate* the durable +human-in-the-loop suspend/resume contract before anything else is built. The +four exit criteria (§7.1 P1) are: + +* (a) kill the box mid-wait and have the task resume after restart; +* (b) submit a duplicate answer and confirm it no-ops; +* (c) submit an answer after the deadline expired and confirm it is rejected + and the task parked; +* (d) two tasks suspended concurrently resume independently to the correct + thread. + +This module is a *simulation* harness, not the production pipeline. There is no +LangGraph runtime on the box yet (that lands in the P1 build proper), so the +harness stands in a minimal, faithful model of the riskiest mechanic — the +``pending_questions`` ledger and the §3.3.1 first-answer-wins / deadline-race / +turn-guarded-resume compare-and-set — *built on the real committed foundation*: + +* :mod:`agent_team.db.schema` — the real ``pending_questions`` ledger DDL and + the real ``answer_question`` / ``expire_question`` / ``supersede_question`` + ``BEGIN IMMEDIATE`` compare-and-set helpers. The harness never re-implements + the atomic statements; it drives the committed ones. +* :mod:`agent_team.state_store` — the real atomic, integrity-checked durable + state store. The "graph checkpoint" each task suspends on is written through + :func:`agent_team.state_store.write_checked` and read back through + :func:`agent_team.state_store.read_checked`, so a simulated "kill the box" + (drop the in-memory harness, reconstruct from disk) exercises real durable + recovery, not a Python dict. +* :mod:`agent_team.task_model` — the real :class:`TaskRecord` / :class:`Phase` + / :class:`TaskStatus` model and its JSON serialization. +* :mod:`agent_team.transport.base` — the real :class:`Transport` ABC and + :class:`QuestionSet` payload; :class:`RecordingTransport` is a faithful + in-memory adapter subclassing the committed contract (no live Slack). + +Nothing here provisions, schedules, or reaches live infrastructure. It is +pre-deployment scaffolding that proves the design's durable mechanic holds. +""" + +from __future__ import annotations + +import json +import sqlite3 +import uuid +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +from agent_team.db.schema import ( + answer_question, + connect, + expire_question, + init_db, + supersede_question, +) +from agent_team.state_store import IntegrityError, read_checked, write_checked +from agent_team.task_model import ( + Phase, + TaskRecord, + TaskStatus, + new_thread_id, + task_from_json, + task_to_json, +) +from agent_team.transport.base import QuestionSet, Transport + +__all__ = [ + "CHECKPOINT_SCHEMA_VERSION", + "PostFailingTransport", + "RecordingTransport", + "ResumeOutcome", + "SimClock", + "SimPipeline", + "SuspendedTask", +] + +# Schema version stamped on the harness's durable "graph checkpoint" sidecars. +# Distinct from the SQL ``SCHEMA_VERSION``; this versions the checkpoint blob +# format the harness round-trips through ``state_store``. +CHECKPOINT_SCHEMA_VERSION: int = 1 + + +class SimClock: + """A monotonically advanceable fake clock for deadline-race tests (§3.3.1). + + Each open question carries a ``deadline_at``. Rather than sleep in tests, + the harness compares a question's deadline against this clock's "now", and + tests advance the clock past a deadline to drive the timer loop. Times are + plain integer ticks (seconds since an arbitrary epoch); the ledger stores + them as ISO-like sortable strings so the durable column is human-readable. + """ + + def __init__(self, start: int = 0) -> None: + self._now = int(start) + + def now(self) -> int: + """Return the current tick.""" + return self._now + + def advance(self, ticks: int) -> int: + """Advance the clock by ``ticks`` and return the new now.""" + if ticks < 0: + raise ValueError("cannot advance the clock backwards") + self._now += int(ticks) + return self._now + + def stamp(self, tick: int | None = None) -> str: + """Render ``tick`` (default: now) as a sortable durable timestamp.""" + value = self._now if tick is None else int(tick) + # Zero-padded so lexical order == numeric order in the ledger column. + return f"t{value:020d}" + + +@dataclass +class RecordingTransport(Transport): + """In-memory :class:`Transport` adapter that records posts (no live Slack). + + A faithful subclass of the committed :class:`agent_team.transport.base. + Transport` contract: :meth:`post_question` embeds the ``question_id`` in + the returned ``channel_ref`` (mirroring the real "the post MUST embed the + question_id" rule) and records the post so recovery/reconcile tests can + inspect delivery. :meth:`parse_answer` normalizes a ``(question_id, + answer, via)`` raw payload, mapping it back via the embedded id. + """ + + name: str = "sim" + posts: list[dict[str, Any]] = field(default_factory=list) + + def post_question( + self, + *, + thread_id: str, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, + ) -> str: + channel_ref = f"{self.name}:{question_id}" + self.posts.append( + { + "thread_id": thread_id, + "question_id": question_id, + "turn": turn, + "questions": list(question_set.questions), + "deadline": deadline, + "channel_ref": channel_ref, + } + ) + return channel_ref + + def parse_answer(self, raw: Any) -> tuple[str, Any, str]: + question_id = raw["question_id"] + answer = raw.get("answer") + via = raw.get("via", self.name) + return question_id, answer, via + + +@dataclass +class PostFailingTransport(RecordingTransport): + """A transport whose first ``post_question`` raises (lost-post simulation). + + Used to exercise the §3.3.1 "if the post fails, the row stays ``open`` with + no ref and a reconcile loop retries idempotently" path. The first post + raises; subsequent posts succeed and record normally. + """ + + fail_times: int = 1 + _attempts: int = 0 + + def post_question( + self, + *, + thread_id: str, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, + ) -> str: + self._attempts += 1 + if self._attempts <= self.fail_times: + raise RuntimeError("simulated transport post failure") + return super().post_question( + thread_id=thread_id, + question_id=question_id, + turn=turn, + question_set=question_set, + deadline=deadline, + ) + + +@dataclass(frozen=True) +class SuspendedTask: + """Handle to a task suspended on a question (returned by :meth:`SimPipeline.submit`).""" + + thread_id: str + question_id: str + turn: int + + +@dataclass(frozen=True) +class ResumeOutcome: + """Result of attempting to resume a task on an answered question. + + ``resumed`` is True iff the turn guard passed and the graph advanced; + ``superseded`` is True iff a stale/redelivered resume was skipped (§3.3.1 + "a resume can never double-apply"). + """ + + thread_id: str + resumed: bool + superseded: bool + new_phase: Phase | None + + +class SimPipeline: + """A minimal, durable simulation of the §3.3.1 suspend/resume mechanic. + + The pipeline owns two real durable stores under ``root``: + + * the SQLite ``pending_questions`` ledger (via the committed + :mod:`agent_team.db.schema`), the single source of truth for the + question lifecycle, and + * one integrity-checked "graph checkpoint" file per task (via the + committed :mod:`agent_team.state_store`), holding the durable + :class:`agent_team.task_model.TaskRecord`. + + A task is submitted, suspends on a question (status ``WAITING_HUMAN``, + ledger row ``open``), and later resumes when a first valid answer wins the + compare-and-set. "Killing the box" is modelled by dropping the in-memory + pipeline and calling :meth:`reopen`, which reconstructs purely from the two + durable stores — proving recovery has no in-memory-only state. + """ + + def __init__(self, root: Path, *, clock: SimClock, transport: Transport) -> None: + self._root = Path(root) + self._db_path = self._root / "agent_team.sqlite" + self._checkpoints = self._root / "checkpoints" + self._clock = clock + self._transport = transport + self._root.mkdir(parents=True, exist_ok=True) + self._checkpoints.mkdir(parents=True, exist_ok=True) + init_db(self._db_path) + + # -- durable checkpoint helpers (real state_store) -------------------- + + def _checkpoint_path(self, thread_id: str) -> Path: + return self._checkpoints / f"{thread_id}.json" + + def _write_checkpoint(self, record: TaskRecord) -> None: + """Persist a task record through the real atomic state-store.""" + write_checked( + self._checkpoint_path(record.thread_id), + task_to_json(record).encode("utf-8"), + schema_version=CHECKPOINT_SCHEMA_VERSION, + ) + + def load_record(self, thread_id: str) -> TaskRecord: + """Read a task record back, integrity-checked (raises IntegrityError).""" + data = read_checked( + self._checkpoint_path(thread_id), + schema_version=CHECKPOINT_SCHEMA_VERSION, + ) + return task_from_json(data.decode("utf-8")) + + # -- ledger helpers (real db.schema connection) ----------------------- + + def _connect(self) -> sqlite3.Connection: + # Use the committed foundation connection helper (WAL + busy_timeout + + # the stashed db path the compare-and-set relies on) rather than a raw + # sqlite3.connect — so concurrent responders genuinely serialize on the + # write lock (§3.3.1) instead of racing without a busy timeout. + return connect(self._db_path) + + def ledger_row(self, question_id: str) -> sqlite3.Row | None: + """Return the durable ``pending_questions`` row for ``question_id``.""" + conn = self._connect() + conn.row_factory = sqlite3.Row + try: + return conn.execute( + "SELECT * FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + finally: + conn.close() + + # -- pipeline operations ---------------------------------------------- + + def submit(self, *, questions: list[str], deadline_in: int) -> SuspendedTask: + """Submit a task; it advances to CLARIFY and suspends on a question. + + Writes the ledger row ``open`` *first*, then posts to the transport and + stores the returned ``channel_ref`` (the §3.3.1 delivery order). If the + post fails the row stays ``open`` with no ref for the reconcile loop to + retry. The durable task record is checkpointed as ``WAITING_HUMAN``. + """ + thread_id = new_thread_id() + question_id = uuid.uuid4().hex + turn = 0 + deadline_at = self._clock.stamp(self._clock.now() + int(deadline_in)) + + record = TaskRecord( + thread_id=thread_id, + status=TaskStatus.WAITING_HUMAN, + current_phase=Phase.CLARIFY, + transport=getattr(self._transport, "name", "sim"), + created_at=self._clock.stamp(), + updated_at=self._clock.stamp(), + ) + self._write_checkpoint(record) + + # Ledger row first (open, no channel_ref yet). + conn = self._connect() + try: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, posted_at, deadline_at) " + "VALUES (?, ?, ?, 'open', ?, ?, ?)", + ( + question_id, + thread_id, + turn, + record.transport, + self._clock.stamp(), + deadline_at, + ), + ) + finally: + conn.close() + + # Then deliver; tolerate a lost post (row stays open, no ref). + self._deliver( + thread_id=thread_id, + question_id=question_id, + turn=turn, + questions=questions, + deadline_at=deadline_at, + ) + return SuspendedTask(thread_id=thread_id, question_id=question_id, turn=turn) + + def _deliver( + self, + *, + thread_id: str, + question_id: str, + turn: int, + questions: list[str], + deadline_at: str, + ) -> str | None: + """Post the question and store the channel_ref; tolerate post failure.""" + question_set = QuestionSet( + thread_id=thread_id, + question_id=question_id, + turn=turn, + questions=list(questions), + ) + try: + channel_ref = self._transport.post_question( + thread_id=thread_id, + question_id=question_id, + turn=turn, + question_set=question_set, + deadline=deadline_at, + ) + except Exception: + # Lost post: row stays open with no ref; reconcile retries later. + return None + conn = self._connect() + try: + conn.execute( + "UPDATE pending_questions SET channel_ref = ? WHERE question_id = ?", + (channel_ref, question_id), + ) + finally: + conn.close() + return channel_ref + + def reconcile(self, *, questions_by_qid: dict[str, list[str]]) -> int: + """Retry delivery for ``open`` rows lacking a ``channel_ref`` (§3.3.1). + + Returns the number of rows for which a (re)delivery now succeeded. + ``questions_by_qid`` supplies the question text per id (the harness + does not persist question text on the ledger, mirroring the design's + ledger schema which carries lifecycle, not prompt bodies). + """ + conn = self._connect() + conn.row_factory = sqlite3.Row + try: + rows = conn.execute( + "SELECT question_id, thread_id, turn, deadline_at " + "FROM pending_questions " + "WHERE status = 'open' AND channel_ref IS NULL" + ).fetchall() + finally: + conn.close() + redelivered = 0 + for row in rows: + ref = self._deliver( + thread_id=row["thread_id"], + question_id=row["question_id"], + turn=row["turn"], + questions=questions_by_qid.get(row["question_id"], []), + deadline_at=row["deadline_at"], + ) + if ref is not None: + redelivered += 1 + return redelivered + + def submit_answer(self, raw: Any) -> bool: + """Normalize ``raw`` via the transport and run the §3.3.1 compare-and-set. + + Returns ``True`` when this answer won the race (ledger rowcount 1 — the + first valid answer, a resume is now eligible) and ``False`` when it + lost (rowcount 0 — duplicate, late, or for a closed question, ignored). + Drives the *committed* :func:`agent_team.db.schema.answer_question` + ``BEGIN IMMEDIATE`` statement; the harness never re-implements it. + """ + question_id, answer, via = self._transport.parse_answer(raw) + conn = self._connect() + try: + return answer_question( + conn, + question_id=question_id, + answer_json=json.dumps(answer), + answered_via=via, + answered_at=self._clock.stamp(), + ) + finally: + conn.close() + + def run_deadline_sweep(self) -> list[str]: + """Expire every overdue ``open`` question and park its task (§3.3.1). + + A timer loop flips overdue ``open`` rows to ``expired`` via the + committed compare-and-set (:func:`agent_team.db.schema. + expire_question`) and applies the park policy: the task record flips to + :attr:`TaskStatus.PARKED` / :attr:`Phase.PARKED`. Returns the list of + ``question_id`` s expired by this sweep. An answer arriving for an + already-expired question will lose its own compare-and-set. + """ + now = self._clock.now() + conn = self._connect() + conn.row_factory = sqlite3.Row + try: + rows = conn.execute( + "SELECT question_id, thread_id, deadline_at " + "FROM pending_questions WHERE status = 'open'" + ).fetchall() + finally: + conn.close() + + expired: list[str] = [] + for row in rows: + if not self._is_overdue(row["deadline_at"], now): + continue + conn2 = self._connect() + try: + won = expire_question(conn2, question_id=row["question_id"]) + finally: + conn2.close() + if won: + expired.append(row["question_id"]) + self._park(row["thread_id"]) + return expired + + @staticmethod + def _is_overdue(deadline_at: str | None, now: int) -> bool: + """Decode a ``SimClock``-stamped deadline and test it against ``now``.""" + if not deadline_at: + return False + try: + deadline_tick = int(deadline_at.lstrip("t")) + except ValueError: + return False + return now >= deadline_tick + + def _park(self, thread_id: str) -> None: + """Flip a task record to PARKED durably (idempotent).""" + record = self.load_record(thread_id) + record.status = TaskStatus.PARKED + record.current_phase = Phase.PARKED + record.updated_at = self._clock.stamp() + self._write_checkpoint(record) + + def resume(self, thread_id: str, question_id: str) -> ResumeOutcome: + """Turn-guarded resume of an ``answered`` question (§3.3.1, single-flight). + + Mirrors the design's resume worker: before advancing the graph it + checks the live checkpoint is still interrupted on this turn. The + durable task record is the checkpoint here, so: + + * if the record is still ``WAITING_HUMAN`` on the answered question, the + graph advances (CLARIFY -> PLAN), the record is checkpointed + ``ACTIVE``, and ``resumed`` is True; + * if the record already advanced (a stale/redelivered resume), the + question is marked ``superseded`` via the committed compare-and-set + and the resume is skipped (``superseded`` True), so a resume can + never double-apply. + """ + row = self.ledger_row(question_id) + if row is None or row["status"] != "answered": + return ResumeOutcome( + thread_id=thread_id, resumed=False, superseded=False, new_phase=None + ) + + record = self.load_record(thread_id) + # Turn guard: only resume if still suspended on this turn/phase. + if ( + record.status is not TaskStatus.WAITING_HUMAN + or record.current_phase is not Phase.CLARIFY + ): + conn = self._connect() + try: + supersede_question(conn, question_id=question_id) + finally: + conn.close() + return ResumeOutcome( + thread_id=thread_id, + resumed=False, + superseded=True, + new_phase=record.current_phase, + ) + + # Apply the won answer into the durable Q&A history and advance a phase. + answer = json.loads(row["answer_json"]) if row["answer_json"] else None + record.qa_history.append( + {"question_id": question_id, "turn": row["turn"], "answer": answer} + ) + record.status = TaskStatus.ACTIVE + record.current_phase = Phase.PLAN + record.updated_at = self._clock.stamp() + self._write_checkpoint(record) + return ResumeOutcome( + thread_id=thread_id, + resumed=True, + superseded=False, + new_phase=Phase.PLAN, + ) + + # -- restart recovery ------------------------------------------------- + + def reopen(self) -> SimPipeline: + """Simulate "kill the box": return a fresh pipeline over the same disk. + + The new pipeline shares the durable SQLite ledger and the + integrity-checked checkpoints but holds *no* in-memory state, so any + recovery must come entirely from disk (§3.3.1 "No in-memory-only + state."). The transport and clock are re-used (a real restart would + re-instantiate adapters; reusing them keeps recorded posts visible to + the assertions). + """ + return SimPipeline(self._root, clock=self._clock, transport=self._transport) + + def startup_sweep( + self, *, questions_by_qid: dict[str, list[str]] + ) -> dict[str, Any]: + """Run the §3.3.1 startup convergence sweep after a restart. + + Concretely: (1) retry delivery for ``open`` rows lacking a ref; (2) + re-enqueue a resume for ``answered`` rows whose task is still suspended + on that turn (idempotent via the turn guard); (3) apply the deadline + policy for overdue ``open`` rows. Returns a summary of what converged. + """ + redelivered = self.reconcile(questions_by_qid=questions_by_qid) + + conn = self._connect() + conn.row_factory = sqlite3.Row + try: + answered = conn.execute( + "SELECT question_id, thread_id FROM pending_questions " + "WHERE status = 'answered'" + ).fetchall() + finally: + conn.close() + + resumed: list[str] = [] + for row in answered: + outcome = self.resume(row["thread_id"], row["question_id"]) + if outcome.resumed: + resumed.append(row["thread_id"]) + + expired = self.run_deadline_sweep() + return { + "redelivered": redelivered, + "resumed": resumed, + "expired": expired, + } + + +def assert_no_integrity_error(pipeline: SimPipeline, thread_id: str) -> TaskRecord: + """Load a record and surface :class:`IntegrityError` as an explicit failure. + + Convenience for tests that want the durable read to be part of the + assertion (the foundation fails closed on corruption rather than returning + junk). + """ + try: + return pipeline.load_record(thread_id) + except IntegrityError as exc: # pragma: no cover - defensive + raise AssertionError( + f"durable checkpoint failed integrity check: {exc}" + ) from exc diff --git a/agent-team/tests/sim/test_p1_exit_criteria.py b/agent-team/tests/sim/test_p1_exit_criteria.py new file mode 100644 index 0000000..d67b533 --- /dev/null +++ b/agent-team/tests/sim/test_p1_exit_criteria.py @@ -0,0 +1,397 @@ +"""P1 exit-criteria simulation tests (design §7.1 P1, demonstrating §3.3.1). + +Phase P1 may begin only once the durable human-in-the-loop suspend/resume +mechanic is *demonstrated*. §7.1 P1 lists four exit criteria; this module is +the executable demonstration of each, driving the committed foundation +(:mod:`agent_team.db.schema` compare-and-set helpers + the atomic, +integrity-checked :mod:`agent_team.state_store`) through the +:mod:`harness.SimPipeline`: + +* (a) kill the box mid-wait and have the task resume after restart; +* (b) submit a duplicate answer and confirm it no-ops; +* (c) submit an answer after the deadline expired and confirm it is rejected + and the task parked; +* (d) two tasks suspended concurrently resume independently to the correct + thread. + +Each criterion has its own test (and a couple of supporting tests for the +delivery/recovery edges §3.3.1 calls out). The tests assert on the *durable* +state — the ledger row status and the integrity-checked task record — so they +verify the real mechanic, not a harness convenience. +""" + +from __future__ import annotations + +import json +import threading + +import pytest + +from harness import ( + PostFailingTransport, + RecordingTransport, + SimClock, + SimPipeline, +) + +from agent_team.state_store import IntegrityError +from agent_team.task_model import Phase, TaskStatus + + +# --------------------------------------------------------------------------- +# Baseline: a single happy-path suspend/resume cycle. +# --------------------------------------------------------------------------- + + +def test_submit_suspends_task_with_open_ledger_row( + pipeline: SimPipeline, transport: RecordingTransport +) -> None: + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + + record = pipeline.load_record(suspended.thread_id) + assert record.status is TaskStatus.WAITING_HUMAN + assert record.current_phase is Phase.CLARIFY + + row = pipeline.ledger_row(suspended.question_id) + assert row is not None + assert row["status"] == "open" + assert row["thread_id"] == suspended.thread_id + # Delivery happened: a channel_ref was stored and it embeds the question id. + assert row["channel_ref"] == f"sim:{suspended.question_id}" + assert ( + transport.posts and transport.posts[0]["question_id"] == suspended.question_id + ) + + +def test_first_answer_wins_and_resumes_to_plan(pipeline: SimPipeline) -> None: + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + + won = pipeline.submit_answer( + {"question_id": suspended.question_id, "answer": "core-api", "via": "slack:U1"} + ) + assert won is True + + outcome = pipeline.resume(suspended.thread_id, suspended.question_id) + assert outcome.resumed is True + assert outcome.new_phase is Phase.PLAN + + record = pipeline.load_record(suspended.thread_id) + assert record.status is TaskStatus.ACTIVE + assert record.current_phase is Phase.PLAN + # The won answer was durably folded into the Q&A history. + assert record.qa_history == [ + {"question_id": suspended.question_id, "turn": 0, "answer": "core-api"} + ] + assert pipeline.ledger_row(suspended.question_id)["status"] == "answered" + + +# --------------------------------------------------------------------------- +# (a) kill the box mid-wait and have the task resume after restart. +# --------------------------------------------------------------------------- + + +def test_a_restart_mid_wait_then_answer_and_resume( + pipeline: SimPipeline, transport: RecordingTransport +) -> None: + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + + # "Kill the box": drop the in-memory pipeline; rebuild purely from disk. + reopened = pipeline.reopen() + + # Durable state survived: ledger row still open, record still WAITING_HUMAN. + row = reopened.ledger_row(suspended.question_id) + assert row is not None and row["status"] == "open" + record = reopened.load_record(suspended.thread_id) + assert record.status is TaskStatus.WAITING_HUMAN + + # The human answers after the restart; the task converges via the sweep. + assert reopened.submit_answer( + {"question_id": suspended.question_id, "answer": "core-api", "via": "slack:U1"} + ) + summary = reopened.startup_sweep( + questions_by_qid={suspended.question_id: ["which repo?"]} + ) + assert summary["resumed"] == [suspended.thread_id] + + resumed_record = reopened.load_record(suspended.thread_id) + assert resumed_record.status is TaskStatus.ACTIVE + assert resumed_record.current_phase is Phase.PLAN + + +def test_a_restart_after_answer_recovers_via_startup_sweep( + pipeline: SimPipeline, +) -> None: + """An answer that won *before* the crash must still resume on restart.""" + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + assert pipeline.submit_answer( + {"question_id": suspended.question_id, "answer": "core-api", "via": "slack:U1"} + ) + + # Crash before the resume worker ran; recover from disk only. + reopened = pipeline.reopen() + # Pre-sweep the record is still suspended (resume never ran). + assert reopened.load_record(suspended.thread_id).status is TaskStatus.WAITING_HUMAN + + summary = reopened.startup_sweep( + questions_by_qid={suspended.question_id: ["which repo?"]} + ) + assert summary["resumed"] == [suspended.thread_id] + assert reopened.load_record(suspended.thread_id).current_phase is Phase.PLAN + + +# --------------------------------------------------------------------------- +# (b) submit a duplicate answer and confirm it no-ops. +# --------------------------------------------------------------------------- + + +def test_b_duplicate_answer_no_ops(pipeline: SimPipeline) -> None: + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + raw = { + "question_id": suspended.question_id, + "answer": "core-api", + "via": "slack:U1", + } + + first = pipeline.submit_answer(raw) + second = pipeline.submit_answer(raw) # exact redelivery / double click + third = pipeline.submit_answer( + { + "question_id": suspended.question_id, + "answer": "other-repo", + "via": "github:U2", + } + ) # a different answer via a second channel + + assert first is True + assert second is False + assert third is False + + # The ledger preserved the *first* answer; later ones never overwrote it. + row = pipeline.ledger_row(suspended.question_id) + assert row["status"] == "answered" + assert json.loads(row["answer_json"]) == "core-api" + assert row["answered_via"] == "slack:U1" + + +def test_b_resume_is_single_apply_under_redelivered_resume( + pipeline: SimPipeline, +) -> None: + """Even if the resume worker is invoked twice, it applies exactly once.""" + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + assert pipeline.submit_answer( + {"question_id": suspended.question_id, "answer": "core-api", "via": "slack:U1"} + ) + + first = pipeline.resume(suspended.thread_id, suspended.question_id) + second = pipeline.resume(suspended.thread_id, suspended.question_id) + + assert first.resumed is True + assert second.resumed is False + assert second.superseded is True # turn guard caught the stale resume + + # The phase advanced exactly one step; the Q&A history has one entry. + record = pipeline.load_record(suspended.thread_id) + assert record.current_phase is Phase.PLAN + assert len(record.qa_history) == 1 + assert pipeline.ledger_row(suspended.question_id)["status"] == "superseded" + + +# --------------------------------------------------------------------------- +# (c) answer after the deadline -> rejected, task parked. +# --------------------------------------------------------------------------- + + +def test_c_late_answer_rejected_and_task_parked( + pipeline: SimPipeline, clock: SimClock +) -> None: + suspended = pipeline.submit(questions=["which repo?"], deadline_in=50) + + # Time passes beyond the deadline; the timer loop expires + parks. + clock.advance(51) + expired = pipeline.run_deadline_sweep() + assert expired == [suspended.question_id] + + parked = pipeline.load_record(suspended.thread_id) + assert parked.status is TaskStatus.PARKED + assert parked.current_phase is Phase.PARKED + + # A late answer loses the compare-and-set against the now-expired row. + won = pipeline.submit_answer( + {"question_id": suspended.question_id, "answer": "core-api", "via": "slack:U1"} + ) + assert won is False + row = pipeline.ledger_row(suspended.question_id) + assert row["status"] == "expired" + assert row["answer_json"] is None + + # A resume attempt on the expired question does nothing. + outcome = pipeline.resume(suspended.thread_id, suspended.question_id) + assert outcome.resumed is False + + +def test_c_deadline_vs_answer_race_answer_first_wins( + pipeline: SimPipeline, clock: SimClock +) -> None: + """If the answer lands before the sweep, the sweep must not expire it.""" + suspended = pipeline.submit(questions=["which repo?"], deadline_in=50) + assert pipeline.submit_answer( + {"question_id": suspended.question_id, "answer": "core-api", "via": "slack:U1"} + ) + + clock.advance(99) # well past the deadline + expired = pipeline.run_deadline_sweep() + # The question is already 'answered', so the sweep finds nothing to expire. + assert expired == [] + assert pipeline.ledger_row(suspended.question_id)["status"] == "answered" + + outcome = pipeline.resume(suspended.thread_id, suspended.question_id) + assert outcome.resumed is True + + +# --------------------------------------------------------------------------- +# (d) two tasks suspended concurrently resume independently to the correct thread. +# --------------------------------------------------------------------------- + + +def test_d_two_concurrent_tasks_resume_to_correct_thread( + pipeline: SimPipeline, +) -> None: + first = pipeline.submit(questions=["repo for A?"], deadline_in=100) + second = pipeline.submit(questions=["repo for B?"], deadline_in=100) + + assert first.thread_id != second.thread_id + assert first.question_id != second.question_id + + # Answer the second task first, with a distinct answer. + assert pipeline.submit_answer( + {"question_id": second.question_id, "answer": "repo-B", "via": "slack:U2"} + ) + assert pipeline.submit_answer( + {"question_id": first.question_id, "answer": "repo-A", "via": "slack:U1"} + ) + + out_a = pipeline.resume(first.thread_id, first.question_id) + out_b = pipeline.resume(second.thread_id, second.question_id) + assert out_a.resumed and out_b.resumed + + rec_a = pipeline.load_record(first.thread_id) + rec_b = pipeline.load_record(second.thread_id) + + # Each thread carries *its own* answer — no cross-contamination. + assert rec_a.qa_history[0]["answer"] == "repo-A" + assert rec_b.qa_history[0]["answer"] == "repo-B" + assert rec_a.current_phase is Phase.PLAN + assert rec_b.current_phase is Phase.PLAN + + +def test_d_concurrent_responders_only_one_wins_per_question( + pipeline: SimPipeline, +) -> None: + """Two threads racing the same question: exactly one compare-and-set wins. + + Exercises the §3.3.1 ``BEGIN IMMEDIATE`` serialization in the committed + ``answer_question`` helper under real OS threads against one SQLite file. + """ + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + + barrier = threading.Barrier(2) + results: list[bool] = [] + errors: list[BaseException] = [] + lock = threading.Lock() + + def race(via: str) -> None: + barrier.wait() + try: + won = pipeline.submit_answer( + {"question_id": suspended.question_id, "answer": via, "via": via} + ) + except BaseException as exc: # noqa: BLE001 - record, must be empty + with lock: + errors.append(exc) + return + with lock: + results.append(won) + + threads = [threading.Thread(target=race, args=(f"slack:U{i}",)) for i in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + # No lock error is tolerated: the committed BEGIN IMMEDIATE + busy_timeout + # must *serialize* the responders, so the single winner is the compare-and + # -set, not a swallowed OperationalError loser (the bug the prior version + # masked). + assert errors == [], f"compare-and-set raised under contention: {errors!r}" + assert sum(1 for r in results if r) == 1 # exactly one winner + assert results.count(False) == 1 # the other genuinely lost the CAS (rowcount 0) + assert pipeline.ledger_row(suspended.question_id)["status"] == "answered" + + +def test_d_concurrent_tasks_survive_restart_independently( + pipeline: SimPipeline, +) -> None: + """Two suspended tasks + a crash: each converges to its own thread.""" + first = pipeline.submit(questions=["repo for A?"], deadline_in=100) + second = pipeline.submit(questions=["repo for B?"], deadline_in=100) + assert pipeline.submit_answer( + {"question_id": first.question_id, "answer": "repo-A", "via": "slack:U1"} + ) + + reopened = pipeline.reopen() + summary = reopened.startup_sweep( + questions_by_qid={ + first.question_id: ["repo for A?"], + second.question_id: ["repo for B?"], + } + ) + # Only the answered task resumes; the still-open one stays suspended. + assert summary["resumed"] == [first.thread_id] + assert reopened.load_record(first.thread_id).current_phase is Phase.PLAN + assert reopened.load_record(second.thread_id).status is TaskStatus.WAITING_HUMAN + + +# --------------------------------------------------------------------------- +# Supporting §3.3.1 edges: lost-post delivery + durable integrity. +# --------------------------------------------------------------------------- + + +def test_lost_post_leaves_open_row_then_reconcile_redelivers( + tmp_path_factory: pytest.TempPathFactory, + clock: SimClock, + post_failing_transport: PostFailingTransport, +) -> None: + pipeline = SimPipeline( + tmp_path_factory.mktemp("lostpost") / "state", + clock=clock, + transport=post_failing_transport, + ) + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + + # The post failed: the row is open with no channel_ref (no in-flight loss). + row = pipeline.ledger_row(suspended.question_id) + assert row["status"] == "open" + assert row["channel_ref"] is None + assert post_failing_transport.posts == [] + + # Reconcile retries idempotently; the second attempt succeeds. + redelivered = pipeline.reconcile( + questions_by_qid={suspended.question_id: ["which repo?"]} + ) + assert redelivered == 1 + row = pipeline.ledger_row(suspended.question_id) + assert row["channel_ref"] == f"sim:{suspended.question_id}" + + +def test_durable_checkpoint_is_integrity_checked( + pipeline: SimPipeline, +) -> None: + """Corrupting the durable checkpoint must fail closed, not return junk.""" + suspended = pipeline.submit(questions=["which repo?"], deadline_in=100) + checkpoint = pipeline._checkpoint_path(suspended.thread_id) # noqa: SLF001 + + # Tamper with the payload after the integrity sidecar was written. + checkpoint.write_bytes(checkpoint.read_bytes() + b"tampered") + + with pytest.raises(IntegrityError): + pipeline.load_record(suspended.thread_id) diff --git a/agent-team/tests/sim/test_p1_graph_integration.py b/agent-team/tests/sim/test_p1_graph_integration.py new file mode 100644 index 0000000..b547588 --- /dev/null +++ b/agent-team/tests/sim/test_p1_graph_integration.py @@ -0,0 +1,258 @@ +"""P1 exit criteria proven against the REAL LangGraph graph (design §7.1, §3.3.1). + +The sibling ``test_p1_exit_criteria.py`` proves the four §7.1 P1 criteria against +a faithful *model* of the ledger/state-store layer. This module proves the same +four criteria against the **actual** mechanic the design's P1 gate requires: + +* the real :mod:`agent_team.graph` ``StateGraph`` with a real + ``interrupt()`` / ``Command(resume=...)`` clarifier human gate, and +* the real ``langgraph.checkpoint.sqlite.SqliteSaver`` durable checkpointer + (design D9), so "kill the box" is modelled by dropping the saver/connection + and rebuilding the graph over the same checkpoint DB file, plus +* the real committed ``pending_questions`` ledger compare-and-set + (``answer_question`` / ``expire_question``) keyed by the graph's own + ``question_id``. + +The integration driver here mirrors what the responder/resume-worker do: write +the ledger row when the graph suspends, win the first-answer-wins compare-and-set +before resuming, and only resume a thread that still has a live interrupt (the +turn guard). Nothing is provisioned or networked; this is pre-deploy scaffolding. +""" + +from __future__ import annotations + +import sqlite3 +from pathlib import Path +from typing import Any + +import pytest + +# The durable SQLite checkpointer is design-required (D9); skip cleanly if the +# optional package is absent so the rest of the suite still runs. +SqliteSaver = pytest.importorskip("langgraph.checkpoint.sqlite").SqliteSaver + +from agent_team.db.schema import ( # noqa: E402 - after importorskip by design + answer_question, + connect, + expire_question, + init_db, +) +from agent_team.graph import ( # noqa: E402 + build_graph, + get_pipeline_state, + pending_question, + resume_task, + start_task, +) +from agent_team.task_model import Phase, TaskStatus # noqa: E402 + + +class _Pipeline: + """Thin integration of the real graph + real SqliteSaver + real ledger. + + Owns two on-disk SQLite files under ``root``: the LangGraph checkpoint DB + (driven by ``SqliteSaver``) and the committed ``pending_questions`` ledger. + The graph/checkpointer can be rebuilt over the same checkpoint file to model + a restart. + """ + + def __init__(self, root: Path) -> None: + self._ckpt_db = root / "checkpoints.sqlite" + self._ledger_db = root / "agent_team.sqlite" + init_db(self._ledger_db) + self._conn: sqlite3.Connection | None = None + self.graph = self._boot() + + def _boot(self) -> Any: + """(Re)build the graph + checkpointer over the same checkpoint DB file.""" + if self._conn is not None: + self._conn.close() + self._conn = sqlite3.connect(str(self._ckpt_db), check_same_thread=False) + saver = SqliteSaver(self._conn) + saver.setup() + return build_graph(checkpointer=saver) + + def restart(self) -> None: + """Model "kill the box": drop the saver/connection, rebuild from disk.""" + self.graph = self._boot() + + def start(self, *, deadline_at: str = "t9999", transport: str = "slack") -> str: + """Start a task, suspend on the clarifier, and ledger the question row.""" + thread_id, _ = start_task(self.graph, transport=transport) + payload = pending_question(self.graph, thread_id=thread_id) + assert payload is not None, "task did not suspend on the human gate" + conn = connect(self._ledger_db) + try: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, posted_at, " + "deadline_at) VALUES (?, ?, ?, 'open', ?, ?, ?)", + ( + payload["question_id"], + thread_id, + payload["turn"], + payload["transport"], + "t0", + deadline_at, + ), + ) + finally: + conn.close() + return thread_id + + def question_id(self, thread_id: str) -> str | None: + payload = pending_question(self.graph, thread_id=thread_id) + return None if payload is None else payload["question_id"] + + def answer(self, question_id: str, value: Any, *, via: str = "slack") -> bool: + """Run the first-answer-wins compare-and-set; True iff this call won.""" + conn = connect(self._ledger_db) + try: + return answer_question( + conn, + question_id=question_id, + answer_json=f'{{"v": "{value}"}}', + answered_via=via, + ) + finally: + conn.close() + + def expire(self, question_id: str) -> bool: + conn = connect(self._ledger_db) + try: + return expire_question(conn, question_id=question_id) + finally: + conn.close() + + def ledger_status(self, question_id: str) -> str | None: + conn = connect(self._ledger_db) + try: + row = conn.execute( + "SELECT status FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + finally: + conn.close() + return None if row is None else row["status"] + + def resume_if_won(self, thread_id: str, question_id: str, value: Any) -> bool: + """Mirror the resume worker: win the CAS, then turn-guarded resume. + + Resumes the graph only if (1) this call won the first-answer-wins + compare-and-set AND (2) the thread still has a live interrupt (the turn + guard — a thread already past the gate is never double-resumed). + Returns True iff the graph was actually advanced. + """ + if not self.answer(question_id, value): + return False + if self.question_id(thread_id) is None: + return False # already advanced; do not double-apply + resume_task(self.graph, thread_id=thread_id, answer={"v": value}) + return True + + def state(self, thread_id: str) -> Any: + return get_pipeline_state(self.graph, thread_id=thread_id) + + +@pytest.fixture() +def pipeline(tmp_path: Path) -> _Pipeline: + return _Pipeline(tmp_path) + + +# --- (a) kill the box mid-wait, resume after restart ------------------------ + + +def test_a_suspend_survives_restart_and_resumes(pipeline: _Pipeline) -> None: + thread_id = pipeline.start() + qid = pipeline.question_id(thread_id) + assert qid is not None + + pipeline.restart() # drop saver + connection; rebuild over the same DB file + + # The interrupt persisted across the restart (durable checkpoint, D9). + assert pipeline.question_id(thread_id) == qid + + assert pipeline.resume_if_won(thread_id, qid, "repo-x") is True + state = pipeline.state(thread_id) + assert state["current_phase"] == Phase.DONE.value + assert state["status"] == TaskStatus.DONE.value + assert state["qa_history"][0]["question_id"] == qid # identity held across resume + assert pipeline.question_id(thread_id) is None + + +# --- (b) duplicate answer no-ops -------------------------------------------- + + +def test_b_duplicate_answer_noops(pipeline: _Pipeline) -> None: + thread_id = pipeline.start() + qid = pipeline.question_id(thread_id) + + # First answer wins the CAS and drives the real resume to completion. + assert pipeline.resume_if_won(thread_id, qid, "first") is True + state_after_first = pipeline.state(thread_id) + assert state_after_first["current_phase"] == Phase.DONE.value + assert len(state_after_first["qa_history"]) == 1 + + # A duplicate answer loses the compare-and-set; no second resume, no change. + assert pipeline.resume_if_won(thread_id, qid, "second") is False + state_after_dup = pipeline.state(thread_id) + assert state_after_dup["qa_history"] == state_after_first["qa_history"] + assert state_after_dup["qa_history"][0]["answer"] == {"v": "first"} + + +# --- (c) answer after deadline rejected, task not resumed ------------------- + + +def test_c_post_deadline_answer_rejected(pipeline: _Pipeline) -> None: + thread_id = pipeline.start(deadline_at="t1") + qid = pipeline.question_id(thread_id) + + # The deadline timer wins the open->expired compare-and-set first. + assert pipeline.expire(qid) is True + assert pipeline.ledger_status(qid) == "expired" + + # A late answer loses its compare-and-set, so no resume fires... + assert pipeline.resume_if_won(thread_id, qid, "too-late") is False + # ...and the task is still suspended on the human gate (never advanced). + assert pipeline.question_id(thread_id) == qid + state = pipeline.state(thread_id) + assert state["current_phase"] == Phase.CLARIFY.value + + +# --- (d) two concurrent tasks resume independently to the correct thread ---- + + +def test_d_two_tasks_resume_to_correct_thread(pipeline: _Pipeline) -> None: + t1 = pipeline.start() + t2 = pipeline.start() + q1, q2 = pipeline.question_id(t1), pipeline.question_id(t2) + assert q1 != q2 # distinct identities per thread + + # Resume each with a distinct answer; each must land on its own thread only. + assert pipeline.resume_if_won(t1, q1, "answer-1") is True + assert pipeline.resume_if_won(t2, q2, "answer-2") is True + + s1, s2 = pipeline.state(t1), pipeline.state(t2) + assert s1["qa_history"][0]["answer"] == {"v": "answer-1"} + assert s2["qa_history"][0]["answer"] == {"v": "answer-2"} + assert s1["qa_history"][0]["question_id"] == q1 + assert s2["qa_history"][0]["question_id"] == q2 + assert s1["status"] == TaskStatus.DONE.value + assert s2["status"] == TaskStatus.DONE.value + + +def test_d_resume_after_completion_does_not_double_apply(pipeline: _Pipeline) -> None: + t1 = pipeline.start() + t2 = pipeline.start() + q1 = pipeline.question_id(t1) + + assert pipeline.resume_if_won(t1, q1, "once") is True + before_t2 = pipeline.state(t2) + + # A stale/redelivered resume for the already-completed t1 is a no-op (turn + # guard: no live interrupt), and never touches t2. + assert pipeline.resume_if_won(t1, q1, "again") is False + assert pipeline.state(t1)["qa_history"] == [ + {"turn": 0, "question_id": q1, "answer": {"v": "once"}} + ] + assert pipeline.state(t2) == before_t2 diff --git a/agent-team/tests/test_billing.py b/agent-team/tests/test_billing.py new file mode 100644 index 0000000..39097b8 --- /dev/null +++ b/agent-team/tests/test_billing.py @@ -0,0 +1,119 @@ +"""Unit tests for agent_team.billing (§3.1).""" + +from __future__ import annotations + +import pytest + +from agent_team import billing +from agent_team.billing import ( + BillingMode, + ClaudeResult, + claude_invoke, + resolve_mode, + set_invoker, +) + + +@pytest.fixture(autouse=True) +def _restore_invoker(): + """Restore the module invoker after each test.""" + original = billing._invoker + yield + billing._invoker = original + + +def test_resolve_mode_default_is_subscription(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("AGENT_TEAM_BILLING_MODE", raising=False) + assert resolve_mode(None) is BillingMode.SUBSCRIPTION + assert resolve_mode({}) is BillingMode.SUBSCRIPTION + + +def test_resolve_mode_from_config_string() -> None: + assert resolve_mode({"billing_mode": "api"}) is BillingMode.API + assert resolve_mode({"billing_mode": "BEDROCK"}) is BillingMode.BEDROCK + + +def test_resolve_mode_from_config_enum() -> None: + assert resolve_mode({"billing_mode": BillingMode.API}) is BillingMode.API + + +def test_resolve_mode_from_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AGENT_TEAM_BILLING_MODE", "bedrock") + assert resolve_mode(None) is BillingMode.BEDROCK + + +def test_resolve_mode_config_overrides_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AGENT_TEAM_BILLING_MODE", "bedrock") + assert resolve_mode({"billing_mode": "api"}) is BillingMode.API + + +def test_resolve_mode_invalid_raises() -> None: + with pytest.raises(ValueError): + resolve_mode({"billing_mode": "carrier-pigeon"}) + + +def test_claude_invoke_delegates_with_resolved_mode() -> None: + captured: dict = {} + + def fake(prompt: str, *, mode: BillingMode, **kw): + captured["prompt"] = prompt + captured["mode"] = mode + captured["kw"] = kw + return ClaudeResult(text="ok", mode=mode) + + set_invoker(fake) + result = claude_invoke("hi", mode=BillingMode.API, temperature=0.2) + assert result.text == "ok" + assert captured["mode"] is BillingMode.API + assert captured["prompt"] == "hi" + assert captured["kw"] == {"temperature": 0.2} + + +def test_subscription_mode_pops_stray_api_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-should-be-hidden") + seen: dict = {} + + def fake(prompt: str, *, mode: BillingMode, **kw): + import os + + seen["key_present"] = "ANTHROPIC_API_KEY" in os.environ + return ClaudeResult(text="ok", mode=mode) + + set_invoker(fake) + claude_invoke("hi", mode=BillingMode.SUBSCRIPTION) + assert seen["key_present"] is False + # Restored after the call. + import os + + assert os.environ.get("ANTHROPIC_API_KEY") == "sk-should-be-hidden" + + +def test_api_mode_does_not_pop_api_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-metered") + seen: dict = {} + + def fake(prompt: str, *, mode: BillingMode, **kw): + import os + + seen["key_present"] = "ANTHROPIC_API_KEY" in os.environ + return ClaudeResult(text="ok", mode=mode) + + set_invoker(fake) + claude_invoke("hi", mode=BillingMode.API) + assert seen["key_present"] is True + + +def test_subscription_with_no_key_is_safe(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + set_invoker(lambda prompt, *, mode, **kw: ClaudeResult(text="ok", mode=mode)) + assert claude_invoke("hi", mode=BillingMode.SUBSCRIPTION).text == "ok" + + +def test_unconfigured_invoker_raises() -> None: + billing._invoker = billing._unconfigured_invoker + with pytest.raises(RuntimeError): + claude_invoke("hi", mode=BillingMode.API) + + +def test_billing_mode_enum_members() -> None: + assert {m.name for m in BillingMode} == {"SUBSCRIPTION", "API", "BEDROCK"} diff --git a/agent-team/tests/test_builders.py b/agent-team/tests/test_builders.py new file mode 100644 index 0000000..23a820e --- /dev/null +++ b/agent-team/tests/test_builders.py @@ -0,0 +1,474 @@ +"""Unit tests for the Plane-2 builders node (design §3.3.2, §7.1 P3). + +Covers the two box-side halves of the §3.3.2 trust boundary this leaf owns: +the trust-control-surface denylist scan (boundary #2) and the diff integrity +hash (boundary #3), plus the LangGraph node's clean-vs-park state transitions. + +The tests import the committed foundation contracts (``billing``, +``state_store``, ``task_model``) verbatim and assert the leaf builds on them +without redefining them. +""" + +from __future__ import annotations + +import pytest + +from agent_team import billing +from agent_team.billing import BillingMode, ClaudeResult +from agent_team.nodes import builders +from agent_team.nodes.builders import ( + BuildError, + TrustBoundaryViolation, + build_candidate_diff, + builders_node, + default_diff_builder, + iter_diff_target_paths, + scan_trust_control_surface, +) +from agent_team.state_store import compute_content_hash +from agent_team.task_model import Phase, TaskStatus + +# --------------------------------------------------------------------------- # +# Diff fixtures +# --------------------------------------------------------------------------- # + +CLEAN_DIFF = """diff --git a/src/app.py b/src/app.py +--- a/src/app.py ++++ b/src/app.py +@@ -1,2 +1,2 @@ +-old = 1 ++new = 2 +""" + +NEW_FILE_DIFF = """diff --git a/src/util/helpers.py b/src/util/helpers.py +new file mode 100644 +--- /dev/null ++++ b/src/util/helpers.py +@@ -0,0 +1 @@ ++def f(): ... +""" + +WORKFLOW_DIFF = """diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml +--- a/.github/workflows/ci.yml ++++ b/.github/workflows/ci.yml +@@ -1 +1 @@ +-on: push ++on: [push, pull_request_target] +""" + + +def _scope_plan(diff_builder, scope=("src",), **extra): + """Build a plan dict whose declared scope is ``scope``.""" + plan = {"title": "t", "scope": list(scope), "phases": ["p1"]} + plan.update(extra) + return plan + + +@pytest.fixture(autouse=True) +def _restore_invoker(): + """Restore the billing invoker after each test (shared module state).""" + original = billing._invoker + yield + billing._invoker = original + + +# --------------------------------------------------------------------------- # +# Foundation imports are used verbatim (no redefinition) +# --------------------------------------------------------------------------- # + + +def test_imports_foundation_contracts_verbatim() -> None: + # The leaf imports, not redefines, the foundation symbols. + assert builders.compute_content_hash is compute_content_hash + assert builders.claude_invoke is billing.claude_invoke + assert builders.Phase is Phase + assert builders.TaskStatus is TaskStatus + + +# --------------------------------------------------------------------------- # +# default_diff_builder uses the §3.1 billing seam +# --------------------------------------------------------------------------- # + + +def test_default_diff_builder_calls_claude_invoke() -> None: + seen: dict = {} + + def fake(prompt: str, *, mode: BillingMode, **kw): + seen["prompt"] = prompt + return ClaudeResult(text=CLEAN_DIFF, mode=mode) + + billing.set_invoker(fake) + out = default_diff_builder(plan={"title": "x", "scope": ["src"]}, config=None) + assert out == CLEAN_DIFF + # The plan title and scope are surfaced into the build prompt. + assert "x" in seen["prompt"] + assert "src" in seen["prompt"] + + +def test_default_diff_builder_fails_loud_when_unwired() -> None: + # Foundation contract: the seam raises until an invoker is bound. + billing._invoker = billing._unconfigured_invoker + with pytest.raises(RuntimeError): + default_diff_builder(plan={"title": "x"}, config=None) + + +# --------------------------------------------------------------------------- # +# Diff parsing / canonicalization +# --------------------------------------------------------------------------- # + + +def test_iter_targets_reads_plus_header() -> None: + targets = iter_diff_target_paths(CLEAN_DIFF) + assert [t.path for t in targets] == ["src/app.py"] + + +def test_iter_targets_ignores_dev_null_source() -> None: + targets = iter_diff_target_paths(NEW_FILE_DIFF) + assert [t.path for t in targets] == ["src/util/helpers.py"] + + +def test_iter_targets_canonicalizes_dot_segments() -> None: + diff = "+++ b/src/./sub/../app.py\n" + targets = iter_diff_target_paths(diff) + assert [t.path for t in targets] == ["src/app.py"] + + +def test_iter_targets_dedupes() -> None: + diff = "+++ b/src/app.py\n+++ b/src/app.py\n" + assert len(iter_diff_target_paths(diff)) == 1 + + +# --------------------------------------------------------------------------- # +# Denylist: each surface (§3.3.2 boundary #2) +# --------------------------------------------------------------------------- # + + +def test_scan_clean_diff_is_empty() -> None: + assert scan_trust_control_surface(CLEAN_DIFF, scope=["src"]) == [] + + +def test_scan_flags_github_workflow() -> None: + v = scan_trust_control_surface(WORKFLOW_DIFF, scope=["src", ".github"]) + assert len(v) == 1 + assert ".github/workflows" in v[0].reason + + +@pytest.mark.parametrize( + "path", + [ + "infra/template.yaml", + "samconfig.toml".replace("toml", "yaml"), + "service/serverless.yml", + "cdk.json", + "iam/read-policy.json", + "policies/deploy.policy.yaml", + "CODEOWNERS", + ".github/CODEOWNERS", + ".github/dependabot.yml", + ".github/settings.yml", + ], +) +def test_scan_flags_trust_control_surface(path: str) -> None: + diff = f"+++ b/{path}\n" + # Declare a wide scope so the only possible failure is the denylist itself. + v = scan_trust_control_surface(diff, scope=[path.split("/")[0], "."]) + assert v, f"expected {path} to be denied" + assert v[0].path == path + + +def test_scan_denylist_beats_scope() -> None: + # A workflow file inside the declared scope is still denied. + v = scan_trust_control_surface(WORKFLOW_DIFF, scope=[".github"]) + assert len(v) == 1 + assert "workflow" in v[0].reason + + +# --------------------------------------------------------------------------- # +# Scope enforcement +# --------------------------------------------------------------------------- # + + +def test_scan_flags_out_of_scope_path() -> None: + diff = "+++ b/other/module.py\n" + v = scan_trust_control_surface(diff, scope=["src"]) + assert len(v) == 1 + assert "declared scope" in v[0].reason + + +def test_empty_scope_rejects_everything() -> None: + # An undeclared scope is a hard stop, not a wildcard. + v = scan_trust_control_surface(CLEAN_DIFF, scope=[]) + assert len(v) == 1 + assert "declared scope" in v[0].reason + + +def test_scope_prefix_match_is_boundary_safe() -> None: + # "src" must not accidentally allow "srcfoo/...". + diff = "+++ b/srcfoo/app.py\n" + v = scan_trust_control_surface(diff, scope=["src"]) + assert len(v) == 1 + assert "declared scope" in v[0].reason + + +# --------------------------------------------------------------------------- # +# Renames cannot launder a forbidden destination +# --------------------------------------------------------------------------- # + + +def test_rename_into_denied_path_is_flagged() -> None: + diff = ( + "diff --git a/src/app.py b/.github/workflows/evil.yml\n" + "similarity index 100%\n" + "rename from src/app.py\n" + "rename to .github/workflows/evil.yml\n" + ) + v = scan_trust_control_surface(diff, scope=["src", ".github"]) + assert len(v) == 1 + assert v[0].path == ".github/workflows/evil.yml" + assert v[0].rename_from == "src/app.py" + + +def test_rename_into_out_of_scope_is_flagged() -> None: + diff = "rename from src/app.py\nrename to other/app.py\n" + v = scan_trust_control_surface(diff, scope=["src"]) + assert len(v) == 1 + assert v[0].path == "other/app.py" + assert v[0].rename_from == "src/app.py" + + +# --------------------------------------------------------------------------- # +# Deletes / mode-changes / copies cannot bypass the scan +# (regression: header-only sections carry their path in ``diff --git``, not the +# ``+++ b/`` body line, so a ``+++``-only scan missed all four of these) +# --------------------------------------------------------------------------- # + + +def test_delete_of_denied_path_is_flagged() -> None: + diff = ( + "diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml\n" + "deleted file mode 100644\n" + "--- a/.github/workflows/ci.yml\n" + "+++ /dev/null\n" + "@@ -1 +0,0 @@\n" + "-on: push\n" + ) + v = scan_trust_control_surface(diff, scope=["src", ".github"]) + assert len(v) == 1 + assert v[0].path == ".github/workflows/ci.yml" + assert "workflow" in v[0].reason + + +def test_mode_change_only_on_denied_path_is_flagged() -> None: + # A chmod with no +++ line at all — only the diff --git header exists. + diff = ( + "diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml\n" + "old mode 100644\n" + "new mode 100755\n" + ) + v = scan_trust_control_surface(diff, scope=["src", ".github"]) + assert len(v) == 1 + assert v[0].path == ".github/workflows/ci.yml" + assert "workflow" in v[0].reason + + +def test_copy_into_denied_path_is_flagged() -> None: + # git ``copy to`` (not ``rename to``) must not launder a forbidden dest. + diff = ( + "diff --git a/src/x.py b/.github/workflows/evil.yml\n" + "similarity index 100%\n" + "copy from src/x.py\n" + "copy to .github/workflows/evil.yml\n" + ) + v = scan_trust_control_surface(diff, scope=["src", ".github"]) + assert [x.path for x in v] == [".github/workflows/evil.yml"] + assert v[0].rename_from == "src/x.py" + assert "workflow" in v[0].reason + + +def test_out_of_scope_delete_is_flagged() -> None: + diff = ( + "diff --git a/secret/key.py b/secret/key.py\n" + "deleted file mode 100644\n" + "--- a/secret/key.py\n" + "+++ /dev/null\n" + "@@ -1 +0,0 @@\n" + "-KEY = 1\n" + ) + v = scan_trust_control_surface(diff, scope=["src"]) + assert len(v) == 1 + assert v[0].path == "secret/key.py" + assert "scope" in v[0].reason + + +# --------------------------------------------------------------------------- # +# Box-side / CI-side denylist parity (security-review: box-side was narrower) +# --------------------------------------------------------------------------- # + + +@pytest.mark.parametrize( + "path", + [ + "infra/main.tf", + "infra/app-stack.ts", + "infra/app_stack.py", + "secrets/deploy.pem", + "signing.key", + "config/policy.json", + "roles/myiam.json", + ".github/actions/build/action.yml", + ], +) +def test_box_side_denylist_covers_ci_surface(path: str) -> None: + # These IaC/key families were caught by the CI guard but slipped the box-side + # scan. The box is the first backstop; it must reject them too (no auto-build). + diff = ( + f"diff --git a/{path} b/{path}\n--- a/{path}\n+++ b/{path}\n" + "@@ -1 +1 @@\n-x\n+y\n" + ) + scope = [path.split("/")[0], "."] + v = scan_trust_control_surface(diff, scope=scope) + assert v, f"expected {path} to be denied box-side" + + +# --------------------------------------------------------------------------- # +# Unsafe paths (absolute / parent-escaping / indirection) +# --------------------------------------------------------------------------- # + + +@pytest.mark.parametrize( + "raw", + [ + "/etc/passwd", + "../../../etc/passwd", + "..", + "src/../../escape.py", + "c:/windows/system32", # absolute-ish via backslash normalization below + ], +) +def test_scan_flags_unsafe_paths(raw: str) -> None: + diff = f"+++ b/{raw}\n" + v = scan_trust_control_surface(diff, scope=["src", "."]) + assert v, f"expected {raw} to be flagged" + assert any("unsafe" in x.reason or "scope" in x.reason for x in v) + + +def test_backslash_separator_is_normalized() -> None: + # A Windows separator must not smuggle a denied path past the POSIX matcher. + diff = "+++ b/.github\\workflows\\ci.yml\n" + v = scan_trust_control_surface(diff, scope=[".github"]) + assert len(v) == 1 + assert "workflow" in v[0].reason + + +# --------------------------------------------------------------------------- # +# build_candidate_diff: hashing (§3.3.2 boundary #3) + errors +# --------------------------------------------------------------------------- # + + +def _stub_builder(diff: str): + def _b(*, plan, config): + return diff + + return _b + + +def test_build_hash_matches_foundation_primitive() -> None: + plan = _scope_plan(None) + outcome = build_candidate_diff(plan, builder=_stub_builder(CLEAN_DIFF)) + assert outcome.diff == CLEAN_DIFF + assert outcome.diff_hash == compute_content_hash(CLEAN_DIFF.encode("utf-8")) + assert outcome.clean is True + + +def test_build_clean_outcome_has_no_violations() -> None: + plan = _scope_plan(None) + outcome = build_candidate_diff(plan, builder=_stub_builder(NEW_FILE_DIFF)) + assert outcome.violations == [] + assert outcome.clean is True + + +def test_build_violation_still_hashes() -> None: + plan = _scope_plan(None, scope=("src", ".github")) + outcome = build_candidate_diff(plan, builder=_stub_builder(WORKFLOW_DIFF)) + # Hash recorded for provenance even though the diff is rejected. + assert outcome.diff_hash == compute_content_hash(WORKFLOW_DIFF.encode("utf-8")) + assert outcome.clean is False + + +def test_build_rejects_non_mapping_plan() -> None: + with pytest.raises(BuildError): + build_candidate_diff(["not", "a", "mapping"], builder=_stub_builder(CLEAN_DIFF)) + + +def test_build_rejects_empty_diff() -> None: + plan = _scope_plan(None) + with pytest.raises(BuildError): + build_candidate_diff(plan, builder=_stub_builder(" \n ")) + + +def test_build_uses_default_builder_when_none() -> None: + billing.set_invoker( + lambda prompt, *, mode, **kw: ClaudeResult(text=CLEAN_DIFF, mode=mode) + ) + outcome = build_candidate_diff(_scope_plan(None)) + assert outcome.diff == CLEAN_DIFF + assert outcome.clean is True + + +# --------------------------------------------------------------------------- # +# builders_node: state transitions +# --------------------------------------------------------------------------- # + + +def test_node_clean_advances_to_verify() -> None: + state = {"plan": _scope_plan(None)} + update = builders_node(state, builder=_stub_builder(CLEAN_DIFF)) + assert update["candidate_diff"] == CLEAN_DIFF + assert update["diff_hash"] == compute_content_hash(CLEAN_DIFF.encode("utf-8")) + assert update["current_phase"] == Phase.VERIFY.value + assert update["status"] == TaskStatus.ACTIVE.value + assert "park_reason" not in update + + +def test_node_violation_parks_for_human_review() -> None: + state = {"plan": _scope_plan(None, scope=("src", ".github"))} + update = builders_node(state, builder=_stub_builder(WORKFLOW_DIFF)) + # Diff + hash recorded for provenance/ALARM, but parked — never auto-built. + assert update["candidate_diff"] == WORKFLOW_DIFF + assert update["current_phase"] == Phase.PARKED.value + assert update["status"] == TaskStatus.PARKED.value + assert "cross-review" in update["park_reason"] + assert "workflow" in update["park_reason"] + + +def test_node_out_of_scope_parks() -> None: + diff = "+++ b/other/x.py\n" + state = {"plan": _scope_plan(None, scope=("src",))} + update = builders_node(state, builder=_stub_builder(diff)) + assert update["status"] == TaskStatus.PARKED.value + assert "declared scope" in update["park_reason"] + + +def test_node_requires_plan() -> None: + with pytest.raises(BuildError): + builders_node({}, builder=_stub_builder(CLEAN_DIFF)) + + +def test_node_uses_default_builder_and_config() -> None: + captured: dict = {} + + def fake(prompt: str, *, mode: BillingMode, **kw): + captured["mode"] = mode + return ClaudeResult(text=CLEAN_DIFF, mode=mode) + + billing.set_invoker(fake) + state = {"plan": _scope_plan(None)} + update = builders_node(state, config={"billing_mode": "api"}) + assert update["current_phase"] == Phase.VERIFY.value + # The config-selected billing mode reached the seam. + assert captured["mode"] is BillingMode.API + + +def test_violation_dataclass_shape() -> None: + v = TrustBoundaryViolation(path="a", reason="b", rename_from="c") + assert (v.path, v.reason, v.rename_from) == ("a", "b", "c") diff --git a/agent-team/tests/test_ci_gate.py b/agent-team/tests/test_ci_gate.py new file mode 100644 index 0000000..4b0611b --- /dev/null +++ b/agent-team/tests/test_ci_gate.py @@ -0,0 +1,345 @@ +"""Unit tests for agent_team.ci_gate — the pure-code pass/fail gate (§3.3.2).""" + +from __future__ import annotations + +import pytest + +from agent_team.ci_gate import ( + DENYLIST_GLOBS, + CiGateError, + GateDecision, + denylist_violations, + diff_touched_paths, + evaluate_ci_gate, + verify_diff_hash, +) +from agent_team.state_store import compute_content_hash + + +def _diff_for(*paths: str) -> str: + """Build a minimal unified diff touching ``paths`` (no rename).""" + chunks = [] + for p in paths: + chunks.append( + f"diff --git a/{p} b/{p}\n--- a/{p}\n+++ b/{p}\n@@ -1 +1 @@\n-old\n+new\n" + ) + return "".join(chunks) + + +def _ledger_hash(diff: str) -> str: + return compute_content_hash(diff.encode("utf-8")) + + +def _good_ci(run_id: str, diff: str, conclusion: str = "success") -> dict: + return { + "run_id": run_id, + "conclusion": conclusion, + "diff_hash": _ledger_hash(diff), + } + + +# --------------------------------------------------------------------------- # +# diff_touched_paths +# --------------------------------------------------------------------------- # + + +def test_touched_paths_basic() -> None: + diff = _diff_for("src/foo.py", "tests/test_foo.py") + assert diff_touched_paths(diff) == ["src/foo.py", "tests/test_foo.py"] + + +def test_touched_paths_strips_git_prefix_and_dedups() -> None: + diff = "diff --git a/pkg/mod.py b/pkg/mod.py\n@@ @@\n+x\n" + assert diff_touched_paths(diff) == ["pkg/mod.py"] + + +def test_touched_paths_includes_rename_lines() -> None: + diff = ( + "diff --git a/safe.txt b/.github/workflows/evil.yml\n" + "similarity index 100%\n" + "rename from safe.txt\n" + "rename to .github/workflows/evil.yml\n" + ) + paths = diff_touched_paths(diff) + assert ".github/workflows/evil.yml" in paths + assert "safe.txt" in paths + + +def test_touched_paths_non_string_raises() -> None: + with pytest.raises(CiGateError): + diff_touched_paths(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- # +# denylist_violations +# --------------------------------------------------------------------------- # + + +def test_clean_diff_has_no_violations() -> None: + diff = _diff_for("src/foo.py", "README.md") + assert denylist_violations(diff) == [] + + +@pytest.mark.parametrize( + "path", + [ + ".github/workflows/ci.yml", + ".github/actions/deploy/action.yml", + ".github/CODEOWNERS", + "CODEOWNERS", + ".github/dependabot.yml", + "infra/cdk.json", + "service/template.yaml", + "deploy/iam/role.json", + "stacks/policies/admin.json", + "modules/main.tf", + ], +) +def test_denylisted_paths_flagged(path: str) -> None: + violations = denylist_violations(_diff_for(path)) + assert violations, f"expected {path!r} to be denylisted" + + +def test_rename_into_workflow_is_flagged() -> None: + diff = ( + "diff --git a/safe.txt b/.github/workflows/evil.yml\n" + "rename from safe.txt\n" + "rename to .github/workflows/evil.yml\n" + ) + violations = denylist_violations(diff) + assert any(".github/workflows/evil.yml" in v for v in violations) + + +def test_path_escape_via_dotdot_flagged() -> None: + diff = "diff --git a/x b/../../etc/passwd\n@@ @@\n+x\n" + violations = denylist_violations(diff) + assert any("escapes repo root" in v for v in violations) + + +def test_allowed_scope_blocks_out_of_scope_path() -> None: + diff = _diff_for("src/in_scope.py", "other/out_of_scope.py") + violations = denylist_violations(diff, allowed_scope=["src/"]) + assert any("declared scope" in v for v in violations) + # In-scope path alone is clean. + assert ( + denylist_violations(_diff_for("src/in_scope.py"), allowed_scope=["src/"]) == [] + ) + + +def test_allowed_scope_exact_file_prefix() -> None: + diff = _diff_for("pkg/exact.py") + assert denylist_violations(diff, allowed_scope=["pkg/exact.py"]) == [] + + +def test_denylist_globs_exported_nonempty() -> None: + assert isinstance(DENYLIST_GLOBS, tuple) + assert ".github/workflows/**" in DENYLIST_GLOBS + + +# --------------------------------------------------------------------------- # +# verify_diff_hash +# --------------------------------------------------------------------------- # + + +def test_hash_matches_ledger() -> None: + diff = _diff_for("a.py") + assert verify_diff_hash(diff, ledger_hash=_ledger_hash(diff)) is True + + +def test_hash_mismatch_ledger() -> None: + diff = _diff_for("a.py") + assert verify_diff_hash(diff, ledger_hash="deadbeef") is False + + +def test_hash_none_ledger_fails() -> None: + diff = _diff_for("a.py") + assert verify_diff_hash(diff, ledger_hash=None) is False + + +def test_hash_ci_verified_must_also_match() -> None: + diff = _diff_for("a.py") + h = _ledger_hash(diff) + assert verify_diff_hash(diff, ledger_hash=h, ci_verified_hash=h) is True + assert verify_diff_hash(diff, ledger_hash=h, ci_verified_hash="nope") is False + + +def test_hash_non_string_raises() -> None: + with pytest.raises(CiGateError): + verify_diff_hash(123, ledger_hash="x") # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- # +# evaluate_ci_gate — the block decision +# --------------------------------------------------------------------------- # + + +def test_gate_pass_on_authenticated_success() -> None: + diff = _diff_for("src/foo.py") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=_good_ci("run-1", diff, "success"), + expected_run_id="run-1", + ) + assert result.decision is GateDecision.PASS + assert result.passed is True + assert result.ci_conclusion == "success" + + +def test_gate_fail_on_authenticated_failure() -> None: + diff = _diff_for("src/foo.py") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=_good_ci("run-1", diff, "failure"), + expected_run_id="run-1", + ) + assert result.decision is GateDecision.FAIL + assert result.ci_conclusion == "failure" + + +@pytest.mark.parametrize( + "conclusion", ["timed_out", "cancelled", "startup_failure", "action_required"] +) +def test_gate_recognised_failures(conclusion: str) -> None: + diff = _diff_for("src/foo.py") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=_good_ci("run-1", diff, conclusion), + expected_run_id="run-1", + ) + assert result.decision is GateDecision.FAIL + + +def test_gate_blocks_denylisted_diff_even_if_ci_success() -> None: + # A diff that touches the trust-control surface BLOCKs regardless of CI. + diff = _diff_for(".github/workflows/ci.yml") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=_good_ci("run-1", diff, "success"), + expected_run_id="run-1", + ) + assert result.decision is GateDecision.BLOCK + assert result.blocked is True + assert any("denylisted" in r for r in result.reasons) + + +def test_gate_blocks_on_hash_mismatch() -> None: + diff = _diff_for("src/foo.py") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash="deadbeef", # does not match recomputed hash + ci_result=_good_ci("run-1", diff, "success"), + expected_run_id="run-1", + ) + assert result.decision is GateDecision.BLOCK + assert any("hash" in r for r in result.reasons) + + +def test_gate_blocks_on_run_id_mismatch() -> None: + diff = _diff_for("src/foo.py") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=_good_ci("run-OTHER", diff, "success"), + expected_run_id="run-1", + ) + assert result.decision is GateDecision.BLOCK + assert any("run-id" in r for r in result.reasons) + + +def test_gate_blocks_on_missing_ci_result() -> None: + diff = _diff_for("src/foo.py") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=None, + expected_run_id="run-1", + ) + assert result.decision is GateDecision.BLOCK + + +@pytest.mark.parametrize("conclusion", [None, "neutral", "skipped", "in_progress"]) +def test_gate_blocks_on_ambiguous_conclusion(conclusion) -> None: + # Anything not an explicit success or recognised failure must NOT silently + # pass — it BLOCKs (refuse-to-proceed). + diff = _diff_for("src/foo.py") + ci = _good_ci("run-1", diff, "success") + ci["conclusion"] = conclusion + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=ci, + expected_run_id="run-1", + ) + assert result.decision is GateDecision.BLOCK + + +def test_gate_ignores_patch_written_success_field() -> None: + # The gate reads only the authenticated conclusion. A patch-controlled + # "passed" flag must not flip a failing run to pass. + diff = _diff_for("src/foo.py") + ci = _good_ci("run-1", diff, "failure") + ci["passed"] = True # attacker-controlled artifact field + ci["success"] = "true" + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=ci, + expected_run_id="run-1", + ) + assert result.decision is GateDecision.FAIL + + +def test_gate_blocks_when_ci_diff_hash_mismatch() -> None: + # CI verified a different diff than the ledger recorded -> BLOCK. + diff = _diff_for("src/foo.py") + ci = _good_ci("run-1", diff, "success") + ci["diff_hash"] = "tampered" + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=ci, + expected_run_id="run-1", + ) + assert result.decision is GateDecision.BLOCK + + +def test_gate_scope_violation_blocks() -> None: + diff = _diff_for("src/foo.py", "unrelated/bar.py") + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=_good_ci("run-1", diff, "success"), + expected_run_id="run-1", + allowed_scope=["src/"], + ) + assert result.decision is GateDecision.BLOCK + assert any("scope" in r for r in result.reasons) + + +def test_gate_empty_run_id_raises() -> None: + diff = _diff_for("src/foo.py") + with pytest.raises(CiGateError): + evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=_ledger_hash(diff), + ci_result=_good_ci("run-1", diff), + expected_run_id="", + ) + + +def test_gate_result_carries_provenance() -> None: + diff = _diff_for("src/foo.py") + h = _ledger_hash(diff) + result = evaluate_ci_gate( + candidate_diff=diff, + ledger_hash=h, + ci_result=_good_ci("run-7", diff, "success"), + expected_run_id="run-7", + ) + assert result.run_id == "run-7" + assert result.diff_hash == h + assert result.reasons # always records why diff --git a/agent-team/tests/test_ci_gate_workflow.py b/agent-team/tests/test_ci_gate_workflow.py new file mode 100644 index 0000000..3e986f8 --- /dev/null +++ b/agent-team/tests/test_ci_gate_workflow.py @@ -0,0 +1,196 @@ +"""Tests for the embedded guard logic in ``ci/agent-team-apply-verify.yml``. + +The §3.3.2 trust-boundary guard (diff-integrity hash, the trust-control-surface +denylist, the symlink-escape reject, declared-scope enforcement) lives as an +inline Python heredoc inside the CI workflow, so it cannot be imported directly. +These tests extract that script from the YAML and execute it as the workflow +does — via env vars and a diff file — asserting the exit code for good and +adversarial diffs. This backs the workflow's correctness claim with a real, +runnable suite instead of an "verified during authoring" assertion, and guards +the symlink / non-UTF-8 / header-only-section fixes against regression. + +It does NOT enable, provision, or run the workflow itself; it only exercises the +pure-code gate the workflow embeds. +""" + +from __future__ import annotations + +import hashlib +import os +import subprocess +import sys +from pathlib import Path + +import pytest + +_WORKFLOW = Path(__file__).resolve().parents[1] / "ci" / "agent-team-apply-verify.yml" + + +def _extract_guard_script() -> str: + """Pull the first ``python3 - <<'PY' ... PY`` heredoc (the guard) from the YAML. + + The body is indented to sit under the YAML ``run:`` block; we strip the + common 10-space lead so it is valid module source. + """ + lines = _WORKFLOW.read_text(encoding="utf-8").splitlines() + start = end = None + for i, line in enumerate(lines): + if start is None and line.strip() == "python3 - <<'PY'": + start = i + 1 + elif start is not None and line.strip() == "PY": + end = i + break + assert start is not None and end is not None, "guard heredoc not found" + body = lines[start:end] + return "\n".join(ln[10:] if ln.startswith(" " * 10) else ln for ln in body) + + +@pytest.fixture(scope="module") +def guard_script(tmp_path_factory: pytest.TempPathFactory) -> Path: + path = tmp_path_factory.mktemp("guard") / "guard.py" + path.write_text(_extract_guard_script(), encoding="utf-8") + return path + + +def _run_guard( + guard_script: Path, + tmp_path: Path, + diff: str | bytes, + scope: str, + *, + bad_hash: bool = False, +) -> int: + raw = diff.encode("utf-8") if isinstance(diff, str) else diff + diff_path = tmp_path / "candidate.diff" + diff_path.write_bytes(raw) + expected = "deadbeef" if bad_hash else hashlib.sha256(raw).hexdigest() + env = dict( + os.environ, + DIFF_PATH=str(diff_path), + EXPECTED_DIFF_HASH=expected, + DECLARED_SCOPE=scope, + ) + result = subprocess.run( + [sys.executable, str(guard_script)], env=env, capture_output=True, text=True + ) + return result.returncode + + +CLEAN = ( + "diff --git a/src/app.py b/src/app.py\n" + "--- a/src/app.py\n+++ b/src/app.py\n@@ -1 +1 @@\n-x\n+y\n" +) +SYMLINK = ( + "diff --git a/src/link b/src/link\nnew file mode 120000\n" + "--- /dev/null\n+++ b/src/link\n@@ -0,0 +1 @@\n+../.github/workflows\n" +) +WORKFLOW_DELETE = ( + "diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml\n" + "deleted file mode 100644\n--- a/.github/workflows/ci.yml\n+++ /dev/null\n" + "@@ -1 +0,0 @@\n-on: push\n" +) +COPY_TO_DENIED = ( + "diff --git a/src/x.py b/.github/workflows/evil.yml\nsimilarity index 100%\n" + "copy from src/x.py\ncopy to .github/workflows/evil.yml\n" +) +NON_UTF8 = ( + b"diff --git a/src/app.py b/src/app.py\n--- a/src/app.py\n" + b"+++ b/src/app.py\n@@ -1 +1 @@\n-x\n+\xff\xfe\n" +) + + +def test_clean_in_scope_diff_passes(guard_script: Path, tmp_path: Path) -> None: + assert _run_guard(guard_script, tmp_path, CLEAN, "src/**") == 0 + + +def test_hash_mismatch_fails(guard_script: Path, tmp_path: Path) -> None: + assert _run_guard(guard_script, tmp_path, CLEAN, "src/**", bad_hash=True) == 2 + + +def test_workflow_delete_is_rejected(guard_script: Path, tmp_path: Path) -> None: + # Header-only section (delete) caught via diff --git, not a +++ body line. + assert ( + _run_guard(guard_script, tmp_path, WORKFLOW_DELETE, "src/**\n.github/**") == 4 + ) + + +def test_copy_into_denied_path_is_rejected(guard_script: Path, tmp_path: Path) -> None: + assert _run_guard(guard_script, tmp_path, COPY_TO_DENIED, "src/**\n.github/**") == 4 + + +def test_symlink_addition_is_rejected(guard_script: Path, tmp_path: Path) -> None: + # The symlink-escape vector: rejected outright (exit 7). + assert _run_guard(guard_script, tmp_path, SYMLINK, "src/**") == 7 + + +def test_non_utf8_diff_fails_closed(guard_script: Path, tmp_path: Path) -> None: + assert _run_guard(guard_script, tmp_path, NON_UTF8, "src/**") == 8 + + +def test_out_of_scope_path_is_rejected(guard_script: Path, tmp_path: Path) -> None: + assert _run_guard(guard_script, tmp_path, CLEAN, "other/**") == 6 + + +def test_unscoped_diff_is_rejected(guard_script: Path, tmp_path: Path) -> None: + assert _run_guard(guard_script, tmp_path, CLEAN, "") == 5 + + +def test_escaping_scope_entries_are_dropped(guard_script: Path, tmp_path: Path) -> None: + # A parent-escaping scope entry must not widen coverage; it is dropped, so a + # diff under it is treated as unscoped. + assert _run_guard(guard_script, tmp_path, CLEAN, "../../etc") == 5 + + +# --- security-review regressions: denylist & scope matcher (was fnmatch) ---- + + +def _modify(path: str) -> str: + return f"diff --git a/{path} b/{path}\n--- a/{path}\n+++ b/{path}\n@@ -1 +1 @@\n-x\n+y\n" + + +@pytest.mark.parametrize( + "root_path", + [ + "template.yaml", + "main.tf", + "cdk.json", + "policy.json", + "id.pem", + "signing.key", + "app-stack.ts", + "infra_stack.py", + ], +) +def test_root_level_iac_is_denied( + guard_script: Path, tmp_path: Path, root_path: str +) -> None: + # Regression: Python fnmatch '**/' is non-recursive, so root-level IaC/secret + # files slipped the denylist. The glob->regex matcher must reject them (exit 4). + assert _run_guard(guard_script, tmp_path, _modify(root_path), ".") == 4 + + +@pytest.mark.parametrize( + "cased_path", ["Template.YAML", "Main.TF", ".github/Workflows/ci.yml"] +) +def test_denylist_is_case_insensitive( + guard_script: Path, tmp_path: Path, cased_path: str +) -> None: + # A case variant of a trust-control filename must not evade the gate. + assert _run_guard(guard_script, tmp_path, _modify(cased_path), ".") == 4 + + +def test_scope_double_star_cannot_widen_to_whole_tree( + guard_script: Path, tmp_path: Path +) -> None: + # Regression: a '**' scope entry made in_scope() true for every path. It must + # reduce to the empty (root) prefix and be dropped -> unscoped (exit 5). + assert _run_guard(guard_script, tmp_path, _modify("any/deep/file.py"), "**") == 5 + + +def test_scope_glob_reduces_to_concrete_prefix( + guard_script: Path, tmp_path: Path +) -> None: + # A legitimate 'src/**' scope still confines to the src/ prefix: in-scope + # passes, a sibling path is rejected. + assert _run_guard(guard_script, tmp_path, _modify("src/app.py"), "src/**") == 0 + assert _run_guard(guard_script, tmp_path, _modify("other/app.py"), "src/**") == 6 diff --git a/agent-team/tests/test_clarifier.py b/agent-team/tests/test_clarifier.py new file mode 100644 index 0000000..b6fd56f --- /dev/null +++ b/agent-team/tests/test_clarifier.py @@ -0,0 +1,419 @@ +"""Unit tests for agent_team.nodes.clarifier (design §3.3, §7.1 P1). + +Covers the clarifier's two responsibilities: the **98% confidence loop** driven +by LangGraph ``interrupt()``/resume (the §7.1 P1 "riskiest mechanic"), and the +**human gate** — only a run clearing the bar advances to planning, a run that +hits the turn cap parks instead of spinning. + +The loop tests drive a real compiled ``StateGraph`` with an in-memory +checkpointer so the suspend/resume + replay semantics are exercised exactly as +they will be on the box (the box swaps in the SQLite checkpointer; the node +itself is checkpointer-agnostic). Confidence and question generation are +injected as deterministic callables, so the tests assert control flow without a +live Claude call. +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import pytest +from langgraph.checkpoint.memory import MemorySaver +from langgraph.graph import END, START, StateGraph +from langgraph.types import Command + +from agent_team.nodes.clarifier import ( + DEFAULT_CONFIDENCE_THRESHOLD, + DEFAULT_MAX_TURNS, + ClarifierConfig, + build_question_set, + make_clarifier_node, +) +from agent_team.task_model import ( + Phase, + PipelineState, + TaskStatus, + new_thread_id, +) +from agent_team.transport import QuestionSet + +# --------------------------------------------------------------------------- # +# Test helpers: deterministic injected assessor / generator. +# --------------------------------------------------------------------------- # + + +def _confidence_after_n_answers(target_turns: int) -> object: + """Assessor that clears the 98% bar once ``target_turns`` answers exist. + + Confidence is 0.0 with no answers and steps to 1.0 once enough answers have + been collected. Lets a test pin exactly how many interrupt turns the loop + should take. + """ + + def assess(qa_history: Sequence[object], state: PipelineState) -> float: + return 1.0 if len(qa_history) >= target_turns else 0.0 + + return assess + + +def _always_questions(*prompts: str): + """Generator returning a fixed question-set every turn.""" + + def generate(qa_history: Sequence[object], state: PipelineState) -> list[str]: + return list(prompts) or ["What is the goal?"] + + return generate + + +def _build_graph(node): + """Compile a one-node graph around ``node`` with an in-memory checkpointer.""" + graph = StateGraph(PipelineState) + graph.add_node("clarify", node) + graph.add_edge(START, "clarify") + graph.add_edge("clarify", END) + return graph.compile(checkpointer=MemorySaver()) + + +def _initial_state(thread_id: str, transport: str = "slack") -> PipelineState: + return { + "thread_id": thread_id, + "status": TaskStatus.ACTIVE.value, + "current_phase": Phase.CLARIFY.value, + "qa_history": [], + "transport": transport, + } + + +# --------------------------------------------------------------------------- # +# build_question_set — the §3.3.1 interrupt payload. +# --------------------------------------------------------------------------- # + + +def test_build_question_set_returns_foundation_questionset() -> None: + qs = build_question_set(thread_id="t1", turn=2, questions=["a", "b"]) + assert isinstance(qs, QuestionSet) + assert qs.thread_id == "t1" + assert qs.turn == 2 + assert qs.questions == ["a", "b"] + assert qs.context == {} + + +def test_build_question_set_mints_unique_question_ids() -> None: + a = build_question_set(thread_id="t", turn=0, questions=["x"]) + b = build_question_set(thread_id="t", turn=0, questions=["x"]) + assert a.question_id != b.question_id + assert len(a.question_id) == 32 # uuid4 hex + + +def test_build_question_set_honours_explicit_question_id() -> None: + qs = build_question_set( + thread_id="t", turn=0, questions=["x"], question_id="fixed-id" + ) + assert qs.question_id == "fixed-id" + + +def test_build_question_set_copies_inputs_defensively() -> None: + questions = ["x"] + context = {"repo": "r"} + qs = build_question_set(thread_id="t", turn=0, questions=questions, context=context) + questions.append("y") + context["repo"] = "mutated" + assert qs.questions == ["x"] + assert qs.context == {"repo": "r"} + + +# --------------------------------------------------------------------------- # +# ClarifierConfig validation. +# --------------------------------------------------------------------------- # + + +def test_config_defaults_match_design_constants() -> None: + cfg = ClarifierConfig() + assert cfg.confidence_threshold == DEFAULT_CONFIDENCE_THRESHOLD == 0.98 + assert cfg.max_turns == DEFAULT_MAX_TURNS + + +@pytest.mark.parametrize("bad", [0.0, -0.1, 1.5]) +def test_config_rejects_threshold_out_of_range(bad: float) -> None: + with pytest.raises(ValueError, match="confidence_threshold"): + ClarifierConfig(confidence_threshold=bad) + + +def test_config_accepts_threshold_of_one() -> None: + assert ClarifierConfig(confidence_threshold=1.0).confidence_threshold == 1.0 + + +@pytest.mark.parametrize("bad", [0, -3]) +def test_config_rejects_nonpositive_max_turns(bad: int) -> None: + with pytest.raises(ValueError, match="max_turns"): + ClarifierConfig(max_turns=bad) + + +# --------------------------------------------------------------------------- # +# Confidence loop: already-confident, no interrupt (human gate opens at once). +# --------------------------------------------------------------------------- # + + +def test_no_interrupt_when_already_confident() -> None: + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(0), # confident immediately + generate_questions=_always_questions("q"), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + result = app.invoke(_initial_state("t1"), cfg) + + assert "__interrupt__" not in result + assert result["current_phase"] == Phase.PLAN.value + assert result["status"] == TaskStatus.ACTIVE.value + assert result["qa_history"] == [] + + +# --------------------------------------------------------------------------- # +# Confidence loop: interrupts until the 98% bar is cleared, then human gate. +# --------------------------------------------------------------------------- # + + +def test_single_turn_loop_clears_bar_and_opens_gate() -> None: + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(1), + generate_questions=_always_questions("What is the goal?"), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + + first = app.invoke(_initial_state("t1"), cfg) + assert "__interrupt__" in first # suspended on the question-set + + final = app.invoke(Command(resume="ship feature X"), cfg) + assert "__interrupt__" not in final + assert final["qa_history"] == ["ship feature X"] + assert final["current_phase"] == Phase.PLAN.value + assert final["status"] == TaskStatus.ACTIVE.value + + +def test_multi_turn_loop_accumulates_answers_until_confident() -> None: + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(3), + generate_questions=_always_questions("q?"), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + + assert "__interrupt__" in app.invoke(_initial_state("t1"), cfg) + assert "__interrupt__" in app.invoke(Command(resume="a1"), cfg) + assert "__interrupt__" in app.invoke(Command(resume="a2"), cfg) + final = app.invoke(Command(resume="a3"), cfg) + + assert "__interrupt__" not in final + assert final["qa_history"] == ["a1", "a2", "a3"] + assert final["current_phase"] == Phase.PLAN.value + assert final["status"] == TaskStatus.ACTIVE.value + + +# --------------------------------------------------------------------------- # +# Interrupt payload shape (§3.3.1): {thread_id, question_id, turn, ...}. +# --------------------------------------------------------------------------- # + + +def test_interrupt_payload_carries_design_fields() -> None: + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(1), + generate_questions=_always_questions("Q1", "Q2"), + config=ClarifierConfig(transport="slack"), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "tid-123"}} + app.invoke(_initial_state("tid-123", transport="ignored"), cfg) + + state = app.get_state(cfg) + assert len(state.interrupts) == 1 + payload = state.interrupts[0].value + + assert payload["thread_id"] == "tid-123" + assert payload["turn"] == 0 + assert payload["transport"] == "slack" # config wins over state + assert payload["deadline"] is None # owned by the ledger/timer seam + assert isinstance(payload["question_id"], str) and payload["question_id"] + + qs = payload["question_set"] + assert isinstance(qs, QuestionSet) + assert qs.questions == ["Q1", "Q2"] + assert qs.question_id == payload["question_id"] + assert qs.thread_id == "tid-123" + + +def test_interrupt_payload_transport_falls_back_to_state() -> None: + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(1), + generate_questions=_always_questions("q"), + config=ClarifierConfig(), # empty transport + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + app.invoke(_initial_state("t1", transport="github"), cfg) + + payload = app.get_state(cfg).interrupts[0].value + assert payload["transport"] == "github" + + +def test_turn_index_increments_across_interrupts() -> None: + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(2), + generate_questions=_always_questions("q"), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + + app.invoke(_initial_state("t1"), cfg) + assert app.get_state(cfg).interrupts[0].value["turn"] == 0 + + app.invoke(Command(resume="a1"), cfg) + assert app.get_state(cfg).interrupts[0].value["turn"] == 1 + + +# --------------------------------------------------------------------------- # +# Turn cap (§7.1): park rather than spin; never advance to planning. +# --------------------------------------------------------------------------- # + + +def test_turn_cap_parks_instead_of_advancing() -> None: + # Never confident; cap at 2 turns. + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(999), + generate_questions=_always_questions("q"), + config=ClarifierConfig(max_turns=2), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + + assert "__interrupt__" in app.invoke(_initial_state("t1"), cfg) + assert "__interrupt__" in app.invoke(Command(resume="a1"), cfg) + final = app.invoke(Command(resume="a2"), cfg) # cap reached + + assert "__interrupt__" not in final + assert final["current_phase"] == Phase.PARKED.value + assert final["status"] == TaskStatus.PARKED.value + assert final["qa_history"] == ["a1", "a2"] + # Human gate stayed shut: never advanced to PLAN. + assert final["current_phase"] != Phase.PLAN.value + + +def test_clearing_bar_on_the_cap_turn_still_opens_gate() -> None: + # Confident exactly on the last allowed turn — gate must open, not park. + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(2), + generate_questions=_always_questions("q"), + config=ClarifierConfig(max_turns=2), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + + app.invoke(_initial_state("t1"), cfg) + app.invoke(Command(resume="a1"), cfg) + final = app.invoke(Command(resume="a2"), cfg) + + assert final["current_phase"] == Phase.PLAN.value + assert final["status"] == TaskStatus.ACTIVE.value + + +# --------------------------------------------------------------------------- # +# Durable resume across a "restart" (§7.1 P1 exit (a)): a fresh app object +# bound to the same checkpointer resumes mid-wait to the right thread. +# --------------------------------------------------------------------------- # + + +def test_resume_after_restart_uses_same_checkpoint() -> None: + saver = MemorySaver() + + def compile_app(): + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(1), + generate_questions=_always_questions("q"), + ) + graph = StateGraph(PipelineState) + graph.add_node("clarify", node) + graph.add_edge(START, "clarify") + graph.add_edge("clarify", END) + return graph.compile(checkpointer=saver) + + cfg = {"configurable": {"thread_id": "t1"}} + app1 = compile_app() + assert "__interrupt__" in app1.invoke(_initial_state("t1"), cfg) + + # Simulate a process restart: a brand-new app object, same checkpointer. + app2 = compile_app() + state = app2.get_state(cfg) + assert state.next == ("clarify",) # still suspended at the node + + final = app2.invoke(Command(resume="late answer"), cfg) + assert final["qa_history"] == ["late answer"] + assert final["current_phase"] == Phase.PLAN.value + + +# --------------------------------------------------------------------------- # +# Parallel-task isolation (§7.1 P1 exit (d)): two threads resume independently. +# --------------------------------------------------------------------------- # + + +def test_two_threads_resume_independently() -> None: + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(1), + generate_questions=_always_questions("q"), + ) + app = _build_graph(node) + cfg_a = {"configurable": {"thread_id": "task-a"}} + cfg_b = {"configurable": {"thread_id": "task-b"}} + + # Suspend both. + app.invoke(_initial_state("task-a"), cfg_a) + app.invoke(_initial_state("task-b"), cfg_b) + + # Resume out of order; each carries its own answer + thread_id. + final_b = app.invoke(Command(resume="answer-b"), cfg_b) + final_a = app.invoke(Command(resume="answer-a"), cfg_a) + + assert final_a["qa_history"] == ["answer-a"] + assert final_b["qa_history"] == ["answer-b"] + assert final_a["current_phase"] == Phase.PLAN.value + assert final_b["current_phase"] == Phase.PLAN.value + + +# --------------------------------------------------------------------------- # +# Resumed task keeps prior Q&A history (a re-opened parked task adds context). +# --------------------------------------------------------------------------- # + + +def test_existing_qa_history_is_preserved_and_extended() -> None: + node = make_clarifier_node( + # Need 2 total answers; one already exists, so one more interrupt. + assess_confidence=_confidence_after_n_answers(2), + generate_questions=_always_questions("q"), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": "t1"}} + seeded = _initial_state("t1") + seeded["qa_history"] = ["prior answer"] + + assert "__interrupt__" in app.invoke(seeded, cfg) + final = app.invoke(Command(resume="new answer"), cfg) + + assert final["qa_history"] == ["prior answer", "new answer"] + assert final["current_phase"] == Phase.PLAN.value + + +# --------------------------------------------------------------------------- # +# The node is callable factory output and integrates with new_thread_id intake. +# --------------------------------------------------------------------------- # + + +def test_node_runs_against_freshly_minted_thread_id() -> None: + tid = new_thread_id() + node = make_clarifier_node( + assess_confidence=_confidence_after_n_answers(0), + generate_questions=_always_questions("q"), + ) + app = _build_graph(node) + cfg = {"configurable": {"thread_id": tid}} + result = app.invoke(_initial_state(tid), cfg) + assert result["current_phase"] == Phase.PLAN.value diff --git a/agent-team/tests/test_claude_code_adapter.py b/agent-team/tests/test_claude_code_adapter.py new file mode 100644 index 0000000..3bfcb3c --- /dev/null +++ b/agent-team/tests/test_claude_code_adapter.py @@ -0,0 +1,310 @@ +"""Unit tests for agent_team.transport.claude_code_adapter (§3.3.1, §7.1 P4).""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from agent_team.transport.base import ( + GITHUB_MARKER_TEMPLATE, + Transport, +) +from agent_team.transport.claude_code_adapter import ( + VIA, + ClaudeCodeAdapter, + ClaudeCodeDeliveryError, + build_channel_ref, + parse_channel_ref, + render_prompt, +) + + +def _question_set(**overrides: Any): + from agent_team.transport.base import QuestionSet + + defaults: dict[str, Any] = { + "thread_id": "t1", + "question_id": "q1", + "turn": 0, + "questions": ["Proceed with the dependency bump?"], + } + defaults.update(overrides) + return QuestionSet(**defaults) + + +class _RecordingSink: + """A fake delivery sink that records its call and returns a session id.""" + + def __init__(self, session_id: str = "sess-abc") -> None: + self.session_id = session_id + self.calls: list[dict[str, Any]] = [] + + def __call__(self, *, session_hint: str, prompt: str) -> str: + self.calls.append({"session_hint": session_hint, "prompt": prompt}) + return self.session_id + + +# --------------------------------------------------------------------------- # +# Contract / typing +# --------------------------------------------------------------------------- # + + +def test_adapter_is_a_transport_subclass() -> None: + assert issubclass(ClaudeCodeAdapter, Transport) + + +def test_adapter_instantiates_without_abstract_error() -> None: + # All abstract methods are implemented, so construction must succeed. + adapter = ClaudeCodeAdapter(_RecordingSink()) + assert isinstance(adapter, Transport) + + +def test_via_constant_value() -> None: + assert VIA == "claude_code" + + +# --------------------------------------------------------------------------- # +# channel_ref helpers +# --------------------------------------------------------------------------- # + + +def test_build_channel_ref_embeds_session_and_question() -> None: + ref = build_channel_ref("sess-xyz", "q9") + assert ref == "claude-session:sess-xyz:q9" + + +def test_parse_channel_ref_roundtrips_build() -> None: + ref = build_channel_ref("sess-xyz", "q9") + assert parse_channel_ref(ref) == ("sess-xyz", "q9") + + +def test_parse_channel_ref_rejects_foreign_ref() -> None: + # A Slack ts or GitHub comment id is not a Claude-Code ref. + assert parse_channel_ref("1718000000.001100") is None + assert parse_channel_ref("issuecomment-12345") is None + + +# --------------------------------------------------------------------------- # +# render_prompt +# --------------------------------------------------------------------------- # + + +def test_render_prompt_embeds_question_marker() -> None: + prompt = render_prompt( + question_id="abc123", + turn=2, + question_set=_question_set(questions=["a?", "b?"]), + deadline="2026-06-18T00:00:00Z", + ) + assert GITHUB_MARKER_TEMPLATE.format(question_id="abc123") in prompt + assert "turn 2" in prompt + assert "1. a?" in prompt + assert "2. b?" in prompt + assert "2026-06-18T00:00:00Z" in prompt + + +def test_render_prompt_surfaces_context_when_present() -> None: + prompt = render_prompt( + question_id="q1", + turn=0, + question_set=_question_set( + context={"repo": "sea-haven/widgets", "summary": "bump urllib3"} + ), + deadline="d", + ) + assert "repo: sea-haven/widgets" in prompt + assert "summary: bump urllib3" in prompt + + +def test_render_prompt_omits_absent_context_keys() -> None: + prompt = render_prompt( + question_id="q1", + turn=0, + question_set=_question_set(context={}), + deadline="d", + ) + assert "repo:" not in prompt + assert "summary:" not in prompt + + +# --------------------------------------------------------------------------- # +# post_question +# --------------------------------------------------------------------------- # + + +def test_post_question_returns_channel_ref_with_question_id() -> None: + sink = _RecordingSink(session_id="sess-42") + adapter = ClaudeCodeAdapter(sink, session_hint="mac-1") + ref = adapter.post_question( + thread_id="t1", + question_id="q1", + turn=0, + question_set=_question_set(), + deadline="2026-06-18T00:00:00Z", + ) + assert ref == "claude-session:sess-42:q1" + # The channel_ref alone must round-trip the question_id (§3.3.1). + assert parse_channel_ref(ref) == ("sess-42", "q1") + + +def test_post_question_passes_rendered_prompt_and_hint_to_sink() -> None: + sink = _RecordingSink() + adapter = ClaudeCodeAdapter(sink, session_hint="mac-7") + adapter.post_question( + thread_id="t1", + question_id="qZ", + turn=1, + question_set=_question_set(question_id="qZ", questions=["go?"]), + deadline="d", + ) + assert len(sink.calls) == 1 + call = sink.calls[0] + assert call["session_hint"] == "mac-7" + assert GITHUB_MARKER_TEMPLATE.format(question_id="qZ") in call["prompt"] + assert "go?" in call["prompt"] + + +def test_post_question_default_sink_raises() -> None: + adapter = ClaudeCodeAdapter() # no delivery injected + with pytest.raises(ClaudeCodeDeliveryError): + adapter.post_question( + thread_id="t1", + question_id="q1", + turn=0, + question_set=_question_set(), + deadline="d", + ) + + +@pytest.mark.parametrize("bad_session", ["", " ", None]) +def test_post_question_blank_session_id_is_failed_post(bad_session: Any) -> None: + def sink(*, session_hint: str, prompt: str) -> Any: + return bad_session + + adapter = ClaudeCodeAdapter(sink) + with pytest.raises(ClaudeCodeDeliveryError): + adapter.post_question( + thread_id="t1", + question_id="q1", + turn=0, + question_set=_question_set(), + deadline="d", + ) + + +def test_post_question_strips_session_id_whitespace() -> None: + adapter = ClaudeCodeAdapter(lambda *, session_hint, prompt: " sess-w ") + ref = adapter.post_question( + thread_id="t1", + question_id="q1", + turn=0, + question_set=_question_set(), + deadline="d", + ) + assert ref == "claude-session:sess-w:q1" + + +# --------------------------------------------------------------------------- # +# parse_answer +# --------------------------------------------------------------------------- # + + +def test_parse_answer_explicit_question_id() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + parsed = adapter.parse_answer({"question_id": "q1", "answer": "yes"}) + assert parsed == ("q1", "yes", VIA) + + +def test_parse_answer_recovers_question_id_from_channel_ref() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + parsed = adapter.parse_answer( + {"channel_ref": build_channel_ref("sess-1", "q7"), "answer": "ship it"} + ) + assert parsed == ("q7", "ship it", VIA) + + +def test_parse_answer_recovers_question_id_from_prompt_marker() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + marker = GITHUB_MARKER_TEMPLATE.format(question_id="qM") + parsed = adapter.parse_answer({"prompt": f"{marker}\nanswered inline", "answer": 3}) + assert parsed == ("qM", 3, VIA) + + +def test_parse_answer_recovers_question_id_from_marker_in_answer_text() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + marker = GITHUB_MARKER_TEMPLATE.format(question_id="qInline") + parsed = adapter.parse_answer({"answer": f"my reply {marker}"}) + assert parsed[0] == "qInline" + assert parsed[2] == VIA + + +def test_parse_answer_explicit_id_takes_priority_over_ref() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + parsed = adapter.parse_answer( + { + "question_id": "explicit", + "channel_ref": build_channel_ref("sess-1", "fromref"), + "answer": "x", + } + ) + assert parsed[0] == "explicit" + + +def test_parse_answer_value_and_text_fallbacks() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + assert adapter.parse_answer({"question_id": "q", "value": "v"})[1] == "v" + assert adapter.parse_answer({"question_id": "q", "text": "t"})[1] == "t" + + +def test_parse_answer_answer_key_wins_over_value_and_text() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + parsed = adapter.parse_answer( + {"question_id": "q", "answer": "A", "value": "V", "text": "T"} + ) + assert parsed[1] == "A" + + +def test_parse_answer_preserves_structured_answer() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + payload = {"question_id": "q", "answer": {"choice": 2, "note": "ok"}} + parsed = adapter.parse_answer(payload) + assert parsed[1] == {"choice": 2, "note": "ok"} + + +def test_parse_answer_missing_answer_is_none() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + parsed = adapter.parse_answer({"question_id": "q"}) + assert parsed == ("q", None, VIA) + + +def test_parse_answer_rejects_unmappable_payload() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + with pytest.raises(ValueError): + adapter.parse_answer({"answer": "no id anywhere"}) + + +def test_parse_answer_rejects_non_mapping() -> None: + adapter = ClaudeCodeAdapter(_RecordingSink()) + with pytest.raises(ValueError): + adapter.parse_answer("not a dict") + + +# --------------------------------------------------------------------------- # +# Round-trip: post then parse the resulting ref +# --------------------------------------------------------------------------- # + + +def test_post_then_parse_roundtrip_via_channel_ref() -> None: + sink = _RecordingSink(session_id="sess-rt") + adapter = ClaudeCodeAdapter(sink) + ref = adapter.post_question( + thread_id="t1", + question_id="qRT", + turn=0, + question_set=_question_set(question_id="qRT"), + deadline="d", + ) + # An inbound answer that echoes only the channel_ref still maps back. + parsed = adapter.parse_answer({"channel_ref": ref, "answer": "done"}) + assert parsed == ("qRT", "done", VIA) diff --git a/agent-team/tests/test_deadline_timer.py b/agent-team/tests/test_deadline_timer.py new file mode 100644 index 0000000..e4f4c5f --- /dev/null +++ b/agent-team/tests/test_deadline_timer.py @@ -0,0 +1,504 @@ +"""Unit tests for agent_team.deadline_timer (design §3.3.1). + +Covers the deadline / no-answer timer loop: overdue selection, the +deterministic answer-vs-timeout race on the ``open`` -> ``expired`` flip, the +PARK and DEFAULT_ANSWER policies, restart idempotency, side-effect isolation, +and the concurrent responder-vs-timer race. +""" + +from __future__ import annotations + +import sqlite3 +import threading +from datetime import datetime, timedelta, timezone +from pathlib import Path + +import pytest + +from agent_team.db.schema import answer_question, connect, init_db +from agent_team.deadline_timer import ( + DeadlinePolicy, + ExpiryAction, + OverdueQuestion, + TimerLoopReport, + overdue_open_questions, + run_deadline_timer, +) + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +_PAST = "2000-01-01T00:00:00+00:00" +_FUTURE = "2999-01-01T00:00:00+00:00" + + +def _iso(dt: datetime) -> str: + return dt.isoformat() + + +def _insert_question( + conn: sqlite3.Connection, + qid: str, + *, + thread_id: str = "thread-1", + turn: int = 0, + status: str = "open", + transport: str = "slack", + deadline_at: str | None = _PAST, +) -> None: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, posted_at, deadline_at) " + "VALUES (?, ?, ?, ?, ?, ?, ?)", + (qid, thread_id, turn, status, transport, _PAST, deadline_at), + ) + + +@pytest.fixture() +def conn(tmp_path: Path) -> sqlite3.Connection: + db = tmp_path / "db.sqlite" + init_db(db) + c = connect(db) + yield c + c.close() + + +def _status(conn: sqlite3.Connection, qid: str) -> str: + return conn.execute( + "SELECT status FROM pending_questions WHERE question_id=?", (qid,) + ).fetchone()["status"] + + +class _Collector: + """Records the questions handed to an injected side-effect callback.""" + + def __init__(self) -> None: + self.calls: list[OverdueQuestion] = [] + + def __call__(self, question: OverdueQuestion) -> None: + self.calls.append(question) + + +# --------------------------------------------------------------------------- +# overdue_open_questions +# --------------------------------------------------------------------------- + + +def test_overdue_selects_only_past_open_with_deadline(conn: sqlite3.Connection) -> None: + _insert_question(conn, "past", deadline_at=_PAST) + _insert_question(conn, "future", deadline_at=_FUTURE) + _insert_question(conn, "no-deadline", deadline_at=None) + _insert_question(conn, "answered", status="answered", deadline_at=_PAST) + _insert_question(conn, "expired", status="expired", deadline_at=_PAST) + + overdue = overdue_open_questions(conn) + assert [q.question_id for q in overdue] == ["past"] + + +def test_overdue_uses_now_cutoff(conn: sqlite3.Connection) -> None: + now = datetime(2026, 6, 17, tzinfo=timezone.utc) + just_past = _iso(now - timedelta(seconds=1)) + just_future = _iso(now + timedelta(seconds=1)) + _insert_question(conn, "before", deadline_at=just_past) + _insert_question(conn, "after", deadline_at=just_future) + + overdue = overdue_open_questions(conn, now=_iso(now)) + assert [q.question_id for q in overdue] == ["before"] + + +def test_overdue_includes_deadline_equal_to_now(conn: sqlite3.Connection) -> None: + now = "2026-06-17T00:00:00+00:00" + _insert_question(conn, "exact", deadline_at=now) + overdue = overdue_open_questions(conn, now=now) + assert [q.question_id for q in overdue] == ["exact"] + + +def test_overdue_ordered_oldest_deadline_first(conn: sqlite3.Connection) -> None: + _insert_question(conn, "newer", deadline_at="2010-01-01T00:00:00+00:00") + _insert_question(conn, "older", deadline_at="2001-01-01T00:00:00+00:00") + overdue = overdue_open_questions(conn) + assert [q.question_id for q in overdue] == ["older", "newer"] + + +def test_overdue_row_fields_mapped(conn: sqlite3.Connection) -> None: + _insert_question( + conn, "q", thread_id="t-42", turn=3, transport="github", deadline_at=_PAST + ) + (q,) = overdue_open_questions(conn) + assert q.question_id == "q" + assert q.thread_id == "t-42" + assert q.turn == 3 + assert q.transport == "github" + assert q.deadline_at == _PAST + + +# --------------------------------------------------------------------------- +# run_deadline_timer — PARK policy (default) +# --------------------------------------------------------------------------- + + +def test_park_policy_expires_and_alarms(conn: sqlite3.Connection) -> None: + _insert_question(conn, "q1") + on_park = _Collector() + + report = run_deadline_timer(conn, on_park=on_park) + + assert _status(conn, "q1") == "expired" + assert [q.question_id for q in on_park.calls] == ["q1"] + assert report.examined == 1 + assert report.parked == 1 + assert report.expired == 1 + assert report.lost_race == 0 + assert report.outcomes[0].action is ExpiryAction.PARKED + assert report.outcomes[0].policy is DeadlinePolicy.PARK + + +def test_default_policy_is_park(conn: sqlite3.Connection) -> None: + _insert_question(conn, "q1") + on_park = _Collector() + # No policy_resolver supplied -> PARK for every overdue question. + run_deadline_timer(conn, on_park=on_park) + assert len(on_park.calls) == 1 + + +def test_multiple_overdue_all_parked(conn: sqlite3.Connection) -> None: + for i in range(5): + _insert_question(conn, f"q{i}") + on_park = _Collector() + + report = run_deadline_timer(conn, on_park=on_park) + + assert report.examined == 5 + assert report.parked == 5 + assert {c.question_id for c in on_park.calls} == {f"q{i}" for i in range(5)} + for i in range(5): + assert _status(conn, f"q{i}") == "expired" + + +def test_future_and_null_deadlines_untouched(conn: sqlite3.Connection) -> None: + _insert_question(conn, "future", deadline_at=_FUTURE) + _insert_question(conn, "none", deadline_at=None) + on_park = _Collector() + + report = run_deadline_timer(conn, on_park=on_park) + + assert report.examined == 0 + assert on_park.calls == [] + assert _status(conn, "future") == "open" + assert _status(conn, "none") == "open" + + +# --------------------------------------------------------------------------- +# run_deadline_timer — DEFAULT_ANSWER policy +# --------------------------------------------------------------------------- + + +def test_default_answer_policy_resumes(conn: sqlite3.Connection) -> None: + _insert_question(conn, "q1") + on_park = _Collector() + resume = _Collector() + + report = run_deadline_timer( + conn, + on_park=on_park, + resume_with_default=resume, + policy_resolver=lambda _q: DeadlinePolicy.DEFAULT_ANSWER, + ) + + assert _status(conn, "q1") == "expired" + assert on_park.calls == [] + assert [q.question_id for q in resume.calls] == ["q1"] + assert report.defaulted == 1 + assert report.parked == 0 + assert report.outcomes[0].action is ExpiryAction.DEFAULTED + + +def test_default_answer_without_callback_raises(conn: sqlite3.Connection) -> None: + _insert_question(conn, "q1") + on_park = _Collector() + + with pytest.raises(ValueError, match="DEFAULT_ANSWER"): + run_deadline_timer( + conn, + on_park=on_park, + policy_resolver=lambda _q: DeadlinePolicy.DEFAULT_ANSWER, + ) + # The row was already flipped to expired before the config error surfaced. + assert _status(conn, "q1") == "expired" + + +def test_mixed_policies_routed_per_question(conn: sqlite3.Connection) -> None: + _insert_question(conn, "park-me") + _insert_question(conn, "default-me") + on_park = _Collector() + resume = _Collector() + + def resolver(q: OverdueQuestion) -> DeadlinePolicy: + return ( + DeadlinePolicy.DEFAULT_ANSWER + if q.question_id == "default-me" + else DeadlinePolicy.PARK + ) + + report = run_deadline_timer( + conn, + on_park=on_park, + resume_with_default=resume, + policy_resolver=resolver, + ) + + assert [c.question_id for c in on_park.calls] == ["park-me"] + assert [c.question_id for c in resume.calls] == ["default-me"] + assert report.parked == 1 + assert report.defaulted == 1 + + +# --------------------------------------------------------------------------- +# Deterministic answer-vs-timeout race (§3.3.1) +# --------------------------------------------------------------------------- + + +def test_answered_first_loses_race_no_policy(conn: sqlite3.Connection) -> None: + """A question answered before the timer runs must NOT be parked/defaulted.""" + _insert_question(conn, "q1") + # Responder wins first. + assert answer_question( + conn, question_id="q1", answer_json="{}", answered_via="slack" + ) + on_park = _Collector() + resume = _Collector() + + report = run_deadline_timer(conn, on_park=on_park, resume_with_default=resume) + + # The row is no longer ``open`` so it is not even returned by the overdue + # query — examined is zero, no policy applied. + assert report.examined == 0 + assert on_park.calls == [] + assert resume.calls == [] + assert _status(conn, "q1") == "answered" + + +def test_lost_race_when_answered_between_select_and_flip( + conn: sqlite3.Connection, monkeypatch: pytest.MonkeyPatch +) -> None: + """If a responder answers a row after it was selected as overdue but before + the timer flips it, the timer's compare-and-set returns False -> LOST_RACE, + and NO no-answer policy is applied (the question was actually answered).""" + import agent_team.deadline_timer as dt + + _insert_question(conn, "q1") + on_park = _Collector() + + real_expire = dt.expire_question + + def racing_expire(c: sqlite3.Connection, *, question_id: str) -> bool: + # Simulate the responder winning the compare-and-set in the window + # between overdue selection and this flip. Uses the SAME connection so + # the ordering is fully deterministic (no thread scheduling needed). + answer_question( + c, question_id=question_id, answer_json="{}", answered_via="slack" + ) + return real_expire(c, question_id=question_id) + + monkeypatch.setattr(dt, "expire_question", racing_expire) + + report = run_deadline_timer(conn, on_park=on_park) + + # The flip lost: no park, row is ``answered`` not ``expired``. + assert on_park.calls == [] + assert report.examined == 1 + assert report.lost_race == 1 + assert report.parked == 0 + assert report.expired == 0 + assert report.outcomes[0].action is ExpiryAction.LOST_RACE + assert _status(conn, "q1") == "answered" + + +# --------------------------------------------------------------------------- +# Restart idempotency +# --------------------------------------------------------------------------- + + +def test_rerun_is_idempotent(conn: sqlite3.Connection) -> None: + """A second pass after a 'crash' processes no already-expired rows.""" + _insert_question(conn, "q1") + on_park = _Collector() + + first = run_deadline_timer(conn, on_park=on_park) + assert first.parked == 1 + + second = run_deadline_timer(conn, on_park=on_park) + # Already expired -> no longer ``open`` -> not selected -> no double park. + assert second.examined == 0 + assert len(on_park.calls) == 1 + assert _status(conn, "q1") == "expired" + + +# --------------------------------------------------------------------------- +# Side-effect isolation +# --------------------------------------------------------------------------- + + +def test_side_effect_error_isolated_row_still_expired( + conn: sqlite3.Connection, +) -> None: + _insert_question(conn, "boom") + _insert_question(conn, "ok") + + def flaky_park(q: OverdueQuestion) -> None: + if q.question_id == "boom": + raise RuntimeError("alarm transport down") + + report = run_deadline_timer(conn, on_park=flaky_park) + + # Both rows are durably expired (flip commits before the side effect). + assert _status(conn, "boom") == "expired" + assert _status(conn, "ok") == "expired" + assert report.errored == 1 + assert report.parked == 1 + errored = next(o for o in report.outcomes if o.action is ExpiryAction.ERRORED) + assert errored.question_id == "boom" + assert "RuntimeError" in (errored.error or "") + assert errored.policy is DeadlinePolicy.PARK + + +def test_error_in_one_row_does_not_abort_batch(conn: sqlite3.Connection) -> None: + for i in range(4): + _insert_question(conn, f"q{i}") + + def park(q: OverdueQuestion) -> None: + if q.question_id == "q1": + raise ValueError("nope") + + report = run_deadline_timer(conn, on_park=park) + + assert report.examined == 4 + assert report.errored == 1 + assert report.parked == 3 + for i in range(4): + assert _status(conn, f"q{i}") == "expired" + + +# --------------------------------------------------------------------------- +# TimerLoopReport counters +# --------------------------------------------------------------------------- + + +def test_empty_report_counters() -> None: + report = TimerLoopReport() + assert report.examined == 0 + assert report.parked == 0 + assert report.defaulted == 0 + assert report.lost_race == 0 + assert report.errored == 0 + assert report.expired == 0 + + +def test_expired_equals_examined_minus_lost_race( + conn: sqlite3.Connection, monkeypatch: pytest.MonkeyPatch +) -> None: + import agent_team.deadline_timer as dt + + _insert_question(conn, "parkable") + _insert_question(conn, "raced") + on_park = _Collector() + + real_expire = dt.expire_question + + def racing_expire(c: sqlite3.Connection, *, question_id: str) -> bool: + if question_id == "raced": + answer_question( + c, question_id="raced", answer_json="{}", answered_via="slack" + ) + return real_expire(c, question_id=question_id) + + monkeypatch.setattr(dt, "expire_question", racing_expire) + + report = run_deadline_timer(conn, on_park=on_park) + + assert report.examined == 2 + assert report.lost_race == 1 + assert report.expired == report.examined - report.lost_race == 1 + + +# --------------------------------------------------------------------------- +# Concurrency: responder thread vs timer thread on the same question +# --------------------------------------------------------------------------- + + +def test_concurrent_timer_and_responder_single_winner(tmp_path: Path) -> None: + """A timer pass and a responder race the same open question: the + BEGIN-IMMEDIATE compare-and-set guarantees exactly one of 'expire'/'answer' + wins, and the timer parks IFF it actually flipped the row to ``expired``. + + Repeated across many rows so the threads genuinely interleave (the barrier + aligns each pair at the start), catching any non-determinism in the race. + """ + db = tmp_path / "race.sqlite" + init_db(db) + + n = 40 + seed = connect(db) + try: + for i in range(n): + _insert_question(seed, f"race{i}", deadline_at=_PAST) + finally: + seed.close() + + park_calls: list[str] = [] + park_lock = threading.Lock() + + def run_one(qid: str) -> None: + barrier = threading.Barrier(2) + + def timer_worker() -> None: + c = connect(db) + try: + + def on_park(q: OverdueQuestion) -> None: + with park_lock: + park_calls.append(q.question_id) + + barrier.wait() + run_deadline_timer(c, on_park=on_park) + finally: + c.close() + + def responder_worker() -> None: + c = connect(db) + try: + barrier.wait() + answer_question( + c, question_id=qid, answer_json="{}", answered_via="slack" + ) + finally: + c.close() + + t1 = threading.Thread(target=timer_worker) + t2 = threading.Thread(target=responder_worker) + t1.start() + t2.start() + t1.join() + t2.join() + + for i in range(n): + run_one(f"race{i}") + + check = connect(db) + try: + rows = { + r["question_id"]: r["status"] + for r in check.execute( + "SELECT question_id, status FROM pending_questions" + ).fetchall() + } + finally: + check.close() + + # Every row ends in exactly one terminal state. + for i in range(n): + assert rows[f"race{i}"] in {"expired", "answered"} + # Every parked question must be one the timer actually expired. + for qid in park_calls: + assert rows[qid] == "expired" diff --git a/agent-team/tests/test_github_adapter.py b/agent-team/tests/test_github_adapter.py new file mode 100644 index 0000000..a663108 --- /dev/null +++ b/agent-team/tests/test_github_adapter.py @@ -0,0 +1,396 @@ +"""Unit tests for agent_team.transport.github_adapter (§3.3.1, §7.1 P4). + +Fully hermetic: the HTTP transport is dependency-injected with an in-memory +fake, so no network call, token, or live infrastructure is exercised. +""" + +from __future__ import annotations + +import json +from typing import Any + +import pytest + +from agent_team.transport.base import ( + GITHUB_MARKER_TEMPLATE, + QuestionSet, + Transport, +) +from agent_team.transport.github_adapter import ( + GITHUB_API_ROOT, + GitHubApiError, + GitHubTransport, + build_marker, + extract_question_id, + render_question_comment, +) + + +# --------------------------------------------------------------------------- # +# Fakes / fixtures +# --------------------------------------------------------------------------- # + + +class FakeHttpPost: + """In-memory ``HttpPost`` double recording calls and returning a scripted + ``(status, data)``.""" + + def __init__(self, status: int = 201, data: dict[str, Any] | None = None) -> None: + self.status = status + self.data = {"id": 987654321} if data is None else data + self.calls: list[dict[str, Any]] = [] + + def __call__( + self, + url: str, + *, + headers: dict[str, str], + json_body: dict[str, Any], + ) -> tuple[int, dict[str, Any]]: + self.calls.append({"url": url, "headers": headers, "json_body": json_body}) + return self.status, self.data + + +def make_transport( + http_post: FakeHttpPost | None = None, + *, + token: str | None = "ghp_fake", +) -> GitHubTransport: + return GitHubTransport( + owner="Sea-Haven-Industries", + repo="agent-team", + issue_number=42, + http_post=http_post or FakeHttpPost(), + token_provider=(lambda: token), + ) + + +def make_question_set() -> QuestionSet: + return QuestionSet( + thread_id="thread-1", + question_id="qid-abc", + turn=3, + questions=["Proceed with the migration?", "Which region?"], + context={"repo": "Sea-Haven-Industries/agent-team", "summary": "DB cutover"}, + ) + + +# --------------------------------------------------------------------------- # +# Contract / typing +# --------------------------------------------------------------------------- # + + +def test_is_transport_subclass() -> None: + assert issubclass(GitHubTransport, Transport) + + +def test_instantiable_concrete_adapter() -> None: + # Transport is abstract; the leaf must implement both abstract methods. + transport = make_transport() + assert isinstance(transport, Transport) + + +# --------------------------------------------------------------------------- # +# Marker helpers +# --------------------------------------------------------------------------- # + + +def test_build_marker_uses_foundation_template() -> None: + assert build_marker("xyz") == GITHUB_MARKER_TEMPLATE.format(question_id="xyz") + assert build_marker("xyz") == "" + + +def test_extract_question_id_roundtrips_build_marker() -> None: + qid = "deadbeef1234" + assert extract_question_id(build_marker(qid)) == qid + + +def test_extract_question_id_found_in_quoted_reply() -> None: + quoted = "> \n> ### Agent-team needs input\n\nYes, go ahead." + assert extract_question_id(quoted) == "q-77" + + +def test_extract_question_id_absent_returns_none() -> None: + assert extract_question_id("just a normal comment") is None + assert extract_question_id("") is None + + +# --------------------------------------------------------------------------- # +# Rendering +# --------------------------------------------------------------------------- # + + +def test_render_embeds_marker_and_questions() -> None: + qs = make_question_set() + body = render_question_comment( + qs, question_id="qid-abc", turn=3, deadline="2026-06-18T00:00:00Z" + ) + assert "" in body + assert extract_question_id(body) == "qid-abc" + assert "1. Proceed with the migration?" in body + assert "2. Which region?" in body + assert "turn 3" in body + assert "2026-06-18T00:00:00Z" in body + # Context surfaced. + assert "Sea-Haven-Industries/agent-team" in body + assert "DB cutover" in body + + +def test_render_handles_empty_questions() -> None: + qs = QuestionSet(thread_id="t", question_id="q", turn=0, questions=[]) + body = render_question_comment( + qs, question_id="q", turn=0, deadline="2026-06-18T00:00:00Z" + ) + assert "(no questions)" in body + assert extract_question_id(body) == "q" + + +# --------------------------------------------------------------------------- # +# post_question +# --------------------------------------------------------------------------- # + + +def test_post_question_returns_comment_id_as_channel_ref() -> None: + http = FakeHttpPost(status=201, data={"id": 555}) + transport = make_transport(http) + ref = transport.post_question( + thread_id="thread-1", + question_id="qid-abc", + turn=3, + question_set=make_question_set(), + deadline="2026-06-18T00:00:00Z", + ) + assert ref == "555" + assert isinstance(ref, str) + + +def test_post_question_targets_correct_issue_endpoint() -> None: + http = FakeHttpPost() + transport = make_transport(http) + transport.post_question( + thread_id="thread-1", + question_id="qid-abc", + turn=1, + question_set=make_question_set(), + deadline="2026-06-18T00:00:00Z", + ) + (call,) = http.calls + assert call["url"] == ( + f"{GITHUB_API_ROOT}/repos/Sea-Haven-Industries/agent-team/issues/42/comments" + ) + + +def test_post_question_body_embeds_question_id() -> None: + http = FakeHttpPost() + transport = make_transport(http) + transport.post_question( + thread_id="thread-1", + question_id="qid-abc", + turn=1, + question_set=make_question_set(), + deadline="2026-06-18T00:00:00Z", + ) + body = http.calls[0]["json_body"]["body"] + assert extract_question_id(body) == "qid-abc" + + +def test_post_question_sends_bearer_auth_header() -> None: + http = FakeHttpPost() + transport = make_transport(http, token="ghp_secret_value") + transport.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=make_question_set(), + deadline="d", + ) + headers = http.calls[0]["headers"] + assert headers["Authorization"] == "Bearer ghp_secret_value" + assert headers["Accept"] == "application/vnd.github+json" + assert headers["X-GitHub-Api-Version"] == "2022-11-28" + + +def test_post_question_raises_on_missing_token() -> None: + transport = make_transport(FakeHttpPost(), token=None) + with pytest.raises(GitHubApiError) as exc: + transport.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=make_question_set(), + deadline="d", + ) + assert exc.value.status == 401 + + +def test_post_question_reads_token_from_env(monkeypatch: pytest.MonkeyPatch) -> None: + http = FakeHttpPost() + monkeypatch.setenv("GITHUB_TOKEN", "ghp_from_env") + transport = GitHubTransport(owner="o", repo="r", issue_number=1, http_post=http) + transport.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=make_question_set(), + deadline="d", + ) + assert http.calls[0]["headers"]["Authorization"] == "Bearer ghp_from_env" + + +def test_post_question_raises_on_non_2xx() -> None: + http = FakeHttpPost(status=422, data={"message": "Validation Failed"}) + transport = make_transport(http) + with pytest.raises(GitHubApiError) as exc: + transport.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=make_question_set(), + deadline="d", + ) + assert exc.value.status == 422 + assert "Validation Failed" in exc.value.body + + +def test_post_question_raises_when_response_missing_id() -> None: + http = FakeHttpPost(status=201, data={"no_id": True}) + transport = make_transport(http) + with pytest.raises(GitHubApiError): + transport.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=make_question_set(), + deadline="d", + ) + + +def test_comments_url_respects_custom_api_root() -> None: + transport = GitHubTransport( + owner="o", + repo="r", + issue_number=7, + http_post=FakeHttpPost(), + api_root="https://ghe.example.com/api/v3/", + token_provider=lambda: "t", + ) + assert transport.comments_url == ( + "https://ghe.example.com/api/v3/repos/o/r/issues/7/comments" + ) + + +# --------------------------------------------------------------------------- # +# parse_answer +# --------------------------------------------------------------------------- # + + +def test_parse_answer_full_webhook_shape() -> None: + transport = make_transport() + raw = { + "comment": { + "body": "\nYes, proceed with us-east-1.", + "user": {"login": "amoussa1229"}, + } + } + qid, answer, via = transport.parse_answer(raw) + assert qid == "qid-abc" + assert answer == "Yes, proceed with us-east-1." + assert via == "github:amoussa1229" + + +def test_parse_answer_flattened_shape() -> None: + transport = make_transport() + raw = { + "body": " looks good", + "user": {"login": "adam"}, + } + qid, answer, via = transport.parse_answer(raw) + assert qid == "q9" + assert answer == "looks good" + assert via == "github:adam" + + +def test_parse_answer_strips_quoted_marker_lines() -> None: + transport = make_transport() + raw = { + "comment": { + "body": ( + "> \n" + "> ### Agent-team needs input (turn 1)\n" + "\n" + "Approved. Use the staging bucket." + ), + "user": {"login": "adam"}, + } + } + qid, answer, via = transport.parse_answer(raw) + assert qid == "q-77" + assert "\nok"}} + ) + assert qid == "q1" + assert via == "github" + + +# --------------------------------------------------------------------------- # +# Round-trip: post then parse the reply maps to the same question_id +# --------------------------------------------------------------------------- # + + +def test_post_then_answer_roundtrip_question_id() -> None: + http = FakeHttpPost() + transport = make_transport(http) + transport.post_question( + thread_id="thread-1", + question_id="round-trip-qid", + turn=2, + question_set=make_question_set(), + deadline="2026-06-18T00:00:00Z", + ) + posted_body = http.calls[0]["json_body"]["body"] + # Simulate a human quoting the posted comment in their reply. + reply_body = "\n".join(f"> {line}" for line in posted_body.splitlines()) + reply_body += "\n\nYes." + qid, answer, via = transport.parse_answer( + {"comment": {"body": reply_body, "user": {"login": "adam"}}} + ) + assert qid == "round-trip-qid" + assert answer == "Yes." + + +def test_default_http_post_is_not_called_in_tests() -> None: + # Sanity: the adapter never falls back to the network when a poster is + # injected (guards against an accidental live call in CI). + http = FakeHttpPost() + transport = make_transport(http) + transport.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=make_question_set(), + deadline="d", + ) + assert len(http.calls) == 1 + + +def test_api_error_str_is_truncated_and_typed() -> None: + err = GitHubApiError(500, "x" * 1000) + assert err.status == 500 + assert isinstance(err, RuntimeError) + # __str__ truncates the body to keep logs bounded. + assert len(str(err)) < 300 + assert json.dumps({"ok": True}) # keep json import meaningful/no-op diff --git a/agent-team/tests/test_graph.py b/agent-team/tests/test_graph.py new file mode 100644 index 0000000..fa6ed27 --- /dev/null +++ b/agent-team/tests/test_graph.py @@ -0,0 +1,296 @@ +"""Unit tests for agent_team.graph (Plane-2 P1 LangGraph wiring; §3.3, §7.1). + +These exercise the P1 skeleton + human gate wiring: + +* the pure node functions (intake/clarify-author/plan) in isolation, +* graph assembly + edge topology, +* the suspend-on-interrupt / resume-with-Command mechanic end to end, +* the thread_id-keyed driver seam (start/resume/get_state/pending_question), +* that the foundation contracts (PipelineState / Phase / TaskStatus / + QuestionSet) are imported verbatim and round-trip through the wiring. + +An in-memory checkpointer is injected (the SQLite checkpointer is the +production store, D9, not constructed in pre-deploy scaffolding). +""" + +from __future__ import annotations + +import pytest + +# InMemorySaver is the modern name; fall back to MemorySaver on older langgraph. +try: # pragma: no cover - import shim + from langgraph.checkpoint.memory import InMemorySaver as _Saver +except ImportError: # pragma: no cover - import shim + from langgraph.checkpoint.memory import MemorySaver as _Saver + +from agent_team import graph as graph_mod +from agent_team.graph import ( + CLARIFY, + INTAKE, + P1_PHASE_SEQUENCE, + PLAN, + build_graph, + build_sqlite_checkpointer, + clarify_node, + get_pipeline_state, + intake_node, + pending_question, + plan_node, + plan_phase, + resume_task, + start_task, + thread_config, +) +from agent_team.task_model import Phase, PipelineState, TaskStatus +from agent_team.transport import QuestionSet + + +@pytest.fixture() +def compiled(): + """A graph compiled with a fresh in-memory checkpointer per test.""" + return build_graph(checkpointer=_Saver()) + + +# --- Module surface / constants. ------------------------------------------- + + +def test_node_name_constants_are_distinct() -> None: + assert len({INTAKE, CLARIFY, PLAN}) == 3 + + +def test_p1_phase_sequence_stops_at_plan() -> None: + # P1 ends at an approved plan — no BUILD/VERIFY in the wired sequence (§7.1). + assert P1_PHASE_SEQUENCE == (Phase.INTAKE, Phase.CLARIFY, Phase.PLAN) + assert Phase.BUILD not in P1_PHASE_SEQUENCE + assert Phase.VERIFY not in P1_PHASE_SEQUENCE + + +# --- Pure node behaviour. --------------------------------------------------- + + +def test_intake_node_activates_and_advances_to_clarify() -> None: + out = intake_node(PipelineState(thread_id="t", current_phase=Phase.INTAKE.value)) + assert out["status"] == TaskStatus.ACTIVE.value + assert out["current_phase"] == Phase.CLARIFY.value + assert out["updated_at"] + + +def test_plan_node_lands_approved_plan_and_finishes() -> None: + out = plan_node(PipelineState(thread_id="t", qa_history=[{"answer": "x"}])) + assert out["status"] == TaskStatus.DONE.value + assert out["current_phase"] == Phase.DONE.value + assert out["plan"]["approved"] is True + + +def test_plan_phase_counts_qa_turns() -> None: + state = PipelineState(qa_history=[{"answer": "a"}, {"answer": "b"}]) + plan = plan_phase(state) + assert plan["qa_turns"] == 2 + assert plan["approved"] is True + + +def test_plan_phase_handles_empty_history() -> None: + assert plan_phase(PipelineState())["qa_turns"] == 0 + + +def test_clarify_node_suspends_rather_than_falling_through() -> None: + # Called bare (no running graph), interrupt() refuses to return a value: + # it raises because there is no runnable context to suspend into. This + # confirms clarify_node genuinely suspends rather than falling through to + # its post-interrupt return. + with pytest.raises(RuntimeError): + clarify_node(PipelineState(thread_id="t", transport="slack")) + + +# --- Graph assembly. -------------------------------------------------------- + + +def test_build_graph_without_checkpointer_compiles() -> None: + # An uncheckpointed graph still compiles (used only for straight-through + # smoke paths); the driver requires a checkpointer for suspend/resume. + assert build_graph() is not None + + +def test_build_graph_with_checkpointer_compiles(compiled) -> None: + assert compiled is not None + + +def test_graph_nodes_present(compiled) -> None: + nodes = set(compiled.get_graph().nodes) + assert {INTAKE, CLARIFY, PLAN} <= nodes + + +# --- Suspend / resume end to end. ------------------------------------------ + + +def test_start_task_suspends_on_human_gate(compiled) -> None: + thread_id, state = start_task(compiled, transport="slack") + # The task ran INTAKE then suspended at CLARIFY's interrupt(). + assert "__interrupt__" in state + payload = pending_question(compiled, thread_id=thread_id) + assert payload is not None + assert payload["thread_id"] == thread_id + assert payload["transport"] == "slack" + assert payload["turn"] == 0 + assert payload["deadline"] + + +def test_pending_question_carries_foundation_questionset(compiled) -> None: + thread_id, _ = start_task(compiled, transport="slack") + payload = pending_question(compiled, thread_id=thread_id) + qset = payload["question_set"] + # Verbatim foundation contract — not a redefinition. + assert isinstance(qset, QuestionSet) + assert qset.thread_id == thread_id + assert qset.question_id == payload["question_id"] + assert qset.turn == 0 + assert qset.questions # non-empty question-set + + +def test_resume_drives_task_to_done(compiled) -> None: + thread_id, _ = start_task(compiled, transport="slack") + final = resume_task(compiled, thread_id=thread_id, answer={"text": "do the thing"}) + assert final["status"] == TaskStatus.DONE.value + assert final["current_phase"] == Phase.DONE.value + assert final["plan"]["approved"] is True + + +def test_answer_is_recorded_in_qa_history(compiled) -> None: + thread_id, _ = start_task(compiled, transport="slack") + answer = {"text": "scope is X"} + final = resume_task(compiled, thread_id=thread_id, answer=answer) + assert len(final["qa_history"]) == 1 + assert final["qa_history"][0]["answer"] == answer + assert final["qa_history"][0]["turn"] == 0 + + +def test_question_id_is_stable_across_resume(compiled) -> None: + # The clarifier node re-executes on resume; the question_id must NOT change + # between the id delivered at suspend (the ledger key) and the one recorded + # in qa_history, or the §3.3.1 identity contract breaks. + thread_id, _ = start_task(compiled, transport="slack") + delivered = pending_question(compiled, thread_id=thread_id)["question_id"] + final = resume_task(compiled, thread_id=thread_id, answer="ok") + assert final["qa_history"][0]["question_id"] == delivered + + +def test_no_pending_question_after_completion(compiled) -> None: + thread_id, _ = start_task(compiled, transport="slack") + resume_task(compiled, thread_id=thread_id, answer="ok") + assert pending_question(compiled, thread_id=thread_id) is None + + +def test_get_pipeline_state_reflects_suspend_then_done(compiled) -> None: + thread_id, _ = start_task(compiled, transport="slack") + mid = get_pipeline_state(compiled, thread_id=thread_id) + # Suspended ON the clarifier gate: INTAKE already advanced the phase to + # CLARIFY, and the clarifier's post-interrupt write (-> PLAN) has NOT yet + # committed because the node is paused at interrupt(). Task is mid-flight. + assert mid["current_phase"] == Phase.CLARIFY.value + assert mid["status"] == TaskStatus.ACTIVE.value + resume_task(compiled, thread_id=thread_id, answer="ok") + done = get_pipeline_state(compiled, thread_id=thread_id) + assert done["status"] == TaskStatus.DONE.value + assert done["current_phase"] == Phase.DONE.value + + +# --- Thread isolation (§3.3.1 P1 exit criterion (d)). ---------------------- + + +def test_two_tasks_suspend_and_resume_independently(compiled) -> None: + t1, _ = start_task(compiled, transport="slack") + t2, _ = start_task(compiled, transport="github") + + assert t1 != t2 + p1 = pending_question(compiled, thread_id=t1) + p2 = pending_question(compiled, thread_id=t2) + assert p1["transport"] == "slack" + assert p2["transport"] == "github" + assert p1["question_id"] != p2["question_id"] + + # Resume only t1; t2 must remain suspended on its own gate. + f1 = resume_task(compiled, thread_id=t1, answer="answer-1") + assert f1["status"] == TaskStatus.DONE.value + assert pending_question(compiled, thread_id=t2) is not None + + f2 = resume_task(compiled, thread_id=t2, answer="answer-2") + assert f2["status"] == TaskStatus.DONE.value + assert f2["qa_history"][0]["answer"] == "answer-2" + + +def test_explicit_thread_id_is_honoured(compiled) -> None: + tid, _ = start_task(compiled, thread_id="fixed-thread", transport="slack") + assert tid == "fixed-thread" + assert pending_question(compiled, thread_id="fixed-thread") is not None + + +# --- Durable resume across a fresh graph object (P1 exit criterion (a)). ---- + + +def test_resume_works_on_a_new_graph_over_shared_checkpointer() -> None: + # Simulates a process restart: a NEW compiled graph object built over the + # SAME checkpointer must resume a task suspended by the first graph object. + saver = _Saver() + g1 = build_graph(checkpointer=saver) + thread_id, _ = start_task(g1, transport="slack") + + g2 = build_graph(checkpointer=saver) # "after restart" + assert pending_question(g2, thread_id=thread_id) is not None + final = resume_task(g2, thread_id=thread_id, answer="post-restart") + assert final["status"] == TaskStatus.DONE.value + assert final["qa_history"][0]["answer"] == "post-restart" + + +# --- Driver-seam helpers. --------------------------------------------------- + + +def test_thread_config_shape() -> None: + assert thread_config("abc") == {"configurable": {"thread_id": "abc"}} + + +def test_start_task_mints_unique_thread_ids(compiled) -> None: + t1, _ = start_task(compiled, transport="slack") + t2, _ = start_task(compiled, transport="slack") + assert t1 != t2 + + +# --- Production checkpointer factory. -------------------------------------- + + +def test_build_sqlite_checkpointer_missing_dep_raises_runtimeerror( + monkeypatch, tmp_path +) -> None: + # When the optional langgraph-checkpoint-sqlite package is absent, the + # factory must fail loudly with a clear RuntimeError, never silently run + # uncheckpointed. Force the ImportError path deterministically. + import builtins + + real_import = builtins.__import__ + + def _blocking_import(name, *args, **kwargs): + if name == "langgraph.checkpoint.sqlite": + raise ImportError("blocked for test") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _blocking_import) + with pytest.raises(RuntimeError, match="SQLite checkpointer"): + build_sqlite_checkpointer(tmp_path / "state.db") + + +def test_build_sqlite_checkpointer_builds_when_dep_present(tmp_path) -> None: + # If the optional package IS installed, the factory returns a checkpointer + # over the DB path. Skip cleanly where it's absent (pre-deploy scaffolding). + pytest.importorskip("langgraph.checkpoint.sqlite") + saver = build_sqlite_checkpointer(tmp_path / "nested" / "state.db") + assert saver is not None + assert (tmp_path / "nested").is_dir() + + +# --- Module import hygiene. ------------------------------------------------- + + +def test_module_imports_without_optional_sqlite_dep() -> None: + # The module-level import of graph must not pull in the optional SQLite + # checkpointer (that import is deferred into build_sqlite_checkpointer). + assert hasattr(graph_mod, "build_graph") + assert hasattr(graph_mod, "build_sqlite_checkpointer") diff --git a/agent-team/tests/test_ledger.py b/agent-team/tests/test_ledger.py new file mode 100644 index 0000000..144c36d --- /dev/null +++ b/agent-team/tests/test_ledger.py @@ -0,0 +1,394 @@ +"""Unit tests for agent_team.ledger — pending-questions ledger ops (§3.3.1).""" + +from __future__ import annotations + +import sqlite3 +import threading +from pathlib import Path + +import pytest + +from agent_team.db.schema import connect, init_db +from agent_team.ledger import ( + QUESTION_STATES, + PendingQuestion, + answer_question, + answered_questions, + count_by_status, + expire_question, + get_question, + list_questions, + open_questions_needing_ref, + overdue_open_questions, + post_question, + set_channel_ref, + supersede_question, +) + + +@pytest.fixture() +def conn(tmp_path: Path) -> sqlite3.Connection: + """A connection to an initialized agent-team DB.""" + db = tmp_path / "ledger.sqlite" + init_db(db) + connection = connect(db) + yield connection + connection.close() + + +# -------------------------------------------------------------------------- +# Re-export contract: the ledger exposes the foundation primitives verbatim. +# -------------------------------------------------------------------------- + + +def test_reexports_are_the_foundation_objects() -> None: + from agent_team.db import schema + + assert answer_question is schema.answer_question + assert expire_question is schema.expire_question + assert supersede_question is schema.supersede_question + assert QUESTION_STATES is schema.QUESTION_STATES + + +# -------------------------------------------------------------------------- +# post_question — write the row `open` first, no channel_ref (delivery step 1). +# -------------------------------------------------------------------------- + + +def test_post_question_writes_open_row_without_ref(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="q1", thread_id="t1", turn=0, transport="slack") + q = get_question(conn, "q1") + assert q is not None + assert q.status == "open" + assert q.channel_ref is None + assert q.transport == "slack" + assert q.posted_at # defaulted to now + assert q.deadline_at is None + + +def test_post_question_records_deadline_and_posted_at( + conn: sqlite3.Connection, +) -> None: + post_question( + conn, + question_id="q1", + thread_id="t1", + turn=2, + transport="github", + deadline_at="2026-06-17T12:00:00+00:00", + posted_at="2026-06-17T11:00:00+00:00", + ) + q = get_question(conn, "q1") + assert q is not None + assert q.turn == 2 + assert q.deadline_at == "2026-06-17T12:00:00+00:00" + assert q.posted_at == "2026-06-17T11:00:00+00:00" + + +def test_post_question_duplicate_id_raises(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="dup", thread_id="t", turn=0, transport="slack") + with pytest.raises(sqlite3.IntegrityError): + post_question(conn, question_id="dup", thread_id="t", turn=1, transport="slack") + + +# -------------------------------------------------------------------------- +# set_channel_ref — delivery step 2, guarded on status='open'. +# -------------------------------------------------------------------------- + + +def test_set_channel_ref_on_open_row(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="q1", thread_id="t", turn=0, transport="slack") + assert set_channel_ref(conn, question_id="q1", channel_ref="1700.0001") is True + assert get_question(conn, "q1").channel_ref == "1700.0001" + + +def test_set_channel_ref_unknown_id_returns_false(conn: sqlite3.Connection) -> None: + assert set_channel_ref(conn, question_id="nope", channel_ref="x") is False + + +def test_set_channel_ref_refuses_non_open(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="q1", thread_id="t", turn=0, transport="slack") + answer_question(conn, question_id="q1", answer_json="{}", answered_via="slack") + # Late post-confirm must not resurrect a ref on an answered question. + assert set_channel_ref(conn, question_id="q1", channel_ref="late") is False + assert get_question(conn, "q1").channel_ref is None + + +# -------------------------------------------------------------------------- +# Compare-and-set ops integrate with ledger-inserted rows (first-answer-wins). +# -------------------------------------------------------------------------- + + +def test_answer_first_wins_on_posted_row(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="q1", thread_id="t", turn=0, transport="slack") + first = answer_question( + conn, question_id="q1", answer_json='{"a":1}', answered_via="slack" + ) + second = answer_question( + conn, question_id="q1", answer_json='{"a":2}', answered_via="github" + ) + assert (first, second) == (True, False) + q = get_question(conn, "q1") + assert q.status == "answered" + assert q.answer_json == '{"a":1}' + assert q.answered_via == "slack" + assert q.answered_at + + +# -------------------------------------------------------------------------- +# open_questions_needing_ref — lost-post reconcile feed. +# -------------------------------------------------------------------------- + + +def test_open_questions_needing_ref(conn: sqlite3.Connection) -> None: + post_question( + conn, + question_id="no-ref", + thread_id="t", + turn=0, + transport="slack", + posted_at="2026-06-17T01:00:00+00:00", + ) + post_question( + conn, + question_id="with-ref", + thread_id="t", + turn=1, + transport="slack", + posted_at="2026-06-17T02:00:00+00:00", + ) + set_channel_ref(conn, question_id="with-ref", channel_ref="ts") + # answered rows (even without a ref) are not delivery-reconcile candidates + post_question( + conn, + question_id="answered", + thread_id="t", + turn=2, + transport="slack", + posted_at="2026-06-17T03:00:00+00:00", + ) + answer_question(conn, question_id="answered", answer_json="{}", answered_via="x") + + ids = [q.question_id for q in open_questions_needing_ref(conn)] + assert ids == ["no-ref"] + + +# -------------------------------------------------------------------------- +# overdue_open_questions + deadline-vs-answer race (§3.3.1). +# -------------------------------------------------------------------------- + + +def test_overdue_open_questions_filters_by_deadline( + conn: sqlite3.Connection, +) -> None: + post_question( + conn, + question_id="past", + thread_id="t", + turn=0, + transport="slack", + deadline_at="2026-06-17T10:00:00+00:00", + ) + post_question( + conn, + question_id="future", + thread_id="t", + turn=1, + transport="slack", + deadline_at="2026-06-17T20:00:00+00:00", + ) + post_question( + conn, + question_id="no-deadline", + thread_id="t", + turn=2, + transport="slack", + ) + overdue = overdue_open_questions(conn, now="2026-06-17T12:00:00+00:00") + assert [q.question_id for q in overdue] == ["past"] + + +def test_overdue_excludes_already_closed(conn: sqlite3.Connection) -> None: + post_question( + conn, + question_id="q", + thread_id="t", + turn=0, + transport="slack", + deadline_at="2026-06-17T10:00:00+00:00", + ) + answer_question(conn, question_id="q", answer_json="{}", answered_via="slack") + assert overdue_open_questions(conn, now="2026-06-17T12:00:00+00:00") == [] + + +def test_answer_after_expire_loses(conn: sqlite3.Connection) -> None: + """Deterministic deadline-vs-answer race: expiry first, then answer loses.""" + post_question( + conn, + question_id="q", + thread_id="t", + turn=0, + transport="slack", + deadline_at="2026-06-17T10:00:00+00:00", + ) + (overdue,) = overdue_open_questions(conn, now="2026-06-17T12:00:00+00:00") + assert expire_question(conn, question_id=overdue.question_id) is True + assert ( + answer_question(conn, question_id="q", answer_json="{}", answered_via="slack") + is False + ) + assert get_question(conn, "q").status == "expired" + + +# -------------------------------------------------------------------------- +# answered_questions — resume-worker / restart-recovery feed. +# -------------------------------------------------------------------------- + + +def test_answered_questions_feed(conn: sqlite3.Connection) -> None: + for qid, tid in (("a", "t1"), ("b", "t2")): + post_question(conn, question_id=qid, thread_id=tid, turn=0, transport="slack") + answer_question(conn, question_id="a", answer_json="{}", answered_via="slack") + # b stays open + answered = answered_questions(conn) + assert [q.question_id for q in answered] == ["a"] + # thread scoping + assert answered_questions(conn, thread_id="t2") == [] + assert [q.question_id for q in answered_questions(conn, thread_id="t1")] == ["a"] + + +# -------------------------------------------------------------------------- +# list_questions — manual CLI feed. +# -------------------------------------------------------------------------- + + +def test_list_questions_orders_oldest_first(conn: sqlite3.Connection) -> None: + post_question( + conn, + question_id="newer", + thread_id="t", + turn=1, + transport="slack", + posted_at="2026-06-17T05:00:00+00:00", + ) + post_question( + conn, + question_id="older", + thread_id="t", + turn=0, + transport="slack", + posted_at="2026-06-17T01:00:00+00:00", + ) + assert [q.question_id for q in list_questions(conn)] == ["older", "newer"] + + +def test_list_questions_status_filter(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="open1", thread_id="t", turn=0, transport="slack") + post_question(conn, question_id="ans1", thread_id="t", turn=1, transport="slack") + answer_question(conn, question_id="ans1", answer_json="{}", answered_via="slack") + assert [q.question_id for q in list_questions(conn, status="open")] == ["open1"] + assert [q.question_id for q in list_questions(conn, status="answered")] == ["ans1"] + + +def test_list_questions_thread_filter(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="a", thread_id="t1", turn=0, transport="slack") + post_question(conn, question_id="b", thread_id="t2", turn=0, transport="slack") + assert [q.question_id for q in list_questions(conn, thread_id="t1")] == ["a"] + + +def test_list_questions_rejects_unknown_status(conn: sqlite3.Connection) -> None: + with pytest.raises(ValueError): + list_questions(conn, status="bogus") + + +# -------------------------------------------------------------------------- +# count_by_status — stable shape over all states. +# -------------------------------------------------------------------------- + + +def test_count_by_status_stable_shape(conn: sqlite3.Connection) -> None: + post_question(conn, question_id="o1", thread_id="t", turn=0, transport="slack") + post_question(conn, question_id="o2", thread_id="t", turn=1, transport="slack") + post_question(conn, question_id="a1", thread_id="t", turn=2, transport="slack") + answer_question(conn, question_id="a1", answer_json="{}", answered_via="slack") + + counts = count_by_status(conn) + assert set(counts) == set(QUESTION_STATES) + assert counts["open"] == 2 + assert counts["answered"] == 1 + assert counts["expired"] == 0 + assert counts["superseded"] == 0 + + +# -------------------------------------------------------------------------- +# get_question + PendingQuestion view. +# -------------------------------------------------------------------------- + + +def test_get_question_missing_returns_none(conn: sqlite3.Connection) -> None: + assert get_question(conn, "ghost") is None + + +def test_pending_question_from_row(conn: sqlite3.Connection) -> None: + post_question( + conn, + question_id="q", + thread_id="t", + turn=3, + transport="github", + deadline_at="2026-06-17T12:00:00+00:00", + ) + set_channel_ref(conn, question_id="q", channel_ref="cref") + q = get_question(conn, "q") + assert isinstance(q, PendingQuestion) + assert (q.question_id, q.thread_id, q.turn, q.transport) == ( + "q", + "t", + 3, + "github", + ) + assert q.channel_ref == "cref" + # frozen dataclass — read snapshot, not mutable. + with pytest.raises(Exception): + q.status = "answered" # type: ignore[misc] + + +# -------------------------------------------------------------------------- +# Concurrency: two threads racing to set the channel_ref via the open guard. +# -------------------------------------------------------------------------- + + +def test_concurrent_answer_single_winner_via_ledger(tmp_path: Path) -> None: + db = tmp_path / "race.sqlite" + init_db(db) + seed = connect(db) + try: + post_question( + seed, question_id="race", thread_id="t", turn=0, transport="slack" + ) + finally: + seed.close() + + results: list[bool] = [] + barrier = threading.Barrier(2) + lock = threading.Lock() + + def worker(via: str) -> None: + c = connect(db) + try: + barrier.wait() + won = answer_question( + c, question_id="race", answer_json='{"v":1}', answered_via=via + ) + with lock: + results.append(won) + finally: + c.close() + + threads = [threading.Thread(target=worker, args=(f"c{i}",)) for i in range(2)] + for t in threads: + t.start() + for t in threads: + t.join() + + assert sorted(results) == [False, True] diff --git a/agent-team/tests/test_operator_cli.py b/agent-team/tests/test_operator_cli.py new file mode 100644 index 0000000..4d21d2d --- /dev/null +++ b/agent-team/tests/test_operator_cli.py @@ -0,0 +1,513 @@ +"""Unit tests for agent_team.operator_cli (§3.3.1, §6.6). + +Covers the two design-mandated invariants — every destructive action is +audit-logged, and destructive actions require an explicit confirmation flag — +plus the ledger lifecycle effects (which delegate to the committed +compare-and-set helpers) and the argparse entrypoint. +""" + +from __future__ import annotations + +import json +import stat +from pathlib import Path + +import pytest + +from agent_team.db.schema import connect, init_db +from agent_team.operator_cli import ( + DESTRUCTIVE_ACTIONS, + AuditEntry, + AuditLog, + ConfirmationRequired, + OperatorCli, + QuestionNotFound, + build_parser, + main, +) + + +# --------------------------------------------------------------------------- # +# fixtures / helpers +# --------------------------------------------------------------------------- # + + +def _insert_open_question( + db_path: Path, + qid: str, + *, + thread_id: str = "thread-1", + turn: int = 0, + status: str = "open", + transport: str = "slack", + channel_ref: str | None = "slack-ts-1", +) -> None: + conn = connect(db_path) + try: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, channel_ref, " + " posted_at) " + "VALUES (?, ?, ?, ?, ?, ?, ?)", + ( + qid, + thread_id, + turn, + status, + transport, + channel_ref, + "2026-06-17T00:00:00", + ), + ) + finally: + conn.close() + + +@pytest.fixture() +def db_path(tmp_path: Path) -> Path: + db = tmp_path / "agent_team.sqlite" + init_db(db) + return db + + +@pytest.fixture() +def audit_path(tmp_path: Path) -> Path: + return tmp_path / "audit" / "operator.jsonl" + + +@pytest.fixture() +def cli(db_path: Path, audit_path: Path) -> OperatorCli: + operator = OperatorCli(db_path, audit_path, actor="tester") + yield operator + operator.close() + + +def _status(db_path: Path, qid: str) -> str: + conn = connect(db_path) + try: + return conn.execute( + "SELECT status FROM pending_questions WHERE question_id=?", (qid,) + ).fetchone()["status"] + finally: + conn.close() + + +# --------------------------------------------------------------------------- # +# AuditEntry / AuditLog +# --------------------------------------------------------------------------- # + + +def test_audit_entry_to_dict_is_json_safe() -> None: + entry = AuditEntry( + timestamp="2026-06-17T00:00:00+00:00", + actor="tester", + action="force-expire", + phase="attempt", + confirmed=True, + question_id="q1", + thread_id="t1", + detail={"k": "v"}, + ) + data = json.loads(entry.to_json()) + assert data["action"] == "force-expire" + assert data["confirmed"] is True + assert data["detail"] == {"k": "v"} + + +def test_audit_entry_is_frozen() -> None: + entry = AuditEntry( + timestamp="t", actor="a", action="x", phase="attempt", confirmed=False + ) + with pytest.raises(Exception): + entry.actor = "other" # type: ignore[misc] + + +def test_audit_log_append_is_jsonl_and_appends(audit_path: Path) -> None: + log = AuditLog(audit_path) + log.append( + AuditEntry( + timestamp="t1", actor="a", action="x", phase="attempt", confirmed=True + ) + ) + log.append( + AuditEntry( + timestamp="t2", actor="a", action="y", phase="outcome", confirmed=True + ) + ) + records = log.read_all() + assert [r["action"] for r in records] == ["x", "y"] + # raw file is one JSON object per line + lines = audit_path.read_text().splitlines() + assert len(lines) == 2 + assert json.loads(lines[0])["timestamp"] == "t1" + + +def test_audit_log_file_is_mode_600(audit_path: Path) -> None: + log = AuditLog(audit_path) + log.append( + AuditEntry( + timestamp="t", actor="a", action="x", phase="attempt", confirmed=True + ) + ) + mode = stat.S_IMODE(audit_path.stat().st_mode) + assert mode == 0o600 + + +def test_audit_log_read_all_empty_when_absent(audit_path: Path) -> None: + assert AuditLog(audit_path).read_all() == [] + + +# --------------------------------------------------------------------------- # +# confirmation gating (§3.3.1: explicit confirmation flag) +# --------------------------------------------------------------------------- # + + +@pytest.mark.parametrize("action", sorted(DESTRUCTIVE_ACTIONS)) +def test_destructive_actions_set_matches_design(action: str) -> None: + assert action in {"force-expire", "answer-on-behalf", "force-resume"} + + +def test_force_expire_without_confirm_raises_and_no_mutation( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1") + with pytest.raises(ConfirmationRequired): + cli.force_expire("q1") + assert _status(db_path, "q1") == "open" # not mutated + + +def test_answer_on_behalf_without_confirm_raises_and_no_mutation( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1") + with pytest.raises(ConfirmationRequired): + cli.answer_on_behalf("q1", {"approve": True}) + assert _status(db_path, "q1") == "open" + + +def test_force_resume_without_confirm_raises_and_no_mutation( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1", status="answered") + with pytest.raises(ConfirmationRequired): + cli.force_resume("q1") + assert _status(db_path, "q1") == "answered" + + +def test_refused_action_is_audit_logged(cli: OperatorCli, db_path: Path) -> None: + _insert_open_question(db_path, "q1") + with pytest.raises(ConfirmationRequired): + cli.force_expire("q1") + records = cli.audit_log.read_all() + phases = [(r["action"], r["phase"]) for r in records] + # attempt + refusal outcome both recorded + assert ("force-expire", "attempt") in phases + assert ("force-expire", "outcome") in phases + refusal = [r for r in records if r["phase"] == "outcome"][0] + assert refusal["confirmed"] is False + assert "refused" in refusal["detail"] + + +# --------------------------------------------------------------------------- # +# destructive actions: audit-logged on success (§3.3.1: audit-logged) +# --------------------------------------------------------------------------- # + + +def test_force_expire_confirmed_mutates_and_audits( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1") + result = cli.force_expire("q1", confirm=True) + assert result.ok is True + assert _status(db_path, "q1") == "expired" + records = cli.audit_log.read_all() + actions = [(r["action"], r["phase"], r["confirmed"]) for r in records] + assert ("force-expire", "attempt", True) in actions + assert ("force-expire", "outcome", True) in actions + outcome = [r for r in records if r["phase"] == "outcome"][0] + assert outcome["question_id"] == "q1" + assert outcome["thread_id"] == "thread-1" + assert outcome["detail"]["changed"] is True + + +def test_answer_on_behalf_confirmed_writes_first_answer_wins( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1") + result = cli.answer_on_behalf("q1", {"approve": True}, confirm=True) + assert result.ok is True + conn = connect(db_path) + try: + row = conn.execute( + "SELECT status, answer_json, answered_via FROM pending_questions " + "WHERE question_id='q1'" + ).fetchone() + finally: + conn.close() + assert row["status"] == "answered" + assert json.loads(row["answer_json"]) == {"approve": True} + assert row["answered_via"] == "operator:tester" + # answer payload + via captured in audit attempt detail + attempt = [ + r + for r in cli.audit_log.read_all() + if r["action"] == "answer-on-behalf" and r["phase"] == "attempt" + ][0] + assert attempt["detail"]["answer"] == {"approve": True} + assert attempt["detail"]["via"] == "operator:tester" + + +def test_answer_on_behalf_late_loses_compare_and_set( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1", status="expired") + result = cli.answer_on_behalf("q1", "yes", confirm=True) + assert result.ok is False # already closed -> ignored + assert _status(db_path, "q1") == "expired" + + +def test_force_resume_confirmed_supersedes_and_records_intent( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1", status="answered") + result = cli.force_resume("q1", confirm=True) + assert result.ok is True + assert result.detail["resume_requested"] is True + assert result.detail["superseded"] is True + assert _status(db_path, "q1") == "superseded" + outcome = [ + r + for r in cli.audit_log.read_all() + if r["action"] == "force-resume" and r["phase"] == "outcome" + ][0] + assert outcome["detail"]["resume_requested"] is True + assert outcome["thread_id"] == "thread-1" + + +def test_force_resume_with_no_open_question_still_records_intent( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1", status="expired") + result = cli.force_resume("q1", confirm=True) + assert result.ok is True + assert result.detail["superseded"] is False + assert result.detail["resume_requested"] is True + + +# --------------------------------------------------------------------------- # +# redeliver (non-destructive, still audited) +# --------------------------------------------------------------------------- # + + +def test_redeliver_clears_channel_ref_and_audits( + cli: OperatorCli, db_path: Path +) -> None: + _insert_open_question(db_path, "q1", channel_ref="slack-ts-99") + result = cli.redeliver("q1") + assert result.ok is True + assert result.detail["prior_channel_ref"] == "slack-ts-99" + conn = connect(db_path) + try: + ref = conn.execute( + "SELECT channel_ref FROM pending_questions WHERE question_id='q1'" + ).fetchone()["channel_ref"] + finally: + conn.close() + assert ref is None + actions = [r["action"] for r in cli.audit_log.read_all()] + assert actions.count("redeliver") == 2 # attempt + outcome + + +def test_redeliver_non_open_is_noop(cli: OperatorCli, db_path: Path) -> None: + _insert_open_question(db_path, "q1", status="answered") + result = cli.redeliver("q1") + assert result.ok is False + + +def test_redeliver_does_not_require_confirm(cli: OperatorCli, db_path: Path) -> None: + # redeliver is NOT in the destructive set, so no confirm needed + assert "redeliver" not in DESTRUCTIVE_ACTIONS + _insert_open_question(db_path, "q1") + cli.redeliver("q1") # must not raise + + +# --------------------------------------------------------------------------- # +# listing (read-only) +# --------------------------------------------------------------------------- # + + +def test_list_questions_defaults_and_filters(cli: OperatorCli, db_path: Path) -> None: + _insert_open_question(db_path, "q1", thread_id="t1", status="open") + _insert_open_question(db_path, "q2", thread_id="t2", status="expired") + _insert_open_question(db_path, "q3", thread_id="t1", status="open") + + open_rows = cli.list_questions(statuses=["open"]) + assert {r["question_id"] for r in open_rows} == {"q1", "q3"} + + t1_rows = cli.list_questions(statuses=["open", "expired"], thread_id="t1") + assert {r["question_id"] for r in t1_rows} == {"q1", "q3"} + + all_rows = cli.list_questions() + assert len(all_rows) == 3 + + +def test_list_is_not_audited_as_mutation(cli: OperatorCli, db_path: Path) -> None: + _insert_open_question(db_path, "q1") + cli.list_questions(statuses=["open"]) + # listing does not write audit rows + assert cli.audit_log.read_all() == [] + + +# --------------------------------------------------------------------------- # +# missing question +# --------------------------------------------------------------------------- # + + +def test_force_expire_missing_question_raises_after_attempt_logged( + cli: OperatorCli, +) -> None: + with pytest.raises(QuestionNotFound): + cli.force_expire("nope", confirm=True) + # the attempt is logged even though the question doesn't exist + attempts = [r for r in cli.audit_log.read_all() if r["phase"] == "attempt"] + assert any(r["question_id"] == "nope" for r in attempts) + + +# --------------------------------------------------------------------------- # +# argparse entrypoint (main) +# --------------------------------------------------------------------------- # + + +def test_build_parser_requires_db_and_audit_log() -> None: + parser = build_parser() + with pytest.raises(SystemExit): + parser.parse_args(["list"]) # missing --db/--audit-log + + +def test_main_list_returns_zero(db_path: Path, audit_path: Path, capsys) -> None: + _insert_open_question(db_path, "q1") + rc = main( + [ + "--db", + str(db_path), + "--audit-log", + str(audit_path), + "list", + "--status", + "open", + ] + ) + assert rc == 0 + out = json.loads(capsys.readouterr().out) + assert out[0]["question_id"] == "q1" + + +def test_main_force_expire_requires_confirm_flag( + db_path: Path, audit_path: Path, capsys +) -> None: + _insert_open_question(db_path, "q1") + rc = main( + ["--db", str(db_path), "--audit-log", str(audit_path), "force-expire", "q1"] + ) + assert rc == 2 # ConfirmationRequired exit code + assert "refused" in capsys.readouterr().err + assert _status(db_path, "q1") == "open" + + +def test_main_force_expire_with_confirm(db_path: Path, audit_path: Path) -> None: + _insert_open_question(db_path, "q1") + rc = main( + [ + "--db", + str(db_path), + "--audit-log", + str(audit_path), + "--actor", + "adam", + "force-expire", + "q1", + "--confirm", + ] + ) + assert rc == 0 + assert _status(db_path, "q1") == "expired" + records = AuditLog(audit_path).read_all() + assert any(r["actor"] == "adam" for r in records) + + +def test_main_answer_on_behalf_parses_json_answer( + db_path: Path, audit_path: Path +) -> None: + _insert_open_question(db_path, "q1") + rc = main( + [ + "--db", + str(db_path), + "--audit-log", + str(audit_path), + "answer-on-behalf", + "q1", + '{"approve": true}', + "--confirm", + ] + ) + assert rc == 0 + conn = connect(db_path) + try: + answer_json = conn.execute( + "SELECT answer_json FROM pending_questions WHERE question_id='q1'" + ).fetchone()["answer_json"] + finally: + conn.close() + assert json.loads(answer_json) == {"approve": True} + + +def test_main_answer_on_behalf_raw_string_answer( + db_path: Path, audit_path: Path +) -> None: + _insert_open_question(db_path, "q1") + rc = main( + [ + "--db", + str(db_path), + "--audit-log", + str(audit_path), + "answer-on-behalf", + "q1", + "approve", + "--confirm", + ] + ) + assert rc == 0 + conn = connect(db_path) + try: + answer_json = conn.execute( + "SELECT answer_json FROM pending_questions WHERE question_id='q1'" + ).fetchone()["answer_json"] + finally: + conn.close() + assert json.loads(answer_json) == "approve" + + +def test_main_missing_question_returns_error_code( + db_path: Path, audit_path: Path, capsys +) -> None: + rc = main( + [ + "--db", + str(db_path), + "--audit-log", + str(audit_path), + "force-resume", + "ghost", + "--confirm", + ] + ) + assert rc == 3 + assert "error" in capsys.readouterr().err + + +def test_main_redeliver_returns_one_when_noop(db_path: Path, audit_path: Path) -> None: + _insert_open_question(db_path, "q1", status="answered") + rc = main(["--db", str(db_path), "--audit-log", str(audit_path), "redeliver", "q1"]) + assert rc == 1 # not open -> ok=False -> exit 1 diff --git a/agent-team/tests/test_planner.py b/agent-team/tests/test_planner.py new file mode 100644 index 0000000..4b6c39a --- /dev/null +++ b/agent-team/tests/test_planner.py @@ -0,0 +1,294 @@ +"""Unit tests for agent_team.nodes.planner (§3.3, §7.1 P2).""" + +from __future__ import annotations + +import json +from typing import Any + +import pytest + +from agent_team import billing +from agent_team.billing import BillingMode, ClaudeResult +from agent_team.nodes.planner import ( + MAX_PLAN_REVISIONS, + PlannerError, + build_plan_prompt, + parse_plan, + plan_node, +) +from agent_team.task_model import Phase, PipelineState, TaskStatus + +# --------------------------------------------------------------------------- # +# Fixtures / helpers +# --------------------------------------------------------------------------- # + +_VALID_PLAN = { + "summary": "Bump the dependency and update the lockfile.", + "phases": [ + {"name": "Phase 1 — bump", "steps": ["edit requirements", "run tests"]}, + {"name": "Phase 2 — verify", "steps": ["open draft PR"]}, + ], +} + + +@pytest.fixture(autouse=True) +def _restore_invoker(): + """Restore the module invoker after each test (mirrors test_billing).""" + original = billing._invoker + yield + billing._invoker = original + + +def _bind_invoker(reply: str) -> list[dict[str, Any]]: + """Bind a fake Claude invoker returning ``reply``; capture its calls.""" + calls: list[dict[str, Any]] = [] + + def _fake(prompt: str, *, mode: BillingMode, **kw: Any) -> ClaudeResult: + calls.append({"prompt": prompt, "mode": mode, "kw": kw}) + return ClaudeResult(text=reply, mode=mode, usage={"input_tokens": 1}) + + billing.set_invoker(_fake) + return calls + + +def _state(**overrides: Any) -> PipelineState: + base: PipelineState = PipelineState( + thread_id="t-1", + status=TaskStatus.ACTIVE.value, + current_phase=Phase.PLAN.value, + qa_history=[], + review_verdicts=[], + ) + base.update(overrides) # type: ignore[typeddict-item] + return base + + +# --------------------------------------------------------------------------- # +# parse_plan +# --------------------------------------------------------------------------- # + + +def test_parse_plan_valid() -> None: + plan = parse_plan(json.dumps(_VALID_PLAN)) + assert plan["summary"].startswith("Bump") + assert len(plan["phases"]) == 2 + assert plan["phases"][0]["name"] == "Phase 1 — bump" + assert plan["phases"][0]["steps"] == ["edit requirements", "run tests"] + + +def test_parse_plan_strips_code_fence() -> None: + fenced = "```json\n" + json.dumps(_VALID_PLAN) + "\n```" + plan = parse_plan(fenced) + assert len(plan["phases"]) == 2 + + +def test_parse_plan_strips_bare_code_fence() -> None: + fenced = "```\n" + json.dumps(_VALID_PLAN) + "\n```" + assert len(parse_plan(fenced)["phases"]) == 2 + + +def test_parse_plan_trims_step_whitespace_and_drops_blanks() -> None: + raw = {"phases": [{"name": "p", "steps": [" a ", "", " "]}]} + plan = parse_plan(json.dumps(raw)) + assert plan["phases"][0]["steps"] == ["a"] + assert plan["summary"] == "" + + +def test_parse_plan_empty_raises() -> None: + with pytest.raises(PlannerError, match="empty"): + parse_plan(" ") + + +def test_parse_plan_invalid_json_raises() -> None: + with pytest.raises(PlannerError, match="valid JSON"): + parse_plan("not json {") + + +def test_parse_plan_non_object_raises() -> None: + with pytest.raises(PlannerError, match="must be an object"): + parse_plan(json.dumps([1, 2, 3])) + + +def test_parse_plan_missing_phases_raises() -> None: + with pytest.raises(PlannerError, match="non-empty 'phases'"): + parse_plan(json.dumps({"summary": "x"})) + + +def test_parse_plan_empty_phases_raises() -> None: + with pytest.raises(PlannerError, match="non-empty 'phases'"): + parse_plan(json.dumps({"phases": []})) + + +def test_parse_plan_phase_missing_name_raises() -> None: + with pytest.raises(PlannerError, match="missing a non-empty 'name'"): + parse_plan(json.dumps({"phases": [{"steps": ["x"]}]})) + + +def test_parse_plan_phase_without_steps_raises() -> None: + with pytest.raises(PlannerError, match="no steps"): + parse_plan(json.dumps({"phases": [{"name": "p", "steps": []}]})) + + +def test_parse_plan_phase_with_only_blank_steps_raises() -> None: + with pytest.raises(PlannerError, match="no non-empty steps"): + parse_plan(json.dumps({"phases": [{"name": "p", "steps": ["", " "]}]})) + + +def test_parse_plan_phase_not_object_raises() -> None: + with pytest.raises(PlannerError, match="phase 1 must be an object"): + parse_plan(json.dumps({"phases": ["just a string"]})) + + +# --------------------------------------------------------------------------- # +# build_plan_prompt +# --------------------------------------------------------------------------- # + + +def test_build_prompt_includes_task_and_qa() -> None: + state = _state( + plan={"task": "Upgrade requests to 2.32"}, + qa_history=[{"question": "Pin exact?", "answer": "Yes, exact."}], + ) + prompt = build_plan_prompt(state) + assert "Upgrade requests to 2.32" in prompt + assert "Pin exact?" in prompt + assert "Yes, exact." in prompt + assert "PLANNER" in prompt + + +def test_build_prompt_task_fallback_from_top_level() -> None: + state = _state(task="Top-level task ask") # type: ignore[typeddict-unknown-key] + assert "Top-level task ask" in build_plan_prompt(state) + + +def test_build_prompt_handles_string_qa_entries() -> None: + state = _state(plan={"task": "t"}, qa_history=["freeform note"]) + assert "freeform note" in build_plan_prompt(state) + + +def test_build_prompt_no_task_uses_placeholder() -> None: + prompt = build_plan_prompt(_state()) + assert "(no task description provided)" in prompt + + +def test_build_prompt_loopback_includes_feedback_and_prior_plan() -> None: + state = _state( + plan={"task": "t", "phases": [{"name": "old", "steps": ["x"]}]}, + review_verdicts=[ + {"decision": "REQUEST_CHANGES", "notes": "Phase 1 missing rollback."} + ], + ) + prompt = build_plan_prompt(state) + assert "Reviewer feedback" in prompt + assert "Phase 1 missing rollback." in prompt + assert "Previous plan" in prompt + assert '"old"' in prompt + + +def test_build_prompt_no_feedback_omits_review_sections() -> None: + prompt = build_plan_prompt(_state(plan={"task": "t"})) + assert "Reviewer feedback" not in prompt + assert "Previous plan" not in prompt + + +# --------------------------------------------------------------------------- # +# plan_node — happy path +# --------------------------------------------------------------------------- # + + +def test_plan_node_produces_plan_and_advances_to_review() -> None: + _bind_invoker(json.dumps(_VALID_PLAN)) + out = plan_node(_state(plan={"task": "do a thing"})) + assert out["current_phase"] == Phase.REVIEW.value + assert out["status"] == TaskStatus.ACTIVE.value + assert out["plan"]["phases"][0]["name"] == "Phase 1 — bump" + assert out["plan"]["revision"] == 0 + + +def test_plan_node_returns_partial_state_only() -> None: + _bind_invoker(json.dumps(_VALID_PLAN)) + out = plan_node(_state(plan={"task": "x"})) + # A node returns only the keys it owns; it must not echo thread_id. + assert set(out.keys()) == {"plan", "current_phase", "status"} + + +def test_plan_node_forwards_config_to_billing_seam() -> None: + calls = _bind_invoker(json.dumps(_VALID_PLAN)) + plan_node(_state(plan={"task": "x"}), config={"billing_mode": "api"}) + assert calls[0]["mode"] is BillingMode.API + + +def test_plan_node_sends_task_into_prompt() -> None: + calls = _bind_invoker(json.dumps(_VALID_PLAN)) + plan_node(_state(plan={"task": "UNIQUE-TASK-MARKER"})) + assert "UNIQUE-TASK-MARKER" in calls[0]["prompt"] + + +def test_plan_node_garbled_reply_raises() -> None: + _bind_invoker("not json at all") + with pytest.raises(PlannerError): + plan_node(_state(plan={"task": "x"})) + + +# --------------------------------------------------------------------------- # +# plan_node — loop-back / convergence bound (§3.3, §6.6) +# --------------------------------------------------------------------------- # + + +def test_plan_node_replans_on_loopback_and_counts_revision() -> None: + _bind_invoker(json.dumps(_VALID_PLAN)) + state = _state( + plan={"task": "x"}, + review_verdicts=[{"decision": "REQUEST_CHANGES", "notes": "fix it"}], + ) + out = plan_node(state) + assert out["current_phase"] == Phase.REVIEW.value + assert out["plan"]["revision"] == 1 + + +def test_plan_node_parks_after_max_revisions() -> None: + # Invoker bound but must NOT be called once we are over the bound. + calls = _bind_invoker(json.dumps(_VALID_PLAN)) + verdicts = [ + {"decision": "REQUEST_CHANGES", "notes": f"round {i}"} + for i in range(MAX_PLAN_REVISIONS) + ] + out = plan_node(_state(plan={"task": "x"}, review_verdicts=verdicts)) + assert out["status"] == TaskStatus.PARKED.value + assert out["current_phase"] == Phase.PARKED.value + assert "plan" not in out + assert calls == [] # no Claude budget burned past the bound + + +def test_plan_node_does_not_park_just_below_bound() -> None: + _bind_invoker(json.dumps(_VALID_PLAN)) + verdicts = [ + {"decision": "REQUEST_CHANGES", "notes": f"round {i}"} + for i in range(MAX_PLAN_REVISIONS - 1) + ] + out = plan_node(_state(plan={"task": "x"}, review_verdicts=verdicts)) + assert out["current_phase"] == Phase.REVIEW.value + assert out["plan"]["revision"] == MAX_PLAN_REVISIONS - 1 + + +def test_plan_node_ignores_non_request_changes_verdicts_for_bound() -> None: + # APPROVE/other verdicts must not count toward the park bound. + _bind_invoker(json.dumps(_VALID_PLAN)) + verdicts = [{"decision": "APPROVE"}] * (MAX_PLAN_REVISIONS + 2) + out = plan_node(_state(plan={"task": "x"}, review_verdicts=verdicts)) + assert out["current_phase"] == Phase.REVIEW.value + assert out["plan"]["revision"] == 0 + + +def test_plan_node_counts_string_request_changes_verdicts() -> None: + _bind_invoker(json.dumps(_VALID_PLAN)) + verdicts: list[Any] = ["REQUEST_CHANGES"] * MAX_PLAN_REVISIONS + out = plan_node(_state(plan={"task": "x"}, review_verdicts=verdicts)) + assert out["status"] == TaskStatus.PARKED.value + + +def test_plan_node_unconfigured_invoker_raises() -> None: + # The foundation seam fails loudly when no invoker is wired. + billing._invoker = billing._unconfigured_invoker + with pytest.raises(RuntimeError, match="no invoker bound"): + plan_node(_state(plan={"task": "x"})) diff --git a/agent-team/tests/test_recovery.py b/agent-team/tests/test_recovery.py new file mode 100644 index 0000000..a0d95cf --- /dev/null +++ b/agent-team/tests/test_recovery.py @@ -0,0 +1,671 @@ +"""Unit tests for agent_team.recovery — the restart-recovery sweep (§3.3.1, §6.7). + +Exercises the three convergence steps (deadline / redeliver / resume-or-supersede), +their idempotency, the first-answer-wins races, post-restore reconciliation, and +per-row error isolation. Uses the committed foundation contracts verbatim +(``agent_team.db.schema`` for the ledger, ``agent_team.transport.base`` for the +transport ABC) — nothing here redefines a foundation interface. +""" + +from __future__ import annotations + +import json +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Any + +import pytest + +from agent_team.db.schema import ( + answer_question, + connect, + init_db, +) +from agent_team.recovery import ( + DeadlineOutcome, + PendingQuestion, + RecoveryReport, + apply_deadline_policy, + load_pending_questions, + redeliver_open_questions, + reenqueue_answered_resumes, + run_restart_recovery, +) +from agent_team.transport.base import QuestionSet, Transport + +UTC = timezone.utc + + +# --------------------------------------------------------------------------- # +# Fixtures & helpers # +# --------------------------------------------------------------------------- # +@pytest.fixture() +def db_path(tmp_path: Path) -> Path: + """A freshly initialized agent-team DB file.""" + path = tmp_path / "agent_team.db" + init_db(path) + return path + + +@pytest.fixture() +def conn(db_path: Path): + """An open connection to the initialized DB (closed at teardown).""" + connection = connect(db_path) + yield connection + connection.close() + + +def _insert( + connection, + *, + question_id: str, + thread_id: str = "t1", + turn: int = 0, + status: str = "open", + transport: str = "slack", + channel_ref: str | None = None, + posted_at: str | None = None, + deadline_at: str | None = None, + answer_json: str | None = None, + answered_via: str | None = None, +) -> None: + """Insert a raw ``pending_questions`` row for a test scenario.""" + connection.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, channel_ref, " + " posted_at, deadline_at, answer_json, answered_at, answered_via) " + "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + ( + question_id, + thread_id, + turn, + status, + transport, + channel_ref, + posted_at, + deadline_at, + answer_json, + None, + answered_via, + ), + ) + + +def _status(connection, question_id: str) -> str: + row = connection.execute( + "SELECT status FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + return row["status"] + + +def _channel_ref(connection, question_id: str) -> str | None: + row = connection.execute( + "SELECT channel_ref FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + return row["channel_ref"] + + +class RecordingTransport(Transport): + """A Transport that records posts and returns a deterministic channel_ref.""" + + def __init__(self, ref: str = "ts-123", *, fail: bool = False) -> None: + self.ref = ref + self.fail = fail + self.posts: list[dict[str, Any]] = [] + + def post_question( + self, + *, + thread_id: str, + question_id: str, + turn: int, + question_set: QuestionSet, + deadline: str, + ) -> str: + if self.fail: + raise RuntimeError("transport down") + self.posts.append( + { + "thread_id": thread_id, + "question_id": question_id, + "turn": turn, + "deadline": deadline, + } + ) + return self.ref + + def parse_answer(self, raw: Any) -> tuple[str, Any, str]: # pragma: no cover + raise NotImplementedError + + +def _park_policy(question: PendingQuestion) -> DeadlineOutcome: + return DeadlineOutcome(question_id=question.question_id, action="parked") + + +# --------------------------------------------------------------------------- # +# load_pending_questions / PendingQuestion # +# --------------------------------------------------------------------------- # +def test_load_pending_questions_filters_by_status(conn) -> None: + _insert(conn, question_id="q-open", status="open") + _insert(conn, question_id="q-ans", status="answered") + opens = load_pending_questions(conn, status="open") + assert [q.question_id for q in opens] == ["q-open"] + + +def test_load_pending_questions_no_filter_returns_all(conn) -> None: + _insert(conn, question_id="q1", status="open") + _insert(conn, question_id="q2", status="answered") + assert len(load_pending_questions(conn)) == 2 + + +def test_load_pending_questions_rejects_unknown_status(conn) -> None: + with pytest.raises(ValueError): + load_pending_questions(conn, status="bogus") + + +def test_pending_question_from_row_maps_columns(conn) -> None: + _insert( + conn, + question_id="q1", + thread_id="thread-x", + turn=3, + status="open", + transport="github", + channel_ref="ref-1", + deadline_at="2026-01-01T00:00:00+00:00", + ) + (q,) = load_pending_questions(conn, status="open") + assert q == PendingQuestion( + question_id="q1", + thread_id="thread-x", + turn=3, + status="open", + transport="github", + channel_ref="ref-1", + posted_at=None, + deadline_at="2026-01-01T00:00:00+00:00", + answer_json=None, + answered_at=None, + answered_via=None, + ) + + +# --------------------------------------------------------------------------- # +# Step 1 — redeliver lost posts # +# --------------------------------------------------------------------------- # +def test_redeliver_posts_open_row_without_ref(conn) -> None: + _insert(conn, question_id="q1", status="open", channel_ref=None) + transport = RecordingTransport(ref="slack-ts-9") + report = RecoveryReport() + redeliver_open_questions( + conn, resolve_transport=lambda _t: transport, report=report + ) + assert report.redelivered == ["q1"] + assert len(transport.posts) == 1 + assert _channel_ref(conn, "q1") == "slack-ts-9" + + +def test_redeliver_skips_row_that_already_has_ref(conn) -> None: + _insert(conn, question_id="q1", status="open", channel_ref="already-here") + transport = RecordingTransport() + report = RecoveryReport() + redeliver_open_questions( + conn, resolve_transport=lambda _t: transport, report=report + ) + assert report.redelivered == [] + assert transport.posts == [] + + +def test_redeliver_defers_when_transport_unreachable(conn) -> None: + _insert(conn, question_id="q1", status="open", channel_ref=None) + report = RecoveryReport() + redeliver_open_questions(conn, resolve_transport=lambda _t: None, report=report) + assert report.redelivery_deferred == ["q1"] + assert report.redelivered == [] + # Row stays open with no ref so a later sweep retries. + assert _status(conn, "q1") == "open" + assert _channel_ref(conn, "q1") is None + + +def test_redeliver_isolates_transport_exception(conn) -> None: + _insert(conn, question_id="q1", status="open", channel_ref=None) + transport = RecordingTransport(fail=True) + report = RecoveryReport() + redeliver_open_questions( + conn, resolve_transport=lambda _t: transport, report=report + ) + assert report.redelivery_deferred == ["q1"] + assert report.errors and report.errors[0][0] == "q1" + + +def test_redeliver_treats_empty_ref_as_deferred(conn) -> None: + _insert(conn, question_id="q1", status="open", channel_ref=None) + transport = RecordingTransport(ref="") + report = RecoveryReport() + redeliver_open_questions( + conn, resolve_transport=lambda _t: transport, report=report + ) + assert report.redelivery_deferred == ["q1"] + assert _channel_ref(conn, "q1") is None + + +# --------------------------------------------------------------------------- # +# Step 2 — re-enqueue resumes / supersede # +# --------------------------------------------------------------------------- # +def test_reenqueue_resumes_when_graph_still_interrupted(conn) -> None: + _insert(conn, question_id="q1", thread_id="t1", turn=2, status="answered") + enqueued: list[tuple[str, str, int]] = [] + report = RecoveryReport() + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda _tid, _turn: True, + enqueue_resume=lambda tid, qid, turn: enqueued.append((tid, qid, turn)) or True, + report=report, + ) + assert report.resumes_enqueued == ["q1"] + assert enqueued == [("t1", "q1", 2)] + # Row remains answered — the resume worker owns the terminal transition. + assert _status(conn, "q1") == "answered" + + +def test_reenqueue_supersedes_when_graph_advanced(conn) -> None: + _insert(conn, question_id="q1", thread_id="t1", turn=2, status="answered") + enqueued: list[Any] = [] + report = RecoveryReport() + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda _tid, _turn: False, + enqueue_resume=lambda *a: enqueued.append(a) or True, + report=report, + ) + assert report.superseded == ["q1"] + assert report.resumes_enqueued == [] + assert enqueued == [] + assert _status(conn, "q1") == "superseded" + + +def test_reenqueue_probe_receives_thread_and_turn(conn) -> None: + _insert(conn, question_id="q1", thread_id="thread-9", turn=7, status="answered") + seen: list[tuple[str, int]] = [] + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda tid, turn: seen.append((tid, turn)) or True, + enqueue_resume=lambda *a: True, + report=RecoveryReport(), + ) + assert seen == [("thread-9", 7)] + + +def test_reenqueue_isolates_probe_exception(conn) -> None: + _insert(conn, question_id="q1", status="answered") + + def boom(_tid: str, _turn: int) -> bool: + raise RuntimeError("probe failed") + + report = RecoveryReport() + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=boom, + enqueue_resume=lambda *a: True, + report=report, + ) + assert report.errors and "probe" in report.errors[0][1] + assert report.resumes_enqueued == [] + + +def test_reenqueue_isolates_enqueue_exception(conn) -> None: + _insert(conn, question_id="q1", status="answered") + + def boom(*_a: Any) -> bool: + raise RuntimeError("queue down") + + report = RecoveryReport() + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=boom, + report=report, + ) + assert report.errors and "resume" in report.errors[0][1] + + +# --------------------------------------------------------------------------- # +# Step 2 — post-restore reconciliation (§6.7) # +# --------------------------------------------------------------------------- # +class _Reconciler: + def __init__(self, *, safe: bool = True, raise_exc: bool = False) -> None: + self.safe = safe + self.raise_exc = raise_exc + self.calls: list[str] = [] + + def reconcile(self, question: PendingQuestion) -> bool: + self.calls.append(question.question_id) + if self.raise_exc: + raise RuntimeError("reconcile blew up") + return self.safe + + +def test_reconciler_allows_resume_when_external_state_consistent(conn) -> None: + _insert(conn, question_id="q1", status="answered") + reconciler = _Reconciler(safe=True) + enqueued: list[Any] = [] + report = RecoveryReport() + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: enqueued.append(a) or True, + report=report, + reconciler=reconciler, + ) + assert reconciler.calls == ["q1"] + assert report.resumes_enqueued == ["q1"] + + +def test_reconciler_holds_resume_when_external_state_diverged(conn) -> None: + _insert(conn, question_id="q1", status="answered") + reconciler = _Reconciler(safe=False) + enqueued: list[Any] = [] + report = RecoveryReport() + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: enqueued.append(a) or True, + report=report, + reconciler=reconciler, + ) + assert report.reconcile_held == ["q1"] + assert report.resumes_enqueued == [] + assert enqueued == [] + + +def test_reconciler_exception_holds_and_records_error(conn) -> None: + _insert(conn, question_id="q1", status="answered") + reconciler = _Reconciler(raise_exc=True) + report = RecoveryReport() + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: True, + report=report, + reconciler=reconciler, + ) + assert report.reconcile_held == ["q1"] + assert report.errors and "reconcile" in report.errors[0][1] + + +def test_reconciler_not_consulted_when_graph_advanced(conn) -> None: + _insert(conn, question_id="q1", status="answered") + reconciler = _Reconciler(safe=True) + reenqueue_answered_resumes( + conn, + is_interrupted_on_turn=lambda *a: False, + enqueue_resume=lambda *a: True, + report=RecoveryReport(), + reconciler=reconciler, + ) + # Superseded path never reaches reconciliation. + assert reconciler.calls == [] + + +# --------------------------------------------------------------------------- # +# Step 3 — deadline policy # +# --------------------------------------------------------------------------- # +def test_deadline_expires_overdue_open_question(conn) -> None: + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + _insert(conn, question_id="q1", status="open", deadline_at=past) + report = RecoveryReport() + apply_deadline_policy(conn, policy=_park_policy, report=report) + assert report.expired == ["q1"] + assert _status(conn, "q1") == "expired" + assert report.deadline_outcomes[0].action == "parked" + + +def test_deadline_leaves_future_question_open(conn) -> None: + future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + _insert(conn, question_id="q1", status="open", deadline_at=future) + report = RecoveryReport() + apply_deadline_policy(conn, policy=_park_policy, report=report) + assert report.expired == [] + assert _status(conn, "q1") == "open" + + +def test_deadline_ignores_row_without_deadline(conn) -> None: + _insert(conn, question_id="q1", status="open", deadline_at=None) + report = RecoveryReport() + apply_deadline_policy(conn, policy=_park_policy, report=report) + assert report.expired == [] + assert _status(conn, "q1") == "open" + + +def test_deadline_respects_injected_now(conn) -> None: + deadline = "2026-06-01T00:00:00+00:00" + _insert(conn, question_id="q1", status="open", deadline_at=deadline) + before = datetime(2026, 5, 1, tzinfo=UTC) + report = RecoveryReport() + apply_deadline_policy(conn, policy=_park_policy, report=report, now=before) + assert report.expired == [] # not yet overdue at injected now + assert _status(conn, "q1") == "open" + + +def test_deadline_treats_naive_timestamp_as_utc(conn) -> None: + past_naive = ( + (datetime.now(UTC) - timedelta(hours=2)).replace(tzinfo=None).isoformat() + ) + _insert(conn, question_id="q1", status="open", deadline_at=past_naive) + report = RecoveryReport() + apply_deadline_policy(conn, policy=_park_policy, report=report) + assert report.expired == ["q1"] + + +def test_deadline_ignores_unparseable_timestamp(conn) -> None: + _insert(conn, question_id="q1", status="open", deadline_at="not-a-date") + report = RecoveryReport() + apply_deadline_policy(conn, policy=_park_policy, report=report) + assert report.expired == [] + assert _status(conn, "q1") == "open" + + +def test_deadline_policy_not_invoked_when_already_answered(conn) -> None: + # An answer that won the race before the sweep: compare-and-set finds no + # open row, so no expiry and no policy call. + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + _insert(conn, question_id="q1", status="answered", deadline_at=past) + invoked: list[str] = [] + report = RecoveryReport() + apply_deadline_policy( + conn, + policy=lambda q: ( + invoked.append(q.question_id) or DeadlineOutcome(q.question_id, "parked") + ), + report=report, + ) + assert report.expired == [] + assert invoked == [] + + +def test_deadline_isolates_policy_exception(conn) -> None: + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + _insert(conn, question_id="q1", status="open", deadline_at=past) + + def boom(_q: PendingQuestion) -> DeadlineOutcome: + raise RuntimeError("policy failed") + + report = RecoveryReport() + apply_deadline_policy(conn, policy=boom, report=report) + # Row still durably expired even though the policy callback failed. + assert report.expired == ["q1"] + assert _status(conn, "q1") == "expired" + assert report.errors and "policy" in report.errors[0][1] + + +# --------------------------------------------------------------------------- # +# Full sweep orchestration # +# --------------------------------------------------------------------------- # +def test_run_restart_recovery_drives_all_three_steps(db_path: Path) -> None: + setup = connect(db_path) + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + # 1: overdue open -> expired + _insert(setup, question_id="q-late", status="open", deadline_at=past) + # 2: open w/o ref -> redelivered + _insert(setup, question_id="q-lost", status="open", channel_ref=None) + # 3: answered, graph still waits -> resume enqueued + _insert(setup, question_id="q-ans", thread_id="ta", turn=1, status="answered") + setup.close() + + transport = RecordingTransport(ref="ts-x") + enqueued: list[Any] = [] + report = run_restart_recovery( + db_path, + resolve_transport=lambda _t: transport, + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: enqueued.append(a) or True, + deadline_policy=_park_policy, + ) + assert report.expired == ["q-late"] + assert report.redelivered == ["q-lost"] + assert report.resumes_enqueued == ["q-ans"] + assert not report.clean + + +def test_run_restart_recovery_expires_before_redelivering(db_path: Path) -> None: + # An overdue open row must be expired by step 1, never redelivered by + # step 2 — proving deadline-first ordering. + setup = connect(db_path) + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + _insert( + setup, + question_id="q1", + status="open", + channel_ref=None, + deadline_at=past, + ) + setup.close() + + transport = RecordingTransport() + report = run_restart_recovery( + db_path, + resolve_transport=lambda _t: transport, + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: True, + deadline_policy=_park_policy, + ) + assert report.expired == ["q1"] + assert report.redelivered == [] + assert transport.posts == [] # never posted an already-expired question + + +def test_run_restart_recovery_clean_when_nothing_pending(db_path: Path) -> None: + report = run_restart_recovery( + db_path, + resolve_transport=lambda _t: RecordingTransport(), + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: True, + deadline_policy=_park_policy, + ) + assert report.clean + + +def test_run_restart_recovery_with_injected_conn(conn) -> None: + _insert(conn, question_id="q1", status="open", channel_ref=None) + report = run_restart_recovery( + ":memory:", # ignored because conn is injected + resolve_transport=lambda _t: RecordingTransport(ref="r"), + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: True, + deadline_policy=_park_policy, + conn=conn, + ) + assert report.redelivered == ["q1"] + # Injected connection is left open for the caller. + assert _status(conn, "q1") == "open" + + +def test_run_restart_recovery_post_restore_holds_diverged_task(db_path: Path) -> None: + setup = connect(db_path) + _insert(setup, question_id="q1", status="answered") + setup.close() + + reconciler = _Reconciler(safe=False) + enqueued: list[Any] = [] + report = run_restart_recovery( + db_path, + resolve_transport=lambda _t: RecordingTransport(), + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: enqueued.append(a) or True, + deadline_policy=_park_policy, + reconciler=reconciler, + ) + assert report.reconcile_held == ["q1"] + assert enqueued == [] + + +# --------------------------------------------------------------------------- # +# Idempotency — re-running the sweep converges, doesn't duplicate # +# --------------------------------------------------------------------------- # +def test_sweep_is_idempotent_for_redelivery(db_path: Path) -> None: + setup = connect(db_path) + _insert(setup, question_id="q1", status="open", channel_ref=None) + setup.close() + + kwargs: dict[str, Any] = dict( + resolve_transport=lambda _t: RecordingTransport(ref="r1"), + is_interrupted_on_turn=lambda *a: True, + enqueue_resume=lambda *a: True, + deadline_policy=_park_policy, + ) + first = run_restart_recovery(db_path, **kwargs) + second = run_restart_recovery(db_path, **kwargs) + assert first.redelivered == ["q1"] + # Second pass: row now has a ref, so nothing to redeliver -> clean. + assert second.redelivered == [] + assert second.clean + + +def test_first_answer_wins_against_concurrent_expiry(conn) -> None: + # A real first-answer-wins race: answer lands, then the deadline sweep + # runs. The compare-and-set protects the answered row from expiry. + past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() + _insert(conn, question_id="q1", status="open", deadline_at=past) + won = answer_question( + conn, + question_id="q1", + answer_json=json.dumps({"ok": True}), + answered_via="slack", + ) + assert won is True + report = RecoveryReport() + apply_deadline_policy(conn, policy=_park_policy, report=report) + assert report.expired == [] + assert _status(conn, "q1") == "answered" + + +# --------------------------------------------------------------------------- # +# RecoveryReport.clean # +# --------------------------------------------------------------------------- # +def test_report_clean_true_for_empty_report() -> None: + assert RecoveryReport().clean is True + + +@pytest.mark.parametrize( + "field_name", + [ + "redelivered", + "redelivery_deferred", + "resumes_enqueued", + "superseded", + "expired", + "reconcile_held", + ], +) +def test_report_not_clean_when_any_action_list_populated(field_name: str) -> None: + report = RecoveryReport() + getattr(report, field_name).append("q1") + assert report.clean is False + + +def test_report_not_clean_when_errors_present() -> None: + report = RecoveryReport() + report.errors.append(("q1", "boom")) + assert report.clean is False diff --git a/agent-team/tests/test_responder.py b/agent-team/tests/test_responder.py new file mode 100644 index 0000000..914b6ae --- /dev/null +++ b/agent-team/tests/test_responder.py @@ -0,0 +1,574 @@ +"""Unit tests for agent_team.responder (§3.3, §3.3.1). + +Covers the notify+resume seam end to end against the real ``pending_questions`` +ledger (an on-disk SQLite DB via the foundation ``init_db``/``connect``): + +* notify: ledger-row-first ordering, channel_ref persisted, lost-post leaves an + open row with no ref; +* submit_answer: first-answer-wins accept + enqueue, duplicate/late no-op, + answer-after-expiry loses the compare-and-set; +* ResumeWorker: turn-guard skip → superseded, happy-path resume, single-flight + per-thread serialization, concurrent different threads; +* deadline_sweep: overdue open → expired, race vs answer; +* recover_open_questions: answered rows re-enqueued. +""" + +from __future__ import annotations + +import sqlite3 +import threading +import time +from pathlib import Path +from typing import Any + +import pytest + +from agent_team.db.schema import ( + answer_question, + connect, + expire_question, + init_db, +) +from agent_team.responder import ( + AnswerOutcome, + GraphHandle, + ResumeJob, + ResumeWorker, + deadline_sweep, + notify_question, + recover_open_questions, + submit_answer, +) +from agent_team.transport.base import NormalizedAnswer, QuestionSet, Transport + + +# --------------------------------------------------------------------------- +# Fixtures + fakes. +# --------------------------------------------------------------------------- + + +@pytest.fixture +def conn(tmp_path: Path) -> sqlite3.Connection: + """A real ledger-backed connection (foundation schema).""" + db_path = tmp_path / "agent-team.db" + init_db(db_path) + connection = connect(db_path) + try: + yield connection + finally: + connection.close() + + +class FakeTransport(Transport): + """In-memory transport recording posts and parsing dict answers. + + ``post_question`` embeds the ``question_id`` in the returned ref (Slack + ``ts`` analogue). ``post_fails`` toggles the lost-post path. + """ + + def __init__(self, *, post_fails: bool = False) -> None: + self.posts: list[dict[str, Any]] = [] + self.post_fails = post_fails + + def post_question( + self, *, thread_id, question_id, turn, question_set, deadline + ) -> str: + if self.post_fails: + raise RuntimeError("transport unreachable") + ref = f"slack-ts-{question_id}" + self.posts.append( + { + "thread_id": thread_id, + "question_id": question_id, + "turn": turn, + "deadline": deadline, + "ref": ref, + } + ) + return ref + + def parse_answer(self, raw) -> tuple[str, Any, str]: + na = NormalizedAnswer( + question_id=raw["callback_id"], answer=raw["value"], via="slack" + ) + return na.question_id, na.answer, na.via + + +class FakeGraph: + """Structural ``GraphHandle``: configurable interrupt turn + resume recorder.""" + + def __init__(self, turns: dict[str, int | None] | None = None) -> None: + # thread_id -> turn it is interrupted on (None = not interrupted). + self.turns: dict[str, int | None] = turns or {} + self.resumed: list[tuple[str, Any]] = [] + self._resume_hook = None + + def interrupted_turn(self, thread_id: str) -> int | None: + return self.turns.get(thread_id) + + def resume(self, thread_id: str, answer: Any) -> Any: + if self._resume_hook is not None: + self._resume_hook(thread_id) + self.resumed.append((thread_id, answer)) + return {"resumed": thread_id} + + +def _question_set( + *, thread_id: str = "t1", question_id: str = "q1", turn: int = 0 +) -> QuestionSet: + return QuestionSet( + thread_id=thread_id, + question_id=question_id, + turn=turn, + questions=["proceed?"], + context={"repo": "x"}, + ) + + +def _row(conn: sqlite3.Connection, qid: str) -> sqlite3.Row: + return conn.execute( + "SELECT * FROM pending_questions WHERE question_id=?", (qid,) + ).fetchone() + + +# --------------------------------------------------------------------------- +# Protocol conformance. +# --------------------------------------------------------------------------- + + +def test_fake_graph_satisfies_protocol() -> None: + assert isinstance(FakeGraph(), GraphHandle) + + +# --------------------------------------------------------------------------- +# notify_question — delivery + lost-post (§3.3.1). +# --------------------------------------------------------------------------- + + +def test_notify_writes_open_row_then_stores_ref(conn: sqlite3.Connection) -> None: + transport = FakeTransport() + qs = _question_set() + + ref = notify_question(conn, transport, qs, deadline="2026-06-18T00:00:00+00:00") + + assert ref == "slack-ts-q1" + row = _row(conn, "q1") + assert row["status"] == "open" + assert row["channel_ref"] == "slack-ts-q1" + assert row["thread_id"] == "t1" + assert row["turn"] == 0 + assert row["transport"] == "FakeTransport" + assert row["deadline_at"] == "2026-06-18T00:00:00+00:00" + assert row["posted_at"] is not None + # The post carried the question_id so an answer can map back. + assert transport.posts[0]["question_id"] == "q1" + + +def test_notify_row_is_written_before_post(conn: sqlite3.Connection) -> None: + """The durable row must exist even while the post is in flight.""" + seen: dict[str, Any] = {} + + class CheckingTransport(FakeTransport): + def post_question(self, *, question_id, **kw): # type: ignore[override] + # At post time the ledger row must already be persisted as open. + seen["row"] = _row(conn, question_id) + return super().post_question(question_id=question_id, **kw) + + notify_question( + conn, CheckingTransport(), _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + assert seen["row"] is not None + assert seen["row"]["status"] == "open" + + +def test_notify_lost_post_leaves_open_row_without_ref( + conn: sqlite3.Connection, +) -> None: + transport = FakeTransport(post_fails=True) + + ref = notify_question( + conn, transport, _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + + # Post failed: no ref returned, row stays open with no ref for reconcile. + assert ref is None + row = _row(conn, "q1") + assert row["status"] == "open" + assert row["channel_ref"] is None + + +# --------------------------------------------------------------------------- +# submit_answer — first-answer-wins (§3.3.1). +# --------------------------------------------------------------------------- + + +def test_submit_answer_first_wins_enqueues_resume( + conn: sqlite3.Connection, +) -> None: + transport = FakeTransport() + notify_question( + conn, transport, _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + + enqueued: list[ResumeJob] = [] + outcome = submit_answer( + conn, + transport, + {"callback_id": "q1", "value": "yes"}, + enqueue_resume=enqueued.append, + ) + + assert isinstance(outcome, AnswerOutcome) + assert outcome.accepted is True + assert outcome.question_id == "q1" + assert outcome.via == "slack" + assert outcome.job is not None + assert outcome.job == ResumeJob( + thread_id="t1", question_id="q1", turn=0, answer="yes" + ) + assert enqueued == [outcome.job] + + row = _row(conn, "q1") + assert row["status"] == "answered" + assert row["answer_json"] == '"yes"' + assert row["answered_via"] == "slack" + assert row["answered_at"] is not None + + +def test_submit_answer_duplicate_is_noop(conn: sqlite3.Connection) -> None: + transport = FakeTransport() + notify_question( + conn, transport, _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + + enqueued: list[ResumeJob] = [] + first = submit_answer( + conn, + transport, + {"callback_id": "q1", "value": "yes"}, + enqueue_resume=enqueued.append, + ) + second = submit_answer( + conn, + transport, + {"callback_id": "q1", "value": "no"}, + enqueue_resume=enqueued.append, + ) + + assert first.accepted is True + assert second.accepted is False + assert second.job is None + # Only the first answer enqueued a resume; the duplicate is ignored. + assert len(enqueued) == 1 + # The stored answer is the first one, never overwritten by the duplicate. + assert _row(conn, "q1")["answer_json"] == '"yes"' + + +def test_submit_answer_after_expiry_loses_race(conn: sqlite3.Connection) -> None: + transport = FakeTransport() + notify_question( + conn, transport, _question_set(), deadline="2000-01-01T00:00:00+00:00" + ) + # Question times out first. + assert expire_question(conn, question_id="q1") is True + + enqueued: list[ResumeJob] = [] + outcome = submit_answer( + conn, + transport, + {"callback_id": "q1", "value": "yes"}, + enqueue_resume=enqueued.append, + ) + + assert outcome.accepted is False + assert enqueued == [] + assert _row(conn, "q1")["status"] == "expired" + + +def test_submit_answer_unknown_question_is_noop(conn: sqlite3.Connection) -> None: + transport = FakeTransport() + enqueued: list[ResumeJob] = [] + outcome = submit_answer( + conn, + transport, + {"callback_id": "nope", "value": "x"}, + enqueue_resume=enqueued.append, + ) + assert outcome.accepted is False + assert enqueued == [] + + +# --------------------------------------------------------------------------- +# ResumeWorker — single-flight, turn-guarded (§3.3.1). +# --------------------------------------------------------------------------- + + +def test_resume_happy_path(conn: sqlite3.Connection) -> None: + notify_question( + conn, FakeTransport(), _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + answer_question(conn, question_id="q1", answer_json='"yes"', answered_via="slack") + + graph = FakeGraph(turns={"t1": 0}) + worker = ResumeWorker(conn, graph) + job = ResumeJob(thread_id="t1", question_id="q1", turn=0, answer="yes") + + assert worker.run(job) is True + assert graph.resumed == [("t1", "yes")] + # Still 'answered' — the worker does not mutate the ledger on success. + assert _row(conn, "q1")["status"] == "answered" + + +def test_resume_stale_turn_supersedes_and_skips(conn: sqlite3.Connection) -> None: + notify_question( + conn, + FakeTransport(), + _question_set(turn=2), + deadline="2026-06-18T00:00:00+00:00", + ) + answer_question(conn, question_id="q1", answer_json='"yes"', answered_via="slack") + + # Graph already advanced to turn 3 (or is not on turn 2 any more). + graph = FakeGraph(turns={"t1": 3}) + worker = ResumeWorker(conn, graph) + job = ResumeJob(thread_id="t1", question_id="q1", turn=2, answer="yes") + + assert worker.run(job) is False + assert graph.resumed == [] + assert _row(conn, "q1")["status"] == "superseded" + + +def test_resume_not_interrupted_supersedes_and_skips( + conn: sqlite3.Connection, +) -> None: + notify_question( + conn, FakeTransport(), _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + answer_question(conn, question_id="q1", answer_json='"yes"', answered_via="slack") + + # Thread not currently interrupted (None). + graph = FakeGraph(turns={"t1": None}) + worker = ResumeWorker(conn, graph) + job = ResumeJob(thread_id="t1", question_id="q1", turn=0, answer="yes") + + assert worker.run(job) is False + assert graph.resumed == [] + assert _row(conn, "q1")["status"] == "superseded" + + +def test_resume_redelivered_job_does_not_double_apply( + conn: sqlite3.Connection, +) -> None: + """A second (redelivered) job for the same answered turn supersedes-skips.""" + notify_question( + conn, FakeTransport(), _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + answer_question(conn, question_id="q1", answer_json='"yes"', answered_via="slack") + + graph = FakeGraph(turns={"t1": 0}) + worker = ResumeWorker(conn, graph) + job = ResumeJob(thread_id="t1", question_id="q1", turn=0, answer="yes") + + assert worker.run(job) is True + + # The real graph advances after a successful resume; model that. + graph.turns["t1"] = 1 + assert worker.run(job) is False + # Resume applied exactly once. + assert graph.resumed == [("t1", "yes")] + + +def test_resume_single_flight_serializes_same_thread( + conn: sqlite3.Connection, +) -> None: + """Two jobs for one thread never resume concurrently (per-thread lock).""" + notify_question( + conn, FakeTransport(), _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + answer_question(conn, question_id="q1", answer_json='"yes"', answered_via="slack") + + graph = FakeGraph(turns={"t1": 0}) + concurrency = {"current": 0, "max": 0} + lock = threading.Lock() + + def hook(_thread_id: str) -> None: + with lock: + concurrency["current"] += 1 + concurrency["max"] = max(concurrency["max"], concurrency["current"]) + time.sleep(0.02) + with lock: + concurrency["current"] -= 1 + + graph._resume_hook = hook + worker = ResumeWorker(conn, graph) + job = ResumeJob(thread_id="t1", question_id="q1", turn=0, answer="yes") + + threads = [threading.Thread(target=worker.run, args=(job,)) for _ in range(5)] + for t in threads: + t.start() + for t in threads: + t.join() + + # Same thread_id => never more than one in-flight resume at a time. + assert concurrency["max"] == 1 + + +def test_resume_different_threads_run_concurrently( + conn: sqlite3.Connection, +) -> None: + """Different thread_ids are NOT serialized against each other.""" + graph = FakeGraph(turns={f"t{i}": 0 for i in range(4)}) + barrier = threading.Barrier(4, timeout=2.0) + reached = {"ok": True} + + def hook(_thread_id: str) -> None: + try: + barrier.wait() + except threading.BrokenBarrierError: + reached["ok"] = False + + graph._resume_hook = hook + worker = ResumeWorker(conn, graph) + jobs = [ + ResumeJob(thread_id=f"t{i}", question_id=f"q{i}", turn=0, answer="y") + for i in range(4) + ] + + threads = [threading.Thread(target=worker.run, args=(j,)) for j in jobs] + for t in threads: + t.start() + for t in threads: + t.join() + + # All four reached the barrier together => they ran concurrently. + assert reached["ok"] is True + assert len(graph.resumed) == 4 + + +# --------------------------------------------------------------------------- +# deadline_sweep — overdue open -> expired (§3.3.1). +# --------------------------------------------------------------------------- + + +def test_deadline_sweep_expires_only_overdue_open( + conn: sqlite3.Connection, +) -> None: + transport = FakeTransport() + notify_question( + conn, + transport, + _question_set(question_id="overdue"), + deadline="2000-01-01T00:00:00+00:00", + ) + notify_question( + conn, + transport, + _question_set(question_id="future"), + deadline="2099-01-01T00:00:00+00:00", + ) + + expired = deadline_sweep(conn, now="2026-06-17T00:00:00+00:00") + + assert expired == ["overdue"] + assert _row(conn, "overdue")["status"] == "expired" + assert _row(conn, "future")["status"] == "open" + + +def test_deadline_sweep_skips_already_answered( + conn: sqlite3.Connection, +) -> None: + transport = FakeTransport() + notify_question( + conn, + transport, + _question_set(question_id="ans"), + deadline="2000-01-01T00:00:00+00:00", + ) + answer_question(conn, question_id="ans", answer_json='"yes"', answered_via="slack") + + expired = deadline_sweep(conn, now="2026-06-17T00:00:00+00:00") + + # Already answered => the deadline race was already lost; not expired. + assert expired == [] + assert _row(conn, "ans")["status"] == "answered" + + +def test_deadline_sweep_ignores_null_deadline(conn: sqlite3.Connection) -> None: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport) " + "VALUES ('q', 't', 0, 'open', 'slack')" + ) + assert deadline_sweep(conn, now="2099-01-01T00:00:00+00:00") == [] + assert _row(conn, "q")["status"] == "open" + + +# --------------------------------------------------------------------------- +# recover_open_questions — startup answered->resume replay (§3.3.1). +# --------------------------------------------------------------------------- + + +def test_recover_reenqueues_answered_rows(conn: sqlite3.Connection) -> None: + transport = FakeTransport() + notify_question( + conn, + transport, + _question_set(thread_id="ta", question_id="qa"), + deadline="2026-06-18T00:00:00+00:00", + ) + notify_question( + conn, + transport, + _question_set(thread_id="tb", question_id="qb", turn=1), + deadline="2026-06-18T00:00:00+00:00", + ) + notify_question( + conn, + transport, + _question_set(thread_id="tc", question_id="qc"), + deadline="2026-06-18T00:00:00+00:00", + ) + + # qa, qb answered before a crash; qc still open. + answer_question(conn, question_id="qa", answer_json='"yes"', answered_via="slack") + answer_question( + conn, question_id="qb", answer_json='{"k": 1}', answered_via="github" + ) + + enqueued: list[ResumeJob] = [] + jobs = recover_open_questions(conn, enqueue_resume=enqueued.append) + + assert jobs == enqueued + by_thread = {j.thread_id: j for j in jobs} + assert set(by_thread) == {"ta", "tb"} + assert by_thread["ta"] == ResumeJob( + thread_id="ta", question_id="qa", turn=0, answer="yes" + ) + assert by_thread["tb"] == ResumeJob( + thread_id="tb", question_id="qb", turn=1, answer={"k": 1} + ) + + +def test_recover_is_idempotent_via_turn_guard(conn: sqlite3.Connection) -> None: + """Re-enqueued recover jobs no-op when the graph already advanced.""" + notify_question( + conn, FakeTransport(), _question_set(), deadline="2026-06-18T00:00:00+00:00" + ) + answer_question(conn, question_id="q1", answer_json='"yes"', answered_via="slack") + + enqueued: list[ResumeJob] = [] + recover_open_questions(conn, enqueue_resume=enqueued.append) + assert len(enqueued) == 1 + + # Graph already past turn 0 (resume happened before the crash record cleared). + graph = FakeGraph(turns={"t1": 1}) + worker = ResumeWorker(conn, graph) + assert worker.run(enqueued[0]) is False + assert graph.resumed == [] + assert _row(conn, "q1")["status"] == "superseded" + + +def test_recover_empty_ledger(conn: sqlite3.Connection) -> None: + enqueued: list[ResumeJob] = [] + assert recover_open_questions(conn, enqueue_resume=enqueued.append) == [] + assert enqueued == [] diff --git a/agent-team/tests/test_resume_worker.py b/agent-team/tests/test_resume_worker.py new file mode 100644 index 0000000..4cdb7c4 --- /dev/null +++ b/agent-team/tests/test_resume_worker.py @@ -0,0 +1,509 @@ +"""Unit tests for agent_team.resume_worker (§3.3.1 single-flight, turn-guarded). + +These tests prove the design's three resume guarantees: + +* turn-guarded: resume applies only while the graph is interrupted on the + answer's turn; a graph that already advanced is superseded and skipped; +* no double-apply: a redelivered job for an already-advanced thread no-ops; +* single-flight: resumes for one ``thread_id`` are serialized while different + threads run concurrently; + +plus the restart-recovery sweep, and an end-to-end check against a real +compiled LangGraph app when ``langgraph`` is importable. +""" + +from __future__ import annotations + +import json +import sqlite3 +import threading +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import pytest + +from agent_team import resume_worker +from agent_team.db.schema import answer_question, init_db, connect +from agent_team.resume_worker import ( + GraphLike, + ResumeOutcome, + ResumeResult, + ResumeWorker, + build_resume_command, + snapshot_interrupt_turns, +) + + +# --------------------------------------------------------------------------- # +# Fakes +# --------------------------------------------------------------------------- # + + +@dataclass +class _FakeInterrupt: + """Mimics a langgraph Interrupt: carries a ``.value`` payload.""" + + value: Any + + +@dataclass +class _FakeSnapshot: + """Mimics a langgraph StateSnapshot's relevant surface.""" + + next: tuple[str, ...] = () + interrupts: tuple[_FakeInterrupt, ...] = () + + +class _FakeGraph: + """A GraphLike test double over an explicit interrupt-turn. + + ``interrupted_turn`` is the turn the graph is currently suspended on, or + ``None`` if it has advanced past every interrupt. ``invoke`` records every + resume payload so double-apply is directly observable, and advances the + graph (clears the interrupt) the way a real resume would. + """ + + def __init__(self, interrupted_turn: int | None) -> None: + self.interrupted_turn = interrupted_turn + self.invocations: list[Any] = [] + self.get_state_calls: list[dict[str, Any]] = [] + self._invoke_hook: Any = None + + def get_state(self, config: dict[str, Any]) -> _FakeSnapshot: + self.get_state_calls.append(config) + if self.interrupted_turn is None: + return _FakeSnapshot(next=(), interrupts=()) + payload = {"turn": self.interrupted_turn, "question_id": "q"} + return _FakeSnapshot( + next=("clarify",), + interrupts=(_FakeInterrupt(value=payload),), + ) + + def invoke(self, command: Any, config: dict[str, Any]) -> Any: + if self._invoke_hook is not None: + self._invoke_hook() + self.invocations.append(command) + # A real resume clears the interrupt and advances the graph. + self.interrupted_turn = None + return {"resumed": True, "command": command} + + +def _insert_question( + conn: sqlite3.Connection, + *, + question_id: str, + thread_id: str, + turn: int, + status: str = "open", +) -> None: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport) " + "VALUES (?, ?, ?, ?, 'slack')", + (question_id, thread_id, turn, status), + ) + + +def _status(conn: sqlite3.Connection, question_id: str) -> str: + row = conn.execute( + "SELECT status FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + return row["status"] + + +@pytest.fixture() +def conn(tmp_path: Path) -> sqlite3.Connection: + db = tmp_path / "agent_team.sqlite" + init_db(db) + connection = connect(db) + yield connection + connection.close() + + +# --------------------------------------------------------------------------- # +# snapshot_interrupt_turns +# --------------------------------------------------------------------------- # + + +def test_snapshot_turns_from_mapping_payload() -> None: + snap = _FakeSnapshot( + next=("n",), + interrupts=(_FakeInterrupt(value={"turn": 7}),), + ) + assert snapshot_interrupt_turns(snap) == {7} + + +def test_snapshot_turns_from_object_payload() -> None: + @dataclass + class _ObjPayload: + turn: int + + snap = _FakeSnapshot(interrupts=(_FakeInterrupt(value=_ObjPayload(turn=2)),)) + assert snapshot_interrupt_turns(snap) == {2} + + +def test_snapshot_turns_empty_when_not_interrupted() -> None: + assert snapshot_interrupt_turns(_FakeSnapshot()) == set() + + +def test_snapshot_turns_ignores_unreadable_turn() -> None: + snap = _FakeSnapshot(interrupts=(_FakeInterrupt(value={"no_turn": 1}),)) + assert snapshot_interrupt_turns(snap) == set() + + +def test_snapshot_turns_ignores_bool_turn() -> None: + # bool is an int subclass; a True/False must not be read as a turn number. + snap = _FakeSnapshot(interrupts=(_FakeInterrupt(value={"turn": True}),)) + assert snapshot_interrupt_turns(snap) == set() + + +def test_snapshot_turns_collects_multiple() -> None: + snap = _FakeSnapshot( + interrupts=( + _FakeInterrupt(value={"turn": 1}), + _FakeInterrupt(value={"turn": 4}), + ) + ) + assert snapshot_interrupt_turns(snap) == {1, 4} + + +# --------------------------------------------------------------------------- # +# resume — turn guard +# --------------------------------------------------------------------------- # + + +def test_resume_applies_when_interrupted_on_turn(conn: sqlite3.Connection) -> None: + _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") + graph = _FakeGraph(interrupted_turn=3) + worker = ResumeWorker(graph, conn) + + result = worker.resume( + thread_id="t1", question_id="q1", turn=3, answer="the answer" + ) + + assert result.outcome is ResumeOutcome.RESUMED + assert result.resumed is True + assert result.graph_result == {"resumed": True, "command": graph.invocations[0]} + assert len(graph.invocations) == 1 + # Question is untouched by the worker on a successful resume (the responder + # already flipped it to answered). + assert _status(conn, "q1") == "answered" + + +def test_resume_superseded_when_graph_advanced(conn: sqlite3.Connection) -> None: + _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") + graph = _FakeGraph(interrupted_turn=None) # already advanced past turn 3 + worker = ResumeWorker(graph, conn) + + result = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="x") + + assert result.outcome is ResumeOutcome.SUPERSEDED + assert result.resumed is False + assert graph.invocations == [] # never invoked -> never applied + assert _status(conn, "q1") == "superseded" + + +def test_resume_superseded_when_interrupted_on_different_turn( + conn: sqlite3.Connection, +) -> None: + # Graph moved on to a *later* interrupt (turn 4); a turn-3 job is stale. + _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") + graph = _FakeGraph(interrupted_turn=4) + worker = ResumeWorker(graph, conn) + + result = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="x") + + assert result.outcome is ResumeOutcome.SUPERSEDED + assert graph.invocations == [] + assert _status(conn, "q1") == "superseded" + + +def test_resume_stale_when_nothing_to_supersede(conn: sqlite3.Connection) -> None: + # Question already expired; graph advanced. Nothing to supersede. + _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="expired") + graph = _FakeGraph(interrupted_turn=None) + worker = ResumeWorker(graph, conn) + + result = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="x") + + assert result.outcome is ResumeOutcome.STALE + assert graph.invocations == [] + assert _status(conn, "q1") == "expired" + + +# --------------------------------------------------------------------------- # +# no double-apply +# --------------------------------------------------------------------------- # + + +def test_redelivered_job_does_not_double_apply(conn: sqlite3.Connection) -> None: + _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") + graph = _FakeGraph(interrupted_turn=3) + worker = ResumeWorker(graph, conn) + + first = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="a") + # A redelivered/duplicate resume job for the same turn arrives. + second = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="a") + + assert first.outcome is ResumeOutcome.RESUMED + assert second.outcome is ResumeOutcome.SUPERSEDED + # The graph was invoked exactly once across both jobs. + assert len(graph.invocations) == 1 + assert _status(conn, "q1") == "superseded" + + +# --------------------------------------------------------------------------- # +# single-flight serialization +# --------------------------------------------------------------------------- # + + +def test_same_thread_resumes_are_serialized(conn: sqlite3.Connection) -> None: + _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") + graph = _FakeGraph(interrupted_turn=3) + + in_invoke = threading.Event() + overlap_detected: list[bool] = [] + concurrency = {"current": 0, "max": 0} + lock = threading.Lock() + + def _hook() -> None: + with lock: + concurrency["current"] += 1 + concurrency["max"] = max(concurrency["max"], concurrency["current"]) + in_invoke.set() + time.sleep(0.05) + with lock: + concurrency["current"] -= 1 + + graph._invoke_hook = _hook + worker = ResumeWorker(graph, conn) + + results: list[ResumeResult] = [] + results_lock = threading.Lock() + + def _run(answer: str) -> None: + r = worker.resume(thread_id="t1", question_id="q1", turn=3, answer=answer) + with results_lock: + results.append(r) + + threads = [threading.Thread(target=_run, args=(f"a{i}",)) for i in range(5)] + for t in threads: + t.start() + for t in threads: + t.join() + + # Lock serialized them: invoke never overlapped. + assert concurrency["max"] == 1 + assert overlap_detected == [] + # Exactly one resumed; the rest were superseded (turn guard) -> no double. + resumed = [r for r in results if r.outcome is ResumeOutcome.RESUMED] + assert len(resumed) == 1 + assert len(graph.invocations) == 1 + + +def test_distinct_threads_use_distinct_locks(conn: sqlite3.Connection) -> None: + graph_a = _FakeGraph(interrupted_turn=1) + graph_b = _FakeGraph(interrupted_turn=1) + # One worker can only hold one graph; emulate isolation by giving each + # thread its own worker over its own graph, sharing the ledger. + _insert_question(conn, question_id="qa", thread_id="ta", turn=1, status="answered") + _insert_question(conn, question_id="qb", thread_id="tb", turn=1, status="answered") + + worker_a = ResumeWorker(graph_a, conn) + worker_b = ResumeWorker(graph_b, conn) + + # Distinct thread_ids must mint distinct locks within a single worker. + single = ResumeWorker(_FakeGraph(interrupted_turn=1), conn) + assert single._lock_for("ta") is not single._lock_for("tb") + assert single._lock_for("ta") is single._lock_for("ta") + + ra = worker_a.resume(thread_id="ta", question_id="qa", turn=1, answer="x") + rb = worker_b.resume(thread_id="tb", question_id="qb", turn=1, answer="y") + assert ra.outcome is ResumeOutcome.RESUMED + assert rb.outcome is ResumeOutcome.RESUMED + + +# --------------------------------------------------------------------------- # +# restart recovery sweep +# --------------------------------------------------------------------------- # + + +def test_recover_resumes_answered_rows_still_interrupted( + conn: sqlite3.Connection, +) -> None: + _insert_question(conn, question_id="q1", thread_id="t1", turn=2, status="open") + answer_question( + conn, question_id="q1", answer_json=json.dumps("ans"), answered_via="slack" + ) + graph = _FakeGraph(interrupted_turn=2) + worker = ResumeWorker(graph, conn) + + results = worker.recover_pending_resumes() + + assert len(results) == 1 + assert results[0].outcome is ResumeOutcome.RESUMED + assert results[0].thread_id == "t1" + # The decoded answer reached the graph as a Command(resume=...). + assert len(graph.invocations) == 1 + + +def test_recover_is_idempotent_when_graph_already_advanced( + conn: sqlite3.Connection, +) -> None: + _insert_question(conn, question_id="q1", thread_id="t1", turn=2, status="open") + answer_question( + conn, question_id="q1", answer_json=json.dumps("ans"), answered_via="slack" + ) + # Graph already advanced (the resume applied before the crash). + graph = _FakeGraph(interrupted_turn=None) + worker = ResumeWorker(graph, conn) + + results = worker.recover_pending_resumes() + + assert len(results) == 1 + assert results[0].outcome is ResumeOutcome.SUPERSEDED + assert graph.invocations == [] # no double-apply across a restart + assert _status(conn, "q1") == "superseded" + + +def test_recover_skips_non_answered_rows(conn: sqlite3.Connection) -> None: + _insert_question(conn, question_id="open1", thread_id="t1", turn=0, status="open") + _insert_question(conn, question_id="exp1", thread_id="t2", turn=0, status="expired") + graph = _FakeGraph(interrupted_turn=0) + worker = ResumeWorker(graph, conn) + + results = worker.recover_pending_resumes() + + assert results == [] + assert graph.invocations == [] + + +def test_recover_processes_answered_oldest_first(conn: sqlite3.Connection) -> None: + # Two answered rows on distinct threads; recovery must visit older first. + _insert_question( + conn, question_id="q_old", thread_id="t_old", turn=0, status="open" + ) + answer_question( + conn, + question_id="q_old", + answer_json=json.dumps("old"), + answered_via="slack", + answered_at="2026-01-01T00:00:00+00:00", + ) + _insert_question( + conn, question_id="q_new", thread_id="t_new", turn=0, status="open" + ) + answer_question( + conn, + question_id="q_new", + answer_json=json.dumps("new"), + answered_via="slack", + answered_at="2026-06-01T00:00:00+00:00", + ) + graph = _FakeGraph(interrupted_turn=0) + worker = ResumeWorker(graph, conn) + + results = worker.recover_pending_resumes() + + assert [r.thread_id for r in results] == ["t_old", "t_new"] + + +# --------------------------------------------------------------------------- # +# answer decoding +# --------------------------------------------------------------------------- # + + +def test_decode_answer_json_roundtrip() -> None: + assert resume_worker._decode_answer(json.dumps({"k": 1})) == {"k": 1} + + +def test_decode_answer_none() -> None: + assert resume_worker._decode_answer(None) is None + + +def test_decode_answer_non_json_passthrough() -> None: + assert resume_worker._decode_answer("not-json{{") == "not-json{{" + + +# --------------------------------------------------------------------------- # +# command builder +# --------------------------------------------------------------------------- # + + +def test_build_resume_command_wraps_answer() -> None: + pytest.importorskip("langgraph") + cmd = build_resume_command("hello") + assert getattr(cmd, "resume", None) == "hello" + + +# --------------------------------------------------------------------------- # +# module contract +# --------------------------------------------------------------------------- # + + +def test_module_exports_public_contract() -> None: + for name in ( + "GraphLike", + "ResumeOutcome", + "ResumeResult", + "ResumeWorker", + "build_resume_command", + "snapshot_interrupt_turns", + ): + assert name in resume_worker.__all__ + assert hasattr(resume_worker, name) + + +def test_graphlike_is_runtime_checkable() -> None: + assert isinstance(_FakeGraph(interrupted_turn=None), GraphLike) + + +# --------------------------------------------------------------------------- # +# end-to-end against a real compiled LangGraph app +# --------------------------------------------------------------------------- # + + +def test_end_to_end_against_real_langgraph(conn: sqlite3.Connection) -> None: + pytest.importorskip("langgraph") + from langgraph.graph import StateGraph, START, END + from langgraph.checkpoint.memory import MemorySaver + from langgraph.types import interrupt + from typing import TypedDict + + class S(TypedDict, total=False): + turn: int + answer: Any + + def clarify(state: S) -> dict[str, Any]: + ans = interrupt({"turn": state.get("turn", 0), "question_id": "q1"}) + return {"answer": ans} + + g = StateGraph(S) + g.add_node("clarify", clarify) + g.add_edge(START, "clarify") + g.add_edge("clarify", END) + app = g.compile(checkpointer=MemorySaver()) + + cfg = {"configurable": {"thread_id": "real-1"}} + app.invoke({"turn": 5}, cfg) # suspends on interrupt at turn 5 + + _insert_question( + conn, question_id="q1", thread_id="real-1", turn=5, status="answered" + ) + worker = ResumeWorker(app, conn) + + first = worker.resume( + thread_id="real-1", question_id="q1", turn=5, answer="confirmed" + ) + assert first.outcome is ResumeOutcome.RESUMED + assert first.graph_result.get("answer") == "confirmed" + + # A redelivered job after the real graph advanced must not double-apply. + second = worker.resume( + thread_id="real-1", question_id="q1", turn=5, answer="confirmed" + ) + assert second.outcome is ResumeOutcome.SUPERSEDED + assert _status(conn, "q1") == "superseded" diff --git a/agent-team/tests/test_review_loop.py b/agent-team/tests/test_review_loop.py new file mode 100644 index 0000000..8666321 --- /dev/null +++ b/agent-team/tests/test_review_loop.py @@ -0,0 +1,352 @@ +"""Unit tests for agent_team.nodes.review_loop (design §3.3, §7.1 P2).""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from agent_team.nodes import review_loop +from agent_team.nodes.review_loop import ( + DEFAULT_MAX_REVIEW_ROUNDS, + ReviewOutcome, + ReviewResult, + ReviewVerdict, + build_review_prompt, + parse_verdict, + review_node, + route_after_review, + set_review_invoker, +) +from agent_team.task_model import Phase, PipelineState, TaskStatus + + +@pytest.fixture(autouse=True) +def _restore_invoker(): + """Restore the module review invoker after each test.""" + original = review_loop._review_invoker + yield + review_loop._review_invoker = original + + +@pytest.fixture(autouse=True) +def _clear_env(monkeypatch: pytest.MonkeyPatch): + """Clear review-loop env knobs so tests are hermetic.""" + monkeypatch.delenv("AGENT_TEAM_MAX_REVIEW_ROUNDS", raising=False) + monkeypatch.delenv("AGENT_TEAM_ORCHESTRATOR_RUN_PY", raising=False) + + +def _state(**overrides: Any) -> PipelineState: + """Build a minimal PipelineState with a plan present.""" + base: PipelineState = { + "thread_id": "t1", + "plan": {"phases": ["P1", "P2"]}, + "review_verdicts": [], + } + base.update(overrides) # type: ignore[typeddict-item] + return base + + +def _invoker_returning(text: str): + """Return a review invoker that always yields ``text`` and records calls.""" + calls: list[dict[str, Any]] = [] + + def invoker(prompt: str, **kw: Any) -> str: + calls.append({"prompt": prompt, "kw": kw}) + return text + + invoker.calls = calls # type: ignore[attr-defined] + return invoker + + +# --------------------------------------------------------------------------- # +# parse_verdict +# --------------------------------------------------------------------------- # + + +def test_parse_verdict_approve() -> None: + assert parse_verdict("VERDICT: APPROVE\nlooks good") is ReviewVerdict.APPROVE + assert parse_verdict("LGTM, ship it") is ReviewVerdict.APPROVE + + +def test_parse_verdict_request_changes() -> None: + assert ( + parse_verdict("VERDICT: REQUEST CHANGES\nmissing rollback") + is ReviewVerdict.REQUEST_CHANGES + ) + assert parse_verdict("BLOCK: unsafe IAM policy") is ReviewVerdict.REQUEST_CHANGES + + +def test_parse_verdict_is_case_insensitive() -> None: + assert parse_verdict("verdict: approve") is ReviewVerdict.APPROVE + + +def test_parse_verdict_fails_closed_on_ambiguous() -> None: + # Neither token present -> REQUEST_CHANGES (fail closed). + assert parse_verdict("hmm, not sure") is ReviewVerdict.REQUEST_CHANGES + assert parse_verdict("") is ReviewVerdict.REQUEST_CHANGES + assert parse_verdict(None) is ReviewVerdict.REQUEST_CHANGES # type: ignore[arg-type] + + +def test_parse_verdict_request_changes_wins_on_conflict() -> None: + # Both tokens present -> REQUEST_CHANGES wins (fail closed). + text = "Some phases APPROVE-able but VERDICT: REQUEST CHANGES overall" + assert parse_verdict(text) is ReviewVerdict.REQUEST_CHANGES + + +# --------------------------------------------------------------------------- # +# build_review_prompt +# --------------------------------------------------------------------------- # + + +def test_build_review_prompt_embeds_plan() -> None: + prompt = build_review_prompt(_state(plan={"phases": ["alpha"]})) + assert "alpha" in prompt + assert "VERDICT: APPROVE" in prompt + assert "VERDICT: REQUEST CHANGES" in prompt + + +def test_build_review_prompt_includes_prior_findings() -> None: + state = _state( + review_verdicts=[{"verdict": "request_changes", "findings": "rollback missing"}] + ) + prompt = build_review_prompt(state) + assert "rollback missing" in prompt + assert "REVISED" in prompt + + +# --------------------------------------------------------------------------- # +# review_node — APPROVE path +# --------------------------------------------------------------------------- # + + +def test_review_node_approve_advances_to_build() -> None: + set_review_invoker(_invoker_returning("VERDICT: APPROVE\nsolid plan")) + update = review_node(_state()) + + assert update["current_phase"] == Phase.BUILD.value + assert update["status"] == TaskStatus.ACTIVE.value + assert update["updated_at"] + + verdicts = update["review_verdicts"] + assert len(verdicts) == 1 + assert verdicts[0]["verdict"] == ReviewVerdict.APPROVE.value + assert verdicts[0]["outcome"] == ReviewOutcome.APPROVED.value + assert verdicts[0]["round_index"] == 1 + assert verdicts[0]["reviewer"] == "cross_reviewer" + + +# --------------------------------------------------------------------------- # +# review_node — REQUEST_CHANGES loop-back vs escalate +# --------------------------------------------------------------------------- # + + +def test_review_node_request_changes_loops_back_under_cap() -> None: + set_review_invoker(_invoker_returning("VERDICT: REQUEST CHANGES\nfix it")) + update = review_node(_state(), config={"max_review_rounds": 3}) + + assert update["current_phase"] == Phase.PLAN.value + assert update["status"] == TaskStatus.ACTIVE.value + assert update["review_verdicts"][-1]["outcome"] == ReviewOutcome.LOOP_BACK.value + + +def test_review_node_escalates_at_round_cap() -> None: + set_review_invoker(_invoker_returning("VERDICT: REQUEST CHANGES\nstill broken")) + # Two prior REQUEST_CHANGES rounds already recorded; cap is 3 -> this is + # round 3 -> escalate. + state = _state( + review_verdicts=[ + {"verdict": "request_changes", "outcome": "loop_back"}, + {"verdict": "request_changes", "outcome": "loop_back"}, + ] + ) + update = review_node(state, config={"max_review_rounds": 3}) + + assert update["current_phase"] == Phase.PARKED.value + assert update["status"] == TaskStatus.PARKED.value + last = update["review_verdicts"][-1] + assert last["outcome"] == ReviewOutcome.ESCALATE.value + assert last["round_index"] == 3 + + +def test_review_node_approve_at_cap_still_advances() -> None: + # Even at the round cap, an APPROVE advances to build (cap only bounds + # REQUEST_CHANGES looping). + set_review_invoker(_invoker_returning("VERDICT: APPROVE")) + state = _state( + review_verdicts=[ + {"verdict": "request_changes", "outcome": "loop_back"}, + {"verdict": "request_changes", "outcome": "loop_back"}, + ] + ) + update = review_node(state, config={"max_review_rounds": 3}) + assert update["current_phase"] == Phase.BUILD.value + assert update["review_verdicts"][-1]["outcome"] == ReviewOutcome.APPROVED.value + + +def test_review_node_appends_to_prior_verdicts() -> None: + set_review_invoker(_invoker_returning("VERDICT: REQUEST CHANGES")) + state = _state( + review_verdicts=[{"verdict": "request_changes", "outcome": "loop_back"}] + ) + update = review_node(state, config={"max_review_rounds": 5}) + assert len(update["review_verdicts"]) == 2 + assert update["review_verdicts"][-1]["round_index"] == 2 + + +# --------------------------------------------------------------------------- # +# review_node — invoker wiring & errors +# --------------------------------------------------------------------------- # + + +def test_review_node_passes_prompt_and_config_to_invoker() -> None: + invoker = _invoker_returning("VERDICT: APPROVE") + set_review_invoker(invoker) + cfg = {"max_review_rounds": 2, "orchestrator_run_py": "/tmp/run.py"} + review_node(_state(plan={"phases": ["zeta"]}), config=cfg) + + assert len(invoker.calls) == 1 # type: ignore[attr-defined] + call = invoker.calls[0] # type: ignore[attr-defined] + assert "zeta" in call["prompt"] + assert call["kw"]["run_py"] == "/tmp/run.py" + assert call["kw"]["config"] is cfg + + +def test_review_node_requires_plan() -> None: + set_review_invoker(_invoker_returning("VERDICT: APPROVE")) + with pytest.raises(ValueError): + review_node({"thread_id": "t1", "review_verdicts": []}) + + +def test_review_node_rejects_bad_round_cap() -> None: + set_review_invoker(_invoker_returning("VERDICT: APPROVE")) + with pytest.raises(ValueError): + review_node(_state(), config={"max_review_rounds": 0}) + with pytest.raises(ValueError): + review_node(_state(), config={"max_review_rounds": "lots"}) + + +def test_review_node_uses_env_round_cap(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AGENT_TEAM_MAX_REVIEW_ROUNDS", "1") + set_review_invoker(_invoker_returning("VERDICT: REQUEST CHANGES")) + # Cap 1 from env -> first REQUEST_CHANGES round escalates immediately. + update = review_node(_state()) + assert update["status"] == TaskStatus.PARKED.value + + +def test_review_node_default_cap_is_three() -> None: + assert DEFAULT_MAX_REVIEW_ROUNDS == 3 + set_review_invoker(_invoker_returning("VERDICT: REQUEST CHANGES")) + # Rounds 1 and 2 loop back under the default cap of 3. + state = _state( + review_verdicts=[{"verdict": "request_changes", "outcome": "loop_back"}] + ) + update = review_node(state) # round 2 + assert update["current_phase"] == Phase.PLAN.value + + +def test_review_node_coerces_non_string_invoker_output() -> None: + set_review_invoker(lambda prompt, **kw: 12345) # type: ignore[return-value] + # No verdict token in "12345" -> fails closed to REQUEST_CHANGES. + update = review_node(_state(), config={"max_review_rounds": 3}) + assert ( + update["review_verdicts"][-1]["verdict"] == ReviewVerdict.REQUEST_CHANGES.value + ) + + +# --------------------------------------------------------------------------- # +# route_after_review +# --------------------------------------------------------------------------- # + + +def test_route_after_review_approved_to_build() -> None: + state: PipelineState = {"review_verdicts": [{"outcome": "approved"}]} + assert route_after_review(state) == "build" + + +def test_route_after_review_loop_back_to_plan() -> None: + state: PipelineState = {"review_verdicts": [{"outcome": "loop_back"}]} + assert route_after_review(state) == "plan" + + +def test_route_after_review_escalate_to_parked() -> None: + state: PipelineState = {"review_verdicts": [{"outcome": "escalate"}]} + assert route_after_review(state) == "parked" + + +def test_route_after_review_no_verdict_parks_fail_closed() -> None: + assert route_after_review({"review_verdicts": []}) == "parked" + assert route_after_review({}) == "parked" + + +def test_route_after_review_unexpected_outcome_parks() -> None: + state: PipelineState = {"review_verdicts": [{"outcome": "weird"}]} + assert route_after_review(state) == "parked" + + +def test_route_uses_most_recent_verdict() -> None: + state: PipelineState = { + "review_verdicts": [{"outcome": "loop_back"}, {"outcome": "approved"}] + } + assert route_after_review(state) == "build" + + +# --------------------------------------------------------------------------- # +# node + router integration (the full loop decision) +# --------------------------------------------------------------------------- # + + +def test_node_then_route_round_trip_approve() -> None: + set_review_invoker(_invoker_returning("VERDICT: APPROVE")) + update = review_node(_state()) + # Merge update back into state (LangGraph reducer would do this). + state = {**_state(), **update} + assert route_after_review(state) == "build" + + +def test_node_then_route_round_trip_loop_back() -> None: + set_review_invoker(_invoker_returning("VERDICT: REQUEST CHANGES")) + update = review_node(_state(), config={"max_review_rounds": 3}) + state = {**_state(), **update} + assert route_after_review(state) == "plan" + + +def test_node_then_route_round_trip_escalate() -> None: + set_review_invoker(_invoker_returning("VERDICT: REQUEST CHANGES")) + update = review_node(_state(), config={"max_review_rounds": 1}) + state = {**_state(), **update} + assert route_after_review(state) == "parked" + + +# --------------------------------------------------------------------------- # +# ReviewResult serialization +# --------------------------------------------------------------------------- # + + +def test_review_result_to_dict_round_trips_fields() -> None: + result = ReviewResult( + verdict=ReviewVerdict.APPROVE, + round_index=2, + outcome=ReviewOutcome.APPROVED, + findings="all good", + created_at="2026-06-17T00:00:00+00:00", + ) + d = result.to_dict() + assert d == { + "verdict": "approve", + "round_index": 2, + "outcome": "approved", + "findings": "all good", + "reviewer": "cross_reviewer", + "created_at": "2026-06-17T00:00:00+00:00", + } + + +# --------------------------------------------------------------------------- # +# default invoker (orchestrator shell-out) — error surface only, no real call +# --------------------------------------------------------------------------- # + + +def test_default_invoker_missing_run_py_raises() -> None: + with pytest.raises(FileNotFoundError): + review_loop._orchestrator_invoker("prompt", run_py="/nonexistent/path/run.py") diff --git a/agent-team/tests/test_run_team.py b/agent-team/tests/test_run_team.py new file mode 100644 index 0000000..1530a81 --- /dev/null +++ b/agent-team/tests/test_run_team.py @@ -0,0 +1,627 @@ +"""Unit tests for the ``run-team.py`` operator CLI (design §3.3.1, §7.1 P1). + +``run-team.py`` is a hyphenated entry script (per the design's "entry CLI +``run-team.py``"), so it cannot be imported by normal ``import`` syntax. These +tests load it via :mod:`importlib` from its file path and exercise the manual +ledger path against the FOUNDATION ``agent_team.db.schema`` ledger. + +The tests assert the §3.3.1 manual-path contract: list open/parked questions, +answer-on-behalf / force-expire / supersede gated behind ``--confirm`` and +audit-logged, first-answer-wins semantics inherited from the foundation +compare-and-set, and read-only commands needing no confirmation. +""" + +from __future__ import annotations + +import importlib.util +import io +import json +from pathlib import Path +from types import ModuleType + +import pytest + +from agent_team.db.schema import QUESTION_STATES, connect, init_db + +# Path to the hyphenated entry CLI (sibling of the agent_team package). +_CLI_PATH = Path(__file__).resolve().parents[1] / "run-team.py" + + +def _load_cli() -> ModuleType: + """Import ``run-team.py`` from its file path as a module.""" + spec = importlib.util.spec_from_file_location("run_team_cli", _CLI_PATH) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.fixture(scope="module") +def cli() -> ModuleType: + """The loaded run-team CLI module (loaded once per test module).""" + return _load_cli() + + +@pytest.fixture() +def db_path(tmp_path: Path) -> Path: + """A fresh, initialized ledger DB for each test.""" + path = tmp_path / "state" / "agent_team.sqlite" + init_db(path) + return path + + +@pytest.fixture() +def audit_log(tmp_path: Path) -> Path: + """Path to a per-test audit log (not created until first destructive op).""" + return tmp_path / "state" / "audit.log.jsonl" + + +def _insert_question( + db_path: Path, + *, + question_id: str, + thread_id: str = "thread-a", + turn: int = 0, + status: str = "open", + transport: str = "slack", + deadline_at: str | None = None, +) -> None: + """Insert a pending_questions row directly for test setup.""" + conn = connect(db_path) + try: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport, posted_at, " + "deadline_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + ( + question_id, + thread_id, + turn, + status, + transport, + "2026-06-17T00:00:00+00:00", + deadline_at, + ), + ) + finally: + conn.close() + + +def _status_of(db_path: Path, question_id: str) -> str | None: + conn = connect(db_path) + try: + row = conn.execute( + "SELECT status FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + finally: + conn.close() + return None if row is None else row["status"] + + +def _run( + cli: ModuleType, + db_path: Path, + audit_log: Path, + *args: str, +) -> tuple[int, str]: + """Invoke ``main`` with the standard global flags, capturing stdout.""" + out = io.StringIO() + argv = ["--db", str(db_path), "--audit-log", str(audit_log), *args] + code = cli.main(argv, out=out) + return code, out.getvalue() + + +# --------------------------------------------------------------------------- # +# Foundation-import / structural assertions +# --------------------------------------------------------------------------- # + + +def test_cli_file_exists_and_is_hyphenated() -> None: + assert _CLI_PATH.name == "run-team.py" + assert _CLI_PATH.is_file() + + +def test_cli_imports_foundation_contracts_verbatim(cli: ModuleType) -> None: + # The CLI must import the foundation, not redefine it. + from agent_team.db import schema as foundation_schema + + assert cli.answer_question is foundation_schema.answer_question + assert cli.expire_question is foundation_schema.expire_question + assert cli.supersede_question is foundation_schema.supersede_question + assert cli.connect is foundation_schema.connect + assert cli.init_db is foundation_schema.init_db + + +def test_build_parser_has_no_side_effects(cli: ModuleType) -> None: + parser = cli.build_parser() + assert parser.prog == "run-team.py" + + +# --------------------------------------------------------------------------- # +# init-db +# --------------------------------------------------------------------------- # + + +def test_init_db_creates_ledger_tables(cli: ModuleType, tmp_path: Path) -> None: + db_path = tmp_path / "state" / "fresh.sqlite" + audit_log = tmp_path / "audit.jsonl" + code, out = _run(cli, db_path, audit_log, "init-db") + assert code == 0 + assert db_path.exists() + conn = connect(db_path) + try: + names = { + r["name"] + for r in conn.execute( + "SELECT name FROM sqlite_master WHERE type='table'" + ).fetchall() + } + finally: + conn.close() + assert "pending_questions" in names + assert "budget_ledger" in names + + +def test_init_db_is_idempotent(cli: ModuleType, tmp_path: Path) -> None: + db_path = tmp_path / "state" / "fresh.sqlite" + audit_log = tmp_path / "audit.jsonl" + assert _run(cli, db_path, audit_log, "init-db")[0] == 0 + assert _run(cli, db_path, audit_log, "init-db")[0] == 0 + + +# --------------------------------------------------------------------------- # +# list / show (read-only, no confirmation) +# --------------------------------------------------------------------------- # + + +def test_list_open_default(cli: ModuleType, db_path: Path, audit_log: Path) -> None: + _insert_question(db_path, question_id="q-open", status="open") + _insert_question(db_path, question_id="q-exp", status="expired") + code, out = _run(cli, db_path, audit_log, "list") + assert code == 0 + payload = json.loads(out) + ids = {row["question_id"] for row in payload} + assert ids == {"q-open"} + + +def test_list_all(cli: ModuleType, db_path: Path, audit_log: Path) -> None: + _insert_question(db_path, question_id="q-open", status="open") + _insert_question(db_path, question_id="q-exp", status="expired") + code, out = _run(cli, db_path, audit_log, "list", "--all") + assert code == 0 + ids = {row["question_id"] for row in json.loads(out)} + assert ids == {"q-open", "q-exp"} + + +def test_list_parked_excludes_open( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q-open", status="open") + _insert_question(db_path, question_id="q-ans", status="answered") + _insert_question(db_path, question_id="q-exp", status="expired") + _insert_question(db_path, question_id="q-sup", status="superseded") + code, out = _run(cli, db_path, audit_log, "list", "--parked") + assert code == 0 + ids = {row["question_id"] for row in json.loads(out)} + assert ids == {"q-ans", "q-exp", "q-sup"} + assert "q-open" not in ids + + +def test_list_status_filter(cli: ModuleType, db_path: Path, audit_log: Path) -> None: + _insert_question(db_path, question_id="q-open", status="open") + _insert_question(db_path, question_id="q-exp", status="expired") + code, out = _run(cli, db_path, audit_log, "list", "--status", "expired") + assert code == 0 + ids = {row["question_id"] for row in json.loads(out)} + assert ids == {"q-exp"} + + +def test_list_empty_returns_empty_array( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + code, out = _run(cli, db_path, audit_log, "list") + assert code == 0 + assert json.loads(out) == [] + + +def test_show_existing(cli: ModuleType, db_path: Path, audit_log: Path) -> None: + _insert_question(db_path, question_id="q1", thread_id="t1", turn=3) + code, out = _run(cli, db_path, audit_log, "show", "q1") + assert code == 0 + row = json.loads(out) + assert row["question_id"] == "q1" + assert row["thread_id"] == "t1" + assert row["turn"] == 3 + + +def test_show_missing_returns_1( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + code, _ = _run(cli, db_path, audit_log, "show", "nope") + assert code == 1 + + +# --------------------------------------------------------------------------- # +# Destructive actions require --confirm and are audit-logged +# --------------------------------------------------------------------------- # + + +def test_expire_without_confirm_refuses_and_does_not_mutate( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run(cli, db_path, audit_log, "expire", "q1") + assert code == 1 + # Unchanged: the guard fired before touching the ledger. + assert _status_of(db_path, "q1") == "open" + assert not audit_log.exists() + + +def test_answer_without_confirm_refuses( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run(cli, db_path, audit_log, "answer", "q1", "--answer", "yes") + assert code == 1 + assert _status_of(db_path, "q1") == "open" + + +def test_supersede_without_confirm_refuses( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run(cli, db_path, audit_log, "supersede", "q1") + assert code == 1 + assert _status_of(db_path, "q1") == "open" + + +def test_expire_with_confirm_flips_status_and_audits( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, out = _run( + cli, db_path, audit_log, "--operator", "adam", "expire", "q1", "--confirm" + ) + assert code == 0 + assert _status_of(db_path, "q1") == "expired" + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + # Attempt is recorded BEFORE the mutation, outcome after, so a mutation can + # never land without a trail (§3.3.1). + assert len(entries) == 2 + assert entries[0]["phase"] == "attempt" + assert "applied" not in entries[0] + assert entries[-1]["phase"] == "outcome" + assert entries[-1]["action"] == "expire" + assert entries[-1]["question_id"] == "q1" + assert entries[-1]["operator"] == "adam" + assert entries[-1]["applied"] is True + assert "ts" in entries[-1] + + +def test_answer_with_confirm_flips_status_records_via( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, out = _run( + cli, + db_path, + audit_log, + "--operator", + "adam", + "answer", + "q1", + "--answer", + '{"choice": "B"}', + "--confirm", + ) + assert code == 0 + assert _status_of(db_path, "q1") == "answered" + conn = connect(db_path) + try: + row = conn.execute( + "SELECT answer_json, answered_via FROM pending_questions " + "WHERE question_id = ?", + ("q1",), + ).fetchone() + finally: + conn.close() + assert row["answer_json"] == '{"choice": "B"}' + assert row["answered_via"] == "cli:adam" + + +def test_answer_explicit_via_overrides_default( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run( + cli, + db_path, + audit_log, + "answer", + "q1", + "--answer", + "x", + "--via", + "slack:U123", + "--confirm", + ) + assert code == 0 + conn = connect(db_path) + try: + row = conn.execute( + "SELECT answered_via FROM pending_questions WHERE question_id = ?", + ("q1",), + ).fetchone() + finally: + conn.close() + assert row["answered_via"] == "slack:U123" + + +def test_supersede_with_confirm_flips_status( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run(cli, db_path, audit_log, "supersede", "q1", "--confirm") + assert code == 0 + assert _status_of(db_path, "q1") == "superseded" + + +# --------------------------------------------------------------------------- # +# First-answer-wins / no-op semantics inherited from the foundation +# --------------------------------------------------------------------------- # + + +def test_answer_already_expired_is_noop_returns_1_but_audits( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="expired") + code, _ = _run( + cli, db_path, audit_log, "answer", "q1", "--answer", "x", "--confirm" + ) + assert code == 1 + # Status unchanged (compare-and-set lost), but the attempt is audited. + assert _status_of(db_path, "q1") == "expired" + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + assert entries[-1]["action"] == "answer" + assert entries[-1]["applied"] is False + + +def test_expire_missing_question_is_noop_returns_1( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + code, _ = _run(cli, db_path, audit_log, "expire", "ghost", "--confirm") + assert code == 1 + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + assert entries[-1]["applied"] is False + + +def test_double_answer_second_is_noop( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + first, _ = _run( + cli, db_path, audit_log, "answer", "q1", "--answer", "a", "--confirm" + ) + second, _ = _run( + cli, db_path, audit_log, "answer", "q1", "--answer", "b", "--confirm" + ) + assert first == 0 + assert second == 1 # first-answer-wins; second is a no-op + conn = connect(db_path) + try: + row = conn.execute( + "SELECT answer_json FROM pending_questions WHERE question_id = ?", + ("q1",), + ).fetchone() + finally: + conn.close() + assert row["answer_json"] == "a" # original answer preserved + + +# --------------------------------------------------------------------------- # +# Audit log durability (append-only, multiple actions) +# --------------------------------------------------------------------------- # + + +def test_audit_log_appends_across_actions( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + _insert_question(db_path, question_id="q2", status="open") + _run(cli, db_path, audit_log, "expire", "q1", "--confirm") + _run(cli, db_path, audit_log, "answer", "q2", "--answer", "y", "--confirm") + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + # Each destructive action writes an attempt + an outcome record (append-only). + assert len(entries) == 4 + actions = [e["action"] for e in entries] + assert actions == ["expire", "expire", "answer", "answer"] + outcomes = [e["action"] for e in entries if e["phase"] == "outcome"] + assert outcomes == ["expire", "answer"] + + +def test_audit_entries_are_valid_json_lines( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + _run(cli, db_path, audit_log, "expire", "q1", "--confirm") + content = audit_log.read_text() + assert content.endswith("\n") + for line in content.splitlines(): + json.loads(line) # raises if any line is not valid JSON + + +# --------------------------------------------------------------------------- # +# argparse-level usage errors +# --------------------------------------------------------------------------- # + + +def test_no_subcommand_is_usage_error(cli: ModuleType) -> None: + with pytest.raises(SystemExit) as exc: + cli.main([]) + assert exc.value.code == 2 + + +def test_unknown_status_choice_is_usage_error(cli: ModuleType) -> None: + with pytest.raises(SystemExit) as exc: + cli.main(["list", "--status", "bogus"]) + assert exc.value.code == 2 + + +def test_answer_requires_answer_flag(cli: ModuleType) -> None: + with pytest.raises(SystemExit) as exc: + cli.main(["answer", "q1", "--confirm"]) + assert exc.value.code == 2 + + +def test_parked_states_derived_from_foundation(cli: ModuleType) -> None: + # The parked-context states are exactly the non-open foundation states. + assert set(cli._PARKED_STATES) == set(QUESTION_STATES) - {"open"} + + +# --------------------------------------------------------------------------- # +# re-deliver + force-resume (design-named operator verbs, §3.3.1 / §6.6) +# --------------------------------------------------------------------------- # + + +def _set_channel_ref(db_path: Path, question_id: str, ref: str) -> None: + conn = connect(db_path) + try: + conn.execute( + "UPDATE pending_questions SET channel_ref = ? WHERE question_id = ?", + (ref, question_id), + ) + finally: + conn.close() + + +def _channel_ref_of(db_path: Path, question_id: str) -> str | None: + conn = connect(db_path) + try: + row = conn.execute( + "SELECT channel_ref FROM pending_questions WHERE question_id = ?", + (question_id,), + ).fetchone() + finally: + conn.close() + return None if row is None else row["channel_ref"] + + +def test_redeliver_clears_channel_ref_and_audits( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + _set_channel_ref(db_path, "q1", "slack:123.456") + code, out = _run(cli, db_path, audit_log, "redeliver", "q1") + assert code == 0 + assert _channel_ref_of(db_path, "q1") is None + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + # Non-destructive but audited: attempt + outcome, no --confirm needed. + assert [e["phase"] for e in entries] == ["attempt", "outcome"] + assert entries[-1]["action"] == "redeliver" + assert entries[-1]["applied"] is True + assert entries[-1]["prior_channel_ref"] == "slack:123.456" + + +def test_redeliver_needs_no_confirm( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + # redeliver is not in the destructive set, so it runs without --confirm. + assert "redeliver" not in cli._DESTRUCTIVE_ACTIONS + + +def test_redeliver_non_open_is_noop_returns_1( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="answered") + code, _ = _run(cli, db_path, audit_log, "redeliver", "q1") + assert code == 1 + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + assert entries[-1]["applied"] is False + + +def test_force_resume_reopens_expired_parked_question( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + # The parked case: an expired question is RE-OPENED so it can be answered, + # NOT superseded (superseding would make it permanently un-resumable). + _insert_question(db_path, question_id="q1", status="expired") + code, out = _run( + cli, db_path, audit_log, "--operator", "adam", "force-resume", "q1", "--confirm" + ) + assert code == 0 + assert _status_of(db_path, "q1") == "open" # reopened, not superseded + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + assert [e["phase"] for e in entries] == ["attempt", "outcome"] + assert entries[-1]["action"] == "force-resume" + assert entries[-1]["resume_requested"] is True + assert entries[-1]["applied"] is True + assert entries[-1]["operator"] == "adam" + + +def test_force_resume_does_not_supersede_answered_row( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + # Regression: force-resume must NOT flip an answered-but-unresumed row out of + # the state the recovery sweep resumes from. It stays 'answered'. + _insert_question(db_path, question_id="q1", status="answered") + code, _ = _run(cli, db_path, audit_log, "force-resume", "q1", "--confirm") + assert code == 0 + assert _status_of(db_path, "q1") == "answered" # untouched, still resumable + + +def test_force_resume_on_open_is_noop( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run(cli, db_path, audit_log, "force-resume", "q1", "--confirm") + assert code == 1 # an open (not parked) question has nothing to force + assert _status_of(db_path, "q1") == "open" + + +def test_force_resume_without_confirm_refuses( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run(cli, db_path, audit_log, "force-resume", "q1") + assert code == 1 + assert _status_of(db_path, "q1") == "open" # unmutated + assert not audit_log.exists() # refused before any audit (confirm-check first) + + +def test_force_resume_is_in_destructive_set(cli: ModuleType) -> None: + assert "force-resume" in cli._DESTRUCTIVE_ACTIONS + + +def test_operator_defaults_to_os_login_not_empty( + cli: ModuleType, db_path: Path, audit_log: Path +) -> None: + # AUTHZ regression: --operator defaulted to "" → non-attributable audit. + # Omitting it must record a real (non-empty) operator identity. + _insert_question(db_path, question_id="q1", status="open") + code, _ = _run(cli, db_path, audit_log, "expire", "q1", "--confirm") + assert code == 0 + entries = [json.loads(line) for line in audit_log.read_text().splitlines()] + assert entries[-1]["operator"] # non-empty + assert cli._default_operator() # helper never returns empty + + +# --------------------------------------------------------------------------- # +# Audit-before-mutate: an unwritable audit path aborts BEFORE the ledger mutates +# (regression: previously the row was mutated, then the audit append crashed, +# leaving a mutation with no record and an uncaught traceback) +# --------------------------------------------------------------------------- # + + +def test_unwritable_audit_path_aborts_before_mutation( + cli: ModuleType, db_path: Path, tmp_path: Path +) -> None: + _insert_question(db_path, question_id="q1", status="open") + # Point the audit log at a path whose parent is a FILE, so the atomic write + # of the attempt record fails with OSError before the mutation runs. + blocker = tmp_path / "not-a-dir" + blocker.write_text("x") + bad_audit = blocker / "audit.jsonl" + code, _ = _run(cli, db_path, bad_audit, "expire", "q1", "--confirm") + assert code == 1 # clean failure, not an uncaught traceback + assert _status_of(db_path, "q1") == "open" # NOT mutated — no trail, no change diff --git a/agent-team/tests/test_schema.py b/agent-team/tests/test_schema.py new file mode 100644 index 0000000..618e5fc --- /dev/null +++ b/agent-team/tests/test_schema.py @@ -0,0 +1,352 @@ +"""Unit tests for agent_team.db.schema (§3.3.1, §6.7).""" + +from __future__ import annotations + +import sqlite3 +import threading +from pathlib import Path + +import pytest + +from agent_team.db.schema import ( + BUDGET_LEDGER_DDL, + PENDING_QUESTIONS_DDL, + QUESTION_STATES, + SCHEMA_VERSION, + answer_question, + connect, + expire_question, + init_db, + migrate, + reopen_question, + supersede_question, +) + + +def _insert_open_question(conn: sqlite3.Connection, qid: str, turn: int = 0) -> None: + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport) " + "VALUES (?, 'thread-1', ?, 'open', 'slack')", + (qid, turn), + ) + + +def test_ddl_constants_are_nonempty_strings() -> None: + assert isinstance(PENDING_QUESTIONS_DDL, str) and PENDING_QUESTIONS_DDL + assert isinstance(BUDGET_LEDGER_DDL, str) and BUDGET_LEDGER_DDL + assert "pending_questions" in PENDING_QUESTIONS_DDL + assert "budget_ledger" in BUDGET_LEDGER_DDL + + +def test_schema_version_is_int() -> None: + assert isinstance(SCHEMA_VERSION, int) + + +def test_question_states_match_ddl_check() -> None: + assert QUESTION_STATES == ("open", "answered", "expired", "superseded") + for state in QUESTION_STATES: + assert f"'{state}'" in PENDING_QUESTIONS_DDL + + +def test_connect_sets_pragmas(tmp_path: Path) -> None: + conn = connect(tmp_path / "db.sqlite") + try: + assert conn.execute("PRAGMA journal_mode").fetchone()[0].lower() == "wal" + assert conn.execute("PRAGMA foreign_keys").fetchone()[0] == 1 + assert conn.execute("PRAGMA busy_timeout").fetchone()[0] >= 1 + finally: + conn.close() + + +def test_init_db_creates_tables(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + conn = connect(db) + try: + names = { + r[0] + for r in conn.execute( + "SELECT name FROM sqlite_master WHERE type='table'" + ).fetchall() + } + finally: + conn.close() + assert {"pending_questions", "budget_ledger", "schema_meta"} <= names + + +def test_init_db_is_idempotent(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + init_db(db) # must not raise + conn = connect(db) + try: + version = conn.execute( + "SELECT schema_version FROM schema_meta WHERE id=1" + ).fetchone()[0] + finally: + conn.close() + assert version == SCHEMA_VERSION + + +def test_pending_questions_status_check_constraint(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + conn = connect(db) + try: + with pytest.raises(sqlite3.IntegrityError): + conn.execute( + "INSERT INTO pending_questions " + "(question_id, thread_id, turn, status, transport) " + "VALUES ('q', 't', 0, 'bogus', 'slack')" + ) + finally: + conn.close() + + +def test_migrate_stamps_version(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + conn = connect(db) + try: + migrate(conn) + version = conn.execute( + "SELECT schema_version FROM schema_meta WHERE id=1" + ).fetchone()[0] + # Tables exist after migrate. + conn.execute("SELECT 1 FROM pending_questions LIMIT 1") + conn.execute("SELECT 1 FROM budget_ledger LIMIT 1") + finally: + conn.close() + assert version == SCHEMA_VERSION + + +def test_answer_question_first_wins(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + conn = connect(db) + try: + _insert_open_question(conn, "q1") + first = answer_question( + conn, question_id="q1", answer_json='{"a":1}', answered_via="slack" + ) + second = answer_question( + conn, question_id="q1", answer_json='{"a":2}', answered_via="github" + ) + assert first is True + assert second is False # duplicate/late loses the compare-and-set + row = conn.execute( + "SELECT status, answer_json, answered_via, answered_at " + "FROM pending_questions WHERE question_id='q1'" + ).fetchone() + finally: + conn.close() + assert row["status"] == "answered" + assert row["answer_json"] == '{"a":1}' # first answer retained + assert row["answered_via"] == "slack" + assert row["answered_at"] + + +def test_answer_after_expire_loses(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + conn = connect(db) + try: + _insert_open_question(conn, "q2") + assert expire_question(conn, question_id="q2") is True + assert ( + answer_question( + conn, question_id="q2", answer_json="{}", answered_via="slack" + ) + is False + ) + status = conn.execute( + "SELECT status FROM pending_questions WHERE question_id='q2'" + ).fetchone()["status"] + finally: + conn.close() + assert status == "expired" + + +def test_expire_only_open(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + conn = connect(db) + try: + _insert_open_question(conn, "q3") + answer_question(conn, question_id="q3", answer_json="{}", answered_via="slack") + # already answered -> cannot expire + assert expire_question(conn, question_id="q3") is False + finally: + conn.close() + + +def test_supersede_open_or_answered(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + conn = connect(db) + try: + _insert_open_question(conn, "q4") + answer_question(conn, question_id="q4", answer_json="{}", answered_via="slack") + assert supersede_question(conn, question_id="q4") is True + # already superseded -> no-op + assert supersede_question(conn, question_id="q4") is False + finally: + conn.close() + + +def test_reopen_question_unparks_expired_only(tmp_path: Path) -> None: + db = tmp_path / "db.sqlite" + init_db(db) + conn = connect(db) + try: + _insert_open_question(conn, "exp") + _insert_open_question(conn, "ans") + assert expire_question(conn, question_id="exp") is True + assert answer_question( + conn, question_id="ans", answer_json="{}", answered_via="t" + ) + # Expired -> reopened. + assert reopen_question(conn, question_id="exp") is True + row = conn.execute( + "SELECT status, deadline_at FROM pending_questions WHERE question_id='exp'" + ).fetchone() + assert row["status"] == "open" + assert row["deadline_at"] is None # no deadline until one is set + # Answered row is NOT reopenable (only expired rows are). + assert reopen_question(conn, question_id="ans") is False + assert ( + conn.execute( + "SELECT status FROM pending_questions WHERE question_id='ans'" + ).fetchone()["status"] + == "answered" + ) + finally: + conn.close() + + +def test_concurrent_answers_single_winner(tmp_path: Path) -> None: + """Two threads racing to answer the same open question: exactly one wins.""" + db = tmp_path / "db.sqlite" + init_db(db) + seed = connect(db) + try: + _insert_open_question(seed, "race") + finally: + seed.close() + + results: list[bool] = [] + barrier = threading.Barrier(2) + lock = threading.Lock() + + def worker(via: str) -> None: + conn = connect(db) + try: + barrier.wait() + won = answer_question( + conn, question_id="race", answer_json='{"v":1}', answered_via=via + ) + with lock: + results.append(won) + finally: + conn.close() + + threads = [threading.Thread(target=worker, args=(f"c{i}",)) for i in range(2)] + for t in threads: + t.start() + for t in threads: + t.join() + + assert sorted(results) == [False, True] + + +def test_shared_connection_concurrent_distinct_questions(tmp_path: Path) -> None: + """Regression: many threads share ONE connection, answering DISTINCT questions. + + The responder and resume worker hold a single connection and call the CAS + helpers from different threads concurrently (``connect()`` sets + ``check_same_thread=False``). A single connection cannot hold two explicit + ``BEGIN IMMEDIATE`` transactions at once, so the previous implementation + raised "cannot start a transaction within a transaction" for all but one + thread. The CAS now runs each write on its own private connection, so every + distinct question is answered with no error. + """ + db = tmp_path / "db.sqlite" + init_db(db) + shared = connect(db) + n = 8 + try: + for i in range(n): + _insert_open_question(shared, f"q{i}") + + barrier = threading.Barrier(n) + lock = threading.Lock() + wins: list[bool] = [] + errors: list[BaseException] = [] + + def worker(qid: str) -> None: + try: + barrier.wait() + won = answer_question( + shared, question_id=qid, answer_json='{"v":1}', answered_via="t" + ) + with lock: + wins.append(won) + except BaseException as exc: # noqa: BLE001 - record for assertion + with lock: + errors.append(exc) + + threads = [threading.Thread(target=worker, args=(f"q{i}",)) for i in range(n)] + for t in threads: + t.start() + for t in threads: + t.join() + finally: + shared.close() + + assert errors == [], f"shared-connection CAS raised: {errors!r}" + assert wins == [True] * n + + +def test_shared_connection_concurrent_same_question_single_winner( + tmp_path: Path, +) -> None: + """Regression: shared connection, many threads racing the SAME question. + + Exactly one first-answer-wins, the rest no-op (rowcount 0), and no thread + raises a transaction-nesting or lock error. + """ + db = tmp_path / "db.sqlite" + init_db(db) + shared = connect(db) + n = 8 + try: + _insert_open_question(shared, "race") + + barrier = threading.Barrier(n) + lock = threading.Lock() + wins: list[bool] = [] + errors: list[BaseException] = [] + + def worker(via: str) -> None: + try: + barrier.wait() + won = answer_question( + shared, question_id="race", answer_json='{"v":1}', answered_via=via + ) + with lock: + wins.append(won) + except BaseException as exc: # noqa: BLE001 - record for assertion + with lock: + errors.append(exc) + + threads = [threading.Thread(target=worker, args=(f"c{i}",)) for i in range(n)] + for t in threads: + t.start() + for t in threads: + t.join() + finally: + shared.close() + + assert errors == [], f"shared-connection CAS raised: {errors!r}" + assert sum(wins) == 1 + assert wins.count(False) == n - 1 diff --git a/agent-team/tests/test_slack_adapter.py b/agent-team/tests/test_slack_adapter.py new file mode 100644 index 0000000..dc9e892 --- /dev/null +++ b/agent-team/tests/test_slack_adapter.py @@ -0,0 +1,358 @@ +"""Unit tests for agent_team.transport.slack_adapter (§3.3.1, §7.1 P1). + +The Slack adapter is the first concrete ``Transport``. These tests prove the +two §3.3.1 contracts in isolation (no network): ``post_question`` embeds the +``question_id`` and returns the message ``ts`` as ``channel_ref``, and +``parse_answer`` normalizes the inbound Slack shapes to ``(question_id, answer, +via)`` for the first-answer-wins compare-and-set. +""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from agent_team.transport.base import QuestionSet, Transport +from agent_team.transport.slack_adapter import ( + CALLBACK_ID_PREFIX, + VIA_SLACK, + SlackPostError, + SlackTransport, + build_callback_id, + build_question_blocks, + parse_callback_id, +) + + +# --------------------------------------------------------------------------- # +# Test doubles # +# --------------------------------------------------------------------------- # + + +class _RecordingPoster: + """A poster that records the message and returns a canned Slack response.""" + + def __init__(self, response: dict[str, Any] | None = None) -> None: + self.response = ( + response + if response is not None + else {"ok": True, "ts": "1700000000.000100"} + ) + self.calls: list[dict[str, Any]] = [] + + def __call__(self, message: dict[str, Any]) -> dict[str, Any]: + self.calls.append(message) + return self.response + + +def _question_set(**overrides: Any) -> QuestionSet: + defaults: dict[str, Any] = { + "thread_id": "thread-1", + "question_id": "q-abc", + "turn": 0, + "questions": ["Which branch?", "Bump major?"], + "context": {"repo": "sea-haven/foo", "summary": "dep bump"}, + } + defaults.update(overrides) + return QuestionSet(**defaults) + + +# --------------------------------------------------------------------------- # +# Contract / typing # +# --------------------------------------------------------------------------- # + + +def test_slack_transport_is_a_transport_subclass() -> None: + assert issubclass(SlackTransport, Transport) + + +def test_slack_transport_is_instantiable_and_concrete() -> None: + # Concrete: implements both abstractmethods, so construction must succeed. + t = SlackTransport(channel="C123", poster=_RecordingPoster()) + assert isinstance(t, Transport) + + +# --------------------------------------------------------------------------- # +# callback_id helpers # +# --------------------------------------------------------------------------- # + + +def test_build_callback_id_embeds_question_id() -> None: + cb = build_callback_id("q-abc") + assert cb == f"{CALLBACK_ID_PREFIX}:q-abc" + assert "q-abc" in cb + + +def test_parse_callback_id_roundtrips() -> None: + assert parse_callback_id(build_callback_id("q-xyz")) == "q-xyz" + + +def test_parse_callback_id_accepts_bare_id() -> None: + assert parse_callback_id("q-bare") == "q-bare" + + +def test_parse_callback_id_rejects_empty() -> None: + with pytest.raises(ValueError): + parse_callback_id("") + + +def test_parse_callback_id_rejects_prefix_only() -> None: + with pytest.raises(ValueError): + parse_callback_id(f"{CALLBACK_ID_PREFIX}:") + + +# --------------------------------------------------------------------------- # +# build_question_blocks # +# --------------------------------------------------------------------------- # + + +def test_build_question_blocks_lists_every_question() -> None: + qs = _question_set(questions=["A?", "B?", "C?"]) + blocks = build_question_blocks(qs, deadline="2026-06-18T00:00:00Z") + section_texts = [b["text"]["text"] for b in blocks if b["type"] == "section"] + assert len(section_texts) == 3 + assert any("A?" in t for t in section_texts) + assert any("C?" in t for t in section_texts) + + +def test_build_question_blocks_surfaces_context_and_deadline() -> None: + qs = _question_set() + blocks = build_question_blocks(qs, deadline="2026-06-18T00:00:00Z") + flat = repr(blocks) + assert "sea-haven/foo" in flat + assert "dep bump" in flat + assert "2026-06-18T00:00:00Z" in flat + + +def test_build_question_blocks_omits_empty_context() -> None: + qs = _question_set(context={}) + blocks = build_question_blocks(qs, deadline="2026-06-18T00:00:00Z") + # No repo/summary context block beyond the trailing deadline context block. + context_blocks = [b for b in blocks if b["type"] == "context"] + assert len(context_blocks) == 1 # only the deadline footer + + +# --------------------------------------------------------------------------- # +# post_question # +# --------------------------------------------------------------------------- # + + +def test_post_question_returns_message_ts_as_channel_ref() -> None: + poster = _RecordingPoster({"ok": True, "ts": "1700000000.000200"}) + t = SlackTransport(channel="C999", poster=poster) + ref = t.post_question( + thread_id="thread-1", + question_id="q-abc", + turn=2, + question_set=_question_set(turn=2), + deadline="2026-06-18T00:00:00Z", + ) + assert ref == "1700000000.000200" + + +def test_post_question_embeds_question_id_in_callback_id() -> None: + poster = _RecordingPoster() + t = SlackTransport(channel="C999", poster=poster) + t.post_question( + thread_id="thread-1", + question_id="q-abc", + turn=0, + question_set=_question_set(), + deadline="2026-06-18T00:00:00Z", + ) + sent = poster.calls[0] + assert sent["callback_id"] == build_callback_id("q-abc") + assert sent["channel"] == "C999" + # The metadata payload also carries the identity for restart recovery. + assert sent["metadata"]["event_payload"]["question_id"] == "q-abc" + assert sent["metadata"]["event_payload"]["thread_id"] == "thread-1" + + +def test_post_question_accepts_nested_message_ts() -> None: + poster = _RecordingPoster({"ok": True, "message": {"ts": "1700000000.000300"}}) + t = SlackTransport(channel="C1", poster=poster) + ref = t.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=_question_set(), + deadline="d", + ) + assert ref == "1700000000.000300" + + +def test_post_question_raises_when_response_missing_ts() -> None: + poster = _RecordingPoster({"ok": True}) # no ts -> cannot record channel_ref + t = SlackTransport(channel="C1", poster=poster) + with pytest.raises(SlackPostError): + t.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=_question_set(), + deadline="d", + ) + + +def test_post_question_default_poster_refuses_network() -> None: + # Foundation ships nothing live: a poster-less transport must not post. + t = SlackTransport(channel="C1") + with pytest.raises(SlackPostError): + t.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=_question_set(), + deadline="d", + ) + + +def test_post_question_normalizes_poster_exception_to_slack_post_error() -> None: + def _boom(_message: dict[str, Any]) -> dict[str, Any]: + raise RuntimeError("connection reset") + + t = SlackTransport(channel="C1", poster=_boom) + with pytest.raises(SlackPostError) as exc: + t.post_question( + thread_id="t", + question_id="q", + turn=0, + question_set=_question_set(), + deadline="d", + ) + assert "connection reset" in str(exc.value) + + +# --------------------------------------------------------------------------- # +# parse_answer # +# --------------------------------------------------------------------------- # + + +def test_parse_answer_button_click() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + payload = { + "type": "block_actions", + "callback_id": build_callback_id("q-abc"), + "actions": [{"action_id": "approve", "value": "yes"}], + } + assert t.parse_answer(payload) == ("q-abc", "yes", VIA_SLACK) + + +def test_parse_answer_select_menu() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + payload = { + "callback_id": build_callback_id("q-sel"), + "actions": [{"action_id": "branch", "selected_option": {"value": "main"}}], + } + assert t.parse_answer(payload) == ("q-sel", "main", VIA_SLACK) + + +def test_parse_answer_multi_select_returns_list() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + payload = { + "callback_id": build_callback_id("q-multi"), + "actions": [ + { + "action_id": "labels", + "selected_options": [{"value": "bug"}, {"value": "ci"}], + } + ], + } + qid, answer, via = t.parse_answer(payload) + assert qid == "q-multi" + assert answer == ["bug", "ci"] + assert via == VIA_SLACK + + +def test_parse_answer_multiple_actions_returns_list() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + payload = { + "callback_id": build_callback_id("q-two"), + "actions": [ + {"action_id": "a", "value": "1"}, + {"action_id": "b", "value": "2"}, + ], + } + _, answer, _ = t.parse_answer(payload) + assert answer == ["1", "2"] + + +def test_parse_answer_plain_text_reply() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + payload = { + "callback_id": build_callback_id("q-text"), + "text": "use the release branch", + } + assert t.parse_answer(payload) == ( + "q-text", + "use the release branch", + VIA_SLACK, + ) + + +def test_parse_answer_uses_metadata_event_payload() -> None: + # No top-level callback_id; the question_id rides in message metadata. + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + payload = { + "message": { + "metadata": { + "event_type": "agent_team_question", + "event_payload": {"question_id": "q-meta", "thread_id": "t"}, + } + }, + "answer": "ok", + } + assert t.parse_answer(payload) == ("q-meta", "ok", VIA_SLACK) + + +def test_parse_answer_bare_question_id_field() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + payload = {"question_id": "q-bare", "answer": 42} + assert t.parse_answer(payload) == ("q-bare", 42, VIA_SLACK) + + +def test_parse_answer_rejects_non_mapping() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + with pytest.raises(ValueError): + t.parse_answer("not a payload") + + +def test_parse_answer_rejects_missing_question_id() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + with pytest.raises(ValueError): + t.parse_answer({"text": "an answer with no id"}) + + +def test_parse_answer_rejects_missing_answer() -> None: + t = SlackTransport(channel="C1", poster=_RecordingPoster()) + with pytest.raises(ValueError): + t.parse_answer({"callback_id": build_callback_id("q-noans")}) + + +# --------------------------------------------------------------------------- # +# round-trip: post then parse # +# --------------------------------------------------------------------------- # + + +def test_post_then_parse_round_trips_question_id() -> None: + poster = _RecordingPoster() + t = SlackTransport(channel="C1", poster=poster) + t.post_question( + thread_id="thread-9", + question_id="q-round", + turn=1, + question_set=_question_set(question_id="q-round", turn=1), + deadline="2026-06-18T00:00:00Z", + ) + sent_callback_id = poster.calls[0]["callback_id"] + + # Simulate Slack echoing the message-level callback_id on an interaction. + inbound = { + "callback_id": sent_callback_id, + "actions": [{"action_id": "approve", "value": "approved"}], + } + qid, answer, via = t.parse_answer(inbound) + assert qid == "q-round" + assert answer == "approved" + assert via == VIA_SLACK diff --git a/agent-team/tests/test_state_store.py b/agent-team/tests/test_state_store.py new file mode 100644 index 0000000..4a6d2ae --- /dev/null +++ b/agent-team/tests/test_state_store.py @@ -0,0 +1,111 @@ +"""Unit tests for agent_team.state_store (§6.7).""" + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +from agent_team import state_store +from agent_team.state_store import ( + IntegrityError, + atomic_write, + compute_content_hash, + read_checked, + write_checked, +) + + +def test_compute_content_hash_is_sha256_hex() -> None: + import hashlib + + data = b"hello world" + assert compute_content_hash(data) == hashlib.sha256(data).hexdigest() + + +def test_compute_content_hash_distinguishes_inputs() -> None: + assert compute_content_hash(b"a") != compute_content_hash(b"b") + + +def test_atomic_write_creates_file_with_exact_bytes(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + payload = b"\x00\x01binary\xff" + atomic_write(target, payload) + assert target.read_bytes() == payload + + +def test_atomic_write_creates_missing_parent_dirs(tmp_path: Path) -> None: + target = tmp_path / "nested" / "deep" / "state.bin" + atomic_write(target, b"x") + assert target.read_bytes() == b"x" + + +def test_atomic_write_overwrites_existing(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + atomic_write(target, b"old") + atomic_write(target, b"new-and-longer") + assert target.read_bytes() == b"new-and-longer" + + +def test_atomic_write_leaves_no_temp_files(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + atomic_write(target, b"data") + leftovers = [p for p in tmp_path.iterdir() if p.name != "state.bin"] + assert leftovers == [] + + +def test_write_then_read_checked_roundtrip(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + payload = json.dumps({"k": "v"}).encode() + write_checked(target, payload, schema_version=3) + assert read_checked(target, schema_version=3) == payload + + +def test_read_checked_schema_version_mismatch_raises(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + write_checked(target, b"data", schema_version=1) + with pytest.raises(IntegrityError): + read_checked(target, schema_version=2) + + +def test_read_checked_content_corruption_raises(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + write_checked(target, b"original", schema_version=1) + # Corrupt the payload without touching the sidecar -> hash mismatch. + target.write_bytes(b"tampered") + with pytest.raises(IntegrityError): + read_checked(target, schema_version=1) + + +def test_read_checked_missing_payload_raises(tmp_path: Path) -> None: + with pytest.raises(IntegrityError): + read_checked(tmp_path / "nope.bin", schema_version=1) + + +def test_read_checked_missing_sidecar_raises(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + # Plain atomic_write writes payload but NOT the integrity sidecar. + atomic_write(target, b"data") + with pytest.raises(IntegrityError): + read_checked(target, schema_version=1) + + +def test_read_checked_garbled_sidecar_raises(tmp_path: Path) -> None: + target = tmp_path / "state.bin" + write_checked(target, b"data", schema_version=1) + meta_path = target.with_name(target.name + ".meta.json") + meta_path.write_bytes(b"not-json{{{") + with pytest.raises(IntegrityError): + read_checked(target, schema_version=1) + + +def test_module_exports_public_contract() -> None: + for name in ( + "IntegrityError", + "atomic_write", + "compute_content_hash", + "read_checked", + ): + assert name in state_store.__all__ + assert hasattr(state_store, name) diff --git a/agent-team/tests/test_task_model.py b/agent-team/tests/test_task_model.py new file mode 100644 index 0000000..28886c0 --- /dev/null +++ b/agent-team/tests/test_task_model.py @@ -0,0 +1,110 @@ +"""Unit tests for agent_team.task_model (§3.3, §3.3.1).""" + +from __future__ import annotations + +from agent_team.task_model import ( + Phase, + PipelineState, + TaskRecord, + TaskStatus, + new_thread_id, + task_from_dict, + task_from_json, + task_to_dict, + task_to_json, +) + + +def test_new_thread_id_unique_hex() -> None: + a = new_thread_id() + b = new_thread_id() + assert a != b + assert len(a) == 32 + int(a, 16) # must be valid hex + + +def test_phase_members() -> None: + assert {p.name for p in Phase} == { + "INTAKE", + "CLARIFY", + "PLAN", + "REVIEW", + "BUILD", + "VERIFY", + "PARKED", + "DONE", + } + + +def test_task_record_defaults() -> None: + rec = TaskRecord( + thread_id="t1", + status=TaskStatus.ACTIVE, + current_phase=Phase.INTAKE, + ) + assert rec.qa_history == [] + assert rec.plan is None + assert rec.review_verdicts == [] + assert rec.candidate_diff is None + assert rec.diff_hash is None + assert rec.ci_results is None + assert rec.transport == "" + + +def test_to_dict_serializes_enums_to_values() -> None: + rec = TaskRecord( + thread_id="t1", + status=TaskStatus.WAITING_HUMAN, + current_phase=Phase.CLARIFY, + ) + data = task_to_dict(rec) + assert data["status"] == "waiting_human" + assert data["current_phase"] == "clarify" + + +def test_roundtrip_dict() -> None: + rec = TaskRecord( + thread_id="t1", + status=TaskStatus.PARKED, + current_phase=Phase.PLAN, + qa_history=[{"q": "x", "a": "y"}], + plan={"phases": [1, 2]}, + review_verdicts=["REQUEST_CHANGES"], + candidate_diff="diff --git a b", + diff_hash="deadbeef", + ci_results={"conclusion": "success"}, + transport="slack", + created_at="2026-06-17T00:00:00Z", + updated_at="2026-06-17T01:00:00Z", + ) + restored = task_from_dict(task_to_dict(rec)) + assert restored == rec + + +def test_roundtrip_json() -> None: + rec = TaskRecord( + thread_id="t2", + status=TaskStatus.DONE, + current_phase=Phase.DONE, + diff_hash="abc", + ) + restored = task_from_json(task_to_json(rec)) + assert restored == rec + assert restored.status is TaskStatus.DONE + assert restored.current_phase is Phase.DONE + + +def test_pipeline_state_keys_mirror_task_record() -> None: + # Every PipelineState key should be a TaskRecord field. + state_keys = set(PipelineState.__annotations__) + record_fields = set(TaskRecord.__dataclass_fields__) + assert state_keys == record_fields + + +def test_pipeline_state_usable_as_dict() -> None: + state: PipelineState = { + "thread_id": "t1", + "status": "active", + "current_phase": "intake", + } + assert state["thread_id"] == "t1" diff --git a/agent-team/tests/test_transport_base.py b/agent-team/tests/test_transport_base.py new file mode 100644 index 0000000..a46364e --- /dev/null +++ b/agent-team/tests/test_transport_base.py @@ -0,0 +1,92 @@ +"""Unit tests for agent_team.transport.base (§3.3.1).""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from agent_team.transport.base import ( + GITHUB_MARKER_TEMPLATE, + NormalizedAnswer, + QuestionSet, + Transport, +) + + +def test_transport_is_abstract() -> None: + with pytest.raises(TypeError): + Transport() # type: ignore[abstract] + + +def test_question_set_fields() -> None: + qs = QuestionSet( + thread_id="t1", + question_id="q1", + turn=2, + questions=["a?", "b?"], + context={"repo": "x"}, + ) + assert qs.thread_id == "t1" + assert qs.question_id == "q1" + assert qs.turn == 2 + assert qs.questions == ["a?", "b?"] + assert qs.context == {"repo": "x"} + + +def test_question_set_context_defaults_empty() -> None: + qs = QuestionSet(thread_id="t", question_id="q", turn=0, questions=[]) + assert qs.context == {} + + +def test_normalized_answer_fields() -> None: + ans = NormalizedAnswer(question_id="q1", answer={"choice": 1}, via="slack") + assert ans.question_id == "q1" + assert ans.answer == {"choice": 1} + assert ans.via == "slack" + + +def test_concrete_subclass_implements_contract() -> None: + class FakeTransport(Transport): + def __init__(self) -> None: + self.posted: dict[str, Any] = {} + + def post_question( + self, *, thread_id, question_id, turn, question_set, deadline + ) -> str: + ref = f"slack-ts-{question_id}" + self.posted = { + "thread_id": thread_id, + "question_id": question_id, + "turn": turn, + "deadline": deadline, + "ref": ref, + } + return ref + + def parse_answer(self, raw) -> tuple[str, Any, str]: + na = NormalizedAnswer( + question_id=raw["callback_id"], answer=raw["value"], via="slack" + ) + return na.question_id, na.answer, na.via + + t = FakeTransport() + qs = QuestionSet(thread_id="t1", question_id="q1", turn=0, questions=["?"]) + ref = t.post_question( + thread_id="t1", + question_id="q1", + turn=0, + question_set=qs, + deadline="2026-06-18T00:00:00Z", + ) + assert ref == "slack-ts-q1" + assert t.posted["question_id"] == "q1" + + parsed = t.parse_answer({"callback_id": "q1", "value": "yes"}) + assert parsed == ("q1", "yes", "slack") + + +def test_github_marker_embeds_question_id() -> None: + marker = GITHUB_MARKER_TEMPLATE.format(question_id="abc123") + assert marker == "" + assert "abc123" in marker diff --git a/agent-team/tests/test_verifier.py b/agent-team/tests/test_verifier.py new file mode 100644 index 0000000..dc362b5 --- /dev/null +++ b/agent-team/tests/test_verifier.py @@ -0,0 +1,204 @@ +"""Unit tests for agent_team.nodes.verifier — the VERIFY node (§3.3, §3.3.2).""" + +from __future__ import annotations + +import pytest + +from agent_team.ci_gate import GateDecision, GateResult +from agent_team.nodes import verifier as verifier_mod +from agent_team.nodes.verifier import ( + DEFAULT_MAX_BUILD_LOOPS, + VerifierConfig, + set_fix_advisor, + verifier_node, +) +from agent_team.state_store import compute_content_hash +from agent_team.task_model import Phase, PipelineState, TaskStatus + + +def _diff_for(*paths: str) -> str: + chunks = [] + for p in paths: + chunks.append(f"diff --git a/{p} b/{p}\n@@ -1 +1 @@\n-old\n+new\n") + return "".join(chunks) + + +def _hash(diff: str) -> str: + return compute_content_hash(diff.encode("utf-8")) + + +def _state(diff: str, ci: dict | None) -> PipelineState: + return { + "thread_id": "t1", + "status": TaskStatus.ACTIVE.value, + "current_phase": Phase.VERIFY.value, + "candidate_diff": diff, + "diff_hash": _hash(diff), + "ci_results": ci, + } + + +@pytest.fixture(autouse=True) +def _reset_advisor(): + """Restore the default null advisor after each test.""" + yield + set_fix_advisor(verifier_mod._null_advisor) + + +# --------------------------------------------------------------------------- # +# PASS path +# --------------------------------------------------------------------------- # + + +def test_pass_advances_to_done() -> None: + diff = _diff_for("src/foo.py") + ci = {"run_id": "r1", "conclusion": "success", "diff_hash": _hash(diff)} + out = verifier_node(_state(diff, ci), VerifierConfig(expected_run_id="r1")) + assert out["status"] == TaskStatus.DONE.value + assert out["current_phase"] == Phase.DONE.value + assert out["ci_results"]["gate_decision"] == "pass" + + +def test_pass_does_not_consult_advisor() -> None: + calls: list = [] + + def advisor(result, state): # pragma: no cover - asserted not called + calls.append(result) + return "hint" + + set_fix_advisor(advisor) + diff = _diff_for("src/foo.py") + ci = {"run_id": "r1", "conclusion": "success", "diff_hash": _hash(diff)} + verifier_node(_state(diff, ci), VerifierConfig(expected_run_id="r1")) + assert calls == [] # the LLM is never asked whether it passed + + +# --------------------------------------------------------------------------- # +# FAIL path — loop back to BUILD +# --------------------------------------------------------------------------- # + + +def test_fail_loops_back_to_build() -> None: + diff = _diff_for("src/foo.py") + ci = {"run_id": "r1", "conclusion": "failure", "diff_hash": _hash(diff)} + out = verifier_node( + _state(diff, ci), VerifierConfig(expected_run_id="r1", build_loops=0) + ) + assert out["status"] == TaskStatus.ACTIVE.value + assert out["current_phase"] == Phase.BUILD.value + + +def test_fail_consults_advisor_for_hint() -> None: + def advisor(result: GateResult, state) -> str: + assert result.decision is GateDecision.FAIL + return "bump the pinned version" + + set_fix_advisor(advisor) + diff = _diff_for("src/foo.py") + ci = {"run_id": "r1", "conclusion": "failure", "diff_hash": _hash(diff)} + out = verifier_node(_state(diff, ci), VerifierConfig(expected_run_id="r1")) + assert out["review_verdicts"][0]["fix_hint"] == "bump the pinned version" + + +def test_fail_parks_when_build_loops_exhausted() -> None: + diff = _diff_for("src/foo.py") + ci = {"run_id": "r1", "conclusion": "failure", "diff_hash": _hash(diff)} + cfg = VerifierConfig( + expected_run_id="r1", + build_loops=DEFAULT_MAX_BUILD_LOOPS - 1, + max_build_loops=DEFAULT_MAX_BUILD_LOOPS, + ) + out = verifier_node(_state(diff, ci), cfg) + assert out["status"] == TaskStatus.PARKED.value + assert out["current_phase"] == Phase.PARKED.value + assert any("max build loops" in r for r in out["review_verdicts"][0]["reasons"]) + + +def test_fail_increments_build_loops_in_verdict() -> None: + diff = _diff_for("src/foo.py") + ci = {"run_id": "r1", "conclusion": "failure", "diff_hash": _hash(diff)} + out = verifier_node( + _state(diff, ci), VerifierConfig(expected_run_id="r1", build_loops=1) + ) + assert out["review_verdicts"][0]["build_loops"] == 2 + + +# --------------------------------------------------------------------------- # +# BLOCK path — park for human + GPT cross-review +# --------------------------------------------------------------------------- # + + +def test_block_on_denylist_parks() -> None: + diff = _diff_for(".github/workflows/ci.yml") + ci = {"run_id": "r1", "conclusion": "success", "diff_hash": _hash(diff)} + out = verifier_node(_state(diff, ci), VerifierConfig(expected_run_id="r1")) + assert out["status"] == TaskStatus.PARKED.value + assert out["current_phase"] == Phase.PARKED.value + assert out["ci_results"]["gate_decision"] == "block" + + +def test_block_on_hash_mismatch_parks() -> None: + diff = _diff_for("src/foo.py") + state = _state(diff, {"run_id": "r1", "conclusion": "success"}) + state["diff_hash"] = "tampered" + out = verifier_node(state, VerifierConfig(expected_run_id="r1")) + assert out["status"] == TaskStatus.PARKED.value + + +def test_block_never_advances_to_done_even_with_advisor() -> None: + set_fix_advisor(lambda result, state: "irrelevant") + diff = _diff_for(".github/workflows/ci.yml") + ci = {"run_id": "r1", "conclusion": "success", "diff_hash": _hash(diff)} + out = verifier_node(_state(diff, ci), VerifierConfig(expected_run_id="r1")) + assert out["status"] != TaskStatus.DONE.value + + +def test_missing_diff_parks() -> None: + state: PipelineState = { + "thread_id": "t1", + "candidate_diff": None, + "diff_hash": None, + "ci_results": None, + } + out = verifier_node(state, VerifierConfig(expected_run_id="r1")) + assert out["status"] == TaskStatus.PARKED.value + assert any("no candidate_diff" in r for r in out["review_verdicts"][0]["reasons"]) + + +# --------------------------------------------------------------------------- # +# Provenance / partial-update shape +# --------------------------------------------------------------------------- # + + +def test_verdict_records_provenance() -> None: + diff = _diff_for("src/foo.py") + ci = {"run_id": "r9", "conclusion": "success", "diff_hash": _hash(diff)} + out = verifier_node(_state(diff, ci), VerifierConfig(expected_run_id="r9")) + verdict = out["review_verdicts"][0] + assert verdict["stage"] == "verify" + assert verdict["run_id"] == "r9" + assert verdict["diff_hash"] == _hash(diff) + assert verdict["ci_conclusion"] == "success" + assert "at" in verdict + + +def test_returns_partial_update_only() -> None: + diff = _diff_for("src/foo.py") + ci = {"run_id": "r1", "conclusion": "success", "diff_hash": _hash(diff)} + out = verifier_node(_state(diff, ci), VerifierConfig(expected_run_id="r1")) + # Node returns only the keys it writes (LangGraph reducer merges the rest). + assert set(out) == { + "status", + "current_phase", + "review_verdicts", + "ci_results", + "updated_at", + } + + +def test_allowed_scope_threaded_to_gate() -> None: + diff = _diff_for("src/foo.py", "elsewhere/bar.py") + ci = {"run_id": "r1", "conclusion": "success", "diff_hash": _hash(diff)} + cfg = VerifierConfig(expected_run_id="r1", allowed_scope=["src/"]) + out = verifier_node(_state(diff, ci), cfg) + assert out["status"] == TaskStatus.PARKED.value # out-of-scope -> BLOCK -> park diff --git a/conftest.py b/conftest.py new file mode 100644 index 0000000..c371733 --- /dev/null +++ b/conftest.py @@ -0,0 +1,12 @@ +"""Root pytest configuration for the orchestrator repo. + +The ``agent-team/`` subproject ships its own ``tests/`` package with its own +``conftest.py`` and path bootstrap. Because both that directory and the repo's +top-level ``tests/`` are named ``tests`` (each with an ``__init__.py``), a single +repo-root ``pytest --collect-only`` maps both to the same ``tests`` package and +fails with ``ImportPathMismatchError``. The agent-team suite is therefore +collected/run by its OWN CI job (``cd agent-team && pytest``), and excluded from +the root collection here. +""" + +collect_ignore_glob = ["agent-team", "agent-team/*"] diff --git a/docs/r720-agent-team-design.md b/docs/r720-agent-team-design.md new file mode 100644 index 0000000..a4e4530 --- /dev/null +++ b/docs/r720-agent-team-design.md @@ -0,0 +1,490 @@ +# R720 Agent Team — Design (v2) + +Status: **DESIGN LOCKED — ready to build. Not built yet.** **v5** folded the final plan-review refinements +(SQLite `BEGIN IMMEDIATE` for the compare-and-set §3.3.1; CI denylist defense-in-depth + authenticated-only +pass/fail gate + compromised-box honesty §3.3.2; contention reserve + parked-task aging §6.6; backup integrity +definition + post-restore reconciliation §6.7; tested rollbacks §7; per-role canaries §6.4; CLI audit/confirm +§3.3.1). The design went through 3 GPT-4.1 `sh-plan-review` cycles (v1, v3, v4); the architecture was stable +throughout and remaining grain is now build-time implementation detail captured in the P1/P3 exit gates. **No +further gate runs by decision; build may begin with Phase 0 / P1.** Earlier status: Drafted 2026-06-17. **v3** reframes around Adam's clarified north star: the R720 +is not just scheduled checkers, it is a self-hosted, human-gated **agentic SDLC pipeline** +(intake -> clarify -> plan -> review -> build -> verify), fed from the Mac harness and tickets, with the +scheduled checkers as one task source. v2's roster work becomes "Plane 1"; the pipeline is "Plane 2" and the +centerpiece. **v4** folds in every v3 plan-review finding: B2 (§3.3.1) and B4 (§3.3.2) resolved with P1/P3 exit +gates; B1 billing realism + contention (§6.6); B3/B5/F5 provisioning + rollback + cross-review gate (§7); F1 +state durability + backup (§6.7); F2 canary update process (§6.4); F3 service-account lifecycle + F1 backups +(§9); F4 runbook incident handling (Phase 6); Q3 escalation ladder (§5); Q1/Q2 in §3.3.1. Ready for a re-run of +`sh-plan-review`. No build until that gate passes and Adam approves. + +Extends the `sh-secrev` pattern (see `security-review/DEPLOY-R720.md`) from a single security sweep into a small +roster of scheduled, unattended agents that check, plan, and (carefully) build, coordinated by a Claude `claude +-p` brain and using GPT / Gemini / DeepSeek where each is the better fit. + +This doc is the plan, not a runbook. It deliberately reuses the security agent's proven substrate. + +## 0. Locked decisions (this revision) + +| # | Decision | Choice | +|---|---|---| +| D1 | Headless `claude -p` under Max | **Permitted for now** (Anthropic pushed the disallowing ToS change to a later, unannounced date). Build the auth path **swappable** so the cutover is a config flip. See `reference_claude_subscription_billing`. | +| D2 | Fixer write path | **Option B**: the always-on box stays read-only; it emits a patch + opens an issue, and a trusted org CI workflow (OIDC) applies the patch on a branch and opens the **draft** PR. No standing write token on the box. | +| D3 | Checker/planner output mode | **Report + ALARM-only to start.** Clean nights post nothing; confirmed criticals alarm Slack; everything else lands in a mode-600 report. No auto-Jira/Notion writes until signal quality is trusted. | +| D4 | Billing/auth resilience | **Build a billing-mode abstraction now** (subscription OAuth default, API-key / Bedrock fallback ready). | +| D5 | aws-posture | **Resident on the box via IAM Roles Anywhere**, with a new **step-ca** internal CA for automated short-lived leaf-cert rotation (no long-lived AWS key on the box). | +| D6 | Confluence write identity | **Dedicated `confluence-bot` Atlassian service account, edit scoped to the IT space only** (Confluence API tokens inherit the whole user's permissions, so a scoped service account is how we bound blast radius). Token in `~/secrev.env` (mode 600). Costs one Confluence seat. | +| D7 | Confluence agent modes + Mermaid | **Scheduled = read + recommend only** (gaps/staleness into the report, per D3, never auto-writes). **On-demand = SSH-invoked from Adam's Mac** for an actual write. Mermaid map edits go through `~/.claude/scripts/confluence_mermaid.py` (ADF-only, dry-run-default, macro-count + revert-diff guarded). | +| D8 | Two-plane architecture | **Plane 1** = scheduled checkers (v2 roster), which also act as a task source. **Plane 2** = the agentic SDLC task pipeline (the centerpiece). Shared substrate. | +| D9 | Pipeline foundation | **LangGraph (the open-source library, runs in-process, NOT SaaS) + a local SQLite checkpointer.** Durable, resumable graph; `interrupt()` for the human gate. Reuses the existing stack. | +| D10 | Human-in-the-loop transport | **Pluggable, all three adapters** (Slack Block Kit, GitHub/Jira ticket comments, Claude Code on the Mac); Adam picks the channel per task. A transport-agnostic notify/resume layer maps answers back to the task thread. | +| D11 | Build/verify execution | **Org CI is the primary sandbox.** Builders emit candidate diffs; the Option-B OIDC workflow builds/tests/security-reviews; the verifier agent reads CI results. Keeps the 4GB box light + read-only. Dedicated builder VM only if fast local loops prove necessary. | +| D12 | Observability (LangSmith) | **Deprecate LangSmith (the SaaS tracer).** Keep LangGraph (framework, local). Local JSONL (`telemetry.py`, already exists) is the default; self-hosted Phoenix optional later. No SaaS dependency. | +| D13 | Task intake | **Mac harness first** (SSH-invoke enqueues onto the box), **GitHub issues next** (phased). | + +## 1. Goal and scope + +Stand up a coordinated team of agents on the existing `sh-secrev` R720 VM that runs unattended on a schedule, +operates across every Sea-Haven-Industries org repo automatically, stays cost-bounded, and reports through one +alarm channel. Claude (subscription OAuth) coordinates and does deep reasoning; Gemini does broad scans; GPT does +adversarial cross-checks; DeepSeek does mechanical code edits. + +In scope: read-mostly checkers, a low-blast-radius planner, a resident AWS posture check (D5), and a CI-applied +draft-PR fixer (D2). Out of scope: anything interactive or needing back-and-forth, and (per the secrev ethos) any +standing **write** credential on the always-on box. + +This does **not** replace the security review agent; it sits beside it and reuses its plumbing. + +## 2. What we reuse vs. what is new + +Reused from `sh-secrev` as-is: the VM, the systemd-timer model, the OAuth billing path, clean-clone +auto-discovery into `~/repo-mirrors` (read-only PAT, scrubbed post-fetch; the team scans the **same mirrors**, it +does not re-clone), the budget primitives (per-call + total caps, fail-toward-over-reporting), ALARM-only Slack, +mode-600 reports, the anti-complacency canary + coverage-rotation idea, the `orchestrator/` rsync deploy, and the +non-Claude provider keys in `~/orchestrator/.env`. + +New: a **coordinator** runner (`agent-team/` dir; entry CLI `run-team.py`, importable modules stay snake_case per +Python rules; the resource/dir name is kebab-case per handbook), per-role prompt/checklist modules each with its +own canary, a **billing-mode abstraction** (D4), the **fixer patch -> CI -> draft-PR** path (D2), and the +**step-ca + Roles Anywhere** setup for aws-posture (D5). + +## 3. Architecture + +``` +systemd timer (shared with secrev — see §8) + │ + ▼ + mirror step (reuse sh-secrev discovery) ──► ~/repo-mirrors (read-only) + │ + ▼ + COORDINATOR (claude -p via billing-mode abstraction; read-only tools + Bash to call orchestrator/gh) + │ shared budget ledger + versioned rotation/coverage state + ├──────────────┬───────────────┬────────────────┬──────────────┬───────────┐ + ▼ ▼ ▼ ▼ ▼ ▼ + drift checker dep/CVE checker doc-drift checker aws-posture planner (each role + (Gemini scan (Claude + (Gemini large- (Sonnet via (Claude) has its own + + Claude judge) GPT tiebreak) context) Roles Anywhere) canary) + │ │ │ │ │ + └──────────────┴───────────────┴────────────────┴──────────────┘ + │ structured JSON per agent + ▼ + coordinator: dedup + prioritize + route + │ + ┌──────────────────────────┼───────────────────────────┐ + ▼ ▼ ▼ + Slack ALARM mode-600 report fix-spec queue (D2) + (confirmed crit) (everything else; │ + no auto-Jira/Notion yet, D3) ▼ + FIXER: DeepSeek edit + Claude + spec + GPT review ──► patch + issue + │ + ▼ org CI (OIDC) applies patch, + opens DRAFT PR, runs pre-push + hooks + CI gates + Claude Code App +``` + +Model assignment (matches how `orchestrator` already splits them): **Claude (subscription)** = coordinator, deep +checks, fix-spec authoring; **Gemini 2.5 Pro** = broad whole-repo scans; **GPT-4.1** = adversarial cross-check / +tiebreak / PR review; **DeepSeek** = mechanical patch writing. + +### 3.1 Billing-mode abstraction (D4) +A single `claude_invoke(...)` seam selects the Claude auth/billing path from config: `subscription` (OAuth token, +default today), `api` (metered `ANTHROPIC_API_KEY`), or `bedrock` (cross-account Bedrock, already used by secrev +for the rare cross-family tiebreak). Switching modes is a config flip, not a code change. The box still pops any +stray `ANTHROPIC_API_KEY` in `subscription` mode so OAuth cannot be silently overridden. + +### 3.2 Relationship to the LangSmith orchestrator (what stays vs what the R720 hosts) + +Two orchestrators coexist after this plan; neither replaces the other. The split is **trigger + Claude billing**, +not capability. + +- **LangSmith orchestrator (Mac-hosted, unchanged).** The existing `orchestrator/` (LangGraph router + memory + retriever + Composio connector + LangSmith tracing on the `orchestration` project) stays the **on-demand, + interactive** delegation path: one task -> retrieve memory -> route -> one agent -> result, API-billed. This is + the CLAUDE.md hybrid-delegation path Claude Code uses for cross-family review, large scans, fast coding, and + connector actions. It stays one-shot and stateless (Q1: not refactored for persistence). +- **R720 orchestrator (the agent-team coordinator, new).** The scheduled, unattended, multi-agent layer: + cadence, the coordinator brain, shared budget ledger, versioned rotation/coverage state, the clean-clone mirror + corpus, and per-role canaries. Claude work here runs **headless under subscription OAuth** (Agent SDK), + billing-mode-swappable (§3.1). + +| Component | Today (LangSmith orchestrator, Mac) | After this plan | +|---|---|---| +| Trigger | On-demand from Claude Code / CLI | + scheduled (systemd timer) and SSH-invoked on-demand, on the R720 | +| Execution shape | One task -> one agent (stateless) | + multi-agent coordination with shared budget + versioned state (R720) | +| Claude billing | Metered API key | **Subscription OAuth on the R720** (swappable to api/bedrock per §3.1) | +| Non-Claude (GPT-4.1 / Gemini / DeepSeek) | `run.py` router, API-billed, LangSmith-traced | **unchanged in shape** — the R720 coordinator calls the **local** `~/orchestrator/run.py` (already rsync'd to the box) for these single-shot sub-tasks, so they keep API billing + LangSmith tracing | +| Memory retriever + embeddings cache | Mac | reused read-only by both (the box's rsync'd copy embeds the same memory store) | +| Composio connector (Slack/Notion/GitHub) | Mac | reused; the R720 routes Slack/Jira/Notion through it. **Confluence stays OUT of the connector** — native Atlassian MCP on the Mac for interactive edits, `confluence_mermaid.py` + REST for the box | +| Observability | LangSmith SaaS tracing (`orchestration` project) | **LangSmith deprecated (D12)** — local JSONL (`telemetry.py`) default, self-hosted Phoenix optional. LangGraph framework stays (it is not SaaS) | +| `models.py` factories + model-ID constants | Mac | shared code (rsync'd); single source of truth for both | + +**What does NOT migrate (stays Mac / interactive):** the daily-driver Claude Code sessions, the hybrid on-demand +delegation, and interactive Confluence edits via the native Atlassian MCP. + +**What is genuinely NEW on the R720 (not a migration — these never existed in the LangSmith orchestrator):** +scheduling, the coordinator + shared state/budget, canary/coverage, and subscription-OAuth Claude. + +Net: the LangSmith orchestrator keeps its job (on-demand routing, non-Claude execution, tracing, connector); the +R720 becomes the host for everything **scheduled, stateful, and subscription-billed**, and it **reuses the +LangSmith orchestrator in place** (the local rsync'd copy) for the non-Claude single-shots rather than +re-implementing them. + +### 3.3 Plane 2 — the agentic SDLC pipeline (the centerpiece) + +A durable, human-gated task pipeline hosted on the R720. A task is a long-lived, resumable record; the +coordinator drives it through stages, asking Adam for input when it is not confident and handing off to the org +CI to actually build and verify. + +``` +INTAKE ─► CLARIFIER ─► PLANNER ─► REVIEW LOOP ─► BUILDERS ─► VERIFIERS ─► draft PR + report + │ │ │ │ │ │ +Mac harness asks Adam phased plan GPT-4.1 + Claude spec org CI builds/tests/ +(SSH-invoke) question- (Claude) multi-model + DeepSeek security-review; +GitHub issue sets until adversarial; edits ─► verifier reads results; +checker 98%+, then loops back candidate loops back to builders +finding HUMAN GATE to planner diff on failure +``` + +**Stages and model per stage:** +- **Intake** — a task enters from the Mac harness (D13, first), a GitHub issue (next), or a Plane-1 checker + finding. It is written as a new task record (LangGraph thread) with a unique `thread_id`. +- **Clarifier (Claude)** — gathers context (repo, memory, handbook), then asks Adam **question-sets until 98%+ + confident**. This is a LangGraph `interrupt()`: the task suspends and checkpoints, a question-set is delivered + over the chosen transport (D10), and the task resumes via `Command(resume=...)` when the answer arrives. The + **human gate**: no progression to build without the clarifier clearing the bar and Adam approving the plan. +- **Planner (Claude)** — produces a phased plan (the format these design docs use). +- **Review loop (GPT-4.1 + optional multi-model)** — adversarial plan review (the `sh-plan-review` / + `cross_reviewer` discipline). Loops back to the planner on REQUEST CHANGES; escalates to Adam if it cannot + converge. +- **Builders (Claude spec + DeepSeek edits)** — turn the approved plan into a **candidate diff**. They do not + write to repos; per D2/D11 they emit the diff for CI. +- **Verifiers (org CI + a Claude/GPT reader)** — CI (Option-B OIDC) applies the diff on a branch, builds, runs + tests + the security review + lint; the verifier agent reads the CI results and either loops back to builders + or advances. Confirmed pass produces a **draft PR** plus a report to Adam. + +**Durable state (D9).** LangGraph (local) + a SQLite checkpointer. Each stage transition is checkpointed, so a +crash, a budget pause, or an overnight wait on a human answer all resume cleanly instead of restarting. The task +record holds: status, current phase, the full Q&A history, the plan, review verdicts, the candidate diff, and CI +results. + +**Human-in-the-loop (D10).** A small transport-agnostic responder service owns the notify+resume seam: it posts +the interrupt's question-set to the channel Adam chose for that task (Slack Block Kit / ticket comment / a Claude +Code session) and maps his reply back to the right `thread_id` to resume it. Adapters are independent so one can +ship first (Slack) and the others follow. + +**Stability + autonomy bounds (the "stable" requirement).** Hard gates, not vibes: (1) no build before the +clarifier hits 98% AND Adam approves the plan; (2) draft PRs only, never auto-merge; (3) the verifier must pass +or the task loops/holds, never ships; (4) per-task budget cap inside the shared nightly cap (§6.1); (5) every +stage checkpointed so failures resume, not restart; (6) a task that stalls (no human answer within a window, or +N failed build loops) parks and ALARMs rather than spinning. Plane-1's canary/coverage discipline applies to the +pipeline's agents too. + +### 3.3.1 Durable human-in-the-loop suspend/resume (resolves B2) + +LangGraph `interrupt()` + the SQLite checkpointer suspend and resume the graph, but the checkpoint alone does not +track the human-interaction lifecycle (delivery, duplicate/late answers, expiry). So the pipeline adds one +durable source of truth, a SQLite `pending_questions` table, and resolves every race with an atomic +compare-and-set against it. This is the riskiest mechanic, so it is specified here and P1 must prove it. + +- **Identity.** Each task is a graph `thread_id`. Each question-set gets a `question_id` (uuid) and a monotonic + `turn` within the task. The interrupt payload carries `{thread_id, question_id, turn, question_set, transport, + deadline}`. +- **Ledger.** `pending_questions(question_id PK, thread_id, turn, status[open|answered|expired|superseded], + transport, channel_ref, posted_at, deadline_at, answer_json, answered_at, answered_via)`. The LangGraph + checkpoint holds graph state; this table holds the question lifecycle and is what delivery, the responder, and + recovery read. +- **Delivery (and lost-post).** On interrupt, write the row `open` first, then post to the chosen transport and + store its `channel_ref` (Slack message ts / issue-comment id / Claude session id). The posted question embeds + the `question_id` (Slack `callback_id`; a `` marker in a GitHub comment). If the post + fails, the row stays `open` with no ref and a reconcile loop retries idempotently. +- **Answer mapping + idempotency (first-answer-wins).** Each transport's inbound adapter normalizes an answer to + `(question_id, answer, via)`. The responder then runs one atomic statement: + `UPDATE pending_questions SET status='answered', answer_json=?, answered_via=? WHERE question_id=? AND + status='open'`. rowcount 1 = first valid answer, enqueue a resume job; rowcount 0 = the question was not open + (already answered/expired/superseded), so the answer is a duplicate or late and is ignored with a "already + closed" reply. This single compare-and-set makes duplicate clicks, transport redelivery, answers via two + channels, and answer-after-timeout all safe. The statement runs inside a `BEGIN IMMEDIATE` transaction (SQLite's + default deferred isolation does not serialize concurrent responders, so the check-and-set must take the write + lock up front). +- **Resume (single-flight, turn-guarded).** A resume worker serializes per `thread_id` and calls + `graph.invoke(Command(resume=answer), {configurable:{thread_id}})`. Before resuming it checks the live + checkpoint is still interrupted on this `turn`; if the graph already advanced (stale/redelivered job) it marks + the question `superseded` and skips. A resume can never double-apply. +- **Deadline / no-answer.** Each open question has `deadline_at`. A timer loop flips overdue `open` rows to + `expired` (same compare-and-set) and applies the task policy: park + ALARM Adam, or apply a defined default + answer. An answer arriving for an already-`expired` question loses the compare-and-set and is ignored. Timeout + vs answer is a deterministic race on flipping `open`. +- **Restart recovery.** All state is durable (both SQLite stores), so a reboot converges via a startup sweep: + retry delivery for `open` rows lacking a ref; re-enqueue resume for `answered` rows whose graph is still + interrupted on that turn (idempotent via the turn guard); run the deadline policy for overdue `open` rows. No + in-memory-only state. +- **Parallel tasks + isolation (Q1).** Each task is its own `thread_id` with its own checkpoint and ledger rows; + the resume worker serializes per thread but runs different threads concurrently within the budget cap. +- **Manual path.** A small CLI over the ledger lets an operator list `open`/`parked` questions, re-deliver, + force-expire, or answer on a task's behalf; a stuck task parks rather than spins. Destructive CLI actions + (force-expire, answer-on-behalf, force-resume) are audit-logged and require an explicit confirmation flag. +- **Transport seam (Q2 fallback).** A `Transport` interface (`post_question(...) -> channel_ref`, + `parse_answer(raw) -> (question_id, answer, via)`) with Slack / GitHub / Claude Code adapters; the ledger + + resume logic are transport-independent. Adam picks the channel per task at intake. If the chosen transport is + unreachable, reconcile retries and, after N failures, falls back to a Slack ALARM pointing at the task. + +### 3.3.2 CI-as-verifier trust boundary (resolves B4) + +The builders are semi-trusted at best: an LLM that read repo content can be wrong or prompt-injected, so **the +candidate diff is treated as untrusted code.** The threat is that executing it in CI with org credentials lets a +bad diff exfiltrate secrets, assume the deploy role, or tamper with other repos. Five boundaries bound it: + +1. **Split CI: untrusted execution is credential-less; privileged steps never see the patch.** The job that + checks out and runs the diff (install/build/test) runs with `permissions: contents: read`, **no secrets, no + OIDC, no write token**, and restricted network egress. The patch executes only here, where there is nothing + to steal and nothing to assume. Any privileged action (the OIDC role, authoritative status, opening the PR) + runs in a **separate job that does not check out or execute patch-controlled code**; it consumes the + build/test report as data only. This is the standard untrusted-code-in-CI ("pwn request") mitigation, so the + workflow must NOT use `pull_request_target` with a checkout of the head ref. +2. **The patch may not touch the trust-control surface.** A box-side check and a CI guard both reject any + candidate diff that modifies `.github/workflows/**`, IAM/policy/permission IaC (CDK/SAM), branch-protection / + `CODEOWNERS` / Dependabot config, or files outside the task's declared scope. Such a diff is escalated to + mandatory human review + GPT cross-review, never auto-built (those files are the mandatory-cross-review + surface regardless). The match is not naive: enforcement is a **CI-side hard fail** (not only the box check), + it resolves symlinks and canonicalizes paths, and it rejects renames into denied paths and build steps that + generate files into them, so a path match cannot be bypassed by indirection. +3. **Diff integrity, box to CI.** The builder hashes the candidate diff and records the hash in the task ledger. + The box has no write token (D2), so the diff reaches CI as a signed artifact / via a short-lived + single-purpose token that can only create the agent branch; CI verifies the diff hash matches the ledger + before applying. A tampered or substituted diff fails the hash check. +4. **Pass/fail is a pure-code gate over authenticated CI results, not the LLM verifier.** Mirroring secrev's "one + pure-code script owns the block decision," a deterministic gate reads the CI run conclusion via the GitHub + Checks/Actions API (authenticated as the box read-only PAT), keyed to the specific run id + diff hash. It + consumes **only** that authenticated, patch-independent conclusion; it never trusts a success/failure file or + artifact the patch could have written. The verifier **agent** only reads failures to propose the next fix; it + cannot declare success. +5. **Branch protection is the final enforcement.** The draft PR targets a protected branch; the locked-down + build/test/security-review are required status checks; merge needs them green plus the security-review, the + Claude Code App review, and human approval. The agent never auto-merges (D2). + +The CI apply/verify workflow is itself IaC + untrusted-input handling, so before it ships it goes through +**`/sh-security-review` and the mandatory GPT-4.1 cross-review** (in P3), and its actions are SHA-pinned per the +handbook Pinning Principle. Provenance (diff hash, CI run id, checks consumed, gate decision) is recorded per +task; a diff that tried to touch the denylist, a build job that attempted egress, or a checks/hash mismatch is an +ALARM. Egress restriction in the untrusted job is enforced at the runner level (a harden-runner-style egress +policy or a self-hosted-runner network deny), but the primary mitigation is that the job holds no secrets, so a +successful exfil yields nothing of value. + +**Threat-model honesty (the diff-hash does not cover a compromised box).** The hash protects the box->CI transit +(tamper/substitution), not a box that is itself compromised: the box is in the trusted computing base for what it +*proposes*. The real backstops against a bad box are the credential-less build job, the trust-control-surface +denylist, and above all the **human review of the draft PR** plus required checks (nothing merges without Adam). +A signed-artifact/attestation step hardens transit further but is not claimed to defend a compromised box. + +## 4. Roster + +### Tier 1 — checkers (read-only) +| Agent | Models | Cadence | Output / gate | +|---|---|---|---| +| **compliance-drift** | Gemini scan + Claude judge | nightly | Drift vs engineering-handbook (naming, secrets placement, CI/CD present, Dependabot, branch protection). Report + Slack ALARM on violations (no auto-Jira yet, D3) | +| **dependency-cve** | Claude + GPT tiebreak | nightly | Cross-ref lockfiles vs advisories org-wide; report + feed fixer. Complements Dependabot | +| **doc-drift** | Gemini (large context) | weekly | Flags repos whose architecture moved but Confluence/README did not | + +### Tier 2 — aws-posture (resident, D5) + planner +| Agent | Models | Cadence | Output | +|---|---|---|---| +| **aws-posture** | Sonnet collectors + judge | weekly | Idle/anomalous spend (≈$330/mo flagged) + reasoning layer over baseline findings. Auths via **Roles Anywhere** (short-lived leaf certs, auto-rotated by step-ca). Complements existing GuardDuty/Security Hub/Config, does not replace them | +| **plan-groomer** | Claude | weekly | Drafts a groomed weekly plan **into the mode-600 report** for now (D3); auto-write to Notion/Jira is a later toggle once trusted | +| **confluence-doc** | Gemini scan + Claude judge | weekly (scheduled) + on-demand | **Scheduled:** diffs repos + AWS inventory + the page-ID map (`project_confluence_migration`) against Confluence, reports doc gaps / stale pages / missing runbooks (recommend-only, D3). **On-demand (SSH-invoked):** performs an actual update, including Mermaid map edits via `confluence_mermaid.py`. Writes as the IT-space-scoped `confluence-bot` (D6). Overlaps the existing `sh-confluence-audit`/`sh-confluence` skills; the box adds unattended cross-repo scope + the tested Mermaid script | + +### Tier 3 — fixer (D2) +| Agent | Models | Trigger | Output | +|---|---|---|---| +| **fixer** | DeepSeek edit + Claude spec + GPT review | on a confirmed, low-risk finding | Emits a patch + opens an issue; org CI applies it and opens a **draft** PR. Never auto-merges; pre-push hooks + CI + Claude Code App gate it (F3) | + +## 5. Coordination model + +Nightly, after the mirror refresh, the coordinator: (1) loads the shared budget ledger and the **versioned** +rotation/coverage state (F1); (2) runs the **canary suite first**, one planted-fault corpus per role, a miss is a +COMPLACENCY ALARM and that role is skipped; (3) fans out the scheduled agents over the mirrors, each with a +per-call cap, all drawing from **one shared total cap** (critical for the Claude subscription draw, §6.1), with +budget-exhausted roles deferred via the rotation pointer (never dropped) and a COVERAGE ALARM if a role slips +past `MAX_CYCLE_NIGHTS`; (4) collects each agent's structured JSON, dedups across agents, prioritizes; (5) routes +per D3: Slack ALARM for confirmed criticals, everything else to the mode-600 report, and fix-specs to the fixer +queue if Tier 3 is enabled; (6) a fully clean night posts nothing. + +**Escalation ladder (resolves Q3).** A confirmed critical that stays unaddressed escalates beyond a one-shot +Slack ALARM: it re-alarms on a backoff each night it persists, and after `N` nights (default 3) the coordinator +opens a tracking Jira ticket (INFRA) so it cannot quietly linger. The same ladder applies to a COMPLACENCY or +COVERAGE alarm that does not clear. Escalation stays ALARM-only in spirit (nothing posts on a clean state). + +The coordinator holds its own state; it does not rely on the orchestrator (one-shot, stateless). It may *call* +`orchestrator/run.py` for GPT/Gemini/DeepSeek single-shot sub-tasks, or call those providers directly (Q1: the +team implements its own coordination; the orchestrator is not refactored for persistence). + +## 6. Constraints and how this revision answers them + +### 6.1 Subscription billing draw — shared cap (B1 resolved by D1/D4) +Headless under Max is permitted for now (D1). Every Claude SDK call still draws from the **same Max pool as +interactive Claude Code**, so the team runs under **one shared nightly cap across all agents**, and pushes volume +to Gemini/GPT/DeepSeek (own-account billing) where quality allows. The billing-mode abstraction (D4) lets us flip +to `api`/`bedrock` when the ToS cutover lands. `ANTHROPIC_API_KEY` stays unset in subscription mode. + +### 6.2 Fixer write path (B2 resolved by D2) +Box stays read-only; CI applies the patch and opens the draft PR. The CI workflow + its OIDC role is **new IAM** +and goes through the **mandatory GPT-4.1 cross-review before it is built** (§7, B3). + +### 6.3 step-ca + Roles Anywhere (D5) +New internal CA (step-ca) issues short-lived leaf certs auto-renewed by a systemd timer; the Roles Anywhere +**trust anchor + the read-only AWS role** are **new IAM** and go through the mandatory cross-review before build +(B3). The leaf is short-lived (self-expiring), which is stronger than the box's long-lived GitHub PAT. + +### 6.4 Anti-complacency per role +Each agent ships with its own canary corpus, versioned in the repo. A role with a failing/stale canary is skipped +with a COMPLACENCY ALARM, never run silently degraded. **Canary update process (resolves F2):** canary corpora +are version-controlled **per agent/role** (no shared global corpus, to avoid cross-role confusion); a change to +any canary goes through a PR + review with a revert point, and because the canary runs every night, a canary edit +that silently weakens recall is itself caught on the next run. + +### 6.5 Host capacity +4GB / 2 vCPU / 40GB. Work is I/O-bound, but more report history + step-ca may pressure disk. Re-check headroom +after Phase 1; size up the Hyper-V VM (snapshot first per `feedback_ec2_replacement_snapshot` discipline) if +needed rather than risking the secrev workload. + +### 6.6 Claude budget realism + contention (resolves B1) +A multi-stage pipeline draws far more Claude than a single sweep (clarifier loop + planner + verifier-read per +task), all on the **same Max pool as Adam's interactive Claude Code**. Three controls: +- **Pre-build measurement (gate before P3 builds anything).** Using the existing `telemetry.py` token capture, + measure the per-stage Claude token draw on a representative task plus the worst-case clarifier loop, then + project per-task and daily aggregate at expected task volume. P1/P2 must emit these numbers before P3 + proceeds; the design measures, it does not assume. +- **One shared daily Claude cap across ALL R720 Claude work** (pipeline + Plane-1 sweeps), tracked in the + persistent budget ledger. Non-Claude stages are pushed to GPT/Gemini/DeepSeek (own-account billing) to keep + the Claude draw down. +- **Interactive-first contention rule.** Adam's interactive Claude Code is never blocked. The box keeps a + reserve headroom; before starting a stage it checks remaining headroom, and if below the reserve it **parks + new pipeline tasks and ALARMs** rather than competing for the pool. A task already mid-flight checkpoints and + pauses at the next stage boundary (never killed). A clarifier is capped at `N` turns per task, then + escalates/parks, so an ambiguous task cannot loop-drain the pool. +- **Fairness + no starvation.** The reserve is a fixed configured fraction of the daily cap (not guessed at + runtime). Parked tasks are FIFO-aged with a `MAX_PARK` window; a task that exceeds it escalates (ALARM, and a + Jira ticket per §5) rather than starving silently, and an operator can force-resume or re-prioritize it via the + CLI. A clarifier parked at the turn cap is resumable the same way: Adam adds context and re-opens it, so it is + never an indefinite stall. + +### 6.7 State durability + backup (resolves F1) +All durable state (the LangGraph SQLite checkpoint, the `pending_questions` ledger, the budget ledger, the +Plane-1 rotation/coverage pointer) is written atomically (write-temp-then-rename), integrity-checked on load, and +included in the nightly offsite backup (mode 600). On corruption the coordinator refuses to proceed silently: the +rotation pointer is rebuildable from the report history, and a corrupt task checkpoint parks that task with an +ALARM rather than restarting it blindly. "Integrity-checked" is concrete: schema-version match + a stored content +hash + a logical-consistency check (e.g. no `answered` question whose graph is already past that turn). After a +restore, a reconciliation step re-syncs against external state (in-flight CI runs, current GitHub PR status) +before any task resumes, so a restored backup cannot act on stale external assumptions. + +## 7. Phased rollout (re-sequenced for provisioning order, rollback, cross-review, and docs-as-you-go) + +- **Phase 0 — substrate factoring (with rollback, B4).** Back up `nightly_sweep.sh` (tag a revert point); + extract discovery/mirror/budget-ledger/rotation/Slack/canary into a shared module used by both secrev and the + team. **Gate:** secrev passes its canary + existing behavior after the refactor, else revert. No team behavior + yet. +- **Phase 1 — one checker end to end.** Build `compliance-drift` + its canary + the mode-600 report path + a + **routing dry-run** (F4) for the Slack alarm. Dry-run on the mirrors. Proves the substrate generalizes. + **Create the `project_r720_agent_team` memory now** (B5/docs-as-you-go). +- **Phase 2 — coordinator + second checker.** Add the coordinator (shared budget, dedup, versioned state F1) and + `dependency-cve`. **Run a forced budget-squeeze dry-run** to prove deferral-not-drop + COVERAGE ALARM (F2). +- **Phase 3 — doc-drift + step-ca/Roles Anywhere + aws-posture.** Stand up step-ca and the Roles Anywhere trust + anchor + read-only AWS role; **cross-review the IAM before building** the agent (B3). Wire aws-posture. Add + doc-drift. +- **Phase 4 — planner + confluence-doc.** `plan-groomer` writing into the report only (D3). For + `confluence-doc`: create the `confluence-bot` service account with IT-space-only edit rights (D6), put its + token in `~/secrev.env`; ship the scheduled gap-detection (read-only, recommend) first, then wire the + on-demand SSH-invoked write path. The Mermaid script (`~/.claude/scripts/confluence_mermaid.py`, already + written and offline-tested) must pass a **live dry-run against page 1540098** (verify it lists all 16 weweave + macros and that a no-op set is clean) before any `--apply`. Notion/Jira auto-write stays a later toggle. +- **Phase 5 — fixer (D2).** Build the org CI apply-and-open-draft-PR workflow; **cross-review its OIDC IAM + before building** (B3). Confirm fixer PRs hit pre-push hooks + CI + Claude Code App (F3). Start with the + narrowest fix class (dep bumps). Draft PRs only. +- **Phase 6 — document as standing infra.** Update Confluence (the team **and** the still-undocumented secrev + host) in the IT host/LAN inventory; write the operator runbook. **The runbook must cover incident handling + (resolves F4):** pipeline stalls, stuck/parked tasks, failed human-in-the-loop resumes, budget exhaustion + mid-pipeline, transport outages, and COMPLACENCY/COVERAGE alarms, each with the manual CLI recovery steps + (§3.3.1) and the escalation ladder (§5). + +**Provisioning + rollback + cross-review gate (resolves B3/B5/F5), applied to every phase:** +- **Cross-review is a hard gate, not a note.** Any new IAM role/policy, trust anchor, OIDC role, or permission + change is provisioned AND passes the mandatory GPT-4.1 cross-review (plus `/sh-security-review` where it + touches the CI / untrusted-input surface) **before** any code that depends on it is built. A phase cannot + start its dependent work until that review is recorded. +- **Every stateful phase has a revert point.** Not just Phase 0: before standing up step-ca, the Roles Anywhere + trust anchor + AWS role, or the CI apply-workflow, capture a documented rollback (remove the role/CA, restore + the prior workflow, revert the cert config) and gate the phase on a successful dry-run. The rollback is + **exercised** (sandbox or simulated teardown/re-provision), not merely written, before the phase is accepted. + VM changes snapshot first per `feedback_ec2_replacement_snapshot`. +- **Docs land with the change, enforced as definition-of-done.** A phase is not "done" until its memory entry + and the relevant Confluence page are updated; that update is a checklist item in the phase, not deferred + (Phase 6 is only the final standing-infra writeup). + +### 7.1 Pipeline track (Plane 2) — depends only on Phase 0 substrate + +This track is largely independent of the Plane-1 checker phases (1-6); both build on the Phase 0 substrate. +Given the north star, **Adam may prioritize this track first.** Sequencing within it: + +- **Phase P1 — skeleton + the human gate.** LangGraph graph + SQLite checkpointer on the box; one trivial task + type; Mac SSH-invoke intake; the **clarifier** with `interrupt()`/resume over **one** transport (Slack first). + Stops at an approved plan, no build yet. This proves durable suspend/resume across a real human answer (the + riskiest mechanic) before anything else. **Exit criteria (must demonstrate §3.3.1):** (a) kill the box + mid-wait and have the task resume after restart; (b) submit a duplicate answer and confirm it no-ops; (c) + submit an answer after the deadline expired and confirm it is rejected and the task parked; (d) two tasks + suspended concurrently resume independently to the correct thread. +- **Phase P2 — planner + review loop.** Wire the planner and the GPT-4.1 review loop (reuse `cross_reviewer`), + including loop-back and the escalate-to-Adam path. +- **Phase P3 — builders + verifier via org CI.** Builders emit a candidate diff; the Option-B OIDC workflow + builds/tests/security-reviews; the verifier reads CI results and produces a draft PR. Start with the narrowest + task class (e.g. a dependency bump or a single-file fix), draft PRs only. **Build the §3.3.2 trust boundary:** + split untrusted/privileged CI jobs, diff-hash integrity, the trust-control-surface denylist, and the pure-code + pass/fail gate. The CI apply/verify workflow + its OIDC role go through **`/sh-security-review` AND the + mandatory GPT-4.1 cross-review** before this phase ships (it is IaC/IAM + untrusted-input handling). +- **Phase P4 — more transports + GitHub intake.** Add the ticket-comment and Claude-Code responder adapters + (D10) and GitHub-issue intake (D13). +- **Phase P5 — checker findings as a task source.** Let a confirmed Plane-1 finding open a pipeline task, closing + the loop between the two planes. + +Observability for both planes (D12): LangSmith stays off; the existing local JSONL (`telemetry.py`) covers the +LangChain/LangGraph path, the Agent-SDK path keeps its own run logs, and self-hosted Phoenix is an optional +later add if per-run trace UI is wanted. + +## 8. Open items folded in (no longer blocking) +- Shared vs separate timer: **shared** with secrev (one discovery/mirror pass, one shared budget); error + isolation handled by per-role try/skip + canary, documented in the runbook (N2). +- Read-only PAT sufficiency (Q2): confirm the existing fine-grained PAT covers all mirrors before Phase 1; it + already clones every non-archived org repo for secrev, so this is a verification step, not a change. + +## 9. Obligations on build (per global instructions) +- **Memory:** `project_r720_agent_team` created in Phase 1; cross-link `project_security_review_agent`, + `project_orchestration_migration`, `reference_claude_subscription_billing`, `feedback_cloudwatch_alarms`. +- **Confluence:** document the team (and secrev) as standing infra (always-on VM holding read-only org PAT + now + a Roles Anywhere AWS identity). +- **Handbook/naming:** kebab-case dirs/resources (`agent-team`), snake_case importable Python modules; secrets in + `.env`/Secrets Manager/`~/secrev.env` (mode 600), never committed; CI/CD for the fixer apply-workflow. +- **Cross-review:** the fixer CI OIDC role and the Roles Anywhere trust anchor + AWS read role each go through the + mandatory GPT-4.1 cross-review before their phase builds. +- **Service-account lifecycle (F3):** the `confluence-bot` token is rotated on a schedule (90 days, calendared + like the GitHub PAT), has a documented revocation step, keeps its edits attributable in Confluence page + history, and is decommissioned if the agent is retired. +- **Backups (F1):** the durable state stores (checkpoint, ledgers, rotation pointer) are included in the nightly + offsite backup. diff --git a/requirements.txt b/requirements.txt index b4c01cf..0f02aa4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,6 @@ langgraph==1.1.10 +# Durable SQLite checkpointer for the R720 agent-team Plane-2 pipeline (design D9). +langgraph-checkpoint-sqlite==3.1.0 langchain-anthropic==1.4.3 langchain-openai==1.2.1 langchain-google-genai==4.2.2