From b33fda1dd3022fee411a50fc3dbbed731336cf45 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Wed, 30 Sep 2026 16:54:32 -0700 Subject: [PATCH] fix(agents): enforce and settle chat adapter invocation fees --- litellm/cost_calculator.py | 9 ++- litellm/main.py | 1 + litellm/proxy/agent_endpoints/a2a_routing.py | 12 +++- litellm/proxy/auth/user_api_key_auth.py | 7 +- .../proxy/auth/test_user_api_key_auth.py | 3 + .../unit/a2a_protocol/test_cost_calculator.py | 68 +++++++++++++++++++ 6 files changed, 96 insertions(+), 4 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 238b7cc3fdd..c022b34282f 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -1529,7 +1529,14 @@ def completion_cost( completion_tokens = token_counter(model=model, text=completion) # Handle A2A calls before model check - A2A doesn't require a model - if call_type in _A2A_CALL_TYPES: + if call_type in _A2A_CALL_TYPES or ( + custom_llm_provider == "a2a" + and litellm_logging_obj is not None + and (litellm_logging_obj.model_call_details.get("litellm_params") or MappingProxyType({})).get( + "cost_per_query" + ) + is not None + ): from litellm.a2a_protocol.cost_calculator import A2ACostCalculator return A2ACostCalculator.calculate_a2a_cost(litellm_logging_obj=litellm_logging_obj) diff --git a/litellm/main.py b/litellm/main.py index 8c9d7f2513d..570366e0792 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5662,6 +5662,7 @@ def completion( preset_cache_key=preset_cache_key, no_log=no_log, cost_per_second=cost_per_second, + cost_per_query=kwargs.get("cost_per_query"), input_cost_per_second=input_cost_per_second, input_cost_per_token=input_cost_per_token, output_cost_per_second=output_cost_per_second, diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index c57315ebc21..1b6f210cf8e 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -5,6 +5,7 @@ Handles routing for A2A agents (models with "a2a/" prefix). Looks up agents in the registry and injects their API base URL. """ +from types import MappingProxyType from typing import Any, Final from fastapi import HTTPException @@ -78,4 +79,13 @@ async def route_a2a_agent_request( data["api_base"] = agent.agent_card_params["url"] verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"]) - return getattr(litellm, f"{route_type}")(**data) + invocation_pricing: Final = ( + MappingProxyType({"cost_per_query": user_api_key_dict.agent_invocation_cost}) + if user_api_key_dict is not None + and user_api_key_dict.agent_invocation_cost is not None + and user_api_key_dict.invoked_agent_policy is not None + and (user_api_key_dict.invoked_agent_policy.litellm_params or MappingProxyType({})).get("cost_per_query") + is not None + else MappingProxyType({}) + ) + return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, **invocation_pricing})) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 278a48911d4..a684b5ebcf3 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3261,8 +3261,11 @@ async def _authorize_authenticated_request( target_name, store, billable=request.method == "POST" - and request_data.get("method") - in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"), + and ( + not RouteChecks.check_route_access(route, ("/a2a/{agent_id}", "/v1/a2a/{agent_id}")) + or request_data.get("method") + in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage") + ), ) await _run_centralized_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index a60d64bd562..d69b3e72183 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -9701,6 +9701,9 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/stream", "params": {}}, True), ("POST", "/a2a/agent", {"method": "message/send", "model": "free-model", "params": {}}, True), ("POST", "/a2a/agent", {"method": "message/stream", "model": "free-model", "params": {}}, True), + ("POST", "/v1/chat/completions", {"model": "a2a/agent", "method": "tasks/get"}, True), + ("POST", "/chat/completions", {"model": "a2a/agent", "method": "tasks/cancel", "stream": True}, True), + ("POST", "/v1/a2a/agent/message/send", {"method": "tasks/get", "params": {}}, True), ], ) async def test_human_agent_discovery_does_not_reserve_target_budget_but_send_and_stream_do( diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py index 8d8ec815f3a..4f2a93fa883 100644 --- a/tests/unit/a2a_protocol/test_cost_calculator.py +++ b/tests/unit/a2a_protocol/test_cost_calculator.py @@ -117,6 +117,7 @@ class CostLogger(CustomLogger): def __init__(self): self.response_cost: Optional[float] = None + self.logged: asyncio.Event = asyncio.Event() super().__init__() async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -125,6 +126,7 @@ class CostLogger(CustomLogger): self.response_cost = ( slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None) ) + self.logged.set() @pytest.mark.asyncio @@ -449,3 +451,69 @@ async def test_asend_message_streaming_triggers_callbacks(): assert callback_logger.agent_id == test_agent_id, ( f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'" ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", (False, True)) +@pytest.mark.parametrize("claimed_fee", (None, 0.0, 99.0)) +async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing( + monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None +) -> None: + import json + from typing import Final + + import httpx + + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request + from litellm.types.agents import AgentResponse + + await _reset_callbacks_and_settle_pending_logs() + logger: Final = CostLogger() + monkeypatch.setattr(litellm, "callbacks", [logger]) + target: Final = AgentResponse( + agent_id="fee-target", agent_name="fee-target", + agent_card_params={"url": "https://agent.test/", "capabilities": {"streaming": True}}, + litellm_params={"cost_per_query": 0.25}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + auth.invoked_agent_policy = target + auth.invoked_agent_id = target.agent_id + auth.agent_invocation_cost = 0.25 + + def reply(request: httpx.Request) -> httpx.Response: + body: Final = json.loads(request.content) + assert body["method"] == ("message/stream" if stream else "message/send") + result: Final = { + "jsonrpc": "2.0", "id": body["id"], + "result": {"kind": "message", "messageId": "reply", "role": "agent", + "parts": [{"kind": "text", "text": "Paid reply"}]}, + } + if stream: + return httpx.Response(200, text=f"data: {json.dumps(result)}\n\n", headers={"content-type": "text/event-stream"}) + return httpx.Response(200, json=result) + + client: Final = AsyncHTTPHandler(transport=httpx.MockTransport(reply)) + try: + pending: Final = await route_a2a_agent_request( + data={"model": "a2a/fee-target", "messages": [{"role": "user", "content": "Hello"}], + "stream": stream, "client": client, + **({"cost_per_query": claimed_fee} if claimed_fee is not None else {})}, + route_type="acompletion", user_api_key_dict=auth, + ) + response: Final = await pending + if stream: + chunks: Final = tuple([chunk async for chunk in response]) + assert any(chunk.choices[0].delta.content == "Paid reply" for chunk in chunks) + else: + assert response.choices[0].message.content == "Paid reply" + await asyncio.wait_for(logger.logged.wait(), timeout=10.0) + await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0) + assert logger.response_cost == pytest.approx(0.25) + finally: + await client.close()