fix(agents): account for completion bridge invocation fees

This commit is contained in:
Joshua Valluru 2026-09-28 19:28:42 -07:00
parent 69b8ada9ac
commit 92e13728f5
3 changed files with 81 additions and 13 deletions

View file

@ -471,13 +471,20 @@ async def asend_message(
if custom_llm_provider:
if request is None:
raise ValueError("request is required for completion bridge")
return await _send_message_via_completion_bridge(
bridge_response: Final = await _send_message_via_completion_bridge(
request=request,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
)
bridge_prompt_tokens, bridge_completion_tokens, _ = await asyncify(
A2ARequestUtils.calculate_usage_from_request_response
)(request=request, response_dict=bridge_response.model_dump(mode="json", exclude_none=True))
_set_usage_on_logging_obj(kwargs, bridge_prompt_tokens, bridge_completion_tokens)
_set_litellm_params_on_logging_obj(kwargs, litellm_params)
_set_agent_id_on_logging_obj(kwargs, agent_id)
return bridge_response
# Standard A2A client flow
if request is None:
@ -692,12 +699,35 @@ async def asend_message_streaming(
request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
)
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
bridge_name: Final = str(litellm_params.get("model") or agent_id or "agent")
existing_logging: Final = kwargs.get("litellm_logging_obj")
bridge_logging: Final = (
existing_logging
if isinstance(existing_logging, Logging)
else _build_streaming_logging_obj(
request=request,
agent_name=bridge_name,
agent_id=agent_id,
litellm_params=litellm_params,
metadata=metadata,
proxy_server_request=proxy_server_request,
)
)
bridge_context: Final = {"litellm_logging_obj": bridge_logging}
_set_litellm_params_on_logging_obj(bridge_context, litellm_params)
_set_agent_id_on_logging_obj(bridge_context, agent_id)
bridge_stream: Final = A2ACompletionBridgeHandler.handle_streaming(
request_id=str(request.id),
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
)
async for chunk in A2AStreamingIterator(
stream=bridge_stream,
request=request,
logging_obj=bridge_logging,
agent_name=bridge_name,
):
yield chunk
return

View file

@ -5,7 +5,7 @@ A2A Streaming Iterator with token tracking and logging support.
import asyncio
from collections.abc import AsyncIterator
from datetime import datetime
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, Generic, TypeVar
import litellm
from litellm._logging import verbose_logger
@ -18,7 +18,10 @@ if TYPE_CHECKING:
from a2a.compat.v0_3.types import SendStreamingMessageRequest, SendStreamingMessageResponse
class A2AStreamingIterator:
_StreamChunk = TypeVar("_StreamChunk", bound="SendStreamingMessageResponse | dict[str, object]")
class A2AStreamingIterator(Generic[_StreamChunk]):
"""
Async iterator for A2A streaming responses with token tracking.
@ -27,7 +30,7 @@ class A2AStreamingIterator:
def __init__(
self,
stream: AsyncIterator["SendStreamingMessageResponse"],
stream: AsyncIterator[_StreamChunk],
request: "SendStreamingMessageRequest",
logging_obj: LiteLLMLoggingObj,
agent_name: str = "unknown",
@ -39,14 +42,14 @@ class A2AStreamingIterator:
self.start_time = datetime.now()
# Collect chunks for token counting
self.chunks: list[SendStreamingMessageResponse] = []
self.chunks: list[_StreamChunk] = []
self.collected_text_parts: list[str] = []
self.final_chunk: SendStreamingMessageResponse | None = None
self.final_chunk: _StreamChunk | None = None
def __aiter__(self):
return self
async def __anext__(self) -> "SendStreamingMessageResponse":
async def __anext__(self) -> _StreamChunk:
try:
chunk: Final = await self.stream.__anext__()
@ -69,20 +72,20 @@ class A2AStreamingIterator:
await self._handle_stream_complete()
raise
def _collect_text_from_chunk(self, chunk: "SendStreamingMessageResponse") -> None:
def _collect_text_from_chunk(self, chunk: _StreamChunk) -> None:
"""Extract text from a streaming chunk and add to collected parts."""
try:
chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True)
text: Final = A2ARequestUtils.extract_text_from_response(chunk_dict)
if text:
self.collected_text_parts.append(text)
except Exception:
verbose_logger.debug("Failed to extract text from A2A streaming chunk")
def _is_completed_chunk(self, chunk: "SendStreamingMessageResponse") -> bool:
def _is_completed_chunk(self, chunk: _StreamChunk) -> bool:
"""Check if chunk indicates stream completion."""
try:
chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True)
result: Final = chunk_dict.get("result", {})
if isinstance(result, dict):
status: Final = result.get("status", {})
@ -160,7 +163,11 @@ class A2AStreamingIterator:
# Add final chunk result if available
if self.final_chunk:
try:
chunk_dict: Final = self.final_chunk.model_dump(mode="json", exclude_none=True)
chunk_dict: Final = (
self.final_chunk
if isinstance(self.final_chunk, dict)
else self.final_chunk.model_dump(mode="json", exclude_none=True)
)
result["result"] = chunk_dict.get("result", {})
except Exception:
pass

View file

@ -539,3 +539,34 @@ def test_streaming_logging_obj_keeps_agent_credentials_out_of_logging_params():
assert logging_obj.litellm_params == expected
assert logging_obj.optional_params == expected
assert logging_obj.model_call_details["litellm_params"] == expected
class _AgentFeeRecorder(CustomLogger):
def __init__(self):
super().__init__()
self.logged = asyncio.Event()
self.fees = ()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if kwargs.get("call_type") in ("asend_message", "asend_message_streaming"):
self.fees = (*self.fees, (kwargs.get("agent_id"), kwargs["standard_logging_object"]["response_cost"]))
self.logged.set()
@pytest.mark.asyncio
@pytest.mark.parametrize("streaming", [False, True])
async def test_completion_bridge_records_one_agent_fee(streaming, monkeypatch):
from litellm.a2a_protocol.main import asend_message_streaming
recorder = _AgentFeeRecorder()
monkeypatch.setattr(litellm, "callbacks", [recorder])
params = {"custom_llm_provider": "openai", "model": "gpt-4o-mini", "mock_response": "hello back", "cost_per_query": 0.01}
if streaming:
request = SendStreamingMessageRequest(id="bridge-stream", params=_request().params)
chunks = [chunk async for chunk in asend_message_streaming(request=request, litellm_params=params, agent_id="budgeted-agent")]
assert chunks[-1]["result"]["final"] is True
else:
response = await asend_message(request=_request(), litellm_params=params, agent_id="budgeted-agent")
assert response.id == "r1"
await asyncio.wait_for(recorder.logged.wait(), timeout=2)
assert recorder.fees == (("budgeted-agent", pytest.approx(0.01)),)