feat(agents/): support budgets + rate limiting on agents + agent sessions

This commit is contained in:
Krrish Dholakia 2026-03-04 19:51:19 -08:00
parent 73224aa7e7
commit bea664389b
23 changed files with 1229 additions and 202 deletions

View file

@ -20,6 +20,7 @@ Add A2A Agents on LiteLLM AI Gateway, Invoke agents in A2A Protocol, track reque
| Logging | ✅ |
| Load Balancing | ✅ |
| Streaming | ✅ |
| [Iteration Budgets](a2a_iteration_budgets) | ✅ |
:::tip

View file

@ -0,0 +1,188 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Agent Iteration Budgets
Control runaway costs from agentic loops with per-session iteration and budget caps.
## Overview
When agents run agentic loops, they can make unbounded LLM calls, causing unexpected costs. LiteLLM provides two controls:
| Control | Description |
|---------|-------------|
| **Max Iterations** | Hard cap on the number of LLM calls per session |
| **Max Budget Per Session** | Dollar cap per session (identified by `x-litellm-trace-id`) |
Both controls require a `session_id` (sent via `x-litellm-trace-id` header or `metadata.session_id`) to track calls within a session.
## Trace-ID Enforcement
LiteLLM supports two independent trace-id flags, configured in `litellm_params` on the agent:
| Flag | Description |
|------|-------------|
| `require_trace_id_on_calls_to_agent` | Requires callers invoking this agent to include `x-litellm-trace-id`. Use when the agent should only be called as a sub-agent with a trace context. Returns **400** if missing. |
| `require_trace_id_on_calls_by_agent` | Requires all LLM/MCP calls made **by** this agent (via its virtual key) to include `x-litellm-trace-id`. This is what enables `max_iterations` and `max_budget_per_session` tracking. Returns **400** if missing. |
## Configuring via UI
When creating an agent in the LiteLLM Admin UI:
1. Navigate to the **Agents** tab and click **Add Agent**
2. In the **Agent Settings** step, expand the **Tracing** section
3. Toggle **Require x-litellm-trace-id on calls BY this agent** to enable session tracking
4. Set **Max Iterations** to cap the number of LLM calls per session
5. Set **Max Budget Per Session ($)** to cap spend per session
The trace-id flags are stored on the agent's `litellm_params`. Budget controls (`max_iterations`, `max_budget_per_session`) are stored in the virtual key's metadata.
## Configuring via API
Set trace-id enforcement on the agent itself:
```bash
curl -X POST 'http://localhost:4000/v1/agents' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"agent_name": "my-research-agent",
"agent_card_params": {
"name": "my-research-agent",
"description": "A research agent with budget controls",
"url": "http://my-agent:8080",
"version": "1.0.0"
},
"litellm_params": {
"require_trace_id_on_calls_to_agent": true,
"require_trace_id_on_calls_by_agent": true
}
}'
```
Budget controls are set on the agent's `litellm_params` (not on individual keys), so they apply across all keys for the agent:
```bash
curl -X POST 'http://localhost:4000/v1/agents' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"agent_name": "my-research-agent",
"agent_card_params": {
"name": "my-research-agent",
"description": "A research agent with budget controls",
"url": "http://my-agent:8080",
"version": "1.0.0"
},
"litellm_params": {
"require_trace_id_on_calls_by_agent": true,
"max_iterations": 25,
"max_budget_per_session": 5.00
}
}'
```
## How It Works
### Session Tracking
Callers identify their session by including a `session_id` in one of these ways:
- **Header**: `x-litellm-trace-id: my-session-123`
- **Metadata**: `{"metadata": {"session_id": "my-session-123"}}`
### Max Iterations
When `max_iterations` is set in agent `litellm_params`:
- Each LLM call for a session increments a counter
- When the counter exceeds `max_iterations`, the request receives a **429 Too Many Requests**
- Counters expire after 1 hour by default (configurable via `LITELLM_MAX_ITERATIONS_TTL` env var)
### Max Budget Per Session
When `max_budget_per_session` is set in agent `litellm_params`:
- After each successful LLM call, the response cost is accumulated for the session
- Before each call, the accumulated spend is checked against the budget
- When spend exceeds the budget, the request receives a **429 Too Many Requests**
- Session spend counters expire after 1 hour by default (configurable via `LITELLM_MAX_BUDGET_PER_SESSION_TTL` env var)
## Example
Create an agent with max 25 iterations and a $5 budget cap:
<Tabs>
<TabItem value="ui" label="Via UI">
1. Go to **Agents** → **Add Agent**
2. Configure your agent (name, model, etc.)
3. In **Agent Settings**, expand the **Tracing** section
4. Toggle on **Require x-litellm-trace-id on calls BY this agent**
5. Set **Max Iterations** to `25`
6. Set **Max Budget Per Session** to `5.00`
7. Proceed to create a new key for the agent
8. Click **Create Agent**
</TabItem>
<TabItem value="api" label="Via API">
```bash
# 1. Create the agent with trace-id enforcement
curl -X POST 'http://localhost:4000/v1/agents' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"agent_name": "my-research-agent",
"agent_card_params": {
"name": "my-research-agent",
"description": "A research agent with budget controls",
"url": "http://my-agent:8080",
"version": "1.0.0"
},
"litellm_params": {
"require_trace_id_on_calls_by_agent": true
}
}'
# 2. Create a key for the agent
curl -X POST 'http://localhost:4000/key/generate' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"agent_id": "<agent_id_from_step_1>",
"key_alias": "my-research-agent-key"
}'
```
</TabItem>
</Tabs>
### Making Calls with Session Tracking
```bash
curl -X POST 'http://localhost:4000/chat/completions' \
-H 'Authorization: Bearer sk-agent-key-xxx' \
-H 'x-litellm-trace-id: session-abc-123' \
-H 'Content-Type: application/json' \
-d '{
"model": "gpt-4o",
"messages": [{"role": "user", "content": "Hello"}]
}'
```
After 25 calls or $5 spent within this session, subsequent requests will receive:
```json
{
"error": {
"message": "Session budget exceeded for session session-abc-123. Current spend: $5.0032, max_budget_per_session: $5.00.",
"type": "budget_exceeded",
"code": 429
}
}
```
## Environment Variables
| Variable | Default | Description |
|----------|---------|-------------|
| `LITELLM_MAX_ITERATIONS_TTL` | `3600` (1 hour) | TTL in seconds for session iteration counters |
| `LITELLM_MAX_BUDGET_PER_SESSION_TTL` | `3600` (1 hour) | TTL in seconds for session budget counters |

View file

