From 69b8ada9accd9e2fcd3bda43a7eea2c1a29272af Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Mon, 28 Sep 2026 16:23:54 -0700
Subject: [PATCH 01/12] feat(agents): agent budgets
---
.../migration.sql | 15 ++
.../litellm_proxy_extras/schema.prisma | 4 +
litellm/proxy/_lazy_openapi_snapshot.json | 114 +++++++++++++
litellm/proxy/_types.py | 1 +
.../proxy/agent_endpoints/agent_registry.py | 12 +-
.../auth/managed_authorization.py | 56 ++++++-
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 +
litellm/proxy/auth/user_api_key_auth.py | 7 +-
.../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 | 22 +++
litellm/proxy/litellm_pre_call_utils.py | 5 +
.../budget_management_endpoints.py | 15 ++
litellm/proxy/proxy_server.py | 15 ++
litellm/proxy/schema.prisma | 4 +
.../spend_tracking/budget_reservation.py | 66 +++++---
.../spend_tracking/spend_counter_batch.py | 18 ++-
litellm/repositories/prisma_protocols.py | 3 +
litellm/repositories/unit_of_work.py | 27 +++-
litellm/types/agents.py | 36 ++++-
litellm/types/proxy/agent_identity.py | 16 ++
schema.prisma | 4 +
.../auth/test_managed_authorization.py | 153 +++++++++++++++++-
.../agent_endpoints/test_agent_registry.py | 2 +-
.../agent_endpoints/test_managed_identity.py | 58 ++++++-
.../proxy/auth/test_auth_checks.py | 27 ++++
.../proxy/auth/test_user_api_key_auth.py | 118 ++++++++++++++
.../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_access_group_management.py | 3 +
.../test_budget_endpoints.py | 30 ++++
.../test_organization_endpoints.py | 1 +
.../test_tag_management_endpoints.py | 8 +-
.../test_team_endpoints.py | 1 +
.../test_management_helpers_utils.py | 1 +
.../spend_tracking/test_budget_reservation.py | 63 ++++++++
.../test_spend_management_endpoints.py | 2 +-
.../proxy/test_litellm_pre_call_utils.py | 16 ++
tests/unit/types/proxy/test_agent_identity.py | 20 +++
tests/unit/types/test_agents.py | 28 ++++
.../_components/AgentIdentityFields.tsx | 35 ++++
.../agents/_components/agent_identity.test.ts | 5 +-
.../agents/_components/agent_identity.ts | 16 +-
.../agents/_components/agent_info.tsx | 16 ++
.../src/components/agents/types.ts | 1 +
ui/litellm-dashboard/src/lib/http/schema.d.ts | 23 +++
54 files changed, 1398 insertions(+), 101 deletions(-)
create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql
create mode 100644 tests/unit/types/proxy/test_agent_identity.py
create mode 100644 tests/unit/types/test_agents.py
diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql
new file mode 100644
index 00000000000..6f1db719855
--- /dev/null
+++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql
@@ -0,0 +1,15 @@
+ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "budget_id" TEXT;
+
+-- CreateIndex
+CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentsTable_budget_id_key" ON "LiteLLM_AgentsTable"("budget_id");
+
+-- AddForeignKey
+DO $$
+BEGIN
+ IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_AgentsTable_budget_id_fkey') THEN
+ ALTER TABLE "LiteLLM_AgentsTable" ADD CONSTRAINT "LiteLLM_AgentsTable_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE;
+ END IF;
+END $$;
+
+
+ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "spend_window" TIMESTAMP(3);
diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
index adfe2a0eee7..6823deadbd3 100644
--- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
+++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
@@ -36,6 +36,7 @@ model LiteLLM_BudgetTable {
model_access_groups LiteLLM_ModelAccessGroupBudgetTable[] // multiple model access groups can have the same budget
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
+ agents LiteLLM_AgentsTable[]
}
// Models on proxy
@@ -83,6 +84,9 @@ model LiteLLM_AgentsTable {
execution_mode String @default("autonomous")
identity LiteLLM_AgentIdentity?
retired_identities LiteLLM_RetiredAgentIdentity[]
+ budget_id String? @unique
+ spend_window DateTime?
+ litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json
index fa4b36a03aa..1c760d8dc6b 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": [
{
@@ -4011,6 +4115,16 @@
"title": "Agent Name",
"type": "string"
},
+ "budget": {
+ "anyOf": [
+ {
+ "$ref": "#/components/schemas/AgentBudgetConfig"
+ },
+ {
+ "type": "null"
+ }
+ ]
+ },
"enabled": {
"title": "Enabled",
"type": "boolean"
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index d2fad212dd9..b74b9a240df 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -4104,6 +4104,7 @@ class SpendLogsRouterMetadata(TypedDict):
class SpendLogsMetadata(TypedDict):
+ billing_agent_counter_key: ReadOnly[NotRequired[str | None]]
actor_agent_id: ReadOnly[NotRequired[str | None]]
target_agent_id: ReadOnly[NotRequired[str | None]]
billing_agent_id: ReadOnly[NotRequired[str | None]]
diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py
index 7929f67720d..7dbe409226a 100644
--- a/litellm/proxy/agent_endpoints/agent_registry.py
+++ b/litellm/proxy/agent_endpoints/agent_registry.py
@@ -654,7 +654,7 @@ class AgentRegistry:
# Create agent in DB
created_agent: Final = await agents_table(prisma_client).create(
data={**create_data, **await _managed_fields(agent, None, created_by, prisma_client)},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
return AgentResponse.model_validate(created_agent.model_dump())
@@ -717,7 +717,7 @@ class AgentRegistry:
"""
try:
existing_record: Final = await agents_table(prisma_client).find_unique(
- where={"agent_id": agent_id}, include={"identity": True}
+ where={"agent_id": agent_id}, include={"identity": True, "litellm_budget_table": True}
)
if existing_record is None:
raise Exception(f"Agent with ID {agent_id} not found")
@@ -771,7 +771,7 @@ class AgentRegistry:
"updated_by": updated_by,
"updated_at": datetime.now(timezone.utc),
},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
if patched_agent is None:
raise ValueError(f"Agent not found, passed agent_id={agent_id}")
@@ -808,7 +808,7 @@ class AgentRegistry:
# caller echoed back redacted (or omitted) rather than persisting
# the marker -- or nothing -- over the real stored credential.
existing_row: Final = await agents_table(prisma_client).find_unique(
- where={"agent_id": agent_id}, include={"identity": True}
+ where={"agent_id": agent_id}, include={"identity": True, "litellm_budget_table": True}
)
existing_litellm_params: Final = parse_agent_litellm_params(
existing_row.litellm_params if existing_row is not None else None
@@ -877,7 +877,7 @@ class AgentRegistry:
prisma_client,
),
},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
if updated_agent is None:
@@ -900,7 +900,7 @@ class AgentRegistry:
try:
agents_from_db: Final = await agents_table(prisma_client).find_many(
order={"created_at": "desc"},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
agents: Final[list[dict[str, object]]] = []
diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
index 17d988127ec..ef437b029fe 100644
--- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py
+++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
@@ -162,6 +162,8 @@ async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore |
raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists"))
return
if not agent.identity_managed:
+ if agent.litellm_budget_table is not None:
+ auth.billing_agent_policy = agent
return
if auth.jwt_claims and auth.managed_agent_context is None:
raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity"))
@@ -202,6 +204,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)])
@@ -234,19 +254,47 @@ 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
auth.invoked_agent_policy = effective
- if auth.agent_id is None and effective.identity_managed:
+ if (
+ billable
+ and auth.agent_id is None
+ and (effective.identity_managed or effective.litellm_budget_table is not None)
+ ):
auth.billing_agent_policy = effective
- raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0
+ pricing: Final = effective.litellm_params or MappingProxyType({})
+ fixed_fee: Final = pricing.get("cost_per_query")
+ billing_policy: Final = auth.billing_agent_policy
+ bounded: Final = (
+ billing_policy is not None
+ and billing_policy.litellm_budget_table is not None
+ and billing_policy.litellm_budget_table.max_budget is not None
+ )
try:
- fee: Final = _INVOCATION_COST.validate_python(raw_fee)
+ fee: Final = _INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0)
+ unbounded_token_price: Final = (
+ billable
+ and bounded
+ and fixed_fee is None
+ and any(
+ _INVOCATION_COST.validate_python(pricing[field]) > 0
+ for field in ("input_cost_per_token", "output_cost_per_token")
+ if pricing.get(field) is not None
+ )
+ )
except ValidationError:
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent invocation price is invalid")
)
+ if unbounded_token_price:
+ raise_identity_failure(
+ AgentIdentityFailure(
+ code="policy_unavailable",
+ message="Budgeted token-priced agent invocations require a fixed cost_per_query before execution",
+ )
+ )
auth.agent_invocation_cost = fee
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 3c8163a8838..ac0ebac7865 100644
--- a/litellm/proxy/agent_endpoints/identity_store.py
+++ b/litellm/proxy/agent_endpoints/identity_store.py
@@ -70,6 +70,7 @@ class AgentIdentityStore:
include: Final[LiteLLM_AgentsTableInclude] = {
"identity": True,
"object_permission": True,
+ "litellm_budget_table": True,
}
row: Final = await self.agents.table.find_unique(where=where, include=include)
if row is None:
diff --git a/litellm/proxy/agent_endpoints/managed_identity.py b/litellm/proxy/agent_endpoints/managed_identity.py
index abab21901ee..87041e92459 100644
--- a/litellm/proxy/agent_endpoints/managed_identity.py
+++ b/litellm/proxy/agent_endpoints/managed_identity.py
@@ -7,8 +7,10 @@ from fastapi import HTTPException
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly
+from litellm.proxy.common_utils.timezone_utils import budget_duration_error, get_budget_reset_time
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import (
+ AgentBudgetConfig,
AgentExecutionMode,
AgentIdentityBinding,
AgentIdentityFailure,
@@ -57,12 +59,29 @@ class IdentityHistoryWrite(TypedDict):
create: ReadOnly[IdentityHistoryEntry]
+class BudgetFields(TypedDict, total=False):
+ max_budget: ReadOnly[float]
+ budget_duration: ReadOnly[str | None]
+ budget_reset_at: ReadOnly[datetime | None]
+ updated_by: ReadOnly[str]
+ created_by: ReadOnly[str]
+
+
+class BudgetRelationWrite(TypedDict, total=False):
+ create: ReadOnly[BudgetFields]
+ update: ReadOnly[BudgetFields]
+ disconnect: ReadOnly[bool]
+
+
class ManagedWriteFields(TypedDict, total=False):
enabled: ReadOnly[bool]
execution_mode: ReadOnly[AgentExecutionMode]
identity_managed: ReadOnly[bool]
identity: ReadOnly[IdentityRelationWrite]
retired_identities: ReadOnly[IdentityHistoryWrite]
+ litellm_budget_table: ReadOnly[BudgetRelationWrite]
+ spend_window: ReadOnly[datetime | None]
+ spend: ReadOnly[float]
def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn:
@@ -109,14 +128,18 @@ def managed_write_fields(
return failure
empty: Final[ManagedWriteFields] = {}
identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty
+ budget_fields: Final = (
+ _budget_write(incoming["budget"], existing, updated_by) if "budget" in incoming else empty
+ )
result: Final[ManagedWriteFields] = {
**({"enabled": incoming["enabled"] is True} if "enabled" in incoming else {}),
**({"execution_mode": mode} if "execution_mode" in incoming else {}),
+ **budget_fields,
**identity_fields,
}
return result
except (ValidationError, ValueError) as exc:
- return AgentIdentityFailure(message=f"Invalid agent identity configuration: {exc}")
+ return AgentIdentityFailure(message=f"Invalid agent identity or budget configuration: {exc}")
def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields:
@@ -165,6 +188,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 3ec430332ee..24b4a8047ec 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -1073,6 +1073,10 @@ async def common_checks(
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
+ if valid_token is not None and not skip_all_budget_checks:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import check_agent_budget
+
+ await check_agent_budget(valid_token)
await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router)
await _check_agent_caller_model_access(
model=_model,
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index 5194f62cf78..278a48911d4 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -2988,6 +2988,7 @@ async def _run_centralized_common_checks(
request=request,
llm_router=llm_router,
team_id=user_api_key_auth_obj.team_id,
+ agent_invocation_cost=user_api_key_auth_obj.agent_invocation_cost,
)
# Pin the metadata variable name (litellm_metadata vs metadata) before
@@ -3161,7 +3162,10 @@ def _should_skip_budget_checks(
request: Request | None,
llm_router: Any | None,
team_id: str | None = None,
+ agent_invocation_cost: float | None = None,
) -> bool:
+ if agent_invocation_cost is not None and agent_invocation_cost > 0:
+ return False
model: Final = _get_model_from_request_context(
request_data=request_data,
route=route,
@@ -3256,7 +3260,8 @@ async def _authorize_authenticated_request(
user_api_key_auth_obj,
target_name,
store,
- billable=request_data.get("method")
+ billable=request.method == "POST"
+ and request_data.get("method")
in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"),
)
await _run_centralized_common_checks(
diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py
index 8a82e253c5c..c512d782db3 100644
--- a/litellm/proxy/common_utils/registry_read_through.py
+++ b/litellm/proxy/common_utils/registry_read_through.py
@@ -177,6 +177,7 @@ async def _resync_agents(agent_id_or_name: str) -> bool:
include_permission: Final[LiteLLM_AgentsTableInclude] = {
"object_permission": True,
"identity": True,
+ "litellm_budget_table": True,
}
async with AGENT_RECONCILE_LOCK:
if _agent_from_registry(agent_id_or_name) is not None:
diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py
index b35b876b475..4c875936b31 100644
--- a/litellm/proxy/common_utils/reset_budget_job.py
+++ b/litellm/proxy/common_utils/reset_budget_job.py
@@ -841,6 +841,7 @@ class ResetBudgetJob:
_queue_budget_linked_resets(uow.projects, cascade, extra=_SPENT_ROWS_WHERE)
_queue_enduser_resets(uow.endusers, cascade)
for budget_id, budget_reset_at in cascade.budget_resets:
+ uow.agents.queue_window_reset(budget_id, budget_reset_at, cascade.rollover_caps.get(budget_id))
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)
async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None:
diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py
index 72553e82283..c88ef2544d6 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(
@@ -1351,15 +1355,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,
)
)
@@ -2341,7 +2349,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
@@ -2914,13 +2926,14 @@ class DBSpendUpdateWriter:
if prisma_client is None:
verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.")
return
- if payload["agent_id"] is None:
+ charged_agent_id: Final = payload.get("billing_agent_id", payload["agent_id"])
+ if charged_agent_id is None:
return
payload_with_agent_id: Final = cast(
SpendLogsPayload,
{
**payload,
- "agent_id": payload["agent_id"],
+ "agent_id": charged_agent_id,
},
)
base_daily_transaction: Final = await self._common_add_spend_log_transaction_to_daily_transaction(
@@ -2929,8 +2942,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 6a2ec120060..c914e04e779 100644
--- a/litellm/proxy/hooks/proxy_track_cost_callback.py
+++ b/litellm/proxy/hooks/proxy_track_cost_callback.py
@@ -2,6 +2,7 @@ import asyncio
import traceback
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime
+from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
import litellm
@@ -387,6 +388,8 @@ class _ProxyDBLogger(CustomLogger):
tags=tags,
response_cost=response_cost,
),
+ billing_agent_id=metadata.get("billing_agent_id"),
+ billing_agent_counter_key=metadata.get("billing_agent_counter_key"),
)
if not charged:
return
@@ -691,6 +694,9 @@ class _IncrementSpendCounters(Protocol):
tags: list[str] | None = None,
request_started_at: datetime | None = None,
model_access_groups: Sequence[str] | None = None,
+ project_id: str | None = None,
+ billing_agent_id: str | None = None,
+ billing_agent_counter_key: str | None = None,
) -> None: ...
@@ -712,6 +718,8 @@ async def _update_database_and_spend_counters(
model_access_groups: Sequence[str] | None = None,
project_id: str | None = None,
update_cache_read_keys: Sequence[str] = (),
+ billing_agent_id: str | None = None,
+ billing_agent_counter_key: str | None = None,
) -> bool:
"""The reservation is reconciled before the spend is persisted, from its own read. One spend counter batch then
spans the database write and the counter update, so the post-call counters are read with a single MGET after the
@@ -754,6 +762,8 @@ async def _update_database_and_spend_counters(
model_access_groups=model_access_groups,
project_id=project_id,
update_cache_read_keys=update_cache_read_keys,
+ billing_agent_id=billing_agent_id,
+ billing_agent_counter_key=billing_agent_counter_key,
)
@@ -775,6 +785,8 @@ async def _update_database_and_spend_counters_in_batch(
model_access_groups: Sequence[str] | None,
project_id: str | None,
update_cache_read_keys: Sequence[str],
+ billing_agent_id: str | None,
+ billing_agent_counter_key: str | None,
) -> bool:
from litellm.proxy.proxy_server import arm_update_cache_read
@@ -823,6 +835,16 @@ async def _update_database_and_spend_counters_in_batch(
request_started_at=start_time,
model_access_groups=model_access_groups,
project_id=project_id,
+ **(
+ MappingProxyType({"billing_agent_id": billing_agent_id})
+ if billing_agent_id is not None
+ else MappingProxyType({})
+ ),
+ **(
+ MappingProxyType({"billing_agent_counter_key": billing_agent_counter_key})
+ if billing_agent_counter_key is not None
+ else MappingProxyType({})
+ ),
)
except Exception:
if budget_reservation is not None:
diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py
index 4188e8ad58a..27d22920591 100644
--- a/litellm/proxy/litellm_pre_call_utils.py
+++ b/litellm/proxy/litellm_pre_call_utils.py
@@ -1672,6 +1672,11 @@ class LiteLLMProxyRequestSetup:
"actor_agent_id": user_api_key_dict.agent_id,
"target_agent_id": user_api_key_dict.invoked_agent_id,
"billing_agent_id": user_api_key_dict.agent_id or user_api_key_dict.invoked_agent_id,
+ "billing_agent_counter_key": (
+ user_api_key_dict.billing_agent_policy.budget_counter_key
+ if user_api_key_dict.billing_agent_policy is not None
+ else None
+ ),
"agent_execution_mode": managed_context.mode if managed_context else None,
"verified_human_user_id": managed_context.user_id if managed_context else None,
}
diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py
index e16ea4a812e..866a747690e 100644
--- a/litellm/proxy/management_endpoints/budget_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py
@@ -28,10 +28,23 @@ from litellm.proxy.management_endpoints.common_utils import (
)
from litellm.proxy.utils import jsonify_object
from litellm.repositories.budget_repository import BudgetRepository
+from litellm.repositories.table_repositories import AgentsRepository
router: Final = APIRouter()
+async def _require_unlinked_agent_budget(budget_id: str, client: object) -> None:
+ from prisma.types import LiteLLM_AgentsTableWhereInput
+
+ where: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id}
+ agent: Final = await AgentsRepository(client, use_writer=True).table.find_first(where=where)
+ if agent is not None:
+ raise HTTPException(
+ status_code=409,
+ detail=f"Manage this agent's budget through PATCH /v1/agents/{agent.agent_id}",
+ )
+
+
@router.post(
"/budget/new",
tags=["budget management"],
@@ -194,6 +207,7 @@ async def update_budget(
}
)
+ await _require_unlinked_agent_budget(budget_obj.budget_id, prisma_client)
response: Final = await BudgetRepository(prisma_client).table.update(
where={"budget_id": budget_obj.budget_id},
data=budget_obj_jsonified,
@@ -357,6 +371,7 @@ async def delete_budget(
detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"},
)
+ await _require_unlinked_agent_budget(data.id, prisma_client)
response: Final = await BudgetRepository(prisma_client).table.delete(where={"budget_id": data.id})
return response
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 0e151199f41..a1e04f5c1c0 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -3022,6 +3022,8 @@ async def increment_spend_counters(
request_started_at: datetime | None = None,
model_access_groups: Sequence[str] | None = None,
project_id: str | None = None,
+ billing_agent_id: str | None = None,
+ billing_agent_counter_key: str | None = None,
):
"""
Atomically increment spend counters for budget enforcement.
@@ -3044,6 +3046,8 @@ async def increment_spend_counters(
tags=tags,
model_access_groups=model_access_groups,
project_id=project_id,
+ billing_agent_id=billing_agent_id,
+ billing_agent_counter_key=billing_agent_counter_key,
),
):
await _increment_spend_counters_batched(
@@ -3058,6 +3062,8 @@ async def increment_spend_counters(
request_started_at=request_started_at,
model_access_groups=model_access_groups,
project_id=project_id,
+ billing_agent_id=billing_agent_id,
+ billing_agent_counter_key=billing_agent_counter_key,
)
@@ -3073,6 +3079,8 @@ async def _increment_spend_counters_batched(
request_started_at: datetime | None,
model_access_groups: Sequence[str] | None,
project_id: str | None = None,
+ billing_agent_id: str | None = None,
+ billing_agent_counter_key: str | None = None,
):
"""Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET, and
the reconcile adjustments go out in the same INCRBYFLOAT pipeline as the counter increments."""
@@ -3248,9 +3256,16 @@ async def _increment_spend_counters_batched(
),
)
+ async def _agent_scope(agent_id: str) -> tuple[PendingSpendIncrement, ...]:
+ counter_key: Final = billing_agent_counter_key or f"spend:agent:{agent_id}"
+ if counter_key in reserved_counter_keys:
+ return ()
+ return (await _prepare_spend_counter_increment(counter_key=counter_key, source_cache_key=[], increment=cost),)
+
scope_coros: Final = tuple(
coro
for coro in (
+ _agent_scope(billing_agent_id) if billing_agent_id is not None else None,
_key_scope(token) if token is not None else None,
_team_scope(team_id) if team_id is not None else None,
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma
index adfe2a0eee7..6823deadbd3 100644
--- a/litellm/proxy/schema.prisma
+++ b/litellm/proxy/schema.prisma
@@ -36,6 +36,7 @@ model LiteLLM_BudgetTable {
model_access_groups LiteLLM_ModelAccessGroupBudgetTable[] // multiple model access groups can have the same budget
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
+ agents LiteLLM_AgentsTable[]
}
// Models on proxy
@@ -83,6 +84,9 @@ model LiteLLM_AgentsTable {
execution_mode String @default("autonomous")
identity LiteLLM_AgentIdentity?
retired_identities LiteLLM_RetiredAgentIdentity[]
+ budget_id String? @unique
+ spend_window DateTime?
+ litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py
index c094e91c6c0..8884a94f4b2 100644
--- a/litellm/proxy/spend_tracking/budget_reservation.py
+++ b/litellm/proxy/spend_tracking/budget_reservation.py
@@ -67,6 +67,7 @@ _COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = {
"Model access group": Litellm_EntityType.MODEL_ACCESS_GROUP.value,
"Organization": Litellm_EntityType.ORGANIZATION.value,
"Project": Litellm_EntityType.PROJECT.value,
+ "Agent": Litellm_EntityType.AGENT.value,
}
@@ -265,7 +266,8 @@ async def reserve_budget_for_request(
return None
if _is_unbilled_route(route):
return None
- if get_model_from_request(request_body, route, llm_router=llm_router) is None:
+ invocation_cost: Final = valid_token.agent_invocation_cost
+ if invocation_cost is None and get_model_from_request(request_body, route, llm_router=llm_router) is None:
return None
counters: Final = await _get_budget_counters(
@@ -283,18 +285,26 @@ async def reserve_budget_for_request(
if not counters:
return None
- input_token_counts: Final = await count_request_input_tokens(
- request_body=request_body,
- route=route,
- llm_router=llm_router,
- raw_body=raw_body,
+ input_token_counts: Final = (
+ await count_request_input_tokens(
+ request_body=request_body,
+ route=route,
+ llm_router=llm_router,
+ raw_body=raw_body,
+ )
+ if invocation_cost is None
+ else MappingProxyType({})
)
- reservation_cost = estimate_request_max_cost(
- request_body=request_body,
- route=route,
- llm_router=llm_router,
- input_token_counts=input_token_counts,
+ reservation_cost = (
+ invocation_cost
+ if invocation_cost is not None
+ else estimate_request_max_cost(
+ request_body=request_body,
+ route=route,
+ llm_router=llm_router,
+ input_token_counts=input_token_counts,
+ )
)
# estimate_request_max_cost still returns None when the model is unknown
# to the cost map (no token-priced cost fields, e.g. image/audio routes).
@@ -316,7 +326,7 @@ async def reserve_budget_for_request(
reservation_cost=reservation_cost,
fail_closed_budget_enforcement=fail_closed_budget_enforcement,
)
- except Exception:
+ except (asyncio.CancelledError, Exception):
await _release_applied_entries_best_effort(
entries=applied_entries,
default_reserved_cost=reservation_cost,
@@ -326,11 +336,15 @@ async def reserve_budget_for_request(
if not applied_entries:
return None
- input_cost: Final = estimate_request_input_cost(
- request_body=request_body,
- route=route,
- llm_router=llm_router,
- input_token_counts=input_token_counts,
+ input_cost: Final = (
+ 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,
@@ -490,7 +504,23 @@ async def _get_budget_counters(
end_user_object: object = None,
apply_user_budget_to_team_keys: bool = False,
) -> list[_BudgetCounter]:
- counters: Final[list[_BudgetCounter]] = []
+ agent: Final = valid_token.billing_agent_policy
+ counters: Final[list[_BudgetCounter]] = (
+ [
+ _BudgetCounter(
+ counter_key=agent.budget_counter_key,
+ source_cache_key=None,
+ max_budget=agent.litellm_budget_table.max_budget,
+ fallback_spend=agent.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 ae24331c236..b3d0e8646e8 100644
--- a/litellm/proxy/spend_tracking/spend_counter_batch.py
+++ b/litellm/proxy/spend_tracking/spend_counter_batch.py
@@ -204,7 +204,13 @@ def _iter_entity_counter_keys(
def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]:
- return frozenset(
+ billing_agent: Final = token.billing_agent_policy
+ charged_agent_id: Final = billing_agent.agent_id if billing_agent is not None else token.agent_id
+ agent_keys: Final = (
+ frozenset((billing_agent.budget_counter_key if billing_agent is not None else f"spend:agent:{charged_agent_id}",))
+ if charged_agent_id is not None else frozenset()
+ )
+ return agent_keys | frozenset(
_iter_entity_counter_keys(
token=token.token,
team_id=token.team_id,
@@ -225,6 +231,8 @@ def post_call_counter_keys(
tags: Sequence[object] | None,
model_access_groups: Sequence[object] | None,
project_id: str | None = None,
+ billing_agent_id: str | None = None,
+ billing_agent_counter_key: str | None = None,
) -> frozenset[str]:
"""Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read."""
entity_keys: Final = frozenset(
@@ -243,7 +251,13 @@ def post_call_counter_keys(
for group in model_access_groups or ()
if group and isinstance(group, str)
)
- return entity_keys | tag_keys | group_keys
+ agent_key: Final = billing_agent_counter_key or (f"spend:agent:{billing_agent_id}" if billing_agent_id is not None else None)
+ return (
+ entity_keys
+ | tag_keys
+ | group_keys
+ | (frozenset((agent_key,)) if agent_key else frozenset())
+ )
def bind_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> None:
diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py
index c42301a9316..d9b25069747 100644
--- a/litellm/repositories/prisma_protocols.py
+++ b/litellm/repositories/prisma_protocols.py
@@ -171,4 +171,7 @@ class PrismaBatch(Protocol):
@property
def litellm_projecttable(self) -> BatchTable: ...
+ @property
+ def litellm_agentstable(self) -> BatchTable: ...
+
async def commit(self) -> None: ...
diff --git a/litellm/repositories/unit_of_work.py b/litellm/repositories/unit_of_work.py
index c09e5eb75d4..ced458d2102 100644
--- a/litellm/repositories/unit_of_work.py
+++ b/litellm/repositories/unit_of_work.py
@@ -19,10 +19,13 @@ from collections.abc import AsyncGenerator, Callable, Mapping
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
-from typing import Final
+from typing import TYPE_CHECKING, Final
from litellm.repositories.prisma_protocols import BatchTable, PrismaBatch
+if TYPE_CHECKING:
+ from prisma.types import LiteLLM_AgentsTableUpdateManyMutationInput, LiteLLM_AgentsTableWhereInput
+
def _spend_reset_data(budget_reset_at: datetime | None, spend_decrement: float) -> Mapping[str, object]:
spend: Final[object] = {"decrement": spend_decrement} # mutable-ok: prisma update payload must be a dict
@@ -78,6 +81,26 @@ class LinkedSpendResetWrites:
)
+@dataclass(frozen=True, slots=True)
+class AgentSpendResetWrites:
+ table: BatchTable
+
+ def queue_window_reset(self, budget_id: str, window: datetime, rollover_cap: float | None) -> None:
+ zero: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = {"spend": 0.0, "spend_window": window}
+ if rollover_cap is None:
+ all_agents: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id}
+ self.table.update_many(where=all_agents, data=zero)
+ return
+ below: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id, "spend": {"lte": rollover_cap}}
+ above: Final[LiteLLM_AgentsTableWhereInput] = {"budget_id": budget_id, "spend": {"gt": rollover_cap}}
+ remainder: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = {
+ "spend": {"decrement": rollover_cap},
+ "spend_window": window,
+ }
+ self.table.update_many(where=below, data=zero)
+ self.table.update_many(where=above, data=remainder)
+
+
@dataclass(frozen=True, slots=True)
class BudgetWindowWrites:
table: BatchTable
@@ -110,6 +133,7 @@ class BudgetCascadeUnitOfWork:
tags: LinkedSpendResetWrites
model_access_groups: LinkedSpendResetWrites
projects: LinkedSpendResetWrites
+ agents: AgentSpendResetWrites
endusers: LinkedSpendResetWrites
budgets: BudgetWindowWrites
@@ -137,6 +161,7 @@ async def budget_cascade_unit_of_work(
tags=LinkedSpendResetWrites(table=batch.litellm_tagtable),
model_access_groups=LinkedSpendResetWrites(table=batch.litellm_modelaccessgroupbudgettable),
projects=LinkedSpendResetWrites(table=batch.litellm_projecttable),
+ agents=AgentSpendResetWrites(table=batch.litellm_agentstable),
endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable),
budgets=BudgetWindowWrites(table=batch.litellm_budgettable),
)
diff --git a/litellm/types/agents.py b/litellm/types/agents.py
index 94adb9f7c4a..964ef94d227 100644
--- a/litellm/types/agents.py
+++ b/litellm/types/agents.py
@@ -1,5 +1,5 @@
from collections.abc import Mapping, Sequence
-from datetime import datetime
+from datetime import datetime, timezone
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, TypeAlias
from urllib.parse import urlsplit
@@ -8,6 +8,8 @@ from typing_extensions import ReadOnly, Required, TypedDict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
from litellm.types.proxy.agent_identity import (
+ AgentBudgetConfig,
+ AgentBudgetState,
AgentExecutionMode,
AgentIdentityBinding,
EntraIdentityConfig,
@@ -15,6 +17,7 @@ from litellm.types.proxy.agent_identity import (
if TYPE_CHECKING:
from a2a.types import SendMessageResponse
+ from prisma.types import LiteLLM_AgentsTableWhereInput
# AgentProvider
@@ -253,6 +256,7 @@ class AgentKillSwitchResult(BaseModel):
class AgentConfig(TypedDict, total=False):
+ budget: ReadOnly[AgentBudgetConfig | None]
identity: ReadOnly[EntraIdentityConfig | None]
enabled: ReadOnly[bool]
execution_mode: ReadOnly[AgentExecutionMode]
@@ -271,6 +275,7 @@ class AgentConfig(TypedDict, total=False):
class PatchAgentRequest(TypedDict, total=False):
+ budget: ReadOnly[AgentBudgetConfig | None]
identity: ReadOnly[EntraIdentityConfig | None]
enabled: ReadOnly[bool]
execution_mode: ReadOnly[AgentExecutionMode]
@@ -311,7 +316,30 @@ class AgentKeySummary(BaseModel):
key_name: str | None = None
+def agent_budget_counter_key(agent_id: str, reset_at: datetime | None) -> str:
+ if reset_at is None:
+ return f"spend:agent:{agent_id}"
+ aware: Final = reset_at if reset_at.tzinfo is not None else reset_at.replace(tzinfo=timezone.utc)
+ window: Final = aware.astimezone(timezone.utc).strftime("%Y%m%dT%H%M%S.%fZ")
+ return f"spend:agent_window:{window}:{agent_id}"
+
+
+def agent_spend_filter(counter_key: str) -> "LiteLLM_AgentsTableWhereInput":
+ if counter_key.startswith("spend:agent_window:"):
+ _, _, raw_window, agent_id = counter_key.split(":", 3)
+ window: Final = datetime.strptime(raw_window, "%Y%m%dT%H%M%S.%fZ").replace(tzinfo=timezone.utc)
+ windowed: Final[LiteLLM_AgentsTableWhereInput] = {"agent_id": agent_id, "spend_window": window}
+ return windowed
+ cumulative: Final[LiteLLM_AgentsTableWhereInput] = {
+ "agent_id": counter_key.removeprefix("spend:agent:"),
+ "spend_window": None,
+ }
+ return cumulative
+
+
class AgentResponse(BaseModel):
+ budget_id: str | None = None
+ litellm_budget_table: AgentBudgetState | None = None
identity: AgentIdentityBinding | None = None
identity_managed: bool = False
enabled: bool = True
@@ -338,6 +366,12 @@ class AgentResponse(BaseModel):
created_by: str | None = None
updated_by: str | None = None
+ @property
+ def budget_counter_key(self) -> str:
+ return agent_budget_counter_key(
+ self.agent_id, self.litellm_budget_table.budget_reset_at if self.litellm_budget_table else None
+ )
+
class ListAgentsResponse(BaseModel):
agents: list[AgentResponse]
diff --git a/litellm/types/proxy/agent_identity.py b/litellm/types/proxy/agent_identity.py
index a7fe0be37e1..2dc43976e0f 100644
--- a/litellm/types/proxy/agent_identity.py
+++ b/litellm/types/proxy/agent_identity.py
@@ -46,6 +46,22 @@ class AgentIdentityBinding(BaseModel):
last_authenticated_at: datetime | None = None
+class AgentBudgetConfig(BaseModel):
+ model_config = ConfigDict(frozen=True, extra="forbid")
+
+ max_budget: float = Field(ge=0, allow_inf_nan=False)
+ budget_duration: str | None = None
+
+
+class AgentBudgetState(BaseModel):
+ model_config = ConfigDict(frozen=True)
+
+ budget_id: str
+ max_budget: float | None = None
+ budget_duration: str | None = None
+ budget_reset_at: datetime | None = None
+
+
class AgentSubject(BaseModel):
model_config = ConfigDict(frozen=True)
diff --git a/schema.prisma b/schema.prisma
index adfe2a0eee7..6823deadbd3 100644
--- a/schema.prisma
+++ b/schema.prisma
@@ -36,6 +36,7 @@ model LiteLLM_BudgetTable {
model_access_groups LiteLLM_ModelAccessGroupBudgetTable[] // multiple model access groups can have the same budget
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
+ agents LiteLLM_AgentsTable[]
}
// Models on proxy
@@ -83,6 +84,9 @@ model LiteLLM_AgentsTable {
execution_mode String @default("autonomous")
identity LiteLLM_AgentIdentity?
retired_identities LiteLLM_RetiredAgentIdentity[]
+ budget_id String? @unique
+ spend_window DateTime?
+ litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
index eee985f0aca..a2c9fca3bd0 100644
--- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
+++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
@@ -98,11 +98,45 @@ def test_caller_cannot_construct_trusted_subject_or_policy() -> None:
assert auth.agent_invocation_cost is None
+@pytest.mark.asyncio
+async def test_agent_budget_accumulates_across_credentials_and_denies_the_next_admission(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ import litellm
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import check_agent_budget
+
+ counters: Final = DualCache()
+ counters.set_cache("spend:agent: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))
+@pytest.mark.parametrize("billable", (True, False))
async def test_invocation_prepares_target_fee_for_the_correct_agent(
monkeypatch: pytest.MonkeyPatch,
autonomous: bool,
+ billable: bool,
) -> None:
from unittest.mock import AsyncMock, MagicMock
@@ -129,11 +163,14 @@ async def test_invocation_prepares_target_fee_for_the_correct_agent(
caller: Final = agent(agent_id="caller", object_permission=permission.model_dump())
auth.managed_agent_policy = caller
auth.billing_agent_policy = caller
- await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
- assert auth.agent_invocation_cost == pytest.approx(0.25)
+ await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable)
+ assert auth.agent_invocation_cost == pytest.approx(0.25 if billable else 0.0)
assert auth.invoked_agent_id == "agent"
- assert auth.billing_agent_policy is not None
- assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent")
+ if autonomous or billable:
+ assert auth.billing_agent_policy is not None
+ assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent")
+ else:
+ assert auth.billing_agent_policy is None
@pytest.mark.asyncio
@@ -167,6 +204,8 @@ async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None:
("/a2a/expensive", {"model": "a2a/cheap"}, "expensive"),
("/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"),
("/v1/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"),
+ ("/a2a/agent", {"model": "a2a/nonexistent"}, "agent"),
+ ("/v1/a2a/agent/message/send", {"model": "a2a/other"}, "agent"),
("/v1/a2a/agent/", {}, "agent"),
("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"),
("/v1/chat/completions", {"model": "a2a/"}, None),
@@ -178,6 +217,23 @@ def test_invocation_routes_resolve_the_same_target(route: str, body: dict[str, o
assert invocation_target(route, body) == expected
+@pytest.mark.asyncio
+@pytest.mark.parametrize("managed", [True, False])
+async def test_aggregate_budget_applies_to_entra_tokens_and_unbound_agent_keys(managed: bool) -> None:
+ policy: Final = agent(
+ identity_managed=managed, identity=BINDING if managed else None,
+ litellm_budget_table={"budget_id": "budget", "max_budget": 0.5},
+ )
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
+ auth: Final = UserAPIKeyAuth(agent_id="agent")
+ if managed:
+ auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ assert auth.billing_agent_policy == policy
+ assert auth.managed_agent_policy == (policy if managed else None)
+
+
@pytest.mark.asyncio
async def test_agent_admission_database_outage_fails_closed() -> None:
database: Final = MagicMock()
@@ -508,6 +564,95 @@ async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_
assert auth.requires_fresh_policy is True
+@pytest.mark.asyncio
+async def test_new_budget_window_isolated_from_inflight_previous_window_charge(monkeypatch: pytest.MonkeyPatch) -> None:
+ import litellm
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import check_agent_budget
+
+ cache: Final = DualCache()
+ old_key: Final = "spend:agent_window:20260101T000000.000000Z:agent"
+ new_key: Final = "spend:agent_window:20260102T000000.000000Z:agent"
+ cache.set_cache("spend:agent:agent", 10.0)
+ cache.set_cache(old_key, 10.0)
+ cache.set_cache(new_key, 0.0)
+ monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
+ auth: Final = UserAPIKeyAuth(agent_id="agent")
+ auth.billing_agent_policy = agent(spend=0, litellm_budget_table={
+ "budget_id": "budget", "max_budget": 1.0, "budget_reset_at": "2026-01-02T00:00:00Z",
+ })
+ await check_agent_budget(auth)
+ await proxy_server.increment_spend_counters(
+ token=None, team_id=None, user_id=None, response_cost=0.5,
+ billing_agent_id="agent", billing_agent_counter_key=old_key,
+ )
+ await check_agent_budget(auth)
+ await proxy_server.increment_spend_counters(
+ token=None, team_id=None, user_id=None, response_cost=1.1,
+ billing_agent_id="agent", billing_agent_counter_key=new_key,
+ )
+ with pytest.raises(litellm.BudgetExceededError):
+ await check_agent_budget(auth)
+ assert await cache.async_get_cache(old_key) == 10.5
+ assert await cache.async_get_cache(new_key) == 1.1
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("autonomous", [False, True])
+@pytest.mark.parametrize(
+ "pricing,billable,bounded,rejected,fee",
+ [
+ ({"input_cost_per_token": 0.01}, True, True, True, None),
+ ({"output_cost_per_token": 0.01}, True, True, True, None),
+ ({"input_cost_per_token": 0.01}, False, True, False, 0.0),
+ ({"input_cost_per_token": 0.01}, True, False, False, 0.0),
+ ({"input_cost_per_token": 0.0, "output_cost_per_token": 0.0}, True, True, False, 0.0),
+ ({"cost_per_query": 0.25, "output_cost_per_token": 0.01}, True, True, False, 0.25),
+ ({"cost_per_query": 0.0, "output_cost_per_token": 0.01}, True, True, False, 0.0),
+ ],
+)
+async def test_budgeted_invocation_requires_a_bounded_price(
+ monkeypatch: pytest.MonkeyPatch,
+ autonomous: bool,
+ pricing: dict[str, float],
+ billable: bool,
+ bounded: bool,
+ rejected: bool,
+ fee: float | None,
+) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+
+ budget: Final = {"budget_id": "budget", "max_budget": 1.0} if bounded else None
+ target_budget: Final = {"budget_id": "target-budget", "max_budget": 1.0} if bounded != autonomous else None
+ target: Final = agent(litellm_params=pricing, litellm_budget_table=target_budget)
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(target)
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="invoke-grant", agents=["agent"])
+ auth: Final = UserAPIKeyAuth(agent_id="caller" if autonomous else None, object_permission=permission)
+ if autonomous:
+ caller: Final = agent(
+ agent_id="caller", object_permission=permission.model_dump(), litellm_budget_table=budget
+ )
+ auth.managed_agent_policy = caller
+ auth.billing_agent_policy = caller
+ if rejected:
+ with pytest.raises(HTTPException) as exc:
+ await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable)
+ assert exc.value.status_code == 503
+ assert "cost_per_query" in str(exc.value.detail)
+ else:
+ await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database), billable=billable)
+ assert auth.agent_invocation_cost == fee
+
+
async def test_jwt_delegation_verification_is_consumed_once_and_cannot_be_supplied_by_a_caller(
monkeypatch: pytest.MonkeyPatch,
) -> None:
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
index 7663f1d30e6..9ed9859717d 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
@@ -1451,7 +1451,7 @@ async def test_agent_listing_preserves_stored_identity_bindings(bound: bool) ->
assert response.identity is None
client.db.litellm_agentstable.find_many.assert_awaited_once_with(
order={"created_at": "desc"},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py
index 17f3cdb52f5..424330e078e 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 353249dddf0..b0fd544c718 100644
--- a/tests/test_litellm/proxy/auth/test_auth_checks.py
+++ b/tests/test_litellm/proxy/auth/test_auth_checks.py
@@ -10276,3 +10276,30 @@ async def test_authoritative_group_grants_propagate_policy_outages(
await _get_agent_ids_from_access_groups(["group"], check_db_only=True)
else:
assert await _get_agent_ids_from_access_groups(["group"]) == []
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("skip", [False, True])
+async def test_explicit_budget_skip_applies_to_agent_budget(skip: bool, monkeypatch: pytest.MonkeyPatch) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.proxy.auth.auth_checks import common_checks
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import AgentBudgetState
+
+ auth: Final = UserAPIKeyAuth()
+ auth.billing_agent_policy = AgentResponse(
+ agent_id="agent", agent_name="Agent", agent_card_params={},
+ litellm_budget_table=AgentBudgetState(budget_id="budget", max_budget=0),
+ )
+ monkeypatch.setattr(proxy_server, "get_current_spend", AsyncMock(return_value=0))
+ checks: Final = common_checks(
+ request_body={"model": "gpt-4"}, team_object=None, user_object=None,
+ end_user_object=None, global_proxy_spend=None, general_settings={},
+ route="/v1/chat/completions", llm_router=None, proxy_logging_obj=MagicMock(),
+ valid_token=auth, request=MagicMock(spec=Request), skip_budget_checks=skip,
+ )
+ if skip:
+ assert await checks is True
+ else:
+ with pytest.raises(litellm.BudgetExceededError, match="Agent budget exceeded"):
+ await checks
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index ef6832ef77b..9d0bc1c3475 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -9688,3 +9688,121 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon
)
assert admitted.authenticated_by_custom_auth is False
assert admitted.via_virtual_key is True
+
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "http_method,route,body,billed",
+ [
+ ("GET", "/a2a/agent/.well-known/agent-card.json", {}, False),
+ ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "tasks/get", "params": {"id": "t"}}, False),
+ ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/send", "params": {}}, True),
+ ("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/stream", "params": {}}, True),
+ ("POST", "/a2a/agent", {"method": "message/send", "model": "free-model", "params": {}}, True),
+ ("POST", "/a2a/agent", {"method": "message/stream", "model": "free-model", "params": {}}, True),
+ ],
+)
+async def test_human_agent_discovery_does_not_reserve_target_budget_but_send_and_stream_do(
+ monkeypatch: pytest.MonkeyPatch, http_method: str, route: str, body: dict, billed: bool
+) -> None:
+ from typing import Final
+
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding
+
+ target: Final = AgentResponse(
+ agent_id="agent",
+ agent_name="Agent",
+ agent_card_params={},
+ identity_managed=True,
+ execution_mode="both",
+ litellm_params={"cost_per_query": 0.25},
+ litellm_budget_table={"budget_id": "agent-budget", "max_budget": 10.0},
+ identity=AgentIdentityBinding(
+ agent_id="agent",
+ provider="microsoft_entra",
+ tenant_id="tenant",
+ client_id="client",
+ service_principal_id="principal",
+ issuer="issuer",
+ revision="current",
+ ),
+ )
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(target)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
+ for name, value in {
+ **_proxy_attrs_for_centralized_checks(),
+ "prisma_client": database,
+ "llm_router": litellm.Router(
+ model_list=[
+ {
+ "model_name": "free-model",
+ "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
+ "model_info": {"input_cost_per_token": 0, "output_cost_per_token": 0},
+ }
+ ]
+ )
+ if body.get("model")
+ else None,
+ "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
+ }.items():
+ monkeypatch.setattr(proxy_server, name, value)
+ request = _alias_request(route, body)
+ request.scope["method"] = http_method
+ auth: Final = UserAPIKeyAuth(
+ api_key="human-key",
+ user_id="human",
+ object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["agent"]),
+ )
+ with patch(
+ "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
+ new_callable=AsyncMock,
+ return_value=None,
+ ) as reserve:
+ assert await _authorize_authenticated_request(auth, request, body, route, "human-key") is None
+ assert auth.invoked_agent_id == "agent"
+ assert auth.agent_invocation_cost == pytest.approx(0.25 if billed else 0.0), (http_method, body.get("method"))
+ if billed:
+ assert auth.billing_agent_policy is not None and auth.billing_agent_policy.agent_id == "agent"
+ else:
+ assert auth.billing_agent_policy is None, (http_method, body.get("method"))
+ reserve.assert_awaited_once()
+ reserved: Final = reserve.call_args.kwargs["valid_token"]
+ assert reserved is auth and (reserved.billing_agent_policy is not None) is billed, (http_method, body.get("method"))
+
+
+@pytest.mark.parametrize("invocation_cost,skipped", [(None, True), (0.0, True), (0.25, False)])
+def test_free_model_only_waives_budgets_without_a_paid_agent_invocation(
+ invocation_cost: float | None, skipped: bool
+) -> None:
+ from typing import Final
+
+ from litellm.proxy.auth.user_api_key_auth import _should_skip_budget_checks
+
+ router: Final = litellm.Router(
+ model_list=[
+ {
+ "model_name": "free-model",
+ "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
+ "model_info": {"input_cost_per_token": 0, "output_cost_per_token": 0},
+ }
+ ]
+ )
+ assert (
+ _should_skip_budget_checks(
+ request_data={"model": "free-model"},
+ route="/chat/completions",
+ request=None,
+ llm_router=router,
+ agent_invocation_cost=invocation_cost,
+ )
+ is skipped
+ )
diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py
index 9e20386bf3d..761c740a370 100644
--- a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py
+++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py
@@ -177,7 +177,7 @@ async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_rep
assert agent.agent_id == agent_id
prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once_with(
where={"agent_id": agent_id},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
@@ -202,7 +202,7 @@ async def test_get_agent_with_read_through_recovers_agent_by_name(clean_agent_re
assert agent.agent_name == agent_name
prisma_client.db.litellm_agentstable.find_unique.assert_awaited_with(
where={"agent_name": agent_name},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
index 131db55ee01..f6b72e32028 100644
--- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
+++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py
@@ -4,22 +4,21 @@ import sys
import types
from datetime import datetime, timedelta, timezone
from datetime import time as dt_time
-from typing import Any, Dict, Final, List, Optional
+from typing import Any, Final
from unittest.mock import AsyncMock, MagicMock
import httpx
import prisma
import pytest
-
-from litellm.proxy._types import LiteLLM_VerificationToken
-from litellm.proxy.common_utils import reset_budget_job as reset_budget_job_module
from litellm.constants import (
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
RESET_BUDGET_JOB_BATCH_SIZE,
RESET_BUDGET_JOB_LOCK_TTL_SECONDS,
RESET_BUDGET_JOB_NAME,
)
+from litellm.proxy._types import LiteLLM_VerificationToken
+from litellm.proxy.common_utils import reset_budget_job as reset_budget_job_module
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob, _RowReset
from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings
@@ -29,12 +28,12 @@ class MockTable:
"""A single prisma table: records reads/writes and replays canned rows."""
def __init__(self):
- self.find_many_calls: List[Dict[str, Any]] = []
- self.update_many_calls: List[Dict[str, Any]] = []
- self._find_many_results: List[Any] = []
- self._find_many_error: Optional[tuple[int, Exception]] = None
+ self.find_many_calls: list[dict[str, Any]] = []
+ self.update_many_calls: list[dict[str, Any]] = []
+ self._find_many_results: list[Any] = []
+ self._find_many_error: tuple[int, Exception] | None = None
- def set_find_many_results(self, results: List[Any]):
+ def set_find_many_results(self, results: list[Any]):
self._find_many_results = results
def set_find_many_error(self, after_reads: int, error: Exception):
@@ -44,10 +43,10 @@ class MockTable:
async def find_many(
self,
- where: Dict[str, Any],
- order: Optional[Dict[str, str]] = None,
- take: Optional[int] = None,
- ) -> List[Any]:
+ where: dict[str, Any],
+ order: dict[str, str] | None = None,
+ take: int | None = None,
+ ) -> list[Any]:
"""Replays canned rows, honouring the keyset cursor + ``take`` a paged
caller relies on: without that a paged walk never advances and the
test would hang instead of failing."""
@@ -63,7 +62,7 @@ class MockTable:
rows.sort(key=lambda row: getattr(row, field, ""), reverse=direction == "desc")
return rows[:take] if take is not None else rows
- async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]:
+ async def update_many(self, where: dict[str, Any], data: dict[str, Any]) -> dict[str, Any]:
self.update_many_calls.append({"where": where, "data": data})
return {"count": 1}
@@ -77,7 +76,7 @@ class MockBatcher:
"""
def __init__(self):
- self.calls: List[Dict[str, Any]] = []
+ self.calls: list[dict[str, Any]] = []
self.committed: bool = False
class _Table:
@@ -103,6 +102,7 @@ class MockBatcher:
self.litellm_tagtable = _Table("tag", self)
self.litellm_modelaccessgroupbudgettable = _Table("model_access_group", self)
self.litellm_projecttable = _Table("project", self)
+ self.litellm_agentstable = _Table("agent", self)
self.litellm_endusertable = _Table("enduser", self)
async def commit(self):
@@ -119,8 +119,9 @@ class MockDB:
self.litellm_tagtable = MockTable()
self.litellm_modelaccessgroupbudgettable = MockTable()
self.litellm_projecttable = MockTable()
- self.batch_calls: List[Dict[str, Any]] = []
- self.batchers: List[MockBatcher] = []
+ self.litellm_agentstable = MockTable()
+ self.batch_calls: list[dict[str, Any]] = []
+ self.batchers: list[MockBatcher] = []
def batch_(self):
batcher = MockBatcher()
@@ -140,21 +141,21 @@ class MockDB:
class MockPrismaClient:
def __init__(self):
- self.data: Dict[str, List[Any]] = {
+ self.data: dict[str, list[Any]] = {
"key": [],
"user": [],
"team": [],
"budget": [],
"enduser": [],
}
- self.updated_data: Dict[str, List[Any]] = {
+ self.updated_data: dict[str, list[Any]] = {
"key": [],
"user": [],
"team": [],
"budget": [],
"enduser": [],
}
- self.get_data_calls: List[Dict[str, Any]] = []
+ self.get_data_calls: list[dict[str, Any]] = []
self.db = MockDB()
async def get_data(self, table_name, query_type, **kwargs):
@@ -246,7 +247,7 @@ def _budget_row(
)
-def _batch_writes(mock_prisma_client, table: str, op: str | None = None) -> List[Dict[str, Any]]:
+def _batch_writes(mock_prisma_client, table: str, op: str | None = None) -> list[dict[str, Any]]:
"""Writes that were committed to the DB, optionally narrowed to one op."""
return [
call
@@ -604,7 +605,7 @@ def test_budget_table_reset_zeroes_spend_on_every_linked_table(
_POSTGRES_MAX_BIND_VARIABLES: Final = 32767
-def _bind_count(where: Dict[str, Any]) -> int:
+def _bind_count(where: dict[str, Any]) -> int:
"""Bind variables one prisma where-clause compiles to: each scalar is one
placeholder and an ``in`` list contributes one per element."""
return sum(len(value["in"]) if isinstance(value, dict) and "in" in value else 1 for value in where.values())
@@ -924,8 +925,8 @@ def test_reset_budget_skips_null_budget_id_endusers_when_default_not_in_reset_li
def _make_reset_budget_windows_job(
monkeypatch,
- key_rows: List[Dict[str, Any]],
- team_rows: List[Dict[str, Any]],
+ key_rows: list[dict[str, Any]],
+ team_rows: list[dict[str, Any]],
):
"""Build a ResetBudgetJob with a fully-mocked prisma client and a fake
`litellm.proxy.proxy_server` module exposing a stub `spend_counter_cache`.
@@ -1699,7 +1700,6 @@ def test_enduser_invalidation_is_paged_and_batched(reset_budget_job, mock_prisma
assert evicted == {f"end_user_id:cust-{i:06d}" for i in range(population)}
-
def test_enduser_invalidation_reports_a_page_read_failure_instead_of_a_clean_finish(
mock_prisma_client, monkeypatch
):
@@ -1983,7 +1983,11 @@ def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mo
budget = _budget_row(budget_id="budget-1", budget_duration="7d")
mock_prisma_client.data["budget"] = [budget]
mock_prisma_client.data["enduser"] = [
- type("EndUser", (), {"spend": 5.0, "litellm_budget_table": budget, "user_id": "enduser-1", "budget_id": "budget-1"})
+ type(
+ "EndUser",
+ (),
+ {"spend": 5.0, "litellm_budget_table": budget, "user_id": "enduser-1", "budget_id": "budget-1"},
+ )
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
@@ -1998,6 +2002,7 @@ def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mo
("tag", "update_many"),
("model_access_group", "update_many"),
("project", "update_many"),
+ ("agent", "update_many"),
("enduser", "update_many"),
("budget", "update_many"),
}
@@ -2200,10 +2205,10 @@ class ChunkedPrismaClient(MockPrismaClient):
seeing rows rather than quietly running out of data.
"""
- def __init__(self, chunks_by_table: Dict[str, List[List[Any]]]):
+ def __init__(self, chunks_by_table: dict[str, list[list[Any]]]):
super().__init__()
self._chunks_by_table = chunks_by_table
- self.fetches_by_table: Dict[str, int] = {}
+ self.fetches_by_table: dict[str, int] = {}
async def get_data(self, table_name, query_type, **kwargs):
self.get_data_calls.append({"table_name": table_name, "query_type": query_type, **kwargs})
@@ -2396,8 +2401,8 @@ class PoisonRow:
class RecordingServiceLogging:
def __init__(self):
- self.success_calls: List[Dict[str, Any]] = []
- self.failure_calls: List[Dict[str, Any]] = []
+ self.success_calls: list[dict[str, Any]] = []
+ self.failure_calls: list[dict[str, Any]] = []
async def async_service_success_hook(self, **kwargs):
self.success_calls.append(kwargs)
@@ -2475,8 +2480,8 @@ class FakePodLockManager:
if self.redis_cache is not None:
self.redis_cache.async_get_cache = AsyncMock(return_value="another-pod" if held_by_other else None)
self._acquired = acquired
- self.acquire_calls: List[Dict[str, str | int | None]] = []
- self.release_calls: List[str] = []
+ self.acquire_calls: list[dict[str, str | int | None]] = []
+ self.release_calls: list[str] = []
@staticmethod
def get_redis_lock_key(cronjob_id: str) -> str:
@@ -2607,21 +2612,21 @@ def test_reset_budget_lease_outlives_one_scheduler_tick(monkeypatch):
assert RESET_BUDGET_JOB_LOCK_TTL_SECONDS > PROXY_BUDGET_RESCHEDULER_MIN_TIME
-def _window_row(source_id_column: str, row_id: str, reset_at: datetime) -> Dict[str, Any]:
+def _window_row(source_id_column: str, row_id: str, reset_at: datetime) -> dict[str, Any]:
return {
source_id_column: row_id,
"budget_limits": [{"budget_duration": "1h", "reset_at": reset_at.isoformat(), "max_budget": 10}],
}
-def _paginating_window_job(monkeypatch, pages_by_table: Dict[str, List[List[Dict[str, Any]]]]):
+def _paginating_window_job(monkeypatch, pages_by_table: dict[str, list[list[dict[str, Any]]]]):
"""Serve each table a canned sequence of pages and record every query.
Returns (job, calls) where calls is a list of (sql, cursor, limit).
"""
prisma_client = MagicMock()
remaining = {table: list(pages) for table, pages in pages_by_table.items()}
- calls: List[Dict[str, Any]] = []
+ calls: list[dict[str, Any]] = []
async def fake_query_raw(query: str, *args, **kwargs):
table = "key" if '"LiteLLM_VerificationToken"' in query else "team"
@@ -2751,7 +2756,7 @@ def test_debug_row_dump_is_deferred_until_a_record_is_emitted():
assert serialized == ["serialized"]
-def _cursor_paginating_window_job(monkeypatch, key_rows: List[Dict[str, Any]]):
+def _cursor_paginating_window_job(monkeypatch, key_rows: list[dict[str, Any]]):
"""Serve real keyset pages out of one ordered table, honouring the cursor.
Unlike the canned-page helper above, this models the database: a page is
@@ -2760,7 +2765,7 @@ def _cursor_paginating_window_job(monkeypatch, key_rows: List[Dict[str, Any]]):
"""
prisma_client = MagicMock()
ordered = sorted(key_rows, key=lambda r: r["token"])
- visited: List[str] = []
+ visited: list[str] = []
async def fake_query_raw(query: str, *args, **kwargs):
if '"LiteLLM_TeamTable"' in query:
@@ -2816,7 +2821,7 @@ class FlakyPrismaClient(MockPrismaClient):
def __init__(self, *, read_failures: int = 0, commit_failures: int = 0, error: Exception | None = None):
super().__init__()
- self.reconnect_reasons: List[str] = []
+ self.reconnect_reasons: list[str] = []
self.read_attempts: int = 0
self.commit_attempts: int = 0
self._read_failures = read_failures
@@ -2994,7 +2999,7 @@ def test_transport_error_on_window_read_reconnects_and_still_resets(monkeypatch)
expired = (datetime.utcnow() - timedelta(minutes=5)).isoformat() + "Z"
key_rows = [{"token": "sk-expired", "budget_limits": [{"budget_duration": "1d", "reset_at": expired}]}]
job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[])
- reconnect_reasons: List[str] = []
+ reconnect_reasons: list[str] = []
good_query_raw = prisma_client.db.query_raw
async def failing_once_query_raw(query: str, *args, **kwargs):
@@ -3019,7 +3024,7 @@ def test_connect_error_on_window_write_reconnects_and_writes(monkeypatch):
expired = (datetime.utcnow() - timedelta(minutes=5)).isoformat() + "Z"
team_rows = [{"team_id": "team-expired", "budget_limits": [{"budget_duration": "1d", "reset_at": expired}]}]
job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=[], team_rows=team_rows)
- reconnect_reasons: List[str] = []
+ reconnect_reasons: list[str] = []
async def failing_once_update(**kwargs) -> None:
if not reconnect_reasons:
@@ -3578,3 +3583,26 @@ def test_reset_deletes_spend_counter_instead_of_seeding(reset_budget_job, mock_p
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:user:carol")
counter_cache.in_memory_cache.set_cache.assert_not_called()
counter_cache.redis_cache.async_set_cache.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("rollover_cap", [None, 1.0])
+async def test_agent_reset_stamps_window_in_each_spend_mutation(rollover_cap):
+ from litellm.proxy.common_utils.reset_budget_job import _BudgetCascade
+
+ client = MockPrismaClient()
+ job = ResetBudgetJob(MagicMock(), client)
+ next_window = datetime(2026, 1, 2, tzinfo=timezone.utc)
+ cascade = _BudgetCascade(
+ budget_ids=("agent-budget",),
+ budget_resets=(("agent-budget", next_window),),
+ rollover_caps={} if rollover_cap is None else {"agent-budget": rollover_cap},
+ )
+ await job._commit_budget_cascade(cascade)
+ agent_writes = [call for call in client.db.batch_calls if call["table"] == "agent"]
+ assert len(agent_writes) == (1 if rollover_cap is None else 2)
+ assert all(call["data"]["spend_window"] == next_window for call in agent_writes)
+ assert agent_writes[0]["data"]["spend"] == 0.0
+ if rollover_cap is not None:
+ assert agent_writes[1]["data"]["spend"] == {"decrement": rollover_cap}
+ assert client.db.batchers[0].committed
diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py
index 7b160c055d2..ca6a9411005 100644
--- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py
+++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py
@@ -930,7 +930,7 @@ async def test_commit_spend_updates_to_db_increments_agent_spend():
mock_batcher.litellm_agentstable.update_many.assert_called_once()
call_kwargs = mock_batcher.litellm_agentstable.update_many.call_args[1]
- assert call_kwargs["where"] == {"agent_id": agent_id}
+ assert call_kwargs["where"] == {"agent_id": agent_id, "spend_window": None}
assert call_kwargs["data"] == {"spend": {"increment": response_cost}}
@@ -1475,7 +1475,8 @@ async def test_add_spend_log_transaction_to_daily_end_user_transaction_skips_whe
@pytest.mark.asyncio
-async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agent_id_and_queues_update():
+@pytest.mark.parametrize("billing_agent", [None, "caller-agent"])
+async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agent_id_and_queues_update(billing_agent):
"""
Ensure agent_id is injected and queued for daily aggregation.
"""
@@ -1487,6 +1488,7 @@ async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agen
payload = {
"request_id": "req-123",
"agent_id": agent_id,
+ "billing_agent_id": billing_agent,
"user": "test-user",
"startTime": "2024-01-01T12:00:00",
"api_key": "test-key",
@@ -1506,14 +1508,19 @@ async def test_add_spend_log_transaction_to_daily_agent_transaction_injects_agen
prisma_client=mock_prisma,
)
+ if billing_agent is None:
+ writer.daily_agent_spend_update_queue.add_update.assert_not_awaited()
+ return
writer.daily_agent_spend_update_queue.add_update.assert_called_once()
call_args = writer.daily_agent_spend_update_queue.add_update.call_args[1]
update_dict = call_args["update"]
assert len(update_dict) == 1
+ charged_agent: Final = billing_agent or agent_id
for key, transaction in update_dict.items():
- assert key == f"{agent_id}_2024-01-01_test-key_gpt-4_openai_"
- assert transaction["agent_id"] == agent_id
+ assert key == f"{charged_agent}_2024-01-01_test-key_gpt-4_openai_"
+ assert transaction["agent_id"] == charged_agent
+ assert transaction["spend"] == 0.3
assert transaction["date"] == "2024-01-01"
assert transaction["api_key"] == "test-key"
assert transaction["model"] == "gpt-4"
@@ -4778,3 +4785,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 84da227c0a6..c1c5b2b570d 100644
--- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py
+++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py
@@ -2745,3 +2745,36 @@ def test_autonomous_agent_cost_tracking_needs_no_human_or_virtual_key(agent_id:
assert _should_track_cost_callback(
user_api_key=None, user_id=None, team_id=None, end_user_id=None, call_type="acompletion", agent_id=agent_id
) is expected
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("billing_agent", [None, "verified-agent"])
+async def test_callback_does_not_charge_a_header_selected_agent(
+ billing_agent: str | None, monkeypatch: pytest.MonkeyPatch
+) -> None:
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy import proxy_server
+
+ cache: Final = DualCache()
+ for key in ("spend:user:human", "spend:agent:header-selected-agent", "spend:agent:verified-agent", "spend:agent_window:20260102T000000.000000Z:verified-agent"):
+ cache.in_memory_cache.set_cache(key=key, value=0.0)
+ logging: Final = MagicMock()
+ logging.db_spend_update_writer.update_database = AsyncMock(return_value=True)
+ logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
+ monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
+ monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging)
+ monkeypatch.setattr(proxy_server, "update_cache", AsyncMock())
+ kwargs: Final = {
+ "call_type": "acompletion", "model": "test-model", "response_cost": 0.01,
+ "litellm_params": {"metadata": {
+ "user_api_key_user_id": "human", "agent_id": "header-selected-agent", "billing_agent_id": billing_agent,
+ "billing_agent_counter_key": "spend:agent_window:20260102T000000.000000Z:verified-agent" if billing_agent else None,
+ }},
+ }
+ await _ProxyDBLogger()._PROXY_track_cost_callback(
+ kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now()
+ )
+ assert cache.in_memory_cache.get_cache(key="spend:user:human") == 0.01
+ assert cache.in_memory_cache.get_cache(key="spend:agent:header-selected-agent") == 0.0
+ assert cache.in_memory_cache.get_cache(key="spend:agent:verified-agent") == 0.0
+ assert cache.in_memory_cache.get_cache(key="spend:agent_window:20260102T000000.000000Z:verified-agent") == (0.01 if billing_agent else 0.0)
diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py
index 59c2921e0d0..2843bb60965 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py
@@ -714,6 +714,9 @@ class _FakePrismaClient:
litellm_modelaccessgroupbudgettable=self.access_group_budget_table,
litellm_proxymodeltable=self.model_table,
)
+ self.writer_db = SimpleNamespace(
+ litellm_agentstable=SimpleNamespace(find_first=AsyncMock(return_value=None)),
+ )
def jsonify_object(self, data):
return dict(data)
diff --git a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py
index 2f3be61d00f..e6af298b8d8 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py
@@ -21,6 +21,8 @@ def client_and_mocks(monkeypatch):
mock_table = MagicMock()
mock_table.create = AsyncMock(side_effect=lambda *, data: data)
mock_table.update = AsyncMock(side_effect=lambda *, where, data: {**where, **data})
+ mock_table.delete = AsyncMock(side_effect=lambda *, where: where)
+ mock_prisma.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(
litellm_budgettable=mock_table,
@@ -46,6 +48,34 @@ def client_and_mocks(monkeypatch):
monkeypatch.setattr(ps, "prisma_client", ps.prisma_client)
+@pytest.mark.parametrize("operation", ["update", "delete"])
+def test_agent_linked_budget_requires_the_agent_management_flow(client_and_mocks, operation):
+ client, prisma, table = client_and_mocks
+ prisma.writer_db.litellm_agentstable.find_first.return_value = types.SimpleNamespace(agent_id="agent-one")
+ admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
+ app.dependency_overrides[ps.user_api_key_auth] = lambda: admin
+ payload = {"budget_id": "agent-budget", "budget_duration": "2h"} if operation == "update" else {"id": "agent-budget"}
+
+ response = client.post(f"/budget/{operation}", json=payload)
+
+ assert response.status_code == 409
+ assert "/v1/agents/agent-one" in response.json()["detail"]
+ table.update.assert_not_awaited()
+ table.delete.assert_not_awaited()
+
+
+def test_unlinked_budget_can_still_be_deleted(client_and_mocks):
+ client, _, table = client_and_mocks
+ admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
+ app.dependency_overrides[ps.user_api_key_auth] = lambda: admin
+
+ response = client.post("/budget/delete", json={"id": "ordinary-budget"})
+
+ assert response.status_code == 200
+ assert response.json()["budget_id"] == "ordinary-budget"
+ table.delete.assert_awaited_once_with(where={"budget_id": "ordinary-budget"})
+
+
@pytest.mark.asyncio
async def test_new_budget_success(client_and_mocks):
client, _, mock_table = client_and_mocks
diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py
index 3c6afa86c45..2bdc81756c7 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py
@@ -1176,6 +1176,7 @@ async def _run_legacy_update_organization(
mock_prisma_client.db.litellm_organizationtable.find_unique = AsyncMock(return_value=existing_org)
mock_prisma_client.db.litellm_organizationtable.update = AsyncMock(return_value=MagicMock())
mock_prisma_client.db.litellm_budgettable.update = AsyncMock()
+ mock_prisma_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock())
diff --git a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py
index 3cfdd345a45..830ffdcad92 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py
@@ -35,6 +35,10 @@ class _BudgetState:
return SimpleNamespace(**self._values)
+def _writer_db_without_agent_budgets() -> SimpleNamespace:
+ return SimpleNamespace(litellm_agentstable=SimpleNamespace(find_first=AsyncMock(return_value=None)))
+
+
class FakeVerificationTokenTable:
"""Stand-in for ``prisma_client.db.litellm_verificationtoken``.
@@ -316,7 +320,7 @@ async def test_update_tag_explicit_null_preserves_general_budget_fields(field):
created_by="admin",
)
mock_db = Mock()
- mock_prisma = SimpleNamespace(db=mock_db)
+ mock_prisma = SimpleNamespace(db=mock_db, writer_db=_writer_db_without_agent_budgets())
mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag)
mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag)
@@ -370,7 +374,7 @@ async def test_update_tag_explicit_null_clears_budget_duration():
created_by="admin",
)
mock_db = Mock()
- mock_prisma = SimpleNamespace(db=mock_db)
+ mock_prisma = SimpleNamespace(db=mock_db, writer_db=_writer_db_without_agent_budgets())
mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag)
mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag)
diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
index 0b866d7f736..5f0488b0c76 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
@@ -2973,6 +2973,7 @@ async def test_upsert_team_member_budget_table_clears_duration_kept_budget(mock_
mock_db_client.db.litellm_budgettable.update = AsyncMock(
side_effect=lambda where, data: SimpleNamespace(**data)
)
+ mock_db_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None)
result = await TeamMemberBudgetHandler.upsert_team_member_budget_table(
team_table=team_table,
diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py
index 922504ecc58..46bac64db19 100644
--- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py
+++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py
@@ -845,6 +845,7 @@ async def test_team_update_reaches_inherited_members_but_not_overridden_ones():
db: Final = _FakeDb()
prisma_client: Final = MagicMock()
prisma_client.db = db
+ prisma_client.writer_db.litellm_agentstable.find_first = AsyncMock(return_value=None)
admin: Final = UserAPIKeyAuth(user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN)
team_id: Final = "team-shared-default"
default_budget: Final = await db.litellm_budgettable.create(data={"budget_id": "team-default", "max_budget": 100.0})
diff --git a/tests/test_litellm/proxy/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/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py
index 3ffb6335ad4..f8cd308145f 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py
@@ -3767,7 +3767,7 @@ class TestSpendLogsPayload:
"model": "gpt-4o",
"user": "",
"team_id": "",
- "metadata": '{"actor_agent_id": null, "target_agent_id": null, "billing_agent_id": null, "agent_execution_mode": null, "verified_human_user_id": null, "applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}',
+ "metadata": '{"actor_agent_id": null, "target_agent_id": null, "billing_agent_id": null, "billing_agent_counter_key": null, "agent_execution_mode": null, "verified_human_user_id": null, "applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}',
"cache_key": "Cache OFF",
"spend": 0.00022500000000000002,
"total_tokens": 30,
diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py
index 84266325226..cda9d066661 100644
--- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py
+++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py
@@ -8523,3 +8523,19 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback():
team_callback_settings_obj=None,
)
assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"}
+@pytest.mark.parametrize("bound", [False, True])
+def test_agent_budget_window_metadata_is_owned_by_authenticated_policy(bound: bool) -> None:
+ from litellm.types.agents import AgentResponse
+
+ auth: Final = UserAPIKeyAuth(agent_id="agent" if bound else None)
+ if bound:
+ auth.billing_agent_policy = AgentResponse(
+ agent_id="agent", agent_name="Agent", agent_card_params={},
+ litellm_budget_table={"budget_id": "budget", "budget_reset_at": "2026-01-02T00:00:00Z"},
+ )
+ result: Final = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
+ {"metadata": {"billing_agent_counter_key": "spend:agent:victim"}}, auth, "metadata"
+ )
+ assert result["metadata"]["billing_agent_counter_key"] == (
+ "spend:agent_window:20260102T000000.000000Z:agent" if bound else None
+ )
diff --git a/tests/unit/types/proxy/test_agent_identity.py b/tests/unit/types/proxy/test_agent_identity.py
new file mode 100644
index 00000000000..ed4b0901da2
--- /dev/null
+++ b/tests/unit/types/proxy/test_agent_identity.py
@@ -0,0 +1,20 @@
+import pytest
+from pydantic import ValidationError
+
+from litellm.types.proxy.agent_identity import AgentBudgetConfig
+
+
+@pytest.mark.parametrize("amount", [-1, float("inf"), float("-inf"), float("nan")])
+def test_agent_budget_rejects_negative_or_nonfinite_caps(amount: float) -> None:
+ with pytest.raises(ValidationError):
+ AgentBudgetConfig(max_budget=amount)
+
+
+@pytest.mark.parametrize("amount", [0, 0.01, 100])
+def test_agent_budget_preserves_a_finite_nonnegative_cap(amount: float) -> None:
+ assert AgentBudgetConfig(max_budget=amount).max_budget == amount
+
+
+def test_agent_budget_rejects_unknown_policy_fields() -> None:
+ with pytest.raises(ValidationError):
+ AgentBudgetConfig.model_validate({"max_budget": 1, "unknown_control": True})
diff --git a/tests/unit/types/test_agents.py b/tests/unit/types/test_agents.py
new file mode 100644
index 00000000000..a583ab9e8a2
--- /dev/null
+++ b/tests/unit/types/test_agents.py
@@ -0,0 +1,28 @@
+from datetime import datetime, timedelta, timezone
+from typing import Final
+
+import pytest
+
+from litellm.types.agents import agent_budget_counter_key, agent_spend_filter
+
+
+@pytest.mark.parametrize("offset", [None, timezone.utc, timezone(timedelta(hours=5, minutes=30))])
+def test_agent_window_key_round_trip_preserves_the_admitted_instant(offset) -> None:
+ instant: Final = datetime(2030, 1, 1, 12, 30, 1, 123000, tzinfo=offset)
+ expected: Final = instant.replace(tzinfo=timezone.utc) if offset is None else instant.astimezone(timezone.utc)
+ key: Final = agent_budget_counter_key("agent:with:colons", instant)
+ assert agent_spend_filter(key) == {"agent_id": "agent:with:colons", "spend_window": expected}
+ assert key == agent_budget_counter_key("agent:with:colons", expected)
+
+
+def test_unbudgeted_key_is_filtered_to_an_unbudgeted_row() -> None:
+ key: Final = agent_budget_counter_key("agent-one", None)
+ assert agent_spend_filter(key) == {"agent_id": "agent-one", "spend_window": None}
+
+
+def test_different_budget_windows_never_share_a_settlement_filter() -> None:
+ first: Final = datetime(2030, 1, 1, tzinfo=timezone.utc)
+ second: Final = first + timedelta(days=1)
+ assert agent_spend_filter(agent_budget_counter_key("agent-one", first)) != agent_spend_filter(
+ agent_budget_counter_key("agent-one", second)
+ )
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx
index 50c60776ff3..8495c99d59f 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx
@@ -256,6 +256,41 @@ export const AgentIdentityFields = ({ accessToken }: { accessToken: string | nul
>
)}
+
>
);
};
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 23045adcf20..2db4f97e128 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
@@ -5,7 +5,7 @@ import type { AgentFormValues, AgentRequestPayload } from "./AgentFormKit";
export type EntraAgentIdentity = components["schemas"]["EntraIdentityConfig"];
type AgentIdentityState = Pick<
components["schemas"]["AgentResponse"],
- "identity" | "enabled" | "execution_mode" | "agent_card_params"
+ "identity" | "enabled" | "execution_mode" | "agent_card_params" | "litellm_budget_table"
>;
export const IDENTITY_UUID_PATTERN = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i;
@@ -47,6 +47,8 @@ export const parseIdentityForForm = (agent?: Partial | null)
...identityFormFields(identity),
execution_mode: agent?.execution_mode ?? "autonomous",
enabled: agent?.enabled ?? true,
+ agent_max_budget: agent?.litellm_budget_table?.max_budget ?? "",
+ agent_budget_duration: agent?.litellm_budget_table?.budget_duration ?? "",
};
};
@@ -98,11 +100,23 @@ export const withAgentIdentity = (
const hasCard = !existing || cardEdited || Object.keys(existing.agent_card_params ?? {}).length > 0;
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 {
...settings,
...(hasCard && agent_card_params ? { agent_card_params } : {}),
...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 c9f154c3ce1..24fd7e9274f 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);
@@ -354,6 +369,7 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT
{agent.agent_id}
{agent.agent_name}
+
{agent.agent_card_params?.name || "-"}
{agent.agent_card_params?.description || "-"}
{agent.agent_card_params?.url || "-"}
diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts
index 92d946c19e1..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 8ffd5c96ab0..a44326e4a7c 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -24597,6 +24597,24 @@ export interface components {
[key: string]: string;
};
};
+ /** AgentBudgetConfig */
+ AgentBudgetConfig: {
+ /** Budget Duration */
+ budget_duration?: string | null;
+ /** Max Budget */
+ max_budget: number;
+ };
+ /** AgentBudgetState */
+ AgentBudgetState: {
+ /** Budget Duration */
+ budget_duration?: string | null;
+ /** Budget Id */
+ budget_id: string;
+ /** Budget Reset At */
+ budget_reset_at?: string | null;
+ /** Max Budget */
+ max_budget?: number | null;
+ };
/**
* AgentCapabilities
* @description Defines optional capabilities supported by an agent.
@@ -24678,6 +24696,7 @@ export interface components {
agent_card_params?: components["schemas"]["AgentCard"];
/** Agent Name */
agent_name: string;
+ budget?: components["schemas"]["AgentBudgetConfig"] | null;
/** Enabled */
enabled?: boolean;
/**
@@ -24964,6 +24983,8 @@ export interface components {
agent_id: string;
/** Agent Name */
agent_name: string;
+ /** Budget Id */
+ budget_id?: string | null;
/** Created At */
created_at?: string | null;
/** Created By */
@@ -24995,6 +25016,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;
@@ -38742,6 +38764,7 @@ export interface components {
agent_card_params?: components["schemas"]["AgentCard"];
/** Agent Name */
agent_name?: string;
+ budget?: components["schemas"]["AgentBudgetConfig"] | null;
/** Enabled */
enabled?: boolean;
/**
From 92e13728f53ad437a4a5fa1d96e13602a299c206 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Mon, 28 Sep 2026 19:28:42 -0700
Subject: [PATCH 02/12] fix(agents): account for completion bridge invocation
fees
---
litellm/a2a_protocol/main.py | 34 ++++++++++++++++++++--
litellm/a2a_protocol/streaming_iterator.py | 29 +++++++++++-------
tests/unit/a2a_protocol/test_main.py | 31 ++++++++++++++++++++
3 files changed, 81 insertions(+), 13 deletions(-)
diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py
index aa41e63b40b..8d733971c70 100644
--- a/litellm/a2a_protocol/main.py
+++ b/litellm/a2a_protocol/main.py
@@ -471,13 +471,20 @@ async def asend_message(
if custom_llm_provider:
if request is None:
raise ValueError("request is required for completion bridge")
- return await _send_message_via_completion_bridge(
+ bridge_response: Final = await _send_message_via_completion_bridge(
request=request,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
)
+ bridge_prompt_tokens, bridge_completion_tokens, _ = await asyncify(
+ A2ARequestUtils.calculate_usage_from_request_response
+ )(request=request, response_dict=bridge_response.model_dump(mode="json", exclude_none=True))
+ _set_usage_on_logging_obj(kwargs, bridge_prompt_tokens, bridge_completion_tokens)
+ _set_litellm_params_on_logging_obj(kwargs, litellm_params)
+ _set_agent_id_on_logging_obj(kwargs, agent_id)
+ return bridge_response
# Standard A2A client flow
if request is None:
@@ -692,12 +699,35 @@ async def asend_message_streaming(
request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
)
- async for chunk in A2ACompletionBridgeHandler.handle_streaming(
+ bridge_name: Final = str(litellm_params.get("model") or agent_id or "agent")
+ existing_logging: Final = kwargs.get("litellm_logging_obj")
+ bridge_logging: Final = (
+ existing_logging
+ if isinstance(existing_logging, Logging)
+ else _build_streaming_logging_obj(
+ request=request,
+ agent_name=bridge_name,
+ agent_id=agent_id,
+ litellm_params=litellm_params,
+ metadata=metadata,
+ proxy_server_request=proxy_server_request,
+ )
+ )
+ bridge_context: Final = {"litellm_logging_obj": bridge_logging}
+ _set_litellm_params_on_logging_obj(bridge_context, litellm_params)
+ _set_agent_id_on_logging_obj(bridge_context, agent_id)
+ bridge_stream: Final = A2ACompletionBridgeHandler.handle_streaming(
request_id=str(request.id),
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
+ )
+ async for chunk in A2AStreamingIterator(
+ stream=bridge_stream,
+ request=request,
+ logging_obj=bridge_logging,
+ agent_name=bridge_name,
):
yield chunk
return
diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py
index 8232d7cf2d8..2a20979501a 100644
--- a/litellm/a2a_protocol/streaming_iterator.py
+++ b/litellm/a2a_protocol/streaming_iterator.py
@@ -5,7 +5,7 @@ A2A Streaming Iterator with token tracking and logging support.
import asyncio
from collections.abc import AsyncIterator
from datetime import datetime
-from typing import TYPE_CHECKING, Final
+from typing import TYPE_CHECKING, Final, Generic, TypeVar
import litellm
from litellm._logging import verbose_logger
@@ -18,7 +18,10 @@ if TYPE_CHECKING:
from a2a.compat.v0_3.types import SendStreamingMessageRequest, SendStreamingMessageResponse
-class A2AStreamingIterator:
+_StreamChunk = TypeVar("_StreamChunk", bound="SendStreamingMessageResponse | dict[str, object]")
+
+
+class A2AStreamingIterator(Generic[_StreamChunk]):
"""
Async iterator for A2A streaming responses with token tracking.
@@ -27,7 +30,7 @@ class A2AStreamingIterator:
def __init__(
self,
- stream: AsyncIterator["SendStreamingMessageResponse"],
+ stream: AsyncIterator[_StreamChunk],
request: "SendStreamingMessageRequest",
logging_obj: LiteLLMLoggingObj,
agent_name: str = "unknown",
@@ -39,14 +42,14 @@ class A2AStreamingIterator:
self.start_time = datetime.now()
# Collect chunks for token counting
- self.chunks: list[SendStreamingMessageResponse] = []
+ self.chunks: list[_StreamChunk] = []
self.collected_text_parts: list[str] = []
- self.final_chunk: SendStreamingMessageResponse | None = None
+ self.final_chunk: _StreamChunk | None = None
def __aiter__(self):
return self
- async def __anext__(self) -> "SendStreamingMessageResponse":
+ async def __anext__(self) -> _StreamChunk:
try:
chunk: Final = await self.stream.__anext__()
@@ -69,20 +72,20 @@ class A2AStreamingIterator:
await self._handle_stream_complete()
raise
- def _collect_text_from_chunk(self, chunk: "SendStreamingMessageResponse") -> None:
+ def _collect_text_from_chunk(self, chunk: _StreamChunk) -> None:
"""Extract text from a streaming chunk and add to collected parts."""
try:
- chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
+ chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True)
text: Final = A2ARequestUtils.extract_text_from_response(chunk_dict)
if text:
self.collected_text_parts.append(text)
except Exception:
verbose_logger.debug("Failed to extract text from A2A streaming chunk")
- def _is_completed_chunk(self, chunk: "SendStreamingMessageResponse") -> bool:
+ def _is_completed_chunk(self, chunk: _StreamChunk) -> bool:
"""Check if chunk indicates stream completion."""
try:
- chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
+ chunk_dict: Final = chunk if isinstance(chunk, dict) else chunk.model_dump(mode="json", exclude_none=True)
result: Final = chunk_dict.get("result", {})
if isinstance(result, dict):
status: Final = result.get("status", {})
@@ -160,7 +163,11 @@ class A2AStreamingIterator:
# Add final chunk result if available
if self.final_chunk:
try:
- chunk_dict: Final = self.final_chunk.model_dump(mode="json", exclude_none=True)
+ chunk_dict: Final = (
+ self.final_chunk
+ if isinstance(self.final_chunk, dict)
+ else self.final_chunk.model_dump(mode="json", exclude_none=True)
+ )
result["result"] = chunk_dict.get("result", {})
except Exception:
pass
diff --git a/tests/unit/a2a_protocol/test_main.py b/tests/unit/a2a_protocol/test_main.py
index 4ba0ef8fa04..1720b521bd7 100644
--- a/tests/unit/a2a_protocol/test_main.py
+++ b/tests/unit/a2a_protocol/test_main.py
@@ -539,3 +539,34 @@ def test_streaming_logging_obj_keeps_agent_credentials_out_of_logging_params():
assert logging_obj.litellm_params == expected
assert logging_obj.optional_params == expected
assert logging_obj.model_call_details["litellm_params"] == expected
+
+
+class _AgentFeeRecorder(CustomLogger):
+ def __init__(self):
+ super().__init__()
+ self.logged = asyncio.Event()
+ self.fees = ()
+
+ async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
+ if kwargs.get("call_type") in ("asend_message", "asend_message_streaming"):
+ self.fees = (*self.fees, (kwargs.get("agent_id"), kwargs["standard_logging_object"]["response_cost"]))
+ self.logged.set()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("streaming", [False, True])
+async def test_completion_bridge_records_one_agent_fee(streaming, monkeypatch):
+ from litellm.a2a_protocol.main import asend_message_streaming
+
+ recorder = _AgentFeeRecorder()
+ monkeypatch.setattr(litellm, "callbacks", [recorder])
+ params = {"custom_llm_provider": "openai", "model": "gpt-4o-mini", "mock_response": "hello back", "cost_per_query": 0.01}
+ if streaming:
+ request = SendStreamingMessageRequest(id="bridge-stream", params=_request().params)
+ chunks = [chunk async for chunk in asend_message_streaming(request=request, litellm_params=params, agent_id="budgeted-agent")]
+ assert chunks[-1]["result"]["final"] is True
+ else:
+ response = await asend_message(request=_request(), litellm_params=params, agent_id="budgeted-agent")
+ assert response.id == "r1"
+ await asyncio.wait_for(recorder.logged.wait(), timeout=2)
+ assert recorder.fees == (("budgeted-agent", pytest.approx(0.01)),)
From bcfa93d95ca3acc31680a5dbcfc63b4f61e4e4a4 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Mon, 28 Sep 2026 21:51:29 -0700
Subject: [PATCH 03/12] fix(agents): retain streamed spend until callback
settlement
---
litellm/a2a_protocol/streaming_iterator.py | 3 ++
.../test_a2a_streaming_iterator.py | 37 +++++++++++++++++++
2 files changed, 40 insertions(+)
diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py
index 2a20979501a..3e69869e0d9 100644
--- a/litellm/a2a_protocol/streaming_iterator.py
+++ b/litellm/a2a_protocol/streaming_iterator.py
@@ -12,6 +12,7 @@ from litellm._logging import verbose_logger
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
from litellm.a2a_protocol.utils import A2ARequestUtils
from litellm.litellm_core_utils.asyncify import asyncify
+from litellm.litellm_core_utils.core_helpers import bind_budget_reservation_to_callbacks
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
if TYPE_CHECKING:
@@ -130,6 +131,8 @@ class A2AStreamingIterator(Generic[_StreamChunk]):
# Build result for logging
result: Final = self._build_logging_result(usage)
+ bind_budget_reservation_to_callbacks(self.logging_obj.litellm_params)
+
# Call success handlers - they will build standard_logging_object
asyncio.create_task(
self.logging_obj.dispatch_success_handlers(
diff --git a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py
index 2e883e91fda..9bf43537737 100644
--- a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py
+++ b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py
@@ -135,3 +135,40 @@ async def test_stream_completion_counts_tokens_off_the_event_loop(monkeypatch):
assert usage.prompt_tokens > 100_000
assert usage.completion_tokens > 100_000
assert_loop_stayed_free(took, lags)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("outcome", ["success", "failure", "cancelled"])
+async def test_stream_reservation_survives_cleanup_only_when_billing_is_scheduled(outcome):
+ from unittest.mock import AsyncMock
+
+ from litellm.proxy.spend_tracking.budget_reservation import release_unbound_budget_reservation
+
+ reservation = {"reserved_cost": 0.01, "entries": [], "finalized": False}
+ logging_obj = SimpleNamespace(
+ litellm_params={"metadata": {"user_api_key_budget_reservation": reservation}},
+ model_call_details={},
+ dispatch_success_handlers=AsyncMock(),
+ )
+
+ async def stream():
+ yield {"result": {"kind": "message", "parts": [{"kind": "text", "text": "hello"}]}}
+ if outcome == "failure":
+ raise RuntimeError("upstream failed")
+ if outcome == "cancelled":
+ raise asyncio.CancelledError()
+
+ iterator = A2AStreamingIterator(
+ stream=stream(),
+ request=SimpleNamespace(params=SimpleNamespace(message={"parts": [{"kind": "text", "text": "hi"}]})),
+ logging_obj=logging_obj,
+ )
+ if outcome == "success":
+ assert len([chunk async for chunk in iterator]) == 1
+ else:
+ with pytest.raises(RuntimeError if outcome == "failure" else asyncio.CancelledError):
+ _ = [chunk async for chunk in iterator]
+ await release_unbound_budget_reservation(reservation)
+ assert reservation["finalized"] is (outcome != "success")
+ await asyncio.sleep(0)
+ assert logging_obj.dispatch_success_handlers.await_count == (1 if outcome == "success" else 0)
From a3ec702e2f7aac6cf4d36f433dd96b22c28a9e41 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Mon, 28 Sep 2026 22:17:12 -0700
Subject: [PATCH 04/12] fix(agents): release unbilled invocation fees on
cancellation
---
litellm/proxy/spend_tracking/budget_reservation.py | 2 +-
.../proxy/spend_tracking/test_budget_reservation.py | 10 +++++++++-
2 files changed, 10 insertions(+), 2 deletions(-)
diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py
index 8884a94f4b2..93cba37a90c 100644
--- a/litellm/proxy/spend_tracking/budget_reservation.py
+++ b/litellm/proxy/spend_tracking/budget_reservation.py
@@ -337,7 +337,7 @@ async def reserve_budget_for_request(
return None
input_cost: Final = (
- invocation_cost
+ 0.0
if invocation_cost is not None
else estimate_request_input_cost(
request_body=request_body,
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 469d8506311..9d64fdb5aab 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py
@@ -309,8 +309,9 @@ async def test_team_member_reservation_counter_adds_temp_increase_to_live_team_d
@pytest.mark.asyncio
@pytest.mark.parametrize("charged_agent", ("caller-agent", "target-agent"))
@pytest.mark.parametrize("window", [None, "2026-01-02T00:00:00Z"])
+@pytest.mark.parametrize("outcome", ["success", "cancelled"])
async def test_agent_invocation_reserves_exact_fee_and_reconciles_without_child_cost(
- monkeypatch: pytest.MonkeyPatch, charged_agent: str, window: str | None,
+ monkeypatch: pytest.MonkeyPatch, charged_agent: str, window: str | None, outcome: str,
) -> None:
from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation
from litellm.types.agents import AgentResponse
@@ -336,6 +337,13 @@ async def test_agent_invocation_reserves_exact_fee_and_reconciles_without_child_
assert reservation["reserved_cost"] == pytest.approx(0.2)
assert [entry["counter_key"] for entry in reservation["entries"]] == [counter_key]
assert await cache.async_get_cache(counter_key) == pytest.approx(0.3)
+ if outcome == "cancelled":
+ from litellm.proxy.spend_tracking.budget_reservation import release_budget_reservation_on_cancel
+
+ await release_budget_reservation_on_cancel(reservation)
+ assert await cache.async_get_cache(counter_key) == pytest.approx(0.1)
+ assert reservation["finalized"] is True
+ return
await proxy_server.increment_spend_counters(
token=None, team_id=None, user_id=None, response_cost=0.2,
billing_agent_id=charged_agent, billing_agent_counter_key=counter_key, budget_reservation=reservation,
From 1529494bedaecb0e216530a2d1fd16aa0cb17f27 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Mon, 28 Sep 2026 23:06:00 -0700
Subject: [PATCH 05/12] fix(agents): keep cancelled stream refunds under one
owner
---
litellm/proxy/common_request_processing.py | 6 +++-
.../proxy/test_budget_reservation.py | 35 +++++++++++++++++++
2 files changed, 40 insertions(+), 1 deletion(-)
diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py
index e7a08711eb1..4005b5bebd8 100644
--- a/litellm/proxy/common_request_processing.py
+++ b/litellm/proxy/common_request_processing.py
@@ -3967,7 +3967,11 @@ class ProxyBaseLLMRequestProcessing:
# Starlette closes on disconnect, so the nested iterator hook (which
# only sees GeneratorExit on GC) cannot own the refund.
client_disconnected = not stream_completed
- if not delivered_chunk and not _withheld_provider_output(response):
+ if (
+ not delivered_chunk
+ and not _withheld_provider_output(response)
+ and user_api_key_dict.agent_invocation_cost is None
+ ):
from litellm.proxy.spend_tracking.budget_reservation import (
release_budget_reservation_on_cancel,
)
diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py
index c8e4df1030f..6d2c10321e8 100644
--- a/tests/test_litellm/proxy/test_budget_reservation.py
+++ b/tests/test_litellm/proxy/test_budget_reservation.py
@@ -3720,3 +3720,38 @@ async def test_unreserved_model_access_group_is_charged_alongside_a_reserved_one
assert counter_cache.in_memory_cache.get_cache(
key=model_access_group_spend_counter_key("starter")
) == pytest.approx(4.2)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("invocation_cost", [0.0, 0.01])
+async def test_agent_stream_cancellation_leaves_refund_to_request_cleanup(spend_counter_state, invocation_cost):
+ from litellm.proxy.middleware.budget_reservation_release_middleware import BudgetReservationReleaseMiddleware
+ from litellm.proxy.spend_tracking.budget_reservation import release_unbound_budget_reservation
+
+ counter_cache, _ = spend_counter_state
+ key = "spend:agent:cancelled-agent"
+ counter_cache.set_cache(key, invocation_cost)
+ reservation = {"reserved_cost": invocation_cost, "input_cost": 0.0, "finalized": False, "entries": [{"counter_key": key, "reserved_cost": invocation_cost}]}
+ auth = UserAPIKeyAuth()
+ auth.agent_invocation_cost = invocation_cost
+ auth.budget_reservation = reservation
+
+ async def cancel_before_chunk(user_api_key_dict, response, request_data):
+ raise asyncio.CancelledError()
+ yield "unreachable"
+
+ generator, logging = _drive_streaming_cancel(auth, cancel_before_chunk)
+
+ async def app(scope, receive, send):
+ try:
+ await anext(generator)
+ finally:
+ assert reservation["finalized"] is False
+ assert await counter_cache.async_get_cache(key) == pytest.approx(invocation_cost)
+
+ middleware = BudgetReservationReleaseMiddleware(app, release_unbound_budget_reservation)
+ with pytest.raises(asyncio.CancelledError):
+ await middleware({"type": "http", "state": {"budget_reservation": reservation}}, AsyncMock(), AsyncMock())
+ assert reservation["finalized"] is True
+ assert await counter_cache.async_get_cache(key) == pytest.approx(0.0)
+ logging._arelease_max_parallel_requests_on_disconnect.assert_awaited_once()
From 8f6d78d652010a2a5eaa7961b77012acb1d4fc07 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Wed, 30 Sep 2026 16:24:27 -0700
Subject: [PATCH 06/12] fix(agents): reserve caller fees and isolate lifetime
budget spend
---
.../migration.sql | 2 +
.../litellm_proxy_extras/schema.prisma | 1 +
litellm/proxy/_lazy_openapi_snapshot.json | 5 ++
.../proxy/agent_endpoints/agent_registry.py | 9 +++
.../auth/managed_authorization.py | 9 ++-
litellm/proxy/agent_endpoints/endpoints.py | 34 +++++++-
.../proxy/agent_endpoints/managed_identity.py | 11 ++-
litellm/proxy/db/db_spend_update_writer.py | 31 +++++++-
litellm/proxy/db/spend_counter_reseed.py | 32 ++++++--
.../proxy/hooks/proxy_track_cost_callback.py | 2 +
litellm/proxy/schema.prisma | 1 +
.../spend_tracking/budget_reservation.py | 2 +-
.../spend_tracking/spend_counter_batch.py | 16 ++--
litellm/types/agents.py | 23 +++++-
schema.prisma | 1 +
.../auth/test_managed_authorization.py | 8 +-
.../agent_endpoints/test_agent_registry.py | 2 +
.../proxy/agent_endpoints/test_endpoints.py | 28 +++++++
.../agent_endpoints/test_managed_identity.py | 43 +++++++++++
.../proxy/auth/test_user_api_key_auth.py | 2 +
.../proxy/db/test_db_spend_update_writer.py | 29 +++++++
.../proxy/db/test_spend_counter_reseed.py | 14 ++++
.../proxy/proxy_server/test_proxy_config.py | 1 +
.../spend_tracking/test_budget_reservation.py | 77 +++++++++++++++++--
tests/unit/types/test_agents.py | 13 ++++
.../agents/_components/agent_identity.test.ts | 23 ++++++
.../agents/_components/agent_identity.ts | 27 ++++---
.../agent_info.integration.test.tsx | 56 ++++++++++++--
.../agents/_components/agent_info.tsx | 4 +-
.../src/components/agents/types.ts | 1 +
ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 ++
31 files changed, 457 insertions(+), 55 deletions(-)
diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql
index 6f1db719855..033cd4101c2 100644
--- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql
+++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220000_agent_budgets/migration.sql
@@ -13,3 +13,5 @@ END $$;
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "spend_window" TIMESTAMP(3);
+
+ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "lifetime_budget_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0;
diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
index 6823deadbd3..fc7325ddac3 100644
--- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
+++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
@@ -86,6 +86,7 @@ model LiteLLM_AgentsTable {
retired_identities LiteLLM_RetiredAgentIdentity[]
budget_id String? @unique
spend_window DateTime?
+ lifetime_budget_spend Float @default(0.0)
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
tpm_limit Int?
rpm_limit Int?
diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json
index 1c760d8dc6b..89bdd6647ce 100644
--- a/litellm/proxy/_lazy_openapi_snapshot.json
+++ b/litellm/proxy/_lazy_openapi_snapshot.json
@@ -3246,6 +3246,11 @@
}
]
},
+ "lifetime_budget_spend": {
+ "default": 0.0,
+ "title": "Lifetime Budget Spend",
+ "type": "number"
+ },
"litellm_budget_table": {
"anyOf": [
{
diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py
index 7dbe409226a..5104d904f26 100644
--- a/litellm/proxy/agent_endpoints/agent_registry.py
+++ b/litellm/proxy/agent_endpoints/agent_registry.py
@@ -91,6 +91,15 @@ class AgentRecord(Protocol):
@property
def spend(self) -> float: ...
+ @property
+ def lifetime_budget_spend(self) -> float: ...
+
+ @property
+ def budget_id(self) -> str | None: ...
+
+ @property
+ def litellm_budget_table(self) -> "prisma_models.LiteLLM_BudgetTable | None": ...
+
def model_dump(self) -> AgentRecordDump: ...
def __iter__(self) -> Iterator[tuple[str, object]]: ...
diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
index ef437b029fe..e254430def3 100644
--- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py
+++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
@@ -214,7 +214,7 @@ async def check_agent_budget(auth: UserAPIKeyAuth) -> None:
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,
+ fallback_spend=agent.budget_spend,
max_budget=budget,
fallback_authoritative=True,
)
@@ -254,7 +254,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 effective.litellm_budget_table is None 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
+ and auth.billing_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"))
diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py
index 5ea8266d5ed..a76ec57b630 100644
--- a/litellm/proxy/agent_endpoints/endpoints.py
+++ b/litellm/proxy/agent_endpoints/endpoints.py
@@ -13,7 +13,7 @@ import os
import uuid
from collections.abc import Mapping, Sequence
from types import MappingProxyType
-from typing import Annotated, Final, TypedDict
+from typing import TYPE_CHECKING, Annotated, Final, TypedDict
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import ValidationError
@@ -77,6 +77,7 @@ from litellm.types.agents import (
)
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.proxy.agent_identity import (
+ AgentBudgetState,
AgentIdentityBinding,
AgentIdentityFailure,
EntraIdentityConfig,
@@ -129,6 +130,12 @@ def _build_merged_agent_card(
)
+if TYPE_CHECKING:
+ from prisma.types import LiteLLM_AgentsTableInclude
+
+_AGENT_BUDGET_INCLUDE: Final["LiteLLM_AgentsTableInclude"] = {"litellm_budget_table": True}
+
+
router: Final = APIRouter()
@@ -374,8 +381,10 @@ async def get_agents(
if agent_ids:
db_agents: Final = await agents_table(prisma_client).find_many(
where={"agent_id": {"in": agent_ids}},
+ include=_AGENT_BUDGET_INCLUDE,
)
- spend_map: Final = {a.agent_id: a.spend for a in db_agents}
+ spend_map: Final = MappingProxyType({a.agent_id: a.spend for a in db_agents})
+ budget_map: Final = MappingProxyType({a.agent_id: a for a in db_agents})
for agent in returned_agents:
matched_spends = tuple(
spend_map[alias_id]
@@ -384,6 +393,14 @@ async def get_agents(
)
if matched_spends:
agent.spend = sum(matched_spends)
+ if (budget_row := budget_map.get(agent.agent_id)) is not None:
+ agent.lifetime_budget_spend = budget_row.lifetime_budget_spend
+ agent.budget_id = budget_row.budget_id
+ agent.litellm_budget_table = (
+ AgentBudgetState.model_validate(budget_row.litellm_budget_table.model_dump())
+ if budget_row.litellm_budget_table is not None
+ else None
+ )
await _attach_keys_to_agents(returned_agents, prisma_client)
# add is_public field to each agent - we do it this way, to allow setting config agents as public
@@ -676,7 +693,7 @@ async def get_agent_by_id(
if agent is None:
agent_row: Final = await agents_table(prisma_client).find_unique(
where={"agent_id": agent_id},
- include={"object_permission": True, "identity": True},
+ include={"object_permission": True, "identity": True, "litellm_budget_table": True},
)
if agent_row is not None:
agent_dict: Final = agent_row.model_dump()
@@ -688,9 +705,18 @@ async def get_agent_by_id(
agent = AgentResponse(**agent_dict)
else:
# Agent found in memory — refresh spend from DB
- db_row: Final = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
+ db_row: Final = await agents_table(prisma_client).find_unique(
+ where={"agent_id": agent_id}, include=_AGENT_BUDGET_INCLUDE
+ )
if db_row is not None:
agent.spend = db_row.spend
+ agent.lifetime_budget_spend = db_row.lifetime_budget_spend
+ agent.budget_id = db_row.budget_id
+ agent.litellm_budget_table = (
+ AgentBudgetState.model_validate(db_row.litellm_budget_table.model_dump())
+ if db_row.litellm_budget_table is not None
+ else None
+ )
if agent is None:
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found")
diff --git a/litellm/proxy/agent_endpoints/managed_identity.py b/litellm/proxy/agent_endpoints/managed_identity.py
index 87041e92459..e2289e3bf47 100644
--- a/litellm/proxy/agent_endpoints/managed_identity.py
+++ b/litellm/proxy/agent_endpoints/managed_identity.py
@@ -82,6 +82,7 @@ class ManagedWriteFields(TypedDict, total=False):
litellm_budget_table: ReadOnly[BudgetRelationWrite]
spend_window: ReadOnly[datetime | None]
spend: ReadOnly[float]
+ lifetime_budget_spend: ReadOnly[float]
def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn:
@@ -202,6 +203,11 @@ def _budget_write(raw: object, existing: AgentResponse | None, updated_by: str)
duration_error: Final = budget_duration_error(budget.budget_duration)
if duration_error is not None:
raise ValueError(duration_error)
+ creating_lifetime: Final = budget.budget_duration is None and (
+ existing is None
+ or existing.litellm_budget_table is None
+ or existing.litellm_budget_table.budget_duration is not None
+ )
fields: Final[BudgetFields] = {
"max_budget": budget.max_budget,
"budget_duration": budget.budget_duration,
@@ -218,6 +224,7 @@ def _budget_write(raw: object, existing: AgentResponse | None, updated_by: str)
}
result: Final[ManagedWriteFields] = {
"spend_window": fields["budget_reset_at"],
+ **({"lifetime_budget_spend": 0.0} if creating_lifetime else {}),
**(
{"spend": 0.0}
if fields["budget_reset_at"] is not None
@@ -229,7 +236,9 @@ def _budget_write(raw: object, existing: AgentResponse | None, updated_by: str)
else {}
),
"litellm_budget_table": (
- {"update": fields} if existing and existing.budget_id else {"create": {**fields, "created_by": updated_by}}
+ {"update": fields}
+ if existing and existing.budget_id and not creating_lifetime
+ else {"create": {**fields, "created_by": updated_by}}
),
}
return result
diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py
index c88ef2544d6..55839620831 100644
--- a/litellm/proxy/db/db_spend_update_writer.py
+++ b/litellm/proxy/db/db_spend_update_writer.py
@@ -86,6 +86,8 @@ from litellm.types.agents import agent_spend_filter
from litellm.types.utils import CallTypes
if TYPE_CHECKING:
+ from prisma.types import LiteLLM_AgentsTableUpdateManyMutationInput, LiteLLM_AgentsTableWhereInput
+
from litellm.proxy.db.autorouter_session_rollup import AutoRouterTurnTransaction
from litellm.proxy.db.baseline_accounting import DailyBaselineAttribution
from litellm.proxy.utils import PrismaClient, ProxyLogging
@@ -163,6 +165,26 @@ def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable)
return _ENTITY_SPEND_TABLES[table_accessor](batcher)
+def _queue_lifetime_agent_spend(table: BatchTable, counter_key: str, response_cost: float) -> None:
+ lifetime_filter: Final = agent_spend_filter(counter_key)
+ lifetime_data: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = {
+ "lifetime_budget_spend": {"increment": response_cost}
+ }
+ history_filter: Final[LiteLLM_AgentsTableWhereInput] = {
+ "agent_id": lifetime_filter.get("agent_id"),
+ "spend_window": None,
+ }
+ history_data: Final[LiteLLM_AgentsTableUpdateManyMutationInput] = {"spend": {"increment": response_cost}}
+ table.update_many(
+ where=lifetime_filter,
+ data=lifetime_data,
+ )
+ table.update_many(
+ where=history_filter,
+ data=history_data,
+ )
+
+
class _SpendBatchManager(Protocol):
async def __aenter__(self) -> _SpendBatch: ...
@@ -1361,7 +1383,7 @@ class DBSpendUpdateWriter:
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:
+ if counter_key is not None and agent_spend_filter(counter_key).get("agent_id") != agent_id:
raise ValueError("Agent spend counter does not match the billed agent")
await self.spend_update_queue.add_update(
@@ -2348,6 +2370,13 @@ class DBSpendUpdateWriter:
entity_id,
response_cost,
)
+ if table_accessor == "litellm_agentstable" and entity_id.startswith(
+ "spend:agent_lifetime:"
+ ):
+ _queue_lifetime_agent_spend(
+ _entity_spend_table(batcher, table_accessor), entity_id, response_cost
+ )
+ continue
_entity_spend_table(batcher, table_accessor).update_many(
where=(
agent_spend_filter(entity_id)
diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py
index fc2fbd4dec8..3222868dad8 100644
--- a/litellm/proxy/db/spend_counter_reseed.py
+++ b/litellm/proxy/db/spend_counter_reseed.py
@@ -42,7 +42,11 @@ from litellm.repositories.verification_token_repository import (
from litellm.types.agents import agent_budget_counter_key
if TYPE_CHECKING:
- from prisma.types import LiteLLM_EndUserTableWhereUniqueInput
+ from prisma.types import (
+ LiteLLM_AgentsTableInclude,
+ LiteLLM_AgentsTableWhereUniqueInput,
+ LiteLLM_EndUserTableWhereUniqueInput,
+ )
from litellm.caching.dual_cache import DualCache
from litellm.proxy.utils import PrismaClient
@@ -172,21 +176,37 @@ class SpendCounterReseed:
return await OrganizationRepository(prisma_client).table.find_unique(
where={"organization_id": counter_key[len("spend:org:") :]}
)
+ if counter_key.startswith("spend:agent_lifetime:"):
+ _, _, budget_id, agent_id = counter_key.split(":", 3)
+ lifetime_where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id}
+ lifetime: Final = await AgentsRepository(prisma_client, use_writer=True).table.find_unique(
+ where=lifetime_where
+ )
+ if lifetime is None:
+ return None
+ return lifetime.model_copy(
+ update=MappingProxyType(
+ {"spend": lifetime.lifetime_budget_spend if lifetime.budget_id == budget_id else 0.0}
+ )
+ )
if counter_key.startswith("spend:agent_window:"):
parts: Final = counter_key.split(":", 3)
if len(parts) != 4:
return None
+ window_where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": parts[3]}
+ window_include: Final[LiteLLM_AgentsTableInclude] = {"litellm_budget_table": True}
row: Final = await AgentsRepository(prisma_client, use_writer=True).table.find_unique(
- where={"agent_id": parts[3]}, include={"litellm_budget_table": True}
+ where=window_where, include=window_include
)
if row is None:
return None
current_key: Final = agent_budget_counter_key(row.agent_id, row.spend_window)
- return row if current_key == counter_key else row.model_copy(update={"spend": 0.0})
+ return row if current_key == counter_key else row.model_copy(update=MappingProxyType({"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:") :]}
- )
+ agent_where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {
+ "agent_id": counter_key[len("spend:agent:") :]
+ }
+ return await AgentsRepository(prisma_client).table.find_unique(where=agent_where)
if counter_key.startswith("spend:project:"):
return await ProjectRepository(prisma_client).table.find_unique(
where={"project_id": counter_key[len("spend:project:") :]}
diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py
index c914e04e779..9cf1e0d6eed 100644
--- a/litellm/proxy/hooks/proxy_track_cost_callback.py
+++ b/litellm/proxy/hooks/proxy_track_cost_callback.py
@@ -742,6 +742,8 @@ async def _update_database_and_spend_counters(
tags=request_tags,
model_access_groups=model_access_groups,
project_id=project_id,
+ billing_agent_id=billing_agent_id,
+ billing_agent_counter_key=billing_agent_counter_key,
)
with spend_counter_batch_scope(spend_counter_cache.redis_cache, counter_keys=counter_keys):
return await _update_database_and_spend_counters_in_batch(
diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma
index 6823deadbd3..fc7325ddac3 100644
--- a/litellm/proxy/schema.prisma
+++ b/litellm/proxy/schema.prisma
@@ -86,6 +86,7 @@ model LiteLLM_AgentsTable {
retired_identities LiteLLM_RetiredAgentIdentity[]
budget_id String? @unique
spend_window DateTime?
+ lifetime_budget_spend Float @default(0.0)
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
tpm_limit Int?
rpm_limit Int?
diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py
index 93cba37a90c..62a6bdbe0fd 100644
--- a/litellm/proxy/spend_tracking/budget_reservation.py
+++ b/litellm/proxy/spend_tracking/budget_reservation.py
@@ -511,7 +511,7 @@ async def _get_budget_counters(
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,
+ fallback_spend=agent.budget_spend,
entity_type="Agent",
entity_id=agent.agent_id,
)
diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py
index b3d0e8646e8..5ae30a5c027 100644
--- a/litellm/proxy/spend_tracking/spend_counter_batch.py
+++ b/litellm/proxy/spend_tracking/spend_counter_batch.py
@@ -207,8 +207,11 @@ def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> fr
billing_agent: Final = token.billing_agent_policy
charged_agent_id: Final = billing_agent.agent_id if billing_agent is not None else token.agent_id
agent_keys: Final = (
- frozenset((billing_agent.budget_counter_key if billing_agent is not None else f"spend:agent:{charged_agent_id}",))
- if charged_agent_id is not None else frozenset()
+ frozenset(
+ (billing_agent.budget_counter_key if billing_agent is not None else f"spend:agent:{charged_agent_id}",)
+ )
+ if charged_agent_id is not None
+ else frozenset()
)
return agent_keys | frozenset(
_iter_entity_counter_keys(
@@ -251,13 +254,10 @@ def post_call_counter_keys(
for group in model_access_groups or ()
if group and isinstance(group, str)
)
- agent_key: Final = billing_agent_counter_key or (f"spend:agent:{billing_agent_id}" if billing_agent_id is not None else None)
- return (
- entity_keys
- | tag_keys
- | group_keys
- | (frozenset((agent_key,)) if agent_key else frozenset())
+ agent_key: Final = billing_agent_counter_key or (
+ f"spend:agent:{billing_agent_id}" if billing_agent_id is not None else None
)
+ return entity_keys | tag_keys | group_keys | (frozenset((agent_key,)) if agent_key else frozenset())
def bind_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> None:
diff --git a/litellm/types/agents.py b/litellm/types/agents.py
index 964ef94d227..d59b6fa6ce9 100644
--- a/litellm/types/agents.py
+++ b/litellm/types/agents.py
@@ -316,7 +316,9 @@ class AgentKeySummary(BaseModel):
key_name: str | None = None
-def agent_budget_counter_key(agent_id: str, reset_at: datetime | None) -> str:
+def agent_budget_counter_key(agent_id: str, reset_at: datetime | None, budget_id: str | None = None) -> str:
+ if reset_at is None and budget_id is not None:
+ return f"spend:agent_lifetime:{budget_id}:{agent_id}"
if reset_at is None:
return f"spend:agent:{agent_id}"
aware: Final = reset_at if reset_at.tzinfo is not None else reset_at.replace(tzinfo=timezone.utc)
@@ -325,6 +327,14 @@ def agent_budget_counter_key(agent_id: str, reset_at: datetime | None) -> str:
def agent_spend_filter(counter_key: str) -> "LiteLLM_AgentsTableWhereInput":
+ if counter_key.startswith("spend:agent_lifetime:"):
+ _, _, budget_id, agent_id = counter_key.split(":", 3)
+ lifetime: Final[LiteLLM_AgentsTableWhereInput] = {
+ "agent_id": agent_id,
+ "budget_id": budget_id,
+ "spend_window": None,
+ }
+ return lifetime
if counter_key.startswith("spend:agent_window:"):
_, _, raw_window, agent_id = counter_key.split(":", 3)
window: Final = datetime.strptime(raw_window, "%Y%m%dT%H%M%S.%fZ").replace(tzinfo=timezone.utc)
@@ -339,6 +349,7 @@ def agent_spend_filter(counter_key: str) -> "LiteLLM_AgentsTableWhereInput":
class AgentResponse(BaseModel):
budget_id: str | None = None
+ lifetime_budget_spend: float = 0.0
litellm_budget_table: AgentBudgetState | None = None
identity: AgentIdentityBinding | None = None
identity_managed: bool = False
@@ -369,9 +380,17 @@ class AgentResponse(BaseModel):
@property
def budget_counter_key(self) -> str:
return agent_budget_counter_key(
- self.agent_id, self.litellm_budget_table.budget_reset_at if self.litellm_budget_table else None
+ self.agent_id,
+ self.litellm_budget_table.budget_reset_at if self.litellm_budget_table else None,
+ self.litellm_budget_table.budget_id if self.litellm_budget_table else None,
)
+ @property
+ def budget_spend(self) -> float:
+ if self.litellm_budget_table is not None and self.litellm_budget_table.budget_duration is None:
+ return self.lifetime_budget_spend
+ return self.spend or 0.0
+
class ListAgentsResponse(BaseModel):
agents: list[AgentResponse]
diff --git a/schema.prisma b/schema.prisma
index 6823deadbd3..fc7325ddac3 100644
--- a/schema.prisma
+++ b/schema.prisma
@@ -86,6 +86,7 @@ model LiteLLM_AgentsTable {
retired_identities LiteLLM_RetiredAgentIdentity[]
budget_id String? @unique
spend_window DateTime?
+ lifetime_budget_spend Float @default(0.0)
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
tpm_limit Int?
rpm_limit Int?
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 a2c9fca3bd0..c9b4d301018 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
@@ -108,7 +108,7 @@ async def test_agent_budget_accumulates_across_credentials_and_denies_the_next_a
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:agent_lifetime:budget:agent", 0.0)
counters.set_cache("spend:key:first", 0.0)
counters.set_cache("spend:key:second", 0.0)
monkeypatch.setattr(proxy_server, "spend_counter_cache", counters)
@@ -117,15 +117,15 @@ async def test_agent_budget_accumulates_across_credentials_and_denies_the_next_a
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"
+ token="first", team_id=None, user_id=None, response_cost=0.3, billing_agent_id="agent", billing_agent_counter_key=policy.budget_counter_key
)
await check_agent_budget(auth)
await proxy_server.increment_spend_counters(
- token="second", team_id=None, user_id=None, response_cost=0.3, billing_agent_id="agent"
+ token="second", team_id=None, user_id=None, response_cost=0.3, billing_agent_id="agent", billing_agent_counter_key=policy.budget_counter_key
)
with pytest.raises(litellm.BudgetExceededError):
await check_agent_budget(auth)
- assert await counters.async_get_cache("spend:agent:agent") == pytest.approx(0.6)
+ assert await counters.async_get_cache("spend:agent_lifetime:budget:agent") == pytest.approx(0.6)
assert await counters.async_get_cache("spend:key:first") == pytest.approx(0.3)
assert await counters.async_get_cache("spend:key:second") == pytest.approx(0.3)
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
index 9ed9859717d..33b9a5d4b9d 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
@@ -1231,6 +1231,7 @@ def _stored_agent_row(values: Mapping[str, object] | SimpleNamespace) -> LiteLLM
"created_by": "admin",
"updated_by": "admin",
"spend": 0,
+ "lifetime_budget_spend": 0,
"identity_managed": False,
"enabled": True,
"execution_mode": "autonomous",
@@ -1431,6 +1432,7 @@ async def test_agent_listing_preserves_stored_identity_bindings(bound: bool) ->
enabled=True,
execution_mode="autonomous",
spend=0.0,
+ lifetime_budget_spend=0.0,
agent_access_groups=[],
access_group_ids=[],
extra_headers=[],
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
index 2cf81892db7..0bf9afa832d 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
@@ -1648,3 +1648,31 @@ def test_invalid_identity_and_untrusted_tenant_cannot_be_registered(
with pytest.raises(HTTPException, match=message) as failure:
agent_endpoints._validate_managed_identity_request(request)
assert failure.value.status_code == 400
+
+
+@pytest.mark.parametrize("cached,path", [(False, "/v1/agents/agent-123"), (True, "/v1/agents/agent-123"), (True, "/v1/agents")])
+def test_agent_budget_readback_refreshes_consumption_and_limit(monkeypatch: pytest.MonkeyPatch, cached: bool, path: str) -> None:
+ from prisma.models import LiteLLM_BudgetTable
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
+
+ row: Final = LiteLLM_AgentsTable.model_construct(
+ agent_id="agent-123", agent_name="Agent", agent_card_params={}, spend=12.5,
+ lifetime_budget_spend=0.75, budget_id="budget", litellm_params=None,
+ litellm_budget_table=LiteLLM_BudgetTable.model_construct(budget_id="budget", max_budget=2.0),
+ )
+ registry: Final = AgentRegistry()
+ if cached:
+ registry.register_agent(_sample_agent_response())
+ table: Final = SimpleNamespace(find_unique=AsyncMock(return_value=row), find_many=AsyncMock(return_value=[row]))
+ database: Final = SimpleNamespace(litellm_agentstable=table, litellm_verificationtoken=SimpleNamespace(find_many=AsyncMock(return_value=[])))
+ monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database, writer_db=database))
+ monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ response: Final = client.get(path)
+ assert response.status_code == 200, response.text
+ payload: Final = response.json()[0] if path == "/v1/agents" else response.json()
+ assert payload["spend"] == 12.5
+ assert payload["lifetime_budget_spend"] == 0.75
+ assert payload["litellm_budget_table"]["max_budget"] == 2.0
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py
index 424330e078e..55b1a0d366e 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py
@@ -311,3 +311,46 @@ def test_budget_write_stamps_the_same_window_on_the_agent_row() -> None:
assert not isinstance(result, AgentIdentityFailure)
assert result["spend_window"] == result["litellm_budget_table"]["create"]["budget_reset_at"]
assert result["spend"] == 0.0
+
+
+def test_new_lifetime_budget_starts_unused_without_erasing_historical_spend() -> None:
+ existing: Final = managed_agent().model_copy(update={"spend": 12.5})
+ result: Final = managed_write_fields({"budget": {"max_budget": 1.0}}, existing, "admin")
+ assert not isinstance(result, AgentIdentityFailure)
+ assert "spend" not in result
+ assert result["lifetime_budget_spend"] == 0.0
+ assert result["litellm_budget_table"]["create"]["max_budget"] == 1.0
+ assert existing.spend == 12.5
+
+
+def test_editing_lifetime_budget_preserves_consumption() -> None:
+ from litellm.types.proxy.agent_identity import AgentBudgetState
+
+ existing: Final = managed_agent().model_copy(update={
+ "spend": 12.5, "lifetime_budget_spend": 0.75, "budget_id": "budget",
+ "litellm_budget_table": AgentBudgetState(budget_id="budget", max_budget=1.0),
+ })
+ result: Final = managed_write_fields({"budget": {"max_budget": 2.0}}, existing, "admin")
+ assert not isinstance(result, AgentIdentityFailure)
+ assert "spend" not in result
+ assert "lifetime_budget_spend" not in result
+ assert result["litellm_budget_table"]["update"]["max_budget"] == 2.0
+
+
+@pytest.mark.parametrize("previous_duration", (None, "1d"))
+def test_recreated_or_converted_lifetime_budget_gets_a_fresh_allowance(previous_duration: str | None) -> None:
+ from litellm.types.proxy.agent_identity import AgentBudgetState
+
+ existing: Final = managed_agent().model_copy(update={
+ "spend": 12.5, "lifetime_budget_spend": 0.75,
+ "budget_id": "previous" if previous_duration else None,
+ "litellm_budget_table": AgentBudgetState(
+ budget_id="previous", max_budget=1.0, budget_duration=previous_duration,
+ ) if previous_duration else None,
+ })
+ result: Final = managed_write_fields({"budget": {"max_budget": 2.0}}, existing, "admin")
+ assert not isinstance(result, AgentIdentityFailure)
+ assert result["lifetime_budget_spend"] == 0.0
+ assert "spend" not in result
+ assert "create" in result["litellm_budget_table"]
+ assert "update" not in result["litellm_budget_table"]
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index 9d0bc1c3475..a60d64bd562 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -9762,6 +9762,8 @@ async def test_human_agent_discovery_does_not_reserve_target_budget_but_send_and
user_id="human",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["agent"]),
)
+ database.get_data = AsyncMock(return_value=auth)
+ proxy_server.proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock(return_value=None)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
new_callable=AsyncMock,
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 ca6a9411005..ae338c518b3 100644
--- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py
+++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py
@@ -4885,3 +4885,32 @@ async def test_agent_settlement_charges_only_the_matching_current_window(capture
"agent_id": "window-agent",
"spend_window": captured_window,
}
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("captured_budget", ("current", "retired"))
+async def test_lifetime_settlement_preserves_history_without_charging_another_budget(captured_budget: str) -> None:
+ from types import SimpleNamespace
+
+ row: Final = SimpleNamespace(
+ agent_id="agent", budget_id="current", spend_window=None, spend=12.5, lifetime_budget_spend=0.25
+ )
+
+ def apply_update(*, where: dict[str, object], data: dict[str, dict[str, float]]) -> None:
+ if all(getattr(row, key) == value for key, value in where.items()):
+ for field, operation in data.items():
+ setattr(row, field, getattr(row, field) + operation["increment"])
+
+ batcher: Final = MagicMock()
+ batcher.litellm_agentstable.update_many.side_effect = apply_update
+ transaction: Final = AsyncMock()
+ transaction.batch_ = MagicMock(return_value=AsyncMock(__aenter__=AsyncMock(return_value=batcher)))
+ client: Final = MagicMock()
+ client.db.tx.return_value = AsyncMock(__aenter__=AsyncMock(return_value=transaction))
+ await DBSpendUpdateWriter._update_entity_spend_in_db(
+ entity_name="Agent", transactions={f"spend:agent_lifetime:{captured_budget}:agent": 0.25},
+ table_accessor="litellm_agentstable", where_field="agent_id", n_retry_times=0,
+ prisma_client=client, proxy_logging_obj=MagicMock(),
+ )
+ assert row.spend == 12.75
+ assert row.lifetime_budget_spend == (0.5 if captured_budget == "current" else 0.25)
diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py
index 391cede1005..8e742c06b51 100644
--- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py
+++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py
@@ -523,3 +523,17 @@ async def test_agent_window_reseed_handles_missing_rows_and_malformed_keys(missi
key: Final = "spend:agent_window:20260102T000000.000000Z:missing" if missing else "spend:agent_window:malformed"
assert await SpendCounterReseed.from_db(client, key) is None
assert lookup.await_count == int(missing)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("budget_id,expected", [("current", 0.25), ("retired", 0.0), ("deleted", None)])
+async def test_lifetime_counter_reseed_excludes_historical_and_other_budget_spend(budget_id: str, expected: float | None) -> None:
+ from prisma.models import LiteLLM_AgentsTable
+
+ row: Final = LiteLLM_AgentsTable.model_construct(
+ agent_id="agent", budget_id="current", spend=12.5, lifetime_budget_spend=0.25,
+ )
+ lookup: Final = AsyncMock(return_value=row if expected is not None else None)
+ client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_agentstable=SimpleNamespace(find_unique=lookup)))
+ assert await SpendCounterReseed.from_db(client, f"spend:agent_lifetime:{budget_id}:agent") == expected
+ lookup.assert_awaited_once_with(where={"agent_id": "agent"})
diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
index 7096bc7c632..b5ca6f5c42d 100644
--- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
+++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
@@ -4759,6 +4759,7 @@ def _agent_db_row(agent_id: str, agent_name: str):
agent_access_groups=[],
access_group_ids=[],
spend=0.0,
+ lifetime_budget_spend=0.0,
identity_managed=False,
enabled=True,
execution_mode="autonomous",
diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py
index 9d64fdb5aab..02686a13a85 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py
@@ -317,13 +317,13 @@ async def test_agent_invocation_reserves_exact_fee_and_reconciles_without_child_
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}"
+ counter_key: Final = f"spend:agent_window:20260102T000000.000000Z:{charged_agent}" if window else f"spend:agent_lifetime:agent-budget:{charged_agent}"
cache.set_cache(counter_key, 0.1)
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
auth: Final = UserAPIKeyAuth(agent_id="caller-agent" if charged_agent == "caller-agent" else None)
auth.billing_agent_policy = AgentResponse(
- agent_id=charged_agent, agent_name="Charged agent", agent_card_params={}, spend=0.1,
- litellm_budget_table={"budget_id": "agent-budget", "max_budget": 0.5, "budget_reset_at": window},
+ agent_id=charged_agent, agent_name="Charged agent", agent_card_params={}, spend=0.1, lifetime_budget_spend=0.1,
+ litellm_budget_table={"budget_id": "agent-budget", "max_budget": 0.5, "budget_reset_at": window, "budget_duration": "1d" if window else None},
)
auth.invoked_agent_id = "target-agent"
auth.agent_invocation_cost = 0.2
@@ -359,11 +359,11 @@ async def test_agent_invocation_over_budget_is_rejected_and_reservation_is_refun
from litellm.types.agents import AgentResponse
cache: Final = DualCache()
- cache.set_cache("spend:agent:agent", 0.4)
+ cache.set_cache("spend:agent_lifetime:agent-budget:agent", 0.4)
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
auth: Final = UserAPIKeyAuth(agent_id="agent")
auth.billing_agent_policy = AgentResponse(
- agent_id="agent", agent_name="Charged agent", agent_card_params={}, spend=0.4,
+ agent_id="agent", agent_name="Charged agent", agent_card_params={}, spend=0.4, lifetime_budget_spend=0.4,
litellm_budget_table={"budget_id": "agent-budget", "max_budget": 0.5},
)
auth.agent_invocation_cost = 0.2
@@ -374,7 +374,7 @@ async def test_agent_invocation_over_budget_is_rejected_and_reservation_is_refun
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)
+ assert await cache.async_get_cache("spend:agent_lifetime:agent-budget:agent") == pytest.approx(0.4)
@pytest.mark.asyncio
@@ -413,3 +413,68 @@ async def test_release_unbound_budget_reservation_leaves_a_bound_one_to_its_call
assert spend_counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(reservation["reserved_cost"])
assert reservation["finalized"] is False
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("outcome", ("failed", "cancelled", "completed"))
+async def test_budgeted_caller_reserves_unmanaged_agent_fees_before_concurrent_admission(
+ spend_counter_cache: DualCache, monkeypatch: pytest.MonkeyPatch, outcome: str
+) -> None:
+ import asyncio
+
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+ from litellm.proxy.spend_tracking.budget_reservation import (
+ reconcile_budget_reservation,
+ release_budget_reservation,
+ release_budget_reservation_on_cancel,
+ )
+ from litellm.types.agents import AgentResponse
+
+ caller: Final = AgentResponse(
+ agent_id="caller", agent_name="Caller", agent_card_params={}, spend=0.0,
+ litellm_budget_table={"budget_id": "caller-budget", "max_budget": 0.5},
+ )
+ target: Final = AgentResponse(
+ agent_id="target", agent_name="Target", agent_card_params={}, litellm_params={"cost_per_query": 0.25},
+ )
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(target)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+
+ async def admit() -> dict[str, object] | None:
+ auth: Final = UserAPIKeyAuth(agent_id="caller", user_role="proxy_admin")
+ auth.billing_agent_policy = caller
+ await prepare_agent_invocation(auth, "target", None)
+ return await reserve_budget_for_request(
+ request_body={"method": "message/send"}, route="/a2a/target", llm_router=None,
+ valid_token=auth, team_object=None, user_object=None, prisma_client=None,
+ user_api_key_cache=UserApiKeyCache(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
+ fail_closed_budget_enforcement=True,
+ )
+
+ results: Final = await asyncio.gather(*(admit() for _ in range(8)), return_exceptions=True)
+ accepted: Final = tuple(result for result in results if isinstance(result, dict))
+ rejected: Final = tuple(result for result in results if isinstance(result, litellm.BudgetExceededError))
+ assert all(result is None or isinstance(result, (dict, litellm.BudgetExceededError)) for result in results), results
+ assert len(accepted) == 2, results
+ assert len(rejected) == 6
+ assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5)
+ first: Final = accepted[0]
+ if outcome == "completed":
+ await proxy_server.increment_spend_counters(
+ token=None, team_id=None, user_id=None, response_cost=0.25,
+ billing_agent_id=caller.agent_id, billing_agent_counter_key=caller.budget_counter_key,
+ budget_reservation=first,
+ )
+ await reconcile_budget_reservation(first, actual_cost=0.25)
+ assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5)
+ with pytest.raises(litellm.BudgetExceededError):
+ await admit()
+ else:
+ release: Final = release_budget_reservation_on_cancel if outcome == "cancelled" else release_budget_reservation
+ await release(first)
+ await release(first)
+ assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.25)
+ assert await admit() is not None
+ assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5)
diff --git a/tests/unit/types/test_agents.py b/tests/unit/types/test_agents.py
index a583ab9e8a2..9acb7c41b2a 100644
--- a/tests/unit/types/test_agents.py
+++ b/tests/unit/types/test_agents.py
@@ -26,3 +26,16 @@ def test_different_budget_windows_never_share_a_settlement_filter() -> None:
assert agent_spend_filter(agent_budget_counter_key("agent-one", first)) != agent_spend_filter(
agent_budget_counter_key("agent-one", second)
)
+
+
+def test_lifetime_budget_consumption_is_separate_from_agent_history() -> None:
+ from litellm.types.agents import AgentResponse
+
+ agent: Final = AgentResponse(
+ agent_id="agent", agent_name="Agent", agent_card_params={}, spend=12.5,
+ lifetime_budget_spend=0.75,
+ litellm_budget_table={"budget_id": "budget", "max_budget": 1.0},
+ )
+ assert agent.budget_spend == 0.75
+ assert agent.spend == 12.5
+ assert agent.budget_counter_key == "spend:agent_lifetime:budget:agent"
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts
index 10ca2114186..f54d6364ddc 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts
@@ -1,5 +1,6 @@
import { describe, expect, it } from "vitest";
import {
+ agentBudgetSpend,
buildIdentityParams,
entraTenantFromIssuer,
parseIdentityForForm,
@@ -84,3 +85,25 @@ describe("agent identity configuration", () => {
expect(entraTenantFromIssuer("https://login.microsoftonline.com/common/v2.0")).toBeNull();
});
});
+
+describe("agent budget consumption", () => {
+ it("keeps historical spend separate from lifetime budget consumption", () => {
+ expect(
+ agentBudgetSpend({
+ spend: 12.5,
+ lifetime_budget_spend: 0.75,
+ litellm_budget_table: { budget_id: "budget", max_budget: 1 },
+ }),
+ ).toBe(0.75);
+ });
+ it("preserves recurring window spend and defaults missing lifetime consumption to zero", () => {
+ expect(
+ agentBudgetSpend({
+ spend: 0.5,
+ lifetime_budget_spend: 9,
+ litellm_budget_table: { budget_id: "budget", max_budget: 1, budget_duration: "1d" },
+ }),
+ ).toBe(0.5);
+ expect(agentBudgetSpend({ spend: 12.5, litellm_budget_table: { budget_id: "budget", max_budget: 1 } })).toBe(0);
+ });
+});
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
index 2db4f97e128..24f25a54f65 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
@@ -100,23 +100,26 @@ export const withAgentIdentity = (
const hasCard = !existing || cardEdited || Object.keys(existing.agent_card_params ?? {}).length > 0;
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 {
...settings,
...(hasCard && agent_card_params ? { agent_card_params } : {}),
...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 } : {}),
+ ...(values.agent_max_budget === undefined
+ ? {}
+ : {
+ budget:
+ values.agent_max_budget === "" || values.agent_max_budget === null
+ ? null
+ : { max_budget: Number(values.agent_max_budget), budget_duration: values.agent_budget_duration || null },
+ }),
};
};
+
+export const agentBudgetSpend = (
+ agent: Pick,
+): number =>
+ agent.litellm_budget_table && !agent.litellm_budget_table.budget_duration
+ ? agent.lifetime_budget_spend ?? 0
+ : agent.spend ?? 0;
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx
index e08cab776c4..3ba83920589 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx
@@ -156,6 +156,44 @@ describe("AgentInfoView update payload", () => {
.mockResolvedValue({} as never);
});
+ it.each(["preserve", "edit", "clear"])(
+ "%s lifetime budget without displaying historical spend as consumption",
+ async (action) => {
+ const budgetedAgent = {
+ ...A2A_AGENT,
+ spend: 12.5,
+ lifetime_budget_spend: 0.5,
+ litellm_budget_table: { budget_id: "budget", max_budget: 1, budget_duration: null },
+ };
+ vi.mocked(networking.getAgentInfo).mockResolvedValue(budgetedAgent);
+ const user = setup();
+ renderView();
+ expect(await screen.findByText("$0.5 / $1")).toBeInTheDocument();
+ await openEditor(user);
+ const limit = screen.getByLabelText("Aggregate Agent Budget ($)");
+ expect(limit).toHaveValue(1);
+ expect(screen.getByLabelText("Budget Reset Period")).toHaveValue("");
+ if (action !== "preserve") {
+ fireEvent.change(limit, { target: { value: action === "edit" ? "2" : "" } });
+ }
+ await save(user);
+ expect(patchedPayload().budget).toEqual(
+ action === "clear" ? null : { max_budget: action === "edit" ? 2 : 1, budget_duration: null },
+ );
+ },
+ );
+
+ it("creates a recurring budget from the budget fields", async () => {
+ const user = setup();
+ renderView();
+ expect(await screen.findByText("No aggregate limit")).toBeInTheDocument();
+ await openEditor(user);
+ fireEvent.change(screen.getByLabelText("Aggregate Agent Budget ($)"), { target: { value: "0.75" } });
+ fireEvent.change(screen.getByLabelText("Budget Reset Period"), { target: { value: "1d" } });
+ await save(user);
+ expect(patchedPayload().budget).toEqual({ max_budget: 0.75, budget_duration: "1d" });
+ });
+
it.each([
{ card: "complete", editCard: false },
{ card: "empty", editCard: false },
@@ -222,7 +260,8 @@ describe("AgentInfoView update payload", () => {
await save(user);
- expect(patchedPayload()).toEqual({
+ const expectedPayload = {
+ budget: null,
agent_name: "my-agent",
agent_card_params: {
protocolVersion: "1.0",
@@ -241,7 +280,8 @@ describe("AgentInfoView update payload", () => {
session_rpm_limit: 444,
object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} },
access_group_ids: [],
- });
+ };
+ expect(patchedPayload()).toEqual(expectedPayload);
});
it("sends the loaded values of every panel the user opens", async () => {
@@ -259,7 +299,8 @@ describe("AgentInfoView update payload", () => {
await save(user);
- expect(patchedPayload()).toEqual({
+ const expectedPayload = {
+ budget: null,
agent_name: "my-agent",
agent_card_params: {
protocolVersion: "1.0",
@@ -283,7 +324,8 @@ describe("AgentInfoView update payload", () => {
session_rpm_limit: 444,
object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} },
access_group_ids: [],
- });
+ };
+ expect(patchedPayload()).toEqual(expectedPayload);
});
it("clamps a rate limit typed below its minimum up to that minimum", async () => {
@@ -342,7 +384,8 @@ describe("AgentInfoView update payload", () => {
await save(user);
- expect(patchedPayload()).toEqual({
+ const expectedPayload = {
+ budget: null,
agent_name: "lg-agent",
agent_card_params: {
protocolVersion: "1.0",
@@ -362,7 +405,8 @@ describe("AgentInfoView update payload", () => {
},
object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} },
access_group_ids: [],
- });
+ };
+ expect(patchedPayload()).toEqual(expectedPayload);
});
it("keeps the agent's existing MCP grants in the update payload", async () => {
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx
index 24fd7e9274f..491496b8a07 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx
@@ -1,6 +1,6 @@
import { AgentIdentityFields } from "./AgentIdentityFields";
import { AgentIdentityDetails } from "./AgentIdentityDetails";
-import { withAgentIdentity } from "./agent_identity";
+import { agentBudgetSpend, withAgentIdentity } from "./agent_identity";
import React, { useState, useEffect, useMemo } from "react";
import { cx } from "@/lib/cva.config";
import { FormProvider, useForm, useWatch } from "react-hook-form";
@@ -82,7 +82,7 @@ const AgentBudgetDetails = ({ agent }: { agent: Agent }) => (
<>
{agent.litellm_budget_table?.max_budget != null
- ? `$${agent.spend ?? 0} / $${agent.litellm_budget_table.max_budget}`
+ ? `$${agentBudgetSpend(agent)} / $${agent.litellm_budget_table.max_budget}`
: "No aggregate limit"}
diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts
index 703d6170764..2df1c53333c 100644
--- a/ui/litellm-dashboard/src/components/agents/types.ts
+++ b/ui/litellm-dashboard/src/components/agents/types.ts
@@ -33,6 +33,7 @@ export interface Agent {
kill_switch?: AgentKillSwitchConfig | null;
keys?: AgentAttachedKey[] | null;
spend?: number;
+ lifetime_budget_spend?: number;
tpm_limit?: number | null;
rpm_limit?: number | null;
session_tpm_limit?: number | null;
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index a44326e4a7c..f199f132812 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -25016,6 +25016,11 @@ export interface components {
/** Keys */
keys?: components["schemas"]["AgentKeySummary"][] | null;
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
+ /**
+ * Lifetime Budget Spend
+ * @default 0
+ */
+ lifetime_budget_spend: number;
litellm_budget_table?: components["schemas"]["AgentBudgetState"] | null;
/** Litellm Params */
litellm_params?: {
From b33fda1dd3022fee411a50fc3dbbed731336cf45 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Wed, 30 Sep 2026 16:54:32 -0700
Subject: [PATCH 07/12] fix(agents): enforce and settle chat adapter invocation
fees
---
litellm/cost_calculator.py | 9 ++-
litellm/main.py | 1 +
litellm/proxy/agent_endpoints/a2a_routing.py | 12 +++-
litellm/proxy/auth/user_api_key_auth.py | 7 +-
.../proxy/auth/test_user_api_key_auth.py | 3 +
.../unit/a2a_protocol/test_cost_calculator.py | 68 +++++++++++++++++++
6 files changed, 96 insertions(+), 4 deletions(-)
diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py
index 238b7cc3fdd..c022b34282f 100644
--- a/litellm/cost_calculator.py
+++ b/litellm/cost_calculator.py
@@ -1529,7 +1529,14 @@ def completion_cost(
completion_tokens = token_counter(model=model, text=completion)
# Handle A2A calls before model check - A2A doesn't require a model
- if call_type in _A2A_CALL_TYPES:
+ if call_type in _A2A_CALL_TYPES or (
+ custom_llm_provider == "a2a"
+ and litellm_logging_obj is not None
+ and (litellm_logging_obj.model_call_details.get("litellm_params") or MappingProxyType({})).get(
+ "cost_per_query"
+ )
+ is not None
+ ):
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
return A2ACostCalculator.calculate_a2a_cost(litellm_logging_obj=litellm_logging_obj)
diff --git a/litellm/main.py b/litellm/main.py
index 8c9d7f2513d..570366e0792 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -5662,6 +5662,7 @@ def completion(
preset_cache_key=preset_cache_key,
no_log=no_log,
cost_per_second=cost_per_second,
+ cost_per_query=kwargs.get("cost_per_query"),
input_cost_per_second=input_cost_per_second,
input_cost_per_token=input_cost_per_token,
output_cost_per_second=output_cost_per_second,
diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py
index c57315ebc21..1b6f210cf8e 100644
--- a/litellm/proxy/agent_endpoints/a2a_routing.py
+++ b/litellm/proxy/agent_endpoints/a2a_routing.py
@@ -5,6 +5,7 @@ Handles routing for A2A agents (models with "a2a/" prefix).
Looks up agents in the registry and injects their API base URL.
"""
+from types import MappingProxyType
from typing import Any, Final
from fastapi import HTTPException
@@ -78,4 +79,13 @@ async def route_a2a_agent_request(
data["api_base"] = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
- return getattr(litellm, f"{route_type}")(**data)
+ invocation_pricing: Final = (
+ MappingProxyType({"cost_per_query": user_api_key_dict.agent_invocation_cost})
+ if user_api_key_dict is not None
+ and user_api_key_dict.agent_invocation_cost is not None
+ and user_api_key_dict.invoked_agent_policy is not None
+ and (user_api_key_dict.invoked_agent_policy.litellm_params or MappingProxyType({})).get("cost_per_query")
+ is not None
+ else MappingProxyType({})
+ )
+ return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, **invocation_pricing}))
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index 278a48911d4..a684b5ebcf3 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -3261,8 +3261,11 @@ async def _authorize_authenticated_request(
target_name,
store,
billable=request.method == "POST"
- and request_data.get("method")
- in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"),
+ and (
+ not RouteChecks.check_route_access(route, ("/a2a/{agent_id}", "/v1/a2a/{agent_id}"))
+ or request_data.get("method")
+ in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage")
+ ),
)
await _run_centralized_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index a60d64bd562..d69b3e72183 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -9701,6 +9701,9 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon
("POST", "/a2a/agent", {"jsonrpc": "2.0", "id": "1", "method": "message/stream", "params": {}}, True),
("POST", "/a2a/agent", {"method": "message/send", "model": "free-model", "params": {}}, True),
("POST", "/a2a/agent", {"method": "message/stream", "model": "free-model", "params": {}}, True),
+ ("POST", "/v1/chat/completions", {"model": "a2a/agent", "method": "tasks/get"}, True),
+ ("POST", "/chat/completions", {"model": "a2a/agent", "method": "tasks/cancel", "stream": True}, True),
+ ("POST", "/v1/a2a/agent/message/send", {"method": "tasks/get", "params": {}}, True),
],
)
async def test_human_agent_discovery_does_not_reserve_target_budget_but_send_and_stream_do(
diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py
index 8d8ec815f3a..4f2a93fa883 100644
--- a/tests/unit/a2a_protocol/test_cost_calculator.py
+++ b/tests/unit/a2a_protocol/test_cost_calculator.py
@@ -117,6 +117,7 @@ class CostLogger(CustomLogger):
def __init__(self):
self.response_cost: Optional[float] = None
+ self.logged: asyncio.Event = asyncio.Event()
super().__init__()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
@@ -125,6 +126,7 @@ class CostLogger(CustomLogger):
self.response_cost = (
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
)
+ self.logged.set()
@pytest.mark.asyncio
@@ -449,3 +451,69 @@ async def test_asend_message_streaming_triggers_callbacks():
assert callback_logger.agent_id == test_agent_id, (
f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("stream", (False, True))
+@pytest.mark.parametrize("claimed_fee", (None, 0.0, 99.0))
+async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing(
+ monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None
+) -> None:
+ import json
+ from typing import Final
+
+ import httpx
+
+ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
+ from litellm.types.agents import AgentResponse
+
+ await _reset_callbacks_and_settle_pending_logs()
+ logger: Final = CostLogger()
+ monkeypatch.setattr(litellm, "callbacks", [logger])
+ target: Final = AgentResponse(
+ agent_id="fee-target", agent_name="fee-target",
+ agent_card_params={"url": "https://agent.test/", "capabilities": {"streaming": True}},
+ litellm_params={"cost_per_query": 0.25},
+ )
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(target)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
+ auth.invoked_agent_policy = target
+ auth.invoked_agent_id = target.agent_id
+ auth.agent_invocation_cost = 0.25
+
+ def reply(request: httpx.Request) -> httpx.Response:
+ body: Final = json.loads(request.content)
+ assert body["method"] == ("message/stream" if stream else "message/send")
+ result: Final = {
+ "jsonrpc": "2.0", "id": body["id"],
+ "result": {"kind": "message", "messageId": "reply", "role": "agent",
+ "parts": [{"kind": "text", "text": "Paid reply"}]},
+ }
+ if stream:
+ return httpx.Response(200, text=f"data: {json.dumps(result)}\n\n", headers={"content-type": "text/event-stream"})
+ return httpx.Response(200, json=result)
+
+ client: Final = AsyncHTTPHandler(transport=httpx.MockTransport(reply))
+ try:
+ pending: Final = await route_a2a_agent_request(
+ data={"model": "a2a/fee-target", "messages": [{"role": "user", "content": "Hello"}],
+ "stream": stream, "client": client,
+ **({"cost_per_query": claimed_fee} if claimed_fee is not None else {})},
+ route_type="acompletion", user_api_key_dict=auth,
+ )
+ response: Final = await pending
+ if stream:
+ chunks: Final = tuple([chunk async for chunk in response])
+ assert any(chunk.choices[0].delta.content == "Paid reply" for chunk in chunks)
+ else:
+ assert response.choices[0].message.content == "Paid reply"
+ await asyncio.wait_for(logger.logged.wait(), timeout=10.0)
+ await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
+ assert logger.response_cost == pytest.approx(0.25)
+ finally:
+ await client.close()
From 95b94520d944afdfe0ab78f164d0ae2ce8f81a5b Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Wed, 30 Sep 2026 17:16:36 -0700
Subject: [PATCH 08/12] fix(agents): reject client supplied invocation pricing
---
litellm/proxy/agent_endpoints/a2a_routing.py | 19 +++++++--------
.../auth/managed_authorization.py | 6 ++---
litellm/proxy/litellm_pre_call_utils.py | 4 +++-
.../proxy/test_pricing_field_strip.py | 5 +++-
.../unit/a2a_protocol/test_cost_calculator.py | 23 ++++++++++++-------
5 files changed, 35 insertions(+), 22 deletions(-)
diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py
index 1b6f210cf8e..29d31251f94 100644
--- a/litellm/proxy/agent_endpoints/a2a_routing.py
+++ b/litellm/proxy/agent_endpoints/a2a_routing.py
@@ -13,6 +13,7 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
+from litellm.proxy.agent_endpoints.auth.managed_authorization import AGENT_INVOCATION_COST
async def route_a2a_agent_request(
@@ -79,13 +80,13 @@ async def route_a2a_agent_request(
data["api_base"] = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
- invocation_pricing: Final = (
- MappingProxyType({"cost_per_query": user_api_key_dict.agent_invocation_cost})
- if user_api_key_dict is not None
- and user_api_key_dict.agent_invocation_cost is not None
- and user_api_key_dict.invoked_agent_policy is not None
- and (user_api_key_dict.invoked_agent_policy.litellm_params or MappingProxyType({})).get("cost_per_query")
- is not None
- else MappingProxyType({})
+ pricing_policy: Final = (
+ user_api_key_dict.invoked_agent_policy
+ if user_api_key_dict is not None and user_api_key_dict.invoked_agent_policy is not None
+ else agent
)
- return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, **invocation_pricing}))
+ configured_fee: Final = (pricing_policy.litellm_params or MappingProxyType({})).get("cost_per_query")
+ invocation_fee: Final = (
+ AGENT_INVOCATION_COST.validate_python(configured_fee) if configured_fee is not None else None
+ )
+ return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, "cost_per_query": invocation_fee}))
diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
index e254430def3..11c8c790ad2 100644
--- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py
+++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
@@ -222,7 +222,7 @@ async def check_agent_budget(auth: UserAPIKeyAuth) -> None:
raise litellm.BudgetExceededError(current_cost=spend, max_budget=budget, message="Agent budget exceeded")
-_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)])
+AGENT_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)])
def invocation_target(route: str, body: Mapping[str, object]) -> str | None:
@@ -280,13 +280,13 @@ async def prepare_agent_invocation(
and billing_policy.litellm_budget_table.max_budget is not None
)
try:
- fee: Final = _INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0)
+ fee: Final = AGENT_INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0)
unbounded_token_price: Final = (
billable
and bounded
and fixed_fee is None
and any(
- _INVOCATION_COST.validate_python(pricing[field]) > 0
+ AGENT_INVOCATION_COST.validate_python(pricing[field]) > 0
for field in ("input_cost_per_token", "output_cost_per_token")
if pricing.get(field) is not None
)
diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py
index 27d22920591..b4272d5e170 100644
--- a/litellm/proxy/litellm_pre_call_utils.py
+++ b/litellm/proxy/litellm_pre_call_utils.py
@@ -362,7 +362,9 @@ _ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY: Final = "allow_client_mess
# not to user-supplied request bodies, so the proxy strips them before they
# reach the call path. Built from the Pydantic model so newly-added pricing
# fields are covered automatically.
-_CLIENT_PRICING_CONTROL_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys())
+_CLIENT_PRICING_CONTROL_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys()) | frozenset(
+ {"cost_per_query"}
+)
# ``model_info`` carries the same pricing fields when read by
# ``use_custom_pricing_for_model``; strip from metadata for the same reason.
# ``standard_logging_guardrail_information`` is proxy-written telemetry summed
diff --git a/tests/test_litellm/proxy/test_pricing_field_strip.py b/tests/test_litellm/proxy/test_pricing_field_strip.py
index a0e25e91f37..87cf39b0d49 100644
--- a/tests/test_litellm/proxy/test_pricing_field_strip.py
+++ b/tests/test_litellm/proxy/test_pricing_field_strip.py
@@ -60,7 +60,7 @@ class TestStripClientPricingOverrides:
# set drifting apart if someone replaces the auto-derivation later.
assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset(
CustomPricingLiteLLMParams.model_fields.keys()
- )
+ ) | {"cost_per_query"}
# Sanity: the obvious top-level pricing fields are in the set.
for field in (
"input_cost_per_token",
@@ -78,6 +78,7 @@ class TestStripClientPricingOverrides:
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"cache_creation_input_token_cost": 0.0,
+ "cost_per_query": -1000.0,
}
_strip_client_pricing_overrides(data)
assert data == {
@@ -192,6 +193,7 @@ class TestStripClientPricingOverrides:
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_root_pricing_fields():
data = {
+ "cost_per_query": -1000.0,
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
"input_cost_per_token": 0.0,
@@ -207,6 +209,7 @@ async def test_add_litellm_data_to_request_strips_root_pricing_fields():
version="test-version",
)
+ assert "cost_per_query" not in updated
assert "input_cost_per_token" not in updated
assert "output_cost_per_token" not in updated
diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py
index 4f2a93fa883..d4087bab2ea 100644
--- a/tests/unit/a2a_protocol/test_cost_calculator.py
+++ b/tests/unit/a2a_protocol/test_cost_calculator.py
@@ -455,9 +455,12 @@ async def test_asend_message_streaming_triggers_callbacks():
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", (False, True))
-@pytest.mark.parametrize("claimed_fee", (None, 0.0, 99.0))
+@pytest.mark.parametrize("claimed_fee", (None, -1000.0, 0.0, 99.0))
+@pytest.mark.parametrize("admitted", (False, True))
+@pytest.mark.parametrize("configured_fee", (None, 0.25))
async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing(
- monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None
+ monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None,
+ admitted: bool, configured_fee: float | None
) -> None:
import json
from typing import Final
@@ -476,15 +479,16 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
target: Final = AgentResponse(
agent_id="fee-target", agent_name="fee-target",
agent_card_params={"url": "https://agent.test/", "capabilities": {"streaming": True}},
- litellm_params={"cost_per_query": 0.25},
+ litellm_params={"cost_per_query": configured_fee} if configured_fee is not None else {},
)
registry: Final = agent_registry.AgentRegistry()
- registry.register_agent(target)
+ registry.register_agent(target.model_copy(update={"litellm_params": {"cost_per_query": 0.5}}) if admitted else target)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
- auth.invoked_agent_policy = target
- auth.invoked_agent_id = target.agent_id
- auth.agent_invocation_cost = 0.25
+ if admitted:
+ auth.invoked_agent_policy = target
+ auth.invoked_agent_id = target.agent_id
+ auth.agent_invocation_cost = configured_fee or 0.0
def reply(request: httpx.Request) -> httpx.Response:
body: Final = json.loads(request.content)
@@ -514,6 +518,9 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
assert response.choices[0].message.content == "Paid reply"
await asyncio.wait_for(logger.logged.wait(), timeout=10.0)
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
- assert logger.response_cost == pytest.approx(0.25)
+ if configured_fee is None:
+ assert logger.response_cost in (None, 0.0)
+ else:
+ assert logger.response_cost == pytest.approx(configured_fee)
finally:
await client.close()
From e3f672064cfa9a39afc8035d505dbda68a9ea4d4 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Wed, 30 Sep 2026 17:35:06 -0700
Subject: [PATCH 09/12] fix(agents): reserve configured fees for key and team
budgets
---
litellm/proxy/agent_endpoints/a2a_routing.py | 11 +-----
.../auth/managed_authorization.py | 11 +++---
.../auth/test_managed_authorization.py | 22 +++++++++++
.../spend_tracking/test_budget_reservation.py | 37 +++++++++++++------
.../unit/a2a_protocol/test_cost_calculator.py | 18 +++++----
5 files changed, 65 insertions(+), 34 deletions(-)
diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py
index 29d31251f94..26fdf83e0b0 100644
--- a/litellm/proxy/agent_endpoints/a2a_routing.py
+++ b/litellm/proxy/agent_endpoints/a2a_routing.py
@@ -13,7 +13,6 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
-from litellm.proxy.agent_endpoints.auth.managed_authorization import AGENT_INVOCATION_COST
async def route_a2a_agent_request(
@@ -80,13 +79,5 @@ async def route_a2a_agent_request(
data["api_base"] = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
- pricing_policy: Final = (
- user_api_key_dict.invoked_agent_policy
- if user_api_key_dict is not None and user_api_key_dict.invoked_agent_policy is not None
- else agent
- )
- configured_fee: Final = (pricing_policy.litellm_params or MappingProxyType({})).get("cost_per_query")
- invocation_fee: Final = (
- AGENT_INVOCATION_COST.validate_python(configured_fee) if configured_fee is not None else None
- )
+ invocation_fee: Final = user_api_key_dict.agent_invocation_cost if user_api_key_dict is not None else None
return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, "cost_per_query": invocation_fee}))
diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
index 11c8c790ad2..9fd1dfa7499 100644
--- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py
+++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
@@ -222,7 +222,7 @@ async def check_agent_budget(auth: UserAPIKeyAuth) -> None:
raise litellm.BudgetExceededError(current_cost=spend, max_budget=budget, message="Agent budget exceeded")
-AGENT_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)])
+_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)])
def invocation_target(route: str, body: Mapping[str, object]) -> str | None:
@@ -254,11 +254,14 @@ async def prepare_agent_invocation(
if target is None and registered_managed:
raise_identity_failure(AgentIdentityFailure(message="Invoked agent no longer exists"))
effective: Final = target if target is not None else registered
+ pricing: Final = effective.litellm_params or MappingProxyType({})
+ fixed_fee: Final = pricing.get("cost_per_query")
if (
not effective.identity_managed
and effective.litellm_budget_table is None
and auth.managed_agent_policy is None
and auth.billing_agent_policy is None
+ and fixed_fee is None
):
return
if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
@@ -271,8 +274,6 @@ async def prepare_agent_invocation(
and (effective.identity_managed or effective.litellm_budget_table is not None)
):
auth.billing_agent_policy = effective
- pricing: Final = effective.litellm_params or MappingProxyType({})
- fixed_fee: Final = pricing.get("cost_per_query")
billing_policy: Final = auth.billing_agent_policy
bounded: Final = (
billing_policy is not None
@@ -280,13 +281,13 @@ async def prepare_agent_invocation(
and billing_policy.litellm_budget_table.max_budget is not None
)
try:
- fee: Final = AGENT_INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0)
+ fee: Final = _INVOCATION_COST.validate_python(fixed_fee if billable and fixed_fee is not None else 0.0)
unbounded_token_price: Final = (
billable
and bounded
and fixed_fee is None
and any(
- AGENT_INVOCATION_COST.validate_python(pricing[field]) > 0
+ _INVOCATION_COST.validate_python(pricing[field]) > 0
for field in ("input_cost_per_token", "output_cost_per_token")
if pricing.get(field) is not None
)
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 c9b4d301018..88160562c9b 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
@@ -722,3 +722,25 @@ def test_registered_inference_routes_have_an_explicit_managed_access_decision(ro
)) or normalized in ("/models", "/cursor/models", "/cursor/v1/models")
concrete: Final = route.split("?")[0].replace("{model}", "model").replace("{model_name:path}", "model")
assert managed_agent_route_allowed(concrete, None) is not unsupported, route
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("fee", (-1.0, "invalid", float("inf"), float("nan")))
+async def test_unmanaged_invocation_rejects_invalid_configured_fees(
+ monkeypatch: pytest.MonkeyPatch, fee: float | str,
+) -> None:
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+
+ target: Final = AgentResponse(
+ agent_id="fee-target", agent_name="Fee target", agent_card_params={}, litellm_params={"cost_per_query": fee},
+ )
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(target)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
+ with pytest.raises(HTTPException) as exc:
+ await prepare_agent_invocation(auth, "fee-target", None)
+ assert exc.value.status_code == 503
+ assert "Agent invocation price is invalid" in str(exc.value.detail)
+ assert auth.agent_invocation_cost is None
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 02686a13a85..84b905af901 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py
@@ -417,8 +417,10 @@ async def test_release_unbound_budget_reservation_leaves_a_bound_one_to_its_call
@pytest.mark.asyncio
@pytest.mark.parametrize("outcome", ("failed", "cancelled", "completed"))
+@pytest.mark.parametrize("budget_owner", ("agent", "key", "team"))
+@pytest.mark.parametrize("route", ("/a2a/target", "/v1/chat/completions"))
async def test_budgeted_caller_reserves_unmanaged_agent_fees_before_concurrent_admission(
- spend_counter_cache: DualCache, monkeypatch: pytest.MonkeyPatch, outcome: str
+ spend_counter_cache: DualCache, monkeypatch: pytest.MonkeyPatch, outcome: str, budget_owner: str, route: str
) -> None:
import asyncio
@@ -442,13 +444,25 @@ async def test_budgeted_caller_reserves_unmanaged_agent_fees_before_concurrent_a
registry.register_agent(target)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ token: Final = "fee-key" if budget_owner == "key" else None
+ team: Final = LiteLLM_TeamTable(team_id="fee-team", max_budget=0.5, spend=0.0) if budget_owner == "team" else None
+ counter_key: Final = (
+ caller.budget_counter_key if budget_owner == "agent"
+ else f"spend:key:{token}" if budget_owner == "key" else "spend:team:fee-team"
+ )
+
async def admit() -> dict[str, object] | None:
- auth: Final = UserAPIKeyAuth(agent_id="caller", user_role="proxy_admin")
- auth.billing_agent_policy = caller
+ auth: Final = UserAPIKeyAuth(
+ agent_id="caller" if budget_owner == "agent" else None, user_role="proxy_admin",
+ token=token, max_budget=0.5 if budget_owner == "key" else None,
+ team_id=team.team_id if team is not None else None,
+ )
+ if budget_owner == "agent":
+ auth.billing_agent_policy = caller
await prepare_agent_invocation(auth, "target", None)
return await reserve_budget_for_request(
- request_body={"method": "message/send"}, route="/a2a/target", llm_router=None,
- valid_token=auth, team_object=None, user_object=None, prisma_client=None,
+ request_body={"method": "message/send", "model": "a2a/target"}, route=route, llm_router=None,
+ valid_token=auth, team_object=team, user_object=None, prisma_client=None,
user_api_key_cache=UserApiKeyCache(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
fail_closed_budget_enforcement=True,
)
@@ -459,22 +473,23 @@ async def test_budgeted_caller_reserves_unmanaged_agent_fees_before_concurrent_a
assert all(result is None or isinstance(result, (dict, litellm.BudgetExceededError)) for result in results), results
assert len(accepted) == 2, results
assert len(rejected) == 6
- assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5)
+ assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.5)
first: Final = accepted[0]
if outcome == "completed":
await proxy_server.increment_spend_counters(
- token=None, team_id=None, user_id=None, response_cost=0.25,
- billing_agent_id=caller.agent_id, billing_agent_counter_key=caller.budget_counter_key,
+ token=token, team_id=team.team_id if team is not None else None, user_id=None, response_cost=0.25,
+ billing_agent_id=caller.agent_id if budget_owner == "agent" else None,
+ billing_agent_counter_key=caller.budget_counter_key if budget_owner == "agent" else None,
budget_reservation=first,
)
await reconcile_budget_reservation(first, actual_cost=0.25)
- assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5)
+ assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.5)
with pytest.raises(litellm.BudgetExceededError):
await admit()
else:
release: Final = release_budget_reservation_on_cancel if outcome == "cancelled" else release_budget_reservation
await release(first)
await release(first)
- assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.25)
+ assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.25)
assert await admit() is not None
- assert await spend_counter_cache.async_get_cache(caller.budget_counter_key) == pytest.approx(0.5)
+ assert await spend_counter_cache.async_get_cache(counter_key) == pytest.approx(0.5)
diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py
index d4087bab2ea..1db6b2bfdce 100644
--- a/tests/unit/a2a_protocol/test_cost_calculator.py
+++ b/tests/unit/a2a_protocol/test_cost_calculator.py
@@ -456,11 +456,11 @@ async def test_asend_message_streaming_triggers_callbacks():
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", (False, True))
@pytest.mark.parametrize("claimed_fee", (None, -1000.0, 0.0, 99.0))
-@pytest.mark.parametrize("admitted", (False, True))
-@pytest.mark.parametrize("configured_fee", (None, 0.25))
+@pytest.mark.parametrize("changed_after_admission", (False, True))
+@pytest.mark.parametrize("configured_fee", (None, 0.0, 0.25))
async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing(
monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None,
- admitted: bool, configured_fee: float | None
+ changed_after_admission: bool, configured_fee: float | None
) -> None:
import json
from typing import Final
@@ -471,6 +471,7 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
from litellm.types.agents import AgentResponse
await _reset_callbacks_and_settle_pending_logs()
@@ -482,13 +483,14 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
litellm_params={"cost_per_query": configured_fee} if configured_fee is not None else {},
)
registry: Final = agent_registry.AgentRegistry()
- registry.register_agent(target.model_copy(update={"litellm_params": {"cost_per_query": 0.5}}) if admitted else target)
+ registry.register_agent(target)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
- if admitted:
- auth.invoked_agent_policy = target
- auth.invoked_agent_id = target.agent_id
- auth.agent_invocation_cost = configured_fee or 0.0
+ await prepare_agent_invocation(auth, "fee-target", None)
+ assert auth.agent_invocation_cost == configured_fee
+ if changed_after_admission:
+ registry.deregister_agent(target.agent_name)
+ registry.register_agent(target.model_copy(update={"litellm_params": {"cost_per_query": 0.5}}))
def reply(request: httpx.Request) -> httpx.Response:
body: Final = json.loads(request.content)
From bc5ec233549f37e4ae6f55a862baa28d8e99b615 Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Wed, 30 Sep 2026 18:19:37 -0700
Subject: [PATCH 10/12] fix(agents): bind dispatch and pricing to resolved
admission
---
.../proxy/agent_endpoints/a2a_endpoints.py | 13 ++-
litellm/proxy/agent_endpoints/a2a_routing.py | 11 +-
.../auth/managed_authorization.py | 64 ++++++++++-
litellm/proxy/auth/user_api_key_auth.py | 44 ++++++--
litellm/proxy/route_llm_request.py | 14 +++
.../auth/test_managed_authorization.py | 38 ++++++-
.../agent_endpoints/test_a2a_endpoints.py | 78 ++++++++++++--
.../proxy/auth/test_user_api_key_auth.py | 100 +++++++++++++++++-
.../proxy/test_route_a2a_models.py | 100 ++++++++++++++++++
.../unit/a2a_protocol/test_cost_calculator.py | 50 ++++++++-
10 files changed, 480 insertions(+), 32 deletions(-)
diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py
index 6ddcd20d919..d66d9835081 100644
--- a/litellm/proxy/agent_endpoints/a2a_endpoints.py
+++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py
@@ -33,6 +33,7 @@ from litellm.proxy.a2a.version_convert import (
normalize_request_params,
normalize_stream_event,
)
+from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
from litellm.proxy.agent_endpoints.databricks_oauth import (
DATABRICKS_OAUTH_PARAM,
resolve_databricks_app_auth_header,
@@ -706,9 +707,10 @@ async def invoke_agent_a2a(
params.pop(key)
# Find the agent
- agent: Final = await _get_agent(agent_id)
- if agent is None:
+ registered: Final = await _get_agent(agent_id)
+ if registered is None:
return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404)
+ agent: Final = agent_invocation_policy(user_api_key_dict, registered)
served_version: Final = _served_version(agent, request, original_method)
@@ -732,7 +734,12 @@ async def invoke_agent_a2a(
agent_name: Final = agent_card_params.get("name", agent_id)
# Get litellm_params (may include custom_llm_provider for completion bridge)
- litellm_params: dict[str, object] = agent.litellm_params or {}
+ litellm_params: dict[
+ str, object
+ ] = { # mutable-ok: A2A SDK and completion bridge accept provider parameters as a dict
+ **(agent.litellm_params or MappingProxyType({})),
+ "cost_per_query": user_api_key_dict.agent_invocation_cost,
+ }
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
# Hand the authenticated key hash to the completion bridge so provider
diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py
index 26fdf83e0b0..9d91fbfff7d 100644
--- a/litellm/proxy/agent_endpoints/a2a_routing.py
+++ b/litellm/proxy/agent_endpoints/a2a_routing.py
@@ -13,12 +13,16 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
+from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
+from litellm.types.agents import AgentResponse
async def route_a2a_agent_request(
data: dict,
route_type: str,
user_api_key_dict: UserAPIKeyAuth | None = None,
+ *,
+ registered_agent: AgentResponse | None = None,
) -> Any | None:
"""
Route A2A agent requests directly to litellm with injected API base.
@@ -47,12 +51,14 @@ async def route_a2a_agent_request(
agent_name: Final = model_name[4:]
# Look up agent in registry
- agent: Final = await get_agent_with_read_through(agent_name)
- if agent is None:
+ registered: Final = registered_agent or await get_agent_with_read_through(agent_name)
+ if registered is None:
verbose_proxy_logger.error("[A2A] Agent '%s' not found in registry", agent_name)
route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
+ agent: Final = agent_invocation_policy(user_api_key_dict, registered)
+
# Verify the caller is permitted to use this agent (admins bypass the check)
is_admin: Final = user_api_key_dict is not None and (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
@@ -76,6 +82,7 @@ async def route_a2a_agent_request(
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
# Inject API base and route to litellm
+ data.pop("litellm_params", None)
data["api_base"] = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
index 9fd1dfa7499..d945c6e1ad2 100644
--- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py
+++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
@@ -99,12 +99,18 @@ def managed_inference_request(
cli_model: str | None,
path_model: object = None,
query_model: object = None,
+ *,
+ auth: UserAPIKeyAuth | None = None,
+ require_model: bool = True,
+ model_group_alias: object = None,
) -> dict[str, object]:
from litellm.proxy.auth.route_checks import RouteChecks
if route in _MANAGED_REALTIME_ROUTES:
model: Final = query_model or body.get("model")
if not isinstance(model, str) or not model:
+ if not require_model:
+ return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
@@ -117,8 +123,32 @@ def managed_inference_request(
endpoint_model: Final = path_model or (
query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None
)
- effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind)
+ from litellm.proxy.common_utils.model_listing_utils import CallerAliases, alias_target
+ from litellm.proxy.litellm_pre_call_utils import (
+ _update_model_if_key_alias_exists,
+ _update_model_if_team_alias_exists,
+ )
+
+ aliased_body: Final = dict(body) # mutable-ok: existing alias helpers rewrite their request copy
+ if auth is not None:
+ _update_model_if_team_alias_exists(aliased_body, auth)
+ _update_model_if_key_alias_exists(aliased_body, auth)
+ selected: Final = resolve_inference_model(aliased_body.get("model"), settings, cli_model, endpoint_model, kind=kind)
+ import litellm
+
+ aliased: Final = (
+ alias_target(selected, CallerAliases((), (litellm.model_alias_map, auth.aliases))) or selected
+ if isinstance(selected, str) and auth is not None
+ else selected
+ )
+ from litellm.router_utils.common_utils import resolve_model_group_alias
+
+ effective: Final = (
+ (resolve_model_group_alias(model_group_alias, aliased) or aliased) if isinstance(aliased, str) else aliased
+ )
if not isinstance(effective, str) or not effective:
+ if not require_model:
+ return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
@@ -256,6 +286,10 @@ async def prepare_agent_invocation(
effective: Final = target if target is not None else registered
pricing: Final = effective.litellm_params or MappingProxyType({})
fixed_fee: Final = pricing.get("cost_per_query")
+ if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
+ raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent"))
+ auth.invoked_agent_id = effective.agent_id
+ auth.invoked_agent_policy = effective
if (
not effective.identity_managed
and effective.litellm_budget_table is None
@@ -264,10 +298,6 @@ async def prepare_agent_invocation(
and fixed_fee 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
- auth.invoked_agent_policy = effective
if (
billable
and auth.agent_id is None
@@ -304,3 +334,27 @@ async def prepare_agent_invocation(
)
)
auth.agent_invocation_cost = fee
+
+
+def agent_invocation_policy(auth: UserAPIKeyAuth | None, registered: AgentResponse) -> AgentResponse:
+ admitted: Final = auth.invoked_agent_policy if auth is not None else None
+ if admitted is not None and auth is not None:
+ matching: Final = auth.invoked_agent_id == registered.agent_id == admitted.agent_id
+ captured_price: Final = (admitted.litellm_params or MappingProxyType({})).get(
+ "cost_per_query"
+ ) is None or auth.agent_invocation_cost is not None
+ if matching and captured_price:
+ return admitted
+ raise_identity_failure(
+ AgentIdentityFailure(code="policy_unavailable", message="Agent dispatch does not match its admission")
+ )
+ if (
+ registered.identity_managed
+ or registered.identity is not None
+ or registered.litellm_budget_table is not None
+ or (registered.litellm_params or MappingProxyType({})).get("cost_per_query") is not None
+ ):
+ raise_identity_failure(
+ AgentIdentityFailure(code="policy_unavailable", message="Agent dispatch requires a matching admission")
+ )
+ return registered
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index a684b5ebcf3..3d13309a682 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -3233,7 +3233,14 @@ async def _authorize_authenticated_request(
prepare_agent_invocation,
)
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
- from litellm.proxy.proxy_server import general_settings, prisma_client, user_model
+ from litellm.proxy.proxy_server import (
+ general_settings,
+ llm_router,
+ prisma_client,
+ proxy_config,
+ proxy_logging_obj,
+ user_model,
+ )
store: Final = AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None
if user_api_key_auth_obj.agent_id is not None:
@@ -3242,19 +3249,34 @@ async def _authorize_authenticated_request(
route, request.method
):
raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
- authorized_data: Final = (
- managed_inference_request(
- route,
- request_data,
- general_settings,
- user_model,
- request.path_params.get("model") or request.path_params.get("model_name"),
- request.query_params.get("model"),
+ router_settings: Final = (
+ await proxy_config.get_hierarchical_router_settings(
+ user_api_key_dict=user_api_key_auth_obj,
+ prisma_client=prisma_client,
+ proxy_logging_obj=proxy_logging_obj,
)
- if user_api_key_auth_obj.managed_agent_policy is not None
+ if llm_router is not None and RouteChecks.is_llm_api_route(route=route)
+ else None
+ )
+ inference_data: Final = managed_inference_request(
+ route,
+ request_data,
+ general_settings,
+ user_model,
+ request.path_params.get("model") or request.path_params.get("model_name"),
+ request.query_params.get("model"),
+ model_group_alias=router_settings.get("model_group_alias")
+ if isinstance(router_settings, Mapping)
+ else None,
+ auth=user_api_key_auth_obj,
+ require_model=user_api_key_auth_obj.managed_agent_policy is not None,
+ )
+ target_name: Final = invocation_target(route, inference_data)
+ authorized_data: Final = (
+ inference_data
+ if target_name is not None or user_api_key_auth_obj.managed_agent_policy is not None
else request_data
)
- target_name: Final = invocation_target(route, authorized_data)
if target_name is not None:
await prepare_agent_invocation(
user_api_key_auth_obj,
diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py
index 42ac74cae33..a2dfc95750c 100644
--- a/litellm/proxy/route_llm_request.py
+++ b/litellm/proxy/route_llm_request.py
@@ -477,6 +477,20 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr
data.pop("enable_tag_filtering", None)
+ if _is_a2a_agent_model(data.get("model")):
+ from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
+ from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
+
+ registered_agent: Final = await get_agent_with_read_through(data["model"][4:])
+ if registered_agent is not None:
+ agent_response: Final = await route_a2a_agent_request(
+ data, route_type, user_api_key_dict=user_api_key_dict, registered_agent=registered_agent
+ )
+ if agent_response is not None:
+ return agent_response
+ if user_api_key_dict is not None and user_api_key_dict.invoked_agent_policy is not None:
+ raise HTTPException(503, "Agent dispatch does not match its admission")
+
team_id: Final = get_team_id_from_data(data)
router_model_names: Final = llm_router.model_names if llm_router is not None else []
is_proxy_admin_without_team: Final = team_id is None and _is_proxy_admin_request(data)
diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
index 88160562c9b..bc84e9ddd03 100644
--- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
+++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
@@ -534,7 +534,8 @@ async def test_unmanaged_agent_invocation_retains_legacy_behavior(monkeypatch: p
await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
assert auth.managed_agent_policy is None
assert auth.billing_agent_policy is None
- assert auth.invoked_agent_id is None
+ assert auth.invoked_agent_id == "agent"
+ assert auth.invoked_agent_policy is not None
@pytest.mark.asyncio
@@ -744,3 +745,38 @@ async def test_unmanaged_invocation_rejects_invalid_configured_fees(
assert exc.value.status_code == 503
assert "Agent invocation price is invalid" in str(exc.value.detail)
assert auth.agent_invocation_cost is None
+
+
+@pytest.mark.parametrize("route", ("/realtime", "/v1/chat/completions", "/v1/files"))
+def test_ordinary_requests_without_models_keep_existing_validation(route: str) -> None:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
+
+ assert managed_inference_request(route, {}, {}, None, require_model=False) == {}
+
+
+@pytest.mark.parametrize("state", (
+ {"identity_managed": True},
+ {"identity": BINDING},
+ {"litellm_budget_table": {"budget_id": "budget", "max_budget": 0.5}},
+ {"litellm_params": {"cost_per_query": 0.25}},
+))
+@pytest.mark.parametrize("has_auth", (False, True))
+def test_protected_agent_dispatch_requires_admission(state: dict[str, object], has_auth: bool) -> None:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
+
+ registered: Final = agent(**{"identity": None, "identity_managed": False, **state})
+ with pytest.raises(HTTPException, match="admission") as exc:
+ agent_invocation_policy(UserAPIKeyAuth() if has_auth else None, registered)
+ assert exc.value.status_code == 503
+
+
+def test_paid_agent_dispatch_requires_the_captured_fee() -> None:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
+
+ policy: Final = agent(identity=None, identity_managed=False, litellm_params={"cost_per_query": 0.25})
+ auth: Final = UserAPIKeyAuth()
+ auth.invoked_agent_id = policy.agent_id
+ auth.invoked_agent_policy = policy
+ with pytest.raises(HTTPException, match="admission") as exc:
+ agent_invocation_policy(auth, policy)
+ assert exc.value.status_code == 503
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
index a5d0d0a3ecc..8f905b2fc96 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
@@ -16,7 +16,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._types import UserAPIKeyAuth
-from litellm.types.agents import AgentCaller
+from litellm.types.agents import AgentCaller, AgentResponse
AddLiteLLMData = Callable[..., Awaitable[dict[str, object]]]
@@ -25,6 +25,8 @@ AddLiteLLMData = Callable[..., Awaitable[dict[str, object]]]
class CapturedAgentCall:
request_id: object
agent_extra_headers: dict[str, str] | None
+ cost_per_query: object
+ api_base: object
@pytest.mark.asyncio
@@ -58,7 +60,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
}
# Mock agent
- mock_agent = MagicMock()
+ mock_agent = _make_agent_mock()
mock_agent.agent_id = "test-agent"
mock_agent.agent_card_params = {
"url": "http://backend-agent:10001",
@@ -211,7 +213,7 @@ async def test_invoke_agent_a2a_handles_none_agent_card_params():
"""
from litellm.proxy._types import UserAPIKeyAuth
- mock_agent = MagicMock()
+ mock_agent = _make_agent_mock()
mock_agent.agent_card_params = None
mock_agent.litellm_params = None
@@ -295,7 +297,7 @@ async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge():
resp.model_dump.return_value = {"jsonrpc": "2.0", "id": "test-id", "result": {}}
return resp
- mock_agent = MagicMock()
+ mock_agent = _make_agent_mock()
mock_agent.agent_id = "lf-agent"
mock_agent.agent_name = "lf-agent"
# No URL: the bridge derives the endpoint from the LangFlow agent config.
@@ -376,6 +378,9 @@ def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock:
agent.litellm_params = {}
agent.static_headers = None
agent.extra_headers = None
+ agent.identity_managed = False
+ agent.identity = None
+ agent.litellm_budget_table = None
return agent
@@ -394,7 +399,7 @@ def _make_request_mock(method: str, params: Mapping[str, object], request_id: ob
def _base_patches(
- agent: MagicMock, add_litellm_data: AddLiteLLMData | None = None
+ agent: MagicMock | AgentResponse, add_litellm_data: AddLiteLLMData | None = None
) -> list[AbstractContextManager[object]]:
return [
patch(
@@ -437,7 +442,7 @@ async def _invoke_message_method(
mock_request: MagicMock,
user_api_key_dict: UserAPIKeyAuth,
add_litellm_data: AddLiteLLMData | None = None,
- agent: MagicMock | None = None,
+ agent: MagicMock | AgentResponse | None = None,
) -> CapturedAgentCall:
from fastapi.responses import JSONResponse
@@ -490,7 +495,12 @@ async def _invoke_message_method(
kwargs: Final = downstream.call_args.kwargs
request_id: Final = kwargs["request"].__dict__["id"] if is_send else kwargs["request_id"]
- return CapturedAgentCall(request_id=request_id, agent_extra_headers=kwargs.get("agent_extra_headers"))
+ return CapturedAgentCall(
+ request_id=request_id,
+ agent_extra_headers=kwargs.get("agent_extra_headers"),
+ cost_per_query=kwargs["litellm_params"].get("cost_per_query"),
+ api_base=kwargs["api_base"],
+ )
@pytest.mark.asyncio
@@ -2689,3 +2699,57 @@ def test_forwarding_headers_minted_bearer_replaces_a_forwarded_authorization_of_
)
assert merged == {"X-Custom": "kept", "Authorization": "Bearer minted-token"}
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("method", ("message/send", "message/stream"))
+@pytest.mark.parametrize("fee", (None, 0.0, 0.25))
+async def test_native_dispatch_keeps_the_admitted_price_and_destination(
+ monkeypatch: pytest.MonkeyPatch,
+ method: str,
+ fee: float | None,
+) -> None:
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+
+ admitted: Final = AgentResponse(
+ agent_id="test-agent",
+ agent_name="test-agent",
+ agent_card_params={"url": "https://admitted.test/"},
+ litellm_params={"cost_per_query": fee} if fee is not None else {},
+ )
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(admitted)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
+ await prepare_agent_invocation(auth, admitted.agent_id, None)
+ changed: Final = admitted.model_copy(
+ update={
+ "litellm_params": {"cost_per_query": 0.75},
+ "agent_card_params": {"url": "https://changed.test/"},
+ }
+ )
+ captured: Final = await _invoke_message_method(
+ method,
+ _make_request_mock(method, _HELLO_MESSAGE_PARAMS),
+ auth,
+ agent=changed,
+ )
+ assert captured.cost_per_query == fee
+ assert captured.api_base == "https://admitted.test/"
+
+
+@pytest.mark.asyncio
+async def test_native_dispatch_returns_not_found_for_a_removed_agent(monkeypatch: pytest.MonkeyPatch) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
+
+ monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry())
+ monkeypatch.setattr(proxy_server, "prisma_client", None)
+ response: Final = await invoke_agent_a2a(
+ agent_id="removed", request=_make_request_mock("message/send", _HELLO_MESSAGE_PARAMS),
+ fastapi_response=MagicMock(), user_api_key_dict=UserAPIKeyAuth(),
+ )
+ assert response.status_code == 404
+ assert json.loads(response.body)["error"]["message"] == "Agent 'removed' not found"
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index d69b3e72183..b514e78984f 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -6267,6 +6267,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it():
"type": "http",
"headers": [(b"content-type", b"application/json")],
"method": "POST",
+ "query_string": b"",
}
)
request._url = URL(url="/chat/completions")
@@ -6321,6 +6322,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder()
"type": "http",
"headers": [(b"content-type", b"application/json")],
"method": "POST",
+ "query_string": b"",
}
)
request._url = URL(url="/chat/completions")
@@ -8403,7 +8405,7 @@ _DDTRACE_AUTH_PROBE = dedent(
async def auth(api_key):
- request = Request(scope={"type": "http", "headers": [], "method": "POST", "path": "/chat/completions"})
+ request = Request(scope={"type": "http", "headers": [], "method": "POST", "path": "/chat/completions", "query_string": b""})
request._url = URL(url="/chat/completions")
try:
await user_api_key_auth(
@@ -9811,3 +9813,99 @@ def test_free_model_only_waives_budgets_without_a_paid_agent_invocation(
)
is skipped
)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "selection",
+ (
+ "body",
+ "query",
+ "path",
+ "cli",
+ "default",
+ "key-alias",
+ "team-alias",
+ "global-alias",
+ "alias-chain",
+ "query-alias",
+ "query-over-alias",
+ "different-agent",
+ "router-alias",
+ "query-router-alias",
+ ),
+)
+async def test_agent_admission_prices_the_model_selected_for_dispatch(
+ monkeypatch: pytest.MonkeyPatch,
+ selection: str,
+) -> None:
+ from typing import Final
+
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
+ from litellm.types.agents import AgentResponse
+
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(
+ AgentResponse(
+ agent_id="paid",
+ agent_name="Paid",
+ agent_card_params={},
+ litellm_params={"cost_per_query": 0.25},
+ )
+ )
+ registry.register_agent(
+ AgentResponse(
+ agent_id="other",
+ agent_name="Other",
+ agent_card_params={},
+ litellm_params={"cost_per_query": 0.75},
+ )
+ )
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ monkeypatch.setattr(litellm, "model_alias_map", {"global": "a2a/paid"})
+ for name, value in {
+ **_proxy_attrs_for_centralized_checks(),
+ "user_model": "a2a/paid" if selection == "cli" else None,
+ "llm_router": litellm.Router(model_list=[]) if "router-alias" in selection else None,
+ "general_settings": {"completion_model": "a2a/paid"} if selection == "default" else {},
+ "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
+ }.items():
+ monkeypatch.setattr(proxy_server, name, value)
+ body: Final = {
+ "model": {
+ "body": "a2a/paid",
+ "key-alias": "alias",
+ "team-alias": "team-alias",
+ "global-alias": "global",
+ "alias-chain": "team-alias",
+ "query-over-alias": "alias",
+ "different-agent": "a2a/other",
+ "router-alias": "router-alias",
+ }.get(selection, "gpt-4o"),
+ "messages": [{"role": "user", "content": "Hello"}],
+ }
+ route: Final = "/openai/deployments/a2a/paid/chat/completions" if selection == "path" else "/v1/chat/completions"
+ request: Final = _alias_request(route, body, path_params={"model": "a2a/paid"} if selection == "path" else {})
+ if selection in ("query", "query-over-alias", "different-agent", "query-alias"):
+ request.scope["query_string"] = b"model=alias" if selection == "query-alias" else b"model=a2a%2Fpaid"
+ if selection == "query-router-alias":
+ request.scope["query_string"] = b"model=router-alias"
+ auth: Final = UserAPIKeyAuth(
+ router_settings={"model_group_alias": {"router-alias": "a2a/paid"}},
+ user_role="proxy_admin",
+ aliases={"alias": "global" if selection == "alias-chain" else "a2a/paid"},
+ team_model_aliases={"team-alias": "alias" if selection == "alias-chain" else "a2a/paid"},
+ )
+ with patch(
+ "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
+ new_callable=AsyncMock,
+ return_value=None,
+ ) as reserve:
+ assert await _authorize_authenticated_request(auth, request, body, route, "test-key") is None
+ assert auth.invoked_agent_id == "paid"
+ assert auth.agent_invocation_cost == pytest.approx(0.25)
+ reserve.assert_awaited_once()
+ assert reserve.call_args.kwargs["valid_token"].agent_invocation_cost == pytest.approx(0.25)
+ assert reserve.call_args.kwargs["request_body"]["model"] == "a2a/paid"
diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py
index 0429dd97a1c..8783ab1b61e 100644
--- a/tests/test_litellm/proxy/test_route_a2a_models.py
+++ b/tests/test_litellm/proxy/test_route_a2a_models.py
@@ -183,3 +183,103 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re
assert call_kwargs["model"] == f"a2a/{agent_name}"
assert call_kwargs["api_base"] == "http://sibling-db-agent.example.com"
prisma_client.db.litellm_agentstable.find_unique.assert_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "override", ("api_key", "api_base", "user_config", "router_settings_override", "deployment", "no-router")
+)
+async def test_registered_agent_dispatch_owns_the_admitted_destination_and_fee(monkeypatch: pytest.MonkeyPatch, override: str) -> None:
+ from typing import Final
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+ from litellm.types.agents import AgentResponse
+
+ agent: Final = AgentResponse(
+ agent_id="paid",
+ agent_name="paid",
+ agent_card_params={"url": "https://registered.test/"},
+ litellm_params={"cost_per_query": 0.25},
+ )
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(agent)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ router: Final = _router_without_models()
+ if override == "deployment":
+ router.is_recognized_model.return_value = True
+ provider: Final = AsyncMock(return_value={"id": "reply"})
+ auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
+ with patch("litellm.acompletion", provider):
+ await prepare_agent_invocation(auth, "paid", None)
+ pending: Final = await route_request(
+ data={
+ "model": "a2a/paid",
+ "messages": [{"role": "user", "content": "Hi"}],
+ "cost_per_query": 99.0,
+ **(
+ {override: {} if override in ("user_config", "router_settings_override") else "override"}
+ if override not in ("deployment", "no-router")
+ else {}
+ ),
+ },
+ llm_router=None if override == "no-router" else router,
+ user_model=None,
+ route_type="acompletion",
+ user_api_key_dict=auth,
+ )
+ assert await pending == {"id": "reply"}
+ provider.assert_awaited_once()
+ assert provider.call_args.kwargs["api_base"] == "https://registered.test/"
+ assert provider.call_args.kwargs["cost_per_query"] == 0.25
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("model", ("a2a/paid", "gpt-4o"))
+async def test_routing_overrides_cannot_dispatch_without_matching_agent_admission(
+ monkeypatch: pytest.MonkeyPatch, model: str,
+) -> None:
+ from typing import Final
+ from fastapi import HTTPException
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+ from litellm.types.agents import AgentResponse
+
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(AgentResponse(
+ agent_id="paid", agent_name="paid", agent_card_params={"url": "https://agent.test/"},
+ litellm_params={"cost_per_query": 0.25},
+ ))
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ monkeypatch.setattr(proxy_server, "prisma_client", None)
+ auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
+ if model == "gpt-4o":
+ await prepare_agent_invocation(auth, "paid", None)
+ provider: Final = Mock(return_value=None)
+ with patch("litellm.acompletion", provider), pytest.raises(HTTPException, match="admission") as exc:
+ await route_request(
+ data={"model": model, "api_base": "https://override.test/", "messages": [{"role": "user", "content": "Hi"}]},
+ llm_router=None, user_model=None, route_type="acompletion", user_api_key_dict=auth,
+ )
+ assert exc.value.status_code == 503
+ provider.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_unregistered_direct_agent_keeps_explicit_endpoint_routing(monkeypatch: pytest.MonkeyPatch) -> None:
+ from typing import Final
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+
+ monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry())
+ monkeypatch.setattr(proxy_server, "prisma_client", None)
+ provider: Final = AsyncMock(return_value={"id": "direct-reply"})
+ with patch("litellm.acompletion", provider):
+ pending: Final = await route_request(
+ data={"model": "a2a/direct", "api_base": "https://direct.test/", "messages": [{"role": "user", "content": "Hi"}]},
+ llm_router=None, user_model=None, route_type="acompletion",
+ )
+ assert await pending == {"id": "direct-reply"}
+ assert provider.call_args.kwargs["api_base"] == "https://direct.test/"
diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py
index 1db6b2bfdce..62dc077054e 100644
--- a/tests/unit/a2a_protocol/test_cost_calculator.py
+++ b/tests/unit/a2a_protocol/test_cost_calculator.py
@@ -456,11 +456,12 @@ async def test_asend_message_streaming_triggers_callbacks():
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", (False, True))
@pytest.mark.parametrize("claimed_fee", (None, -1000.0, 0.0, 99.0))
+@pytest.mark.parametrize("fee_field", ("cost_per_query", "litellm_params"))
@pytest.mark.parametrize("changed_after_admission", (False, True))
@pytest.mark.parametrize("configured_fee", (None, 0.0, 0.25))
async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing(
monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None,
- changed_after_admission: bool, configured_fee: float | None
+ changed_after_admission: bool, configured_fee: float | None, fee_field: str
) -> None:
import json
from typing import Final
@@ -509,7 +510,8 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
pending: Final = await route_a2a_agent_request(
data={"model": "a2a/fee-target", "messages": [{"role": "user", "content": "Hello"}],
"stream": stream, "client": client,
- **({"cost_per_query": claimed_fee} if claimed_fee is not None else {})},
+ **({fee_field: claimed_fee if fee_field == "cost_per_query" else {"cost_per_query": claimed_fee}}
+ if claimed_fee is not None else {})},
route_type="acompletion", user_api_key_dict=auth,
)
response: Final = await pending
@@ -526,3 +528,47 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
assert logger.response_cost == pytest.approx(configured_fee)
finally:
await client.close()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("admitted_target", (None, "other"))
+async def test_chat_agent_dispatch_rejects_missing_or_different_admission(
+ monkeypatch: pytest.MonkeyPatch,
+ admitted_target: str | None,
+) -> None:
+ from typing import Final
+ from unittest.mock import Mock
+
+ from fastapi import HTTPException
+
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+ from litellm.types.agents import AgentResponse
+
+ registry: Final = agent_registry.AgentRegistry()
+ for name in ("paid", "other"):
+ registry.register_agent(
+ AgentResponse(
+ agent_id=name,
+ agent_name=name,
+ agent_card_params={"url": "https://agent.test/"},
+ litellm_params={"cost_per_query": 0.25},
+ )
+ )
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
+ if admitted_target is not None:
+ await prepare_agent_invocation(auth, admitted_target, None)
+ provider: Final = Mock(return_value=None)
+ monkeypatch.setattr(litellm, "acompletion", provider)
+ with pytest.raises(HTTPException) as exc:
+ await route_a2a_agent_request(
+ data={"model": "a2a/paid", "messages": [{"role": "user", "content": "Hello"}]},
+ route_type="acompletion",
+ user_api_key_dict=auth,
+ )
+ assert exc.value.status_code == 503
+ assert "admission" in str(exc.value.detail).lower()
+ provider.assert_not_called()
From cbec06c18342d2ea05e8be3d0f8eb565c9ee3c1f Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Wed, 30 Sep 2026 18:41:55 -0700
Subject: [PATCH 11/12] fix(agents): preserve token pricing and legacy request
handling
---
.../proxy/agent_endpoints/a2a_endpoints.py | 4 +-
litellm/proxy/agent_endpoints/a2a_routing.py | 10 +++--
litellm/proxy/auth/user_api_key_auth.py | 2 +-
.../agent_endpoints/test_a2a_endpoints.py | 40 +++++++++++++++++++
.../test_llm_pass_through_endpoints.py | 2 +-
5 files changed, 51 insertions(+), 7 deletions(-)
diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py
index d66d9835081..23d11a5a9c5 100644
--- a/litellm/proxy/agent_endpoints/a2a_endpoints.py
+++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py
@@ -738,7 +738,9 @@ async def invoke_agent_a2a(
str, object
] = { # mutable-ok: A2A SDK and completion bridge accept provider parameters as a dict
**(agent.litellm_params or MappingProxyType({})),
- "cost_per_query": user_api_key_dict.agent_invocation_cost,
+ "cost_per_query": user_api_key_dict.agent_invocation_cost
+ if (agent.litellm_params or MappingProxyType({})).get("cost_per_query") is not None
+ else None,
}
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py
index 9d91fbfff7d..96c0b0ceb5a 100644
--- a/litellm/proxy/agent_endpoints/a2a_routing.py
+++ b/litellm/proxy/agent_endpoints/a2a_routing.py
@@ -82,9 +82,11 @@ async def route_a2a_agent_request(
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
# Inject API base and route to litellm
- data.pop("litellm_params", None)
- data["api_base"] = agent.agent_card_params["url"]
- verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])
+ api_base: Final = agent.agent_card_params["url"]
+ verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, api_base)
invocation_fee: Final = user_api_key_dict.agent_invocation_cost if user_api_key_dict is not None else None
- return getattr(litellm, f"{route_type}")(**MappingProxyType({**data, "cost_per_query": invocation_fee}))
+ provider_data: Final = MappingProxyType({key: value for key, value in data.items() if key != "litellm_params"})
+ return getattr(litellm, f"{route_type}")(
+ **MappingProxyType({**provider_data, "api_base": api_base, "cost_per_query": invocation_fee})
+ )
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index 3d13309a682..f3cc82eb895 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -3264,7 +3264,7 @@ async def _authorize_authenticated_request(
general_settings,
user_model,
request.path_params.get("model") or request.path_params.get("model_name"),
- request.query_params.get("model"),
+ _safe_get_request_query_params(request).get("model"),
model_group_alias=router_settings.get("model_group_alias")
if isinstance(router_settings, Mapping)
else None,
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
index 8f905b2fc96..1e04a6df3a5 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
@@ -27,6 +27,7 @@ class CapturedAgentCall:
agent_extra_headers: dict[str, str] | None
cost_per_query: object
api_base: object
+ pricing: Mapping[str, object]
@pytest.mark.asyncio
@@ -500,6 +501,7 @@ async def _invoke_message_method(
agent_extra_headers=kwargs.get("agent_extra_headers"),
cost_per_query=kwargs["litellm_params"].get("cost_per_query"),
api_base=kwargs["api_base"],
+ pricing=kwargs["litellm_params"],
)
@@ -2753,3 +2755,41 @@ async def test_native_dispatch_returns_not_found_for_a_removed_agent(monkeypatch
)
assert response.status_code == 404
assert json.loads(response.body)["error"]["message"] == "Agent 'removed' not found"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("method", ("message/send", "message/stream"))
+@pytest.mark.parametrize("fixed_fee", (None, 0.0, 0.25))
+async def test_unbudgeted_managed_agent_keeps_token_pricing_without_a_fixed_fee(
+ monkeypatch: pytest.MonkeyPatch, method: str, fixed_fee: float | None,
+) -> None:
+ from litellm import Usage
+ from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+ from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
+
+ policy: Final = AgentResponse(
+ agent_id="test-agent", agent_name="test-agent", identity_managed=True,
+ agent_card_params={"url": "https://agent.test/"},
+ litellm_params={"input_cost_per_token": 0.02, "output_cost_per_token": 0.03,
+ **({"cost_per_query": fixed_fee} if fixed_fee is not None else {})},
+ )
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(policy)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
+ auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
+ with ExitStack() as stack:
+ for context in _base_patches(policy):
+ stack.enter_context(context)
+ await prepare_agent_invocation(auth, policy.agent_id, AgentIdentityStore.from_client(database))
+ captured: Final = await _invoke_message_method(
+ method, _make_request_mock(method, _HELLO_MESSAGE_PARAMS), auth, agent=policy,
+ )
+ logging: Final = MagicMock(model_call_details={
+ "litellm_params": captured.pricing,
+ "usage": Usage(prompt_tokens=10, completion_tokens=2, total_tokens=12),
+ })
+ assert A2ACostCalculator.calculate_a2a_cost(logging) == pytest.approx(0.26 if fixed_fee is None else fixed_fee)
diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
index 4577d578263..72364904f93 100644
--- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
+++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
@@ -2270,7 +2270,7 @@ class TestGigachatProxyRoute:
mock_request.headers = {"content-type": "application/json"}
mock_request.query_params = {}
mock_fastapi_response = MagicMock(spec=Response)
- mock_user_api_key_dict = MagicMock()
+ mock_user_api_key_dict = UserAPIKeyAuth()
mock_llm_router.allm_passthrough_route = AsyncMock(
return_value=httpx.Response(200, json={"response": "success"})
)
From 867f097a00b6a53e603c1040682c7fa3e7ae102f Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Wed, 30 Sep 2026 18:57:18 -0700
Subject: [PATCH 12/12] test(agents): use real policy data in header forwarding
fixtures
---
.../agent_endpoints/test_agent_headers.py | 27 ++++++++++---------
1 file changed, 15 insertions(+), 12 deletions(-)
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py
index e894f4ad69a..4d6077b3472 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py
@@ -14,23 +14,26 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
+from litellm.types.agents import AgentResponse
+
# ---------------------------------------------------------------------------
-# Helper: build a minimal mock agent
+# Helper: build a minimal agent
# ---------------------------------------------------------------------------
def _make_mock_agent(
- static_headers=None,
- extra_headers=None,
- url="http://backend-agent:10001",
-):
- mock_agent = MagicMock()
- mock_agent.agent_id = "agent-123"
- mock_agent.agent_card_params = {"url": url, "name": "Test Agent"}
- mock_agent.litellm_params = {}
- mock_agent.static_headers = static_headers or {}
- mock_agent.extra_headers = extra_headers or []
- return mock_agent
+ static_headers: dict[str, str] | None = None,
+ extra_headers: list[str] | None = None,
+ url: str = "http://backend-agent:10001",
+) -> AgentResponse:
+ return AgentResponse(
+ agent_id="agent-123",
+ agent_name="Test Agent",
+ agent_card_params={"url": url, "name": "Test Agent"},
+ litellm_params={},
+ static_headers=static_headers or {},
+ extra_headers=extra_headers or [],
+ )
def _make_mock_request(extra_headers=None, method="message/send"):