diff --git a/docs/my-website/docs/guides/security_settings.md b/docs/my-website/docs/guides/security_settings.md index 7995f6c3c9c..d6397a7c197 100644 --- a/docs/my-website/docs/guides/security_settings.md +++ b/docs/my-website/docs/guides/security_settings.md @@ -117,10 +117,52 @@ litellm_settings: ```bash export SSL_CERTIFICATE="/path/to/certificate.pem" ``` + -## 5. Use HTTP_PROXY environment variable +## 5. Configure ECDH Curve for SSL/TLS Performance + +The `ssl_ecdh_curve` setting allows you to configure the Elliptic Curve Diffie-Hellman (ECDH) curve used for SSL/TLS key exchange. This is particularly useful for disabling Post-Quantum Cryptography (PQC) to improve performance in environments where PQC is not required. + +**Use Case:** Some OpenSSL 3.x systems enable PQC by default, which can slow down TLS handshakes. Setting the ECDH curve to `X25519` disables PQC and can significantly improve connection performance. + + + + +```python +import litellm +litellm.ssl_ecdh_curve = "X25519" # Disables PQC for better performance +``` + + + + +```yaml +litellm_settings: + ssl_ecdh_curve: "X25519" +``` + + + + +```bash +export SSL_ECDH_CURVE="X25519" +``` + + + + +**Common Valid Curves:** + +- `X25519` - Modern, fast curve (recommended for disabling PQC) +- `prime256v1` - NIST P-256 curve +- `secp384r1` - NIST P-384 curve +- `secp521r1` - NIST P-521 curve + +**Note:** If an invalid curve name is provided or if your Python/OpenSSL version doesn't support this feature, LiteLLM will log a warning and continue with default curves. + +## 6. Use HTTP_PROXY environment variable Both httpx and aiohttp libraries use `urllib.request.getproxies` from environment variables. Before client initialization, you may set proxy (and optional SSL_CERT_FILE) by setting the environment variables: diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index f832526ecda..0dbecad9d4a 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -753,6 +753,7 @@ router_settings: | SPEND_LOGS_URL | URL for retrieving spend logs | SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000 | SSL_CERTIFICATE | Path to the SSL certificate file +| SSL_ECDH_CURVE | ECDH curve for SSL/TLS key exchange (e.g., 'X25519' to disable PQC). | SSL_SECURITY_LEVEL | [BETA] Security level for SSL/TLS connections. E.g. `DEFAULT@SECLEVEL=1` | SSL_VERIFY | Flag to enable or disable SSL certificate verification | SSL_CERT_FILE | Path to the SSL certificate file for custom CA bundle diff --git a/docs/my-website/docs/proxy/guardrails/pillar_security.md b/docs/my-website/docs/proxy/guardrails/pillar_security.md index c730da5b416..a5a416839f6 100644 --- a/docs/my-website/docs/proxy/guardrails/pillar_security.md +++ b/docs/my-website/docs/proxy/guardrails/pillar_security.md @@ -38,13 +38,17 @@ model_list: api_key: os.environ/OPENAI_API_KEY guardrails: - - guardrail_name: "pillar-minitor-everything" # you can change my name + - guardrail_name: "pillar-monitor-everything" # you can change my name litellm_params: guardrail: pillar mode: [pre_call, post_call] # Monitor both input and output api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "monitor" # Log threats but allow requests + persist_session: true # Keep conversations visible in Pillar dashboard + async_mode: false # Request synchronous verdicts + include_scanners: true # Return scanner category breakdown + include_evidence: true # Include detailed findings for triage default_on: true # Enable for all requests general_settings: @@ -104,10 +108,14 @@ guardrails: api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "block" # Block malicious requests + persist_session: true # Keep records for investigation + async_mode: false # Require an immediate verdict + include_scanners: true # Understand which rule triggered + include_evidence: true # Capture concrete evidence default_on: true # Enable for all requests general_settings: - master_key: "your-master-key-here" + master_key: "YOUR_LITELLM_PROXY_MASTER_KEY" litellm_settings: set_verbose: true @@ -136,10 +144,14 @@ guardrails: api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "monitor" # Log threats but allow requests + persist_session: false # Skip dashboard storage for low latency + async_mode: false # Still receive results inline + include_scanners: false # Minimal payload for performance + include_evidence: false # Omit details to keep responses light default_on: true # Enable for all requests general_settings: - master_key: "your-secure-master-key-here" + master_key: "YOUR_LITELLM_PROXY_MASTER_KEY" litellm_settings: set_verbose: true # Enable detailed logging @@ -169,10 +181,14 @@ guardrails: api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "block" # Block threats on input and output + persist_session: true # Preserve conversations in Pillar dashboard + async_mode: false # Require synchronous approval + include_scanners: true # Inspect which scanners fired + include_evidence: true # Include detailed evidence for auditing default_on: true # Enable for all requests general_settings: - master_key: "your-secure-master-key-here" + master_key: "YOUR_LITELLM_PROXY_MASTER_KEY" litellm_settings: set_verbose: true # Enable detailed logging @@ -229,19 +245,139 @@ Logs the violation but allows the request to proceed: on_flagged_action: "monitor" ``` +## Advanced Configuration + +**Quick takeaways** +- Every request still runs *all* Pillar scanners; these options only change what comes back. +- Choose richer responses when you need audit trails, lighter responses when latency or cost matters. +- Blocking is controlled by LiteLLM’s `on_flagged_action` configuration—Pillar headers do not change block/monitor behaviour. + +Pillar Security executes the full scanner suite on each call. The settings below tune the Protect response headers LiteLLM sends, letting you balance fidelity, retention, and latency. + +### Response Control + +#### Data Retention (`persist_session`) +```yaml +persist_session: false # Default: true +``` +- **Why**: Controls whether Pillar stores session data for dashboard visibility. +- **Set false for**: Ephemeral testing, privacy-sensitive interactions. +- **Set true for**: Production monitoring, compliance, historical review (default behaviour). +- **Impact**: `false` means the conversation will *not* appear in the Pillar dashboard. + +#### Response Detail Level +The following toggles grow the payload size without changing detection behaviour. + +```yaml +include_scanners: true # → plr_scanners (default true in LiteLLM) +include_evidence: true # → plr_evidence (default true in LiteLLM) +``` + +- **Minimal response** (`include_scanners=false`, `include_evidence=false`) + ```json + { + "session_id": "abc-123", + "flagged": true + } + ``` + Use when you only care about whether Pillar detected a threat. + + > **📝 Note:** `flagged: true` means Pillar’s scanners recommend blocking. Pillar only reports this verdict—LiteLLM enforces your policy via the `on_flagged_action` configuration (no Pillar header controls it): + > - `on_flagged_action: "block"` → LiteLLM raises a 400 guardrail error + > - `on_flagged_action: "monitor"` → LiteLLM logs the threat but still returns the LLM response + +- **Scanner breakdown** (`include_scanners=true`) + ```json + { + "session_id": "abc-123", + "flagged": true, + "scanners": { + "jailbreak": true, + "prompt_injection": false, + "pii": false, + "secret": false, + "toxic_language": false + /* ... more categories ... */ + } + } + ``` + Use when you need to know which categories triggered. + +- **Full context** (both toggles true) + ```json + { + "session_id": "abc-123", + "flagged": true, + "scanners": { /* ... */ }, + "evidence": [ + { + "category": "jailbreak", + "type": "prompt_injection", + "evidence": "Ignore previous instructions", + "metadata": { "start_idx": 0, "end_idx": 28 } + } + ] + } + ``` + Ideal for debugging, audit logs, or compliance exports. + +### Processing Mode (`async_mode`) +```yaml +async_mode: true # Default: false +``` +- **Why**: Queue the request for background processing instead of waiting for a synchronous verdict. +- **Response shape**: + ```json + { + "status": "queued", + "session_id": "abc-123", + "position": 1 + } + ``` +- **Set true for**: Large batch jobs, latency-tolerant pipelines. +- **Set false for**: Real-time user flows (default). +- ⚠️ **Note**: Async mode returns only a 202 queue acknowledgment (no flagged verdict). LiteLLM treats that as “no block,” so the pre-call hook always allows the request. Use async mode only for post-call or monitor-only workflows where delayed review is acceptable. + +### Complete Examples + +```yaml +guardrails: + # Production: full fidelity & dashboard visibility + - guardrail_name: "pillar-production" + litellm_params: + guardrail: pillar + mode: [pre_call, post_call] + persist_session: true + include_scanners: true + include_evidence: true + on_flagged_action: "block" + + # Testing: lightweight, no persistence + - guardrail_name: "pillar-testing" + litellm_params: + guardrail: pillar + mode: pre_call + persist_session: false + include_scanners: false + include_evidence: false + on_flagged_action: "monitor" +``` + +Keep in mind that LiteLLM forwards these values as the documented `plr_*` headers, so any direct HTTP integrations outside the proxy can reuse the same guidance. + ## Examples -**Safe requset** +**Safe request** ```bash # Test with safe content curl -X POST "http://localhost:4000/v1/chat/completions" \ -H "Content-Type: application/json" \ - -H "Authorization: Bearer your-master-key-here" \ + -H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \ -d '{ "model": "gpt-4.1-mini", "messages": [{"role": "user", "content": "Hello! Can you tell me a joke?"}], @@ -300,7 +436,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \ ```bash curl -X POST "http://localhost:4000/v1/chat/completions" \ -H "Content-Type: application/json" \ - -H "Authorization: Bearer your-master-key-here" \ + -H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \ -d '{ "model": "gpt-4.1-mini", "messages": [ @@ -350,7 +486,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \ ```bash curl -X POST "http://localhost:4000/v1/chat/completions" \ -H "Content-Type: application/json" \ - -H "Authorization: Bearer your-master-key-here" \ + -H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \ -d '{ "model": "gpt-4.1-mini", "messages": [ @@ -405,4 +541,4 @@ Feel free to contact us at support@pillar.security - [Pillar Security API Docs](https://docs.pillar.security/docs/api/introduction) - [Pillar Security Dashboard](https://app.pillar.security) - [Pillar Security Website](https://pillar.security) -- [LiteLLM Docs](https://docs.litellm.ai) \ No newline at end of file +- [LiteLLM Docs](https://docs.litellm.ai) diff --git a/litellm/__init__.py b/litellm/__init__.py index 162f07c56ec..02bfaaf72f9 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -263,6 +263,7 @@ use_client: bool = False ssl_verify: Union[str, bool] = True ssl_security_level: Optional[str] = None ssl_certificate: Optional[str] = None +ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance disable_streaming_logging: bool = False disable_token_counter: bool = False disable_add_transform_inline_image_block: bool = False diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index a3ad2c67272..accdddbc4dd 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1,6 +1,7 @@ import asyncio import os import ssl +import sys import time from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Union @@ -114,6 +115,28 @@ def get_ssl_configuration( # but falls back to widely compatible ones custom_ssl_context.set_ciphers(DEFAULT_SSL_CIPHERS) + # Configure ECDH curve for key exchange (e.g., to disable PQC and improve performance) + # Set SSL_ECDH_CURVE env var or litellm.ssl_ecdh_curve to 'X25519' to disable PQC + # Common valid curves: X25519, prime256v1, secp384r1, secp521r1 + ssl_ecdh_curve = os.getenv("SSL_ECDH_CURVE", litellm.ssl_ecdh_curve) + if ssl_ecdh_curve and isinstance(ssl_ecdh_curve, str): + try: + custom_ssl_context.set_ecdh_curve(ssl_ecdh_curve) + verbose_logger.debug(f"SSL ECDH curve set to: {ssl_ecdh_curve}") + except AttributeError: + verbose_logger.warning( + f"SSL ECDH curve configuration not supported. " + f"Python version: {sys.version.split()[0]}, OpenSSL version: {ssl.OPENSSL_VERSION}. " + f"Requested curve: {ssl_ecdh_curve}. Continuing with default curves." + ) + except ValueError as e: + # Invalid curve name + verbose_logger.warning( + f"Invalid SSL ECDH curve name: '{ssl_ecdh_curve}'. {e}. " + f"Common valid curves: X25519, prime256v1, secp384r1, secp521r1. " + f"Continuing with default curves (including PQC)." + ) + # Use our custom SSL context instead of the original ssl_verify value return custom_ssl_context diff --git a/litellm/llms/openai/image_edit/__init__.py b/litellm/llms/openai/image_edit/__init__.py new file mode 100644 index 00000000000..c1898326b72 --- /dev/null +++ b/litellm/llms/openai/image_edit/__init__.py @@ -0,0 +1,26 @@ +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig + +from .dalle2_transformation import DallE2ImageEditConfig +from .transformation import OpenAIImageEditConfig + +__all__ = ["OpenAIImageEditConfig", "DallE2ImageEditConfig", "get_openai_image_edit_config"] + + +def get_openai_image_edit_config(model: str) -> BaseImageEditConfig: + """ + Get the appropriate OpenAI image edit config based on the model. + + Args: + model: The model name (e.g., "dall-e-2", "gpt-image-1") + + Returns: + The appropriate config instance for the model + """ + model_normalized = model.lower().replace("-", "").replace("_", "") + + if model_normalized == "dalle2": + return DallE2ImageEditConfig() + else: + # Default to standard OpenAI config for gpt-image-1 and other models + return OpenAIImageEditConfig() + diff --git a/litellm/llms/openai/image_edit/dalle2_transformation.py b/litellm/llms/openai/image_edit/dalle2_transformation.py new file mode 100644 index 00000000000..37e92be17a8 --- /dev/null +++ b/litellm/llms/openai/image_edit/dalle2_transformation.py @@ -0,0 +1,101 @@ +from io import BufferedReader +from typing import TYPE_CHECKING, Any, Dict, List, Tuple, cast + +from httpx._types import RequestFiles + +import litellm +from litellm.images.utils import ImageEditRequestUtils +from litellm.types.images.main import ImageEditRequestParams +from litellm.types.llms.openai import FileTypes +from litellm.types.router import GenericLiteLLMParams + +from .transformation import OpenAIImageEditConfig + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class DallE2ImageEditConfig(OpenAIImageEditConfig): + """ + DALL-E-2 specific configuration for image edit API. + + DALL-E-2 only supports editing a single image (not an array). + Uses "image" field name instead of "image[]". + """ + + def transform_image_edit_request( + self, + model: str, + prompt: str, + image: FileTypes, + image_edit_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[Dict, RequestFiles]: + """ + Transform image edit request for DALL-E-2. + + DALL-E-2 only accepts a single image with field name "image" (not "image[]"). + """ + request = ImageEditRequestParams( + model=model, + image=image, + prompt=prompt, + **image_edit_optional_request_params, + ) + request_dict = cast(Dict, request) + + ######################################################### + # Separate images and masks as `files` and send other parameters as `data` + ######################################################### + _image_list = request_dict.get("image") + _mask = request_dict.get("mask") + data_without_files = { + k: v for k, v in request_dict.items() if k not in ["image", "mask"] + } + files_list: List[Tuple[str, Any]] = [] + + # Handle image parameter - DALL-E-2 only supports single image + if _image_list is not None: + image_list = ( + [_image_list] if not isinstance(_image_list, list) else _image_list + ) + + # Validate only one image is provided + if len(image_list) > 1: + raise litellm.BadRequestError( + message="DALL-E-2 only supports editing a single image. Please provide one image.", + model=model, + llm_provider="openai", + ) + + # Use "image" field name (singular) for DALL-E-2 + for _image in image_list: + if _image is not None: + self._add_image_to_files( + files_list=files_list, + image=_image, + field_name="image", + ) + + # Handle mask parameter if provided + if _mask is not None: + # Handle case where mask can be a list (extract first mask) + if isinstance(_mask, list): + _mask = _mask[0] if _mask else None + + if _mask is not None: + mask_content_type: str = ImageEditRequestUtils.get_image_content_type( + _mask + ) + if isinstance(_mask, BufferedReader): + files_list.append(("mask", (_mask.name, _mask, mask_content_type))) + else: + files_list.append(("mask", ("mask.png", _mask, mask_content_type))) + + return data_without_files, files_list + diff --git a/litellm/llms/openai/image_edit/transformation.py b/litellm/llms/openai/image_edit/transformation.py index be960641154..1b90d96fa92 100644 --- a/litellm/llms/openai/image_edit/transformation.py +++ b/litellm/llms/openai/image_edit/transformation.py @@ -27,6 +27,11 @@ else: class OpenAIImageEditConfig(BaseImageEditConfig): + """ + Base configuration for OpenAI image edit API. + Used for models like gpt-image-1 that support multiple images. + """ + def get_supported_openai_params(self, model: str) -> list: """ All OpenAI Image Edits params are supported @@ -57,6 +62,20 @@ class OpenAIImageEditConfig(BaseImageEditConfig): """No mapping applied since inputs are in OpenAI spec already""" return dict(image_edit_optional_params) + def _add_image_to_files( + self, + files_list: List[Tuple[str, Any]], + image: Any, + field_name: str, + ) -> None: + """Add an image to the files list with appropriate content type""" + image_content_type = ImageEditRequestUtils.get_image_content_type(image) + + if isinstance(image, BufferedReader): + files_list.append((field_name, (image.name, image, image_content_type))) + else: + files_list.append((field_name, ("image.png", image, image_content_type))) + def transform_image_edit_request( self, model: str, @@ -67,9 +86,10 @@ class OpenAIImageEditConfig(BaseImageEditConfig): headers: dict, ) -> Tuple[Dict, RequestFiles]: """ - No transform applied since inputs are in OpenAI spec already + Transform image edit request to OpenAI API format. - This handles buffered readers as images to be sent as multipart/form-data for OpenAI + Handles multipart/form-data for images. Uses "image[]" field name + to support multiple images (e.g., for gpt-image-1). """ request = ImageEditRequestParams( model=model, @@ -94,19 +114,14 @@ class OpenAIImageEditConfig(BaseImageEditConfig): image_list = ( [_image_list] if not isinstance(_image_list, list) else _image_list ) + for _image in image_list: if _image is not None: - image_content_type: str = ( - ImageEditRequestUtils.get_image_content_type(_image) + self._add_image_to_files( + files_list=files_list, + image=_image, + field_name="image[]", ) - if isinstance(_image, BufferedReader): - files_list.append( - ("image[]", (_image.name, _image, image_content_type)) - ) - else: - files_list.append( - ("image[]", ("image.png", _image, image_content_type)) - ) # Handle mask parameter if provided if _mask is not None: # Handle case where mask can be a list (extract first mask) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index caca89d23f8..8cab321c4dd 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -17521,6 +17521,25 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "openrouter/anthropic/claude-sonnet-4.5": { + "input_cost_per_image": 0.0048, + "input_cost_per_token": 3e-06, + "input_cost_per_token_above_200k_tokens": 6e-06, + "output_cost_per_token_above_200k_tokens": 2.25e-05, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 159 + }, "openrouter/bytedance/ui-tars-1.5-7b": { "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", @@ -21896,8 +21915,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, @@ -21913,8 +21932,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 33d432e8695..807b895bd03 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -1,6 +1,6 @@ import json import re -from typing import Any, Dict, List, Optional +from typing import Any, Collection, Dict, List, Optional import orjson from fastapi import Request, UploadFile, status @@ -149,7 +149,7 @@ def _safe_get_request_headers(request: Optional[Request]) -> dict: def check_file_size_under_limit( request_data: dict, file: UploadFile, - router_model_names: List[str], + router_model_names: Collection[str], ) -> bool: """ Check if any files passed in request are under max_file_size_mb diff --git a/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py index 556c22b9495..29ede085ed6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py @@ -23,6 +23,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" if not guardrail_name: raise ValueError("Pillar guardrail name is required") + optional_params = getattr(litellm_params, "optional_params", None) + _pillar_callback = PillarGuardrail( guardrail_name=guardrail_name, api_key=litellm_params.api_key, @@ -30,12 +32,34 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" on_flagged_action=getattr(litellm_params, "on_flagged_action", "monitor"), event_hook=litellm_params.mode, default_on=litellm_params.default_on, + async_mode=_get_config_value( + litellm_params, optional_params, "async_mode" + ), + persist_session=_get_config_value( + litellm_params, optional_params, "persist_session" + ), + include_scanners=_get_config_value( + litellm_params, optional_params, "include_scanners" + ), + include_evidence=_get_config_value( + litellm_params, optional_params, "include_evidence" + ), ) litellm.logging_callback_manager.add_litellm_callback(_pillar_callback) return _pillar_callback +def _get_config_value(litellm_params, optional_params, attribute_name): + """Return guardrail configuration value prioritising optional params when present.""" + + if optional_params is not None: + value = getattr(optional_params, attribute_name, None) + if value is not None: + return value + return getattr(litellm_params, attribute_name, None) + + guardrail_initializer_registry = { SupportedGuardrailIntegrations.PILLAR.value: initialize_guardrail, } diff --git a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py index f4741aa8e00..e19125ed031 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py @@ -69,6 +69,10 @@ class PillarGuardrail(CustomGuardrail): api_key: Optional[str] = None, api_base: Optional[str] = None, on_flagged_action: Optional[str] = None, + async_mode: Optional[bool] = None, + persist_session: Optional[bool] = None, + include_scanners: Optional[bool] = None, + include_evidence: Optional[bool] = None, **kwargs, ) -> None: """ @@ -110,6 +114,31 @@ class PillarGuardrail(CustomGuardrail): f"Pillar Guardrail: Initialized with on_flagged_action: {self.on_flagged_action}" ) + self.async_mode = self._resolve_bool_config( + provided_value=async_mode, + env_var="PILLAR_ASYNC", + default=None, + setting_name="async_mode", + ) + self.persist_session = self._resolve_bool_config( + provided_value=persist_session, + env_var="PILLAR_PERSIST", + default=None, + setting_name="persist_session", + ) + self.include_scanners = self._resolve_bool_config( + provided_value=include_scanners, + env_var="PILLAR_INCLUDE_SCANNERS", + default=True, + setting_name="include_scanners", + ) + self.include_evidence = self._resolve_bool_config( + provided_value=include_evidence, + env_var="PILLAR_INCLUDE_EVIDENCE", + default=True, + setting_name="include_evidence", + ) + # Define supported event hooks supported_event_hooks = [ GuardrailEventHooks.pre_call, @@ -347,12 +376,74 @@ class PillarGuardrail(CustomGuardrail): "Content-Type": "application/json", } - # Add Pillar-specific headers for enhanced response data - headers["plr_evidence"] = "true" - headers["plr_scanners"] = "true" + # Add Pillar-specific headers based on configuration + self._set_bool_header(headers, "plr_scanners", self.include_scanners) + self._set_bool_header(headers, "plr_evidence", self.include_evidence) + self._set_bool_header(headers, "plr_async", self.async_mode) + self._set_bool_header(headers, "plr_persist", self.persist_session) return headers + def _set_bool_header( + self, headers: Dict[str, str], header_name: str, value: Optional[bool] + ) -> None: + """Apply a boolean value as a lowercase string HTTP header when provided.""" + + if value is None: + return + headers[header_name] = "true" if value else "false" + + def _resolve_bool_config( + self, + provided_value: Optional[Union[bool, str, int]], + env_var: Optional[str], + default: Optional[bool], + setting_name: str, + ) -> Optional[bool]: + """Resolve configuration precedence: explicit value -> environment -> default.""" + + if provided_value is not None: + try: + return self._parse_bool_value(provided_value) + except ValueError: + verbose_proxy_logger.warning( + "Pillar Guardrail: Invalid boolean value '%s' for %s, falling back to default.", + provided_value, + setting_name, + ) + return default + + if env_var: + env_value = os.getenv(env_var) + if env_value is not None: + try: + return self._parse_bool_value(env_value) + except ValueError: + verbose_proxy_logger.warning( + "Pillar Guardrail: Invalid boolean env value '%s' for %s, falling back to default.", + env_value, + env_var, + ) + return default + + return default + + @staticmethod + def _parse_bool_value(value: Union[bool, str, int]) -> bool: + """Normalise various truthy/falsey inputs to a strict boolean.""" + + if isinstance(value, bool): + return value + if isinstance(value, int): + return bool(value) + + value_str = str(value).strip().lower() + if value_str in {"true", "1", "yes", "y", "on"}: + return True + if value_str in {"false", "0", "no", "n", "off"}: + return False + raise ValueError(f"Unrecognised boolean value: {value}") + def _extract_model_and_provider(self, data: dict) -> Tuple[str, str]: """ Extract the model and provider from the request data. diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index b9110afa222..dec399c7f74 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -881,9 +881,9 @@ async def insert_sso_user( if user_defined_values.get("max_budget") is None: user_defined_values["max_budget"] = litellm.max_internal_user_budget if user_defined_values.get("budget_duration") is None: - user_defined_values[ - "budget_duration" - ] = litellm.internal_user_budget_duration + user_defined_values["budget_duration"] = ( + litellm.internal_user_budget_duration + ) if user_defined_values["user_role"] is None: user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY @@ -1194,9 +1194,9 @@ class SSOAuthenticationHandler: generic_authorization_endpoint and "okta" in generic_authorization_endpoint ): - redirect_params[ - "state" - ] = uuid.uuid4().hex # set state param for okta - required + redirect_params["state"] = ( + uuid.uuid4().hex + ) # set state param for okta - required # Handle PKCE (Proof Key for Code Exchange) if enabled # Set GENERIC_CLIENT_USE_PKCE=true to enable PKCE for enhanced OAuth security @@ -1849,9 +1849,9 @@ class MicrosoftSSOHandler: # if user is trying to get the raw sso response for debugging, return the raw sso response if return_raw_sso_response: - original_msft_result[ - MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY - ] = user_team_ids + original_msft_result[MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY] = ( + user_team_ids + ) original_msft_result["app_roles"] = app_roles return original_msft_result or {} @@ -1909,7 +1909,8 @@ class MicrosoftSSOHandler: decoded_token = jwt.decode(id_token, options={"verify_signature": False}) # Extract app_roles claim from the token - roles = decoded_token.get("app_roles", []) + ## check for both 'roles' and 'app_roles' claims + roles = decoded_token.get("app_roles", []) or decoded_token.get("roles", []) if roles and isinstance(roles, list): verbose_proxy_logger.debug( @@ -1967,9 +1968,9 @@ class MicrosoftSSOHandler: # Fetch user membership from Microsoft Graph API all_group_ids = [] - next_link: Optional[ - str - ] = MicrosoftSSOHandler.graph_api_user_groups_endpoint + next_link: Optional[str] = ( + MicrosoftSSOHandler.graph_api_user_groups_endpoint + ) auth_headers = {"Authorization": f"Bearer {access_token}"} page_count = 0 diff --git a/litellm/router.py b/litellm/router.py index b136dcbae83..569a78e3f3b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -189,7 +189,7 @@ class RoutingArgs(enum.Enum): class Router: - model_names: List = [] + model_names: set = set() cache_responses: Optional[bool] = False default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour tenacity = None @@ -1065,7 +1065,7 @@ class Router: self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) request_priority = kwargs.get("priority") or self.default_priority - start_time = time.time() + start_time = time.perf_counter() _is_prompt_management_model = self._is_prompt_management_model(model) if _is_prompt_management_model: @@ -1078,7 +1078,7 @@ class Router: response = await self.schedule_acompletion(**kwargs) else: response = await self.async_function_with_fallbacks(**kwargs) - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( @@ -1253,7 +1253,7 @@ class Router: input_kwargs_for_streaming_fallback["model"] = model parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) - start_time = time.time() + start_time = time.perf_counter() deployment = await self.async_get_available_deployment( model=model, messages=messages, @@ -1262,7 +1262,7 @@ class Router: ) _timeout_debug_deployment_dict = deployment - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( @@ -1842,8 +1842,8 @@ class Router: await self.scheduler.add_request(request=item) ## POLL QUEUE - end_time = time.time() + self.timeout - curr_time = time.time() + end_time = time.monotonic() + self.timeout + curr_time = time.monotonic() poll_interval = self.scheduler.polling_interval # poll every 3ms make_request = False @@ -1860,7 +1860,7 @@ class Router: break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) - curr_time = time.time() + curr_time = time.monotonic() if make_request: try: @@ -1904,8 +1904,8 @@ class Router: await self.scheduler.add_request(request=item) ## POLL QUEUE - end_time = time.time() + self.timeout - curr_time = time.time() + end_time = time.monotonic() + self.timeout + curr_time = time.monotonic() poll_interval = self.scheduler.polling_interval # poll every 3ms make_request = False @@ -1922,7 +1922,7 @@ class Router: break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) - curr_time = time.time() + curr_time = time.monotonic() if make_request: try: @@ -4929,22 +4929,25 @@ class Router: - hash - use hash as id """ - concat_str = model_group + # Optimized: Use list and join instead of string concatenation in loop + # This avoids creating many temporary string objects (O(n) vs O(n²) complexity) + parts = [model_group] for k, v in litellm_params.items(): if isinstance(k, str): - concat_str += k + parts.append(k) elif isinstance(k, dict): - concat_str += json.dumps(k) + parts.append(json.dumps(k)) else: - concat_str += str(k) + parts.append(str(k)) if isinstance(v, str): - concat_str += v + parts.append(v) elif isinstance(v, dict): - concat_str += json.dumps(v) + parts.append(json.dumps(v)) else: - concat_str += str(v) + parts.append(str(v)) + concat_str = "".join(parts) hash_object = hashlib.sha256(concat_str.encode()) return hash_object.hexdigest() @@ -5168,7 +5171,7 @@ class Router: verbose_router_logger.debug( f"\nInitialized Model List {self.get_model_names()}" ) - self.model_names = [m["model_name"] for m in model_list] + self.model_names = {m["model_name"] for m in model_list} # Build model_name index for O(1) lookups self._build_model_name_index(self.model_list) @@ -5374,7 +5377,7 @@ class Router: self._add_model_to_list_and_index_map( model=_deployment, model_id=deployment.model_info.id ) - self.model_names.append(deployment.model_name) + self.model_names.add(deployment.model_name) return deployment def _update_deployment_indices_after_removal( @@ -5533,9 +5536,15 @@ class Router: Returns -> Deployment or None Raise Exception -> if model found in invalid format + + Optimized with O(1) index lookup instead of O(n) linear scan. """ - for model in self.model_list: - if model["model_name"] == model_group_name: + # O(1) lookup in model_name index + if model_group_name in self.model_name_to_deployment_indices: + indices = self.model_name_to_deployment_indices[model_group_name] + if indices: + # Return first deployment for this model_name + model = self.model_list[indices[0]] if isinstance(model, dict): return Deployment(**model) elif isinstance(model, Deployment): @@ -5645,11 +5654,13 @@ class Router: Returns - dict: the model in list with 'model_name', 'litellm_params', Optional['model_info'] - None: could not find deployment in list + + Optimized with O(1) index lookup instead of O(n) linear scan. """ - for model in self.model_list: - if "model_info" in model and "id" in model["model_info"]: - if id == model["model_info"]["id"]: - return model + # O(1) lookup via model_id_to_deployment_index_map + if id in self.model_id_to_deployment_index_map: + idx = self.model_id_to_deployment_index_map[id] + return self.model_list[idx] return None def get_model_group(self, id: str) -> Optional[List]: @@ -6183,17 +6194,33 @@ class Router: if 'model_name' is none, returns all. Returns list of model id's. + + Optimized with O(1) or O(k) index lookup when model_name provided, + instead of O(n) linear scan. """ ids = [] - for model in self.model_list: - if "model_info" in model and "id" in model["model_info"]: - id = model["model_info"]["id"] - if exclude_team_models and model["model_info"].get("team_id"): - continue - if model_name is not None and model["model_name"] == model_name: - ids.append(id) - elif model_name is None: - ids.append(id) + + if model_name is not None: + # O(1) lookup in model_name index, then O(k) iteration where k = deployments for this model_name + if model_name in self.model_name_to_deployment_indices: + indices = self.model_name_to_deployment_indices[model_name] + for idx in indices: + model = self.model_list[idx] + if "model_info" in model and "id" in model["model_info"]: + if exclude_team_models and model["model_info"].get("team_id"): + continue + ids.append(model["model_info"]["id"]) + else: + # When model_name is None, return all model IDs + # Use the index map keys for O(n) where n = total deployments + for model_id in self.model_id_to_deployment_index_map.keys(): + idx = self.model_id_to_deployment_index_map[model_id] + model = self.model_list[idx] + if "model_info" in model and "id" in model["model_info"]: + if exclude_team_models and model["model_info"].get("team_id"): + continue + ids.append(model_id) + return ids def has_model_id(self, candidate_id: str) -> bool: @@ -6271,7 +6298,9 @@ class Router: model_name=model_name, model=model, team_id=team_id ): if model_alias is not None: - alias_model = copy.deepcopy(model) + # Optimized: Use shallow copy since we only modify top-level model_name + # This is much faster than deepcopy for nested dict structures + alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: @@ -6285,7 +6314,8 @@ class Router: model_name=model_name, model=model, team_id=team_id ): if model_alias is not None: - alias_model = copy.deepcopy(model) + # Optimized: Use shallow copy since we only modify top-level model_name + alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: @@ -7084,7 +7114,7 @@ class Router: if isinstance(healthy_deployments, dict): return healthy_deployments - start_time = time.time() + start_time = time.perf_counter() if ( self.routing_strategy == "usage-based-routing-v2" and self.lowesttpm_logger_v2 is not None @@ -7151,7 +7181,7 @@ class Router: f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}" ) - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index e75f969892e..115cce9fb94 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -360,6 +360,22 @@ class PillarGuardrailConfigModel(BaseModel): default="monitor", description="Action to take when content is flagged: 'block' (raise exception) or 'monitor' (log only)", ) + async_mode: Optional[bool] = Field( + default=None, + description="Set to True to request asynchronous analysis (sets `plr_async` header). Defaults to provider behaviour when omitted.", + ) + persist_session: Optional[bool] = Field( + default=None, + description="Controls Pillar session persistence (sets `plr_persist` header). Set to False to disable persistence.", + ) + include_scanners: Optional[bool] = Field( + default=True, + description="Include scanner category summaries in responses (sets `plr_scanners` header).", + ) + include_evidence: Optional[bool] = Field( + default=True, + description="Include detailed evidence payloads in responses (sets `plr_evidence` header).", + ) class NomaGuardrailConfigModel(BaseModel): diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py b/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py index e18f8dfb20e..4d0c9ed1cc5 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py @@ -15,6 +15,22 @@ class PillarGuardrailConfigModelOptionalParams(BaseModel): default="monitor", description="Action to take when content is flagged: 'block' (raise exception) or 'monitor' (log only). If not provided, the `PILLAR_ON_FLAGGED_ACTION` environment variable is checked, defaults to 'monitor'.", ) + async_mode: Optional[bool] = Field( + default=None, + description="Set to True to request asynchronous analysis (sets `plr_async` header).", + ) + persist_session: Optional[bool] = Field( + default=None, + description="Set to False to disable session persistence (sets `plr_persist` header).", + ) + include_scanners: Optional[bool] = Field( + default=True, + description="Include scanner summaries in response payloads (sets `plr_scanners` header).", + ) + include_evidence: Optional[bool] = Field( + default=True, + description="Include detailed evidence objects in response payloads (sets `plr_evidence` header).", + ) class PillarGuardrailConfigModel( diff --git a/litellm/utils.py b/litellm/utils.py index 45cc901a1a0..f017543ceaa 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7559,11 +7559,9 @@ class ProviderConfigManager: provider: LlmProviders, ) -> Optional[BaseImageEditConfig]: if LlmProviders.OPENAI == provider: - from litellm.llms.openai.image_edit.transformation import ( - OpenAIImageEditConfig, - ) + from litellm.llms.openai.image_edit import get_openai_image_edit_config - return OpenAIImageEditConfig() + return get_openai_image_edit_config(model=model) elif LlmProviders.AZURE == provider: from litellm.llms.azure.image_edit.transformation import ( AzureImageEditConfig, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index caca89d23f8..8cab321c4dd 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -17521,6 +17521,25 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "openrouter/anthropic/claude-sonnet-4.5": { + "input_cost_per_image": 0.0048, + "input_cost_per_token": 3e-06, + "input_cost_per_token_above_200k_tokens": 6e-06, + "output_cost_per_token_above_200k_tokens": 2.25e-05, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 159 + }, "openrouter/bytedance/ui-tars-1.5-7b": { "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", @@ -21896,8 +21915,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, @@ -21913,8 +21932,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, diff --git a/test_image_edit.png b/test_image_edit.png index 0380114f241..0386d2af106 100644 Binary files a/test_image_edit.png and b/test_image_edit.png differ diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index 74562c4648f..90544f747bb 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -9,6 +9,7 @@ import base64 from io import BytesIO from unittest.mock import patch, AsyncMock import json +from abc import ABC, abstractmethod sys.path.insert( 0, os.path.abspath("../..") @@ -30,6 +31,72 @@ class TestCustomLogger(CustomLogger): self.standard_logging_payload = kwargs.get("standard_logging_object", None) pass + +class BaseLLMImageEditTest(ABC): + """ + Abstract base test class that enforces a common test across all image edit test classes. + """ + + @property + def image_edit_function(self): + return litellm.image_edit + + @property + def async_image_edit_function(self): + return litellm.aimage_edit + + @abstractmethod + def get_base_image_edit_call_args(self) -> dict: + """Must return the base image edit call args""" + pass + + @pytest.fixture(autouse=True) + def _handle_rate_limits(self): + """Fixture to handle rate limit errors for all test methods""" + try: + yield + except litellm.RateLimitError: + pytest.skip("Rate limit exceeded") + except litellm.InternalServerError: + pytest.skip("Model is overloaded") + + @pytest.mark.parametrize("sync_mode", [True, False]) + @pytest.mark.flaky(retries=3, delay=2) + @pytest.mark.asyncio + async def test_openai_image_edit_litellm_sdk(self, sync_mode): + """ + Test image edit functionality with both sync and async modes. + """ + litellm._turn_on_debug() + try: + prompt = """ + Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. + """ + + call_args = self.get_base_image_edit_call_args() + call_args["prompt"] = prompt + + if sync_mode: + result = self.image_edit_function(**call_args) + else: + result = await self.async_image_edit_function(**call_args) + + print("result from image edit", result) + + # Validate the response meets expected schema + ImageResponse.model_validate(result) + + if isinstance(result, ImageResponse) and result.data: + image_base64 = result.data[0].b64_json + if image_base64: + image_bytes = base64.b64decode(image_base64) + + # Save the image to a file + with open("test_image_edit.png", "wb") as f: + f.write(image_bytes) + except litellm.ContentPolicyViolationError as e: + pass + # Get the current directory of the file being run pwd = os.path.dirname(os.path.realpath(__file__)) @@ -49,45 +116,31 @@ def get_test_images_as_bytesio(): bytesio_images.append(BytesIO(image_bytes)) return bytesio_images -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.flaky(retries=3, delay=2) -@pytest.mark.asyncio -async def test_openai_image_edit_litellm_sdk(sync_mode): - from litellm import image_edit, aimage_edit - litellm._turn_on_debug() - try: - prompt = """ - Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. - """ - if sync_mode: - result = image_edit( - prompt=prompt, - model="gpt-image-1", - image=TEST_IMAGES, - ) - else: - result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=TEST_IMAGES, - ) - print("result from image edit", result) +class TestOpenAIImageEditGPTImage1(BaseLLMImageEditTest): + """ + Concrete implementation of BaseLLMImageEditTest for OpenAI image edits. + """ - # Validate the response meets expected schema - ImageResponse.model_validate(result) - - if isinstance(result, ImageResponse) and result.data: - image_base64 = result.data[0].b64_json - if image_base64: - image_bytes = base64.b64decode(image_base64) + def get_base_image_edit_call_args(self) -> dict: + """Return base call args for OpenAI image edit""" + return { + "model": "gpt-image-1", + "image": TEST_IMAGES, + } - # Save the image to a file - with open("test_image_edit.png", "wb") as f: - f.write(image_bytes) - except litellm.ContentPolicyViolationError as e: - pass +class TestOpenAIImageEditDallE2(BaseLLMImageEditTest): + """ + Concrete implementation of BaseLLMImageEditTest for OpenAI DALL-E-2 image edits. + DALL-E-2 only supports a single image (not an array). + """ + def get_base_image_edit_call_args(self) -> dict: + """Return base call args for OpenAI DALL-E-2 image edit (single image only)""" + return { + "model": "dall-e-2", + "image": SINGLE_TEST_IMAGE, + } @pytest.mark.flaky(retries=3, delay=2) diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index a31c4d8210f..c2339d9eec5 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -1411,7 +1411,8 @@ def test_generate_model_id_with_deployment_model_name(model_list): "Expected TypeError when model_group is None - this confirms our fix is needed" ) except TypeError as e: - assert "unsupported operand type(s) for +=" in str(e) + # After optimization, error message changed but still fails appropriately on None + assert "unsupported operand type(s) for +=" in str(e) or "expected str instance, NoneType found" in str(e) print(f"✓ Correctly failed with None model_group (as expected): {e}") except Exception as e: pytest.fail(f"Unexpected error with None model_group: {e}") diff --git a/tests/router_unit_tests/test_router_index_management.py b/tests/router_unit_tests/test_router_index_management.py index 04ea9214991..0c313e9ef4e 100644 --- a/tests/router_unit_tests/test_router_index_management.py +++ b/tests/router_unit_tests/test_router_index_management.py @@ -1,6 +1,8 @@ import sys import os import pytest +import ast +import ast sys.path.insert( 0, os.path.abspath("../..") @@ -177,3 +179,97 @@ class TestRouterIndexManagement: # Verify: New entry is added assert "claude-3" in router.model_name_to_deployment_indices assert router.model_name_to_deployment_indices["claude-3"] == [0] + + def test_no_linear_scans_in_router(self): + """ + Static analysis test to ensure Router doesn't use O(n) linear scans. + + Scans router.py for 'in self.model_list' pattern which indicates + inefficient O(n) iteration instead of using index-based O(1) lookups. + + Methods should use: + - model_id_to_deployment_index_map for O(1) model_id lookups + - model_name_to_deployment_indices for O(1) + O(k) model_name lookups + """ + # Methods that are allowed to iterate through self.model_list + ALLOWED_METHODS = [ + "_get_deployment_by_litellm_model", # Edge case: lookup by litellm_params.model (not indexed) + ] + + # Get path to router.py + router_file = os.path.join( + os.path.dirname(os.path.dirname(os.path.dirname(__file__))), + "litellm", + "router.py" + ) + + # Read the file + with open(router_file, 'r') as f: + content = f.read() + + # Parse with AST + tree = ast.parse(content) + + # Find violations + violations = [] + ignore_methods = set(ALLOWED_METHODS) + + for node in ast.walk(tree): + if isinstance(node, ast.FunctionDef): + method_name = node.name + + # Skip ignored methods + if method_name in ignore_methods: + continue + + # Get source for this method + try: + method_source = ast.get_source_segment(content, node) + if not method_source: + continue + + # Check for the anti-pattern: "in self.model_list" + # This catches: for x in self.model_list, if x in self.model_list, etc. + if "in self.model_list" in method_source: + # Extract the specific line for better error reporting + lines = method_source.split('\n') + pattern_line = None + for line in lines: + if "in self.model_list" in line: + pattern_line = line.strip() + break + + violations.append({ + "method": method_name, + "line": node.lineno, + "pattern": pattern_line or "in self.model_list" + }) + except Exception: + # Skip if we can't get source segment + pass + + # Assert no violations + if violations: + error_msg = "\n".join([ + f" - {v['method']}() at line {v['line']}: {v['pattern']}" + for v in violations + ]) + + pytest.fail( + f"\n{'='*70}\n" + f"Found O(n) linear scan pattern in router.py:\n\n" + f"{error_msg}\n\n" + f"These methods should use index maps instead:\n" + f" - model_id_to_deployment_index_map (for model_id lookups)\n" + f" - model_name_to_deployment_indices (for model_name lookups)\n\n" + f"If a method legitimately needs O(n) iteration, add it to\n" + f"ALLOWED_METHODS in this test method.\n" + f"{'='*70}\n" + ) + def test_model_names_is_set(self): + """Verify that model_names uses a set for O(1) lookups, not a list (O(n))""" + router = Router(model_list=[]) + + assert isinstance(router.model_names, set), ( + f"model_names should be a set for O(1) lookups, but got {type(router.model_names)}" + ) diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 0e31699fd83..09fc31d18b7 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -399,3 +399,41 @@ async def test_session_validation(): mock_valid_session = MockClientSession() transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore assert transport3.client is mock_valid_session # Should reuse session + + +@pytest.mark.parametrize( + "env_curve,litellm_curve,expected_curve,should_call", + [ + # env_curve: SSL_ECDH_CURVE env var | litellm_curve: litellm.ssl_ecdh_curve variable + # expected_curve: curve that should be set | should_call: whether set_ecdh_curve() should be called + + # Valid configurations + ("X25519", None, "X25519", True), # Env var only + ("prime256v1", None, "prime256v1", True), # Different valid curve + (None, "secp384r1", "secp384r1", True), # litellm variable only + ("X25519", "secp521r1", "X25519", True), # Env var takes precedence + # Empty/None configurations - should skip + ("", None, None, False), # Empty string - skip configuration + (None, None, None, False), # None value - skip configuration + ] +) +def test_ssl_ecdh_curve(env_curve, litellm_curve, expected_curve, should_call, monkeypatch): + """Test SSL ECDH curve configuration with valid curves and precedence""" + with patch.dict(os.environ, clear=True): + if env_curve: + monkeypatch.setenv("SSL_ECDH_CURVE", env_curve) + + original_value = litellm.ssl_ecdh_curve + try: + litellm.ssl_ecdh_curve = litellm_curve + + with patch.object(ssl.SSLContext, 'set_ecdh_curve') as mock_set_curve: + ssl_context = get_ssl_configuration() + + if should_call: + mock_set_curve.assert_called_once_with(expected_curve) + else: + mock_set_curve.assert_not_called() + assert isinstance(ssl_context, ssl.SSLContext) + finally: + litellm.ssl_ecdh_curve = original_value diff --git a/tests/guardrails_tests/test_pillar_guardrails.py b/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py similarity index 90% rename from tests/guardrails_tests/test_pillar_guardrails.py rename to tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py index aeb2227f9b9..67030a8161f 100644 --- a/tests/guardrails_tests/test_pillar_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py @@ -8,8 +8,12 @@ and following LiteLLM testing patterns and best practices. # Standard library imports import os import sys +from typing import Dict from unittest.mock import Mock, patch +# Add parent directory to path for imports +sys.path.insert(0, os.path.abspath("../../..")) + # Third-party imports import pytest from fastapi.exceptions import HTTPException @@ -26,9 +30,6 @@ from litellm.proxy.guardrails.guardrail_hooks.pillar import ( ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 -# Add parent directory to path for imports -sys.path.insert(0, os.path.abspath("../..")) - # ============================================================================ # FIXTURES @@ -221,6 +222,18 @@ def mock_llm_response(): return mock_response +@pytest.fixture +def pillar_async_response(): + """Fixture providing an asynchronous Pillar API queue response.""" + return Response( + json={"status": "queued", "session_id": "async-session", "position": 1}, + status_code=202, + request=Request( + method="POST", url="https://api.pillar.security/api/v1/protect" + ), + ) + + @pytest.fixture def mock_llm_response_with_tools(): """Fixture providing a mock LLM response with tool calls.""" @@ -440,6 +453,55 @@ async def test_post_call_hook_with_tool_calls( assert result == mock_llm_response_with_tools +# ========================================================================= +# HEADER CONFIGURATION TESTS +# ========================================================================= + + +@pytest.mark.asyncio +async def test_pre_call_hook_custom_header_overrides( + sample_request_data, + user_api_key_dict, + dual_cache, + pillar_async_response, +): + """Ensure configuration values translate into correct Protect headers.""" + + guardrail = PillarGuardrail( + guardrail_name="pillar-header-test", + api_key="test-pillar-key", + api_base="https://api.pillar.security", + on_flagged_action="monitor", + persist_session=False, + async_mode=True, + include_scanners=False, + include_evidence=False, + ) + + captured_headers: Dict[str, str] = {} + + async def _mock_post(*args, **kwargs): + captured_headers.update(kwargs.get("headers", {})) + return pillar_async_response + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new=_mock_post, + ): + result = await guardrail.async_pre_call_hook( + data=sample_request_data, + cache=dual_cache, + user_api_key_dict=user_api_key_dict, + call_type="completion", + ) + + assert result == sample_request_data + assert captured_headers.get("plr_persist") == "false" + assert captured_headers.get("plr_async") == "true" + assert captured_headers.get("plr_scanners") == "false" + assert captured_headers.get("plr_evidence") == "false" + + # ============================================================================ # EDGE CASE TESTS # ============================================================================ diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index d8a41553565..0f403ae5d65 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -99,7 +99,9 @@ def test_microsoft_sso_handler_with_empty_response(): # Test with None response # Act - result = MicrosoftSSOHandler.openid_from_response(response=None, team_ids=[], user_role=None) + result = MicrosoftSSOHandler.openid_from_response( + response=None, team_ids=[], user_role=None + ) # Assert assert isinstance(result, CustomOpenID) @@ -789,11 +791,11 @@ class TestCLISSOCallbackFunction: "not-sk-key", "sk", # too short ] - + for invalid_key in invalid_keys: # This should fail validation before any database operations # We can test this by checking if the key starts with 'sk-' - if not invalid_key or not invalid_key.startswith('sk-'): + if not invalid_key or not invalid_key.startswith("sk-"): # This would trigger the validation error assert True # Validation works as expected @@ -806,14 +808,14 @@ class TestCLIPollingFunction: # Test key format validation logic invalid_keys = [ "invalid-key", - "not-sk-key", + "not-sk-key", "", "sk", # too short ] - + for invalid_key in invalid_keys: # Validation logic: key must start with 'sk-' - if not invalid_key.startswith('sk-'): + if not invalid_key.startswith("sk-"): # This would trigger the validation error in the actual function assert True # Validation works as expected @@ -827,7 +829,7 @@ class TestAuthCallbackRouting: # Test CLI state detection logic cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123" - + # This mimics the logic in auth_callback if cli_state and cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"): # Extract the key ID from the state @@ -839,18 +841,20 @@ class TestAuthCallbackRouting: def test_non_cli_state_routing(self): """Test that non-CLI states don't trigger CLI routing""" from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX - + non_cli_states = [ "regular_oauth_state", - "some_random_string", + "some_random_string", None, "", - "not_session_token:something" + "not_session_token:something", ] - + for state in non_cli_states: # This mimics the routing logic in auth_callback - should_route_to_cli = state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") + should_route_to_cli = state and state.startswith( + f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:" + ) assert not should_route_to_cli, f"State '{state}' should not route to CLI" @@ -864,9 +868,9 @@ class TestGoogleLoginCLIIntegration: # Test the CLI state generation logic used in google_login source = "litellm-cli" key = "sk-test123" - + cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key) - + assert cli_state is not None assert cli_state.startswith("litellm-session-token:") assert "sk-test123" in cli_state @@ -882,10 +886,12 @@ class TestGoogleLoginCLIIntegration: (None, "sk-test123"), ("wrong-source", "sk-test123"), ] - + for source, key in test_cases: cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key) - assert cli_state is None, f"CLI state should not be generated for source='{source}', key='{key}'" + assert ( + cli_state is None + ), f"CLI state should not be generated for source='{source}', key='{key}'" class TestSSOHandlerIntegration: @@ -896,13 +902,24 @@ class TestSSOHandlerIntegration: from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler # Test that SSO handler is used when client IDs are provided - assert SSOAuthenticationHandler.should_use_sso_handler(google_client_id="test") is True - assert SSOAuthenticationHandler.should_use_sso_handler(microsoft_client_id="test") is True - assert SSOAuthenticationHandler.should_use_sso_handler(generic_client_id="test") is True - + assert ( + SSOAuthenticationHandler.should_use_sso_handler(google_client_id="test") + is True + ) + assert ( + SSOAuthenticationHandler.should_use_sso_handler(microsoft_client_id="test") + is True + ) + assert ( + SSOAuthenticationHandler.should_use_sso_handler(generic_client_id="test") + is True + ) + # Test that SSO handler is not used when no client IDs are provided assert SSOAuthenticationHandler.should_use_sso_handler() is False - assert SSOAuthenticationHandler.should_use_sso_handler(None, None, None) is False + assert ( + SSOAuthenticationHandler.should_use_sso_handler(None, None, None) is False + ) def test_get_redirect_url_for_sso(self): """Test the redirect URL generation for SSO""" @@ -911,13 +928,12 @@ class TestSSOHandlerIntegration: # Mock request object mock_request = MagicMock() mock_request.base_url = "https://test.litellm.ai/" - + # Test redirect URL generation redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( - request=mock_request, - sso_callback_route="sso/callback" + request=mock_request, sso_callback_route="sso/callback" ) - + assert redirect_url.startswith("https://test.litellm.ai") assert "sso/callback" in redirect_url @@ -928,21 +944,25 @@ class TestUISSO_FunctionsExistence: def test_cli_sso_callback_exists(self): """Test that cli_sso_callback function exists""" from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback + assert callable(cli_sso_callback) def test_cli_poll_key_exists(self): """Test that cli_poll_key function exists""" from litellm.proxy.management_endpoints.ui_sso import cli_poll_key + assert callable(cli_poll_key) def test_auth_callback_exists(self): """Test that auth_callback function exists""" from litellm.proxy.management_endpoints.ui_sso import auth_callback + assert callable(auth_callback) def test_google_login_exists(self): """Test that google_login function exists""" from litellm.proxy.management_endpoints.ui_sso import google_login + assert callable(google_login) def test_sso_authentication_handler_exists(self): @@ -951,9 +971,9 @@ class TestUISSO_FunctionsExistence: # Check that the class exists assert SSOAuthenticationHandler is not None - + # Check that the new _get_cli_state method exists - assert hasattr(SSOAuthenticationHandler, '_get_cli_state') + assert hasattr(SSOAuthenticationHandler, "_get_cli_state") assert callable(SSOAuthenticationHandler._get_cli_state) @@ -963,9 +983,11 @@ class TestSSOStateHandling: def test_get_cli_state_valid(self): """Test generating CLI state with valid parameters""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - - state = SSOAuthenticationHandler._get_cli_state(source="litellm-cli", key="sk-test123") - + + state = SSOAuthenticationHandler._get_cli_state( + source="litellm-cli", key="sk-test123" + ) + assert state is not None assert state.startswith("litellm-session-token:") assert "sk-test123" in state @@ -973,37 +995,39 @@ class TestSSOStateHandling: def test_get_cli_state_invalid_source(self): """Test generating CLI state with invalid source""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - - state = SSOAuthenticationHandler._get_cli_state(source="invalid_source", key="sk-test123") - + + state = SSOAuthenticationHandler._get_cli_state( + source="invalid_source", key="sk-test123" + ) + assert state is None def test_get_cli_state_no_key(self): """Test generating CLI state without key""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state(source="litellm-cli", key=None) - + assert state is None def test_get_cli_state_no_source(self): """Test generating CLI state without source""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state(source=None, key="sk-test123") - + assert state is None def test_get_cli_state_with_existing_key(self): """Test generating CLI state with existing_key embedded in state parameter""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state( - source="litellm-cli", + source="litellm-cli", key="sk-new-key-123", - existing_key="sk-existing-key-456" + existing_key="sk-existing-key-456", ) - + assert state is not None assert state.startswith("litellm-session-token:") assert "sk-new-key-123" in state @@ -1014,13 +1038,11 @@ class TestSSOStateHandling: def test_get_cli_state_without_existing_key(self): """Test generating CLI state without existing_key""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state( - source="litellm-cli", - key="sk-new-key-789", - existing_key=None + source="litellm-cli", key="sk-new-key-789", existing_key=None ) - + assert state is not None assert state.startswith("litellm-session-token:") assert "sk-new-key-789" in state @@ -1039,7 +1061,7 @@ class TestStateRouting: # Test CLI state format cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123" assert cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") - + # Test extraction of key from state key_id = cli_state.split(":", 1)[1] assert key_id == "sk-test123" @@ -1049,13 +1071,15 @@ class TestStateRouting: from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX # State format: {PREFIX}:{key}:{existing_key} - cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-456:sk-existing-key-789" - + cli_state = ( + f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-456:sk-existing-key-789" + ) + # Parse as done in auth_callback state_parts = cli_state.split(":", 2) # Split into max 3 parts key_id = state_parts[1] if len(state_parts) > 1 else None existing_key = state_parts[2] if len(state_parts) > 2 else None - + assert key_id == "sk-new-key-456" assert existing_key == "sk-existing-key-789" @@ -1065,12 +1089,12 @@ class TestStateRouting: # State format: {PREFIX}:{key} cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-999" - + # Parse as done in auth_callback state_parts = cli_state.split(":", 2) # Split into max 3 parts key_id = state_parts[1] if len(state_parts) > 1 else None existing_key = state_parts[2] if len(state_parts) > 2 else None - + assert key_id == "sk-new-key-999" assert existing_key is None @@ -1084,9 +1108,9 @@ class TestStateRouting: "some_random_string", None, "", - "not_session_token:something" + "not_session_token:something", ] - + for state in test_states: if state: assert not state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") @@ -1105,10 +1129,10 @@ class TestHTMLIntegration: # Test that function exists and is callable assert callable(render_cli_sso_success_page) - + # Test that it returns expected type html = render_cli_sso_success_page() - + assert isinstance(html, str) assert len(html) > 0 @@ -1125,11 +1149,19 @@ class TestCustomUISSO: # Mock request mock_request = MagicMock() mock_request.base_url = "https://test.example.com/" - + # Mock user_custom_ui_sso_sign_in_handler to exist but make enterprise import fail with patch("litellm.proxy.proxy_server.premium_user", True): - with patch("litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", MagicMock()): - with patch.dict('sys.modules', {'enterprise.litellm_enterprise.proxy.auth.custom_sso_handler': None}): + with patch( + "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", + MagicMock(), + ): + with patch.dict( + "sys.modules", + { + "enterprise.litellm_enterprise.proxy.auth.custom_sso_handler": None + }, + ): # Temporarily mock the google_login function call to test the import error path async def mock_google_login(): # This mimics the relevant part of google_login that would trigger the import error @@ -1137,13 +1169,19 @@ class TestCustomUISSO: from enterprise.litellm_enterprise.proxy.auth.custom_sso_handler import ( EnterpriseCustomSSOHandler, ) + return "success" except ImportError: - raise ValueError("Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise.") - + raise ValueError( + "Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise." + ) + # Test that the ValueError is raised with the correct message import pytest - with pytest.raises(ValueError, match="Enterprise features are not available"): + + with pytest.raises( + ValueError, match="Enterprise features are not available" + ): asyncio.run(mock_google_login()) @pytest.mark.asyncio @@ -1195,8 +1233,10 @@ class TestCustomUISSO: return_value=mock_redirect_response, ) as mock_get_redirect: # Act - result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( - request=mock_request + result = ( + await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( + request=mock_request + ) ) # Assert @@ -1241,12 +1281,14 @@ class TestCustomUISSO: async def handle_custom_ui_sso_sign_in(self, request: Request) -> OpenID: self.method_called = True self.received_request = request - + # Parse headers like the actual implementation would request_headers_dict = dict(request.headers) return OpenID( id=request_headers_dict.get("x-litellm-user-id", "default_user"), - email=request_headers_dict.get("x-litellm-user-email", "default@test.com"), + email=request_headers_dict.get( + "x-litellm-user-email", "default@test.com" + ), first_name="Custom", last_name="Handler", display_name="Custom Handler Test", @@ -1260,7 +1302,7 @@ class TestCustomUISSO: # Mock request with custom headers mock_request = MagicMock(spec=Request) mock_request.headers = { - "x-litellm-user-id": "custom_test_user_456", + "x-litellm-user-id": "custom_test_user_456", "x-litellm-user-email": "custom@example.com", "x-forwarded-for": "10.0.0.1", } @@ -1281,8 +1323,10 @@ class TestCustomUISSO: return_value=mock_redirect_response, ) as mock_get_redirect: # Act - result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( - request=mock_request + result = ( + await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( + request=mock_request + ) ) # Assert that our custom handler was executed @@ -1292,7 +1336,7 @@ class TestCustomUISSO: # Verify the redirect response was called with the OpenID from our custom handler mock_get_redirect.assert_called_once() call_args = mock_get_redirect.call_args.kwargs - + # Verify the OpenID object has the expected values from our custom handler openid_result = call_args["result"] assert openid_result.id == "custom_test_user_456" @@ -1323,25 +1367,30 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock(spec=Request) - + # Test data existing_key = "sk-existing-key-123" new_key = "sk-new-key-456" - + # Mock the regenerate helper function - with patch("litellm.proxy.management_endpoints.ui_sso._regenerate_cli_key") as mock_regenerate, \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success"): - + with patch( + "litellm.proxy.management_endpoints.ui_sso._regenerate_cli_key" + ) as mock_regenerate, patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", + return_value="Success", + ): + # Act result = await cli_sso_callback( - request=mock_request, - key=new_key, - existing_key=existing_key + request=mock_request, key=new_key, existing_key=existing_key ) - + # Assert - mock_regenerate.assert_called_once_with(existing_key=existing_key, new_key=new_key, user_id=None) + mock_regenerate.assert_called_once_with( + existing_key=existing_key, new_key=new_key, user_id=None + ) assert result.status_code == 200 assert "Success" in result.body.decode() @@ -1352,22 +1401,25 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock(spec=Request) - + # Test data new_key = "sk-new-key-789" - + # Mock the create helper function - with patch("litellm.proxy.management_endpoints.ui_sso._create_new_cli_key") as mock_create, \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success"): - + with patch( + "litellm.proxy.management_endpoints.ui_sso._create_new_cli_key" + ) as mock_create, patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", + return_value="Success", + ): + # Act result = await cli_sso_callback( - request=mock_request, - key=new_key, - existing_key=None + request=mock_request, key=new_key, existing_key=None ) - + # Assert mock_create.assert_called_once_with(key=new_key, user_id=None) assert result.status_code == 200 @@ -1381,32 +1433,42 @@ class TestCLIKeyRegenerationFlow: # Mock request (no query params needed - existing_key is in state) mock_request = MagicMock(spec=Request) - + # CLI state with existing_key embedded: {PREFIX}:{key}:{existing_key} cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-session-key-456:sk-existing-cli-key-123" - + # Mock the CLI callback and required proxy server components mock_result = {"user_id": "test-user", "email": "test@example.com"} - - with patch("litellm.proxy.management_endpoints.ui_sso.cli_sso_callback") as mock_cli_callback, \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.proxy_server.master_key", "test-master-key"), \ - patch("litellm.proxy.proxy_server.general_settings", {}), \ - patch("litellm.proxy.proxy_server.jwt_handler", MagicMock()), \ - patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), \ - patch.dict(os.environ, {"GOOGLE_CLIENT_ID": "test-google-id"}, clear=True), \ - patch("litellm.proxy.management_endpoints.ui_sso.GoogleSSOHandler.get_google_callback_response", return_value=mock_result): + + with patch( + "litellm.proxy.management_endpoints.ui_sso.cli_sso_callback" + ) as mock_cli_callback, patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.proxy_server.master_key", "test-master-key" + ), patch( + "litellm.proxy.proxy_server.general_settings", {} + ), patch( + "litellm.proxy.proxy_server.jwt_handler", MagicMock() + ), patch( + "litellm.proxy.proxy_server.user_api_key_cache", MagicMock() + ), patch.dict( + os.environ, {"GOOGLE_CLIENT_ID": "test-google-id"}, clear=True + ), patch( + "litellm.proxy.management_endpoints.ui_sso.GoogleSSOHandler.get_google_callback_response", + return_value=mock_result, + ): mock_cli_callback.return_value = MagicMock() - + # Act await auth_callback(request=mock_request, state=cli_state) - + # Assert - existing_key should be extracted from state parameter mock_cli_callback.assert_called_once_with( request=mock_request, key="sk-new-session-key-456", existing_key="sk-existing-cli-key-123", - result=mock_result + result=mock_result, ) def test_get_redirect_url_does_not_include_existing_key_in_url(self): @@ -1416,15 +1478,17 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock() mock_request.base_url = "https://test.litellm.ai/" - - with patch("litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"): + + with patch( + "litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai" + ): # Test with existing_key - should NOT be in URL redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( request=mock_request, sso_callback_route="sso/callback", - existing_key="sk-existing-123" + existing_key="sk-existing-123", ) - + # existing_key should NOT be in the URL assert "https://test.litellm.ai/sso/callback" == redirect_url assert "existing_key" not in redirect_url @@ -1436,44 +1500,181 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock() mock_request.base_url = "https://test.litellm.ai/" - - with patch("litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"): + + with patch( + "litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai" + ): # Test without existing_key redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( - request=mock_request, - sso_callback_route="sso/callback" + request=mock_request, sso_callback_route="sso/callback" ) - + assert "https://test.litellm.ai/sso/callback" == redirect_url @pytest.mark.asyncio async def test_cli_sso_callback_regenerate_vs_create_flow(self): """Test CLI SSO callback calls regenerate_key_fn when existing_key provided, generate_key_helper_fn when not""" from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback - + mock_request = MagicMock(spec=Request) - - with patch("litellm.proxy.management_endpoints.key_management_endpoints.regenerate_key_fn") as mock_regenerate, \ - patch("litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn") as mock_generate, \ - patch("litellm.proxy._types.UserAPIKeyAuth.get_litellm_cli_user_api_key_auth"), \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success"): - + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.regenerate_key_fn" + ) as mock_regenerate, patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn" + ) as mock_generate, patch( + "litellm.proxy._types.UserAPIKeyAuth.get_litellm_cli_user_api_key_auth" + ), patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", + return_value="Success", + ): + # Test regeneration path - await cli_sso_callback(mock_request, key="sk-new-123", existing_key="sk-existing-456") + await cli_sso_callback( + mock_request, key="sk-new-123", existing_key="sk-existing-456" + ) mock_regenerate.assert_called_once() mock_generate.assert_not_called() - + # Reset mocks mock_regenerate.reset_mock() mock_generate.reset_mock() - + # Test creation path await cli_sso_callback(mock_request, key="sk-new-789", existing_key=None) mock_regenerate.assert_not_called() mock_generate.assert_called_once() +class TestGetAppRolesFromIdToken: + """Test the get_app_roles_from_id_token method""" + + def test_roles_picked_when_app_roles_not_exists(self): + """Test that 'roles' is picked when 'app_roles' doesn't exist""" + import jwt + + # Create a token with only 'roles' claim + token_payload = { + "sub": "user123", + "email": "test@example.com", + "roles": ["Admin", "User", "Developer"], + } + + # Create a mock JWT token + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload) as mock_jwt_decode: + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert + assert result == ["Admin", "User", "Developer"] + mock_jwt_decode.assert_called_once_with( + mock_token, options={"verify_signature": False} + ) + + def test_app_roles_picked_when_both_exist(self): + """Test that 'app_roles' takes precedence when both 'app_roles' and 'roles' exist""" + import jwt + + # Create a token with both 'app_roles' and 'roles' claims + token_payload = { + "sub": "user123", + "email": "test@example.com", + "app_roles": ["AppAdmin", "AppUser"], + "roles": ["RoleAdmin", "RoleUser"], + } + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert - app_roles should be picked, not roles + assert result == ["AppAdmin", "AppUser"] + + def test_roles_picked_when_app_roles_is_empty(self): + """Test that 'roles' is picked when 'app_roles' exists but is empty""" + import jwt + + # Create a token with empty 'app_roles' and populated 'roles' + token_payload = { + "sub": "user123", + "email": "test@example.com", + "app_roles": [], + "roles": ["Admin", "User"], + } + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert - roles should be picked since app_roles is empty + assert result == ["Admin", "User"] + + def test_empty_list_when_neither_exists(self): + """Test that empty list is returned when neither 'app_roles' nor 'roles' exist""" + import jwt + + # Create a token without roles claims + token_payload = {"sub": "user123", "email": "test@example.com"} + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert + assert result == [] + + def test_empty_list_when_no_token_provided(self): + """Test that empty list is returned when no token is provided""" + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(None) + + # Assert + assert result == [] + + def test_empty_list_when_roles_not_a_list(self): + """Test that empty list is returned when roles is not a list""" + import jwt + + # Create a token with non-list roles + token_payload = { + "sub": "user123", + "email": "test@example.com", + "roles": "Admin", # String instead of list + } + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert + assert result == [] + + def test_error_handling_on_jwt_decode_exception(self): + """Test that exceptions during JWT decode are handled gracefully""" + import jwt + + mock_token = "invalid.jwt.token" + + with patch("jwt.decode", side_effect=Exception("Invalid token")): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert - should return empty list on error + assert result == [] + + class TestProcessSSOJWTAccessToken: """Test the process_sso_jwt_access_token helper function""" @@ -1496,10 +1697,12 @@ class TestProcessSSOJWTAccessToken: "sub": "1234567890", "name": "John Doe", "iat": 1516239022, - "groups": ["team1", "team2", "team3"] + "groups": ["team1", "team2", "team3"], } - def test_process_sso_jwt_access_token_with_valid_token(self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload): + def test_process_sso_jwt_access_token_with_valid_token( + self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload + ): """Test processing a valid JWT access token with team extraction""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1513,7 +1716,7 @@ class TestProcessSSOJWTAccessToken: last_name="User", display_name="Test User", provider="generic", - team_ids=[] + team_ids=[], ) with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode: @@ -1521,7 +1724,7 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert @@ -1529,14 +1732,18 @@ class TestProcessSSOJWTAccessToken: mock_jwt_decode.assert_called_once_with( sample_jwt_token, options={"verify_signature": False} ) - + # Verify team IDs were extracted from JWT - mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(sample_jwt_payload) - + mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with( + sample_jwt_payload + ) + # Verify team IDs were set on the result object assert result.team_ids == ["team1", "team2", "team3"] - def test_process_sso_jwt_access_token_with_existing_team_ids(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_with_existing_team_ids( + self, mock_jwt_handler, sample_jwt_token + ): """Test that existing team IDs are not overwritten""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1551,7 +1758,7 @@ class TestProcessSSOJWTAccessToken: last_name="User", display_name="Test User", provider="generic", - team_ids=existing_team_ids + team_ids=existing_team_ids, ) with patch("jwt.decode") as mock_jwt_decode: @@ -1559,51 +1766,53 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert # JWT should still be decoded mock_jwt_decode.assert_called_once() - + # But team IDs should NOT be extracted since they already exist mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - + # Existing team IDs should remain unchanged assert result.team_ids == existing_team_ids - def test_process_sso_jwt_access_token_with_dict_result(self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload): + def test_process_sso_jwt_access_token_with_dict_result( + self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload + ): """Test processing with a dictionary result object""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, ) # Create a dictionary result without team_ids - result = { - "id": "test_user", - "email": "test@example.com", - "name": "Test User" - } + result = {"id": "test_user", "email": "test@example.com", "name": "Test User"} with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode: # Act process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert mock_jwt_decode.assert_called_once_with( sample_jwt_token, options={"verify_signature": False} ) - mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(sample_jwt_payload) - + mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with( + sample_jwt_payload + ) + # Verify team_ids was added to the dict as a key assert "team_ids" in result assert result["team_ids"] == ["team1", "team2", "team3"] - def test_process_sso_jwt_access_token_with_dict_existing_team_ids(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_with_dict_existing_team_ids( + self, mock_jwt_handler, sample_jwt_token + ): """Test that existing team IDs in dictionary are not overwritten""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1615,7 +1824,7 @@ class TestProcessSSOJWTAccessToken: "id": "test_user", "email": "test@example.com", "name": "Test User", - "team_ids": existing_team_ids + "team_ids": existing_team_ids, } with patch("jwt.decode") as mock_jwt_decode: @@ -1623,16 +1832,16 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert # JWT should still be decoded mock_jwt_decode.assert_called_once() - + # But team IDs should NOT be extracted since they already exist mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - + # Existing team IDs should remain unchanged assert result["team_ids"] == existing_team_ids @@ -1642,20 +1851,14 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token, ) - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) # Test with None access token with patch("jwt.decode") as mock_jwt_decode: process_sso_jwt_access_token( - access_token_str=None, - sso_jwt_handler=mock_jwt_handler, - result=result + access_token_str=None, sso_jwt_handler=mock_jwt_handler, result=result ) - + # Assert nothing was processed mock_jwt_decode.assert_not_called() mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() @@ -1664,11 +1867,9 @@ class TestProcessSSOJWTAccessToken: # Test with empty string access token with patch("jwt.decode") as mock_jwt_decode: process_sso_jwt_access_token( - access_token_str="", - sso_jwt_handler=mock_jwt_handler, - result=result + access_token_str="", sso_jwt_handler=mock_jwt_handler, result=result ) - + # Assert nothing was processed mock_jwt_decode.assert_not_called() mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() @@ -1680,25 +1881,21 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token, ) - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) with patch("jwt.decode") as mock_jwt_decode: # Act process_sso_jwt_access_token( - access_token_str=sample_jwt_token, - sso_jwt_handler=None, - result=result + access_token_str=sample_jwt_token, sso_jwt_handler=None, result=result ) # Assert nothing was processed mock_jwt_decode.assert_not_called() assert result.team_ids == [] - def test_process_sso_jwt_access_token_no_result(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_no_result( + self, mock_jwt_handler, sample_jwt_token + ): """Test that nothing happens when result is None""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1709,32 +1906,32 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=None + result=None, ) # Assert nothing was processed mock_jwt_decode.assert_not_called() mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - def test_process_sso_jwt_access_token_jwt_decode_exception(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_jwt_decode_exception( + self, mock_jwt_handler, sample_jwt_token + ): """Test that JWT decode exceptions are not caught (should propagate up)""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, ) - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) - with patch("jwt.decode", side_effect=Exception("JWT decode error")) as mock_jwt_decode: + with patch( + "jwt.decode", side_effect=Exception("JWT decode error") + ) as mock_jwt_decode: # Act & Assert with pytest.raises(Exception, match="JWT decode error"): process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Verify JWT decode was attempted @@ -1742,7 +1939,9 @@ class TestProcessSSOJWTAccessToken: # But team extraction should not have been called mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - def test_process_sso_jwt_access_token_empty_team_ids_from_jwt(self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload): + def test_process_sso_jwt_access_token_empty_team_ids_from_jwt( + self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload + ): """Test processing when JWT handler returns empty team IDs""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1751,24 +1950,22 @@ class TestProcessSSOJWTAccessToken: # Configure mock to return empty team IDs mock_jwt_handler.get_team_ids_from_jwt.return_value = [] - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode: # Act process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert mock_jwt_decode.assert_called_once() - mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(sample_jwt_payload) - + mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with( + sample_jwt_payload + ) + # Even empty team IDs should be set assert result.team_ids == []