mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Revert "feat: Add built-in migration lock to prevent concurrent Prisma migrat…" (#18719)
This reverts commit 9f68081f6d.
This commit is contained in:
parent
9c544949f8
commit
7004734528
87 changed files with 567 additions and 5827 deletions
|
|
@ -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 libsndfile
|
||||
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip
|
||||
|
||||
WORKDIR /app
|
||||
# Copy the current directory contents into the container at /app
|
||||
|
|
|
|||
|
|
@ -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_allow_all_ui.png')}
|
||||
img={require('../img/mcp_ui.png')}
|
||||
style={{width: '80%', display: 'block', margin: '1rem auto'}}
|
||||
alt="MCP server configuration in Admin UI"
|
||||
/>
|
||||
|
|
@ -634,18 +634,3 @@ 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.
|
||||
|
|
|
|||
|
|
@ -1,283 +0,0 @@
|
|||
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
|
||||
|
|
@ -111,7 +111,6 @@ 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
|
||||
|
|
@ -231,7 +230,6 @@ 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. |
|
||||
|
|
@ -671,7 +669,6 @@ 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
|
||||
|
|
@ -710,7 +707,6 @@ 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
|
||||
|
|
@ -778,7 +774,6 @@ 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
|
||||
|
|
@ -893,4 +888,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
|
||||
|
|
@ -591,68 +591,3 @@ 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
|
||||
)
|
||||
```
|
||||
|
|
|
|||
|
|
@ -1,104 +0,0 @@
|
|||
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.
|
Before Width: | Height: | Size: 135 KiB |
|
|
@ -541,14 +541,7 @@ const sidebars = {
|
|||
},
|
||||
"realtime",
|
||||
"rerank",
|
||||
{
|
||||
type: "category",
|
||||
label: "/responses",
|
||||
items: [
|
||||
"response_api",
|
||||
"response_api_compact",
|
||||
]
|
||||
},
|
||||
"response_api",
|
||||
{
|
||||
type: "category",
|
||||
label: "/search",
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ 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:
|
||||
|
|
@ -19,103 +18,6 @@ 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."""
|
||||
|
|
@ -444,50 +346,19 @@ class ProxyExtrasDBManager:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def setup_database(
|
||||
use_migrate: bool = False, redis_cache: Optional[RedisCache] = None
|
||||
) -> bool:
|
||||
def setup_database(use_migrate: bool = False) -> bool:
|
||||
"""
|
||||
Set up the database using either prisma migrate or prisma db push
|
||||
Uses migrations from litellm-proxy-extras package.
|
||||
In multi-instance environment, use redis lock to prevent concurrent execution.
|
||||
Uses migrations from litellm-proxy-extras package
|
||||
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -197,7 +197,6 @@ 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
|
||||
|
|
@ -276,7 +275,6 @@ 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
|
||||
|
||||
|
|
@ -1442,8 +1440,6 @@ 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
|
||||
|
|
|
|||
|
|
@ -255,8 +255,6 @@ LLM_CONFIG_NAMES = (
|
|||
"GithubCopilotEmbeddingConfig",
|
||||
"NebiusConfig",
|
||||
"WandbConfig",
|
||||
"GigaChatConfig",
|
||||
"GigaChatEmbeddingConfig",
|
||||
"DashScopeChatConfig",
|
||||
"MoonshotChatConfig",
|
||||
"DockerModelRunnerChatConfig",
|
||||
|
|
@ -646,8 +644,6 @@ _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"),
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ Handler for transforming /chat/completions api requests to litellm.responses req
|
|||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
|
|
@ -23,7 +22,6 @@ 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
|
||||
|
|
@ -693,26 +691,19 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if isinstance(reasoning_effort, dict):
|
||||
return Reasoning(**reasoning_effort) # type: ignore[typeddict-item]
|
||||
|
||||
# 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 string is passed, map with summary="detailed"
|
||||
if reasoning_effort == "none":
|
||||
return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") # type: ignore
|
||||
return Reasoning(effort="none", summary="detailed") # type: ignore
|
||||
elif reasoning_effort == "high":
|
||||
return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high")
|
||||
return Reasoning(effort="high", summary="detailed")
|
||||
elif reasoning_effort == "xhigh":
|
||||
return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") # type: ignore[typeddict-item]
|
||||
return Reasoning(effort="xhigh", summary="detailed") # type: ignore[typeddict-item]
|
||||
elif reasoning_effort == "medium":
|
||||
return Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium")
|
||||
return Reasoning(effort="medium", summary="detailed")
|
||||
elif reasoning_effort == "low":
|
||||
return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low")
|
||||
return Reasoning(effort="low", summary="detailed")
|
||||
elif reasoning_effort == "minimal":
|
||||
return Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal")
|
||||
return Reasoning(effort="minimal", summary="detailed")
|
||||
return None
|
||||
|
||||
def _transform_response_format_to_text_format(
|
||||
|
|
|
|||
|
|
@ -375,7 +375,6 @@ LITELLM_CHAT_PROVIDERS = [
|
|||
"perplexity",
|
||||
"mistral",
|
||||
"groq",
|
||||
"gigachat",
|
||||
"nvidia_nim",
|
||||
"cerebras",
|
||||
"baseten",
|
||||
|
|
|
|||
|
|
@ -187,12 +187,6 @@
|
|||
"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"
|
||||
|
|
|
|||
|
|
@ -50,42 +50,6 @@ 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__(
|
||||
|
|
@ -793,8 +757,8 @@ class LangFuseLogger:
|
|||
cache_creation_input_tokens = (
|
||||
_usage_obj.get("cache_creation_input_tokens") or 0
|
||||
)
|
||||
cache_read_input_tokens = _extract_cache_read_input_tokens(
|
||||
_usage_obj
|
||||
cache_read_input_tokens = (
|
||||
_usage_obj.get("cache_read_input_tokens") or 0
|
||||
)
|
||||
|
||||
usage = {
|
||||
|
|
|
|||
|
|
@ -40,7 +40,6 @@ 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()
|
||||
|
|
@ -49,7 +48,6 @@ 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
|
||||
|
|
@ -78,7 +76,6 @@ 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 = (
|
||||
|
|
@ -89,13 +86,11 @@ 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(
|
||||
|
|
@ -370,11 +365,8 @@ 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:
|
||||
|
|
@ -426,7 +418,6 @@ 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:
|
||||
|
|
@ -475,9 +466,6 @@ 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
|
||||
|
|
@ -503,16 +491,13 @@ 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=headers,
|
||||
headers={"x-api-key": langsmith_api_key},
|
||||
)
|
||||
|
||||
return response.json()
|
||||
|
|
|
|||
|
|
@ -196,88 +196,50 @@ 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
|
||||
|
||||
def create_tracer_provider():
|
||||
provider = TracerProvider(resource=_get_litellm_resource())
|
||||
provider.add_span_processor(self._get_span_processor())
|
||||
return provider
|
||||
# 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
|
||||
|
||||
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,
|
||||
)
|
||||
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__,
|
||||
)
|
||||
|
||||
# Grab our tracer from the TracerProvider (not from global context)
|
||||
# This ensures we use the provided TracerProvider (e.g., for testing)
|
||||
|
|
@ -295,24 +257,39 @@ class OpenTelemetry(CustomLogger):
|
|||
return
|
||||
|
||||
from opentelemetry import metrics
|
||||
from opentelemetry.sdk.metrics import MeterProvider
|
||||
from opentelemetry.sdk.metrics import Histogram, MeterProvider
|
||||
|
||||
def create_meter_provider():
|
||||
metric_reader = self._get_metric_reader()
|
||||
return MeterProvider(
|
||||
metric_readers=[metric_reader], resource=_get_litellm_resource()
|
||||
# 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,
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
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 = meter_provider.get_meter(__name__)
|
||||
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)
|
||||
|
||||
self._operation_duration_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38
|
||||
|
|
@ -350,26 +327,22 @@ class OpenTelemetry(CustomLogger):
|
|||
if not self.config.enable_events:
|
||||
return
|
||||
|
||||
from opentelemetry._logs import get_logger_provider, set_logger_provider
|
||||
from opentelemetry._logs import set_logger_provider
|
||||
from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider
|
||||
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
|
||||
|
||||
def create_logger_provider():
|
||||
provider = OTLoggerProvider(resource=_get_litellm_resource())
|
||||
# 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
|
||||
log_exporter = self._get_log_exporter()
|
||||
provider.add_log_record_processor(
|
||||
BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
|
||||
)
|
||||
return provider
|
||||
if log_exporter:
|
||||
logger_provider.add_log_record_processor(
|
||||
BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
set_logger_provider(logger_provider)
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
self._handle_success(kwargs, response_obj, start_time, end_time)
|
||||
|
|
@ -971,15 +944,6 @@ 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
|
||||
|
||||
|
|
@ -1843,8 +1807,7 @@ class OpenTelemetry(CustomLogger):
|
|||
)
|
||||
return self.OTEL_EXPORTER
|
||||
|
||||
otel_logs_exporter = os.getenv("OTEL_LOGS_EXPORTER")
|
||||
if self.OTEL_EXPORTER == "console" or otel_logs_exporter == "console":
|
||||
if self.OTEL_EXPORTER == "console":
|
||||
from opentelemetry.sdk._logs.export import ConsoleLogExporter
|
||||
|
||||
verbose_logger.debug(
|
||||
|
|
@ -1891,67 +1854,6 @@ 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]:
|
||||
|
|
|
|||
|
|
@ -2000,56 +2000,24 @@ class CustomStreamWrapper:
|
|||
)
|
||||
## Map to OpenAI Exception
|
||||
try:
|
||||
mapped_exception = exception_type(
|
||||
raise exception_type(
|
||||
model=self.model,
|
||||
custom_llm_provider=self.custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
except Exception as mapping_error:
|
||||
mapped_exception = mapping_error
|
||||
except Exception as e:
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
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,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _strip_sse_data_from_chunk(chunk: Optional[str]) -> Optional[str]:
|
||||
|
|
|
|||
|
|
@ -242,30 +242,3 @@ 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 ######
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -91,7 +91,6 @@ 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,
|
||||
|
|
@ -851,9 +850,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
sync_httpx_client = _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
sync_httpx_client = _get_httpx_client()
|
||||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
|
|
@ -899,8 +896,7 @@ 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),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider)
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
|
@ -2008,10 +2004,6 @@ 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:
|
||||
|
|
@ -2068,18 +2060,6 @@ 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,
|
||||
|
|
@ -2117,8 +2097,6 @@ 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(
|
||||
|
|
@ -2128,8 +2106,6 @@ 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
|
||||
|
|
@ -2213,18 +2189,6 @@ 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,
|
||||
|
|
@ -2263,8 +2227,6 @@ 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
|
||||
|
|
@ -2275,8 +2237,6 @@ 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
|
||||
|
|
@ -3566,174 +3526,6 @@ 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
|
||||
|
|
@ -8496,4 +8288,4 @@ class BaseLLMHTTPHandler:
|
|||
return skills_api_provider_config.transform_delete_skill_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
)
|
||||
|
|
@ -1,23 +0,0 @@
|
|||
"""
|
||||
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",
|
||||
]
|
||||
|
|
@ -1,241 +0,0 @@
|
|||
"""
|
||||
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
|
||||
|
|
@ -1,12 +0,0 @@
|
|||
"""
|
||||
GigaChat Chat Module
|
||||
"""
|
||||
|
||||
from .transformation import GigaChatConfig, GigaChatError
|
||||
from .streaming import GigaChatModelResponseIterator
|
||||
|
||||
__all__ = [
|
||||
"GigaChatConfig",
|
||||
"GigaChatError",
|
||||
"GigaChatModelResponseIterator",
|
||||
]
|
||||
|
|
@ -1,134 +0,0 @@
|
|||
"""
|
||||
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
|
||||
|
|
@ -1,473 +0,0 @@
|
|||
"""
|
||||
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,
|
||||
)
|
||||
|
|
@ -1,7 +0,0 @@
|
|||
"""
|
||||
GigaChat Embedding Module
|
||||
"""
|
||||
|
||||
from .transformation import GigaChatEmbeddingConfig
|
||||
|
||||
__all__ = ["GigaChatEmbeddingConfig"]
|
||||
|
|
@ -1,212 +0,0 @@
|
|||
"""
|
||||
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,
|
||||
)
|
||||
|
|
@ -1,211 +0,0 @@
|
|||
"""
|
||||
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
|
||||
|
|
@ -500,69 +500,3 @@ 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
|
||||
|
|
|
|||
|
|
@ -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, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -168,6 +168,7 @@ class VertexBase:
|
|||
)
|
||||
|
||||
def _credentials_from_default_auth(self, scopes):
|
||||
|
||||
import google.auth as google_auth
|
||||
|
||||
return google_auth.default(scopes=scopes)
|
||||
|
|
@ -391,7 +392,7 @@ class VertexBase:
|
|||
Returns
|
||||
token, url
|
||||
"""
|
||||
version: Optional[Literal["v1", "v1beta1"]] = None
|
||||
version: Optional[Literal["v1beta1", "v1"]] = None
|
||||
if custom_llm_provider == "gemini":
|
||||
url, endpoint = _get_gemini_url(
|
||||
mode=mode,
|
||||
|
|
@ -414,7 +415,7 @@ class VertexBase:
|
|||
stream=stream,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_api_version=cast(Literal["v1", "v1beta1"], version),
|
||||
vertex_api_version=version,
|
||||
)
|
||||
|
||||
return self._check_custom_proxy(
|
||||
|
|
|
|||
|
|
@ -2141,49 +2141,6 @@ 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
|
||||
|
|
@ -5267,28 +5224,6 @@ 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
|
||||
|
|
|
|||
|
|
@ -15831,68 +15831,6 @@
|
|||
"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",
|
||||
|
|
@ -32154,4 +32092,3 @@
|
|||
"mode": "chat"
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -15,7 +15,6 @@ 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"],
|
||||
|
|
@ -382,30 +381,13 @@ async def callback(code: str, state: str):
|
|||
# ------------------------------
|
||||
# Optional .well-known endpoints for MCP + OAuth discovery
|
||||
# ------------------------------
|
||||
"""
|
||||
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/{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": [
|
||||
(
|
||||
|
|
@ -419,25 +401,14 @@ 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 [],
|
||||
}
|
||||
|
||||
"""
|
||||
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/{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)
|
||||
|
||||
|
|
@ -452,21 +423,16 @@ 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"],
|
||||
"scopes_supported": mcp_server.scopes if mcp_server else [],
|
||||
"grant_types_supported": ["authorization_code", "refresh_token"],
|
||||
"grant_types_supported": ["authorization_code"],
|
||||
"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" if mcp_server_name else f"{request_base_url}/register",
|
||||
"registration_endpoint": f"{request_base_url}/{mcp_server_name}/register",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
async def _fetch_server_tools(server_id: str) -> List[MCPTool]:
|
||||
"""Fetch tools from a single server with error handling."""
|
||||
for server_id in allowed_mcp_servers:
|
||||
server = self.get_mcp_server_by_id(server_id)
|
||||
if server is None:
|
||||
verbose_logger.warning(f"MCP Server {server_id} not found")
|
||||
return []
|
||||
continue
|
||||
|
||||
# Get server-specific auth header if available
|
||||
server_auth_header = None
|
||||
|
|
@ -685,21 +685,15 @@ class MCPServerManager:
|
|||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
)
|
||||
return tools
|
||||
list_tools_result.extend(tools)
|
||||
verbose_logger.info(
|
||||
f"Successfully fetched {len(tools)} tools from server {server.name}"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Failed to list tools from server {server.name}: {str(e)}. Continuing with other servers."
|
||||
)
|
||||
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
|
||||
]
|
||||
# Continue with other servers instead of failing completely
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully fetched {len(list_tools_result)} tools total from all servers"
|
||||
|
|
@ -2009,9 +2003,6 @@ 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
|
||||
|
|
@ -2293,7 +2284,14 @@ class MCPServerManager:
|
|||
# Check all accessible servers
|
||||
target_server_ids = allowed_server_ids
|
||||
|
||||
return await self._run_health_checks(target_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
|
||||
|
||||
async def get_all_allowed_mcp_servers(
|
||||
self,
|
||||
|
|
@ -2308,6 +2306,8 @@ 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,56 +2319,40 @@ class MCPServerManager:
|
|||
verbose_logger.warning(f"MCP Server {server_id} not found in registry")
|
||||
continue
|
||||
|
||||
mcp_server_table = self._build_mcp_server_table(server)
|
||||
# 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,
|
||||
)
|
||||
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.
|
||||
|
|
@ -2376,34 +2360,5 @@ 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()
|
||||
|
|
|
|||
|
|
@ -709,8 +709,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
extra_headers: Optional[Dict[str, str]] = None
|
||||
if server.auth_type == MCPAuth.oauth2:
|
||||
# Copy to avoid mutating the original dict (important for parallel fetching)
|
||||
extra_headers = oauth2_headers.copy() if oauth2_headers else None
|
||||
extra_headers = oauth2_headers
|
||||
|
||||
if server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
|
|
@ -756,10 +755,11 @@ if MCP_AVAILABLE:
|
|||
# Decide whether to add prefix based on number of allowed servers
|
||||
add_prefix = not (len(allowed_mcp_servers) == 1)
|
||||
|
||||
async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]:
|
||||
"""Fetch and filter tools from a single server with error handling."""
|
||||
# Get tools from each allowed server
|
||||
all_tools = []
|
||||
for server in allowed_mcp_servers:
|
||||
if server is None:
|
||||
return []
|
||||
continue
|
||||
|
||||
server_auth_header, extra_headers = _prepare_mcp_server_headers(
|
||||
server=server,
|
||||
|
|
@ -786,24 +786,16 @@ 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)}"
|
||||
)
|
||||
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]
|
||||
# Continue with other servers instead of failing completely
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully fetched {len(all_tools)} tools total from all MCP servers"
|
||||
|
|
|
|||
|
|
@ -1908,9 +1908,6 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase):
|
|||
}
|
||||
|
||||
|
||||
UserMCPManagementMode = Literal["restricted", "view_all"]
|
||||
|
||||
|
||||
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
||||
"""
|
||||
Documents all the fields supported by `general_settings` in config.yaml
|
||||
|
|
@ -2028,10 +2025,6 @@ 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):
|
||||
|
|
|
|||
|
|
@ -323,7 +323,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"acreate_batch",
|
||||
"aretrieve_batch",
|
||||
"alist_batches",
|
||||
|
|
@ -485,7 +484,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"atext_completion",
|
||||
"aimage_edit",
|
||||
"alist_input_items",
|
||||
|
|
|
|||
|
|
@ -154,11 +154,7 @@ class PrismaManager:
|
|||
|
||||
prisma_dir = PrismaManager._get_prisma_dir()
|
||||
|
||||
from litellm.proxy.proxy_server import redis_usage_cache
|
||||
|
||||
return ProxyExtrasDBManager.setup_database(
|
||||
use_migrate=use_migrate, redis_cache=redis_usage_cache
|
||||
)
|
||||
return ProxyExtrasDBManager.setup_database(use_migrate=use_migrate)
|
||||
else:
|
||||
# Use prisma db push with increased timeout
|
||||
subprocess.run(
|
||||
|
|
|
|||
|
|
@ -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.ULID()) # type: ignore
|
||||
return str(ulid.new()) # type: ignore
|
||||
else:
|
||||
verbose_proxy_logger.debug("ULID library not available, using UUID")
|
||||
return str(uuid.uuid4())
|
||||
|
|
|
|||
|
|
@ -32,8 +32,8 @@ from fastapi import (
|
|||
from fastapi.responses import JSONResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
validate_and_normalize_mcp_server_payload,
|
||||
|
|
@ -67,6 +67,7 @@ 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,
|
||||
|
|
@ -75,10 +76,8 @@ 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
|
||||
|
|
@ -303,20 +302,6 @@ 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",
|
||||
|
|
@ -334,26 +319,18 @@ if MCP_AVAILABLE:
|
|||
```
|
||||
"""
|
||||
|
||||
user_mcp_management_mode = _get_user_mcp_management_mode()
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
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()
|
||||
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())
|
||||
|
||||
# augment the mcp servers with public status
|
||||
if litellm.public_mcp_servers is not None:
|
||||
|
|
@ -395,17 +372,6 @@ 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[
|
||||
|
|
|
|||
|
|
@ -698,88 +698,6 @@ 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)],
|
||||
|
|
|
|||
|
|
@ -25,7 +25,6 @@ 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",
|
||||
|
|
@ -117,7 +116,6 @@ async def route_request(
|
|||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"acreate_response_reply",
|
||||
"alist_input_items",
|
||||
"_arealtime", # private function for realtime API
|
||||
|
|
|
|||
|
|
@ -1361,205 +1361,3 @@ 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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
|
@ -12,9 +11,6 @@ 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
|
||||
|
|
@ -26,8 +22,7 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
|
||||
class BaseResponsesAPIStreamingIterator:
|
||||
|
|
@ -45,8 +40,6 @@ 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
|
||||
|
|
@ -54,25 +47,21 @@ class BaseResponsesAPIStreamingIterator:
|
|||
self.finished = False
|
||||
self.responses_api_provider_config = responses_api_provider_config
|
||||
self.completed_response: Optional[ResponsesAPIStreamingResponse] = None
|
||||
self.start_time = getattr(logging_obj, "start_time", datetime.now())
|
||||
self.start_time = datetime.now()
|
||||
|
||||
# track request context for hooks
|
||||
# set request kwargs
|
||||
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 the stream wrapper in litellm/litellm_core_utils/streaming_handler.py
|
||||
# This matches ths 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,
|
||||
|
|
@ -113,21 +102,13 @@ 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
|
||||
|
|
@ -168,159 +149,11 @@ 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):
|
||||
"""
|
||||
|
|
@ -335,8 +168,6 @@ 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,
|
||||
|
|
@ -345,8 +176,6 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj,
|
||||
litellm_metadata,
|
||||
custom_llm_provider,
|
||||
request_data,
|
||||
call_type,
|
||||
)
|
||||
self.stream_iterator = response.aiter_lines()
|
||||
|
||||
|
|
@ -374,21 +203,16 @@ 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,
|
||||
|
|
@ -405,7 +229,6 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
self._run_post_success_hooks(end_time=datetime.now())
|
||||
|
||||
|
||||
class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
||||
|
|
@ -421,8 +244,6 @@ 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,
|
||||
|
|
@ -431,8 +252,6 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj,
|
||||
litellm_metadata,
|
||||
custom_llm_provider,
|
||||
request_data,
|
||||
call_type,
|
||||
)
|
||||
self.stream_iterator = response.iter_lines()
|
||||
|
||||
|
|
@ -460,21 +279,16 @@ 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,
|
||||
|
|
@ -490,7 +304,6 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
self._run_post_success_hooks(end_time=datetime.now())
|
||||
|
||||
|
||||
class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
||||
|
|
@ -511,8 +324,6 @@ 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,
|
||||
|
|
@ -521,8 +332,6 @@ 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
|
||||
|
|
|
|||
|
|
@ -713,23 +713,6 @@ 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
|
||||
|
|
@ -829,9 +812,6 @@ 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"
|
||||
)
|
||||
|
|
@ -3944,7 +3924,6 @@ class Router:
|
|||
"anthropic_messages",
|
||||
"aresponses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"responses",
|
||||
"aget_responses",
|
||||
"adelete_responses",
|
||||
|
|
@ -4173,7 +4152,6 @@ class Router:
|
|||
elif call_type in (
|
||||
"aget_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"adelete_responses",
|
||||
"alist_input_items",
|
||||
):
|
||||
|
|
|
|||
|
|
@ -31,7 +31,6 @@ class LangsmithCredentialsObject(TypedDict):
|
|||
LANGSMITH_API_KEY: Optional[str]
|
||||
LANGSMITH_PROJECT: Optional[str]
|
||||
LANGSMITH_BASE_URL: str
|
||||
LANGSMITH_TENANT_ID: Optional[str]
|
||||
|
||||
|
||||
class LangsmithQueueObject(TypedDict):
|
||||
|
|
@ -53,7 +52,6 @@ class CredentialsKey(NamedTuple):
|
|||
api_key: str
|
||||
project: str
|
||||
base_url: str
|
||||
tenant_id: Optional[str]
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
|
|||
|
|
@ -2677,7 +2677,6 @@ 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]
|
||||
|
|
@ -2947,7 +2946,6 @@ class LlmProviders(str, Enum):
|
|||
MISTRAL = "mistral"
|
||||
MILVUS = "milvus"
|
||||
GROQ = "groq"
|
||||
GIGACHAT = "gigachat"
|
||||
NVIDIA_NIM = "nvidia_nim"
|
||||
CEREBRAS = "cerebras"
|
||||
AI21_CHAT = "ai21_chat"
|
||||
|
|
|
|||
|
|
@ -7521,8 +7521,6 @@ 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 (
|
||||
|
|
@ -7718,8 +7716,6 @@ 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,
|
||||
|
|
|
|||
|
|
@ -15831,68 +15831,6 @@
|
|||
"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",
|
||||
|
|
@ -32154,4 +32092,3 @@
|
|||
"mode": "chat"
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -28,8 +28,7 @@
|
|||
"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",
|
||||
"compact": "Supports /responses/compact endpoint"
|
||||
"delete_container_file": "Supports DELETE /containers/{id}/files/{file_id} endpoint"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
|
@ -1520,7 +1519,6 @@
|
|||
"retrieve_container_file": true,
|
||||
"retrieve_container_file_content": true,
|
||||
"delete_container_file": true,
|
||||
"compact": true,
|
||||
"a2a": true,
|
||||
"interactions": true
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.25.0 ; python_version >= "3.10" # for MCP server
|
||||
mcp==1.23.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
|
||||
|
|
|
|||
|
|
@ -61,19 +61,13 @@ async def test_bedrock_apply_guardrail_blocked():
|
|||
guardrailVersion="DRAFT",
|
||||
)
|
||||
|
||||
# Mock the make_bedrock_api_request method to raise an exception for blocked content
|
||||
# Mock the make_bedrock_api_request method
|
||||
with patch.object(
|
||||
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
|
||||
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
|
||||
) as mock_api_request:
|
||||
# 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": "",
|
||||
},
|
||||
)
|
||||
# Mock a blocked response from Bedrock
|
||||
mock_response = {"action": "BLOCKED", "reason": "Content violates policy"}
|
||||
mock_api_request.return_value = mock_response
|
||||
|
||||
# Test the apply_guardrail method should raise an exception
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
|
|
@ -83,9 +77,8 @@ async def test_bedrock_apply_guardrail_blocked():
|
|||
input_type="request",
|
||||
)
|
||||
|
||||
# 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)
|
||||
assert "Content blocked by Bedrock guardrail" in str(exc_info.value)
|
||||
assert "Content violates policy" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -260,15 +253,7 @@ 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 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",
|
||||
},
|
||||
)
|
||||
mock_api.return_value = {"action": "BLOCKED", "reason": "policy"}
|
||||
|
||||
with pytest.raises(Exception, match="policy") as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
|
|
@ -280,8 +265,7 @@ 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]]
|
||||
# The apply_guardrail method wraps the original exception in a generic Exception
|
||||
assert "Bedrock guardrail failed:" in str(exc_info.value)
|
||||
assert "Content blocked by Bedrock guardrail" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_filters_latest_user_message_when_enabled():
|
||||
|
|
|
|||
|
|
@ -1,14 +1,11 @@
|
|||
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, MigrationLockManager
|
||||
|
||||
from litellm_proxy_extras.utils import ProxyExtrasDBManager
|
||||
|
||||
|
||||
def test_custom_prisma_dir(monkeypatch):
|
||||
|
|
@ -30,279 +27,101 @@ def test_custom_prisma_dir(monkeypatch):
|
|||
assert os.path.exists(migrations_dir)
|
||||
|
||||
|
||||
class TestMigrationLockManager:
|
||||
"""Test cases for MigrationLockManager"""
|
||||
class TestPermissionErrorDetection:
|
||||
"""Test cases for permission error detection in Prisma migrations"""
|
||||
|
||||
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_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_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_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
|
||||
|
||||
result = lock_manager.acquire_lock()
|
||||
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
|
||||
|
||||
assert result is True
|
||||
assert lock_manager.lock_acquired is True
|
||||
mock_redis.set_cache.assert_called_once()
|
||||
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
|
||||
|
||||
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_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
|
||||
|
||||
result = lock_manager.acquire_lock()
|
||||
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
|
||||
|
||||
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()
|
||||
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
|
||||
|
||||
|
||||
class TestProxyExtrasDBManagerMigrationLock:
|
||||
"""Test cases for ProxyExtrasDBManager with migration locking"""
|
||||
class TestIdempotentErrorDetection:
|
||||
"""Test cases for idempotent error detection in Prisma migrations"""
|
||||
|
||||
@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_already_exists(self):
|
||||
"""Test detection of generic 'already exists' error"""
|
||||
error_message = "object 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_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
|
||||
|
||||
# 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_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
|
||||
|
||||
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_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
|
||||
|
||||
@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_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
|
||||
|
||||
# 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_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
|
||||
|
||||
# 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_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
|
||||
|
||||
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()
|
||||
|
||||
@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
|
||||
class TestErrorClassificationPriority:
|
||||
"""Test cases to ensure errors are correctly classified"""
|
||||
|
||||
# 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_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
|
||||
|
||||
assert result is True
|
||||
# Redis가 없을 때는 락 보호 없이 마이그레이션을 실행해야 함
|
||||
mock_subprocess.assert_called_once()
|
||||
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
|
||||
|
||||
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
|
||||
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
|
||||
|
|
|
|||
|
|
@ -1814,49 +1814,3 @@ 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"
|
||||
|
|
|
|||
|
|
@ -1,165 +0,0 @@
|
|||
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
|
||||
|
|
@ -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, False),
|
||||
(False, True, True),
|
||||
(True, True, True),
|
||||
(False, False, False),
|
||||
],
|
||||
|
|
|
|||
|
|
@ -15,7 +15,6 @@ 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:
|
||||
|
|
@ -726,7 +725,6 @@ 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(
|
||||
{
|
||||
|
|
@ -769,7 +767,6 @@ 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(
|
||||
{
|
||||
|
|
@ -826,7 +823,6 @@ 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(
|
||||
{
|
||||
|
|
@ -899,7 +895,6 @@ 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(
|
||||
{
|
||||
|
|
@ -928,7 +923,6 @@ 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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -1,349 +0,0 @@
|
|||
"""
|
||||
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
|
||||
|
|
@ -286,7 +286,7 @@ def test_completion_claude_3_empty_response():
|
|||
},
|
||||
]
|
||||
try:
|
||||
response = litellm.completion(model="claude-3-7-sonnet-20250219", messages=messages)
|
||||
response = litellm.completion(model="claude-3-opus-20240229", 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-7-sonnet-20250219",
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
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-7-sonnet-20250219", "anthropic.claude-3-sonnet-20240229-v1:0"],
|
||||
["anthropic/claude-3-opus-20240229", "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-7-sonnet-20250219", None, None),
|
||||
("claude-3-opus-20240229", 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-7-sonnet-20250219",
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
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-7-sonnet-20250219",
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
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-7-sonnet-20250219",
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
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-7-sonnet-20250219",
|
||||
"anthropic/claude-3-opus-20240229",
|
||||
],
|
||||
) #
|
||||
def test_completion_base64(model):
|
||||
|
|
|
|||
|
|
@ -1418,7 +1418,7 @@ def test_bedrock_claude_3_streaming():
|
|||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"claude-3-7-sonnet-20250219",
|
||||
"claude-3-opus-20240229",
|
||||
"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-7-sonnet-20250219",
|
||||
model="claude-3-opus-20240229",
|
||||
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):
|
||||
|
|
|
|||
|
|
@ -47,19 +47,6 @@ 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():
|
||||
|
|
@ -73,7 +60,6 @@ async def test_group_batches_by_credentials():
|
|||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -83,7 +69,6 @@ async def test_group_batches_by_credentials():
|
|||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -110,7 +95,6 @@ async def test_group_batches_by_credentials_multiple_credentials():
|
|||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -120,7 +104,6 @@ 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,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -130,7 +113,6 @@ 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,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -145,57 +127,6 @@ 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():
|
||||
|
|
@ -270,43 +201,10 @@ 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_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():
|
||||
async def test_langsmith_key_based_logging(mocker):
|
||||
"""
|
||||
In key based logging langsmith_api_key and langsmith_project are passed directly to litellm.acompletion
|
||||
"""
|
||||
|
|
@ -321,11 +219,10 @@ async def test_langsmith_key_based_logging():
|
|||
mock_response.text = ""
|
||||
mock_async_httpx_handler.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
mock_get_client = patch(
|
||||
mock_get_client = mocker.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
|
||||
|
|
@ -356,8 +253,6 @@ async def test_langsmith_key_based_logging():
|
|||
|
||||
# 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"]
|
||||
|
|
@ -449,8 +344,6 @@ async def test_langsmith_key_based_logging():
|
|||
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}")
|
||||
|
|
|
|||
|
|
@ -65,6 +65,31 @@ 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.
|
||||
|
|
|
|||
|
|
@ -79,73 +79,6 @@ 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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1011,87 +1011,39 @@ def test_multiple_tool_calls_in_single_choice():
|
|||
|
||||
def test_map_reasoning_effort_adds_summary_detailed():
|
||||
"""
|
||||
Test that _map_reasoning_effort behavior with reasoning_auto_summary flag.
|
||||
Test that _map_reasoning_effort adds summary="detailed" when user provides reasoning_effort as a string.
|
||||
|
||||
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.
|
||||
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.
|
||||
"""
|
||||
import os
|
||||
|
||||
import litellm
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
# Test all string effort levels - DEFAULT BEHAVIOR (no summary)
|
||||
# Test all string effort levels
|
||||
effort_levels = ["none", "low", "medium", "high", "xhigh", "minimal"]
|
||||
|
||||
# Save original flag value
|
||||
original_flag = litellm.reasoning_auto_summary
|
||||
original_env = os.environ.get("LITELLM_REASONING_AUTO_SUMMARY")
|
||||
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' for effort={effort}"
|
||||
|
||||
print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed'")
|
||||
|
||||
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"]
|
||||
|
||||
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)")
|
||||
|
||||
# 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")
|
||||
|
||||
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"]
|
||||
# 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'")
|
||||
|
|
|
|||
|
|
@ -1,90 +0,0 @@
|
|||
"""
|
||||
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
|
||||
|
|
@ -172,86 +172,6 @@ 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
|
||||
|
|
@ -700,6 +620,7 @@ 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)
|
||||
|
|
@ -766,8 +687,11 @@ 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()
|
||||
|
|
@ -1396,23 +1320,6 @@ 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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -691,29 +691,6 @@ 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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -1,355 +0,0 @@
|
|||
"""
|
||||
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"])
|
||||
|
||||
|
|
@ -10,13 +10,11 @@ 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():
|
||||
|
|
@ -1607,39 +1605,6 @@ 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.
|
||||
|
|
|
|||
|
|
@ -354,7 +354,7 @@ async def test_register_client_remote_registration_success():
|
|||
|
||||
request_payload = {
|
||||
"client_name": "Litellm Proxy",
|
||||
"grant_types": ["authorization_code", "refresh_token"],
|
||||
"grant_types": ["authorization_code"],
|
||||
"response_types": ["code"],
|
||||
"token_endpoint_auth_method": "client_secret_post",
|
||||
}
|
||||
|
|
@ -556,33 +556,9 @@ 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)
|
||||
|
|
@ -592,14 +568,13 @@ 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_oauth",
|
||||
mcp_server_name="test_server",
|
||||
)
|
||||
|
||||
# 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
|
||||
|
|
@ -609,33 +584,9 @@ 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)
|
||||
|
|
@ -645,15 +596,13 @@ 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_oauth",
|
||||
mcp_server_name="test_server",
|
||||
)
|
||||
|
||||
# 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
|
||||
|
|
|
|||
|
|
@ -594,26 +594,7 @@ 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"""
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
Unit tests for Qualifire guardrail integration.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -140,98 +139,76 @@ class TestQualifireGuardrailEvaluateKwargs:
|
|||
@pytest.mark.asyncio
|
||||
async def test_evaluate_called_with_prompt_injections(self):
|
||||
"""Test that evaluate is called with prompt_injections enabled."""
|
||||
# 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,
|
||||
)
|
||||
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."""
|
||||
# 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,
|
||||
)
|
||||
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:
|
||||
|
|
|
|||
|
|
@ -40,7 +40,6 @@ 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"}
|
||||
|
|
@ -60,10 +59,6 @@ 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(
|
||||
|
|
@ -101,7 +96,6 @@ 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"}
|
||||
|
|
@ -121,10 +115,6 @@ 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(
|
||||
|
|
|
|||
|
|
@ -231,40 +231,6 @@ 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):
|
||||
"""
|
||||
|
|
@ -1130,51 +1096,6 @@ 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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -742,16 +742,18 @@ 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
|
||||
|
||||
# Override the FastAPI dependency with a proper mock
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
user_id="test-user-123",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
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,
|
||||
)
|
||||
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()
|
||||
|
|
@ -759,11 +761,7 @@ class TestProxySettingEndpoints:
|
|||
|
||||
payload = {"disable_model_add_for_internal_users": True}
|
||||
|
||||
try:
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
finally:
|
||||
# Clean up the dependency override
|
||||
app.dependency_overrides.clear()
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
|
|
@ -782,16 +780,18 @@ 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
|
||||
|
||||
# Override the FastAPI dependency with a proper mock
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
user_id="test-user-123",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
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,
|
||||
)
|
||||
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,11 +802,7 @@ class TestProxySettingEndpoints:
|
|||
"unsupported_flag": True,
|
||||
}
|
||||
|
||||
try:
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
finally:
|
||||
# Clean up the dependency override
|
||||
app.dependency_overrides.clear()
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
|
|
|
|||
|
|
@ -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,15 +27,13 @@ export interface SSOSettingsValues {
|
|||
proxy_base_url: string | null;
|
||||
user_email: string | null;
|
||||
ui_access_mode: string | null;
|
||||
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[];
|
||||
role_mappings: {
|
||||
provider: string;
|
||||
group_claim: string;
|
||||
default_role: "internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer";
|
||||
roles: {
|
||||
[key: string]: string[];
|
||||
};
|
||||
};
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -18,7 +18,6 @@ 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";
|
||||
|
|
@ -275,30 +274,6 @@ 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}
|
||||
|
|
|
|||
|
|
@ -189,16 +189,11 @@ const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormS
|
|||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.use_role_mappings !== currentValues.use_role_mappings ||
|
||||
prevValues.sso_provider !== currentValues.sso_provider
|
||||
}
|
||||
shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings}
|
||||
>
|
||||
{({ getFieldValue }) => {
|
||||
const useRoleMappings = getFieldValue("use_role_mappings");
|
||||
const provider = getFieldValue("sso_provider");
|
||||
const supportsRoleMappings = provider === "okta" || provider === "generic";
|
||||
return useRoleMappings && supportsRoleMappings ? (
|
||||
return useRoleMappings ? (
|
||||
<Form.Item
|
||||
label="Group Claim"
|
||||
name="group_claim"
|
||||
|
|
@ -212,16 +207,11 @@ const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormS
|
|||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.use_role_mappings !== currentValues.use_role_mappings ||
|
||||
prevValues.sso_provider !== currentValues.sso_provider
|
||||
}
|
||||
shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings}
|
||||
>
|
||||
{({ getFieldValue }) => {
|
||||
const useRoleMappings = getFieldValue("use_role_mappings");
|
||||
const provider = getFieldValue("sso_provider");
|
||||
const supportsRoleMappings = provider === "okta" || provider === "generic";
|
||||
return useRoleMappings && supportsRoleMappings ? (
|
||||
return useRoleMappings ? (
|
||||
<>
|
||||
<Form.Item label="Default Role" name="default_role" initialValue="Internal User">
|
||||
<Select>
|
||||
|
|
|
|||
|
|
@ -1,60 +1,20 @@
|
|||
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(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<DeleteSSOSettingsModal isVisible={true} onCancel={onCancel} onSuccess={onSuccess} />
|
||||
</QueryClientProvider>,
|
||||
<DeleteSSOSettingsModal isVisible={true} onCancel={onCancel} onSuccess={onSuccess} accessToken="test-token" />,
|
||||
);
|
||||
|
||||
expect(screen.getByText("Confirm Clear SSO Settings")).toBeInTheDocument();
|
||||
expect(
|
||||
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.",
|
||||
),
|
||||
screen.getByText("Are you sure you want to clear all SSO settings? This action cannot be undone."),
|
||||
).toBeInTheDocument();
|
||||
expect(screen.getByText("Users will no longer be able to login using SSO after this change.")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,66 +1,79 @@
|
|||
import { useEditSSOSettings } from "@/app/(dashboard)/hooks/sso/useEditSSOSettings";
|
||||
import { useSSOSettings } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
import { Modal } from "antd";
|
||||
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 }) => {
|
||||
const { data: ssoSettings } = useSSOSettings();
|
||||
const { mutateAsync: editSSOSettings, isPending: isEditingSSOSettings } = useEditSSOSettings();
|
||||
|
||||
const DeleteSSOSettingsModal: React.FC<DeleteSSOSettingsModalProps> = ({
|
||||
isVisible,
|
||||
onCancel,
|
||||
onSuccess,
|
||||
accessToken,
|
||||
}) => {
|
||||
// Handle clearing SSO settings
|
||||
const handleClearSSO = async () => {
|
||||
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,
|
||||
};
|
||||
if (!accessToken) {
|
||||
NotificationsManager.fromBackend("No access token available");
|
||||
return;
|
||||
}
|
||||
|
||||
await editSSOSettings(clearSettings, {
|
||||
onSuccess: () => {
|
||||
NotificationsManager.success("SSO settings cleared successfully");
|
||||
onCancel();
|
||||
onSuccess();
|
||||
},
|
||||
onError: (error) => {
|
||||
NotificationsManager.fromBackend("Failed to clear SSO settings: " + parseErrorMessage(error));
|
||||
},
|
||||
});
|
||||
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));
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<DeleteResourceModal
|
||||
isOpen={isVisible}
|
||||
<Modal
|
||||
title="Confirm Clear SSO Settings"
|
||||
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}
|
||||
visible={isVisible}
|
||||
onOk={handleClearSSO}
|
||||
confirmLoading={isEditingSSOSettings}
|
||||
/>
|
||||
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>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -1,92 +0,0 @@
|
|||
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();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,74 +0,0 @@
|
|||
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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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 { detectSSOProvider } from "./utils";
|
||||
import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./constants";
|
||||
|
||||
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,8 +26,24 @@ 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}>
|
||||
|
|
@ -169,52 +185,46 @@ export default function SSOSettings() {
|
|||
{isLoading ? (
|
||||
<SSOSettingsLoadingSkeleton />
|
||||
) : (
|
||||
<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>
|
||||
</>
|
||||
)}
|
||||
<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>
|
||||
|
||||
{isSSOConfigured ? (
|
||||
renderSSOSettings()
|
||||
) : (
|
||||
<SSOSettingsEmptyPlaceholder onAdd={() => setIsAddModalVisible(true)} />
|
||||
)}
|
||||
</Space>
|
||||
</Card>
|
||||
{isRoleMappingsEnabled && <RoleMappings roleMappings={ssoSettings?.values.role_mappings} />}
|
||||
</Space>
|
||||
<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>
|
||||
)}
|
||||
|
||||
<DeleteSSOSettingsModal
|
||||
isVisible={isDeleteModalVisible}
|
||||
onCancel={() => setIsDeleteModalVisible(false)}
|
||||
onSuccess={() => refetch()}
|
||||
accessToken={accessToken}
|
||||
/>
|
||||
|
||||
<AddSSOSettingsModal
|
||||
|
|
|
|||
|
|
@ -13,10 +13,3 @@ 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",
|
||||
};
|
||||
|
|
|
|||
|
|
@ -55,7 +55,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
other_field: "value",
|
||||
};
|
||||
|
||||
|
|
@ -84,7 +83,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -102,7 +100,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -124,7 +121,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -146,7 +142,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -165,7 +160,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -180,7 +174,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -193,7 +186,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -206,7 +198,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -219,7 +210,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -232,7 +222,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "unknown_role",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -244,7 +233,6 @@ describe("processSSOSettingsPayload", () => {
|
|||
const formValues = {
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
|
|||
|
|
@ -1,5 +1,3 @@
|
|||
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
|
||||
|
|
@ -20,7 +18,7 @@ export const processSSOSettingsPayload = (formValues: Record<string, any>): Reco
|
|||
...rest,
|
||||
};
|
||||
|
||||
// Add role mappings only if use_role_mappings is checked AND provider supports role mappings
|
||||
// Add role mappings if use_role_mappings is checked
|
||||
if (use_role_mappings) {
|
||||
// Helper function to split comma-separated string into array
|
||||
const splitTeams = (teams: string | undefined): string[] => {
|
||||
|
|
@ -54,20 +52,3 @@ 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;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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", "refresh_token"],
|
||||
grant_types: ["authorization_code"],
|
||||
response_types: ["code"],
|
||||
token_endpoint_auth_method:
|
||||
temporaryPayload.credentials && temporaryPayload.credentials.client_secret ? "client_secret_post" : "none",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue