diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index 61a9ea01b46..4cd3ba39ab2 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -5,12 +5,25 @@ 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 +_JSON_FIELDS = ("agent_card_params", "litellm_params") + + +def _parse_json_fields(d: Dict[str, Any]) -> Dict[str, Any]: + """Prisma may return JSON columns as strings — parse them back to dicts.""" + for field in _JSON_FIELDS: + val = d.get(field) + if isinstance(val, str): + d[field] = json.loads(val) + return d + + class AgentRegistry: def __init__(self): self.agent_list: List[AgentResponse] = [] @@ -139,7 +152,12 @@ 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"): + 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 @@ -150,12 +168,16 @@ class AgentRegistry: include={"object_permission": True}, ) - created_agent_dict = created_agent.model_dump() + created_agent_dict = _parse_json_fields(created_agent.model_dump()) if created_agent.object_permission is not None: try: - created_agent_dict["object_permission"] = created_agent.object_permission.model_dump() + created_agent_dict[ + "object_permission" + ] = created_agent.object_permission.model_dump() except Exception: - created_agent_dict["object_permission"] = created_agent.object_permission.dict() + created_agent_dict[ + "object_permission" + ] = created_agent.object_permission.dict() return AgentResponse(**created_agent_dict) # type: ignore except Exception as e: raise Exception(f"Error adding agent to DB: {str(e)}") @@ -196,7 +218,6 @@ class AgentRegistry: The patched agent """ try: - existing_agent = await prisma_client.db.litellm_agentstable.find_unique( where={"agent_id": agent_id} ) @@ -219,7 +240,12 @@ class AgentRegistry: augment_agent.get("agent_card_params") ) - for rate_field in ("tpm_limit", "rpm_limit", "session_tpm_limit", "session_rpm_limit"): + 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: @@ -227,12 +253,10 @@ class AgentRegistry: existing_object_permission_id = existing_agent.get( "object_permission_id" ) - object_permission_id = ( - await handle_update_object_permission_common( - agent_copy, - existing_object_permission_id, - prisma_client, - ) + object_permission_id = await handle_update_object_permission_common( + agent_copy, + existing_object_permission_id, + prisma_client, ) if object_permission_id is not None: update_data["object_permission_id"] = object_permission_id @@ -246,12 +270,16 @@ class AgentRegistry: }, include={"object_permission": True}, ) - patched_agent_dict = patched_agent.model_dump() + patched_agent_dict = _parse_json_fields(patched_agent.model_dump()) if patched_agent.object_permission is not None: try: - patched_agent_dict["object_permission"] = patched_agent.object_permission.model_dump() + patched_agent_dict[ + "object_permission" + ] = patched_agent.object_permission.model_dump() except Exception: - patched_agent_dict["object_permission"] = patched_agent.object_permission.dict() + patched_agent_dict[ + "object_permission" + ] = patched_agent.object_permission.dict() return AgentResponse(**patched_agent_dict) # type: ignore except Exception as e: raise Exception(f"Error patching agent in DB: {str(e)}") @@ -297,7 +325,12 @@ class AgentRegistry: "updated_at": datetime.now(timezone.utc), } - for rate_field in ("tpm_limit", "rpm_limit", "session_tpm_limit", "session_rpm_limit"): + 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 @@ -312,12 +345,10 @@ class AgentRegistry: else None ) agent_copy = dict(agent) - object_permission_id = ( - await handle_update_object_permission_common( - agent_copy, - existing_object_permission_id, - prisma_client, - ) + object_permission_id = await handle_update_object_permission_common( + agent_copy, + existing_object_permission_id, + prisma_client, ) if object_permission_id is not None: update_data["object_permission_id"] = object_permission_id @@ -329,12 +360,16 @@ class AgentRegistry: include={"object_permission": True}, ) - updated_agent_dict = updated_agent.model_dump() + updated_agent_dict = _parse_json_fields(updated_agent.model_dump()) if updated_agent.object_permission is not None: try: - updated_agent_dict["object_permission"] = updated_agent.object_permission.model_dump() + updated_agent_dict[ + "object_permission" + ] = updated_agent.object_permission.model_dump() except Exception: - updated_agent_dict["object_permission"] = updated_agent.object_permission.dict() + updated_agent_dict[ + "object_permission" + ] = updated_agent.object_permission.dict() return AgentResponse(**updated_agent_dict) # type: ignore except Exception as e: raise Exception(f"Error updating agent in DB: {str(e)}") @@ -354,11 +389,13 @@ class AgentRegistry: agents: List[Dict[str, Any]] = [] for agent in agents_from_db: - agent_dict = dict(agent) + agent_dict = _parse_json_fields(dict(agent)) # object_permission is eagerly loaded via include above if agent.object_permission is not None: try: - agent_dict["object_permission"] = agent.object_permission.model_dump() + agent_dict[ + "object_permission" + ] = agent.object_permission.model_dump() except Exception: agent_dict["object_permission"] = agent.object_permission.dict() agents.append(agent_dict) diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 40480bec3e2..9baa8bbccc0 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -14,16 +14,19 @@ from fastapi import APIRouter, Depends, HTTPException, Request import litellm from litellm._logging import verbose_proxy_logger -from litellm.proxy._types import (CommonProxyErrors, LitellmUserRoles, - UserAPIKeyAuth) +from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.management_endpoints.common_daily_activity import \ - get_daily_activity -from litellm.types.agents import (AgentConfig, AgentMakePublicResponse, - AgentResponse, MakeAgentsPublicRequest, - PatchAgentRequest) -from litellm.types.proxy.management_endpoints.common_daily_activity import \ - SpendAnalyticsPaginatedResponse +from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity +from litellm.types.agents import ( + AgentConfig, + AgentMakePublicResponse, + AgentResponse, + MakeAgentsPublicRequest, + PatchAgentRequest, +) +from litellm.types.proxy.management_endpoints.common_daily_activity import ( + SpendAnalyticsPaginatedResponse, +) router = APIRouter() @@ -66,10 +69,10 @@ async def get_agents( Returns: List[AgentResponse] """ - from litellm.proxy.agent_endpoints.agent_registry import \ - global_agent_registry - from litellm.proxy.agent_endpoints.auth.agent_permission_handler import \ - AgentRequestHandler + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + ) try: returned_agents: List[AgentResponse] = [] @@ -96,27 +99,14 @@ async def get_agents( agent for agent in all_agents if agent.agent_id in allowed_agent_ids ] - # Fetch current spend from DB for all returned agents - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is not None: - agent_ids = [agent.agent_id for agent in returned_agents] - if agent_ids: - db_agents = await prisma_client.db.litellm_agentstable.find_many( - where={"agent_id": {"in": agent_ids}}, - ) - spend_map = {a.agent_id: a.spend for a in db_agents} - for agent in returned_agents: - if agent.agent_id in spend_map: - agent.spend = spend_map[agent.agent_id] - # add is_public field to each agent - we do it this way, to allow setting config agents as public for agent in returned_agents: if agent.litellm_params is None: agent.litellm_params = {} - agent.litellm_params["is_public"] = ( - litellm.public_agent_groups is not None - and (agent.agent_id in litellm.public_agent_groups) + agent.litellm_params[ + "is_public" + ] = litellm.public_agent_groups is not None and ( + agent.agent_id in litellm.public_agent_groups ) return returned_agents @@ -135,8 +125,9 @@ async def get_agents( #### CRUD ENDPOINTS FOR AGENTS #### -from litellm.proxy.agent_endpoints.agent_registry import \ - global_agent_registry as AGENT_REGISTRY +from litellm.proxy.agent_endpoints.agent_registry import ( + global_agent_registry as AGENT_REGISTRY, +) @router.post( @@ -265,25 +256,21 @@ async def get_agent_by_id(agent_id: str): include={"object_permission": True}, ) if agent_row is not None: - agent_dict = agent_row.model_dump() + from litellm.proxy.agent_endpoints.agent_registry import ( + _parse_json_fields, + ) + + agent_dict = _parse_json_fields(agent_row.model_dump()) if agent_row.object_permission is not None: try: - agent_dict["object_permission"] = ( - agent_row.object_permission.model_dump() - ) + agent_dict[ + "object_permission" + ] = agent_row.object_permission.model_dump() except Exception: - agent_dict["object_permission"] = ( - agent_row.object_permission.dict() - ) + agent_dict[ + "object_permission" + ] = agent_row.object_permission.dict() agent = AgentResponse(**agent_dict) # type: ignore - else: - # Agent found in memory — refresh spend from DB - db_row = await prisma_client.db.litellm_agentstable.find_unique( - where={"agent_id": agent_id} - ) - if db_row is not None: - agent.spend = db_row.spend - if agent is None: raise HTTPException( status_code=404, detail=f"Agent with ID {agent_id} not found" @@ -584,8 +571,9 @@ async def make_agent_public( try: # Update the public model groups import litellm - from litellm.proxy.agent_endpoints.agent_registry import \ - global_agent_registry as AGENT_REGISTRY + from litellm.proxy.agent_endpoints.agent_registry import ( + global_agent_registry as AGENT_REGISTRY, + ) from litellm.proxy.proxy_server import proxy_config # Check if user has admin permissions @@ -606,7 +594,11 @@ async def make_agent_public( where={"agent_id": agent_id} ) if agent is not None: - agent = AgentResponse(**agent.model_dump()) # type: ignore + from litellm.proxy.agent_endpoints.agent_registry import ( + _parse_json_fields, + ) + + agent = AgentResponse(**_parse_json_fields(agent.model_dump())) # type: ignore if agent is None: raise HTTPException( @@ -700,8 +692,9 @@ async def make_agents_public( try: # Update the public model groups import litellm - from litellm.proxy.agent_endpoints.agent_registry import \ - global_agent_registry as AGENT_REGISTRY + from litellm.proxy.agent_endpoints.agent_registry import ( + global_agent_registry as AGENT_REGISTRY, + ) from litellm.proxy.proxy_server import proxy_config # Load existing config @@ -728,7 +721,11 @@ async def make_agents_public( where={"agent_id": agent_id} ) if agent is not None: - agent = AgentResponse(**agent.model_dump()) # type: ignore + from litellm.proxy.agent_endpoints.agent_registry import ( + _parse_json_fields, + ) + + agent = AgentResponse(**_parse_json_fields(agent.model_dump())) # type: ignore if agent is None: raise HTTPException( diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 28b1e6601b1..6d91124512e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -856,9 +856,7 @@ class DBSpendUpdateWriter: or {} ), len( - db_spend_update_transactions.get( - "agent_list_transactions" - ) + db_spend_update_transactions.get("agent_list_transactions") or {} ), ) @@ -1345,7 +1343,9 @@ class DBSpendUpdateWriter: ) ### UPDATE AGENT TABLE ### - agent_list_transactions = db_spend_update_transactions["agent_list_transactions"] + agent_list_transactions = db_spend_update_transactions[ + "agent_list_transactions" + ] await DBSpendUpdateWriter._update_entity_spend_in_db( entity_name="Agent", transactions=agent_list_transactions, @@ -1356,6 +1356,17 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, ) + # Also update in-memory agent registry so GET /v1/agents returns fresh spend + if agent_list_transactions: + from litellm.proxy.agent_endpoints.agent_registry import ( + global_agent_registry, + ) + + for agent_id, response_cost in agent_list_transactions.items(): + agent = global_agent_registry.get_agent_by_id(agent_id=agent_id) + if agent is not None: + agent.spend = (agent.spend or 0) + response_cost + @staticmethod async def _update_entity_spend_in_db( entity_name: str, @@ -1615,14 +1626,14 @@ class DBSpendUpdateWriter: # Add cache-related fields if they exist if "cache_read_input_tokens" in transaction: - common_data["cache_read_input_tokens"] = ( - transaction.get("cache_read_input_tokens", 0) - ) + common_data[ + "cache_read_input_tokens" + ] = transaction.get("cache_read_input_tokens", 0) if "cache_creation_input_tokens" in transaction: - common_data["cache_creation_input_tokens"] = ( - transaction.get( - "cache_creation_input_tokens", 0 - ) + common_data[ + "cache_creation_input_tokens" + ] = transaction.get( + "cache_creation_input_tokens", 0 ) if entity_type == "tag" and "request_id" in transaction: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b5bbb4237c1..3131776f942 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -909,6 +909,29 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + # Agent rate limits + _agent_id = user_api_key_dict.agent_id + if _agent_id is not 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 not None and ( + _agent.tpm_limit is not None or _agent.rpm_limit is not None + ): + descriptors.append( + RateLimitDescriptor( + key="agent", + value=_agent_id, + rate_limit={ + "requests_per_unit": _agent.rpm_limit, + "tokens_per_unit": _agent.tpm_limit, + "window_size": self.window_size, + }, + ) + ) + # Model rate limits requested_model = data.get("model", None) self._add_model_per_key_rate_limit_descriptor( @@ -1503,6 +1526,20 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + # Agent TPM + user_api_key_agent_id = standard_logging_metadata.get( + "user_api_key_agent_id" + ) + if user_api_key_agent_id: + pipeline_operations.extend( + self._create_pipeline_operations( + key="agent", + value=user_api_key_agent_id, + rate_limit_type="tokens", + total_tokens=total_tokens, + ) + ) + # Model-specific TPM if model_group and user_api_key: pipeline_operations.extend( @@ -1555,9 +1592,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): from litellm.types.caching import RedisPipelineIncrementOperation try: - litellm_parent_otel_span: Union[Span, None] = ( - _get_parent_otel_span_from_kwargs(kwargs) - ) + litellm_parent_otel_span: Union[ + Span, None + ] = _get_parent_otel_span_from_kwargs(kwargs) # Get metadata from standard_logging_object - this correctly handles both # 'metadata' and 'litellm_metadata' fields from litellm_params standard_logging_object = kwargs.get("standard_logging_object") or {} diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index b612c6883f8..35bbcfb0926 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -196,12 +196,12 @@ def _get_dynamic_logging_metadata( user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig ) -> Optional[TeamCallbackMetadata]: callback_settings_obj: Optional[TeamCallbackMetadata] = None - key_dynamic_logging_settings: Optional[dict] = ( - KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) - ) - team_dynamic_logging_settings: Optional[dict] = ( - KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) - ) + key_dynamic_logging_settings: Optional[ + dict + ] = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) + team_dynamic_logging_settings: Optional[ + dict + ] = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) ######################################################################################### # Key-based callbacks ######################################################################################### @@ -602,7 +602,6 @@ class LiteLLMProxyRequestSetup: "x-litellm-session-id" ) - if agent_id_from_header: metadata_from_headers["agent_id"] = agent_id_from_header verbose_proxy_logger.debug( @@ -637,6 +636,7 @@ class LiteLLMProxyRequestSetup: user_api_key_org_id=user_api_key_dict.org_id, user_api_key_team_alias=user_api_key_dict.team_alias, user_api_key_end_user_id=user_api_key_dict.end_user_id, + user_api_key_agent_id=user_api_key_dict.agent_id, user_api_key_user_email=user_api_key_dict.user_email, user_api_key_request_route=user_api_key_dict.request_route, user_api_key_budget_reset_at=( @@ -728,11 +728,11 @@ class LiteLLMProxyRequestSetup: ## KEY-LEVEL SPEND LOGS / TAGS if "tags" in key_metadata and key_metadata["tags"] is not None: - data[_metadata_variable_name]["tags"] = ( - LiteLLMProxyRequestSetup._merge_tags( - request_tags=data[_metadata_variable_name].get("tags"), - tags_to_add=key_metadata["tags"], - ) + data[_metadata_variable_name][ + "tags" + ] = LiteLLMProxyRequestSetup._merge_tags( + request_tags=data[_metadata_variable_name].get("tags"), + tags_to_add=key_metadata["tags"], ) if "disable_global_guardrails" in key_metadata and isinstance( key_metadata["disable_global_guardrails"], bool @@ -1007,9 +1007,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data[_metadata_variable_name]["litellm_api_version"] = version if general_settings is not None: - data[_metadata_variable_name]["global_max_parallel_requests"] = ( - general_settings.get("global_max_parallel_requests", None) - ) + data[_metadata_variable_name][ + "global_max_parallel_requests" + ] = general_settings.get("global_max_parallel_requests", None) ### KEY-LEVEL Controls key_metadata = user_api_key_dict.metadata @@ -1097,14 +1097,14 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ] = user_api_key_dict.user_max_budget data[_metadata_variable_name]["user_api_key_metadata"] = user_api_key_dict.metadata - data[_metadata_variable_name]["user_api_key_team_metadata"] = ( - user_api_key_dict.team_metadata + data[_metadata_variable_name][ + "user_api_key_team_metadata" + ] = user_api_key_dict.team_metadata + data[_metadata_variable_name]["user_api_key_object_permission_id"] = getattr( + user_api_key_dict, "object_permission_id", None ) - data[_metadata_variable_name]["user_api_key_object_permission_id"] = ( - getattr(user_api_key_dict, "object_permission_id", None) - ) - data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = ( - getattr(user_api_key_dict, "team_object_permission_id", None) + data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = getattr( + user_api_key_dict, "team_object_permission_id", None ) data[_metadata_variable_name]["headers"] = _headers data[_metadata_variable_name]["endpoint"] = str(request.url) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 3becc6b41df..d12a398d32f 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1683,7 +1683,6 @@ class StreamingChatCompletionChunk(OpenAIChatCompletionChunk): super().__init__(**kwargs) - class ModelResponseBase(OpenAIObject): id: str """A unique identifier for the completion.""" @@ -2440,6 +2439,7 @@ class StandardLoggingUserAPIKeyMetadata(TypedDict): user_api_key_user_email: Optional[str] user_api_key_team_alias: Optional[str] user_api_key_end_user_id: Optional[str] + user_api_key_agent_id: Optional[str] user_api_key_request_route: Optional[str] user_api_key_auth_metadata: Optional[Dict[str, str]] 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 0cec0331f43..2c813c1e99c 100644 --- a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx @@ -390,7 +390,7 @@ const AddAgentForm: React.FC = ({
{!requireTraceIdOutbound && (
- Enable "Require x-litellm-trace-id on calls BY this agent" in Tracing to configure budgets and rate limits. + Enable "Require x-litellm-trace-id on calls BY this agent" in Tracing to configure per-session budgets and rate limits.
)} @@ -431,10 +431,10 @@ const AddAgentForm: React.FC = ({

- + - +