fix(agents): reserve configured fees for key and team budgets

This commit is contained in:
Joshua Valluru 2026-09-30 17:35:06 -07:00
parent 95b94520d9
commit e3f672064c
5 changed files with 65 additions and 34 deletions

View file

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

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

View file

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

View file

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

View file

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