@ -10,6 +10,8 @@ import TabItem from '@theme/TabItem';
**Team member budgets**: Set individual spending limits within the team's shared budget
**Agent budgets**: Set rate limits (tpm/rpm) and session-level caps (iterations, dollar budget) on agents [**Jump**](#agents)
***If a key belongs to a team, the team budget is applied, not the user's personal budget.***
:::
@ -420,6 +422,109 @@ Expected response on failure
</Tabs>
### Agents
Set budgets and rate limits on agents registered with LiteLLM's [Agent Gateway](../a2a.md). You can control:
- **Per-agent rate limits**: `tpm_limit` and `rpm_limit` on the agent itself
- **Per-session rate limits**: `session_tpm_limit` and `session_rpm_limit` applied per session
- **Per-session iteration cap**: `max_iterations` in agent `litellm_params`
- **Per-session budget cap**: `max_budget_per_session` in agent `litellm_params`
<Tabs>
<TabItem value="agent-rate-limits" label="Agent Rate Limits">
Set `tpm_limit` and `rpm_limit` on the agent to cap total throughput across all sessions.
```bash
curl -X POST 'http://localhost:4000/v1/agents' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"agent_name": "my-research-agent",
"agent_card_params": {
"name": "my-research-agent",
"description": "A research agent",
"url": "http://my-agent:8080",
"version": "1.0.0"
},
"tpm_limit": 100000,
"rpm_limit": 100
}'
```
</TabItem>
<TabItem value="session-rate-limits" label="Session Rate Limits">
Set `session_tpm_limit` and `session_rpm_limit` to cap throughput per individual session.
```bash
curl -X POST 'http://localhost:4000/v1/agents' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"agent_name": "my-research-agent",
"agent_card_params": {
"name": "my-research-agent",
"description": "A research agent",
"url": "http://my-agent:8080",
"version": "1.0.0"
},
"session_tpm_limit": 50000,
"session_rpm_limit": 50
}'
```
</TabItem>
<TabItem value="session-budgets" label="Session Budgets">
Set `max_iterations` and `max_budget_per_session` in agent `litellm_params` to cap individual sessions. Requires `require_trace_id_on_calls_by_agent` so LiteLLM can track calls per session.
```bash
curl -X POST 'http://localhost:4000/v1/agents' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"agent_name": "my-research-agent",
"agent_card_params": {
"name": "my-research-agent",
"description": "A research agent",
"url": "http://my-agent:8080",
"version": "1.0.0"
},
"litellm_params": {
"require_trace_id_on_calls_by_agent": true,
"max_iterations": 25,
"max_budget_per_session": 5.00
}
}'
```
When a session exceeds the limit, requests receive a **429 Too Many Requests** response.
See the [Agent Iteration Budgets](../a2a_iteration_budgets) guide for full details.
</TabItem>
</Tabs>
:::info
You can also update rate limits on existing agents using `PATCH /v1/agents/{agent_id}`:
```bash
curl -X PATCH 'http://localhost:4000/v1/agents/<agent_id>' \
-H 'Authorization: Bearer sk-1234' \
-H 'Content-Type: application/json' \
-d '{
"tpm_limit": 200000,
"rpm_limit": 200,
"session_tpm_limit": 50000,
"session_rpm_limit": 50
}'
```
:::
### Customers
Use this to budget `user` passed to `/chat/completions`, **without needing to create a key for every user**
@ -685,6 +790,31 @@ These headers indicate:
- 1 request remaining for the GPT-4 model for key=`sk-ulGNRXWtv7M0lFnnsQk0wQ`
- 179 tokens remaining for the GPT-4 model for key=`sk-ulGNRXWtv7M0lFnnsQk0wQ`
</TabItem>
<TabItem value="per-agent" label="Per Agent">
Set rate limits on agents registered with the [Agent Gateway](../a2a.md).
**Agent-level limits** cap total throughput across all sessions:
```shell
curl -X POST 'http://0.0.0.0:4000/v1/agents' \
--header 'Authorization: Bearer sk-1234' \
--header 'Content-Type: application/json' \
--data '{"agent_name": "my-agent", "agent_card_params": {"name": "my-agent", "description": "My agent", "url": "http://my-agent:8080", "version": "1.0.0"}, "tpm_limit": 100000, "rpm_limit": 100}'
```
**Session-level limits** cap throughput per individual session:
```shell
curl -X POST 'http://0.0.0.0:4000/v1/agents' \
--header 'Authorization: Bearer sk-1234' \
--header 'Content-Type: application/json' \
--data '{"agent_name": "my-agent", "agent_card_params": {"name": "my-agent", "description": "My agent", "url": "http://my-agent:8080", "version": "1.0.0"}, "session_tpm_limit": 50000, "session_rpm_limit": 50}'
```
You can also set **max_iterations** (call count cap) and **max_budget_per_session** (dollar cap) per session via `litellm_params`. See [Agent Iteration Budgets](../a2a_iteration_budgets) for details.
</TabItem>
<TabItem value="per-end-user" label="For customers">

View file

@ -539,7 +539,8 @@ const sidebars = {
"a2a",
"a2a_invoking_agents",
"a2a_cost_tracking",
"a2a_agent_permissions"
"a2a_agent_permissions",
"a2a_iteration_budgets"
],
},
"assistants",

View file

@ -0,0 +1,5 @@
-- AlterTable
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN "tpm_limit" INTEGER;
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN "rpm_limit" INTEGER;
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN "session_tpm_limit" INTEGER;
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN "session_rpm_limit" INTEGER;

View file

