feat: Add built-in migration lock to prevent concurrent Prisma migrate deploy (#14440)

* feat: prisma migrate deploy with lock
Author: Mini Jeong <mini.jeong@navercorp.com>

* fix: use redis cache from proxy server
Author: Mini Jeong <mini.jeong@navercorp.com>

* fix: add type checks and fix unit tests for migration lock

- Add DATABASE_URL validation in _create_baseline_migration() and _resolve_all_migrations()
- Fix MyPy type errors by adding None checks before using database_url in subprocess calls
- Add _resolve_all_migrations mock to failing unit tests to prevent filesystem errors
- Apply Black formatting to modified files

Fixes:
- MyPy type errors: database_url could be None when passed to subprocess
- Unit test failures: _resolve_all_migrations tried to create directories in read-only /test path

* fix: resolve MyPy type error in vertex_ai vertex_llm_base

Fix MyPy type checking error where vertex_api_version parameter type
was incompatible with function signature expectation.

* fix: Return 403 exception when calling GET responses api

* fix: added new step into rotate master key function for processing credentials table

* Add redisvl in requirements.txt

* fix: fixed the issue of handling root paths when processing Discovery protected resource metadata and authorization server metadata URLs.

* fix: added additional grant type into oauth_authorization_server response for fixing mcp auth register bad request issue

* fix: added RFC RECOMMENDED property(scopes_supported) to protected resource and authorization server metadata

* fix: removed initialize the tool name to MCP server name mapping(oauth2) on startup for avoiding 401 error

* fix: upgraded mcp sdk depency version for fixing ClosedResourceError

* Use already configured opentelemetry providers

Users that instrument using opentelemetry-instrument can now setup exporters as per their environment.

* Handle all protocols for all telemetry

* Add more tests

* feat(mcp): parallelize tool fetching from multiple MCP servers (#18627)

* feat(mcp): parallelize tool fetching from multiple MCP servers

Replace sequential tool fetching with asyncio.gather() to reduce
client timeouts when using multiple MCP servers.

Changes:
- mcp_server_manager.py: list_tools() now fetches tools in parallel
- server.py: _get_tools_from_mcp_servers() now fetches tools in parallel

Real-world impact (7 MCP servers example):
- Sequential: ~4.5+ seconds (exceeds typical 5-second client timeouts)
- Parallel: ~1.2 seconds (max of all servers)

Fixes #18626

* fix: copy oauth2_headers to avoid shared dict mutation in parallel tasks

* feat: add display_name, model_vendor, and model_version metadata

* added the option of adding langsmith tenant id in the env (#18623)

* fix(router): Validate routing_strategy at startup to fail fast with helpful error. (#18624)

Invalid routing_strategy values (e.g., "simple" instead of "simple-shuffle") previously failed silently, causing confusing "No deployments available" errors downstream. This change adds upfront validation in routing_strategy_init() to:

- Check if the provided strategy matches valid string values or RoutingStrategy enum
- Raise a clear ValueError listing valid options if invalid
- Fail fast at startup instead of at request time

Fixes behavior reported in #11330 where users had to debug cryptic errors.

Valid strategies: simple-shuffle, least-busy, usage-based-routing, latency-based-routing, cost-based-routing, usage-based-routing-v2

Co-authored-by: Flibbert E. Gibbitz <flibbertygibbitz@runelabs.ai>

* Add libsndfile to database Docker image for audio processing (#18612)

The litellm-database Docker image was missing the libsndfile system
library, which is required by the soundfile Python package for audio
file processing. This caused failures when using audio transcription
endpoints that attempt to calculate audio duration.

This adds libsndfile to the runtime dependencies in Dockerfile.database,
consistent with Dockerfile.alpine which already includes this library.

* Fix: Map Gemini cached_tokens to Langfuse cache_read_input_tokens (#18614)

* Fix: Map Gemini cached_tokens to Langfuse cache_read_input_tokens

Fixes #18520

## Problem
Langfuse integration was not capturing cached tokens from Gemini models.
Gemini returns cached tokens in `usage.prompt_tokens_details.cached_tokens`,
but Langfuse only read from top-level `usage.cache_read_input_tokens`
(which only Anthropic populates).

## Solution
Updated langfuse.py to check both locations:
1. First check top-level cache_read_input_tokens (for Anthropic)
2. Then check prompt_tokens_details.cached_tokens (for Gemini, OpenAI, others)

This ensures all providers' cached tokens are properly reported to Langfuse.

## Changes
- Modified litellm/integrations/langfuse/langfuse.py (lines 742-761)
- Added 3 unit tests in tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py
- All existing Langfuse tests still pass (11/11)

## Testing
- test_cached_tokens_extraction: Verifies Gemini cached_tokens extraction
- test_cached_tokens_not_present: Backward compatibility (no cached_tokens)
- test_cached_tokens_is_zero: Edge case when cached_tokens = 0

* Refactor: Extract cache token logic into helper function

Address review feedback from @officer47p

- Created _extract_cache_read_input_tokens() helper function
- Reduces code bloat in _log_langfuse_v2 method
- Improves testability and reusability
- All tests still passing (11/11)

* Adding Role Mappings

* Fixing Edit SSO Settings Modal

* feat: add user_mcp_management_mode for view_all visibility

* Fixing tests

* fix: missing mcp_allow_all_ui.png

* docs: add user_mcp_management_mode

* Align responses API streaming hooks with chat pipeline

* Clarify responses API streaming context

* Address review comments

* feat: Add GigaChat provider support (#18564)

* feat: Add GigaChat provider support

Add native support for GigaChat API (Sber AI, Russia's leading LLM).

Supported features:
- 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

Closes #18515

* fix: resolve mypy type errors in GigaChat handler

- Fix _prepare_file_data return type (use 3-tuple for cleaner type flow)
- Add type annotations for lists in _process_content_parts methods
- Add type annotations in _collapse_user_messages
- Use ChatCompletionToolCallChunk for proper tool_use typing
- Add type: ignore[override] for astreaming async generator

* refactor(gigachat): migrate to BaseConfig pattern

* fix: remove unused imports

* fix: resolve mypy type errors

* fix: mypy type errors

* refactor: address review feedback for GigaChat provider

- Remove singleton pattern, reuse litellm HTTPHandler
- Move constants/errors to transformation files, delete common_utils.py
- Add models to model_prices_and_context_window.json
- Fix ssl_verify not passed to HTTP client for embeddings

* docs: update GigaChat documentation with ssl_verify requirement

* Revert "Add redisvl in requirements.txt"

* Put reasoning summary behind feat flag

* fix: model eol

* fix: anthropic claude-3-opus-20240229 EOL

* Revert "fix: model eol"

This reverts commit 5aa1665d79.

* Fix: ImportError: qualifire package is required for QualifireGuardrail. Install it with: pip install qualifire

* fix: test_secret_manager_failure_does_not_block_email

* fix: test_update_ui_settings_allowlisted_value

* fix: test_aaamodel_prices_and_context_window_json_is_valid

* fix: test_all_models_have_display_name

* fix: async def test_bedrock_apply_guardrail_blocked()

* fix: test_databricks_embeddings[True]

* fix:test_anthropic_beta_header

* fix:test_api_error_handling

* fix:mypy mcp management

* Revert "feat(model_cost): add display_name, model_vendor, and model_version metadata to model entries"

* [Feat] New API Endpoint - Responses API (v1/responses/compact) (#18697)

* init transform_compact_response_api_request

* init acompact_responses

* init async_compact_response_api_handler in llm http handler

* init transform_compact_response_api_request for openai

* init acompact_responses

* fix acompact_responses

* add OAI Compact API

* docs responses API Compact

* code qa checks

* test_openai_compact_responses_api

* fix mypy linting

* fix: remove display name

* Add the LITELLM_REASONING_AUTO_SUMMARY in doc

* fix model map

* [UI] - Feat add request provider form on UI (#18704)

* add request provider form

* fix link to github

* add button

* fix link

* fix(streaming): normalize status code extraction to prevent 4xx errors from triggering mid-stream fallback (#18698)

在流式处理错误时,添加状态码标准化逻辑,确保 4xx 客户端错误直接抛出而不是被包装成 MidStreamFallbackError。

- 新增 _normalize_status_code 函数用于从异常对象提取状态码
- 优先从异常的 status_code 属性获取,其次从 response.status_code 获取
- 当映射异常或原始异常的状态码在 400-499 范围内时,直接抛出映射异常
- 添加单元测试验证 Vertex AI 400 错误正确抛出为 BadRequestError
- 确保流式处理中的客户端错误能够正确传播,而不会触发回退机制

---------

Co-authored-by: Eric84626 <lixiannan@gmail.com>
Co-authored-by: Eric84626 <97266539+Eric84626@users.noreply.github.com>
Co-authored-by: Sameer Kankute <sameer@berri.ai>
Co-authored-by: mangabits <1457532+mangabits@users.noreply.github.com>
Co-authored-by: Costa Tsaousis <costa@tsaousis.gr>
Co-authored-by: Nik <nikolas.garza5@gmail.com>
Co-authored-by: Shivam Rawat <161387515+shivamrawat1@users.noreply.github.com>
Co-authored-by: FlibbertyGibbitz <seth@evenkeelconsultingllc.com>
Co-authored-by: Flibbert E. Gibbitz <flibbertygibbitz@runelabs.ai>
Co-authored-by: Cesar Garcia <128240629+Chesars@users.noreply.github.com>
Co-authored-by: yuneng-jiang <yuneng.jiang@gmail.com>
Co-authored-by: Yuta Saito <uc4w6c@bma.biglobe.ne.jp>
Co-authored-by: LingXuanYin <3546599908@qq.com>
Co-authored-by: YutaSaito <36355491+uc4w6c@users.noreply.github.com>
Co-authored-by: 0717376 <103773680+0717376@users.noreply.github.com>
Co-authored-by: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Co-authored-by: Kris Xia <xiajiayi0506@gmail.com>
Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
minijeong-log 2026-01-07 03:16:24 +09:00 • committed by GitHub
parent 9e6714fe1b
commit 9f68081f6d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
87 changed files with 5826 additions and 566 deletions

View file

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

View file

@ -108,7 +108,7 @@ Some MCP servers are meant to be shared broadly—think internal knowledge bases
3. Toggle **Allow All LiteLLM Keys** on.
<Image
img={require('../img/mcp_ui.png')}
img={require('../img/mcp_allow_all_ui.png')}
style={{width: '80%', display: 'block', margin: '1rem auto'}}
alt="MCP server configuration in Admin UI"
/>
@ -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.
<iframe width="840" height="500" src="https://www.loom.com/embed/7464d444c3324078892367272fe50745" frameborder="0" webkitallowfullscreen mozallowfullscreen allowfullscreen></iframe>
## 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.

View file

@ -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/<any-model-on-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
<Tabs>
<TabItem value="Curl" label="Curl Request">
```shell
curl --location 'http://0.0.0.0:4000/chat/completions' \
--header 'Content-Type: application/json' \
--data '{
"model": "gigachat",
"messages": [
{
"role": "user",
"content": "Hello!"
}
]
}'
```
</TabItem>
<TabItem value="openai" label="OpenAI v1.0.0+">
```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)
```
</TabItem>
</Tabs>
## 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

View file

@ -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
| ZSCALER_AI_GUARD_URL | Base URL for Zscaler AI Guard API. Default is https://api.us1.zseclipse.net/v1/detection/execute-policy

View file

@ -591,3 +591,68 @@ Expected Response
</TabItem>
</Tabs>
## 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:
<Tabs>
<TabItem value="sdk" label="SDK">
```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"
)
```
</TabItem>
<TabItem value="env" label="Environment Variable">
```bash
# Set environment variable
export LITELLM_REASONING_AUTO_SUMMARY=true
# Or in your .env file
LITELLM_REASONING_AUTO_SUMMARY=true
```
</TabItem>
<TabItem value="proxy" label="Proxy Config">
```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
```
</TabItem>
</Tabs>
### 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
)
```

View file

@ -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
<Tabs>
<TabItem value="curl" label="Curl">
```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"
}'
```
</TabItem>
<TabItem value="openai-sdk" label="OpenAI Python SDK">
```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())
```
</TabItem>
</Tabs>
## 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
}
}
```

Binary file not shown.

After

Width:  |  Height:  |  Size: 135 KiB

View file

@ -541,7 +541,14 @@ const sidebars = {
},
"realtime",
"rerank",
"response_api",
{
type: "category",
label: "/responses",
items: [
"response_api",
"response_api_compact",
]
},
{
type: "category",
label: "/search",

View file

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

View file

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

View file

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

View file

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

View file

@ -375,6 +375,7 @@ LITELLM_CHAT_PROVIDERS = [
"perplexity",
"mistral",
"groq",
"gigachat",
"nvidia_nim",
"cerebras",
"baseten",

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -0,0 +1,12 @@
"""
GigaChat Chat Module
"""
from .transformation import GigaChatConfig, GigaChatError
from .streaming import GigaChatModelResponseIterator
__all__ = [
"GigaChatConfig",
"GigaChatError",
"GigaChatModelResponseIterator",
]

View file

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

View file

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

View file

@ -0,0 +1,7 @@
"""
GigaChat Embedding Module
"""
from .transformation import GigaChatEmbeddingConfig
__all__ = ["GigaChatEmbeddingConfig"]

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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[];
};
}

View file

@ -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<ModelDashboardProps> = ({ premiumUser, te
)}
</div>
</div>
{/* Missing Provider Banner */}
<div className="mb-4 px-4 py-3 bg-blue-50 rounded-lg border border-blue-100 flex items-center gap-4">
<div className="flex-shrink-0 w-10 h-10 bg-white rounded-full flex items-center justify-center border border-blue-200">
<PlusCircleOutlined style={{ fontSize: '18px', color: '#6366f1' }} />
</div>
<div className="flex-1 min-w-0">
<h4 className="text-gray-900 font-semibold text-sm m-0">Missing a provider?</h4>
<p className="text-gray-500 text-xs m-0 mt-0.5">
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.
</p>
</div>
<a
href="https://models.litellm.ai/?request=true"
target="_blank"
rel="noopener noreferrer"
className="flex-shrink-0 inline-flex items-center gap-2 px-4 py-2 bg-[#6366f1] hover:bg-[#5558e3] text-white text-sm font-medium rounded-lg transition-colors"
>
Request Provider
<svg xmlns="http://www.w3.org/2000/svg" className="h-4 w-4" fill="none" viewBox="0 0 24 24" stroke="currentColor" strokeWidth={2}>
<path strokeLinecap="round" strokeLinejoin="round" d="M10 6H6a2 2 0 00-2 2v10a2 2 0 002 2h10a2 2 0 002-2v-4M14 4h6m0 0v6m0-6L10 14" />
</svg>
</a>
</div>
{selectedModelId && !isLoading ? (
<ModelInfoView
modelId={selectedModelId}

View file

@ -189,11 +189,16 @@ const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormS
<Form.Item
noStyle
shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings}
shouldUpdate={(prevValues, currentValues) =>
prevValues.use_role_mappings !== currentValues.use_role_mappings ||
prevValues.sso_provider !== currentValues.sso_provider
}
>
{({ getFieldValue }) => {
const useRoleMappings = getFieldValue("use_role_mappings");
return useRoleMappings ? (
const provider = getFieldValue("sso_provider");
const supportsRoleMappings = provider === "okta" || provider === "generic";
return useRoleMappings && supportsRoleMappings ? (
<Form.Item
label="Group Claim"
name="group_claim"
@ -207,11 +212,16 @@ const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormS
<Form.Item
noStyle
shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings}
shouldUpdate={(prevValues, currentValues) =>
prevValues.use_role_mappings !== currentValues.use_role_mappings ||
prevValues.sso_provider !== currentValues.sso_provider
}
>
{({ getFieldValue }) => {
const useRoleMappings = getFieldValue("use_role_mappings");
return useRoleMappings ? (
const provider = getFieldValue("sso_provider");
const supportsRoleMappings = provider === "okta" || provider === "generic";
return useRoleMappings && supportsRoleMappings ? (
<>
<Form.Item label="Default Role" name="default_role" initialValue="Internal User">
<Select>

View file

@ -1,20 +1,60 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import DeleteSSOSettingsModal from "./DeleteSSOSettingsModal";
vi.mock("@/app/(dashboard)/hooks/sso/useSSOSettings", () => ({
useSSOSettings: vi.fn(() => ({
data: {
values: {
google_client_id: "test-client-id",
},
},
})),
}));
vi.mock("@/app/(dashboard)/hooks/sso/useEditSSOSettings", () => ({
useEditSSOSettings: vi.fn(() => ({
mutateAsync: vi.fn(),
isPending: false,
})),
}));
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
default: vi.fn(() => ({
accessToken: "test-token",
userId: "test-user-id",
userRole: "proxy_admin",
})),
}));
const createQueryClient = () =>
new QueryClient({
defaultOptions: {
queries: {
retry: false,
gcTime: 0,
},
},
});
describe("DeleteSSOSettingsModal", () => {
it("should render", () => {
const onCancel = vi.fn();
const onSuccess = vi.fn();
const queryClient = createQueryClient();
render(
<DeleteSSOSettingsModal isVisible={true} onCancel={onCancel} onSuccess={onSuccess} accessToken="test-token" />,
<QueryClientProvider client={queryClient}>
<DeleteSSOSettingsModal isVisible={true} onCancel={onCancel} onSuccess={onSuccess} />
</QueryClientProvider>,
);
expect(screen.getByText("Confirm Clear SSO Settings")).toBeInTheDocument();
expect(
screen.getByText("Are you sure you want to clear all SSO settings? This action cannot be undone."),
screen.getByText(
"Are you sure you want to clear all SSO settings? Users will no longer be able to login using SSO after this change.",
),
).toBeInTheDocument();
expect(screen.getByText("Users will no longer be able to login using SSO after this change.")).toBeInTheDocument();
});
});

View file

@ -1,79 +1,66 @@
import { Modal } from "antd";
import { useEditSSOSettings } from "@/app/(dashboard)/hooks/sso/useEditSSOSettings";
import { useSSOSettings } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
import React from "react";
import DeleteResourceModal from "../../../../common_components/DeleteResourceModal";
import NotificationsManager from "../../../../molecules/notifications_manager";
import { updateSSOSettings } from "../../../../networking";
import { parseErrorMessage } from "../../../../shared/errorUtils";
import { detectSSOProvider } from "../utils";
interface DeleteSSOSettingsModalProps {
isVisible: boolean;
onCancel: () => void;
onSuccess: () => void;
accessToken: string | null;
}
const DeleteSSOSettingsModal: React.FC<DeleteSSOSettingsModalProps> = ({
isVisible,
onCancel,
onSuccess,
accessToken,
}) => {
const DeleteSSOSettingsModal: React.FC<DeleteSSOSettingsModalProps> = ({ isVisible, onCancel, onSuccess }) => {
const { data: ssoSettings } = useSSOSettings();
const { mutateAsync: editSSOSettings, isPending: isEditingSSOSettings } = useEditSSOSettings();
// Handle clearing SSO settings
const handleClearSSO = async () => {
if (!accessToken) {
NotificationsManager.fromBackend("No access token available");
return;
}
const clearSettings = {
google_client_id: null,
google_client_secret: null,
microsoft_client_id: null,
microsoft_client_secret: null,
microsoft_tenant: null,
generic_client_id: null,
generic_client_secret: null,
generic_authorization_endpoint: null,
generic_token_endpoint: null,
generic_userinfo_endpoint: null,
proxy_base_url: null,
user_email: null,
sso_provider: null,
role_mappings: null,
};
try {
// Clear all SSO settings
const clearSettings = {
google_client_id: null,
google_client_secret: null,
microsoft_client_id: null,
microsoft_client_secret: null,
microsoft_tenant: null,
generic_client_id: null,
generic_client_secret: null,
generic_authorization_endpoint: null,
generic_token_endpoint: null,
generic_userinfo_endpoint: null,
proxy_base_url: null,
user_email: null,
sso_provider: null,
};
await updateSSOSettings(accessToken, clearSettings);
NotificationsManager.success("SSO settings cleared successfully");
// Close modal and trigger success callback
onCancel();
onSuccess();
} catch (error) {
console.error("Failed to clear SSO settings:", error);
NotificationsManager.fromBackend("Failed to clear SSO settings: " + parseErrorMessage(error));
}
await editSSOSettings(clearSettings, {
onSuccess: () => {
NotificationsManager.success("SSO settings cleared successfully");
onCancel();
onSuccess();
},
onError: (error) => {
NotificationsManager.fromBackend("Failed to clear SSO settings: " + parseErrorMessage(error));
},
});
};
return (
<Modal
<DeleteResourceModal
isOpen={isVisible}
title="Confirm Clear SSO Settings"
visible={isVisible}
onOk={handleClearSSO}
alertMessage="This action cannot be undone."
message="Are you sure you want to clear all SSO settings? Users will no longer be able to login using SSO after this change."
resourceInformationTitle="SSO Settings"
resourceInformation={[
{ label: "Provider", value: (ssoSettings?.values && detectSSOProvider(ssoSettings?.values)) || "Generic" },
]}
onCancel={onCancel}
okText="Yes, Clear"
cancelText="Cancel"
okButtonProps={{
danger: true,
style: {
backgroundColor: "#dc2626",
borderColor: "#dc2626",
},
}}
>
<p>Are you sure you want to clear all SSO settings? This action cannot be undone.</p>
<p>Users will no longer be able to login using SSO after this change.</p>
</Modal>
onOk={handleClearSSO}
confirmLoading={isEditingSSOSettings}
/>
);
};

View file

@ -0,0 +1,92 @@
import type { RoleMappings as RoleMappingsType } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
import { screen } from "@testing-library/react";
import { describe, expect, it } from "vitest";
import { renderWithProviders } from "../../../../../tests/test-utils";
import RoleMappings from "./RoleMappings";
describe("RoleMappings", () => {
it("should render successfully", () => {
const roleMappings: RoleMappingsType = {
provider: "generic",
group_claim: "groups",
default_role: "internal_user",
roles: {
proxy_admin: ["admin-group"],
proxy_admin_viewer: [],
internal_user: ["user-group"],
internal_user_viewer: [],
},
};
renderWithProviders(<RoleMappings roleMappings={roleMappings} />);
expect(screen.getByText("Role Mappings")).toBeInTheDocument();
});
it("should return null when roleMappings is undefined", () => {
const { container } = renderWithProviders(<RoleMappings roleMappings={undefined} />);
expect(container.firstChild).toBeNull();
});
it("should display Group Claim and Default Role with correct values and display names", () => {
const testCases: Array<{ role: RoleMappingsType["default_role"]; displayName: string; groupClaim: string }> = [
{ role: "internal_user_viewer", displayName: "Internal Viewer", groupClaim: "custom-groups-1" },
{ role: "internal_user", displayName: "Internal User", groupClaim: "custom-groups-2" },
{ role: "proxy_admin_viewer", displayName: "Proxy Admin Viewer", groupClaim: "custom-groups-3" },
{ role: "proxy_admin", displayName: "Proxy Admin", groupClaim: "custom-groups-4" },
];
testCases.forEach(({ role, displayName, groupClaim }) => {
const roleMappings: RoleMappingsType = {
provider: "generic",
group_claim: groupClaim,
default_role: role,
roles: {
proxy_admin: [],
proxy_admin_viewer: [],
internal_user: [],
internal_user_viewer: [],
},
};
const { unmount } = renderWithProviders(<RoleMappings roleMappings={roleMappings} />);
expect(screen.getByText("Group Claim")).toBeInTheDocument();
expect(screen.getByText(groupClaim)).toBeInTheDocument();
expect(screen.getByText("Default Role")).toBeInTheDocument();
const displayNameElements = screen.getAllByText(displayName);
expect(displayNameElements.length).toBeGreaterThan(0);
unmount();
});
});
it("should display table with roles, groups as Tags when mapped, and 'No groups mapped' when empty", () => {
const roleMappings: RoleMappingsType = {
provider: "generic",
group_claim: "groups",
default_role: "internal_user",
roles: {
proxy_admin: ["admin-group-1", "admin-group-2", "admin-group-3"],
proxy_admin_viewer: ["viewer-group"],
internal_user: ["user-group"],
internal_user_viewer: [],
},
};
renderWithProviders(<RoleMappings roleMappings={roleMappings} />);
expect(screen.getByText("Role")).toBeInTheDocument();
expect(screen.getByText("Mapped Groups")).toBeInTheDocument();
expect(screen.getAllByText("Proxy Admin").length).toBeGreaterThan(0);
expect(screen.getAllByText("Proxy Admin Viewer").length).toBeGreaterThan(0);
expect(screen.getAllByText("Internal User").length).toBeGreaterThan(0);
expect(screen.getAllByText("Internal Viewer").length).toBeGreaterThan(0);
expect(screen.getByText("admin-group-1")).toBeInTheDocument();
expect(screen.getByText("admin-group-2")).toBeInTheDocument();
expect(screen.getByText("admin-group-3")).toBeInTheDocument();
expect(screen.getByText("viewer-group")).toBeInTheDocument();
expect(screen.getByText("user-group")).toBeInTheDocument();
expect(screen.getByText("No groups mapped")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,74 @@
import type { RoleMappings as RoleMappingsType } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
import { Card, Divider, Table, Tag, Typography } from "antd";
import { Users } from "lucide-react";
import { defaultRoleDisplayNames } from "./constants";
const { Title, Text } = Typography;
export default function RoleMappings({ roleMappings }: { roleMappings: RoleMappingsType | undefined }) {
if (!roleMappings) {
return null;
}
const roleMappingsColumns = [
{
title: "Role",
dataIndex: "role",
key: "role",
render: (text: string) => <Text strong>{defaultRoleDisplayNames[text]}</Text>,
},
{
title: "Mapped Groups",
dataIndex: "groups",
key: "groups",
render: (groups: string[]) => (
<>
{groups.length > 0 ? (
groups.map((group, index) => (
<Tag key={index} color="blue">
{group}
</Tag>
))
) : (
<Text className="text-gray-400 italic">No groups mapped</Text>
)}
</>
),
},
];
return (
<Card>
<div className="flex items-center gap-3">
<Users className="w-6 h-6 text-gray-400 mb-2" />
<Title level={3}>Role Mappings</Title>
</div>
<div className="space-y-8">
<div className="grid grid-cols-2 gap-4">
<div>
<Title level={5}>Group Claim</Title>
<div>
<Text code>{roleMappings.group_claim}</Text>
</div>
</div>
<div>
<Title level={5}>Default Role</Title>
<div>
<Text strong>{defaultRoleDisplayNames[roleMappings.default_role]}</Text>
</div>
</div>
</div>
<Divider />
<Table
columns={roleMappingsColumns}
dataSource={Object.entries(roleMappings.roles).map(([role, groups]) => ({
role,
groups,
}))}
pagination={false}
bordered
size="small"
className="w-full"
/>
</div>
</Card>
);
}

View file

@ -1,23 +1,23 @@
"use client";
import { useSSOSettings, type SSOSettingsValues } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { Button, Card, Descriptions, Space, Typography } from "antd";
import { Edit, Shield, Trash2 } from "lucide-react";
import { useState } from "react";
import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./constants";
import AddSSOSettingsModal from "./Modals/AddSSOSettingsModal";
import DeleteSSOSettingsModal from "./Modals/DeleteSSOSettingsModal";
import EditSSOSettingsModal from "./Modals/EditSSOSettingsModal";
import RedactableField from "./RedactableField";
import RoleMappings from "./RoleMappings";
import SSOSettingsEmptyPlaceholder from "./SSOSettingsEmptyPlaceholder";
import SSOSettingsLoadingSkeleton from "./SSOSettingsLoadingSkeleton";
import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./constants";
import { detectSSOProvider } from "./utils";
const { Title, Text } = Typography;
export default function SSOSettings() {
const { data: ssoSettings, refetch, isLoading } = useSSOSettings();
const { accessToken } = useAuthorized();
const [isDeleteModalVisible, setIsDeleteModalVisible] = useState(false);
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
const [isEditModalVisible, setIsEditModalVisible] = useState(false);
@ -26,24 +26,8 @@ export default function SSOSettings() {
Boolean(ssoSettings?.values.microsoft_client_id) ||
Boolean(ssoSettings?.values.generic_client_id);
// Determine the SSO provider based on the configuration
const detectSSOProvider = (values: SSOSettingsValues): string | null => {
if (values.google_client_id) return "google";
if (values.microsoft_client_id) return "microsoft";
if (values.generic_client_id) {
// Check if it looks like Okta/Auth0 based on endpoints
if (
values.generic_authorization_endpoint?.includes("okta") ||
values.generic_authorization_endpoint?.includes("auth0")
) {
return "okta";
}
return "generic";
}
return null;
};
const selectedProvider = ssoSettings?.values ? detectSSOProvider(ssoSettings.values) : null;
const isRoleMappingsEnabled = Boolean(ssoSettings?.values.role_mappings);
const renderEndpointValue = (value?: string | null) => (
<Text className="font-mono text-gray-600 text-sm" copyable={!!value}>
@ -185,46 +169,52 @@ export default function SSOSettings() {
{isLoading ? (
<SSOSettingsLoadingSkeleton />
) : (
<Card>
<Space direction="vertical" size="large" className="w-full">
{/* Header Section */}
<div className="flex items-center justify-between">
<div className="flex items-center gap-3">
<Shield className="w-6 h-6 text-gray-400" />
<div>
<Title level={3}>SSO Configuration</Title>
<Text type="secondary">Manage Single Sign-On authentication settings</Text>
<Space direction="vertical" size="large" className="w-full">
<Card>
<Space direction="vertical" size="large" className="w-full">
{/* Header Section */}
<div className="flex items-center justify-between">
<div className="flex items-center gap-3">
<Shield className="w-6 h-6 text-gray-400" />
<div>
<Title level={3}>SSO Configuration</Title>
<Text type="secondary">Manage Single Sign-On authentication settings</Text>
</div>
</div>
<div className="flex items-center gap-3">
{isSSOConfigured && (
<>
<Button icon={<Edit className="w-4 h-4" />} onClick={() => setIsEditModalVisible(true)}>
Edit SSO Settings
</Button>
<Button
danger
icon={<Trash2 className="w-4 h-4" />}
onClick={() => setIsDeleteModalVisible(true)}
>
Delete SSO Settings
</Button>
</>
)}
</div>
</div>
<div className="flex items-center gap-3">
{isSSOConfigured && (
<>
<Button icon={<Edit className="w-4 h-4" />} onClick={() => setIsEditModalVisible(true)}>
Edit SSO Settings
</Button>
<Button danger icon={<Trash2 className="w-4 h-4" />} onClick={() => setIsDeleteModalVisible(true)}>
Delete SSO Settings
</Button>
</>
)}
</div>
</div>
{isSSOConfigured ? (
renderSSOSettings()
) : (
<SSOSettingsEmptyPlaceholder onAdd={() => setIsAddModalVisible(true)} />
)}
</Space>
</Card>
{isSSOConfigured ? (
renderSSOSettings()
) : (
<SSOSettingsEmptyPlaceholder onAdd={() => setIsAddModalVisible(true)} />
)}
</Space>
</Card>
{isRoleMappingsEnabled && <RoleMappings roleMappings={ssoSettings?.values.role_mappings} />}
</Space>
)}
<DeleteSSOSettingsModal
isVisible={isDeleteModalVisible}
onCancel={() => setIsDeleteModalVisible(false)}
onSuccess={() => refetch()}
accessToken={accessToken}
/>
<AddSSOSettingsModal

View file

@ -13,3 +13,10 @@ export const ssoProviderDisplayNames: Record<string, string> = {
okta: "Okta / Auth0 SSO",
generic: "Generic SSO",
};
export const defaultRoleDisplayNames: Record<string, string> = {
internal_user_viewer: "Internal Viewer",
internal_user: "Internal User",
proxy_admin_viewer: "Proxy Admin Viewer",
proxy_admin: "Proxy Admin",
};

View file

@ -55,6 +55,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "proxy_admin",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
other_field: "value",
};
@ -83,6 +84,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "internal_user",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -100,6 +102,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "internal_user_viewer",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -121,6 +124,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "proxy_admin_viewer",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -142,6 +146,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "internal_user",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -160,6 +165,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "internal_user",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -174,6 +180,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "internal_user_viewer",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -186,6 +193,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "internal_user",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -198,6 +206,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "proxy_admin_viewer",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -210,6 +219,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "proxy_admin",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -222,6 +232,7 @@ describe("processSSOSettingsPayload", () => {
default_role: "unknown_role",
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);
@ -233,6 +244,7 @@ describe("processSSOSettingsPayload", () => {
const formValues = {
group_claim: "groups",
use_role_mappings: true,
sso_provider: "generic",
};
const result = processSSOSettingsPayload(formValues);

View file

@ -1,3 +1,5 @@
import { SSOSettingsValues } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
/**
* Processes SSO settings form values and transforms them into the payload format expected by the API
* Handles role mappings transformation and field extraction
@ -18,7 +20,7 @@ export const processSSOSettingsPayload = (formValues: Record<string, any>): Reco
...rest,
};
// Add role mappings if use_role_mappings is checked
// Add role mappings only if use_role_mappings is checked AND provider supports role mappings
if (use_role_mappings) {
// Helper function to split comma-separated string into array
const splitTeams = (teams: string | undefined): string[] => {
@ -52,3 +54,20 @@ export const processSSOSettingsPayload = (formValues: Record<string, any>): Reco
return payload;
};
// Determine the SSO provider based on the configuration
export const detectSSOProvider = (values: SSOSettingsValues): string | null => {
if (values.google_client_id) return "google";
if (values.microsoft_client_id) return "microsoft";
if (values.generic_client_id) {
// Check if it looks like Okta/Auth0 based on endpoints
if (
values.generic_authorization_endpoint?.includes("okta") ||
values.generic_authorization_endpoint?.includes("auth0")
) {
return "okta";
}
return "generic";
}
return null;
};

View file

@ -136,7 +136,7 @@ export const useMcpOAuthFlow = ({
if (!hasPreconfiguredCredentials) {
const registration = await registerMcpOAuthClient(accessToken, serverId, {
client_name: temporaryPayload.alias || temporaryPayload.server_name || serverId,
grant_types: ["authorization_code"],
grant_types: ["authorization_code", "refresh_token"],
response_types: ["code"],
token_endpoint_auth_method:
temporaryPayload.credentials && temporaryPayload.credentials.client_secret ? "client_secret_post" : "none",