mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #22151 from BerriAI/litellm_fix_cicd_26_02
[Fix] CICD 26/02/26
This commit is contained in:
commit
4d68151d03
32 changed files with 286 additions and 193 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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"]}], []
|
||||
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ ignored_keys = [
|
|||
"endTime",
|
||||
"completionStartTime",
|
||||
"endTime",
|
||||
"request_duration_ms",
|
||||
"metadata.model_map_information",
|
||||
"metadata.usage_object",
|
||||
"metadata.cold_storage_object_key",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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, ...)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue