diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database
index 0e804cbfd12..9a4e9a315ea 100644
--- a/docker/Dockerfile.database
+++ b/docker/Dockerfile.database
@@ -48,7 +48,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
# Install runtime dependencies
-RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip
+RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip libsndfile
WORKDIR /app
# Copy the current directory contents into the container at /app
diff --git a/docs/my-website/docs/mcp_control.md b/docs/my-website/docs/mcp_control.md
index f6ab4a73087..a7d66a6b7fc 100644
--- a/docs/my-website/docs/mcp_control.md
+++ b/docs/my-website/docs/mcp_control.md
@@ -108,7 +108,7 @@ Some MCP servers are meant to be shared broadly—think internal knowledge bases
3. Toggle **Allow All LiteLLM Keys** on.
@@ -634,3 +634,18 @@ Control which tools different teams can access from the same MCP server. For exa
This video shows how to set allowed tools for a Key, Team, or Organization.
+
+
+## Dashboard View Modes
+
+Proxy admins can also control what non-admins see inside the MCP dashboard via `general_settings.user_mcp_management_mode`:
+
+- `restricted` *(default)* – users only see servers that their team explicitly has access to.
+- `view_all` – every dashboard user can see the full MCP server list.
+
+```yaml title="Config example"
+general_settings:
+ user_mcp_management_mode: view_all
+```
+
+This is useful when you want discoverability for MCP offerings without granting additional execution privileges.
diff --git a/docs/my-website/docs/providers/gigachat.md b/docs/my-website/docs/providers/gigachat.md
new file mode 100644
index 00000000000..13eec298c25
--- /dev/null
+++ b/docs/my-website/docs/providers/gigachat.md
@@ -0,0 +1,283 @@
+import Tabs from '@theme/Tabs';
+import TabItem from '@theme/TabItem';
+
+# GigaChat
+https://developers.sber.ru/docs/ru/gigachat/api/overview
+
+GigaChat is Sber AI's large language model, Russia's leading LLM provider.
+
+:::tip
+
+**We support ALL GigaChat models, just set `model=gigachat/` as a prefix when sending litellm requests**
+
+:::
+
+:::warning
+
+GigaChat API uses self-signed SSL certificates. You must pass `ssl_verify=False` in your requests.
+
+:::
+
+## Supported Features
+
+| Feature | Supported |
+|---------|-----------|
+| Chat Completion | Yes |
+| Streaming | Yes |
+| Async | Yes |
+| Function Calling / Tools | Yes |
+| Structured Output (JSON Schema) | Yes (via function call emulation) |
+| Image Input | Yes (base64 and URL) - GigaChat-2-Max, GigaChat-2-Pro only |
+| Embeddings | Yes |
+
+## API Key
+
+GigaChat uses OAuth authentication. Set your credentials as environment variables:
+
+```python
+import os
+
+# Required: Set credentials (base64-encoded client_id:client_secret)
+os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
+
+# Optional: Set scope (default is GIGACHAT_API_PERS for personal use)
+os.environ['GIGACHAT_SCOPE'] = "GIGACHAT_API_PERS" # or GIGACHAT_API_B2B for business
+```
+
+Get your credentials at: https://developers.sber.ru/studio/
+
+## Sample Usage
+
+```python
+from litellm import completion
+import os
+
+os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
+
+response = completion(
+ model="gigachat/GigaChat-2-Max",
+ messages=[
+ {"role": "user", "content": "Hello from LiteLLM!"}
+ ],
+ ssl_verify=False, # Required for GigaChat
+)
+print(response)
+```
+
+## Sample Usage - Streaming
+
+```python
+from litellm import completion
+import os
+
+os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
+
+response = completion(
+ model="gigachat/GigaChat-2-Max",
+ messages=[
+ {"role": "user", "content": "Hello from LiteLLM!"}
+ ],
+ stream=True,
+ ssl_verify=False, # Required for GigaChat
+)
+
+for chunk in response:
+ print(chunk)
+```
+
+## Sample Usage - Function Calling
+
+```python
+from litellm import completion
+import os
+
+os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
+
+tools = [{
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "description": "Get weather for a city",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "city": {"type": "string", "description": "City name"}
+ },
+ "required": ["city"]
+ }
+ }
+}]
+
+response = completion(
+ model="gigachat/GigaChat-2-Max",
+ messages=[{"role": "user", "content": "What's the weather in Moscow?"}],
+ tools=tools,
+ ssl_verify=False, # Required for GigaChat
+)
+print(response)
+```
+
+## Sample Usage - Structured Output
+
+GigaChat supports structured output via JSON schema (emulated through function calling):
+
+```python
+from litellm import completion
+import os
+
+os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
+
+response = completion(
+ model="gigachat/GigaChat-2-Max",
+ messages=[{"role": "user", "content": "Extract info: John is 30 years old"}],
+ response_format={
+ "type": "json_schema",
+ "json_schema": {
+ "name": "person",
+ "schema": {
+ "type": "object",
+ "properties": {
+ "name": {"type": "string"},
+ "age": {"type": "integer"}
+ }
+ }
+ }
+ },
+ ssl_verify=False, # Required for GigaChat
+)
+print(response) # Returns JSON: {"name": "John", "age": 30}
+```
+
+## Sample Usage - Image Input
+
+GigaChat supports image input via base64 or URL (GigaChat-2-Max and GigaChat-2-Pro only):
+
+```python
+from litellm import completion
+import os
+
+os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
+
+response = completion(
+ model="gigachat/GigaChat-2-Max", # Vision requires GigaChat-2-Max or GigaChat-2-Pro
+ messages=[{
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "What's in this image?"},
+ {"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}
+ ]
+ }],
+ ssl_verify=False, # Required for GigaChat
+)
+print(response)
+```
+
+## Sample Usage - Embeddings
+
+```python
+from litellm import embedding
+import os
+
+os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
+
+response = embedding(
+ model="gigachat/Embeddings",
+ input=["Hello world", "How are you?"],
+ ssl_verify=False, # Required for GigaChat
+)
+print(response)
+```
+
+## Usage with LiteLLM Proxy
+
+### 1. Set GigaChat Models on config.yaml
+
+```yaml
+model_list:
+ - model_name: gigachat
+ litellm_params:
+ model: gigachat/GigaChat-2-Max
+ api_key: "os.environ/GIGACHAT_CREDENTIALS"
+ ssl_verify: false
+ - model_name: gigachat-lite
+ litellm_params:
+ model: gigachat/GigaChat-2-Lite
+ api_key: "os.environ/GIGACHAT_CREDENTIALS"
+ ssl_verify: false
+ - model_name: gigachat-embeddings
+ litellm_params:
+ model: gigachat/Embeddings
+ api_key: "os.environ/GIGACHAT_CREDENTIALS"
+ ssl_verify: false
+```
+
+### 2. Start Proxy
+
+```bash
+litellm --config config.yaml
+```
+
+### 3. Test it
+
+
+
+
+```shell
+curl --location 'http://0.0.0.0:4000/chat/completions' \
+--header 'Content-Type: application/json' \
+--data '{
+ "model": "gigachat",
+ "messages": [
+ {
+ "role": "user",
+ "content": "Hello!"
+ }
+ ]
+}'
+```
+
+
+
+```python
+import openai
+client = openai.OpenAI(
+ api_key="anything",
+ base_url="http://0.0.0.0:4000"
+)
+
+response = client.chat.completions.create(
+ model="gigachat",
+ messages=[{"role": "user", "content": "Hello!"}]
+)
+print(response)
+```
+
+
+
+## Supported Models
+
+### Chat Models
+
+| Model Name | Context Window | Vision | Description |
+|------------|----------------|--------|-------------|
+| gigachat/GigaChat-2-Lite | 128K | No | Fast, lightweight model |
+| gigachat/GigaChat-2-Pro | 128K | Yes | Professional model with vision |
+| gigachat/GigaChat-2-Max | 128K | Yes | Maximum capability model |
+
+### Embedding Models
+
+| Model Name | Max Input | Dimensions | Description |
+|------------|-----------|------------|-------------|
+| gigachat/Embeddings | 512 | 1024 | Standard embeddings |
+| gigachat/Embeddings-2 | 512 | 1024 | Updated embeddings |
+| gigachat/EmbeddingsGigaR | 4096 | 2560 | High-dimensional embeddings |
+
+:::note
+Available models may vary depending on your API access level (personal or business).
+:::
+
+## Limitations
+
+- Only one function call per request (GigaChat API limitation)
+- Maximum 1 image per message, 10 images total per conversation
+- GigaChat API uses self-signed SSL certificates - `ssl_verify=False` is required
diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md
index 87064d442ae..f4359a86ba9 100644
--- a/docs/my-website/docs/proxy/config_settings.md
+++ b/docs/my-website/docs/proxy/config_settings.md
@@ -111,6 +111,7 @@ general_settings:
master_key: string
maximum_spend_logs_retention_period: 30d # The maximum time to retain spend logs before deletion.
maximum_spend_logs_retention_interval: 1d # interval in which the spend log cleanup task should run in.
+ user_mcp_management_mode: restricted # or "view_all"
# Database Settings
database_url: string
@@ -230,6 +231,7 @@ router_settings:
| image_generation_model | str | The default model to use for image generation - ignores model set in request |
| store_model_in_db | boolean | If true, enables storing model + credential information in the DB. |
| supported_db_objects | List[str] | Fine-grained control over which object types to load from the database when `store_model_in_db` is True. Available types: `"models"`, `"mcp"`, `"guardrails"`, `"vector_stores"`, `"pass_through_endpoints"`, `"prompts"`, `"model_cost_map"`. If not set, all object types are loaded (default behavior). Example: `supported_db_objects: ["mcp"]` to only load MCP servers from DB. |
+| user_mcp_management_mode | string | Controls what non-admins can see on the MCP dashboard. `restricted` (default) only lists MCP servers that the user’s teams are explicitly allowed to access. `view_all` lets every user see the full MCP server list. Tool list/call always respects per-key permissions, so users still cannot run MCP calls without access. |
| store_prompts_in_spend_logs | boolean | If true, allows prompts and responses to be stored in the spend logs table. |
| max_request_size_mb | int | The maximum size for requests in MB. Requests above this size will be rejected. |
| max_response_size_mb | int | The maximum size for responses in MB. LLM Responses above this size will not be sent. |
@@ -669,6 +671,7 @@ router_settings:
| LANGSMITH_DEFAULT_RUN_NAME | Default name for Langsmith run
| LANGSMITH_PROJECT | Project name for Langsmith integration
| LANGSMITH_SAMPLING_RATE | Sampling rate for Langsmith logging
+| LANGSMITH_TENANT_ID | Tenant ID for Langsmith multi-tenant deployments
| LANGTRACE_API_KEY | API key for Langtrace service
| LASSO_API_BASE | Base URL for Lasso API
| LASSO_API_KEY | API key for Lasso service
@@ -707,6 +710,7 @@ router_settings:
| LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development)
| LITELLM_NON_ROOT | Flag to run LiteLLM in non-root mode for enhanced security in Docker containers
| LITELLM_RATE_LIMIT_WINDOW_SIZE | Rate limit window size for LiteLLM. Default is 60
+| LITELLM_REASONING_AUTO_SUMMARY | If set to "true", automatically enables detailed reasoning summaries for reasoning models (e.g., o1, o3-mini, deepseek-reasoner). When enabled, adds `summary: "detailed"` to reasoning effort configurations. Default is "false"
| LITELLM_SALT_KEY | Salt key for encryption in LiteLLM
| LITELLM_SSL_CIPHERS | SSL/TLS cipher configuration for faster handshakes. Controls cipher suite preferences for OpenSSL connections.
| LITELLM_SECRET_AWS_KMS_LITELLM_LICENSE | AWS KMS encrypted license for LiteLLM
@@ -774,6 +778,7 @@ router_settings:
| OTEL_EXPORTER_OTLP_HEADERS | Headers for OpenTelemetry requests
| OTEL_SERVICE_NAME | Service name identifier for OpenTelemetry
| OTEL_TRACER_NAME | Tracer name for OpenTelemetry tracing
+| OTEL_LOGS_EXPORTER | Exporter type for OpenTelemetry logs (e.g., console)
| PAGERDUTY_API_KEY | API key for PagerDuty Alerting
| PANW_PRISMA_AIRS_API_KEY | API key for PANW Prisma AIRS service
| PANW_PRISMA_AIRS_API_BASE | Base URL for PANW Prisma AIRS service
@@ -888,4 +893,4 @@ router_settings:
| DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL | Time-to-live in seconds for health check lock in shared health check mode. Default is 60 (1 minute)
| ZSCALER_AI_GUARD_API_KEY | API key for Zscaler AI Guard service
| ZSCALER_AI_GUARD_POLICY_ID | Policy ID for Zscaler AI Guard guardrails
-| ZSCALER_AI_GUARD_URL | Base URL for Zscaler AI Guard API. Default is https://api.us1.zseclipse.net/v1/detection/execute-policy
\ No newline at end of file
+| ZSCALER_AI_GUARD_URL | Base URL for Zscaler AI Guard API. Default is https://api.us1.zseclipse.net/v1/detection/execute-policy
diff --git a/docs/my-website/docs/reasoning_content.md b/docs/my-website/docs/reasoning_content.md
index fca3df638c7..04c6d7ee6cc 100644
--- a/docs/my-website/docs/reasoning_content.md
+++ b/docs/my-website/docs/reasoning_content.md
@@ -591,3 +591,68 @@ Expected Response
+
+## OpenAI Responses API - Auto-Summary Control
+
+When using OpenAI Responses API models (like `gpt-5`) via `/chat/completions` with `reasoning_effort`, you can control whether `summary="detailed"` is automatically added to the reasoning parameter.
+
+### Enabling Auto-Summary
+
+You can enable automatic `summary="detailed"` in two ways:
+
+
+
+
+```python
+import litellm
+
+# Enable auto-summary globally
+litellm.reasoning_auto_summary = True
+
+response = litellm.completion(
+ model="openai/responses/gpt-5-mini",
+ messages=[{"role": "user", "content": "What is the capital of France?"}],
+ reasoning_effort="low", # Will automatically add summary="detailed"
+)
+```
+
+
+
+
+
+```bash
+# Set environment variable
+export LITELLM_REASONING_AUTO_SUMMARY=true
+
+# Or in your .env file
+LITELLM_REASONING_AUTO_SUMMARY=true
+```
+
+
+
+
+
+```yaml
+litellm_settings:
+ reasoning_auto_summary: true # Enable auto-summary for all requests
+
+model_list:
+ - model_name: gpt-5-mini
+ litellm_params:
+ model: openai/responses/gpt-5-mini
+```
+
+
+
+
+### Manual Control (Recommended)
+
+For fine-grained control, pass `reasoning_effort` as a dictionary:
+
+```python
+response = litellm.completion(
+ model="openai/responses/gpt-5-mini",
+ messages=[{"role": "user", "content": "What is the capital of France?"}],
+ reasoning_effort={"effort": "low", "summary": "detailed"}, # Explicit control
+)
+```
diff --git a/docs/my-website/docs/response_api_compact.md b/docs/my-website/docs/response_api_compact.md
new file mode 100644
index 00000000000..f5caa32ea33
--- /dev/null
+++ b/docs/my-website/docs/response_api_compact.md
@@ -0,0 +1,104 @@
+import Tabs from '@theme/Tabs';
+import TabItem from '@theme/TabItem';
+
+# /responses/compact
+
+Compress conversation history using OpenAI's `/responses/compact` endpoint.
+
+| Feature | Supported |
+|---------|-----------|
+| Supported LiteLLM Versions | 1.72.0+ |
+| Supported Providers | `openai` |
+
+## Usage
+
+### LiteLLM Python SDK
+
+```python showLineNumbers title="Compact Response"
+import litellm
+
+response = litellm.compact_responses(
+ model="openai/gpt-4o",
+ input=[{"role": "user", "content": "Hello, how are you?"}],
+ instructions="Be helpful",
+ previous_response_id="resp_abc123" # optional
+)
+
+print(response.id)
+print(response.object) # "response.compaction"
+print(response.output)
+```
+
+### LiteLLM Proxy
+
+
+
+
+```bash showLineNumbers title="Compact Request"
+curl http://localhost:4000/v1/responses/compact \
+ -H "Content-Type: application/json" \
+ -H "Authorization: Bearer sk-1234" \
+ -d '{
+ "model": "openai/gpt-4o",
+ "input": [{"role": "user", "content": "Hello"}],
+ "instructions": "Be helpful"
+ }'
+```
+
+
+
+
+```python showLineNumbers title="Compact with OpenAI SDK"
+import httpx
+
+response = httpx.post(
+ "http://localhost:4000/v1/responses/compact",
+ headers={"Authorization": "Bearer sk-1234"},
+ json={
+ "model": "openai/gpt-4o",
+ "input": [{"role": "user", "content": "Hello"}],
+ "instructions": "Be helpful"
+ }
+)
+
+print(response.json())
+```
+
+
+
+
+## Request Parameters
+
+| Parameter | Type | Required | Description |
+|-----------|------|----------|-------------|
+| `model` | string | Yes | Model to use for compaction |
+| `input` | string or array | Yes | Input messages to compact |
+| `instructions` | string | No | System instructions |
+| `previous_response_id` | string | No | ID of previous response to continue from |
+
+## Response Format
+
+```json
+{
+ "id": "resp_abc123",
+ "object": "response.compaction",
+ "created_at": 1734366691,
+ "output": [
+ {
+ "type": "message",
+ "role": "assistant",
+ "content": [...]
+ },
+ {
+ "type": "compaction",
+ "encrypted_content": "..."
+ }
+ ],
+ "usage": {
+ "input_tokens": 100,
+ "output_tokens": 50,
+ "total_tokens": 150
+ }
+}
+```
+
diff --git a/docs/my-website/img/mcp_allow_all_ui.png b/docs/my-website/img/mcp_allow_all_ui.png
new file mode 100644
index 00000000000..f074deb801e
Binary files /dev/null and b/docs/my-website/img/mcp_allow_all_ui.png differ
diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js
index 3793aec037f..11effc82fe7 100644
--- a/docs/my-website/sidebars.js
+++ b/docs/my-website/sidebars.js
@@ -541,7 +541,14 @@ const sidebars = {
},
"realtime",
"rerank",
- "response_api",
+ {
+ type: "category",
+ label: "/responses",
+ items: [
+ "response_api",
+ "response_api_compact",
+ ]
+ },
{
type: "category",
label: "/search",
diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py
index 7ffbe95be13..a1fa4ef8440 100644
--- a/litellm-proxy-extras/litellm_proxy_extras/utils.py
+++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py
@@ -10,6 +10,7 @@ from pathlib import Path
from typing import Optional
from litellm_proxy_extras._logging import logger
+from litellm.caching.redis_cache import RedisCache
def str_to_bool(value: Optional[str]) -> bool:
@@ -18,6 +19,103 @@ def str_to_bool(value: Optional[str]) -> bool:
return value.lower() in ("true", "1", "t", "y", "yes")
+class MigrationLockManager:
+ """Redis-based lock manager for database migrations"""
+
+ MIGRATION_LOCK_KEY = "migration_lock"
+ LOCK_TTL_SECONDS = 300 # 5 minutes TTL
+
+ def __init__(self, redis_cache: Optional[RedisCache] = None):
+ self.redis_cache = redis_cache
+ self.lock_acquired = False
+ self.pod_id = f"pod_{os.getpid()}_{int(time.time())}"
+
+ def _get_redis_lock_key(self) -> str:
+ """Get Redis lock key for migration"""
+ return f"migration_lock:{self.MIGRATION_LOCK_KEY}"
+
+ def acquire_lock(self) -> bool:
+ """Acquire migration lock"""
+ if self.redis_cache is None:
+ logger.warning(
+ "Redis cache is not available, running migration without lock protection"
+ )
+ self.lock_acquired = True
+ return True
+
+ try:
+ lock_key = self._get_redis_lock_key()
+
+ # Redis SET with NX (only if not exists) and EX (expiration)
+ acquired = self.redis_cache.set_cache(
+ key=lock_key, value=self.pod_id, nx=True, ttl=self.LOCK_TTL_SECONDS
+ )
+
+ if acquired:
+ self.lock_acquired = True
+ logger.info(f"Migration lock acquired by pod {self.pod_id}")
+ return True
+ else:
+ logger.info("Migration lock is already held by another pod")
+ return False
+
+ except Exception as e:
+ logger.warning(f"Failed to acquire migration lock: {e}")
+ return False
+
+ def wait_for_lock_release(
+ self, check_interval: int = 5, max_wait: int = 300
+ ) -> bool:
+ """Wait for another process to release the lock"""
+ if self.redis_cache is None:
+ logger.warning("Redis cache is not available, cannot wait for lock")
+ return False
+
+ logger.info(f"Waiting for migration lock to be released (max {max_wait}s)...")
+ start_time = time.time()
+
+ while time.time() - start_time < max_wait:
+ # Try to acquire lock using the public acquire_lock method
+ if self.acquire_lock():
+ logger.info(
+ f"Migration lock acquired after waiting by pod {self.pod_id}"
+ )
+ return True
+
+ time.sleep(check_interval)
+
+ logger.warning(f"Failed to acquire migration lock within {max_wait} seconds")
+ return False
+
+ def release_lock(self):
+ """Release migration lock"""
+ if not self.lock_acquired or self.redis_cache is None:
+ return
+
+ try:
+ lock_key = self._get_redis_lock_key()
+
+ # Verify current pod owns the lock
+ current_value = self.redis_cache.get_cache(lock_key)
+ if current_value and str(current_value) == self.pod_id:
+ self.redis_cache.delete_cache(lock_key)
+ logger.info(f"Migration lock released by pod {self.pod_id}")
+ else:
+ logger.warning(f"Pod {self.pod_id} cannot release lock (not owner)")
+
+ except Exception as e:
+ logger.warning(f"Failed to release migration lock: {e}")
+ finally:
+ self.lock_acquired = False
+
+ def __enter__(self):
+ """Context manager entry - acquire lock when entering with statement"""
+ self.acquire_lock()
+ return self
+
+ def __exit__(self, exc_type, exc_val, exc_tb):
+ """Context manager exit - release lock when exiting with statement"""
+ self.release_lock()
def _get_prisma_env() -> dict:
"""Get environment variables for Prisma, handling offline mode if configured."""
@@ -346,19 +444,50 @@ class ProxyExtrasDBManager:
)
@staticmethod
- def setup_database(use_migrate: bool = False) -> bool:
+ def setup_database(
+ use_migrate: bool = False, redis_cache: Optional[RedisCache] = None
+ ) -> bool:
"""
Set up the database using either prisma migrate or prisma db push
- Uses migrations from litellm-proxy-extras package
+ Uses migrations from litellm-proxy-extras package.
+ In multi-instance environment, use redis lock to prevent concurrent execution.
Args:
schema_path (str): Path to the Prisma schema file
use_migrate (bool): Whether to use prisma migrate instead of db push
+ redis_cache: Redis cache instance for distributed locking
Returns:
bool: True if setup was successful, False otherwise
"""
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
+
+ database_url = os.getenv("DATABASE_URL")
+ if not database_url:
+ logger.error("DATABASE_URL environment variable is not set")
+ return False
+
+ # Use MigrationLockManager to prevent concurrent migration execution
+ with MigrationLockManager(redis_cache) as lock_manager:
+ # Lock is already acquired in __enter__, check if it was successful
+ if not lock_manager.lock_acquired:
+ # Cannot acquire lock, another process is running migration
+ logger.info(
+ "Another pod is running migration, waiting for completion..."
+ )
+
+ # Wait for other process to complete migration
+ if not lock_manager.wait_for_lock_release():
+ logger.error("Failed to acquire migration lock after waiting")
+ return False
+
+ # Successfully acquired lock, proceed with migration
+ logger.info("Acquired migration lock, proceeding with migration")
+ return ProxyExtrasDBManager._execute_migration(use_migrate, schema_path)
+
+ @staticmethod
+ def _execute_migration(use_migrate: bool, schema_path: str) -> bool:
+ """Execute the actual migration"""
for attempt in range(4):
original_dir = os.getcwd()
migrations_dir = ProxyExtrasDBManager._get_prisma_dir()
diff --git a/litellm/__init__.py b/litellm/__init__.py
index dfe959bf747..7f7ee21f692 100644
--- a/litellm/__init__.py
+++ b/litellm/__init__.py
@@ -197,6 +197,7 @@ retry = True
api_key: Optional[str] = None
openai_key: Optional[str] = None
groq_key: Optional[str] = None
+gigachat_key: Optional[str] = None
databricks_key: Optional[str] = None
openai_like_key: Optional[str] = None
azure_key: Optional[str] = None
@@ -275,6 +276,7 @@ banned_keywords_list: Optional[Union[str, List]] = None
llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all"
guardrail_name_config_map: Dict[str, GuardrailItem] = {}
include_cost_in_streaming_usage: bool = False
+reasoning_auto_summary: bool = False
### PROMPTS ####
from litellm.types.prompts.init_prompts import PromptSpec
@@ -1440,6 +1442,8 @@ if TYPE_CHECKING:
from .llms.github_copilot.chat.transformation import GithubCopilotConfig as GithubCopilotConfig
from .llms.github_copilot.responses.transformation import GithubCopilotResponsesAPIConfig as GithubCopilotResponsesAPIConfig
from .llms.github_copilot.embedding.transformation import GithubCopilotEmbeddingConfig as GithubCopilotEmbeddingConfig
+ from .llms.gigachat.chat.transformation import GigaChatConfig as GigaChatConfig
+ from .llms.gigachat.embedding.transformation import GigaChatEmbeddingConfig as GigaChatEmbeddingConfig
from .llms.nebius.chat.transformation import NebiusConfig as NebiusConfig
from .llms.wandb.chat.transformation import WandbConfig as WandbConfig
from .llms.dashscope.chat.transformation import DashScopeChatConfig as DashScopeChatConfig
diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py
index 8c8266d5ca4..26133ebc222 100644
--- a/litellm/_lazy_imports_registry.py
+++ b/litellm/_lazy_imports_registry.py
@@ -255,6 +255,8 @@ LLM_CONFIG_NAMES = (
"GithubCopilotEmbeddingConfig",
"NebiusConfig",
"WandbConfig",
+ "GigaChatConfig",
+ "GigaChatEmbeddingConfig",
"DashScopeChatConfig",
"MoonshotChatConfig",
"DockerModelRunnerChatConfig",
@@ -644,6 +646,8 @@ _LLM_CONFIGS_IMPORT_MAP = {
"GithubCopilotEmbeddingConfig": (".llms.github_copilot.embedding.transformation", "GithubCopilotEmbeddingConfig"),
"NebiusConfig": (".llms.nebius.chat.transformation", "NebiusConfig"),
"WandbConfig": (".llms.wandb.chat.transformation", "WandbConfig"),
+ "GigaChatConfig": (".llms.gigachat.chat.transformation", "GigaChatConfig"),
+ "GigaChatEmbeddingConfig": (".llms.gigachat.embedding.transformation", "GigaChatEmbeddingConfig"),
"DashScopeChatConfig": (".llms.dashscope.chat.transformation", "DashScopeChatConfig"),
"MoonshotChatConfig": (".llms.moonshot.chat.transformation", "MoonshotChatConfig"),
"DockerModelRunnerChatConfig": (".llms.docker_model_runner.chat.transformation", "DockerModelRunnerChatConfig"),
diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py
index 5b206317b29..a89efc4e82b 100644
--- a/litellm/completion_extras/litellm_responses_transformation/transformation.py
+++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py
@@ -3,6 +3,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req
"""
import json
+import os
from typing import (
TYPE_CHECKING,
Any,
@@ -22,6 +23,7 @@ from typing import (
from openai.types.responses.tool_param import FunctionToolParam
from pydantic import BaseModel
+import litellm
from litellm import ModelResponse
from litellm._logging import verbose_logger
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
@@ -691,19 +693,26 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
if isinstance(reasoning_effort, dict):
return Reasoning(**reasoning_effort) # type: ignore[typeddict-item]
- # If string is passed, map with summary="detailed"
+ # Check if auto-summary is enabled via flag or environment variable
+ # Priority: litellm.reasoning_auto_summary flag > LITELLM_REASONING_AUTO_SUMMARY env var
+ auto_summary_enabled = (
+ litellm.reasoning_auto_summary
+ or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
+ )
+
+ # If string is passed, map with optional summary based on flag/env var
if reasoning_effort == "none":
- return Reasoning(effort="none", summary="detailed") # type: ignore
+ return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") # type: ignore
elif reasoning_effort == "high":
- return Reasoning(effort="high", summary="detailed")
+ return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high")
elif reasoning_effort == "xhigh":
- return Reasoning(effort="xhigh", summary="detailed") # type: ignore[typeddict-item]
+ return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") # type: ignore[typeddict-item]
elif reasoning_effort == "medium":
- return Reasoning(effort="medium", summary="detailed")
+ return Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium")
elif reasoning_effort == "low":
- return Reasoning(effort="low", summary="detailed")
+ return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low")
elif reasoning_effort == "minimal":
- return Reasoning(effort="minimal", summary="detailed")
+ return Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal")
return None
def _transform_response_format_to_text_format(
diff --git a/litellm/constants.py b/litellm/constants.py
index e8524a87c41..1cd2da549ca 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -375,6 +375,7 @@ LITELLM_CHAT_PROVIDERS = [
"perplexity",
"mistral",
"groq",
+ "gigachat",
"nvidia_nim",
"cerebras",
"baseten",
diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json
index 88f7908e9a2..6b30b6b736e 100644
--- a/litellm/integrations/callback_configs.json
+++ b/litellm/integrations/callback_configs.json
@@ -187,6 +187,12 @@
"ui_name": "Sampling Rate",
"description": "Sampling rate for logging (0.0 to 1.0, default: 1.0)",
"required": false
+ },
+ "langsmith_tenant_id": {
+ "type": "text",
+ "ui_name": "Tenant ID",
+ "description": "LangSmith tenant ID for organization-scoped API keys (required when using org-scoped keys)",
+ "required": false
}
},
"description": "Langsmith Logging Integration"
diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py
index f0f355b4895..7e62613a7e4 100644
--- a/litellm/integrations/langfuse/langfuse.py
+++ b/litellm/integrations/langfuse/langfuse.py
@@ -50,6 +50,42 @@ else:
Langfuse = Any
+def _extract_cache_read_input_tokens(usage_obj) -> int:
+ """
+ Extract cache_read_input_tokens from usage object.
+
+ Checks both:
+ 1. Top-level cache_read_input_tokens (Anthropic format)
+ 2. prompt_tokens_details.cached_tokens (Gemini, OpenAI format)
+
+ See: https://github.com/BerriAI/litellm/issues/18520
+
+ Args:
+ usage_obj: Usage object from LLM response
+
+ Returns:
+ int: Number of cached tokens read, defaults to 0
+ """
+ cache_read_input_tokens = usage_obj.get("cache_read_input_tokens") or 0
+
+ # Check prompt_tokens_details.cached_tokens (used by Gemini and other providers)
+ if hasattr(usage_obj, "prompt_tokens_details"):
+ prompt_tokens_details = getattr(usage_obj, "prompt_tokens_details", None)
+ if (
+ prompt_tokens_details is not None
+ and hasattr(prompt_tokens_details, "cached_tokens")
+ ):
+ cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
+ if (
+ cached_tokens is not None
+ and isinstance(cached_tokens, (int, float))
+ and cached_tokens > 0
+ ):
+ cache_read_input_tokens = cached_tokens
+
+ return cache_read_input_tokens
+
+
class LangFuseLogger:
# Class variables or attributes
def __init__(
@@ -757,8 +793,8 @@ class LangFuseLogger:
cache_creation_input_tokens = (
_usage_obj.get("cache_creation_input_tokens") or 0
)
- cache_read_input_tokens = (
- _usage_obj.get("cache_read_input_tokens") or 0
+ cache_read_input_tokens = _extract_cache_read_input_tokens(
+ _usage_obj
)
usage = {
diff --git a/litellm/integrations/langsmith.py b/litellm/integrations/langsmith.py
index cc9b361b69d..570b78f2927 100644
--- a/litellm/integrations/langsmith.py
+++ b/litellm/integrations/langsmith.py
@@ -40,6 +40,7 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_project: Optional[str] = None,
langsmith_base_url: Optional[str] = None,
langsmith_sampling_rate: Optional[float] = None,
+ langsmith_tenant_id: Optional[str] = None,
**kwargs,
):
self.flush_lock = asyncio.Lock()
@@ -48,6 +49,7 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_api_key=langsmith_api_key,
langsmith_project=langsmith_project,
langsmith_base_url=langsmith_base_url,
+ langsmith_tenant_id=langsmith_tenant_id,
)
self.sampling_rate: float = (
langsmith_sampling_rate
@@ -76,6 +78,7 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_api_key: Optional[str] = None,
langsmith_project: Optional[str] = None,
langsmith_base_url: Optional[str] = None,
+ langsmith_tenant_id: Optional[str] = None,
) -> LangsmithCredentialsObject:
_credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY")
_credentials_project = (
@@ -86,11 +89,13 @@ class LangsmithLogger(CustomBatchLogger):
or os.getenv("LANGSMITH_BASE_URL")
or "https://api.smith.langchain.com"
)
+ _credentials_tenant_id = langsmith_tenant_id or os.getenv("LANGSMITH_TENANT_ID")
return LangsmithCredentialsObject(
LANGSMITH_API_KEY=_credentials_api_key,
LANGSMITH_BASE_URL=_credentials_base_url,
LANGSMITH_PROJECT=_credentials_project,
+ LANGSMITH_TENANT_ID=_credentials_tenant_id,
)
def _prepare_log_data(
@@ -365,8 +370,11 @@ class LangsmithLogger(CustomBatchLogger):
"""
langsmith_api_base = credentials["LANGSMITH_BASE_URL"]
langsmith_api_key = credentials["LANGSMITH_API_KEY"]
+ langsmith_tenant_id = credentials.get("LANGSMITH_TENANT_ID")
url = self._add_endpoint_to_url(langsmith_api_base, "runs/batch")
headers = {"x-api-key": langsmith_api_key}
+ if langsmith_tenant_id:
+ headers["x-tenant-id"] = langsmith_tenant_id
elements_to_log = [queue_object["data"] for queue_object in queue_objects]
try:
@@ -418,6 +426,7 @@ class LangsmithLogger(CustomBatchLogger):
api_key=credentials["LANGSMITH_API_KEY"],
project=credentials["LANGSMITH_PROJECT"],
base_url=credentials["LANGSMITH_BASE_URL"],
+ tenant_id=credentials.get("LANGSMITH_TENANT_ID"),
)
if key not in log_queue_by_credentials:
@@ -466,6 +475,9 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_base_url=standard_callback_dynamic_params.get(
"langsmith_base_url", None
),
+ langsmith_tenant_id=standard_callback_dynamic_params.get(
+ "langsmith_tenant_id", None
+ ),
)
else:
credentials = self.default_credentials
@@ -491,13 +503,16 @@ class LangsmithLogger(CustomBatchLogger):
def get_run_by_id(self, run_id):
langsmith_api_key = self.default_credentials["LANGSMITH_API_KEY"]
-
langsmith_api_base = self.default_credentials["LANGSMITH_BASE_URL"]
+ langsmith_tenant_id = self.default_credentials.get("LANGSMITH_TENANT_ID")
url = f"{langsmith_api_base}/runs/{run_id}"
+ headers = {"x-api-key": langsmith_api_key}
+ if langsmith_tenant_id:
+ headers["x-tenant-id"] = langsmith_tenant_id
response = litellm.module_level_client.get(
url=url,
- headers={"x-api-key": langsmith_api_key},
+ headers=headers,
)
return response.json()
diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py
index 12e60bc25bb..a7d2326d938 100644
--- a/litellm/integrations/opentelemetry.py
+++ b/litellm/integrations/opentelemetry.py
@@ -196,50 +196,88 @@ class OpenTelemetry(CustomLogger):
litellm.service_callback.append(self)
setattr(proxy_server, "open_telemetry_logger", self)
+ def _get_or_create_provider(
+ self,
+ provider,
+ provider_name: str,
+ get_existing_provider_fn,
+ sdk_provider_class,
+ create_new_provider_fn,
+ set_provider_fn,
+ ):
+ """
+ Generic helper to get or create an OpenTelemetry provider (Tracer, Meter, or Logger).
+
+ Args:
+ provider: The provider instance passed to the init function (can be None)
+ provider_name: Name for logging (e.g., "TracerProvider")
+ get_existing_provider_fn: Function to get the existing global provider
+ sdk_provider_class: The SDK provider class to check for (e.g., TracerProvider from SDK)
+ create_new_provider_fn: Function to create a new provider instance
+ set_provider_fn: Function to set the provider globally
+
+ Returns:
+ The provider to use (either existing, new, or explicitly provided)
+ """
+ if provider is not None:
+ # Provider explicitly provided (e.g., for testing)
+ # Do NOT call set_provider_fn - the caller is responsible for managing global state
+ # If they want it to be global, they've already set it before passing it to us
+ verbose_logger.debug(
+ "OpenTelemetry: Using provided TracerProvider: %s",
+ type(provider).__name__,
+ )
+ return provider
+
+ # Check if a provider is already set globally
+ try:
+ existing_provider = get_existing_provider_fn()
+
+ # If a real SDK provider exists (set by another SDK like Langfuse), use it
+ # This uses a positive check for SDK providers instead of a negative check for proxy providers
+ if isinstance(existing_provider, sdk_provider_class):
+ verbose_logger.debug(
+ "OpenTelemetry: Using existing %s: %s",
+ provider_name,
+ type(existing_provider).__name__,
+ )
+ provider = existing_provider
+ # Don't call set_provider to preserve existing context
+ else:
+ # Default proxy provider or unknown type, create our own
+ verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name)
+ provider = create_new_provider_fn()
+ set_provider_fn(provider)
+ except Exception as e:
+ # Fallback: create a new provider if something goes wrong
+ verbose_logger.debug(
+ "OpenTelemetry: Exception checking existing %s, creating new one: %s",
+ provider_name,
+ str(e),
+ )
+ provider = create_new_provider_fn()
+ set_provider_fn(provider)
+
+ return provider
+
def _init_tracing(self, tracer_provider):
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.trace import SpanKind
- # use provided tracer or create a new one
- if tracer_provider is None:
- # Check if a TracerProvider is already set globally (e.g., by Langfuse SDK)
- try:
- from opentelemetry.trace import ProxyTracerProvider
+ def create_tracer_provider():
+ provider = TracerProvider(resource=_get_litellm_resource())
+ provider.add_span_processor(self._get_span_processor())
+ return provider
- existing_provider = trace.get_tracer_provider()
-
- # If an actual provider exists (not the default proxy), use it
- if not isinstance(existing_provider, ProxyTracerProvider):
- verbose_logger.debug(
- "OpenTelemetry: Using existing TracerProvider: %s",
- type(existing_provider).__name__,
- )
- tracer_provider = existing_provider
- # Don't call set_tracer_provider to preserve existing context
- else:
- # No real provider exists yet, create our own
- verbose_logger.debug("OpenTelemetry: Creating new TracerProvider")
- tracer_provider = TracerProvider(resource=_get_litellm_resource())
- tracer_provider.add_span_processor(self._get_span_processor())
- trace.set_tracer_provider(tracer_provider)
- except Exception as e:
- # Fallback: create a new provider if something goes wrong
- verbose_logger.debug(
- "OpenTelemetry: Exception checking existing provider, creating new one: %s",
- str(e),
- )
- tracer_provider = TracerProvider(resource=_get_litellm_resource())
- tracer_provider.add_span_processor(self._get_span_processor())
- trace.set_tracer_provider(tracer_provider)
- else:
- # Tracer provider explicitly provided (e.g., for testing)
- # Do NOT call set_tracer_provider - the caller is responsible for managing global state
- # If they want it to be global, they've already set it before passing it to us
- verbose_logger.debug(
- "OpenTelemetry: Using provided TracerProvider: %s",
- type(tracer_provider).__name__,
- )
+ tracer_provider = self._get_or_create_provider(
+ provider=tracer_provider,
+ provider_name="TracerProvider",
+ get_existing_provider_fn=trace.get_tracer_provider,
+ sdk_provider_class=TracerProvider,
+ create_new_provider_fn=create_tracer_provider,
+ set_provider_fn=trace.set_tracer_provider,
+ )
# Grab our tracer from the TracerProvider (not from global context)
# This ensures we use the provided TracerProvider (e.g., for testing)
@@ -257,39 +295,24 @@ class OpenTelemetry(CustomLogger):
return
from opentelemetry import metrics
- from opentelemetry.sdk.metrics import Histogram, MeterProvider
+ from opentelemetry.sdk.metrics import MeterProvider
- # Only create OTLP infrastructure if no custom meter provider is provided
- if meter_provider is None:
- from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
- OTLPMetricExporter,
- )
- from opentelemetry.sdk.metrics.export import (
- AggregationTemporality,
- PeriodicExportingMetricReader,
+ def create_meter_provider():
+ metric_reader = self._get_metric_reader()
+ return MeterProvider(
+ metric_readers=[metric_reader], resource=_get_litellm_resource()
)
- normalized_endpoint = self._normalize_otel_endpoint(
- self.config.endpoint, "metrics"
- )
- _metric_exporter = OTLPMetricExporter(
- endpoint=normalized_endpoint,
- headers=OpenTelemetry._get_headers_dictionary(self.config.headers),
- preferred_temporality={Histogram: AggregationTemporality.DELTA},
- )
- _metric_reader = PeriodicExportingMetricReader(
- _metric_exporter, export_interval_millis=10000
- )
+ meter_provider = self._get_or_create_provider(
+ provider=meter_provider,
+ provider_name="MeterProvider",
+ get_existing_provider_fn=metrics.get_meter_provider,
+ sdk_provider_class=MeterProvider,
+ create_new_provider_fn=create_meter_provider,
+ set_provider_fn=metrics.set_meter_provider,
+ )
- meter_provider = MeterProvider(
- metric_readers=[_metric_reader], resource=_get_litellm_resource()
- )
- meter = meter_provider.get_meter(__name__)
- else:
- # Use the provided meter provider as-is, without creating additional OTLP infrastructure
- meter = meter_provider.get_meter(__name__)
-
- metrics.set_meter_provider(meter_provider)
+ meter = meter_provider.get_meter(__name__)
self._operation_duration_histogram = meter.create_histogram(
name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38
@@ -327,22 +350,26 @@ class OpenTelemetry(CustomLogger):
if not self.config.enable_events:
return
- from opentelemetry._logs import set_logger_provider
+ from opentelemetry._logs import get_logger_provider, set_logger_provider
from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
- # set up log pipeline
- if logger_provider is None:
- litellm_resource = _get_litellm_resource()
- logger_provider = OTLoggerProvider(resource=litellm_resource)
- # Only add OTLP exporter if we created the logger provider ourselves
+ def create_logger_provider():
+ provider = OTLoggerProvider(resource=_get_litellm_resource())
log_exporter = self._get_log_exporter()
- if log_exporter:
- logger_provider.add_log_record_processor(
- BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
- )
+ provider.add_log_record_processor(
+ BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
+ )
+ return provider
- set_logger_provider(logger_provider)
+ self._get_or_create_provider(
+ provider=logger_provider,
+ provider_name="LoggerProvider",
+ get_existing_provider_fn=get_logger_provider,
+ sdk_provider_class=OTLoggerProvider,
+ create_new_provider_fn=create_logger_provider,
+ set_provider_fn=set_logger_provider,
+ )
def log_success_event(self, kwargs, response_obj, start_time, end_time):
self._handle_success(kwargs, response_obj, start_time, end_time)
@@ -944,6 +971,15 @@ class OpenTelemetry(CustomLogger):
if not self.config.enable_events:
return
+ # NOTE: Semantic logs (gen_ai.content.prompt/completion events) have compatibility issues
+ # with OTEL SDK >= 1.39.0 due to breaking changes in PR #4676:
+ # - LogRecord moved from opentelemetry.sdk._logs to opentelemetry.sdk._logs._internal
+ # - LogRecord constructor no longer accepts 'resource' parameter (now inherited from LoggerProvider)
+ # - LogData class was removed entirely
+ # These logs work correctly in OTEL SDK < 1.39.0 but may fail in >= 1.39.0.
+ # See: https://github.com/open-telemetry/opentelemetry-python/pull/4676
+ # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords
+
from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider
from opentelemetry.sdk._logs import LogRecord as SdkLogRecord
@@ -1807,7 +1843,8 @@ class OpenTelemetry(CustomLogger):
)
return self.OTEL_EXPORTER
- if self.OTEL_EXPORTER == "console":
+ otel_logs_exporter = os.getenv("OTEL_LOGS_EXPORTER")
+ if self.OTEL_EXPORTER == "console" or otel_logs_exporter == "console":
from opentelemetry.sdk._logs.export import ConsoleLogExporter
verbose_logger.debug(
@@ -1854,6 +1891,67 @@ class OpenTelemetry(CustomLogger):
return ConsoleLogExporter()
+ def _get_metric_reader(self):
+ """
+ Get the appropriate metric reader based on the configuration.
+ """
+ from opentelemetry.sdk.metrics import Histogram
+ from opentelemetry.sdk.metrics.export import (
+ AggregationTemporality,
+ ConsoleMetricExporter,
+ PeriodicExportingMetricReader,
+ )
+
+ verbose_logger.debug(
+ "OpenTelemetry Logger, initializing metric reader\nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
+ self.OTEL_EXPORTER,
+ self.OTEL_ENDPOINT,
+ self.OTEL_HEADERS,
+ )
+
+ _split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)
+ normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "metrics")
+
+ if self.OTEL_EXPORTER == "console":
+ exporter = ConsoleMetricExporter()
+ return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
+
+ elif (
+ self.OTEL_EXPORTER == "otlp_http"
+ or self.OTEL_EXPORTER == "http/protobuf"
+ or self.OTEL_EXPORTER == "http/json"
+ ):
+ from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
+ OTLPMetricExporter,
+ )
+
+ exporter = OTLPMetricExporter(
+ endpoint=normalized_endpoint,
+ headers=_split_otel_headers,
+ preferred_temporality={Histogram: AggregationTemporality.DELTA},
+ )
+ return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
+
+ elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc":
+ from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
+ OTLPMetricExporter,
+ )
+
+ exporter = OTLPMetricExporter(
+ endpoint=normalized_endpoint,
+ headers=_split_otel_headers,
+ preferred_temporality={Histogram: AggregationTemporality.DELTA},
+ )
+ return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
+
+ else:
+ verbose_logger.warning(
+ "OpenTelemetry: Unknown metric exporter '%s', defaulting to console. Supported: console, otlp_http, otlp_grpc",
+ self.OTEL_EXPORTER,
+ )
+ exporter = ConsoleMetricExporter()
+ return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
+
def _normalize_otel_endpoint(
self, endpoint: Optional[str], signal_type: str
) -> Optional[str]:
diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py
index d92af417175..6baaae7ae3f 100644
--- a/litellm/litellm_core_utils/streaming_handler.py
+++ b/litellm/litellm_core_utils/streaming_handler.py
@@ -2000,24 +2000,56 @@ class CustomStreamWrapper:
)
## Map to OpenAI Exception
try:
- raise exception_type(
+ mapped_exception = exception_type(
model=self.model,
custom_llm_provider=self.custom_llm_provider,
original_exception=e,
completion_kwargs={},
extra_kwargs={},
)
- except Exception as e:
- from litellm.exceptions import MidStreamFallbackError
+ except Exception as mapping_error:
+ mapped_exception = mapping_error
- raise MidStreamFallbackError(
- message=str(e),
- model=self.model,
- llm_provider=self.custom_llm_provider or "anthropic",
- original_exception=e,
- generated_content=self.response_uptil_now,
- is_pre_first_chunk=not self.sent_first_chunk,
- )
+ def _normalize_status_code(exc: Exception) -> Optional[int]:
+ """
+ Best-effort status_code extraction.
+ Uses status_code on the exception, then falls back to the response.
+ """
+ try:
+ code = getattr(exc, "status_code", None)
+ if code is not None:
+ return int(code)
+ except Exception:
+ pass
+
+ response = getattr(exc, "response", None)
+ if response is not None:
+ try:
+ status_code = getattr(response, "status_code", None)
+ if status_code is not None:
+ return int(status_code)
+ except Exception:
+ pass
+ return None
+
+ mapped_status_code = _normalize_status_code(mapped_exception)
+ original_status_code = _normalize_status_code(e)
+
+ if mapped_status_code is not None and 400 <= mapped_status_code < 500:
+ raise mapped_exception
+ if original_status_code is not None and 400 <= original_status_code < 500:
+ raise mapped_exception
+
+ from litellm.exceptions import MidStreamFallbackError
+
+ raise MidStreamFallbackError(
+ message=str(mapped_exception),
+ model=self.model,
+ llm_provider=self.custom_llm_provider or "anthropic",
+ original_exception=mapped_exception,
+ generated_content=self.response_uptil_now,
+ is_pre_first_chunk=not self.sent_first_chunk,
+ )
@staticmethod
def _strip_sse_data_from_chunk(chunk: Optional[str]) -> Optional[str]:
diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py
index facabbda72a..7a4da985528 100644
--- a/litellm/llms/base_llm/responses/transformation.py
+++ b/litellm/llms/base_llm/responses/transformation.py
@@ -242,3 +242,30 @@ class BaseResponsesAPIConfig(ABC):
#########################################################
########## END CANCEL RESPONSE API TRANSFORMATION #######
#########################################################
+
+ #########################################################
+ ########## COMPACT RESPONSE API TRANSFORMATION ##########
+ #########################################################
+ @abstractmethod
+ def transform_compact_response_api_request(
+ self,
+ model: str,
+ input: Union[str, ResponseInputParam],
+ response_api_optional_request_params: Dict,
+ api_base: str,
+ litellm_params: GenericLiteLLMParams,
+ headers: dict,
+ ) -> Tuple[str, Dict]:
+ pass
+
+ @abstractmethod
+ def transform_compact_response_api_response(
+ self,
+ raw_response: httpx.Response,
+ logging_obj: LiteLLMLoggingObj,
+ ) -> ResponsesAPIResponse:
+ pass
+
+ #########################################################
+ ########## END COMPACT RESPONSE API TRANSFORMATION ######
+ #########################################################
diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py
index 34ea598a655..ea740400664 100644
--- a/litellm/llms/custom_httpx/llm_http_handler.py
+++ b/litellm/llms/custom_httpx/llm_http_handler.py
@@ -91,6 +91,7 @@ from litellm.types.rerank import RerankResponse
from litellm.types.responses.main import DeleteResponseResult
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
+ CallTypes,
EmbeddingResponse,
FileTypes,
LiteLLMBatch,
@@ -850,7 +851,9 @@ class BaseLLMHTTPHandler:
)
if client is None or not isinstance(client, HTTPHandler):
- sync_httpx_client = _get_httpx_client()
+ sync_httpx_client = _get_httpx_client(
+ params={"ssl_verify": litellm_params.get("ssl_verify", None)}
+ )
else:
sync_httpx_client = client
@@ -896,7 +899,8 @@ class BaseLLMHTTPHandler:
) -> EmbeddingResponse:
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
- llm_provider=litellm.LlmProviders(custom_llm_provider)
+ llm_provider=litellm.LlmProviders(custom_llm_provider),
+ params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
else:
async_httpx_client = client
@@ -2004,6 +2008,10 @@ class BaseLLMHTTPHandler:
"""
Handles responses API requests.
When _is_async=True, returns a coroutine instead of making the call directly.
+
+ Keeps the pre-transform request context for streaming so post-call hooks/metadata
+ (added for Responses API parity with chat) receive the original params instead of
+ the provider-shaped body that caused them to be skipped before.
"""
if _is_async:
@@ -2060,6 +2068,18 @@ class BaseLLMHTTPHandler:
if extra_body:
data.update(extra_body)
+ # Preserve the OpenAI-style request context (not sent to the provider) for streaming
+ # hooks/metadata; the streaming iterator now consumes this to run deployment hooks
+ # with the same info as chat, including litellm_params.
+ request_context: Dict[str, Any] = {"input": input}
+ try:
+ request_context.update(response_api_optional_request_params)
+ except Exception:
+ pass
+ # Needed by streaming callbacks/metadata helpers to reconstruct api_base/model_id
+ # but never included in the outbound provider payload.
+ request_context["litellm_params"] = dict(litellm_params)
+
## LOGGING
logging_obj.pre_call(
input=input,
@@ -2097,6 +2117,8 @@ class BaseLLMHTTPHandler:
responses_api_provider_config=responses_api_provider_config,
litellm_metadata=litellm_metadata,
custom_llm_provider=custom_llm_provider,
+ request_data=request_context,
+ call_type=CallTypes.responses.value,
)
return SyncResponsesAPIStreamingIterator(
@@ -2106,6 +2128,8 @@ class BaseLLMHTTPHandler:
responses_api_provider_config=responses_api_provider_config,
litellm_metadata=litellm_metadata,
custom_llm_provider=custom_llm_provider,
+ request_data=request_context,
+ call_type=CallTypes.responses.value,
)
else:
# For non-streaming requests
@@ -2189,6 +2213,18 @@ class BaseLLMHTTPHandler:
if extra_body:
data.update(extra_body)
+ # Preserve the OpenAI-style request context (not sent to the provider) for streaming
+ # hooks/metadata; the streaming iterator now consumes this to run deployment hooks
+ # with the same info as chat, including litellm_params.
+ request_context: Dict[str, Any] = {"input": input}
+ try:
+ request_context.update(response_api_optional_request_params)
+ except Exception:
+ pass
+ # Needed by streaming callbacks/metadata helpers to reconstruct api_base/model_id
+ # but never included in the outbound provider payload.
+ request_context["litellm_params"] = dict(litellm_params)
+
## LOGGING
logging_obj.pre_call(
input=input,
@@ -2227,6 +2263,8 @@ class BaseLLMHTTPHandler:
responses_api_provider_config=responses_api_provider_config,
litellm_metadata=litellm_metadata,
custom_llm_provider=custom_llm_provider,
+ request_data=request_context,
+ call_type=CallTypes.responses.value,
)
# Return the streaming iterator
@@ -2237,6 +2275,8 @@ class BaseLLMHTTPHandler:
responses_api_provider_config=responses_api_provider_config,
litellm_metadata=litellm_metadata,
custom_llm_provider=custom_llm_provider,
+ request_data=request_context,
+ call_type=CallTypes.responses.value,
)
else:
# For non-streaming, proceed as before
@@ -3526,6 +3566,174 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
+ def compact_response_api_handler(
+ self,
+ model: str,
+ input: Union[str, "ResponseInputParam"],
+ responses_api_provider_config: BaseResponsesAPIConfig,
+ response_api_optional_request_params: Dict,
+ litellm_params: GenericLiteLLMParams,
+ logging_obj: LiteLLMLoggingObj,
+ custom_llm_provider: Optional[str],
+ extra_headers: Optional[Dict[str, Any]] = None,
+ extra_body: Optional[Dict[str, Any]] = None,
+ timeout: Optional[Union[float, httpx.Timeout]] = None,
+ client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
+ _is_async: bool = False,
+ shared_session: Optional["ClientSession"] = None,
+ ) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]:
+ """
+ Handler for the compact responses API.
+ """
+ if _is_async:
+ return self.async_compact_response_api_handler(
+ model=model,
+ input=input,
+ responses_api_provider_config=responses_api_provider_config,
+ response_api_optional_request_params=response_api_optional_request_params,
+ litellm_params=litellm_params,
+ logging_obj=logging_obj,
+ custom_llm_provider=custom_llm_provider,
+ extra_headers=extra_headers,
+ extra_body=extra_body,
+ timeout=timeout,
+ client=client,
+ shared_session=shared_session,
+ )
+ if client is None or not isinstance(client, HTTPHandler):
+ sync_httpx_client = _get_httpx_client(
+ params={"ssl_verify": litellm_params.get("ssl_verify", None)}
+ )
+ else:
+ sync_httpx_client = client
+
+ headers = responses_api_provider_config.validate_environment(
+ headers=extra_headers or {}, model=model, litellm_params=litellm_params
+ )
+
+ if extra_headers:
+ headers.update(extra_headers)
+
+ api_base = responses_api_provider_config.get_complete_url(
+ api_base=litellm_params.api_base,
+ litellm_params=dict(litellm_params),
+ )
+
+ url, data = responses_api_provider_config.transform_compact_response_api_request(
+ model=model,
+ input=input,
+ response_api_optional_request_params=response_api_optional_request_params,
+ api_base=api_base,
+ litellm_params=litellm_params,
+ headers=headers,
+ )
+
+ ## LOGGING
+ logging_obj.pre_call(
+ input=input,
+ api_key="",
+ additional_args={
+ "complete_input_dict": data,
+ "api_base": url,
+ "headers": headers,
+ },
+ )
+
+ try:
+ response = sync_httpx_client.post(
+ url=url, headers=headers, json=data, timeout=timeout
+ )
+
+ except Exception as e:
+ raise self._handle_error(
+ e=e,
+ provider_config=responses_api_provider_config,
+ )
+
+ return responses_api_provider_config.transform_compact_response_api_response(
+ raw_response=response,
+ logging_obj=logging_obj,
+ )
+
+ async def async_compact_response_api_handler(
+ self,
+ model: str,
+ input: Union[str, "ResponseInputParam"],
+ responses_api_provider_config: BaseResponsesAPIConfig,
+ response_api_optional_request_params: Dict,
+ litellm_params: GenericLiteLLMParams,
+ logging_obj: LiteLLMLoggingObj,
+ custom_llm_provider: Optional[str],
+ extra_headers: Optional[Dict[str, Any]] = None,
+ extra_body: Optional[Dict[str, Any]] = None,
+ timeout: Optional[Union[float, httpx.Timeout]] = None,
+ client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
+ _is_async: bool = False,
+ shared_session: Optional["ClientSession"] = None,
+ ) -> ResponsesAPIResponse:
+ """
+ Async version of the compact response API handler.
+ """
+ if client is None or not isinstance(client, AsyncHTTPHandler):
+ verbose_logger.debug(
+ f"Creating HTTP client for compact_response with shared_session: {id(shared_session) if shared_session else None}"
+ )
+ async_httpx_client = get_async_httpx_client(
+ llm_provider=litellm.LlmProviders(custom_llm_provider),
+ params={"ssl_verify": litellm_params.get("ssl_verify", None)},
+ shared_session=shared_session,
+ )
+ else:
+ async_httpx_client = client
+
+ headers = responses_api_provider_config.validate_environment(
+ headers=extra_headers or {}, model=model, litellm_params=litellm_params
+ )
+
+ if extra_headers:
+ headers.update(extra_headers)
+
+ api_base = responses_api_provider_config.get_complete_url(
+ api_base=litellm_params.api_base,
+ litellm_params=dict(litellm_params),
+ )
+
+ url, data = responses_api_provider_config.transform_compact_response_api_request(
+ model=model,
+ input=input,
+ response_api_optional_request_params=response_api_optional_request_params,
+ api_base=api_base,
+ litellm_params=litellm_params,
+ headers=headers,
+ )
+
+ ## LOGGING
+ logging_obj.pre_call(
+ input=input,
+ api_key="",
+ additional_args={
+ "complete_input_dict": data,
+ "api_base": url,
+ "headers": headers,
+ },
+ )
+
+ try:
+ response = await async_httpx_client.post(
+ url=url, headers=headers, json=data, timeout=timeout
+ )
+
+ except Exception as e:
+ raise self._handle_error(
+ e=e,
+ provider_config=responses_api_provider_config,
+ )
+
+ return responses_api_provider_config.transform_compact_response_api_response(
+ raw_response=response,
+ logging_obj=logging_obj,
+ )
+
def list_files(self):
"""
Lists all files
@@ -8288,4 +8496,4 @@ class BaseLLMHTTPHandler:
return skills_api_provider_config.transform_delete_skill_response(
raw_response=response,
logging_obj=logging_obj,
- )
\ No newline at end of file
+ )
diff --git a/litellm/llms/gigachat/__init__.py b/litellm/llms/gigachat/__init__.py
new file mode 100644
index 00000000000..3ddbd7864d9
--- /dev/null
+++ b/litellm/llms/gigachat/__init__.py
@@ -0,0 +1,23 @@
+"""
+GigaChat Provider for LiteLLM
+
+GigaChat is Sber AI's large language model (Russia's leading LLM).
+Supports:
+- Chat completions (sync/async)
+- Streaming (sync/async)
+- Function calling / Tools
+- Structured output via JSON schema (emulated through function calls)
+- Image input (base64 and URL)
+- Embeddings
+
+API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/overview
+"""
+
+from .chat.transformation import GigaChatConfig, GigaChatError
+from .embedding.transformation import GigaChatEmbeddingConfig
+
+__all__ = [
+ "GigaChatConfig",
+ "GigaChatEmbeddingConfig",
+ "GigaChatError",
+]
diff --git a/litellm/llms/gigachat/authenticator.py b/litellm/llms/gigachat/authenticator.py
new file mode 100644
index 00000000000..e61015a4a21
--- /dev/null
+++ b/litellm/llms/gigachat/authenticator.py
@@ -0,0 +1,241 @@
+"""
+GigaChat OAuth Authenticator
+
+Handles OAuth 2.0 token management for GigaChat API.
+Based on official GigaChat SDK authentication flow.
+"""
+
+import time
+import uuid
+from typing import Optional, Tuple
+
+import httpx
+
+from litellm._logging import verbose_logger
+from litellm.caching.caching import InMemoryCache
+from litellm.llms.base_llm.chat.transformation import BaseLLMException
+from litellm.llms.custom_httpx.http_handler import (
+ HTTPHandler,
+ _get_httpx_client,
+ get_async_httpx_client,
+)
+from litellm.secret_managers.main import get_secret_str
+from litellm.types.utils import LlmProviders
+
+# GigaChat OAuth endpoint
+GIGACHAT_AUTH_URL = "https://ngw.devices.sberbank.ru:9443/api/v2/oauth"
+
+# Default scope for personal API access
+GIGACHAT_SCOPE = "GIGACHAT_API_PERS"
+
+# Token expiry buffer in milliseconds (refresh token 60s before expiry)
+TOKEN_EXPIRY_BUFFER_MS = 60000
+
+# Cache for access tokens
+_token_cache = InMemoryCache()
+
+
+class GigaChatAuthError(BaseLLMException):
+ """GigaChat authentication error."""
+
+ pass
+
+
+def _get_credentials() -> Optional[str]:
+ """Get GigaChat credentials from environment."""
+ return get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY")
+
+
+def _get_auth_url() -> str:
+ """Get GigaChat auth URL from environment or use default."""
+ return get_secret_str("GIGACHAT_AUTH_URL") or GIGACHAT_AUTH_URL
+
+
+def _get_scope() -> str:
+ """Get GigaChat scope from environment or use default."""
+ return get_secret_str("GIGACHAT_SCOPE") or GIGACHAT_SCOPE
+
+
+def _get_http_client() -> HTTPHandler:
+ """Get cached httpx client with SSL verification disabled."""
+ return _get_httpx_client(params={"ssl_verify": False})
+
+
+def get_access_token(
+ credentials: Optional[str] = None,
+ scope: Optional[str] = None,
+ auth_url: Optional[str] = None,
+) -> str:
+ """
+ Get valid access token, using cache if available.
+
+ Args:
+ credentials: Base64-encoded credentials (client_id:client_secret)
+ scope: API scope (GIGACHAT_API_PERS, GIGACHAT_API_CORP, etc.)
+ auth_url: OAuth endpoint URL
+
+ Returns:
+ Access token string
+
+ Raises:
+ GigaChatAuthError: If authentication fails
+ """
+ credentials = credentials or _get_credentials()
+ if not credentials:
+ raise GigaChatAuthError(
+ status_code=401,
+ message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.",
+ )
+
+ scope = scope or _get_scope()
+ auth_url = auth_url or _get_auth_url()
+
+ # Check cache
+ cache_key = f"gigachat_token:{credentials[:16]}"
+ cached = _token_cache.get_cache(cache_key)
+ if cached:
+ token, expires_at = cached
+ # Check if token is still valid (with buffer)
+ if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS:
+ verbose_logger.debug("Using cached GigaChat access token")
+ return token
+
+ # Request new token
+ token, expires_at = _request_token_sync(credentials, scope, auth_url)
+
+ # Cache token
+ ttl_seconds = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
+ if ttl_seconds > 0:
+ _token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds)
+
+ return token
+
+
+async def get_access_token_async(
+ credentials: Optional[str] = None,
+ scope: Optional[str] = None,
+ auth_url: Optional[str] = None,
+) -> str:
+ """Async version of get_access_token."""
+ credentials = credentials or _get_credentials()
+ if not credentials:
+ raise GigaChatAuthError(
+ status_code=401,
+ message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.",
+ )
+
+ scope = scope or _get_scope()
+ auth_url = auth_url or _get_auth_url()
+
+ # Check cache
+ cache_key = f"gigachat_token:{credentials[:16]}"
+ cached = _token_cache.get_cache(cache_key)
+ if cached:
+ token, expires_at = cached
+ if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS:
+ verbose_logger.debug("Using cached GigaChat access token")
+ return token
+
+ # Request new token
+ token, expires_at = await _request_token_async(credentials, scope, auth_url)
+
+ # Cache token
+ ttl_seconds = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
+ if ttl_seconds > 0:
+ _token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds)
+
+ return token
+
+
+def _request_token_sync(
+ credentials: str,
+ scope: str,
+ auth_url: str,
+) -> Tuple[str, int]:
+ """
+ Request new access token from GigaChat OAuth endpoint (sync).
+
+ Returns:
+ Tuple of (access_token, expires_at_ms)
+ """
+ headers = {
+ "Authorization": f"Basic {credentials}",
+ "RqUID": str(uuid.uuid4()),
+ "Content-Type": "application/x-www-form-urlencoded",
+ }
+ data = {"scope": scope}
+
+ verbose_logger.debug(f"Requesting GigaChat access token from {auth_url}")
+
+ try:
+ client = _get_http_client()
+ response = client.post(auth_url, headers=headers, data=data, timeout=30)
+ response.raise_for_status()
+ return _parse_token_response(response)
+ except httpx.HTTPStatusError as e:
+ raise GigaChatAuthError(
+ status_code=e.response.status_code,
+ message=f"GigaChat authentication failed: {e.response.text}",
+ )
+ except httpx.RequestError as e:
+ raise GigaChatAuthError(
+ status_code=500,
+ message=f"GigaChat authentication request failed: {str(e)}",
+ )
+
+
+async def _request_token_async(
+ credentials: str,
+ scope: str,
+ auth_url: str,
+) -> Tuple[str, int]:
+ """Async version of _request_token_sync."""
+ headers = {
+ "Authorization": f"Basic {credentials}",
+ "RqUID": str(uuid.uuid4()),
+ "Content-Type": "application/x-www-form-urlencoded",
+ }
+ data = {"scope": scope}
+
+ verbose_logger.debug(f"Requesting GigaChat access token from {auth_url}")
+
+ try:
+ client = get_async_httpx_client(
+ llm_provider=LlmProviders.GIGACHAT,
+ params={"ssl_verify": False},
+ )
+ response = await client.post(auth_url, headers=headers, data=data, timeout=30)
+ response.raise_for_status()
+ return _parse_token_response(response)
+ except httpx.HTTPStatusError as e:
+ raise GigaChatAuthError(
+ status_code=e.response.status_code,
+ message=f"GigaChat authentication failed: {e.response.text}",
+ )
+ except httpx.RequestError as e:
+ raise GigaChatAuthError(
+ status_code=500,
+ message=f"GigaChat authentication request failed: {str(e)}",
+ )
+
+
+def _parse_token_response(response: httpx.Response) -> Tuple[str, int]:
+ """Parse OAuth token response."""
+ data = response.json()
+
+ # GigaChat returns either 'tok'/'exp' or 'access_token'/'expires_at'
+ access_token = data.get("tok") or data.get("access_token")
+ expires_at = data.get("exp") or data.get("expires_at")
+
+ if not access_token:
+ raise GigaChatAuthError(
+ status_code=500,
+ message=f"Invalid token response: {data}",
+ )
+
+ # expires_at is in milliseconds
+ if isinstance(expires_at, str):
+ expires_at = int(expires_at)
+
+ verbose_logger.debug("GigaChat access token obtained successfully")
+ return access_token, expires_at
diff --git a/litellm/llms/gigachat/chat/__init__.py b/litellm/llms/gigachat/chat/__init__.py
new file mode 100644
index 00000000000..3e030497a1a
--- /dev/null
+++ b/litellm/llms/gigachat/chat/__init__.py
@@ -0,0 +1,12 @@
+"""
+GigaChat Chat Module
+"""
+
+from .transformation import GigaChatConfig, GigaChatError
+from .streaming import GigaChatModelResponseIterator
+
+__all__ = [
+ "GigaChatConfig",
+ "GigaChatError",
+ "GigaChatModelResponseIterator",
+]
diff --git a/litellm/llms/gigachat/chat/streaming.py b/litellm/llms/gigachat/chat/streaming.py
new file mode 100644
index 00000000000..3565559e43c
--- /dev/null
+++ b/litellm/llms/gigachat/chat/streaming.py
@@ -0,0 +1,134 @@
+"""
+GigaChat Streaming Response Handler
+"""
+
+import json
+import uuid
+from typing import Any, Optional
+
+from litellm.types.llms.openai import ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk
+from litellm.types.utils import GenericStreamingChunk
+
+
+class GigaChatModelResponseIterator:
+ """Iterator for GigaChat streaming responses."""
+
+ def __init__(
+ self,
+ streaming_response: Any,
+ sync_stream: bool,
+ json_mode: Optional[bool] = False,
+ ):
+ self.streaming_response = streaming_response
+ self.response_iterator = self.streaming_response
+ self.json_mode = json_mode
+
+ def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
+ """Parse a single streaming chunk from GigaChat."""
+ text = ""
+ tool_use: Optional[ChatCompletionToolCallChunk] = None
+ is_finished = False
+ finish_reason: Optional[str] = None
+
+ choices = chunk.get("choices", [])
+ if not choices:
+ return GenericStreamingChunk(
+ text="",
+ tool_use=None,
+ is_finished=False,
+ finish_reason="",
+ usage=None,
+ index=0,
+ )
+
+ choice = choices[0]
+ delta = choice.get("delta", {})
+ finish_reason = choice.get("finish_reason")
+
+ # Extract text content
+ text = delta.get("content", "") or ""
+
+ # Handle function_call in stream
+ if finish_reason == "function_call" and delta.get("function_call"):
+ func_call = delta["function_call"]
+ args = func_call.get("arguments", {})
+
+ if isinstance(args, dict):
+ args = json.dumps(args, ensure_ascii=False)
+
+ tool_use = ChatCompletionToolCallChunk(
+ id=f"call_{uuid.uuid4().hex[:24]}",
+ type="function",
+ function=ChatCompletionToolCallFunctionChunk(
+ name=func_call.get("name", ""),
+ arguments=args,
+ ),
+ index=0,
+ )
+ finish_reason = "tool_calls"
+
+ if finish_reason is not None:
+ is_finished = True
+
+ return GenericStreamingChunk(
+ text=text,
+ tool_use=tool_use,
+ is_finished=is_finished,
+ finish_reason=finish_reason or "",
+ usage=None,
+ index=choice.get("index", 0),
+ )
+
+ def __iter__(self):
+ return self
+
+ def __next__(self) -> GenericStreamingChunk:
+ try:
+ chunk = self.response_iterator.__next__()
+ if isinstance(chunk, str):
+ # Parse SSE format: data: {...}
+ if chunk.startswith("data: "):
+ chunk = chunk[6:]
+ if chunk.strip() == "[DONE]":
+ raise StopIteration
+ try:
+ chunk = json.loads(chunk)
+ except json.JSONDecodeError:
+ return GenericStreamingChunk(
+ text="",
+ tool_use=None,
+ is_finished=False,
+ finish_reason="",
+ usage=None,
+ index=0,
+ )
+ return self.chunk_parser(chunk)
+ except StopIteration:
+ raise
+
+ def __aiter__(self):
+ return self
+
+ async def __anext__(self) -> GenericStreamingChunk:
+ try:
+ chunk = await self.response_iterator.__anext__()
+ if isinstance(chunk, str):
+ # Parse SSE format
+ if chunk.startswith("data: "):
+ chunk = chunk[6:]
+ if chunk.strip() == "[DONE]":
+ raise StopAsyncIteration
+ try:
+ chunk = json.loads(chunk)
+ except json.JSONDecodeError:
+ return GenericStreamingChunk(
+ text="",
+ tool_use=None,
+ is_finished=False,
+ finish_reason="",
+ usage=None,
+ index=0,
+ )
+ return self.chunk_parser(chunk)
+ except StopAsyncIteration:
+ raise
diff --git a/litellm/llms/gigachat/chat/transformation.py b/litellm/llms/gigachat/chat/transformation.py
new file mode 100644
index 00000000000..4ce333a1309
--- /dev/null
+++ b/litellm/llms/gigachat/chat/transformation.py
@@ -0,0 +1,473 @@
+"""
+GigaChat Chat Transformation
+
+Transforms OpenAI-format requests to GigaChat format and back.
+"""
+
+import json
+import time
+import uuid
+from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional, Union
+
+import httpx
+
+from litellm._logging import verbose_logger
+from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
+from litellm.secret_managers.main import get_secret_str
+from litellm.types.llms.openai import AllMessageValues
+from litellm.types.utils import Choices, Message, ModelResponse, Usage
+
+from ..authenticator import get_access_token
+from ..file_handler import upload_file_sync
+
+if TYPE_CHECKING:
+ from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
+
+ LiteLLMLoggingObj = _LiteLLMLoggingObj
+else:
+ LiteLLMLoggingObj = Any
+
+# GigaChat API endpoint
+GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1"
+
+
+class GigaChatError(BaseLLMException):
+ """GigaChat API error."""
+
+ pass
+
+
+class GigaChatConfig(BaseConfig):
+ """
+ Configuration class for GigaChat API.
+
+ GigaChat is Sber's (Russia's largest bank) LLM API.
+
+ Supported parameters:
+ temperature: Sampling temperature (0-2, default 0.87)
+ top_p: Nucleus sampling parameter
+ max_tokens: Maximum tokens to generate
+ repetition_penalty: Repetition penalty factor
+ profanity_check: Enable content filtering
+ stream: Enable streaming
+ """
+
+ temperature: Optional[float] = None
+ top_p: Optional[float] = None
+ max_tokens: Optional[int] = None
+ repetition_penalty: Optional[float] = None
+ profanity_check: Optional[bool] = None
+
+ def __init__(
+ self,
+ temperature: Optional[float] = None,
+ top_p: Optional[float] = None,
+ max_tokens: Optional[int] = None,
+ repetition_penalty: Optional[float] = None,
+ profanity_check: Optional[bool] = None,
+ ) -> None:
+ locals_ = locals().copy()
+ for key, value in locals_.items():
+ if key != "self" and value is not None:
+ setattr(self.__class__, key, value)
+ # Instance variables for current request context
+ self._current_credentials: Optional[str] = None
+ self._current_api_base: Optional[str] = None
+
+ def get_complete_url(
+ self,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ model: str,
+ optional_params: dict,
+ litellm_params: dict,
+ stream: Optional[bool] = None,
+ ) -> str:
+ """Get complete API URL for chat completions."""
+ base = api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL
+ return f"{base}/chat/completions"
+
+ def validate_environment(
+ self,
+ headers: dict,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ api_key: Optional[str] = None,
+ api_base: Optional[str] = None,
+ ) -> dict:
+ """
+ Set up headers with OAuth token.
+ """
+ # Get access token
+ credentials = api_key or get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY")
+ access_token = get_access_token(credentials=credentials)
+
+ # Store credentials for image uploads
+ self._current_credentials = credentials
+ self._current_api_base = api_base
+
+ headers["Authorization"] = f"Bearer {access_token}"
+ headers["Content-Type"] = "application/json"
+ headers["Accept"] = "application/json"
+
+ return headers
+
+ def get_supported_openai_params(self, model: str) -> List[str]:
+ """Return list of supported OpenAI parameters."""
+ return [
+ "stream",
+ "temperature",
+ "top_p",
+ "max_tokens",
+ "max_completion_tokens",
+ "stop",
+ "tools",
+ "tool_choice",
+ "functions",
+ "function_call",
+ "response_format",
+ ]
+
+ def map_openai_params(
+ self,
+ non_default_params: dict,
+ optional_params: dict,
+ model: str,
+ drop_params: bool,
+ ) -> dict:
+ """Map OpenAI parameters to GigaChat parameters."""
+ for param, value in non_default_params.items():
+ if param == "stream":
+ optional_params["stream"] = value
+ elif param == "temperature":
+ # GigaChat: temperature 0 means use top_p=0 instead
+ if value == 0:
+ optional_params["top_p"] = 0
+ else:
+ optional_params["temperature"] = value
+ elif param == "top_p":
+ optional_params["top_p"] = value
+ elif param in ("max_tokens", "max_completion_tokens"):
+ optional_params["max_tokens"] = value
+ elif param == "stop":
+ # GigaChat doesn't support stop sequences
+ pass
+ elif param == "tools":
+ # Convert tools to functions format
+ optional_params["functions"] = self._convert_tools_to_functions(value)
+ elif param == "tool_choice":
+ if isinstance(value, dict) and value.get("function"):
+ optional_params["function_call"] = {"name": value["function"]["name"]}
+ elif value == "auto":
+ pass # Default behavior
+ elif value == "required":
+ # GigaChat doesn't have 'required', handled differently
+ pass
+ elif param == "functions":
+ optional_params["functions"] = value
+ elif param == "function_call":
+ optional_params["function_call"] = value
+ elif param == "response_format":
+ # Handle structured output via function calling
+ if value.get("type") == "json_schema":
+ json_schema = value.get("json_schema", {})
+ schema_name = json_schema.get("name", "structured_output")
+ schema = json_schema.get("schema", {})
+
+ function_def = {
+ "name": schema_name,
+ "description": f"Output structured response: {schema_name}",
+ "parameters": schema,
+ }
+
+ if "functions" not in optional_params:
+ optional_params["functions"] = []
+ optional_params["functions"].append(function_def)
+ optional_params["function_call"] = {"name": schema_name}
+ optional_params["_structured_output"] = True
+
+ return optional_params
+
+ def _convert_tools_to_functions(self, tools: List[dict]) -> List[dict]:
+ """Convert OpenAI tools format to GigaChat functions format."""
+ functions = []
+ for tool in tools:
+ if tool.get("type") == "function":
+ func = tool.get("function", {})
+ functions.append({
+ "name": func.get("name", ""),
+ "description": func.get("description", ""),
+ "parameters": func.get("parameters", {}),
+ })
+ return functions
+
+ def _upload_image(self, image_url: str) -> Optional[str]:
+ """
+ Upload image to GigaChat and return file_id.
+
+ Args:
+ image_url: URL or base64 data URL of the image
+
+ Returns:
+ file_id string or None if upload failed
+ """
+ try:
+ return upload_file_sync(
+ image_url=image_url,
+ credentials=self._current_credentials,
+ api_base=self._current_api_base,
+ )
+ except Exception as e:
+ verbose_logger.error(f"Failed to upload image: {e}")
+ return None
+
+ def transform_request(
+ self,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ headers: dict,
+ ) -> dict:
+ """Transform OpenAI request to GigaChat format."""
+ # Transform messages
+ giga_messages = self._transform_messages(messages)
+
+ # Build request
+ request_data = {
+ "model": model.replace("gigachat/", ""),
+ "messages": giga_messages,
+ }
+
+ # Add optional params
+ for key in ["temperature", "top_p", "max_tokens", "stream",
+ "repetition_penalty", "profanity_check"]:
+ if key in optional_params:
+ request_data[key] = optional_params[key]
+
+ # Add functions if present
+ if "functions" in optional_params:
+ request_data["functions"] = optional_params["functions"]
+ if "function_call" in optional_params:
+ request_data["function_call"] = optional_params["function_call"]
+
+ return request_data
+
+ def _transform_messages(self, messages: List[AllMessageValues]) -> List[dict]:
+ """Transform OpenAI messages to GigaChat format."""
+ transformed = []
+
+ for i, msg in enumerate(messages):
+ message = dict(msg)
+
+ # Remove unsupported fields
+ message.pop("name", None)
+
+ # Transform roles
+ role = message.get("role", "user")
+ if role == "developer":
+ message["role"] = "system"
+ elif role == "system" and i > 0:
+ # GigaChat only allows system message as first message
+ message["role"] = "user"
+ elif role == "tool":
+ message["role"] = "function"
+ content = message.get("content", "")
+ if not isinstance(content, str):
+ message["content"] = json.dumps(content, ensure_ascii=False)
+
+ # Handle None content
+ if message.get("content") is None:
+ message["content"] = ""
+
+ # Handle list content (multimodal) - extract text and images
+ content = message.get("content")
+ if isinstance(content, list):
+ texts = []
+ attachments = []
+ for part in content:
+ if isinstance(part, dict):
+ if part.get("type") == "text":
+ texts.append(part.get("text", ""))
+ elif part.get("type") == "image_url":
+ # Extract image URL and upload to GigaChat
+ image_url = part.get("image_url", {})
+ if isinstance(image_url, str):
+ url = image_url
+ else:
+ url = image_url.get("url", "")
+ if url:
+ file_id = self._upload_image(url)
+ if file_id:
+ attachments.append(file_id)
+ message["content"] = "\n".join(texts) if texts else ""
+ if attachments:
+ message["attachments"] = attachments
+
+ # Transform tool_calls to function_call
+ tool_calls = message.get("tool_calls")
+ if tool_calls and isinstance(tool_calls, list) and len(tool_calls) > 0:
+ tool_call = tool_calls[0]
+ func = tool_call.get("function", {})
+ args = func.get("arguments", "{}")
+ if isinstance(args, str):
+ try:
+ args = json.loads(args)
+ except json.JSONDecodeError:
+ args = {}
+ message["function_call"] = {
+ "name": func.get("name", ""),
+ "arguments": args,
+ }
+ message.pop("tool_calls", None)
+
+ transformed.append(message)
+
+ # Collapse consecutive user messages
+ return self._collapse_user_messages(transformed)
+
+ def _collapse_user_messages(self, messages: List[dict]) -> List[dict]:
+ """Collapse consecutive user messages into one."""
+ collapsed: List[dict] = []
+ prev_user_msg: Optional[dict] = None
+ content_parts: List[str] = []
+
+ for msg in messages:
+ if msg.get("role") == "user" and prev_user_msg is not None:
+ content_parts.append(msg.get("content", ""))
+ else:
+ if content_parts and prev_user_msg:
+ prev_user_msg["content"] = "\n".join(
+ [prev_user_msg.get("content", "")] + content_parts
+ )
+ content_parts = []
+ collapsed.append(msg)
+ prev_user_msg = msg if msg.get("role") == "user" else None
+
+ if content_parts and prev_user_msg:
+ prev_user_msg["content"] = "\n".join(
+ [prev_user_msg.get("content", "")] + content_parts
+ )
+
+ return collapsed
+
+ def transform_response(
+ self,
+ model: str,
+ raw_response: httpx.Response,
+ model_response: ModelResponse,
+ logging_obj: LiteLLMLoggingObj,
+ request_data: dict,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ encoding: Any,
+ api_key: Optional[str] = None,
+ json_mode: Optional[bool] = None,
+ ) -> ModelResponse:
+ """Transform GigaChat response to OpenAI format."""
+ try:
+ response_json = raw_response.json()
+ except Exception:
+ raise GigaChatError(
+ status_code=raw_response.status_code,
+ message=f"Invalid JSON response: {raw_response.text}",
+ )
+
+ is_structured_output = optional_params.get("_structured_output", False)
+
+ choices = []
+ for choice in response_json.get("choices", []):
+ message_data = choice.get("message", {})
+ finish_reason = choice.get("finish_reason", "stop")
+
+ # Transform function_call to tool_calls or content
+ if finish_reason == "function_call" and message_data.get("function_call"):
+ func_call = message_data["function_call"]
+ args = func_call.get("arguments", {})
+
+ if is_structured_output:
+ # Convert to content for structured output
+ if isinstance(args, dict):
+ content = json.dumps(args, ensure_ascii=False)
+ else:
+ content = str(args)
+ message_data["content"] = content
+ message_data.pop("function_call", None)
+ message_data.pop("functions_state_id", None)
+ finish_reason = "stop"
+ else:
+ # Convert to tool_calls format
+ if isinstance(args, dict):
+ args = json.dumps(args, ensure_ascii=False)
+ message_data["tool_calls"] = [{
+ "id": f"call_{uuid.uuid4().hex[:24]}",
+ "type": "function",
+ "function": {
+ "name": func_call.get("name", ""),
+ "arguments": args,
+ }
+ }]
+ message_data.pop("function_call", None)
+ finish_reason = "tool_calls"
+
+ # Clean up GigaChat-specific fields
+ message_data.pop("functions_state_id", None)
+
+ choices.append(
+ Choices(
+ index=choice.get("index", 0),
+ message=Message(
+ role=message_data.get("role", "assistant"),
+ content=message_data.get("content"),
+ tool_calls=message_data.get("tool_calls"),
+ ),
+ finish_reason=finish_reason,
+ )
+ )
+
+ # Build usage
+ usage_data = response_json.get("usage", {})
+ usage = Usage(
+ prompt_tokens=usage_data.get("prompt_tokens", 0),
+ completion_tokens=usage_data.get("completion_tokens", 0),
+ total_tokens=usage_data.get("total_tokens", 0),
+ )
+
+ model_response.id = response_json.get("id", f"chatcmpl-{uuid.uuid4().hex[:12]}")
+ model_response.created = response_json.get("created", int(time.time()))
+ model_response.model = model
+ model_response.choices = choices # type: ignore
+ setattr(model_response, "usage", usage)
+
+ return model_response
+
+ def get_error_class(
+ self,
+ error_message: str,
+ status_code: int,
+ headers: Union[dict, httpx.Headers],
+ ) -> BaseLLMException:
+ """Return GigaChat error class."""
+ return GigaChatError(
+ status_code=status_code,
+ message=error_message,
+ headers=headers,
+ )
+
+ def get_model_response_iterator(
+ self,
+ streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
+ sync_stream: bool,
+ json_mode: Optional[bool] = False,
+ ):
+ """Return streaming response iterator."""
+ from .streaming import GigaChatModelResponseIterator
+
+ return GigaChatModelResponseIterator(
+ streaming_response=streaming_response,
+ sync_stream=sync_stream,
+ json_mode=json_mode,
+ )
diff --git a/litellm/llms/gigachat/embedding/__init__.py b/litellm/llms/gigachat/embedding/__init__.py
new file mode 100644
index 00000000000..af237e49aab
--- /dev/null
+++ b/litellm/llms/gigachat/embedding/__init__.py
@@ -0,0 +1,7 @@
+"""
+GigaChat Embedding Module
+"""
+
+from .transformation import GigaChatEmbeddingConfig
+
+__all__ = ["GigaChatEmbeddingConfig"]
diff --git a/litellm/llms/gigachat/embedding/transformation.py b/litellm/llms/gigachat/embedding/transformation.py
new file mode 100644
index 00000000000..0da6565050e
--- /dev/null
+++ b/litellm/llms/gigachat/embedding/transformation.py
@@ -0,0 +1,212 @@
+"""
+GigaChat Embedding Transformation
+
+Transforms OpenAI /v1/embeddings format to GigaChat format.
+API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/rest/post-embeddings
+"""
+
+import types
+from typing import List, Optional, Tuple, Union
+
+import httpx
+
+from litellm import LlmProviders
+from litellm.llms.base_llm.chat.transformation import BaseLLMException
+from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
+from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
+from litellm.types.utils import EmbeddingResponse
+
+from ..authenticator import get_access_token
+
+# GigaChat API endpoint
+GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1"
+
+
+class GigaChatEmbeddingError(BaseLLMException):
+ """GigaChat Embedding API error."""
+
+ pass
+
+
+class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
+ """
+ Configuration class for GigaChat Embeddings API.
+
+ GigaChat embeddings endpoint: POST /api/v1/embeddings
+ """
+
+ def __init__(self) -> None:
+ pass
+
+ @classmethod
+ def get_config(cls):
+ return {
+ k: v
+ for k, v in cls.__dict__.items()
+ if not k.startswith("__")
+ and not isinstance(
+ v,
+ (
+ types.FunctionType,
+ types.BuiltinFunctionType,
+ classmethod,
+ staticmethod,
+ ),
+ )
+ and v is not None
+ }
+
+ def get_supported_openai_params(self, model: str) -> List[str]:
+ """GigaChat embeddings don't support additional parameters."""
+ return []
+
+ def map_openai_params(
+ self,
+ non_default_params: dict,
+ optional_params: dict,
+ model: str,
+ drop_params: bool,
+ ) -> dict:
+ """Map OpenAI params to GigaChat format (no special mapping needed)."""
+ return optional_params
+
+ def _get_openai_compatible_provider_info(
+ self,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ ) -> Tuple[str, Optional[str], Optional[str]]:
+ """
+ Returns provider info for GigaChat.
+
+ Returns:
+ Tuple of (custom_llm_provider, api_base, dynamic_api_key)
+ """
+ api_base = api_base or GIGACHAT_BASE_URL
+ return LlmProviders.GIGACHAT.value, api_base, api_key
+
+ def get_complete_url(
+ self,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ model: str,
+ optional_params: dict,
+ litellm_params: dict,
+ stream: Optional[bool] = None,
+ ) -> str:
+ """Get the complete URL for embeddings endpoint."""
+ base = api_base or GIGACHAT_BASE_URL
+ return f"{base}/embeddings"
+
+ def transform_embedding_request(
+ self,
+ model: str,
+ input: AllEmbeddingInputValues,
+ optional_params: dict,
+ headers: dict,
+ ) -> dict:
+ """
+ Transform OpenAI embedding request to GigaChat format.
+
+ GigaChat format:
+ {
+ "model": "Embeddings",
+ "input": ["text1", "text2", ...]
+ }
+ """
+ # Normalize input to list
+ if isinstance(input, str):
+ input_list: list = [input]
+ elif isinstance(input, list):
+ input_list = input
+ else:
+ input_list = [input]
+
+ # Remove gigachat/ prefix from model if present
+ if model.startswith("gigachat/"):
+ model = model[9:]
+
+ return {
+ "model": model,
+ "input": input_list,
+ }
+
+ def transform_embedding_response(
+ self,
+ model: str,
+ raw_response: httpx.Response,
+ model_response: EmbeddingResponse,
+ logging_obj: LiteLLMLoggingObj,
+ api_key: Optional[str],
+ request_data: dict,
+ optional_params: dict,
+ litellm_params: dict,
+ ) -> EmbeddingResponse:
+ """
+ Transform GigaChat embedding response to OpenAI format.
+
+ GigaChat returns:
+ {
+ "object": "list",
+ "data": [{"object": "embedding", "embedding": [...], "index": 0, "usage": {...}}],
+ "model": "Embeddings"
+ }
+ """
+ response_json = raw_response.json()
+
+ # Log response
+ logging_obj.post_call(
+ input=request_data.get("input"),
+ api_key=api_key,
+ additional_args={"complete_input_dict": request_data},
+ original_response=response_json,
+ )
+
+ # Calculate total tokens from individual embeddings
+ total_tokens = 0
+ if "data" in response_json:
+ for emb in response_json["data"]:
+ if "usage" in emb and "prompt_tokens" in emb["usage"]:
+ total_tokens += emb["usage"]["prompt_tokens"]
+ # Remove usage from individual embeddings (not part of OpenAI format)
+ if "usage" in emb:
+ del emb["usage"]
+
+ # Set overall usage
+ response_json["usage"] = {
+ "prompt_tokens": total_tokens,
+ "total_tokens": total_tokens,
+ }
+
+ return EmbeddingResponse(**response_json)
+
+ def validate_environment(
+ self,
+ headers: dict,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ api_key: Optional[str] = None,
+ api_base: Optional[str] = None,
+ ) -> dict:
+ """
+ Set up headers with OAuth token for GigaChat.
+ """
+ # Get access token via OAuth
+ access_token = get_access_token(api_key)
+
+ default_headers = {
+ "Content-Type": "application/json",
+ "Authorization": f"Bearer {access_token}",
+ }
+ return {**default_headers, **headers}
+
+ def get_error_class(
+ self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
+ ) -> BaseLLMException:
+ """Return GigaChat-specific error class."""
+ return GigaChatEmbeddingError(
+ status_code=status_code,
+ message=error_message,
+ )
diff --git a/litellm/llms/gigachat/file_handler.py b/litellm/llms/gigachat/file_handler.py
new file mode 100644
index 00000000000..200428a747a
--- /dev/null
+++ b/litellm/llms/gigachat/file_handler.py
@@ -0,0 +1,211 @@
+"""
+GigaChat File Handler
+
+Handles file uploads to GigaChat API for image processing.
+GigaChat requires files to be uploaded first, then referenced by file_id.
+"""
+
+import base64
+import hashlib
+import re
+import uuid
+from typing import Dict, Optional, Tuple
+
+from litellm._logging import verbose_logger
+from litellm.llms.custom_httpx.http_handler import (
+ _get_httpx_client,
+ get_async_httpx_client,
+)
+from litellm.types.utils import LlmProviders
+
+from .authenticator import get_access_token, get_access_token_async
+
+# GigaChat API endpoint
+GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1"
+
+# Simple in-memory cache for file IDs
+_file_cache: Dict[str, str] = {}
+
+
+def _get_url_hash(url: str) -> str:
+ """Generate hash for URL to use as cache key."""
+ return hashlib.sha256(url.encode()).hexdigest()
+
+
+def _parse_data_url(data_url: str) -> Optional[Tuple[bytes, str, str]]:
+ """
+ Parse data URL (base64 image).
+
+ Returns:
+ Tuple of (content_bytes, content_type, extension) or None
+ """
+ match = re.match(r"data:([^;]+);base64,(.+)", data_url)
+ if not match:
+ return None
+
+ content_type = match.group(1)
+ base64_data = match.group(2)
+ content_bytes = base64.b64decode(base64_data)
+ ext = content_type.split("/")[-1].split(";")[0] or "jpg"
+
+ return content_bytes, content_type, ext
+
+
+def _download_image_sync(url: str) -> Tuple[bytes, str, str]:
+ """Download image from URL synchronously."""
+ client = _get_httpx_client(params={"ssl_verify": False})
+ response = client.get(url)
+ response.raise_for_status()
+
+ content_type = response.headers.get("content-type", "image/jpeg")
+ ext = content_type.split("/")[-1].split(";")[0] or "jpg"
+
+ return response.content, content_type, ext
+
+
+async def _download_image_async(url: str) -> Tuple[bytes, str, str]:
+ """Download image from URL asynchronously."""
+ client = get_async_httpx_client(
+ llm_provider=LlmProviders.GIGACHAT,
+ params={"ssl_verify": False},
+ )
+ response = await client.get(url)
+ response.raise_for_status()
+
+ content_type = response.headers.get("content-type", "image/jpeg")
+ ext = content_type.split("/")[-1].split(";")[0] or "jpg"
+
+ return response.content, content_type, ext
+
+
+def upload_file_sync(
+ image_url: str,
+ credentials: Optional[str] = None,
+ api_base: Optional[str] = None,
+) -> Optional[str]:
+ """
+ Upload file to GigaChat and return file_id (sync).
+
+ Args:
+ image_url: URL or base64 data URL of the image
+ credentials: GigaChat credentials for auth
+ api_base: Optional custom API base URL
+
+ Returns:
+ file_id string or None if upload failed
+ """
+ url_hash = _get_url_hash(image_url)
+
+ # Check cache
+ if url_hash in _file_cache:
+ verbose_logger.debug(f"Image found in cache: {url_hash[:16]}...")
+ return _file_cache[url_hash]
+
+ try:
+ # Get image data
+ parsed = _parse_data_url(image_url)
+ if parsed:
+ content_bytes, content_type, ext = parsed
+ verbose_logger.debug("Decoded base64 image")
+ else:
+ verbose_logger.debug(f"Downloading image from URL: {image_url[:80]}...")
+ content_bytes, content_type, ext = _download_image_sync(image_url)
+
+ filename = f"{uuid.uuid4()}.{ext}"
+
+ # Get access token
+ access_token = get_access_token(credentials)
+
+ # Upload to GigaChat
+ base_url = api_base or GIGACHAT_BASE_URL
+ upload_url = f"{base_url}/files"
+
+ client = _get_httpx_client(params={"ssl_verify": False})
+ response = client.post(
+ upload_url,
+ headers={"Authorization": f"Bearer {access_token}"},
+ files={"file": (filename, content_bytes, content_type)},
+ data={"purpose": "general"},
+ timeout=60,
+ )
+ response.raise_for_status()
+ result = response.json()
+
+ file_id = result.get("id")
+ if file_id:
+ _file_cache[url_hash] = file_id
+ verbose_logger.debug(f"File uploaded successfully, file_id: {file_id}")
+
+ return file_id
+
+ except Exception as e:
+ verbose_logger.error(f"Error uploading file to GigaChat: {e}")
+ return None
+
+
+async def upload_file_async(
+ image_url: str,
+ credentials: Optional[str] = None,
+ api_base: Optional[str] = None,
+) -> Optional[str]:
+ """
+ Upload file to GigaChat and return file_id (async).
+
+ Args:
+ image_url: URL or base64 data URL of the image
+ credentials: GigaChat credentials for auth
+ api_base: Optional custom API base URL
+
+ Returns:
+ file_id string or None if upload failed
+ """
+ url_hash = _get_url_hash(image_url)
+
+ # Check cache
+ if url_hash in _file_cache:
+ verbose_logger.debug(f"Image found in cache: {url_hash[:16]}...")
+ return _file_cache[url_hash]
+
+ try:
+ # Get image data
+ parsed = _parse_data_url(image_url)
+ if parsed:
+ content_bytes, content_type, ext = parsed
+ verbose_logger.debug("Decoded base64 image")
+ else:
+ verbose_logger.debug(f"Downloading image from URL: {image_url[:80]}...")
+ content_bytes, content_type, ext = await _download_image_async(image_url)
+
+ filename = f"{uuid.uuid4()}.{ext}"
+
+ # Get access token
+ access_token = await get_access_token_async(credentials)
+
+ # Upload to GigaChat
+ base_url = api_base or GIGACHAT_BASE_URL
+ upload_url = f"{base_url}/files"
+
+ client = get_async_httpx_client(
+ llm_provider=LlmProviders.GIGACHAT,
+ params={"ssl_verify": False},
+ )
+ response = await client.post(
+ upload_url,
+ headers={"Authorization": f"Bearer {access_token}"},
+ files={"file": (filename, content_bytes, content_type)},
+ data={"purpose": "general"},
+ timeout=60,
+ )
+ response.raise_for_status()
+ result = response.json()
+
+ file_id = result.get("id")
+ if file_id:
+ _file_cache[url_hash] = file_id
+ verbose_logger.debug(f"File uploaded successfully, file_id: {file_id}")
+
+ return file_id
+
+ except Exception as e:
+ verbose_logger.error(f"Error uploading file to GigaChat: {e}")
+ return None
diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py
index 96598c1dfe6..cc2439b431a 100644
--- a/litellm/llms/openai/responses/transformation.py
+++ b/litellm/llms/openai/responses/transformation.py
@@ -500,3 +500,69 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
response._hidden_params["headers"] = raw_response_headers
return response
+
+ #########################################################
+ ########## COMPACT RESPONSE API TRANSFORMATION ##########
+ #########################################################
+ def transform_compact_response_api_request(
+ self,
+ model: str,
+ input: Union[str, ResponseInputParam],
+ response_api_optional_request_params: Dict,
+ api_base: str,
+ litellm_params: GenericLiteLLMParams,
+ headers: dict,
+ ) -> Tuple[str, Dict]:
+ """
+ Transform the compact response API request into a URL and data
+
+ OpenAI API expects the following request
+ - POST /v1/responses/compact
+ """
+ url = f"{api_base}/compact"
+
+ input = self._validate_input_param(input)
+ data = dict(
+ ResponsesAPIRequestParams(
+ model=model, input=input, **response_api_optional_request_params
+ )
+ )
+
+ return url, data
+
+ def transform_compact_response_api_response(
+ self,
+ raw_response: httpx.Response,
+ logging_obj: LiteLLMLoggingObj,
+ ) -> ResponsesAPIResponse:
+ """
+ Transform the compact response API response into a ResponsesAPIResponse
+ """
+ try:
+ logging_obj.post_call(
+ original_response=raw_response.text,
+ additional_args={"complete_input_dict": {}},
+ )
+ raw_response_json = raw_response.json()
+ raw_response_json["created_at"] = _safe_convert_created_field(
+ raw_response_json["created_at"]
+ )
+ except Exception:
+ raise OpenAIError(
+ message=raw_response.text, status_code=raw_response.status_code
+ )
+ raw_response_headers = dict(raw_response.headers)
+ processed_headers = process_response_headers(raw_response_headers)
+
+ try:
+ response = ResponsesAPIResponse(**raw_response_json)
+ except Exception:
+ verbose_logger.debug(
+ f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct"
+ )
+ response = ResponsesAPIResponse.model_construct(**raw_response_json)
+
+ response._hidden_params["additional_headers"] = processed_headers
+ response._hidden_params["headers"] = raw_response_headers
+
+ return response
diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py
index 826f151df35..d23c698cd7a 100644
--- a/litellm/llms/vertex_ai/vertex_llm_base.py
+++ b/litellm/llms/vertex_ai/vertex_llm_base.py
@@ -6,7 +6,7 @@ Handles Authentication and generating request urls for Vertex AI and Google AI S
import json
import os
-from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
+from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple, cast
import litellm
from litellm._logging import verbose_logger
@@ -168,7 +168,6 @@ class VertexBase:
)
def _credentials_from_default_auth(self, scopes):
-
import google.auth as google_auth
return google_auth.default(scopes=scopes)
@@ -392,7 +391,7 @@ class VertexBase:
Returns
token, url
"""
- version: Optional[Literal["v1beta1", "v1"]] = None
+ version: Optional[Literal["v1", "v1beta1"]] = None
if custom_llm_provider == "gemini":
url, endpoint = _get_gemini_url(
mode=mode,
@@ -415,7 +414,7 @@ class VertexBase:
stream=stream,
vertex_project=vertex_project,
vertex_location=vertex_location,
- vertex_api_version=version,
+ vertex_api_version=cast(Literal["v1", "v1beta1"], version),
)
return self._check_custom_proxy(
diff --git a/litellm/main.py b/litellm/main.py
index f4f27eb5841..e8a8b504d96 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -2141,6 +2141,49 @@ def completion( # type: ignore # noqa: PLR0915
logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
client=client,
)
+ elif custom_llm_provider == "gigachat":
+ # GigaChat - Sber AI's LLM (Russia)
+ api_key = (
+ api_key
+ or litellm.api_key
+ or litellm.gigachat_key
+ or get_secret("GIGACHAT_API_KEY")
+ or get_secret("GIGACHAT_CREDENTIALS")
+ )
+
+ headers = headers or litellm.headers or {}
+
+ ## COMPLETION CALL
+ try:
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=_get_encoding(),
+ stream=stream,
+ provider_config=provider_config,
+ )
+ except Exception as e:
+ ## LOGGING - log the original exception returned
+ logging.post_call(
+ input=messages,
+ api_key=api_key,
+ original_response=str(e),
+ additional_args={"headers": headers},
+ )
+ raise e
+
elif custom_llm_provider == "sap":
headers = headers or litellm.headers
## LOAD CONFIG - if set
@@ -5224,6 +5267,28 @@ def embedding( # noqa: PLR0915
aembedding=aembedding,
litellm_params={},
)
+ elif custom_llm_provider == "gigachat":
+ api_key = (
+ api_key
+ or litellm.api_key
+ or litellm.gigachat_key
+ or get_secret_str("GIGACHAT_CREDENTIALS")
+ or get_secret_str("GIGACHAT_API_KEY")
+ )
+ response = base_llm_http_handler.embedding(
+ model=model,
+ input=input,
+ custom_llm_provider=custom_llm_provider,
+ api_base=api_base,
+ api_key=api_key,
+ logging_obj=logging,
+ timeout=timeout,
+ model_response=EmbeddingResponse(),
+ optional_params=optional_params,
+ client=client,
+ aembedding=aembedding,
+ litellm_params={"ssl_verify": kwargs.get("ssl_verify", None)},
+ )
else:
raise LiteLLMUnknownProvider(
model=model, custom_llm_provider=custom_llm_provider
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index e823dd5dc6b..c7a2f60856d 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -15831,6 +15831,68 @@
"max_tokens": 8191,
"mode": "embedding"
},
+ "gigachat/GigaChat-2-Lite": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0,
+ "supports_function_calling": true,
+ "supports_system_messages": true
+ },
+ "gigachat/GigaChat-2-Max": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0,
+ "supports_function_calling": true,
+ "supports_system_messages": true,
+ "supports_vision": true
+ },
+ "gigachat/GigaChat-2-Pro": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0,
+ "supports_function_calling": true,
+ "supports_system_messages": true,
+ "supports_vision": true
+ },
+ "gigachat/Embeddings": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 512,
+ "max_tokens": 512,
+ "mode": "embedding",
+ "output_cost_per_token": 0.0,
+ "output_vector_size": 1024
+ },
+ "gigachat/Embeddings-2": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 512,
+ "max_tokens": 512,
+ "mode": "embedding",
+ "output_cost_per_token": 0.0,
+ "output_vector_size": 1024
+ },
+ "gigachat/EmbeddingsGigaR": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 4096,
+ "max_tokens": 4096,
+ "mode": "embedding",
+ "output_cost_per_token": 0.0,
+ "output_vector_size": 2560
+ },
"google.gemma-3-12b-it": {
"input_cost_per_token": 9e-08,
"litellm_provider": "bedrock_converse",
@@ -32092,3 +32154,4 @@
"mode": "chat"
}
}
+
diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
index ffa17a5b7c4..ded591a8f53 100644
--- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
+++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
@@ -15,6 +15,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.types.mcp_server.mcp_server_manager import MCPServer
+from litellm.proxy.utils import get_server_root_path
router = APIRouter(
tags=["mcp"],
@@ -381,13 +382,30 @@ async def callback(code: str, state: str):
# ------------------------------
# Optional .well-known endpoints for MCP + OAuth discovery
# ------------------------------
-@router.get("/.well-known/oauth-protected-resource/{mcp_server_name}/mcp")
+"""
+ Per SEP-985, the client MUST:
+ 1. Try resource_metadata from WWW-Authenticate header (if present)
+ 2. Fall back to path-based well-known URI: /.well-known/oauth-protected-resource/{path}
+ (
+ If the resource identifier value contains a path or query component, any terminating slash (/)
+ following the host component MUST be removed before inserting /.well-known/ and the well-known
+ URI path suffix between the host component and the path(include root path) and/or query components.
+ https://datatracker.ietf.org/doc/html/rfc9728#section-3.1)
+ 3. Fall back to root-based well-known URI: /.well-known/oauth-protected-resource
+"""
+@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp")
@router.get("/.well-known/oauth-protected-resource")
async def oauth_protected_resource_mcp(
request: Request, mcp_server_name: Optional[str] = None
):
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
# Get the correct base URL considering X-Forwarded-* headers
request_base_url = get_request_base_url(request)
+ mcp_server: Optional[MCPServer] = None
+ if mcp_server_name:
+ mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name)
return {
"authorization_servers": [
(
@@ -401,14 +419,25 @@ async def oauth_protected_resource_mcp(
if mcp_server_name
else f"{request_base_url}/mcp"
), # this is what Claude will call
+ "scopes_supported": mcp_server.scopes if mcp_server else [],
}
-
-@router.get("/.well-known/oauth-authorization-server/{mcp_server_name}")
+"""
+ https://datatracker.ietf.org/doc/html/rfc8414#section-3.1
+ RFC 8414: Path-aware OAuth discovery
+ If the issuer identifier value contains a path component, any
+ terminating "/" MUST be removed before inserting "/.well-known/" and
+ the well-known URI suffix between the host component and the path(include root path)
+ component.
+"""
+@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}")
@router.get("/.well-known/oauth-authorization-server")
async def oauth_authorization_server_mcp(
request: Request, mcp_server_name: Optional[str] = None
):
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
# Get the correct base URL considering X-Forwarded-* headers
request_base_url = get_request_base_url(request)
@@ -423,16 +452,21 @@ async def oauth_authorization_server_mcp(
else f"{request_base_url}/token"
)
+ mcp_server: Optional[MCPServer] = None
+ if mcp_server_name:
+ mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name)
+
return {
"issuer": request_base_url, # point to your proxy
"authorization_endpoint": authorization_endpoint,
"token_endpoint": token_endpoint,
"response_types_supported": ["code"],
- "grant_types_supported": ["authorization_code"],
+ "scopes_supported": mcp_server.scopes if mcp_server else [],
+ "grant_types_supported": ["authorization_code", "refresh_token"],
"code_challenge_methods_supported": ["S256"],
"token_endpoint_auth_methods_supported": ["client_secret_post"],
# Claude expects a registration endpoint, even if we just fake it
- "registration_endpoint": f"{request_base_url}/{mcp_server_name}/register",
+ "registration_endpoint": f"{request_base_url}/{mcp_server_name}/register" if mcp_server_name else f"{request_base_url}/register",
}
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index a5ac966062e..3a548e203c5 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -660,14 +660,14 @@ class MCPServerManager:
"""
allowed_mcp_servers = await self.get_allowed_mcp_servers(user_api_key_auth)
- list_tools_result: List[MCPTool] = []
verbose_logger.debug("SERVER MANAGER LISTING TOOLS")
- for server_id in allowed_mcp_servers:
+ async def _fetch_server_tools(server_id: str) -> List[MCPTool]:
+ """Fetch tools from a single server with error handling."""
server = self.get_mcp_server_by_id(server_id)
if server is None:
verbose_logger.warning(f"MCP Server {server_id} not found")
- continue
+ return []
# Get server-specific auth header if available
server_auth_header = None
@@ -685,15 +685,21 @@ class MCPServerManager:
server=server,
mcp_auth_header=server_auth_header,
)
- list_tools_result.extend(tools)
- verbose_logger.info(
- f"Successfully fetched {len(tools)} tools from server {server.name}"
- )
+ return tools
except Exception as e:
verbose_logger.warning(
f"Failed to list tools from server {server.name}: {str(e)}. Continuing with other servers."
)
- # Continue with other servers instead of failing completely
+ return []
+
+ # Fetch tools from all servers in parallel
+ tasks = [_fetch_server_tools(server_id) for server_id in allowed_mcp_servers]
+ results = await asyncio.gather(*tasks)
+
+ # Flatten results into single list
+ list_tools_result: List[MCPTool] = [
+ tool for tools in results for tool in tools
+ ]
verbose_logger.info(
f"Successfully fetched {len(list_tools_result)} tools total from all servers"
@@ -2003,6 +2009,9 @@ class MCPServerManager:
Note: This now handles prefixed tool names
"""
for server in self.get_registry().values():
+ if server.auth_type == MCPAuth.oauth2:
+ # Skip OAuth2 servers for now as they may require user-specific tokens
+ continue
tools = await self._get_tools_from_server(server)
for tool in tools:
# The tool.name here is already prefixed from _get_tools_from_server
@@ -2284,14 +2293,7 @@ class MCPServerManager:
# Check all accessible servers
target_server_ids = allowed_server_ids
- # Run health checks concurrently
- tasks = [self.health_check_server(server_id) for server_id in target_server_ids]
- results = await asyncio.gather(*tasks)
-
- # Filter out None results (servers that were not found)
- list_mcp_servers = [server for server in results if server is not None]
-
- return list_mcp_servers
+ return await self._run_health_checks(target_server_ids)
async def get_all_allowed_mcp_servers(
self,
@@ -2306,8 +2308,6 @@ class MCPServerManager:
Returns:
List of MCP server objects without health status
"""
- from datetime import datetime
-
# Get allowed server IDs
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
@@ -2319,40 +2319,56 @@ class MCPServerManager:
verbose_logger.warning(f"MCP Server {server_id} not found in registry")
continue
- # Build LiteLLM_MCPServerTable without health check
- mcp_server_table = LiteLLM_MCPServerTable(
- server_id=server.server_id,
- server_name=server.server_name,
- alias=server.alias,
- description=(
- server.mcp_info.get("description") if server.mcp_info else None
- ),
- url=server.url,
- transport=server.transport,
- auth_type=server.auth_type,
- created_at=datetime.now(),
- updated_at=datetime.now(),
- teams=[],
- mcp_access_groups=server.access_groups or [],
- allowed_tools=server.allowed_tools or [],
- extra_headers=server.extra_headers or [],
- mcp_info=server.mcp_info,
- static_headers=server.static_headers,
- status=None, # No health check performed
- last_health_check=None, # No health check performed
- health_check_error=None,
- command=getattr(server, "command", None),
- args=getattr(server, "args", None) or [],
- env=getattr(server, "env", None) or {},
- authorization_url=server.authorization_url,
- token_url=server.token_url,
- registration_url=server.registration_url,
- allow_all_keys=server.allow_all_keys,
- )
+ mcp_server_table = self._build_mcp_server_table(server)
list_mcp_servers.append(mcp_server_table)
return list_mcp_servers
+ def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable:
+ from datetime import datetime
+
+ return LiteLLM_MCPServerTable(
+ server_id=server.server_id,
+ server_name=server.server_name,
+ alias=server.alias,
+ description=(
+ server.mcp_info.get("description") if server.mcp_info else None
+ ),
+ url=server.url,
+ transport=server.transport,
+ auth_type=server.auth_type,
+ created_at=datetime.now(),
+ updated_at=datetime.now(),
+ teams=[],
+ mcp_access_groups=server.access_groups or [],
+ allowed_tools=server.allowed_tools or [],
+ extra_headers=server.extra_headers or [],
+ mcp_info=server.mcp_info,
+ static_headers=server.static_headers,
+ status=None, # No health check performed
+ last_health_check=None, # No health check performed
+ health_check_error=None,
+ command=getattr(server, "command", None),
+ args=getattr(server, "args", None) or [],
+ env=getattr(server, "env", None) or {},
+ authorization_url=server.authorization_url,
+ token_url=server.token_url,
+ registration_url=server.registration_url,
+ allow_all_keys=server.allow_all_keys,
+ )
+
+ async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]:
+ """Return all MCP servers from registry without applying access controls."""
+
+ registry = self.get_registry()
+ if not registry:
+ return []
+
+ servers: List[LiteLLM_MCPServerTable] = []
+ for server in registry.values():
+ servers.append(self._build_mcp_server_table(server))
+ return servers
+
async def reload_servers_from_database(self):
"""
Public method to reload all MCP servers from database into registry.
@@ -2360,5 +2376,34 @@ class MCPServerManager:
"""
await self._add_mcp_servers_from_db_to_in_memory_registry()
+ async def get_all_mcp_servers_with_health_unfiltered(
+ self, server_ids: Optional[List[str]] = None
+ ) -> List[LiteLLM_MCPServerTable]:
+ """Return health info for all servers in registry regardless of user access."""
+
+ registry = self.get_registry()
+ if not registry:
+ return []
+
+ if server_ids:
+ target_server_ids = [sid for sid in server_ids if sid in registry]
+ else:
+ target_server_ids = list(registry.keys())
+
+ if not target_server_ids:
+ return []
+
+ return await self._run_health_checks(target_server_ids)
+
+ async def _run_health_checks(
+ self, target_server_ids: List[str]
+ ) -> List[LiteLLM_MCPServerTable]:
+ if not target_server_ids:
+ return []
+
+ tasks = [self.health_check_server(server_id) for server_id in target_server_ids]
+ results = await asyncio.gather(*tasks)
+ return [server for server in results if server is not None]
+
global_mcp_server_manager: MCPServerManager = MCPServerManager()
diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py
index e00fdbfb930..9c7001266f0 100644
--- a/litellm/proxy/_experimental/mcp_server/server.py
+++ b/litellm/proxy/_experimental/mcp_server/server.py
@@ -709,7 +709,8 @@ if MCP_AVAILABLE:
extra_headers: Optional[Dict[str, str]] = None
if server.auth_type == MCPAuth.oauth2:
- extra_headers = oauth2_headers
+ # Copy to avoid mutating the original dict (important for parallel fetching)
+ extra_headers = oauth2_headers.copy() if oauth2_headers else None
if server.extra_headers and raw_headers:
if extra_headers is None:
@@ -755,11 +756,10 @@ if MCP_AVAILABLE:
# Decide whether to add prefix based on number of allowed servers
add_prefix = not (len(allowed_mcp_servers) == 1)
- # Get tools from each allowed server
- all_tools = []
- for server in allowed_mcp_servers:
+ async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]:
+ """Fetch and filter tools from a single server with error handling."""
if server is None:
- continue
+ return []
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
@@ -786,16 +786,24 @@ if MCP_AVAILABLE:
user_api_key_auth=user_api_key_auth,
)
- all_tools.extend(filtered_tools)
-
verbose_logger.debug(
f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering"
)
+ return filtered_tools
except Exception as e:
verbose_logger.exception(
f"Error getting tools from server {server.name}: {str(e)}"
)
- # Continue with other servers instead of failing completely
+ return []
+
+ # Fetch tools from all servers in parallel
+ tasks = [
+ _fetch_and_filter_server_tools(server) for server in allowed_mcp_servers
+ ]
+ results = await asyncio.gather(*tasks)
+
+ # Flatten results into single list
+ all_tools: List[MCPTool] = [tool for tools in results for tool in tools]
verbose_logger.info(
f"Successfully fetched {len(all_tools)} tools total from all MCP servers"
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index b77cc40d6dc..954c26e2cb2 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -1908,6 +1908,9 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase):
}
+UserMCPManagementMode = Literal["restricted", "view_all"]
+
+
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"""
Documents all the fields supported by `general_settings` in config.yaml
@@ -2025,6 +2028,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).",
)
+ user_mcp_management_mode: Optional[UserMCPManagementMode] = Field(
+ None,
+ description="Controls how non-admin users interact with MCP servers in the dashboard. 'restricted' shows only accessible servers, 'view_all' lists every server in read-only mode.",
+ )
class ConfigYAML(LiteLLMPydanticObjectBase):
diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py
index 34049a44c8c..537b48f06ed 100644
--- a/litellm/proxy/common_request_processing.py
+++ b/litellm/proxy/common_request_processing.py
@@ -319,6 +319,7 @@ class ProxyBaseLLMRequestProcessing:
"aget_responses",
"adelete_responses",
"acancel_responses",
+ "acompact_responses",
"acreate_batch",
"aretrieve_batch",
"alist_batches",
@@ -457,6 +458,7 @@ class ProxyBaseLLMRequestProcessing:
"aget_responses",
"adelete_responses",
"acancel_responses",
+ "acompact_responses",
"atext_completion",
"aimage_edit",
"alist_input_items",
diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py
index 406ddceabf5..fe27af78d58 100644
--- a/litellm/proxy/db/prisma_client.py
+++ b/litellm/proxy/db/prisma_client.py
@@ -154,7 +154,11 @@ class PrismaManager:
prisma_dir = PrismaManager._get_prisma_dir()
- return ProxyExtrasDBManager.setup_database(use_migrate=use_migrate)
+ from litellm.proxy.proxy_server import redis_usage_cache
+
+ return ProxyExtrasDBManager.setup_database(
+ use_migrate=use_migrate, redis_cache=redis_usage_cache
+ )
else:
# Use prisma db push with increased timeout
subprocess.run(
diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py
index ea8f1b0a97f..5850103132c 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py
@@ -118,7 +118,7 @@ class LassoGuardrail(CustomGuardrail):
Falls back to UUID if ULID library is not available.
"""
if ULID_AVAILABLE and ulid is not None:
- return str(ulid.new()) # type: ignore
+ return str(ulid.ULID()) # type: ignore
else:
verbose_proxy_logger.debug("ULID library not available, using UUID")
return str(uuid.uuid4())
diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
index a871a6637a2..47793c8fc8e 100644
--- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
@@ -32,8 +32,8 @@ from fastapi import (
from fastapi.responses import JSONResponse
import litellm
-from litellm._uuid import uuid
from litellm._logging import verbose_logger, verbose_proxy_logger
+from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.proxy._experimental.mcp_server.utils import (
validate_and_normalize_mcp_server_payload,
@@ -67,7 +67,6 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
build_effective_auth_contexts,
)
- from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
LitellmUserRoles,
@@ -76,8 +75,10 @@ if MCP_AVAILABLE:
SpecialMCPServerName,
UpdateMCPServerRequest,
UserAPIKeyAuth,
+ UserMCPManagementMode,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
+ from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
from litellm.types.mcp import MCPCredentials
@@ -302,6 +303,20 @@ if MCP_AVAILABLE:
return {"access_groups": access_groups_list}
## FastAPI Routes
+ def _get_user_mcp_management_mode() -> UserMCPManagementMode:
+ proxy_general_settings: dict = {}
+ try:
+ from litellm.proxy.proxy_server import (
+ general_settings as proxy_general_settings,
+ )
+ except Exception:
+ pass
+
+ mode = proxy_general_settings.get("user_mcp_management_mode")
+ if mode == "view_all":
+ return "view_all"
+ return "restricted"
+
@router.get(
"/server",
description="Returns the mcp server list with associated teams",
@@ -319,18 +334,26 @@ if MCP_AVAILABLE:
```
"""
- auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
+ user_mcp_management_mode = _get_user_mcp_management_mode()
- aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
- for auth_context in auth_contexts:
- servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
- user_api_key_auth=auth_context
+ if user_mcp_management_mode == "view_all":
+ servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered()
+ redacted_mcp_servers = _redact_mcp_credentials_list(servers)
+ else:
+ auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
+
+ aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
+ for auth_context in auth_contexts:
+ servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
+ user_api_key_auth=auth_context
+ )
+ for server in servers:
+ if server.server_id not in aggregated_servers:
+ aggregated_servers[server.server_id] = server
+
+ redacted_mcp_servers = _redact_mcp_credentials_list(
+ aggregated_servers.values()
)
- for server in servers:
- if server.server_id not in aggregated_servers:
- aggregated_servers[server.server_id] = server
-
- redacted_mcp_servers = _redact_mcp_credentials_list(aggregated_servers.values())
# augment the mcp servers with public status
if litellm.public_mcp_servers is not None:
@@ -372,6 +395,17 @@ if MCP_AVAILABLE:
--header 'Authorization: Bearer your_api_key_here'
```
"""
+ user_mcp_management_mode = _get_user_mcp_management_mode()
+
+ if user_mcp_management_mode == "view_all":
+ servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered(
+ server_ids=server_ids
+ )
+ return [
+ {"server_id": server.server_id, "status": server.status}
+ for server in servers
+ ]
+
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
server_status_map: Dict[
diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py
index 623e8408862..ec1bc5497bd 100644
--- a/litellm/proxy/response_api_endpoints/endpoints.py
+++ b/litellm/proxy/response_api_endpoints/endpoints.py
@@ -698,6 +698,88 @@ async def get_response_input_items(
)
+@router.post(
+ "/v1/responses/compact",
+ dependencies=[Depends(user_api_key_auth)],
+ tags=["responses"],
+)
+@router.post(
+ "/responses/compact",
+ dependencies=[Depends(user_api_key_auth)],
+ tags=["responses"],
+)
+@router.post(
+ "/openai/v1/responses/compact",
+ dependencies=[Depends(user_api_key_auth)],
+ tags=["responses"],
+)
+async def compact_response(
+ request: Request,
+ fastapi_response: Response,
+ user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
+):
+ """
+ Compact a response by running a compaction pass over a conversation.
+
+ Returns encrypted, opaque items that can be used to reduce context size.
+
+ Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/compact
+
+ ```bash
+ curl -X POST http://localhost:4000/v1/responses/compact \
+ -H "Content-Type: application/json" \
+ -H "Authorization: Bearer sk-1234" \
+ -d '{
+ "model": "gpt-4o",
+ "input": [{"role": "user", "content": "Hello"}]
+ }'
+ ```
+ """
+ from litellm.proxy.proxy_server import (
+ _read_request_body,
+ general_settings,
+ llm_router,
+ proxy_config,
+ proxy_logging_obj,
+ select_data_generator,
+ user_api_base,
+ user_max_tokens,
+ user_model,
+ user_request_timeout,
+ user_temperature,
+ version,
+ )
+
+ data = await _read_request_body(request=request)
+ processor = ProxyBaseLLMRequestProcessing(data=data)
+ try:
+ return await processor.base_process_llm_request(
+ request=request,
+ fastapi_response=fastapi_response,
+ user_api_key_dict=user_api_key_dict,
+ route_type="acompact_responses",
+ proxy_logging_obj=proxy_logging_obj,
+ llm_router=llm_router,
+ general_settings=general_settings,
+ proxy_config=proxy_config,
+ select_data_generator=select_data_generator,
+ model=None,
+ user_model=user_model,
+ user_temperature=user_temperature,
+ user_request_timeout=user_request_timeout,
+ user_max_tokens=user_max_tokens,
+ user_api_base=user_api_base,
+ version=version,
+ )
+ except Exception as e:
+ raise await processor._handle_llm_api_exception(
+ e=e,
+ user_api_key_dict=user_api_key_dict,
+ proxy_logging_obj=proxy_logging_obj,
+ version=version,
+ )
+
+
@router.post(
"/v1/responses/{response_id}/cancel",
dependencies=[Depends(user_api_key_auth)],
diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py
index fd00cfc1c0a..a321e25a9a5 100644
--- a/litellm/proxy/route_llm_request.py
+++ b/litellm/proxy/route_llm_request.py
@@ -25,6 +25,7 @@ ROUTE_ENDPOINT_MAPPING = {
"alist_input_items": "/responses/{response_id}/input_items",
"aimage_edit": "/images/edits",
"acancel_responses": "/responses/{response_id}/cancel",
+ "acompact_responses": "/responses/compact",
"aocr": "/ocr",
"asearch": "/search",
"avideo_generation": "/videos",
@@ -116,6 +117,7 @@ async def route_request(
"aget_responses",
"adelete_responses",
"acancel_responses",
+ "acompact_responses",
"acreate_response_reply",
"alist_input_items",
"_arealtime", # private function for realtime API
diff --git a/litellm/responses/main.py b/litellm/responses/main.py
index e837346df23..8177b177fe6 100644
--- a/litellm/responses/main.py
+++ b/litellm/responses/main.py
@@ -1361,3 +1361,205 @@ def cancel_responses(
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
+
+
+@client
+async def acompact_responses(
+ input: Union[str, ResponseInputParam],
+ model: str,
+ instructions: Optional[str] = None,
+ previous_response_id: Optional[str] = None,
+ # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
+ # The extra values given here take precedence over values defined on the client or passed to this method.
+ extra_headers: Optional[Dict[str, Any]] = None,
+ extra_query: Optional[Dict[str, Any]] = None,
+ extra_body: Optional[Dict[str, Any]] = None,
+ timeout: Optional[Union[float, httpx.Timeout]] = None,
+ # LiteLLM specific params,
+ custom_llm_provider: Optional[str] = None,
+ **kwargs,
+) -> ResponsesAPIResponse:
+ """
+ Async version of the POST Compact Responses API
+
+ POST /v1/responses/compact endpoint in the responses API
+
+ Runs a compaction pass over a conversation, returning encrypted, opaque items.
+ """
+ local_vars = locals()
+ try:
+ loop = asyncio.get_event_loop()
+ kwargs["acompact_responses"] = True
+
+ # get custom llm provider so we can use this for mapping exceptions
+ if custom_llm_provider is None:
+ _, custom_llm_provider, _, _ = litellm.get_llm_provider(
+ model=model, api_base=local_vars.get("base_url", None)
+ )
+
+ func = partial(
+ compact_responses,
+ input=input,
+ model=model,
+ instructions=instructions,
+ previous_response_id=previous_response_id,
+ extra_headers=extra_headers,
+ extra_query=extra_query,
+ extra_body=extra_body,
+ timeout=timeout,
+ custom_llm_provider=custom_llm_provider,
+ **kwargs,
+ )
+
+ ctx = contextvars.copy_context()
+ func_with_context = partial(ctx.run, func)
+ init_response = await loop.run_in_executor(None, func_with_context)
+
+ if asyncio.iscoroutine(init_response):
+ response = await init_response
+ else:
+ response = init_response
+
+ # Update the responses_api_response_id with the model_id
+ if isinstance(response, ResponsesAPIResponse):
+ response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
+ responses_api_response=response,
+ litellm_metadata=kwargs.get("litellm_metadata", {}),
+ custom_llm_provider=custom_llm_provider,
+ )
+
+ return response
+ except Exception as e:
+ raise litellm.exception_type(
+ model=model,
+ custom_llm_provider=custom_llm_provider,
+ original_exception=e,
+ completion_kwargs=local_vars,
+ extra_kwargs=kwargs,
+ )
+
+
+@client
+def compact_responses(
+ input: Union[str, ResponseInputParam],
+ model: str,
+ instructions: Optional[str] = None,
+ previous_response_id: Optional[str] = None,
+ # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
+ # The extra values given here take precedence over values defined on the client or passed to this method.
+ extra_headers: Optional[Dict[str, Any]] = None,
+ extra_query: Optional[Dict[str, Any]] = None,
+ extra_body: Optional[Dict[str, Any]] = None,
+ timeout: Optional[Union[float, httpx.Timeout]] = None,
+ # LiteLLM specific params,
+ custom_llm_provider: Optional[str] = None,
+ **kwargs,
+) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]:
+ """
+ Synchronous version of the POST Compact Responses API
+
+ POST /v1/responses/compact endpoint in the responses API
+
+ Runs a compaction pass over a conversation, returning encrypted, opaque items.
+ """
+ local_vars = locals()
+ try:
+ litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
+ litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
+ _is_async = kwargs.pop("acompact_responses", False) is True
+
+ # get llm provider logic
+ litellm_params = GenericLiteLLMParams(**kwargs)
+
+ (
+ model,
+ custom_llm_provider,
+ dynamic_api_key,
+ dynamic_api_base,
+ ) = litellm.get_llm_provider(
+ model=model,
+ custom_llm_provider=custom_llm_provider,
+ api_base=litellm_params.api_base,
+ api_key=litellm_params.api_key,
+ )
+
+ if custom_llm_provider is None:
+ raise ValueError("custom_llm_provider is required but passed as None")
+
+ # get provider config
+ responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
+ ProviderConfigManager.get_provider_responses_api_config(
+ model=model,
+ provider=litellm.LlmProviders(custom_llm_provider),
+ )
+ )
+
+ if responses_api_provider_config is None:
+ raise ValueError(
+ f"COMPACT responses is not supported for {custom_llm_provider}"
+ )
+
+ local_vars.update(kwargs)
+
+ # Build optional params for compact endpoint
+ response_api_optional_params: ResponsesAPIOptionalRequestParams = (
+ ResponsesAPIRequestUtils.get_requested_response_api_optional_param(
+ local_vars
+ )
+ )
+
+ # Get optional parameters for the responses API
+ responses_api_request_params: Dict = (
+ ResponsesAPIRequestUtils.get_optional_params_responses_api(
+ model=model,
+ responses_api_provider_config=responses_api_provider_config,
+ response_api_optional_params=response_api_optional_params,
+ allowed_openai_params=None,
+ )
+ )
+
+ # Pre Call logging
+ litellm_logging_obj.update_environment_variables(
+ model=model,
+ optional_params=dict(responses_api_request_params),
+ litellm_params={
+ **responses_api_request_params,
+ "litellm_call_id": litellm_call_id,
+ },
+ custom_llm_provider=custom_llm_provider,
+ )
+
+ # Call the handler with _is_async flag instead of directly calling the async handler
+ response = base_llm_http_handler.compact_response_api_handler(
+ model=model,
+ input=input,
+ responses_api_provider_config=responses_api_provider_config,
+ response_api_optional_request_params=responses_api_request_params,
+ litellm_params=litellm_params,
+ logging_obj=litellm_logging_obj,
+ custom_llm_provider=custom_llm_provider,
+ extra_headers=extra_headers,
+ extra_body=extra_body,
+ timeout=timeout or request_timeout,
+ _is_async=_is_async,
+ client=kwargs.get("client"),
+ shared_session=kwargs.get("shared_session"),
+ )
+
+ # Update the responses_api_response_id with the model_id
+ if isinstance(response, ResponsesAPIResponse):
+ response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
+ responses_api_response=response,
+ litellm_metadata=kwargs.get("litellm_metadata", {}),
+ custom_llm_provider=custom_llm_provider,
+ )
+
+ return response
+ except Exception as e:
+ raise litellm.exception_type(
+ model=model,
+ custom_llm_provider=custom_llm_provider,
+ original_exception=e,
+ completion_kwargs=local_vars,
+ extra_kwargs=kwargs,
+ )
diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py
index 0407776029d..0b838f916e2 100644
--- a/litellm/responses/streaming_iterator.py
+++ b/litellm/responses/streaming_iterator.py
@@ -1,5 +1,6 @@
import asyncio
import json
+import traceback
from datetime import datetime
from typing import Any, Dict, Optional
@@ -11,6 +12,9 @@ from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_base
+from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
+ update_response_metadata,
+)
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.utils import ResponsesAPIRequestUtils
@@ -22,7 +26,8 @@ from litellm.types.llms.openai import (
ResponsesAPIStreamEvents,
ResponsesAPIStreamingResponse,
)
-from litellm.utils import CustomStreamWrapper
+from litellm.types.utils import CallTypes
+from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook
class BaseResponsesAPIStreamingIterator:
@@ -40,6 +45,8 @@ class BaseResponsesAPIStreamingIterator:
logging_obj: LiteLLMLoggingObj,
litellm_metadata: Optional[Dict[str, Any]] = None,
custom_llm_provider: Optional[str] = None,
+ request_data: Optional[Dict[str, Any]] = None,
+ call_type: Optional[str] = None,
):
self.response = response
self.model = model
@@ -47,21 +54,25 @@ class BaseResponsesAPIStreamingIterator:
self.finished = False
self.responses_api_provider_config = responses_api_provider_config
self.completed_response: Optional[ResponsesAPIStreamingResponse] = None
- self.start_time = datetime.now()
+ self.start_time = getattr(logging_obj, "start_time", datetime.now())
- # set request kwargs
+ # track request context for hooks
self.litellm_metadata = litellm_metadata
self.custom_llm_provider = custom_llm_provider
+ self.request_data: Dict[str, Any] = request_data or {}
+ self.call_type: Optional[str] = call_type
# set hidden params for response headers (e.g., x-litellm-model-id)
- # This matches ths stream wrapper in litellm/litellm_core_utils/streaming_handler.py
+ # This matches the stream wrapper in litellm/litellm_core_utils/streaming_handler.py
_api_base = get_api_base(
model=model or "",
optional_params=self.logging_obj.model_call_details.get(
"litellm_params", {}
),
)
- _model_info: Dict = litellm_metadata.get("model_info", {}) if litellm_metadata else {}
+ _model_info: Dict = (
+ litellm_metadata.get("model_info", {}) if litellm_metadata else {}
+ )
self._hidden_params = {
"model_id": _model_info.get("id", None),
"api_base": _api_base,
@@ -102,13 +113,21 @@ class BaseResponsesAPIStreamingIterator:
# if "response" in parsed_chunk, then encode litellm specific information like custom_llm_provider
response_object = getattr(openai_responses_api_chunk, "response", None)
if response_object:
- response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
- responses_api_response=response_object,
- litellm_metadata=self.litellm_metadata,
- custom_llm_provider=self.custom_llm_provider,
+ response = (
+ ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
+ responses_api_response=response_object,
+ litellm_metadata=self.litellm_metadata,
+ custom_llm_provider=self.custom_llm_provider,
+ )
)
setattr(openai_responses_api_chunk, "response", response)
+ # Allow callbacks to modify chunk before returning
+ openai_responses_api_chunk = run_async_function(
+ async_function=self._call_post_streaming_deployment_hook,
+ chunk=openai_responses_api_chunk,
+ )
+
# Store the completed response
if (
openai_responses_api_chunk
@@ -149,11 +168,159 @@ class BaseResponsesAPIStreamingIterator:
except json.JSONDecodeError:
# If we can't parse the chunk, continue
return None
+ except Exception as e:
+ # Ensure failures trigger failure hooks
+ self._handle_failure(e)
+ raise
def _handle_logging_completed_response(self):
"""Base implementation - should be overridden by subclasses"""
pass
+ async def _call_post_streaming_deployment_hook(self, chunk):
+ """
+ Allow callbacks to modify streaming chunks before returning (parity with chat).
+ """
+ try:
+ # Align with chat pipeline: use logging_obj model_call_details + call_type
+ typed_call_type: Optional[CallTypes] = None
+ if self.call_type is not None:
+ try:
+ typed_call_type = CallTypes(self.call_type)
+ except ValueError:
+ typed_call_type = None
+ if typed_call_type is None:
+ try:
+ typed_call_type = CallTypes(getattr(self.logging_obj, "call_type", None))
+ except Exception:
+ typed_call_type = None
+
+ request_data = self.request_data or getattr(
+ self.logging_obj, "model_call_details", {}
+ )
+ callbacks = getattr(litellm, "callbacks", None) or []
+ hooks_ran = False
+ for callback in callbacks:
+ if hasattr(callback, "async_post_call_streaming_deployment_hook"):
+ hooks_ran = True
+ result = await callback.async_post_call_streaming_deployment_hook(
+ request_data=request_data,
+ response_chunk=chunk,
+ call_type=typed_call_type,
+ )
+ if result is not None:
+ chunk = result
+ if hooks_ran:
+ setattr(chunk, "_post_streaming_hooks_ran", True)
+ return chunk
+ except Exception:
+ return chunk
+
+ async def call_post_streaming_hooks_for_testing(self, chunk):
+ """
+ Helper to invoke streaming deployment hooks explicitly (used in tests).
+ """
+ return await self._call_post_streaming_deployment_hook(chunk)
+
+ def _run_post_success_hooks(self, end_time: datetime):
+ """
+ Run post-call deployment hooks and update metadata similar to chat pipeline.
+ """
+ if self.completed_response is None:
+ return
+
+ request_payload: Dict[str, Any] = {}
+ if isinstance(self.request_data, dict):
+ request_payload.update(self.request_data)
+ try:
+ if hasattr(self.logging_obj, "model_call_details"):
+ request_payload.update(self.logging_obj.model_call_details)
+ except Exception:
+ pass
+ if "litellm_params" not in request_payload:
+ try:
+ request_payload["litellm_params"] = getattr(
+ self.logging_obj, "model_call_details", {}
+ ).get("litellm_params", {})
+ except Exception:
+ request_payload["litellm_params"] = {}
+
+ try:
+ update_response_metadata(
+ result=self.completed_response,
+ logging_obj=self.logging_obj,
+ model=self.model,
+ kwargs=request_payload,
+ start_time=self.start_time,
+ end_time=end_time,
+ )
+ except Exception:
+ # Non-blocking
+ pass
+
+ try:
+ typed_call_type: Optional[CallTypes] = None
+ if self.call_type is not None:
+ try:
+ typed_call_type = CallTypes(self.call_type)
+ except ValueError:
+ typed_call_type = None
+ except Exception:
+ typed_call_type = None
+ if typed_call_type is None:
+ try:
+ typed_call_type = CallTypes.responses
+ except Exception:
+ typed_call_type = None
+
+ try:
+ # Call synchronously; async hook will be executed via asyncio.run in a new loop
+ run_async_function(
+ async_function=async_post_call_success_deployment_hook,
+ request_data=request_payload,
+ response=self.completed_response,
+ call_type=typed_call_type,
+ )
+ except Exception:
+ pass
+
+ def _handle_failure(self, exception: Exception):
+ """
+ Trigger failure handlers before bubbling the exception.
+ """
+ traceback_exception = traceback.format_exc()
+ try:
+ run_async_function(
+ async_function=self.logging_obj.async_failure_handler,
+ exception=exception,
+ traceback_exception=traceback_exception,
+ start_time=self.start_time,
+ end_time=datetime.now(),
+ )
+ except Exception:
+ pass
+
+ try:
+ executor.submit(
+ self.logging_obj.failure_handler,
+ exception,
+ traceback_exception,
+ self.start_time,
+ datetime.now(),
+ )
+ except Exception:
+ pass
+
+
+async def call_post_streaming_hooks_for_testing(iterator, chunk):
+ """
+ Module-level helper for tests to ensure hooks can be invoked even if the iterator is wrapped.
+ """
+ hook_fn = getattr(iterator, "_call_post_streaming_deployment_hook", None)
+ if hook_fn is None:
+ return chunk
+ return await hook_fn(chunk)
+
class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
"""
@@ -168,6 +335,8 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
logging_obj: LiteLLMLoggingObj,
litellm_metadata: Optional[Dict[str, Any]] = None,
custom_llm_provider: Optional[str] = None,
+ request_data: Optional[Dict[str, Any]] = None,
+ call_type: Optional[str] = None,
):
super().__init__(
response,
@@ -176,6 +345,8 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
logging_obj,
litellm_metadata,
custom_llm_provider,
+ request_data,
+ call_type,
)
self.stream_iterator = response.aiter_lines()
@@ -203,16 +374,21 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
except httpx.HTTPError as e:
# Handle HTTP errors
self.finished = True
+ self._handle_failure(e)
+ raise e
+ except Exception as e:
+ self.finished = True
+ self._handle_failure(e)
raise e
def _handle_logging_completed_response(self):
"""Handle logging for completed responses in async context"""
# Create a deep copy for logging to avoid modifying the response object that will be returned to the user
- # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
+ # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
import copy
logging_response = copy.deepcopy(self.completed_response)
-
+
asyncio.create_task(
self.logging_obj.async_success_handler(
result=logging_response,
@@ -229,6 +405,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
start_time=self.start_time,
end_time=datetime.now(),
)
+ self._run_post_success_hooks(end_time=datetime.now())
class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
@@ -244,6 +421,8 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
logging_obj: LiteLLMLoggingObj,
litellm_metadata: Optional[Dict[str, Any]] = None,
custom_llm_provider: Optional[str] = None,
+ request_data: Optional[Dict[str, Any]] = None,
+ call_type: Optional[str] = None,
):
super().__init__(
response,
@@ -252,6 +431,8 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
logging_obj,
litellm_metadata,
custom_llm_provider,
+ request_data,
+ call_type,
)
self.stream_iterator = response.iter_lines()
@@ -279,16 +460,21 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
except httpx.HTTPError as e:
# Handle HTTP errors
self.finished = True
+ self._handle_failure(e)
+ raise e
+ except Exception as e:
+ self.finished = True
+ self._handle_failure(e)
raise e
def _handle_logging_completed_response(self):
"""Handle logging for completed responses in sync context"""
# Create a deep copy for logging to avoid modifying the response object that will be returned to the user
- # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
+ # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
import copy
logging_response = copy.deepcopy(self.completed_response)
-
+
run_async_function(
async_function=self.logging_obj.async_success_handler,
result=logging_response,
@@ -304,6 +490,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
start_time=self.start_time,
end_time=datetime.now(),
)
+ self._run_post_success_hooks(end_time=datetime.now())
class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
@@ -324,6 +511,8 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
logging_obj: LiteLLMLoggingObj,
litellm_metadata: Optional[Dict[str, Any]] = None,
custom_llm_provider: Optional[str] = None,
+ request_data: Optional[Dict[str, Any]] = None,
+ call_type: Optional[str] = None,
):
super().__init__(
response=response,
@@ -332,6 +521,8 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
logging_obj=logging_obj,
litellm_metadata=litellm_metadata,
custom_llm_provider=custom_llm_provider,
+ request_data=request_data,
+ call_type=call_type,
)
# one-time transform
diff --git a/litellm/router.py b/litellm/router.py
index 6821ab9e6c6..d980b5f74d8 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -713,6 +713,23 @@ class Router:
self, routing_strategy: Union[RoutingStrategy, str], routing_strategy_args: dict
):
verbose_router_logger.info(f"Routing strategy: {routing_strategy}")
+
+ # Validate routing_strategy value to fail fast with helpful error
+ # See: https://github.com/BerriAI/litellm/issues/11330
+ # Derive valid strategies from RoutingStrategy enum + "simple-shuffle" (default, not in enum)
+ valid_strategy_strings = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
+
+ if routing_strategy is not None:
+ is_valid_string = isinstance(routing_strategy, str) and routing_strategy in valid_strategy_strings
+ is_valid_enum = isinstance(routing_strategy, RoutingStrategy)
+ if not is_valid_string and not is_valid_enum:
+ raise ValueError(
+ f"Invalid routing_strategy: '{routing_strategy}'. "
+ f"Valid options: {valid_strategy_strings}. "
+ f"Check 'router_settings.routing_strategy' in your config.yaml "
+ f"or the 'routing_strategy' parameter if using the Router SDK directly."
+ )
+
if (
routing_strategy == RoutingStrategy.LEAST_BUSY.value
or routing_strategy == RoutingStrategy.LEAST_BUSY
@@ -812,6 +829,9 @@ class Router:
self.acancel_responses = self.factory_function(
litellm.acancel_responses, call_type="acancel_responses"
)
+ self.acompact_responses = self.factory_function(
+ litellm.acompact_responses, call_type="acompact_responses"
+ )
self.adelete_responses = self.factory_function(
litellm.adelete_responses, call_type="adelete_responses"
)
@@ -3924,6 +3944,7 @@ class Router:
"anthropic_messages",
"aresponses",
"acancel_responses",
+ "acompact_responses",
"responses",
"aget_responses",
"adelete_responses",
@@ -4152,6 +4173,7 @@ class Router:
elif call_type in (
"aget_responses",
"acancel_responses",
+ "acompact_responses",
"adelete_responses",
"alist_input_items",
):
diff --git a/litellm/types/integrations/langsmith.py b/litellm/types/integrations/langsmith.py
index 23f760ecf32..9c026a117fd 100644
--- a/litellm/types/integrations/langsmith.py
+++ b/litellm/types/integrations/langsmith.py
@@ -31,6 +31,7 @@ class LangsmithCredentialsObject(TypedDict):
LANGSMITH_API_KEY: Optional[str]
LANGSMITH_PROJECT: Optional[str]
LANGSMITH_BASE_URL: str
+ LANGSMITH_TENANT_ID: Optional[str]
class LangsmithQueueObject(TypedDict):
@@ -52,6 +53,7 @@ class CredentialsKey(NamedTuple):
api_key: str
project: str
base_url: str
+ tenant_id: Optional[str]
@dataclass
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 3eec67d9d26..784c8403c3f 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -2677,6 +2677,7 @@ class StandardCallbackDynamicParams(TypedDict, total=False):
langsmith_project: Optional[str]
langsmith_base_url: Optional[str]
langsmith_sampling_rate: Optional[float]
+ langsmith_tenant_id: Optional[str]
# Humanloop dynamic params
humanloop_api_key: Optional[str]
@@ -2946,6 +2947,7 @@ class LlmProviders(str, Enum):
MISTRAL = "mistral"
MILVUS = "milvus"
GROQ = "groq"
+ GIGACHAT = "gigachat"
NVIDIA_NIM = "nvidia_nim"
CEREBRAS = "cerebras"
AI21_CHAT = "ai21_chat"
diff --git a/litellm/utils.py b/litellm/utils.py
index 6b9aeca0934..fbbaa94f7a1 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -7521,6 +7521,8 @@ class ProviderConfigManager:
return litellm.CompactifAIChatConfig()
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
return litellm.GithubCopilotConfig()
+ elif litellm.LlmProviders.GIGACHAT == provider:
+ return litellm.GigaChatConfig()
elif litellm.LlmProviders.RAGFLOW == provider:
return litellm.RAGFlowConfig()
elif (
@@ -7716,6 +7718,8 @@ class ProviderConfigManager:
return litellm.CometAPIEmbeddingConfig()
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
return litellm.GithubCopilotEmbeddingConfig()
+ elif litellm.LlmProviders.GIGACHAT == provider:
+ return litellm.GigaChatEmbeddingConfig()
elif litellm.LlmProviders.SAGEMAKER == provider:
from litellm.llms.sagemaker.embedding.transformation import (
SagemakerEmbeddingConfig,
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index e823dd5dc6b..c7a2f60856d 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -15831,6 +15831,68 @@
"max_tokens": 8191,
"mode": "embedding"
},
+ "gigachat/GigaChat-2-Lite": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0,
+ "supports_function_calling": true,
+ "supports_system_messages": true
+ },
+ "gigachat/GigaChat-2-Max": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0,
+ "supports_function_calling": true,
+ "supports_system_messages": true,
+ "supports_vision": true
+ },
+ "gigachat/GigaChat-2-Pro": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "chat",
+ "output_cost_per_token": 0.0,
+ "supports_function_calling": true,
+ "supports_system_messages": true,
+ "supports_vision": true
+ },
+ "gigachat/Embeddings": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 512,
+ "max_tokens": 512,
+ "mode": "embedding",
+ "output_cost_per_token": 0.0,
+ "output_vector_size": 1024
+ },
+ "gigachat/Embeddings-2": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 512,
+ "max_tokens": 512,
+ "mode": "embedding",
+ "output_cost_per_token": 0.0,
+ "output_vector_size": 1024
+ },
+ "gigachat/EmbeddingsGigaR": {
+ "input_cost_per_token": 0.0,
+ "litellm_provider": "gigachat",
+ "max_input_tokens": 4096,
+ "max_tokens": 4096,
+ "mode": "embedding",
+ "output_cost_per_token": 0.0,
+ "output_vector_size": 2560
+ },
"google.gemma-3-12b-it": {
"input_cost_per_token": 9e-08,
"litellm_provider": "bedrock_converse",
@@ -32092,3 +32154,4 @@
"mode": "chat"
}
}
+
diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json
index 45ee47c01bc..bc5dea7b97c 100644
--- a/provider_endpoints_support.json
+++ b/provider_endpoints_support.json
@@ -28,7 +28,8 @@
"list_container_files": "Supports GET /containers/{id}/files endpoint",
"retrieve_container_file": "Supports GET /containers/{id}/files/{file_id} endpoint",
"retrieve_container_file_content": "Supports GET /containers/{id}/files/{file_id}/content endpoint",
- "delete_container_file": "Supports DELETE /containers/{id}/files/{file_id} endpoint"
+ "delete_container_file": "Supports DELETE /containers/{id}/files/{file_id} endpoint",
+ "compact": "Supports /responses/compact endpoint"
}
}
},
@@ -1519,6 +1520,7 @@
"retrieve_container_file": true,
"retrieve_container_file_content": true,
"delete_container_file": true,
+ "compact": true,
"a2a": true,
"interactions": true
}
diff --git a/requirements.txt b/requirements.txt
index 06a7c17336c..249b899b86b 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -20,7 +20,7 @@ google-cloud-aiplatform==1.47.0 # for vertex ai calls
google-cloud-iam==2.19.1 # for GCP IAM Redis authentication
google-genai==1.22.0
anthropic[vertex]==0.54.0
-mcp==1.23.0 ; python_version >= "3.10" # for MCP server
+mcp==1.25.0 ; python_version >= "3.10" # for MCP server
google-generativeai==0.5.0 # for vertex ai calls
async_generator==1.10.0 # for async ollama calls
langfuse==2.59.7 # for langfuse self-hosted logging
diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py
index 9a96919da87..60d4f479733 100644
--- a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py
+++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py
@@ -61,13 +61,19 @@ async def test_bedrock_apply_guardrail_blocked():
guardrailVersion="DRAFT",
)
- # Mock the make_bedrock_api_request method
+ # Mock the make_bedrock_api_request method to raise an exception for blocked content
with patch.object(
- guardrail, "make_bedrock_api_request", new_callable=AsyncMock
+ guardrail, "make_bedrock_api_request", new_callable=AsyncMock
) as mock_api_request:
- # Mock a blocked response from Bedrock
- mock_response = {"action": "BLOCKED", "reason": "Content violates policy"}
- mock_api_request.return_value = mock_response
+ # Mock the method to raise an HTTPException as it would for blocked content
+ from fastapi import HTTPException
+ mock_api_request.side_effect = HTTPException(
+ status_code=400,
+ detail={
+ "error": "Violated guardrail policy",
+ "bedrock_guardrail_response": "",
+ },
+ )
# Test the apply_guardrail method should raise an exception
with pytest.raises(Exception) as exc_info:
@@ -77,8 +83,9 @@ async def test_bedrock_apply_guardrail_blocked():
input_type="request",
)
- assert "Content blocked by Bedrock guardrail" in str(exc_info.value)
- assert "Content violates policy" in str(exc_info.value)
+ # The apply_guardrail method wraps the original exception in a generic Exception
+ assert "Bedrock guardrail failed:" in str(exc_info.value)
+ assert "Violated guardrail policy" in str(exc_info.value)
@pytest.mark.asyncio
@@ -253,7 +260,15 @@ async def test_bedrock_apply_guardrail_filters_request_messages_when_flag_enable
with patch.object(
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
) as mock_api:
- mock_api.return_value = {"action": "BLOCKED", "reason": "policy"}
+ # Mock the method to raise an HTTPException as it would for blocked content
+ from fastapi import HTTPException
+ mock_api.side_effect = HTTPException(
+ status_code=400,
+ detail={
+ "error": "Violated guardrail policy",
+ "bedrock_guardrail_response": "policy",
+ },
+ )
with pytest.raises(Exception, match="policy") as exc_info:
await guardrail.apply_guardrail(
@@ -265,7 +280,8 @@ async def test_bedrock_apply_guardrail_filters_request_messages_when_flag_enable
assert mock_api.called
_, kwargs = mock_api.call_args
assert kwargs["messages"] == [request_messages[-1]]
- assert "Content blocked by Bedrock guardrail" in str(exc_info.value)
+ # The apply_guardrail method wraps the original exception in a generic Exception
+ assert "Bedrock guardrail failed:" in str(exc_info.value)
def test_bedrock_guardrail_filters_latest_user_message_when_enabled():
diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
index 5714cd5c487..c981ccd8dac 100644
--- a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
+++ b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
@@ -1,11 +1,14 @@
import os
import sys
+import time
+from unittest.mock import Mock, patch
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
-from litellm_proxy_extras.utils import ProxyExtrasDBManager
+from litellm_proxy_extras.utils import ProxyExtrasDBManager, MigrationLockManager
+
def test_custom_prisma_dir(monkeypatch):
@@ -27,101 +30,279 @@ def test_custom_prisma_dir(monkeypatch):
assert os.path.exists(migrations_dir)
-class TestPermissionErrorDetection:
- """Test cases for permission error detection in Prisma migrations"""
+class TestMigrationLockManager:
+ """Test cases for MigrationLockManager"""
- def test_is_permission_error_postgres_42501(self):
- """Test detection of PostgreSQL 42501 error code (insufficient privilege)"""
- error_message = "Database error code: 42501 - permission denied for table users"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is True
+ def test_acquire_lock_without_redis(self):
+ """Test lock acquisition when Redis is not available"""
+ lock_manager = MigrationLockManager()
+ result = lock_manager.acquire_lock()
+ assert result is True # Should return True when Redis is not available
+ assert lock_manager.lock_acquired is True # Redis 없을 때도 lock_acquired는 True
- def test_is_permission_error_must_be_owner(self):
- """Test detection of 'must be owner of table' error"""
- error_message = "ERROR: must be owner of table my_table"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is True
+ def test_acquire_lock_with_redis_success(self):
+ """Test successful lock acquisition with Redis"""
+ mock_redis = Mock()
+ mock_redis.set_cache.return_value = True
+ lock_manager = MigrationLockManager(mock_redis)
- def test_is_permission_error_permission_denied_schema(self):
- """Test detection of 'permission denied for schema' error"""
- error_message = "permission denied for schema public"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is True
+ result = lock_manager.acquire_lock()
- def test_is_permission_error_permission_denied_table(self):
- """Test detection of 'permission denied for table' error"""
- error_message = "permission denied for table my_table"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is True
+ assert result is True
+ assert lock_manager.lock_acquired is True
+ mock_redis.set_cache.assert_called_once()
- def test_is_permission_error_must_be_owner_schema(self):
- """Test detection of 'must be owner of schema' error"""
- error_message = "must be owner of schema public"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is True
+ def test_acquire_lock_with_redis_failure(self):
+ """Test failed lock acquisition with Redis"""
+ mock_redis = Mock()
+ mock_redis.set_cache.return_value = False
+ lock_manager = MigrationLockManager(mock_redis)
- def test_is_permission_error_case_insensitive(self):
- """Test that permission error detection is case insensitive"""
- error_message = "PERMISSION DENIED FOR TABLE my_table"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is True
+ result = lock_manager.acquire_lock()
- def test_is_permission_error_negative(self):
- """Test that non-permission errors are not detected as permission errors"""
- error_message = "column 'id' already exists"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is False
+ assert result is False
+ assert lock_manager.lock_acquired is False
+ mock_redis.set_cache.assert_called_once()
+
+ def test_acquire_lock_with_redis_exception(self):
+ """Test lock acquisition with Redis exception"""
+ mock_redis = Mock()
+ mock_redis.set_cache.side_effect = Exception("Redis error")
+ lock_manager = MigrationLockManager(mock_redis)
+
+ result = lock_manager.acquire_lock()
+
+ assert result is False
+ assert lock_manager.lock_acquired is False
+
+ def test_wait_for_lock_release_success(self):
+ """Test successful waiting for lock release"""
+ mock_redis = Mock()
+ # First call returns False (lock held), second call returns True (lock acquired)
+ mock_redis.set_cache.side_effect = [False, True]
+ lock_manager = MigrationLockManager(mock_redis)
+
+ result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=1)
+
+ assert result is True
+ assert lock_manager.lock_acquired is True
+ assert mock_redis.set_cache.call_count == 2
+
+ def test_wait_for_lock_release_timeout(self):
+ """Test timeout while waiting for lock release"""
+ mock_redis = Mock()
+ mock_redis.set_cache.return_value = False # Lock always held
+ lock_manager = MigrationLockManager(mock_redis)
+
+ result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=0.2)
+
+ assert result is False
+ assert lock_manager.lock_acquired is False
+
+ def test_release_lock_not_acquired(self):
+ """Test releasing lock when not acquired"""
+ mock_redis = Mock()
+ lock_manager = MigrationLockManager(mock_redis)
+
+ lock_manager.release_lock()
+
+ mock_redis.get_cache.assert_not_called()
+ mock_redis.delete_cache.assert_not_called()
+
+ def test_release_lock_success(self):
+ """Test successful lock release"""
+ mock_redis = Mock()
+ mock_redis.get_cache.return_value = "pod_123_456"
+ lock_manager = MigrationLockManager(mock_redis)
+ lock_manager.pod_id = "pod_123_456"
+ lock_manager.lock_acquired = True
+
+ lock_manager.release_lock()
+
+ mock_redis.get_cache.assert_called_once()
+ mock_redis.delete_cache.assert_called_once()
+ assert lock_manager.lock_acquired is False
+
+ def test_release_lock_wrong_owner(self):
+ """Test releasing lock when not the owner"""
+ mock_redis = Mock()
+ mock_redis.get_cache.return_value = "pod_999_999" # Different pod
+ lock_manager = MigrationLockManager(mock_redis)
+ lock_manager.pod_id = "pod_123_456"
+ lock_manager.lock_acquired = True
+
+ lock_manager.release_lock()
+
+ mock_redis.get_cache.assert_called_once()
+ mock_redis.delete_cache.assert_not_called()
+ assert lock_manager.lock_acquired is False
+
+ def test_context_manager(self):
+ """Test MigrationLockManager as context manager"""
+ mock_redis = Mock()
+ mock_redis.set_cache.return_value = True
+ # Mock get_cache to return the same pod_id for successful release
+ mock_redis.get_cache.return_value = "pod_123_456"
+
+ lock_manager = MigrationLockManager(mock_redis)
+ lock_manager.pod_id = "pod_123_456" # Set consistent pod_id
+
+ with lock_manager:
+ assert lock_manager.lock_acquired is True
+
+ # Should call release_lock when exiting context
+ mock_redis.get_cache.assert_called_once()
+ mock_redis.delete_cache.assert_called_once()
-class TestIdempotentErrorDetection:
- """Test cases for idempotent error detection in Prisma migrations"""
+class TestProxyExtrasDBManagerMigrationLock:
+ """Test cases for ProxyExtrasDBManager with migration locking"""
- def test_is_idempotent_error_already_exists(self):
- """Test detection of generic 'already exists' error"""
- error_message = "object already exists"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
+ @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations")
+ @patch("litellm_proxy_extras.utils.subprocess.run")
+ @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
+ @patch("os.chdir")
+ def test_setup_database_with_redis_lock_success(
+ self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations
+ ):
+ """Test successful database setup with Redis lock"""
+ # Setup mocks
+ mock_get_prisma_dir.return_value = "/test/prisma"
+ mock_subprocess.return_value = Mock(stdout="Migration completed", stderr="")
+ mock_resolve_migrations.return_value = None
- def test_is_idempotent_error_column_already_exists(self):
- """Test detection of 'column already exists' error"""
- error_message = "column 'email' already exists"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
+ # Mock Redis cache
+ mock_redis = Mock()
+ mock_redis.set_cache.return_value = True # Lock acquired successfully
+ mock_redis.get_cache.return_value = "pod_123_456"
- def test_is_idempotent_error_duplicate_key(self):
- """Test detection of duplicate key violation error"""
- error_message = "duplicate key value violates unique constraint"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
+ # Set DATABASE_URL
+ with patch.dict(
+ os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
+ ):
+ result = ProxyExtrasDBManager.setup_database(
+ use_migrate=True, redis_cache=mock_redis
+ )
- def test_is_idempotent_error_relation_already_exists(self):
- """Test detection of 'relation already exists' error"""
- error_message = "relation 'users_pkey' already exists"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
+ assert result is True
+ # set_cache is called once in acquire_lock (__enter__ calls acquire_lock)
+ assert mock_redis.set_cache.call_count == 1
+ mock_subprocess.assert_called_once()
- def test_is_idempotent_error_constraint_already_exists(self):
- """Test detection of 'constraint already exists' error"""
- error_message = "constraint 'fk_user_id' already exists"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
+ @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations")
+ @patch("litellm_proxy_extras.utils.subprocess.run")
+ @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
+ @patch("os.chdir")
+ def test_setup_database_with_redis_lock_wait_and_skip(
+ self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations
+ ):
+ """Test database setup when lock is held by another pod, then acquired after waiting"""
+ # Setup mocks
+ mock_get_prisma_dir.return_value = "/test/prisma"
+ mock_resolve_migrations.return_value = None
- def test_is_idempotent_error_case_insensitive(self):
- """Test that idempotent error detection is case insensitive"""
- error_message = "COLUMN 'ID' ALREADY EXISTS"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
+ # Mock Redis cache - first call fails, second call succeeds
+ mock_redis = Mock()
+ mock_redis.set_cache.side_effect = [False, True] # First fails, then succeeds
+ mock_redis.get_cache.return_value = "pod_123_456"
- def test_is_idempotent_error_negative(self):
- """Test that non-idempotent errors are not detected as idempotent errors"""
- error_message = "Database error code: 42501 - permission denied"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False
+ # Set DATABASE_URL
+ with patch.dict(
+ os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
+ ):
+ result = ProxyExtrasDBManager.setup_database(
+ use_migrate=True, redis_cache=mock_redis
+ )
+ assert result is True # Should return True after waiting and acquiring lock
+ # set_cache is called 2 times: once in __enter__, once in wait_for_lock_release
+ assert mock_redis.set_cache.call_count == 2
+ # Proceed for case handling in case of migration failure
+ mock_subprocess.assert_called_once()
-class TestErrorClassificationPriority:
- """Test cases to ensure errors are correctly classified"""
+ @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations")
+ @patch("litellm_proxy_extras.utils.subprocess.run")
+ @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
+ @patch("os.chdir")
+ def test_setup_database_without_redis(
+ self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations
+ ):
+ """Test database setup without Redis cache"""
+ # Setup mocks
+ mock_get_prisma_dir.return_value = "/test/prisma"
+ mock_subprocess.return_value = Mock(stdout="Migration completed", stderr="")
+ mock_resolve_migrations.return_value = None
- def test_permission_error_not_classified_as_idempotent(self):
- """Ensure permission errors are not mistakenly classified as idempotent"""
- error_message = "Database error code: 42501 - must be owner of table users"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is True
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False
+ # Set DATABASE_URL
+ with patch.dict(
+ os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
+ ):
+ result = ProxyExtrasDBManager.setup_database(
+ use_migrate=True, redis_cache=None
+ )
- def test_idempotent_error_not_classified_as_permission(self):
- """Ensure idempotent errors are not mistakenly classified as permission errors"""
- error_message = "column 'created_at' already exists"
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
- assert ProxyExtrasDBManager._is_permission_error(error_message) is False
+ assert result is True
+ # Redis가 없을 때는 락 보호 없이 마이그레이션을 실행해야 함
+ mock_subprocess.assert_called_once()
- def test_unknown_error_classified_as_neither(self):
- """Ensure unknown errors are classified as neither permission nor idempotent"""
- error_message = "connection timeout"
- assert ProxyExtrasDBManager._is_permission_error(error_message) is False
- assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False
+ def test_setup_database_no_database_url(self):
+ """Test database setup without DATABASE_URL"""
+ with patch.dict(os.environ, {}, clear=True):
+ result = ProxyExtrasDBManager.setup_database(
+ use_migrate=True, redis_cache=None
+ )
+
+ assert result is False
+
+ @patch("litellm_proxy_extras.utils.subprocess.run")
+ @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
+ @patch("os.chdir")
+ @patch.object(
+ MigrationLockManager, "LOCK_TTL_SECONDS", 1
+ ) # Set short TTL for testing
+ def test_setup_database_lock_timeout(
+ self, mock_chdir, mock_get_prisma_dir, mock_subprocess
+ ):
+ """Test database setup when lock acquisition times out"""
+ # Setup mocks
+ mock_get_prisma_dir.return_value = "/test/prisma"
+
+ # Mock Redis cache - always fails to acquire lock
+ mock_redis = Mock()
+ mock_redis.set_cache.return_value = False # Always fails
+
+ # Set DATABASE_URL
+ with patch.dict(
+ os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
+ ):
+ # Patch the wait_for_lock_release method to use shorter timeout
+ with patch.object(
+ MigrationLockManager, "wait_for_lock_release"
+ ) as mock_wait:
+ mock_wait.return_value = False # Simulate timeout
+
+ result = ProxyExtrasDBManager.setup_database(
+ use_migrate=True, redis_cache=mock_redis
+ )
+
+ assert result is False # Should return False after timeout
+ mock_subprocess.assert_not_called() # Should not run migration
+ # Verify that wait_for_lock_release was called with default parameters
+ mock_wait.assert_called_once_with()
+
+ def test_wait_for_lock_release_actual_timeout(self):
+ """Test actual timeout behavior of wait_for_lock_release with real timing"""
+ mock_redis = Mock()
+ mock_redis.set_cache.return_value = False # Always fails to acquire lock
+ lock_manager = MigrationLockManager(mock_redis)
+
+ # Test with very short timeout to verify actual timeout behavior
+ start_time = time.time()
+ result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=0.5)
+ end_time = time.time()
+
+ assert result is False # Should timeout
+ assert end_time - start_time >= 0.5 # Should wait at least the max_wait time
+ assert end_time - start_time < 1.0 # But not too much longer
+ # Should have called set_cache multiple times during the wait
+ assert mock_redis.set_cache.call_count > 1
diff --git a/tests/llm_responses_api_testing/test_openai_responses_api.py b/tests/llm_responses_api_testing/test_openai_responses_api.py
index 7553c670774..5f35d6837c0 100644
--- a/tests/llm_responses_api_testing/test_openai_responses_api.py
+++ b/tests/llm_responses_api_testing/test_openai_responses_api.py
@@ -1814,3 +1814,49 @@ async def test_extra_body_merges_with_request_data(extra_body_mock_response_data
assert "temperature" in request_body
assert "custom_field" in request_body
assert request_body["custom_field"] == "custom_value"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("sync_mode", [True, False])
+async def test_openai_compact_responses_api(sync_mode):
+ """
+ Test the compact_responses API for OpenAI.
+
+ This test verifies that the compact_responses endpoint works correctly
+ for compressing conversation history.
+ """
+ litellm._turn_on_debug()
+ litellm.set_verbose = True
+
+ input_messages = [
+ {"role": "user", "content": "Hello, how are you?"},
+ {"role": "assistant", "content": "I'm doing well, thank you for asking!"},
+ {"role": "user", "content": "What is the weather like today?"},
+ ]
+
+ try:
+ if sync_mode:
+ response = litellm.compact_responses(
+ model="openai/gpt-4o",
+ input=input_messages,
+ instructions="Be helpful and concise",
+ )
+ else:
+ response = await litellm.acompact_responses(
+ model="openai/gpt-4o",
+ input=input_messages,
+ instructions="Be helpful and concise",
+ )
+ except litellm.InternalServerError:
+ pytest.skip("Skipping test due to InternalServerError")
+ except litellm.BadRequestError as e:
+ # compact_responses may not be available for all models/accounts
+ pytest.skip(f"Skipping test due to BadRequestError: {e}")
+
+ print("compact_responses response=", json.dumps(response, indent=4, default=str))
+
+ # Validate response structure
+ assert response is not None
+ assert "id" in response, "Response should have an 'id' field"
+ assert "output" in response, "Response should have an 'output' field"
+ assert isinstance(response["output"], list), "Output should be a list"
diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py
new file mode 100644
index 00000000000..8c0f7dab2af
--- /dev/null
+++ b/tests/llm_responses_api_testing/test_responses_hooks.py
@@ -0,0 +1,165 @@
+import asyncio
+from datetime import datetime
+from types import SimpleNamespace
+
+import httpx
+import pytest
+
+import litellm
+from litellm.integrations.custom_logger import CustomLogger
+from litellm.responses import streaming_iterator as streaming_module
+from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
+from litellm.types.llms.openai import ResponsesAPIStreamEvents
+from litellm.types.utils import CallTypes
+
+
+class _FakeLoggingObj:
+ def __init__(self):
+ self.success_calls = 0
+ self.async_success_calls = 0
+ self.failure_calls = 0
+ self.async_failure_calls = 0
+ self.start_time = datetime.now()
+ self.model_call_details = {"litellm_params": {}}
+
+ # Signature alignment with Logging handlers
+ def success_handler(self, *args, **kwargs):
+ self.success_calls += 1
+
+ async def async_success_handler(self, *args, **kwargs):
+ self.async_success_calls += 1
+
+ def failure_handler(self, *args, **kwargs):
+ self.failure_calls += 1
+
+ async def async_failure_handler(self, *args, **kwargs):
+ self.async_failure_calls += 1
+
+
+@pytest.mark.asyncio
+async def test_responses_streaming_triggers_hooks(monkeypatch):
+ """
+ Ensure streaming iterator fires success + post-call hooks for responses API.
+ """
+ hook_calls = {"post_call": 0, "metadata": 0}
+ seen = {}
+
+ async def fake_post_call(request_data, response, call_type):
+ hook_calls["post_call"] += 1
+ seen["request_data"] = request_data
+ seen["call_type"] = call_type
+
+ def fake_update_metadata(**kwargs):
+ hook_calls["metadata"] += 1
+
+ monkeypatch.setattr(
+ streaming_module,
+ "async_post_call_success_deployment_hook",
+ fake_post_call,
+ )
+ monkeypatch.setattr(
+ streaming_module,
+ "update_response_metadata",
+ fake_update_metadata,
+ )
+
+ logging_obj = _FakeLoggingObj()
+
+ iterator = ResponsesAPIStreamingIterator(
+ response=httpx.Response(200),
+ model="test-model",
+ responses_api_provider_config=SimpleNamespace(), # not used in this test
+ logging_obj=logging_obj,
+ request_data={"foo": "bar", "litellm_params": {}},
+ call_type=CallTypes.responses.value,
+ )
+
+ # Simulate completed streaming event
+ iterator.completed_response = SimpleNamespace(
+ type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, response=SimpleNamespace()
+ )
+
+ iterator._handle_logging_completed_response()
+ await asyncio.sleep(0.2) # allow async tasks to run
+
+ assert logging_obj.success_calls == 1
+ assert logging_obj.async_success_calls == 1
+ assert hook_calls["post_call"] == 1
+ assert hook_calls["metadata"] == 1
+ assert seen["request_data"]["foo"] == "bar"
+ assert seen["request_data"].get("litellm_params") is not None
+ assert seen["call_type"] == CallTypes.responses
+
+
+@pytest.mark.asyncio
+async def test_responses_streaming_calls_post_streaming_deployment_hook(monkeypatch):
+ """
+ Ensure per-chunk streaming deployment hook can modify chunks.
+ """
+
+ class _HookLogger(CustomLogger):
+ async def async_post_call_streaming_deployment_hook(
+ self, request_data, response_chunk, call_type
+ ):
+ response_chunk.tagged = True
+ return response_chunk
+
+ # Set callbacks to our fake hook
+ original_callbacks = litellm.callbacks
+ litellm.callbacks = [_HookLogger()]
+
+ logging_obj = _FakeLoggingObj()
+
+ class _StubConfig:
+ def transform_streaming_response(self, **kwargs):
+ return SimpleNamespace(
+ type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None
+ )
+
+ iterator = ResponsesAPIStreamingIterator(
+ response=httpx.Response(200),
+ model="test-model",
+ responses_api_provider_config=_StubConfig(),
+ logging_obj=logging_obj,
+ request_data={"foo": "bar"},
+ call_type=CallTypes.responses.value,
+ )
+
+ # Call hook helper directly to verify chunk is modified/flagged
+ chunk = SimpleNamespace(type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None)
+ chunk = await streaming_module.call_post_streaming_hooks_for_testing(iterator, chunk)
+ assert getattr(chunk, "_post_streaming_hooks_ran", False) is True
+ assert getattr(chunk, "tagged", False) is True
+
+ # reset callbacks
+ litellm.callbacks = original_callbacks
+
+
+@pytest.mark.asyncio
+async def test_responses_streaming_failure_triggers_failure_handlers():
+ """
+ If transform raises, failure handlers should be called.
+ """
+
+ class _FailConfig:
+ def transform_streaming_response(self, **kwargs):
+ raise ValueError("boom")
+
+ logging_obj = _FakeLoggingObj()
+
+ iterator = ResponsesAPIStreamingIterator(
+ response=httpx.Response(200),
+ model="test-model",
+ responses_api_provider_config=_FailConfig(),
+ logging_obj=logging_obj,
+ request_data={"foo": "bar"},
+ call_type=CallTypes.responses.value,
+ )
+
+ with pytest.raises(ValueError):
+ iterator._process_chunk('{"delta": "chunk"}')
+
+ # allow failure callbacks to run
+ await asyncio.sleep(0.2)
+ assert logging_obj.failure_calls >= 1
+ assert logging_obj.async_failure_calls >= 1
diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py
index 7c849650bf6..ab5709cd72d 100644
--- a/tests/llm_translation/test_anthropic_completion.py
+++ b/tests/llm_translation/test_anthropic_completion.py
@@ -385,7 +385,7 @@ def test_anthropic_tool_use(tool_type, tool_config, message_content):
"computer_tool_used, prompt_caching_set, expected_beta_header",
[
(True, False, True),
- (False, True, True),
+ (False, True, False),
(True, True, True),
(False, False, False),
],
diff --git a/tests/llm_translation/test_databricks.py b/tests/llm_translation/test_databricks.py
index 40fc712f2b7..3013d00288f 100644
--- a/tests/llm_translation/test_databricks.py
+++ b/tests/llm_translation/test_databricks.py
@@ -15,6 +15,7 @@ import litellm
from litellm.exceptions import BadRequestError
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.utils import CustomStreamWrapper
+from litellm._version import version
from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest
try:
@@ -725,6 +726,7 @@ def test_embeddings_with_sync_http_handler(monkeypatch):
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
+ "User-Agent": f"litellm/{version}",
},
data=json.dumps(
{
@@ -767,6 +769,7 @@ def test_embeddings_with_async_http_handler(monkeypatch):
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
+ "User-Agent": f"litellm/{version}",
},
data=json.dumps(
{
@@ -823,6 +826,7 @@ def test_embeddings_uses_databricks_sdk_if_api_key_and_base_not_specified(monkey
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
+ "User-Agent": f"litellm/{version}",
},
data=json.dumps(
{
@@ -895,6 +899,7 @@ async def test_databricks_embeddings(sync_mode, monkeypatch):
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
+ "User-Agent": f"litellm/{version}",
},
data=json.dumps(
{
@@ -923,6 +928,7 @@ async def test_databricks_embeddings(sync_mode, monkeypatch):
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
+ "User-Agent": f"litellm/{version}",
},
data=json.dumps(
{
diff --git a/tests/llm_translation/test_gigachat.py b/tests/llm_translation/test_gigachat.py
new file mode 100644
index 00000000000..80bf51b4646
--- /dev/null
+++ b/tests/llm_translation/test_gigachat.py
@@ -0,0 +1,349 @@
+"""
+Tests for GigaChat LiteLLM Provider
+
+Tests message transformation, parameter handling, and response transformation.
+Run with: pytest tests/llm_translation/test_gigachat.py -v
+"""
+
+import json
+import pytest
+from unittest.mock import Mock, MagicMock
+
+
+class TestGigaChatMessageTransformation:
+ """Tests for message transformation (OpenAI -> GigaChat format)"""
+
+ @pytest.fixture
+ def config(self):
+ from litellm.llms.gigachat.chat.transformation import GigaChatConfig
+ return GigaChatConfig()
+
+ def test_simple_user_message(self, config):
+ """Basic user message should pass through"""
+ messages = [{"role": "user", "content": "Hello"}]
+ result = config._transform_messages(messages)
+
+ assert len(result) == 1
+ assert result[0]["role"] == "user"
+ assert result[0]["content"] == "Hello"
+
+ def test_developer_role_to_system(self, config):
+ """Developer role should be converted to system"""
+ messages = [{"role": "developer", "content": "You are helpful"}]
+ result = config._transform_messages(messages)
+
+ assert result[0]["role"] == "system"
+
+ def test_system_after_first_becomes_user(self, config):
+ """System message after first position should become user"""
+ messages = [
+ {"role": "assistant", "content": "Response"},
+ {"role": "system", "content": "Additional instruction"},
+ ]
+ result = config._transform_messages(messages)
+
+ assert result[0]["role"] == "assistant"
+ assert result[1]["role"] == "user" # system after first becomes user
+
+ def test_tool_role_to_function(self, config):
+ """Tool role should be converted to function"""
+ messages = [{"role": "tool", "content": "result data"}]
+ result = config._transform_messages(messages)
+
+ assert result[0]["role"] == "function"
+
+ def test_tool_calls_to_function_call(self, config):
+ """tool_calls should be converted to function_call"""
+ messages = [{
+ "role": "assistant",
+ "content": "",
+ "tool_calls": [{
+ "id": "call_123",
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "arguments": '{"city": "Moscow"}'
+ }
+ }]
+ }]
+ result = config._transform_messages(messages)
+
+ assert "function_call" in result[0]
+ assert result[0]["function_call"]["name"] == "get_weather"
+ assert result[0]["function_call"]["arguments"] == {"city": "Moscow"}
+ assert "tool_calls" not in result[0]
+
+ def test_none_content_becomes_empty_string(self, config):
+ """None content should become empty string"""
+ messages = [{"role": "assistant", "content": None}]
+ result = config._transform_messages(messages)
+
+ assert result[0]["content"] == ""
+
+ def test_name_field_removed(self, config):
+ """name field should be removed (not supported by GigaChat)"""
+ messages = [{"role": "user", "content": "Hi", "name": "John"}]
+ result = config._transform_messages(messages)
+
+ assert "name" not in result[0]
+
+
+class TestGigaChatCollapseUserMessages:
+ """Tests for collapsing consecutive user messages"""
+
+ @pytest.fixture
+ def config(self):
+ from litellm.llms.gigachat.chat.transformation import GigaChatConfig
+ return GigaChatConfig()
+
+ def test_no_collapse_single_message(self, config):
+ """Single message should not be changed"""
+ messages = [{"role": "user", "content": "Hello"}]
+ result = config._collapse_user_messages(messages)
+
+ assert len(result) == 1
+ assert result[0]["content"] == "Hello"
+
+ def test_collapse_consecutive_user_messages(self, config):
+ """Consecutive user messages should be collapsed"""
+ messages = [
+ {"role": "user", "content": "First"},
+ {"role": "user", "content": "Second"},
+ {"role": "user", "content": "Third"},
+ ]
+ result = config._collapse_user_messages(messages)
+
+ assert len(result) == 1
+ assert "First" in result[0]["content"]
+ assert "Second" in result[0]["content"]
+ assert "Third" in result[0]["content"]
+
+ def test_no_collapse_with_assistant_between(self, config):
+ """Messages with assistant between should not be collapsed"""
+ messages = [
+ {"role": "user", "content": "First"},
+ {"role": "assistant", "content": "Response"},
+ {"role": "user", "content": "Second"},
+ ]
+ result = config._collapse_user_messages(messages)
+
+ assert len(result) == 3
+
+
+class TestGigaChatToolsTransformation:
+ """Tests for tools -> functions conversion"""
+
+ @pytest.fixture
+ def config(self):
+ from litellm.llms.gigachat.chat.transformation import GigaChatConfig
+ return GigaChatConfig()
+
+ def test_single_tool_conversion(self, config):
+ """Single tool should be converted correctly"""
+ tools = [{
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "description": "Get weather for a city",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "city": {"type": "string"}
+ }
+ }
+ }
+ }]
+ result = config._convert_tools_to_functions(tools)
+
+ assert len(result) == 1
+ assert result[0]["name"] == "get_weather"
+ assert result[0]["description"] == "Get weather for a city"
+
+ def test_multiple_tools_conversion(self, config):
+ """Multiple tools should all be converted"""
+ tools = [
+ {"type": "function", "function": {"name": "func1", "description": "First", "parameters": {"type": "object", "properties": {}}}},
+ {"type": "function", "function": {"name": "func2", "description": "Second", "parameters": {"type": "object", "properties": {}}}},
+ ]
+ result = config._convert_tools_to_functions(tools)
+
+ assert len(result) == 2
+ assert result[0]["name"] == "func1"
+ assert result[1]["name"] == "func2"
+
+
+class TestGigaChatParamsTransformation:
+ """Tests for parameter transformation"""
+
+ @pytest.fixture
+ def config(self):
+ from litellm.llms.gigachat.chat.transformation import GigaChatConfig
+ return GigaChatConfig()
+
+ def test_temperature_zero_becomes_top_p_zero(self, config):
+ """temperature=0 should become top_p=0"""
+ params = {"temperature": 0}
+ result = config.map_openai_params(
+ non_default_params=params,
+ optional_params={},
+ model="GigaChat",
+ drop_params=False,
+ )
+
+ assert "top_p" in result
+ assert result["top_p"] == 0
+ assert "temperature" not in result
+
+ def test_temperature_nonzero_preserved(self, config):
+ """Non-zero temperature should be preserved"""
+ params = {"temperature": 0.7}
+ result = config.map_openai_params(
+ non_default_params=params,
+ optional_params={},
+ model="GigaChat",
+ drop_params=False,
+ )
+
+ assert result["temperature"] == 0.7
+
+ def test_max_completion_tokens_to_max_tokens(self, config):
+ """max_completion_tokens should become max_tokens"""
+ params = {"max_completion_tokens": 100}
+ result = config.map_openai_params(
+ non_default_params=params,
+ optional_params={},
+ model="GigaChat",
+ drop_params=False,
+ )
+
+ assert result["max_tokens"] == 100
+
+ def test_structured_output_via_json_schema(self, config):
+ """json_schema response_format should trigger structured output mode"""
+ params = {
+ "response_format": {
+ "type": "json_schema",
+ "json_schema": {
+ "name": "person",
+ "schema": {
+ "type": "object",
+ "properties": {
+ "name": {"type": "string"},
+ "age": {"type": "integer"}
+ }
+ }
+ }
+ }
+ }
+ result = config.map_openai_params(
+ non_default_params=params,
+ optional_params={},
+ model="GigaChat",
+ drop_params=False,
+ )
+
+ assert "_structured_output" in result
+ assert result["_structured_output"] is True
+ assert "function_call" in result
+ assert result["function_call"]["name"] == "person"
+
+
+class TestGigaChatProviderRegistration:
+ """Tests for provider registration in LiteLLM"""
+
+ def test_gigachat_in_provider_list(self):
+ """GigaChat should be in provider list"""
+ from litellm.types.utils import LlmProviders
+
+ assert hasattr(LlmProviders, "GIGACHAT")
+ assert LlmProviders.GIGACHAT.value == "gigachat"
+
+ def test_gigachat_in_chat_providers(self):
+ """GigaChat should be in LITELLM_CHAT_PROVIDERS"""
+ from litellm.constants import LITELLM_CHAT_PROVIDERS
+
+ assert "gigachat" in LITELLM_CHAT_PROVIDERS
+
+ def test_gigachat_key_exists(self):
+ """gigachat_key should be available"""
+ import litellm
+
+ assert hasattr(litellm, "gigachat_key")
+
+ def test_gigachat_config_exists(self):
+ """GigaChatConfig should be available"""
+ import litellm
+
+ assert hasattr(litellm, "GigaChatConfig")
+
+
+class TestGigaChatTransformRequest:
+ """Tests for request transformation"""
+
+ @pytest.fixture
+ def config(self):
+ from litellm.llms.gigachat.chat.transformation import GigaChatConfig
+ return GigaChatConfig()
+
+ def test_basic_request(self, config):
+ """Basic request should be transformed correctly"""
+ messages = [{"role": "user", "content": "Hello"}]
+ result = config.transform_request(
+ model="gigachat/GigaChat",
+ messages=messages,
+ optional_params={},
+ litellm_params={},
+ headers={},
+ )
+
+ assert result["model"] == "GigaChat"
+ assert len(result["messages"]) == 1
+ assert result["messages"][0]["role"] == "user"
+
+ def test_request_with_temperature(self, config):
+ """Request with temperature should include it"""
+ messages = [{"role": "user", "content": "Hello"}]
+ result = config.transform_request(
+ model="gigachat/GigaChat",
+ messages=messages,
+ optional_params={"temperature": 0.7},
+ litellm_params={},
+ headers={},
+ )
+
+ assert result["temperature"] == 0.7
+
+ def test_request_with_functions(self, config):
+ """Request with functions should include them"""
+ messages = [{"role": "user", "content": "Hello"}]
+ functions = [{"name": "test", "description": "Test", "parameters": {}}]
+ result = config.transform_request(
+ model="gigachat/GigaChat",
+ messages=messages,
+ optional_params={"functions": functions},
+ litellm_params={},
+ headers={},
+ )
+
+ assert "functions" in result
+ assert len(result["functions"]) == 1
+
+
+class TestGigaChatSupportedParams:
+ """Tests for supported parameters"""
+
+ @pytest.fixture
+ def config(self):
+ from litellm.llms.gigachat.chat.transformation import GigaChatConfig
+ return GigaChatConfig()
+
+ def test_supported_params(self, config):
+ """Check supported parameters list"""
+ supported = config.get_supported_openai_params("GigaChat")
+
+ assert "temperature" in supported
+ assert "max_tokens" in supported
+ assert "max_completion_tokens" in supported
+ assert "tools" in supported
+ assert "response_format" in supported
+ assert "stream" in supported
diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py
index 2b01b4c2a12..5e92c10fbdc 100644
--- a/tests/local_testing/test_completion.py
+++ b/tests/local_testing/test_completion.py
@@ -286,7 +286,7 @@ def test_completion_claude_3_empty_response():
},
]
try:
- response = litellm.completion(model="claude-3-opus-20240229", messages=messages)
+ response = litellm.completion(model="claude-3-7-sonnet-20250219", messages=messages)
print(response)
except litellm.InternalServerError as e:
pytest.skip(f"InternalServerError - {str(e)}")
@@ -313,7 +313,7 @@ def test_completion_claude_3():
try:
# test without max tokens
response = completion(
- model="anthropic/claude-3-opus-20240229",
+ model="anthropic/claude-3-7-sonnet-20250219",
messages=messages,
)
# Add any assertions, here to check response args
@@ -326,7 +326,7 @@ def test_completion_claude_3():
@pytest.mark.parametrize(
"model",
- ["anthropic/claude-3-opus-20240229", "anthropic.claude-3-sonnet-20240229-v1:0"],
+ ["anthropic/claude-3-7-sonnet-20250219", "anthropic.claude-3-sonnet-20240229-v1:0"],
)
def test_completion_claude_3_function_call(model):
litellm.set_verbose = True
@@ -411,7 +411,7 @@ def test_completion_claude_3_function_call(model):
"model, api_key, api_base",
[
("gpt-3.5-turbo", None, None),
- ("claude-3-opus-20240229", None, None),
+ ("claude-3-7-sonnet-20250219", None, None),
("anthropic.claude-3-sonnet-20240229-v1:0", None, None),
# (
# "azure_ai/command-r-plus",
@@ -512,7 +512,7 @@ async def test_anthropic_no_content_error():
try:
litellm.drop_params = True
response = await litellm.acompletion(
- model="anthropic/claude-3-opus-20240229",
+ model="anthropic/claude-3-7-sonnet-20250219",
api_key=os.getenv("ANTHROPIC_API_KEY"),
messages=[
{
@@ -630,7 +630,7 @@ def test_completion_claude_3_multi_turn_conversations():
]
try:
response = completion(
- model="anthropic/claude-3-opus-20240229",
+ model="anthropic/claude-3-7-sonnet-20250219",
messages=messages,
)
print(response)
@@ -644,7 +644,7 @@ def test_completion_claude_3_stream():
try:
# test without max tokens
response = completion(
- model="anthropic/claude-3-opus-20240229",
+ model="anthropic/claude-3-7-sonnet-20250219",
messages=messages,
max_tokens=10,
stream=True,
@@ -669,7 +669,7 @@ def encode_image(image_path):
[
"gpt-4o",
"azure/gpt-4.1-mini",
- "anthropic/claude-3-opus-20240229",
+ "anthropic/claude-3-7-sonnet-20250219",
],
) #
def test_completion_base64(model):
diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py
index 0d9f84a301c..b9b5d0fdb07 100644
--- a/tests/local_testing/test_streaming.py
+++ b/tests/local_testing/test_streaming.py
@@ -1418,7 +1418,7 @@ def test_bedrock_claude_3_streaming():
@pytest.mark.parametrize(
"model",
[
- "claude-3-opus-20240229",
+ "claude-3-7-sonnet-20250219",
"cohere.command-r-plus-v1:0", # bedrock
"gpt-3.5-turbo",
],
@@ -2914,7 +2914,7 @@ def test_completion_claude_3_function_call_with_streaming():
try:
# test without max tokens
response = completion(
- model="claude-3-opus-20240229",
+ model="claude-3-7-sonnet-20250219",
messages=messages,
tools=tools,
tool_choice="required",
@@ -2946,7 +2946,7 @@ def test_completion_claude_3_function_call_with_streaming():
"model",
[
"gemini/gemini-2.5-flash-lite",
- ], # "claude-3-opus-20240229"
+ ],
) #
@pytest.mark.asyncio
async def test_acompletion_function_call_with_streaming(model):
diff --git a/tests/logging_callback_tests/test_langsmith_unit_test.py b/tests/logging_callback_tests/test_langsmith_unit_test.py
index e63ce9f8b38..bde2b944579 100644
--- a/tests/logging_callback_tests/test_langsmith_unit_test.py
+++ b/tests/logging_callback_tests/test_langsmith_unit_test.py
@@ -47,6 +47,19 @@ async def test_get_credentials_from_env():
credentials = logger.get_credentials_from_env()
assert credentials["LANGSMITH_BASE_URL"] == "https://api.smith.langchain.com"
+ # Test with tenant_id
+ credentials = logger.get_credentials_from_env(
+ langsmith_tenant_id="test-tenant-id"
+ )
+ assert credentials["LANGSMITH_TENANT_ID"] == "test-tenant-id"
+
+ # Test tenant_id from environment variable
+ import os
+ os.environ["LANGSMITH_TENANT_ID"] = "env-tenant-id"
+ credentials = logger.get_credentials_from_env()
+ assert credentials["LANGSMITH_TENANT_ID"] == "env-tenant-id"
+ del os.environ["LANGSMITH_TENANT_ID"]
+
@pytest.mark.asyncio
async def test_group_batches_by_credentials():
@@ -60,6 +73,7 @@ async def test_group_batches_by_credentials():
"LANGSMITH_API_KEY": "key1",
"LANGSMITH_PROJECT": "proj1",
"LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": None,
},
)
@@ -69,6 +83,7 @@ async def test_group_batches_by_credentials():
"LANGSMITH_API_KEY": "key1",
"LANGSMITH_PROJECT": "proj1",
"LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": None,
},
)
@@ -95,6 +110,7 @@ async def test_group_batches_by_credentials_multiple_credentials():
"LANGSMITH_API_KEY": "key1",
"LANGSMITH_PROJECT": "proj1",
"LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": None,
},
)
@@ -104,6 +120,7 @@ async def test_group_batches_by_credentials_multiple_credentials():
"LANGSMITH_API_KEY": "key2", # Different API key
"LANGSMITH_PROJECT": "proj1",
"LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": None,
},
)
@@ -113,6 +130,7 @@ async def test_group_batches_by_credentials_multiple_credentials():
"LANGSMITH_API_KEY": "key1",
"LANGSMITH_PROJECT": "proj2", # Different project
"LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": None,
},
)
@@ -127,6 +145,57 @@ async def test_group_batches_by_credentials_multiple_credentials():
assert len(batch_group.queue_objects) == 1 # Each group should have one object
+@pytest.mark.asyncio
+async def test_group_batches_by_credentials_with_tenant_id():
+
+ # Test that different tenant_ids create separate groups
+ logger = LangsmithLogger(langsmith_api_key="test-key")
+
+ queue_obj1 = LangsmithQueueObject(
+ data={"test": "data1"},
+ credentials={
+ "LANGSMITH_API_KEY": "key1",
+ "LANGSMITH_PROJECT": "proj1",
+ "LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": "tenant1",
+ },
+ )
+
+ queue_obj2 = LangsmithQueueObject(
+ data={"test": "data2"},
+ credentials={
+ "LANGSMITH_API_KEY": "key1",
+ "LANGSMITH_PROJECT": "proj1",
+ "LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": "tenant2", # Different tenant_id
+ },
+ )
+
+ queue_obj3 = LangsmithQueueObject(
+ data={"test": "data3"},
+ credentials={
+ "LANGSMITH_API_KEY": "key1",
+ "LANGSMITH_PROJECT": "proj1",
+ "LANGSMITH_BASE_URL": "url1",
+ "LANGSMITH_TENANT_ID": "tenant1", # Same as queue_obj1
+ },
+ )
+
+ logger.log_queue = [queue_obj1, queue_obj2, queue_obj3]
+
+ grouped = logger._group_batches_by_credentials()
+
+ # Should have two groups: one for tenant1 (queue_obj1 and queue_obj3), one for tenant2 (queue_obj2)
+ assert len(grouped) == 2
+ for key, batch_group in grouped.items():
+ assert isinstance(key, CredentialsKey)
+ assert key.tenant_id in ["tenant1", "tenant2"]
+ if key.tenant_id == "tenant1":
+ assert len(batch_group.queue_objects) == 2
+ else:
+ assert len(batch_group.queue_objects) == 1
+
+
# Test make_dot_order
@pytest.mark.asyncio
async def test_make_dot_order():
@@ -201,10 +270,43 @@ async def test_async_send_batch():
call_args = logger.async_httpx_client.post.call_args
assert "runs/batch" in call_args[1]["url"]
assert "x-api-key" in call_args[1]["headers"]
+ # tenant_id should not be in headers if not provided
+ assert "x-tenant-id" not in call_args[1]["headers"]
@pytest.mark.asyncio
-async def test_langsmith_key_based_logging(mocker):
+async def test_async_send_batch_with_tenant_id():
+ logger = LangsmithLogger(
+ langsmith_api_key="test-key",
+ langsmith_tenant_id="test-tenant-id"
+ )
+
+ # Mock the httpx client
+ mock_response = AsyncMock()
+ mock_response.status_code = 200
+ logger.async_httpx_client = AsyncMock()
+ logger.async_httpx_client.post.return_value = mock_response
+
+ # Add test data to queue
+ logger.log_queue = [
+ LangsmithQueueObject(
+ data={"test": "data"}, credentials=logger.default_credentials
+ )
+ ]
+
+ await logger.async_send_batch()
+
+ # Verify the API call includes tenant_id header
+ logger.async_httpx_client.post.assert_called_once()
+ call_args = logger.async_httpx_client.post.call_args
+ assert "runs/batch" in call_args[1]["url"]
+ assert "x-api-key" in call_args[1]["headers"]
+ assert "x-tenant-id" in call_args[1]["headers"]
+ assert call_args[1]["headers"]["x-tenant-id"] == "test-tenant-id"
+
+
+@pytest.mark.asyncio
+async def test_langsmith_key_based_logging():
"""
In key based logging langsmith_api_key and langsmith_project are passed directly to litellm.acompletion
"""
@@ -219,10 +321,11 @@ async def test_langsmith_key_based_logging(mocker):
mock_response.text = ""
mock_async_httpx_handler.post = AsyncMock(return_value=mock_response)
- mock_get_client = mocker.patch(
+ mock_get_client = patch(
"litellm.integrations.langsmith.get_async_httpx_client",
return_value=mock_async_httpx_handler
)
+ mock_get_client.start()
litellm.set_verbose = True
litellm.DEFAULT_FLUSH_INTERVAL_SECONDS = 1
@@ -253,6 +356,8 @@ async def test_langsmith_key_based_logging(mocker):
# Check headers contain the correct API key
assert call_args[1]["headers"]["x-api-key"] == "fake_key_project2"
+ # tenant_id should not be in headers if not provided
+ assert "x-tenant-id" not in call_args[1]["headers"]
# Verify the request body contains the expected data
request_body = call_args[1]["json"]
@@ -344,6 +449,8 @@ async def test_langsmith_key_based_logging(mocker):
actual_body["post"][0]["session_name"]
== expected_body["post"][0]["session_name"]
)
+
+ mock_get_client.stop()
except Exception as e:
pytest.fail(f"Error occurred: {e}")
diff --git a/tests/logging_callback_tests/test_opentelemetry_unit_tests.py b/tests/logging_callback_tests/test_opentelemetry_unit_tests.py
index 3d0682d9033..04f8abe64de 100644
--- a/tests/logging_callback_tests/test_opentelemetry_unit_tests.py
+++ b/tests/logging_callback_tests/test_opentelemetry_unit_tests.py
@@ -65,31 +65,6 @@ class TestOpentelemetryUnitTests(BaseLoggingCallbackTest):
# External spans should only be closed by their creators
parent_otel_span.end.assert_not_called()
- def test_init_tracing_respects_existing_tracer_provider(self):
- """
- Unit test: _init_tracing() should respect existing TracerProvider.
-
- When a TracerProvider already exists (e.g., set by Langfuse SDK),
- LiteLLM should use it instead of creating a new one.
- """
- from opentelemetry import trace
- from opentelemetry.sdk.trace import TracerProvider
- from litellm.integrations.opentelemetry import OpenTelemetry
-
- # Setup: Create and set an existing TracerProvider
- tracer_provider = TracerProvider()
- trace.set_tracer_provider(tracer_provider)
- existing_provider = trace.get_tracer_provider()
-
- # Act: Initialize OpenTelemetry integration (should detect existing provider)
- otel_integration = OpenTelemetry()
-
- # Assert: The existing provider should still be active
- current_provider = trace.get_tracer_provider()
- assert current_provider is existing_provider, (
- "Existing TracerProvider should be respected and not overridden"
- )
-
def test_get_span_context_detects_active_span(self):
"""
Unit test: _get_span_context() should auto-detect active spans from global context.
diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py
index 40e223ffe07..45aae3b9aee 100644
--- a/tests/router_unit_tests/test_router_helper_utils.py
+++ b/tests/router_unit_tests/test_router_helper_utils.py
@@ -79,6 +79,73 @@ def test_routing_strategy_init(model_list):
)
+def test_routing_strategy_init_invalid_strategy(model_list):
+ """Test that invalid routing_strategy raises ValueError with helpful message.
+
+ See: https://github.com/BerriAI/litellm/issues/11330
+ Invalid strategies like 'simple' (without '-shuffle') should fail fast
+ with a clear error, not silently cause 'No deployments available' errors.
+ """
+ router = Router(model_list=model_list)
+
+ # Test common mistake: "simple" instead of "simple-shuffle"
+ with pytest.raises(ValueError) as exc_info:
+ router.routing_strategy_init(
+ routing_strategy="simple",
+ routing_strategy_args={}
+ )
+
+ # Verify error message is helpful
+ error_msg = str(exc_info.value)
+ assert "Invalid routing_strategy" in error_msg
+ assert "simple" in error_msg
+ assert "simple-shuffle" in error_msg # Suggests the correct option
+ # Verify error message tells user WHERE to fix it
+ assert "config.yaml" in error_msg
+ assert "router_settings.routing_strategy" in error_msg
+ assert "Router SDK" in error_msg
+
+ # Test completely invalid strategy
+ with pytest.raises(ValueError) as exc_info:
+ router.routing_strategy_init(
+ routing_strategy="not-a-real-strategy",
+ routing_strategy_args={}
+ )
+ assert "Invalid routing_strategy" in str(exc_info.value)
+
+
+def test_routing_strategy_init_valid_string_strategies(model_list):
+ """Test that all valid string routing strategies work without error.
+
+ Valid strategies are derived from RoutingStrategy enum values plus 'simple-shuffle'.
+ """
+ from litellm.types.router import RoutingStrategy
+
+ router = Router(model_list=model_list)
+
+ # All strategies from enum + simple-shuffle (default, not in enum)
+ valid_strategies = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
+
+ for strategy in valid_strategies:
+ # Should not raise
+ router.routing_strategy_init(
+ routing_strategy=strategy, routing_strategy_args={}
+ )
+
+
+def test_routing_strategy_init_valid_enum_strategies(model_list):
+ """Test that RoutingStrategy enum values work without error."""
+ from litellm.types.router import RoutingStrategy
+
+ router = Router(model_list=model_list)
+
+ for strategy in RoutingStrategy:
+ # Should not raise when passing enum directly
+ router.routing_strategy_init(
+ routing_strategy=strategy, routing_strategy_args={}
+ )
+
+
def test_print_deployment(model_list):
"""Test if the api key is masked correctly"""
diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
index 6490352c39b..596398e639f 100644
--- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
+++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py
@@ -1011,39 +1011,87 @@ def test_multiple_tool_calls_in_single_choice():
def test_map_reasoning_effort_adds_summary_detailed():
"""
- Test that _map_reasoning_effort adds summary="detailed" when user provides reasoning_effort as a string.
+ Test that _map_reasoning_effort behavior with reasoning_auto_summary flag.
- This ensures that when users pass reasoning_effort in the completions API for OpenAI responses/models,
- the transformation automatically includes summary="detailed" in the reasoning parameter.
+ By default (flag=False), summary should NOT be added to avoid:
+ 1. Breaking for users without verified OpenAI orgs (400 errors)
+ 2. Making requests more expensive by including summary reasoning tokens
+
+ When flag is enabled (flag=True or env var), summary="detailed" is added.
"""
+ import os
+
+ import litellm
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
handler = LiteLLMResponsesTransformationHandler()
- # Test all string effort levels
+ # Test all string effort levels - DEFAULT BEHAVIOR (no summary)
effort_levels = ["none", "low", "medium", "high", "xhigh", "minimal"]
- for effort in effort_levels:
- result = handler._map_reasoning_effort(effort)
+ # Save original flag value
+ original_flag = litellm.reasoning_auto_summary
+ original_env = os.environ.get("LITELLM_REASONING_AUTO_SUMMARY")
+
+ try:
+ # Test 1: Default behavior (flag=False, no env var) - NO summary
+ litellm.reasoning_auto_summary = False
+ if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
+ del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
- assert result is not None, f"Result should not be None for effort={effort}"
- assert result["effort"] == effort, f"Effort should be {effort}"
- assert result["summary"] == "detailed", f"Summary should be 'detailed' for effort={effort}"
+ for effort in effort_levels:
+ result = handler._map_reasoning_effort(effort)
+
+ assert result is not None, f"Result should not be None for effort={effort}"
+ assert result["effort"] == effort, f"Effort should be {effort}"
+ assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}"
+
+ print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)")
- print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed'")
+ # Test 2: With flag enabled - summary IS added
+ litellm.reasoning_auto_summary = True
+
+ for effort in effort_levels:
+ result = handler._map_reasoning_effort(effort)
+
+ assert result is not None, f"Result should not be None for effort={effort}"
+ assert result["effort"] == effort, f"Effort should be {effort}"
+ assert result["summary"] == "detailed", f"Summary should be 'detailed' when flag is enabled for effort={effort}"
+
+ print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)")
+
+ # Test 3: With env var enabled (flag disabled) - summary IS added
+ litellm.reasoning_auto_summary = False
+ os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true"
+
+ result = handler._map_reasoning_effort("high")
+ assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled"
+ print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly")
+
+ # Test 4: Dict input is passed through as-is (no modification)
+ litellm.reasoning_auto_summary = False
+ if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
+ del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
+
+ dict_input = {"effort": "high", "summary": "custom_summary"}
+ result_dict = handler._map_reasoning_effort(dict_input)
+ assert result_dict["effort"] == "high"
+ assert result_dict["summary"] == "custom_summary"
+ print("✓ Dict input is passed through without modification")
+
+ # Test 5: None/unknown values return None
+ result_unknown = handler._map_reasoning_effort("unknown_value")
+ assert result_unknown is None
+ print("✓ Unknown reasoning_effort values return None")
+
+ print("✓ All reasoning_effort behaviors work correctly with flag/env var control")
- # Test that dict input is passed through as-is (no modification)
- dict_input = {"effort": "high", "summary": "custom_summary"}
- result_dict = handler._map_reasoning_effort(dict_input)
- assert result_dict["effort"] == "high"
- assert result_dict["summary"] == "custom_summary"
- print("✓ Dict input is passed through without modification")
-
- # Test that None/unknown values return None
- result_unknown = handler._map_reasoning_effort("unknown_value")
- assert result_unknown is None
- print("✓ Unknown reasoning_effort values return None")
-
- print("✓ All reasoning_effort string values correctly map to summary='detailed'")
+ finally:
+ # Restore original values
+ litellm.reasoning_auto_summary = original_flag
+ if original_env is not None:
+ os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = original_env
+ elif "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
+ del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
diff --git a/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py b/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py
new file mode 100644
index 00000000000..e717840ec95
--- /dev/null
+++ b/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py
@@ -0,0 +1,90 @@
+"""
+Test for Langfuse integration with Gemini cached_tokens bug
+https://github.com/BerriAI/litellm/issues/18520
+"""
+import pytest
+from litellm.types.utils import PromptTokensDetailsWrapper, Usage
+
+
+def test_cached_tokens_extraction():
+ """
+ Test that we can extract cached_tokens from prompt_tokens_details.
+ This is the core logic fix for https://github.com/BerriAI/litellm/issues/18520
+ """
+ # Create usage object like Gemini returns
+ usage = Usage(
+ prompt_tokens=20209,
+ prompt_tokens_details=PromptTokensDetailsWrapper(
+ cached_tokens=20203,
+ text_tokens=6,
+ ),
+ completion_tokens=541,
+ )
+
+ # Simulate the logic from langfuse.py lines 745-757 (after the fix)
+ cache_read_input_tokens = 0 # Default value
+
+ # Check prompt_tokens_details.cached_tokens (the fix)
+ if hasattr(usage, "prompt_tokens_details"):
+ prompt_tokens_details = getattr(usage, "prompt_tokens_details", None)
+ if (
+ prompt_tokens_details is not None
+ and hasattr(prompt_tokens_details, "cached_tokens")
+ ):
+ cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
+ if cached_tokens is not None and cached_tokens > 0:
+ cache_read_input_tokens = cached_tokens
+
+ # Verify the fix works
+ assert cache_read_input_tokens == 20203, f"Expected 20203, got {cache_read_input_tokens}"
+
+
+def test_cached_tokens_not_present():
+ """Test backward compatibility when cached_tokens is not present"""
+ # Usage without prompt_tokens_details
+ usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ )
+
+ cache_read_input_tokens = 0
+
+ if hasattr(usage, "prompt_tokens_details"):
+ prompt_tokens_details = getattr(usage, "prompt_tokens_details", None)
+ if (
+ prompt_tokens_details is not None
+ and hasattr(prompt_tokens_details, "cached_tokens")
+ ):
+ cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
+ if cached_tokens is not None and cached_tokens > 0:
+ cache_read_input_tokens = cached_tokens
+
+ # Should remain 0
+ assert cache_read_input_tokens == 0
+
+
+def test_cached_tokens_is_zero():
+ """Test when cached_tokens is explicitly 0"""
+ usage = Usage(
+ prompt_tokens=100,
+ prompt_tokens_details=PromptTokensDetailsWrapper(
+ cached_tokens=0,
+ text_tokens=100,
+ ),
+ completion_tokens=50,
+ )
+
+ cache_read_input_tokens = 0
+
+ if hasattr(usage, "prompt_tokens_details"):
+ prompt_tokens_details = getattr(usage, "prompt_tokens_details", None)
+ if (
+ prompt_tokens_details is not None
+ and hasattr(prompt_tokens_details, "cached_tokens")
+ ):
+ cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
+ if cached_tokens is not None and cached_tokens > 0:
+ cache_read_input_tokens = cached_tokens
+
+ # Should remain 0 when cached_tokens is 0
+ assert cache_read_input_tokens == 0
diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py
index 5d648f601f6..6c17570e135 100644
--- a/tests/test_litellm/integrations/test_opentelemetry.py
+++ b/tests/test_litellm/integrations/test_opentelemetry.py
@@ -172,6 +172,86 @@ class TestOpenTelemetryCostBreakdown(unittest.TestCase):
assert ("gen_ai.cost.original_cost", 0.004) not in call_args_list
+class TestOpenTelemetryProviderInitialization(unittest.TestCase):
+ """Test suite for verifying provider initialization respects existing providers"""
+
+ def test_init_tracing_respects_existing_tracer_provider(self):
+ """
+ Unit test: _init_tracing() should respect existing TracerProvider.
+
+ When a TracerProvider already exists (e.g., set by Langfuse SDK),
+ LiteLLM should use it instead of creating a new one.
+ """
+ from opentelemetry import trace
+ from opentelemetry.sdk.trace import TracerProvider
+
+ # Setup: Create and set an existing TracerProvider
+ tracer_provider = TracerProvider()
+ trace.set_tracer_provider(tracer_provider)
+ existing_provider = trace.get_tracer_provider()
+
+ # Act: Initialize OpenTelemetry integration (should detect existing provider)
+ otel_integration = OpenTelemetry()
+
+ # Assert: The existing provider should still be active
+ current_provider = trace.get_tracer_provider()
+ assert current_provider is existing_provider, (
+ "Existing TracerProvider should be respected and not overridden"
+ )
+
+ @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True)
+ def test_init_metrics_respects_existing_meter_provider(self):
+ """
+ Unit test: _init_metrics() should respect existing MeterProvider.
+
+ When a MeterProvider already exists (e.g., set by Langfuse SDK),
+ LiteLLM should use it instead of creating a new one.
+ """
+ from opentelemetry import metrics
+ from opentelemetry.sdk.metrics import MeterProvider
+
+ # Create and set an existing MeterProvider
+ meter_provider = MeterProvider()
+ metrics.set_meter_provider(meter_provider)
+ existing_provider = metrics.get_meter_provider()
+
+ # Act: Initialize OpenTelemetry integration (should detect existing provider)
+ config = OpenTelemetryConfig.from_env()
+ otel_integration = OpenTelemetry(config=config)
+
+ # Assert: The existing provider should still be active
+ current_provider = metrics.get_meter_provider()
+ assert current_provider is existing_provider, (
+ "Existing MeterProvider should be respected and not overridden"
+ )
+
+ @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS": "true"}, clear=True)
+ def test_init_logs_respects_existing_logger_provider(self):
+ """
+ Unit test: _init_logs() should respect existing LoggerProvider.
+
+ When a LoggerProvider already exists (e.g., set by Langfuse SDK),
+ LiteLLM should use it instead of creating a new one.
+ """
+ from opentelemetry._logs import get_logger_provider, set_logger_provider
+ from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider
+
+ # Create and set an existing LoggerProvider
+ logger_provider = OTLoggerProvider()
+ set_logger_provider(logger_provider)
+ existing_provider = get_logger_provider()
+
+ # Act: Initialize OpenTelemetry integration (should detect existing provider)
+ config = OpenTelemetryConfig.from_env()
+ otel_integration = OpenTelemetry(config=config)
+
+ # Assert: The existing provider should still be active
+ current_provider = get_logger_provider()
+ assert current_provider is existing_provider, (
+ "Existing LoggerProvider should be respected and not overridden"
+ )
+
+
class TestOpenTelemetry(unittest.TestCase):
POLL_INTERVAL = 0.05
POLL_TIMEOUT = 2.0
@@ -620,7 +700,6 @@ class TestOpenTelemetry(unittest.TestCase):
self.assertEqual(attributes.get("extra.attr"), "extra-value")
-
def test_handle_success_spans_only(self):
# make sure neither events nor metrics is on
os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None)
@@ -687,11 +766,8 @@ class TestOpenTelemetry(unittest.TestCase):
logs = log_exporter.get_finished_logs()
self.assertFalse(logs, "Did not expect any logs")
+ @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True)
def test_handle_success_spans_and_metrics(self):
- # only metrics on
- os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None)
- os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_METRICS"] = "true"
-
# ─── build in‐memory OTEL providers/exporters ─────────────────────────────
span_exporter = InMemorySpanExporter()
tracer_provider = TracerProvider()
@@ -1320,6 +1396,23 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
)
self.assertEqual(normalized, "http://collector:4317/v1/logs")
+ def test_get_metric_reader_uses_http_exporter_for_http_protobuf(self):
+ """Test that http/protobuf protocol uses OTLPMetricExporterHTTP"""
+ from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
+ OTLPMetricExporter,
+ )
+ from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader
+
+ config = OpenTelemetryConfig(
+ exporter="http/protobuf", endpoint="http://collector:4318"
+ )
+ otel = OpenTelemetry(config=config)
+
+ reader = otel._get_metric_reader()
+
+ self.assertIsInstance(reader, PeriodicExportingMetricReader)
+ self.assertIsInstance(reader._exporter, OTLPMetricExporter)
+
class TestOpenTelemetryExternalSpan(unittest.TestCase):
"""
diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
index 6a528fef8f0..ec2f528a35d 100644
--- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
+++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
@@ -691,6 +691,29 @@ async def test_streaming_completion_start_time(logging_obj: Logging):
)
+@pytest.mark.asyncio
+async def test_vertex_streaming_bad_request_not_midstream(logging_obj: Logging):
+ """Ensure Vertex bad request errors surface as 400, not mid-stream fallbacks."""
+ from litellm.llms.vertex_ai.common_utils import VertexAIError
+
+ async def _raise_bad_request(**kwargs):
+ raise VertexAIError(status_code=400, message="invalid maxOutputTokens", headers=None)
+
+ response = CustomStreamWrapper(
+ completion_stream=None,
+ model="gemini-3-pro-preview",
+ logging_obj=logging_obj,
+ custom_llm_provider="vertex_ai_beta",
+ make_call=_raise_bad_request,
+ )
+
+ with pytest.raises(litellm.BadRequestError) as excinfo:
+ await response.__anext__()
+
+ assert getattr(excinfo.value, "status_code", None) == 400
+ assert "invalid maxOutputTokens" in str(excinfo.value)
+
+
def test_streaming_handler_with_created_time_propagation(
initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging
):
diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py b/tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py
new file mode 100644
index 00000000000..0f369fbb8b9
--- /dev/null
+++ b/tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py
@@ -0,0 +1,355 @@
+"""
+Test cases for functionCall args serialization in Vertex AI Gemini.
+
+This test file specifically tests the edge cases where Vertex AI might return
+functionCall args in unexpected formats that could lead to invalid JSON strings
+like: {"x":"x"}{"a":"a"}
+"""
+import json
+from typing import List, Optional
+
+import pytest
+
+from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
+ VertexGeminiConfig,
+)
+from litellm.types.llms.vertex_ai import HttpxPartType
+
+
+class TestFunctionCallArgsSerialization:
+ """Test cases for functionCall args serialization edge cases."""
+
+ def test_normal_dict_args(self):
+ """Test normal case: args is a dict."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": {"location": "Boston", "unit": "celsius"},
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ assert tools[0]["function"]["name"] == "get_weather"
+
+ # Verify arguments is a valid JSON string
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+ # Should be valid JSON
+ parsed = json.loads(arguments)
+ assert parsed == {"location": "Boston", "unit": "celsius"}
+
+ def test_none_args(self):
+ """Test case: args is None."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": None,
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ # Should serialize None to "null" or empty dict
+ assert isinstance(arguments, str)
+ parsed = json.loads(arguments)
+ # json.dumps(None) returns "null"
+ assert parsed is None or parsed == {}
+
+ def test_args_as_string_valid_json(self):
+ """Test case: args is already a valid JSON string."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": '{"location": "Boston"}', # String, not dict
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ # If args is a string, json.dumps will double-encode it
+ # This would result in: "{\"location\": \"Boston\"}"
+ assert isinstance(arguments, str)
+ # This is the problematic case - string gets double-encoded
+ # The result would be a JSON string containing a JSON string
+ parsed = json.loads(arguments)
+ # If it's double-encoded, parsed would be a string, not a dict
+ if isinstance(parsed, str):
+ # Double-encoded case
+ inner_parsed = json.loads(parsed)
+ assert inner_parsed == {"location": "Boston"}
+ else:
+ # Normal case (shouldn't happen if args is string)
+ assert parsed == {"location": "Boston"}
+
+ def test_args_as_string_invalid_json_concatenated(self):
+ """Test case: args is a string with concatenated JSON objects (the bug case).
+
+ When args is a string like '{"x":"x"}{"a":"a"}', json.dumps() will serialize it
+ as a JSON string, resulting in: "{\"x\":\"x\"}{\"a\":\"a\"}"
+ This is a valid JSON string (the outer quotes), but the content inside is invalid JSON.
+ When you try to parse the inner content, it fails.
+ """
+ # This simulates the case where Vertex might return something like:
+ # args = '{"x":"x"}{"a":"a"}' # Two JSON objects concatenated
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": '{"x":"x"}{"a":"a"}', # Invalid concatenated JSON
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+
+ # json.dumps() on a string will escape it, so we get:
+ # arguments = '"{\\"x\\":\\"x\\"}{\\"a\\":\\"a\\"}"'
+ # This is a valid JSON string (the outer quotes), but the inner content is invalid
+ parsed_outer = json.loads(arguments)
+ assert isinstance(parsed_outer, str)
+
+ # The inner string is invalid JSON (two objects concatenated)
+ # This is the bug: the inner content cannot be parsed as valid JSON
+ with pytest.raises(json.JSONDecodeError):
+ json.loads(parsed_outer)
+
+ # The arguments string would be: "{\"x\":\"x\"}{\"a\":\"a\"}"
+ # Which when parsed gives: '{"x":"x"}{"a":"a"}' (invalid JSON)
+
+ def test_args_as_array(self):
+ """Test case: args is an array (unexpected but possible)."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": [{"x": "x"}, {"a": "a"}], # Array of objects
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+ # Should serialize array correctly
+ parsed = json.loads(arguments)
+ assert parsed == [{"x": "x"}, {"a": "a"}]
+
+ def test_args_missing_key(self):
+ """Test case: args key is missing from functionCall.
+
+ This will raise a KeyError because the code directly accesses part["functionCall"]["args"]
+ without checking if the key exists. This is a bug that should be fixed.
+ """
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ # args key missing
+ }
+ }
+ ]
+
+ # This should raise KeyError because args key is missing
+ with pytest.raises(KeyError):
+ VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ def test_multiple_function_calls(self):
+ """Test case: multiple function calls in parts."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": {"location": "Boston"},
+ }
+ },
+ {
+ "functionCall": {
+ "name": "get_time",
+ "args": {"timezone": "EST"},
+ }
+ },
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 2
+ assert tools[0]["function"]["name"] == "get_weather"
+ assert tools[1]["function"]["name"] == "get_time"
+
+ # Both should have valid JSON arguments
+ args1 = json.loads(tools[0]["function"]["arguments"])
+ args2 = json.loads(tools[1]["function"]["arguments"])
+ assert args1 == {"location": "Boston"}
+ assert args2 == {"timezone": "EST"}
+
+ def test_args_with_vertex_protobuf_format(self):
+ """Test case: args in Vertex protobuf format with string_value, etc."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": {
+ "location": {"string_value": "Boston, MA"},
+ "unit": {"string_value": "celsius"},
+ },
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+ # Should serialize the nested structure correctly
+ parsed = json.loads(arguments)
+ assert "location" in parsed
+ assert "unit" in parsed
+
+ def test_args_as_empty_dict(self):
+ """Test case: args is an empty dict."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": {},
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+ parsed = json.loads(arguments)
+ assert parsed == {}
+
+ def test_args_with_special_characters(self):
+ """Test case: args contains special characters that need escaping."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": {
+ "location": 'Boston, MA "downtown"',
+ "note": "Line 1\nLine 2",
+ },
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+ # Should handle special characters correctly
+ parsed = json.loads(arguments)
+ assert parsed["location"] == 'Boston, MA "downtown"'
+ assert parsed["note"] == "Line 1\nLine 2"
+
+ def test_args_as_list_of_strings_that_look_like_json(self):
+ """Test case: args is a list containing strings that look like JSON objects."""
+ # This could potentially cause issues if not handled correctly
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "get_weather",
+ "args": ['{"x":"x"}', '{"a":"a"}'], # List of JSON strings
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+ # Should serialize list correctly
+ parsed = json.loads(arguments)
+ assert isinstance(parsed, list)
+ assert parsed == ['{"x":"x"}', '{"a":"a"}']
+
+ def test_args_as_dict_with_nested_structures(self):
+ """Test case: args contains nested dicts and lists."""
+ parts: List[HttpxPartType] = [
+ {
+ "functionCall": {
+ "name": "complex_function",
+ "args": {
+ "nested": {"key": "value"},
+ "list": [1, 2, 3],
+ "mixed": [{"a": 1}, {"b": 2}],
+ },
+ }
+ }
+ ]
+
+ function, tools, idx = VertexGeminiConfig._transform_parts(
+ parts=parts, cumulative_tool_call_idx=0, is_function_call=False
+ )
+
+ assert tools is not None
+ assert len(tools) == 1
+ arguments = tools[0]["function"]["arguments"]
+ assert isinstance(arguments, str)
+ parsed = json.loads(arguments)
+ assert parsed["nested"] == {"key": "value"}
+ assert parsed["list"] == [1, 2, 3]
+ assert parsed["mixed"] == [{"a": 1}, {"b": 2}]
+
+
+if __name__ == "__main__":
+ pytest.main([__file__, "-v"])
+
diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
index 91a28ee6ec9..d09de3a0f26 100644
--- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
+++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
@@ -10,11 +10,13 @@ from pydantic import BaseModel
import litellm
from litellm import ModelResponse, completion
+from litellm.llms.vertex_ai.common_utils import VertexAIError
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
from litellm.types.llms.vertex_ai import UsageMetadata
from litellm.types.utils import ChoiceLogprobs, Usage
+from litellm.utils import CustomStreamWrapper
def test_top_logprobs():
@@ -1605,6 +1607,39 @@ def test_vertex_ai_annotation_streaming_events():
assert "Weather information" in annotation["url_citation"]["title"]
+@pytest.mark.asyncio
+async def test_vertex_ai_streaming_bad_request_is_not_wrapped():
+ class DummyLogging:
+ def __init__(self):
+ self.model_call_details = {"litellm_params": {}}
+ self.optional_params = {}
+ self.messages = []
+ self.completion_start_time = None
+ self.stream_options = None
+
+ def failure_handler(self, *args, **kwargs):
+ return None
+
+ async def async_failure_handler(self, *args, **kwargs):
+ return None
+
+ async def failing_make_call(client=None, **kwargs):
+ raise VertexAIError(status_code=400, message="bad input", headers={})
+
+ stream = CustomStreamWrapper(
+ completion_stream=None,
+ make_call=failing_make_call,
+ model="gemini-3-pro-preview",
+ logging_obj=DummyLogging(),
+ custom_llm_provider="vertex_ai_beta",
+ )
+
+ with pytest.raises(litellm.BadRequestError) as exc_info:
+ await stream.__anext__()
+
+ assert getattr(exc_info.value, "status_code", None) == 400
+
+
def test_vertex_ai_annotation_conversion():
"""
Test the conversion of Vertex AI grounding metadata to OpenAI annotations.
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
index 6df9abd3fee..4c5723b8284 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
@@ -354,7 +354,7 @@ async def test_register_client_remote_registration_success():
request_payload = {
"client_name": "Litellm Proxy",
- "grant_types": ["authorization_code"],
+ "grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
"token_endpoint_auth_method": "client_secret_post",
}
@@ -556,9 +556,33 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto():
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
oauth_protected_resource_mcp,
)
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+ from litellm.types.mcp import MCPAuth
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+ from litellm.proxy._types import MCPTransport
from fastapi import Request
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
+ # Clear registry
+ global_mcp_server_manager.registry.clear()
+
+ # Create mock OAuth2 server
+ oauth2_server = MCPServer(
+ server_id="test_oauth_server",
+ name="test_oauth",
+ server_name="test_oauth",
+ alias="test_oauth",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id="test_client_id",
+ client_secret="test_client_secret",
+ authorization_url="https://provider.com/oauth/authorize",
+ token_url="https://provider.com/oauth/token",
+ scopes=["read", "write"],
+ )
+ global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
# Mock request with http base_url but X-Forwarded-Proto: https
mock_request = MagicMock(spec=Request)
@@ -568,13 +592,14 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto():
# Call the endpoint
response = await oauth_protected_resource_mcp(
request=mock_request,
- mcp_server_name="test_server",
+ mcp_server_name="test_oauth",
)
# Verify response uses HTTPS URLs
assert response["authorization_servers"][0].startswith(
"https://litellm.example.com/"
)
+ assert response["scopes_supported"] == oauth2_server.scopes
@pytest.mark.asyncio
@@ -584,9 +609,33 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto():
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
oauth_authorization_server_mcp,
)
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+ from litellm.types.mcp import MCPAuth
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+ from litellm.proxy._types import MCPTransport
from fastapi import Request
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
+ # Clear registry
+ global_mcp_server_manager.registry.clear()
+
+ # Create mock OAuth2 server
+ oauth2_server = MCPServer(
+ server_id="test_oauth_server",
+ name="test_oauth",
+ server_name="test_oauth",
+ alias="test_oauth",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id="test_client_id",
+ client_secret="test_client_secret",
+ authorization_url="https://provider.com/oauth/authorize",
+ token_url="https://provider.com/oauth/token",
+ scopes=["read", "write"],
+ )
+ global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
# Mock request with http base_url but X-Forwarded-Proto: https
mock_request = MagicMock(spec=Request)
@@ -596,13 +645,15 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto():
# Call the endpoint
response = await oauth_authorization_server_mcp(
request=mock_request,
- mcp_server_name="test_server",
+ mcp_server_name="test_oauth",
)
# Verify response uses HTTPS URLs
assert response["authorization_endpoint"].startswith("https://litellm.example.com/")
assert response["token_endpoint"].startswith("https://litellm.example.com/")
assert response["registration_endpoint"].startswith("https://litellm.example.com/")
+ assert response["grant_types_supported"] == ["authorization_code", "refresh_token"]
+ assert response["scopes_supported"] == oauth2_server.scopes
@pytest.mark.asyncio
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
index 6491e11024a..d59b3f04ef5 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
@@ -594,7 +594,26 @@ class TestMCPServerManager:
assert (
server.registration_url == "https://discovered.example.com/register"
)
+ @pytest.mark.asyncio
+ async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self):
+ manager = MCPServerManager()
+ config = {
+ "example": {
+ "url": "https://example.com/mcp",
+ "transport": MCPTransport.http,
+ "auth_type": MCPAuth.oauth2,
+ "scopes": ["config"],
+ "authorization_url": "https://config.example.com/auth",
+ }
+ }
+
+ await manager.load_servers_from_config(config)
+
+ # Initialize the tool mapping
+ await manager._initialize_tool_name_to_mcp_server_name_mapping()
+ assert manager.tool_name_to_mcp_server_name_mapping == {}
+
@pytest.mark.asyncio
async def test_list_tools_handles_missing_server_alias(self):
"""Test that list_tools handles servers without alias gracefully"""
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py
index 6d6129b17bb..35ed49a84ed 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py
@@ -2,6 +2,7 @@
Unit tests for Qualifire guardrail integration.
"""
+import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@@ -139,76 +140,98 @@ class TestQualifireGuardrailEvaluateKwargs:
@pytest.mark.asyncio
async def test_evaluate_called_with_prompt_injections(self):
"""Test that evaluate is called with prompt_injections enabled."""
- from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
- QualifireGuardrail,
- )
+ # Mock the qualifire module and its types
+ mock_qualifire_types = MagicMock()
+ mock_llm_message = MagicMock()
+ mock_llm_tool_call = MagicMock()
+ mock_message_instance = MagicMock()
+ mock_llm_message.return_value = mock_message_instance
+
+ mock_qualifire_types.LLMMessage = mock_llm_message
+ mock_qualifire_types.LLMToolCall = mock_llm_tool_call
+
+ with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}):
+ from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
+ QualifireGuardrail,
+ )
- guardrail = QualifireGuardrail(
- api_key="test_key",
- prompt_injections=True,
- guardrail_name="test_guardrail",
- )
+ guardrail = QualifireGuardrail(
+ api_key="test_key",
+ prompt_injections=True,
+ guardrail_name="test_guardrail",
+ )
- # Mock the client
- mock_client = MagicMock()
- mock_result = MagicMock()
- mock_result.score = 100
- mock_result.status = "completed"
- mock_result.evaluationResults = []
- mock_client.evaluate.return_value = mock_result
- guardrail._client = mock_client
+ # Mock the client
+ mock_client = MagicMock()
+ mock_result = MagicMock()
+ mock_result.score = 100
+ mock_result.status = "completed"
+ mock_result.evaluationResults = []
+ mock_client.evaluate.return_value = mock_result
+ guardrail._client = mock_client
- messages = [{"role": "user", "content": "Hello, world!"}]
+ messages = [{"role": "user", "content": "Hello, world!"}]
- await guardrail._run_qualifire_check(
- messages=messages, output=None, dynamic_params={}
- )
+ await guardrail._run_qualifire_check(
+ messages=messages, output=None, dynamic_params={}
+ )
- # Verify evaluate was called with correct kwargs
- mock_client.evaluate.assert_called_once()
- call_kwargs = mock_client.evaluate.call_args[1]
- assert call_kwargs["prompt_injections"] is True
- assert "messages" in call_kwargs
+ # Verify evaluate was called with correct kwargs
+ mock_client.evaluate.assert_called_once()
+ call_kwargs = mock_client.evaluate.call_args[1]
+ assert call_kwargs["prompt_injections"] is True
+ assert "messages" in call_kwargs
@pytest.mark.asyncio
async def test_evaluate_called_with_multiple_checks(self):
"""Test that evaluate is called with multiple checks enabled."""
- from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
- QualifireGuardrail,
- )
+ # Mock the qualifire module and its types
+ mock_qualifire_types = MagicMock()
+ mock_llm_message = MagicMock()
+ mock_llm_tool_call = MagicMock()
+ mock_message_instance = MagicMock()
+ mock_llm_message.return_value = mock_message_instance
+
+ mock_qualifire_types.LLMMessage = mock_llm_message
+ mock_qualifire_types.LLMToolCall = mock_llm_tool_call
+
+ with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}):
+ from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
+ QualifireGuardrail,
+ )
- guardrail = QualifireGuardrail(
- api_key="test_key",
- prompt_injections=True,
- pii_check=True,
- hallucinations_check=True,
- assertions=["Output must be valid JSON"],
- guardrail_name="test_guardrail",
- )
+ guardrail = QualifireGuardrail(
+ api_key="test_key",
+ prompt_injections=True,
+ pii_check=True,
+ hallucinations_check=True,
+ assertions=["Output must be valid JSON"],
+ guardrail_name="test_guardrail",
+ )
- # Mock the client
- mock_client = MagicMock()
- mock_result = MagicMock()
- mock_result.score = 100
- mock_result.status = "completed"
- mock_result.evaluationResults = []
- mock_client.evaluate.return_value = mock_result
- guardrail._client = mock_client
+ # Mock the client
+ mock_client = MagicMock()
+ mock_result = MagicMock()
+ mock_result.score = 100
+ mock_result.status = "completed"
+ mock_result.evaluationResults = []
+ mock_client.evaluate.return_value = mock_result
+ guardrail._client = mock_client
- messages = [{"role": "user", "content": "Hello, world!"}]
+ messages = [{"role": "user", "content": "Hello, world!"}]
- await guardrail._run_qualifire_check(
- messages=messages, output="Test output", dynamic_params={}
- )
+ await guardrail._run_qualifire_check(
+ messages=messages, output="Test output", dynamic_params={}
+ )
- # Verify evaluate was called with correct kwargs
- mock_client.evaluate.assert_called_once()
- call_kwargs = mock_client.evaluate.call_args[1]
- assert call_kwargs["prompt_injections"] is True
- assert call_kwargs["pii_check"] is True
- assert call_kwargs["hallucinations_check"] is True
- assert call_kwargs["assertions"] == ["Output must be valid JSON"]
- assert call_kwargs["output"] == "Test output"
+ # Verify evaluate was called with correct kwargs
+ mock_client.evaluate.assert_called_once()
+ call_kwargs = mock_client.evaluate.call_args[1]
+ assert call_kwargs["prompt_injections"] is True
+ assert call_kwargs["pii_check"] is True
+ assert call_kwargs["hallucinations_check"] is True
+ assert call_kwargs["assertions"] == ["Output must be valid JSON"]
+ assert call_kwargs["output"] == "Test output"
class TestQualifireGuardrailCheckIfFlagged:
diff --git a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py
index 011031c1e4f..97c1733a935 100644
--- a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py
+++ b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py
@@ -40,6 +40,7 @@ class TestKeyManagementEventHooksIndependentOperations:
mock_data = MagicMock()
mock_data.key_alias = "test-key-alias"
mock_data.team_id = None
+ mock_data.send_invite_email = True
mock_response = MagicMock()
mock_response.model_dump.return_value = {"key": "sk-test", "token": "test-token"}
@@ -59,6 +60,10 @@ class TestKeyManagementEventHooksIndependentOperations:
KeyManagementEventHooks,
"_store_virtual_key_in_secret_manager",
side_effect=mock_store_secret,
+ ), patch.object(
+ KeyManagementEventHooks,
+ "_is_email_sending_enabled",
+ return_value=True,
), patch(
"litellm.store_audit_logs", False
), patch(
@@ -96,6 +101,7 @@ class TestKeyManagementEventHooksIndependentOperations:
mock_data = MagicMock()
mock_data.key_alias = "test-key-alias"
mock_data.team_id = None
+ mock_data.send_invite_email = True
mock_response = MagicMock()
mock_response.model_dump.return_value = {"key": "sk-test", "token": "test-token"}
@@ -115,6 +121,10 @@ class TestKeyManagementEventHooksIndependentOperations:
KeyManagementEventHooks,
"_store_virtual_key_in_secret_manager",
side_effect=mock_store_secret_raises,
+ ), patch.object(
+ KeyManagementEventHooks,
+ "_is_email_sending_enabled",
+ return_value=True,
), patch(
"litellm.store_audit_logs", False
), patch(
diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
index cc268ab9925..f2bae2cb14a 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
@@ -231,6 +231,40 @@ class TestListMCPServers:
assert server.url == "https://mcp.deepwiki.com/mcp"
assert server.transport == "http"
+ @pytest.mark.asyncio
+ async def test_list_mcp_servers_view_all_mode(self):
+ """Users should see all MCP servers when view_all mode is enabled."""
+
+ mock_user_auth = generate_mock_user_api_key_auth(
+ user_role=LitellmUserRoles.INTERNAL_USER
+ )
+
+ mock_servers = [
+ generate_mock_mcp_server_db_record(server_id="server-1", alias="One"),
+ generate_mock_mcp_server_db_record(server_id="server-2", alias="Two"),
+ ]
+
+ mock_manager = MagicMock()
+ mock_manager.get_all_mcp_servers_unfiltered = AsyncMock(
+ return_value=mock_servers
+ )
+
+ with patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
+ return_value="view_all",
+ ), patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
+ mock_manager,
+ ):
+ from litellm.proxy.management_endpoints.mcp_management_endpoints import (
+ fetch_all_mcp_servers,
+ )
+
+ result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
+
+ assert len(result) == 2
+ assert {server.server_id for server in result} == {"server-1", "server-2"}
+
@pytest.mark.asyncio
async def test_list_mcp_servers_combined_config_and_db(self):
"""
@@ -1096,6 +1130,51 @@ class TestHealthCheckServers:
assert result[0]["server_id"] == "server-1"
assert result[0]["status"] == "healthy"
+ @pytest.mark.asyncio
+ async def test_health_check_view_all_mode(self):
+ """view_all mode should return health info for all MCP servers."""
+
+ from litellm.proxy.management_endpoints.mcp_management_endpoints import (
+ health_check_servers,
+ )
+
+ mock_user_auth = generate_mock_user_api_key_auth(
+ user_role=LitellmUserRoles.INTERNAL_USER
+ )
+
+ health_result_one = generate_mock_mcp_server_db_record(
+ server_id="server-1", alias="One"
+ )
+ health_result_one.status = "healthy"
+
+ health_result_two = generate_mock_mcp_server_db_record(
+ server_id="server-2", alias="Two"
+ )
+ health_result_two.status = "unhealthy"
+
+ mock_manager = MagicMock()
+ mock_manager.get_all_mcp_servers_with_health_unfiltered = AsyncMock(
+ return_value=[health_result_one, health_result_two]
+ )
+
+ with patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
+ return_value="view_all",
+ ), patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
+ mock_manager,
+ ):
+ result = await health_check_servers(
+ server_ids=None,
+ user_api_key_dict=mock_user_auth,
+ )
+
+ assert len(result) == 2
+ assert result[0]["server_id"] == "server-1"
+ assert result[0]["status"] == "healthy"
+ assert result[1]["server_id"] == "server-2"
+ assert result[1]["status"] == "unhealthy"
+
@pytest.mark.asyncio
async def test_health_check_unauthorized_servers(self):
"""
diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py
index 8fdfd6897a8..ad4f53dac4b 100644
--- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py
+++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py
@@ -742,18 +742,16 @@ class TestProxySettingEndpoints:
):
"""Test updating UI settings with an allowlisted field"""
from unittest.mock import AsyncMock, MagicMock
+ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
+ from litellm.proxy._types import UserAPIKeyAuth
- class MockUser:
- def __init__(self, user_role):
- self.user_role = user_role
-
- async def mock_admin_auth():
- return MockUser(LitellmUserRoles.PROXY_ADMIN)
-
- monkeypatch.setattr(
- "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.user_api_key_auth",
- mock_admin_auth,
+ # Override the FastAPI dependency with a proper mock
+ mock_user_auth = UserAPIKeyAuth(
+ user_id="test-user-123",
+ user_role=LitellmUserRoles.PROXY_ADMIN,
)
+ app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
+
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
mock_prisma = MagicMock()
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
@@ -761,7 +759,11 @@ class TestProxySettingEndpoints:
payload = {"disable_model_add_for_internal_users": True}
- response = client.patch("/update/ui_settings", json=payload)
+ try:
+ response = client.patch("/update/ui_settings", json=payload)
+ finally:
+ # Clean up the dependency override
+ app.dependency_overrides.clear()
assert response.status_code == 200
data = response.json()
@@ -780,18 +782,16 @@ class TestProxySettingEndpoints:
):
"""Test non-allowlisted UI settings are ignored on update"""
from unittest.mock import AsyncMock, MagicMock
+ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
+ from litellm.proxy._types import UserAPIKeyAuth
- class MockUser:
- def __init__(self, user_role):
- self.user_role = user_role
-
- async def mock_admin_auth():
- return MockUser(LitellmUserRoles.PROXY_ADMIN)
-
- monkeypatch.setattr(
- "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.user_api_key_auth",
- mock_admin_auth,
+ # Override the FastAPI dependency with a proper mock
+ mock_user_auth = UserAPIKeyAuth(
+ user_id="test-user-123",
+ user_role=LitellmUserRoles.PROXY_ADMIN,
)
+ app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
+
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
mock_prisma = MagicMock()
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
@@ -802,7 +802,11 @@ class TestProxySettingEndpoints:
"unsupported_flag": True,
}
- response = client.patch("/update/ui_settings", json=payload)
+ try:
+ response = client.patch("/update/ui_settings", json=payload)
+ finally:
+ # Clean up the dependency override
+ app.dependency_overrides.clear()
assert response.status_code == 200
data = response.json()
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts
index 3e09c3c2ca8..f03f3977115 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts
@@ -1,7 +1,7 @@
+import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
+import { getSSOSettings } from "@/components/networking";
import { useQuery, UseQueryResult } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
-import { getSSOSettings } from "@/components/networking";
-import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
export interface SSOFieldSchema {
description: string;
@@ -27,13 +27,15 @@ export interface SSOSettingsValues {
proxy_base_url: string | null;
user_email: string | null;
ui_access_mode: string | null;
- role_mappings: {
- provider: string;
- group_claim: string;
- default_role: "internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer";
- roles: {
- [key: string]: string[];
- };
+ role_mappings: RoleMappings;
+}
+
+export interface RoleMappings {
+ provider: string;
+ group_claim: string;
+ default_role: "internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer";
+ roles: {
+ [key: string]: string[];
};
}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx
index 65418179595..a8b1d2cddc9 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx
@@ -18,6 +18,7 @@ import { useQueryClient } from "@tanstack/react-query";
import { Col, Grid, Icon, Tab, TabGroup, TabList, TabPanel, TabPanels, Text } from "@tremor/react";
import type { UploadProps } from "antd";
import { Form, Typography } from "antd";
+import { PlusCircleOutlined } from "@ant-design/icons";
import React, { useEffect, useMemo, useState } from "react";
import AddModelTab from "../../../components/add_model/add_model_tab";
import HealthCheckComponent from "../../../components/model_dashboard/HealthCheckComponent";
@@ -274,6 +275,30 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te
)}
+
+ {/* Missing Provider Banner */}
+
+
+
+
+
+
Missing a provider?
+
+ The LiteLLM engineering team is constantly adding support for new LLM models, providers, endpoints. If you don't see the one you need, let us know and we'll prioritize it.
+