Merge pull request #22151 from BerriAI/litellm_fix_cicd_26_02

[Fix] CICD 26/02/26
This commit is contained in:
Sameer Kankute 2026-02-26 13:30:23 +05:30 • committed by GitHub
commit 4d68151d03
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
32 changed files with 286 additions and 193 deletions

View file

@ -796,6 +796,7 @@ router_settings:
| PYROSCOPE_SERVER_ADDRESS | Pyroscope server URL to send profiles to. Required when LITELLM_ENABLE_PYROSCOPE is true. No default.
| PYROSCOPE_SAMPLE_RATE | Optional. Sample rate for Pyroscope profiling (integer). No default; when unset, the pyroscope-io library default is used.
| LITELLM_MASTER_KEY | Master key for proxy authentication
| LITELLM_MAX_ITERATIONS_TTL | TTL in seconds for session iteration counters used by the max-iterations limiter. Default is 3600 (1 hour)
| LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development)
| LITELLM_NON_ROOT | Flag to run LiteLLM in non-root mode for enhanced security in Docker containers
| LITELLM_RATE_LIMIT_WINDOW_SIZE | Rate limit window size for LiteLLM. Default is 60
@ -991,6 +992,7 @@ router_settings:
| TOGETHER_AI_EMBEDDING_150_M | Size parameter for Together AI 150M embedding model. Default is 150
| TOGETHER_AI_EMBEDDING_350_M | Size parameter for Together AI 350M embedding model. Default is 350
| TOOL_CHOICE_OBJECT_TOKEN_COUNT | Token count for tool choice objects. Default is 4
| TOOL_POLICY_CACHE_TTL_SECONDS | TTL in seconds for caching tool policy guardrail results. Default is 60
| UI_LOGO_PATH | Path to the logo image used in the UI
| UI_PASSWORD | Password for accessing the UI
| UI_USERNAME | Username for accessing the UI

View file

@ -1,7 +1,7 @@
import asyncio
import concurrent.futures
import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
import litellm
from litellm._logging import verbose_logger
@ -92,7 +92,7 @@ class RealTimeStreaming:
message_obj = message
else:
message_obj = json.loads(message)
self._collect_tool_calls_from_response_done(message_obj)
self._collect_tool_calls_from_response_done(cast(dict, message_obj))
try:
if (
not isinstance(message, dict)
@ -159,7 +159,7 @@ class RealTimeStreaming:
event_type
== "conversation.item.input_audio_transcription.completed"
):
transcript = event_obj.get("transcript", "")
transcript = cast(str, event_obj.get("transcript", ""))
if transcript:
self.input_messages.append(
{"role": "user", "content": transcript}
@ -174,7 +174,7 @@ class RealTimeStreaming:
try:
if event_obj.get("type") != "response.done":
return
response = event_obj.get("response", {})
response = cast(Dict[str, Any], event_obj.get("response", {}))
for item in response.get("output", []):
if item.get("type") == "function_call":
self.tool_calls.append(
@ -428,11 +428,11 @@ class RealTimeStreaming:
== "conversation.item.input_audio_transcription.completed"
):
transcript = event.get("transcript", "")
self._collect_user_input_from_backend_event(event)
self._collect_user_input_from_backend_event(cast(dict, event))
self.store_message(event_str)
await self.websocket.send_text(event_str)
blocked = await self.run_realtime_guardrails(
transcript, item_id=event.get("item_id")
cast(str, transcript), item_id=cast(Optional[str], event.get("item_id"))
)
if not blocked:
await self._send_to_backend(

View file

@ -11,7 +11,6 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.bedrock.chat.invoke_transformations.amazon_qwen3_transformation import (
AmazonQwen3Config,
)
@ -19,7 +18,7 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
LiteLLMLoggingObj,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
from litellm.types.utils import ModelResponse, Usage
class AmazonQwen2Config(AmazonQwen3Config):
@ -80,10 +79,14 @@ class AmazonQwen2Config(AmazonQwen3Config):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
setattr(
model_response,
"usage",
Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
),
)
return model_response

View file

@ -10,14 +10,13 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
LiteLLMLoggingObj,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
from litellm.types.utils import ModelResponse, Usage
class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
@ -202,10 +201,14 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
setattr(
model_response,
"usage",
Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
),
)
return model_response

View file

@ -1,4 +1,4 @@
from typing import Dict, List, Optional, Set, Tuple
from typing import Dict, List, Optional, Set, Tuple, cast
from fastapi import HTTPException
from starlette.datastructures import Headers
@ -539,7 +539,7 @@ class MCPRequestHandler:
allowed_tools = team_tools
else:
# No team restrictions → use key restrictions
allowed_tools = key_tools
allowed_tools = cast(List[str], key_tools)
# Intersect with agent's tool permissions if agent_id is set
if user_api_key_auth.agent_id:

View file

