fix(agents): reject client supplied invocation pricing

This commit is contained in:
Joshua Valluru 2026-09-30 17:16:36 -07:00
parent b33fda1dd3
commit 95b94520d9
5 changed files with 35 additions and 22 deletions

View file

@ -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}))

View file

@ -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
)

View file

@ -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

View file

@ -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

View file

@ -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()