From bea664389bcb06427b43a4b5a2372287e2c03144 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 4 Mar 2026 19:51:19 -0800 Subject: [PATCH] feat(agents/): support budgets + rate limiting on agents + agent sessions --- docs/my-website/docs/a2a.md | 1 + docs/my-website/docs/a2a_iteration_budgets.md | 188 ++++++++++++ docs/my-website/docs/proxy/users.md | 130 +++++++++ docs/my-website/sidebars.js | 3 +- .../migration.sql | 5 + .../litellm_proxy_extras/schema.prisma | 4 + .../proxy/agent_endpoints/a2a_endpoints.py | 53 ++-- .../proxy/agent_endpoints/agent_registry.py | 20 +- litellm/proxy/hooks/__init__.py | 4 + .../hooks/max_budget_per_session_limiter.py | 271 ++++++++++++++++++ litellm/proxy/hooks/max_iterations_limiter.py | 26 +- litellm/proxy/schema.prisma | 4 + litellm/proxy/utils.py | 161 +++++------ litellm/types/agents.py | 12 + .../test_max_budget_per_session_limiter.py | 165 +++++++++++ .../hooks/test_max_iterations_limiter.py | 158 ++++++---- .../src/components/agents/add_agent_form.tsx | 144 +++++++++- .../src/components/agents/agent_config.ts | 9 + .../components/agents/agent_form_fields.tsx | 16 -- .../src/components/agents/agent_info.tsx | 25 +- .../agents/dynamic_agent_form_fields.tsx | 9 +- .../src/components/agents/types.ts | 4 + .../src/components/networking.tsx | 19 +- 23 files changed, 1229 insertions(+), 202 deletions(-) create mode 100644 docs/my-website/docs/a2a_iteration_budgets.md create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260305000000_add_rate_limits_to_agents/migration.sql create mode 100644 litellm/proxy/hooks/max_budget_per_session_limiter.py create mode 100644 tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py diff --git a/docs/my-website/docs/a2a.md b/docs/my-website/docs/a2a.md index b1166a7809c..9c86d0de383 100644 --- a/docs/my-website/docs/a2a.md +++ b/docs/my-website/docs/a2a.md @@ -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 diff --git a/docs/my-website/docs/a2a_iteration_budgets.md b/docs/my-website/docs/a2a_iteration_budgets.md new file mode 100644 index 00000000000..47beca3470f --- /dev/null +++ b/docs/my-website/docs/a2a_iteration_budgets.md @@ -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: + + + + +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** + + + + +```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": "", + "key_alias": "my-research-agent-key" + }' +``` + + + + +### 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 | diff --git a/docs/my-website/docs/proxy/users.md b/docs/my-website/docs/proxy/users.md index 8517db51a8f..58813eaf49e 100644 --- a/docs/my-website/docs/proxy/users.md +++ b/docs/my-website/docs/proxy/users.md @@ -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 +### 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` + + + + +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 + }' +``` + + + + +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 + }' +``` + + + + +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. + + + + +:::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/' \ + -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` + + + +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. + diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index f7487d24b12..0efab4edf50 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -539,7 +539,8 @@ const sidebars = { "a2a", "a2a_invoking_agents", "a2a_cost_tracking", - "a2a_agent_permissions" + "a2a_agent_permissions", + "a2a_iteration_budgets" ], }, "assistants", diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260305000000_add_rate_limits_to_agents/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260305000000_add_rate_limits_to_agents/migration.sql new file mode 100644 index 00000000000..3cd8ca638a4 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260305000000_add_rate_limits_to_agents/migration.sql @@ -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; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 638bc63285a..ed4f12b9154 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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") diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 6bcee14f29e..808954cdc96 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -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( diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index 159c9fb93d9..61a9ea01b46 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -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} diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index 1d1e559d4be..790ebcd8791 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -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 ## diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py new file mode 100644 index 00000000000..a981207f000 --- /dev/null +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -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:}: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 diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index 8d481f6b261..acfe45c9b90 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 638bc63285a..ed4f12b9154 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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") diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index afcdd9d0c50..62ed59af402 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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: diff --git a/litellm/types/agents.py b/litellm/types/agents.py index 92ab437c370..d2da82bbc68 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -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 diff --git a/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py new file mode 100644 index 00000000000..879e2d65c7a --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_max_budget_per_session_limiter.py @@ -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 diff --git a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py index deb1c483b87..20928ef46d5 100644 --- a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py @@ -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 diff --git a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx index 10fbfa4615a..0cec0331f43 100644 --- a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx @@ -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 = ({ const [createdKeyValue, setCreatedKeyValue] = useState(null); const [assignedKeyAlias, setAssignedKeyAlias] = useState(null); + // Tracing & guardrails state + const [requireTraceIdInbound, setRequireTraceIdInbound] = useState(false); + const [requireTraceIdOutbound, setRequireTraceIdOutbound] = useState(false); + const [maxIterations, setMaxIterations] = useState(null); + const [maxBudgetPerSession, setMaxBudgetPerSession] = useState(null); + // Fetch agent type metadata on mount useEffect(() => { const fetchMetadata = async () => { @@ -218,6 +224,19 @@ const AddAgentForm: React.FC = ({ } } + // 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 = ({ setCreatedAgentName(""); setCreatedKeyValue(null); setAssignedKeyAlias(null); + setRequireTraceIdInbound(false); + setRequireTraceIdOutbound(false); + setMaxIterations(null); + setMaxBudgetPerSession(null); onClose(); }; @@ -315,6 +338,122 @@ const AddAgentForm: React.FC = ({ )} + + Tracing, + children: ( +
+
+
+ + Require x-litellm-trace-id on calls TO this agent + +

