From 5b64f51c1ebc22f6f7cfe2c677659404c3529ac8 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:12:00 -0700 Subject: [PATCH] feat(agents): enforce aggregate agent budgets --- .../migration.sql | 15 +++ .../litellm_proxy_extras/schema.prisma | 4 + litellm/proxy/_lazy_openapi_snapshot.json | 114 +++++++++++++++++ litellm/proxy/_types.py | 1 + .../proxy/agent_endpoints/agent_registry.py | 12 +- .../auth/managed_authorization.py | 24 +++- litellm/proxy/agent_endpoints/endpoints.py | 6 +- .../proxy/agent_endpoints/identity_store.py | 1 + .../proxy/agent_endpoints/managed_identity.py | 72 ++++++++++- litellm/proxy/auth/auth_checks.py | 4 + .../common_utils/registry_read_through.py | 1 + .../proxy/common_utils/reset_budget_job.py | 1 + litellm/proxy/db/db_spend_update_writer.py | 27 ++-- litellm/proxy/db/spend_counter_reseed.py | 17 +++ .../proxy/hooks/proxy_track_cost_callback.py | 18 +++ litellm/proxy/litellm_pre_call_utils.py | 5 + .../budget_management_endpoints.py | 15 +++ litellm/proxy/proxy_server.py | 15 +++ litellm/proxy/schema.prisma | 4 + .../spend_tracking/budget_reservation.py | 66 +++++++--- .../spend_tracking/spend_counter_batch.py | 22 +++- litellm/repositories/prisma_protocols.py | 3 + litellm/repositories/unit_of_work.py | 27 +++- litellm/types/agents.py | 36 +++++- litellm/types/proxy/agent_identity.py | 16 +++ schema.prisma | 4 + .../auth/test_managed_authorization.py | 83 +++++++++++++ .../agent_endpoints/test_agent_registry.py | 2 +- .../agent_endpoints/test_managed_identity.py | 58 ++++++++- .../proxy/auth/test_auth_checks.py | 27 ++++ .../test_registry_read_through.py | 4 +- .../common_utils/test_reset_budget_job.py | 108 ++++++++++------ .../proxy/db/test_db_spend_update_writer.py | 115 +++++++++++++++++- .../proxy/db/test_spend_counter_reseed.py | 43 +++++++ .../hooks/test_proxy_track_cost_callback.py | 33 +++++ .../test_budget_endpoints.py | 30 +++++ .../spend_tracking/test_budget_reservation.py | 63 ++++++++++ .../proxy/test_litellm_pre_call_utils.py | 18 +++ tests/unit/types/proxy/test_agent_identity.py | 20 +++ tests/unit/types/test_agents.py | 28 +++++ .../_components/AgentIdentityFields.tsx | 35 ++++++ .../agents/_components/agent_identity.test.ts | 5 +- .../agents/_components/agent_identity.ts | 19 ++- .../agents/_components/agent_info.tsx | 16 +++ .../src/components/agents/types.ts | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 23 ++++ 46 files changed, 1170 insertions(+), 91 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql create mode 100644 tests/unit/types/proxy/test_agent_identity.py create mode 100644 tests/unit/types/test_agents.py 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..6f1db719855 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql @@ -0,0 +1,15 @@ +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); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 8535d1004f7..92dfe53e29e 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,9 @@ model LiteLLM_AgentsTable { execution_mode String @default("autonomous") identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] + budget_id String? @unique + spend_window DateTime? + 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/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index f374ba09d07..2133e0450f7 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,16 @@ } ] }, + "litellm_budget_table": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentBudgetState" + }, + { + "type": "null" + } + ] + }, "litellm_params": { "anyOf": [ { @@ -4035,6 +4139,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 641be7078b7..cb00cc427d7 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4084,6 +4084,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/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index ef7128a0e8e..2e70fc6cee6 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -632,7 +632,7 @@ class AgentRegistry: # Create agent in DB created_agent: Final = await agents_table(prisma_client).create( data={**create_data, **_managed_fields(agent, None, created_by)}, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) return AgentResponse.model_validate(created_agent.model_dump()) @@ -695,7 +695,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") @@ -747,7 +747,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}") @@ -777,7 +777,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 @@ -845,7 +845,7 @@ class AgentRegistry: updated_by, ), }, - include={"object_permission": True, "identity": True}, + include={"object_permission": True, "identity": True, "litellm_budget_table": True}, ) if updated_agent is None: @@ -868,7 +868,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 abcef353eef..48980c780a0 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -148,6 +148,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")) @@ -183,6 +185,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.spend or 0.0, + 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)]) @@ -215,12 +235,12 @@ 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: + if not effective.identity_managed and effective.litellm_budget_table is None and auth.managed_agent_policy is None: return 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 - if auth.agent_id is None and effective.identity_managed: + if 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 try: diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index e1b2ac63d51..5ea8266d5ed 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -766,7 +766,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 +873,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 +965,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 87512c779fe..1a568b7a487 100644 --- a/litellm/proxy/agent_endpoints/identity_store.py +++ b/litellm/proxy/agent_endpoints/identity_store.py @@ -63,6 +63,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 260b74fcbd1..a520fbdc2ce 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, @@ -66,12 +68,29 @@ class IdentityHistoryWrite(TypedDict): connectOrCreate: ReadOnly[IdentityHistoryConnect] +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] def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn: @@ -118,14 +137,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: @@ -183,6 +206,53 @@ 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) + 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"], + **( + {"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 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 11c7dbd64cf..d3c5ea97275 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1084,6 +1084,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/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 17e6152bef6..0cd4eff85a4 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -82,6 +82,7 @@ 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: @@ -1068,12 +1069,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( @@ -1341,15 +1345,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)["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, ) ) @@ -2331,7 +2339,11 @@ class DBSpendUpdateWriter: response_cost, ) _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 @@ -2904,13 +2916,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( @@ -2919,8 +2932,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..fc2fbd4dec8 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,6 +39,7 @@ 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 @@ -170,6 +172,21 @@ class SpendCounterReseed: return await OrganizationRepository(prisma_client).table.find_unique( where={"organization_id": counter_key[len("spend:org:") :]} ) + if counter_key.startswith("spend:agent_window:"): + parts: Final = counter_key.split(":", 3) + if len(parts) != 4: + return None + row: Final = await AgentsRepository(prisma_client, use_writer=True).table.find_unique( + where={"agent_id": parts[3]}, include={"litellm_budget_table": True} + ) + 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={"spend": 0.0}) + if counter_key.startswith("spend:agent:"): + return await AgentsRepository(prisma_client).table.find_unique( + where={"agent_id": counter_key[len("spend:agent:") :]} + ) 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 e735f4a51f0..64c7fd6ec29 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 @@ -378,6 +379,8 @@ class _ProxyDBLogger(CustomLogger): request_tags=tags, model_access_groups=model_access_groups, project_id=project_id, + billing_agent_id=metadata.get("billing_agent_id"), + billing_agent_counter_key=metadata.get("billing_agent_counter_key"), ) if not charged: return @@ -682,6 +685,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: ... @@ -702,6 +708,8 @@ async def _update_database_and_spend_counters( request_tags: list[str] | 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, ) -> bool: if budget_reservation is not None: await _reconcile_budget_reservation_before_db_update( @@ -751,6 +759,16 @@ async def _update_database_and_spend_counters( 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 a698881189a..b122f82fbea 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1667,6 +1667,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 b3b20a16234..7299793921d 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3001,6 +3001,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. @@ -3023,6 +3025,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( @@ -3037,6 +3041,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, ) @@ -3052,6 +3058,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.""" reserved_counter_keys: Final = await _reconcile_budget_reservation_for_counter_update( @@ -3223,9 +3231,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/schema.prisma b/litellm/proxy/schema.prisma index 8535d1004f7..92dfe53e29e 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,9 @@ model LiteLLM_AgentsTable { execution_mode String @default("autonomous") identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] + budget_id String? @unique + spend_window DateTime? + 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 e28fa2c06a4..5c89c87fb99 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,19 +285,27 @@ 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({}) ) current_spend_by_counter_key: Final[dict[str, float]] = {} - 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). @@ -346,7 +356,7 @@ async def reserve_budget_for_request( fail_closed_budget_enforcement=fail_closed_budget_enforcement, ) continue - except Exception: + except (asyncio.CancelledError, Exception): await _release_applied_entries_best_effort( entries=applied_entries, default_reserved_cost=reservation_cost, @@ -356,11 +366,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 = ( + invocation_cost + 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, @@ -500,7 +514,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.spend or 0.0, + 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 ddb074ae023..efd1135e3b7 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -157,6 +157,10 @@ def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) yield f"spend:end_user:{end_user_id}" if token.org_id is not None: yield f"spend:org:{token.org_id}" + 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 + if charged_agent_id is not None: + yield billing_agent.budget_counter_key if billing_agent is not None else f"spend:agent:{charged_agent_id}" if token.project_id is not None: yield project_spend_counter_key(token.project_id) @@ -174,10 +178,19 @@ 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 = admission_counter_keys( - UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id, project_id=project_id), + UserAPIKeyAuth( + token=token, + team_id=team_id, + user_id=user_id, + org_id=org_id, + project_id=project_id, + agent_id=billing_agent_id if billing_agent_counter_key is None else None, + ), end_user_id, ) tag_keys: Final = frozenset(f"spend:tag:{tag}" for tag in tags or () if tag and isinstance(tag, str)) @@ -186,7 +199,12 @@ 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 + return ( + entity_keys + | tag_keys + | group_keys + | (frozenset((billing_agent_counter_key,)) if billing_agent_counter_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..964ef94d227 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,30 @@ class AgentKeySummary(BaseModel): key_name: str | None = None +def agent_budget_counter_key(agent_id: str, reset_at: datetime | None) -> str: + 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_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 + litellm_budget_table: AgentBudgetState | None = None identity: AgentIdentityBinding | None = None identity_managed: bool = False enabled: bool = True @@ -338,6 +366,12 @@ 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 + ) + 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 8535d1004f7..92dfe53e29e 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,9 @@ model LiteLLM_AgentsTable { execution_mode String @default("autonomous") identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] + budget_id String? @unique + spend_window DateTime? + 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 2e37f0f4462..e0cda7ce62d 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 @@ -95,6 +95,38 @@ 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: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" + ) + 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" + ) + with pytest.raises(litellm.BudgetExceededError): + await check_agent_budget(auth) + assert await counters.async_get_cache("spend:agent: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)) async def test_invocation_prepares_target_fee_for_the_correct_agent( @@ -172,6 +204,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() @@ -482,3 +531,37 @@ async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None: assert auth.managed_agent_policy == agent() assert auth.billing_agent_policy == agent() assert auth.user_id is None + + +@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 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 ba160b8e1b4..81c2e29115a 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -1451,7 +1451,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_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py index 45fe4b0655f..25146b47875 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,10 @@ 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 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 88833f38270..69785f5924d 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -10135,3 +10135,30 @@ async def test_managed_agent_model_policy_checks_dispatched_model( with pytest.raises((HTTPException, ModelAccessDeniedProxyException)) as failure: await checks assert str(getattr(failure.value, "status_code", getattr(failure.value, "code", None))) == "403" + + +@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/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 7abb6e1ef92..1ff080a82d5 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" @@ -4739,3 +4746,103 @@ 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, + } 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..391cede1005 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,46 @@ 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) 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 d7401554d58..63dcd2c66c5 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 @@ -2743,3 +2743,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_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/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py index 9df8e6f4d67..469d8506311 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,69 @@ 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"]) +async def test_agent_invocation_reserves_exact_fee_and_reconciles_without_child_cost( + monkeypatch: pytest.MonkeyPatch, charged_agent: str, window: str | None, +) -> 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:{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, + litellm_budget_table={"budget_id": "agent-budget", "max_budget": 0.5, "budget_reset_at": window}, + ) + 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) + 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: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, + 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:agent") == pytest.approx(0.4) + + @pytest.mark.asyncio async def test_reservation_starts_unbound_to_any_callback(): reservation: Final = await _reserve("/v1/responses") 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 ee42042bb8e..6b08260e99a 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8470,3 +8470,21 @@ async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custo for name, value in secrets.items(): assert updated["secret_fields"]["raw_headers"][name.lower()] == value assert request.headers[name] == value + + +@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/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..a583ab9e8a2 --- /dev/null +++ b/tests/unit/types/test_agents.py @@ -0,0 +1,28 @@ +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) + ) 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 43538784145..d0e38cd1d62 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx @@ -254,6 +254,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..10ca2114186 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 @@ -45,7 +45,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 +53,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 +64,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 = { 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 34986f75679..1f6adb6ae96 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 @@ -3,7 +3,10 @@ import type { components } from "@/lib/http/schema"; import type { AgentFormValues, AgentRequestPayload } from "./AgentFormKit"; export type EntraAgentIdentity = components["schemas"]["EntraIdentityConfig"]; -type AgentIdentityState = Pick; +type AgentIdentityState = Pick< + components["schemas"]["AgentResponse"], + "identity" | "enabled" | "execution_mode" | "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; @@ -44,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 ?? "", }; }; @@ -92,10 +97,22 @@ export const withAgentIdentity = ( ): AgentRequestPayload => { const identityFields = buildIdentityParams(values, existing?.identity); const managed = values.identity_provider === "microsoft_entra" || Boolean(readAgentIdentity(existing?.identity)); + const budgetIsSet = + values.agent_max_budget !== undefined && values.agent_max_budget !== "" && values.agent_max_budget !== null; + const budgetWasSet = existing?.litellm_budget_table?.max_budget != null; return { ...payload, ...identityFields, ...(managed && values.execution_mode !== undefined ? { execution_mode: values.execution_mode } : {}), ...(managed && values.enabled !== undefined ? { enabled: values.enabled } : {}), + ...(budgetIsSet + ? { + budget: { + max_budget: Number(values.agent_max_budget), + budget_duration: values.agent_budget_duration || null, + }, + } + : {}), + ...(!budgetIsSet && budgetWasSet && values.agent_max_budget !== undefined ? { budget: null } : {}), }; }; 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 d4c05d0ebac..a05ded386de 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 @@ -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 + ? `$${agent.spend ?? 0} / $${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); @@ -349,6 +364,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..703d6170764 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; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 51d0165330b..05dc42d1bc9 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -24162,6 +24162,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. @@ -24243,6 +24261,7 @@ export interface components { agent_card_params?: components["schemas"]["AgentCard"]; /** Agent Name */ agent_name: string; + budget?: components["schemas"]["AgentBudgetConfig"] | null; /** Enabled */ enabled?: boolean; /** @@ -24529,6 +24548,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 */ @@ -24560,6 +24581,7 @@ export interface components { /** Keys */ keys?: components["schemas"]["AgentKeySummary"][] | null; kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null; + litellm_budget_table?: components["schemas"]["AgentBudgetState"] | null; /** Litellm Params */ litellm_params?: { [key: string]: unknown; @@ -37770,6 +37792,7 @@ export interface components { agent_card_params?: components["schemas"]["AgentCard"]; /** Agent Name */ agent_name?: string; + budget?: components["schemas"]["AgentBudgetConfig"] | null; /** Enabled */ enabled?: boolean; /**