@ -1,40 +1,59 @@
import enum
import json
from datetime import datetime
from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Literal,
Optional, Union)
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union
import httpx
from pydantic import (BaseModel, ConfigDict, Field, Json, field_validator,
model_validator)
from pydantic import (
BaseModel,
ConfigDict,
Field,
Json,
field_validator,
model_validator,
)
from typing_extensions import Required, TypedDict
from litellm._uuid import uuid
from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.openai import (AllMessageValues, OpenAIFileObject,
ResponsesAPIResponse)
from litellm.types.mcp import (MCPAuth, MCPAuthType, MCPCredentials,
MCPTransport, MCPTransportType)
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIFileObject,
ResponsesAPIResponse,
)
from litellm.types.mcp import (
MCPAuthType,
MCPCredentials,
MCPTransport,
MCPTransportType,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.router import RouterErrors, UpdateRouterConfig
from litellm.types.secret_managers.main import KeyManagementSystem
from litellm.types.utils import (CallTypes, CostBreakdown, EmbeddingResponse,
GenericBudgetConfigType, ImageResponse,
LiteLLMBatch, LiteLLMFineTuningJob,
LiteLLMPydanticObjectBase, ModelResponse,
ProviderField, StandardCallbackDynamicParams,
StandardLoggingGuardrailInformation,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse)
from litellm.types.utils import (
CallTypes,
CostBreakdown,
EmbeddingResponse,
GenericBudgetConfigType,
ImageResponse,
LiteLLMBatch,
LiteLLMFineTuningJob,
LiteLLMPydanticObjectBase,
ModelResponse,
ProviderField,
StandardCallbackDynamicParams,
StandardLoggingGuardrailInformation,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse,
)
from litellm.types.videos.main import VideoObject
from .types_utils.utils import (get_instance_fn,
validate_custom_validate_return_type)
from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -2349,8 +2368,7 @@ class UserAPIKeyAuth(
This is used to track number of requests/spend for health check calls.
"""
from litellm.constants import \
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
return cls(
api_key=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
@ -2382,8 +2400,7 @@ class UserAPIKeyAuth(
This is used to track actions performed by automated system jobs.
"""
from litellm.constants import \
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
return cls(
api_key=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
@ -2774,8 +2791,7 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):
@model_validator(mode="after")
def mask_api_keys(self):
from litellm.litellm_core_utils.sensitive_data_masker import \
SensitiveDataMasker
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
masker = SensitiveDataMasker(sensitive_patterns={"key"})

View file

@ -143,7 +143,8 @@ def _safe_get_request_headers(request: Optional[Request]) -> dict:
"""
if request is None:
return {}
cached = getattr(request.state, "_cached_headers", None)
state = getattr(request, "state", None)
cached = getattr(state, "_cached_headers", None)
if cached is not None:
return cached
try:
@ -154,7 +155,8 @@ def _safe_get_request_headers(request: Optional[Request]) -> dict:
)
headers = {}
try:
request.state._cached_headers = headers
if state is not None:
state._cached_headers = headers
except Exception:
pass # request.state may not be available in all contexts
return headers

View file

@ -312,7 +312,7 @@ class DBSpendUpdateWriter:
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
litellm_proxy_budget_name: Optional[str],
payload_copy: dict,
payload_copy: SpendLogsPayload,
request_tags: Optional[Any],
):
"""

View file

@ -11,7 +11,6 @@ from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
AsyncGenerator,
Dict,
List,
Literal,
@ -26,6 +25,7 @@ from fastapi import HTTPException
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
log_guardrail_information,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -37,7 +37,6 @@ from litellm.types.utils import (
GenericGuardrailAPIInputs,
GuardrailStatus,
GuardrailTracingDetail,
ModelResponseStream,
)
if TYPE_CHECKING:
@ -538,6 +537,7 @@ class BlockCodeExecutionGuardrail(CustomGuardrail):
detection_info={"language": language},
)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -69,10 +69,12 @@ always_block_keywords:
severity: "high"
exceptions:
- "explainability"
- "improve explainability"
- "add explainability"
- "interpretability"
- "model card"
- "audit trail"
- "with audit trail"
- "add audit trail"
- "explain what"
- "explain how"
- "what is"

View file

@ -364,7 +364,8 @@ async def new_user(
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- team_id: Optional[str] - [DEPRECATED PARAM] The team id of the user. Default is None.
- agent_id: Optional[str] - The agent id associated with the user.
- team_id: Optional[str] - [DEPRECATED PARAM] The team id of the user. Default is None.
- duration: Optional[str] - Duration for the key auto-created on `/user/new`. Default is None.
- key_alias: Optional[str] - Alias for the key auto-created on `/user/new`. Default is None.
- sso_user_id: Optional[str] - The id of the user in the SSO provider.
@ -1075,7 +1076,8 @@ async def user_update(
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- team_id: Optional[str] - [DEPRECATED PARAM] The team id of the user. Default is None.
- agent_id: Optional[str] - The agent id associated with the user.
- team_id: Optional[str] - [DEPRECATED PARAM] The team id of the user. Default is None.
- duration: Optional[str] - [NOT IMPLEMENTED].
- key_alias: Optional[str] - [NOT IMPLEMENTED].
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - internal user-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.

View file

@ -1070,6 +1070,7 @@ async def generate_key_fn(
- key: Optional[str] - User defined key value. If not set, a 16-digit unique sk-key is created for you.
- team_id: Optional[str] - The team id of the key
- user_id: Optional[str] - The user id of the key
- agent_id: Optional[str] - The agent id associated with the key.
- organization_id: Optional[str] - The organization id of the key. If not set, and team_id is set, the organization id will be the same as the team id. If conflict, an error will be raised.
- project_id: Optional[str] - The project id of the key. When set, models and max_budget are validated against the project's limits.
- budget_id: Optional[str] - The budget id associated with the key. Created by calling `/budget/new`.
@ -1749,6 +1750,7 @@ async def update_key_fn(
- key_alias: Optional[str] - User-friendly key alias
- user_id: Optional[str] - User ID associated with key
- team_id: Optional[str] - Team ID associated with key
- agent_id: Optional[str] - The agent id associated with the key.
- budget_id: Optional[str] - The budget id associated with the key. Created by calling `/budget/new`.
- models: Optional[list] - Model_name's a user is allowed to call
- tags: Optional[List[str]] - Tags for organizing keys (Enterprise only)

View file

@ -5,7 +5,9 @@ usage/spend data by querying the aggregated daily activity endpoints.
import json
from datetime import date
from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional
from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional, cast
from typing_extensions import TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
@ -14,8 +16,6 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from typing_extensions import TypedDict
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
@ -492,17 +492,17 @@ async def _process_tool_call(
"tool_label": handler["label"],
"arguments": fn_args,
}
yield _sse({**tool_event_base, "status": "running"})
yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "running"}))
try:
tool_result = await _execute_tool_call(
handler, fn_name, fn_args, user_id, is_admin
)
yield _sse({**tool_event_base, "status": "complete"})
yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "complete"}))
except Exception as e:
verbose_proxy_logger.error("Tool %s failed: %s", fn_name, e)
tool_result = f"Error fetching {handler['label']}. Please try again."
yield _sse({**tool_event_base, "status": "error"})
yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "error"}))
chat_messages.append(
{"role": "tool", "tool_call_id": tc.id, "content": tool_result}

View file

@ -42,7 +42,7 @@ def _build_litellm_metadata(kwargs: dict) -> dict:
@wrapper_client
async def _arealtime(
async def _arealtime( # noqa: PLR0915
model: str,
websocket: Any, # fastapi websocket
api_base: Optional[str] = None,

View file

@ -7053,7 +7053,7 @@ class Router:
user_model_info = deployment.get("model_info") or {}
if model_info is not None:
model_info.update(user_model_info)
model_info.update(cast(ModelInfo, user_model_info))
return model_info

View file

@ -3,7 +3,7 @@ from dataclasses import dataclass
from enum import Enum
from typing import Any, Dict, List, Literal, Optional, Tuple
from pydantic import BaseModel, Field
from pydantic import BaseModel, Field, field_validator
from typing_extensions import Annotated
import litellm
@ -721,6 +721,13 @@ class UserAPIKeyLabelValues(BaseModel):
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value)
] = None
@field_validator("stream", mode="before")
@classmethod
def coerce_stream_to_str(cls, v: Any) -> Optional[str]:
if v is None:
return None
return str(v)
class PrometheusMetricsConfig(BaseModel):
"""Configuration for filtering Prometheus metrics"""

View file

@ -105,6 +105,7 @@ google-cloud-aiplatform: >=1.47.0 # Unknown license
mcp: >=1.5.0 # Unknown license
google-generativeai: >=0.5.0 # Unknown license
async_generator: >=1.10.0 # Unknown license
wheel: >=0.40.0 # MIT License - https://github.com/pypa/wheel/blob/main/LICENSE.txt
langfuse: >=2.45.0 # Unknown license
prometheus_client: >=0.20.0 # Unknown license
ddtrace: >=2.19.0 # Unknown license

View file

@ -12,6 +12,8 @@ from litellm.proxy.guardrails.guardrail_hooks.presidio import _OPTIONAL_Presidio
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import StandardLoggingPayload, StandardLoggingGuardrailInformation
from litellm.types.guardrails import GuardrailEventHooks
from litellm.proxy._types import UserAPIKeyAuth
from litellm.caching.caching import DualCache
from typing import Optional
@ -64,9 +66,13 @@ async def test_standard_logging_payload_includes_guardrail_information():
# Create mock response objects
mock_analyze_resp = MagicMock()
mock_analyze_resp.status = 200
mock_analyze_resp.content_type = "application/json"
mock_analyze_resp.json = AsyncMock(return_value=mock_analyze_response)
mock_anonymize_resp = MagicMock()
mock_anonymize_resp.status = 200
mock_anonymize_resp.content_type = "application/json"
mock_anonymize_resp.json = AsyncMock(return_value=mock_anonymize_response)
# Mock the aiohttp ClientSession with global call tracking
@ -85,7 +91,7 @@ async def test_standard_logging_payload_includes_guardrail_information():
async def close(self):
self.closed = True
def post(self, url, json=None):
def post(self, url, json=None, **kwargs):
class MockResponse:
def __init__(self, response_obj):
self.response_obj = response_obj
@ -116,8 +122,8 @@ async def test_standard_logging_payload_includes_guardrail_information():
with patch("aiohttp.ClientSession", MockClientSession):
await presidio_guard.async_pre_call_hook(
user_api_key_dict={},
cache=None,
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=request_data,
call_type="acompletion"
)
@ -136,11 +142,11 @@ async def test_standard_logging_payload_includes_guardrail_information():
assert len(test_custom_logger.standard_logging_payload["guardrail_information"]) > 0
guardrail_info = test_custom_logger.standard_logging_payload["guardrail_information"][0]
assert guardrail_info["guardrail_name"] == "presidio_guard"
assert guardrail_info["guardrail_mode"] == GuardrailEventHooks.pre_call
assert guardrail_info.get("guardrail_name") == "presidio_guard"
assert guardrail_info.get("guardrail_mode") == GuardrailEventHooks.pre_call
# assert that the guardrail_response is a response from presidio analyze
presidio_response = guardrail_info["guardrail_response"]
presidio_response = guardrail_info.get("guardrail_response")
assert isinstance(presidio_response, list)
for response_item in presidio_response:
assert "analysis_explanation" in response_item
@ -150,12 +156,14 @@ async def test_standard_logging_payload_includes_guardrail_information():
assert "entity_type" in response_item
# assert that the duration is not None
assert guardrail_info["duration"] is not None
assert guardrail_info["duration"] > 0
duration = guardrail_info.get("duration")
assert duration is not None
assert duration > 0
# assert that we get the count of masked entities
assert guardrail_info["masked_entity_count"] is not None
assert guardrail_info["masked_entity_count"]["PHONE_NUMBER"] == 1
masked_entity_count = guardrail_info.get("masked_entity_count")
assert masked_entity_count is not None
assert masked_entity_count["PHONE_NUMBER"] == 1
@ -201,8 +209,8 @@ async def test_langfuse_trace_includes_guardrail_information():
"metadata": {},
}
await presidio_guard.async_pre_call_hook(
user_api_key_dict={},
cache=None,
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=request_data,
call_type="acompletion"
)
@ -310,7 +318,7 @@ async def test_bedrock_guardrail_status_blocked():
try:
await bedrock_guard.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=None,
cache=DualCache(),
data=request_data,
call_type="completion"
)
@ -331,8 +339,8 @@ async def test_bedrock_guardrail_status_blocked():
# Verify guardrail information fields (guardrail_information is now a list)
guardrail_info = test_custom_logger.standard_logging_payload["guardrail_information"][0]
assert guardrail_info["guardrail_status"] == "guardrail_intervened"
assert guardrail_info["guardrail_provider"] == "bedrock"
assert guardrail_info.get("guardrail_status") == "guardrail_intervened"
assert guardrail_info.get("guardrail_provider") == "bedrock"
# Verify the new typed status fields
# guardrail_status should be "guardrail_intervened" when content is blocked
@ -395,7 +403,7 @@ async def test_bedrock_guardrail_status_success():
with patch.object(bedrock_guard, 'should_run_guardrail', return_value=True):
await bedrock_guard.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=None,
cache=DualCache(),
data=request_data,
call_type="completion"
)
@ -411,8 +419,8 @@ async def test_bedrock_guardrail_status_success():
assert len(test_custom_logger.standard_logging_payload["guardrail_information"]) > 0
guardrail_info = test_custom_logger.standard_logging_payload["guardrail_information"][0]
assert guardrail_info["guardrail_status"] == "success"
assert guardrail_info["guardrail_provider"] == "bedrock"
assert guardrail_info.get("guardrail_status") == "success"
assert guardrail_info.get("guardrail_provider") == "bedrock"
# Check status fields
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
@ -469,7 +477,7 @@ async def test_bedrock_guardrail_status_failure():
try:
await bedrock_guard.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=None,
cache=DualCache(),
data=request_data,
call_type="completion"
)
@ -488,8 +496,8 @@ async def test_bedrock_guardrail_status_failure():
assert len(test_custom_logger.standard_logging_payload["guardrail_information"]) > 0
guardrail_info = test_custom_logger.standard_logging_payload["guardrail_information"][0]
assert guardrail_info["guardrail_status"] == "guardrail_failed_to_respond"
assert guardrail_info["guardrail_provider"] == "bedrock"
assert guardrail_info.get("guardrail_status") == "guardrail_failed_to_respond"
assert guardrail_info.get("guardrail_provider") == "bedrock"
# Check status fields
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
@ -554,7 +562,7 @@ async def test_noma_guardrail_status_blocked():
try:
await noma_guard.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=None,
cache=DualCache(),
data=request_data,
call_type="completion"
)
@ -572,8 +580,8 @@ async def test_noma_guardrail_status_blocked():
assert len(test_custom_logger.standard_logging_payload["guardrail_information"]) > 0
guardrail_info = test_custom_logger.standard_logging_payload["guardrail_information"][0]
assert guardrail_info["guardrail_status"] == "guardrail_intervened"
assert guardrail_info["guardrail_provider"] == "noma"
assert guardrail_info.get("guardrail_status") == "guardrail_intervened"
assert guardrail_info.get("guardrail_provider") == "noma"
# Check status fields
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
@ -632,7 +640,7 @@ async def test_noma_guardrail_status_success():
with patch.object(noma_guard, 'should_run_guardrail', return_value=True):
await noma_guard.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=None,
cache=DualCache(),
data=request_data,
call_type="completion"
)
@ -648,8 +656,8 @@ async def test_noma_guardrail_status_success():
assert len(test_custom_logger.standard_logging_payload["guardrail_information"]) > 0
guardrail_info = test_custom_logger.standard_logging_payload["guardrail_information"][0]
assert guardrail_info["guardrail_status"] == "success"
assert guardrail_info["guardrail_provider"] == "noma"
assert guardrail_info.get("guardrail_status") == "success"
assert guardrail_info.get("guardrail_provider") == "noma"
# Check status fields
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
@ -679,8 +687,8 @@ def test_guardrail_status_fields_computation():
guardrail_information=intervened_info,
error_str=None
)
assert status_fields_intervened["llm_api_status"] == "success"
assert status_fields_intervened["guardrail_status"] == "guardrail_intervened"
assert status_fields_intervened.get("llm_api_status") == "success"
assert status_fields_intervened.get("guardrail_status") == "guardrail_intervened"
# Test legacy blocked status (for backward compatibility)
blocked_info = [{"guardrail_status": "blocked"}]
@ -689,8 +697,8 @@ def test_guardrail_status_fields_computation():
guardrail_information=blocked_info,
error_str=None
)
assert status_fields_blocked["llm_api_status"] == "success"
assert status_fields_blocked["guardrail_status"] == "guardrail_intervened"
assert status_fields_blocked.get("llm_api_status") == "success"
assert status_fields_blocked.get("guardrail_status") == "guardrail_intervened"
# Test success status
success_info = [{"guardrail_status": "success"}]
@ -699,8 +707,8 @@ def test_guardrail_status_fields_computation():
guardrail_information=success_info,
error_str=None
)
assert status_fields_success["llm_api_status"] == "success"
assert status_fields_success["guardrail_status"] == "success"
assert status_fields_success.get("llm_api_status") == "success"
assert status_fields_success.get("guardrail_status") == "success"
# Test guardrail_failed_to_respond status
failed_info = [{"guardrail_status": "guardrail_failed_to_respond"}]
@ -709,8 +717,8 @@ def test_guardrail_status_fields_computation():
guardrail_information=failed_info,
error_str=None
)
assert status_fields_failed["llm_api_status"] == "failure"
assert status_fields_failed["guardrail_status"] == "guardrail_failed_to_respond"
assert status_fields_failed.get("llm_api_status") == "failure"
assert status_fields_failed.get("guardrail_status") == "guardrail_failed_to_respond"
# Test legacy failure status (for backward compatibility)
failure_info = [{"guardrail_status": "failure"}]
@ -719,8 +727,8 @@ def test_guardrail_status_fields_computation():
guardrail_information=failure_info,
error_str=None
)
assert status_fields_failure["llm_api_status"] == "failure"
assert status_fields_failure["guardrail_status"] == "guardrail_failed_to_respond"
assert status_fields_failure.get("llm_api_status") == "failure"
assert status_fields_failure.get("guardrail_status") == "guardrail_failed_to_respond"
# Test no guardrail run
no_guardrail = None
@ -729,5 +737,5 @@ def test_guardrail_status_fields_computation():
guardrail_information=no_guardrail,
error_str=None
)
assert status_fields_no_guardrail["llm_api_status"] == "success"
assert status_fields_no_guardrail["guardrail_status"] == "not_run"
assert status_fields_no_guardrail.get("llm_api_status") == "success"
assert status_fields_no_guardrail.get("guardrail_status") == "not_run"

