From 71e1cc900caea10094540663a31d71fa198c9450 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Sat, 26 Sep 2026 12:12:52 -0700
Subject: [PATCH] feat(agents): enforce budgets across admission and deferred
accounting
---
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 +++
.../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 | 12 ++
.../auth/test_managed_authorization.py | 83 +++++++++++++
.../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 +++
.../_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 ++++
38 files changed, 1055 insertions(+), 89 deletions(-)
diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json
index 259d13c619c..ed7460ac815 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 f9551bdfc4d..8eee8d9cc01 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -4076,6 +4076,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 5de9d3d73aa..f38b51977b7 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -2999,6 +2999,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.
@@ -3021,6 +3023,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(
@@ -3035,6 +3039,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,
)
@@ -3050,6 +3056,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(
@@ -3221,9 +3229,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/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 5564105de7d..964ef94d227 100644
--- a/litellm/types/agents.py
+++ b/litellm/types/agents.py
@@ -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,
@@ -254,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]
@@ -272,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]
@@ -334,6 +338,8 @@ def agent_spend_filter(counter_key: str) -> "LiteLLM_AgentsTableWhereInput":
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
@@ -360,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/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_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/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
>
)}
+
>
);
};
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 4a6f3463ebe..c3a356a6098 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;
@@ -37765,6 +37787,7 @@ export interface components {
agent_card_params?: components["schemas"]["AgentCard"];
/** Agent Name */
agent_name?: string;
+ budget?: components["schemas"]["AgentBudgetConfig"] | null;
/** Enabled */
enabled?: boolean;
/**