fix issues with endpoints and edge cases

This commit is contained in:
Harshit28j 2026-03-07 07:00:39 +05:30
parent bb8b8dba43
commit 244763f1fe
7 changed files with 204 additions and 122 deletions

View file

@ -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)

View file

@ -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(

View file

@ -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:

View file

@ -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 {}

View file

@ -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)

View file

@ -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]]

View file

@ -390,7 +390,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
<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.
Enable &quot;Require x-litellm-trace-id on calls BY this agent&quot; in Tracing to configure per-session budgets and rate limits.
</div>
)}
@ -431,10 +431,10 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
</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} />
<InputNumber className="w-full" min={0} placeholder="e.g. 100000" />
</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} />
<InputNumber className="w-full" min={0} placeholder="e.g. 100" />
</Form.Item>
</div>