From 95b94520d944afdfe0ab78f164d0ae2ce8f81a5b Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Wed, 30 Sep 2026 17:16:36 -0700 Subject: [PATCH] fix(agents): reject client supplied invocation pricing --- litellm/proxy/agent_endpoints/a2a_routing.py | 19 +++++++-------- .../auth/managed_authorization.py | 6 ++--- litellm/proxy/litellm_pre_call_utils.py | 4 +++- .../proxy/test_pricing_field_strip.py | 5 +++- .../unit/a2a_protocol/test_cost_calculator.py | 23 ++++++++++++------- 5 files changed, 35 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 1b6f210cf8e..29d31251f94 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -13,6 +13,7 @@ from fastapi import HTTPException import litellm from litellm._logging import verbose_proxy_logger from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.agent_endpoints.auth.managed_authorization import AGENT_INVOCATION_COST async def route_a2a_agent_request( @@ -79,13 +80,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"]) - 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({}) + pricing_policy: Final = ( + user_api_key_dict.invoked_agent_policy + if user_api_key_dict is not None and user_api_key_dict.invoked_agent_policy is not None + else agent ) - return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, **invocation_pricing})) + configured_fee: Final = (pricing_policy.litellm_params or MappingProxyType({})).get("cost_per_query") + invocation_fee: Final = ( + AGENT_INVOCATION_COST.validate_python(configured_fee) if configured_fee is not None else None + ) + return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, "cost_per_query": invocation_fee})) diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index e254430def3..11c8c790ad2 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -222,7 +222,7 @@ async def check_agent_budget(auth: UserAPIKeyAuth) -> None: raise litellm.BudgetExceededError(current_cost=spend, max_budget=budget, message="Agent budget exceeded") -_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)]) +AGENT_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)]) def invocation_target(route: str, body: Mapping[str, object]) -> str | None: @@ -280,13 +280,13 @@ async def prepare_agent_invocation( and billing_policy.litellm_budget_table.max_budget is not None ) try: - fee: Final = _INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0) + fee: Final = AGENT_INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0) unbounded_token_price: Final = ( billable and bounded and fixed_fee is None and any( - _INVOCATION_COST.validate_python(pricing[field]) > 0 + AGENT_INVOCATION_COST.validate_python(pricing[field]) > 0 for field in ("input_cost_per_token", "output_cost_per_token") if pricing.get(field) is not None ) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 27d22920591..b4272d5e170 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -362,7 +362,9 @@ _ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY: Final = "allow_client_mess # not to user-supplied request bodies, so the proxy strips them before they # reach the call path. Built from the Pydantic model so newly-added pricing # fields are covered automatically. -_CLIENT_PRICING_CONTROL_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys()) +_CLIENT_PRICING_CONTROL_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys()) | frozenset( + {"cost_per_query"} +) # ``model_info`` carries the same pricing fields when read by # ``use_custom_pricing_for_model``; strip from metadata for the same reason. # ``standard_logging_guardrail_information`` is proxy-written telemetry summed diff --git a/tests/test_litellm/proxy/test_pricing_field_strip.py b/tests/test_litellm/proxy/test_pricing_field_strip.py index a0e25e91f37..87cf39b0d49 100644 --- a/tests/test_litellm/proxy/test_pricing_field_strip.py +++ b/tests/test_litellm/proxy/test_pricing_field_strip.py @@ -60,7 +60,7 @@ class TestStripClientPricingOverrides: # set drifting apart if someone replaces the auto-derivation later. assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset( CustomPricingLiteLLMParams.model_fields.keys() - ) + ) | {"cost_per_query"} # Sanity: the obvious top-level pricing fields are in the set. for field in ( "input_cost_per_token", @@ -78,6 +78,7 @@ class TestStripClientPricingOverrides: "input_cost_per_token": 0.0, "output_cost_per_token": 0.0, "cache_creation_input_token_cost": 0.0, + "cost_per_query": -1000.0, } _strip_client_pricing_overrides(data) assert data == { @@ -192,6 +193,7 @@ class TestStripClientPricingOverrides: @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_root_pricing_fields(): data = { + "cost_per_query": -1000.0, "model": "gpt-4", "messages": [{"role": "user", "content": "hi"}], "input_cost_per_token": 0.0, @@ -207,6 +209,7 @@ async def test_add_litellm_data_to_request_strips_root_pricing_fields(): version="test-version", ) + assert "cost_per_query" not in updated assert "input_cost_per_token" not in updated assert "output_cost_per_token" not in updated diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py index 4f2a93fa883..d4087bab2ea 100644 --- a/tests/unit/a2a_protocol/test_cost_calculator.py +++ b/tests/unit/a2a_protocol/test_cost_calculator.py @@ -455,9 +455,12 @@ async def test_asend_message_streaming_triggers_callbacks(): @pytest.mark.asyncio @pytest.mark.parametrize("stream", (False, True)) -@pytest.mark.parametrize("claimed_fee", (None, 0.0, 99.0)) +@pytest.mark.parametrize("claimed_fee", (None, -1000.0, 0.0, 99.0)) +@pytest.mark.parametrize("admitted", (False, True)) +@pytest.mark.parametrize("configured_fee", (None, 0.25)) async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing( - monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None + monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None, + admitted: bool, configured_fee: float | None ) -> None: import json from typing import Final @@ -476,15 +479,16 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing 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}, + litellm_params={"cost_per_query": configured_fee} if configured_fee is not None else {}, ) registry: Final = agent_registry.AgentRegistry() - registry.register_agent(target) + registry.register_agent(target.model_copy(update={"litellm_params": {"cost_per_query": 0.5}}) if admitted else 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 + if admitted: + auth.invoked_agent_policy = target + auth.invoked_agent_id = target.agent_id + auth.agent_invocation_cost = configured_fee or 0.0 def reply(request: httpx.Request) -> httpx.Response: body: Final = json.loads(request.content) @@ -514,6 +518,9 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing 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) + if configured_fee is None: + assert logger.response_cost in (None, 0.0) + else: + assert logger.response_cost == pytest.approx(configured_fee) finally: await client.close()