+ Only accept this agent being invoked with a trace-id (e.g. when used as a sub-agent). +

+
+ +
+ +
+
+ + Require x-litellm-trace-id on calls BY this agent + +

+ Requires LLM/MCP calls made by this agent to include x-litellm-trace-id for session tracking. +

+
+ { + setRequireTraceIdOutbound(checked); + if (!checked) { + setMaxIterations(null); + setMaxBudgetPerSession(null); + } + }} + /> +
+
+ ), + }, + { + key: "budgets_and_rate_limits", + label: Budgets & Rate Limits, + children: ( +
+ {!requireTraceIdOutbound && ( +
+ Enable "Require x-litellm-trace-id on calls BY this agent" in Tracing to configure budgets and rate limits. +
+ )} + +
Session Budgets
+
+
+ + setMaxIterations(val)} + /> +

Hard cap on LLM calls per session

+
+
+ + setMaxBudgetPerSession(val)} + /> +

Max spend per trace before returning 429

+
+
+ + + +
Agent Rate Limits
+

+ Global rate limits applied across all callers of this agent. +

+
+ + + + + + +
+ +
Per-Session Rate Limits
+

+ Rate limits per session (x-litellm-trace-id). Each session gets its own counters. +

+
+ + + + + + +
+
+ ), + }, + ]} /> ); @@ -456,6 +595,7 @@ const AddAgentForm: React.FC = ({ ) : null} + ); @@ -643,7 +783,7 @@ const AddAgentForm: React.FC = ({ {/* Step indicator */} - + diff --git a/ui/litellm-dashboard/src/components/agents/agent_config.ts b/ui/litellm-dashboard/src/components/agents/agent_config.ts index 3fc80373025..2e8570f86dd 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_config.ts +++ b/ui/litellm-dashboard/src/components/agents/agent_config.ts @@ -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, }; }; diff --git a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx index 8e2574efec2..e3ed70fcaf0 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx +++ b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx @@ -189,22 +189,6 @@ const AgentFormFields: React.FC = ({ showAgentName = true, )} - {/* Tracing */} - {shouldShow(AGENT_FORM_CONFIG.tracing.key) && ( - - {AGENT_FORM_CONFIG.tracing.fields.map((field) => ( - - - - ))} - - )}
); diff --git a/ui/litellm-dashboard/src/components/agents/agent_info.tsx b/ui/litellm-dashboard/src/components/agents/agent_info.tsx index deb4f900377..b41e318a766 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_info.tsx +++ b/ui/litellm-dashboard/src/components/agents/agent_info.tsx @@ -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 = ({ {agent.agent_card_params?.documentationUrl && ( {agent.agent_card_params.documentationUrl} )} + {agent.tpm_limit ?? "Unlimited"} + {agent.rpm_limit ?? "Unlimited"} + {agent.session_tpm_limit ?? "Unlimited"} + {agent.session_rpm_limit ?? "Unlimited"} {formatDate(agent.created_at)} {formatDate(agent.updated_at)} @@ -295,6 +299,25 @@ const AgentInfoView: React.FC = ({ )} + + Rate Limits +
+ + + + + + +
+
+ + + + + + +
+
{ setIsEditing(false); diff --git a/ui/litellm-dashboard/src/components/agents/dynamic_agent_form_fields.tsx b/ui/litellm-dashboard/src/components/agents/dynamic_agent_form_fields.tsx index 55a4a62953a..1138b0de730 100644 --- a/ui/litellm-dashboard/src/components/agents/dynamic_agent_form_fields.tsx +++ b/ui/litellm-dashboard/src/components/agents/dynamic_agent_form_fields.tsx @@ -118,7 +118,7 @@ export const buildDynamicAgentData = ( litellmParams.model = model; } - return { + const agentData: Record = { 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; diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts index 0e789e0376e..3e27177c815 100644 --- a/ui/litellm-dashboard/src/components/agents/types.ts +++ b/ui/litellm-dashboard/src/components/agents/types.ts @@ -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; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index e231f7c43e8..a9efe93f777 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -935,19 +935,24 @@ export const keyCreateForAgentCall = async ( agentId: string, keyAlias: string, models: string[], + metadata?: Record, ) => { const url = proxyBaseUrl ? `${proxyBaseUrl}/key/generate` : `/key/generate`; + const body: Record = { + 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; agent_card_params?: Record; + tpm_limit?: number | null; + rpm_limit?: number | null; + session_tpm_limit?: number | null; + session_rpm_limit?: number | null; }, ) => { try {