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).
This commit is contained in:
parent
b853598579
commit
270ce93b2a
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