@ -67,6 +67,10 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
session_rpm_limit Int?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -38,7 +38,8 @@ def _jsonrpc_error(
def _get_agent(agent_id: str):
"""Look up an agent by ID or name. Returns None if not found."""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.agent_endpoints.agent_registry import \
global_agent_registry
agent = global_agent_registry.get_agent_by_id(agent_id=agent_id)
if agent is None:
@ -46,6 +47,26 @@ def _get_agent(agent_id: str):
return agent
def _enforce_inbound_trace_id(agent: Any, request: Request) -> None:
"""Raise 400 if agent requires x-litellm-trace-id on inbound calls and it is missing."""
agent_litellm_params = agent.litellm_params or {}
if not agent_litellm_params.get("require_trace_id_on_calls_to_agent"):
return
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
headers_dict = dict(request.headers)
trace_id = get_chain_id_from_headers(headers_dict)
if not trace_id:
raise HTTPException(
status_code=400,
detail=(
f"Agent '{agent.agent_id}' requires x-litellm-trace-id header "
"on all inbound requests."
),
)
async def _handle_stream_message(
api_base: Optional[str],
request_id: str,
@ -113,9 +134,8 @@ async def _handle_stream_message(
and request_data is not None
and proxy_logging_obj is not None
):
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
from litellm.proxy.common_request_processing import \
ProxyBaseLLMRequestProcessing
def _ndjson_chunk(chunk: Any) -> str:
if hasattr(chunk, "model_dump"):
@ -215,9 +235,8 @@ async def get_agent_card(
The URL in the agent card is rewritten to point to the LiteLLM proxy,
so all subsequent A2A calls go through LiteLLM for logging and cost tracking.
"""
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
AgentRequestHandler,
)
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import \
AgentRequestHandler
try:
agent = _get_agent(agent_id)
@ -281,15 +300,10 @@ async def invoke_agent_a2a(
"""
from litellm.a2a_protocol import asend_message
from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
AgentRequestHandler,
)
from litellm.proxy.proxy_server import (
general_settings,
proxy_config,
proxy_logging_obj,
version,
)
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import \
AgentRequestHandler
from litellm.proxy.proxy_server import (general_settings, proxy_config,
proxy_logging_obj, version)
body = {}
try:
@ -342,6 +356,8 @@ async def invoke_agent_a2a(
detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.",
)
_enforce_inbound_trace_id(agent, request)
# Get backend URL and agent name
agent_url = agent.agent_card_params.get("url")
agent_name = agent.agent_card_params.get("name", agent_id)
@ -370,9 +386,8 @@ async def invoke_agent_a2a(
)
# Add litellm data (user_api_key, user_id, team_id, etc.)
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
from litellm.proxy.common_request_processing import \
ProxyBaseLLMRequestProcessing
processor = ProxyBaseLLMRequestProcessing(data=body)
data, logging_obj = await processor.common_processing_pre_call_logic(

View file

@ -5,9 +5,8 @@ from typing import Any, Dict, List, Optional
import litellm
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy.management_helpers.object_permission_utils import (
handle_update_object_permission_common,
)
from litellm.proxy.management_helpers.object_permission_utils import \
handle_update_object_permission_common
from litellm.proxy.utils import PrismaClient
from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest
@ -140,6 +139,11 @@ class AgentRegistry:
if object_permission_id is not None:
create_data["object_permission_id"] = object_permission_id
for rate_field in ("tpm_limit", "rpm_limit", "session_tpm_limit", "session_rpm_limit"):
_val = agent.get(rate_field)
if _val is not None:
create_data[rate_field] = _val
# Create agent in DB
created_agent = await prisma_client.db.litellm_agentstable.create(
data=create_data,
@ -214,6 +218,10 @@ class AgentRegistry:
update_data["agent_card_params"] = safe_dumps(
augment_agent.get("agent_card_params")
)
for rate_field in ("tpm_limit", "rpm_limit", "session_tpm_limit", "session_rpm_limit"):
if rate_field in agent:
update_data[rate_field] = agent.get(rate_field)
if agent.get("object_permission") is not None:
agent_copy = dict(augment_agent)
existing_object_permission_id = existing_agent.get(
@ -288,6 +296,12 @@ class AgentRegistry:
"updated_by": updated_by,
"updated_at": datetime.now(timezone.utc),
}
for rate_field in ("tpm_limit", "rpm_limit", "session_tpm_limit", "session_rpm_limit"):
_val = agent.get(rate_field)
if _val is not None:
update_data[rate_field] = _val
if agent.get("object_permission") is not None:
existing_agent = await prisma_client.db.litellm_agentstable.find_unique(
where={"agent_id": agent_id}

View file

@ -5,6 +5,8 @@ from . import *
from .cache_control_check import _PROXY_CacheControlCheck
from .litellm_skills import SkillsInjectionHook
from .max_budget_limiter import _PROXY_MaxBudgetLimiter
from .max_budget_per_session_limiter import _PROXY_MaxBudgetPerSessionHandler
from .max_iterations_limiter import _PROXY_MaxIterationsHandler
from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
from .responses_id_security import ResponsesIDSecurity
@ -23,6 +25,8 @@ PROXY_HOOKS = {
"cache_control_check": _PROXY_CacheControlCheck,
"responses_id_security": ResponsesIDSecurity,
"litellm_skills": SkillsInjectionHook,
"max_iterations_limiter": _PROXY_MaxIterationsHandler,
"max_budget_per_session_limiter": _PROXY_MaxBudgetPerSessionHandler,
}
## FEATURE FLAG HOOKS ##

View file

@ -0,0 +1,271 @@
"""
Per-Session Budget Limiter for LiteLLM Proxy.
Enforces a dollar-amount cap per session (identified by `session_id` /
`x-litellm-trace-id`). After each successful LLM call the response cost is
accumulated against the session. When the accumulated spend exceeds
`max_budget_per_session` (configured in agent litellm_params), subsequent
requests for that session receive a 429.
Note: trace-id enforcement (require_trace_id_on_calls_by_agent) is handled
separately in auth_checks.py at the agent level, not in this hook.
Works across multiple proxy instances via DualCache (in-memory + Redis).
Follows the same pattern as max_iterations_limiter.py.
"""
import os
from typing import TYPE_CHECKING, Any, Optional, Union
from fastapi import HTTPException
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
InternalUsageCache = _InternalUsageCache
else:
InternalUsageCache = Any
# Redis Lua script for atomic float increment with TTL.
# INCRBYFLOAT returns the new value as a string.
# Only sets EXPIRE on first call (when prior value was nil).
MAX_BUDGET_SESSION_INCREMENT_SCRIPT = """
local key = KEYS[1]
local amount = ARGV[1]
local ttl = tonumber(ARGV[2])
local existed = redis.call('EXISTS', key)
local new_val = redis.call('INCRBYFLOAT', key, amount)
if existed == 0 then
redis.call('EXPIRE', key, ttl)
end
return new_val
"""
# Default TTL for session budget counters (1 hour)
DEFAULT_MAX_BUDGET_PER_SESSION_TTL = 3600
class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
"""
Pre-call hook that enforces max_budget_per_session.
Configuration (set in agent litellm_params):
- max_budget_per_session: dollar cap per session_id
Cache key pattern:
{session_budget:<session_id>}:spend
"""
def __init__(self, internal_usage_cache: InternalUsageCache):
self.internal_usage_cache = internal_usage_cache
self.ttl = int(
os.getenv(
"LITELLM_MAX_BUDGET_PER_SESSION_TTL",
DEFAULT_MAX_BUDGET_PER_SESSION_TTL,
)
)
if self.internal_usage_cache.dual_cache.redis_cache is not None:
self.increment_script = (
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
MAX_BUDGET_SESSION_INCREMENT_SCRIPT
)
)
else:
self.increment_script = None
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict,
call_type: str,
) -> Optional[Union[Exception, str, dict]]:
"""
Before each LLM call, check if max_budget_per_session is set and
whether accumulated spend exceeds the budget (429 if so).
"""
max_budget = self._get_max_budget_per_session(user_api_key_dict)
session_id = self._get_session_id(data)
if max_budget is None or session_id is None:
return None
max_budget = float(max_budget)
cache_key = self._make_cache_key(session_id)
current_spend = await self._get_current_spend(cache_key)
verbose_proxy_logger.debug(
"MaxBudgetPerSessionHandler: session_id=%s, spend=%.4f, max=%.2f",
session_id,
current_spend,
max_budget,
)
if current_spend >= max_budget:
raise HTTPException(
status_code=429,
detail=(
f"Session budget exceeded for session {session_id}. "
f"Current spend: ${current_spend:.4f}, "
f"max_budget_per_session: ${max_budget:.2f}."
),
)
return None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
After a successful LLM call, increment the session spend by the response cost.
"""
try:
litellm_params = kwargs.get("litellm_params") or {}
metadata = litellm_params.get("metadata") or {}
session_id = metadata.get("session_id")
if session_id is None:
return
agent_id = metadata.get("agent_id")
if agent_id is None:
return
from litellm.proxy.agent_endpoints.agent_registry import (
global_agent_registry,
)
agent = global_agent_registry.get_agent_by_id(agent_id=str(agent_id))
if agent is None:
return
agent_litellm_params = agent.litellm_params or {}
max_budget = agent_litellm_params.get("max_budget_per_session")
if max_budget is None:
return
response_cost = kwargs.get("response_cost") or 0.0
if response_cost <= 0:
return
cache_key = self._make_cache_key(str(session_id))
await self._increment_spend(cache_key, float(response_cost))
verbose_proxy_logger.debug(
"MaxBudgetPerSessionHandler: incremented session %s spend by %.6f",
session_id,
response_cost,
)
except Exception as e:
verbose_proxy_logger.warning(
"MaxBudgetPerSessionHandler: error in async_log_success_event: %s",
str(e),
)
def _get_session_id(self, data: dict) -> Optional[str]:
"""Extract session_id from request metadata."""
metadata = data.get("metadata") or {}
session_id = metadata.get("session_id")
if session_id is not None:
return str(session_id)
litellm_metadata = data.get("litellm_metadata") or {}
session_id = litellm_metadata.get("session_id")
if session_id is not None:
return str(session_id)
return None
def _get_max_budget_per_session(
self, user_api_key_dict: UserAPIKeyAuth
) -> Optional[float]:
"""Extract max_budget_per_session from agent litellm_params."""
agent_id = user_api_key_dict.agent_id
if agent_id is None:
return None
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
agent = global_agent_registry.get_agent_by_id(agent_id=agent_id)
if agent is None:
return None
litellm_params = agent.litellm_params or {}
max_budget = litellm_params.get("max_budget_per_session")
if max_budget is not None:
return float(max_budget)
return None
def _make_cache_key(self, session_id: str) -> str:
return f"{{session_budget:{session_id}}}:spend"
async def _get_current_spend(self, cache_key: str) -> float:
"""Read current accumulated spend for a session."""
if (
self.internal_usage_cache.dual_cache.redis_cache is not None
):
try:
result = await self.internal_usage_cache.dual_cache.redis_cache.async_get_cache(
key=cache_key
)
if result is not None:
return float(result)
return 0.0
except Exception as e:
verbose_proxy_logger.warning(
"MaxBudgetPerSessionHandler: Redis GET failed, "
"falling back to in-memory: %s",
str(e),
)
result = await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
)
if result is not None:
return float(result)
return 0.0
async def _increment_spend(self, cache_key: str, amount: float) -> float:
"""Atomically increment the session spend and return the new value."""
if self.increment_script is not None:
try:
result = await self.increment_script(
keys=[cache_key],
args=[str(amount), self.ttl],
)
return float(result)
except Exception as e:
verbose_proxy_logger.warning(
"MaxBudgetPerSessionHandler: Redis INCRBYFLOAT failed, "
"falling back to in-memory: %s",
str(e),
)
return await self._in_memory_increment_spend(cache_key, amount)
async def _in_memory_increment_spend(
self, cache_key: str, amount: float
) -> float:
current = await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
)
new_value = (float(current) if current is not None else 0.0) + amount
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=new_value,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
return new_value

View file

@ -4,7 +4,7 @@ Max Iterations Limiter for LiteLLM Proxy.
Enforces a per-session cap on the number of LLM calls an agentic loop can make.
Callers send a `session_id` with each request (via `x-litellm-session-id` header
or `metadata.session_id`), and this hook counts calls per session. When the count
exceeds `max_iterations` (configured in key/team metadata), returns 429.
exceeds `max_iterations` (configured in agent litellm_params), returns 429.
Works across multiple proxy instances via DualCache (in-memory + Redis).
Follows the same pattern as parallel_request_limiter_v3.py.
@ -52,8 +52,8 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
Pre-call hook that enforces max_iterations per session.
Configuration:
- max_iterations: set in key metadata via /key/generate or /key/update
e.g. metadata={"max_iterations": 25}
- max_iterations: set in agent litellm_params
e.g. litellm_params={"max_iterations": 25}
- session_id: sent by caller via x-litellm-session-id header or
metadata.session_id in request body
@ -93,14 +93,13 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
Check session iteration count before making the API call.
Extracts session_id from request metadata and max_iterations from
key metadata. If the session has exceeded max_iterations, raises 429.
agent litellm_params. If the session has exceeded max_iterations, raises 429.
"""
# Extract session_id from request data
session_id = self._get_session_id(data)
if session_id is None:
return None
# Extract max_iterations from key metadata
max_iterations = self._get_max_iterations(user_api_key_dict)
if max_iterations is None:
return None
@ -151,9 +150,20 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
def _get_max_iterations(
self, user_api_key_dict: UserAPIKeyAuth
) -> Optional[int]:
"""Extract max_iterations from key metadata."""
metadata = user_api_key_dict.metadata or {}
max_iterations = metadata.get("max_iterations")
"""Extract max_iterations from agent litellm_params."""
agent_id = user_api_key_dict.agent_id
if agent_id is None:
return None
from litellm.proxy.agent_endpoints.agent_registry import \
global_agent_registry
agent = global_agent_registry.get_agent_by_id(agent_id=agent_id)
if agent is None:
return None
litellm_params = agent.litellm_params or {}
max_iterations = litellm_params.get("max_iterations")
if max_iterations is not None:
return int(max_iterations)
return None

View file

@ -67,6 +67,10 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
session_rpm_limit Int?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")

View file

@ -10,44 +10,27 @@ import traceback
from datetime import date, datetime, timedelta, timezone
from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from typing import (
TYPE_CHECKING,
Any,
Dict,
List,
Literal,
Optional,
Union,
cast,
overload,
)
from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union,
cast, overload)
from litellm import _custom_logger_compatible_callbacks_literal
from litellm.constants import DEFAULT_MODEL_CREATED_AT_TIME, MAX_TEAM_LIST_LIMIT
from litellm.proxy._types import (
DB_CONNECTION_ERROR_TYPES,
CommonProxyErrors,
ProxyErrorTypes,
ProxyException,
SpendLogsMetadata,
SpendLogsPayload,
)
from litellm.constants import (DEFAULT_MODEL_CREATED_AT_TIME,
MAX_TEAM_LIST_LIMIT)
from litellm.proxy._types import (DB_CONNECTION_ERROR_TYPES, CommonProxyErrors,
ProxyErrorTypes, ProxyException,
SpendLogsMetadata, SpendLogsPayload)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypes, CallTypesLiteral
try:
from litellm_enterprise.enterprise_callbacks.send_emails.base_email import (
BaseEmailLogger,
)
from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import (
ResendEmailLogger,
)
from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import (
SendGridEmailLogger,
)
from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import (
SMTPEmailLogger,
)
from litellm_enterprise.enterprise_callbacks.send_emails.base_email import \
BaseEmailLogger
from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import \
ResendEmailLogger
from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import \
SendGridEmailLogger
from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import \
SMTPEmailLogger
except ImportError:
BaseEmailLogger = None # type: ignore
SendGridEmailLogger = None # type: ignore
@ -66,70 +49,60 @@ from fastapi import HTTPException, status
import litellm
import litellm.litellm_core_utils
import litellm.litellm_core_utils.litellm_logging
from litellm import (
EmbeddingResponse,
ImageResponse,
ModelResponse,
ModelResponseStream,
Router,
)
from litellm import (EmbeddingResponse, ImageResponse, ModelResponse,
ModelResponseStream, Router)
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.exceptions import RejectedRequestError
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.integrations.custom_guardrail import (CustomGuardrail,
ModifyResponseException)
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
from litellm.integrations.SlackAlerting.utils import \
_add_langfuse_trace_id_to_alert
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import (
AlertType,
CallInfo,
LiteLLM_VerificationTokenView,
Member,
UserAPIKeyAuth,
)
from litellm.proxy._types import (AlertType, CallInfo,
LiteLLM_VerificationTokenView, Member,
UserAPIKeyAuth)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.db.create_views import (
create_missing_views,
should_create_missing_views,
)
from litellm.proxy.db.create_views import (create_missing_views,
should_create_missing_views)
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.db.log_db_metrics import log_db_metrics
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import \
UnifiedLLMGuardrails
from litellm.proxy.hooks import PROXY_HOOKS, get_proxy_hook
from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck
from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.hooks.max_budget_per_session_limiter import \
_PROXY_MaxBudgetPerSessionHandler
from litellm.proxy.hooks.max_iterations_limiter import \
_PROXY_MaxIterationsHandler
from litellm.proxy.hooks.parallel_request_limiter import \
_PROXY_MaxParallelRequestsHandler
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
from litellm.secret_managers.main import str_to_bool
from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES
from litellm.types.mcp import (
MCPDuringCallResponseObject,
MCPPreCallRequestObject,
MCPPreCallResponseObject,
)
from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult
from litellm.types.mcp import (MCPDuringCallResponseObject,
MCPPreCallRequestObject,
MCPPreCallResponseObject)
from litellm.types.proxy.policy_engine.pipeline_types import \
PipelineExecutionResult
from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
Span = Union[_Span, Any]
else:
@ -309,6 +282,12 @@ class ProxyLogging:
self.internal_usage_cache
)
self.max_budget_limiter = _PROXY_MaxBudgetLimiter()
self.max_iterations_handler = _PROXY_MaxIterationsHandler(
self.internal_usage_cache
)
self.max_budget_per_session_handler = _PROXY_MaxBudgetPerSessionHandler(
self.internal_usage_cache
)
self.cache_control_check = _PROXY_CacheControlCheck()
self.alerting: Optional[List] = None
self.alerting_threshold: float = 300 # default to 5 min. threshold
@ -460,6 +439,8 @@ class ProxyLogging:
def _init_litellm_callbacks(self, llm_router: Optional[Router] = None):
self._add_proxy_hooks(llm_router)
litellm.logging_callback_manager.add_litellm_callback(self.max_iterations_handler)
litellm.logging_callback_manager.add_litellm_callback(self.max_budget_per_session_handler)
litellm.logging_callback_manager.add_litellm_callback(self.service_logging_obj) # type: ignore
# Track string callbacks and their initialized instances so we can
@ -1072,10 +1053,9 @@ class ProxyLogging:
"""Process prompt template if applicable."""
from litellm.proxy.prompts.prompt_endpoints import (
construct_versioned_prompt_id,
get_latest_version_prompt_id,
)
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
construct_versioned_prompt_id, get_latest_version_prompt_id)
from litellm.proxy.prompts.prompt_registry import \
IN_MEMORY_PROMPT_REGISTRY
from litellm.utils import get_non_default_completion_params
if prompt_version is None:
@ -1125,9 +1105,8 @@ class ProxyLogging:
def _process_guardrail_metadata(self, data: dict) -> None:
"""Process guardrails from metadata and add to applied_guardrails."""
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
from litellm.proxy.common_utils.callback_utils import \
add_guardrail_to_applied_guardrails_header
metadata_standard = data.get("metadata") or {}
metadata_litellm = data.get("litellm_metadata") or {}
@ -2024,7 +2003,8 @@ class ProxyLogging:
if isinstance(response, (ModelResponse, ModelResponseStream)):
response_str = litellm.get_response_string(response_obj=response)
elif isinstance(response, dict) and self.is_a2a_streaming_response(response):
from litellm.llms.a2a.common_utils import extract_text_from_a2a_response
from litellm.llms.a2a.common_utils import \
extract_text_from_a2a_response
response_str = extract_text_from_a2a_response(response)
if response_str is not None:
@ -2038,7 +2018,8 @@ class ProxyLogging:
_callback: Optional[CustomLogger] = None
if isinstance(callback, CustomGuardrail):
# Main - V2 Guardrails implementation
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.guardrails import \
GuardrailEventHooks
## CHECK FOR MODEL-LEVEL GUARDRAILS (cached per-request)
if not _guardrail_data_computed:
@ -3583,8 +3564,8 @@ class PrismaClient:
def _get_engine_pid(self) -> int:
try:
engine = self.db._original_prisma._engine # type: ignore[attr-defined]
if engine is not None and engine.process is not None:
return engine.process.pid
if engine is not None and engine.process is not None: # type: ignore[union-attr]
return engine.process.pid # type: ignore[union-attr]
except (AttributeError, TypeError):
pass
return 0
@ -4675,9 +4656,8 @@ async def update_spend_logs_job(
# Guardrail/policy usage tracking (same batch, outside spend-logs update)
try:
from litellm.proxy.guardrails.usage_tracking import (
process_spend_logs_guardrail_usage,
)
from litellm.proxy.guardrails.usage_tracking import \
process_spend_logs_guardrail_usage
await process_spend_logs_guardrail_usage(
prisma_client=prisma_client,
logs_to_process=logs_to_process,
@ -4703,10 +4683,8 @@ async def _monitor_spend_logs_queue(
db_writer_client: Optional HTTP handler for external spend logs endpoint
proxy_logging_obj: Proxy logging object
"""
from litellm.constants import (
SPEND_LOG_QUEUE_POLL_INTERVAL,
SPEND_LOG_QUEUE_SIZE_THRESHOLD,
)
from litellm.constants import (SPEND_LOG_QUEUE_POLL_INTERVAL,
SPEND_LOG_QUEUE_SIZE_THRESHOLD)
threshold = SPEND_LOG_QUEUE_SIZE_THRESHOLD
base_interval = SPEND_LOG_QUEUE_POLL_INTERVAL
@ -5227,12 +5205,11 @@ async def get_available_models_for_user(
List of model names available to the user
"""
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.auth.model_checks import (
get_complete_model_list,
get_key_models,
get_team_models,
)
from litellm.proxy.management_endpoints.team_endpoints import validate_membership
from litellm.proxy.auth.model_checks import (get_complete_model_list,
get_key_models,
get_team_models)
from litellm.proxy.management_endpoints.team_endpoints import \
validate_membership
# Get proxy model list and access groups
if llm_router is None:

View file

@ -179,6 +179,10 @@ class AgentConfig(TypedDict, total=False):
agent_card_params: Required[AgentCard]
litellm_params: Dict[str, Any] # allow for any future litellm params
object_permission: AgentObjectPermission
tpm_limit: Optional[int]
rpm_limit: Optional[int]
session_tpm_limit: Optional[int]
session_rpm_limit: Optional[int]
class PatchAgentRequest(TypedDict, total=False):
@ -186,6 +190,10 @@ class PatchAgentRequest(TypedDict, total=False):
agent_card_params: AgentCard
litellm_params: Dict[str, Any]
object_permission: AgentObjectPermission
tpm_limit: Optional[int]
rpm_limit: Optional[int]
session_tpm_limit: Optional[int]
session_rpm_limit: Optional[int]
# Request/Response models for CRUD endpoints
@ -198,6 +206,10 @@ class AgentResponse(BaseModel):
agent_card_params: Dict[str, Any]
object_permission: Optional[Dict[str, Any]] = None
spend: Optional[float] = None
tpm_limit: Optional[int] = None
rpm_limit: Optional[int] = None
session_tpm_limit: Optional[int] = None
session_rpm_limit: Optional[int] = None
created_at: Optional[datetime] = None
updated_at: Optional[datetime] = None
created_by: Optional[str] = None

View file

@ -0,0 +1,165 @@
"""
Unit Tests for the per-session budget limiter for the proxy.
Tests that session-scoped budget tracking works correctly:
- Enforces max_budget_per_session per session_id (read from agent litellm_params)
- Different sessions have independent budgets
- Requests under budget pass through
- Requests without agent_id pass through
"""
from unittest.mock import patch
import pytest
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.max_budget_per_session_limiter import (
_PROXY_MaxBudgetPerSessionHandler,
)
from litellm.proxy.utils import InternalUsageCache
from litellm.types.agents import AgentResponse
def _make_mock_agent(max_budget_per_session: float) -> AgentResponse:
return AgentResponse(
agent_id="agent-budget-123",
agent_name="budget-agent",
litellm_params={"max_budget_per_session": max_budget_per_session},
agent_card_params={"name": "budget-agent", "version": "1.0.0"},
)
@pytest.mark.asyncio
async def test_budget_per_session_under_budget_passes():
"""
Requests under budget should pass through without error.
"""
local_cache = DualCache()
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-budget",
agent_id="agent-budget-123",
)
mock_agent = _make_mock_agent(max_budget_per_session=5.0)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-budget-1"}},
call_type="",
)
assert result is None
@pytest.mark.asyncio
async def test_budget_per_session_exceeds_budget():
"""
After accumulating spend beyond max_budget_per_session, the next
pre-call check should raise 429.
"""
local_cache = DualCache()
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-budget",
agent_id="agent-budget-123",
)
session_id = "session-over-budget"
cache_key = handler._make_cache_key(session_id)
await handler._increment_spend(cache_key, 1.50)
mock_agent = _make_mock_agent(max_budget_per_session=1.0)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": session_id}},
call_type="",
)
assert exc_info.value.status_code == 429
assert "budget exceeded" in str(exc_info.value.detail).lower()
@pytest.mark.asyncio
async def test_budget_per_session_independent_sessions():
"""
Different session_ids have independent budget counters.
Exhausting session A does not affect session B.
"""
local_cache = DualCache()
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-budget",
agent_id="agent-budget-123",
)
cache_key_a = handler._make_cache_key("session-A")
await handler._increment_spend(cache_key_a, 3.0)
mock_agent = _make_mock_agent(max_budget_per_session=2.0)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
# Session A should be blocked
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-A"}},
call_type="",
)
assert exc_info.value.status_code == 429
# Session B should still pass
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-B"}},
call_type="",
)
assert result is None
@pytest.mark.asyncio
async def test_no_agent_id_passes():
"""
When no agent_id is set on the key, all requests pass through.
"""
local_cache = DualCache()
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-no-agent",
)
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "any-session"}},
call_type="",
)
assert result is None

