diff --git a/agent_core/completion.py b/agent_core/completion.py new file mode 100644 index 0000000..a58efcb --- /dev/null +++ b/agent_core/completion.py @@ -0,0 +1,149 @@ +"""Provider-neutral completion signals, before display/history normalization.""" +from __future__ import annotations + +import inspect +from collections.abc import Callable +from typing import Any, cast + +from agent_core.llm import LLMResponse, StreamDelta + + +def get_recovery_hook(client: Any, name: str) -> Callable[..., Any] | None: + """Require explicit opt-in; transparent __getattr__ must not unwrap proxies.""" + if not callable(inspect.getattr_static(client, name, None)): + return None + hook = getattr(client, name, None) + return cast("Callable[..., Any]", hook) if callable(hook) else None + + +def _mapping(value: Any) -> dict[str, Any]: + return cast("dict[str, Any]", value) if isinstance(value, dict) else {} + + +def response_rejection_reason(response: Any) -> str: + """Return the explicit provider rejection marker, never infer from silence.""" + metadata = _mapping(getattr(response, "response_metadata", None)) + extra = _mapping(getattr(response, "additional_kwargs", None)) + for mapping in (metadata, extra): + details = _mapping(mapping.get("stop_details")) + if mapping.get("refusal") or ( + details.get("type") == "refusal" + ): + return "refusal" + if getattr(response, "refusal", None): + return "refusal" + content = getattr(response, "content", None) + if isinstance(content, list) and any( + _mapping(block).get("type") == "refusal" for block in cast("list[Any]", content) + ): + return "refusal" + reason = str(getattr(response, "finish_reason", "") or "").strip().lower() + if reason in ("refusal", "content_filter"): + return reason + reason = str(metadata.get("stop_reason") or "").strip().lower() + if reason in ("refusal", "content_filter"): + return reason + return "" + + +def reported_usage(response: Any) -> dict[str, Any] | None: + """Usage eligible for real accounting; missing/estimated data stays diagnostic. + + Unmarked legacy mappings retain their historical reported interpretation. + This does not rewrite the original response or its estimate provenance. + """ + metadata = _mapping(getattr(response, "response_metadata", None)) + source = getattr(response, "usage_source", "") or metadata.get("usage_source", "") + if source == "estimated": + return None + candidates = [getattr(response, "usage", None), getattr(response, "usage_metadata", None)] + candidates.extend((metadata.get("token_usage"), metadata.get("usage"))) + for candidate in candidates: + usage = _mapping(candidate) + if usage and not usage.get("estimated"): + return usage + return None + + +def usage_count(usage: dict[str, Any], *keys: str) -> int | None: + """First reported count, including zero; None means absent/unparseable.""" + for key in keys: + value = usage.get(key) + if value is None: + continue + try: + count = int(value) + if count >= 0: + return count + except (TypeError, ValueError, OverflowError): + continue + return None + + +def _has_reported_usage(response: Any) -> bool: + metadata = _mapping(getattr(response, "response_metadata", None)) + source = getattr(response, "usage_source", "") or ( + metadata.get("usage_source", "") + ) + if source == "estimated": + return False + candidates = [getattr(response, "usage", None), getattr(response, "usage_metadata", None)] + candidates.extend((metadata.get("token_usage"), metadata.get("usage"))) + # Real reported zeros count as a signal. Estimates and wrapper placeholders + # do not: products should retain their provenance instead of inventing a + # provider report. Unmarked legacy usage is treated conservatively. + return any( + usage and not (_mapping(usage).get("estimated")) + for usage in candidates + ) + + +def is_wholly_empty_response(response: Any) -> bool: + """No content, tool calls, reasoning, reported usage or explicit rejection.""" + if getattr(response, "tool_calls", None) or response_rejection_reason(response): + return False + if str(getattr(response, "reasoning_content", "") or "").strip(): + return False + extra = _mapping(getattr(response, "additional_kwargs", None)) + if str(extra.get("reasoning_content") or "").strip(): + return False + if _has_reported_usage(response): + return False + metadata = _mapping(getattr(response, "response_metadata", None)) + if metadata.get("stop_details"): + return False + content = getattr(response, "content", None) + if isinstance(content, str): + return not content.strip() + if isinstance(content, list): + for block in cast("list[Any]", content): + if isinstance(block, str): + if block.strip(): + return False + elif isinstance(block, dict): + block = cast("dict[str, Any]", block) + if ( + str(block.get("text") or block.get("content") or "").strip() + or block.get("type") not in ("text", "") + ): + return False + elif block is not None: + return False + return True + return content is None + + +def stream_delta_has_completion_signal(delta: StreamDelta) -> bool: + """Signals for provider wrappers that cannot assemble signed history.""" + if any(call.get("name") for call in delta.tool_call_deltas): + return True + return not is_wholly_empty_response(LLMResponse( + content=delta.reasoning_blocks or delta.content, + reasoning_content=delta.reasoning_content, + usage=delta.usage, usage_source=delta.usage_source, + finish_reason=delta.finish_reason, + response_metadata={ + "stop_details": delta.stop_details, "stop_reason": delta.stop_reason, + "refusal": delta.refusal, + }, + )) diff --git a/agent_core/components/middleware/llm/loop_detection.py b/agent_core/components/middleware/llm/loop_detection.py index 5ffd0a0..2298a37 100644 --- a/agent_core/components/middleware/llm/loop_detection.py +++ b/agent_core/components/middleware/llm/loop_detection.py @@ -86,6 +86,8 @@ def _history_for(self, key: tuple[str, str, str]) -> deque[str]: async def after_llm( self, ctx: LLMCallContext, response: LLMResponse ) -> LLMResponse: + if ctx.metadata.get("_llm_stream_incomplete"): + return response key = self._scope_key(ctx) if key is None: return response diff --git a/agent_core/components/middleware/llm/proxy.py b/agent_core/components/middleware/llm/proxy.py index bfaae3e..269c741 100644 --- a/agent_core/components/middleware/llm/proxy.py +++ b/agent_core/components/middleware/llm/proxy.py @@ -2,12 +2,14 @@ from __future__ import annotations +import copy import itertools import logging import time from collections.abc import AsyncIterator from typing import Any +from agent_core.completion import get_recovery_hook from agent_core.components.middleware.llm.base import ( LLMCallContext, ) @@ -18,12 +20,70 @@ from agent_core.llm import LLMResponse, StreamDelta from agent_core.messages import Message from agent_core.runtime.async_utils import await_bounded, closing_stream +from agent_core.stream_tools import StreamToolCalls logger = logging.getLogger(__name__) __all__ = ["LLMProxy"] _HOOK_TIMEOUT_S = 5.0 +class _StreamTerminal: + """Terminal metadata of one streamed attempt, folded like the loop's + stream assembler (``_streaming._assembled_response``) so ``after_llm`` + sees usage, finish_reason and refusal signals, not just text.""" + + def __init__(self) -> None: + self.usage: dict[str, int] = {} + self.usage_source = "" + self.finish_reason = "" + self.model = "" + self.provider = "" + self.refusal = "" + self.stop_details: dict[str, Any] = {} + self.stop_reason = "" + self.tool_calls = StreamToolCalls() + + def feed(self, delta: StreamDelta) -> None: + # Last non-empty wins, except refusal text, which streams in pieces. + if delta.usage: + self.usage = delta.usage + self.usage_source = delta.usage_source + if delta.finish_reason: + self.finish_reason = delta.finish_reason + if delta.model: + self.model = delta.model + if delta.provider: + self.provider = delta.provider + self.refusal += delta.refusal + if delta.stop_details: + self.stop_details = delta.stop_details + if delta.stop_reason: + self.stop_reason = delta.stop_reason + # Keep known usage/provenance even if a malformed tool fragment fails. + self.tool_calls.feed(delta.tool_call_deltas) + + def response(self, content: str, reasoning: str) -> LLMResponse: + metadata: dict[str, Any] = {} + if self.provider: + metadata["provider_actually_used"] = self.provider + if self.refusal: + metadata["refusal"] = self.refusal + if self.stop_details: + metadata["stop_details"] = self.stop_details + if self.stop_reason: + metadata["stop_reason"] = self.stop_reason + return LLMResponse( + content=content, + tool_calls=self.tool_calls.complete(), + reasoning_content=reasoning, + finish_reason=self.finish_reason, + model=self.model, + usage=self.usage, + usage_source=self.usage_source, + response_metadata=metadata, + ) + + async def _log_llm_exception( ctx: LLMCallContext, error: Exception, @@ -77,6 +137,20 @@ def __init__( self.model = getattr(inner, "model", "") or "" self._counter = itertools.count(1) + def for_logical_call(self) -> LLMProxy: + prepare = get_recovery_hook(self.inner, "for_logical_call") + if prepare is None: + return self + # Preserve middleware, role, subclasses and the shared atomic counter. + # Only the wrapped provider's routing cursor belongs to this call. + proxy = copy.copy(self) + proxy.inner = prepare() + return proxy + + def advance_empty_completion(self, error: BaseException) -> bool: + advance = get_recovery_hook(self.inner, "advance_empty_completion") + return bool(advance(error)) if advance is not None else False + @property def call_counter(self) -> int: """Last-issued call index (peek without advancing). @@ -187,12 +261,15 @@ async def stream( start_time = time.time() full_content = "" full_reasoning = "" + terminal = _StreamTerminal() + stream_completed = False stream_error: Exception | None = None try: while True: start_time = time.time() full_content = "" full_reasoning = "" + terminal = _StreamTerminal() stream_error = None any_chunk_yielded = False try: @@ -212,6 +289,7 @@ async def stream( continue full_content += delta.content or "" full_reasoning += delta.reasoning_content or "" + terminal.feed(delta) any_chunk_yielded = True # Per-chunk middleware hook. A middleware returning # True (e.g. StreamRepetitionDetector noticing a @@ -235,6 +313,7 @@ async def stream( len(full_content), ) break + stream_completed = True break except Exception as e: stream_error = e @@ -264,6 +343,9 @@ async def stream( (time.time() - start_time) * 1000, ) ctx.metadata["duration_ms"] = duration_ms + # Accounting may use reported usage from failed requests, but loop + # detection must not treat an abandoned proposal as a delivered turn. + ctx.metadata["_llm_stream_incomplete"] = not stream_completed if stream_error: ctx.metadata["error"] = str(stream_error) # Stream after hooks are passive (their result is discarded), so @@ -273,10 +355,7 @@ async def stream( # response (e.g. output repair) and stay unbounded. await self.chain.run_after( ctx, - LLMResponse( - content=full_content, - reasoning_content=full_reasoning, - ), + terminal.response(full_content, full_reasoning), timeout_s=_HOOK_TIMEOUT_S, ) diff --git a/agent_core/components/middleware/llm/token_accounting.py b/agent_core/components/middleware/llm/token_accounting.py index 8ce44ee..ae79802 100644 --- a/agent_core/components/middleware/llm/token_accounting.py +++ b/agent_core/components/middleware/llm/token_accounting.py @@ -3,6 +3,7 @@ import logging from typing import Any, cast +from agent_core.completion import reported_usage, usage_count from agent_core.components.middleware.llm.base import ( LLMCallContext, LLMMiddleware, @@ -88,31 +89,16 @@ def get_usage(self, task_id: str) -> dict[str, int]: return dict(self._usage.get(task_id, {"input": 0, "output": 0, "total": 0, "llm_calls": 0})) def _extract_usage(self, response: LLMResponse) -> tuple[int, int, int, int]: - """Extract (input, output, cache_read, cache_creation) token counts. - - Native :class:`LLMResponse` carries a single flat ``usage`` dict in - OpenAI-wire shape regardless of provider — the infra clients - normalise both OpenAI (``prompt_tokens`` / ``completion_tokens`` / - ``prompt_tokens_details.cached_tokens``) and Anthropic - (``input_tokens`` / ``output_tokens`` / ``cache_read_input_tokens``) - into ``{prompt_tokens, completion_tokens, total_tokens, - cached_tokens}``. Cache-creation tokens are not surfaced by the - native clients, so that count is always 0. - """ - raw_usage: object = getattr(response, "usage", None) - if not isinstance(raw_usage, dict): + """Read canonical usage with legacy aliases, preserving reported zeros.""" + usage = reported_usage(response) + if usage is None: return 0, 0, 0, 0 - usage = cast("dict[str, Any]", raw_usage) - - inp = usage.get("prompt_tokens", usage.get("input_tokens", 0)) or 0 - out = usage.get("completion_tokens", usage.get("output_tokens", 0)) or 0 - cache_read = ( - usage.get("cached_tokens") - or usage.get("cache_read_input_tokens") - or 0 + return ( + usage_count(usage, "prompt_tokens", "input_tokens") or 0, + usage_count(usage, "completion_tokens", "output_tokens") or 0, + usage_count(usage, "cache_read_tokens", "cached_tokens", "cache_read_input_tokens") or 0, + usage_count(usage, "cache_write_tokens", "cache_creation_tokens", "cache_creation_input_tokens") or 0, ) - cache_create = usage.get("cache_creation_input_tokens", 0) or 0 - return int(inp), int(out), int(cache_read), int(cache_create) async def after_llm( self, ctx: LLMCallContext, response: LLMResponse @@ -120,7 +106,7 @@ async def after_llm( input_tokens, output_tokens, cache_read, cache_create = self._extract_usage(response) total = input_tokens + output_tokens - if total == 0: + if total == 0 and cache_read == 0 and cache_create == 0: return response task_id = ctx.task_id or "unknown" diff --git a/agent_core/components/middleware/llm/tracing.py b/agent_core/components/middleware/llm/tracing.py index 663b555..e4cd915 100644 --- a/agent_core/components/middleware/llm/tracing.py +++ b/agent_core/components/middleware/llm/tracing.py @@ -65,6 +65,11 @@ async def after_llm( "duration_ms": duration_ms, "usage": usage, } + source = getattr(response, "usage_source", "") or rm.get("usage_source", "") + if source == "estimated" or (isinstance(usage, dict) and cast("dict[str, Any]", usage).get("estimated")): + metadata["usage_source"] = "estimated" + elif source: + metadata["usage_source"] = source # A fallback chain stamps these on the response when the # call fell through to a secondary model. Surface them in # the trace metadata so observability can flag failover diff --git a/agent_core/components/middleware/rate_limit.py b/agent_core/components/middleware/rate_limit.py index f007f64..43e7882 100644 --- a/agent_core/components/middleware/rate_limit.py +++ b/agent_core/components/middleware/rate_limit.py @@ -12,7 +12,9 @@ import asyncio import logging import time +import uuid +from agent_core.completion import reported_usage, usage_count from agent_core.components.middleware.llm.base import LLMCallContext, LLMMiddleware from agent_core.llm import LLMResponse from agent_core.messages import Message, text_of @@ -59,6 +61,10 @@ def _refill(self) -> None: ) self._last_refill = now + def reservation_size(self, estimated_tokens: int) -> int: + """Amount actually reserved, capped at one full minute's capacity.""" + return min(max(estimated_tokens, 0), int(self._tpm)) + async def acquire(self, estimated_tokens: int = 0) -> float: """Acquire rate limit capacity. Returns wait time in seconds. @@ -69,7 +75,7 @@ async def acquire(self, estimated_tokens: int = 0) -> float: # A single request cannot reserve more than a full minute's token # capacity. Capping avoids an infinite wait for an oversized prompt; # the provider remains the authority on whether that request is valid. - requested_tokens = min(max(estimated_tokens, 0), int(self._tpm)) + requested_tokens = self.reservation_size(estimated_tokens) while True: async with self._lock: @@ -128,6 +134,7 @@ def __init__( ) -> None: self._bucket = TokenBucket(requests_per_min, tokens_per_min) self._estimate_key = "_rate_limit_estimated_tokens" + self._reserved_key = f"_rate_limit_reserved_tokens_{uuid.uuid4().hex}" @property def name(self) -> str: @@ -144,6 +151,9 @@ async def before_llm( len(str(text_of(m.get("content")))) for m in messages ) // 4 ctx.metadata[self._estimate_key] = estimated + # acquire caps oversized reservations at one full token bucket. Correct + # against what was reserved, not the uncapped prompt estimate. + ctx.metadata[self._reserved_key] = self._bucket.reservation_size(estimated) wait = await self._bucket.acquire(estimated) if wait > 0: @@ -157,22 +167,24 @@ async def after_llm( response: LLMResponse, ) -> LLMResponse: """Correct token bucket with actual usage from response.""" - usage = ( - response.usage - or response.response_metadata.get("token_usage") - or response.response_metadata.get("usage") - or {} - ) - actual_total = ( - usage.get("total_tokens") - or usage.get("prompt_tokens", 0) - + usage.get("completion_tokens", 0) - or usage.get("input_tokens", 0) - + usage.get("output_tokens", 0) - ) - estimated = ctx.metadata.get(self._estimate_key, 0) - - if actual_total and estimated: - self._bucket.adjust(actual_total, estimated) + usage = reported_usage(response) + if usage is None: + # Keep the original reservation when provider usage is unknown. + return response + actual_total = usage_count(usage, "total_tokens") + if actual_total is None: + inp = usage_count(usage, "prompt_tokens", "input_tokens") + out = usage_count(usage, "completion_tokens", "output_tokens") + if inp is None and out is None: + return response + actual_total = (inp or 0) + (out or 0) + reserved = usage_count(ctx.metadata, self._reserved_key) + if reserved is None: + # Compatibility for contexts created by older before hooks. + legacy_estimate = usage_count(ctx.metadata, self._estimate_key) + if legacy_estimate is not None: + reserved = self._bucket.reservation_size(legacy_estimate) + if reserved is not None: + self._bucket.adjust(actual_total, reserved) return response diff --git a/agent_core/errors.py b/agent_core/errors.py index 5a8ca19..3369d80 100644 --- a/agent_core/errors.py +++ b/agent_core/errors.py @@ -75,35 +75,21 @@ def __init__( class LLMEmptyCompletion(LLMError): - """A streamed call completed cleanly and produced nothing at all. - - Not a stall (chunks did arrive, or the stream closed without ever going - silent) and not a runaway (no reasoning was spent either). The provider - simply ended the stream with no visible text, no tool call, no reasoning - and — the tell — no usage, so there is nothing to bill and nothing to act - on. :func:`agent_core.runtime.retriable.is_empty_completion` has always - named this case ("reasoning-runaway, all-tokens-in-thinking, or an empty - stream") and routes it through ``is_retriable_with_fallback``, which - resamples on the same key and then advances the chain; what was missing is - the raise, so the assembled empty response reached the loop as a VALID - answer instead. - - Why that is worse than an error: an empty reply with no tool call is - shaped exactly like "the model chose to stop talking", so under - ``no_tool_behavior="stop"`` the run ENDS. Measured on ApodexHarness's - 2026-10-06 GDPval batch over llm-hub: 10 of 19 streamed trials died this - way, six of them inside five minutes (one on turn 1 after 37s), each - discarding every turn of work it had already done. The message wording is - matched by ``_EMPTY_COMPLETION_PATTERNS`` so existing classification picks - it up without a registry edit. + """A completion produced no output, reported usage or rejection signal. + + ``partial_response`` retains the original accounting/provenance for + discarded-attempt events. ``chunks_seen`` is zero for a chat response. """ - def __init__(self, *, chunks_seen: int, elapsed_s: float) -> None: + def __init__( + self, *, chunks_seen: int, elapsed_s: float, partial_response: Any = None, + ) -> None: self.chunks_seen = int(chunks_seen) self.elapsed_s = float(elapsed_s) + self.partial_response = partial_response super().__init__( - "empty completion: stream returned no content, no tool call and " - f"no usage (chunks_seen={self.chunks_seen}, " + "empty completion: response returned no content, no tool call and " + f"no reported usage (chunks_seen={self.chunks_seen}, " f"elapsed={self.elapsed_s:.1f}s)", ) diff --git a/agent_core/llm.py b/agent_core/llm.py index 31cb502..03471d8 100644 --- a/agent_core/llm.py +++ b/agent_core/llm.py @@ -20,6 +20,9 @@ class LLMResponse: model: str = "" usage: dict[str, int] = field(default_factory=dict[str, int]) response_metadata: dict[str, Any] = field(default_factory=dict[str, Any]) + # "provider" for reported usage, "estimated" for host-generated counts. + # Empty keeps conservative compatibility with unmarked legacy clients. + usage_source: str = "" @dataclass @@ -78,6 +81,9 @@ class StreamDelta: # HTTP progress without model output (e.g. Anthropic SSE pings). Keeps # stall guards alive without committing a provider fallback leg. transport_activity: bool = False + usage_source: str = "" + # Refusal text has its own SDK channel on OpenAI Chat and Responses. + refusal: str = "" @runtime_checkable diff --git a/agent_core/loop_types.py b/agent_core/loop_types.py index 61f19c1..fc0b28b 100644 --- a/agent_core/loop_types.py +++ b/agent_core/loop_types.py @@ -172,14 +172,10 @@ class LoopConfig: # a tool-less turn is the model choosing to stop, a truncated one is the # model being stopped, so a truncation must not spend the nudge budget. truncation_max_continuations: int = 2 - # Resamples per consecutive empty-response episode: no text, no tool - # call, no reasoning and no usage. A non-empty reply resets this budget. - # Separate from both budgets above because it is a - # third signal — truncation is "stopped mid-sentence", a tool-less turn is - # "chose to stop", and this is "the transport returned nothing", which is - # not the model's statement at all. Two is enough: the failure is - # per-request, and a provider that returns blank three times running is - # down, not slow. + # Same-leg resamples after a wholly empty response, in both transports. + # Applies per logical call and per serving fallback leg, independently of + # max_llm_retries (the generic failure allowance). All resamples and chain + # advances share logical_call_timeout_s and the run's wall deadline. empty_completion_max_retries: int = 2 max_llm_retries: int = 5 # Fixed retry delay; None uses exponential backoff. diff --git a/agent_core/providers/anthropic.py b/agent_core/providers/anthropic.py index 839f77a..c415bd7 100644 --- a/agent_core/providers/anthropic.py +++ b/agent_core/providers/anthropic.py @@ -269,6 +269,7 @@ async def stream( started = time.monotonic() stream, reset = await self._create_message(kwargs) activity = StreamActivity() + usage_estimated = False async with contextlib.aclosing(stream_events_with_activity(stream, activity)) as events: async for event in events: if event is None: @@ -286,6 +287,7 @@ async def stream( model = getattr(msg, "model", "") or model u = getattr(msg, "usage", None) if u is not None: + usage_estimated = usage_estimated or bool(getattr(u, "estimated", False)) input_tokens = getattr(u, "input_tokens", input_tokens) cr = getattr(u, "cache_read_input_tokens", None) if cr is not None: @@ -383,6 +385,7 @@ async def stream( stop_details = details u = getattr(event, "usage", None) if u is not None: + usage_estimated = usage_estimated or bool(getattr(u, "estimated", False)) ot = getattr(u, "output_tokens", None) if ot is not None: output_tokens = ot @@ -429,6 +432,7 @@ async def stream( block["input"] = {} activity.mark_output() yield StreamDelta( + usage_source="estimated" if usage_estimated else "provider", usage=_anthropic_usage_dict( input_tokens, output_tokens, @@ -1055,6 +1059,9 @@ def _to_llm_response(raw: Any) -> LLMResponse: _anthropic_reasoning_tokens(usage), ) + if getattr(usage, "estimated", False): + usage_dict["estimated"] = True + # ``stop_reason`` is normalised for ``finish_reason`` (``max_tokens`` → # ``length``), which is what the loop's truncation checks need. The RAW # value is kept alongside it: ``refusal`` survives normalisation today, @@ -1075,6 +1082,7 @@ def _to_llm_response(raw: Any) -> LLMResponse: finish_reason=normalize_finish_reason(getattr(raw, "stop_reason", "")), model=getattr(raw, "model", "") or "", usage=usage_dict, + usage_source="estimated" if usage_dict.get("estimated") else "provider" if usage_dict else "", response_metadata=metadata, ) diff --git a/agent_core/providers/fallback.py b/agent_core/providers/fallback.py index 8308150..cde9466 100644 --- a/agent_core/providers/fallback.py +++ b/agent_core/providers/fallback.py @@ -18,6 +18,7 @@ / "too many requests". - ``5xx`` — exception with a ``status_code`` attribute in [500, 599], or message starting with "5". +- ``empty_completion`` — a typed transport blank with no completion signals. - ``any_error`` — match anything (use this on the last entry to guarantee at least one fallback). @@ -57,9 +58,16 @@ import random import time from collections.abc import AsyncIterator, Awaitable, Callable, Sequence -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace +from functools import partial from typing import Any, Literal +from agent_core.completion import ( + get_recovery_hook, + is_wholly_empty_response, + stream_delta_has_completion_signal, +) +from agent_core.errors import LLMEmptyCompletion from agent_core.llm import LLMResponse, StreamDelta from agent_core.messages import Message @@ -95,7 +103,7 @@ async def _noop_event(_name: str, _payload: dict[str, Any]) -> None: return None -FallbackTrigger = Literal["timeout", "rate_limit", "5xx", "any_error"] +FallbackTrigger = Literal["timeout", "rate_limit", "5xx", "any_error", "empty_completion"] @dataclass @@ -155,6 +163,8 @@ def _trigger_matches(trigger: FallbackTrigger, exc: BaseException) -> bool: """Decide whether ``exc`` matches the given trigger keyword.""" if trigger == "any_error": return True + if trigger == "empty_completion": + return isinstance(exc, LLMEmptyCompletion) name = type(exc).__name__.lower() msg = str(exc).lower() @@ -200,6 +210,34 @@ def _trigger_matches(trigger: FallbackTrigger, exc: BaseException) -> bool: return False +async def _nonempty_stream( + client: Any, messages: list[Message], **kwargs: Any, +) -> AsyncIterator[StreamDelta]: + """Do not commit candidate-empty deltas that could contaminate a new leg.""" + pending: list[StreamDelta] = [] + committed = False + chunks_seen = 0 + started = time.monotonic() + async with closing_stream(client.stream(messages, **kwargs)) as inner: + async for delta in inner: + if delta.transport_activity: + yield delta + continue + chunks_seen += 1 + if committed: + yield delta + elif stream_delta_has_completion_signal(delta): + committed = True + for buffered in pending: + yield buffered + pending.clear() + yield delta + else: + pending.append(delta) + if not committed: + raise LLMEmptyCompletion(chunks_seen=chunks_seen, elapsed_s=time.monotonic() - started) + + class CooldownFallbackLLM: """Retry a primary client, then use a fallback client during cooldown. @@ -417,7 +455,10 @@ async def chat( for attempt in range(self.max_retries): try: await self._emit("request", leg="primary", attempt=attempt + 1) - return await self.primary.chat(messages, **kwargs) + response = await self.primary.chat(messages, **kwargs) + if is_wholly_empty_response(response): + raise LLMEmptyCompletion(chunks_seen=0, elapsed_s=0, partial_response=response) + return response except Exception as error: last_error = error await self._emit( @@ -426,7 +467,7 @@ async def chat( attempt=attempt + 1, error=str(error), ) - if not self._retryable(error): + if not self._retryable(error) and not isinstance(error, LLMEmptyCompletion): break if attempt + 1 >= self.max_retries: # Last attempt: sleeping here only delays the degrade and @@ -495,7 +536,7 @@ async def stream( "request", leg="fallback", mode="cooldown", streaming=True, ) try: - async with closing_stream(self.fallback.stream(messages, **kwargs)) as inner_stream: + async with closing_stream(_nonempty_stream(self.fallback, messages, **kwargs)) as inner_stream: async for delta in inner_stream: yield delta except Exception as fallback_error: @@ -518,7 +559,7 @@ async def stream( attempt=attempt + 1, streaming=True, ) - async with closing_stream(self.primary.stream(messages, **kwargs)) as inner_stream: + async with closing_stream(_nonempty_stream(self.primary, messages, **kwargs)) as inner_stream: async for delta in inner_stream: if not delta.transport_activity: yielded = True @@ -533,7 +574,7 @@ async def stream( streaming=True, error=str(error), ) - if not self._retryable(error): + if not self._retryable(error) and not isinstance(error, LLMEmptyCompletion): break if yielded and is_truncated_stream(error): # The consumer already has part of this turn. Appending a @@ -581,7 +622,7 @@ async def stream( "request", leg="fallback", mode="degraded", streaming=True, ) try: - async with closing_stream(self.fallback.stream(messages, **kwargs)) as inner_stream: + async with closing_stream(_nonempty_stream(self.fallback, messages, **kwargs)) as inner_stream: async for delta in inner_stream: yield delta except Exception as fallback_error: @@ -626,6 +667,9 @@ class LLMFallbackChain: entries: list[FallbackEntry] = field(default_factory=list) default_triggers: tuple[FallbackTrigger, ...] = ("any_error",) model: str = field(init=False, default="") + _empty_recovery_managed: bool = field(default=False, init=False, repr=False) + _start_index: int = field(default=0, init=False, repr=False) + _serving_index: int = field(default=0, init=False, repr=False) def __post_init__(self) -> None: if not self.entries: @@ -639,6 +683,32 @@ def __post_init__(self) -> None: if entry.triggers is None: entry.triggers = self.default_triggers + def for_logical_call(self) -> LLMFallbackChain: + """Return an independent cursor; never mutate the cached/shared chain.""" + entries = [] + for entry in self.entries: + prepare = get_recovery_hook(entry.model, "for_logical_call") + model = prepare() if callable(prepare) else entry.model + entries.append(replace(entry, model=model)) + cursor = replace(self, entries=entries) + cursor._empty_recovery_managed = True + return cursor + + def advance_empty_completion(self, error: BaseException) -> bool: + """Advance after the caller spent its same-leg empty retry budget.""" + index = self._serving_index + nested_advance = get_recovery_hook(self.entries[index].model, "advance_empty_completion") + if self._empty_recovery_managed and callable(nested_advance) and nested_advance(error): + return True + if ( + not self._empty_recovery_managed + or index + 1 >= len(self.entries) + or not self.entries[index].matches(error) + ): + return False + self._start_index = index + 1 + return True + @classmethod def from_models( cls, @@ -670,6 +740,11 @@ async def chat( ) -> LLMResponse: last_exc: BaseException | None = None for idx, entry in enumerate(self.entries): + if idx < self._start_index: + continue + self._serving_index = idx + if self._empty_recovery_managed: + self._start_index = idx try: result = await entry.model.chat( messages, @@ -679,6 +754,8 @@ async def chat( extra_headers=extra_headers, timeout=timeout, ) + if not self._empty_recovery_managed and is_wholly_empty_response(result): + raise LLMEmptyCompletion(chunks_seen=0, elapsed_s=0, partial_response=result) _stamp_metadata(result, idx, entry.model, entry.provider) return result except Exception as exc: @@ -707,9 +784,17 @@ async def stream( # has committed to that entry. last_exc: BaseException | None = None for idx, entry in enumerate(self.entries): + if idx < self._start_index: + continue + self._serving_index = idx + if self._empty_recovery_managed: + self._start_index = idx yielded_any = False try: - async with closing_stream(entry.model.stream( + stream_client = entry.model.stream + if not self._empty_recovery_managed: + stream_client = partial(_nonempty_stream, entry.model) + async with closing_stream(stream_client( messages, tools=tools, temperature=temperature, diff --git a/agent_core/providers/openai_chat.py b/agent_core/providers/openai_chat.py index fe49ba6..dabe824 100644 --- a/agent_core/providers/openai_chat.py +++ b/agent_core/providers/openai_chat.py @@ -400,6 +400,11 @@ async def stream( stream = await self._open_stream(kwargs) activity = StreamActivity() + # ``refusal: ""`` is a refusal only when the turn carries nothing else; + # some gateways send it next to ordinary output. Hold that marker until + # normal EOF, when the whole turn is known even without finish_reason. + empty_refusal_seen = False + output_seen = False watch_done_sentinel(stream, activity) started = time.monotonic() chunks_seen = 0 @@ -418,7 +423,7 @@ async def stream( # streaming usage/billing read 0). if chunk_usage or chunk_model: activity.mark_output() - yield StreamDelta(usage=chunk_usage, model=chunk_model) + yield StreamDelta(usage=chunk_usage, usage_source="estimated" if chunk_usage.get("estimated") else "provider", model=chunk_model) continue choice = chunk.choices[0] delta = choice.delta @@ -426,13 +431,23 @@ async def stream( if finish_reason: last_finish_reason = finish_reason activity.mark_output() + refusal = getattr(delta, "refusal", None) + tool_call_deltas = _tool_call_deltas(getattr(delta, "tool_calls", None)) + empty_refusal_seen = empty_refusal_seen or refusal == "" + reasoning = _reasoning_text(delta) + output_seen = output_seen or bool( + getattr(delta, "content", None) or refusal or tool_call_deltas or reasoning + ) yield StreamDelta( - content=getattr(delta, "content", None) or "", - reasoning_content=_reasoning_text(delta), - tool_call_deltas=_tool_call_deltas(getattr(delta, "tool_calls", None)), + content=getattr(delta, "content", None) or refusal or "", + refusal=refusal or "", + stop_details={"type": "refusal"} if refusal else {}, + reasoning_content=reasoning, + tool_call_deltas=tool_call_deltas, finish_reason=finish_reason, model=chunk_model, usage=chunk_usage, + usage_source="estimated" if chunk_usage.get("estimated") else "provider" if chunk_usage else "", ) # Some compatible gateways omit finish_reason, while the SDK consumes # [DONE] without yielding it. Either signal proves a completed turn. @@ -449,6 +464,11 @@ async def stream( logger.warning("Accepting stream without terminator (%s=0): %s", STREAM_TERMINATOR_ENV, error) + # Do not put this in finally: errors and consumer cancellation must + # propagate, not manufacture a refusal from an incomplete turn. + if empty_refusal_seen and not output_seen: + yield StreamDelta(stop_details={"type": "refusal"}) + async def _open_stream(self, kwargs: dict[str, Any]) -> Any: """Open the stream, retrying once past a rejected ``reasoning_effort``.""" try: @@ -536,6 +556,8 @@ def _usage_dict(usage: Any) -> dict[str, int]: reasoning = getattr(ctd, "reasoning_tokens", None) if ctd is not None else None if reasoning is not None: out["reasoning_tokens"] = int(reasoning) + if getattr(usage, "estimated", False): + out["estimated"] = True return out @@ -561,7 +583,14 @@ def _to_llm_response(raw: Any) -> LLMResponse: }, }) usage_dict = _usage_dict(getattr(raw, "usage", None)) - content = getattr(msg, "content", "") or "" + refusal = getattr(msg, "refusal", None) + # ``refusal: ""`` alongside real output is gateway noise, not a decline. + reasoning = _reasoning_text(msg) + has_refusal = bool(refusal) or ( + refusal == "" and not (getattr(msg, "content", None) or tool_calls or reasoning) + ) + refusal = refusal or "" + content = getattr(msg, "content", "") or refusal # Mirror the streaming path (llm_client._stream_llm_response): a Qwen # ``\n\n`` separator remnant is either the whole of ``content`` # (whitespace-only → drop) or leads the real answer (``\n\nAnswer…`` → @@ -572,11 +601,15 @@ def _to_llm_response(raw: Any) -> LLMResponse: return LLMResponse( content=content, tool_calls=tool_calls, - reasoning_content=_reasoning_text(msg), + reasoning_content=reasoning, finish_reason=getattr(choice, "finish_reason", "") or "", model=getattr(raw, "model", "") or "", usage=usage_dict, - response_metadata={"id": getattr(raw, "id", "")}, + usage_source="estimated" if usage_dict.get("estimated") else "provider" if usage_dict else "", + response_metadata={ + "id": getattr(raw, "id", ""), + **({"refusal": refusal, "stop_details": {"type": "refusal"}} if has_refusal else {}), + }, ) diff --git a/agent_core/providers/openai_responses.py b/agent_core/providers/openai_responses.py index a76b2c9..d1da3e2 100644 --- a/agent_core/providers/openai_responses.py +++ b/agent_core/providers/openai_responses.py @@ -196,6 +196,7 @@ async def stream( kwargs["stream"] = True stream = await self._client.responses.create(**kwargs) activity = StreamActivity() + refusal_seen = False started = time.monotonic() events_seen = 0 last_event = "" @@ -211,6 +212,11 @@ async def stream( if etype == "response.output_text.delta": activity.mark_output() yield StreamDelta(content=getattr(event, "delta", "") or "") + elif etype == "response.refusal.delta": + refusal = getattr(event, "delta", "") or "" + refusal_seen = refusal_seen or bool(refusal) + activity.mark_output() + yield StreamDelta(content=refusal, refusal=refusal) elif etype in ( "response.reasoning_summary_text.delta", "response.reasoning_text.delta", @@ -232,6 +238,10 @@ async def stream( elif etype in ("response.completed", "response.incomplete"): saw_terminal = True resp = getattr(event, "response", None) + parsed = _parse_responses_output(resp) + refusal = parsed.response_metadata.get("refusal", "") + if refusal and not refusal_seen: + yield StreamDelta(content=refusal, refusal=refusal) usage = _responses_usage_dict(getattr(resp, "usage", None)) reason = normalize_finish_reason( _get(_get(resp, "incomplete_details", None) or {}, "reason", ""), @@ -239,6 +249,8 @@ async def stream( activity.mark_output() yield StreamDelta( usage=usage, + stop_details=parsed.response_metadata.get("stop_details", {}), + usage_source="estimated" if usage.get("estimated") else "provider" if usage else "", model=getattr(resp, "model", "") or "", finish_reason=reason or "stop", ) @@ -411,6 +423,7 @@ def _parse_responses_output(raw: Any) -> LLMResponse: if str(_get(raw, "status", "") or "") == "failed": raise _response_failure(raw, fallback="Responses request failed") blocks_out: list[dict[str, Any]] = [] + refusal_parts: list[str] = [] text_parts: list[str] = [] summary_parts: list[str] = [] tool_calls: list[ToolCall] = [] @@ -441,6 +454,11 @@ def _parse_responses_output(raw: Any) -> LLMResponse: txt = _get(part, "text", "") or "" text_parts.append(txt) blocks_out.append({"type": "text", "text": txt}) + elif _get(part, "type", "") == "refusal": + refusal = _get(part, "refusal", "") or "" + refusal_parts.append(refusal) + text_parts.append(refusal) + blocks_out.append({"type": "text", "text": refusal}) elif itype == "function_call": call_id = _get(item, "call_id", "") or _get(item, "id", "") or "" call_name = _get(item, "name", "") or "" @@ -485,7 +503,11 @@ def _parse_responses_output(raw: Any) -> LLMResponse: ), model=_get(raw, "model", "") or "", usage=_responses_usage_dict(_get(raw, "usage", None)), - response_metadata={"id": _get(raw, "id", "")}, + usage_source="estimated" if _get(_get(raw, "usage", None), "estimated", False) else "provider" if _get(raw, "usage", None) is not None else "", + response_metadata={ + "id": _get(raw, "id", ""), + **({"refusal": "\n".join(refusal_parts), "stop_details": {"type": "refusal"}} if refusal_parts else {}), + }, ) @@ -528,6 +550,8 @@ def _responses_usage_dict(usage: Any) -> dict[str, int]: # — exist to avoid. if reasoning is not None: out["reasoning_tokens"] = int(reasoning) + if _get(usage, "estimated", False): + out["estimated"] = True return out diff --git a/agent_core/runtime/loop/_bind.py b/agent_core/runtime/loop/_bind.py index ede08e7..16be204 100644 --- a/agent_core/runtime/loop/_bind.py +++ b/agent_core/runtime/loop/_bind.py @@ -6,6 +6,7 @@ from dataclasses import dataclass, replace from typing import Any +from agent_core.completion import get_recovery_hook from agent_core.llm import LLMResponse from agent_core.messages import Message @@ -30,6 +31,15 @@ class _BoundLLM: extra_headers: dict[str, str] | None = None max_tokens: int | None = None + def for_logical_call(self) -> _BoundLLM: + """Allocate optional client-owned routing state for this call only.""" + prepare = get_recovery_hook(self.client, "for_logical_call") + return replace(self, client=prepare()) if callable(prepare) else self + + def advance_empty_completion(self, error: BaseException) -> bool: + advance = get_recovery_hook(self.client, "advance_empty_completion") + return bool(advance(error)) if callable(advance) else False + @property def model(self) -> str: return getattr(self.client, "model", "") or "" diff --git a/agent_core/runtime/loop/_call.py b/agent_core/runtime/loop/_call.py index 8bac92f..f788a45 100644 --- a/agent_core/runtime/loop/_call.py +++ b/agent_core/runtime/loop/_call.py @@ -12,6 +12,7 @@ from agent_core.errors import ( LLMCallExhausted, LLMDeadlineExceeded, + LLMEmptyCompletion, LLMReasoningRunaway, LLMStreamStalled, ) @@ -37,7 +38,7 @@ from ..async_utils import await_bounded, hold_until_settled from ._bind import _ensure_bound -from ._response import _visible_response_text, extract_usage +from ._response import _visible_response_text, extract_usage, is_wholly_empty_response from ._runaway import ( _RUNAWAY_BACKOFF_S, _RUNAWAY_EXPAND_ENABLED, @@ -215,6 +216,7 @@ async def call_llm( context_token_limit_hint: int | None = None, wall_deadline_remaining: Callable[[], float | None] | None = None, chain_fallback_active: Callable[[], bool] | None = None, + empty_completion_max_retries: int = 2, ) -> LLMResponse | None: """Call ``llm.chat`` (``llm.stream`` when ``on_delta`` is set) with exponential backoff on transient errors. @@ -238,6 +240,15 @@ async def call_llm( / 5xx / 429 / proxy-wrap): raises (reason=``exhausted``) carrying the last transient exception. + ``empty_completion_max_retries`` is the same-leg resampling allowance + after a transport blank in either mode. It does not consume the generic + ``max_retries`` allowance. After that budget is spent, a prepared native + fallback cursor can advance to a matching next leg and gets a fresh empty + allowance; a configured external chain instead receives ``chain_advance``. + Otherwise exhaustion raises with reason ``empty_completion``. All of these + attempts share this call's original logical deadline. Empty attempts are + discarded/failed before any accepted event or history insertion. + Streaming calls additionally run under the inter-chunk stall watchdog (:class:`LLMStreamStalled`, ``LLM_STREAM_STALL_S`` under any supported prefix, default 180 s): a stream that @@ -422,7 +433,7 @@ def _response_attempt_fields(response: LLMResponse) -> dict[str, Any]: # ``llm_active`` may be re-bound with a reduced max_tokens after a # repeated reasoning runaway — error retries then reuse the bound # variant too, which is fine (the cap only applies post-runaway). - llm_active = _ensure_bound(llm) + llm_active = _ensure_bound(llm).for_logical_call() messages_active = messages retry_thinking: ThinkingRetryOverride | None = None guard_timeout_s = reasoning_only_timeout_s @@ -433,7 +444,10 @@ def _response_attempt_fields(response: LLMResponse) -> dict[str, Any]: # from this number, so reusing it would emit two ``finished`` events # under one id and double-count that attempt's usage downstream. physical_attempt_index = 0 - for attempt in range(max_retries): + empty_completions = 0 + empty_budget = max(int(empty_completion_max_retries), 0) + attempt = 0 + while attempt < max_retries: attempt_thinking = retry_thinking # Budget closure: clamp this attempt to the remaining wall (when # a deadline is stamped) so a retry chain can never outlive the @@ -447,6 +461,7 @@ def _response_attempt_fields(response: LLMResponse) -> dict[str, Any]: attempt_index = physical_attempt_index attempt_started = time.monotonic() attempt_first_delta: float | None = None + response_ended_at: float | None = None active_cap = ( getattr(llm_active, "max_tokens", None) or max_completion_tokens_hint @@ -499,10 +514,15 @@ async def _attempt_delta( attempt_delta = _attempt_delta async def _chat_active(read_timeout: float) -> LLMResponse: + chat_started = time.monotonic() with thinking_retry_override(attempt_thinking): - return await llm_active.chat( - messages_active, timeout=read_timeout, - ) + response = await llm_active.chat(messages_active, timeout=read_timeout) + if is_wholly_empty_response(response): + raise LLMEmptyCompletion( + chunks_seen=0, elapsed_s=time.monotonic() - chat_started, + partial_response=response, + ) + return response async def _stream_active(read_timeout: float) -> LLMResponse: if attempt_delta is None: @@ -556,6 +576,40 @@ async def _finish_attempt( event.update(_response_attempt_fields(response)) await _emit_attempt(event) + async def _recover_empty_completion(exc: LLMEmptyCompletion) -> None: + nonlocal last_exc, empty_completions + last_exc = exc + empty_completions += 1 + advance = False + if empty_completions > empty_budget: + advance = llm_active.advance_empty_completion(exc) + if advance: + empty_completions = 0 + else: + reason = "chain_advance" if _chain_fallback_active() else "empty_completion" + await _finish_attempt( + outcome=ATTEMPT_FAILED, reason="empty_completion", + recovery_action=reason, error=exc, + response=exc.partial_response, + ) + raise LLMCallExhausted(exc, reason) from exc + delay = 0 if advance else _transient_backoff(empty_completions - 1) + remaining, deadline_reason = _nearest_deadline() + if remaining is not None and delay + _WALL_DEADLINE_FLOOR_S > remaining: + deadline_exc = LLMDeadlineExceeded(deadline_reason, "no budget for empty resample") + await _finish_attempt( + outcome=ATTEMPT_FAILED, reason=deadline_reason, + recovery_action="abandon_retry", error=deadline_exc, + response=exc.partial_response, + ) + raise LLMCallExhausted(deadline_exc, deadline_reason, prior_exc=exc) from exc + await _finish_attempt( + outcome=ATTEMPT_DISCARDED, reason="empty_completion", + recovery_action="chain_advance" if advance else "retry_same_key", + error=exc, response=exc.partial_response, + ) + await asyncio.sleep(delay) + retry_reason = "transient_error" retry_error: BaseException | None = None try: @@ -646,6 +700,8 @@ async def _finish_attempt( ) recovered: LLMResponse | None = None recovery_error: BaseException | None = None + replay_started: float | None = None + replay_index = 0 try: # Clamp to whatever is left of THIS attempt's own # budget as well as the wall/logical deadline: the @@ -671,6 +727,25 @@ async def _finish_attempt( f"{effective_timeout:.0f}s attempt budget " f"left for the replay", ) + physical_attempt_index += 1 + replay_index = physical_attempt_index + replay_started = time.monotonic() + await _emit_attempt({ + "phase": "started", "attempt_index": replay_index, + "max_tokens": active_cap, + "thinking_mode": attempt_thinking.mode if attempt_thinking else "profile_default", + "thinking_budget": attempt_thinking.thinking_budget if attempt_thinking else None, + }) + # Observation hooks have their own grace period; + # re-clamp before the actual replay request starts. + recovery_timeout = min( + _effective_timeout_or_deadline_exhausted( + attempt=attempt, reason="stream_empty_tool_arguments", + )[0], + max(attempt_deadline - time.monotonic(), 0.0), + ) + if _stream_recovery_budget_too_small(recovery_timeout, float(effective_timeout)): + raise TimeoutError("no useful replay budget remains after observation") recovered = await await_bounded( _chat_active(recovery_timeout), recovery_timeout, @@ -685,6 +760,23 @@ async def _finish_attempt( recovery_error = exc if recovered is None: + response_ended_at = stream_ended_at + if replay_started is not None: + replay_event: dict[str, Any] = { + "phase": "finished", "attempt_index": replay_index, + "outcome": ATTEMPT_FAILED, + "reason": "empty_completion" if isinstance(recovery_error, LLMEmptyCompletion) else "stream_empty_args_replay", + "recovery_action": "keep_streamed_response", + "duration_ms": int((time.monotonic() - replay_started) * 1000), + "ttft_ms": None, "max_tokens": active_cap, + "thinking_mode": attempt_thinking.mode if attempt_thinking else "profile_default", + "thinking_budget": attempt_thinking.thinking_budget if attempt_thinking else None, + "error_type": type(recovery_error).__name__, + } + partial = getattr(recovery_error, "partial_response", None) + if partial is not None: + replay_event.update(_response_attempt_fields(partial)) + await _emit_attempt(replay_event) logger.warning( "Non-streaming replay failed (%s: %s); keeping " "the streamed response with blank tool " @@ -719,23 +811,9 @@ async def _finish_attempt( response=streamed_response, ended_at=stream_ended_at, ) - physical_attempt_index += 1 - attempt_index = physical_attempt_index - attempt_started = stream_ended_at + attempt_index = replay_index + attempt_started = replay_started if replay_started is not None else stream_ended_at attempt_first_delta = None - await _emit_attempt({ - "phase": "started", - "attempt_index": attempt_index, - "max_tokens": active_cap, - "thinking_mode": ( - attempt_thinking.mode - if attempt_thinking else "profile_default" - ), - "thinking_budget": ( - attempt_thinking.thinking_budget - if attempt_thinking else None - ), - }) response = recovered response.response_metadata = { **(response.response_metadata or {}), @@ -829,6 +907,7 @@ async def _finish_attempt( prior_turn_runaway, ) await asyncio.sleep(_RUNAWAY_BACKOFF_S) + attempt += 1 continue if runaway_state is not None: runaway_state["consecutive_turns"] = ( @@ -867,8 +946,12 @@ async def _finish_attempt( reason="", recovery_action="accepted", response=response, + ended_at=response_ended_at, ) return response + except LLMEmptyCompletion as exc: + await _recover_empty_completion(exc) + continue except LLMCallExhausted as exc: await _finish_attempt( outcome=ATTEMPT_FAILED, @@ -953,6 +1036,7 @@ async def _finish_attempt( prior_turn_runaway, ) await asyncio.sleep(_RUNAWAY_BACKOFF_S) + attempt += 1 continue if runaway_state is not None: runaway_state["consecutive_turns"] = ( @@ -1055,6 +1139,13 @@ async def _finish_attempt( ) backoff = _transient_backoff(attempt) except Exception as exc: + if is_empty_completion(exc): + empty_exc = LLMEmptyCompletion( + chunks_seen=0, elapsed_s=time.monotonic() - attempt_started, + ) + empty_exc.__cause__ = exc + await _recover_empty_completion(empty_exc) + continue last_exc = exc retry_error = exc retry_reason = "transient_error" @@ -1090,6 +1181,7 @@ async def _finish_attempt( turn, attempt + 1, max_retries, next_cap, retry_thinking.thinking_budget, ) + attempt += 1 continue # Chain-aware shortcut: model_not_found / overload / credit # / safety_filter is deterministic on this (provider, input). @@ -1240,6 +1332,8 @@ async def _finish_attempt( error=retry_error, ) + attempt += 1 + logger.error("LLM call failed after %d retries (turn=%d)", max_retries, turn) # Should always have an exception captured here — every except clause # sets last_exc. Defensive RuntimeError covers a hypothetical diff --git a/agent_core/runtime/loop/_response.py b/agent_core/runtime/loop/_response.py index 210591c..c621e73 100644 --- a/agent_core/runtime/loop/_response.py +++ b/agent_core/runtime/loop/_response.py @@ -6,6 +6,7 @@ from collections.abc import Mapping from typing import Any +from agent_core.completion import is_wholly_empty_response as is_wholly_empty_response from agent_core.llm import LLMResponse from agent_core.loop_types import UsageMetadata from agent_core.messages import Message @@ -94,49 +95,6 @@ def extract_leaked_reasoning(response: Any) -> str: return value if isinstance(value, str) else "" -def is_wholly_empty_response(response: Any) -> bool: - """No content, tool calls, reasoning or usage on the original response. - - Inspect before normalization: hiding reasoning or dropping opaque blocks - for display does not make an upstream response empty. Legacy wrappers may - carry reasoning and usage on their auxiliary metadata channels. - """ - if getattr(response, "tool_calls", None): - return False - reasoning = getattr(response, "reasoning_content", "") or "" - if str(reasoning).strip(): - return False - extra = getattr(response, "additional_kwargs", None) or {} - if isinstance(extra, dict) and str(extra.get("reasoning_content") or "").strip(): - return False - if getattr(response, "usage", None) or getattr(response, "usage_metadata", None): - return False - metadata = getattr(response, "response_metadata", None) or {} - if isinstance(metadata, dict) and ( - metadata.get("token_usage") or metadata.get("usage") or metadata.get("stop_details") - ): - return False - content = getattr(response, "content", None) - if isinstance(content, str): - return not content.strip() - if isinstance(content, list): - for block in content: - if isinstance(block, str): - if block.strip(): - return False - elif isinstance(block, dict): - if ( - str(block.get("text") or block.get("content") or "").strip() - or block.get("type") not in ("text", "") - ): - return False - elif block is not None: - # Unknown payloads are not evidence of an empty response. - return False - return True - return content is None - - def _pick_int(*candidates: Any) -> int: """Return the first non-zero int-coercible candidate, else 0. @@ -271,6 +229,8 @@ def extract_usage(response: Any) -> UsageMetadata | None: # invariant on the ``UsageMetadata`` return type. "reasoning_tokens": int(usage.get("reasoning_tokens", 0) or 0), } + if usage.get("estimated") or getattr(response, "usage_source", "") == "estimated" or rmd.get("usage_source") == "estimated": + out_dict["estimated"] = True return out_dict rmd = getattr(response, "response_metadata", None) or {} diff --git a/agent_core/runtime/loop/_streaming.py b/agent_core/runtime/loop/_streaming.py index 3bc6df8..22f1d51 100644 --- a/agent_core/runtime/loop/_streaming.py +++ b/agent_core/runtime/loop/_streaming.py @@ -15,8 +15,9 @@ LLMStreamStalled, ) from agent_core.llm import LLMResponse -from agent_core.messages import Message, ToolCall +from agent_core.messages import Message from agent_core.runtime.async_utils import await_bounded +from agent_core.stream_tools import StreamToolCalls from ._response import is_wholly_empty_response from ._runaway import _env_float, _env_int @@ -243,13 +244,12 @@ async def _stream_llm_response( accumulated = "" thinking_accum = "" delta_index = 0 - # Typed as ToolCall, not dict[str, Any]: the slots below are assembled - # in the wire shape LLMResponse.tool_calls declares, and the literal - # keeps ToolCall's fixed {id, type, function} key order. - tool_call_acc: dict[int, ToolCall] = {} + tool_calls = StreamToolCalls() # Terminal metadata streamed late by the provider — kept so the assembled # LLMResponse carries usage/finish_reason/model (else streaming runs report # 0 usage and observers never see finish_reason="length"). + final_usage_source = "" + final_refusal = "" final_usage: dict[str, int] = {} final_finish_reason = "" final_model = "" @@ -298,6 +298,8 @@ def _assembled_response() -> LLMResponse: response_metadata = ( {"provider_actually_used": final_provider} if final_provider else {} ) + if final_refusal: + response_metadata["refusal"] = final_refusal if final_stop_details: response_metadata["stop_details"] = final_stop_details if final_stop_reason: @@ -327,15 +329,11 @@ def _assembled_response() -> LLMResponse: # NAMED calls to tools that do not exist are deliberately kept: those # reach the executor and come back as "unknown tool 'x'", which the # model can read and act on. - complete_tool_calls = [ - tool_call_acc[k] - for k in sorted(tool_call_acc) - if tool_call_acc[k]["function"]["name"] - ] + complete_tool_calls = tool_calls.complete() # Warn rather than drop silently — a provider emitting these # consistently is a real upstream defect, and this is the only place # that can still see it. - dropped = len(tool_call_acc) - len(complete_tool_calls) + dropped = tool_calls.dropped_count if dropped: logger.warning( "dropped %d streamed tool_call(s) with no function name", dropped, @@ -373,6 +371,7 @@ def _assembled_response() -> LLMResponse: tool_calls=complete_tool_calls, reasoning_content=thinking_accum, usage=final_usage, + usage_source=final_usage_source, finish_reason=final_finish_reason, model=final_model, response_metadata=response_metadata, @@ -400,7 +399,7 @@ async def _close_chunk_stream() -> None: continue else: chunks_seen += 1 - raw_visible = delta.content or "" + raw_visible = delta.content or getattr(delta, "refusal", "") or "" typed_thinking = delta.reasoning_content or "" tc_chunks = delta.tool_call_deltas or [] # Capture terminal metadata as it arrives (usage on the @@ -408,6 +407,9 @@ async def _close_chunk_stream() -> None: # content chunk). Last non-empty wins. if getattr(delta, "usage", None): final_usage = delta.usage + final_usage_source = getattr(delta, "usage_source", "") + if getattr(delta, "refusal", ""): + final_refusal += delta.refusal if getattr(delta, "finish_reason", ""): final_finish_reason = delta.finish_reason if getattr(delta, "model", ""): @@ -443,18 +445,7 @@ async def _close_chunk_stream() -> None: ) else "" # Stitch streamed tool-call deltas by index — id set # once, name/arguments appended — into wire-shaped slots. - for tcd in tc_chunks: - idx = tcd.get("index") or 0 - slot = tool_call_acc.setdefault(idx, { - "id": "", "type": "function", - "function": {"name": "", "arguments": ""}, - }) - if tcd.get("id"): - slot["id"] = tcd["id"] - if tcd.get("name"): - slot["function"]["name"] += tcd["name"] - if tcd.get("arguments"): - slot["function"]["arguments"] += tcd["arguments"] + tool_calls.feed(tc_chunks) if visible or thinking or tc_chunks: if visible: accumulated += visible @@ -618,6 +609,7 @@ async def _close_chunk_stream() -> None: raise LLMEmptyCompletion( chunks_seen=chunks_seen, elapsed_s=time.monotonic() - stream_started, + partial_response=assembled, ) return assembled diff --git a/agent_core/runtime/loop/agent_loop.py b/agent_core/runtime/loop/agent_loop.py index 8659192..b428029 100644 --- a/agent_core/runtime/loop/agent_loop.py +++ b/agent_core/runtime/loop/agent_loop.py @@ -16,6 +16,7 @@ from dataclasses import dataclass, field from typing import Any, cast +from agent_core.completion import response_rejection_reason from agent_core.llm import LLMClient from agent_core.loop_types import ( AgentLoopResult, @@ -65,7 +66,6 @@ extract_leaked_reasoning, extract_usage, is_truncated_with_text, - is_wholly_empty_response, ) from agent_core.runtime.loop.model_profile import ( DefaultThinkingParser, @@ -440,7 +440,6 @@ async def _run_loop_inner( total_tool_calls = 0 no_tool_retries = 0 truncation_continuations = 0 - empty_completions = 0 truncated_text_parts: list[str] = [] last_input_tokens = 0 @@ -537,28 +536,6 @@ async def _run_loop_inner( if stop_reason: break - # Reject transport blanks before history normalization or observers - # can treat them as model output. Resampling must replay the original - # history, not an empty/whitespace assistant prefill. - if is_wholly_empty_response(response): - empty_completions += 1 - if empty_completions <= cfg.empty_completion_max_retries: - logger.warning( - "turn=%d response carried no content, no tool call, no " - "reasoning and no usage — resampling (%d/%d)", - turn, empty_completions, cfg.empty_completion_max_retries, - ) - turn -= 1 - continue - stop_reason = "empty_completion" - logger.warning( - "turn=%d still empty after %d resample(s) — stopping", - turn, cfg.empty_completion_max_retries, - ) - break - # Each consecutive empty episode has its own recovery budget. - empty_completions = 0 - ( parsed_calls, ctx, stop_reason, continue_to_next_turn, skip_tool_execution, @@ -577,6 +554,14 @@ async def _run_loop_inner( turn -= 1 continue + rejection = response_rejection_reason(response) + if rejection: + _answer_unexecuted_tool_calls( + messages, "the provider rejected this turn; this tool call was not executed.", + ) + stop_reason = rejection + break + # A reply the output cap cut off is the one case where we KNOW the model # was not finished, and it must never reach the ``no_tool`` branch below: # truncation and "the model chose to stop talking" are opposite signals @@ -922,8 +907,16 @@ async def _on_delta( context_token_limit_hint=cfg.context_token_limit, wall_deadline_remaining=runtime_hooks.wall_deadline_remaining, chain_fallback_active=runtime_hooks.chain_fallback_active, + empty_completion_max_retries=cfg.empty_completion_max_retries, ) except LLMCallExhausted as exhausted: + if exhausted.reason == "empty_completion": + metadata["llm_error"] = str(exhausted.last_exc) + metadata["llm_error_reason"] = exhausted.reason + return ( + None, "empty_completion", first_delta_at, llm_call_started, + time.perf_counter(), call_id, current_attempt_id, current_attempt_index, + ) if exhausted.reason == "wall_deadline": logger.warning( "agent_loop: wall deadline reached mid-turn %d; ending with wall_deadline for salvage: %s", diff --git a/agent_core/runtime/retriable.py b/agent_core/runtime/retriable.py index 01f0140..ef629d3 100644 --- a/agent_core/runtime/retriable.py +++ b/agent_core/runtime/retriable.py @@ -465,8 +465,8 @@ def is_empty_completion(err: BaseException) -> bool: layer but produced zero content tokens (reasoning-runaway, all-tokens-in-thinking, or an empty stream). Routed through :func:`is_retriable_with_fallback` so the caller retries the same key - first (a temperature>0 resample frequently recovers) and then advances - the chain to a different provider, which always recovers. + first (a temperature>0 resample can recover) and then may advance + a configured chain. Recovery remains bounded and can exhaust. """ blob = _stringify(err) return any(p.search(blob) for p in _EMPTY_COMPLETION_PATTERNS) @@ -496,7 +496,7 @@ def is_retriable_with_fallback(err: BaseException) -> bool: ``empty_completion`` and ``truncated_stream`` also route here: they aren't deterministic on the input (a temp>0 resample may recover), but the caller's same-key retry budget runs first, and advancing the chain afterwards is the - guaranteed recovery — so it belongs to the same predicate. + configured recovery — so it belongs to the same predicate. """ return ( is_overloaded_error(err) diff --git a/agent_core/stream_tools.py b/agent_core/stream_tools.py new file mode 100644 index 0000000..0e7abda --- /dev/null +++ b/agent_core/stream_tools.py @@ -0,0 +1,34 @@ +"""Shared wire-format assembly for streamed native tool calls.""" +from __future__ import annotations + +from typing import Any + +from agent_core.messages import ToolCall + + +class StreamToolCalls: + def __init__(self) -> None: + self._slots: dict[int, ToolCall] = {} + + def feed(self, deltas: list[dict[str, Any]]) -> None: + for delta in deltas: + index = delta.get("index") or 0 + slot = self._slots.setdefault(index, { + "id": "", "type": "function", "function": {"name": "", "arguments": ""}, + }) + if delta.get("id"): + slot["id"] = delta["id"] + if delta.get("name"): + slot["function"]["name"] += delta["name"] + if delta.get("arguments"): + slot["function"]["arguments"] += delta["arguments"] + + def complete(self) -> list[ToolCall]: + # Keep provider order by index; unexecutable nameless slots never reach + # history or middleware fingerprints. Argument validation stays with + # the caller's existing repair/parser path. + return [self._slots[i] for i in sorted(self._slots) if self._slots[i]["function"]["name"]] + + @property + def dropped_count(self) -> int: + return sum(not slot["function"]["name"] for slot in self._slots.values()) diff --git a/changes/unified-empty-recovery.breaking.md b/changes/unified-empty-recovery.breaking.md new file mode 100644 index 0000000..91b9104 --- /dev/null +++ b/changes/unified-empty-recovery.breaking.md @@ -0,0 +1,9 @@ +Empty completion recovery now runs inside one logical `call_llm` invocation in both chat and streaming modes. `empty_completion_max_retries` (default 2, also available to direct callers) supplies same-leg resamples independently of the generic `max_retries` allowance; the original logical/wall deadlines bound backoff and native fallback advances. Empty attempts are emitted as discarded/failed with reason `empty_completion`, including separately identified failed tool-argument replays, and the loop stops with `empty_completion` after terminal exhaustion. Native and nested `LLMFallbackChain` routing uses independent per-call cursors, preserves bindings and middleware, respects trigger barriers, and supports an explicit `empty_completion` trigger. Direct fallback streams buffer candidate-empty deltas until a real signal commits the leg, preventing discarded tool-argument fragments from contaminating a fallback response. The cooldown wrapper detects blanks before applying its existing policy. Consumers relying on streaming retry counts, stop/error reasons or accepted-attempt metrics should adopt these unified semantics. + +`LLMResponse` and `StreamDelta` gain optional `usage_source` provenance; adapters mark provider reports, while wrappers should mark synthetic counts as `estimated` or retain `estimated: true` in usage maps. Reported zero usage remains distinct from estimates. OpenAI Chat/Responses refusal fields and events are preserved, and explicit refusals/content filters now terminate under their own stop reasons even with a nudge policy, without executing accompanying tool calls. Product wrappers opting into native delayed-assembly failover must explicitly implement `for_logical_call` and `advance_empty_completion` while preserving themselves; dynamic attribute forwarding alone is not opt-in. + + +Empty OpenAI refusal placeholders beside text, tool calls or reasoning are ignored. Streams infer an otherwise empty refusal only at normal EOF, including clean closes without a finish-reason chunk; exceptions and cancellation preserve their original failure semantics. + + +Streamed middleware now receives terminal usage/provenance, provider/model and rejection metadata plus named native tool calls. This restores streamed token/cost accounting and loop detection. Estimates stay diagnostic and do not charge real cost, task budgets or billing aggregates; canonical cache read/write fields preserve zeros and cache-only reported calls. Rate correction handles zero estimates and real zero usage against each limiter's actual capped reservation, while unknown/estimated usage keeps the reservation. Failed or consumer-closed tool proposals are excluded from repeat detection even when their authentic reported usage is accounted. Reported streaming usage/cost and budget consumption will increase from the old undercount; downstream budgets/alerts should be recalibrated. The legacy CostSink signature remains unchanged; cache-aware aggregators receive the separate read/write fields. diff --git a/docs/llm-runtime-boundary.md b/docs/llm-runtime-boundary.md index d6e0e55..e20b437 100644 --- a/docs/llm-runtime-boundary.md +++ b/docs/llm-runtime-boundary.md @@ -107,41 +107,115 @@ When automatic mode suppresses deltas something asked for, the loop logs a warning. It used to be silent, which meant a profile configuring the reasoning-only watchdog on a gated protocol got no watchdog and no sign of it. -## A blank reply is a failure, not an answer - -Three signals arrive in the same shape — a reply with no tool call — and need -opposite treatment: - -| signal | what it means | handling | -|---|---|---| -| truncation (`finish_reason="length"` with text) | stopped mid-sentence | continue from the partial text (`truncation_max_continuations`) | -| a tool-less turn with text | the model chose to stop | `no_tool_behavior` decides | -| **nothing whatsoever** | the transport returned nothing | **resample** (`empty_completion_max_retries`), then `stopped_by="empty_completion"` | - -"Nothing whatsoever" means no visible text, no tool call, no reasoning AND no -usage. The usage test is what separates it from a real 0-token reply, which -still reports prompt tokens; a reasoning-only stream likewise stays with the -runaway guard. Only the wholly blank response is treated as a fault. - -It is caught twice. `_stream_llm_response` raises `LLMEmptyCompletion` so a -blank stream never becomes an `LLMResponse`; `is_empty_completion` then -resamples on the same key and advances the chain, which is the recovery it has -always documented for "an empty stream". `run_agent_loop` repeats the test for -anything that reaches it by another path — a non-streamed reply, or one a -product wrapper rebuilt. Both guards use the same predicate on the original -response, before reasoning extraction or display normalization. Signed or -opaque reasoning blocks, reported usage (even zero-valued), and structured -refusal details are not transport blanks. - -The loop rejects a blank before appending it to history or notifying response -observers. A resample replays the existing conversation without an empty or -whitespace assistant prefill; exhausting the retry budget likewise leaves -history intact. The budget applies to consecutive blank responses and resets -as soon as a non-empty response arrives, so separate failures in a productive -run each get their own recovery allowance. - -Why it is worth two guards: without them a blank reply takes the no-tool exit, -so the run ends with `stopped_by="no_tool"` and a trajectory that reads as a -clean finish, while every turn of work already done is discarded. On -ApodexHarness's 2026-10-06 GDPval batch over llm-hub that cost 10 of 19 -streamed trials, six of them inside five minutes and one on turn 1 after 37s. +## Completion classification and recovery + +The same raw-response classifier is used for chat and assembled streams, before +history normalization or accepted-attempt events: + +| Signal | Handling | +|---|---| +| Visible content, executable call, reasoning or opaque signed blocks | Preserve the response; existing tool/truncation/runaway policies apply. | +| Provider-reported usage, including reported zeros | Preserve it. Usage establishes a real response, not a useful answer. | +| `refusal`, `content_filter`, native refusal text/blocks or structured refusal details | Preserve the provider evidence and stop with `refusal` or `content_filter`, even under a nudge policy. Tools from that turn are not executed; recorded ids get synthetic replies for safe replay. | +| No content, calls, reasoning, reported usage or explicit rejection | Raise `LLMEmptyCompletion` and recover inside `call_llm`. | + +`LLMResponse.usage_source` and `StreamDelta.usage_source` are optional provenance +channels. Adapters stamp `provider` for reported usage. A wrapper that invents +counts should stamp `estimated`, or retain an `estimated: true` marker on its +usage map. Estimates are retained as estimated in attempt accounting but cannot +turn a transport blank into a valid response. Unmarked legacy usage is treated +conservatively as reported; there is no reliable way to reconstruct provenance +once a wrapper discards it. The classifier also reads legacy `usage_metadata` +and raw `response_metadata.token_usage` / `usage` channels. + +OpenAI-compatible gateways may send `refusal=""` as a placeholder beside text, +tool calls or reasoning (`reasoning_content` / `reasoning`). That placeholder is +not a decline when the turn carries output. Streaming inference waits for normal +EOF, including streams without `finish_reason`; errors and consumer cancellation +never turn a pending marker into a refusal. Protocol terminator validation runs +before this inference, so a truncated transport cannot become a completed decline. +Non-empty refusal text and explicit +provider rejection reasons retain their usual semantics. + +### One recovery owner in both transports + +`empty_completion_max_retries` defaults to two same-leg resamples. This allowance +is independent of `max_llm_retries`, which bounds generic failures and the +existing reasoning recovery. Resamples use the same retry backoff configuration, +conversation and bindings. A fresh logical call gets a fresh empty allowance; +resamples stay inside the current call and never spend additional loop turns. +They do not restart `logical_call_timeout_s`: admission, requests, backoff and +fallback all share the original deadline, with the run wall deadline taking +precedence when earlier. + +Empty requests produce exactly one finished attempt event, with reason +`empty_completion` and outcome `discarded` when recovery continues, or `failed` +when it stops. They never enter assistant history or response observers. The +loop translates terminal empty exhaustion to `stopped_by="empty_completion"`; +deadline exhaustion keeps its own error/deadline reason. The opportunistic chat +replay used to repair streamed tool arguments also gets a separate attempt id +before its request starts. A blank/failed replay is recorded as failed while the +original streamed response remains deliverable; each request retains its own +latency and accounting, rather than charging replay time to the original. + +### Fallback routing and wrappers + +`LLMFallbackChain` allocates a private cursor for each logical call. After the +serving leg spends its empty allowance, it advances only if that leg's triggers +match: `any_error` or the optional `empty_completion` trigger. An explicit empty +trigger tuple is a barrier. Each subsequent leg gets the same allowance, bounded +by the finite chain and the shared deadline. An ordinary HTTP failure that +already selected a later leg pins empty resamples to that actual leg. + +Cursors survive tool, temperature, token and session bindings, nested provider +stamping chains and middleware proxies. Cached clients are never moved to another +leg, so concurrent calls and the next turn start independently. Direct native +chain callers detect blanks inside the chain and apply their configured triggers. +Candidate-empty deltas (including nameless tool argument fragments) are buffered +until a real signal commits that leg, so discarded fragments cannot corrupt the +next leg's tool calls. Transport heartbeats still pass through immediately. +`CooldownFallbackLLM` similarly detects them before its existing retry/degrade +policy and preserves its cooldown and tracing events. + +Transparent product wrappers opt into delayed-assembly recovery by explicitly +implementing `for_logical_call()` and `advance_empty_completion(error)`. The +first returns a client preserving the wrapper around fresh inner routing state; +the second returns whether the current call advanced. `LLMProxy` implements both +and keeps its middleware, role and shared atomic counter. A `__getattr__` +forwarder alone is not opt-in, because invoking the inner preparation method +could silently bypass the wrapper. Wrappers without these hooks still receive +empty classification and same-key recovery; products using an external chain +can inject `chain_fallback_active` to receive the terminal error for routing. + + +### Streamed middleware consumers + +`LLMProxy.stream` supplies `after_llm` with terminal usage/provenance, model, +provider and rejection metadata, plus assembled **named native tool calls**. +The proxy and loop use the same indexed tool accumulator: names/arguments are +concatenated in provider order, nameless slots are dropped, and each retry +starts with fresh state. Heartbeats are never model output. The middleware +response remains a passive view; signed-block replay stays owned by the loop. + +Token accounting charges only reported usage. `usage_source="estimated"` and +`estimated: true` maps remain visible in tracing/attempt diagnostics, but never +enter `CostSink`, task budget charges, billing events or `UsageAggregator`. +Unmarked legacy usage retains its reported interpretation. Reported cache reads +and writes use `cache_read_tokens` / `cache_write_tokens` first, including real +zeros, with legacy read/creation aliases as fallbacks. Cache-only reported calls +still reach cache-aware usage aggregation; the existing four-argument `CostSink` +contract and prompt/completion budget units remain unchanged. + +Rate correction uses reported totals (or reported prompt/completion counts), +including zero. Missing, invalid or estimated usage leaves the admission +reservation in place. Corrections use the **capped reservation actually taken**, +not the original prompt estimate, and records are isolated by limiter instance +so multiple quota layers cannot overwrite one another's state. Request quota is +reserved independently of whether token usage is later available. + +`LoopDetectionMiddleware` now receives native streamed tool calls, so completed +repeated calls can trigger the existing strategy-switch hint. Failed or +consumer-closed stream proposals do not enter loop history; authentic usage +from those requests is still eligible for accounting. Streaming cost/token +reports therefore increase from the previously omitted values. Consumers should +recheck budgets and alerts calibrated against the old undercount. diff --git a/tests/test_anthropic_latest_models.py b/tests/test_anthropic_latest_models.py index 1b553d0..f20fe0f 100644 --- a/tests/test_anthropic_latest_models.py +++ b/tests/test_anthropic_latest_models.py @@ -330,7 +330,7 @@ async def chat(self, messages, **kwargs): llm = BlankFirst() result = await run_agent_loop( system_prompt="s", user_message="u", llm=llm, tools=[], - config=LoopConfig(max_turns=1, max_llm_retries=1), + config=LoopConfig(max_turns=1, max_llm_retries=1, retry_wait_fixed=0), model_profile=ModelProfile( model_id=client.model, provider="anthropic", protocol="anthropic", thinking_format="content_block", diff --git a/tests/test_completion_provider_signals.py b/tests/test_completion_provider_signals.py new file mode 100644 index 0000000..08b7eaf --- /dev/null +++ b/tests/test_completion_provider_signals.py @@ -0,0 +1,312 @@ +"""Real SDK serialization/parsing over local HTTP mocks; no API calls.""" +from __future__ import annotations + +import asyncio +import json + +import httpx +import pytest +from openai import AsyncOpenAI, DefaultAsyncHttpxClient + +from agent_core.loop_types import LoopConfig, LoopPolicy +from agent_core.providers.openai_chat import OpenAIClient +from agent_core.providers.openai_responses import OpenAIResponsesClient +from agent_core.runtime.loop.agent_loop import run_agent_loop +from agent_core.runtime.loop.model_profile import ModelProfile + +if issubclass(DefaultAsyncHttpxClient, httpx.AsyncClient): + sdk_httpx = httpx +else: + import httpx2 as sdk_httpx + + +def sse(events): + return "".join(f"data: {json.dumps(event)}\n\n" for event in events) + "data: [DONE]\n\n" + + +async def run(client, streaming, protocol): + return await run_agent_loop(system_prompt="s", user_message="u", llm=client, tools=[], + config=LoopConfig(max_turns=2, max_llm_retries=1, stream_llm_tokens=streaming, + retry_wait_fixed=0, loop_policy=LoopPolicy(no_tool_behavior="nudge")), + model_profile=ModelProfile(model_id="m", provider="openai", protocol=protocol, + thinking_format="content_block" if protocol == "responses" else "none")) + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("refusal,finish,expected", [ + ("request declined", "stop", "refusal"), + ("", "stop", "refusal"), + (None, "content_filter", "content_filter"), +]) +async def test_chat_refusal_and_filter_without_usage_survive_sdk_boundary(streaming, refusal, finish, expected): + requests = [] + def respond(request): + requests.append(json.loads(request.content)) + if streaming: + events = [{ + "id": "id", "object": "chat.completion.chunk", "created": 1, "model": "m", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": None, "refusal": refusal}, "finish_reason": finish}], + }] + return sdk_httpx.Response(200, text=sse(events), headers={"content-type": "text/event-stream"}) + return sdk_httpx.Response(200, json={ + "id": "id", "object": "chat.completion", "created": 1, "model": "m", + "choices": [{"index": 0, "message": {"role": "assistant", "content": None, "refusal": refusal}, "finish_reason": finish}], + }) + client = OpenAIClient("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + result = await run(client, streaming, "chat_completions") + assert len(requests) == 1 + assert result.stopped_by == expected + assert result.final_content == (refusal or "") + assert result.turns_used == 1 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_chat_empty_refusal_beside_content_is_not_a_refusal(streaming): + def respond(request): + if streaming: + events = [{ + "id": "id", "object": "chat.completion.chunk", "created": 1, "model": "m", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": None, "refusal": ""}, "finish_reason": None}], + }, { + "id": "id", "object": "chat.completion.chunk", "created": 1, "model": "m", + "choices": [{"index": 0, "delta": {"content": "hello", "refusal": ""}, "finish_reason": "stop"}], + }] + return sdk_httpx.Response(200, text=sse(events), headers={"content-type": "text/event-stream"}) + return sdk_httpx.Response(200, json={ + "id": "id", "object": "chat.completion", "created": 1, "model": "m", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hello", "refusal": ""}, "finish_reason": "stop"}], + }) + client = OpenAIClient("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + result = await run(client, streaming, "chat_completions") + assert result.stopped_by != "refusal" + assert result.final_content == "hello" + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("delta_events", [False, True]) +@pytest.mark.parametrize("refusal", ["request declined", ""]) +async def test_responses_refusal_parts_and_events_without_usage_survive_sdk_boundary(streaming, delta_events, refusal): + requests = [] + response = { + "id": "resp", "object": "response", "created_at": 1, "model": "m", "status": "completed", + "output": [{"type": "message", "id": "msg", "role": "assistant", "status": "completed", + "content": [{"type": "refusal", "refusal": refusal}]}], + "usage": None, + } + def respond(request): + requests.append(json.loads(request.content)) + if streaming: + events = [{"type": "response.refusal.delta", "delta": refusal, "output_index": 0, + "content_index": 0, "item_id": "msg", "sequence_number": 1}] if delta_events else [] + events.append({"type": "response.completed", "response": response, "sequence_number": 2}) + return sdk_httpx.Response(200, text=sse(events), headers={"content-type": "text/event-stream"}) + return sdk_httpx.Response(200, json=response) + client = OpenAIResponsesClient("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + result = await run(client, streaming, "responses") + assert len(requests) == 1 + assert result.stopped_by == "refusal" + assert result.final_content == refusal + assert result.turns_used == 1 + + +@pytest.mark.parametrize("protocol", ["chat_completions", "responses"]) +@pytest.mark.parametrize("streaming", [False, True]) +async def test_gateway_estimated_usage_marker_survives_real_sdk_normalization(protocol, streaming): + requests = [] + def respond(request): + requests.append(json.loads(request.content)) + blank = len(requests) == 1 + text = "" if blank else "answer" + if protocol == "chat_completions": + payload = { + "id": "id", "object": "chat.completion.chunk" if streaming else "chat.completion", + "created": 1, "model": "m", "choices": [{"index": 0, "finish_reason": "stop", + "delta" if streaming else "message": {"role": "assistant", "content": text}}], + "usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0, "estimated": True} if blank else None, + } + events = [payload] + else: + payload = { + "id": "resp", "object": "response", "created_at": 1, "model": "m", "status": "completed", + "output": [{"type": "message", "id": "msg", "role": "assistant", "status": "completed", + "content": [{"type": "output_text", "text": text, "annotations": []}]}], + "usage": {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0, "estimated": True} if blank else None, + } + events = [{"type": "response.output_text.delta", "delta": text, "sequence_number": 1}, + {"type": "response.completed", "response": payload, "sequence_number": 2}] + if streaming: + return sdk_httpx.Response(200, text=sse(events), headers={"content-type": "text/event-stream"}) + return sdk_httpx.Response(200, json=payload) + cls = OpenAIClient if protocol == "chat_completions" else OpenAIResponsesClient + client = cls("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + result = await run_agent_loop(system_prompt="s", user_message="u", llm=client, tools=[], + config=LoopConfig(max_turns=1, max_llm_retries=1, stream_llm_tokens=streaming, + retry_wait_fixed=0, loop_policy=LoopPolicy(no_tool_behavior="stop")), + model_profile=ModelProfile(model_id="m", provider="openai", protocol=protocol, + thinking_format="content_block" if protocol == "responses" else "none")) + assert len(requests) == 2 + assert result.final_content == "answer" + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("reasoning_field", ["reasoning_content", "reasoning"]) +@pytest.mark.parametrize("finish", ["stop", None]) +async def test_empty_refusal_beside_reasoning_continues_to_the_answer(streaming, reasoning_field, finish): + requests = [] + def respond(request): + requests.append(json.loads(request.content)) + first = len(requests) == 1 + message = {"role": "assistant", "content": None, "refusal": "", reasoning_field: "real thinking"} if first else {"role": "assistant", "content": "answer"} + payload = {"id": "id", "created": 1, "model": "m", + "object": "chat.completion.chunk" if streaming else "chat.completion", + "choices": [{"index": 0, "delta" if streaming else "message": message, + "finish_reason": finish if first else "stop"}]} + return sdk_httpx.Response(200, text=sse([payload]), headers={"content-type": "text/event-stream"}) if streaming else sdk_httpx.Response(200, json=payload) + client = OpenAIClient("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + result = await run_agent_loop(system_prompt="s", user_message="u", llm=client, tools=[], + config=LoopConfig(max_turns=2, max_llm_retries=1, stream_llm_tokens=streaming, + retry_wait_fixed=0, loop_policy=LoopPolicy(no_tool_behavior="nudge")), + model_profile=ModelProfile(model_id="m", provider="openai", thinking_format="reasoning_content")) + assert len(requests) == 2 + assert result.final_content == "answer" + assert result.stopped_by == "no_tool" + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("finish", ["stop", None]) +@pytest.mark.parametrize("with_usage", [False, True]) +async def test_empty_refusal_at_clean_eof_is_preserved_with_or_without_finish_reason(streaming, finish, with_usage): + requests = [] + def respond(request): + requests.append(json.loads(request.content)) + payload = {"id": "id", "created": 1, "model": "m", + "object": "chat.completion.chunk" if streaming else "chat.completion", + "choices": [{"index": 0, "delta" if streaming else "message": { + "role": "assistant", "content": None, "refusal": ""}, "finish_reason": finish}], + "usage": {"prompt_tokens": 10, "completion_tokens": 0, "total_tokens": 10} if with_usage else None} + return sdk_httpx.Response(200, text=sse([payload]), headers={"content-type": "text/event-stream"}) if streaming else sdk_httpx.Response(200, json=payload) + client = OpenAIClient("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + result = await run(client, streaming, "chat_completions") + assert len(requests) == 1 + assert result.stopped_by == "refusal" + assert result.final_content == "" + + +@pytest.mark.parametrize("output", ["text", "reasoning", "tool"]) +async def test_empty_refusal_waits_for_the_entire_stream_before_inference(output): + from agent_core.completion import response_rejection_reason + from agent_core.runtime.loop._streaming import _stream_llm_response + + actual = {"content": "answer"} if output == "text" else {"reasoning_content": "thinking"} if output == "reasoning" else {"tool_calls": [{"index": 0, "id": "call", "type": "function", "function": {"name": "echo", "arguments": "{}"}}]} + def respond(request): + chunks = [{"refusal": "", "content": None}, actual] + events = [{"id": "id", "created": 1, "model": "m", "object": "chat.completion.chunk", + "choices": [{"index": 0, "delta": delta, "finish_reason": "stop" if index == 0 else None}]} + for index, delta in enumerate(chunks)] + return sdk_httpx.Response(200, text=sse(events), headers={"content-type": "text/event-stream"}) + client = OpenAIClient("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + async def noop(*_args, **_kwargs): + pass + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + response = await _stream_llm_response(client, [], 10, noop) + assert response_rejection_reason(response) == "" + if output == "text": + assert response.content == "answer" + elif output == "reasoning": + assert response.reasoning_content == "thinking" + else: + assert response.tool_calls[0]["function"]["name"] == "echo" + + +@pytest.mark.parametrize("abort", ["error", "close", "cancel"]) +async def test_incomplete_stream_does_not_infer_refusal_on_error_or_consumer_close(abort): + from types import SimpleNamespace + + async def events(): + yield SimpleNamespace(choices=[SimpleNamespace(delta=SimpleNamespace(refusal="", content=None), finish_reason=None)], usage=None, model="m") + if abort == "cancel": + await asyncio.Event().wait() + raise RuntimeError("transport broke") + client = OpenAIClient("m", api_key="test") + async def open_stream(_kwargs): + return events() + client._open_stream = open_stream + seen = [] + try: + stream = client.stream([]) + if abort == "error": + with pytest.raises(RuntimeError, match="transport broke"): + async for delta in stream: + seen.append(delta) + elif abort == "close": + seen.append(await anext(stream)) + await stream.aclose() + else: + ready = asyncio.Event() + async def consume(): + async for delta in stream: + seen.append(delta) + ready.set() + task = asyncio.create_task(consume()) + await asyncio.wait_for(ready.wait(), 1) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + await client._client.close() + assert len(seen) == 1 + assert all(not delta.stop_details for delta in seen) + + +@pytest.mark.parametrize("protocol", ["chat_completions", "responses"]) +async def test_truncated_refusal_stream_is_never_accepted_as_a_completed_decline(protocol, monkeypatch): + from agent_core.errors import LLMOpenAITruncatedStream + + monkeypatch.setenv("AGENT_CORE_STREAM_REQUIRE_TERMINATOR", "1") + def respond(request): + event = {"id": "id", "object": "chat.completion.chunk", "created": 1, "model": "m", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": None, "refusal": ""}, "finish_reason": None}]} + if protocol == "responses": + event = {"type": "response.refusal.delta", "delta": "declined", "output_index": 0, + "content_index": 0, "item_id": "msg", "sequence_number": 1} + # Neither [DONE]/finish_reason nor a Responses terminal event arrives. + return sdk_httpx.Response(200, text=f"data: {json.dumps(event)}\n\n", headers={"content-type": "text/event-stream"}) + cls = OpenAIClient if protocol == "chat_completions" else OpenAIResponsesClient + client = cls("m", api_key="test", base_url="https://openai.invalid") + await client._client.close() + seen = [] + async with AsyncOpenAI(api_key="test", base_url="https://openai.invalid", max_retries=0, + http_client=DefaultAsyncHttpxClient(transport=sdk_httpx.MockTransport(respond))) as sdk: + client._client = sdk + with pytest.raises(LLMOpenAITruncatedStream): + async for delta in client.stream([]): + seen.append(delta) + if protocol == "chat_completions": + assert all(not delta.stop_details for delta in seen) diff --git a/tests/test_completion_recovery.py b/tests/test_completion_recovery.py new file mode 100644 index 0000000..3c44eca --- /dev/null +++ b/tests/test_completion_recovery.py @@ -0,0 +1,500 @@ +"""Recovery parity and routing contracts across chat/stream transports.""" +from __future__ import annotations + +import asyncio +from copy import deepcopy +from types import SimpleNamespace + +import pytest + +from agent_core.completion import is_wholly_empty_response +from agent_core.errors import LLMCallExhausted +from agent_core.llm import LLMResponse, StreamDelta +from agent_core.loop_types import LoopConfig, LoopPolicy +from agent_core.messages import user_msg +from agent_core.providers.fallback import FallbackEntry, LLMFallbackChain, with_provider_stamp +from agent_core.runtime.loop._bind import ( + bind_max_tokens, + bind_session_id, + bind_temperature, + bind_tools, +) +from agent_core.runtime.loop._call import call_llm +from agent_core.runtime.loop.agent_loop import run_agent_loop + + +class Script: + def __init__(self, actions, *, model="m", delay=0): + self.actions = list(actions) + self.model = model + self.delay = delay + self.calls = [] + self.cancelled = False + + async def _next(self, messages, kwargs): + self.calls.append((deepcopy(messages), dict(kwargs))) + try: + await asyncio.sleep(self.delay) + except asyncio.CancelledError: + self.cancelled = True + raise + action = self.actions.pop(0) + if isinstance(action, BaseException): + raise action + return action + + async def chat(self, messages, **kwargs): + return await self._next(messages, kwargs) + + async def stream(self, messages, **kwargs): + response = await self._next(messages, kwargs) + if isinstance(response, list): + for delta in response: + yield delta + return + yield StreamDelta( + content=response.content if isinstance(response.content, str) else "", + reasoning_blocks=response.content if isinstance(response.content, list) else [], + reasoning_content=response.reasoning_content, + usage=response.usage, usage_source=response.usage_source or response.response_metadata.get("usage_source", ""), + finish_reason=response.finish_reason, model=self.model, + refusal=response.response_metadata.get("refusal", ""), + stop_details=response.response_metadata.get("stop_details", {}), + ) + + +async def noop(*_args, **_kwargs): + pass + + +async def invoke(llm, streaming, events=None, **kwargs): + async def record(event): + if events is not None: + events.append(dict(event)) + + return await call_llm( + llm, [user_msg("u")], timeout=10, max_retries=kwargs.pop("max_retries", 1), turn=1, + on_delta=noop if streaming else None, retry_wait_fixed=0, + on_attempt=record, **kwargs, + ) + + +def finished(events): + return [e for e in events if e["phase"] == "finished"] + + +def assert_balanced(events, count): + assert [e["attempt_index"] for e in events if e["phase"] == "started"] == list(range(1, count + 1)) + assert [e["attempt_index"] for e in finished(events)] == list(range(1, count + 1)) + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_blank_recovery_is_one_logical_call_with_discarded_attempts(streaming): + llm = Script([LLMResponse(), LLMResponse(content=" \n"), LLMResponse(content="answer")]) + events = [] + response = await invoke(llm, streaming, events) + assert response.content == "answer" + assert len(llm.calls) == 3 + assert all(messages == [user_msg("u")] for messages, _ in llm.calls) + assert_balanced(events, 3) + assert [(e["outcome"], e["reason"]) for e in finished(events)] == [ + ("discarded", "empty_completion"), ("discarded", "empty_completion"), ("accepted", ""), + ] + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("budget", [0, 1, 2]) +async def test_empty_budget_is_independent_of_generic_retry_allowance(streaming, budget): + llm = Script([LLMResponse()] * (budget + 1)) + events = [] + with pytest.raises(LLMCallExhausted) as caught: + await invoke(llm, streaming, events, empty_completion_max_retries=budget) + assert caught.value.reason == "empty_completion" + assert len(llm.calls) == budget + 1 + assert_balanced(events, budget + 1) + assert finished(events)[-1]["outcome"] == "failed" + assert all(e["reason"] == "empty_completion" for e in finished(events)) + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_mixed_generic_and_blank_errors_do_not_extend_generic_allowance(streaming): + llm = Script([LLMResponse(), RuntimeError("network reset"), LLMResponse(), LLMResponse(content="answer")]) + events = [] + assert (await invoke(llm, streaming, events, max_retries=2)).content == "answer" + assert_balanced(events, 4) + assert [e["reason"] for e in finished(events)] == ["empty_completion", "transient_error", "empty_completion", ""] + llm = Script([LLMResponse(), RuntimeError("network reset"), LLMResponse(content="unreachable")]) + with pytest.raises(LLMCallExhausted) as caught: + await invoke(llm, streaming) + assert caught.value.reason == "exhausted" + assert len(llm.calls) == 2 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_native_chain_advances_after_same_leg_resamples_preserving_bindings(streaming): + primary = Script([LLMResponse()] * 3 + [LLMResponse(content="next call")], model="primary") + secondary = Script([LLMResponse(content="fallback")], model="secondary") + chain = LLMFallbackChain([ + FallbackEntry(primary, provider="primary-vendor"), FallbackEntry(secondary, provider="secondary-vendor"), + ]) + schema = {"type": "function", "function": {"name": "echo", "parameters": {"type": "object"}}} + llm = bind_max_tokens(bind_temperature(bind_session_id(bind_tools(chain, [schema]), "session"), .7), 123) + events = [] + result = await invoke(llm, streaming, events) + assert result.content == "fallback" + assert result.response_metadata["provider_actually_used"] == "secondary-vendor" + if not streaming: + assert result.response_metadata["fallback_used"] == 1 + assert len(primary.calls) == 3 + assert len(secondary.calls) == 1 + for _, kw in primary.calls + secondary.calls: + assert kw["tools"] == [schema] + assert kw["temperature"] == .7 + assert kw["max_tokens"] == 123 + assert kw["extra_headers"] == {"x-upstream-session-id": "session"} + assert_balanced(events, 4) + assert finished(events)[2]["recovery_action"] == "chain_advance" + # Cached chain and next logical call always start independently. + assert (await invoke(llm, streaming)).content == "next call" + assert chain._start_index == 0 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_provider_stamp_wrapper_preserves_nested_chain_recovery(streaming): + first = Script([LLMResponse()] * 3) + second = Script([LLMResponse(content="answer")]) + llm = with_provider_stamp(LLMFallbackChain.from_models([first, second]), "vendor") + assert (await invoke(llm, streaming)).content == "answer" + assert (len(first.calls), len(second.calls)) == (3, 1) + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("triggers", [(), ("timeout",)]) +async def test_serving_leg_trigger_barrier_is_respected(streaming, triggers): + first = Script([LLMResponse()] * 3) + second = Script([LLMResponse(content="unreachable")]) + llm = LLMFallbackChain([FallbackEntry(first, triggers=triggers), FallbackEntry(second)]) + with pytest.raises(LLMCallExhausted) as caught: + await invoke(llm, streaming) + assert caught.value.reason == "empty_completion" + assert len(second.calls) == 0 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_ordinary_failure_then_blank_routes_from_actual_serving_leg(streaming): + first = Script([RuntimeError("upstream error")]) + second = Script([LLMResponse()] * 3) + third = Script([LLMResponse(content="answer")]) + llm = LLMFallbackChain.from_models([first, second, third]) + assert (await invoke(llm, streaming)).content == "answer" + assert [len(client.calls) for client in (first, second, third)] == [1, 3, 1] + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_shared_native_chain_has_independent_concurrent_call_cursors(streaming): + class Primary: + model = "primary" + async def chat(self, messages, **_kwargs): + await asyncio.sleep(.001) + return LLMResponse() if messages[-1]["content"] == "bad" else LLMResponse(content="primary") + async def stream(self, messages, **kwargs): + r = await self.chat(messages, **kwargs) + yield StreamDelta(content=r.content) + secondary = Script([LLMResponse(content="secondary")]) + chain = LLMFallbackChain.from_models([Primary(), secondary]) + async def one(text): + return await call_llm(chain, [user_msg(text)], 10, 1, 1, + on_delta=noop if streaming else None, retry_wait_fixed=0) + bad, good = await asyncio.gather(one("bad"), one("good")) + assert (bad.content, good.content) == ("secondary", "primary") + assert len(secondary.calls) == 1 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_empty_episode_and_fallback_share_one_logical_deadline(streaming, monkeypatch): + monkeypatch.setattr("agent_core.runtime.loop._call._WALL_DEADLINE_FLOOR_S", .001) + primary = Script([LLMResponse()] * 3, delay=.02) + secondary = Script([LLMResponse(content="unreachable")]) + events = [] + with pytest.raises(LLMCallExhausted) as caught: + await invoke(LLMFallbackChain.from_models([primary, secondary]), streaming, events, + logical_call_timeout_s=.05) + assert caught.value.reason == "logical_call_deadline" + assert len(secondary.calls) == 0 + assert len(primary.calls) <= 3 + assert finished(events)[-1]["reason"] == "logical_call_deadline" + assert all(e["outcome"] != "accepted" for e in finished(events)) + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_run_wall_deadline_precedes_logical_deadline(streaming, monkeypatch): + monkeypatch.setattr("agent_core.runtime.loop._call._WALL_DEADLINE_FLOOR_S", .001) + llm = Script([LLMResponse()] * 3) + with pytest.raises(LLMCallExhausted) as caught: + await invoke(llm, streaming, logical_call_timeout_s=10, + wall_deadline_remaining=lambda: .05 if not llm.calls else 0) + assert caught.value.reason == "wall_deadline" + assert len(llm.calls) <= 1 + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("response", [ + LLMResponse(usage={"prompt_tokens": 0, "completion_tokens": 0, "estimated": True}), + LLMResponse(usage={"prompt_tokens": 100}, usage_source="estimated"), + LLMResponse(usage={"prompt_tokens": 100}, response_metadata={"usage_source": "estimated"}), +]) +async def test_synthetic_usage_does_not_hide_a_transport_blank(streaming, response): + llm = Script([response, LLMResponse(content="answer")]) + assert (await invoke(llm, streaming)).content == "answer" + assert len(llm.calls) == 2 + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("response", [ + LLMResponse(usage={"prompt_tokens": 0, "completion_tokens": 0}, usage_source="provider"), + LLMResponse(reasoning_content="thinking"), + LLMResponse(content=[{"type": "redacted_thinking", "data": "opaque"}]), +]) +async def test_actual_usage_and_reasoning_do_not_trigger_empty_recovery(streaming, response): + llm = Script([response]) + assert await invoke(llm, streaming) is not None + assert len(llm.calls) == 1 + + +@pytest.mark.parametrize("response", [ + SimpleNamespace(content="", usage_metadata={"input_tokens": 0, "estimated": True}), + SimpleNamespace(content="", response_metadata={"token_usage": {"prompt_tokens": 0, "estimated": True}}), +]) +def test_legacy_wrapper_usage_estimates_are_not_provider_evidence(response): + assert is_wholly_empty_response(response) + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("reason", ["refusal", "content_filter"]) +async def test_explicit_rejections_stop_even_with_nudge_policy_and_missing_usage(streaming, reason): + llm = Script([LLMResponse(finish_reason=reason)]) + result = await run_agent_loop(system_prompt="s", user_message="u", llm=llm, tools=[], + config=LoopConfig(max_turns=5, max_llm_retries=1, stream_llm_tokens=streaming, + loop_policy=LoopPolicy(no_tool_behavior="nudge"), retry_wait_fixed=0)) + assert result.stopped_by == reason + assert result.turns_used == 1 + assert len(llm.calls) == 1 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_outer_chain_receives_exhausted_blank_after_same_key_budget(streaming): + llm = Script([LLMResponse()] * 3) + events = [] + with pytest.raises(LLMCallExhausted) as caught: + await invoke(llm, streaming, events, chain_fallback_active=lambda: True) + assert caught.value.reason == "chain_advance" + assert len(llm.calls) == 3 + assert_balanced(events, 3) + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_direct_native_chain_recovers_without_runtime_cursor(streaming): + blank = Script([LLMResponse()]) + good = Script([LLMResponse(content="answer")]) + chain = LLMFallbackChain.from_models([blank, good]) + if streaming: + chunks = [d async for d in chain.stream([user_msg("u")])] + assert chunks[-1].content == "answer" + else: + assert (await chain.chat([user_msg("u")])).content == "answer" + assert (len(blank.calls), len(good.calls)) == (1, 1) + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_legacy_empty_errors_share_the_same_recovery_policy(streaming): + llm = Script([RuntimeError("no generation chunks returned"), LLMResponse(content="answer")]) + events = [] + assert (await invoke(llm, streaming, events)).content == "answer" + assert [e["reason"] for e in finished(events)] == ["empty_completion", ""] + assert_balanced(events, 2) + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_middleware_proxy_preserves_hooks_and_shared_counter_across_recovery(streaming): + from agent_core.components.middleware.llm.proxy import LLMProxy + + class Middleware: + def __init__(self): + self.before = [] + self.after = [] + self.chunks = [] + async def run_before(self, ctx, messages): + self.before.append((ctx.call_index, ctx.role_id)) + return messages + async def run_after(self, ctx, response, **kwargs): + self.after.append(ctx.call_index) + return response + async def run_on_chunk(self, ctx, chunk, content): + self.chunks.append(ctx.call_index) + return False + async def run_on_llm_error(self, ctx, error, attempt): + return False + + primary = Script([LLMResponse()] * 3) + secondary = Script([LLMResponse(content="answer"), LLMResponse(content="next")]) + middleware = Middleware() + proxy = LLMProxy(with_provider_stamp(LLMFallbackChain.from_models([primary, secondary]), "p"), middleware, "role") + assert (await invoke(proxy, streaming)).content == "answer" + assert middleware.before == [(1, "role"), (2, "role"), (3, "role"), (4, "role")] + assert middleware.after == [1, 2, 3, 4] + assert proxy.call_counter == 4 + if streaming: + assert middleware.chunks == [1, 2, 3, 4] + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_unknown_transparent_wrapper_is_not_silently_unwrapped(streaming): + class Wrapper: + def __init__(self, inner): + self.inner = inner + self.calls = 0 + def __getattr__(self, name): + return getattr(self.inner, name) + async def chat(self, messages, **kwargs): + self.calls += 1 + return await self.inner.chat(messages, **kwargs) + async def stream(self, messages, **kwargs): + self.calls += 1 + async for delta in self.inner.stream(messages, **kwargs): + yield delta + wrapper = Wrapper(with_provider_stamp(Script([LLMResponse(content="answer")]), "p")) + assert (await invoke(wrapper, streaming)).content == "answer" + assert wrapper.calls == 1 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_cooldown_wrapper_detects_blanks_before_its_existing_degrade_policy(streaming): + from agent_core.providers.fallback import CooldownFallbackLLM + + async def sleep(_delay): + pass + primary = Script([LLMResponse()] * 2) + secondary = Script([LLMResponse(content="fallback"), LLMResponse(content="cooldown")]) + events = [] + async def hook(name, payload): + events.append((name, payload)) + llm = CooldownFallbackLLM(primary, secondary, max_retries=2, cooldown_seconds=60, + sleep=sleep, event_hook=hook) + assert (await invoke(llm, streaming)).content == "fallback" + assert (await invoke(llm, streaming)).content == "cooldown" + assert len(primary.calls) == 2 + assert len(secondary.calls) == 2 + assert sum(name == "error" for name, _ in events) == 2 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_explicit_empty_completion_trigger_routes_without_enabling_all_errors(streaming): + primary = Script([LLMResponse()] * 3) + secondary = Script([LLMResponse(content="answer")]) + llm = LLMFallbackChain([FallbackEntry(primary, triggers=("empty_completion",)), FallbackEntry(secondary)]) + assert (await invoke(llm, streaming)).content == "answer" + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_exhausted_last_leg_is_bounded_and_every_attempt_finishes(streaming): + primary = Script([LLMResponse()] * 3) + secondary = Script([LLMResponse()] * 3) + events = [] + with pytest.raises(LLMCallExhausted) as caught: + await invoke(LLMFallbackChain.from_models([primary, secondary]), streaming, events) + assert caught.value.reason == "empty_completion" + assert (len(primary.calls), len(secondary.calls)) == (3, 3) + assert_balanced(events, 6) + assert [e["recovery_action"] for e in finished(events)] == [ + "retry_same_key", "retry_same_key", "chain_advance", "retry_same_key", "retry_same_key", "empty_completion", + ] + + +async def test_estimated_usage_remains_labeled_on_discarded_attempt(): + llm = Script([LLMResponse(usage={"prompt_tokens": 100}, usage_source="estimated"), LLMResponse(content="answer")]) + events = [] + await invoke(llm, False, events) + assert finished(events)[0]["usage"]["estimated"] is True + + +async def test_rejection_does_not_execute_tools_and_keeps_history_replayable(): + class Tool: + name = "echo" + calls = 0 + def to_openai_schema(self): + return {"type": "function", "function": {"name": "echo", "parameters": {"type": "object"}}} + async def ainvoke(self, args): + self.calls += 1 + return "unreachable" + tool = Tool() + llm = Script([LLMResponse(finish_reason="content_filter", tool_calls=[{ + "id": "call", "type": "function", "function": {"name": "echo", "arguments": "{}"}, + }])]) + result = await run_agent_loop(system_prompt="s", user_message="u", llm=llm, tools=[tool], + config=LoopConfig(max_turns=2, max_llm_retries=1, retry_wait_fixed=0)) + assert result.stopped_by == "content_filter" + assert tool.calls == 0 + assert result.messages[-1]["role"] == "tool" + assert result.messages[-1]["tool_call_id"] == "call" + + +@pytest.mark.parametrize("wrapper", ["chain", "cooldown"]) +async def test_direct_empty_leg_fragments_cannot_corrupt_fallback_tool_arguments(wrapper): + from agent_core.providers.fallback import CooldownFallbackLLM + from agent_core.runtime.loop._streaming import _stream_llm_response + + primary = Script([[StreamDelta(tool_call_deltas=[{"index": 0, "id": "broken", "arguments": '{"x":'}])]]) + secondary = Script([[StreamDelta(tool_call_deltas=[{ + "index": 0, "id": "good", "name": "tool", "arguments": '{"x":1}', + }])]]) + llm = LLMFallbackChain.from_models([primary, secondary]) if wrapper == "chain" else CooldownFallbackLLM(primary, secondary, max_retries=1) + result = await _stream_llm_response(llm, [user_msg("u")], 10, noop) + assert result.tool_calls == [{ + "id": "good", "type": "function", "function": {"name": "tool", "arguments": '{"x":1}'}, + }] + + +@pytest.mark.parametrize("wrapper", ["chain", "cooldown"]) +async def test_candidate_empty_fragments_survive_if_the_same_leg_later_gets_a_name(wrapper): + from agent_core.providers.fallback import CooldownFallbackLLM + from agent_core.runtime.loop._streaming import _stream_llm_response + + primary = Script([[ + StreamDelta(tool_call_deltas=[{"index": 0, "id": "good", "arguments": '{"x":'}]), + StreamDelta(tool_call_deltas=[{"index": 0, "name": "tool", "arguments": '1}'}]), + ]]) + secondary = Script([LLMResponse(content="unreachable")]) + llm = LLMFallbackChain.from_models([primary, secondary]) if wrapper == "chain" else CooldownFallbackLLM(primary, secondary, max_retries=1) + result = await _stream_llm_response(llm, [user_msg("u")], 10, noop) + assert result.tool_calls[0]["function"]["arguments"] == '{"x":1}' + assert len(secondary.calls) == 0 + + +@pytest.mark.parametrize("replay", ["blank", "success"]) +async def test_opportunistic_tool_argument_replay_has_its_own_attempt_events(replay): + schema = {"type": "function", "function": {"name": "tool", "parameters": { + "type": "object", "required": ["x"], "properties": {"x": {"type": "integer"}}, + }}} + recovered = LLMResponse() if replay == "blank" else LLMResponse(tool_calls=[{ + "id": "good", "type": "function", "function": {"name": "tool", "arguments": '{"x":1}'}, + }]) + llm = bind_tools(Script([ + [StreamDelta(tool_call_deltas=[{"index": 0, "id": "original", "name": "tool", "arguments": "{}"}])], + recovered, + ]), [schema]) + events = [] + result = await invoke(llm, True, events) + assert [e["attempt_index"] for e in events if e["phase"] == "started"] == [1, 2] + by_index = {e["attempt_index"]: e for e in finished(events)} + assert len(finished(events)) == 2 + if replay == "blank": + assert by_index[1]["outcome"] == "accepted" + assert by_index[2]["outcome"] == "failed" + assert by_index[2]["reason"] == "empty_completion" + assert result.tool_calls[0]["id"] == "original" + assert result.response_metadata["stream_empty_args_fallback"] is False + else: + assert by_index[1]["outcome"] == "discarded" + assert by_index[2]["outcome"] == "accepted" + assert result.tool_calls[0]["id"] == "good" + assert result.response_metadata["stream_empty_args_fallback"] is True diff --git a/tests/test_empty_completion_guard.py b/tests/test_empty_completion_guard.py index 62ad6bf..5fda92b 100644 --- a/tests/test_empty_completion_guard.py +++ b/tests/test_empty_completion_guard.py @@ -113,7 +113,7 @@ def stream(self, *_a: Any, **_k: Any): def _cfg(**kw: Any) -> LoopConfig: return LoopConfig( - max_turns=4, max_llm_retries=1, + max_turns=4, max_llm_retries=1, retry_wait_fixed=0, loop_policy=LoopPolicy(no_tool_behavior="stop"), **kw, ) @@ -247,7 +247,7 @@ async def test_nonempty_raw_responses_are_not_resampled(response, thinking_forma model_profile=ModelProfile(model_id="m", provider="p", thinking_format=thinking_format), ) assert len(llm.requests) == 1 - assert result.stopped_by == "no_tool" + assert result.stopped_by == ("refusal" if response.finish_reason == "refusal" else "no_tool") assert result.turns_used == 1 @@ -265,7 +265,7 @@ async def test_blank_resamples_on_the_only_available_turn() -> None: llm = _SequenceLLM([LLMResponse(), LLMResponse(), LLMResponse(content="finished")]) result = await run_agent_loop( system_prompt="s", user_message="u", llm=llm, tools=[], - config=LoopConfig(max_turns=1, max_llm_retries=1), + config=LoopConfig(max_turns=1, max_llm_retries=1, retry_wait_fixed=0), ) assert result.final_content == "finished" assert result.turns_used == 1 diff --git a/tests/test_llm_proxy_stream_lifecycle.py b/tests/test_llm_proxy_stream_lifecycle.py index 0dec950..c945e3e 100644 --- a/tests/test_llm_proxy_stream_lifecycle.py +++ b/tests/test_llm_proxy_stream_lifecycle.py @@ -267,3 +267,47 @@ async def after_llm(self, ctx, response): assert sorted(steps) == ["research:llm:1", "research:llm:2"] finally: reset_current_execution_scope(token) + + +async def test_stream_after_llm_sees_terminal_metadata() -> None: + """Usage/cost and refusal-aware middleware need the streamed terminal + metadata, not just the concatenated text.""" + + seen: list[LLMResponse] = [] + + class _Capture(LLMMiddleware): + name = "capture" + + async def after_llm( + self, ctx: LLMCallContext, response: LLMResponse, + ) -> LLMResponse: + seen.append(response) + return response + + class _Refusing: + model = "test-model" + + async def stream(self, messages, **kwargs): + yield StreamDelta(transport_activity=True, usage={"total_tokens": 99}) + yield StreamDelta(content="I can", refusal="I can", model="m-1") + yield StreamDelta(content="not", refusal="not", stop_details={"type": "refusal"}, + finish_reason="stop", provider="vendor", stop_reason="refusal") + yield StreamDelta(usage={"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + usage_source="provider") + + proxy = LLMProxy(inner=_Refusing(), chain=_chain(_Capture()), role_id="r") + async for _ in proxy.stream([user_msg("hi")]): + pass + + [response] = seen + assert response.content == "I cannot" + assert response.finish_reason == "stop" + assert response.model == "m-1" + assert response.usage == {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5} + assert response.usage_source == "provider" + assert response.response_metadata == { + "provider_actually_used": "vendor", + "refusal": "I cannot", + "stop_details": {"type": "refusal"}, + "stop_reason": "refusal", + } diff --git a/tests/test_stream_middleware_accounting.py b/tests/test_stream_middleware_accounting.py new file mode 100644 index 0000000..ee7959f --- /dev/null +++ b/tests/test_stream_middleware_accounting.py @@ -0,0 +1,342 @@ +"""Real middleware consumers must receive authoritative streamed completion data.""" +from __future__ import annotations + +from contextlib import contextmanager +from copy import deepcopy + +import pytest + +from agent_core.components.middleware.llm.base import LLMCallContext, LLMMiddlewareChain +from agent_core.components.middleware.llm.loop_detection import LoopDetectionMiddleware +from agent_core.components.middleware.llm.proxy import LLMProxy +from agent_core.components.middleware.llm.token_accounting import TokenAccountingMiddleware +from agent_core.components.middleware.llm.tracing import LLMTracingMiddleware +from agent_core.components.middleware.rate_limit import RateLimitMiddleware +from agent_core.execution_context import ( + ExecutionScope, + reset_current_execution_scope, + set_current_execution_scope, +) +from agent_core.llm import LLMResponse, StreamDelta +from agent_core.messages import system_msg, user_msg +from agent_core.models.task_budget import BudgetState, TaskBudget +from agent_core.runtime.loop._call import call_llm +from agent_core.runtime.loop._streaming import _stream_llm_response + + +class Script: + model = "model" + def __init__(self, actions): + self.actions = list(actions) + self.requests = [] + async def _next(self, messages): + self.requests.append(deepcopy(messages)) + action = self.actions.pop(0) + if isinstance(action, Exception): + raise action + return action + async def chat(self, messages, **kwargs): + return await self._next(messages) + async def stream(self, messages, **kwargs): + action = await self._next(messages) + if isinstance(action, list): + for delta in action: + if isinstance(delta, Exception): + raise delta + yield delta + else: + yield StreamDelta(content=action.content, model=self.model, + usage=action.usage, usage_source=action.usage_source or action.response_metadata.get("usage_source", "")) + + +class Cost: + def __init__(self): + self.records = [] + def record(self, *args): + self.records.append(args) + return 0.0 + def get_summary(self, task_id): + return {"calls": len(self.records)} + + +class Aggregator: + def __init__(self): + self.records = [] + def record_llm_call(self, **kwargs): + self.records.append(kwargs) + + +class Events: + def __init__(self): + self.records = [] + async def append(self, **kwargs): + self.records.append(kwargs) + + +class Trace: + def __init__(self): + self.records = [] + async def log_llm_call(self, **kwargs): + self.records.append(kwargs) + + +@contextmanager +def scoped(budget=None, task="task"): + scope = ExecutionScope(task_id=task, phase_id="phase", role_id="role", + metadata={"budget_state": budget} if budget is not None else {}) + token = set_current_execution_scope(scope) + try: + yield + finally: + reset_current_execution_scope(token) + + +def chain(*middlewares): + result = LLMMiddlewareChain() + for mw in middlewares: + result.add(mw) + return result + + +async def consume(proxy, streaming, messages=None): + messages = messages or [user_msg("hi")] + if not streaming: + return await proxy.chat(messages) + return [delta async for delta in proxy.stream(messages)] + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("marker", ["source", "flag", "metadata"]) +async def test_estimates_never_enter_cost_budget_events_or_usage_aggregator(streaming, marker): + usage = {"prompt_tokens": 100, "completion_tokens": 50} + response = LLMResponse(content="answer", usage=usage) + if marker == "source": + response.usage_source = "estimated" + elif marker == "flag": + response.usage["estimated"] = True + else: + response.response_metadata["usage_source"] = "estimated" + cost, aggregator, events, trace = Cost(), Aggregator(), Events(), Trace() + budget = BudgetState(allocated=TaskBudget(max_tokens=20)) + accounting = TokenAccountingMiddleware(events, cost_sink=cost, usage_aggregator=aggregator) + proxy = LLMProxy(Script([response]), chain(accounting, LLMTracingMiddleware(trace))) + with scoped(budget): + await consume(proxy, streaming) + assert accounting.get_usage("task")["llm_calls"] == 0 + assert cost.records == aggregator.records == events.records == [] + assert budget.tokens_used == budget.llm_calls_used == 0 + assert not budget.exhausted + assert len(trace.records) == 1 + assert trace.records[0]["metadata"]["usage_source"] == "estimated" + assert trace.records[0]["metadata"]["usage"]["prompt_tokens"] == 100 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_empty_estimated_attempts_do_not_poison_real_cost_or_budget(streaming): + client = Script([LLMResponse(usage={"prompt_tokens": 100}, usage_source="estimated"), + LLMResponse(usage={"prompt_tokens": 200, "estimated": True}), + LLMResponse(content="answer", model="model", usage={"prompt_tokens": 10, "completion_tokens": 5}, usage_source="provider")]) + cost, aggregator, events = Cost(), Aggregator(), Events() + accounting = TokenAccountingMiddleware(events, cost_sink=cost, usage_aggregator=aggregator) + budget = BudgetState(allocated=TaskBudget(max_tokens=20)) + proxy = LLMProxy(client, chain(accounting)) + async def noop(*args, **kwargs): + pass + with scoped(budget): + result = await call_llm(proxy, [user_msg("hi")], 5, 1, 1, + on_delta=noop if streaming else None, retry_wait_fixed=0) + assert result.content == "answer" + assert len(client.requests) == 3 + assert accounting.get_usage("task") == {"input": 10, "output": 5, "total": 15, "llm_calls": 1} + assert cost.records == [("task", "model", 10, 5)] + assert len(aggregator.records) == len(events.records) == 1 + assert budget.tokens_used == 15 and budget.llm_calls_used == 1 + assert not budget.exhausted + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("read,write", [(10, 20), (0, 20), (10, 0)]) +async def test_canonical_cache_reads_and_writes_reach_real_consumers_once(streaming, read, write): + usage = {"prompt_tokens": 100, "completion_tokens": 5, "cache_read_tokens": read, + "cache_write_tokens": write, "cached_tokens": read + write, "cache_creation_tokens": write} + cost, aggregator, events = Cost(), Aggregator(), Events() + accounting = TokenAccountingMiddleware(events, cost_sink=cost, usage_aggregator=aggregator, scene="test") + budget = BudgetState() + proxy = LLMProxy(Script([LLMResponse(content="answer", model="model", usage=usage, usage_source="provider")]), chain(accounting)) + with scoped(budget): + await consume(proxy, streaming) + [record] = aggregator.records + assert record["cache_read_tokens"] == read and record["cache_write_tokens"] == write + assert record["scene"] == "test" + [event] = events.records + assert event["payload"]["this_call"]["cache_read"] == read + assert event["payload"]["this_call"]["cache_creation"] == write + assert cost.records == [("task", "model", 100, 5)] + assert budget.tokens_used == 105 and budget.llm_calls_used == 1 + + +@pytest.mark.parametrize("usage,expected", [ + ({"input_tokens": 10, "output_tokens": 5, "cache_read_input_tokens": 2, "cache_creation_input_tokens": 3}, (10, 5, 2, 3)), + ({"prompt_tokens": 10, "completion_tokens": 5, "cached_tokens": 2, "cache_creation_tokens": 3}, (10, 5, 2, 3)), + ({"prompt_tokens": 10, "completion_tokens": 5, "cache_read_tokens": 0, "cache_write_tokens": 0, + "cached_tokens": 999, "cache_creation_tokens": 999}, (10, 5, 0, 0)), +]) +def test_accounting_cache_aliases_and_canonical_zero_precedence(usage, expected): + assert TokenAccountingMiddleware()._extract_usage(LLMResponse(usage=usage)) == expected + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("prompt,usage,source,expected", [ + ("hi", {"prompt_tokens": 10, "completion_tokens": 90}, "provider", 900), + ("x" * 40, {"total_tokens": 0}, "provider", 1000), + ("x" * 40, {}, "", 990), + ("x" * 40, {"total_tokens": 500}, "estimated", 990), + ("x" * 40, {"total_tokens": 500, "estimated": True}, "", 990), + ("x" * 40, {"total_tokens": None}, "provider", 990), + ("x" * 40, {"total_tokens": -1}, "provider", 990), + ("x" * 8000, {"total_tokens": 100}, "provider", 900), +], ids=["zero-estimate", "reported-zero", "missing", "estimated-source", "estimated-flag", "null", "negative", "capped"]) +async def test_rate_bucket_corrects_real_usage_against_actual_reservation(streaming, prompt, usage, source, expected): + rate = RateLimitMiddleware(tokens_per_min=1000) + rate._bucket._refill = lambda: None + proxy = LLMProxy(Script([LLMResponse(content="answer", usage=usage, usage_source=source)]), chain(rate)) + await consume(proxy, streaming, [user_msg(prompt)]) + assert rate._bucket._token_tokens == expected + assert rate._bucket._request_tokens == 59 + + +async def test_rate_after_without_before_does_not_invent_a_reservation(): + rate = RateLimitMiddleware(tokens_per_min=1000) + await rate.after_llm(LLMCallContext(), LLMResponse(usage={"total_tokens": 100})) + assert rate._bucket._token_tokens == 1000 + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_real_loop_detector_gets_tool_calls_and_injects_the_next_hint(streaming): + call = {"id": "call", "type": "function", "function": {"name": "echo", "arguments": '{"x":1}'}} + actions = [LLMResponse(tool_calls=[deepcopy(call)]) for _ in range(3)] + if streaming: + actions = [[StreamDelta(tool_call_deltas=[{"index": 0, "id": "call", "name": "ec", "arguments": '{"x":'}]), + StreamDelta(tool_call_deltas=[{"index": 0, "name": "ho", "arguments": '1}'}])] for _ in range(3)] + client = Script(actions) + detection = LoopDetectionMiddleware(trigger_count=2) + proxy = LLMProxy(client, chain(detection), role_id="role") + with scoped(): + for _ in range(3): + await consume(proxy, streaming, [system_msg("system"), user_msg("hi")]) + assert "[Loop detected]" not in client.requests[1][0]["content"] + assert "[Loop detected]" in client.requests[2][0]["content"] + + +async def test_stream_tool_assembly_matches_runtime_and_drops_nameless_slots(): + seen = [] + from agent_core.components.middleware.llm.base import LLMMiddleware + class Capture(LLMMiddleware): + name = "capture" + async def after_llm(self, ctx, response): + seen.append(response) + return response + client = Script([[ + StreamDelta(tool_call_deltas=[{"index": 2, "id": "b", "name": "second", "arguments": "{"}]), + StreamDelta(tool_call_deltas=[{"index": 1, "id": "broken", "arguments": "{"}]), + StreamDelta(tool_call_deltas=[{"index": 0, "id": "a", "name": "fi", "arguments": '{"x":'}]), + StreamDelta(tool_call_deltas=[{"index": 2, "arguments": "}"}, {"index": 0, "name": "rst", "arguments": '1}'}]), + ]]) + async def noop(*args, **kwargs): + pass + response = await _stream_llm_response(LLMProxy(client, chain(Capture())), [], 5, noop) + assert seen[0].tool_calls == response.tool_calls + assert [c["function"]["name"] for c in response.tool_calls] == ["first", "second"] + assert [c["function"]["arguments"] for c in response.tool_calls] == ['{"x":1}', '{}'] + + +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("rewrite_prompt", [False, True]) +async def test_multiple_limiters_keep_independent_capped_reservations(streaming, rewrite_prompt): + from agent_core.components.middleware.llm.base import LLMMiddleware + class Rewrite(LLMMiddleware): + name = "rewrite" + async def before_llm(self, ctx, messages): + return [user_msg("x" * 40)] + first = RateLimitMiddleware(tokens_per_min=100) + second = RateLimitMiddleware(tokens_per_min=20) + first._bucket._refill = second._bucket._refill = lambda: None + nodes = (first, Rewrite(), second) if rewrite_prompt else (first, second) + proxy = LLMProxy(Script([LLMResponse(content="answer", usage={"total_tokens": 10}, usage_source="provider")]), chain(*nodes)) + await consume(proxy, streaming, [user_msg("x" * 400)]) + assert first._bucket._token_tokens == 90 + assert second._bucket._token_tokens == 10 + assert first._reserved_key != second._reserved_key + + +async def test_failed_stream_proposals_do_not_inject_false_loop_hints(): + delta = StreamDelta(tool_call_deltas=[{"index": 0, "id": "call", "name": "echo", "arguments": "{}"}]) + client = Script([[delta, RuntimeError("reset")], [delta, RuntimeError("reset")], [delta]]) + detection = LoopDetectionMiddleware(trigger_count=2) + proxy = LLMProxy(client, chain(detection), role_id="role") + with scoped(): + for _ in range(2): + with pytest.raises(RuntimeError, match="reset"): + await consume(proxy, True, [system_msg("system"), user_msg("hi")]) + await consume(proxy, True, [system_msg("system"), user_msg("hi")]) + assert all("[Loop detected]" not in request[0]["content"] for request in client.requests) + assert len(detection._histories[("task", "role", "phase")]) == 1 + + +async def test_consumer_closed_stream_proposal_does_not_enter_loop_history(): + client = Script([[StreamDelta(tool_call_deltas=[{"index": 0, "id": "call", "name": "echo", "arguments": "{}"}])]]) + detection = LoopDetectionMiddleware(trigger_count=2) + proxy = LLMProxy(client, chain(detection), role_id="role") + with scoped(): + stream = proxy.stream([system_msg("system"), user_msg("hi")]) + await anext(stream) + await stream.aclose() + assert not detection._histories and not detection._pending_hints + + +@pytest.mark.parametrize("streaming", [False, True]) +async def test_reported_cache_only_call_is_not_lost_to_zero_base_tokens(streaming): + usage = {"prompt_tokens": 0, "completion_tokens": 0, "cache_read_tokens": 10, "cache_write_tokens": 20, + "cached_tokens": 30, "cache_creation_tokens": 20} + aggregator, events = Aggregator(), Events() + accounting = TokenAccountingMiddleware(events, usage_aggregator=aggregator) + proxy = LLMProxy(Script([LLMResponse(content="", usage=usage, usage_source="provider")]), chain(accounting)) + with scoped(): + await consume(proxy, streaming) + assert accounting.get_usage("task")["llm_calls"] == 1 + [record] = aggregator.records + assert record["cache_read_tokens"] == 10 and record["cache_write_tokens"] == 20 + assert len(events.records) == 1 + + +async def test_reported_usage_from_failed_stream_is_billed_without_loop_history(): + cost, aggregator = Cost(), Aggregator() + budget = BudgetState() + accounting = TokenAccountingMiddleware(cost_sink=cost, usage_aggregator=aggregator) + detection = LoopDetectionMiddleware(trigger_count=2) + client = Script([[StreamDelta(model="model", usage_source="provider", usage={"prompt_tokens": 10, "completion_tokens": 5}, + tool_call_deltas=[{"index": 0, "id": "call", "name": "echo", "arguments": "{}"}]), RuntimeError("reset")]]) + proxy = LLMProxy(client, chain(accounting, detection), role_id="role") + with scoped(budget), pytest.raises(RuntimeError, match="reset"): + await consume(proxy, True) + assert cost.records == [("task", "model", 10, 5)] + assert budget.tokens_used == 15 and budget.llm_calls_used == 1 + assert len(aggregator.records) == 1 + assert not detection._histories + + +async def test_legacy_rate_context_correction_respects_the_reservation_cap(): + rate = RateLimitMiddleware(tokens_per_min=1000) + rate._bucket._refill = lambda: None + await rate._bucket.acquire(5000) + ctx = LLMCallContext(metadata={"_rate_limit_estimated_tokens": 5000}) + await rate.after_llm(ctx, LLMResponse(usage={"total_tokens": 10}, usage_source="provider")) + assert rate._bucket._token_tokens == 990 + + +@pytest.mark.parametrize("estimate,expected", [(None, 1000), ("bad", 1000), (0, 900)]) +async def test_legacy_rate_context_distinguishes_unknown_from_zero(estimate, expected): + rate = RateLimitMiddleware(tokens_per_min=1000) + ctx = LLMCallContext(metadata={"_rate_limit_estimated_tokens": estimate}) + await rate.after_llm(ctx, LLMResponse(usage={"total_tokens": 100}, usage_source="provider")) + assert rate._bucket._token_tokens == expected