diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index aa41e63b40b..8d733971c70 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -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 diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py index 8232d7cf2d8..2a20979501a 100644 --- a/litellm/a2a_protocol/streaming_iterator.py +++ b/litellm/a2a_protocol/streaming_iterator.py @@ -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 diff --git a/tests/unit/a2a_protocol/test_main.py b/tests/unit/a2a_protocol/test_main.py index 4ba0ef8fa04..1720b521bd7 100644 --- a/tests/unit/a2a_protocol/test_main.py +++ b/tests/unit/a2a_protocol/test_main.py @@ -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)),)