From 0c0dbd1659a2f44832bc248f0a8cb42814ce5da7 Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:19:10 +0800 Subject: [PATCH 1/9] feat: add portable metering and tool guardrails --- agent_core/runtime/loop/guardrails.py | 189 ++++++++++++++++++ agent_core/runtime/loop/tool_call_repair.py | 173 ++++++++++++++++ agent_core/runtime/usage_meter.py | 210 ++++++++++++++++++++ tests/test_tool_guardrails.py | 77 +++++++ tests/test_usage_meter.py | 90 +++++++++ 5 files changed, 739 insertions(+) create mode 100644 agent_core/runtime/loop/guardrails.py create mode 100644 agent_core/runtime/loop/tool_call_repair.py create mode 100644 agent_core/runtime/usage_meter.py create mode 100644 tests/test_tool_guardrails.py create mode 100644 tests/test_usage_meter.py diff --git a/agent_core/runtime/loop/guardrails.py b/agent_core/runtime/loop/guardrails.py new file mode 100644 index 0000000..f559b17 --- /dev/null +++ b/agent_core/runtime/loop/guardrails.py @@ -0,0 +1,189 @@ +"""Portable middleware for duplicate-call and repeated-loop guardrails.""" + +from __future__ import annotations + +import hashlib +import json +import logging +from collections import deque +from collections.abc import Mapping +from typing import Any, cast + +from agent_core.protocols import ExecutionMiddleware, ToolCallContext + +logger = logging.getLogger(__name__) + +DEFAULT_DUPLICATE_THRESHOLDS: dict[str, int] = { + "web_search": 6, + "web_fetch": 5, + "bash": 6, + "file_editor_str_replace": 5, + "file_editor_view": 5, + "file_editor_create": 5, + "read_file": 5, + "read_text": 5, + "write_file": 5, + "grep_search": 5, + "glob_search": 5, + "view_image": 4, + "delegate_subtask": 3, + "collect_results": 3, + "abort_task": 3, +} + + +def _args_fingerprint(tool_name: str, args: Mapping[str, Any]) -> str: + raw = json.dumps({"t": tool_name, "a": args}, sort_keys=True) + return hashlib.md5(raw.encode()).hexdigest()[:12] + + +class GuardrailsMiddleware(ExecutionMiddleware): + """Block pathological repetition while allowing hosts to tune tool policy.""" + + def __init__( + self, + search_warn_threshold: int = 50, + max_loop_hints: int = 3, + *, + duplicate_thresholds: Mapping[str, int] | None = None, + default_duplicate_threshold: int = 5, + search_tool_name: str = "web_search", + ignored_tools: frozenset[str] = frozenset({"tool_search"}), + ) -> None: + self._search_warn_threshold = search_warn_threshold + self._max_loop_hints = max_loop_hints + self._duplicate_thresholds = dict( + duplicate_thresholds or DEFAULT_DUPLICATE_THRESHOLDS, + ) + self._default_duplicate_threshold = default_duplicate_threshold + self._search_tool_name = search_tool_name + self._ignored_tools = ignored_tools + self._recent_calls: dict[str, deque[str]] = {} + self._search_counts: dict[str, int] = {} + self._search_warned: set[str] = set() + self._loop_hint_counts: dict[str, int] = {} + self._stats = { + "duplicate_blocks": 0, + "search_warnings": 0, + "loop_escalation_blocks": 0, + "total_checks": 0, + } + + @property + def stats(self) -> dict[str, int]: + return dict(self._stats) + + def notify_loop_hint(self, task_id: str) -> None: + count = self._loop_hint_counts.get(task_id, 0) + 1 + self._loop_hint_counts[task_id] = count + logger.info( + "Guardrails loop hint #%d for task %s (hard block at %d)", + count, + task_id, + self._max_loop_hints, + ) + + def reset_loop_hints(self, task_id: str) -> None: + self._loop_hint_counts.pop(task_id, None) + + async def before_tool_call(self, ctx: ToolCallContext) -> ToolCallContext: + self._stats["total_checks"] += 1 + task_id = ctx.task_id + tool_name = ctx.tool_name + if tool_name in self._ignored_tools: + return ctx + + fingerprint = _args_fingerprint(tool_name, ctx.tool_args) + recent = self._recent_calls.setdefault(task_id, deque(maxlen=20)) + consecutive = 0 + for previous in reversed(recent): + if previous != fingerprint: + break + consecutive += 1 + threshold = self._duplicate_thresholds.get( + tool_name, + self._default_duplicate_threshold, + ) + if consecutive >= threshold: + self._stats["duplicate_blocks"] += 1 + ctx.metadata["blocked"] = True + ctx.metadata["block_reason"] = ( + f"Blocked: {tool_name} called {consecutive + 1} times " + "consecutively with identical arguments. Try different " + "parameters or a different approach." + ) + return ctx + recent.append(fingerprint) + + if tool_name == self._search_tool_name: + count = self._search_counts.get(task_id, 0) + 1 + self._search_counts[task_id] = count + if count >= self._search_warn_threshold and task_id not in self._search_warned: + self._search_warned.add(task_id) + self._stats["search_warnings"] += 1 + logger.warning( + "Search count %d reached warning threshold %d for task %s", + count, + self._search_warn_threshold, + task_id, + ) + + hint_count = self._loop_hint_counts.get(task_id, 0) + if hint_count >= self._max_loop_hints and fingerprint in list(recent)[:-1]: + self._stats["loop_escalation_blocks"] += 1 + ctx.metadata["blocked"] = True + ctx.metadata["block_reason"] = ( + "Blocked: repeated tool call pattern detected after " + f"{hint_count} loop warnings. Try a completely different " + "approach or conclude with available evidence." + ) + return ctx + + def cleanup_task(self, task_id: str) -> None: + self._recent_calls.pop(task_id, None) + self._search_counts.pop(task_id, None) + self._search_warned.discard(task_id) + self._loop_hint_counts.pop(task_id, None) + + +def check_budget_exhausted( + working_state: Mapping[str, Any], + token_usage: Mapping[str, int] | None, +) -> str | None: + """Return a finalization hint when the allocated token budget is spent.""" + empty: Mapping[str, Any] = {} + plan_value: object = working_state.get("execution_plan") + plan: Mapping[str, Any] = ( + cast("Mapping[str, Any]", plan_value) + if isinstance(plan_value, Mapping) + else empty + ) + budget_value: object = plan.get("budget") + budget: Mapping[str, Any] = ( + cast("Mapping[str, Any]", budget_value) + if isinstance(budget_value, Mapping) + else empty + ) + allocated_value: object = budget.get("allocated") + allocated: Mapping[str, Any] = ( + cast("Mapping[str, Any]", allocated_value) + if isinstance(allocated_value, Mapping) + else budget + ) + if token_usage: + maximum_value: object = allocated.get("max_tokens", 500_000) + maximum = int(maximum_value) if isinstance(maximum_value, int | str) else 500_000 + used = token_usage.get("total", 0) + if used >= maximum: + return ( + f"Token budget exhausted ({used:,}/{maximum:,}). " + "Provide final analysis with available evidence." + ) + return None + + +__all__ = [ + "DEFAULT_DUPLICATE_THRESHOLDS", + "GuardrailsMiddleware", + "check_budget_exhausted", +] diff --git a/agent_core/runtime/loop/tool_call_repair.py b/agent_core/runtime/loop/tool_call_repair.py new file mode 100644 index 0000000..159547c --- /dev/null +++ b/agent_core/runtime/loop/tool_call_repair.py @@ -0,0 +1,173 @@ +"""Configurable, best-effort repair of common LLM tool-call mistakes.""" + +from __future__ import annotations + +import json +import re +from collections.abc import Callable, Mapping +from typing import Any + +from agent_core.protocols import ExecutionMiddleware, ToolCallContext + +ToolRepair = Callable[[str, dict[str, Any]], dict[str, Any]] + +DEFAULT_KEY_ALIASES: dict[str, dict[str, str]] = { + "web_search": { + "search_query": "query", "q": "query", "text": "query", + "results": "num_results", "count": "num_results", "n": "num_results", + }, + "web_fetch": {"link": "url", "page": "url", "site": "url"}, + "read_file": {"filename": "path", "file": "path", "file_path": "path"}, + "read_text": {"filename": "path", "file": "path", "file_path": "path"}, + "write_file": { + "filename": "path", "file": "path", "file_path": "path", + "text": "content", "body": "content", + }, + "bash": {"cmd": "command", "shell": "command", "run": "command"}, + "grep_search": { + "regex": "pattern", "search": "pattern", "query": "pattern", + "directory": "path", "glob": "glob_filter", "lines": "context_lines", + }, + "glob_search": { + "glob": "pattern", "file_pattern": "pattern", "directory": "path", + }, + "delegate_subtask": {"tasks": "subtasks", "questions": "subtasks"}, + "file_editor_str_replace": { + "file": "path", "file_path": "path", "before": "old_string", + "old": "old_string", "after": "new_string", "new": "new_string", + }, + "file_editor_view": {"file": "path", "file_path": "path"}, + "file_editor_create": { + "file": "path", "file_path": "path", "file_text": "content", + }, + "view_image": {"file": "path", "file_path": "path", "image": "path"}, +} + +DEFAULT_TYPE_COERCIONS: dict[str, dict[str, type[Any]]] = { + "web_search": {"num_results": int}, + "grep_search": {"context_lines": int, "max_results": int}, +} + + +def repair_truncated_json(raw: str) -> str | None: + """Close simple truncated strings/brackets and return valid JSON text.""" + raw = raw.strip() + if not raw: + return None + try: + json.loads(raw) + return raw + except (json.JSONDecodeError, ValueError): + pass + text = re.sub(r",\s*([}\]])", r"\1", raw) + text = re.sub(r",\s*$", "", text) + try: + json.loads(text) + return text + except (json.JSONDecodeError, ValueError): + pass + in_string = False + for index, char in enumerate(text): + if char == '"' and (index == 0 or text[index - 1] != "\\"): + in_string = not in_string + if in_string: + text += '"' + text = re.sub(r",\s*$", "", text) + stack: list[str] = [] + in_string = False + for index, char in enumerate(text): + if char == '"' and (index == 0 or text[index - 1] != "\\"): + in_string = not in_string + elif not in_string and char == "{": + stack.append("}") + elif not in_string and char == "[": + stack.append("]") + elif not in_string and char in "}]" and stack and stack[-1] == char: + stack.pop() + text += "".join(reversed(stack)) + try: + json.loads(text) + return text + except (json.JSONDecodeError, ValueError): + return None + + +def _default_tool_repair(tool_name: str, args: dict[str, Any]) -> dict[str, Any]: + if tool_name == "web_search" and isinstance(args.get("query"), list): + args["query"] = " ".join(str(item) for item in args["query"]) + if tool_name == "web_fetch": + url = args.get("url") + if isinstance(url, str) and url.startswith("<") and url.endswith(">"): + args["url"] = url[1:-1] + if tool_name == "bash" and isinstance(args.get("command"), str): + match = re.match( + r"^```(?:bash|sh|shell)?\s*\n(.*?)```\s*$", + args["command"], + re.DOTALL, + ) + if match: + args["command"] = match.group(1).strip() + if tool_name == "delegate_subtask" and isinstance(args.get("subtasks"), str): + raw = args["subtasks"] + try: + parsed = json.loads(raw) + args["subtasks"] = parsed if isinstance(parsed, list) else raw + except (json.JSONDecodeError, ValueError): + args["subtasks"] = [{"question": raw}] + if tool_name == "grep_search" and isinstance(args.get("pattern"), list): + args["pattern"] = "|".join(str(item) for item in args["pattern"]) + return args + + +class ToolCallRepairMiddleware(ExecutionMiddleware): + """Normalize aliases, primitive types, whitespace, and host repairs.""" + + def __init__( + self, + *, + key_aliases: Mapping[str, Mapping[str, str]] | None = None, + type_coercions: Mapping[str, Mapping[str, type[Any]]] | None = None, + tool_repair: ToolRepair = _default_tool_repair, + ) -> None: + self._key_aliases = key_aliases or DEFAULT_KEY_ALIASES + self._type_coercions = type_coercions or DEFAULT_TYPE_COERCIONS + self._tool_repair = tool_repair + self._stats = { + "key_renames": 0, + "type_coercions": 0, + "whitespace_strips": 0, + "total_calls": 0, + } + + @property + def stats(self) -> dict[str, int]: + return dict(self._stats) + + async def before_tool_call(self, ctx: ToolCallContext) -> ToolCallContext: + self._stats["total_calls"] += 1 + args = ctx.tool_args + for old, new in self._key_aliases.get(ctx.tool_name, {}).items(): + if old in args and new not in args: + args[new] = args.pop(old) + self._stats["key_renames"] += 1 + for key, expected in self._type_coercions.get(ctx.tool_name, {}).items(): + if key in args and not isinstance(args[key], expected): + try: + args[key] = expected(args[key]) + self._stats["type_coercions"] += 1 + except (TypeError, ValueError): + pass + for key, value in args.items(): + if isinstance(value, str) and value != value.strip(): + args[key] = value.strip() + self._stats["whitespace_strips"] += 1 + ctx.tool_args = self._tool_repair(ctx.tool_name, args) + return ctx + + +__all__ = [ + "DEFAULT_KEY_ALIASES", + "DEFAULT_TYPE_COERCIONS", + "ToolCallRepairMiddleware", + "repair_truncated_json", +] diff --git a/agent_core/runtime/usage_meter.py b/agent_core/runtime/usage_meter.py new file mode 100644 index 0000000..f714297 --- /dev/null +++ b/agent_core/runtime/usage_meter.py @@ -0,0 +1,210 @@ +"""Task-local, thread-safe metering for external APIs, tools, and raw LLMs. + +The meter records transport-level facts without knowing product prices, +billing policy, or telemetry backends. Hosts bind one meter per execution +context and may supply an LLM recorder callback for their own aggregation. +""" + +from __future__ import annotations + +import contextlib +import contextvars +import threading +import time +from collections.abc import Callable +from typing import Any + +_BASE_FIELDS: tuple[str, ...] = ( + "requests", + "cache_hits", + "retries", + "errors", +) + + +class ExternalAPIMeter: + """Accumulate external request, tool-call, gauge, and span measurements.""" + + def __init__(self, *, llm_recorder: Callable[..., None] | None = None) -> None: + self._lock = threading.Lock() + self._providers: dict[str, dict[str, float]] = {} + self._tool_counts: dict[str, int] = {} + self._open_spans: dict[tuple[str, str], tuple[float, str]] = {} + self._gauges: dict[tuple[str, str], float] = {} + self._llm_recorder = llm_recorder + + def record_api_request( + self, + provider: str, + *, + requests: int = 1, + cache_hits: int = 0, + retries: int = 0, + errors: int = 0, + **extra: float, + ) -> None: + """Fold one or more wire-level events into a provider slot.""" + with self._lock: + slot = self._slot_locked(provider) + slot["requests"] += requests + slot["cache_hits"] += cache_hits + slot["retries"] += retries + slot["errors"] += errors + for field, value in extra.items(): + slot[field] = slot.get(field, 0.0) + float(value) + + def record_tool_call(self, name: str) -> None: + with self._lock: + self._tool_counts[name] = self._tool_counts.get(name, 0) + 1 + + def set_gauge(self, provider: str, field: str, value: float) -> None: + """Record a non-additive value using conservative max-wins folding.""" + with self._lock: + key = (provider, field) + self._gauges[key] = max(self._gauges.get(key, 0.0), float(value)) + + def record_llm_usage(self, **kwargs: Any) -> None: + """Forward usage to the host callback without breaking the measured call.""" + if self._llm_recorder is not None: + with contextlib.suppress(Exception): + self._llm_recorder(**kwargs) + + def open_span( + self, + provider: str, + key: str, + *, + field: str = "sandbox_seconds", + ) -> None: + """Start an idempotent monotonic-clock span.""" + with self._lock: + self._open_spans.setdefault( + (provider, key), + (time.monotonic(), field), + ) + + def close_span(self, provider: str, key: str) -> None: + """Close a span and fold its elapsed time into the provider slot.""" + with self._lock: + span = self._open_spans.pop((provider, key), None) + if span is None: + return + started, field = span + slot = self._slot_locked(provider) + slot[field] = slot.get(field, 0.0) + (time.monotonic() - started) + + def snapshot(self) -> dict[str, Any]: + """Return an isolated snapshot including elapsed still-open spans.""" + with self._lock: + now = time.monotonic() + apis: dict[str, dict[str, Any]] = { + provider: { + key: round(value, 2) if isinstance(value, float) else value + for key, value in slot.items() + } + for provider, slot in self._providers.items() + } + for (provider, _key), (started, field) in self._open_spans.items(): + view = apis.setdefault( + provider, + dict.fromkeys(_BASE_FIELDS, 0.0), + ) + view[field] = round( + float(view.get(field, 0.0)) + now - started, + 2, + ) + view["spans_open"] = int(view.get("spans_open", 0)) + 1 + for (provider, field), value in self._gauges.items(): + view = apis.setdefault( + provider, + dict.fromkeys(_BASE_FIELDS, 0.0), + ) + view[field] = round(value, 2) + return { + "external_apis": apis, + "tools": dict(self._tool_counts), + } + + def _slot_locked(self, provider: str) -> dict[str, float]: + slot = self._providers.get(provider) + if slot is None: + slot = dict.fromkeys(_BASE_FIELDS, 0.0) + self._providers[provider] = slot + return slot + + +_CURRENT_METER: contextvars.ContextVar[ExternalAPIMeter | None] = ( + contextvars.ContextVar("agent_core_usage_meter", default=None) +) + + +def bind_usage_meter( + meter: ExternalAPIMeter, +) -> contextvars.Token[ExternalAPIMeter | None]: + """Bind ``meter`` to the current context and return its reset token.""" + return _CURRENT_METER.set(meter) + + +def reset_usage_meter( + token: contextvars.Token[ExternalAPIMeter | None], +) -> None: + _CURRENT_METER.reset(token) + + +def get_usage_meter() -> ExternalAPIMeter | None: + return _CURRENT_METER.get() + + +def record_api_request(provider: str, **kwargs: Any) -> None: + meter = get_usage_meter() + if meter is not None: + meter.record_api_request(provider, **kwargs) + + +def record_tool_call(name: str) -> None: + meter = get_usage_meter() + if meter is not None: + meter.record_tool_call(name) + + +def record_llm_usage(**kwargs: Any) -> None: + meter = get_usage_meter() + if meter is not None: + meter.record_llm_usage(**kwargs) + + +def open_meter_span( + provider: str, + key: str, + *, + field: str = "sandbox_seconds", +) -> None: + meter = get_usage_meter() + if meter is not None: + meter.open_span(provider, key, field=field) + + +def close_meter_span(provider: str, key: str) -> None: + meter = get_usage_meter() + if meter is not None: + meter.close_span(provider, key) + + +def set_meter_gauge(provider: str, field: str, value: float) -> None: + meter = get_usage_meter() + if meter is not None: + meter.set_gauge(provider, field, value) + + +__all__ = [ + "ExternalAPIMeter", + "bind_usage_meter", + "close_meter_span", + "get_usage_meter", + "open_meter_span", + "record_api_request", + "record_llm_usage", + "record_tool_call", + "reset_usage_meter", + "set_meter_gauge", +] diff --git a/tests/test_tool_guardrails.py b/tests/test_tool_guardrails.py new file mode 100644 index 0000000..2c67037 --- /dev/null +++ b/tests/test_tool_guardrails.py @@ -0,0 +1,77 @@ +from __future__ import annotations + +import asyncio +import json + +from agent_core.protocols import ToolCallContext +from agent_core.runtime.loop.guardrails import GuardrailsMiddleware +from agent_core.runtime.loop.tool_call_repair import ( + ToolCallRepairMiddleware, + repair_truncated_json, +) + + +def _context(tool: str, args: dict[str, object]) -> ToolCallContext: + return ToolCallContext( + task_id="task", + phase_id="phase", + role_id="role", + tool_name=tool, + tool_args=args, + metadata={}, + ) + + +def test_guardrails_thresholds_are_configurable() -> None: + middleware = GuardrailsMiddleware(duplicate_thresholds={"custom": 1}) + first = asyncio.run(middleware.before_tool_call(_context("custom", {"x": 1}))) + second = asyncio.run(middleware.before_tool_call(_context("custom", {"x": 1}))) + assert not first.metadata.get("blocked") + assert second.metadata["blocked"] is True + + +def test_guardrails_escalate_repeated_nonconsecutive_patterns() -> None: + middleware = GuardrailsMiddleware(max_loop_hints=1) + asyncio.run(middleware.before_tool_call(_context("custom", {"x": 1}))) + asyncio.run(middleware.before_tool_call(_context("custom", {"x": 2}))) + middleware.notify_loop_hint("task") + result = asyncio.run( + middleware.before_tool_call(_context("custom", {"x": 1})), + ) + assert result.metadata["blocked"] is True + middleware.cleanup_task("task") + + +def test_tool_repair_normalizes_aliases_types_and_special_cases() -> None: + middleware = ToolCallRepairMiddleware() + search = asyncio.run( + middleware.before_tool_call( + _context("web_search", {"q": ["one", "two"], "n": "3"}), + ), + ) + assert search.tool_args == {"query": "one two", "num_results": 3} + bash = asyncio.run( + middleware.before_tool_call( + _context("bash", {"cmd": "```bash\necho ok\n```"}), + ), + ) + assert bash.tool_args == {"command": "echo ok"} + + +def test_tool_repair_accepts_host_alias_tables() -> None: + middleware = ToolCallRepairMiddleware( + key_aliases={"host_tool": {"old": "new"}}, + type_coercions={}, + tool_repair=lambda _name, args: args, + ) + result = asyncio.run( + middleware.before_tool_call(_context("host_tool", {"old": " value "})), + ) + assert result.tool_args == {"new": "value"} + + +def test_repair_truncated_json() -> None: + repaired = repair_truncated_json('{"items": [1, {"ok": true') + assert repaired is not None + assert json.loads(repaired) == {"items": [1, {"ok": True}]} + assert repair_truncated_json("not json") is None diff --git a/tests/test_usage_meter.py b/tests/test_usage_meter.py new file mode 100644 index 0000000..c19d3c9 --- /dev/null +++ b/tests/test_usage_meter.py @@ -0,0 +1,90 @@ +from __future__ import annotations + +from unittest.mock import patch + +from agent_core.runtime.usage_meter import ( + ExternalAPIMeter, + bind_usage_meter, + close_meter_span, + get_usage_meter, + open_meter_span, + record_api_request, + record_llm_usage, + record_tool_call, + reset_usage_meter, + set_meter_gauge, +) + + +def test_meter_records_requests_tools_and_isolates_snapshots() -> None: + meter = ExternalAPIMeter() + meter.record_api_request( + "search", + requests=2, + cache_hits=1, + retries=1, + bytes_read=2.5, + ) + meter.record_tool_call("web_search") + first = meter.snapshot() + first["external_apis"]["search"]["requests"] = 99 + assert meter.snapshot() == { + "external_apis": { + "search": { + "requests": 2.0, + "cache_hits": 1.0, + "retries": 1.0, + "errors": 0.0, + "bytes_read": 2.5, + }, + }, + "tools": {"web_search": 1}, + } + + +def test_open_spans_and_gauges_are_monotonic_and_idempotent() -> None: + meter = ExternalAPIMeter() + with patch("agent_core.runtime.usage_meter.time.monotonic") as clock: + clock.side_effect = [10.0, 12.5, 14.0, 16.0] + meter.open_span("sandbox", "one") + meter.open_span("sandbox", "one") + meter.set_gauge("sandbox", "ttl", 30) + meter.set_gauge("sandbox", "ttl", 20) + current = meter.snapshot()["external_apis"]["sandbox"] + assert current["sandbox_seconds"] == 4.0 + assert current["spans_open"] == 1 + assert current["ttl"] == 30.0 + meter.close_span("sandbox", "one") + assert meter.snapshot()["external_apis"]["sandbox"][ + "sandbox_seconds" + ] == 6.0 + + +def test_context_helpers_are_noop_safe_and_resettable() -> None: + calls: list[dict[str, object]] = [] + meter = ExternalAPIMeter(llm_recorder=lambda **kw: calls.append(kw)) + token = bind_usage_meter(meter) + try: + assert get_usage_meter() is meter + record_api_request("api", errors=1) + record_tool_call("tool") + record_llm_usage(model="m", prompt_tokens=2) + set_meter_gauge("api", "limit", 5) + with patch( + "agent_core.runtime.usage_meter.time.monotonic", + side_effect=[1.0, 2.0], + ): + open_meter_span("api", "span", field="seconds") + close_meter_span("api", "span") + finally: + reset_usage_meter(token) + assert get_usage_meter() is None + assert calls == [{"model": "m", "prompt_tokens": 2}] + assert meter.snapshot()["tools"] == {"tool": 1} + + +def test_llm_recorder_failure_is_suppressed() -> None: + def fail(**_kwargs: object) -> None: + raise RuntimeError("accounting unavailable") + + ExternalAPIMeter(llm_recorder=fail).record_llm_usage(model="m") From cf5cb77e0c2c0c5336447a00aa482a8d263070e0 Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:21:10 +0800 Subject: [PATCH 2/9] feat: share legacy cooldown fallback engine --- agent_core/providers/fallback.py | 224 ++++++++++++++++++++++++++++++- tests/test_cooldown_fallback.py | 89 ++++++++++++ 2 files changed, 312 insertions(+), 1 deletion(-) create mode 100644 tests/test_cooldown_fallback.py diff --git a/agent_core/providers/fallback.py b/agent_core/providers/fallback.py index 1fb0bd6..348d2b9 100644 --- a/agent_core/providers/fallback.py +++ b/agent_core/providers/fallback.py @@ -54,7 +54,9 @@ # pyright: basic, reportPrivateImportUsage=false import asyncio import logging -from collections.abc import AsyncIterator, Sequence +import random +import time +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from dataclasses import dataclass, field from typing import Any, Literal @@ -64,12 +66,33 @@ logger = logging.getLogger(__name__) __all__ = [ + "CooldownFallbackLLM", "FallbackEntry", "FallbackTrigger", "LLMFallbackChain", "with_provider_stamp", ] +type FallbackEventHook = Callable[[str, dict[str, Any]], Awaitable[None]] + +_LEGACY_RETRYABLE_KEYWORDS = frozenset({ + "timeout", "timed out", "429", "500", "502", "503", "504", "529", + "overloaded", "rate limit", "rate_limit", "server error", + "connection reset", "connection error", "econnreset", "gateway timeout", + "model_dump", "model_not_found", +}) + + +def legacy_retryable(error: Exception) -> bool: + """Match the historical two-model fallback wrapper's retry policy.""" + return isinstance(error, AttributeError) or any( + keyword in str(error).lower() for keyword in _LEGACY_RETRYABLE_KEYWORDS + ) + + +async def _noop_event(_name: str, _payload: dict[str, Any]) -> None: + return None + FallbackTrigger = Literal["timeout", "rate_limit", "5xx", "any_error"] @@ -171,6 +194,205 @@ def _trigger_matches(trigger: FallbackTrigger, exc: BaseException) -> bool: return False +class CooldownFallbackLLM: + """Retry a primary client, then use a fallback client during cooldown. + + This preserves the legacy two-model fallback behavior used by both host + products. Product tracing is an optional async event hook, keeping the + retry state machine independent of execution scopes and trace backends. + """ + + def __init__( + self, + primary: Any, + fallback: Any, + max_retries: int = 2, + cooldown_seconds: int = 60, + *, + retryable: Callable[[Exception], bool] = legacy_retryable, + event_hook: FallbackEventHook = _noop_event, + clock: Callable[[], float] = time.time, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + jitter: Callable[[], float] = random.random, + ) -> None: + self.primary = primary + self.fallback = fallback + self.max_retries = max_retries + self.cooldown_seconds = cooldown_seconds + self.model: str = _model_id(primary) + self._retryable = retryable + self._event_hook = event_hook + self._clock = clock + self._sleep = sleep + self._jitter = jitter + self._cooldown_until = 0.0 + + @property + def model_name(self) -> str: + return f"fallback({_model_id(self.primary)})" + + async def _emit(self, name: str, **payload: Any) -> None: + try: + await self._event_hook(name, payload) + except Exception: + logger.debug("Fallback event hook failed", exc_info=True) + + def _kwargs( + self, + *, + tools: list[dict[str, Any]] | None, + temperature: float | None, + max_tokens: int | None, + extra_headers: dict[str, str] | None, + timeout: float | None, + ) -> dict[str, Any]: + return { + "tools": tools, + "temperature": temperature, + "max_tokens": max_tokens, + "extra_headers": extra_headers, + "timeout": timeout, + } + + async def chat( + self, + messages: list[Message], + *, + tools: list[dict[str, Any]] | None = None, + temperature: float | None = None, + max_tokens: int | None = None, + extra_headers: dict[str, str] | None = None, + timeout: float | None = None, + ) -> LLMResponse: + kwargs = self._kwargs( + tools=tools, + temperature=temperature, + max_tokens=max_tokens, + extra_headers=extra_headers, + timeout=timeout, + ) + if self._clock() < self._cooldown_until: + await self._emit( + "degrade", + reason="primary_cooldown", + degrade_from=_model_id(self.primary), + degrade_to=_model_id(self.fallback), + ) + return await self.fallback.chat(messages, **kwargs) + + last_error: Exception | None = None + for attempt in range(self.max_retries): + try: + await self._emit("request", leg="primary", attempt=attempt + 1) + return await self.primary.chat(messages, **kwargs) + except Exception as error: + last_error = error + await self._emit( + "error", + leg="primary", + attempt=attempt + 1, + error=str(error), + ) + if not self._retryable(error): + break + delay = min(0.5 * (2**attempt), 8.0) + self._jitter() * 0.25 + await self._emit("retry", attempt=attempt + 1, delay_s=delay) + await self._sleep(delay) + + self._cooldown_until = self._clock() + self.cooldown_seconds + await self._emit( + "degrade", + reason="primary_exhausted", + degrade_from=_model_id(self.primary), + degrade_to=_model_id(self.fallback), + cooldown_seconds=self.cooldown_seconds, + ) + try: + return await self.fallback.chat(messages, **kwargs) + except Exception as fallback_error: + await self._emit("error", leg="fallback", error=str(fallback_error)) + if last_error is None: + raise + raise last_error from fallback_error + + async def stream( + self, + messages: list[Message], + *, + tools: list[dict[str, Any]] | None = None, + temperature: float | None = None, + max_tokens: int | None = None, + extra_headers: dict[str, str] | None = None, + timeout: float | None = None, + ) -> AsyncIterator[StreamDelta]: + kwargs = self._kwargs( + tools=tools, + temperature=temperature, + max_tokens=max_tokens, + extra_headers=extra_headers, + timeout=timeout, + ) + if self._clock() < self._cooldown_until: + await self._emit("degrade", reason="primary_stream_cooldown") + async for delta in self.fallback.stream(messages, **kwargs): + yield delta + return + + last_error: Exception | None = None + for attempt in range(self.max_retries): + yielded = False + try: + await self._emit( + "request", + leg="primary", + attempt=attempt + 1, + streaming=True, + ) + async for delta in self.primary.stream(messages, **kwargs): + yielded = True + yield delta + return + except Exception as error: + last_error = error + await self._emit( + "error", + leg="primary", + attempt=attempt + 1, + streaming=True, + error=str(error), + ) + # The historical wrapper retried even after yielding. Preserve + # that behavior here; callers choosing rewind-safe semantics + # should use LLMFallbackChain instead. + if not self._retryable(error): + break + delay = min(0.5 * (2**attempt), 8.0) + self._jitter() * 0.25 + await self._emit( + "retry", + attempt=attempt + 1, + delay_s=delay, + streaming=True, + yielded=yielded, + ) + await self._sleep(delay) + + self._cooldown_until = self._clock() + self.cooldown_seconds + await self._emit("degrade", reason="primary_stream_exhausted") + try: + async for delta in self.fallback.stream(messages, **kwargs): + yield delta + except Exception as fallback_error: + await self._emit( + "error", + leg="fallback", + streaming=True, + error=str(fallback_error), + ) + if last_error is None: + raise + raise last_error from fallback_error + + @dataclass class LLMFallbackChain: """Ordered list of :class:`LLMClient` entries with per-entry triggers. diff --git a/tests/test_cooldown_fallback.py b/tests/test_cooldown_fallback.py new file mode 100644 index 0000000..aa5302b --- /dev/null +++ b/tests/test_cooldown_fallback.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +from collections.abc import AsyncIterator + +import pytest + +from agent_core.llm import LLMResponse, StreamDelta +from agent_core.messages import Message, user_msg +from agent_core.providers.fallback import CooldownFallbackLLM, legacy_retryable + + +class ScriptedLLM: + def __init__(self, model: str, script: list[LLMResponse | Exception]) -> None: + self.model = model + self.script = script + self.calls = 0 + self.kwargs: list[dict[str, object]] = [] + + async def chat(self, _messages: list[Message], **kwargs: object) -> LLMResponse: + self.kwargs.append(kwargs) + item = self.script[min(self.calls, len(self.script) - 1)] + self.calls += 1 + if isinstance(item, Exception): + raise item + return item + + async def stream( + self, + messages: list[Message], + **kwargs: object, + ) -> AsyncIterator[StreamDelta]: + response = await self.chat(messages, **kwargs) + yield StreamDelta(content=str(response.content)) + + +@pytest.mark.asyncio +async def test_cooldown_fallback_retries_then_skips_primary() -> None: + primary = ScriptedLLM("primary", [TimeoutError("timeout")]) + fallback = ScriptedLLM("fallback", [LLMResponse(content="ok")]) + sleeps: list[float] = [] + + async def sleep(delay: float) -> None: + sleeps.append(delay) + + clock = iter([10.0, 10.0, 11.0]).__next__ + llm = CooldownFallbackLLM( + primary, + fallback, + max_retries=1, + cooldown_seconds=60, + clock=clock, + sleep=sleep, + jitter=lambda: 0.0, + ) + assert (await llm.chat([user_msg("one")])).content == "ok" + assert (await llm.chat([user_msg("two")])).content == "ok" + assert primary.calls == 1 + assert fallback.calls == 2 + assert sleeps == [0.5] + + +@pytest.mark.asyncio +async def test_cooldown_fallback_preserves_kwargs_and_original_error() -> None: + original = TimeoutError("primary failed") + primary = ScriptedLLM("primary", [original]) + fallback = ScriptedLLM("fallback", [RuntimeError("fallback failed")]) + llm = CooldownFallbackLLM( + primary, + fallback, + max_retries=1, + cooldown_seconds=0, + sleep=lambda _delay: _completed(), + jitter=lambda: 0.0, + ) + tools = [{"type": "function", "function": {"name": "search"}}] + with pytest.raises(TimeoutError, match="primary failed"): + await llm.chat([user_msg("one")], tools=tools) + assert primary.kwargs[0]["tools"] == tools + assert fallback.kwargs[0]["tools"] == tools + + +async def _completed() -> None: + return None + + +def test_legacy_retryable_contract() -> None: + assert legacy_retryable(TimeoutError("timed out")) + assert legacy_retryable(AttributeError("model_dump")) + assert not legacy_retryable(ValueError("invalid json")) From 88672ac412b416bb34637e88ec087b8d3b1cfad4 Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:22:40 +0800 Subject: [PATCH 3/9] feat: add configurable auxiliary LLM factory --- agent_core/providers/aux_builder.py | 184 ++++++++++++++++++++++++++++ tests/test_aux_builder.py | 81 ++++++++++++ 2 files changed, 265 insertions(+) create mode 100644 agent_core/providers/aux_builder.py create mode 100644 tests/test_aux_builder.py diff --git a/agent_core/providers/aux_builder.py b/agent_core/providers/aux_builder.py new file mode 100644 index 0000000..7c63302 --- /dev/null +++ b/agent_core/providers/aux_builder.py @@ -0,0 +1,184 @@ +"""Configurable factory for profile-defined auxiliary LLM clients.""" + +from __future__ import annotations + +# pyright: basic +import logging +from collections.abc import Callable, Mapping +from typing import Any, cast + +logger = logging.getLogger(__name__) + +type ClientFactory = Callable[..., Any] +type ProviderTypeResolver = Callable[[str], str] +type SessionHeadersResolver = Callable[[str, Mapping[str, Any]], Mapping[str, str]] +type ClientDecorator = Callable[[Any, str, str], Any] + +_DUMMY_KEY_WARNED: set[tuple[str, str]] = set() + + +def _resolve_api_key(section: Mapping[str, Any], provider: str) -> str: + key = section.get("api_key") or "" + if key and str(key).strip(): + return str(key) + model = str(section.get("model") or "") + cache_key = (provider, model) + if cache_key not in _DUMMY_KEY_WARNED: + _DUMMY_KEY_WARNED.add(cache_key) + logger.warning( + "Aux LLM provider=%r model=%r has no api_key; using the legacy " + "'dummy' token", + provider or "", + model or "", + ) + return "dummy" + + +def _thinking_budget(section: Mapping[str, Any]) -> object | None: + for key in ("thinking_budget", "thinking_budget_tokens"): + value = section.get(key) + if value is not None: + return value + thinking = section.get("thinking") + if isinstance(thinking, Mapping): + for key in ("budget", "budget_tokens", "max_tokens"): + value = thinking.get(key) + if value is not None: + return value + return None + + +def _anthropic_thinking(section: Mapping[str, Any]) -> dict[str, Any] | None: + raw = section.get("thinking") + if not isinstance(raw, Mapping): + return None + thinking = {str(key): value for key, value in raw.items()} + kind = str(thinking.get("type") or "").strip().lower() + if kind in {"", "disabled", "off", "none", "false"}: + return None + return thinking + + +def _headers(value: object) -> dict[str, str]: + if not isinstance(value, Mapping): + return {} + return { + str(key): str(item) + for key, item in value.items() + if item is not None + } + + +class AuxLLMFactory: + """Build OpenAI-compatible, Anthropic, or Bedrock auxiliary clients. + + Host products provide catalog resolution, concrete constructors, session + headers, and post-build decoration. All request-shape normalization stays + here so profile-driven DAG/report/summary clients cannot drift. + """ + + def __init__( + self, + *, + openai_factory: ClientFactory, + anthropic_factory: ClientFactory, + provider_type: ProviderTypeResolver, + session_headers: SessionHeadersResolver | None = None, + decorate: ClientDecorator | None = None, + ) -> None: + self._openai_factory = openai_factory + self._anthropic_factory = anthropic_factory + self._provider_type = provider_type + self._session_headers = session_headers + self._decorate = decorate + + def build(self, section: Mapping[str, Any]) -> Any: + provider = str( + section.get("_provider_label") or section.get("provider") or "", + ) + provider_type = self._provider_type(provider).lower() + model = str(section.get("model") or "") + if provider_type in {"anthropic", "bedrock"}: + if "claude" not in model.strip().lower(): + raise ValueError( + f"provider {provider!r} uses {provider_type!r} transport, " + f"which requires a Claude model; got {model!r}", + ) + client = self._build_anthropic( + section, + provider, + bedrock=provider_type == "bedrock", + ) + else: + client = self._build_openai(section, provider) + if self._decorate is not None: + client = self._decorate(client, provider, model) + return client + + def _build_anthropic( + self, + section: Mapping[str, Any], + provider: str, + *, + bedrock: bool, + ) -> Any: + kwargs: dict[str, Any] = { + "model": section["model"], + "api_key": _resolve_api_key(section, provider), + "temperature": section.get("temperature", 0.0), + "timeout": float(section.get("llm_timeout_s", 120)), + "thinking": _anthropic_thinking(section), + "effort": str(section.get("effort") or ""), + "bedrock": bedrock, + } + if section.get("base_url"): + kwargs["base_url"] = section["base_url"] + default_headers = _headers(section.get("extra_headers")) + if default_headers: + kwargs["default_headers"] = default_headers + maximum = section.get("max_completion_tokens") or section.get("max_tokens") + if maximum is not None: + kwargs["max_tokens"] = int(maximum) + return self._anthropic_factory(**kwargs) + + def _build_openai( + self, + section: Mapping[str, Any], + provider: str, + ) -> Any: + kwargs: dict[str, Any] = { + "model": section["model"], + "api_key": _resolve_api_key(section, provider), + "base_url": section.get("base_url") or None, + "temperature": section.get("temperature", 0.0), + "timeout": float(section.get("llm_timeout_s", 120)), + } + maximum = section.get("max_completion_tokens") or section.get("max_tokens") + if maximum is not None: + kwargs["max_completion_tokens"] = int(maximum) + + extra_body_value = section.get("extra_body") + extra_body = dict(extra_body_value) if isinstance(extra_body_value, Mapping) else {} + template_value = extra_body.get("chat_template_kwargs") + template = dict(template_value) if isinstance(template_value, Mapping) else {} + if section.get("enable_thinking"): + template["enable_thinking"] = True + template.setdefault("preserve_thinking", False) + budget = _thinking_budget(section) + if budget is not None: + template["thinking_budget"] = int(cast("Any", budget)) + if template: + extra_body["chat_template_kwargs"] = template + if extra_body: + kwargs["extra_body"] = extra_body + + default_headers: dict[str, str] = {} + if self._session_headers is not None: + default_headers.update(self._session_headers(provider, section)) + default_headers.update(_headers(section.get("extra_headers"))) + if default_headers: + kwargs["default_headers"] = default_headers + return self._openai_factory(**kwargs) + + +__all__ = ["AuxLLMFactory"] diff --git a/tests/test_aux_builder.py b/tests/test_aux_builder.py new file mode 100644 index 0000000..dadca1d --- /dev/null +++ b/tests/test_aux_builder.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +from typing import Any + +import pytest + +from agent_core.providers.aux_builder import AuxLLMFactory + + +def _capture(kind: str, calls: list[tuple[str, dict[str, Any]]]): + def build(**kwargs: Any) -> dict[str, Any]: + calls.append((kind, kwargs)) + return kwargs + + return build + + +def test_openai_aux_request_shape_and_host_hooks() -> None: + calls: list[tuple[str, dict[str, Any]]] = [] + factory = AuxLLMFactory( + openai_factory=_capture("openai", calls), + anthropic_factory=_capture("anthropic", calls), + provider_type=lambda _provider: "openai_compat", + session_headers=lambda _provider, section: { + "x-session": str(section.get("session_suffix", "")), + }, + decorate=lambda client, provider, _model: {**client, "provider": provider}, + ) + client = factory.build({ + "provider": "gateway", + "model": "qwen", + "api_key": "key", + "enable_thinking": True, + "thinking_budget": 123, + "session_suffix": ":dag", + "extra_headers": {"X-Inspection": "on"}, + "max_tokens": 99, + }) + assert calls[0][0] == "openai" + assert client["provider"] == "gateway" + assert client["max_completion_tokens"] == 99 + assert client["default_headers"] == { + "x-session": ":dag", + "X-Inspection": "on", + } + assert client["extra_body"]["chat_template_kwargs"] == { + "enable_thinking": True, + "preserve_thinking": False, + "thinking_budget": 123, + } + + +def test_anthropic_and_bedrock_request_shape() -> None: + calls: list[tuple[str, dict[str, Any]]] = [] + factory = AuxLLMFactory( + openai_factory=_capture("openai", calls), + anthropic_factory=_capture("anthropic", calls), + provider_type=lambda _provider: "bedrock", + ) + result = factory.build({ + "provider": "aws", + "model": "global.anthropic.claude-sonnet", + "api_key": "key", + "thinking": {"type": "adaptive"}, + "effort": "high", + "max_completion_tokens": 500, + }) + assert calls[0][0] == "anthropic" + assert result["bedrock"] is True + assert result["thinking"] == {"type": "adaptive"} + assert result["max_tokens"] == 500 + + +def test_native_anthropic_rejects_non_claude_model() -> None: + factory = AuxLLMFactory( + openai_factory=lambda **_kwargs: None, + anthropic_factory=lambda **_kwargs: None, + provider_type=lambda _provider: "anthropic", + ) + with pytest.raises(ValueError, match="requires a Claude model"): + factory.build({"provider": "anthropic", "model": "qwen"}) From e6bdcb9d47b164eaa6e5d496820ef0f343ec3a7e Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:25:45 +0800 Subject: [PATCH 4/9] feat: share summary LLM execution engine --- agent_core/providers/summary.py | 233 ++++++++++++++++++++++++++++++++ tests/test_summary_engine.py | 40 ++++++ 2 files changed, 273 insertions(+) create mode 100644 agent_core/providers/summary.py create mode 100644 tests/test_summary_engine.py diff --git a/agent_core/providers/summary.py b/agent_core/providers/summary.py new file mode 100644 index 0000000..ec7fc7e --- /dev/null +++ b/agent_core/providers/summary.py @@ -0,0 +1,233 @@ +"""Provider-neutral execution engine for OpenAI-compatible summary LLMs.""" + +from __future__ import annotations + +# pyright: basic, reportMissingImports=false +import asyncio +import hashlib +import logging +from collections.abc import Callable, Sequence +from typing import Any, NotRequired, TypedDict + +import httpx + +logger = logging.getLogger(__name__) + + +class SummaryCandidate(TypedDict): + endpoint: str + model: str + api_key: NotRequired[str] + provider: NotRequired[str] + extra_headers: NotRequired[dict[str, str]] + + +type UsageRecorder = Callable[..., None] + +EXTRACT_INFO_PROMPT = """You are given a piece of content and the requirement of information to extract. Your task is to extract the information specifically requested. Be precise and focus exclusively on the requested information. + +INFORMATION TO EXTRACT: +{focus} + +INSTRUCTIONS: +1. Extract the information relevant to the focus above. +2. If the exact information is not found, extract the most closely related details. +3. Be specific and include exact details when available. +4. Clearly organize the extracted information for easy understanding. +5. Do not include general summaries or unrelated content. + +CONTENT TO ANALYZE: +{content} + +EXTRACTED INFORMATION:""" + + +def normalize_summary_endpoint(base_url: str) -> str: + endpoint = (base_url or "").rstrip("/") + if endpoint and not endpoint.endswith("/chat/completions"): + endpoint += "/chat/completions" + return endpoint + + +def describe_summary_candidates(candidates: Sequence[SummaryCandidate]) -> str: + """Render candidates without exposing usable credentials.""" + if not candidates: + return "(no summary LLM candidates)" + lines: list[str] = [] + for index, candidate in enumerate(candidates, 1): + key = str(candidate.get("api_key") or "") + fingerprint = ( + f"len={len(key)} #{hashlib.sha256(key.encode()).hexdigest()[:12]}" + if key + else "unset" + ) + lines.append( + f"{index}. provider={candidate.get('provider') or '?'}" + f" model={candidate.get('model') or '?'}" + f" endpoint={candidate.get('endpoint') or '?'}" + f" api_key={fingerprint}", + ) + return "\n".join(lines) + + +def build_summary_payload(model: str, prompt: str) -> dict[str, Any]: + """Build the shared Chat Completions request shape for extraction.""" + lowered = model.lower() + if "gpt-5" in lowered or "gpt5" in lowered: + return { + "model": model, + "max_completion_tokens": 8192, + "messages": [{"role": "user", "content": prompt}], + "reasoning_effort": "minimal", + "service_tier": "flex", + } + payload: dict[str, Any] = { + "model": model, + "max_tokens": 8192, + "messages": [{"role": "user", "content": prompt}], + "temperature": 1.0, + } + if any(key in lowered for key in ("qwen", "apodex", "sglang", "397b")): + payload["chat_template_kwargs"] = {"enable_thinking": False} + return payload + + +def truncate_summary_fallback(content: str, limit: int = 20_000) -> str: + if len(content) > limit: + return content[:limit] + "\n\n[Content truncated...]" + return content + + +class SummaryLLMEngine: + """Retry ordered raw-HTTP summary candidates and degrade to truncation.""" + + def __init__( + self, + *, + max_retries: int = 4, + truncate_step: int = 40_960, + request_timeout: float = 300, + fallback_limit: int = 20_000, + usage_recorder: UsageRecorder | None = None, + sleep: Callable[[float], Any] = asyncio.sleep, + ) -> None: + self.max_retries = max_retries + self.truncate_step = truncate_step + self.request_timeout = request_timeout + self.fallback_limit = fallback_limit + self.usage_recorder = usage_recorder + self.sleep = sleep + + async def summarize( + self, + content: str, + focus: str, + candidates: Sequence[SummaryCandidate], + ) -> str: + for index, candidate in enumerate(candidates): + summary = await self._summarize_one( + candidate, + content, + focus, + candidate_index=index, + ) + if summary: + return summary + if index + 1 < len(candidates): + logger.warning( + "Summary candidate %d (%s) exhausted; falling back to %s", + index + 1, + candidate.get("model"), + candidates[index + 1].get("model"), + ) + return truncate_summary_fallback(content, self.fallback_limit) + + async def _summarize_one( + self, + candidate: SummaryCandidate, + content: str, + focus: str, + *, + candidate_index: int, + ) -> str: + endpoint = candidate["endpoint"] + model = candidate["model"] + prompt = EXTRACT_INFO_PROMPT.format(focus=focus, content=content) + payload = build_summary_payload(model, prompt) + headers = {"Content-Type": "application/json"} + api_key = candidate.get("api_key") or "" + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + headers.update(candidate.get("extra_headers") or {}) + + current_content = content + for attempt in range(self.max_retries): + try: + async with httpx.AsyncClient(timeout=self.request_timeout) as client: + response = await client.post( + endpoint, + headers=headers, + json=payload, + ) + body = response.text + if response.status_code >= 400 and ( + "maximum context length" in body + or "longer than the model's context length" in body + ): + remove = self.truncate_step * (attempt + 1) + if remove >= len(current_content): + return "" + current_content = content[:-remove] + "[...truncated]" + payload["messages"][0]["content"] = EXTRACT_INFO_PROMPT.format( + focus=focus, + content=current_content, + ) + continue + response.raise_for_status() + data = response.json() + self._record_usage(candidate, data.get("usage") or {}) + summary = ( + data.get("choices", [{}])[0] + .get("message", {}) + .get("content", "") + ) + if summary: + return str(summary) + logger.warning( + "Empty summary response on attempt %d (candidate %d)", + attempt + 1, + candidate_index + 1, + ) + return "" + except Exception as error: + logger.warning("Summary LLM attempt %d failed: %s", attempt + 1, error) + if attempt + 1 < self.max_retries: + await self.sleep(float(attempt + 1)) + return "" + + def _record_usage( + self, + candidate: SummaryCandidate, + usage: dict[str, Any], + ) -> None: + if not usage or self.usage_recorder is None: + return + details = usage.get("prompt_tokens_details") or {} + self.usage_recorder( + model=candidate.get("model", ""), + provider=candidate.get("provider") or "summary_llm", + prompt_tokens=int(usage.get("prompt_tokens", 0) or 0), + completion_tokens=int(usage.get("completion_tokens", 0) or 0), + cache_read_tokens=int(details.get("cached_tokens", 0) or 0), + ) + + +__all__ = [ + "EXTRACT_INFO_PROMPT", + "SummaryCandidate", + "SummaryLLMEngine", + "build_summary_payload", + "describe_summary_candidates", + "normalize_summary_endpoint", + "truncate_summary_fallback", +] diff --git a/tests/test_summary_engine.py b/tests/test_summary_engine.py new file mode 100644 index 0000000..2577519 --- /dev/null +++ b/tests/test_summary_engine.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +from agent_core.providers.summary import ( + build_summary_payload, + describe_summary_candidates, + normalize_summary_endpoint, + truncate_summary_fallback, +) + + +def test_summary_endpoint_normalization() -> None: + assert normalize_summary_endpoint("https://host/v1/") == ( + "https://host/v1/chat/completions" + ) + assert normalize_summary_endpoint("https://host/v1/chat/completions") == ( + "https://host/v1/chat/completions" + ) + + +def test_summary_payload_dialects() -> None: + assert "max_completion_tokens" in build_summary_payload("gpt-5", "prompt") + qwen = build_summary_payload("qwen-3", "prompt") + assert qwen["chat_template_kwargs"] == {"enable_thinking": False} + assert qwen["temperature"] == 1.0 + + +def test_candidate_description_redacts_api_key() -> None: + rendered = describe_summary_candidates([{ + "endpoint": "https://host/v1/chat/completions", + "model": "model", + "provider": "provider", + "api_key": "top-secret-key", + }]) + assert "top-secret-key" not in rendered + assert "len=14" in rendered + + +def test_truncate_fallback() -> None: + assert truncate_summary_fallback("short", 10) == "short" + assert truncate_summary_fallback("01234567890", 10).startswith("0123456789") From 1a05d2c2868c79726f1864440b12d0f90e6ad95f Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:26:49 +0800 Subject: [PATCH 5/9] fix: merge workflow defaults without null blanking --- agent_core/scheduling/workflow_defaults.py | 19 ++++++++++++++++--- tests/test_scheduler.py | 17 +++++++++++++++++ 2 files changed, 33 insertions(+), 3 deletions(-) diff --git a/agent_core/scheduling/workflow_defaults.py b/agent_core/scheduling/workflow_defaults.py index dc944f3..aad6894 100644 --- a/agent_core/scheduling/workflow_defaults.py +++ b/agent_core/scheduling/workflow_defaults.py @@ -11,12 +11,25 @@ def register_workflow_defaults( pipeline_ids: str | tuple[str, ...], defaults: Mapping[str, Any], + *, + merge: bool = True, + include_none: bool = False, ) -> None: - """Register an immutable snapshot of product-owned workflow defaults.""" + """Register product defaults without accidentally blanking prior values. + + Registrations merge by default because hosts commonly contribute defaults + from several composition modules. ``None`` means "no product default" and + is ignored unless ``include_none`` is explicitly requested. + """ ids = (pipeline_ids,) if isinstance(pipeline_ids, str) else pipeline_ids - snapshot = dict(defaults) + snapshot = { + key: value + for key, value in defaults.items() + if include_none or value is not None + } for pipeline_id in ids: - _WORKFLOW_DEFAULTS[pipeline_id] = snapshot + current = _WORKFLOW_DEFAULTS.get(pipeline_id, {}) if merge else {} + _WORKFLOW_DEFAULTS[pipeline_id] = {**current, **snapshot} def clear_workflow_defaults() -> None: diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index 56b5c3f..566520d 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -15,6 +15,7 @@ ) from agent_core.scheduling.workflow_defaults import ( clear_workflow_defaults, + get_workflow_default, register_workflow_defaults, ) from agent_core.types import TaskId, TaskStatus @@ -145,6 +146,22 @@ def test_agent_team_thrash_is_an_incomplete_stop_reason(): clear_workflow_defaults() +def test_workflow_defaults_merge_and_ignore_missing_values(): + clear_workflow_defaults() + register_workflow_defaults("research", {"wall": 900, "mode": "soft"}) + register_workflow_defaults("research", {"wall": None, "retries": 2}) + assert get_workflow_default("research", "wall") == 900 + assert get_workflow_default("research", "mode") == "soft" + assert get_workflow_default("research", "retries") == 2 + register_workflow_defaults( + "research", + {"wall": None}, + include_none=True, + ) + assert get_workflow_default("research", "wall") is None + clear_workflow_defaults() + + @pytest.mark.asyncio async def test_scheduler_completes_only_with_terminal_output(): pm = _ProcessManager() From f2fa94ed89c14ba9cbf62b0f0e4f764928537f46 Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:28:38 +0800 Subject: [PATCH 6/9] docs: define final shared runtime boundary --- README.md | 15 +++++++++++---- docs/provider-substrate-boundary.md | 9 ++++++++- 2 files changed, 19 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index c19cf40..b82276d 100644 --- a/README.md +++ b/README.md @@ -24,6 +24,8 @@ Version `0.1.x` contains the converged foundation layer: - failure-isolated async event dispatch and composable fail-closed tool permission policy; - provider-neutral LLM response, stream, and client contracts; +- profile-driven auxiliary-client construction, summary execution, legacy + cooldown fallback, and task-local external-API metering; - loop configuration, lifecycle contexts, observer protocol, intervention merging, and observer dispatch helpers. - streamed tool-call recovery checks for missing required arguments. @@ -72,10 +74,12 @@ Code belongs in AgentCore when it: - accepts product behavior through typed inputs or explicit hooks; - has tests that run without either product repository installed. -Provider clients, session affinity, durable process/event-store implementations, +Provider catalogs, credentials, endpoint selection, host session-affinity +context, billing policy, durable process/event-store implementations, user/session retention policy, checkpoint association, workflow node implementations, sandbox mounting/authorization, UI history, and -product-specific observers remain in their product. +product-specific observers remain in their product. AgentCore owns the +provider transports and product-neutral affinity lifecycle safeguards. ## Development @@ -160,7 +164,10 @@ edit both products' core copies, that is evidence it belongs here. 9. **Provider substrate** (complete in AgentCore): shared provider transports, fallback engine, prompt cache, usage metadata normalization, and non-blocking streaming behind product configuration and session-affinity adapters. -10. Remove product compatibility facades after downstream imports have moved to - `agent_core`. +10. **Shared runtime closeout**: portable usage metering, configurable aux-LLM + construction, summary execution, tool-call repair/guardrails, and safe + workflow-default merging. +11. Remove product compatibility facades after downstream imports have moved to + `agent_core` and the compatibility window has elapsed. Each slice must leave product CI green and must not depend on a floating branch. diff --git a/docs/provider-substrate-boundary.md b/docs/provider-substrate-boundary.md index bf0f16f..33ec875 100644 --- a/docs/provider-substrate-boundary.md +++ b/docs/provider-substrate-boundary.md @@ -11,10 +11,17 @@ all host products: non-blocking diagnostic streams. Hosts continue to own provider catalogs, credentials, endpoint selection, -deployment-specific headers, billing meters, traces, and UI-facing provider +deployment-specific headers, billing policy/sinks, traces, and UI-facing provider metadata. Session affinity is supplied to `OpenAIClient` through a `SessionQueryResolver`; AgentCore never interprets a host header by itself. +AgentCore additionally owns the product-neutral mechanics for profile-defined +auxiliary clients, raw-HTTP summary execution, cooldown fallback, and task-local +usage accumulation. Hosts inject provider-type lookup, concrete constructors, +session headers, decorators, candidate configuration, and billing/trace sinks. +The shared meter records quantities only; it does not assign prices or decide +which events are billable. + The SDK response boundary is intentionally dynamic. Third-party OpenAI and Anthropic response classes vary by SDK and compatible gateway, so those files use basic Pyright checking locally while the public constructors, AgentCore From e7745fc1df5a608ed03536c9bb66a3b0d02848f1 Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:33:30 +0800 Subject: [PATCH 7/9] fix: preserve mirothinker summary payload compatibility --- agent_core/providers/summary.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/agent_core/providers/summary.py b/agent_core/providers/summary.py index ec7fc7e..b417a00 100644 --- a/agent_core/providers/summary.py +++ b/agent_core/providers/summary.py @@ -87,7 +87,10 @@ def build_summary_payload(model: str, prompt: str) -> dict[str, Any]: "messages": [{"role": "user", "content": prompt}], "temperature": 1.0, } - if any(key in lowered for key in ("qwen", "apodex", "sglang", "397b")): + if any( + key in lowered + for key in ("qwen", "apodex", "mirothinker", "sglang", "397b") + ): payload["chat_template_kwargs"] = {"enable_thinking": False} return payload From 95915a97b4eb941ea88ee97db880c89b59ce42ea Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 18:40:16 +0800 Subject: [PATCH 8/9] feat: allow product-specific auxiliary key policy --- agent_core/providers/aux_builder.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/agent_core/providers/aux_builder.py b/agent_core/providers/aux_builder.py index 7c63302..6138c91 100644 --- a/agent_core/providers/aux_builder.py +++ b/agent_core/providers/aux_builder.py @@ -13,6 +13,7 @@ type ProviderTypeResolver = Callable[[str], str] type SessionHeadersResolver = Callable[[str, Mapping[str, Any]], Mapping[str, str]] type ClientDecorator = Callable[[Any, str, str], Any] +type APIKeyResolver = Callable[[Mapping[str, Any], str], str] _DUMMY_KEY_WARNED: set[tuple[str, str]] = set() @@ -85,12 +86,14 @@ def __init__( provider_type: ProviderTypeResolver, session_headers: SessionHeadersResolver | None = None, decorate: ClientDecorator | None = None, + api_key_resolver: APIKeyResolver = _resolve_api_key, ) -> None: self._openai_factory = openai_factory self._anthropic_factory = anthropic_factory self._provider_type = provider_type self._session_headers = session_headers self._decorate = decorate + self._api_key_resolver = api_key_resolver def build(self, section: Mapping[str, Any]) -> Any: provider = str( @@ -124,7 +127,7 @@ def _build_anthropic( ) -> Any: kwargs: dict[str, Any] = { "model": section["model"], - "api_key": _resolve_api_key(section, provider), + "api_key": self._api_key_resolver(section, provider), "temperature": section.get("temperature", 0.0), "timeout": float(section.get("llm_timeout_s", 120)), "thinking": _anthropic_thinking(section), @@ -148,7 +151,7 @@ def _build_openai( ) -> Any: kwargs: dict[str, Any] = { "model": section["model"], - "api_key": _resolve_api_key(section, provider), + "api_key": self._api_key_resolver(section, provider), "base_url": section.get("base_url") or None, "temperature": section.get("temperature", 0.0), "timeout": float(section.get("llm_timeout_s", 120)), From 29a995286f45abd1e7d7a1a82568d873ddbd8ef0 Mon Sep 17 00:00:00 2001 From: zhanghanduo Date: Wed, 2 Sep 2026 20:59:47 +0800 Subject: [PATCH 9/9] fix: harden shared runtime extraction after review MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Review fixes for the modules this branch extracted, plus four provider defects the review surfaced in code the extraction depends on. Extracted-module fixes: - tool_call_repair: stop whitespace-stripping literal content args. Every string arg was stripped, so file_editor_str_replace lost the indentation and trailing newline of old_string/new_string (breaking exact matching) and write_file dropped trailing newlines. Adds LITERAL_CONTENT_KEYS. - tool_call_repair: honour an explicitly empty key_aliases/type_coercions as an opt-out instead of falling back to the default table. - summary: classify permanent 4xx as non-retryable. A bad API key burned max_retries attempts per candidate (8 requests, 12s of sleep); now 2 requests and no backoff. Injectable via `retryable`. - summary: measure the context-length retry ladder against the original content, not the already-shortened text, which abandoned a candidate one truncation step early. - guardrails: read total_tokens, the key providers and UsageMetadata actually emit. "total" appears nowhere else in agent_core, so the budget-exhausted hint never fired. Both keys are now accepted. - guardrails: coerce untyped budget config instead of raising ValueError out of an advisory check. - guardrails: tolerate non-JSON tool args when fingerprinting, and pass usedforsecurity=False so md5 imports on FIPS builds. - fallback: skip the backoff sleep after the final attempt, which only delayed the degrade and emitted a retry event no retry followed. - fallback: stop replaying already-yielded stream deltas to the consumer; set replay_partial_stream=True for the historical behavior. - usage_meter: serialize wire counters as ints rather than floats. - protocols: document the tool-call veto contract. GuardrailsMiddleware only marks metadata["blocked"], and nothing in AgentCore dispatches tools, so an unchecked flag meant silent non-enforcement. Adds BLOCKED_KEY/BLOCK_REASON_KEY and ctx.block()/is_blocked/block_reason. - export the new classes from the providers/runtime/runtime.loop package roots, matching the existing re-export convention. Provider fixes (behavior changes — see below): - finish_reason: new shared normalizer. _runaway matches the literal "length" and deliberately has no token-count fallback once visible text is present, but Anthropic sends stop_reason="max_tokens" and the Responses API sends status="incomplete" plus incomplete_details.reason="max_output_tokens". Truncation recovery was dead for both protocols. Only truncation markers are rewritten; tool_use/end_turn/stop pass through unchanged. - openai_responses: also emit a finish_reason on response.incomplete, which the streaming path ignored while hardcoding "stop". - openai_chat: compute session affinity on every call. A scoped client withholds default_query from the SDK because cached clients outlive their task, making the per-call path the only source of affinity and the only staleness check — but it was gated on extra_headers, so chat(messages) pinned no replica at all. - _streaming: keep the visible answer when a streamed block list carries no text block. A text block is only recorded on a content_block_start of type text, so a gateway omitting that event returned thinking-only blocks and silently dropped the answer. Three of these change observable behavior rather than only fixing a latent bug, and matter to downstream repos pinning this commit: truncation continuation now fires for anthropic/bedrock and responses; scoped clients now send an affinity query on every request; and CooldownFallbackLLM.stream no longer duplicates partial output. test_cooldown_fallback's `assert sleeps == [0.5]` pinned the wasted post-final-attempt sleep and is updated to the corrected contract. Validation: Ruff clean, Pyright 0 errors, 1108 passed (was 1074). Co-Authored-By: Claude Opus 5 (1M context) --- agent_core/protocols.py | 24 +++ agent_core/providers/__init__.py | 39 ++++- agent_core/providers/anthropic.py | 7 +- agent_core/providers/aux_builder.py | 29 +++- agent_core/providers/fallback.py | 27 ++- agent_core/providers/finish_reason.py | 65 +++++++ agent_core/providers/openai_chat.py | 40 +++-- agent_core/providers/openai_responses.py | 18 +- agent_core/providers/summary.py | 28 ++- agent_core/runtime/__init__.py | 22 +++ agent_core/runtime/loop/__init__.py | 18 ++ agent_core/runtime/loop/_streaming.py | 30 ++++ agent_core/runtime/loop/guardrails.py | 56 ++++-- agent_core/runtime/loop/tool_call_repair.py | 35 +++- agent_core/runtime/usage_meter.py | 27 +-- docs/provider-substrate-boundary.md | 38 ++++ tests/test_aux_builder.py | 56 ++++++ tests/test_cooldown_fallback.py | 112 +++++++++++- tests/test_finish_reason_normalization.py | 126 ++++++++++++++ tests/test_session_affinity_per_call.py | 78 +++++++++ tests/test_stream_block_text_preservation.py | 82 +++++++++ tests/test_summary_engine.py | 173 +++++++++++++++++++ tests/test_tool_guardrails.py | 84 +++++++++ tests/test_usage_meter.py | 21 +++ 24 files changed, 1177 insertions(+), 58 deletions(-) create mode 100644 agent_core/providers/finish_reason.py create mode 100644 tests/test_finish_reason_normalization.py create mode 100644 tests/test_session_affinity_per_call.py create mode 100644 tests/test_stream_block_text_preservation.py diff --git a/agent_core/protocols.py b/agent_core/protocols.py index 9be32d8..56bed3f 100644 --- a/agent_core/protocols.py +++ b/agent_core/protocols.py @@ -98,6 +98,15 @@ def __post_init__(self) -> None: self.start_time = time.time() +# Reserved ``ToolCallContext.metadata`` keys. A ``before_tool_call`` +# middleware sets ``BLOCKED_KEY`` to veto a call; the host dispatcher MUST +# check it after running the middleware chain and, when set, skip the tool and +# return ``BLOCK_REASON_KEY`` to the model as the tool result. AgentCore does +# not dispatch tools itself, so an unchecked flag means silent non-enforcement. +BLOCKED_KEY = "blocked" +BLOCK_REASON_KEY = "block_reason" + + @dataclass class ToolCallContext: task_id: str @@ -107,6 +116,19 @@ class ToolCallContext: tool_args: dict[str, Any] = field(default_factory=dict[str, Any]) metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + def block(self, reason: str) -> None: + """Veto this tool call. See ``BLOCKED_KEY`` for the host contract.""" + self.metadata[BLOCKED_KEY] = True + self.metadata[BLOCK_REASON_KEY] = reason + + @property + def is_blocked(self) -> bool: + return bool(self.metadata.get(BLOCKED_KEY)) + + @property + def block_reason(self) -> str: + return str(self.metadata.get(BLOCK_REASON_KEY) or "") + class ExecutionMiddleware: async def before_phase(self, ctx: PhaseContext) -> PhaseContext: @@ -191,6 +213,8 @@ def reload(self) -> None: ... __all__ = [ + "BLOCKED_KEY", + "BLOCK_REASON_KEY", "EventReader", "EventSink", "ExecutionMiddleware", diff --git a/agent_core/providers/__init__.py b/agent_core/providers/__init__.py index a87bda3..181368b 100644 --- a/agent_core/providers/__init__.py +++ b/agent_core/providers/__init__.py @@ -1,7 +1,19 @@ """Provider transports and provider-neutral client wrappers.""" from agent_core.providers.anthropic import AnthropicClient -from agent_core.providers.fallback import FallbackEntry, LLMFallbackChain +from agent_core.providers.aux_builder import AuxLLMFactory +from agent_core.providers.fallback import ( + CooldownFallbackLLM, + FallbackEntry, + LLMFallbackChain, + legacy_retryable, +) +from agent_core.providers.finish_reason import ( + FINISH_REASON_LENGTH, + TRUNCATION_MARKERS, + normalize_finish_reason, + responses_finish_reason, +) from agent_core.providers.nonblocking_stream import NonBlockingStream, nonblocking_stderr from agent_core.providers.openai_chat import ( OpenAIClient, @@ -21,10 +33,25 @@ provider_label, thinking_format_for_protocol, ) +from agent_core.providers.summary import ( + EXTRACT_INFO_PROMPT, + SummaryCandidate, + SummaryLLMEngine, + build_summary_payload, + default_summary_retryable, + describe_summary_candidates, + normalize_summary_endpoint, + truncate_summary_fallback, +) __all__ = [ + "EXTRACT_INFO_PROMPT", + "FINISH_REASON_LENGTH", + "TRUNCATION_MARKERS", "AnthropicClient", "AnthropicPromptCacheAdapter", + "AuxLLMFactory", + "CooldownFallbackLLM", "FallbackEntry", "LLMFallbackChain", "NonBlockingStream", @@ -32,12 +59,22 @@ "OpenAIResponsesClient", "SessionQueryResolver", "SessionScopeResolver", + "SummaryCandidate", + "SummaryLLMEngine", "build_protocol_client", + "build_summary_payload", "configure_session_query_resolver", "configure_session_scope_resolver", + "default_summary_retryable", + "describe_summary_candidates", + "legacy_retryable", "maybe_wrap_for_prompt_cache", "nonblocking_stderr", + "normalize_finish_reason", + "normalize_summary_endpoint", "protocol_of", "provider_label", + "responses_finish_reason", "thinking_format_for_protocol", + "truncate_summary_fallback", ] diff --git a/agent_core/providers/anthropic.py b/agent_core/providers/anthropic.py index c1237bf..a0cb8c9 100644 --- a/agent_core/providers/anthropic.py +++ b/agent_core/providers/anthropic.py @@ -29,6 +29,7 @@ from agent_core.llm import LLMClient, LLMResponse, StreamDelta from agent_core.messages import Message, ToolCall, text_of +from agent_core.providers.finish_reason import normalize_finish_reason logger = logging.getLogger(__name__) @@ -297,7 +298,9 @@ async def stream( cache_write, reasoning_tokens, ), - finish_reason=stop_reason, + # ``max_tokens`` must reach the runaway/truncation checks as + # ``length``; other stop reasons pass through untouched. + finish_reason=normalize_finish_reason(stop_reason), model=model, reasoning_blocks=ordered if _has_thinking(ordered) else [], ) @@ -711,7 +714,7 @@ def _to_llm_response(raw: Any) -> LLMResponse: content=content, tool_calls=tool_calls, reasoning_content="\n".join(thinking_parts), - finish_reason=getattr(raw, "stop_reason", "") or "", + finish_reason=normalize_finish_reason(getattr(raw, "stop_reason", "")), model=getattr(raw, "model", "") or "", usage=usage_dict, response_metadata={"id": getattr(raw, "id", "")}, diff --git a/agent_core/providers/aux_builder.py b/agent_core/providers/aux_builder.py index 6138c91..44a752e 100644 --- a/agent_core/providers/aux_builder.py +++ b/agent_core/providers/aux_builder.py @@ -5,7 +5,7 @@ # pyright: basic import logging from collections.abc import Callable, Mapping -from typing import Any, cast +from typing import Any logger = logging.getLogger(__name__) @@ -35,6 +35,15 @@ def _resolve_api_key(section: Mapping[str, Any], provider: str) -> str: return "dummy" +def _as_int(value: object) -> int: + """Coerce a profile-supplied budget, tolerating numeric strings.""" + if isinstance(value, bool): + return int(value) + if isinstance(value, int): + return value + return int(float(str(value).strip().replace("_", ""))) + + def _thinking_budget(section: Mapping[str, Any]) -> object | None: for key in ("thinking_budget", "thinking_budget_tokens"): value = section.get(key) @@ -164,12 +173,18 @@ def _build_openai( extra_body = dict(extra_body_value) if isinstance(extra_body_value, Mapping) else {} template_value = extra_body.get("chat_template_kwargs") template = dict(template_value) if isinstance(template_value, Mapping) else {} - if section.get("enable_thinking"): - template["enable_thinking"] = True - template.setdefault("preserve_thinking", False) - budget = _thinking_budget(section) - if budget is not None: - template["thinking_budget"] = int(cast("Any", budget)) + # Distinguish "absent" from an explicit false: models such as Qwen3 and + # SGLang default thinking on, so a profile disabling it must emit the + # key rather than fall through silently. + enable_thinking = section.get("enable_thinking") + if enable_thinking is not None: + template["enable_thinking"] = bool(enable_thinking) + if enable_thinking: + template.setdefault("preserve_thinking", False) + if enable_thinking is None or enable_thinking: + budget = _thinking_budget(section) + if budget is not None: + template["thinking_budget"] = _as_int(budget) if template: extra_body["chat_template_kwargs"] = template if extra_body: diff --git a/agent_core/providers/fallback.py b/agent_core/providers/fallback.py index 348d2b9..6517458 100644 --- a/agent_core/providers/fallback.py +++ b/agent_core/providers/fallback.py @@ -70,6 +70,7 @@ "FallbackEntry", "FallbackTrigger", "LLMFallbackChain", + "legacy_retryable", "with_provider_stamp", ] @@ -200,6 +201,11 @@ class CooldownFallbackLLM: This preserves the legacy two-model fallback behavior used by both host products. Product tracing is an optional async event hook, keeping the retry state machine independent of execution scopes and trace backends. + + ``stream`` stops retrying the primary once any delta has reached the + consumer, because a retry would emit those deltas twice. Set + ``replay_partial_stream=True`` to restore the historical duplicating + behavior; prefer :class:`LLMFallbackChain` for rewind-safe semantics. """ def __init__( @@ -214,6 +220,7 @@ def __init__( clock: Callable[[], float] = time.time, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, jitter: Callable[[], float] = random.random, + replay_partial_stream: bool = False, ) -> None: self.primary = primary self.fallback = fallback @@ -225,6 +232,7 @@ def __init__( self._clock = clock self._sleep = sleep self._jitter = jitter + self._replay_partial_stream = replay_partial_stream self._cooldown_until = 0.0 @property @@ -295,6 +303,10 @@ async def chat( ) if not self._retryable(error): break + if attempt + 1 >= self.max_retries: + # Last attempt: sleeping here only delays the degrade and + # emits a ``retry`` event that no retry follows. + break delay = min(0.5 * (2**attempt), 8.0) + self._jitter() * 0.25 await self._emit("retry", attempt=attempt + 1, delay_s=delay) await self._sleep(delay) @@ -361,11 +373,20 @@ async def stream( streaming=True, error=str(error), ) - # The historical wrapper retried even after yielding. Preserve - # that behavior here; callers choosing rewind-safe semantics - # should use LLMFallbackChain instead. if not self._retryable(error): break + if yielded and not self._replay_partial_stream: + # Deltas already reached the consumer; retrying the primary + # would duplicate them. Degrade to the fallback leg instead. + await self._emit( + "abandon_stream_retry", + attempt=attempt + 1, + streaming=True, + yielded=True, + ) + break + if attempt + 1 >= self.max_retries: + break delay = min(0.5 * (2**attempt), 8.0) + self._jitter() * 0.25 await self._emit( "retry", diff --git a/agent_core/providers/finish_reason.py b/agent_core/providers/finish_reason.py new file mode 100644 index 0000000..289c1e2 --- /dev/null +++ b/agent_core/providers/finish_reason.py @@ -0,0 +1,65 @@ +"""Normalization of provider-specific completion-stop signals. + +``LLMResponse.finish_reason`` is a provider-neutral field, but each transport +names the "output token cap was hit" case differently: OpenAI Chat says +``length``, Anthropic Messages says ``max_tokens``, and the OpenAI Responses API +reports ``status="incomplete"`` with ``incomplete_details.reason +="max_output_tokens"``. + +``agent_core.runtime.loop._runaway`` tests for exactly ``"length"`` — it is the +only evidence that can distinguish a reply cut off mid-sentence from a complete +one, and there is deliberately no token-count fallback once visible text is +present. An unmapped marker therefore silently disables truncation recovery for +that transport, so every client funnels its stop signal through here. + +Only the truncation markers are rewritten. Every other value is passed through +unchanged, because ``tool_use``/``end_turn``/``stop`` carry provider-meaningful +detail that hosts and tests read directly. +""" + +from __future__ import annotations + +from typing import Any + +FINISH_REASON_LENGTH = "length" + +# Values, lowercased, that mean "the output token cap stopped generation". +TRUNCATION_MARKERS: frozenset[str] = frozenset({ + "length", # OpenAI Chat Completions (already normalized) + "max_completion_tokens", + "max_output_tokens", # OpenAI Responses incomplete_details.reason + "max_tokens", # Anthropic Messages stop_reason + "model_length", + "output_limit", +}) + + +def normalize_finish_reason(value: Any) -> str: + """Map a transport's truncation marker to ``"length"``; pass others through.""" + text = str(value or "").strip() + if not text: + return "" + if text.lower() in TRUNCATION_MARKERS: + return FINISH_REASON_LENGTH + return text + + +def responses_finish_reason(status: Any, incomplete_reason: Any = None) -> str: + """Fold a Responses-API ``status`` plus ``incomplete_details.reason``. + + ``status`` alone is never ``"length"`` — it is ``completed`` / ``incomplete`` + / ``failed`` — so a gpt-5 turn stopped at ``max_output_tokens`` is invisible + without the nested reason. + """ + reason = normalize_finish_reason(incomplete_reason) + if reason: + return reason + return normalize_finish_reason(status) + + +__all__ = [ + "FINISH_REASON_LENGTH", + "TRUNCATION_MARKERS", + "normalize_finish_reason", + "responses_finish_reason", +] diff --git a/agent_core/providers/openai_chat.py b/agent_core/providers/openai_chat.py index 03e78c9..2c9af0a 100644 --- a/agent_core/providers/openai_chat.py +++ b/agent_core/providers/openai_chat.py @@ -303,13 +303,19 @@ async def chat( kwargs["max_completion_tokens"] = eff_mt if extra_headers: kwargs["extra_headers"] = extra_headers - # EAS UCH affinity keys on the URL query parameter (the header - # alone never pins an upstream replica) — mirror the per-call - # session id into ``extra_query``; per-call wins over any - # construction-time ``default_query``. - session_query = self._session_query(extra_headers) - if session_query: - kwargs["extra_query"] = session_query + # EAS UCH affinity keys on the URL query parameter (the header alone + # never pins an upstream replica) — mirror the session id into + # ``extra_query``; per-call wins over construction-time affinity. + # + # Computed on EVERY call, not only when ``extra_headers`` is present: a + # scoped client deliberately withholds ``default_query`` from the SDK + # (see ``__init__``) because cached clients outlive their task, so this + # is the only path that re-supplies its affinity — and the only one that + # runs the staleness check. Gating it on ``extra_headers`` left + # ``chat(messages)`` pinning no replica at all. + session_query = self._session_query(extra_headers) + if session_query: + kwargs["extra_query"] = session_query extra_body = self._effective_extra_body() if extra_body: kwargs["extra_body"] = extra_body @@ -360,13 +366,19 @@ async def stream( kwargs["max_completion_tokens"] = eff_mt if extra_headers: kwargs["extra_headers"] = extra_headers - # EAS UCH affinity keys on the URL query parameter (the header - # alone never pins an upstream replica) — mirror the per-call - # session id into ``extra_query``; per-call wins over any - # construction-time ``default_query``. - session_query = self._session_query(extra_headers) - if session_query: - kwargs["extra_query"] = session_query + # EAS UCH affinity keys on the URL query parameter (the header alone + # never pins an upstream replica) — mirror the session id into + # ``extra_query``; per-call wins over construction-time affinity. + # + # Computed on EVERY call, not only when ``extra_headers`` is present: a + # scoped client deliberately withholds ``default_query`` from the SDK + # (see ``__init__``) because cached clients outlive their task, so this + # is the only path that re-supplies its affinity — and the only one that + # runs the staleness check. Gating it on ``extra_headers`` left + # ``chat(messages)`` pinning no replica at all. + session_query = self._session_query(extra_headers) + if session_query: + kwargs["extra_query"] = session_query extra_body = self._effective_extra_body() if extra_body: kwargs["extra_body"] = extra_body diff --git a/agent_core/providers/openai_responses.py b/agent_core/providers/openai_responses.py index ae49f60..4f26a37 100644 --- a/agent_core/providers/openai_responses.py +++ b/agent_core/providers/openai_responses.py @@ -37,6 +37,10 @@ from agent_core.llm import LLMClient, LLMResponse, StreamDelta from agent_core.messages import Message, ToolCall, text_of +from agent_core.providers.finish_reason import ( + normalize_finish_reason, + responses_finish_reason, +) logger = logging.getLogger(__name__) @@ -157,13 +161,16 @@ async def stream( "response.reasoning_text.delta", ): yield StreamDelta(reasoning_content=getattr(event, "delta", "") or "") - elif etype == "response.completed": + elif etype in ("response.completed", "response.incomplete"): resp = getattr(event, "response", None) usage = _responses_usage_dict(getattr(resp, "usage", None)) + reason = normalize_finish_reason( + _get(_get(resp, "incomplete_details", None) or {}, "reason", ""), + ) yield StreamDelta( usage=usage, model=getattr(resp, "model", "") or "", - finish_reason="stop", + finish_reason=reason or "stop", ) @@ -383,7 +390,12 @@ def _parse_responses_output(raw: Any) -> LLMResponse: content=content, tool_calls=tool_calls, reasoning_content="\n".join(summary_parts), - finish_reason=_get(raw, "status", "") or "", + # ``status`` is completed/incomplete/failed — never ``length`` — so the + # nested ``incomplete_details.reason`` carries the truncation signal. + finish_reason=responses_finish_reason( + _get(raw, "status", ""), + _get(_get(raw, "incomplete_details", None) or {}, "reason", ""), + ), model=_get(raw, "model", "") or "", usage=_responses_usage_dict(_get(raw, "usage", None)), response_metadata={"id": _get(raw, "id", "")}, diff --git a/agent_core/providers/summary.py b/agent_core/providers/summary.py index b417a00..858597a 100644 --- a/agent_core/providers/summary.py +++ b/agent_core/providers/summary.py @@ -23,6 +23,19 @@ class SummaryCandidate(TypedDict): type UsageRecorder = Callable[..., None] +type RetryPredicate = Callable[[Exception], bool] + +# 4xx statuses that never succeed on retry. 408/409/425/429 are excluded +# because those are explicitly transient. +PERMANENT_STATUS_CODES: frozenset[int] = frozenset({ + 400, 401, 402, 403, 404, 405, 406, 410, 413, 414, 415, 422, 501, +}) + + +def default_summary_retryable(error: Exception) -> bool: + """Retry transport faults and 5xx, but not permanent request errors.""" + status = getattr(getattr(error, "response", None), "status_code", None) + return not (isinstance(status, int) and status in PERMANENT_STATUS_CODES) EXTRACT_INFO_PROMPT = """You are given a piece of content and the requirement of information to extract. Your task is to extract the information specifically requested. Be precise and focus exclusively on the requested information. @@ -113,6 +126,7 @@ def __init__( fallback_limit: int = 20_000, usage_recorder: UsageRecorder | None = None, sleep: Callable[[float], Any] = asyncio.sleep, + retryable: RetryPredicate = default_summary_retryable, ) -> None: self.max_retries = max_retries self.truncate_step = truncate_step @@ -120,6 +134,7 @@ def __init__( self.fallback_limit = fallback_limit self.usage_recorder = usage_recorder self.sleep = sleep + self.retryable = retryable async def summarize( self, @@ -178,7 +193,10 @@ async def _summarize_one( or "longer than the model's context length" in body ): remove = self.truncate_step * (attempt + 1) - if remove >= len(current_content): + # ``remove`` indexes into the original ``content``, so the + # exhaustion test must measure ``content`` too; comparing + # against the already-shortened text bails a step early. + if remove >= len(content): return "" current_content = content[:-remove] + "[...truncated]" payload["messages"][0]["content"] = EXTRACT_INFO_PROMPT.format( @@ -204,6 +222,12 @@ async def _summarize_one( return "" except Exception as error: logger.warning("Summary LLM attempt %d failed: %s", attempt + 1, error) + if not self.retryable(error): + logger.warning( + "Summary candidate %d failed permanently; not retrying", + candidate_index + 1, + ) + return "" if attempt + 1 < self.max_retries: await self.sleep(float(attempt + 1)) return "" @@ -227,9 +251,11 @@ def _record_usage( __all__ = [ "EXTRACT_INFO_PROMPT", + "PERMANENT_STATUS_CODES", "SummaryCandidate", "SummaryLLMEngine", "build_summary_payload", + "default_summary_retryable", "describe_summary_candidates", "normalize_summary_endpoint", "truncate_summary_fallback", diff --git a/agent_core/runtime/__init__.py b/agent_core/runtime/__init__.py index 50b41d7..707ce0d 100644 --- a/agent_core/runtime/__init__.py +++ b/agent_core/runtime/__init__.py @@ -7,12 +7,34 @@ from_config_map, from_execution_policy, ) +from agent_core.runtime.usage_meter import ( + ExternalAPIMeter, + bind_usage_meter, + close_meter_span, + get_usage_meter, + open_meter_span, + record_api_request, + record_llm_usage, + record_tool_call, + reset_usage_meter, + set_meter_gauge, +) __all__ = [ "EventBus", + "ExternalAPIMeter", "Handler", "SpillStore", "ToolPermissionContext", + "bind_usage_meter", + "close_meter_span", "from_config_map", "from_execution_policy", + "get_usage_meter", + "open_meter_span", + "record_api_request", + "record_llm_usage", + "record_tool_call", + "reset_usage_meter", + "set_meter_gauge", ] diff --git a/agent_core/runtime/loop/__init__.py b/agent_core/runtime/loop/__init__.py index 7ab7d38..de43ef7 100644 --- a/agent_core/runtime/loop/__init__.py +++ b/agent_core/runtime/loop/__init__.py @@ -10,6 +10,11 @@ DefaultMessageCompactor, ) from agent_core.runtime.loop.compact_llm import LLMSummaryCompactor +from agent_core.runtime.loop.guardrails import ( + DEFAULT_DUPLICATE_THRESHOLDS, + GuardrailsMiddleware, + check_budget_exhausted, +) from agent_core.runtime.loop.llm_client import ( RUNAWAY_STATE_KEY, TRUNCATION_CONTINUATION_GUIDANCE, @@ -54,6 +59,12 @@ DefaultToolCallParser, MultiFormatToolCallParser, ) +from agent_core.runtime.loop.tool_call_repair import ( + DEFAULT_KEY_ALIASES, + LITERAL_CONTENT_KEYS, + ToolCallRepairMiddleware, + repair_truncated_json, +) from agent_core.runtime.loop.tool_exec import ( DefaultToolResultPostProcessor, ToolExecutionHooks, @@ -62,7 +73,10 @@ __all__ = [ "COMPACTION_TRIGGER_RATIO", + "DEFAULT_DUPLICATE_THRESHOLDS", + "DEFAULT_KEY_ALIASES", "DEFAULT_TRIGGER_RATIO", + "LITERAL_CONTENT_KEYS", "RUNAWAY_STATE_KEY", "TRUNCATION_CONTINUATION_GUIDANCE", "AgentLoopHooks", @@ -71,6 +85,7 @@ "DefaultThinkingParser", "DefaultToolCallParser", "DefaultToolResultPostProcessor", + "GuardrailsMiddleware", "HistoryPolicy", "InputTokenGauge", "InputTokenThresholdPolicy", @@ -87,12 +102,14 @@ "TaskBoundaryTrimmer", "ThinkTagSplitter", "TieredCompactor", + "ToolCallRepairMiddleware", "ToolExecutionHooks", "bind_max_tokens", "bind_session_id", "bind_temperature", "bind_tools", "call_llm", + "check_budget_exhausted", "check_context_budget", "compaction_trigger_tokens", "configure_model_registry", @@ -102,6 +119,7 @@ "extract_model_name", "extract_usage", "is_truncated_with_text", + "repair_truncated_json", "run_agent_loop", "usage_input_tokens", "usage_output_tokens", diff --git a/agent_core/runtime/loop/_streaming.py b/agent_core/runtime/loop/_streaming.py index 3643867..53a1aaa 100644 --- a/agent_core/runtime/loop/_streaming.py +++ b/agent_core/runtime/loop/_streaming.py @@ -171,6 +171,14 @@ def flush(self) -> tuple[str, str]: return (remainder, "") +def _has_text_block(blocks: list[dict[str, Any]]) -> bool: + """True when a streamed block list carries usable visible text.""" + return any( + block.get("type") == "text" and str(block.get("text") or "").strip() + for block in blocks + ) + + async def _stream_llm_response( llm: Any, messages: list[Message], @@ -320,6 +328,28 @@ def _assembled_response() -> LLMResponse: # lstrip/think-tag normalisation above is deliberately NOT applied to # them — replay requires the bytes Anthropic signed, unmodified. assembled_content: Any = final_reasoning_blocks or visible_content + # ...unless the blocks turn out NOT to hold that text. A text block is + # only recorded when a ``content_block_start`` of type ``text`` is seen + # (see ``AnthropicClient.stream``); a ``text_delta`` for an index that + # never opened reaches ``accumulated`` but no block. A gateway that + # omits or renames that event while thinking blocks arrive normally + # would then hand back thinking-only blocks and silently drop the + # model's answer. Append it as a text block instead: thinking + # signatures stay byte-exact, and text blocks carry no signature. + if ( + final_reasoning_blocks + and visible_content.strip() + and not _has_text_block(final_reasoning_blocks) + ): + logger.warning( + "streamed block list carried no text block but %d chars of " + "visible content arrived; appending it to preserve the answer", + len(visible_content), + ) + assembled_content = [ + *final_reasoning_blocks, + {"type": "text", "text": visible_content}, + ] return LLMResponse( content=assembled_content, tool_calls=complete_tool_calls, diff --git a/agent_core/runtime/loop/guardrails.py b/agent_core/runtime/loop/guardrails.py index f559b17..ff1f452 100644 --- a/agent_core/runtime/loop/guardrails.py +++ b/agent_core/runtime/loop/guardrails.py @@ -33,12 +33,22 @@ def _args_fingerprint(tool_name: str, args: Mapping[str, Any]) -> str: - raw = json.dumps({"t": tool_name, "a": args}, sort_keys=True) - return hashlib.md5(raw.encode()).hexdigest()[:12] + # ``default=str`` keeps a non-JSON argument (set, dataclass, Path) from + # turning a guardrail check into an unhandled TypeError mid tool call. + raw = json.dumps({"t": tool_name, "a": args}, sort_keys=True, default=str) + # Fingerprinting only; ``usedforsecurity=False`` keeps this importable on + # FIPS-enabled builds where plain md5 is unavailable. + return hashlib.md5(raw.encode(), usedforsecurity=False).hexdigest()[:12] class GuardrailsMiddleware(ExecutionMiddleware): - """Block pathological repetition while allowing hosts to tune tool policy.""" + """Block pathological repetition while allowing hosts to tune tool policy. + + This middleware only *marks* a call via :meth:`ToolCallContext.block`. + The host dispatcher must honour ``ctx.is_blocked`` after the middleware + chain and feed ``ctx.block_reason`` back to the model instead of executing + the tool; otherwise the guardrail is inert. + """ def __init__( self, @@ -106,11 +116,10 @@ async def before_tool_call(self, ctx: ToolCallContext) -> ToolCallContext: ) if consecutive >= threshold: self._stats["duplicate_blocks"] += 1 - ctx.metadata["blocked"] = True - ctx.metadata["block_reason"] = ( + ctx.block( f"Blocked: {tool_name} called {consecutive + 1} times " "consecutively with identical arguments. Try different " - "parameters or a different approach." + "parameters or a different approach.", ) return ctx recent.append(fingerprint) @@ -131,11 +140,10 @@ async def before_tool_call(self, ctx: ToolCallContext) -> ToolCallContext: hint_count = self._loop_hint_counts.get(task_id, 0) if hint_count >= self._max_loop_hints and fingerprint in list(recent)[:-1]: self._stats["loop_escalation_blocks"] += 1 - ctx.metadata["blocked"] = True - ctx.metadata["block_reason"] = ( + ctx.block( "Blocked: repeated tool call pattern detected after " f"{hint_count} loop warnings. Try a completely different " - "approach or conclude with available evidence." + "approach or conclude with available evidence.", ) return ctx @@ -146,6 +154,25 @@ def cleanup_task(self, task_id: str) -> None: self._loop_hint_counts.pop(task_id, None) +DEFAULT_MAX_TOKENS = 500_000 + + +def _as_int(value: object, default: int) -> int: + """Coerce untyped product config, falling back instead of raising.""" + if isinstance(value, bool) or value is None: + return default + if isinstance(value, int): + return value + if isinstance(value, float): + return int(value) + if isinstance(value, str): + try: + return int(float(value.strip().replace("_", ""))) + except ValueError: + return default + return default + + def check_budget_exhausted( working_state: Mapping[str, Any], token_usage: Mapping[str, int] | None, @@ -171,9 +198,13 @@ def check_budget_exhausted( else budget ) if token_usage: - maximum_value: object = allocated.get("max_tokens", 500_000) - maximum = int(maximum_value) if isinstance(maximum_value, int | str) else 500_000 - used = token_usage.get("total", 0) + maximum = _as_int(allocated.get("max_tokens"), DEFAULT_MAX_TOKENS) + # Providers and ``UsageMetadata`` emit ``total_tokens``; ``total`` is + # accepted for host mappings that use the shorter name. + used = _as_int( + token_usage.get("total_tokens", token_usage.get("total")), + 0, + ) if used >= maximum: return ( f"Token budget exhausted ({used:,}/{maximum:,}). " @@ -184,6 +215,7 @@ def check_budget_exhausted( __all__ = [ "DEFAULT_DUPLICATE_THRESHOLDS", + "DEFAULT_MAX_TOKENS", "GuardrailsMiddleware", "check_budget_exhausted", ] diff --git a/agent_core/runtime/loop/tool_call_repair.py b/agent_core/runtime/loop/tool_call_repair.py index 159547c..8299143 100644 --- a/agent_core/runtime/loop/tool_call_repair.py +++ b/agent_core/runtime/loop/tool_call_repair.py @@ -48,6 +48,22 @@ "grep_search": {"context_lines": int, "max_results": int}, } +# Argument names whose value is literal file/document content. Leading indent +# and trailing newlines are semantically significant for exact-match editors, +# so these are never whitespace-stripped. Keys are matched across every tool +# because host tool catalogs reuse the same names. +LITERAL_CONTENT_KEYS: frozenset[str] = frozenset({ + "content", + "file_text", + "new_str", + "new_string", + "old_str", + "old_string", + "patch", + "replacement", + "text", +}) + def repair_truncated_json(raw: str) -> str | None: """Close simple truncated strings/brackets and return valid JSON text.""" @@ -128,9 +144,21 @@ def __init__( key_aliases: Mapping[str, Mapping[str, str]] | None = None, type_coercions: Mapping[str, Mapping[str, type[Any]]] | None = None, tool_repair: ToolRepair = _default_tool_repair, + literal_content_keys: frozenset[str] | None = None, ) -> None: - self._key_aliases = key_aliases or DEFAULT_KEY_ALIASES - self._type_coercions = type_coercions or DEFAULT_TYPE_COERCIONS + # ``is None`` rather than falsiness: an empty mapping is a host + # explicitly disabling the default table, not a request for it. + self._key_aliases = ( + DEFAULT_KEY_ALIASES if key_aliases is None else key_aliases + ) + self._type_coercions = ( + DEFAULT_TYPE_COERCIONS if type_coercions is None else type_coercions + ) + self._literal_content_keys = ( + LITERAL_CONTENT_KEYS + if literal_content_keys is None + else literal_content_keys + ) self._tool_repair = tool_repair self._stats = { "key_renames": 0, @@ -158,6 +186,8 @@ async def before_tool_call(self, ctx: ToolCallContext) -> ToolCallContext: except (TypeError, ValueError): pass for key, value in args.items(): + if key in self._literal_content_keys: + continue if isinstance(value, str) and value != value.strip(): args[key] = value.strip() self._stats["whitespace_strips"] += 1 @@ -168,6 +198,7 @@ async def before_tool_call(self, ctx: ToolCallContext) -> ToolCallContext: __all__ = [ "DEFAULT_KEY_ALIASES", "DEFAULT_TYPE_COERCIONS", + "LITERAL_CONTENT_KEYS", "ToolCallRepairMiddleware", "repair_truncated_json", ] diff --git a/agent_core/runtime/usage_meter.py b/agent_core/runtime/usage_meter.py index f714297..7204a0d 100644 --- a/agent_core/runtime/usage_meter.py +++ b/agent_core/runtime/usage_meter.py @@ -22,6 +22,18 @@ ) +def _empty_slot() -> dict[str, Any]: + """A provider view for gauge/span-only providers: counts stay integral.""" + return dict.fromkeys(_BASE_FIELDS, 0) + + +def _render(field: str, value: float) -> Any: + """Report the four wire counters as ints; durations keep 2 decimals.""" + if field in _BASE_FIELDS: + return int(value) + return round(value, 2) + + class ExternalAPIMeter: """Accumulate external request, tool-call, gauge, and span measurements.""" @@ -98,27 +110,18 @@ def snapshot(self) -> dict[str, Any]: with self._lock: now = time.monotonic() apis: dict[str, dict[str, Any]] = { - provider: { - key: round(value, 2) if isinstance(value, float) else value - for key, value in slot.items() - } + provider: {key: _render(key, value) for key, value in slot.items()} for provider, slot in self._providers.items() } for (provider, _key), (started, field) in self._open_spans.items(): - view = apis.setdefault( - provider, - dict.fromkeys(_BASE_FIELDS, 0.0), - ) + view = apis.setdefault(provider, _empty_slot()) view[field] = round( float(view.get(field, 0.0)) + now - started, 2, ) view["spans_open"] = int(view.get("spans_open", 0)) + 1 for (provider, field), value in self._gauges.items(): - view = apis.setdefault( - provider, - dict.fromkeys(_BASE_FIELDS, 0.0), - ) + view = apis.setdefault(provider, _empty_slot()) view[field] = round(value, 2) return { "external_apis": apis, diff --git a/docs/provider-substrate-boundary.md b/docs/provider-substrate-boundary.md index 33ec875..105b37f 100644 --- a/docs/provider-substrate-boundary.md +++ b/docs/provider-substrate-boundary.md @@ -29,3 +29,41 @@ messages, and the rest of the runtime remain strict. Task pause checks follow the same rule: AgentCore owns safe polling semantics, while a host injects the task-status loader and its missing-task exception. + +## Tool-call middleware contract + +`GuardrailsMiddleware` and `ToolCallRepairMiddleware` are `before_tool_call` +middleware. AgentCore does not dispatch tools, so enforcement is the host's: + +- After running the middleware chain, check `ctx.is_blocked`. When set, do not + execute the tool — return `ctx.block_reason` to the model as the tool result. + The reserved metadata keys are `protocols.BLOCKED_KEY` and + `protocols.BLOCK_REASON_KEY`; set them through `ctx.block(reason)`. +- Call `GuardrailsMiddleware.cleanup_task(task_id)` when a task ends. The + middleware keeps per-task fingerprint, search-count, and loop-hint state that + is only released there. +- `ToolCallRepairMiddleware` strips surrounding whitespace from string + arguments, except keys in `LITERAL_CONTENT_KEYS` (file contents, str-replace + needles). Hosts whose tools carry literal text under other names must extend + that set, or exact-match edits will be corrupted. +- Passing `key_aliases={}` or `type_coercions={}` disables the default table; + omit the argument to inherit it. + +## Completion-stop normalization + +`agent_core.runtime.loop._runaway` detects a reply the output cap cut off by +matching `LLMResponse.finish_reason == "length"` exactly, and deliberately has +no token-count fallback once visible text is present — an explicit +`finish_reason` is the only evidence that can carry that case. Each transport +names it differently, so every client routes its stop signal through +`providers.finish_reason.normalize_finish_reason`: + +| Transport | Raw signal | Normalized | +|---|---|---| +| OpenAI Chat | `finish_reason="length"` | `length` | +| Anthropic Messages | `stop_reason="max_tokens"` | `length` | +| OpenAI Responses | `status="incomplete"` + `incomplete_details.reason="max_output_tokens"` | `length` | + +Only truncation markers are rewritten. `tool_use`, `end_turn`, and `stop` pass +through unchanged because hosts read them directly. A new transport that skips +this normalization silently disables truncation recovery for its protocol. diff --git a/tests/test_aux_builder.py b/tests/test_aux_builder.py index dadca1d..6516ae7 100644 --- a/tests/test_aux_builder.py +++ b/tests/test_aux_builder.py @@ -79,3 +79,59 @@ def test_native_anthropic_rejects_non_claude_model() -> None: ) with pytest.raises(ValueError, match="requires a Claude model"): factory.build({"provider": "anthropic", "model": "qwen"}) + + +def test_explicit_enable_thinking_false_is_emitted() -> None: + """Qwen3/SGLang default thinking on: disabling it must send the key.""" + calls: list[tuple[str, dict[str, Any]]] = [] + factory = AuxLLMFactory( + openai_factory=_capture("openai", calls), + anthropic_factory=_capture("anthropic", calls), + provider_type=lambda _provider: "openai_compat", + ) + client = factory.build({ + "provider": "gateway", + "model": "qwen3", + "api_key": "key", + "enable_thinking": False, + # A budget without an enable flag must not leak through. + "thinking_budget": 512, + }) + template = client["extra_body"]["chat_template_kwargs"] + assert template == {"enable_thinking": False} + + +def test_absent_enable_thinking_leaves_the_key_unset() -> None: + calls: list[tuple[str, dict[str, Any]]] = [] + factory = AuxLLMFactory( + openai_factory=_capture("openai", calls), + anthropic_factory=_capture("anthropic", calls), + provider_type=lambda _provider: "openai_compat", + ) + client = factory.build({ + "provider": "gateway", + "model": "qwen3", + "api_key": "key", + "thinking_budget": 512, + }) + template = client["extra_body"]["chat_template_kwargs"] + assert "enable_thinking" not in template + assert template["thinking_budget"] == 512 + + +def test_numeric_string_thinking_budget_is_coerced() -> None: + calls: list[tuple[str, dict[str, Any]]] = [] + factory = AuxLLMFactory( + openai_factory=_capture("openai", calls), + anthropic_factory=_capture("anthropic", calls), + provider_type=lambda _provider: "openai_compat", + ) + client = factory.build({ + "provider": "gateway", + "model": "qwen3", + "api_key": "key", + "enable_thinking": True, + "thinking": {"budget_tokens": "1024"}, + }) + template = client["extra_body"]["chat_template_kwargs"] + assert template["thinking_budget"] == 1024 diff --git a/tests/test_cooldown_fallback.py b/tests/test_cooldown_fallback.py index aa5302b..a254a0a 100644 --- a/tests/test_cooldown_fallback.py +++ b/tests/test_cooldown_fallback.py @@ -56,7 +56,9 @@ async def sleep(delay: float) -> None: assert (await llm.chat([user_msg("two")])).content == "ok" assert primary.calls == 1 assert fallback.calls == 2 - assert sleeps == [0.5] + # max_retries=1 means a single attempt: no backoff sleep, since no retry + # follows it. See test_no_backoff_sleep_after_the_final_attempt. + assert sleeps == [] @pytest.mark.asyncio @@ -87,3 +89,111 @@ def test_legacy_retryable_contract() -> None: assert legacy_retryable(TimeoutError("timed out")) assert legacy_retryable(AttributeError("model_dump")) assert not legacy_retryable(ValueError("invalid json")) + + +@pytest.mark.asyncio +async def test_no_backoff_sleep_after_the_final_attempt() -> None: + """The last attempt must not sleep — no retry follows it.""" + primary = ScriptedLLM("primary", [TimeoutError("timeout")]) + fallback = ScriptedLLM("fallback", [LLMResponse(content="ok")]) + sleeps: list[float] = [] + events: list[str] = [] + + async def sleep(delay: float) -> None: + sleeps.append(delay) + + async def hook(name: str, _payload: dict[str, object]) -> None: + events.append(name) + + llm = CooldownFallbackLLM( + primary, + fallback, + max_retries=1, + cooldown_seconds=60, + clock=lambda: 0.0, + sleep=sleep, + jitter=lambda: 0.0, + event_hook=hook, + ) + assert (await llm.chat([user_msg("one")])).content == "ok" + assert sleeps == [] + assert "retry" not in events + + # Two attempts sleep exactly once, between them. + primary2 = ScriptedLLM( + "primary", + [TimeoutError("timeout"), TimeoutError("timeout")], + ) + sleeps.clear() + llm2 = CooldownFallbackLLM( + primary2, + ScriptedLLM("fallback", [LLMResponse(content="ok")]), + max_retries=2, + clock=lambda: 0.0, + sleep=sleep, + jitter=lambda: 0.0, + ) + assert (await llm2.chat([user_msg("two")])).content == "ok" + assert sleeps == [0.5] + assert primary2.calls == 2 + + +class PartialStream: + model = "primary" + + def __init__(self) -> None: + self.starts = 0 + + async def stream( + self, + _messages: list[Message], + **_kwargs: object, + ) -> AsyncIterator[StreamDelta]: + self.starts += 1 + yield StreamDelta(content="part1 ") + yield StreamDelta(content="part2 ") + raise TimeoutError("timeout mid-stream") + + +class OkStream: + model = "fallback" + + async def stream( + self, + _messages: list[Message], + **_kwargs: object, + ) -> AsyncIterator[StreamDelta]: + yield StreamDelta(content="FALLBACK") + + +@pytest.mark.asyncio +async def test_stream_does_not_replay_already_yielded_deltas() -> None: + primary = PartialStream() + llm = CooldownFallbackLLM( + primary, + OkStream(), + max_retries=3, + clock=lambda: 0.0, + sleep=lambda _d: _completed(), + jitter=lambda: 0.0, + ) + seen = [delta.content async for delta in llm.stream([user_msg("x")])] + assert seen == ["part1 ", "part2 ", "FALLBACK"] + assert primary.starts == 1 + + +@pytest.mark.asyncio +async def test_stream_replay_is_opt_in() -> None: + primary = PartialStream() + llm = CooldownFallbackLLM( + primary, + OkStream(), + max_retries=2, + clock=lambda: 0.0, + sleep=lambda _d: _completed(), + jitter=lambda: 0.0, + replay_partial_stream=True, + ) + seen = [delta.content async for delta in llm.stream([user_msg("x")])] + assert seen == ["part1 ", "part2 ", "part1 ", "part2 ", "FALLBACK"] + assert primary.starts == 2 diff --git a/tests/test_finish_reason_normalization.py b/tests/test_finish_reason_normalization.py new file mode 100644 index 0000000..dcd6062 --- /dev/null +++ b/tests/test_finish_reason_normalization.py @@ -0,0 +1,126 @@ +"""Truncation signals must reach the runaway/continuation checks as ``length``.""" + +from __future__ import annotations + +from typing import Any + +from agent_core.llm import LLMResponse +from agent_core.providers.anthropic import _to_llm_response +from agent_core.providers.finish_reason import ( + normalize_finish_reason, + responses_finish_reason, +) +from agent_core.providers.openai_responses import _parse_responses_output +from agent_core.runtime.loop._runaway import ( + _is_runaway_response, + is_truncated_with_text, +) + + +def test_only_truncation_markers_are_rewritten() -> None: + assert normalize_finish_reason("max_tokens") == "length" + assert normalize_finish_reason("max_output_tokens") == "length" + assert normalize_finish_reason("length") == "length" + # Provider-meaningful values are preserved: hosts and tests read them. + assert normalize_finish_reason("tool_use") == "tool_use" + assert normalize_finish_reason("end_turn") == "end_turn" + assert normalize_finish_reason("stop") == "stop" + assert normalize_finish_reason("") == "" + assert normalize_finish_reason(None) == "" + + +def test_responses_status_folds_the_nested_reason() -> None: + # ``status`` alone can never express truncation. + assert responses_finish_reason("incomplete", "max_output_tokens") == "length" + assert responses_finish_reason("completed", None) == "completed" + assert responses_finish_reason("incomplete", "content_filter") == "content_filter" + + +class _Raw: + def __init__(self, **kw: Any) -> None: + self.__dict__.update(kw) + + +def test_anthropic_max_tokens_reaches_the_truncation_check() -> None: + raw = _Raw( + content=[_Raw(type="text", text="a half sentence that stops")], + stop_reason="max_tokens", + model="claude-x", + usage=None, + id="m1", + ) + response = _to_llm_response(raw) + assert response.finish_reason == "length" + # The whole point: recovery now fires for anthropic/bedrock. + assert is_truncated_with_text(response) is True + + +def test_anthropic_tool_use_is_still_passed_through() -> None: + raw = _Raw( + content=[_Raw(type="text", text="calling")], + stop_reason="tool_use", + model="claude-x", + usage=None, + id="m2", + ) + assert _to_llm_response(raw).finish_reason == "tool_use" + + +def test_responses_incomplete_reaches_the_truncation_check() -> None: + raw = { + "output": [{ + "type": "message", + "content": [{"type": "output_text", "text": "cut off here"}], + }], + "status": "incomplete", + "incomplete_details": {"reason": "max_output_tokens"}, + "model": "gpt-5", + "id": "r1", + } + response = _parse_responses_output(raw) + assert response.finish_reason == "length" + assert is_truncated_with_text(response) is True + + +def test_responses_empty_truncated_reply_is_a_runaway() -> None: + raw = { + "output": [], + "status": "incomplete", + "incomplete_details": {"reason": "max_output_tokens"}, + "model": "gpt-5", + "id": "r2", + } + response = _parse_responses_output(raw) + assert response.finish_reason == "length" + assert _is_runaway_response(response) is True + + +def test_completed_response_is_not_treated_as_truncated() -> None: + raw = { + "output": [{ + "type": "message", + "content": [{"type": "output_text", "text": "a full answer."}], + }], + "status": "completed", + "model": "gpt-5", + "id": "r3", + } + response = _parse_responses_output(raw) + assert response.finish_reason == "completed" + assert is_truncated_with_text(response) is False + assert isinstance(response, LLMResponse) + + +def test_unmapped_marker_would_be_invisible() -> None: + """Documents the coupling this normalization exists to satisfy. + + ``_runaway`` matches the literal string, so a raw provider marker reaching + ``LLMResponse`` unmapped disables truncation recovery entirely. + """ + raw_marker = LLMResponse(content="cut off", finish_reason="max_tokens") + assert is_truncated_with_text(raw_marker) is False + normalized = LLMResponse( + content="cut off", + finish_reason=normalize_finish_reason("max_tokens"), + ) + assert is_truncated_with_text(normalized) is True diff --git a/tests/test_session_affinity_per_call.py b/tests/test_session_affinity_per_call.py new file mode 100644 index 0000000..a8c6276 --- /dev/null +++ b/tests/test_session_affinity_per_call.py @@ -0,0 +1,78 @@ +"""A scoped client must pin a replica even when the caller sends no headers.""" + +from __future__ import annotations + +from typing import Any + +from agent_core.providers.openai_chat import OpenAIClient + + +def _client(scope: str) -> OpenAIClient: + return OpenAIClient( + "m", + api_key="k", + base_url="https://host/v1", + default_headers={"x-session-id": "sess-1"}, + # A scoped resolver makes __init__ withhold ``default_query`` from the + # SDK, so the per-call path is the only source of affinity. + session_query_resolver=lambda headers: ( + {"session": headers["x-session-id"]} + if headers and "x-session-id" in headers + else {} + ), + session_scope_resolver=lambda: scope, + ) + + +def test_scoped_client_withholds_sdk_default_query() -> None: + client = _client("task-1") + assert client._client._custom_query in (None, {}) + + +def test_affinity_is_supplied_without_extra_headers() -> None: + """The regression: gating on ``extra_headers`` pinned no replica at all.""" + client = _client("task-1") + query = client._session_query(None) + assert query == {"session": "sess-1"} + + +def test_per_call_headers_win_over_construction_time() -> None: + client = _client("task-1") + query = client._session_query({"x-session-id": "sess-2"}) + assert query == {"session": "sess-2"} + + +def test_stale_scope_drops_construction_time_affinity() -> None: + client = _client("task-1") + # The task that built the client is gone; the cached client outlived it. + client._session_scope_resolver = lambda: "task-2" + assert client._session_query(None) == {} + + +def test_kwargs_carry_extra_query_with_no_extra_headers( + monkeypatch: Any, +) -> None: + """End-to-end: the request kwargs actually receive ``extra_query``.""" + client = _client("task-1") + captured: dict[str, Any] = {} + + async def fake_create(**kwargs: Any) -> Any: + captured.update(kwargs) + raise RuntimeError("stop after capture") + + monkeypatch.setattr( + client._client.chat.completions, + "create", + fake_create, + ) + import asyncio + + import pytest + + from agent_core.messages import user_msg + + with pytest.raises(Exception, match="stop after capture"): + asyncio.run(client.chat([user_msg("hi")])) + # The caller passed no session header, yet the request still pins a replica. + assert captured["extra_query"] == {"session": "sess-1"} + assert "x-session-id" not in captured.get("extra_headers", {}) diff --git a/tests/test_stream_block_text_preservation.py b/tests/test_stream_block_text_preservation.py new file mode 100644 index 0000000..af31919 --- /dev/null +++ b/tests/test_stream_block_text_preservation.py @@ -0,0 +1,82 @@ +"""A thinking-only block list must not swallow the model's visible answer.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncIterator +from typing import Any + +from agent_core.llm import StreamDelta +from agent_core.messages import Message, user_msg +from agent_core.runtime.loop._streaming import _stream_llm_response + + +class _Client: + """Emits text deltas plus a block list that omits the text block. + + Mirrors a gateway that renames/omits the ``content_block_start`` of type + ``text`` while thinking blocks arrive normally: ``AnthropicClient.stream`` + only records a text block on that event, so the delta text reaches + ``accumulated`` but never the block list. + """ + + def __init__(self, blocks: list[dict[str, Any]]) -> None: + self._blocks = blocks + + async def stream( + self, + _messages: list[Message], + **_kwargs: Any, + ) -> AsyncIterator[StreamDelta]: + yield StreamDelta(content="The answer ") + yield StreamDelta(content="is 42.") + yield StreamDelta( + reasoning_blocks=self._blocks, + finish_reason="end_turn", + model="claude-x", + ) + + +async def _noop(*_args: Any, **_kwargs: Any) -> None: + return None + + +def _run(blocks: list[dict[str, Any]]) -> Any: + return asyncio.run( + _stream_llm_response( + _Client(blocks), + [user_msg("q")], + timeout=5.0, + on_delta=_noop, + ), + ) + + +THINKING = {"type": "thinking", "thinking": "hmm", "signature": "sig-abc"} + + +def test_thinking_only_blocks_keep_the_visible_answer() -> None: + response = _run([THINKING]) + assert isinstance(response.content, list) + # Thinking block preserved byte-exact for replay... + assert response.content[0] == THINKING + # ...and the answer is no longer lost. + assert response.content[-1] == {"type": "text", "text": "The answer is 42."} + + +def test_blocks_that_already_carry_text_are_untouched() -> None: + blocks = [THINKING, {"type": "text", "text": "The answer is 42."}] + response = _run(blocks) + assert response.content == blocks + # No duplicated text block. + assert sum(b.get("type") == "text" for b in response.content) == 1 + + +def test_empty_text_block_does_not_count_as_carrying_text() -> None: + response = _run([THINKING, {"type": "text", "text": " "}]) + assert response.content[-1] == {"type": "text", "text": "The answer is 42."} + + +def test_no_blocks_still_yields_the_flat_string() -> None: + response = _run([]) + assert response.content == "The answer is 42." diff --git a/tests/test_summary_engine.py b/tests/test_summary_engine.py index 2577519..66d3fcf 100644 --- a/tests/test_summary_engine.py +++ b/tests/test_summary_engine.py @@ -38,3 +38,176 @@ def test_candidate_description_redacts_api_key() -> None: def test_truncate_fallback() -> None: assert truncate_summary_fallback("short", 10) == "short" assert truncate_summary_fallback("01234567890", 10).startswith("0123456789") + + +import asyncio # noqa: E402 +from typing import Any # noqa: E402 + +import httpx # noqa: E402 +import pytest # noqa: E402 + +from agent_core.providers import summary as summary_mod # noqa: E402 +from agent_core.providers.summary import ( # noqa: E402 + SummaryLLMEngine, + default_summary_retryable, +) + + +class _Response: + def __init__(self, status: int, body: str) -> None: + self.status_code = status + self.text = body + self._body = body + + def json(self) -> Any: + import json + + return json.loads(self._body) + + def raise_for_status(self) -> None: + if self.status_code >= 400: + raise httpx.HTTPStatusError( + f"{self.status_code}", + request=httpx.Request("POST", "https://host/v1/chat/completions"), + response=httpx.Response(self.status_code, text=self._body), + ) + + +def _install(monkeypatch: pytest.MonkeyPatch, responses: list[_Response]) -> list[int]: + posts: list[int] = [] + + class FakeClient: + def __init__(self, **_kwargs: Any) -> None: + pass + + async def __aenter__(self) -> FakeClient: + return self + + async def __aexit__(self, *_args: Any) -> bool: + return False + + async def post(self, *_args: Any, **_kwargs: Any) -> _Response: + posts.append(1) + return responses[min(len(posts) - 1, len(responses) - 1)] + + monkeypatch.setattr(summary_mod.httpx, "AsyncClient", FakeClient) + return posts + + +def test_permanent_status_is_not_retried(monkeypatch: pytest.MonkeyPatch) -> None: + """A bad key must not burn max_retries attempts on every candidate.""" + posts = _install(monkeypatch, [_Response(401, '{"error":"bad key"}')]) + sleeps: list[float] = [] + + async def sleep(delay: float) -> None: + sleeps.append(delay) + + engine = SummaryLLMEngine(sleep=sleep, fallback_limit=10) + out = asyncio.run( + engine.summarize( + "some long content", + "focus", + [ + {"endpoint": "https://a/v1/chat/completions", "model": "m1"}, + {"endpoint": "https://b/v1/chat/completions", "model": "m2"}, + ], + ), + ) + # One request per candidate, no backoff, then the truncation fallback. + assert len(posts) == 2 + assert sleeps == [] + assert out.endswith("[Content truncated...]") + + +def test_transient_status_is_retried(monkeypatch: pytest.MonkeyPatch) -> None: + posts = _install(monkeypatch, [_Response(503, "overloaded")]) + sleeps: list[float] = [] + + async def sleep(delay: float) -> None: + sleeps.append(delay) + + engine = SummaryLLMEngine(max_retries=3, sleep=sleep, fallback_limit=10) + asyncio.run( + engine.summarize( + "content", + "focus", + [{"endpoint": "https://a/v1/chat/completions", "model": "m1"}], + ), + ) + assert len(posts) == 3 + assert sleeps == [1.0, 2.0] + + +def test_successful_summary_records_usage(monkeypatch: pytest.MonkeyPatch) -> None: + body = ( + '{"choices":[{"message":{"content":"the summary"}}],' + '"usage":{"prompt_tokens":11,"completion_tokens":5,' + '"prompt_tokens_details":{"cached_tokens":7}}}' + ) + _install(monkeypatch, [_Response(200, body)]) + recorded: list[dict[str, Any]] = [] + engine = SummaryLLMEngine(usage_recorder=lambda **kw: recorded.append(kw)) + out = asyncio.run( + engine.summarize( + "content", + "focus", + [{ + "endpoint": "https://a/v1/chat/completions", + "model": "m1", + "provider": "prov", + "api_key": "k", + }], + ), + ) + assert out == "the summary" + assert recorded == [{ + "model": "m1", + "provider": "prov", + "prompt_tokens": 11, + "completion_tokens": 5, + "cache_read_tokens": 7, + }] + + +def test_context_length_truncation_uses_the_full_budget( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The retry ladder must measure the original content, not the shortened one.""" + posts = _install( + monkeypatch, + [_Response(400, "maximum context length exceeded")], + ) + engine = SummaryLLMEngine( + max_retries=4, + truncate_step=40_960, + sleep=lambda _d: _noop(), + fallback_limit=10, + ) + asyncio.run( + engine.summarize( + "x" * 100_000, + "focus", + [{"endpoint": "https://a/v1/chat/completions", "model": "m1"}], + ), + ) + # 40 960 / 81 920 both fit inside 100 000; the third step (122 880) does not. + assert len(posts) == 3 + + +async def _noop() -> None: + return None + + +def test_default_retryable_classification() -> None: + def err(status: int) -> httpx.HTTPStatusError: + return httpx.HTTPStatusError( + "boom", + request=httpx.Request("POST", "https://h/"), + response=httpx.Response(status), + ) + + assert not default_summary_retryable(err(401)) + assert not default_summary_retryable(err(404)) + assert default_summary_retryable(err(429)) + assert default_summary_retryable(err(503)) + assert default_summary_retryable(TimeoutError("timeout")) diff --git a/tests/test_tool_guardrails.py b/tests/test_tool_guardrails.py index 2c67037..10df491 100644 --- a/tests/test_tool_guardrails.py +++ b/tests/test_tool_guardrails.py @@ -75,3 +75,87 @@ def test_repair_truncated_json() -> None: assert repaired is not None assert json.loads(repaired) == {"items": [1, {"ok": True}]} assert repair_truncated_json("not json") is None + + +def test_literal_content_args_are_not_whitespace_stripped() -> None: + """Indent and trailing newlines are semantic for exact-match editors.""" + middleware = ToolCallRepairMiddleware() + result = asyncio.run( + middleware.before_tool_call( + _context( + "file_editor_str_replace", + { + "path": " a.py ", + "old_string": " def foo():\n pass\n", + "new_string": " def foo():\n return 1\n", + }, + ), + ), + ) + assert result.tool_args["old_string"] == " def foo():\n pass\n" + assert result.tool_args["new_string"] == " def foo():\n return 1\n" + # Non-content args are still normalized. + assert result.tool_args["path"] == "a.py" + + written = asyncio.run( + middleware.before_tool_call( + _context("write_file", {"path": "x.py", "content": "print(1)\n"}), + ), + ) + assert written.tool_args["content"] == "print(1)\n" + + +def test_empty_host_tables_disable_defaults() -> None: + """An empty mapping is an explicit opt-out, not a request for defaults.""" + middleware = ToolCallRepairMiddleware( + key_aliases={}, + type_coercions={}, + tool_repair=lambda _name, args: args, + ) + result = asyncio.run( + middleware.before_tool_call(_context("web_search", {"q": "hi", "n": "3"})), + ) + assert result.tool_args == {"q": "hi", "n": "3"} + + +def test_block_helper_sets_the_documented_contract_keys() -> None: + from agent_core.protocols import BLOCK_REASON_KEY, BLOCKED_KEY + + middleware = GuardrailsMiddleware(duplicate_thresholds={"x": 1}) + asyncio.run(middleware.before_tool_call(_context("x", {"a": 1}))) + blocked = asyncio.run(middleware.before_tool_call(_context("x", {"a": 1}))) + assert blocked.is_blocked is True + assert blocked.metadata[BLOCKED_KEY] is True + assert "identical arguments" in blocked.block_reason + assert blocked.metadata[BLOCK_REASON_KEY] == blocked.block_reason + + +def test_fingerprint_tolerates_non_json_arguments() -> None: + middleware = GuardrailsMiddleware() + result = asyncio.run(middleware.before_tool_call(_context("x", {"a": {1, 2}}))) + assert not result.is_blocked + + +def test_budget_exhausted_reads_total_tokens_and_survives_junk_config() -> None: + from agent_core.runtime.loop.guardrails import ( + DEFAULT_MAX_TOKENS, + check_budget_exhausted, + ) + + plan = {"execution_plan": {"budget": {"allocated": {"max_tokens": 100}}}} + # The key every provider and UsageMetadata actually emits. + assert check_budget_exhausted(plan, {"total_tokens": 150}) is not None + assert check_budget_exhausted(plan, {"total": 150}) is not None + assert check_budget_exhausted(plan, {"total_tokens": 50}) is None + # Non-numeric config falls back instead of raising out of an advisory check. + junk = {"execution_plan": {"budget": {"max_tokens": "unlimited"}}} + assert check_budget_exhausted(junk, {"total_tokens": 10}) is None + assert ( + check_budget_exhausted(junk, {"total_tokens": DEFAULT_MAX_TOKENS + 1}) + is not None + ) + # Numeric strings and floats are honoured. + assert check_budget_exhausted( + {"execution_plan": {"budget": {"max_tokens": "100"}}}, + {"total_tokens": 150}, + ) is not None diff --git a/tests/test_usage_meter.py b/tests/test_usage_meter.py index c19d3c9..d81d94f 100644 --- a/tests/test_usage_meter.py +++ b/tests/test_usage_meter.py @@ -88,3 +88,24 @@ def fail(**_kwargs: object) -> None: raise RuntimeError("accounting unavailable") ExternalAPIMeter(llm_recorder=fail).record_llm_usage(model="m") + + +def test_wire_counters_serialize_as_integers() -> None: + """A billing/JSON payload should not carry ``requests: 3.0``.""" + meter = ExternalAPIMeter() + meter.record_api_request("serper", requests=3, cache_hits=1) + meter.set_gauge("sandbox", "peak_mb", 512) + snapshot = meter.snapshot() + serper = snapshot["external_apis"]["serper"] + assert serper["requests"] == 3 + assert all( + isinstance(serper[field], int) + for field in ("requests", "cache_hits", "retries", "errors") + ) + # Gauge-only providers keep the same integral base shape. + sandbox = snapshot["external_apis"]["sandbox"] + assert isinstance(sandbox["requests"], int) + assert sandbox["peak_mb"] == 512.0 + # Duration-style extras stay fractional. + meter.record_api_request("serper", latency_s=1.234) + assert meter.snapshot()["external_apis"]["serper"]["latency_s"] == 1.23