From 270ce93b2aeb763cc51ebff5ca267933a8cd5012 Mon Sep 17 00:00:00 2001 From: Adam Moussa Date: Thu, 18 Jun 2026 12:56:42 -0400 Subject: [PATCH] feat(agent-team): bind P1 clarifier to real Claude via subscription-OAuth invoker Adds the billing-seam invoker (claude_agent_sdk subscription-OAuth, deferred import, API/Bedrock paths) and the Claude-backed clarifier callables (ConfidenceAssessor/QuestionGenerator, one call/turn memoized on (thread_id,len,content-hash), fail-safe to 0.0 so garbage never clears the 98% human gate). --- agent-team/agent_team/invoker.py | 300 ++++++++++++ agent-team/agent_team/nodes/clarifier_llm.py | 466 +++++++++++++++++++ agent-team/tests/test_clarifier_llm.py | 335 +++++++++++++ agent-team/tests/test_invoker.py | 255 ++++++++++ 4 files changed, 1356 insertions(+) create mode 100644 agent-team/agent_team/invoker.py create mode 100644 agent-team/agent_team/nodes/clarifier_llm.py create mode 100644 agent-team/tests/test_clarifier_llm.py create mode 100644 agent-team/tests/test_invoker.py diff --git a/agent-team/agent_team/invoker.py b/agent-team/agent_team/invoker.py new file mode 100644 index 0000000..58d11b3 --- /dev/null +++ b/agent-team/agent_team/invoker.py @@ -0,0 +1,300 @@ +"""Real Claude invokers for the ``claude_invoke`` billing seam (design §3.1). + +:mod:`agent_team.billing` owns mode selection and the subscription-mode env +hygiene, but delegates the actual SDK call to a pluggable invoker bound via +:func:`agent_team.billing.set_invoker`. This module supplies that invoker: a +single function matching the ``billing.Invoker`` signature +``(prompt, *, mode, **kw) -> ClaudeResult`` that dispatches on +:class:`~agent_team.billing.BillingMode`: + +* ``SUBSCRIPTION`` — the R720 default. Runs the Claude Agent SDK headless over + the subscription OAuth token (``CLAUDE_CODE_OAUTH_TOKEN``), mirroring the + canonical pattern in ``security-review/run_headless.py``. ``billing`` has + already popped any stray ``ANTHROPIC_API_KEY`` for the duration of the call, + so we only assert the OAuth token is present. +* ``API`` — a thin metered call through the ``anthropic`` SDK. +* ``BEDROCK`` — the rare cross-family tiebreak path; not wired for P1, so it + raises :class:`NotImplementedError` honestly (a later config-flip wires it). + +Deferred-import rationale: neither ``claude_agent_sdk`` nor ``anthropic`` is +installed in the test/Mac scaffolding environment, so importing either at +module load would raise :class:`ModuleNotFoundError` and break a clean import. +Following the deferred-import pattern of +:func:`agent_team.graph.build_sqlite_checkpointer`, the SDK imports live inside +the functions that actually call them and raise a clear :class:`RuntimeError` +when the package is missing. The SDK callables are also injectable (``_query``, +``_client``) so the real path stays unit-testable without the SDKs installed, +mirroring how the codebase keeps SDK calls injectable (see +:func:`agent_team.billing.set_invoker` and +:func:`agent_team.resume_worker.build_resume_command`). +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Callable + +from agent_team.billing import BillingMode, ClaudeResult, set_invoker + +__all__ = [ + "API_MODEL", + "api_invoker", + "bind_invoker", + "bind_subscription_invoker", + "subscription_invoker", +] + +# Metered model for the API path (the rare opt-in billing mode). +API_MODEL = "claude-sonnet-4-6" + +# Default per-call agent budget for the headless subscription path, in USD. +_DEFAULT_BUDGET_USD = 2.0 +_DEFAULT_MAX_TURNS = 40 + + +# --------------------------------------------------------------------------- # +# Subscription path (Claude Agent SDK, headless OAuth) +# --------------------------------------------------------------------------- # + + +def _require_agent_sdk() -> Any: + """Import and return ``claude_agent_sdk`` or raise a clear RuntimeError. + + Deferred so this module imports cleanly where the SDK is absent (the + Mac/test scaffold). Mirrors graph.build_sqlite_checkpointer. + """ + try: + import claude_agent_sdk + except ImportError as exc: # pragma: no cover - depends on optional dep + raise RuntimeError( + "claude_agent_sdk is unavailable; install it to use the " + "subscription billing path (the R720 default). Tests inject a fake " + "query via the _query parameter." + ) from exc + return claude_agent_sdk + + +async def _collect_subscription_text( + prompt: str, + *, + max_turns: int, + budget_usd: float, + model: str | None, + _query: Callable[..., Any] | None = None, + _options_cls: Callable[..., Any] | None = None, +) -> tuple[str, dict[str, Any], list[Any]]: + """Drive one headless Agent SDK turn; return (text, usage, raw_messages). + + ``_query``/``_options_cls`` default to the real ``claude_agent_sdk`` + callables (lazily imported) but are injectable so tests can supply a fake + async ``query`` without the SDK installed. Text extraction mirrors + ``run_headless.py``: prefer the terminal ``ResultMessage.result``, falling + back to concatenated ``AssistantMessage`` text blocks. + """ + if _query is None or _options_cls is None: + sdk = _require_agent_sdk() + if _query is None: + _query = sdk.query + if _options_cls is None: + _options_cls = sdk.ClaudeAgentOptions + + opts = _options_cls( + permission_mode="bypassPermissions", + setting_sources=[], # hermetic: ignore user/project/local config + CLAUDE.md + model=model, + max_turns=max_turns, + max_budget_usd=budget_usd, + ) + + texts: list[str] = [] + result_text: str | None = None + messages: list[Any] = [] + usage: dict[str, Any] = {} + async for msg in _query(prompt=prompt, options=opts): + messages.append(msg) + name = type(msg).__name__ + if name == "AssistantMessage": + for block in getattr(msg, "content", []) or []: + text = getattr(block, "text", None) + if text: + texts.append(text) + elif name == "ResultMessage": + result_text = getattr(msg, "result", None) + cost = getattr(msg, "total_cost_usd", None) + if cost is not None: + usage["total_cost_usd"] = float(cost) + sdk_usage = getattr(msg, "usage", None) + if isinstance(sdk_usage, dict): + usage.update(sdk_usage) + elif sdk_usage is not None: + usage["usage"] = sdk_usage + + return (result_text or "\n".join(texts)), usage, messages + + +def subscription_invoker( + prompt: str, + *, + mode: BillingMode, + max_turns: int = _DEFAULT_MAX_TURNS, + budget_usd: float = _DEFAULT_BUDGET_USD, + model: str | None = None, + _query: Callable[..., Any] | None = None, + _options_cls: Callable[..., Any] | None = None, + **kw: Any, +) -> ClaudeResult: + """Invoke Claude headless over the subscription OAuth token (§3.1). + + Asserts ``CLAUDE_CODE_OAUTH_TOKEN`` is present (the metered key is already + popped by :func:`agent_team.billing.claude_invoke` in subscription mode) and + refuses to run without it, naming ``~/secrev.env`` as the source. The Agent + SDK ``query()`` coroutine is bridged to this sync seam with + :func:`asyncio.run`; the box path is synchronous, but we fail clearly rather + than silently if invoked from inside a running event loop. + + ``_query``/``_options_cls`` are injection seams for tests; production leaves + them ``None`` so the real ``claude_agent_sdk`` callables are used. + """ + import os + + if not os.environ.get("CLAUDE_CODE_OAUTH_TOKEN"): + raise RuntimeError( + "subscription_invoker requires CLAUDE_CODE_OAUTH_TOKEN to be set " + "(source ~/secrev.env). Refusing to run the subscription OAuth path " + "without it." + ) + + coro = _collect_subscription_text( + prompt, + max_turns=max_turns, + budget_usd=budget_usd, + model=model, + _query=_query, + _options_cls=_options_cls, + ) + + try: + asyncio.get_running_loop() + except RuntimeError: + text, usage, messages = asyncio.run(coro) + else: # pragma: no cover - the box path is synchronous + coro.close() + raise RuntimeError( + "subscription_invoker cannot bridge asyncio.run from within a " + "running event loop; call claude_invoke from synchronous code." + ) + + return ClaudeResult(text=text, mode=mode, usage=usage, raw=messages) + + +# --------------------------------------------------------------------------- # +# API path (anthropic SDK, metered) +# --------------------------------------------------------------------------- # + + +def _require_anthropic() -> Any: + """Import and return the ``anthropic`` module or raise a clear RuntimeError.""" + try: + import anthropic + except ImportError as exc: # pragma: no cover - depends on optional dep + raise RuntimeError( + "anthropic is unavailable; install it to use the API billing path. " + "Tests inject a fake client via the _client parameter." + ) from exc + return anthropic + + +def _extract_text(message: Any) -> str: + """Join the text blocks of an anthropic Messages response.""" + parts: list[str] = [] + for block in getattr(message, "content", []) or []: + text = getattr(block, "text", None) + if text: + parts.append(text) + return "".join(parts) + + +def api_invoker( + prompt: str, + *, + mode: BillingMode, + model: str = API_MODEL, + max_tokens: int = 4096, + _client: Any | None = None, + **kw: Any, +) -> ClaudeResult: + """Invoke Claude through the metered ``anthropic`` SDK (§3.1, API mode). + + ``_client`` is an injection seam for tests; production leaves it ``None`` so + a real ``anthropic.Anthropic()`` is constructed (reading + ``ANTHROPIC_API_KEY`` from the environment, as the SDK does by default). + """ + if _client is None: + anthropic = _require_anthropic() + _client = anthropic.Anthropic() + + message = _client.messages.create( + model=model, + max_tokens=max_tokens, + messages=[{"role": "user", "content": prompt}], + ) + + usage_obj = getattr(message, "usage", None) + if usage_obj is None: + usage: dict[str, Any] = {} + elif isinstance(usage_obj, dict): + usage = dict(usage_obj) + elif hasattr(usage_obj, "model_dump"): + usage = usage_obj.model_dump() + else: + usage = { + "input_tokens": getattr(usage_obj, "input_tokens", None), + "output_tokens": getattr(usage_obj, "output_tokens", None), + } + + return ClaudeResult( + text=_extract_text(message), mode=mode, usage=usage, raw=message + ) + + +# --------------------------------------------------------------------------- # +# Dispatch + binding +# --------------------------------------------------------------------------- # + + +def real_invoker(prompt: str, *, mode: BillingMode, **kw: Any) -> ClaudeResult: + """Dispatch to the per-mode real invoker (the ``billing.Invoker`` contract). + + ``BEDROCK`` is the rare cross-family tiebreak path and is not wired for P1; + it raises :class:`NotImplementedError` honestly. Flipping it on is a later + config change, not a code rewrite of the seam. + """ + if mode is BillingMode.SUBSCRIPTION: + return subscription_invoker(prompt, mode=mode, **kw) + if mode is BillingMode.API: + return api_invoker(prompt, mode=mode, **kw) + if mode is BillingMode.BEDROCK: + raise NotImplementedError( + "BEDROCK billing is the rare cross-family tiebreak path and is not " + "wired for P1; enable it later via config-flip once the cross-account " + "Bedrock transport is provisioned." + ) + raise NotImplementedError(f"no invoker for billing mode {mode!r}") + + +def bind_invoker(invoker: Callable[..., ClaudeResult] | None = None) -> None: + """Bind a real invoker into the billing seam in one line at startup. + + Defaults to :func:`real_invoker` (mode-dispatching). Not called at import + time so importing this module has no global side effects. + """ + set_invoker(invoker or real_invoker) + + +def bind_subscription_invoker() -> None: + """Bind the mode-dispatching real invoker (subscription is the default mode). + + Convenience for the common R720 startup: one call wires + :func:`agent_team.billing.claude_invoke` to the real Claude path. + """ + set_invoker(real_invoker) diff --git a/agent-team/agent_team/nodes/clarifier_llm.py b/agent-team/agent_team/nodes/clarifier_llm.py new file mode 100644 index 0000000..cd6604d --- /dev/null +++ b/agent-team/agent_team/nodes/clarifier_llm.py @@ -0,0 +1,466 @@ +"""Claude-backed clarifier callables — the real §3.3 / §7.1 P1 bindings. + +:mod:`agent_team.nodes.clarifier` owns the *loop* (the LangGraph +``interrupt()``/resume 98% gate) but deliberately injects the two reasoning +seams so the loop stays pure and testable: + +* ``ConfidenceAssessor = Callable[[Sequence[object], PipelineState], float]`` +* ``QuestionGenerator = Callable[[Sequence[object], PipelineState], list[str]]`` + +This module supplies the **real, Claude-backed** implementations of those two +callables. It calls Claude only through the committed +:func:`agent_team.billing.claude_invoke` seam (§3.1) — never a raw SDK — so the +billing-mode hygiene and the budget ledger stay in one place. + +The naive binding is wasteful: the clarifier loop calls ``assess_confidence`` +and then ``generate_questions`` separately on the same turn, so two independent +implementations would make **two** Claude calls per turn for what is really one +reasoning step. :class:`ClaudeClarifier` instead makes **one** Claude call per +turn and serves both methods from the memoized result. The memo is keyed on the +Q&A history length, so a new answer (history grows) recomputes, while the +back-to-back assess/generate pair within one turn reuses the same call. + +Defensive parsing is a hard requirement here because the model output is +UNTRUSTED and this is the **human gate** (§3.3): a parse failure must *never* +clear the gate. The parser fails SAFE — a missing/garbled confidence defaults to +``0.0`` (so the loop keeps asking rather than falsely advancing to planning), +and a missing question-set below threshold falls back to a single generic +clarifying question (so the loop still has something to ask). +""" + +from __future__ import annotations + +import hashlib +import json +import re +from collections.abc import Callable, Sequence +from typing import Any + +from agent_team.billing import ClaudeResult, claude_invoke +from agent_team.nodes.clarifier import ( + DEFAULT_CONFIDENCE_THRESHOLD, + ConfidenceAssessor, + QuestionGenerator, +) +from agent_team.task_model import PipelineState + +__all__ = [ + "FALLBACK_QUESTION", + "ClaudeClarifier", + "build_claude_clarifier_callables", +] + +# The signature the billing seam exposes: ``claude_invoke(prompt, *, mode=None, +# config=None, **kw) -> ClaudeResult``. Injected so tests pass a fake, mirroring +# the injection pattern used across this codebase (billing.set_invoker, the +# clarifier loop's injected callables, etc.). +ClaudeInvoke = Callable[..., ClaudeResult] + +# Used when the model is below the confidence bar but supplied no usable +# question-set. The loop must always have something to ask rather than spin or +# falsely advance, so we substitute a generic clarifier prompt. +FALLBACK_QUESTION = ( + "Could you share more about the goal, scope, and constraints of this task " + "so I can be sure I understand it well enough to plan?" +) + +# Default system framing handed to Claude. Kept as a module constant so callers +# can override via the ``system`` constructor hook without forking the class. +_DEFAULT_SYSTEM = ( + "You are the CLARIFIER stage of an agentic SDLC pipeline and the human " + "gate before any planning happens. Your job is to decide whether the " + "requirement is understood well enough to plan, drawing conceptually on " + "the repo, prior memory, and the engineering handbook. Be rigorous: only " + "report high confidence when the goal, scope, and constraints are " + "genuinely unambiguous." +) + + +def _turn_cache_key( + qa_history: Sequence[object], state: PipelineState +) -> tuple[str, int, str]: + """Build the task-scoped memo key for one clarifier turn. + + Binds the ``thread_id`` (task isolation), the history length (turn index), + and a content hash of the Q&A so far. The thread id is the load-bearing + part: one :class:`ClaudeClarifier` instance is shared by the long-lived + graph node across every task, so keying on length alone would let one + task's cached confidence satisfy another task's gate with no model call. + The content hash is belt-and-suspenders so an in-place edit of the same- + length history (should one ever occur) also invalidates the memo. + """ + thread_id = str(state.get("thread_id", "") if isinstance(state, dict) else "") + try: + digest_src = json.dumps(list(qa_history), sort_keys=True, default=repr) + except (TypeError, ValueError): + digest_src = repr(list(qa_history)) + content_hash = hashlib.sha1(digest_src.encode("utf-8")).hexdigest() + return (thread_id, len(qa_history), content_hash) + + +class ClaudeClarifier: + """One Claude call per turn, serving both clarifier callables (§3.3, §7.1 P1). + + Construct with an optional ``invoke`` callable (defaults to + :func:`agent_team.billing.claude_invoke`) so tests inject a fake and the + real wiring goes through the billing seam. ``model`` / ``config`` are passed + through to the invoker, and ``system`` overrides the prompt framing. + + The single call per turn is memoized on a task-scoped key + (``thread_id`` + history length + content hash, see :func:`_turn_cache_key`): + calling :meth:`assess_confidence` then :meth:`generate_questions` for the + same turn of the same task reuses one Claude call; appending an answer (the + history grows) or a different task entering the shared node invalidates the + memo and the next assess triggers a fresh call. The thread-scoping is what + stops one task's cached confidence from clearing another task's human gate. + + :meth:`assess_confidence` and :meth:`generate_questions` are bound methods + that match :data:`~agent_team.nodes.clarifier.ConfidenceAssessor` and + :data:`~agent_team.nodes.clarifier.QuestionGenerator` exactly, so they wire + straight into :func:`~agent_team.nodes.clarifier.make_clarifier_node`. + """ + + def __init__( + self, + *, + invoke: ClaudeInvoke | None = None, + model: str | None = None, + config: Any = None, + system: str = _DEFAULT_SYSTEM, + confidence_threshold: float = DEFAULT_CONFIDENCE_THRESHOLD, + ) -> None: + self._invoke: ClaudeInvoke = invoke if invoke is not None else claude_invoke + self._model = model + self._config = config + self._system = system + self._confidence_threshold = confidence_threshold + # Memo of the single per-turn call. The key is task-scoped, NOT just the + # history length: one ClaudeClarifier instance serves every task through + # the long-lived graph node, so a key of len(qa_history) alone would let + # one task's cached high confidence clear ANOTHER task's human gate with + # no Claude call (a fail-OPEN cross-task collision). The key therefore + # binds (thread_id, history-length, content-hash) so the memo isolates + # per task/thread and still recomputes when the Q&A changes. + self._cache_key: tuple[str, int, str] | None = None + self._cache: dict[str, Any] | None = None + + # ------------------------------------------------------------------ # + # Public callables — exact ConfidenceAssessor / QuestionGenerator types. + # ------------------------------------------------------------------ # + + def assess_confidence( + self, qa_history: Sequence[object], state: PipelineState + ) -> float: + """Return the current 0..1 confidence the requirement is understood. + + Matches :data:`~agent_team.nodes.clarifier.ConfidenceAssessor`. Serves + the memoized per-turn Claude call; fails SAFE to ``0.0`` on any parse + trouble so a garbled response never clears the human gate. + """ + return float(self._turn(qa_history, state)["confidence"]) + + def generate_questions( + self, qa_history: Sequence[object], state: PipelineState + ) -> list[str]: + """Return the next ordered question-set. + + Matches :data:`~agent_team.nodes.clarifier.QuestionGenerator`. Reuses + the same memoized call as :meth:`assess_confidence` for this turn, and + always returns a non-empty list (the loop must have something to ask). + """ + return list(self._turn(qa_history, state)["questions"]) + + # ------------------------------------------------------------------ # + # Internals: the single per-turn call + memo. + # ------------------------------------------------------------------ # + + def _turn( + self, qa_history: Sequence[object], state: PipelineState + ) -> dict[str, Any]: + """Return the parsed result for this turn, making at most one Claude call. + + Memoized on ``(thread_id, len(qa_history), content-hash)``: the + assess/generate pair within one turn of one task shares a call; once an + answer is appended (history grows) or a different task/thread enters the + shared node, the key changes and a fresh call is made. Keying on the + thread id is what prevents one task's cached confidence from clearing + another task's human gate (the fail-OPEN collision the review caught). + """ + key = _turn_cache_key(qa_history, state) + if self._cache_key == key and self._cache is not None: + return self._cache + + prompt = self._build_prompt(qa_history, state) + result = self._invoke(prompt, model=self._model, config=self._config) + parsed = self._parse(getattr(result, "text", "")) + + self._cache_key = key + self._cache = parsed + return parsed + + def _build_prompt(self, qa_history: Sequence[object], state: PipelineState) -> str: + """Assemble the clarifier prompt from the Q&A history and task state. + + Pure string assembly over the graph state (§3.3) — no I/O — so the + prompt shape is directly unit-testable. + """ + description = _task_description(state) + repo = _state_field(state, "repo") + context = _state_field(state, "context") + qa = _format_qa_history(qa_history) + threshold_pct = int(round(self._confidence_threshold * 100)) + + sections: list[str] = [ + self._system, + "", + "## Task", + description or "(no task description provided)", + ] + if repo: + sections += ["", "## Repository", repo] + if context: + sections += ["", "## Additional context", context] + sections += [ + "", + "## Clarifier Q&A so far (oldest first)", + qa or "(no questions answered yet)", + "", + "## Your job", + ( + f"Decide whether you are at least {threshold_pct}% confident the " + "requirement is understood well enough to plan. If you are NOT, " + "produce the next ordered set of clarifying questions to ask the " + "human. Ask only what is genuinely needed; order them most " + "important first." + ), + "", + "## Output format", + ( + "Respond with ONLY a strict JSON object and no prose outside it, " + 'with keys: "confidence" (a float in [0, 1]), "questions" (a list ' + "of strings; empty only when you are confident enough to plan), " + 'and "rationale" (a short string). Example: ' + '{"confidence": 0.42, "questions": ["..."], "rationale": "..."}' + ), + ] + return "\n".join(sections) + + def _parse(self, text: str) -> dict[str, Any]: + """Parse the UNTRUSTED model reply into ``{confidence, questions, rationale}``. + + Fails SAFE at every step (§3.3 human gate): + + * confidence missing/unparseable -> ``0.0`` (keep asking, never clear + the gate on a garbled reply); + * confidence out of range -> clamped into ``[0, 1]``; + * questions missing/empty while below threshold -> a single generic + fallback question so the loop always has something to ask. + + A parse error is swallowed into the fail-safe default rather than + raised, so a bad reply degrades to "ask again", never to "advance". + """ + data = _extract_json_object(text) + + confidence = _coerce_confidence(data.get("confidence") if data else None) + questions = _coerce_questions(data.get("questions") if data else None) + rationale = "" + if data is not None: + raw_rationale = data.get("rationale") + if isinstance(raw_rationale, str): + rationale = raw_rationale.strip() + + if not questions and confidence < self._confidence_threshold: + # Below the bar but no usable question-set: substitute a generic + # clarifier so the loop still asks rather than spinning or advancing. + questions = [FALLBACK_QUESTION] + + return { + "confidence": confidence, + "questions": questions, + "rationale": rationale, + } + + +def build_claude_clarifier_callables( + *, + invoke: ClaudeInvoke | None = None, + model: str | None = None, + config: Any = None, + system: str = _DEFAULT_SYSTEM, + confidence_threshold: float = DEFAULT_CONFIDENCE_THRESHOLD, +) -> tuple[ConfidenceAssessor, QuestionGenerator]: + """Build the ``(assess_confidence, generate_questions)`` pair for wiring. + + Returns the two bound methods of a single shared :class:`ClaudeClarifier`, + ready to hand straight to + :func:`~agent_team.nodes.clarifier.make_clarifier_node`. Because both + callables share one instance, they share the per-turn memo, so the loop + makes one Claude call per turn rather than two. + """ + clarifier = ClaudeClarifier( + invoke=invoke, + model=model, + config=config, + system=system, + confidence_threshold=confidence_threshold, + ) + return clarifier.assess_confidence, clarifier.generate_questions + + +# --------------------------------------------------------------------------- # +# Module-level helpers (pure; no I/O). +# --------------------------------------------------------------------------- # + + +def _state_field(state: PipelineState, key: str) -> str: + """Pull a string field from the (untyped-extra) graph state, defensively.""" + value = state.get(key) # type: ignore[call-overload] + if isinstance(value, str) and value.strip(): + return value.strip() + return "" + + +def _task_description(state: PipelineState) -> str: + """Pull the task description out of the graph state (mirrors planner.py). + + Looks in the conventional places (the ``plan`` dict, then a top-level + ``task``/``description`` key) and falls back to an empty string so a + malformed state surfaces as an empty prompt section, never a ``KeyError``. + """ + 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() + for key in ("task", "description"): + desc = _state_field(state, key) + if desc: + return desc + return "" + + +def _format_qa_history(qa_history: Sequence[object]) -> str: + """Render the clarifier Q&A history (oldest first) into prompt text. + + Each entry may be a ``{"question": ..., "answer": ...}`` mapping or a plain + string (the raw resume value the loop appends); both are handled so this + does not couple to a single record shape. + """ + 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) + + +# A fenced ```json ... ``` block, if the model wrapped its JSON in Markdown. +_FENCE_RE = re.compile( + r"```(?:json)?\s*\n?(?P.*?)\n?\s*```", + flags=re.DOTALL | re.IGNORECASE, +) + + +def _extract_json_object(text: str) -> dict[str, Any] | None: + """Extract a JSON object from UNTRUSTED model output, or ``None``. + + Tolerates the common ways a model deviates from "JSON only": a leading + apology or trailing prose, and ```json fences. Tries, in order, the whole + string, the contents of a fenced block, then the first ``{...}`` span found + by brace matching. Returns ``None`` (never raises) when nothing parses to a + JSON object, so the caller can fail SAFE. + """ + if not isinstance(text, str) or not text.strip(): + return None + + candidates: list[str] = [text.strip()] + + fence = _FENCE_RE.search(text) + if fence: + candidates.append(fence.group("body").strip()) + + span = _first_brace_span(text) + if span is not None: + candidates.append(span) + + for candidate in candidates: + if not candidate: + continue + try: + parsed = json.loads(candidate) + except (json.JSONDecodeError, ValueError): + continue + if isinstance(parsed, dict): + return parsed + return None + + +def _first_brace_span(text: str) -> str | None: + """Return the first balanced ``{...}`` span in ``text`` (string-aware).""" + start = text.find("{") + if start == -1: + return None + depth = 0 + in_string = False + escaped = False + for idx in range(start, len(text)): + ch = text[idx] + if in_string: + if escaped: + escaped = False + elif ch == "\\": + escaped = True + elif ch == '"': + in_string = False + continue + if ch == '"': + in_string = True + elif ch == "{": + depth += 1 + elif ch == "}": + depth -= 1 + if depth == 0: + return text[start : idx + 1] + return None + + +def _coerce_confidence(value: Any) -> float: + """Coerce the model's confidence into a clamped ``[0, 1]`` float. + + Missing or unparseable -> ``0.0`` (fail SAFE: keep asking, never clear the + gate). Out-of-range values are clamped rather than rejected. + """ + try: + confidence = float(value) + except (TypeError, ValueError): + return 0.0 + if confidence != confidence: # NaN guard + return 0.0 + if confidence < 0.0: + return 0.0 + if confidence > 1.0: + return 1.0 + return confidence + + +def _coerce_questions(value: Any) -> list[str]: + """Coerce the model's question-set into a clean list of non-empty strings. + + Anything that is not a list of usable strings collapses to an empty list, + which the parser then fills with the generic fallback when below threshold. + """ + if not isinstance(value, list): + return [] + questions: list[str] = [] + for item in value: + if isinstance(item, str): + text = item.strip() + if text: + questions.append(text) + return questions diff --git a/agent-team/tests/test_clarifier_llm.py b/agent-team/tests/test_clarifier_llm.py new file mode 100644 index 0000000..1448709 --- /dev/null +++ b/agent-team/tests/test_clarifier_llm.py @@ -0,0 +1,335 @@ +"""Unit tests for agent_team.nodes.clarifier_llm (§3.3, §7.1 P1). + +The Claude-backed clarifier callables are exercised with a FAKE invoke that +returns canned :class:`~agent_team.billing.ClaudeResult` text — no network. The +load-bearing properties under test: + +* **One call per turn (memoization).** ``assess_confidence`` then + ``generate_questions`` on the same turn must reuse a single Claude call. +* **Fail SAFE (the human gate).** A garbled / non-JSON reply must yield + confidence ``0.0`` (never >= the 0.98 bar) and a non-empty fallback question. +* **Defensive parsing.** ```json fences and surrounding prose still parse, and + out-of-range confidence is clamped to ``[0, 1]``. +* **Integration smoke.** The callables wire into the real + :func:`~agent_team.nodes.clarifier.make_clarifier_node` and clear the gate + once confidence rises across turns. +""" + +from __future__ import annotations + +import json +from typing import Any + +from langgraph.checkpoint.memory import MemorySaver +from langgraph.graph import END, START, StateGraph +from langgraph.types import Command + +from agent_team.billing import BillingMode, ClaudeResult +from agent_team.nodes.clarifier import ( + DEFAULT_CONFIDENCE_THRESHOLD, + make_clarifier_node, +) +from agent_team.nodes.clarifier_llm import ( + FALLBACK_QUESTION, + ClaudeClarifier, + build_claude_clarifier_callables, +) +from agent_team.task_model import Phase, PipelineState, TaskStatus + +# --------------------------------------------------------------------------- # +# Fakes / helpers +# --------------------------------------------------------------------------- # + + +class _FakeInvoke: + """A fake billing.claude_invoke that returns canned text and counts calls. + + ``replies`` may be a single string (returned every call) or a list of + strings (consumed one per call, last one repeating) so a test can simulate + rising confidence across turns. + """ + + def __init__(self, replies: str | list[str]) -> None: + self._replies = [replies] if isinstance(replies, str) else list(replies) + self.calls: list[dict[str, Any]] = [] + + def __call__(self, prompt: str, **kw: Any) -> ClaudeResult: + idx = min(len(self.calls), len(self._replies) - 1) + text = self._replies[idx] + self.calls.append({"prompt": prompt, "kw": kw}) + return ClaudeResult(text=text, mode=BillingMode.SUBSCRIPTION) + + +def _state(**overrides: Any) -> PipelineState: + base: PipelineState = PipelineState( + thread_id="t-1", + status=TaskStatus.ACTIVE.value, + current_phase=Phase.CLARIFY.value, + qa_history=[], + transport="slack", + ) + base.update(overrides) # type: ignore[typeddict-item] + return base + + +def _json(confidence: Any, questions: Any, rationale: str = "because") -> str: + return json.dumps( + {"confidence": confidence, "questions": questions, "rationale": rationale} + ) + + +# --------------------------------------------------------------------------- # +# High confidence: assess returns ~value AND the call is reused (memoization). +# --------------------------------------------------------------------------- # + + +def test_high_confidence_parsed() -> None: + fake = _FakeInvoke(_json(0.99, [])) + clar = ClaudeClarifier(invoke=fake) + assert clar.assess_confidence([], _state()) == 0.99 + + +def test_single_call_per_turn_memoized() -> None: + fake = _FakeInvoke(_json(0.99, [])) + clar = ClaudeClarifier(invoke=fake) + qa: list[object] = [] + + # Both methods called for the same turn -> exactly ONE Claude call. + conf = clar.assess_confidence(qa, _state()) + questions = clar.generate_questions(qa, _state()) + + assert conf == 0.99 + assert questions == [] # confident, no questions needed + assert len(fake.calls) == 1 + + +def test_memo_recomputes_when_history_grows() -> None: + fake = _FakeInvoke([_json(0.10, ["q1"]), _json(0.99, [])]) + clar = ClaudeClarifier(invoke=fake) + + # Turn 0: one answer-less call. + assert clar.assess_confidence([], _state()) == 0.10 + assert clar.generate_questions([], _state()) == ["q1"] + assert len(fake.calls) == 1 + + # Turn 1: history grew -> a fresh call, now confident. + assert clar.assess_confidence(["a1"], _state()) == 0.99 + assert clar.generate_questions(["a1"], _state()) == [] + assert len(fake.calls) == 2 + + +def test_memo_isolates_across_tasks_no_cross_gate_clear() -> None: + """A second task at the same history length must NOT reuse task A's memo. + + Regression for the fail-OPEN collision: one ClaudeClarifier instance serves + every task through the shared graph node, so keying the memo on history + length alone would let task A's cached 0.99 clear task B's human gate with + no model call. Keying on thread_id forces a fresh assessment for task B. + """ + fake = _FakeInvoke([_json(0.99, []), _json(0.10, ["need more from B"])]) + clar = ClaudeClarifier(invoke=fake) + + # Task A (thread t-A), empty history -> confident, cached. + assert clar.assess_confidence([], _state(thread_id="t-A")) == 0.99 + assert len(fake.calls) == 1 + + # Task B (thread t-B), SAME empty history/length -> must re-assess, NOT + # inherit A's cache, so its low confidence holds and the gate stays shut. + assert clar.assess_confidence([], _state(thread_id="t-B")) == 0.10 + assert clar.generate_questions([], _state(thread_id="t-B")) == ["need more from B"] + assert len(fake.calls) == 2 # a real second call happened for task B + + +# --------------------------------------------------------------------------- # +# Low confidence: below the bar, questions are returned. +# --------------------------------------------------------------------------- # + + +def test_low_confidence_returns_questions() -> None: + fake = _FakeInvoke(_json(0.40, ["What is the scope?", "Which repo?"])) + clar = ClaudeClarifier(invoke=fake) + + assert clar.assess_confidence([], _state()) < DEFAULT_CONFIDENCE_THRESHOLD + assert clar.generate_questions([], _state()) == [ + "What is the scope?", + "Which repo?", + ] + assert len(fake.calls) == 1 + + +# --------------------------------------------------------------------------- # +# Malformed output: fail SAFE (0.0 confidence, non-empty fallback questions). +# --------------------------------------------------------------------------- # + + +def test_malformed_output_fails_safe() -> None: + fake = _FakeInvoke("I'm sorry, I cannot help with that. ") + clar = ClaudeClarifier(invoke=fake) + + assert clar.assess_confidence([], _state()) == 0.0 + questions = clar.generate_questions([], _state()) + assert questions == [FALLBACK_QUESTION] + assert questions # non-empty + + +def test_garbage_never_clears_the_gate() -> None: + # The critical safety property: garbage must never read >= 0.98. + for garbage in ["", " ", "not json", "{broken", "[1,2,3]", "null", "42"]: + clar = ClaudeClarifier(invoke=_FakeInvoke(garbage)) + conf = clar.assess_confidence([], _state()) + assert conf < DEFAULT_CONFIDENCE_THRESHOLD + assert conf == 0.0 + assert clar.generate_questions([], _state()) == [FALLBACK_QUESTION] + + +def test_missing_confidence_key_defaults_zero() -> None: + fake = _FakeInvoke(json.dumps({"questions": ["q?"], "rationale": "x"})) + clar = ClaudeClarifier(invoke=fake) + assert clar.assess_confidence([], _state()) == 0.0 + # Questions present in the reply are kept as-is. + assert clar.generate_questions([], _state()) == ["q?"] + + +def test_low_confidence_empty_questions_gets_fallback() -> None: + # Below the bar but model gave no questions -> generic fallback so the loop + # always has something to ask. + fake = _FakeInvoke(_json(0.20, [])) + clar = ClaudeClarifier(invoke=fake) + assert clar.generate_questions([], _state()) == [FALLBACK_QUESTION] + + +# --------------------------------------------------------------------------- # +# Defensive parsing: fences and surrounding prose still parse. +# --------------------------------------------------------------------------- # + + +def test_json_in_code_fence_is_parsed() -> None: + fenced = "```json\n" + _json(0.97, ["q?"]) + "\n```" + clar = ClaudeClarifier(invoke=_FakeInvoke(fenced)) + assert clar.assess_confidence([], _state()) == 0.97 + assert clar.generate_questions([], _state()) == ["q?"] + + +def test_json_wrapped_in_prose_is_parsed() -> None: + prose = ( + "Sure! Here is my assessment:\n" + + _json(0.55, ["Clarify the deadline?"]) + + "\nLet me know if that helps." + ) + clar = ClaudeClarifier(invoke=_FakeInvoke(prose)) + assert clar.assess_confidence([], _state()) == 0.55 + assert clar.generate_questions([], _state()) == ["Clarify the deadline?"] + + +def test_plain_json_fence_without_lang_is_parsed() -> None: + fenced = "```\n" + _json(0.33, ["q?"]) + "\n```" + clar = ClaudeClarifier(invoke=_FakeInvoke(fenced)) + assert clar.assess_confidence([], _state()) == 0.33 + + +# --------------------------------------------------------------------------- # +# Confidence clamping into [0, 1]. +# --------------------------------------------------------------------------- # + + +def test_confidence_above_one_is_clamped() -> None: + clar = ClaudeClarifier(invoke=_FakeInvoke(_json(1.5, []))) + assert clar.assess_confidence([], _state()) == 1.0 + + +def test_confidence_below_zero_is_clamped() -> None: + clar = ClaudeClarifier(invoke=_FakeInvoke(_json(-0.2, ["q?"]))) + assert clar.assess_confidence([], _state()) == 0.0 + + +def test_confidence_as_string_is_coerced() -> None: + clar = ClaudeClarifier(invoke=_FakeInvoke(_json("0.88", ["q?"]))) + assert clar.assess_confidence([], _state()) == 0.88 + + +# --------------------------------------------------------------------------- # +# Prompt assembly pulls task/repo/context out of state. +# --------------------------------------------------------------------------- # + + +def test_prompt_includes_task_repo_and_qa() -> None: + fake = _FakeInvoke(_json(0.99, [])) + clar = ClaudeClarifier(invoke=fake) + state = _state(task="Add a webhook verifier", repo="agent-team") + clar.assess_confidence(["prior answer"], state) + + prompt = fake.calls[0]["prompt"] + assert "Add a webhook verifier" in prompt + assert "agent-team" in prompt + assert "prior answer" in prompt + + +# --------------------------------------------------------------------------- # +# Factory returns the exact ConfidenceAssessor / QuestionGenerator pair. +# --------------------------------------------------------------------------- # + + +def test_factory_returns_shared_memoized_pair() -> None: + fake = _FakeInvoke(_json(0.45, ["q?"])) + assess, generate = build_claude_clarifier_callables(invoke=fake) + + # Both come from one shared instance -> one call serves both this turn. + assert assess([], _state()) == 0.45 + assert generate([], _state()) == ["q?"] + assert len(fake.calls) == 1 + + +# --------------------------------------------------------------------------- # +# Integration smoke: wire into the real make_clarifier_node, gate clears. +# --------------------------------------------------------------------------- # + + +def _build_app(node): + graph = StateGraph(PipelineState) + graph.add_node("clarify", node) + graph.add_edge(START, "clarify") + graph.add_edge("clarify", END) + return graph.compile(checkpointer=MemorySaver()) + + +def test_node_clears_gate_when_confidence_rises() -> None: + # Turn 0 (no answers): low confidence, asks. Turn 1 (one answer): confident. + fake = _FakeInvoke([_json(0.20, ["What is the goal?"]), _json(0.99, [])]) + assess, generate = build_claude_clarifier_callables(invoke=fake) + node = make_clarifier_node(assess_confidence=assess, generate_questions=generate) + app = _build_app(node) + cfg = {"configurable": {"thread_id": "t-1"}} + + first = app.invoke(_state(thread_id="t-1"), 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_node_parks_when_garbage_never_clears_gate() -> None: + # A model that only ever emits garbage must NEVER open the gate; the loop + # asks until the turn cap and parks (human gate stays shut). + fake = _FakeInvoke("garbage, no json") + assess, generate = build_claude_clarifier_callables(invoke=fake) + from agent_team.nodes.clarifier import ClarifierConfig + + node = make_clarifier_node( + assess_confidence=assess, + generate_questions=generate, + config=ClarifierConfig(max_turns=2), + ) + app = _build_app(node) + cfg = {"configurable": {"thread_id": "t-1"}} + + assert "__interrupt__" in app.invoke(_state(thread_id="t-1"), cfg) + assert "__interrupt__" in app.invoke(Command(resume="a1"), cfg) + final = app.invoke(Command(resume="a2"), cfg) + + assert "__interrupt__" not in final + assert final["current_phase"] == Phase.PARKED.value + assert final["status"] == TaskStatus.PARKED.value + assert final["current_phase"] != Phase.PLAN.value diff --git a/agent-team/tests/test_invoker.py b/agent-team/tests/test_invoker.py new file mode 100644 index 0000000..b1873ce --- /dev/null +++ b/agent-team/tests/test_invoker.py @@ -0,0 +1,255 @@ +"""Unit tests for agent_team.invoker (§3.1) — all mocked, no network/SDK. + +These tests prove the module imports cleanly without ``claude_agent_sdk`` or +``anthropic`` installed, and exercise each billing path through the injected +SDK seams (``_query``/``_options_cls``/``_client``) so nothing real is called. +""" + +from __future__ import annotations + +import pytest + +from agent_team import billing, invoker +from agent_team.billing import BillingMode, ClaudeResult + + +@pytest.fixture(autouse=True) +def _restore_invoker(): + """Restore the module-global billing invoker after each test.""" + original = billing._invoker + yield + billing._invoker = original + + +# --------------------------------------------------------------------------- # +# Fakes mirroring the Agent SDK / anthropic message shapes +# --------------------------------------------------------------------------- # + + +class _FakeTextBlock: + def __init__(self, text: str) -> None: + self.text = text + + +# Class names mirror the real Agent SDK message types — the invoker dispatches +# on ``type(msg).__name__``, so these MUST be named AssistantMessage / +# ResultMessage to be recognised. +class AssistantMessage: + def __init__(self, text: str) -> None: + self.content = [_FakeTextBlock(text)] + + +class ResultMessage: + def __init__(self, result: str, cost: float = 0.42) -> None: + self.result = result + self.total_cost_usd = cost + self.usage = {"input_tokens": 11, "output_tokens": 7} + + +def _fake_options(**kwargs): + """Stand-in for ClaudeAgentOptions: just record the kwargs.""" + return dict(kwargs) + + +def _make_fake_query(messages): + """Build an async ``query(prompt=..., options=...)`` yielding ``messages``.""" + + async def _query(*, prompt, options): + for msg in messages: + yield msg + + return _query + + +class _FakeAnthropicUsage: + def __init__(self) -> None: + self.input_tokens = 12 + self.output_tokens = 5 + + def model_dump(self) -> dict: + return {"input_tokens": self.input_tokens, "output_tokens": self.output_tokens} + + +class _FakeAnthropicMessage: + def __init__(self, text: str) -> None: + self.content = [_FakeTextBlock(text)] + self.usage = _FakeAnthropicUsage() + + +class _FakeMessages: + def __init__(self, text: str) -> None: + self._text = text + self.calls: list[dict] = [] + + def create(self, **kwargs): + self.calls.append(kwargs) + return _FakeAnthropicMessage(self._text) + + +class _FakeAnthropicClient: + def __init__(self, text: str) -> None: + self.messages = _FakeMessages(text) + + +# --------------------------------------------------------------------------- # +# Clean import without the SDKs +# --------------------------------------------------------------------------- # + + +def test_module_imports_without_sdks() -> None: + import importlib + import sys + + # claude_agent_sdk / anthropic are not installed in this env. + assert "claude_agent_sdk" not in sys.modules + mod = importlib.import_module("agent_team.invoker") + assert hasattr(mod, "subscription_invoker") + assert hasattr(mod, "api_invoker") + + +# --------------------------------------------------------------------------- # +# Subscription path +# --------------------------------------------------------------------------- # + + +def test_subscription_invoker_returns_result( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("CLAUDE_CODE_OAUTH_TOKEN", "oauth-tok") + messages = [ + AssistantMessage("partial "), + ResultMessage("final answer"), + ] + result = invoker.subscription_invoker( + "hello", + mode=BillingMode.SUBSCRIPTION, + _query=_make_fake_query(messages), + _options_cls=_fake_options, + ) + assert isinstance(result, ClaudeResult) + assert result.text == "final answer" + assert result.mode is BillingMode.SUBSCRIPTION + # usage populated from ResultMessage cost + usage dict. + assert result.usage["total_cost_usd"] == pytest.approx(0.42) + assert result.usage["input_tokens"] == 11 + # raw carries the collected message stream. + assert result.raw == messages + + +def test_subscription_invoker_falls_back_to_assistant_text( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("CLAUDE_CODE_OAUTH_TOKEN", "oauth-tok") + messages = [AssistantMessage("a"), AssistantMessage("b")] + result = invoker.subscription_invoker( + "hi", + mode=BillingMode.SUBSCRIPTION, + _query=_make_fake_query(messages), + _options_cls=_fake_options, + ) + assert result.text == "a\nb" + + +def test_subscription_invoker_requires_oauth_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) + with pytest.raises(RuntimeError, match="secrev.env"): + invoker.subscription_invoker( + "hi", + mode=BillingMode.SUBSCRIPTION, + _query=_make_fake_query([]), + _options_cls=_fake_options, + ) + + +# --------------------------------------------------------------------------- # +# API path +# --------------------------------------------------------------------------- # + + +def test_api_invoker_returns_result() -> None: + client = _FakeAnthropicClient("api text") + result = invoker.api_invoker("ask", mode=BillingMode.API, _client=client) + assert result.text == "api text" + assert result.mode is BillingMode.API + assert result.usage == {"input_tokens": 12, "output_tokens": 5} + # The model + prompt were threaded into the SDK call. + assert client.messages.calls[0]["model"] == invoker.API_MODEL + assert client.messages.calls[0]["messages"] == [{"role": "user", "content": "ask"}] + + +# --------------------------------------------------------------------------- # +# Bedrock path +# --------------------------------------------------------------------------- # + + +def test_bedrock_raises_not_implemented() -> None: + with pytest.raises(NotImplementedError, match="BEDROCK"): + invoker.real_invoker("hi", mode=BillingMode.BEDROCK) + + +# --------------------------------------------------------------------------- # +# real_invoker dispatch — SUBSCRIPTION + API branches forward kwargs to the leaf +# --------------------------------------------------------------------------- # + + +def test_real_invoker_dispatches_subscription_branch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """``real_invoker`` routes SUBSCRIPTION to ``subscription_invoker`` (kwargs fwd).""" + monkeypatch.setenv("CLAUDE_CODE_OAUTH_TOKEN", "oauth-tok") + messages = [AssistantMessage("partial "), ResultMessage("sub final")] + result = invoker.real_invoker( + "hello", + mode=BillingMode.SUBSCRIPTION, + _query=_make_fake_query(messages), + _options_cls=_fake_options, + ) + assert isinstance(result, ClaudeResult) + assert result.mode is BillingMode.SUBSCRIPTION + assert result.text == "sub final" # the ResultMessage leaf ran + assert result.raw == messages + + +def test_real_invoker_dispatches_api_branch() -> None: + """``real_invoker`` routes API to ``api_invoker`` with the injected client.""" + client = _FakeAnthropicClient("api branch text") + result = invoker.real_invoker("ask", mode=BillingMode.API, _client=client) + assert result.mode is BillingMode.API + assert result.text == "api branch text" # the anthropic-client leaf ran + # The prompt was threaded through to the injected client. + assert client.messages.calls[0]["messages"] == [{"role": "user", "content": "ask"}] + + +# --------------------------------------------------------------------------- # +# Binding into the billing seam +# --------------------------------------------------------------------------- # + + +def test_bind_subscription_invoker_sets_billing_invoker() -> None: + invoker.bind_subscription_invoker() + assert billing._invoker is invoker.real_invoker + + +def test_bound_invoker_drives_claude_invoke( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Bind a fake invoker so claude_invoke routes through it end to end. + captured: dict = {} + + def fake(prompt: str, *, mode: BillingMode, **kw): + captured["prompt"] = prompt + captured["mode"] = mode + return ClaudeResult(text="routed", mode=mode) + + invoker.bind_invoker(fake) + result = billing.claude_invoke("q", mode=BillingMode.API) + assert result.text == "routed" + assert captured["mode"] is BillingMode.API + assert captured["prompt"] == "q" + + +def test_bind_invoker_defaults_to_real_invoker() -> None: + invoker.bind_invoker() + assert billing._invoker is invoker.real_invoker