diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql new file mode 100644 index 00000000000..033cd4101c2 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql @@ -0,0 +1,17 @@ +ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "budget_id" TEXT; + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentsTable_budget_id_key" ON "LiteLLM_AgentsTable"("budget_id"); + +-- AddForeignKey +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_AgentsTable_budget_id_fkey') THEN + ALTER TABLE "LiteLLM_AgentsTable" ADD CONSTRAINT "LiteLLM_AgentsTable_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE; + END IF; +END $$; + + +ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "spend_window" TIMESTAMP(3); + +ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "lifetime_budget_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index adfe2a0eee7..fc7325ddac3 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -36,6 +36,7 @@ model LiteLLM_BudgetTable { model_access_groups LiteLLM_ModelAccessGroupBudgetTable[] // multiple model access groups can have the same budget team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization + agents LiteLLM_AgentsTable[] } // Models on proxy @@ -83,6 +84,10 @@ model LiteLLM_AgentsTable { execution_mode String @default("autonomous") identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] + budget_id String? @unique + spend_window DateTime? + lifetime_budget_spend Float @default(0.0) + litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id]) tpm_limit Int? rpm_limit Int? session_tpm_limit Int? diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index aa41e63b40b..8d733971c70 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -471,13 +471,20 @@ async def asend_message( if custom_llm_provider: if request is None: raise ValueError("request is required for completion bridge") - return await _send_message_via_completion_bridge( + bridge_response: Final = await _send_message_via_completion_bridge( request=request, custom_llm_provider=custom_llm_provider, api_base=api_base, litellm_params=litellm_params, agent_extra_headers=agent_extra_headers, ) + bridge_prompt_tokens, bridge_completion_tokens, _ = await asyncify( + A2ARequestUtils.calculate_usage_from_request_response + )(request=request, response_dict=bridge_response.model_dump(mode="json", exclude_none=True)) + _set_usage_on_logging_obj(kwargs, bridge_prompt_tokens, bridge_completion_tokens) + _set_litellm_params_on_logging_obj(kwargs, litellm_params) + _set_agent_id_on_logging_obj(kwargs, agent_id) + return bridge_response # Standard A2A client flow if request is None: @@ -692,12 +699,35 @@ async def asend_message_streaming( request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params) ) - async for chunk in A2ACompletionBridgeHandler.handle_streaming( + bridge_name: Final = str(litellm_params.get("model") or agent_id or "agent") + existing_logging: Final = kwargs.get("litellm_logging_obj") + bridge_logging: Final = ( + existing_logging + if isinstance(existing_logging, Logging) + else _build_streaming_logging_obj( + request=request, + agent_name=bridge_name, + agent_id=agent_id, + litellm_params=litellm_params, + metadata=metadata, + proxy_server_request=proxy_server_request, + ) + ) + bridge_context: Final = {"litellm_logging_obj": bridge_logging} + _set_litellm_params_on_logging_obj(bridge_context, litellm_params) + _set_agent_id_on_logging_obj(bridge_context, agent_id) + bridge_stream: Final = A2ACompletionBridgeHandler.handle_streaming( request_id=str(request.id), params=params, litellm_params=litellm_params, api_base=api_base, agent_extra_headers=agent_extra_headers, + ) + async for chunk in A2AStreamingIterator( + stream=bridge_stream, + request=request, + logging_obj=bridge_logging, + agent_name=bridge_name, ): yield chunk return diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py index 8232d7cf2d8..3e69869e0d9 100644 --- a/litellm/a2a_protocol/streaming_iterator.py +++ b/litellm/a2a_protocol/streaming_iterator.py @@ -5,20 +5,24 @@ A2A Streaming Iterator with token tracking and logging support. import asyncio from collections.abc import AsyncIterator from datetime import datetime -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Generic, TypeVar import litellm from litellm._logging import verbose_logger from litellm.a2a_protocol.cost_calculator import A2ACostCalculator from litellm.a2a_protocol.utils import A2ARequestUtils from litellm.litellm_core_utils.asyncify import asyncify +from litellm.litellm_core_utils.core_helpers import bind_budget_reservation_to_callbacks from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj if TYPE_CHECKING: from a2a.compat.v0_3.types import SendStreamingMessageRequest, SendStreamingMessageResponse -class A2AStreamingIterator: +_StreamChunk = TypeVar("_StreamChunk", bound="SendStreamingMessageResponse | dict[str, object]") + + +class A2AStreamingIterator(Generic[_StreamChunk]): """ Async iterator for A2A streaming responses with token tracking. @@ -27,7 +31,7 @@ class A2AStreamingIterator: def __init__( self, - stream: AsyncIterator["SendStreamingMessageResponse"], + stream: AsyncIterator[_StreamChunk], request: "SendStreamingMessageRequest", logging_obj: LiteLLMLoggingObj, agent_name: str = "unknown", @@ -39,14 +43,14 @@ class A2AStreamingIterator: self.start_time = datetime.now() # Collect chunks for token counting - self.chunks: list[SendStreamingMessageResponse] = [] + self.chunks: list[_StreamChunk] = [] self.collected_text_parts: list[str] = [] - self.final_chunk: SendStreamingMessageResponse | None = None + self.final_chunk: _StreamChunk | None = None def __aiter__(self): return self - async def __anext__(self) -> "SendStreamingMessageResponse": + async def __anext__(self) -> _StreamChunk: try: chunk: Final = await self.stream.__anext__() @@ -69,20 +73,20 @@ class A2AStreamingIterator: await self._handle_stream_complete() raise - def _collect_text_from_chunk(self, chunk: "SendStreamingMessageResponse") -> None: + def _collect_text_from_chunk(self, chunk: _StreamChunk) -> None: """Extract text from a streaming chunk and add to collected parts.""" try: - chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} + chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True) text: Final = A2ARequestUtils.extract_text_from_response(chunk_dict) if text: self.collected_text_parts.append(text) except Exception: verbose_logger.debug("Failed to extract text from A2A streaming chunk") - def _is_completed_chunk(self, chunk: "SendStreamingMessageResponse") -> bool: + def _is_completed_chunk(self, chunk: _StreamChunk) -> bool: """Check if chunk indicates stream completion.""" try: - chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} + chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True) result: Final = chunk_dict.get("result", {}) if isinstance(result, dict): status: Final = result.get("status", {}) @@ -127,6 +131,8 @@ class A2AStreamingIterator: # Build result for logging result: Final = self._build_logging_result(usage) + bind_budget_reservation_to_callbacks(self.logging_obj.litellm_params) + # Call success handlers - they will build standard_logging_object asyncio.create_task( self.logging_obj.dispatch_success_handlers( @@ -160,7 +166,11 @@ class A2AStreamingIterator: # Add final chunk result if available if self.final_chunk: try: - chunk_dict: Final = self.final_chunk.model_dump(mode="json", exclude_none=True) + chunk_dict: Final = ( + self.final_chunk + if isinstance(self.final_chunk, dict) + else self.final_chunk.model_dump(mode="json", exclude_none=True) + ) result["result"] = chunk_dict.get("result", {}) except Exception: pass diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 238b7cc3fdd..c022b34282f 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -1529,7 +1529,14 @@ def completion_cost( completion_tokens = token_counter(model=model, text=completion) # Handle A2A calls before model check - A2A doesn't require a model - if call_type in _A2A_CALL_TYPES: + if call_type in _A2A_CALL_TYPES or ( + custom_llm_provider == "a2a" + and litellm_logging_obj is not None + and (litellm_logging_obj.model_call_details.get("litellm_params") or MappingProxyType({})).get( + "cost_per_query" + ) + is not None + ): from litellm.a2a_protocol.cost_calculator import A2ACostCalculator return A2ACostCalculator.calculate_a2a_cost(litellm_logging_obj=litellm_logging_obj) diff --git a/litellm/main.py b/litellm/main.py index 8c9d7f2513d..570366e0792 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5662,6 +5662,7 @@ def completion( preset_cache_key=preset_cache_key, no_log=no_log, cost_per_second=cost_per_second, + cost_per_query=kwargs.get("cost_per_query"), input_cost_per_second=input_cost_per_second, input_cost_per_token=input_cost_per_token, output_cost_per_second=output_cost_per_second, diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index fa4b36a03aa..89bdd6647ce 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -2091,6 +2091,79 @@ "title": "APIKeySecurityScheme", "type": "object" }, + "AgentBudgetConfig": { + "additionalProperties": false, + "properties": { + "budget_duration": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Budget Duration" + }, + "max_budget": { + "minimum": 0.0, + "title": "Max Budget", + "type": "number" + } + }, + "required": [ + "max_budget" + ], + "title": "AgentBudgetConfig", + "type": "object" + }, + "AgentBudgetState": { + "properties": { + "budget_duration": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Budget Duration" + }, + "budget_id": { + "title": "Budget Id", + "type": "string" + }, + "budget_reset_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Budget Reset At" + }, + "max_budget": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Max Budget" + } + }, + "required": [ + "budget_id" + ], + "title": "AgentBudgetState", + "type": "object" + }, "AgentCapabilities": { "description": "Defines optional capabilities supported by an agent.", "properties": { @@ -2378,6 +2451,16 @@ "title": "Agent Name", "type": "string" }, + "budget": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentBudgetConfig" + }, + { + "type": "null" + } + ] + }, "enabled": { "title": "Enabled", "type": "boolean" @@ -3056,6 +3139,17 @@ "title": "Agent Name", "type": "string" }, + "budget_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Budget Id" + }, "created_at": { "anyOf": [ { @@ -3152,6 +3246,21 @@ } ] }, + "lifetime_budget_spend": { + "default": 0.0, + "title": "Lifetime Budget Spend", + "type": "number" + }, + "litellm_budget_table": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentBudgetState" + }, + { + "type": "null" + } + ] + }, "litellm_params": { "anyOf": [ { @@ -4011,6 +4120,16 @@ "title": "Agent Name", "type": "string" }, + "budget": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentBudgetConfig" + }, + { + "type": "null" + } + ] + }, "enabled": { "title": "Enabled", "type": "boolean" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d2fad212dd9..b74b9a240df 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4104,6 +4104,7 @@ class SpendLogsRouterMetadata(TypedDict): class SpendLogsMetadata(TypedDict): + billing_agent_counter_key: ReadOnly[NotRequired[str | None]] actor_agent_id: ReadOnly[NotRequired[str | None]] target_agent_id: ReadOnly[NotRequired[str | None]] billing_agent_id: ReadOnly[NotRequired[str | None]] diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 6ddcd20d919..23d11a5a9c5 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy.a2a.version_convert import ( normalize_request_params, normalize_stream_event, ) +from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy from litellm.proxy.agent_endpoints.databricks_oauth import ( DATABRICKS_OAUTH_PARAM, resolve_databricks_app_auth_header, @@ -706,9 +707,10 @@ async def invoke_agent_a2a( params.pop(key) # Find the agent - agent: Final = await _get_agent(agent_id) - if agent is None: + registered: Final = await _get_agent(agent_id) + if registered is None: return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404) + agent: Final = agent_invocation_policy(user_api_key_dict, registered) served_version: Final = _served_version(agent, request, original_method) @@ -732,7 +734,14 @@ async def invoke_agent_a2a( agent_name: Final = agent_card_params.get("name", agent_id) # Get litellm_params (may include custom_llm_provider for completion bridge) - litellm_params: dict[str, object] = agent.litellm_params or {} + litellm_params: dict[ + str, object + ] = { # mutable-ok: A2A SDK and completion bridge accept provider parameters as a dict + **(agent.litellm_params or MappingProxyType({})), + "cost_per_query": user_api_key_dict.agent_invocation_cost + if (agent.litellm_params or MappingProxyType({})).get("cost_per_query") is not None + else None, + } custom_llm_provider: Final = litellm_params.get("custom_llm_provider") # Hand the authenticated key hash to the completion bridge so provider diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index c57315ebc21..96c0b0ceb5a 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -5,6 +5,7 @@ Handles routing for A2A agents (models with "a2a/" prefix). Looks up agents in the registry and injects their API base URL. """ +from types import MappingProxyType from typing import Any, Final from fastapi import HTTPException @@ -12,12 +13,16 @@ 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_policy +from litellm.types.agents import AgentResponse async def route_a2a_agent_request( data: dict, route_type: str, user_api_key_dict: UserAPIKeyAuth | None = None, + *, + registered_agent: AgentResponse | None = None, ) -> Any | None: """ Route A2A agent requests directly to litellm with injected API base. @@ -46,12 +51,14 @@ async def route_a2a_agent_request( agent_name: Final = model_name[4:] # Look up agent in registry - agent: Final = await get_agent_with_read_through(agent_name) - if agent is None: + registered: Final = registered_agent or await get_agent_with_read_through(agent_name) + if registered is None: verbose_proxy_logger.error("[A2A] Agent '%s' not found in registry", agent_name) route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False) + agent: Final = agent_invocation_policy(user_api_key_dict, registered) + # Verify the caller is permitted to use this agent (admins bypass the check) is_admin: Final = user_api_key_dict is not None and ( user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN @@ -75,7 +82,11 @@ async def route_a2a_agent_request( raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False) # Inject API base and route to litellm - data["api_base"] = agent.agent_card_params["url"] - verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"]) + api_base: Final = agent.agent_card_params["url"] + verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, api_base) - return getattr(litellm, f"{route_type}")(**data) + invocation_fee: Final = user_api_key_dict.agent_invocation_cost if user_api_key_dict is not None else None + provider_data: Final = MappingProxyType({key: value for key, value in data.items() if key != "litellm_params"}) + return getattr(litellm, f"{route_type}")( + **MappingProxyType({**provider_data, "api_base": api_base, "cost_per_query": invocation_fee}) + ) diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index 7929f67720d..5104d904f26 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -91,6 +91,15 @@ class AgentRecord(Protocol): @property def spend(self) -> float: ... + @property + def lifetime_budget_spend(self) -> float: ... + + @property + def budget_id(self) -> str | None: ... + + @property + def litellm_budget_table(self) -> "prisma_models.LiteLLM_BudgetTable | None": ... + def model_dump(self) -> AgentRecordDump: ... def __iter__(self) -> Iterator[tuple[str, object]]: ... @@ -654,7 +663,7 @@ class AgentRegistry: # Create agent in DB created_agent: Final = await agents_table(prisma_client).create( data={**create_data, **await _managed_fields(agent, None, created_by, prisma_client)}, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) return AgentResponse.model_validate(created_agent.model_dump()) @@ -717,7 +726,7 @@ class AgentRegistry: """ try: existing_record: Final = await agents_table(prisma_client).find_unique( - where={"agent_id": agent_id}, include={"identity": True} + where={"agent_id": agent_id}, include={"identity": True, "litellm_budget_table": True} ) if existing_record is None: raise Exception(f"Agent with ID {agent_id} not found") @@ -771,7 +780,7 @@ class AgentRegistry: "updated_by": updated_by, "updated_at": datetime.now(timezone.utc), }, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) if patched_agent is None: raise ValueError(f"Agent not found, passed agent_id={agent_id}") @@ -808,7 +817,7 @@ class AgentRegistry: # caller echoed back redacted (or omitted) rather than persisting # the marker -- or nothing -- over the real stored credential. existing_row: Final = await agents_table(prisma_client).find_unique( - where={"agent_id": agent_id}, include={"identity": True} + where={"agent_id": agent_id}, include={"identity": True, "litellm_budget_table": True} ) existing_litellm_params: Final = parse_agent_litellm_params( existing_row.litellm_params if existing_row is not None else None @@ -877,7 +886,7 @@ class AgentRegistry: prisma_client, ), }, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) if updated_agent is None: @@ -900,7 +909,7 @@ class AgentRegistry: try: agents_from_db: Final = await agents_table(prisma_client).find_many( order={"created_at": "desc"}, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) agents: Final[list[dict[str, object]]] = [] diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 17d988127ec..d945c6e1ad2 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -99,12 +99,18 @@ def managed_inference_request( cli_model: str | None, path_model: object = None, query_model: object = None, + *, + auth: UserAPIKeyAuth | None = None, + require_model: bool = True, + model_group_alias: object = None, ) -> dict[str, object]: from litellm.proxy.auth.route_checks import RouteChecks if route in _MANAGED_REALTIME_ROUTES: model: Final = query_model or body.get("model") if not isinstance(model, str) or not model: + if not require_model: + return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata raise_identity_failure( AgentIdentityFailure(message="Managed inference requires an explicit or configured model") ) @@ -117,8 +123,32 @@ def managed_inference_request( endpoint_model: Final = path_model or ( query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None ) - effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind) + from litellm.proxy.common_utils.model_listing_utils import CallerAliases, alias_target + from litellm.proxy.litellm_pre_call_utils import ( + _update_model_if_key_alias_exists, + _update_model_if_team_alias_exists, + ) + + aliased_body: Final = dict(body) # mutable-ok: existing alias helpers rewrite their request copy + if auth is not None: + _update_model_if_team_alias_exists(aliased_body, auth) + _update_model_if_key_alias_exists(aliased_body, auth) + selected: Final = resolve_inference_model(aliased_body.get("model"), settings, cli_model, endpoint_model, kind=kind) + import litellm + + aliased: Final = ( + alias_target(selected, CallerAliases((), (litellm.model_alias_map, auth.aliases))) or selected + if isinstance(selected, str) and auth is not None + else selected + ) + from litellm.router_utils.common_utils import resolve_model_group_alias + + effective: Final = ( + (resolve_model_group_alias(model_group_alias, aliased) or aliased) if isinstance(aliased, str) else aliased + ) if not isinstance(effective, str) or not effective: + if not require_model: + return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata raise_identity_failure( AgentIdentityFailure(message="Managed inference requires an explicit or configured model") ) @@ -162,6 +192,8 @@ async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists")) return if not agent.identity_managed: + if agent.litellm_budget_table is not None: + auth.billing_agent_policy = agent return if auth.jwt_claims and auth.managed_agent_context is None: raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity")) @@ -202,6 +234,24 @@ def actor_admission_failure( return None +async def check_agent_budget(auth: UserAPIKeyAuth) -> None: + import litellm + from litellm.proxy.proxy_server import get_current_spend + + agent: Final = auth.billing_agent_policy + if agent is None or agent.litellm_budget_table is None or agent.litellm_budget_table.max_budget is None: + return + budget: Final = agent.litellm_budget_table.max_budget + spend: Final = await get_current_spend( + counter_key=agent.budget_counter_key, + fallback_spend=agent.budget_spend, + max_budget=budget, + fallback_authoritative=True, + ) + if spend >= budget: + 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)]) @@ -234,19 +284,77 @@ 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 - if not effective.identity_managed and auth.managed_agent_policy is None: - return + pricing: Final = effective.litellm_params or MappingProxyType({}) + fixed_fee: Final = pricing.get("cost_per_query") if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth): raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent")) auth.invoked_agent_id = effective.agent_id auth.invoked_agent_policy = effective - if auth.agent_id is None and effective.identity_managed: + 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 ( + billable + and auth.agent_id is None + and (effective.identity_managed or effective.litellm_budget_table is not None) + ): auth.billing_agent_policy = effective - raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0 + billing_policy: Final = auth.billing_agent_policy + bounded: Final = ( + billing_policy is not None + and billing_policy.litellm_budget_table is not None + and billing_policy.litellm_budget_table.max_budget is not None + ) try: - fee: Final = _INVOCATION_COST.validate_python(raw_fee) + 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( + _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 + ) + ) except ValidationError: raise_identity_failure( AgentIdentityFailure(code="policy_unavailable", message="Agent invocation price is invalid") ) + if unbounded_token_price: + raise_identity_failure( + AgentIdentityFailure( + code="policy_unavailable", + message="Budgeted token-priced agent invocations require a fixed cost_per_query before execution", + ) + ) auth.agent_invocation_cost = fee + + +def agent_invocation_policy(auth: UserAPIKeyAuth | None, registered: AgentResponse) -> AgentResponse: + admitted: Final = auth.invoked_agent_policy if auth is not None else None + if admitted is not None and auth is not None: + matching: Final = auth.invoked_agent_id == registered.agent_id == admitted.agent_id + captured_price: Final = (admitted.litellm_params or MappingProxyType({})).get( + "cost_per_query" + ) is None or auth.agent_invocation_cost is not None + if matching and captured_price: + return admitted + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Agent dispatch does not match its admission") + ) + if ( + registered.identity_managed + or registered.identity is not None + or registered.litellm_budget_table is not None + or (registered.litellm_params or MappingProxyType({})).get("cost_per_query") is not None + ): + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Agent dispatch requires a matching admission") + ) + return registered diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index e1b2ac63d51..a76ec57b630 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -13,7 +13,7 @@ import os import uuid from collections.abc import Mapping, Sequence from types import MappingProxyType -from typing import Annotated, Final, TypedDict +from typing import TYPE_CHECKING, Annotated, Final, TypedDict from fastapi import APIRouter, Depends, HTTPException, Query, Request from pydantic import ValidationError @@ -77,6 +77,7 @@ from litellm.types.agents import ( ) from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.proxy.agent_identity import ( + AgentBudgetState, AgentIdentityBinding, AgentIdentityFailure, EntraIdentityConfig, @@ -129,6 +130,12 @@ def _build_merged_agent_card( ) +if TYPE_CHECKING: + from prisma.types import LiteLLM_AgentsTableInclude + +_AGENT_BUDGET_INCLUDE: Final["LiteLLM_AgentsTableInclude"] = {"litellm_budget_table": True} + + router: Final = APIRouter() @@ -374,8 +381,10 @@ async def get_agents( if agent_ids: db_agents: Final = await agents_table(prisma_client).find_many( where={"agent_id": {"in": agent_ids}}, + include=_AGENT_BUDGET_INCLUDE, ) - spend_map: Final = {a.agent_id: a.spend for a in db_agents} + spend_map: Final = MappingProxyType({a.agent_id: a.spend for a in db_agents}) + budget_map: Final = MappingProxyType({a.agent_id: a for a in db_agents}) for agent in returned_agents: matched_spends = tuple( spend_map[alias_id] @@ -384,6 +393,14 @@ async def get_agents( ) if matched_spends: agent.spend = sum(matched_spends) + if (budget_row := budget_map.get(agent.agent_id)) is not None: + agent.lifetime_budget_spend = budget_row.lifetime_budget_spend + agent.budget_id = budget_row.budget_id + agent.litellm_budget_table = ( + AgentBudgetState.model_validate(budget_row.litellm_budget_table.model_dump()) + if budget_row.litellm_budget_table is not None + else None + ) await _attach_keys_to_agents(returned_agents, prisma_client) # add is_public field to each agent - we do it this way, to allow setting config agents as public @@ -676,7 +693,7 @@ async def get_agent_by_id( if agent is None: agent_row: Final = await agents_table(prisma_client).find_unique( where={"agent_id": agent_id}, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) if agent_row is not None: agent_dict: Final = agent_row.model_dump() @@ -688,9 +705,18 @@ async def get_agent_by_id( agent = AgentResponse(**agent_dict) else: # Agent found in memory — refresh spend from DB - db_row: Final = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id}) + db_row: Final = await agents_table(prisma_client).find_unique( + where={"agent_id": agent_id}, include=_AGENT_BUDGET_INCLUDE + ) if db_row is not None: agent.spend = db_row.spend + agent.lifetime_budget_spend = db_row.lifetime_budget_spend + agent.budget_id = db_row.budget_id + agent.litellm_budget_table = ( + AgentBudgetState.model_validate(db_row.litellm_budget_table.model_dump()) + if db_row.litellm_budget_table is not None + else None + ) if agent is None: raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found") @@ -766,7 +792,7 @@ async def update_agent( try: # Check if agent exists existing_agent = await agents_table(prisma_client).find_unique( - where={"agent_id": agent_id}, include={"identity": True} + where={"agent_id": agent_id}, include={"identity": True, "litellm_budget_table": True} ) if existing_agent is not None: existing_agent = existing_agent.model_dump() @@ -873,7 +899,7 @@ async def patch_agent( try: # Check if agent exists existing_agent = await agents_table(prisma_client).find_unique( - where={"agent_id": agent_id}, include={"identity": True} + where={"agent_id": agent_id}, include={"identity": True, "litellm_budget_table": True} ) if existing_agent is not None: existing_agent = existing_agent.model_dump() @@ -965,7 +991,7 @@ async def delete_agent( try: # Check if agent exists existing_agent = await agents_table(prisma_client).find_unique( - where={"agent_id": agent_id}, include={"identity": True} + where={"agent_id": agent_id}, include={"identity": True, "litellm_budget_table": True} ) if existing_agent is not None: existing_agent = dict[str, object](existing_agent) diff --git a/litellm/proxy/agent_endpoints/identity_store.py b/litellm/proxy/agent_endpoints/identity_store.py index 3c8163a8838..ac0ebac7865 100644 --- a/litellm/proxy/agent_endpoints/identity_store.py +++ b/litellm/proxy/agent_endpoints/identity_store.py @@ -70,6 +70,7 @@ class AgentIdentityStore: include: Final[LiteLLM_AgentsTableInclude] = { "identity": True, "object_permission": True, + "litellm_budget_table": True, } row: Final = await self.agents.table.find_unique(where=where, include=include) if row is None: diff --git a/litellm/proxy/agent_endpoints/managed_identity.py b/litellm/proxy/agent_endpoints/managed_identity.py index abab21901ee..e2289e3bf47 100644 --- a/litellm/proxy/agent_endpoints/managed_identity.py +++ b/litellm/proxy/agent_endpoints/managed_identity.py @@ -7,8 +7,10 @@ from fastapi import HTTPException from pydantic import TypeAdapter, ValidationError from typing_extensions import ReadOnly +from litellm.proxy.common_utils.timezone_utils import budget_duration_error, get_budget_reset_time from litellm.types.agents import AgentResponse from litellm.types.proxy.agent_identity import ( + AgentBudgetConfig, AgentExecutionMode, AgentIdentityBinding, AgentIdentityFailure, @@ -57,12 +59,30 @@ class IdentityHistoryWrite(TypedDict): create: ReadOnly[IdentityHistoryEntry] +class BudgetFields(TypedDict, total=False): + max_budget: ReadOnly[float] + budget_duration: ReadOnly[str | None] + budget_reset_at: ReadOnly[datetime | None] + updated_by: ReadOnly[str] + created_by: ReadOnly[str] + + +class BudgetRelationWrite(TypedDict, total=False): + create: ReadOnly[BudgetFields] + update: ReadOnly[BudgetFields] + disconnect: ReadOnly[bool] + + class ManagedWriteFields(TypedDict, total=False): enabled: ReadOnly[bool] execution_mode: ReadOnly[AgentExecutionMode] identity_managed: ReadOnly[bool] identity: ReadOnly[IdentityRelationWrite] retired_identities: ReadOnly[IdentityHistoryWrite] + litellm_budget_table: ReadOnly[BudgetRelationWrite] + spend_window: ReadOnly[datetime | None] + spend: ReadOnly[float] + lifetime_budget_spend: ReadOnly[float] def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn: @@ -109,14 +129,18 @@ def managed_write_fields( return failure empty: Final[ManagedWriteFields] = {} identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty + budget_fields: Final = ( + _budget_write(incoming["budget"], existing, updated_by) if "budget" in incoming else empty + ) result: Final[ManagedWriteFields] = { **({"enabled": incoming["enabled"] is True} if "enabled" in incoming else {}), **({"execution_mode": mode} if "execution_mode" in incoming else {}), + **budget_fields, **identity_fields, } return result except (ValidationError, ValueError) as exc: - return AgentIdentityFailure(message=f"Invalid agent identity configuration: {exc}") + return AgentIdentityFailure(message=f"Invalid agent identity or budget configuration: {exc}") def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields: @@ -165,6 +189,61 @@ def _identity_write(identity: EntraIdentityConfig | None, existing: AgentRespons return result +def _budget_write(raw: object, existing: AgentResponse | None, updated_by: str) -> ManagedWriteFields: + if raw is None: + if existing and existing.budget_id: + disconnected: Final[ManagedWriteFields] = { + "litellm_budget_table": {"disconnect": True}, + "spend_window": None, + } + return disconnected + empty: Final[ManagedWriteFields] = {} + return empty + budget: Final = AgentBudgetConfig.model_validate(raw) + duration_error: Final = budget_duration_error(budget.budget_duration) + if duration_error is not None: + raise ValueError(duration_error) + creating_lifetime: Final = budget.budget_duration is None and ( + existing is None + or existing.litellm_budget_table is None + or existing.litellm_budget_table.budget_duration is not None + ) + fields: Final[BudgetFields] = { + "max_budget": budget.max_budget, + "budget_duration": budget.budget_duration, + "updated_by": updated_by, + "budget_reset_at": ( + existing.litellm_budget_table.budget_reset_at + if existing + and existing.litellm_budget_table + and existing.litellm_budget_table.budget_duration == budget.budget_duration + else get_budget_reset_time(budget.budget_duration) + if budget.budget_duration + else None + ), + } + result: Final[ManagedWriteFields] = { + "spend_window": fields["budget_reset_at"], + **({"lifetime_budget_spend": 0.0} if creating_lifetime else {}), + **( + {"spend": 0.0} + if fields["budget_reset_at"] is not None + and ( + existing is None + or existing.litellm_budget_table is None + or existing.litellm_budget_table.budget_reset_at != fields["budget_reset_at"] + ) + else {} + ), + "litellm_budget_table": ( + {"update": fields} + if existing and existing.budget_id and not creating_lifetime + else {"create": {**fields, "created_by": updated_by}} + ), + } + return result + + def classify_agent_subject( binding: AgentIdentityBinding, claims: Mapping[str, object], diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3ec430332ee..24b4a8047ec 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1073,6 +1073,10 @@ async def common_checks( key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) + if valid_token is not None and not skip_all_budget_checks: + from litellm.proxy.agent_endpoints.auth.managed_authorization import check_agent_budget + + await check_agent_budget(valid_token) await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router) await _check_agent_caller_model_access( model=_model, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5194f62cf78..f3cc82eb895 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2988,6 +2988,7 @@ async def _run_centralized_common_checks( request=request, llm_router=llm_router, team_id=user_api_key_auth_obj.team_id, + agent_invocation_cost=user_api_key_auth_obj.agent_invocation_cost, ) # Pin the metadata variable name (litellm_metadata vs metadata) before @@ -3161,7 +3162,10 @@ def _should_skip_budget_checks( request: Request | None, llm_router: Any | None, team_id: str | None = None, + agent_invocation_cost: float | None = None, ) -> bool: + if agent_invocation_cost is not None and agent_invocation_cost > 0: + return False model: Final = _get_model_from_request_context( request_data=request_data, route=route, @@ -3229,7 +3233,14 @@ async def _authorize_authenticated_request( prepare_agent_invocation, ) from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore - from litellm.proxy.proxy_server import general_settings, prisma_client, user_model + from litellm.proxy.proxy_server import ( + general_settings, + llm_router, + prisma_client, + proxy_config, + proxy_logging_obj, + user_model, + ) store: Final = AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None if user_api_key_auth_obj.agent_id is not None: @@ -3238,26 +3249,45 @@ async def _authorize_authenticated_request( route, request.method ): raise HTTPException(403, "Agent identities can only access inference and agent discovery routes") - authorized_data: Final = ( - managed_inference_request( - route, - request_data, - general_settings, - user_model, - request.path_params.get("model") or request.path_params.get("model_name"), - request.query_params.get("model"), + router_settings: Final = ( + await proxy_config.get_hierarchical_router_settings( + user_api_key_dict=user_api_key_auth_obj, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, ) - if user_api_key_auth_obj.managed_agent_policy is not None + if llm_router is not None and RouteChecks.is_llm_api_route(route=route) + else None + ) + inference_data: Final = managed_inference_request( + route, + request_data, + general_settings, + user_model, + request.path_params.get("model") or request.path_params.get("model_name"), + _safe_get_request_query_params(request).get("model"), + model_group_alias=router_settings.get("model_group_alias") + if isinstance(router_settings, Mapping) + else None, + auth=user_api_key_auth_obj, + require_model=user_api_key_auth_obj.managed_agent_policy is not None, + ) + target_name: Final = invocation_target(route, inference_data) + authorized_data: Final = ( + inference_data + if target_name is not None or user_api_key_auth_obj.managed_agent_policy is not None else request_data ) - target_name: Final = invocation_target(route, authorized_data) if target_name is not None: await prepare_agent_invocation( user_api_key_auth_obj, target_name, store, - billable=request_data.get("method") - in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"), + billable=request.method == "POST" + and ( + not RouteChecks.check_route_access(route, ("/a2a/{agent_id}", "/v1/a2a/{agent_id}")) + or request_data.get("method") + in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage") + ), ) await _run_centralized_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e7a08711eb1..4005b5bebd8 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3967,7 +3967,11 @@ class ProxyBaseLLMRequestProcessing: # Starlette closes on disconnect, so the nested iterator hook (which # only sees GeneratorExit on GC) cannot own the refund. client_disconnected = not stream_completed - if not delivered_chunk and not _withheld_provider_output(response): + if ( + not delivered_chunk + and not _withheld_provider_output(response) + and user_api_key_dict.agent_invocation_cost is None + ): from litellm.proxy.spend_tracking.budget_reservation import ( release_budget_reservation_on_cancel, ) diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py index 8a82e253c5c..c512d782db3 100644 --- a/litellm/proxy/common_utils/registry_read_through.py +++ b/litellm/proxy/common_utils/registry_read_through.py @@ -177,6 +177,7 @@ async def _resync_agents(agent_id_or_name: str) -> bool: include_permission: Final[LiteLLM_AgentsTableInclude] = { "object_permission": True, "identity": True, + "litellm_budget_table": True, } async with AGENT_RECONCILE_LOCK: if _agent_from_registry(agent_id_or_name) is not None: diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index b35b876b475..4c875936b31 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -841,6 +841,7 @@ class ResetBudgetJob: _queue_budget_linked_resets(uow.projects, cascade, extra=_SPENT_ROWS_WHERE) _queue_enduser_resets(uow.endusers, cascade) for budget_id, budget_reset_at in cascade.budget_resets: + uow.agents.queue_window_reset(budget_id, budget_reset_at, cascade.rollover_caps.get(budget_id)) uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at) async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None: diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 72553e82283..55839620831 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -82,9 +82,12 @@ from litellm.proxy.spend_tracking.savings import ( ) from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.repositories.prisma_protocols import BatchTable +from litellm.types.agents import agent_spend_filter from litellm.types.utils import CallTypes if TYPE_CHECKING: + from prisma.types import LiteLLM_AgentsTableUpdateManyMutationInput, LiteLLM_AgentsTableWhereInput + from litellm.proxy.db.autorouter_session_rollup import AutoRouterTurnTransaction from litellm.proxy.db.baseline_accounting import DailyBaselineAttribution from litellm.proxy.utils import PrismaClient, ProxyLogging @@ -162,6 +165,26 @@ def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable) return _ENTITY_SPEND_TABLES[table_accessor](batcher) +def _queue_lifetime_agent_spend(table: BatchTable, counter_key: str, response_cost: float) -> None: + lifetime_filter: Final = agent_spend_filter(counter_key) + lifetime_data: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = { + "lifetime_budget_spend": {"increment": response_cost} + } + history_filter: Final[LiteLLM_AgentsTableWhereInput] = { + "agent_id": lifetime_filter.get("agent_id"), + "spend_window": None, + } + history_data: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = {"spend": {"increment": response_cost}} + table.update_many( + where=lifetime_filter, + data=lifetime_data, + ) + table.update_many( + where=history_filter, + data=history_data, + ) + + class _SpendBatchManager(Protocol): async def __aenter__(self) -> _SpendBatch: ... @@ -1068,12 +1091,15 @@ class DBSpendUpdateWriter: router=get_llm_router(), ) - _agent_id_for_spend: Final = payload_copy.get("agent_id") + _agent_id_for_spend: Final = payload_copy.get("billing_agent_id", payload_copy.get("agent_id")) try: + spend_metadata: Final = _SPEND_METADATA_ADAPTER.validate_json(payload_copy.get("metadata") or "{}") + captured_counter: Final = spend_metadata.get("billing_agent_counter_key") await self._update_agent_db( response_cost=response_cost, agent_id=_agent_id_for_spend, prisma_client=prisma_client, + counter_key=captured_counter if isinstance(captured_counter, str) else None, ) except Exception: verbose_proxy_logger.debug( @@ -1351,15 +1377,19 @@ class DBSpendUpdateWriter: response_cost: float | None, agent_id: str | None, prisma_client: PrismaClient | None, + *, + counter_key: str | None = None, ): try: if agent_id is None or prisma_client is None: return + if counter_key is not None and agent_spend_filter(counter_key).get("agent_id") != agent_id: + raise ValueError("Agent spend counter does not match the billed agent") await self.spend_update_queue.add_update( update=SpendUpdateQueueItem( entity_type=Litellm_EntityType.AGENT, - entity_id=agent_id, + entity_id=counter_key or agent_id, response_cost=response_cost, ) ) @@ -2340,8 +2370,19 @@ class DBSpendUpdateWriter: entity_id, response_cost, ) + if table_accessor == "litellm_agentstable" and entity_id.startswith( + "spend:agent_lifetime:" + ): + _queue_lifetime_agent_spend( + _entity_spend_table(batcher, table_accessor), entity_id, response_cost + ) + continue _entity_spend_table(batcher, table_accessor).update_many( - where={where_field: entity_id}, + where=( + agent_spend_filter(entity_id) + if table_accessor == "litellm_agentstable" + else {where_field: entity_id} # mutable-ok: Prisma filter + ), data={"spend": {"increment": response_cost}}, ) break @@ -2914,13 +2955,14 @@ class DBSpendUpdateWriter: if prisma_client is None: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") return - if payload["agent_id"] is None: + charged_agent_id: Final = payload.get("billing_agent_id", payload["agent_id"]) + if charged_agent_id is None: return payload_with_agent_id: Final = cast( SpendLogsPayload, { **payload, - "agent_id": payload["agent_id"], + "agent_id": charged_agent_id, }, ) base_daily_transaction: Final = await self._common_add_spend_log_transaction_to_daily_transaction( @@ -2929,8 +2971,8 @@ class DBSpendUpdateWriter: if base_daily_transaction is None: return endpoint_str: Final = base_daily_transaction.get("endpoint") or "" - daily_transaction_key = f"{payload['agent_id']}_{base_daily_transaction['date']}_{payload_with_agent_id['api_key']}_{payload_with_agent_id['model']}_{payload_with_agent_id['custom_llm_provider']}_{endpoint_str}" - daily_transaction: Final = DailyAgentSpendTransaction(agent_id=payload["agent_id"], **base_daily_transaction) + daily_transaction_key = f"{charged_agent_id}_{base_daily_transaction['date']}_{payload_with_agent_id['api_key']}_{payload_with_agent_id['model']}_{payload_with_agent_id['custom_llm_provider']}_{endpoint_str}" + daily_transaction: Final = DailyAgentSpendTransaction(agent_id=charged_agent_id, **base_daily_transaction) await self.daily_agent_spend_update_queue.add_update(update={daily_transaction_key: daily_transaction}) async def add_spend_log_transaction_to_daily_tag_transaction( diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index f8e102d2682..3222868dad8 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -28,6 +28,7 @@ from litellm.proxy.spend_tracking.spend_counter_batch import read_batched_spend_ from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import ( + AgentsRepository, BudgetWindowSpendRepository, EndUserRepository, SpendLogsRepository, @@ -38,9 +39,14 @@ from litellm.repositories.user_repository import UserRepository from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, ) +from litellm.types.agents import agent_budget_counter_key if TYPE_CHECKING: - from prisma.types import LiteLLM_EndUserTableWhereUniqueInput + from prisma.types import ( + LiteLLM_AgentsTableInclude, + LiteLLM_AgentsTableWhereUniqueInput, + LiteLLM_EndUserTableWhereUniqueInput, + ) from litellm.caching.dual_cache import DualCache from litellm.proxy.utils import PrismaClient @@ -170,6 +176,37 @@ class SpendCounterReseed: return await OrganizationRepository(prisma_client).table.find_unique( where={"organization_id": counter_key[len("spend:org:") :]} ) + if counter_key.startswith("spend:agent_lifetime:"): + _, _, budget_id, agent_id = counter_key.split(":", 3) + lifetime_where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id} + lifetime: Final = await AgentsRepository(prisma_client, use_writer=True).table.find_unique( + where=lifetime_where + ) + if lifetime is None: + return None + return lifetime.model_copy( + update=MappingProxyType( + {"spend": lifetime.lifetime_budget_spend if lifetime.budget_id == budget_id else 0.0} + ) + ) + if counter_key.startswith("spend:agent_window:"): + parts: Final = counter_key.split(":", 3) + if len(parts) != 4: + return None + window_where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": parts[3]} + window_include: Final[LiteLLM_AgentsTableInclude] = {"litellm_budget_table": True} + row: Final = await AgentsRepository(prisma_client, use_writer=True).table.find_unique( + where=window_where, include=window_include + ) + if row is None: + return None + current_key: Final = agent_budget_counter_key(row.agent_id, row.spend_window) + return row if current_key == counter_key else row.model_copy(update=MappingProxyType({"spend": 0.0})) + if counter_key.startswith("spend:agent:"): + agent_where: Final[LiteLLM_AgentsTableWhereUniqueInput] = { + "agent_id": counter_key[len("spend:agent:") :] + } + return await AgentsRepository(prisma_client).table.find_unique(where=agent_where) if counter_key.startswith("spend:project:"): return await ProjectRepository(prisma_client).table.find_unique( where={"project_id": counter_key[len("spend:project:") :]} diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 6a2ec120060..9cf1e0d6eed 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -2,6 +2,7 @@ import asyncio import traceback from collections.abc import Callable, Mapping, Sequence from datetime import datetime +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Protocol, cast import litellm @@ -387,6 +388,8 @@ class _ProxyDBLogger(CustomLogger): tags=tags, response_cost=response_cost, ), + billing_agent_id=metadata.get("billing_agent_id"), + billing_agent_counter_key=metadata.get("billing_agent_counter_key"), ) if not charged: return @@ -691,6 +694,9 @@ class _IncrementSpendCounters(Protocol): tags: list[str] | None = None, request_started_at: datetime | None = None, model_access_groups: Sequence[str] | None = None, + project_id: str | None = None, + billing_agent_id: str | None = None, + billing_agent_counter_key: str | None = None, ) -> None: ... @@ -712,6 +718,8 @@ async def _update_database_and_spend_counters( model_access_groups: Sequence[str] | None = None, project_id: str | None = None, update_cache_read_keys: Sequence[str] = (), + billing_agent_id: str | None = None, + billing_agent_counter_key: str | None = None, ) -> bool: """The reservation is reconciled before the spend is persisted, from its own read. One spend counter batch then spans the database write and the counter update, so the post-call counters are read with a single MGET after the @@ -734,6 +742,8 @@ async def _update_database_and_spend_counters( tags=request_tags, model_access_groups=model_access_groups, project_id=project_id, + billing_agent_id=billing_agent_id, + billing_agent_counter_key=billing_agent_counter_key, ) with spend_counter_batch_scope(spend_counter_cache.redis_cache, counter_keys=counter_keys): return await _update_database_and_spend_counters_in_batch( @@ -754,6 +764,8 @@ async def _update_database_and_spend_counters( model_access_groups=model_access_groups, project_id=project_id, update_cache_read_keys=update_cache_read_keys, + billing_agent_id=billing_agent_id, + billing_agent_counter_key=billing_agent_counter_key, ) @@ -775,6 +787,8 @@ async def _update_database_and_spend_counters_in_batch( model_access_groups: Sequence[str] | None, project_id: str | None, update_cache_read_keys: Sequence[str], + billing_agent_id: str | None, + billing_agent_counter_key: str | None, ) -> bool: from litellm.proxy.proxy_server import arm_update_cache_read @@ -823,6 +837,16 @@ async def _update_database_and_spend_counters_in_batch( request_started_at=start_time, model_access_groups=model_access_groups, project_id=project_id, + **( + MappingProxyType({"billing_agent_id": billing_agent_id}) + if billing_agent_id is not None + else MappingProxyType({}) + ), + **( + MappingProxyType({"billing_agent_counter_key": billing_agent_counter_key}) + if billing_agent_counter_key is not None + else MappingProxyType({}) + ), ) except Exception: if budget_reservation is not None: diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 4188e8ad58a..b4272d5e170 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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 @@ -1672,6 +1674,11 @@ class LiteLLMProxyRequestSetup: "actor_agent_id": user_api_key_dict.agent_id, "target_agent_id": user_api_key_dict.invoked_agent_id, "billing_agent_id": user_api_key_dict.agent_id or user_api_key_dict.invoked_agent_id, + "billing_agent_counter_key": ( + user_api_key_dict.billing_agent_policy.budget_counter_key + if user_api_key_dict.billing_agent_policy is not None + else None + ), "agent_execution_mode": managed_context.mode if managed_context else None, "verified_human_user_id": managed_context.user_id if managed_context else None, } diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index e16ea4a812e..866a747690e 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -28,10 +28,23 @@ from litellm.proxy.management_endpoints.common_utils import ( ) from litellm.proxy.utils import jsonify_object from litellm.repositories.budget_repository import BudgetRepository +from litellm.repositories.table_repositories import AgentsRepository router: Final = APIRouter() +async def _require_unlinked_agent_budget(budget_id: str, client: object) -> None: + from prisma.types import LiteLLM_AgentsTableWhereInput + + where: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id} + agent: Final = await AgentsRepository(client, use_writer=True).table.find_first(where=where) + if agent is not None: + raise HTTPException( + status_code=409, + detail=f"Manage this agent's budget through PATCH /v1/agents/{agent.agent_id}", + ) + + @router.post( "/budget/new", tags=["budget management"], @@ -194,6 +207,7 @@ async def update_budget( } ) + await _require_unlinked_agent_budget(budget_obj.budget_id, prisma_client) response: Final = await BudgetRepository(prisma_client).table.update( where={"budget_id": budget_obj.budget_id}, data=budget_obj_jsonified, @@ -357,6 +371,7 @@ async def delete_budget( detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, ) + await _require_unlinked_agent_budget(data.id, prisma_client) response: Final = await BudgetRepository(prisma_client).table.delete(where={"budget_id": data.id}) return response diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0e151199f41..a1e04f5c1c0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3022,6 +3022,8 @@ async def increment_spend_counters( request_started_at: datetime | None = None, model_access_groups: Sequence[str] | None = None, project_id: str | None = None, + billing_agent_id: str | None = None, + billing_agent_counter_key: str | None = None, ): """ Atomically increment spend counters for budget enforcement. @@ -3044,6 +3046,8 @@ async def increment_spend_counters( tags=tags, model_access_groups=model_access_groups, project_id=project_id, + billing_agent_id=billing_agent_id, + billing_agent_counter_key=billing_agent_counter_key, ), ): await _increment_spend_counters_batched( @@ -3058,6 +3062,8 @@ async def increment_spend_counters( request_started_at=request_started_at, model_access_groups=model_access_groups, project_id=project_id, + billing_agent_id=billing_agent_id, + billing_agent_counter_key=billing_agent_counter_key, ) @@ -3073,6 +3079,8 @@ async def _increment_spend_counters_batched( request_started_at: datetime | None, model_access_groups: Sequence[str] | None, project_id: str | None = None, + billing_agent_id: str | None = None, + billing_agent_counter_key: str | None = None, ): """Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET, and the reconcile adjustments go out in the same INCRBYFLOAT pipeline as the counter increments.""" @@ -3248,9 +3256,16 @@ async def _increment_spend_counters_batched( ), ) + async def _agent_scope(agent_id: str) -> tuple[PendingSpendIncrement, ...]: + counter_key: Final = billing_agent_counter_key or f"spend:agent:{agent_id}" + if counter_key in reserved_counter_keys: + return () + return (await _prepare_spend_counter_increment(counter_key=counter_key, source_cache_key=[], increment=cost),) + scope_coros: Final = tuple( coro for coro in ( + _agent_scope(billing_agent_id) if billing_agent_id is not None else None, _key_scope(token) if token is not None else None, _team_scope(team_id) if team_id is not None else None, _team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 42ac74cae33..a2dfc95750c 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -477,6 +477,20 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr data.pop("enable_tag_filtering", None) + if _is_a2a_agent_model(data.get("model")): + from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request + from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through + + registered_agent: Final = await get_agent_with_read_through(data["model"][4:]) + if registered_agent is not None: + agent_response: Final = await route_a2a_agent_request( + data, route_type, user_api_key_dict=user_api_key_dict, registered_agent=registered_agent + ) + if agent_response is not None: + return agent_response + if user_api_key_dict is not None and user_api_key_dict.invoked_agent_policy is not None: + raise HTTPException(503, "Agent dispatch does not match its admission") + team_id: Final = get_team_id_from_data(data) router_model_names: Final = llm_router.model_names if llm_router is not None else [] is_proxy_admin_without_team: Final = team_id is None and _is_proxy_admin_request(data) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index adfe2a0eee7..fc7325ddac3 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -36,6 +36,7 @@ model LiteLLM_BudgetTable { model_access_groups LiteLLM_ModelAccessGroupBudgetTable[] // multiple model access groups can have the same budget team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization + agents LiteLLM_AgentsTable[] } // Models on proxy @@ -83,6 +84,10 @@ model LiteLLM_AgentsTable { execution_mode String @default("autonomous") identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] + budget_id String? @unique + spend_window DateTime? + lifetime_budget_spend Float @default(0.0) + litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id]) tpm_limit Int? rpm_limit Int? session_tpm_limit Int? diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index c094e91c6c0..62a6bdbe0fd 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -67,6 +67,7 @@ _COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = { "Model access group": Litellm_EntityType.MODEL_ACCESS_GROUP.value, "Organization": Litellm_EntityType.ORGANIZATION.value, "Project": Litellm_EntityType.PROJECT.value, + "Agent": Litellm_EntityType.AGENT.value, } @@ -265,7 +266,8 @@ async def reserve_budget_for_request( return None if _is_unbilled_route(route): return None - if get_model_from_request(request_body, route, llm_router=llm_router) is None: + invocation_cost: Final = valid_token.agent_invocation_cost + if invocation_cost is None and get_model_from_request(request_body, route, llm_router=llm_router) is None: return None counters: Final = await _get_budget_counters( @@ -283,18 +285,26 @@ async def reserve_budget_for_request( if not counters: return None - input_token_counts: Final = await count_request_input_tokens( - request_body=request_body, - route=route, - llm_router=llm_router, - raw_body=raw_body, + input_token_counts: Final = ( + await count_request_input_tokens( + request_body=request_body, + route=route, + llm_router=llm_router, + raw_body=raw_body, + ) + if invocation_cost is None + else MappingProxyType({}) ) - reservation_cost = estimate_request_max_cost( - request_body=request_body, - route=route, - llm_router=llm_router, - input_token_counts=input_token_counts, + reservation_cost = ( + invocation_cost + if invocation_cost is not None + else estimate_request_max_cost( + request_body=request_body, + route=route, + llm_router=llm_router, + input_token_counts=input_token_counts, + ) ) # estimate_request_max_cost still returns None when the model is unknown # to the cost map (no token-priced cost fields, e.g. image/audio routes). @@ -316,7 +326,7 @@ async def reserve_budget_for_request( reservation_cost=reservation_cost, fail_closed_budget_enforcement=fail_closed_budget_enforcement, ) - except Exception: + except (asyncio.CancelledError, Exception): await _release_applied_entries_best_effort( entries=applied_entries, default_reserved_cost=reservation_cost, @@ -326,11 +336,15 @@ async def reserve_budget_for_request( if not applied_entries: return None - input_cost: Final = estimate_request_input_cost( - request_body=request_body, - route=route, - llm_router=llm_router, - input_token_counts=input_token_counts, + input_cost: Final = ( + 0.0 + if invocation_cost is not None + else estimate_request_input_cost( + request_body=request_body, + route=route, + llm_router=llm_router, + input_token_counts=input_token_counts, + ) ) budget_reservation: Final = { "reserved_cost": reservation_cost, @@ -490,7 +504,23 @@ async def _get_budget_counters( end_user_object: object = None, apply_user_budget_to_team_keys: bool = False, ) -> list[_BudgetCounter]: - counters: Final[list[_BudgetCounter]] = [] + agent: Final = valid_token.billing_agent_policy + counters: Final[list[_BudgetCounter]] = ( + [ + _BudgetCounter( + counter_key=agent.budget_counter_key, + source_cache_key=None, + max_budget=agent.litellm_budget_table.max_budget, + fallback_spend=agent.budget_spend, + entity_type="Agent", + entity_id=agent.agent_id, + ) + ] + if agent is not None + and agent.litellm_budget_table is not None + and agent.litellm_budget_table.max_budget is not None + else [] + ) if valid_token.token is not None: if valid_token.max_budget is not None and valid_token.max_budget > 0: diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py index ae24331c236..5ae30a5c027 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -204,7 +204,16 @@ def _iter_entity_counter_keys( def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]: - return frozenset( + billing_agent: Final = token.billing_agent_policy + charged_agent_id: Final = billing_agent.agent_id if billing_agent is not None else token.agent_id + agent_keys: Final = ( + frozenset( + (billing_agent.budget_counter_key if billing_agent is not None else f"spend:agent:{charged_agent_id}",) + ) + if charged_agent_id is not None + else frozenset() + ) + return agent_keys | frozenset( _iter_entity_counter_keys( token=token.token, team_id=token.team_id, @@ -225,6 +234,8 @@ def post_call_counter_keys( tags: Sequence[object] | None, model_access_groups: Sequence[object] | None, project_id: str | None = None, + billing_agent_id: str | None = None, + billing_agent_counter_key: str | None = None, ) -> frozenset[str]: """Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read.""" entity_keys: Final = frozenset( @@ -243,7 +254,10 @@ def post_call_counter_keys( for group in model_access_groups or () if group and isinstance(group, str) ) - return entity_keys | tag_keys | group_keys + agent_key: Final = billing_agent_counter_key or ( + f"spend:agent:{billing_agent_id}" if billing_agent_id is not None else None + ) + return entity_keys | tag_keys | group_keys | (frozenset((agent_key,)) if agent_key else frozenset()) def bind_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> None: diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index c42301a9316..d9b25069747 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -171,4 +171,7 @@ class PrismaBatch(Protocol): @property def litellm_projecttable(self) -> BatchTable: ... + @property + def litellm_agentstable(self) -> BatchTable: ... + async def commit(self) -> None: ... diff --git a/litellm/repositories/unit_of_work.py b/litellm/repositories/unit_of_work.py index c09e5eb75d4..ced458d2102 100644 --- a/litellm/repositories/unit_of_work.py +++ b/litellm/repositories/unit_of_work.py @@ -19,10 +19,13 @@ from collections.abc import AsyncGenerator, Callable, Mapping from contextlib import asynccontextmanager from dataclasses import dataclass from datetime import datetime -from typing import Final +from typing import TYPE_CHECKING, Final from litellm.repositories.prisma_protocols import BatchTable, PrismaBatch +if TYPE_CHECKING: + from prisma.types import LiteLLM_AgentsTableUpdateManyMutationInput, LiteLLM_AgentsTableWhereInput + def _spend_reset_data(budget_reset_at: datetime | None, spend_decrement: float) -> Mapping[str, object]: spend: Final[object] = {"decrement": spend_decrement} # mutable-ok: prisma update payload must be a dict @@ -78,6 +81,26 @@ class LinkedSpendResetWrites: ) +@dataclass(frozen=True, slots=True) +class AgentSpendResetWrites: + table: BatchTable + + def queue_window_reset(self, budget_id: str, window: datetime, rollover_cap: float | None) -> None: + zero: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = {"spend": 0.0, "spend_window": window} + if rollover_cap is None: + all_agents: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id} + self.table.update_many(where=all_agents, data=zero) + return + below: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id, "spend": {"lte": rollover_cap}} + above: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id, "spend": {"gt": rollover_cap}} + remainder: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = { + "spend": {"decrement": rollover_cap}, + "spend_window": window, + } + self.table.update_many(where=below, data=zero) + self.table.update_many(where=above, data=remainder) + + @dataclass(frozen=True, slots=True) class BudgetWindowWrites: table: BatchTable @@ -110,6 +133,7 @@ class BudgetCascadeUnitOfWork: tags: LinkedSpendResetWrites model_access_groups: LinkedSpendResetWrites projects: LinkedSpendResetWrites + agents: AgentSpendResetWrites endusers: LinkedSpendResetWrites budgets: BudgetWindowWrites @@ -137,6 +161,7 @@ async def budget_cascade_unit_of_work( tags=LinkedSpendResetWrites(table=batch.litellm_tagtable), model_access_groups=LinkedSpendResetWrites(table=batch.litellm_modelaccessgroupbudgettable), projects=LinkedSpendResetWrites(table=batch.litellm_projecttable), + agents=AgentSpendResetWrites(table=batch.litellm_agentstable), endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable), budgets=BudgetWindowWrites(table=batch.litellm_budgettable), ) diff --git a/litellm/types/agents.py b/litellm/types/agents.py index 94adb9f7c4a..d59b6fa6ce9 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -1,5 +1,5 @@ from collections.abc import Mapping, Sequence -from datetime import datetime +from datetime import datetime, timezone from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, TypeAlias from urllib.parse import urlsplit @@ -8,6 +8,8 @@ from typing_extensions import ReadOnly, Required, TypedDict from litellm.types.llms.base import LiteLLMPydanticObjectBase from litellm.types.proxy.agent_identity import ( + AgentBudgetConfig, + AgentBudgetState, AgentExecutionMode, AgentIdentityBinding, EntraIdentityConfig, @@ -15,6 +17,7 @@ from litellm.types.proxy.agent_identity import ( if TYPE_CHECKING: from a2a.types import SendMessageResponse + from prisma.types import LiteLLM_AgentsTableWhereInput # AgentProvider @@ -253,6 +256,7 @@ class AgentKillSwitchResult(BaseModel): class AgentConfig(TypedDict, total=False): + budget: ReadOnly[AgentBudgetConfig | None] identity: ReadOnly[EntraIdentityConfig | None] enabled: ReadOnly[bool] execution_mode: ReadOnly[AgentExecutionMode] @@ -271,6 +275,7 @@ class AgentConfig(TypedDict, total=False): class PatchAgentRequest(TypedDict, total=False): + budget: ReadOnly[AgentBudgetConfig | None] identity: ReadOnly[EntraIdentityConfig | None] enabled: ReadOnly[bool] execution_mode: ReadOnly[AgentExecutionMode] @@ -311,7 +316,41 @@ class AgentKeySummary(BaseModel): key_name: str | None = None +def agent_budget_counter_key(agent_id: str, reset_at: datetime | None, budget_id: str | None = None) -> str: + if reset_at is None and budget_id is not None: + return f"spend:agent_lifetime:{budget_id}:{agent_id}" + if reset_at is None: + return f"spend:agent:{agent_id}" + aware: Final = reset_at if reset_at.tzinfo is not None else reset_at.replace(tzinfo=timezone.utc) + window: Final = aware.astimezone(timezone.utc).strftime("%Y%m%dT%H%M%S.%fZ") + return f"spend:agent_window:{window}:{agent_id}" + + +def agent_spend_filter(counter_key: str) -> "LiteLLM_AgentsTableWhereInput": + if counter_key.startswith("spend:agent_lifetime:"): + _, _, budget_id, agent_id = counter_key.split(":", 3) + lifetime: Final[LiteLLM_AgentsTableWhereInput] = { + "agent_id": agent_id, + "budget_id": budget_id, + "spend_window": None, + } + return lifetime + if counter_key.startswith("spend:agent_window:"): + _, _, raw_window, agent_id = counter_key.split(":", 3) + window: Final = datetime.strptime(raw_window, "%Y%m%dT%H%M%S.%fZ").replace(tzinfo=timezone.utc) + windowed: Final[LiteLLM_AgentsTableWhereInput] = {"agent_id": agent_id, "spend_window": window} + return windowed + cumulative: Final[LiteLLM_AgentsTableWhereInput] = { + "agent_id": counter_key.removeprefix("spend:agent:"), + "spend_window": None, + } + return cumulative + + class AgentResponse(BaseModel): + budget_id: str | None = None + lifetime_budget_spend: float = 0.0 + litellm_budget_table: AgentBudgetState | None = None identity: AgentIdentityBinding | None = None identity_managed: bool = False enabled: bool = True @@ -338,6 +377,20 @@ class AgentResponse(BaseModel): created_by: str | None = None updated_by: str | None = None + @property + def budget_counter_key(self) -> str: + return agent_budget_counter_key( + self.agent_id, + self.litellm_budget_table.budget_reset_at if self.litellm_budget_table else None, + self.litellm_budget_table.budget_id if self.litellm_budget_table else None, + ) + + @property + def budget_spend(self) -> float: + if self.litellm_budget_table is not None and self.litellm_budget_table.budget_duration is None: + return self.lifetime_budget_spend + return self.spend or 0.0 + class ListAgentsResponse(BaseModel): agents: list[AgentResponse] diff --git a/litellm/types/proxy/agent_identity.py b/litellm/types/proxy/agent_identity.py index a7fe0be37e1..2dc43976e0f 100644 --- a/litellm/types/proxy/agent_identity.py +++ b/litellm/types/proxy/agent_identity.py @@ -46,6 +46,22 @@ class AgentIdentityBinding(BaseModel): last_authenticated_at: datetime | None = None +class AgentBudgetConfig(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + max_budget: float = Field(ge=0, allow_inf_nan=False) + budget_duration: str | None = None + + +class AgentBudgetState(BaseModel): + model_config = ConfigDict(frozen=True) + + budget_id: str + max_budget: float | None = None + budget_duration: str | None = None + budget_reset_at: datetime | None = None + + class AgentSubject(BaseModel): model_config = ConfigDict(frozen=True) diff --git a/schema.prisma b/schema.prisma index adfe2a0eee7..fc7325ddac3 100644 --- a/schema.prisma +++ b/schema.prisma @@ -36,6 +36,7 @@ model LiteLLM_BudgetTable { model_access_groups LiteLLM_ModelAccessGroupBudgetTable[] // multiple model access groups can have the same budget team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization + agents LiteLLM_AgentsTable[] } // Models on proxy @@ -83,6 +84,10 @@ model LiteLLM_AgentsTable { execution_mode String @default("autonomous") identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] + budget_id String? @unique + spend_window DateTime? + lifetime_budget_spend Float @default(0.0) + litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id]) tpm_limit Int? rpm_limit Int? session_tpm_limit Int? 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 eee985f0aca..bc84e9ddd03 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 @@ -98,11 +98,45 @@ def test_caller_cannot_construct_trusted_subject_or_policy() -> None: assert auth.agent_invocation_cost is None +@pytest.mark.asyncio +async def test_agent_budget_accumulates_across_credentials_and_denies_the_next_admission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints.auth.managed_authorization import check_agent_budget + + counters: Final = DualCache() + counters.set_cache("spend:agent_lifetime:budget:agent", 0.0) + counters.set_cache("spend:key:first", 0.0) + counters.set_cache("spend:key:second", 0.0) + monkeypatch.setattr(proxy_server, "spend_counter_cache", counters) + policy: Final = agent(spend=0, litellm_budget_table={"budget_id": "budget", "max_budget": 0.5}) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.billing_agent_policy = policy + await check_agent_budget(auth) + await proxy_server.increment_spend_counters( + token="first", team_id=None, user_id=None, response_cost=0.3, billing_agent_id="agent", billing_agent_counter_key=policy.budget_counter_key + ) + await check_agent_budget(auth) + await proxy_server.increment_spend_counters( + token="second", team_id=None, user_id=None, response_cost=0.3, billing_agent_id="agent", billing_agent_counter_key=policy.budget_counter_key + ) + with pytest.raises(litellm.BudgetExceededError): + await check_agent_budget(auth) + assert await counters.async_get_cache("spend:agent_lifetime:budget:agent") == pytest.approx(0.6) + assert await counters.async_get_cache("spend:key:first") == pytest.approx(0.3) + assert await counters.async_get_cache("spend:key:second") == pytest.approx(0.3) + + @pytest.mark.asyncio @pytest.mark.parametrize("autonomous", (True, False)) +@pytest.mark.parametrize("billable", (True, False)) async def test_invocation_prepares_target_fee_for_the_correct_agent( monkeypatch: pytest.MonkeyPatch, autonomous: bool, + billable: bool, ) -> None: from unittest.mock import AsyncMock, MagicMock @@ -129,11 +163,14 @@ async def test_invocation_prepares_target_fee_for_the_correct_agent( caller: Final = agent(agent_id="caller", object_permission=permission.model_dump()) auth.managed_agent_policy = caller auth.billing_agent_policy = caller - await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database)) - assert auth.agent_invocation_cost == pytest.approx(0.25) + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable) + assert auth.agent_invocation_cost == pytest.approx(0.25 if billable else 0.0) assert auth.invoked_agent_id == "agent" - assert auth.billing_agent_policy is not None - assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent") + if autonomous or billable: + assert auth.billing_agent_policy is not None + assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent") + else: + assert auth.billing_agent_policy is None @pytest.mark.asyncio @@ -167,6 +204,8 @@ async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None: ("/a2a/expensive", {"model": "a2a/cheap"}, "expensive"), ("/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"), ("/v1/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"), + ("/a2a/agent", {"model": "a2a/nonexistent"}, "agent"), + ("/v1/a2a/agent/message/send", {"model": "a2a/other"}, "agent"), ("/v1/a2a/agent/", {}, "agent"), ("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"), ("/v1/chat/completions", {"model": "a2a/"}, None), @@ -178,6 +217,23 @@ def test_invocation_routes_resolve_the_same_target(route: str, body: dict[str, o assert invocation_target(route, body) == expected +@pytest.mark.asyncio +@pytest.mark.parametrize("managed", [True, False]) +async def test_aggregate_budget_applies_to_entra_tokens_and_unbound_agent_keys(managed: bool) -> None: + policy: Final = agent( + identity_managed=managed, identity=BINDING if managed else None, + litellm_budget_table={"budget_id": "budget", "max_budget": 0.5}, + ) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + auth: Final = UserAPIKeyAuth(agent_id="agent") + if managed: + auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous") + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.billing_agent_policy == policy + assert auth.managed_agent_policy == (policy if managed else None) + + @pytest.mark.asyncio async def test_agent_admission_database_outage_fails_closed() -> None: database: Final = MagicMock() @@ -478,7 +534,8 @@ async def test_unmanaged_agent_invocation_retains_legacy_behavior(monkeypatch: p await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database)) assert auth.managed_agent_policy is None assert auth.billing_agent_policy is None - assert auth.invoked_agent_id is None + assert auth.invoked_agent_id == "agent" + assert auth.invoked_agent_policy is not None @pytest.mark.asyncio @@ -508,6 +565,95 @@ async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_ assert auth.requires_fresh_policy is True +@pytest.mark.asyncio +async def test_new_budget_window_isolated_from_inflight_previous_window_charge(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints.auth.managed_authorization import check_agent_budget + + cache: Final = DualCache() + old_key: Final = "spend:agent_window:20260101T000000.000000Z:agent" + new_key: Final = "spend:agent_window:20260102T000000.000000Z:agent" + cache.set_cache("spend:agent:agent", 10.0) + cache.set_cache(old_key, 10.0) + cache.set_cache(new_key, 0.0) + monkeypatch.setattr(proxy_server, "spend_counter_cache", cache) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.billing_agent_policy = agent(spend=0, litellm_budget_table={ + "budget_id": "budget", "max_budget": 1.0, "budget_reset_at": "2026-01-02T00:00:00Z", + }) + await check_agent_budget(auth) + await proxy_server.increment_spend_counters( + token=None, team_id=None, user_id=None, response_cost=0.5, + billing_agent_id="agent", billing_agent_counter_key=old_key, + ) + await check_agent_budget(auth) + await proxy_server.increment_spend_counters( + token=None, team_id=None, user_id=None, response_cost=1.1, + billing_agent_id="agent", billing_agent_counter_key=new_key, + ) + with pytest.raises(litellm.BudgetExceededError): + await check_agent_budget(auth) + assert await cache.async_get_cache(old_key) == 10.5 + assert await cache.async_get_cache(new_key) == 1.1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("autonomous", [False, True]) +@pytest.mark.parametrize( + "pricing,billable,bounded,rejected,fee", + [ + ({"input_cost_per_token": 0.01}, True, True, True, None), + ({"output_cost_per_token": 0.01}, True, True, True, None), + ({"input_cost_per_token": 0.01}, False, True, False, 0.0), + ({"input_cost_per_token": 0.01}, True, False, False, 0.0), + ({"input_cost_per_token": 0.0, "output_cost_per_token": 0.0}, True, True, False, 0.0), + ({"cost_per_query": 0.25, "output_cost_per_token": 0.01}, True, True, False, 0.25), + ({"cost_per_query": 0.0, "output_cost_per_token": 0.01}, True, True, False, 0.0), + ], +) +async def test_budgeted_invocation_requires_a_bounded_price( + monkeypatch: pytest.MonkeyPatch, + autonomous: bool, + pricing: dict[str, float], + billable: bool, + bounded: bool, + rejected: bool, + fee: float | None, +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + budget: Final = {"budget_id": "budget", "max_budget": 1.0} if bounded else None + target_budget: Final = {"budget_id": "target-budget", "max_budget": 1.0} if bounded != autonomous else None + target: Final = agent(litellm_params=pricing, litellm_budget_table=target_budget) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr(proxy_server, "prisma_client", database) + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="invoke-grant", agents=["agent"]) + auth: Final = UserAPIKeyAuth(agent_id="caller" if autonomous else None, object_permission=permission) + if autonomous: + caller: Final = agent( + agent_id="caller", object_permission=permission.model_dump(), litellm_budget_table=budget + ) + auth.managed_agent_policy = caller + auth.billing_agent_policy = caller + if rejected: + with pytest.raises(HTTPException) as exc: + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable) + assert exc.value.status_code == 503 + assert "cost_per_query" in str(exc.value.detail) + else: + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable) + assert auth.agent_invocation_cost == fee + + async def test_jwt_delegation_verification_is_consumed_once_and_cannot_be_supplied_by_a_caller( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -577,3 +723,60 @@ 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 + + +@pytest.mark.parametrize("route", ("/realtime", "/v1/chat/completions", "/v1/files")) +def test_ordinary_requests_without_models_keep_existing_validation(route: str) -> None: + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request + + assert managed_inference_request(route, {}, {}, None, require_model=False) == {} + + +@pytest.mark.parametrize("state", ( + {"identity_managed": True}, + {"identity": BINDING}, + {"litellm_budget_table": {"budget_id": "budget", "max_budget": 0.5}}, + {"litellm_params": {"cost_per_query": 0.25}}, +)) +@pytest.mark.parametrize("has_auth", (False, True)) +def test_protected_agent_dispatch_requires_admission(state: dict[str, object], has_auth: bool) -> None: + from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy + + registered: Final = agent(**{"identity": None, "identity_managed": False, **state}) + with pytest.raises(HTTPException, match="admission") as exc: + agent_invocation_policy(UserAPIKeyAuth() if has_auth else None, registered) + assert exc.value.status_code == 503 + + +def test_paid_agent_dispatch_requires_the_captured_fee() -> None: + from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy + + policy: Final = agent(identity=None, identity_managed=False, litellm_params={"cost_per_query": 0.25}) + auth: Final = UserAPIKeyAuth() + auth.invoked_agent_id = policy.agent_id + auth.invoked_agent_policy = policy + with pytest.raises(HTTPException, match="admission") as exc: + agent_invocation_policy(auth, policy) + assert exc.value.status_code == 503 diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index a5d0d0a3ecc..1e04a6df3a5 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -16,7 +16,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from litellm.proxy._types import UserAPIKeyAuth -from litellm.types.agents import AgentCaller +from litellm.types.agents import AgentCaller, AgentResponse AddLiteLLMData = Callable[..., Awaitable[dict[str, object]]] @@ -25,6 +25,9 @@ AddLiteLLMData = Callable[..., Awaitable[dict[str, object]]] class CapturedAgentCall: request_id: object agent_extra_headers: dict[str, str] | None + cost_per_query: object + api_base: object + pricing: Mapping[str, object] @pytest.mark.asyncio @@ -58,7 +61,7 @@ async def test_invoke_agent_a2a_adds_litellm_data(): } # Mock agent - mock_agent = MagicMock() + mock_agent = _make_agent_mock() mock_agent.agent_id = "test-agent" mock_agent.agent_card_params = { "url": "http://backend-agent:10001", @@ -211,7 +214,7 @@ async def test_invoke_agent_a2a_handles_none_agent_card_params(): """ from litellm.proxy._types import UserAPIKeyAuth - mock_agent = MagicMock() + mock_agent = _make_agent_mock() mock_agent.agent_card_params = None mock_agent.litellm_params = None @@ -295,7 +298,7 @@ async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge(): resp.model_dump.return_value = {"jsonrpc": "2.0", "id": "test-id", "result": {}} return resp - mock_agent = MagicMock() + mock_agent = _make_agent_mock() mock_agent.agent_id = "lf-agent" mock_agent.agent_name = "lf-agent" # No URL: the bridge derives the endpoint from the LangFlow agent config. @@ -376,6 +379,9 @@ def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock: agent.litellm_params = {} agent.static_headers = None agent.extra_headers = None + agent.identity_managed = False + agent.identity = None + agent.litellm_budget_table = None return agent @@ -394,7 +400,7 @@ def _make_request_mock(method: str, params: Mapping[str, object], request_id: ob def _base_patches( - agent: MagicMock, add_litellm_data: AddLiteLLMData | None = None + agent: MagicMock | AgentResponse, add_litellm_data: AddLiteLLMData | None = None ) -> list[AbstractContextManager[object]]: return [ patch( @@ -437,7 +443,7 @@ async def _invoke_message_method( mock_request: MagicMock, user_api_key_dict: UserAPIKeyAuth, add_litellm_data: AddLiteLLMData | None = None, - agent: MagicMock | None = None, + agent: MagicMock | AgentResponse | None = None, ) -> CapturedAgentCall: from fastapi.responses import JSONResponse @@ -490,7 +496,13 @@ async def _invoke_message_method( kwargs: Final = downstream.call_args.kwargs request_id: Final = kwargs["request"].__dict__["id"] if is_send else kwargs["request_id"] - return CapturedAgentCall(request_id=request_id, agent_extra_headers=kwargs.get("agent_extra_headers")) + return CapturedAgentCall( + request_id=request_id, + agent_extra_headers=kwargs.get("agent_extra_headers"), + cost_per_query=kwargs["litellm_params"].get("cost_per_query"), + api_base=kwargs["api_base"], + pricing=kwargs["litellm_params"], + ) @pytest.mark.asyncio @@ -2689,3 +2701,95 @@ def test_forwarding_headers_minted_bearer_replaces_a_forwarded_authorization_of_ ) assert merged == {"X-Custom": "kept", "Authorization": "Bearer minted-token"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ("message/send", "message/stream")) +@pytest.mark.parametrize("fee", (None, 0.0, 0.25)) +async def test_native_dispatch_keeps_the_admitted_price_and_destination( + monkeypatch: pytest.MonkeyPatch, + method: str, + fee: float | None, +) -> None: + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + admitted: Final = AgentResponse( + agent_id="test-agent", + agent_name="test-agent", + agent_card_params={"url": "https://admitted.test/"}, + litellm_params={"cost_per_query": fee} if fee is not None else {}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(admitted) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + await prepare_agent_invocation(auth, admitted.agent_id, None) + changed: Final = admitted.model_copy( + update={ + "litellm_params": {"cost_per_query": 0.75}, + "agent_card_params": {"url": "https://changed.test/"}, + } + ) + captured: Final = await _invoke_message_method( + method, + _make_request_mock(method, _HELLO_MESSAGE_PARAMS), + auth, + agent=changed, + ) + assert captured.cost_per_query == fee + assert captured.api_base == "https://admitted.test/" + + +@pytest.mark.asyncio +async def test_native_dispatch_returns_not_found_for_a_removed_agent(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry()) + monkeypatch.setattr(proxy_server, "prisma_client", None) + response: Final = await invoke_agent_a2a( + agent_id="removed", request=_make_request_mock("message/send", _HELLO_MESSAGE_PARAMS), + fastapi_response=MagicMock(), user_api_key_dict=UserAPIKeyAuth(), + ) + assert response.status_code == 404 + assert json.loads(response.body)["error"]["message"] == "Agent 'removed' not found" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ("message/send", "message/stream")) +@pytest.mark.parametrize("fixed_fee", (None, 0.0, 0.25)) +async def test_unbudgeted_managed_agent_keeps_token_pricing_without_a_fixed_fee( + monkeypatch: pytest.MonkeyPatch, method: str, fixed_fee: float | None, +) -> None: + from litellm import Usage + from litellm.a2a_protocol.cost_calculator import A2ACostCalculator + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + + policy: Final = AgentResponse( + agent_id="test-agent", agent_name="test-agent", identity_managed=True, + agent_card_params={"url": "https://agent.test/"}, + litellm_params={"input_cost_per_token": 0.02, "output_cost_per_token": 0.03, + **({"cost_per_query": fixed_fee} if fixed_fee is not None else {})}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(policy) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + with ExitStack() as stack: + for context in _base_patches(policy): + stack.enter_context(context) + await prepare_agent_invocation(auth, policy.agent_id, AgentIdentityStore.from_client(database)) + captured: Final = await _invoke_message_method( + method, _make_request_mock(method, _HELLO_MESSAGE_PARAMS), auth, agent=policy, + ) + logging: Final = MagicMock(model_call_details={ + "litellm_params": captured.pricing, + "usage": Usage(prompt_tokens=10, completion_tokens=2, total_tokens=12), + }) + assert A2ACostCalculator.calculate_a2a_cost(logging) == pytest.approx(0.26 if fixed_fee is None else fixed_fee) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py index e894f4ad69a..4d6077b3472 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py @@ -14,23 +14,26 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from litellm.types.agents import AgentResponse + # --------------------------------------------------------------------------- -# Helper: build a minimal mock agent +# Helper: build a minimal agent # --------------------------------------------------------------------------- def _make_mock_agent( - static_headers=None, - extra_headers=None, - url="http://backend-agent:10001", -): - mock_agent = MagicMock() - mock_agent.agent_id = "agent-123" - mock_agent.agent_card_params = {"url": url, "name": "Test Agent"} - mock_agent.litellm_params = {} - mock_agent.static_headers = static_headers or {} - mock_agent.extra_headers = extra_headers or [] - return mock_agent + static_headers: dict[str, str] | None = None, + extra_headers: list[str] | None = None, + url: str = "http://backend-agent:10001", +) -> AgentResponse: + return AgentResponse( + agent_id="agent-123", + agent_name="Test Agent", + agent_card_params={"url": url, "name": "Test Agent"}, + litellm_params={}, + static_headers=static_headers or {}, + extra_headers=extra_headers or [], + ) def _make_mock_request(extra_headers=None, method="message/send"): diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py index 7663f1d30e6..33b9a5d4b9d 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -1231,6 +1231,7 @@ def _stored_agent_row(values: Mapping[str, object] | SimpleNamespace) -> LiteLLM "created_by": "admin", "updated_by": "admin", "spend": 0, + "lifetime_budget_spend": 0, "identity_managed": False, "enabled": True, "execution_mode": "autonomous", @@ -1431,6 +1432,7 @@ async def test_agent_listing_preserves_stored_identity_bindings(bound: bool) -> enabled=True, execution_mode="autonomous", spend=0.0, + lifetime_budget_spend=0.0, agent_access_groups=[], access_group_ids=[], extra_headers=[], @@ -1451,7 +1453,7 @@ async def test_agent_listing_preserves_stored_identity_bindings(bound: bool) -> assert response.identity is None client.db.litellm_agentstable.find_many.assert_awaited_once_with( order={"created_at": "desc"}, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index 2cf81892db7..0bf9afa832d 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -1648,3 +1648,31 @@ def test_invalid_identity_and_untrusted_tenant_cannot_be_registered( with pytest.raises(HTTPException, match=message) as failure: agent_endpoints._validate_managed_identity_request(request) assert failure.value.status_code == 400 + + +@pytest.mark.parametrize("cached,path", [(False, "/v1/agents/agent-123"), (True, "/v1/agents/agent-123"), (True, "/v1/agents")]) +def test_agent_budget_readback_refreshes_consumption_and_limit(monkeypatch: pytest.MonkeyPatch, cached: bool, path: str) -> None: + from prisma.models import LiteLLM_BudgetTable + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry + + row: Final = LiteLLM_AgentsTable.model_construct( + agent_id="agent-123", agent_name="Agent", agent_card_params={}, spend=12.5, + lifetime_budget_spend=0.75, budget_id="budget", litellm_params=None, + litellm_budget_table=LiteLLM_BudgetTable.model_construct(budget_id="budget", max_budget=2.0), + ) + registry: Final = AgentRegistry() + if cached: + registry.register_agent(_sample_agent_response()) + table: Final = SimpleNamespace(find_unique=AsyncMock(return_value=row), find_many=AsyncMock(return_value=[row])) + database: Final = SimpleNamespace(litellm_agentstable=table, litellm_verificationtoken=SimpleNamespace(find_many=AsyncMock(return_value=[]))) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database, writer_db=database)) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + response: Final = client.get(path) + assert response.status_code == 200, response.text + payload: Final = response.json()[0] if path == "/v1/agents" else response.json() + assert payload["spend"] == 12.5 + assert payload["lifetime_budget_spend"] == 0.75 + assert payload["litellm_budget_table"]["max_budget"] == 2.0 diff --git a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py index 17f3cdb52f5..55b1a0d366e 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py @@ -218,7 +218,56 @@ def test_invalid_identity_configuration_returns_a_public_validation_failure(inco result: Final = managed_write_fields(incoming, None, "admin") assert isinstance(result, AgentIdentityFailure) assert result.code == "identity_denied" - assert result.message.startswith("Invalid agent identity configuration:") + assert result.message.startswith("Invalid agent identity or budget configuration:") + + +@pytest.mark.parametrize("budget", [{"max_budget": -1}, {"max_budget": float("inf")}, {"max_budget": 1, "budget_duration": "0d"}, {"max_budget": 1, "budget_duration": "0s"}, {"max_budget": 1, "budget_duration": "-1d"}, {"max_budget": 1, "budget_duration": ""}]) +def test_invalid_budget_changes_are_rejected(budget: dict[str, object]) -> None: + result: Final = managed_write_fields({"budget": budget}, managed_agent(), "admin") + assert isinstance(result, AgentIdentityFailure) + assert "Invalid agent identity or budget configuration" in result.message + + +def test_budget_updates_preserve_current_window_until_duration_changes() -> None: + from datetime import datetime, timezone + + from litellm.types.proxy.agent_identity import AgentBudgetState + + reset: Final = datetime(2027, 1, 1, tzinfo=timezone.utc) + agent: Final = managed_agent().model_copy( + update={ + "budget_id": "budget", + "litellm_budget_table": AgentBudgetState( + budget_id="budget", max_budget=1, budget_duration="1d", budget_reset_at=reset + ), + } + ) + same: Final = managed_write_fields({"budget": {"max_budget": 2, "budget_duration": "1d"}}, agent, "admin") + assert not isinstance(same, AgentIdentityFailure) + assert same["litellm_budget_table"]["update"]["budget_reset_at"] == reset + assert same["litellm_budget_table"]["update"]["max_budget"] == 2 + assert same["spend_window"] == reset + assert "spend" not in same + changed: Final = managed_write_fields({"budget": {"max_budget": 2, "budget_duration": "1h"}}, agent, "admin") + assert not isinstance(changed, AgentIdentityFailure) + assert changed["litellm_budget_table"]["update"]["budget_reset_at"] != reset + assert changed["spend_window"] == changed["litellm_budget_table"]["update"]["budget_reset_at"] + assert changed["spend"] == 0.0 + removed: Final = managed_write_fields({"budget": None}, agent, "admin") + assert removed == {"litellm_budget_table": {"disconnect": True}, "spend_window": None} + assert managed_write_fields({"budget": None}, managed_agent(), "admin") == {} + + +def test_new_agent_budget_is_created_with_administrator_attribution() -> None: + result: Final = managed_write_fields({"budget": {"max_budget": 0}}, None, "admin") + assert not isinstance(result, AgentIdentityFailure) + assert result["litellm_budget_table"]["create"] == { + "max_budget": 0, + "budget_duration": None, + "budget_reset_at": None, + "created_by": "admin", + "updated_by": "admin", + } @pytest.mark.parametrize("roles", ["Agent.Invoke", [42], None]) @@ -255,3 +304,53 @@ def test_empty_requirements_do_not_make_a_scope_less_human_token_valid(scope: ob binding: Final = BINDING.model_copy(update={"required_scopes": ()}) result: Final = classify_agent_subject(binding, claims(oid=HUMAN, scp=scope), "both") assert isinstance(result, AgentIdentityFailure) + + +def test_budget_write_stamps_the_same_window_on_the_agent_row() -> None: + result: Final = managed_write_fields({"budget": {"max_budget": 1.0, "budget_duration": "1d"}}, None, "admin") + assert not isinstance(result, AgentIdentityFailure) + assert result["spend_window"] == result["litellm_budget_table"]["create"]["budget_reset_at"] + assert result["spend"] == 0.0 + + +def test_new_lifetime_budget_starts_unused_without_erasing_historical_spend() -> None: + existing: Final = managed_agent().model_copy(update={"spend": 12.5}) + result: Final = managed_write_fields({"budget": {"max_budget": 1.0}}, existing, "admin") + assert not isinstance(result, AgentIdentityFailure) + assert "spend" not in result + assert result["lifetime_budget_spend"] == 0.0 + assert result["litellm_budget_table"]["create"]["max_budget"] == 1.0 + assert existing.spend == 12.5 + + +def test_editing_lifetime_budget_preserves_consumption() -> None: + from litellm.types.proxy.agent_identity import AgentBudgetState + + existing: Final = managed_agent().model_copy(update={ + "spend": 12.5, "lifetime_budget_spend": 0.75, "budget_id": "budget", + "litellm_budget_table": AgentBudgetState(budget_id="budget", max_budget=1.0), + }) + result: Final = managed_write_fields({"budget": {"max_budget": 2.0}}, existing, "admin") + assert not isinstance(result, AgentIdentityFailure) + assert "spend" not in result + assert "lifetime_budget_spend" not in result + assert result["litellm_budget_table"]["update"]["max_budget"] == 2.0 + + +@pytest.mark.parametrize("previous_duration", (None, "1d")) +def test_recreated_or_converted_lifetime_budget_gets_a_fresh_allowance(previous_duration: str | None) -> None: + from litellm.types.proxy.agent_identity import AgentBudgetState + + existing: Final = managed_agent().model_copy(update={ + "spend": 12.5, "lifetime_budget_spend": 0.75, + "budget_id": "previous" if previous_duration else None, + "litellm_budget_table": AgentBudgetState( + budget_id="previous", max_budget=1.0, budget_duration=previous_duration, + ) if previous_duration else None, + }) + result: Final = managed_write_fields({"budget": {"max_budget": 2.0}}, existing, "admin") + assert not isinstance(result, AgentIdentityFailure) + assert result["lifetime_budget_spend"] == 0.0 + assert "spend" not in result + assert "create" in result["litellm_budget_table"] + assert "update" not in result["litellm_budget_table"] diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 353249dddf0..b0fd544c718 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -10276,3 +10276,30 @@ async def test_authoritative_group_grants_propagate_policy_outages( await _get_agent_ids_from_access_groups(["group"], check_db_only=True) else: assert await _get_agent_ids_from_access_groups(["group"]) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("skip", [False, True]) +async def test_explicit_budget_skip_applies_to_agent_budget(skip: bool, monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.auth.auth_checks import common_checks + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentBudgetState + + auth: Final = UserAPIKeyAuth() + auth.billing_agent_policy = AgentResponse( + agent_id="agent", agent_name="Agent", agent_card_params={}, + litellm_budget_table=AgentBudgetState(budget_id="budget", max_budget=0), + ) + monkeypatch.setattr(proxy_server, "get_current_spend", AsyncMock(return_value=0)) + checks: Final = common_checks( + request_body={"model": "gpt-4"}, team_object=None, user_object=None, + end_user_object=None, global_proxy_spend=None, general_settings={}, + route="/v1/chat/completions", llm_router=None, proxy_logging_obj=MagicMock(), + valid_token=auth, request=MagicMock(spec=Request), skip_budget_checks=skip, + ) + if skip: + assert await checks is True + else: + with pytest.raises(litellm.BudgetExceededError, match="Agent budget exceeded"): + await checks diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index ef6832ef77b..b514e78984f 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -6267,6 +6267,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it(): "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "query_string": b"", } ) request._url = URL(url="/chat/completions") @@ -6321,6 +6322,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "query_string": b"", } ) request._url = URL(url="/chat/completions") @@ -8403,7 +8405,7 @@ _DDTRACE_AUTH_PROBE = dedent( async def auth(api_key): - request = Request(scope={"type": "http", "headers": [], "method": "POST", "path": "/chat/completions"}) + request = Request(scope={"type": "http", "headers": [], "method": "POST", "path": "/chat/completions", "query_string": b""}) request._url = URL(url="/chat/completions") try: await user_api_key_auth( @@ -9688,3 +9690,222 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon ) assert admitted.authenticated_by_custom_auth is False assert admitted.via_virtual_key is True + + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "http_method,route,body,billed", + [ + ("GET", "/a2a/agent/.well-known/agent-card.json", {}, False), + ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "tasks/get", "params": {"id": "t"}}, False), + ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/send", "params": {}}, True), + ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/stream", "params": {}}, True), + ("POST", "/a2a/agent", {"method": "message/send", "model": "free-model", "params": {}}, True), + ("POST", "/a2a/agent", {"method": "message/stream", "model": "free-model", "params": {}}, True), + ("POST", "/v1/chat/completions", {"model": "a2a/agent", "method": "tasks/get"}, True), + ("POST", "/chat/completions", {"model": "a2a/agent", "method": "tasks/cancel", "stream": True}, True), + ("POST", "/v1/a2a/agent/message/send", {"method": "tasks/get", "params": {}}, True), + ], +) +async def test_human_agent_discovery_does_not_reserve_target_budget_but_send_and_stream_do( + monkeypatch: pytest.MonkeyPatch, http_method: str, route: str, body: dict, billed: bool +) -> None: + from typing import Final + + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + target: Final = AgentResponse( + agent_id="agent", + agent_name="Agent", + agent_card_params={}, + identity_managed=True, + execution_mode="both", + litellm_params={"cost_per_query": 0.25}, + litellm_budget_table={"budget_id": "agent-budget", "max_budget": 10.0}, + identity=AgentIdentityBinding( + agent_id="agent", + provider="microsoft_entra", + tenant_id="tenant", + client_id="client", + service_principal_id="principal", + issuer="issuer", + revision="current", + ), + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + for name, value in { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": database, + "llm_router": litellm.Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"}, + "model_info": {"input_cost_per_token": 0, "output_cost_per_token": 0}, + } + ] + ) + if body.get("model") + else None, + "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + }.items(): + monkeypatch.setattr(proxy_server, name, value) + request = _alias_request(route, body) + request.scope["method"] = http_method + auth: Final = UserAPIKeyAuth( + api_key="human-key", + user_id="human", + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["agent"]), + ) + database.get_data = AsyncMock(return_value=auth) + proxy_server.proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock(return_value=None) + with patch( + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + new_callable=AsyncMock, + return_value=None, + ) as reserve: + assert await _authorize_authenticated_request(auth, request, body, route, "human-key") is None + assert auth.invoked_agent_id == "agent" + assert auth.agent_invocation_cost == pytest.approx(0.25 if billed else 0.0), (http_method, body.get("method")) + if billed: + assert auth.billing_agent_policy is not None and auth.billing_agent_policy.agent_id == "agent" + else: + assert auth.billing_agent_policy is None, (http_method, body.get("method")) + reserve.assert_awaited_once() + reserved: Final = reserve.call_args.kwargs["valid_token"] + assert reserved is auth and (reserved.billing_agent_policy is not None) is billed, (http_method, body.get("method")) + + +@pytest.mark.parametrize("invocation_cost,skipped", [(None, True), (0.0, True), (0.25, False)]) +def test_free_model_only_waives_budgets_without_a_paid_agent_invocation( + invocation_cost: float | None, skipped: bool +) -> None: + from typing import Final + + from litellm.proxy.auth.user_api_key_auth import _should_skip_budget_checks + + router: Final = litellm.Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"}, + "model_info": {"input_cost_per_token": 0, "output_cost_per_token": 0}, + } + ] + ) + assert ( + _should_skip_budget_checks( + request_data={"model": "free-model"}, + route="/chat/completions", + request=None, + llm_router=router, + agent_invocation_cost=invocation_cost, + ) + is skipped + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "selection", + ( + "body", + "query", + "path", + "cli", + "default", + "key-alias", + "team-alias", + "global-alias", + "alias-chain", + "query-alias", + "query-over-alias", + "different-agent", + "router-alias", + "query-router-alias", + ), +) +async def test_agent_admission_prices_the_model_selected_for_dispatch( + monkeypatch: pytest.MonkeyPatch, + selection: str, +) -> None: + from typing import Final + + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + from litellm.types.agents import AgentResponse + + registry: Final = agent_registry.AgentRegistry() + registry.register_agent( + AgentResponse( + agent_id="paid", + agent_name="Paid", + agent_card_params={}, + litellm_params={"cost_per_query": 0.25}, + ) + ) + registry.register_agent( + AgentResponse( + agent_id="other", + agent_name="Other", + agent_card_params={}, + litellm_params={"cost_per_query": 0.75}, + ) + ) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr(litellm, "model_alias_map", {"global": "a2a/paid"}) + for name, value in { + **_proxy_attrs_for_centralized_checks(), + "user_model": "a2a/paid" if selection == "cli" else None, + "llm_router": litellm.Router(model_list=[]) if "router-alias" in selection else None, + "general_settings": {"completion_model": "a2a/paid"} if selection == "default" else {}, + "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + }.items(): + monkeypatch.setattr(proxy_server, name, value) + body: Final = { + "model": { + "body": "a2a/paid", + "key-alias": "alias", + "team-alias": "team-alias", + "global-alias": "global", + "alias-chain": "team-alias", + "query-over-alias": "alias", + "different-agent": "a2a/other", + "router-alias": "router-alias", + }.get(selection, "gpt-4o"), + "messages": [{"role": "user", "content": "Hello"}], + } + route: Final = "/openai/deployments/a2a/paid/chat/completions" if selection == "path" else "/v1/chat/completions" + request: Final = _alias_request(route, body, path_params={"model": "a2a/paid"} if selection == "path" else {}) + if selection in ("query", "query-over-alias", "different-agent", "query-alias"): + request.scope["query_string"] = b"model=alias" if selection == "query-alias" else b"model=a2a%2Fpaid" + if selection == "query-router-alias": + request.scope["query_string"] = b"model=router-alias" + auth: Final = UserAPIKeyAuth( + router_settings={"model_group_alias": {"router-alias": "a2a/paid"}}, + user_role="proxy_admin", + aliases={"alias": "global" if selection == "alias-chain" else "a2a/paid"}, + team_model_aliases={"team-alias": "alias" if selection == "alias-chain" else "a2a/paid"}, + ) + with patch( + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + new_callable=AsyncMock, + return_value=None, + ) as reserve: + assert await _authorize_authenticated_request(auth, request, body, route, "test-key") is None + assert auth.invoked_agent_id == "paid" + assert auth.agent_invocation_cost == pytest.approx(0.25) + reserve.assert_awaited_once() + assert reserve.call_args.kwargs["valid_token"].agent_invocation_cost == pytest.approx(0.25) + assert reserve.call_args.kwargs["request_body"]["model"] == "a2a/paid" diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py index 9e20386bf3d..761c740a370 100644 --- a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py +++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py @@ -177,7 +177,7 @@ async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_rep assert agent.agent_id == agent_id prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once_with( where={"agent_id": agent_id}, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) @@ -202,7 +202,7 @@ async def test_get_agent_with_read_through_recovers_agent_by_name(clean_agent_re assert agent.agent_name == agent_name prisma_client.db.litellm_agentstable.find_unique.assert_awaited_with( where={"agent_name": agent_name}, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 131db55ee01..f6b72e32028 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -4,22 +4,21 @@ import sys import types from datetime import datetime, timedelta, timezone from datetime import time as dt_time -from typing import Any, Dict, Final, List, Optional +from typing import Any, Final from unittest.mock import AsyncMock, MagicMock import httpx import prisma import pytest - -from litellm.proxy._types import LiteLLM_VerificationToken -from litellm.proxy.common_utils import reset_budget_job as reset_budget_job_module from litellm.constants import ( PROXY_BUDGET_RESCHEDULER_MIN_TIME, RESET_BUDGET_JOB_BATCH_SIZE, RESET_BUDGET_JOB_LOCK_TTL_SECONDS, RESET_BUDGET_JOB_NAME, ) +from litellm.proxy._types import LiteLLM_VerificationToken +from litellm.proxy.common_utils import reset_budget_job as reset_budget_job_module from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob, _RowReset from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings @@ -29,12 +28,12 @@ class MockTable: """A single prisma table: records reads/writes and replays canned rows.""" def __init__(self): - self.find_many_calls: List[Dict[str, Any]] = [] - self.update_many_calls: List[Dict[str, Any]] = [] - self._find_many_results: List[Any] = [] - self._find_many_error: Optional[tuple[int, Exception]] = None + self.find_many_calls: list[dict[str, Any]] = [] + self.update_many_calls: list[dict[str, Any]] = [] + self._find_many_results: list[Any] = [] + self._find_many_error: tuple[int, Exception] | None = None - def set_find_many_results(self, results: List[Any]): + def set_find_many_results(self, results: list[Any]): self._find_many_results = results def set_find_many_error(self, after_reads: int, error: Exception): @@ -44,10 +43,10 @@ class MockTable: async def find_many( self, - where: Dict[str, Any], - order: Optional[Dict[str, str]] = None, - take: Optional[int] = None, - ) -> List[Any]: + where: dict[str, Any], + order: dict[str, str] | None = None, + take: int | None = None, + ) -> list[Any]: """Replays canned rows, honouring the keyset cursor + ``take`` a paged caller relies on: without that a paged walk never advances and the test would hang instead of failing.""" @@ -63,7 +62,7 @@ class MockTable: rows.sort(key=lambda row: getattr(row, field, ""), reverse=direction == "desc") return rows[:take] if take is not None else rows - async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: + async def update_many(self, where: dict[str, Any], data: dict[str, Any]) -> dict[str, Any]: self.update_many_calls.append({"where": where, "data": data}) return {"count": 1} @@ -77,7 +76,7 @@ class MockBatcher: """ def __init__(self): - self.calls: List[Dict[str, Any]] = [] + self.calls: list[dict[str, Any]] = [] self.committed: bool = False class _Table: @@ -103,6 +102,7 @@ class MockBatcher: self.litellm_tagtable = _Table("tag", self) self.litellm_modelaccessgroupbudgettable = _Table("model_access_group", self) self.litellm_projecttable = _Table("project", self) + self.litellm_agentstable = _Table("agent", self) self.litellm_endusertable = _Table("enduser", self) async def commit(self): @@ -119,8 +119,9 @@ class MockDB: self.litellm_tagtable = MockTable() self.litellm_modelaccessgroupbudgettable = MockTable() self.litellm_projecttable = MockTable() - self.batch_calls: List[Dict[str, Any]] = [] - self.batchers: List[MockBatcher] = [] + self.litellm_agentstable = MockTable() + self.batch_calls: list[dict[str, Any]] = [] + self.batchers: list[MockBatcher] = [] def batch_(self): batcher = MockBatcher() @@ -140,21 +141,21 @@ class MockDB: class MockPrismaClient: def __init__(self): - self.data: Dict[str, List[Any]] = { + self.data: dict[str, list[Any]] = { "key": [], "user": [], "team": [], "budget": [], "enduser": [], } - self.updated_data: Dict[str, List[Any]] = { + self.updated_data: dict[str, list[Any]] = { "key": [], "user": [], "team": [], "budget": [], "enduser": [], } - self.get_data_calls: List[Dict[str, Any]] = [] + self.get_data_calls: list[dict[str, Any]] = [] self.db = MockDB() async def get_data(self, table_name, query_type, **kwargs): @@ -246,7 +247,7 @@ def _budget_row( ) -def _batch_writes(mock_prisma_client, table: str, op: str | None = None) -> List[Dict[str, Any]]: +def _batch_writes(mock_prisma_client, table: str, op: str | None = None) -> list[dict[str, Any]]: """Writes that were committed to the DB, optionally narrowed to one op.""" return [ call @@ -604,7 +605,7 @@ def test_budget_table_reset_zeroes_spend_on_every_linked_table( _POSTGRES_MAX_BIND_VARIABLES: Final = 32767 -def _bind_count(where: Dict[str, Any]) -> int: +def _bind_count(where: dict[str, Any]) -> int: """Bind variables one prisma where-clause compiles to: each scalar is one placeholder and an ``in`` list contributes one per element.""" return sum(len(value["in"]) if isinstance(value, dict) and "in" in value else 1 for value in where.values()) @@ -924,8 +925,8 @@ def test_reset_budget_skips_null_budget_id_endusers_when_default_not_in_reset_li def _make_reset_budget_windows_job( monkeypatch, - key_rows: List[Dict[str, Any]], - team_rows: List[Dict[str, Any]], + key_rows: list[dict[str, Any]], + team_rows: list[dict[str, Any]], ): """Build a ResetBudgetJob with a fully-mocked prisma client and a fake `litellm.proxy.proxy_server` module exposing a stub `spend_counter_cache`. @@ -1699,7 +1700,6 @@ def test_enduser_invalidation_is_paged_and_batched(reset_budget_job, mock_prisma assert evicted == {f"end_user_id:cust-{i:06d}" for i in range(population)} - def test_enduser_invalidation_reports_a_page_read_failure_instead_of_a_clean_finish( mock_prisma_client, monkeypatch ): @@ -1983,7 +1983,11 @@ def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mo budget = _budget_row(budget_id="budget-1", budget_duration="7d") mock_prisma_client.data["budget"] = [budget] mock_prisma_client.data["enduser"] = [ - type("EndUser", (), {"spend": 5.0, "litellm_budget_table": budget, "user_id": "enduser-1", "budget_id": "budget-1"}) + type( + "EndUser", + (), + {"spend": 5.0, "litellm_budget_table": budget, "user_id": "enduser-1", "budget_id": "budget-1"}, + ) ] asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) @@ -1998,6 +2002,7 @@ def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mo ("tag", "update_many"), ("model_access_group", "update_many"), ("project", "update_many"), + ("agent", "update_many"), ("enduser", "update_many"), ("budget", "update_many"), } @@ -2200,10 +2205,10 @@ class ChunkedPrismaClient(MockPrismaClient): seeing rows rather than quietly running out of data. """ - def __init__(self, chunks_by_table: Dict[str, List[List[Any]]]): + def __init__(self, chunks_by_table: dict[str, list[list[Any]]]): super().__init__() self._chunks_by_table = chunks_by_table - self.fetches_by_table: Dict[str, int] = {} + self.fetches_by_table: dict[str, int] = {} async def get_data(self, table_name, query_type, **kwargs): self.get_data_calls.append({"table_name": table_name, "query_type": query_type, **kwargs}) @@ -2396,8 +2401,8 @@ class PoisonRow: class RecordingServiceLogging: def __init__(self): - self.success_calls: List[Dict[str, Any]] = [] - self.failure_calls: List[Dict[str, Any]] = [] + self.success_calls: list[dict[str, Any]] = [] + self.failure_calls: list[dict[str, Any]] = [] async def async_service_success_hook(self, **kwargs): self.success_calls.append(kwargs) @@ -2475,8 +2480,8 @@ class FakePodLockManager: if self.redis_cache is not None: self.redis_cache.async_get_cache = AsyncMock(return_value="another-pod" if held_by_other else None) self._acquired = acquired - self.acquire_calls: List[Dict[str, str | int | None]] = [] - self.release_calls: List[str] = [] + self.acquire_calls: list[dict[str, str | int | None]] = [] + self.release_calls: list[str] = [] @staticmethod def get_redis_lock_key(cronjob_id: str) -> str: @@ -2607,21 +2612,21 @@ def test_reset_budget_lease_outlives_one_scheduler_tick(monkeypatch): assert RESET_BUDGET_JOB_LOCK_TTL_SECONDS > PROXY_BUDGET_RESCHEDULER_MIN_TIME -def _window_row(source_id_column: str, row_id: str, reset_at: datetime) -> Dict[str, Any]: +def _window_row(source_id_column: str, row_id: str, reset_at: datetime) -> dict[str, Any]: return { source_id_column: row_id, "budget_limits": [{"budget_duration": "1h", "reset_at": reset_at.isoformat(), "max_budget": 10}], } -def _paginating_window_job(monkeypatch, pages_by_table: Dict[str, List[List[Dict[str, Any]]]]): +def _paginating_window_job(monkeypatch, pages_by_table: dict[str, list[list[dict[str, Any]]]]): """Serve each table a canned sequence of pages and record every query. Returns (job, calls) where calls is a list of (sql, cursor, limit). """ prisma_client = MagicMock() remaining = {table: list(pages) for table, pages in pages_by_table.items()} - calls: List[Dict[str, Any]] = [] + calls: list[dict[str, Any]] = [] async def fake_query_raw(query: str, *args, **kwargs): table = "key" if '"LiteLLM_VerificationToken"' in query else "team" @@ -2751,7 +2756,7 @@ def test_debug_row_dump_is_deferred_until_a_record_is_emitted(): assert serialized == ["serialized"] -def _cursor_paginating_window_job(monkeypatch, key_rows: List[Dict[str, Any]]): +def _cursor_paginating_window_job(monkeypatch, key_rows: list[dict[str, Any]]): """Serve real keyset pages out of one ordered table, honouring the cursor. Unlike the canned-page helper above, this models the database: a page is @@ -2760,7 +2765,7 @@ def _cursor_paginating_window_job(monkeypatch, key_rows: List[Dict[str, Any]]): """ prisma_client = MagicMock() ordered = sorted(key_rows, key=lambda r: r["token"]) - visited: List[str] = [] + visited: list[str] = [] async def fake_query_raw(query: str, *args, **kwargs): if '"LiteLLM_TeamTable"' in query: @@ -2816,7 +2821,7 @@ class FlakyPrismaClient(MockPrismaClient): def __init__(self, *, read_failures: int = 0, commit_failures: int = 0, error: Exception | None = None): super().__init__() - self.reconnect_reasons: List[str] = [] + self.reconnect_reasons: list[str] = [] self.read_attempts: int = 0 self.commit_attempts: int = 0 self._read_failures = read_failures @@ -2994,7 +2999,7 @@ def test_transport_error_on_window_read_reconnects_and_still_resets(monkeypatch) expired = (datetime.utcnow() - timedelta(minutes=5)).isoformat() + "Z" key_rows = [{"token": "sk-expired", "budget_limits": [{"budget_duration": "1d", "reset_at": expired}]}] job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[]) - reconnect_reasons: List[str] = [] + reconnect_reasons: list[str] = [] good_query_raw = prisma_client.db.query_raw async def failing_once_query_raw(query: str, *args, **kwargs): @@ -3019,7 +3024,7 @@ def test_connect_error_on_window_write_reconnects_and_writes(monkeypatch): expired = (datetime.utcnow() - timedelta(minutes=5)).isoformat() + "Z" team_rows = [{"team_id": "team-expired", "budget_limits": [{"budget_duration": "1d", "reset_at": expired}]}] job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=[], team_rows=team_rows) - reconnect_reasons: List[str] = [] + reconnect_reasons: list[str] = [] async def failing_once_update(**kwargs) -> None: if not reconnect_reasons: @@ -3578,3 +3583,26 @@ def test_reset_deletes_spend_counter_instead_of_seeding(reset_budget_job, mock_p counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:user:carol") counter_cache.in_memory_cache.set_cache.assert_not_called() counter_cache.redis_cache.async_set_cache.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("rollover_cap", [None, 1.0]) +async def test_agent_reset_stamps_window_in_each_spend_mutation(rollover_cap): + from litellm.proxy.common_utils.reset_budget_job import _BudgetCascade + + client = MockPrismaClient() + job = ResetBudgetJob(MagicMock(), client) + next_window = datetime(2026, 1, 2, tzinfo=timezone.utc) + cascade = _BudgetCascade( + budget_ids=("agent-budget",), + budget_resets=(("agent-budget", next_window),), + rollover_caps={} if rollover_cap is None else {"agent-budget": rollover_cap}, + ) + await job._commit_budget_cascade(cascade) + agent_writes = [call for call in client.db.batch_calls if call["table"] == "agent"] + assert len(agent_writes) == (1 if rollover_cap is None else 2) + assert all(call["data"]["spend_window"] == next_window for call in agent_writes) + assert agent_writes[0]["data"]["spend"] == 0.0 + if rollover_cap is not None: + assert agent_writes[1]["data"]["spend"] == {"decrement": rollover_cap} + assert client.db.batchers[0].committed diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 7b160c055d2..ae338c518b3 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -930,7 +930,7 @@ async def test_commit_spend_updates_to_db_increments_agent_spend(): mock_batcher.litellm_agentstable.update_many.assert_called_once() call_kwargs = mock_batcher.litellm_agentstable.update_many.call_args[1] - assert call_kwargs["where"] == {"agent_id": agent_id} + assert call_kwargs["where"] == {"agent_id": agent_id, "spend_window": None} assert call_kwargs["data"] == {"spend": {"increment": response_cost}} @@ -1475,7 +1475,8 @@ async def test_add_spend_log_transaction_to_daily_end_user_transaction_skips_whe @pytest.mark.asyncio -async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agent_id_and_queues_update(): +@pytest.mark.parametrize("billing_agent", [None, "caller-agent"]) +async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agent_id_and_queues_update(billing_agent): """ Ensure agent_id is injected and queued for daily aggregation. """ @@ -1487,6 +1488,7 @@ async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agen payload = { "request_id": "req-123", "agent_id": agent_id, + "billing_agent_id": billing_agent, "user": "test-user", "startTime": "2024-01-01T12:00:00", "api_key": "test-key", @@ -1506,14 +1508,19 @@ async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agen prisma_client=mock_prisma, ) + if billing_agent is None: + writer.daily_agent_spend_update_queue.add_update.assert_not_awaited() + return writer.daily_agent_spend_update_queue.add_update.assert_called_once() call_args = writer.daily_agent_spend_update_queue.add_update.call_args[1] update_dict = call_args["update"] assert len(update_dict) == 1 + charged_agent: Final = billing_agent or agent_id for key, transaction in update_dict.items(): - assert key == f"{agent_id}_2024-01-01_test-key_gpt-4_openai_" - assert transaction["agent_id"] == agent_id + assert key == f"{charged_agent}_2024-01-01_test-key_gpt-4_openai_" + assert transaction["agent_id"] == charged_agent + assert transaction["spend"] == 0.3 assert transaction["date"] == "2024-01-01" assert transaction["api_key"] == "test-key" assert transaction["model"] == "gpt-4" @@ -4778,3 +4785,132 @@ async def test_shutdown_drain_that_lands_before_the_interrupted_tag_commit_resol assert redis_buffer.restored == [drained], "a tag batch whose COMMIT came back failed must be restored to Redis" (upsert,) = _daily_upserts(final_db, "LiteLLM_DailyTagSpend") assert _row_values(upsert, "api_requests") == [1] + + +@pytest.mark.asyncio +async def test_agent_spend_queue_keeps_admission_windows_separate(): + from litellm.types.agents import agent_budget_counter_key + + writer = DBSpendUpdateWriter() + client = MagicMock() + old_key = agent_budget_counter_key("window-agent", datetime(2026, 1, 1, tzinfo=timezone.utc)) + new_key = agent_budget_counter_key("window-agent", datetime(2026, 1, 2, tzinfo=timezone.utc)) + await writer._update_agent_db(0.4, "window-agent", client, counter_key=old_key) + await writer._update_agent_db(0.1, "window-agent", client, counter_key=new_key) + await writer._update_agent_db(0.2, "window-agent", client, counter_key=new_key) + transactions = await writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() + assert transactions["agent_list_transactions"] == {old_key: 0.4, new_key: pytest.approx(0.3)} + + +@pytest.mark.asyncio +async def test_agent_admission_window_survives_logging_payload_and_background_queue() -> None: + from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload + from litellm.types.agents import agent_budget_counter_key + + now: Final = datetime.now(timezone.utc) + counter: Final = agent_budget_counter_key("window-agent", now) + payload: Final = get_logging_payload( + kwargs={ + "model": "demo-model", + "litellm_params": { + "metadata": { + "billing_agent_id": "window-agent", + "billing_agent_counter_key": counter, + } + }, + }, + response_obj={}, + start_time=now, + end_time=now, + ) + assert json.loads(payload["metadata"])["billing_agent_counter_key"] == counter + writer: Final = DBSpendUpdateWriter() + await writer._batch_database_updates( + response_cost=0.4, + user_id=None, + hashed_token=None, + team_id=None, + org_id=None, + end_user_id=None, + prisma_client=MagicMock(), + litellm_proxy_budget_name=None, + payload=payload, + ) + transactions: Final = await writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() + assert transactions["agent_list_transactions"] == {counter: 0.4} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("counter", ["spend:agent:another-agent", "spend:agent_window:malformed:window-agent"]) +async def test_invalid_agent_window_cannot_charge_another_agent(counter: str) -> None: + writer: Final = DBSpendUpdateWriter() + with pytest.raises(ValueError, match="does not match"): + await writer._update_agent_db(0.4, "window-agent", MagicMock(), counter_key=counter) + transactions: Final = await writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() + assert transactions["agent_list_transactions"] == {} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "captured_window", [None, datetime(2026, 1, 1, tzinfo=timezone.utc), datetime(2026, 1, 2, tzinfo=timezone.utc)] +) +async def test_agent_settlement_charges_only_the_matching_current_window(captured_window): + from litellm.types.agents import agent_budget_counter_key + + active_window = datetime(2026, 1, 2, tzinfo=timezone.utc) + row = {"agent_id": "window-agent", "spend_window": active_window, "spend": 0.2} + + def apply_update(*, where, data): + if all(row[key] == value for key, value in where.items()): + row["spend"] += data["spend"]["increment"] + + batcher = MagicMock() + batcher.litellm_agentstable.update_many.side_effect = apply_update + transaction = AsyncMock() + transaction.batch_ = MagicMock(return_value=AsyncMock(__aenter__=AsyncMock(return_value=batcher))) + client = MagicMock() + client.db.tx.return_value = AsyncMock(__aenter__=AsyncMock(return_value=transaction)) + key = agent_budget_counter_key("window-agent", captured_window) + await DBSpendUpdateWriter._update_entity_spend_in_db( + entity_name="Agent", + transactions={key: 0.4}, + table_accessor="litellm_agentstable", + where_field="agent_id", + n_retry_times=0, + prisma_client=client, + proxy_logging_obj=MagicMock(), + ) + assert row["spend"] == pytest.approx(0.6 if captured_window == active_window else 0.2) + assert batcher.litellm_agentstable.update_many.call_args.kwargs["where"] == { + "agent_id": "window-agent", + "spend_window": captured_window, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("captured_budget", ("current", "retired")) +async def test_lifetime_settlement_preserves_history_without_charging_another_budget(captured_budget: str) -> None: + from types import SimpleNamespace + + row: Final = SimpleNamespace( + agent_id="agent", budget_id="current", spend_window=None, spend=12.5, lifetime_budget_spend=0.25 + ) + + def apply_update(*, where: dict[str, object], data: dict[str, dict[str, float]]) -> None: + if all(getattr(row, key) == value for key, value in where.items()): + for field, operation in data.items(): + setattr(row, field, getattr(row, field) + operation["increment"]) + + batcher: Final = MagicMock() + batcher.litellm_agentstable.update_many.side_effect = apply_update + transaction: Final = AsyncMock() + transaction.batch_ = MagicMock(return_value=AsyncMock(__aenter__=AsyncMock(return_value=batcher))) + client: Final = MagicMock() + client.db.tx.return_value = AsyncMock(__aenter__=AsyncMock(return_value=transaction)) + await DBSpendUpdateWriter._update_entity_spend_in_db( + entity_name="Agent", transactions={f"spend:agent_lifetime:{captured_budget}:agent": 0.25}, + table_accessor="litellm_agentstable", where_field="agent_id", n_retry_times=0, + prisma_client=client, proxy_logging_obj=MagicMock(), + ) + assert row.spend == 12.75 + assert row.lifetime_budget_spend == (0.5 if captured_budget == "current" else 0.25) diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index ab931277313..8e742c06b51 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -480,3 +480,60 @@ async def test_from_db_still_never_reads_the_end_user_row(): assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:end_user:customer-42") is None assert prisma.db.litellm_endusertable.where_clauses == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("spend", [None, 0.0, 7.25]) +async def test_agent_counter_reseed_uses_persisted_agent_spend(spend: float | None) -> None: + row: Final = None if spend is None else SimpleNamespace(agent_id="agent-1", spend=spend) + table: Final = _FakeFindUniqueTable(row) + client: Final = SimpleNamespace(db=SimpleNamespace(litellm_agentstable=table)) + assert await SpendCounterReseed.from_db(prisma_client=client, counter_key="spend:agent:agent-1") == spend + assert table.where_clauses == [{"agent_id": "agent-1"}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("window,expected", [("20260102T000000.000000Z", 0.3), ("20260101T000000.000000Z", 0.0)]) +async def test_agent_window_reseed_cannot_load_another_windows_spend(window: str, expected: float) -> None: + from datetime import datetime, timezone + from prisma.models import LiteLLM_AgentsTable + + row: Final = LiteLLM_AgentsTable.model_construct( + agent_id="agent:with:colons", + spend=0.3, + spend_window=datetime(2026, 1, 2, tzinfo=timezone.utc), + ) + writer: Final = AsyncMock(return_value=row) + replica: Final = AsyncMock(return_value=row.model_copy(update={"spend": 99.0})) + client: Final = SimpleNamespace( + writer_db=SimpleNamespace(litellm_agentstable=SimpleNamespace(find_unique=writer)), + db=SimpleNamespace(litellm_agentstable=SimpleNamespace(find_unique=replica)), + ) + result: Final = await SpendCounterReseed.from_db(client, f"spend:agent_window:{window}:agent:with:colons") + assert result == expected + replica.assert_not_awaited() + writer.assert_awaited_once_with(where={"agent_id": "agent:with:colons"}, include={"litellm_budget_table": True}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("missing", [True, False]) +async def test_agent_window_reseed_handles_missing_rows_and_malformed_keys(missing: bool) -> None: + lookup: Final = AsyncMock(return_value=None) + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_agentstable=SimpleNamespace(find_unique=lookup))) + key: Final = "spend:agent_window:20260102T000000.000000Z:missing" if missing else "spend:agent_window:malformed" + assert await SpendCounterReseed.from_db(client, key) is None + assert lookup.await_count == int(missing) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("budget_id,expected", [("current", 0.25), ("retired", 0.0), ("deleted", None)]) +async def test_lifetime_counter_reseed_excludes_historical_and_other_budget_spend(budget_id: str, expected: float | None) -> None: + from prisma.models import LiteLLM_AgentsTable + + row: Final = LiteLLM_AgentsTable.model_construct( + agent_id="agent", budget_id="current", spend=12.5, lifetime_budget_spend=0.25, + ) + lookup: Final = AsyncMock(return_value=row if expected is not None else None) + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_agentstable=SimpleNamespace(find_unique=lookup))) + assert await SpendCounterReseed.from_db(client, f"spend:agent_lifetime:{budget_id}:agent") == expected + lookup.assert_awaited_once_with(where={"agent_id": "agent"}) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 84da227c0a6..c1c5b2b570d 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -2745,3 +2745,36 @@ def test_autonomous_agent_cost_tracking_needs_no_human_or_virtual_key(agent_id: assert _should_track_cost_callback( user_api_key=None, user_id=None, team_id=None, end_user_id=None, call_type="acompletion", agent_id=agent_id ) is expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("billing_agent", [None, "verified-agent"]) +async def test_callback_does_not_charge_a_header_selected_agent( + billing_agent: str | None, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + + cache: Final = DualCache() + for key in ("spend:user:human", "spend:agent:header-selected-agent", "spend:agent:verified-agent", "spend:agent_window:20260102T000000.000000Z:verified-agent"): + cache.in_memory_cache.set_cache(key=key, value=0.0) + logging: Final = MagicMock() + logging.db_spend_update_writer.update_database = AsyncMock(return_value=True) + logging.slack_alerting_instance.customer_spend_alert = AsyncMock() + monkeypatch.setattr(proxy_server, "spend_counter_cache", cache) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging) + monkeypatch.setattr(proxy_server, "update_cache", AsyncMock()) + kwargs: Final = { + "call_type": "acompletion", "model": "test-model", "response_cost": 0.01, + "litellm_params": {"metadata": { + "user_api_key_user_id": "human", "agent_id": "header-selected-agent", "billing_agent_id": billing_agent, + "billing_agent_counter_key": "spend:agent_window:20260102T000000.000000Z:verified-agent" if billing_agent else None, + }}, + } + await _ProxyDBLogger()._PROXY_track_cost_callback( + kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now() + ) + assert cache.in_memory_cache.get_cache(key="spend:user:human") == 0.01 + assert cache.in_memory_cache.get_cache(key="spend:agent:header-selected-agent") == 0.0 + assert cache.in_memory_cache.get_cache(key="spend:agent:verified-agent") == 0.0 + assert cache.in_memory_cache.get_cache(key="spend:agent_window:20260102T000000.000000Z:verified-agent") == (0.01 if billing_agent else 0.0) diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index 59c2921e0d0..2843bb60965 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -714,6 +714,9 @@ class _FakePrismaClient: litellm_modelaccessgroupbudgettable=self.access_group_budget_table, litellm_proxymodeltable=self.model_table, ) + self.writer_db = SimpleNamespace( + litellm_agentstable=SimpleNamespace(find_first=AsyncMock(return_value=None)), + ) def jsonify_object(self, data): return dict(data) diff --git a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py index 2f3be61d00f..e6af298b8d8 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py @@ -21,6 +21,8 @@ def client_and_mocks(monkeypatch): mock_table = MagicMock() mock_table.create = AsyncMock(side_effect=lambda *, data: data) mock_table.update = AsyncMock(side_effect=lambda *, where, data: {**where, **data}) + mock_table.delete = AsyncMock(side_effect=lambda *, where: where) + mock_prisma.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None) mock_prisma.db = types.SimpleNamespace( litellm_budgettable=mock_table, @@ -46,6 +48,34 @@ def client_and_mocks(monkeypatch): monkeypatch.setattr(ps, "prisma_client", ps.prisma_client) +@pytest.mark.parametrize("operation", ["update", "delete"]) +def test_agent_linked_budget_requires_the_agent_management_flow(client_and_mocks, operation): + client, prisma, table = client_and_mocks + prisma.writer_db.litellm_agentstable.find_first.return_value = types.SimpleNamespace(agent_id="agent-one") + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + app.dependency_overrides[ps.user_api_key_auth] = lambda: admin + payload = {"budget_id": "agent-budget", "budget_duration": "2h"} if operation == "update" else {"id": "agent-budget"} + + response = client.post(f"/budget/{operation}", json=payload) + + assert response.status_code == 409 + assert "/v1/agents/agent-one" in response.json()["detail"] + table.update.assert_not_awaited() + table.delete.assert_not_awaited() + + +def test_unlinked_budget_can_still_be_deleted(client_and_mocks): + client, _, table = client_and_mocks + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + app.dependency_overrides[ps.user_api_key_auth] = lambda: admin + + response = client.post("/budget/delete", json={"id": "ordinary-budget"}) + + assert response.status_code == 200 + assert response.json()["budget_id"] == "ordinary-budget" + table.delete.assert_awaited_once_with(where={"budget_id": "ordinary-budget"}) + + @pytest.mark.asyncio async def test_new_budget_success(client_and_mocks): client, _, mock_table = client_and_mocks diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index 3c6afa86c45..2bdc81756c7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -1176,6 +1176,7 @@ async def _run_legacy_update_organization( mock_prisma_client.db.litellm_organizationtable.find_unique = AsyncMock(return_value=existing_org) mock_prisma_client.db.litellm_organizationtable.update = AsyncMock(return_value=MagicMock()) mock_prisma_client.db.litellm_budgettable.update = AsyncMock() + mock_prisma_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) diff --git a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py index 3cfdd345a45..830ffdcad92 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py @@ -35,6 +35,10 @@ class _BudgetState: return SimpleNamespace(**self._values) +def _writer_db_without_agent_budgets() -> SimpleNamespace: + return SimpleNamespace(litellm_agentstable=SimpleNamespace(find_first=AsyncMock(return_value=None))) + + class FakeVerificationTokenTable: """Stand-in for ``prisma_client.db.litellm_verificationtoken``. @@ -316,7 +320,7 @@ async def test_update_tag_explicit_null_preserves_general_budget_fields(field): created_by="admin", ) mock_db = Mock() - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, writer_db=_writer_db_without_agent_budgets()) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) @@ -370,7 +374,7 @@ async def test_update_tag_explicit_null_clears_budget_duration(): created_by="admin", ) mock_db = Mock() - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, writer_db=_writer_db_without_agent_budgets()) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 0b866d7f736..5f0488b0c76 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -2973,6 +2973,7 @@ async def test_upsert_team_member_budget_table_clears_duration_kept_budget(mock_ mock_db_client.db.litellm_budgettable.update = AsyncMock( side_effect=lambda where, data: SimpleNamespace(**data) ) + mock_db_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None) result = await TeamMemberBudgetHandler.upsert_team_member_budget_table( team_table=team_table, diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index 922504ecc58..46bac64db19 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -845,6 +845,7 @@ async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): db: Final = _FakeDb() prisma_client: Final = MagicMock() prisma_client.db = db + prisma_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None) admin: Final = UserAPIKeyAuth(user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN) team_id: Final = "team-shared-default" default_budget: Final = await db.litellm_budgettable.create(data={"budget_id": "team-default", "max_budget": 100.0}) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 4577d578263..72364904f93 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -2270,7 +2270,7 @@ class TestGigachatProxyRoute: mock_request.headers = {"content-type": "application/json"} mock_request.query_params = {} mock_fastapi_response = MagicMock(spec=Response) - mock_user_api_key_dict = MagicMock() + mock_user_api_key_dict = UserAPIKeyAuth() mock_llm_router.allm_passthrough_route = AsyncMock( return_value=httpx.Response(200, json={"response": "success"}) ) diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 7096bc7c632..b5ca6f5c42d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -4759,6 +4759,7 @@ def _agent_db_row(agent_id: str, agent_name: str): agent_access_groups=[], access_group_ids=[], spend=0.0, + lifetime_budget_spend=0.0, identity_managed=False, enabled=True, execution_mode="autonomous", 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 9df8e6f4d67..84b905af901 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -306,6 +306,77 @@ async def test_team_member_reservation_counter_adds_temp_increase_to_live_team_d assert counter.fallback_spend == 0.5 +@pytest.mark.asyncio +@pytest.mark.parametrize("charged_agent", ("caller-agent", "target-agent")) +@pytest.mark.parametrize("window", [None, "2026-01-02T00:00:00Z"]) +@pytest.mark.parametrize("outcome", ["success", "cancelled"]) +async def test_agent_invocation_reserves_exact_fee_and_reconciles_without_child_cost( + monkeypatch: pytest.MonkeyPatch, charged_agent: str, window: str | None, outcome: str, +) -> None: + from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation + from litellm.types.agents import AgentResponse + + cache: Final = DualCache() + counter_key: Final = f"spend:agent_window:20260102T000000.000000Z:{charged_agent}" if window else f"spend:agent_lifetime:agent-budget:{charged_agent}" + cache.set_cache(counter_key, 0.1) + monkeypatch.setattr(proxy_server, "spend_counter_cache", cache) + auth: Final = UserAPIKeyAuth(agent_id="caller-agent" if charged_agent == "caller-agent" else None) + auth.billing_agent_policy = AgentResponse( + agent_id=charged_agent, agent_name="Charged agent", agent_card_params={}, spend=0.1, lifetime_budget_spend=0.1, + litellm_budget_table={"budget_id": "agent-budget", "max_budget": 0.5, "budget_reset_at": window, "budget_duration": "1d" if window else None}, + ) + auth.invoked_agent_id = "target-agent" + auth.agent_invocation_cost = 0.2 + reservation: Final = await reserve_budget_for_request( + request_body={"jsonrpc": "2.0", "method": "message/send"}, route="/a2a/target-agent", + llm_router=None, valid_token=auth, team_object=None, 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, + ) + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.2) + assert [entry["counter_key"] for entry in reservation["entries"]] == [counter_key] + assert await cache.async_get_cache(counter_key) == pytest.approx(0.3) + if outcome == "cancelled": + from litellm.proxy.spend_tracking.budget_reservation import release_budget_reservation_on_cancel + + await release_budget_reservation_on_cancel(reservation) + assert await cache.async_get_cache(counter_key) == pytest.approx(0.1) + assert reservation["finalized"] is True + return + await proxy_server.increment_spend_counters( + token=None, team_id=None, user_id=None, response_cost=0.2, + billing_agent_id=charged_agent, billing_agent_counter_key=counter_key, budget_reservation=reservation, + ) + await reconcile_budget_reservation(reservation, actual_cost=4.0) + assert await cache.async_get_cache(counter_key) == pytest.approx(0.3) + + +@pytest.mark.asyncio +async def test_agent_invocation_over_budget_is_rejected_and_reservation_is_refunded( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.types.agents import AgentResponse + + cache: Final = DualCache() + cache.set_cache("spend:agent_lifetime:agent-budget:agent", 0.4) + monkeypatch.setattr(proxy_server, "spend_counter_cache", cache) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.billing_agent_policy = AgentResponse( + agent_id="agent", agent_name="Charged agent", agent_card_params={}, spend=0.4, lifetime_budget_spend=0.4, + litellm_budget_table={"budget_id": "agent-budget", "max_budget": 0.5}, + ) + auth.agent_invocation_cost = 0.2 + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body={"method": "message/send"}, route="/a2a/target-agent", llm_router=None, + valid_token=auth, team_object=None, 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, + ) + assert await cache.async_get_cache("spend:agent_lifetime:agent-budget:agent") == pytest.approx(0.4) + + @pytest.mark.asyncio async def test_reservation_starts_unbound_to_any_callback(): reservation: Final = await _reserve("/v1/responses") @@ -342,3 +413,83 @@ async def test_release_unbound_budget_reservation_leaves_a_bound_one_to_its_call assert spend_counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(reservation["reserved_cost"]) assert reservation["finalized"] is False + + +@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, budget_owner: str, route: str +) -> None: + import asyncio + + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + from litellm.proxy.spend_tracking.budget_reservation import ( + reconcile_budget_reservation, + release_budget_reservation, + release_budget_reservation_on_cancel, + ) + from litellm.types.agents import AgentResponse + + caller: Final = AgentResponse( + agent_id="caller", agent_name="Caller", agent_card_params={}, spend=0.0, + litellm_budget_table={"budget_id": "caller-budget", "max_budget": 0.5}, + ) + target: Final = AgentResponse( + agent_id="target", agent_name="Target", agent_card_params={}, litellm_params={"cost_per_query": 0.25}, + ) + registry: Final = agent_registry.AgentRegistry() + 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" 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", "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, + ) + + results: Final = await asyncio.gather(*(admit() for _ in range(8)), return_exceptions=True) + accepted: Final = tuple(result for result in results if isinstance(result, dict)) + rejected: Final = tuple(result for result in results if isinstance(result, litellm.BudgetExceededError)) + 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(counter_key) == pytest.approx(0.5) + first: Final = accepted[0] + if outcome == "completed": + await proxy_server.increment_spend_counters( + 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(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(counter_key) == pytest.approx(0.25) + assert await admit() is not None + assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.5) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 3ffb6335ad4..f8cd308145f 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -3767,7 +3767,7 @@ class TestSpendLogsPayload: "model": "gpt-4o", "user": "", "team_id": "", - "metadata": '{"actor_agent_id": null, "target_agent_id": null, "billing_agent_id": null, "agent_execution_mode": null, "verified_human_user_id": null, "applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}', + "metadata": '{"actor_agent_id": null, "target_agent_id": null, "billing_agent_id": null, "billing_agent_counter_key": null, "agent_execution_mode": null, "verified_human_user_id": null, "applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}', "cache_key": "Cache OFF", "spend": 0.00022500000000000002, "total_tokens": 30, diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index c8e4df1030f..6d2c10321e8 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -3720,3 +3720,38 @@ async def test_unreserved_model_access_group_is_charged_alongside_a_reserved_one assert counter_cache.in_memory_cache.get_cache( key=model_access_group_spend_counter_key("starter") ) == pytest.approx(4.2) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("invocation_cost", [0.0, 0.01]) +async def test_agent_stream_cancellation_leaves_refund_to_request_cleanup(spend_counter_state, invocation_cost): + from litellm.proxy.middleware.budget_reservation_release_middleware import BudgetReservationReleaseMiddleware + from litellm.proxy.spend_tracking.budget_reservation import release_unbound_budget_reservation + + counter_cache, _ = spend_counter_state + key = "spend:agent:cancelled-agent" + counter_cache.set_cache(key, invocation_cost) + reservation = {"reserved_cost": invocation_cost, "input_cost": 0.0, "finalized": False, "entries": [{"counter_key": key, "reserved_cost": invocation_cost}]} + auth = UserAPIKeyAuth() + auth.agent_invocation_cost = invocation_cost + auth.budget_reservation = reservation + + async def cancel_before_chunk(user_api_key_dict, response, request_data): + raise asyncio.CancelledError() + yield "unreachable" + + generator, logging = _drive_streaming_cancel(auth, cancel_before_chunk) + + async def app(scope, receive, send): + try: + await anext(generator) + finally: + assert reservation["finalized"] is False + assert await counter_cache.async_get_cache(key) == pytest.approx(invocation_cost) + + middleware = BudgetReservationReleaseMiddleware(app, release_unbound_budget_reservation) + with pytest.raises(asyncio.CancelledError): + await middleware({"type": "http", "state": {"budget_reservation": reservation}}, AsyncMock(), AsyncMock()) + assert reservation["finalized"] is True + assert await counter_cache.async_get_cache(key) == pytest.approx(0.0) + logging._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 84266325226..cda9d066661 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8523,3 +8523,19 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback(): team_callback_settings_obj=None, ) assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"} +@pytest.mark.parametrize("bound", [False, True]) +def test_agent_budget_window_metadata_is_owned_by_authenticated_policy(bound: bool) -> None: + from litellm.types.agents import AgentResponse + + auth: Final = UserAPIKeyAuth(agent_id="agent" if bound else None) + if bound: + auth.billing_agent_policy = AgentResponse( + agent_id="agent", agent_name="Agent", agent_card_params={}, + litellm_budget_table={"budget_id": "budget", "budget_reset_at": "2026-01-02T00:00:00Z"}, + ) + result: Final = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( + {"metadata": {"billing_agent_counter_key": "spend:agent:victim"}}, auth, "metadata" + ) + assert result["metadata"]["billing_agent_counter_key"] == ( + "spend:agent_window:20260102T000000.000000Z:agent" if bound else None + ) diff --git a/tests/test_litellm/proxy/test_pricing_field_strip.py b/tests/test_litellm/proxy/test_pricing_field_strip.py index a0e25e91f37..87cf39b0d49 100644 --- a/tests/test_litellm/proxy/test_pricing_field_strip.py +++ b/tests/test_litellm/proxy/test_pricing_field_strip.py @@ -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 diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 0429dd97a1c..8783ab1b61e 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -183,3 +183,103 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re assert call_kwargs["model"] == f"a2a/{agent_name}" assert call_kwargs["api_base"] == "http://sibling-db-agent.example.com" prisma_client.db.litellm_agentstable.find_unique.assert_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "override", ("api_key", "api_base", "user_config", "router_settings_override", "deployment", "no-router") +) +async def test_registered_agent_dispatch_owns_the_admitted_destination_and_fee(monkeypatch: pytest.MonkeyPatch, override: str) -> None: + from typing import Final + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + from litellm.types.agents import AgentResponse + + agent: Final = AgentResponse( + agent_id="paid", + agent_name="paid", + agent_card_params={"url": "https://registered.test/"}, + litellm_params={"cost_per_query": 0.25}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(agent) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + router: Final = _router_without_models() + if override == "deployment": + router.is_recognized_model.return_value = True + provider: Final = AsyncMock(return_value={"id": "reply"}) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + with patch("litellm.acompletion", provider): + await prepare_agent_invocation(auth, "paid", None) + pending: Final = await route_request( + data={ + "model": "a2a/paid", + "messages": [{"role": "user", "content": "Hi"}], + "cost_per_query": 99.0, + **( + {override: {} if override in ("user_config", "router_settings_override") else "override"} + if override not in ("deployment", "no-router") + else {} + ), + }, + llm_router=None if override == "no-router" else router, + user_model=None, + route_type="acompletion", + user_api_key_dict=auth, + ) + assert await pending == {"id": "reply"} + provider.assert_awaited_once() + assert provider.call_args.kwargs["api_base"] == "https://registered.test/" + assert provider.call_args.kwargs["cost_per_query"] == 0.25 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ("a2a/paid", "gpt-4o")) +async def test_routing_overrides_cannot_dispatch_without_matching_agent_admission( + monkeypatch: pytest.MonkeyPatch, model: str, +) -> None: + from typing import Final + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + from litellm.types.agents import AgentResponse + + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(AgentResponse( + agent_id="paid", agent_name="paid", agent_card_params={"url": "https://agent.test/"}, + litellm_params={"cost_per_query": 0.25}, + )) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr(proxy_server, "prisma_client", None) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + if model == "gpt-4o": + await prepare_agent_invocation(auth, "paid", None) + provider: Final = Mock(return_value=None) + with patch("litellm.acompletion", provider), pytest.raises(HTTPException, match="admission") as exc: + await route_request( + data={"model": model, "api_base": "https://override.test/", "messages": [{"role": "user", "content": "Hi"}]}, + llm_router=None, user_model=None, route_type="acompletion", user_api_key_dict=auth, + ) + assert exc.value.status_code == 503 + provider.assert_not_called() + + +@pytest.mark.asyncio +async def test_unregistered_direct_agent_keeps_explicit_endpoint_routing(monkeypatch: pytest.MonkeyPatch) -> None: + from typing import Final + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + + monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry()) + monkeypatch.setattr(proxy_server, "prisma_client", None) + provider: Final = AsyncMock(return_value={"id": "direct-reply"}) + with patch("litellm.acompletion", provider): + pending: Final = await route_request( + data={"model": "a2a/direct", "api_base": "https://direct.test/", "messages": [{"role": "user", "content": "Hi"}]}, + llm_router=None, user_model=None, route_type="acompletion", + ) + assert await pending == {"id": "direct-reply"} + assert provider.call_args.kwargs["api_base"] == "https://direct.test/" diff --git a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py index 2e883e91fda..9bf43537737 100644 --- a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py +++ b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py @@ -135,3 +135,40 @@ async def test_stream_completion_counts_tokens_off_the_event_loop(monkeypatch): assert usage.prompt_tokens > 100_000 assert usage.completion_tokens > 100_000 assert_loop_stayed_free(took, lags) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outcome", ["success", "failure", "cancelled"]) +async def test_stream_reservation_survives_cleanup_only_when_billing_is_scheduled(outcome): + from unittest.mock import AsyncMock + + from litellm.proxy.spend_tracking.budget_reservation import release_unbound_budget_reservation + + reservation = {"reserved_cost": 0.01, "entries": [], "finalized": False} + logging_obj = SimpleNamespace( + litellm_params={"metadata": {"user_api_key_budget_reservation": reservation}}, + model_call_details={}, + dispatch_success_handlers=AsyncMock(), + ) + + async def stream(): + yield {"result": {"kind": "message", "parts": [{"kind": "text", "text": "hello"}]}} + if outcome == "failure": + raise RuntimeError("upstream failed") + if outcome == "cancelled": + raise asyncio.CancelledError() + + iterator = A2AStreamingIterator( + stream=stream(), + request=SimpleNamespace(params=SimpleNamespace(message={"parts": [{"kind": "text", "text": "hi"}]})), + logging_obj=logging_obj, + ) + if outcome == "success": + assert len([chunk async for chunk in iterator]) == 1 + else: + with pytest.raises(RuntimeError if outcome == "failure" else asyncio.CancelledError): + _ = [chunk async for chunk in iterator] + await release_unbound_budget_reservation(reservation) + assert reservation["finalized"] is (outcome != "success") + await asyncio.sleep(0) + assert logging_obj.dispatch_success_handlers.await_count == (1 if outcome == "success" else 0) diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py index 8d8ec815f3a..62dc077054e 100644 --- a/tests/unit/a2a_protocol/test_cost_calculator.py +++ b/tests/unit/a2a_protocol/test_cost_calculator.py @@ -117,6 +117,7 @@ class CostLogger(CustomLogger): def __init__(self): self.response_cost: Optional[float] = None + self.logged: asyncio.Event = asyncio.Event() super().__init__() async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -125,6 +126,7 @@ class CostLogger(CustomLogger): self.response_cost = ( slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None) ) + self.logged.set() @pytest.mark.asyncio @@ -449,3 +451,124 @@ async def test_asend_message_streaming_triggers_callbacks(): assert callback_logger.agent_id == test_agent_id, ( f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'" ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", (False, True)) +@pytest.mark.parametrize("claimed_fee", (None, -1000.0, 0.0, 99.0)) +@pytest.mark.parametrize("fee_field", ("cost_per_query", "litellm_params")) +@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, + changed_after_admission: bool, configured_fee: float | None, fee_field: str +) -> None: + import json + from typing import Final + + import httpx + + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + 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() + logger: Final = CostLogger() + monkeypatch.setattr(litellm, "callbacks", [logger]) + 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": configured_fee} if configured_fee is not None else {}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + 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) + assert body["method"] == ("message/stream" if stream else "message/send") + result: Final = { + "jsonrpc": "2.0", "id": body["id"], + "result": {"kind": "message", "messageId": "reply", "role": "agent", + "parts": [{"kind": "text", "text": "Paid reply"}]}, + } + if stream: + return httpx.Response(200, text=f"data: {json.dumps(result)}\n\n", headers={"content-type": "text/event-stream"}) + return httpx.Response(200, json=result) + + client: Final = AsyncHTTPHandler(transport=httpx.MockTransport(reply)) + try: + pending: Final = await route_a2a_agent_request( + data={"model": "a2a/fee-target", "messages": [{"role": "user", "content": "Hello"}], + "stream": stream, "client": client, + **({fee_field: claimed_fee if fee_field == "cost_per_query" else {"cost_per_query": claimed_fee}} + if claimed_fee is not None else {})}, + route_type="acompletion", user_api_key_dict=auth, + ) + response: Final = await pending + if stream: + chunks: Final = tuple([chunk async for chunk in response]) + assert any(chunk.choices[0].delta.content == "Paid reply" for chunk in chunks) + else: + 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) + 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() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("admitted_target", (None, "other")) +async def test_chat_agent_dispatch_rejects_missing_or_different_admission( + monkeypatch: pytest.MonkeyPatch, + admitted_target: str | None, +) -> None: + from typing import Final + from unittest.mock import Mock + + from fastapi import HTTPException + + 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 + + registry: Final = agent_registry.AgentRegistry() + for name in ("paid", "other"): + registry.register_agent( + AgentResponse( + agent_id=name, + agent_name=name, + agent_card_params={"url": "https://agent.test/"}, + litellm_params={"cost_per_query": 0.25}, + ) + ) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + auth: Final = UserAPIKeyAuth(user_role="proxy_admin") + if admitted_target is not None: + await prepare_agent_invocation(auth, admitted_target, None) + provider: Final = Mock(return_value=None) + monkeypatch.setattr(litellm, "acompletion", provider) + with pytest.raises(HTTPException) as exc: + await route_a2a_agent_request( + data={"model": "a2a/paid", "messages": [{"role": "user", "content": "Hello"}]}, + route_type="acompletion", + user_api_key_dict=auth, + ) + assert exc.value.status_code == 503 + assert "admission" in str(exc.value.detail).lower() + provider.assert_not_called() diff --git a/tests/unit/a2a_protocol/test_main.py b/tests/unit/a2a_protocol/test_main.py index 4ba0ef8fa04..1720b521bd7 100644 --- a/tests/unit/a2a_protocol/test_main.py +++ b/tests/unit/a2a_protocol/test_main.py @@ -539,3 +539,34 @@ def test_streaming_logging_obj_keeps_agent_credentials_out_of_logging_params(): assert logging_obj.litellm_params == expected assert logging_obj.optional_params == expected assert logging_obj.model_call_details["litellm_params"] == expected + + +class _AgentFeeRecorder(CustomLogger): + def __init__(self): + super().__init__() + self.logged = asyncio.Event() + self.fees = () + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + if kwargs.get("call_type") in ("asend_message", "asend_message_streaming"): + self.fees = (*self.fees, (kwargs.get("agent_id"), kwargs["standard_logging_object"]["response_cost"])) + self.logged.set() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +async def test_completion_bridge_records_one_agent_fee(streaming, monkeypatch): + from litellm.a2a_protocol.main import asend_message_streaming + + recorder = _AgentFeeRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + params = {"custom_llm_provider": "openai", "model": "gpt-4o-mini", "mock_response": "hello back", "cost_per_query": 0.01} + if streaming: + request = SendStreamingMessageRequest(id="bridge-stream", params=_request().params) + chunks = [chunk async for chunk in asend_message_streaming(request=request, litellm_params=params, agent_id="budgeted-agent")] + assert chunks[-1]["result"]["final"] is True + else: + response = await asend_message(request=_request(), litellm_params=params, agent_id="budgeted-agent") + assert response.id == "r1" + await asyncio.wait_for(recorder.logged.wait(), timeout=2) + assert recorder.fees == (("budgeted-agent", pytest.approx(0.01)),) diff --git a/tests/unit/types/proxy/test_agent_identity.py b/tests/unit/types/proxy/test_agent_identity.py new file mode 100644 index 00000000000..ed4b0901da2 --- /dev/null +++ b/tests/unit/types/proxy/test_agent_identity.py @@ -0,0 +1,20 @@ +import pytest +from pydantic import ValidationError + +from litellm.types.proxy.agent_identity import AgentBudgetConfig + + +@pytest.mark.parametrize("amount", [-1, float("inf"), float("-inf"), float("nan")]) +def test_agent_budget_rejects_negative_or_nonfinite_caps(amount: float) -> None: + with pytest.raises(ValidationError): + AgentBudgetConfig(max_budget=amount) + + +@pytest.mark.parametrize("amount", [0, 0.01, 100]) +def test_agent_budget_preserves_a_finite_nonnegative_cap(amount: float) -> None: + assert AgentBudgetConfig(max_budget=amount).max_budget == amount + + +def test_agent_budget_rejects_unknown_policy_fields() -> None: + with pytest.raises(ValidationError): + AgentBudgetConfig.model_validate({"max_budget": 1, "unknown_control": True}) diff --git a/tests/unit/types/test_agents.py b/tests/unit/types/test_agents.py new file mode 100644 index 00000000000..9acb7c41b2a --- /dev/null +++ b/tests/unit/types/test_agents.py @@ -0,0 +1,41 @@ +from datetime import datetime, timedelta, timezone +from typing import Final + +import pytest + +from litellm.types.agents import agent_budget_counter_key, agent_spend_filter + + +@pytest.mark.parametrize("offset", [None, timezone.utc, timezone(timedelta(hours=5, minutes=30))]) +def test_agent_window_key_round_trip_preserves_the_admitted_instant(offset) -> None: + instant: Final = datetime(2030, 1, 1, 12, 30, 1, 123000, tzinfo=offset) + expected: Final = instant.replace(tzinfo=timezone.utc) if offset is None else instant.astimezone(timezone.utc) + key: Final = agent_budget_counter_key("agent:with:colons", instant) + assert agent_spend_filter(key) == {"agent_id": "agent:with:colons", "spend_window": expected} + assert key == agent_budget_counter_key("agent:with:colons", expected) + + +def test_unbudgeted_key_is_filtered_to_an_unbudgeted_row() -> None: + key: Final = agent_budget_counter_key("agent-one", None) + assert agent_spend_filter(key) == {"agent_id": "agent-one", "spend_window": None} + + +def test_different_budget_windows_never_share_a_settlement_filter() -> None: + first: Final = datetime(2030, 1, 1, tzinfo=timezone.utc) + second: Final = first + timedelta(days=1) + assert agent_spend_filter(agent_budget_counter_key("agent-one", first)) != agent_spend_filter( + agent_budget_counter_key("agent-one", second) + ) + + +def test_lifetime_budget_consumption_is_separate_from_agent_history() -> None: + from litellm.types.agents import AgentResponse + + agent: Final = AgentResponse( + agent_id="agent", agent_name="Agent", agent_card_params={}, spend=12.5, + lifetime_budget_spend=0.75, + litellm_budget_table={"budget_id": "budget", "max_budget": 1.0}, + ) + assert agent.budget_spend == 0.75 + assert agent.spend == 12.5 + assert agent.budget_counter_key == "spend:agent_lifetime:budget:agent" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx index 50c60776ff3..8495c99d59f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx @@ -256,6 +256,41 @@ export const AgentIdentityFields = ({ accessToken }: { accessToken: string | nul )} +
+

Agent Budget

+ + {({ value, onChange, ref, ...control }) => ( + + )} + + + {({ value, onChange, ref, ...control }) => ( + + )} + +
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts index 0639e7a6dd4..f54d6364ddc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts @@ -1,5 +1,6 @@ import { describe, expect, it } from "vitest"; import { + agentBudgetSpend, buildIdentityParams, entraTenantFromIssuer, parseIdentityForForm, @@ -45,7 +46,7 @@ describe("agent identity configuration", () => { it("rejects incomplete submissions", () => { expect(() => buildIdentityParams({ identity_provider: "microsoft_entra" })).toThrow("Enter valid Entra"); }); - it("submits identity as top-level settings without changing runtime parameters", () => { + it("submits identity and budget as top-level settings without changing runtime parameters", () => { const formValues = { identity_provider: "microsoft_entra", identity_tenant_id: identity.tenant_id, @@ -53,6 +54,8 @@ describe("agent identity configuration", () => { identity_service_principal_id: identity.service_principal_id, execution_mode: "both", enabled: false, + agent_max_budget: 0, + agent_budget_duration: "1d", }; const payload = withAgentIdentity({ litellm_params: { model: "runtime" } }, formValues); expect(payload.litellm_params).toEqual({ model: "runtime" }); @@ -62,6 +65,7 @@ describe("agent identity configuration", () => { }); expect(payload.execution_mode).toBe("both"); expect(payload.enabled).toBe(false); + expect(payload.budget).toEqual({ max_budget: 0, budget_duration: "1d" }); }); it("requires a service principal for autonomous execution", () => { const values = { @@ -81,3 +85,25 @@ describe("agent identity configuration", () => { expect(entraTenantFromIssuer("https://login.microsoftonline.com/common/v2.0")).toBeNull(); }); }); + +describe("agent budget consumption", () => { + it("keeps historical spend separate from lifetime budget consumption", () => { + expect( + agentBudgetSpend({ + spend: 12.5, + lifetime_budget_spend: 0.75, + litellm_budget_table: { budget_id: "budget", max_budget: 1 }, + }), + ).toBe(0.75); + }); + it("preserves recurring window spend and defaults missing lifetime consumption to zero", () => { + expect( + agentBudgetSpend({ + spend: 0.5, + lifetime_budget_spend: 9, + litellm_budget_table: { budget_id: "budget", max_budget: 1, budget_duration: "1d" }, + }), + ).toBe(0.5); + expect(agentBudgetSpend({ spend: 12.5, litellm_budget_table: { budget_id: "budget", max_budget: 1 } })).toBe(0); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts index 23045adcf20..24f25a54f65 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts @@ -5,7 +5,7 @@ import type { AgentFormValues, AgentRequestPayload } from "./AgentFormKit"; export type EntraAgentIdentity = components["schemas"]["EntraIdentityConfig"]; type AgentIdentityState = Pick< components["schemas"]["AgentResponse"], - "identity" | "enabled" | "execution_mode" | "agent_card_params" + "identity" | "enabled" | "execution_mode" | "agent_card_params" | "litellm_budget_table" >; export const IDENTITY_UUID_PATTERN = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i; @@ -47,6 +47,8 @@ export const parseIdentityForForm = (agent?: Partial | null) ...identityFormFields(identity), execution_mode: agent?.execution_mode ?? "autonomous", enabled: agent?.enabled ?? true, + agent_max_budget: agent?.litellm_budget_table?.max_budget ?? "", + agent_budget_duration: agent?.litellm_budget_table?.budget_duration ?? "", }; }; @@ -104,5 +106,20 @@ export const withAgentIdentity = ( ...identityFields, ...(managed && values.execution_mode !== undefined ? { execution_mode: values.execution_mode } : {}), ...(managed && values.enabled !== undefined ? { enabled: values.enabled } : {}), + ...(values.agent_max_budget === undefined + ? {} + : { + budget: + values.agent_max_budget === "" || values.agent_max_budget === null + ? null + : { max_budget: Number(values.agent_max_budget), budget_duration: values.agent_budget_duration || null }, + }), }; }; + +export const agentBudgetSpend = ( + agent: Pick, +): number => + agent.litellm_budget_table && !agent.litellm_budget_table.budget_duration + ? agent.lifetime_budget_spend ?? 0 + : agent.spend ?? 0; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx index e08cab776c4..3ba83920589 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx @@ -156,6 +156,44 @@ describe("AgentInfoView update payload", () => { .mockResolvedValue({} as never); }); + it.each(["preserve", "edit", "clear"])( + "%s lifetime budget without displaying historical spend as consumption", + async (action) => { + const budgetedAgent = { + ...A2A_AGENT, + spend: 12.5, + lifetime_budget_spend: 0.5, + litellm_budget_table: { budget_id: "budget", max_budget: 1, budget_duration: null }, + }; + vi.mocked(networking.getAgentInfo).mockResolvedValue(budgetedAgent); + const user = setup(); + renderView(); + expect(await screen.findByText("$0.5 / $1")).toBeInTheDocument(); + await openEditor(user); + const limit = screen.getByLabelText("Aggregate Agent Budget ($)"); + expect(limit).toHaveValue(1); + expect(screen.getByLabelText("Budget Reset Period")).toHaveValue(""); + if (action !== "preserve") { + fireEvent.change(limit, { target: { value: action === "edit" ? "2" : "" } }); + } + await save(user); + expect(patchedPayload().budget).toEqual( + action === "clear" ? null : { max_budget: action === "edit" ? 2 : 1, budget_duration: null }, + ); + }, + ); + + it("creates a recurring budget from the budget fields", async () => { + const user = setup(); + renderView(); + expect(await screen.findByText("No aggregate limit")).toBeInTheDocument(); + await openEditor(user); + fireEvent.change(screen.getByLabelText("Aggregate Agent Budget ($)"), { target: { value: "0.75" } }); + fireEvent.change(screen.getByLabelText("Budget Reset Period"), { target: { value: "1d" } }); + await save(user); + expect(patchedPayload().budget).toEqual({ max_budget: 0.75, budget_duration: "1d" }); + }); + it.each([ { card: "complete", editCard: false }, { card: "empty", editCard: false }, @@ -222,7 +260,8 @@ describe("AgentInfoView update payload", () => { await save(user); - expect(patchedPayload()).toEqual({ + const expectedPayload = { + budget: null, agent_name: "my-agent", agent_card_params: { protocolVersion: "1.0", @@ -241,7 +280,8 @@ describe("AgentInfoView update payload", () => { session_rpm_limit: 444, object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} }, access_group_ids: [], - }); + }; + expect(patchedPayload()).toEqual(expectedPayload); }); it("sends the loaded values of every panel the user opens", async () => { @@ -259,7 +299,8 @@ describe("AgentInfoView update payload", () => { await save(user); - expect(patchedPayload()).toEqual({ + const expectedPayload = { + budget: null, agent_name: "my-agent", agent_card_params: { protocolVersion: "1.0", @@ -283,7 +324,8 @@ describe("AgentInfoView update payload", () => { session_rpm_limit: 444, object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} }, access_group_ids: [], - }); + }; + expect(patchedPayload()).toEqual(expectedPayload); }); it("clamps a rate limit typed below its minimum up to that minimum", async () => { @@ -342,7 +384,8 @@ describe("AgentInfoView update payload", () => { await save(user); - expect(patchedPayload()).toEqual({ + const expectedPayload = { + budget: null, agent_name: "lg-agent", agent_card_params: { protocolVersion: "1.0", @@ -362,7 +405,8 @@ describe("AgentInfoView update payload", () => { }, object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} }, access_group_ids: [], - }); + }; + expect(patchedPayload()).toEqual(expectedPayload); }); it("keeps the agent's existing MCP grants in the update payload", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx index c9f154c3ce1..491496b8a07 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx @@ -1,6 +1,6 @@ import { AgentIdentityFields } from "./AgentIdentityFields"; import { AgentIdentityDetails } from "./AgentIdentityDetails"; -import { withAgentIdentity } from "./agent_identity"; +import { agentBudgetSpend, withAgentIdentity } from "./agent_identity"; import React, { useState, useEffect, useMemo } from "react"; import { cx } from "@/lib/cva.config"; import { FormProvider, useForm, useWatch } from "react-hook-form"; @@ -78,6 +78,21 @@ const DetailItem: React.FC<{ label: React.ReactNode; children: React.ReactNode } ); +const AgentBudgetDetails = ({ agent }: { agent: Agent }) => ( + <> + + {agent.litellm_budget_table?.max_budget != null + ? `$${agentBudgetSpend(agent)} / $${agent.litellm_budget_table.max_budget}` + : "No aggregate limit"} + + + {agent.litellm_budget_table?.budget_reset_at + ? new Date(agent.litellm_budget_table.budget_reset_at).toLocaleString() + : "No scheduled reset"} + + +); + const AgentInfoView: React.FC = ({ agentId, onClose, accessToken, isAdmin }) => { const [agent, setAgent] = useState(null); const [selectedKey, setSelectedKey] = useState(null); @@ -354,6 +369,7 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT {agent.agent_id} {agent.agent_name} + {agent.agent_card_params?.name || "-"} {agent.agent_card_params?.description || "-"} {agent.agent_card_params?.url || "-"} diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts index 92d946c19e1..2df1c53333c 100644 --- a/ui/litellm-dashboard/src/components/agents/types.ts +++ b/ui/litellm-dashboard/src/components/agents/types.ts @@ -15,6 +15,7 @@ export interface Agent { identity_managed?: boolean; enabled?: boolean; execution_mode?: components["schemas"]["AgentResponse"]["execution_mode"]; + litellm_budget_table?: components["schemas"]["AgentBudgetState"] | null; jwt_auth_configured?: boolean; agent_id: string; agent_name: string; @@ -32,6 +33,7 @@ export interface Agent { kill_switch?: AgentKillSwitchConfig | null; keys?: AgentAttachedKey[] | null; spend?: number; + lifetime_budget_spend?: number; tpm_limit?: number | null; rpm_limit?: number | null; session_tpm_limit?: number | null; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 8ffd5c96ab0..f199f132812 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -24597,6 +24597,24 @@ export interface components { [key: string]: string; }; }; + /** AgentBudgetConfig */ + AgentBudgetConfig: { + /** Budget Duration */ + budget_duration?: string | null; + /** Max Budget */ + max_budget: number; + }; + /** AgentBudgetState */ + AgentBudgetState: { + /** Budget Duration */ + budget_duration?: string | null; + /** Budget Id */ + budget_id: string; + /** Budget Reset At */ + budget_reset_at?: string | null; + /** Max Budget */ + max_budget?: number | null; + }; /** * AgentCapabilities * @description Defines optional capabilities supported by an agent. @@ -24678,6 +24696,7 @@ export interface components { agent_card_params?: components["schemas"]["AgentCard"]; /** Agent Name */ agent_name: string; + budget?: components["schemas"]["AgentBudgetConfig"] | null; /** Enabled */ enabled?: boolean; /** @@ -24964,6 +24983,8 @@ export interface components { agent_id: string; /** Agent Name */ agent_name: string; + /** Budget Id */ + budget_id?: string | null; /** Created At */ created_at?: string | null; /** Created By */ @@ -24995,6 +25016,12 @@ export interface components { /** Keys */ keys?: components["schemas"]["AgentKeySummary"][] | null; kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null; + /** + * Lifetime Budget Spend + * @default 0 + */ + lifetime_budget_spend: number; + litellm_budget_table?: components["schemas"]["AgentBudgetState"] | null; /** Litellm Params */ litellm_params?: { [key: string]: unknown; @@ -38742,6 +38769,7 @@ export interface components { agent_card_params?: components["schemas"]["AgentCard"]; /** Agent Name */ agent_name?: string; + budget?: components["schemas"]["AgentBudgetConfig"] | null; /** Enabled */ enabled?: boolean; /**