From e3f672064cfa9a39afc8035d505dbda68a9ea4d4 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Wed, 30 Sep 2026 17:35:06 -0700 Subject: [PATCH] fix(agents): reserve configured fees for key and team budgets --- litellm/proxy/agent_endpoints/a2a_routing.py | 11 +----- .../auth/managed_authorization.py | 11 +++--- .../auth/test_managed_authorization.py | 22 +++++++++++ .../spend_tracking/test_budget_reservation.py | 37 +++++++++++++------ .../unit/a2a_protocol/test_cost_calculator.py | 18 +++++---- 5 files changed, 65 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 29d31251f94..26fdf83e0b0 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -13,7 +13,6 @@ 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( @@ -80,13 +79,5 @@ 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"]) - 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 - ) - 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 - ) + 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})) diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 11c8c790ad2..9fd1dfa7499 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") -AGENT_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)]) +_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)]) def invocation_target(route: str, body: Mapping[str, object]) -> str | None: @@ -254,11 +254,14 @@ async def prepare_agent_invocation( if target is None and registered_managed: raise_identity_failure(AgentIdentityFailure(message="Invoked agent no longer exists")) effective: Final = target if target is not None else registered + pricing: Final = effective.litellm_params or MappingProxyType({}) + fixed_fee: Final = pricing.get("cost_per_query") if ( not effective.identity_managed and effective.litellm_budget_table is None and auth.managed_agent_policy is None and auth.billing_agent_policy is None + and fixed_fee is None ): return if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth): @@ -271,8 +274,6 @@ async def prepare_agent_invocation( and (effective.identity_managed or effective.litellm_budget_table is not None) ): auth.billing_agent_policy = effective - pricing: Final = effective.litellm_params or MappingProxyType({}) - fixed_fee: Final = pricing.get("cost_per_query") billing_policy: Final = auth.billing_agent_policy bounded: Final = ( billing_policy is not None @@ -280,13 +281,13 @@ async def prepare_agent_invocation( and billing_policy.litellm_budget_table.max_budget is not None ) try: - fee: Final = AGENT_INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0) + fee: Final = _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( - AGENT_INVOCATION_COST.validate_python(pricing[field]) > 0 + _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/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py index c9b4d301018..88160562c9b 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -722,3 +722,25 @@ def test_registered_inference_routes_have_an_explicit_managed_access_decision(ro )) or normalized in ("/models", "/cursor/models", "/cursor/v1/models") concrete: Final = route.split("?")[0].replace("{model}", "model").replace("{model_name:path}", "model") assert managed_agent_route_allowed(concrete, None) is not unsupported, route + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fee", (-1.0, "invalid", float("inf"), float("nan"))) +async def test_unmanaged_invocation_rejects_invalid_configured_fees( + monkeypatch: pytest.MonkeyPatch, fee: float | str, +) -> None: + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + target: Final = AgentResponse( + agent_id="fee-target", agent_name="Fee target", agent_card_params={}, litellm_params={"cost_per_query": fee}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + with pytest.raises(HTTPException) as exc: + await prepare_agent_invocation(auth, "fee-target", None) + assert exc.value.status_code == 503 + assert "Agent invocation price is invalid" in str(exc.value.detail) + assert auth.agent_invocation_cost is None diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py index 02686a13a85..84b905af901 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -417,8 +417,10 @@ async def test_release_unbound_budget_reservation_leaves_a_bound_one_to_its_call @pytest.mark.asyncio @pytest.mark.parametrize("outcome", ("failed", "cancelled", "completed")) +@pytest.mark.parametrize("budget_owner", ("agent", "key", "team")) +@pytest.mark.parametrize("route", ("/a2a/target", "/v1/chat/completions")) async def test_budgeted_caller_reserves_unmanaged_agent_fees_before_concurrent_admission( - spend_counter_cache: DualCache, monkeypatch: pytest.MonkeyPatch, outcome: str + spend_counter_cache: DualCache, monkeypatch: pytest.MonkeyPatch, outcome: str, budget_owner: str, route: str ) -> None: import asyncio @@ -442,13 +444,25 @@ async def test_budgeted_caller_reserves_unmanaged_agent_fees_before_concurrent_a registry.register_agent(target) monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + token: Final = "fee-key" if budget_owner == "key" else None + team: Final = LiteLLM_TeamTable(team_id="fee-team", max_budget=0.5, spend=0.0) if budget_owner == "team" else None + counter_key: Final = ( + caller.budget_counter_key if budget_owner == "agent" + else f"spend:key:{token}" if budget_owner == "key" else "spend:team:fee-team" + ) + async def admit() -> dict[str, object] | None: - auth: Final = UserAPIKeyAuth(agent_id="caller", user_role="proxy_admin") - auth.billing_agent_policy = caller + auth: Final = UserAPIKeyAuth( + agent_id="caller" if budget_owner == "agent" else None, user_role="proxy_admin", + token=token, max_budget=0.5 if budget_owner == "key" else None, + team_id=team.team_id if team is not None else None, + ) + if budget_owner == "agent": + auth.billing_agent_policy = caller await prepare_agent_invocation(auth, "target", None) return await reserve_budget_for_request( - request_body={"method": "message/send"}, route="/a2a/target", llm_router=None, - valid_token=auth, team_object=None, user_object=None, prisma_client=None, + request_body={"method": "message/send", "model": "a2a/target"}, route=route, llm_router=None, + valid_token=auth, team_object=team, user_object=None, prisma_client=None, user_api_key_cache=UserApiKeyCache(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), fail_closed_budget_enforcement=True, ) @@ -459,22 +473,23 @@ async def test_budgeted_caller_reserves_unmanaged_agent_fees_before_concurrent_a assert all(result is None or isinstance(result, (dict, litellm.BudgetExceededError)) for result in results), results assert len(accepted) == 2, results assert len(rejected) == 6 - assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5) + assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.5) first: Final = accepted[0] if outcome == "completed": await proxy_server.increment_spend_counters( - token=None, team_id=None, user_id=None, response_cost=0.25, - billing_agent_id=caller.agent_id, billing_agent_counter_key=caller.budget_counter_key, + token=token, team_id=team.team_id if team is not None else None, user_id=None, response_cost=0.25, + billing_agent_id=caller.agent_id if budget_owner == "agent" else None, + billing_agent_counter_key=caller.budget_counter_key if budget_owner == "agent" else None, budget_reservation=first, ) await reconcile_budget_reservation(first, actual_cost=0.25) - assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5) + assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.5) with pytest.raises(litellm.BudgetExceededError): await admit() else: release: Final = release_budget_reservation_on_cancel if outcome == "cancelled" else release_budget_reservation await release(first) await release(first) - assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.25) + assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.25) assert await admit() is not None - assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5) + assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.5) diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py index d4087bab2ea..1db6b2bfdce 100644 --- a/tests/unit/a2a_protocol/test_cost_calculator.py +++ b/tests/unit/a2a_protocol/test_cost_calculator.py @@ -456,11 +456,11 @@ async def test_asend_message_streaming_triggers_callbacks(): @pytest.mark.asyncio @pytest.mark.parametrize("stream", (False, True)) @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)) +@pytest.mark.parametrize("changed_after_admission", (False, True)) +@pytest.mark.parametrize("configured_fee", (None, 0.0, 0.25)) async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing( monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None, - admitted: bool, configured_fee: float | None + changed_after_admission: bool, configured_fee: float | None ) -> None: import json from typing import Final @@ -471,6 +471,7 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing 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.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation from litellm.types.agents import AgentResponse await _reset_callbacks_and_settle_pending_logs() @@ -482,13 +483,14 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing litellm_params={"cost_per_query": configured_fee} if configured_fee is not None else {}, ) registry: Final = agent_registry.AgentRegistry() - registry.register_agent(target.model_copy(update={"litellm_params": {"cost_per_query": 0.5}}) if admitted else target) + registry.register_agent(target) monkeypatch.setattr(agent_registry, "global_agent_registry", registry) auth: Final = UserAPIKeyAuth(user_role="proxy_admin") - if admitted: - auth.invoked_agent_policy = target - auth.invoked_agent_id = target.agent_id - auth.agent_invocation_cost = configured_fee or 0.0 + await prepare_agent_invocation(auth, "fee-target", None) + assert auth.agent_invocation_cost == configured_fee + if changed_after_admission: + registry.deregister_agent(target.agent_name) + registry.register_agent(target.model_copy(update={"litellm_params": {"cost_per_query": 0.5}})) def reply(request: httpx.Request) -> httpx.Response: body: Final = json.loads(request.content)