diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 8dbebad884e..b694549cf40 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -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 diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 6ba1b48c647..8df41aea4a3 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -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( diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py index 2abcc679eef..0260eeafe63 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py @@ -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 diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py index 12333623f51..6eddcccd631 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 860569d24cb..6e78458cc0e 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 88cfa5f6c11..de1609baf62 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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"}) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 8d179a9caed..04d46ecaeb8 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -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 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 37fac8b56ed..0c25424ceaa 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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], ): """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py index 77f2eaa27d7..e76a02a6e4d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py +++ b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/policy_templates/sg_mas_transparency_explainability.yaml b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/policy_templates/sg_mas_transparency_explainability.yaml index 0a08abd86e8..4b3555b49df 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/policy_templates/sg_mas_transparency_explainability.yaml +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/policy_templates/sg_mas_transparency_explainability.yaml @@ -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" diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 36fc65de810..e535ccaaa46 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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. diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index c1165ab26d0..b489369071f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py index f156be7d2cc..4de29e04092 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -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} diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index df49d4c54b2..ac597fc623d 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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, diff --git a/litellm/router.py b/litellm/router.py index 3a6c514989d..cbe5b414040 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 482b87085dd..2c75276d9ca 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -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""" diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini index e6e9d761ad5..376d2859ffa 100644 --- a/tests/code_coverage_tests/liccheck.ini +++ b/tests/code_coverage_tests/liccheck.ini @@ -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 diff --git a/tests/guardrails_tests/test_tracing_guardrails.py b/tests/guardrails_tests/test_tracing_guardrails.py index 8e7ce27bc28..f119d6df3db 100644 --- a/tests/guardrails_tests/test_tracing_guardrails.py +++ b/tests/guardrails_tests/test_tracing_guardrails.py @@ -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" \ No newline at end of file + assert status_fields_no_guardrail.get("llm_api_status") == "success" + assert status_fields_no_guardrail.get("guardrail_status") == "not_run" \ No newline at end of file diff --git a/tests/litellm_utils_tests/test_health_check.py b/tests/litellm_utils_tests/test_health_check.py index b459c3cfc99..963fe4f5b9f 100644 --- a/tests/litellm_utils_tests/test_health_check.py +++ b/tests/litellm_utils_tests/test_health_check.py @@ -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"]}], [] diff --git a/tests/logging_callback_tests/test_gcs_pub_sub.py b/tests/logging_callback_tests/test_gcs_pub_sub.py index 8ffbc8eedd5..540fb59ab01 100644 --- a/tests/logging_callback_tests/test_gcs_pub_sub.py +++ b/tests/logging_callback_tests/test_gcs_pub_sub.py @@ -36,6 +36,7 @@ ignored_keys = [ "endTime", "completionStartTime", "endTime", + "request_duration_ms", "metadata.model_map_information", "metadata.usage_object", "metadata.cold_storage_object_key", diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py index 0e681cb1e02..f529fe85ab8 100644 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py @@ -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") diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index d236f46e5c5..a1484bc263b 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -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) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index c0f16c8b953..0ac3637b380 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 05df3c2dcbb..fb71adc1085 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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, ...) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 98161402c45..fdc821ef8bb 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -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", diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index d2fdb157c8d..eb4749549c2 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -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() diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index e439dfd693c..2aecc2ec2e5 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index 4c91d0ae91e..354698b02fe 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -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( diff --git a/tests/test_litellm/proxy/test_model_dump_with_preserved_fields.py b/tests/test_litellm/proxy/test_model_dump_with_preserved_fields.py index 3001c87ebed..316c1e879cc 100644 --- a/tests/test_litellm/proxy/test_model_dump_with_preserved_fields.py +++ b/tests/test_litellm/proxy/test_model_dump_with_preserved_fields.py @@ -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(): """ diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index c839c22de5f..c6b2015984e 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -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 diff --git a/tests/test_litellm/proxy/test_shared_health_check.py b/tests/test_litellm/proxy/test_shared_health_check.py index 82deebc424a..0212d87baab 100644 --- a/tests/test_litellm/proxy/test_shared_health_check.py +++ b/tests/test_litellm/proxy/test_shared_health_check.py @@ -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 diff --git a/tests/test_litellm/router_utils/test_router_utils_common_utils.py b/tests/test_litellm/router_utils/test_router_utils_common_utils.py index 8ff1ba45cc2..587b6a97b56 100644 --- a/tests/test_litellm/router_utils/test_router_utils_common_utils.py +++ b/tests/test_litellm/router_utils/test_router_utils_common_utils.py @@ -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