View file

@ -2,10 +2,12 @@
Unit Tests for the max iterations limiter for the proxy.
Tests that session-scoped iteration counting works correctly:
- Enforces max_iterations per session_id
- Enforces max_iterations per session_id (read from agent litellm_params)
- Different sessions have independent counters
"""
from unittest.mock import MagicMock, patch
import pytest
from fastapi import HTTPException
@ -13,6 +15,16 @@ from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler
from litellm.proxy.utils import InternalUsageCache
from litellm.types.agents import AgentResponse
def _make_mock_agent(max_iterations: int) -> AgentResponse:
return AgentResponse(
agent_id="agent-test-123",
agent_name="test-agent",
litellm_params={"max_iterations": max_iterations},
agent_card_params={"name": "test-agent", "version": "1.0.0"},
)
@pytest.mark.asyncio
@ -28,28 +40,36 @@ async def test_max_iterations_basic_enforcement():
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-1234", metadata={"max_iterations": 3}
api_key="sk-test-key-1234",
agent_id="agent-test-123",
)
# First 3 requests should succeed
for i in range(3):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-abc"}},
call_type="",
)
mock_agent = _make_mock_agent(max_iterations=3)
# 4th request should fail with 429
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-abc"}},
call_type="",
)
assert exc_info.value.status_code == 429
assert "max_iterations" in str(exc_info.value.detail).lower()
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
# First 3 requests should succeed
for i in range(3):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-abc"}},
call_type="",
)
# 4th request should fail with 429
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-abc"}},
call_type="",
)
assert exc_info.value.status_code == 429
assert "max_iterations" in str(exc_info.value.detail).lower()
@pytest.mark.asyncio
@ -65,42 +85,72 @@ async def test_max_iterations_different_sessions_independent():
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-5678", metadata={"max_iterations": 2}
api_key="sk-test-key-5678",
agent_id="agent-test-123",
)
# Session A: 2 calls succeed
for _ in range(2):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-A"}},
call_type="",
)
mock_agent = _make_mock_agent(max_iterations=2)
# Session B: 2 calls succeed (independent counter)
for _ in range(2):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-B"}},
call_type="",
)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = mock_agent
# Session A: 3rd call fails
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-A"}},
call_type="",
)
assert exc_info.value.status_code == 429
# Session A: 2 calls succeed
for _ in range(2):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-A"}},
call_type="",
)
# Session B: 3rd call also fails
with pytest.raises(HTTPException):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-B"}},
call_type="",
)
# Session B: 2 calls succeed (independent counter)
for _ in range(2):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-B"}},
call_type="",
)
# Session A: 3rd call fails
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-A"}},
call_type="",
)
assert exc_info.value.status_code == 429
# Session B: 3rd call also fails
with pytest.raises(HTTPException):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-B"}},
call_type="",
)
@pytest.mark.asyncio
async def test_max_iterations_no_agent_id_passes():
"""
When no agent_id is set on the key, all requests pass through.
"""
local_cache = DualCache()
handler = _PROXY_MaxIterationsHandler(
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-no-agent",
)
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-any"}},
call_type="",
)
assert result is None

View file

@ -1,5 +1,5 @@
import React, { useState, useEffect } from "react";
import { Modal, Form, message, Select, Input, Steps, Radio, Tag, Divider } from "antd";
import { Modal, Form, message, Select, Input, Steps, Radio, Tag, Divider, Switch, InputNumber, Collapse } from "antd";
import { Button } from "@tremor/react";
import { CheckCircleFilled, KeyOutlined, RobotOutlined, AppstoreOutlined, InfoCircleOutlined } from "@ant-design/icons";
import CreatedKeyDisplay from "../shared/CreatedKeyDisplay";
@ -60,6 +60,12 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
const [createdKeyValue, setCreatedKeyValue] = useState<string | null>(null);
const [assignedKeyAlias, setAssignedKeyAlias] = useState<string | null>(null);
// Tracing & guardrails state
const [requireTraceIdInbound, setRequireTraceIdInbound] = useState(false);
const [requireTraceIdOutbound, setRequireTraceIdOutbound] = useState(false);
const [maxIterations, setMaxIterations] = useState<number | null>(null);
const [maxBudgetPerSession, setMaxBudgetPerSession] = useState<number | null>(null);
// Fetch agent type metadata on mount
useEffect(() => {
const fetchMetadata = async () => {
@ -218,6 +224,19 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
}
}
// Wire trace-id flags and budget controls into agent litellm_params (before create call)
if (requireTraceIdInbound || requireTraceIdOutbound) {
if (!agentData.litellm_params) agentData.litellm_params = {};
if (requireTraceIdInbound) {
agentData.litellm_params.require_trace_id_on_calls_to_agent = true;
}
if (requireTraceIdOutbound) {
agentData.litellm_params.require_trace_id_on_calls_by_agent = true;
if (maxIterations) agentData.litellm_params.max_iterations = maxIterations;
if (maxBudgetPerSession) agentData.litellm_params.max_budget_per_session = maxBudgetPerSession;
}
}
const agentResponse = await createAgentCall(accessToken, agentData);
const agentId: string = agentResponse.agent_id;
const agentName: string = agentResponse.agent_name || values.agent_name || agentId;
@ -267,6 +286,10 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
setCreatedAgentName("");
setCreatedKeyValue(null);
setAssignedKeyAlias(null);
setRequireTraceIdInbound(false);
setRequireTraceIdOutbound(false);
setMaxIterations(null);
setMaxBudgetPerSession(null);
onClose();
};
@ -315,6 +338,122 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
</div>
)}
</Form.Item>
<Collapse ghost className="mt-6" items={[
{
key: "tracing",
label: <span className="text-sm font-medium text-gray-700">Tracing</span>,
children: (
<div className="space-y-4">
<div className="flex items-center justify-between">
<div>
<span className="text-sm font-medium text-gray-700">
Require x-litellm-trace-id on calls TO this agent
</span>
<p className="text-xs text-gray-500 mt-1">
Only accept this agent being invoked with a trace-id (e.g. when used as a sub-agent).
</p>
</div>
<Switch
checked={requireTraceIdInbound}
onChange={setRequireTraceIdInbound}
/>
</div>
<div className="flex items-center justify-between">
<div>
<span className="text-sm font-medium text-gray-700">
Require x-litellm-trace-id on calls BY this agent
</span>
<p className="text-xs text-gray-500 mt-1">
Requires LLM/MCP calls made by this agent to include x-litellm-trace-id for session tracking.
</p>
</div>
<Switch
checked={requireTraceIdOutbound}
onChange={(checked) => {
setRequireTraceIdOutbound(checked);
if (!checked) {
setMaxIterations(null);
setMaxBudgetPerSession(null);
}
}}
/>
</div>
</div>
),
},
{
key: "budgets_and_rate_limits",
label: <span className="text-sm font-medium text-gray-700">Budgets &amp; Rate Limits</span>,
children: (
<div className="space-y-4">
{!requireTraceIdOutbound && (
<div className="p-3 bg-yellow-50 border border-yellow-200 rounded-lg text-sm text-yellow-800">
Enable &quot;Require x-litellm-trace-id on calls BY this agent&quot; in Tracing to configure budgets and rate limits.
</div>
)}
<div className="text-sm font-medium text-gray-700">Session Budgets</div>
<div className="grid grid-cols-2 gap-4">
<div>
<label className="text-sm text-gray-600 block mb-1">Max Iterations</label>
<InputNumber
className="w-full"
min={1}
placeholder="e.g. 25"
disabled={!requireTraceIdOutbound}
value={maxIterations}
onChange={(val) => setMaxIterations(val)}
/>
<p className="text-xs text-gray-400 mt-1">Hard cap on LLM calls per session</p>
</div>
<div>
<label className="text-sm text-gray-600 block mb-1">Max Budget Per Session ($)</label>
<InputNumber
className="w-full"
min={0.01}
step={0.5}
placeholder="e.g. 5.00"
disabled={!requireTraceIdOutbound}
value={maxBudgetPerSession}
onChange={(val) => setMaxBudgetPerSession(val)}
/>
<p className="text-xs text-gray-400 mt-1">Max spend per trace before returning 429</p>
</div>
</div>
<Divider className="my-2" />
<div className="text-sm font-medium text-gray-700">Agent Rate Limits</div>
<p className="text-xs text-gray-500">
Global rate limits applied across all callers of this agent.
</p>
<div className="grid grid-cols-2 gap-4">
<Form.Item label="TPM Limit" name="tpm_limit" className="mb-0">
<InputNumber className="w-full" min={0} placeholder="e.g. 100000" disabled={!requireTraceIdOutbound} />
</Form.Item>
<Form.Item label="RPM Limit" name="rpm_limit" className="mb-0">
<InputNumber className="w-full" min={0} placeholder="e.g. 100" disabled={!requireTraceIdOutbound} />
</Form.Item>
</div>
<div className="text-sm font-medium text-gray-700 mt-4">Per-Session Rate Limits</div>
<p className="text-xs text-gray-500">
Rate limits per session (x-litellm-trace-id). Each session gets its own counters.
</p>
<div className="grid grid-cols-2 gap-4">
<Form.Item label="Session TPM Limit" name="session_tpm_limit" className="mb-0">
<InputNumber className="w-full" min={0} placeholder="e.g. 10000" disabled={!requireTraceIdOutbound} />
</Form.Item>
<Form.Item label="Session RPM Limit" name="session_rpm_limit" className="mb-0">
<InputNumber className="w-full" min={0} placeholder="e.g. 20" disabled={!requireTraceIdOutbound} />
</Form.Item>
</div>
</div>
),
},
]} />
</div>
);
@ -456,6 +595,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
<DynamicAgentFormFields agentTypeInfo={selectedAgentTypeInfo} />
) : null}
</div>
</>
);
@ -643,7 +783,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
{/* Step indicator */}
<Steps current={currentStep} size="small" className="mb-8">
<Step title="Configure" />
<Step title="MCP Tools" />
<Step title="Agent Settings" />
<Step title="Assign Key" />
<Step title="Ready" />
</Steps>

View file

@ -283,6 +283,11 @@ export const buildAgentDataFromForm = (values: any, existingAgent?: any) => {
agentData.litellm_params = params;
}
if (values.tpm_limit != null) agentData.tpm_limit = values.tpm_limit;
if (values.rpm_limit != null) agentData.rpm_limit = values.rpm_limit;
if (values.session_tpm_limit != null) agentData.session_tpm_limit = values.session_tpm_limit;
if (values.session_rpm_limit != null) agentData.session_rpm_limit = values.session_rpm_limit;
return agentData;
};
@ -316,5 +321,9 @@ export const parseAgentForForm = (agent: any) => {
cost_per_query: agent.litellm_params?.cost_per_query,
input_cost_per_token: agent.litellm_params?.input_cost_per_token,
output_cost_per_token: agent.litellm_params?.output_cost_per_token,
tpm_limit: agent.tpm_limit,
rpm_limit: agent.rpm_limit,
session_tpm_limit: agent.session_tpm_limit,
session_rpm_limit: agent.session_rpm_limit,
};
};

View file

@ -189,22 +189,6 @@ const AgentFormFields: React.FC<AgentFormFieldsProps> = ({ showAgentName = true,
</Panel>
)}
{/* Tracing */}
{shouldShow(AGENT_FORM_CONFIG.tracing.key) && (
<Panel header={AGENT_FORM_CONFIG.tracing.title} key={AGENT_FORM_CONFIG.tracing.key}>
{AGENT_FORM_CONFIG.tracing.fields.map((field) => (
<Form.Item
key={field.name}
label={field.label}
name={field.name}
valuePropName="checked"
tooltip={field.tooltip}
>
<Switch />
</Form.Item>
))}
</Panel>
)}
</Collapse>
</>
);

View file

@ -1,6 +1,6 @@
import React, { useState, useEffect } from "react";
import { Card, Title, Text, Button as TremorButton, Tab, TabGroup, TabList, TabPanel, TabPanels} from "@tremor/react";
import { Form, Input, Button as AntButton, message, Spin, Descriptions } from "antd";
import { Form, Input, InputNumber, Button as AntButton, message, Spin, Descriptions, Divider } from "antd";
import { ArrowLeftIcon } from "@heroicons/react/outline";
import { getAgentInfo, patchAgentCall, getAgentCreateMetadata, AgentCreateInfo } from "../networking";
import { Agent } from "./types";
@ -201,6 +201,10 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({
{agent.agent_card_params?.documentationUrl && (
<Descriptions.Item label="Documentation URL">{agent.agent_card_params.documentationUrl}</Descriptions.Item>
)}
<Descriptions.Item label="TPM Limit">{agent.tpm_limit ?? "Unlimited"}</Descriptions.Item>
<Descriptions.Item label="RPM Limit">{agent.rpm_limit ?? "Unlimited"}</Descriptions.Item>
<Descriptions.Item label="Session TPM Limit">{agent.session_tpm_limit ?? "Unlimited"}</Descriptions.Item>
<Descriptions.Item label="Session RPM Limit">{agent.session_rpm_limit ?? "Unlimited"}</Descriptions.Item>
<Descriptions.Item label="Created At">{formatDate(agent.created_at)}</Descriptions.Item>
<Descriptions.Item label="Updated At">{formatDate(agent.updated_at)}</Descriptions.Item>
</Descriptions>
@ -295,6 +299,25 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({
<AgentFormFields showAgentName={true} />
)}
<Divider />
<Title className="mb-4">Rate Limits</Title>
<div className="grid grid-cols-2 gap-4">
<Form.Item label="TPM Limit" name="tpm_limit">
<InputNumber className="w-full" min={0} placeholder="Unlimited" />
</Form.Item>
<Form.Item label="RPM Limit" name="rpm_limit">
<InputNumber className="w-full" min={0} placeholder="Unlimited" />
</Form.Item>
</div>
<div className="grid grid-cols-2 gap-4">
<Form.Item label="Session TPM Limit" name="session_tpm_limit">
<InputNumber className="w-full" min={0} placeholder="Unlimited" />
</Form.Item>
<Form.Item label="Session RPM Limit" name="session_rpm_limit">
<InputNumber className="w-full" min={0} placeholder="Unlimited" />
</Form.Item>
</div>
<div className="flex justify-end gap-2 mt-6">
<AntButton onClick={() => {
setIsEditing(false);

View file

@ -118,7 +118,7 @@ export const buildDynamicAgentData = (
litellmParams.model = model;
}
return {
const agentData: Record<string, any> = {
agent_name: values.agent_name,
agent_card_params: {
protocolVersion: "1.0",
@ -140,6 +140,13 @@ export const buildDynamicAgentData = (
},
litellm_params: litellmParams,
};
if (values.tpm_limit != null) agentData.tpm_limit = values.tpm_limit;
if (values.rpm_limit != null) agentData.rpm_limit = values.rpm_limit;
if (values.session_tpm_limit != null) agentData.session_tpm_limit = values.session_tpm_limit;
if (values.session_rpm_limit != null) agentData.session_rpm_limit = values.session_rpm_limit;
return agentData;
};
export default DynamicAgentFormFields;

View file

@ -24,6 +24,10 @@ export interface Agent {
};
object_permission?: AgentObjectPermission;
spend?: number;
tpm_limit?: number | null;
rpm_limit?: number | null;
session_tpm_limit?: number | null;
session_rpm_limit?: number | null;
created_at?: string;
updated_at?: string;
created_by?: string;

View file

@ -935,19 +935,24 @@ export const keyCreateForAgentCall = async (
agentId: string,
keyAlias: string,
models: string[],
metadata?: Record<string, any>,
) => {
const url = proxyBaseUrl ? `${proxyBaseUrl}/key/generate` : `/key/generate`;
const body: Record<string, any> = {
agent_id: agentId,
key_alias: keyAlias,
models: models.length > 0 ? models : [],
};
if (metadata && Object.keys(metadata).length > 0) {
body.metadata = metadata;
}
const response = await fetch(url, {
method: "POST",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify({
agent_id: agentId,
key_alias: keyAlias,
models: models.length > 0 ? models : [],
}),
body: JSON.stringify(body),
});
if (!response.ok) {
@ -8453,6 +8458,10 @@ export const patchAgentCall = async (
agent_name?: string;
litellm_params?: Record<string, any>;
agent_card_params?: Record<string, any>;
tpm_limit?: number | null;
rpm_limit?: number | null;
session_tpm_limit?: number | null;
session_rpm_limit?: number | null;
},
) => {
try {