agent-team Plane-2: bind P1+P2 to real models, live transport, coordinator #12
4 changed files with 1356 additions and 0 deletions
300
agent-team/agent_team/invoker.py
Normal file
300
agent-team/agent_team/invoker.py
Normal file
|
|
@ -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)
|
||||
466
agent-team/agent_team/nodes/clarifier_llm.py
Normal file
466
agent-team/agent_team/nodes/clarifier_llm.py
Normal file
|
|
@ -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<body>.*?)\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
|
||||
335
agent-team/tests/test_clarifier_llm.py
Normal file
335
agent-team/tests/test_clarifier_llm.py
Normal file
|
|
@ -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. <no json here>")
|
||||
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
|
||||
255
agent-team/tests/test_invoker.py
Normal file
255
agent-team/tests/test_invoker.py
Normal file
|
|
@ -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
|
||||
Reference in a new issue