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 == []