mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): reserve configured fees for key and team budgets
This commit is contained in:
parent
95b94520d9
commit
e3f672064c
5 changed files with 65 additions and 34 deletions
|
|
@ -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}))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue