mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): reject client supplied invocation pricing
This commit is contained in:
parent
b33fda1dd3
commit
95b94520d9
5 changed files with 35 additions and 22 deletions
|
|
@ -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}))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue