Merge branch 'main' into litellm_sso_add_pkce

This commit is contained in:
Ishaan Jaff 2025-10-16 15:48:01 -07:00 • committed by GitHub
commit f27f2d4803
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
25 changed files with 1327 additions and 321 deletions

View file

@ -117,10 +117,52 @@ litellm_settings:
```bash
export SSL_CERTIFICATE="/path/to/certificate.pem"
```
</TabItem>
</Tabs>
## 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.
<Tabs>
<TabItem value="sdk" label="SDK">
```python
import litellm
litellm.ssl_ecdh_curve = "X25519" # Disables PQC for better performance
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```yaml
litellm_settings:
ssl_ecdh_curve: "X25519"
```
</TabItem>
<TabItem value="env_var" label="Environment Variables">
```bash
export SSL_ECDH_CURVE="X25519"
```
</TabItem>
</Tabs>
**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:

View file

@ -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

View file

@ -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
<Tabs>
<TabItem value="safe" label="Simple Safe Request">
**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)
- [LiteLLM Docs](https://docs.litellm.ai)

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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

View file

@ -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)

View file

@ -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,

View file

@ -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

View file

@ -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,
}

View file

@ -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.

View file

@ -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

View file

@ -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(

View file

@ -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):

View file

@ -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(

View file

@ -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,

View file

@ -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,

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.5 MiB

After

Width:  |  Height:  |  Size: 1.9 MiB

View file

@ -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)

View file

@ -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}")

View file

@ -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)}"
)

View file

@ -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

View file

@ -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
# ============================================================================

File diff suppressed because it is too large Load diff