diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 37bf7c34f02..4c23ecfed54 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -129,6 +129,33 @@ def _set_agent_id_on_logging_obj( litellm_logging_obj.model_call_details["agent_id"] = agent_id +_A2A_COST_PARAM_KEYS = ("cost_per_query", "input_cost_per_token", "output_cost_per_token") + + +def _set_litellm_params_on_logging_obj( + kwargs: dict[str, Any], + litellm_params: dict[str, Any], +) -> None: + """ + Merge the agent's pricing params into model_call_details["litellm_params"] + so A2ACostCalculator can read them. + + The non-streaming path reuses the proxy-built logging object, whose + litellm_params already carries metadata / proxy_server_request / user-key + context, so merge the pricing keys in rather than replacing the dict. + """ + logging_obj = kwargs.get("litellm_logging_obj") + if logging_obj is None: + return + + cost_params = {key: litellm_params[key] for key in _A2A_COST_PARAM_KEYS if litellm_params.get(key) is not None} + if not cost_params: + return + + existing = logging_obj.model_call_details.get("litellm_params") or {} + logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params} + + def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str: """ Extract agent info and set model/custom_llm_provider for cost tracking. @@ -477,6 +504,9 @@ async def asend_message( completion_tokens=completion_tokens, ) + # Merge agent pricing params into the logging obj so cost is calculated + _set_litellm_params_on_logging_obj(kwargs=kwargs, litellm_params=litellm_params) + # Set agent_id on logging obj for SpendLogs tracking _set_agent_id_on_logging_obj(kwargs=kwargs, agent_id=agent_id) diff --git a/litellm/a2a_protocol/utils.py b/litellm/a2a_protocol/utils.py index 0dbd1eefc63..ce5a168c3ac 100644 --- a/litellm/a2a_protocol/utils.py +++ b/litellm/a2a_protocol/utils.py @@ -121,8 +121,13 @@ class A2ARequestUtils: Returns: Tuple of (prompt_tokens, completion_tokens, total_tokens) """ - # Count input tokens + # Count input tokens. Dump the message to a dict first so extraction hits + # the dict branch — request-side parts are a2a-sdk Part RootModels whose + # kind/text live on part.root, which the object branch cannot read. This + # mirrors how the response side already works (it operates on model_dump). input_message = A2ARequestUtils.get_input_message_from_request(request) + if input_message is not None and hasattr(input_message, "model_dump"): + input_message = input_message.model_dump(mode="json") input_text = A2ARequestUtils.extract_text_from_message(input_message) prompt_tokens = A2ARequestUtils.count_tokens(input_text) diff --git a/tests/test_litellm/a2a_protocol/__init__.py b/tests/test_litellm/a2a_protocol/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/a2a_protocol/test_cost_calculator.py b/tests/test_litellm/a2a_protocol/test_cost_calculator.py index bf03562f8ae..a29f012170f 100644 --- a/tests/test_litellm/a2a_protocol/test_cost_calculator.py +++ b/tests/test_litellm/a2a_protocol/test_cost_calculator.py @@ -103,9 +103,7 @@ async def _mock_stream_messages(a2a_client: Any, request: Any) -> AsyncIterator[ kind="message", ) for _ in range(2): - yield SendStreamingMessageResponse( - root=SendStreamingMessageSuccessResponse(id=request.id, result=msg) - ) + yield SendStreamingMessageResponse(root=SendStreamingMessageSuccessResponse(id=request.id, result=msg)) class CostLogger(CustomLogger): @@ -119,9 +117,7 @@ class CostLogger(CustomLogger): slp = kwargs.get("standard_logging_object") if slp: self.response_cost = ( - slp.get("response_cost") - if isinstance(slp, dict) - else getattr(slp, "response_cost", None) + slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None) ) @@ -160,6 +156,43 @@ async def test_asend_message_uses_cost_per_query(): assert cost_logger.response_cost == 0.05 +@pytest.mark.asyncio +async def test_asend_message_uses_cost_per_query_from_litellm_params_dict(): + """ + Proxy passes agent pricing as the litellm_params dict param (not top-level + kwargs). Regression for cost_per_query landing at $0 on the native path. + """ + from litellm.a2a_protocol import asend_message + + litellm.logging_callback_manager._reset_all_callbacks() + cost_logger = CostLogger() + litellm.callbacks = [cost_logger] + + mock_client = MagicMock() + mock_client._litellm_agent_card = MagicMock() + mock_client._litellm_agent_card.name = "test-agent" + + mock_request = _make_send_message_request("test-123") + + with patch( + "litellm.a2a_protocol.main._execute_a2a_send_with_retry", + new=_mock_execute_a2a_send, + ): + await asend_message( + a2a_client=mock_client, + request=mock_request, + litellm_params={ + "cost_per_query": 0.5, + "input_cost_per_token": 0.099999, + "output_cost_per_token": 0.1, + }, + ) + + await asyncio.sleep(0.1) + + assert cost_logger.response_cost == 0.5 + + class TokenAndCostLogger(CustomLogger): """Custom logger to capture both token counts and cost.""" @@ -173,19 +206,13 @@ class TokenAndCostLogger(CustomLogger): slp = kwargs.get("standard_logging_object") if slp: self.response_cost = ( - slp.get("response_cost") - if isinstance(slp, dict) - else getattr(slp, "response_cost", None) + slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None) ) self.prompt_tokens = ( - slp.get("prompt_tokens") - if isinstance(slp, dict) - else getattr(slp, "prompt_tokens", None) + slp.get("prompt_tokens") if isinstance(slp, dict) else getattr(slp, "prompt_tokens", None) ) self.completion_tokens = ( - slp.get("completion_tokens") - if isinstance(slp, dict) - else getattr(slp, "completion_tokens", None) + slp.get("completion_tokens") if isinstance(slp, dict) else getattr(slp, "completion_tokens", None) ) @@ -207,9 +234,7 @@ async def test_asend_message_uses_input_output_cost_per_token(): mock_client._litellm_agent_card = MagicMock() mock_client._litellm_agent_card.name = "test-agent" - mock_request = _make_send_message_request( - "test-123", user_text="Hello, what can you do?" - ) + mock_request = _make_send_message_request("test-123", user_text="Hello, what can you do?") # Define specific cost per token values input_cost_per_token = 0.00001 # $0.01 per 1000 tokens @@ -246,15 +271,11 @@ async def test_asend_message_uses_input_output_cost_per_token(): assert response_cost is not None, "response_cost should be captured" # Calculate expected cost - expected_cost = (prompt_tokens * input_cost_per_token) + ( - completion_tokens * output_cost_per_token - ) + expected_cost = (prompt_tokens * input_cost_per_token) + (completion_tokens * output_cost_per_token) print(f"expected_cost: {expected_cost}") # Verify exact cost calculation - assert ( - response_cost == expected_cost - ), f"response_cost {response_cost} should equal expected {expected_cost}" + assert response_cost == expected_cost, f"response_cost {response_cost} should equal expected {expected_cost}" class AgentIdLogger(CustomLogger): @@ -305,9 +326,9 @@ async def test_asend_message_passes_agent_id_to_callback(): await asyncio.sleep(0.1) # Verify agent_id was passed to callback - assert ( - agent_id_logger.agent_id == test_agent_id - ), f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'" + assert agent_id_logger.agent_id == test_agent_id, ( + f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'" + ) class MetadataLogger(CustomLogger): @@ -418,9 +439,7 @@ async def test_asend_message_streaming_triggers_callbacks(): assert len(chunks) == 2 # Verify callbacks WERE triggered after stream completed - assert ( - callback_logger.kwargs is not None - ), "Streaming should trigger callbacks after completion" - assert ( - callback_logger.agent_id == test_agent_id - ), f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'" + assert callback_logger.kwargs is not None, "Streaming should trigger callbacks after completion" + assert callback_logger.agent_id == test_agent_id, ( + f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'" + ) diff --git a/tests/test_litellm/a2a_protocol/test_utils.py b/tests/test_litellm/a2a_protocol/test_utils.py new file mode 100644 index 00000000000..8a219211c62 --- /dev/null +++ b/tests/test_litellm/a2a_protocol/test_utils.py @@ -0,0 +1,41 @@ +"""Tests for litellm/a2a_protocol/utils.py token/usage extraction.""" + +import pytest + +pytest.importorskip("a2a.compat.v0_3.types") + +from a2a.compat.v0_3.types import MessageSendParams, SendMessageRequest + +from litellm.a2a_protocol.utils import A2ARequestUtils + + +def _request(user_text: str) -> SendMessageRequest: + return SendMessageRequest( + id="r1", + params=MessageSendParams( + message={ + "messageId": "m1", + "role": "user", + "parts": [{"kind": "text", "text": user_text}], + } + ), + ) + + +def test_calculate_usage_counts_input_tokens_from_request_object(): + """Regression: request-side Part is a RootModel; input tokens must be counted.""" + request = _request("count these input tokens please") + response_dict = { + "result": { + "kind": "message", + "parts": [{"kind": "text", "text": "ok"}], + } + } + + prompt_tokens, completion_tokens, total_tokens = A2ARequestUtils.calculate_usage_from_request_response( + request=request, response_dict=response_dict + ) + + assert prompt_tokens > 0 + assert completion_tokens > 0 + assert total_tokens == prompt_tokens + completion_tokens