Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
149 changes: 149 additions & 0 deletions agent_core/completion.py
Original file line number Diff line number Diff line change
@@ -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,
},
))
2 changes: 2 additions & 0 deletions agent_core/components/middleware/llm/loop_detection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
87 changes: 83 additions & 4 deletions agent_core/components/middleware/llm/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand All @@ -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,
Expand Down Expand Up @@ -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).
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -235,6 +313,7 @@ async def stream(
len(full_content),
)
break
stream_completed = True
break
except Exception as e:
stream_error = e
Expand Down Expand Up @@ -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
Expand All @@ -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,
)

Expand Down
34 changes: 10 additions & 24 deletions agent_core/components/middleware/llm/token_accounting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -88,39 +89,24 @@ 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
) -> LLMResponse:
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"
Expand Down
5 changes: 5 additions & 0 deletions agent_core/components/middleware/llm/tracing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading