From cbec06c18342d2ea05e8be3d0f8eb565c9ee3c1f Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Wed, 30 Sep 2026 18:41:55 -0700 Subject: [PATCH] fix(agents): preserve token pricing and legacy request handling --- .../proxy/agent_endpoints/a2a_endpoints.py | 4 +- litellm/proxy/agent_endpoints/a2a_routing.py | 10 +++-- litellm/proxy/auth/user_api_key_auth.py | 2 +- .../agent_endpoints/test_a2a_endpoints.py | 40 +++++++++++++++++++ .../test_llm_pass_through_endpoints.py | 2 +- 5 files changed, 51 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index d66d9835081..23d11a5a9c5 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -738,7 +738,9 @@ async def invoke_agent_a2a( str, object ] = { # mutable-ok: A2A SDK and completion bridge accept provider parameters as a dict **(agent.litellm_params or MappingProxyType({})), - "cost_per_query": user_api_key_dict.agent_invocation_cost, + "cost_per_query": user_api_key_dict.agent_invocation_cost + if (agent.litellm_params or MappingProxyType({})).get("cost_per_query") is not None + else None, } custom_llm_provider: Final = litellm_params.get("custom_llm_provider") diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 9d91fbfff7d..96c0b0ceb5a 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -82,9 +82,11 @@ async def route_a2a_agent_request( raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False) # Inject API base and route to litellm - data.pop("litellm_params", None) - data["api_base"] = agent.agent_card_params["url"] - verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"]) + api_base: Final = agent.agent_card_params["url"] + verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, api_base) invocation_fee: Final = user_api_key_dict.agent_invocation_cost if user_api_key_dict is not None else None - return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, "cost_per_query": invocation_fee})) + provider_data: Final = MappingProxyType({key: value for key, value in data.items() if key != "litellm_params"}) + return getattr(litellm, f"{route_type}")( + **MappingProxyType({**provider_data, "api_base": api_base, "cost_per_query": invocation_fee}) + ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 3d13309a682..f3cc82eb895 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3264,7 +3264,7 @@ async def _authorize_authenticated_request( general_settings, user_model, request.path_params.get("model") or request.path_params.get("model_name"), - request.query_params.get("model"), + _safe_get_request_query_params(request).get("model"), model_group_alias=router_settings.get("model_group_alias") if isinstance(router_settings, Mapping) else None, diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 8f905b2fc96..1e04a6df3a5 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -27,6 +27,7 @@ class CapturedAgentCall: agent_extra_headers: dict[str, str] | None cost_per_query: object api_base: object + pricing: Mapping[str, object] @pytest.mark.asyncio @@ -500,6 +501,7 @@ async def _invoke_message_method( agent_extra_headers=kwargs.get("agent_extra_headers"), cost_per_query=kwargs["litellm_params"].get("cost_per_query"), api_base=kwargs["api_base"], + pricing=kwargs["litellm_params"], ) @@ -2753,3 +2755,41 @@ async def test_native_dispatch_returns_not_found_for_a_removed_agent(monkeypatch ) assert response.status_code == 404 assert json.loads(response.body)["error"]["message"] == "Agent 'removed' not found" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ("message/send", "message/stream")) +@pytest.mark.parametrize("fixed_fee", (None, 0.0, 0.25)) +async def test_unbudgeted_managed_agent_keeps_token_pricing_without_a_fixed_fee( + monkeypatch: pytest.MonkeyPatch, method: str, fixed_fee: float | None, +) -> None: + from litellm import Usage + from litellm.a2a_protocol.cost_calculator import A2ACostCalculator + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + + policy: Final = AgentResponse( + agent_id="test-agent", agent_name="test-agent", identity_managed=True, + agent_card_params={"url": "https://agent.test/"}, + litellm_params={"input_cost_per_token": 0.02, "output_cost_per_token": 0.03, + **({"cost_per_query": fixed_fee} if fixed_fee is not None else {})}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(policy) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + with ExitStack() as stack: + for context in _base_patches(policy): + stack.enter_context(context) + await prepare_agent_invocation(auth, policy.agent_id, AgentIdentityStore.from_client(database)) + captured: Final = await _invoke_message_method( + method, _make_request_mock(method, _HELLO_MESSAGE_PARAMS), auth, agent=policy, + ) + logging: Final = MagicMock(model_call_details={ + "litellm_params": captured.pricing, + "usage": Usage(prompt_tokens=10, completion_tokens=2, total_tokens=12), + }) + assert A2ACostCalculator.calculate_a2a_cost(logging) == pytest.approx(0.26 if fixed_fee is None else fixed_fee) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 4577d578263..72364904f93 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -2270,7 +2270,7 @@ class TestGigachatProxyRoute: mock_request.headers = {"content-type": "application/json"} mock_request.query_params = {} mock_fastapi_response = MagicMock(spec=Response) - mock_user_api_key_dict = MagicMock() + mock_user_api_key_dict = UserAPIKeyAuth() mock_llm_router.allm_passthrough_route = AsyncMock( return_value=httpx.Response(200, json={"response": "success"}) )