View file

@ -497,7 +497,7 @@ async def test_perform_health_check_filters_by_model_id():
captured_list = []
async def mock_perform_health_check(m_list, details=True):
async def mock_perform_health_check(m_list, details=True, **kwargs):
captured_list.append(m_list)
return [{"model": "gpt-4", "api_key": m_list[0]["litellm_params"]["api_key"]}], []

View file

@ -36,6 +36,7 @@ ignored_keys = [
"endTime",
"completionStartTime",
"endTime",
"request_duration_ms",
"metadata.model_map_information",
"metadata.usage_object",
"metadata.cold_storage_object_key",

View file

@ -68,6 +68,8 @@ def mock_request():
self.request_body = request_body or {}
# Add url attribute that the actual code expects
self.url = "http://localhost:8000/test"
# Add state attribute that FastAPI requests have
self.state = type("State", (), {})()
async def body(self) -> bytes:
return bytes(json.dumps(self.request_body), "utf-8")

View file

@ -75,6 +75,7 @@ async def test_form_data_parsing():
mock_request.form = AsyncMock(return_value=test_data)
mock_request.headers = {"content-type": "application/x-www-form-urlencoded"}
mock_request.scope = {}
mock_request.state._cached_headers = None
# Parse the form data
result = await _read_request_body(mock_request)
@ -123,6 +124,7 @@ async def test_form_data_with_json_metadata():
mock_request.form = AsyncMock(return_value=test_data)
mock_request.headers = {"content-type": "multipart/form-data"}
mock_request.scope = {}
mock_request.state._cached_headers = None
# Parse the form data
result = await _read_request_body(mock_request)
@ -163,6 +165,7 @@ async def test_form_data_with_invalid_json_metadata():
mock_request.form = AsyncMock(return_value=test_data)
mock_request.headers = {"content-type": "multipart/form-data"}
mock_request.scope = {}
mock_request.state._cached_headers = None
# Should raise JSONDecodeError when trying to parse invalid JSON metadata
with pytest.raises(json.JSONDecodeError):
@ -189,6 +192,7 @@ async def test_form_data_without_metadata():
mock_request.form = AsyncMock(return_value=test_data)
mock_request.headers = {"content-type": "application/x-www-form-urlencoded"}
mock_request.scope = {}
mock_request.state._cached_headers = None
# Parse the form data
result = await _read_request_body(mock_request)
@ -219,6 +223,7 @@ async def test_form_data_with_empty_metadata():
mock_request.form = AsyncMock(return_value=test_data)
mock_request.headers = {"content-type": "multipart/form-data"}
mock_request.scope = {}
mock_request.state._cached_headers = None
# Parse the form data
result = await _read_request_body(mock_request)
@ -256,6 +261,7 @@ async def test_form_data_with_dict_metadata():
mock_request.form = AsyncMock(return_value=test_data)
mock_request.headers = {"content-type": "multipart/form-data"}
mock_request.scope = {}
mock_request.state._cached_headers = None
# Parse the form data
result = await _read_request_body(mock_request)
@ -286,6 +292,7 @@ async def test_form_data_with_none_metadata():
mock_request.form = AsyncMock(return_value=test_data)
mock_request.headers = {"content-type": "multipart/form-data"}
mock_request.scope = {}
mock_request.state._cached_headers = None
# Parse the form data
result = await _read_request_body(mock_request)

View file

@ -13,6 +13,7 @@ sys.path.insert(
from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_endpoints import (
CreateGuardrailRequest,
PatchGuardrailRequest,
@ -25,6 +26,8 @@ from litellm.proxy.guardrails.guardrail_endpoints import (
patch_guardrail,
update_guardrail,
)
MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
from litellm.proxy.guardrails.guardrail_registry import (
IN_MEMORY_GUARDRAIL_HANDLER,
InMemoryGuardrailHandler,
@ -700,15 +703,15 @@ async def test_create_guardrail_endpoint(
# Run the test
if expected_exception:
with pytest.raises(expected_exception) as exc_info:
await create_guardrail(MOCK_CREATE_REQUEST)
await create_guardrail(MOCK_CREATE_REQUEST, user_api_key_dict=MOCK_ADMIN_USER)
if scenario == "database_failure":
assert "Database error" in str(exc_info.value.detail)
elif scenario == "no_prisma_client":
assert "Prisma client not initialized" in str(exc_info.value.detail)
else:
result = await create_guardrail(MOCK_CREATE_REQUEST)
result = await create_guardrail(MOCK_CREATE_REQUEST, user_api_key_dict=MOCK_ADMIN_USER)
assert result["guardrail_id"] == expected_result
assert result["guardrail_name"] == "Test DB Guardrail"
@ -789,15 +792,15 @@ async def test_update_guardrail_endpoint(
# Run the test
if expected_exception:
with pytest.raises(expected_exception) as exc_info:
await update_guardrail("test-guardrail-id", MOCK_UPDATE_REQUEST)
await update_guardrail("test-guardrail-id", MOCK_UPDATE_REQUEST, user_api_key_dict=MOCK_ADMIN_USER)
if scenario == "database_failure":
assert "Database error" in str(exc_info.value.detail)
elif scenario == "no_prisma_client":
assert "Prisma client not initialized" in str(exc_info.value.detail)
else:
result = await update_guardrail("test-guardrail-id", MOCK_UPDATE_REQUEST)
result = await update_guardrail("test-guardrail-id", MOCK_UPDATE_REQUEST, user_api_key_dict=MOCK_ADMIN_USER)
assert result["guardrail_id"] == expected_result
assert result["guardrail_name"] == "Test DB Guardrail"
@ -883,15 +886,15 @@ async def test_patch_guardrail_endpoint(
# Run the test
if expected_exception:
with pytest.raises(expected_exception) as exc_info:
await patch_guardrail("test-guardrail-id", MOCK_PATCH_REQUEST)
await patch_guardrail("test-guardrail-id", MOCK_PATCH_REQUEST, user_api_key_dict=MOCK_ADMIN_USER)
if scenario == "database_failure":
assert "Database error" in str(exc_info.value.detail)
elif scenario == "no_prisma_client":
assert "Prisma client not initialized" in str(exc_info.value.detail)
else:
result = await patch_guardrail("test-guardrail-id", MOCK_PATCH_REQUEST)
result = await patch_guardrail("test-guardrail-id", MOCK_PATCH_REQUEST, user_api_key_dict=MOCK_ADMIN_USER)
assert result["guardrail_id"] == expected_result
assert result["guardrail_name"] == "Test DB Guardrail"
@ -947,9 +950,9 @@ async def test_delete_guardrail_endpoint(
if expected_exception:
with pytest.raises(expected_exception):
await delete_guardrail(guardrail_id=expected_result)
await delete_guardrail(guardrail_id=expected_result, user_api_key_dict=MOCK_ADMIN_USER)
else:
result = await delete_guardrail(guardrail_id=expected_result)
result = await delete_guardrail(guardrail_id=expected_result, user_api_key_dict=MOCK_ADMIN_USER)
assert result == MOCK_DB_GUARDRAIL

View file

@ -6194,15 +6194,14 @@ async def test_generate_key_helper_fn_agent_id():
)
mock_prisma_client.insert_data = mock_insert
with patch.object(km, "prisma_client", mock_prisma_client):
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
await generate_key_helper_fn(
request_type="key",
agent_id="test-agent-456",
key_alias="test-agent-key",
models=[],
table_name="key",
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
await generate_key_helper_fn(
request_type="key",
agent_id="test-agent-456",
key_alias="test-agent-key",
models=[],
table_name="key",
)
assert mock_insert.called, "insert_data was never called"
# insert_data is called as insert_data(data=key_data, ...)

View file

@ -224,6 +224,7 @@ class TestVertexAIPassThroughHandler:
# Mock request
mock_request = Mock()
mock_request.state = None # Prevent Mock from returning a truthy _cached_headers
mock_request.method = "POST"
mock_request.headers = {
"Authorization": "Bearer test-creds",
@ -323,6 +324,7 @@ class TestVertexAIPassThroughHandler:
# Mock request
mock_request = Mock()
mock_request.state = None # Prevent Mock from returning a truthy _cached_headers
mock_request.method = "POST"
mock_request.headers = {
"Authorization": "Bearer test-creds",
@ -905,6 +907,7 @@ class TestVertexAIDiscoveryPassThroughHandler:
# Mock request
mock_request = Mock()
mock_request.state = None # Prevent Mock from returning a truthy _cached_headers
mock_request.method = "POST"
mock_request.headers = {
"Authorization": "Bearer test-key",
@ -1479,10 +1482,11 @@ class TestForwardHeaders:
# Create a mock request with custom headers
mock_request = MagicMock(spec=Request)
mock_request.state = None # Prevent MagicMock from returning a truthy _cached_headers
mock_request.method = "POST"
mock_request.url = MagicMock()
mock_request.url.path = "/test/endpoint"
# User headers that should be forwarded
user_headers = {
"x-custom-header": "custom-value",

View file

@ -253,6 +253,9 @@ async def test_vertex_passthrough_forwards_anthropic_beta_header():
"content-length": "1234", # Should be removed
"host": "localhost:4000", # Should be removed
})
# Prevent MagicMock from auto-creating a truthy _cached_headers attribute,
# which would short-circuit _safe_get_request_headers before reading .headers
mock_request.state._cached_headers = None
# Create mock vertex credentials
mock_vertex_credentials = MagicMock()

View file

@ -1650,58 +1650,64 @@ async def test_global_spend_keys_endpoint_limit_validation(client, monkeypatch):
# Create a simple mock for prisma client with empty response
mock_prisma_client = MagicMock()
mock_db = MagicMock()
mock_query_raw = MagicMock()
mock_query_raw.return_value = asyncio.Future()
mock_query_raw.return_value.set_result([])
mock_query_raw = AsyncMock(return_value=[])
mock_db.query_raw = mock_query_raw
mock_prisma_client.db = mock_db
# Apply the mock to the prisma_client module
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Call the endpoint without specifying a limit
no_limit_response = client.get("/global/spend/keys")
assert no_limit_response.status_code == 200
mock_query_raw.assert_called_once_with('SELECT * FROM "Last30dKeysBySpend";')
# Reset the mock for the next test
mock_query_raw.reset_mock()
# Test with valid input
normal_limit = "10"
good_input_response = client.get(f"/global/spend/keys?limit={normal_limit}")
assert good_input_response.status_code == 200
# Verify the mock was called with the correct parameters
mock_query_raw.assert_called_once_with(
'SELECT * FROM "Last30dKeysBySpend" LIMIT $1 ;', 10
# Override auth to bypass API key validation
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user"
)
# Reset the mock for the next test
mock_query_raw.reset_mock()
# Test with SQL injection payload
sql_injection_limit = "10; DROP TABLE spend_logs; --"
response = client.get(f"/global/spend/keys?limit={sql_injection_limit}")
# Verify the response is a validation error (422)
assert response.status_code == 422
# Verify the mock was not called with the SQL injection payload
# This confirms that the validation happens before the database query
mock_query_raw.assert_not_called()
# Reset the mock for the next test
mock_query_raw.reset_mock()
# Test with non-numeric input
non_numeric_limit = "abc"
response = client.get(f"/global/spend/keys?limit={non_numeric_limit}")
assert response.status_code == 422
mock_query_raw.assert_not_called()
mock_query_raw.reset_mock()
# Test with negative number
negative_limit = "-5"
response = client.get(f"/global/spend/keys?limit={negative_limit}")
assert response.status_code == 422
mock_query_raw.assert_not_called()
mock_query_raw.reset_mock()
# Test with zero
zero_limit = "0"
response = client.get(f"/global/spend/keys?limit={zero_limit}")
assert response.status_code == 422
mock_query_raw.assert_not_called()
mock_query_raw.reset_mock()
try:
# Call the endpoint without specifying a limit
no_limit_response = client.get("/global/spend/keys")
assert no_limit_response.status_code == 200
mock_query_raw.assert_called_once_with('SELECT * FROM "Last30dKeysBySpend";')
# Reset the mock for the next test
mock_query_raw.reset_mock()
# Test with valid input
normal_limit = "10"
good_input_response = client.get(f"/global/spend/keys?limit={normal_limit}")
assert good_input_response.status_code == 200
# Verify the mock was called with the correct parameters
mock_query_raw.assert_called_once_with(
'SELECT * FROM "Last30dKeysBySpend" LIMIT $1 ;', 10
)
# Reset the mock for the next test
mock_query_raw.reset_mock()
# Test with SQL injection payload
sql_injection_limit = "10; DROP TABLE spend_logs; --"
response = client.get(f"/global/spend/keys?limit={sql_injection_limit}")
# Verify the response is a validation error (422)
assert response.status_code == 422
# Verify the mock was not called with the SQL injection payload
# This confirms that the validation happens before the database query
mock_query_raw.assert_not_called()
# Reset the mock for the next test
mock_query_raw.reset_mock()
# Test with non-numeric input
non_numeric_limit = "abc"
response = client.get(f"/global/spend/keys?limit={non_numeric_limit}")
assert response.status_code == 422
mock_query_raw.assert_not_called()
mock_query_raw.reset_mock()
# Test with negative number
negative_limit = "-5"
response = client.get(f"/global/spend/keys?limit={negative_limit}")
assert response.status_code == 422
mock_query_raw.assert_not_called()
mock_query_raw.reset_mock()
# Test with zero
zero_limit = "0"
response = client.get(f"/global/spend/keys?limit={zero_limit}")
assert response.status_code == 422
mock_query_raw.assert_not_called()
mock_query_raw.reset_mock()
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio

View file

@ -480,7 +480,7 @@ async def test_perform_health_check_and_save_passes_model_id_to_perform_health_c
healthy = [{"model": "gpt-4"}]
unhealthy = []
async def mock_perform_health_check(model_list, model=None, cli_model=None, details=True, model_id=None):
async def mock_perform_health_check(model_list, model=None, cli_model=None, details=True, model_id=None, max_concurrency=None):
return healthy, unhealthy
with patch(

View file

@ -242,7 +242,7 @@ def test_full_output_structure_non_streaming():
)
result = model_dump_with_preserved_fields(response, exclude_unset=True)
# Top-level keys
# Top-level keys (usage is None when not explicitly set and excluded by exclude_unset=True)
assert set(result.keys()) == {
"id",
"choices",
@ -250,7 +250,6 @@ def test_full_output_structure_non_streaming():
"model",
"object",
"system_fingerprint",
"usage",
}
assert result["object"] == "chat.completion"
assert result["model"] == "gpt-4.1"
@ -270,12 +269,6 @@ def test_full_output_structure_non_streaming():
assert msg["content"] == "Hello!"
assert msg["role"] == "assistant"
# Usage structure
usage = result["usage"]
assert "prompt_tokens" in usage
assert "completion_tokens" in usage
assert "total_tokens" in usage
def test_full_output_structure_tool_calls():
"""

View file

@ -331,7 +331,8 @@ class TestProxyInitializationHelpers:
@patch("uvicorn.run")
@patch("builtins.print")
def test_max_requests_before_restart_flag(self, mock_print, mock_uvicorn_run):
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
def test_max_requests_before_restart_flag(self, mock_setup_db, mock_print, mock_uvicorn_run):
"""Test that the max_requests_before_restart flag is passed to uvicorn as limit_max_requests"""
from click.testing import CliRunner
@ -344,7 +345,10 @@ class TestProxyInitializationHelpers:
mock_key_mgmt = MagicMock()
mock_save_worker_config = MagicMock()
clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")}
with patch.dict(
os.environ, clean_env, clear=True,
), patch.dict(
"sys.modules",
{
"proxy_server": MagicMock(
@ -367,7 +371,7 @@ class TestProxyInitializationHelpers:
run_server, ["--local", "--max_requests_before_restart", "123"]
)
assert result.exit_code == 0
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
mock_uvicorn_run.assert_called_once()
# Check that uvicorn.run was called with limit_max_requests parameter

View file

@ -1,10 +1,13 @@
import asyncio
import json
import pytest
import time
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy.health_check_utils.shared_health_check_manager import SharedHealthCheckManager
import pytest
from litellm.proxy.health_check_utils.shared_health_check_manager import (
SharedHealthCheckManager,
)
class TestSharedHealthCheckManager:
@ -272,7 +275,7 @@ class TestSharedHealthCheckManager:
)
# Should call perform_health_check and cache results
mock_perform.assert_called_once_with(model_list=model_list, details=True)
mock_perform.assert_called_once_with(model_list=model_list, details=True, max_concurrency=None)
assert healthy == expected_healthy
assert unhealthy == expected_unhealthy
@ -329,7 +332,7 @@ class TestSharedHealthCheckManager:
# Should fall back to local health check
mock_sleep.assert_called_once_with(2)
mock_perform.assert_called_once_with(model_list=model_list, details=True)
mock_perform.assert_called_once_with(model_list=model_list, details=True, max_concurrency=None)
assert healthy == expected_healthy
assert unhealthy == expected_unhealthy

View file

@ -3,6 +3,7 @@ from unittest.mock import Mock
import pytest
from litellm import Router
from litellm.router_utils.common_utils import (
_deployment_supports_web_search,
filter_team_based_models,
@ -340,3 +341,22 @@ class TestFilterWebSearchDeployments:
result = filter_web_search_deployments(deployment, request_kwargs)
# Should return the dict unchanged, not filter it
assert result == deployment
def test_invalidate_model_group_info_cache():
"""Test that _invalidate_model_group_info_cache clears the LRU cache."""
router = Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake-key"},
}
]
)
# Populate the cache
router._cached_get_model_group_info("gpt-4")
assert router._cached_get_model_group_info.cache_info().currsize > 0
# Invalidate and verify cache is cleared
router._invalidate_model_group_info_cache()
assert router._cached_get_model_group_info.cache_info().currsize == 0