mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat: Add built-in migration lock to prevent concurrent Prisma migrate deploy (#14440)
* feat: prisma migrate deploy with lock
Author: Mini Jeong <mini.jeong@navercorp.com>
* fix: use redis cache from proxy server
Author: Mini Jeong <mini.jeong@navercorp.com>
* fix: add type checks and fix unit tests for migration lock
- Add DATABASE_URL validation in _create_baseline_migration() and _resolve_all_migrations()
- Fix MyPy type errors by adding None checks before using database_url in subprocess calls
- Add _resolve_all_migrations mock to failing unit tests to prevent filesystem errors
- Apply Black formatting to modified files
Fixes:
- MyPy type errors: database_url could be None when passed to subprocess
- Unit test failures: _resolve_all_migrations tried to create directories in read-only /test path
* fix: resolve MyPy type error in vertex_ai vertex_llm_base
Fix MyPy type checking error where vertex_api_version parameter type
was incompatible with function signature expectation.
* fix: Return 403 exception when calling GET responses api
* fix: added new step into rotate master key function for processing credentials table
* Add redisvl in requirements.txt
* fix: fixed the issue of handling root paths when processing Discovery protected resource metadata and authorization server metadata URLs.
* fix: added additional grant type into oauth_authorization_server response for fixing mcp auth register bad request issue
* fix: added RFC RECOMMENDED property(scopes_supported) to protected resource and authorization server metadata
* fix: removed initialize the tool name to MCP server name mapping(oauth2) on startup for avoiding 401 error
* fix: upgraded mcp sdk depency version for fixing ClosedResourceError
* Use already configured opentelemetry providers
Users that instrument using opentelemetry-instrument can now setup exporters as per their environment.
* Handle all protocols for all telemetry
* Add more tests
* feat(mcp): parallelize tool fetching from multiple MCP servers (#18627)
* feat(mcp): parallelize tool fetching from multiple MCP servers
Replace sequential tool fetching with asyncio.gather() to reduce
client timeouts when using multiple MCP servers.
Changes:
- mcp_server_manager.py: list_tools() now fetches tools in parallel
- server.py: _get_tools_from_mcp_servers() now fetches tools in parallel
Real-world impact (7 MCP servers example):
- Sequential: ~4.5+ seconds (exceeds typical 5-second client timeouts)
- Parallel: ~1.2 seconds (max of all servers)
Fixes #18626
* fix: copy oauth2_headers to avoid shared dict mutation in parallel tasks
* feat: add display_name, model_vendor, and model_version metadata
* added the option of adding langsmith tenant id in the env (#18623)
* fix(router): Validate routing_strategy at startup to fail fast with helpful error. (#18624)
Invalid routing_strategy values (e.g., "simple" instead of "simple-shuffle") previously failed silently, causing confusing "No deployments available" errors downstream. This change adds upfront validation in routing_strategy_init() to:
- Check if the provided strategy matches valid string values or RoutingStrategy enum
- Raise a clear ValueError listing valid options if invalid
- Fail fast at startup instead of at request time
Fixes behavior reported in #11330 where users had to debug cryptic errors.
Valid strategies: simple-shuffle, least-busy, usage-based-routing, latency-based-routing, cost-based-routing, usage-based-routing-v2
Co-authored-by: Flibbert E. Gibbitz <flibbertygibbitz@runelabs.ai>
* Add libsndfile to database Docker image for audio processing (#18612)
The litellm-database Docker image was missing the libsndfile system
library, which is required by the soundfile Python package for audio
file processing. This caused failures when using audio transcription
endpoints that attempt to calculate audio duration.
This adds libsndfile to the runtime dependencies in Dockerfile.database,
consistent with Dockerfile.alpine which already includes this library.
* Fix: Map Gemini cached_tokens to Langfuse cache_read_input_tokens (#18614)
* Fix: Map Gemini cached_tokens to Langfuse cache_read_input_tokens
Fixes #18520
## Problem
Langfuse integration was not capturing cached tokens from Gemini models.
Gemini returns cached tokens in `usage.prompt_tokens_details.cached_tokens`,
but Langfuse only read from top-level `usage.cache_read_input_tokens`
(which only Anthropic populates).
## Solution
Updated langfuse.py to check both locations:
1. First check top-level cache_read_input_tokens (for Anthropic)
2. Then check prompt_tokens_details.cached_tokens (for Gemini, OpenAI, others)
This ensures all providers' cached tokens are properly reported to Langfuse.
## Changes
- Modified litellm/integrations/langfuse/langfuse.py (lines 742-761)
- Added 3 unit tests in tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py
- All existing Langfuse tests still pass (11/11)
## Testing
- test_cached_tokens_extraction: Verifies Gemini cached_tokens extraction
- test_cached_tokens_not_present: Backward compatibility (no cached_tokens)
- test_cached_tokens_is_zero: Edge case when cached_tokens = 0
* Refactor: Extract cache token logic into helper function
Address review feedback from @officer47p
- Created _extract_cache_read_input_tokens() helper function
- Reduces code bloat in _log_langfuse_v2 method
- Improves testability and reusability
- All tests still passing (11/11)
* Adding Role Mappings
* Fixing Edit SSO Settings Modal
* feat: add user_mcp_management_mode for view_all visibility
* Fixing tests
* fix: missing mcp_allow_all_ui.png
* docs: add user_mcp_management_mode
* Align responses API streaming hooks with chat pipeline
* Clarify responses API streaming context
* Address review comments
* feat: Add GigaChat provider support (#18564)
* feat: Add GigaChat provider support
Add native support for GigaChat API (Sber AI, Russia's leading LLM).
Supported features:
- Chat completions (sync/async)
- Streaming (sync/async)
- Function calling / Tools
- Structured output via JSON schema (emulated through function calls)
- Image input (base64 and URL)
- Embeddings
Closes #18515
* fix: resolve mypy type errors in GigaChat handler
- Fix _prepare_file_data return type (use 3-tuple for cleaner type flow)
- Add type annotations for lists in _process_content_parts methods
- Add type annotations in _collapse_user_messages
- Use ChatCompletionToolCallChunk for proper tool_use typing
- Add type: ignore[override] for astreaming async generator
* refactor(gigachat): migrate to BaseConfig pattern
* fix: remove unused imports
* fix: resolve mypy type errors
* fix: mypy type errors
* refactor: address review feedback for GigaChat provider
- Remove singleton pattern, reuse litellm HTTPHandler
- Move constants/errors to transformation files, delete common_utils.py
- Add models to model_prices_and_context_window.json
- Fix ssl_verify not passed to HTTP client for embeddings
* docs: update GigaChat documentation with ssl_verify requirement
* Revert "Add redisvl in requirements.txt"
* Put reasoning summary behind feat flag
* fix: model eol
* fix: anthropic claude-3-opus-20240229 EOL
* Revert "fix: model eol"
This reverts commit 5aa1665d79.
* Fix: ImportError: qualifire package is required for QualifireGuardrail. Install it with: pip install qualifire
* fix: test_secret_manager_failure_does_not_block_email
* fix: test_update_ui_settings_allowlisted_value
* fix: test_aaamodel_prices_and_context_window_json_is_valid
* fix: test_all_models_have_display_name
* fix: async def test_bedrock_apply_guardrail_blocked()
* fix: test_databricks_embeddings[True]
* fix:test_anthropic_beta_header
* fix:test_api_error_handling
* fix:mypy mcp management
* Revert "feat(model_cost): add display_name, model_vendor, and model_version metadata to model entries"
* [Feat] New API Endpoint - Responses API (v1/responses/compact) (#18697)
* init transform_compact_response_api_request
* init acompact_responses
* init async_compact_response_api_handler in llm http handler
* init transform_compact_response_api_request for openai
* init acompact_responses
* fix acompact_responses
* add OAI Compact API
* docs responses API Compact
* code qa checks
* test_openai_compact_responses_api
* fix mypy linting
* fix: remove display name
* Add the LITELLM_REASONING_AUTO_SUMMARY in doc
* fix model map
* [UI] - Feat add request provider form on UI (#18704)
* add request provider form
* fix link to github
* add button
* fix link
* fix(streaming): normalize status code extraction to prevent 4xx errors from triggering mid-stream fallback (#18698)
在流式处理错误时,添加状态码标准化逻辑,确保 4xx 客户端错误直接抛出而不是被包装成 MidStreamFallbackError。
- 新增 _normalize_status_code 函数用于从异常对象提取状态码
- 优先从异常的 status_code 属性获取,其次从 response.status_code 获取
- 当映射异常或原始异常的状态码在 400-499 范围内时,直接抛出映射异常
- 添加单元测试验证 Vertex AI 400 错误正确抛出为 BadRequestError
- 确保流式处理中的客户端错误能够正确传播,而不会触发回退机制
---------
Co-authored-by: Eric84626 <lixiannan@gmail.com>
Co-authored-by: Eric84626 <97266539+Eric84626@users.noreply.github.com>
Co-authored-by: Sameer Kankute <sameer@berri.ai>
Co-authored-by: mangabits <1457532+mangabits@users.noreply.github.com>
Co-authored-by: Costa Tsaousis <costa@tsaousis.gr>
Co-authored-by: Nik <nikolas.garza5@gmail.com>
Co-authored-by: Shivam Rawat <161387515+shivamrawat1@users.noreply.github.com>
Co-authored-by: FlibbertyGibbitz <seth@evenkeelconsultingllc.com>
Co-authored-by: Flibbert E. Gibbitz <flibbertygibbitz@runelabs.ai>
Co-authored-by: Cesar Garcia <128240629+Chesars@users.noreply.github.com>
Co-authored-by: yuneng-jiang <yuneng.jiang@gmail.com>
Co-authored-by: Yuta Saito <uc4w6c@bma.biglobe.ne.jp>
Co-authored-by: LingXuanYin <3546599908@qq.com>
Co-authored-by: YutaSaito <36355491+uc4w6c@users.noreply.github.com>
Co-authored-by: 0717376 <103773680+0717376@users.noreply.github.com>
Co-authored-by: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Co-authored-by: Kris Xia <xiajiayi0506@gmail.com>
Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
parent
9e6714fe1b
commit
9f68081f6d
87 changed files with 5826 additions and 566 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
|
||||
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip libsndfile
|
||||
|
||||
WORKDIR /app
|
||||
# Copy the current directory contents into the container at /app
|
||||
|
|
|
|||
|
|
@ -108,7 +108,7 @@ Some MCP servers are meant to be shared broadly—think internal knowledge bases
|
|||
3. Toggle **Allow All LiteLLM Keys** on.
|
||||
|
||||
<Image
|
||||
img={require('../img/mcp_ui.png')}
|
||||
img={require('../img/mcp_allow_all_ui.png')}
|
||||
style={{width: '80%', display: 'block', margin: '1rem auto'}}
|
||||
alt="MCP server configuration in Admin UI"
|
||||
/>
|
||||
|
|
@ -634,3 +634,18 @@ Control which tools different teams can access from the same MCP server. For exa
|
|||
This video shows how to set allowed tools for a Key, Team, or Organization.
|
||||
|
||||
<iframe width="840" height="500" src="https://www.loom.com/embed/7464d444c3324078892367272fe50745" frameborder="0" webkitallowfullscreen mozallowfullscreen allowfullscreen></iframe>
|
||||
|
||||
|
||||
## Dashboard View Modes
|
||||
|
||||
Proxy admins can also control what non-admins see inside the MCP dashboard via `general_settings.user_mcp_management_mode`:
|
||||
|
||||
- `restricted` *(default)* – users only see servers that their team explicitly has access to.
|
||||
- `view_all` – every dashboard user can see the full MCP server list.
|
||||
|
||||
```yaml title="Config example"
|
||||
general_settings:
|
||||
user_mcp_management_mode: view_all
|
||||
```
|
||||
|
||||
This is useful when you want discoverability for MCP offerings without granting additional execution privileges.
|
||||
|
|
|
|||
283
docs/my-website/docs/providers/gigachat.md
Normal file
283
docs/my-website/docs/providers/gigachat.md
Normal file
|
|
@ -0,0 +1,283 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# GigaChat
|
||||
https://developers.sber.ru/docs/ru/gigachat/api/overview
|
||||
|
||||
GigaChat is Sber AI's large language model, Russia's leading LLM provider.
|
||||
|
||||
:::tip
|
||||
|
||||
**We support ALL GigaChat models, just set `model=gigachat/<any-model-on-gigachat>` as a prefix when sending litellm requests**
|
||||
|
||||
:::
|
||||
|
||||
:::warning
|
||||
|
||||
GigaChat API uses self-signed SSL certificates. You must pass `ssl_verify=False` in your requests.
|
||||
|
||||
:::
|
||||
|
||||
## Supported Features
|
||||
|
||||
| Feature | Supported |
|
||||
|---------|-----------|
|
||||
| Chat Completion | Yes |
|
||||
| Streaming | Yes |
|
||||
| Async | Yes |
|
||||
| Function Calling / Tools | Yes |
|
||||
| Structured Output (JSON Schema) | Yes (via function call emulation) |
|
||||
| Image Input | Yes (base64 and URL) - GigaChat-2-Max, GigaChat-2-Pro only |
|
||||
| Embeddings | Yes |
|
||||
|
||||
## API Key
|
||||
|
||||
GigaChat uses OAuth authentication. Set your credentials as environment variables:
|
||||
|
||||
```python
|
||||
import os
|
||||
|
||||
# Required: Set credentials (base64-encoded client_id:client_secret)
|
||||
os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
|
||||
|
||||
# Optional: Set scope (default is GIGACHAT_API_PERS for personal use)
|
||||
os.environ['GIGACHAT_SCOPE'] = "GIGACHAT_API_PERS" # or GIGACHAT_API_B2B for business
|
||||
```
|
||||
|
||||
Get your credentials at: https://developers.sber.ru/studio/
|
||||
|
||||
## Sample Usage
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
|
||||
|
||||
response = completion(
|
||||
model="gigachat/GigaChat-2-Max",
|
||||
messages=[
|
||||
{"role": "user", "content": "Hello from LiteLLM!"}
|
||||
],
|
||||
ssl_verify=False, # Required for GigaChat
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
## Sample Usage - Streaming
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
|
||||
|
||||
response = completion(
|
||||
model="gigachat/GigaChat-2-Max",
|
||||
messages=[
|
||||
{"role": "user", "content": "Hello from LiteLLM!"}
|
||||
],
|
||||
stream=True,
|
||||
ssl_verify=False, # Required for GigaChat
|
||||
)
|
||||
|
||||
for chunk in response:
|
||||
print(chunk)
|
||||
```
|
||||
|
||||
## Sample Usage - Function Calling
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
|
||||
|
||||
tools = [{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a city",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": {"type": "string", "description": "City name"}
|
||||
},
|
||||
"required": ["city"]
|
||||
}
|
||||
}
|
||||
}]
|
||||
|
||||
response = completion(
|
||||
model="gigachat/GigaChat-2-Max",
|
||||
messages=[{"role": "user", "content": "What's the weather in Moscow?"}],
|
||||
tools=tools,
|
||||
ssl_verify=False, # Required for GigaChat
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
## Sample Usage - Structured Output
|
||||
|
||||
GigaChat supports structured output via JSON schema (emulated through function calling):
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
|
||||
|
||||
response = completion(
|
||||
model="gigachat/GigaChat-2-Max",
|
||||
messages=[{"role": "user", "content": "Extract info: John is 30 years old"}],
|
||||
response_format={
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "person",
|
||||
"schema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"age": {"type": "integer"}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
ssl_verify=False, # Required for GigaChat
|
||||
)
|
||||
print(response) # Returns JSON: {"name": "John", "age": 30}
|
||||
```
|
||||
|
||||
## Sample Usage - Image Input
|
||||
|
||||
GigaChat supports image input via base64 or URL (GigaChat-2-Max and GigaChat-2-Pro only):
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
|
||||
|
||||
response = completion(
|
||||
model="gigachat/GigaChat-2-Max", # Vision requires GigaChat-2-Max or GigaChat-2-Pro
|
||||
messages=[{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What's in this image?"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}
|
||||
]
|
||||
}],
|
||||
ssl_verify=False, # Required for GigaChat
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
## Sample Usage - Embeddings
|
||||
|
||||
```python
|
||||
from litellm import embedding
|
||||
import os
|
||||
|
||||
os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here"
|
||||
|
||||
response = embedding(
|
||||
model="gigachat/Embeddings",
|
||||
input=["Hello world", "How are you?"],
|
||||
ssl_verify=False, # Required for GigaChat
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
## Usage with LiteLLM Proxy
|
||||
|
||||
### 1. Set GigaChat Models on config.yaml
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gigachat
|
||||
litellm_params:
|
||||
model: gigachat/GigaChat-2-Max
|
||||
api_key: "os.environ/GIGACHAT_CREDENTIALS"
|
||||
ssl_verify: false
|
||||
- model_name: gigachat-lite
|
||||
litellm_params:
|
||||
model: gigachat/GigaChat-2-Lite
|
||||
api_key: "os.environ/GIGACHAT_CREDENTIALS"
|
||||
ssl_verify: false
|
||||
- model_name: gigachat-embeddings
|
||||
litellm_params:
|
||||
model: gigachat/Embeddings
|
||||
api_key: "os.environ/GIGACHAT_CREDENTIALS"
|
||||
ssl_verify: false
|
||||
```
|
||||
|
||||
### 2. Start Proxy
|
||||
|
||||
```bash
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
### 3. Test it
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="Curl" label="Curl Request">
|
||||
|
||||
```shell
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "gigachat",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello!"
|
||||
}
|
||||
]
|
||||
}'
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="openai" label="OpenAI v1.0.0+">
|
||||
|
||||
```python
|
||||
import openai
|
||||
client = openai.OpenAI(
|
||||
api_key="anything",
|
||||
base_url="http://0.0.0.0:4000"
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="gigachat",
|
||||
messages=[{"role": "user", "content": "Hello!"}]
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Supported Models
|
||||
|
||||
### Chat Models
|
||||
|
||||
| Model Name | Context Window | Vision | Description |
|
||||
|------------|----------------|--------|-------------|
|
||||
| gigachat/GigaChat-2-Lite | 128K | No | Fast, lightweight model |
|
||||
| gigachat/GigaChat-2-Pro | 128K | Yes | Professional model with vision |
|
||||
| gigachat/GigaChat-2-Max | 128K | Yes | Maximum capability model |
|
||||
|
||||
### Embedding Models
|
||||
|
||||
| Model Name | Max Input | Dimensions | Description |
|
||||
|------------|-----------|------------|-------------|
|
||||
| gigachat/Embeddings | 512 | 1024 | Standard embeddings |
|
||||
| gigachat/Embeddings-2 | 512 | 1024 | Updated embeddings |
|
||||
| gigachat/EmbeddingsGigaR | 4096 | 2560 | High-dimensional embeddings |
|
||||
|
||||
:::note
|
||||
Available models may vary depending on your API access level (personal or business).
|
||||
:::
|
||||
|
||||
## Limitations
|
||||
|
||||
- Only one function call per request (GigaChat API limitation)
|
||||
- Maximum 1 image per message, 10 images total per conversation
|
||||
- GigaChat API uses self-signed SSL certificates - `ssl_verify=False` is required
|
||||
|
|
@ -111,6 +111,7 @@ general_settings:
|
|||
master_key: string
|
||||
maximum_spend_logs_retention_period: 30d # The maximum time to retain spend logs before deletion.
|
||||
maximum_spend_logs_retention_interval: 1d # interval in which the spend log cleanup task should run in.
|
||||
user_mcp_management_mode: restricted # or "view_all"
|
||||
|
||||
# Database Settings
|
||||
database_url: string
|
||||
|
|
@ -230,6 +231,7 @@ router_settings:
|
|||
| image_generation_model | str | The default model to use for image generation - ignores model set in request |
|
||||
| store_model_in_db | boolean | If true, enables storing model + credential information in the DB. |
|
||||
| supported_db_objects | List[str] | Fine-grained control over which object types to load from the database when `store_model_in_db` is True. Available types: `"models"`, `"mcp"`, `"guardrails"`, `"vector_stores"`, `"pass_through_endpoints"`, `"prompts"`, `"model_cost_map"`. If not set, all object types are loaded (default behavior). Example: `supported_db_objects: ["mcp"]` to only load MCP servers from DB. |
|
||||
| user_mcp_management_mode | string | Controls what non-admins can see on the MCP dashboard. `restricted` (default) only lists MCP servers that the user’s teams are explicitly allowed to access. `view_all` lets every user see the full MCP server list. Tool list/call always respects per-key permissions, so users still cannot run MCP calls without access. |
|
||||
| store_prompts_in_spend_logs | boolean | If true, allows prompts and responses to be stored in the spend logs table. |
|
||||
| max_request_size_mb | int | The maximum size for requests in MB. Requests above this size will be rejected. |
|
||||
| max_response_size_mb | int | The maximum size for responses in MB. LLM Responses above this size will not be sent. |
|
||||
|
|
@ -669,6 +671,7 @@ router_settings:
|
|||
| LANGSMITH_DEFAULT_RUN_NAME | Default name for Langsmith run
|
||||
| LANGSMITH_PROJECT | Project name for Langsmith integration
|
||||
| LANGSMITH_SAMPLING_RATE | Sampling rate for Langsmith logging
|
||||
| LANGSMITH_TENANT_ID | Tenant ID for Langsmith multi-tenant deployments
|
||||
| LANGTRACE_API_KEY | API key for Langtrace service
|
||||
| LASSO_API_BASE | Base URL for Lasso API
|
||||
| LASSO_API_KEY | API key for Lasso service
|
||||
|
|
@ -707,6 +710,7 @@ router_settings:
|
|||
| LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development)
|
||||
| LITELLM_NON_ROOT | Flag to run LiteLLM in non-root mode for enhanced security in Docker containers
|
||||
| LITELLM_RATE_LIMIT_WINDOW_SIZE | Rate limit window size for LiteLLM. Default is 60
|
||||
| LITELLM_REASONING_AUTO_SUMMARY | If set to "true", automatically enables detailed reasoning summaries for reasoning models (e.g., o1, o3-mini, deepseek-reasoner). When enabled, adds `summary: "detailed"` to reasoning effort configurations. Default is "false"
|
||||
| LITELLM_SALT_KEY | Salt key for encryption in LiteLLM
|
||||
| LITELLM_SSL_CIPHERS | SSL/TLS cipher configuration for faster handshakes. Controls cipher suite preferences for OpenSSL connections.
|
||||
| LITELLM_SECRET_AWS_KMS_LITELLM_LICENSE | AWS KMS encrypted license for LiteLLM
|
||||
|
|
@ -774,6 +778,7 @@ router_settings:
|
|||
| OTEL_EXPORTER_OTLP_HEADERS | Headers for OpenTelemetry requests
|
||||
| OTEL_SERVICE_NAME | Service name identifier for OpenTelemetry
|
||||
| OTEL_TRACER_NAME | Tracer name for OpenTelemetry tracing
|
||||
| OTEL_LOGS_EXPORTER | Exporter type for OpenTelemetry logs (e.g., console)
|
||||
| PAGERDUTY_API_KEY | API key for PagerDuty Alerting
|
||||
| PANW_PRISMA_AIRS_API_KEY | API key for PANW Prisma AIRS service
|
||||
| PANW_PRISMA_AIRS_API_BASE | Base URL for PANW Prisma AIRS service
|
||||
|
|
@ -888,4 +893,4 @@ router_settings:
|
|||
| DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL | Time-to-live in seconds for health check lock in shared health check mode. Default is 60 (1 minute)
|
||||
| ZSCALER_AI_GUARD_API_KEY | API key for Zscaler AI Guard service
|
||||
| ZSCALER_AI_GUARD_POLICY_ID | Policy ID for Zscaler AI Guard guardrails
|
||||
| ZSCALER_AI_GUARD_URL | Base URL for Zscaler AI Guard API. Default is https://api.us1.zseclipse.net/v1/detection/execute-policy
|
||||
| ZSCALER_AI_GUARD_URL | Base URL for Zscaler AI Guard API. Default is https://api.us1.zseclipse.net/v1/detection/execute-policy
|
||||
|
|
|
|||
|
|
@ -591,3 +591,68 @@ Expected Response
|
|||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## OpenAI Responses API - Auto-Summary Control
|
||||
|
||||
When using OpenAI Responses API models (like `gpt-5`) via `/chat/completions` with `reasoning_effort`, you can control whether `summary="detailed"` is automatically added to the reasoning parameter.
|
||||
|
||||
### Enabling Auto-Summary
|
||||
|
||||
You can enable automatic `summary="detailed"` in two ways:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# Enable auto-summary globally
|
||||
litellm.reasoning_auto_summary = True
|
||||
|
||||
response = litellm.completion(
|
||||
model="openai/responses/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
reasoning_effort="low", # Will automatically add summary="detailed"
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="env" label="Environment Variable">
|
||||
|
||||
```bash
|
||||
# Set environment variable
|
||||
export LITELLM_REASONING_AUTO_SUMMARY=true
|
||||
|
||||
# Or in your .env file
|
||||
LITELLM_REASONING_AUTO_SUMMARY=true
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="Proxy Config">
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
reasoning_auto_summary: true # Enable auto-summary for all requests
|
||||
|
||||
model_list:
|
||||
- model_name: gpt-5-mini
|
||||
litellm_params:
|
||||
model: openai/responses/gpt-5-mini
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Manual Control (Recommended)
|
||||
|
||||
For fine-grained control, pass `reasoning_effort` as a dictionary:
|
||||
|
||||
```python
|
||||
response = litellm.completion(
|
||||
model="openai/responses/gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
reasoning_effort={"effort": "low", "summary": "detailed"}, # Explicit control
|
||||
)
|
||||
```
|
||||
|
|
|
|||
104
docs/my-website/docs/response_api_compact.md
Normal file
104
docs/my-website/docs/response_api_compact.md
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# /responses/compact
|
||||
|
||||
Compress conversation history using OpenAI's `/responses/compact` endpoint.
|
||||
|
||||
| Feature | Supported |
|
||||
|---------|-----------|
|
||||
| Supported LiteLLM Versions | 1.72.0+ |
|
||||
| Supported Providers | `openai` |
|
||||
|
||||
## Usage
|
||||
|
||||
### LiteLLM Python SDK
|
||||
|
||||
```python showLineNumbers title="Compact Response"
|
||||
import litellm
|
||||
|
||||
response = litellm.compact_responses(
|
||||
model="openai/gpt-4o",
|
||||
input=[{"role": "user", "content": "Hello, how are you?"}],
|
||||
instructions="Be helpful",
|
||||
previous_response_id="resp_abc123" # optional
|
||||
)
|
||||
|
||||
print(response.id)
|
||||
print(response.object) # "response.compaction"
|
||||
print(response.output)
|
||||
```
|
||||
|
||||
### LiteLLM Proxy
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="curl" label="Curl">
|
||||
|
||||
```bash showLineNumbers title="Compact Request"
|
||||
curl http://localhost:4000/v1/responses/compact \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "openai/gpt-4o",
|
||||
"input": [{"role": "user", "content": "Hello"}],
|
||||
"instructions": "Be helpful"
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="openai-sdk" label="OpenAI Python SDK">
|
||||
|
||||
```python showLineNumbers title="Compact with OpenAI SDK"
|
||||
import httpx
|
||||
|
||||
response = httpx.post(
|
||||
"http://localhost:4000/v1/responses/compact",
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
json={
|
||||
"model": "openai/gpt-4o",
|
||||
"input": [{"role": "user", "content": "Hello"}],
|
||||
"instructions": "Be helpful"
|
||||
}
|
||||
)
|
||||
|
||||
print(response.json())
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Request Parameters
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
|-----------|------|----------|-------------|
|
||||
| `model` | string | Yes | Model to use for compaction |
|
||||
| `input` | string or array | Yes | Input messages to compact |
|
||||
| `instructions` | string | No | System instructions |
|
||||
| `previous_response_id` | string | No | ID of previous response to continue from |
|
||||
|
||||
## Response Format
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "resp_abc123",
|
||||
"object": "response.compaction",
|
||||
"created_at": 1734366691,
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [...]
|
||||
},
|
||||
{
|
||||
"type": "compaction",
|
||||
"encrypted_content": "..."
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"total_tokens": 150
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
BIN
docs/my-website/img/mcp_allow_all_ui.png
Normal file
BIN
docs/my-website/img/mcp_allow_all_ui.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 135 KiB |
|
|
@ -541,7 +541,14 @@ const sidebars = {
|
|||
},
|
||||
"realtime",
|
||||
"rerank",
|
||||
"response_api",
|
||||
{
|
||||
type: "category",
|
||||
label: "/responses",
|
||||
items: [
|
||||
"response_api",
|
||||
"response_api_compact",
|
||||
]
|
||||
},
|
||||
{
|
||||
type: "category",
|
||||
label: "/search",
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from pathlib import Path
|
|||
from typing import Optional
|
||||
|
||||
from litellm_proxy_extras._logging import logger
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
|
||||
def str_to_bool(value: Optional[str]) -> bool:
|
||||
|
|
@ -18,6 +19,103 @@ def str_to_bool(value: Optional[str]) -> bool:
|
|||
return value.lower() in ("true", "1", "t", "y", "yes")
|
||||
|
||||
|
||||
class MigrationLockManager:
|
||||
"""Redis-based lock manager for database migrations"""
|
||||
|
||||
MIGRATION_LOCK_KEY = "migration_lock"
|
||||
LOCK_TTL_SECONDS = 300 # 5 minutes TTL
|
||||
|
||||
def __init__(self, redis_cache: Optional[RedisCache] = None):
|
||||
self.redis_cache = redis_cache
|
||||
self.lock_acquired = False
|
||||
self.pod_id = f"pod_{os.getpid()}_{int(time.time())}"
|
||||
|
||||
def _get_redis_lock_key(self) -> str:
|
||||
"""Get Redis lock key for migration"""
|
||||
return f"migration_lock:{self.MIGRATION_LOCK_KEY}"
|
||||
|
||||
def acquire_lock(self) -> bool:
|
||||
"""Acquire migration lock"""
|
||||
if self.redis_cache is None:
|
||||
logger.warning(
|
||||
"Redis cache is not available, running migration without lock protection"
|
||||
)
|
||||
self.lock_acquired = True
|
||||
return True
|
||||
|
||||
try:
|
||||
lock_key = self._get_redis_lock_key()
|
||||
|
||||
# Redis SET with NX (only if not exists) and EX (expiration)
|
||||
acquired = self.redis_cache.set_cache(
|
||||
key=lock_key, value=self.pod_id, nx=True, ttl=self.LOCK_TTL_SECONDS
|
||||
)
|
||||
|
||||
if acquired:
|
||||
self.lock_acquired = True
|
||||
logger.info(f"Migration lock acquired by pod {self.pod_id}")
|
||||
return True
|
||||
else:
|
||||
logger.info("Migration lock is already held by another pod")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to acquire migration lock: {e}")
|
||||
return False
|
||||
|
||||
def wait_for_lock_release(
|
||||
self, check_interval: int = 5, max_wait: int = 300
|
||||
) -> bool:
|
||||
"""Wait for another process to release the lock"""
|
||||
if self.redis_cache is None:
|
||||
logger.warning("Redis cache is not available, cannot wait for lock")
|
||||
return False
|
||||
|
||||
logger.info(f"Waiting for migration lock to be released (max {max_wait}s)...")
|
||||
start_time = time.time()
|
||||
|
||||
while time.time() - start_time < max_wait:
|
||||
# Try to acquire lock using the public acquire_lock method
|
||||
if self.acquire_lock():
|
||||
logger.info(
|
||||
f"Migration lock acquired after waiting by pod {self.pod_id}"
|
||||
)
|
||||
return True
|
||||
|
||||
time.sleep(check_interval)
|
||||
|
||||
logger.warning(f"Failed to acquire migration lock within {max_wait} seconds")
|
||||
return False
|
||||
|
||||
def release_lock(self):
|
||||
"""Release migration lock"""
|
||||
if not self.lock_acquired or self.redis_cache is None:
|
||||
return
|
||||
|
||||
try:
|
||||
lock_key = self._get_redis_lock_key()
|
||||
|
||||
# Verify current pod owns the lock
|
||||
current_value = self.redis_cache.get_cache(lock_key)
|
||||
if current_value and str(current_value) == self.pod_id:
|
||||
self.redis_cache.delete_cache(lock_key)
|
||||
logger.info(f"Migration lock released by pod {self.pod_id}")
|
||||
else:
|
||||
logger.warning(f"Pod {self.pod_id} cannot release lock (not owner)")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to release migration lock: {e}")
|
||||
finally:
|
||||
self.lock_acquired = False
|
||||
|
||||
def __enter__(self):
|
||||
"""Context manager entry - acquire lock when entering with statement"""
|
||||
self.acquire_lock()
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Context manager exit - release lock when exiting with statement"""
|
||||
self.release_lock()
|
||||
|
||||
def _get_prisma_env() -> dict:
|
||||
"""Get environment variables for Prisma, handling offline mode if configured."""
|
||||
|
|
@ -346,19 +444,50 @@ class ProxyExtrasDBManager:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def setup_database(use_migrate: bool = False) -> bool:
|
||||
def setup_database(
|
||||
use_migrate: bool = False, redis_cache: Optional[RedisCache] = None
|
||||
) -> bool:
|
||||
"""
|
||||
Set up the database using either prisma migrate or prisma db push
|
||||
Uses migrations from litellm-proxy-extras package
|
||||
Uses migrations from litellm-proxy-extras package.
|
||||
In multi-instance environment, use redis lock to prevent concurrent execution.
|
||||
|
||||
Args:
|
||||
schema_path (str): Path to the Prisma schema file
|
||||
use_migrate (bool): Whether to use prisma migrate instead of db push
|
||||
redis_cache: Redis cache instance for distributed locking
|
||||
|
||||
Returns:
|
||||
bool: True if setup was successful, False otherwise
|
||||
"""
|
||||
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
|
||||
|
||||
database_url = os.getenv("DATABASE_URL")
|
||||
if not database_url:
|
||||
logger.error("DATABASE_URL environment variable is not set")
|
||||
return False
|
||||
|
||||
# Use MigrationLockManager to prevent concurrent migration execution
|
||||
with MigrationLockManager(redis_cache) as lock_manager:
|
||||
# Lock is already acquired in __enter__, check if it was successful
|
||||
if not lock_manager.lock_acquired:
|
||||
# Cannot acquire lock, another process is running migration
|
||||
logger.info(
|
||||
"Another pod is running migration, waiting for completion..."
|
||||
)
|
||||
|
||||
# Wait for other process to complete migration
|
||||
if not lock_manager.wait_for_lock_release():
|
||||
logger.error("Failed to acquire migration lock after waiting")
|
||||
return False
|
||||
|
||||
# Successfully acquired lock, proceed with migration
|
||||
logger.info("Acquired migration lock, proceeding with migration")
|
||||
return ProxyExtrasDBManager._execute_migration(use_migrate, schema_path)
|
||||
|
||||
@staticmethod
|
||||
def _execute_migration(use_migrate: bool, schema_path: str) -> bool:
|
||||
"""Execute the actual migration"""
|
||||
for attempt in range(4):
|
||||
original_dir = os.getcwd()
|
||||
migrations_dir = ProxyExtrasDBManager._get_prisma_dir()
|
||||
|
|
|
|||
|
|
@ -197,6 +197,7 @@ retry = True
|
|||
api_key: Optional[str] = None
|
||||
openai_key: Optional[str] = None
|
||||
groq_key: Optional[str] = None
|
||||
gigachat_key: Optional[str] = None
|
||||
databricks_key: Optional[str] = None
|
||||
openai_like_key: Optional[str] = None
|
||||
azure_key: Optional[str] = None
|
||||
|
|
@ -275,6 +276,7 @@ banned_keywords_list: Optional[Union[str, List]] = None
|
|||
llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all"
|
||||
guardrail_name_config_map: Dict[str, GuardrailItem] = {}
|
||||
include_cost_in_streaming_usage: bool = False
|
||||
reasoning_auto_summary: bool = False
|
||||
### PROMPTS ####
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
|
||||
|
|
@ -1440,6 +1442,8 @@ if TYPE_CHECKING:
|
|||
from .llms.github_copilot.chat.transformation import GithubCopilotConfig as GithubCopilotConfig
|
||||
from .llms.github_copilot.responses.transformation import GithubCopilotResponsesAPIConfig as GithubCopilotResponsesAPIConfig
|
||||
from .llms.github_copilot.embedding.transformation import GithubCopilotEmbeddingConfig as GithubCopilotEmbeddingConfig
|
||||
from .llms.gigachat.chat.transformation import GigaChatConfig as GigaChatConfig
|
||||
from .llms.gigachat.embedding.transformation import GigaChatEmbeddingConfig as GigaChatEmbeddingConfig
|
||||
from .llms.nebius.chat.transformation import NebiusConfig as NebiusConfig
|
||||
from .llms.wandb.chat.transformation import WandbConfig as WandbConfig
|
||||
from .llms.dashscope.chat.transformation import DashScopeChatConfig as DashScopeChatConfig
|
||||
|
|
|
|||
|
|
@ -255,6 +255,8 @@ LLM_CONFIG_NAMES = (
|
|||
"GithubCopilotEmbeddingConfig",
|
||||
"NebiusConfig",
|
||||
"WandbConfig",
|
||||
"GigaChatConfig",
|
||||
"GigaChatEmbeddingConfig",
|
||||
"DashScopeChatConfig",
|
||||
"MoonshotChatConfig",
|
||||
"DockerModelRunnerChatConfig",
|
||||
|
|
@ -644,6 +646,8 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
"GithubCopilotEmbeddingConfig": (".llms.github_copilot.embedding.transformation", "GithubCopilotEmbeddingConfig"),
|
||||
"NebiusConfig": (".llms.nebius.chat.transformation", "NebiusConfig"),
|
||||
"WandbConfig": (".llms.wandb.chat.transformation", "WandbConfig"),
|
||||
"GigaChatConfig": (".llms.gigachat.chat.transformation", "GigaChatConfig"),
|
||||
"GigaChatEmbeddingConfig": (".llms.gigachat.embedding.transformation", "GigaChatEmbeddingConfig"),
|
||||
"DashScopeChatConfig": (".llms.dashscope.chat.transformation", "DashScopeChatConfig"),
|
||||
"MoonshotChatConfig": (".llms.moonshot.chat.transformation", "MoonshotChatConfig"),
|
||||
"DockerModelRunnerChatConfig": (".llms.docker_model_runner.chat.transformation", "DockerModelRunnerChatConfig"),
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req
|
|||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
|
|
@ -22,6 +23,7 @@ from typing import (
|
|||
from openai.types.responses.tool_param import FunctionToolParam
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm import ModelResponse
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
|
|
@ -691,19 +693,26 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if isinstance(reasoning_effort, dict):
|
||||
return Reasoning(**reasoning_effort) # type: ignore[typeddict-item]
|
||||
|
||||
# If string is passed, map with summary="detailed"
|
||||
# Check if auto-summary is enabled via flag or environment variable
|
||||
# Priority: litellm.reasoning_auto_summary flag > LITELLM_REASONING_AUTO_SUMMARY env var
|
||||
auto_summary_enabled = (
|
||||
litellm.reasoning_auto_summary
|
||||
or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
|
||||
)
|
||||
|
||||
# If string is passed, map with optional summary based on flag/env var
|
||||
if reasoning_effort == "none":
|
||||
return Reasoning(effort="none", summary="detailed") # type: ignore
|
||||
return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") # type: ignore
|
||||
elif reasoning_effort == "high":
|
||||
return Reasoning(effort="high", summary="detailed")
|
||||
return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high")
|
||||
elif reasoning_effort == "xhigh":
|
||||
return Reasoning(effort="xhigh", summary="detailed") # type: ignore[typeddict-item]
|
||||
return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") # type: ignore[typeddict-item]
|
||||
elif reasoning_effort == "medium":
|
||||
return Reasoning(effort="medium", summary="detailed")
|
||||
return Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium")
|
||||
elif reasoning_effort == "low":
|
||||
return Reasoning(effort="low", summary="detailed")
|
||||
return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low")
|
||||
elif reasoning_effort == "minimal":
|
||||
return Reasoning(effort="minimal", summary="detailed")
|
||||
return Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal")
|
||||
return None
|
||||
|
||||
def _transform_response_format_to_text_format(
|
||||
|
|
|
|||
|
|
@ -375,6 +375,7 @@ LITELLM_CHAT_PROVIDERS = [
|
|||
"perplexity",
|
||||
"mistral",
|
||||
"groq",
|
||||
"gigachat",
|
||||
"nvidia_nim",
|
||||
"cerebras",
|
||||
"baseten",
|
||||
|
|
|
|||
|
|
@ -187,6 +187,12 @@
|
|||
"ui_name": "Sampling Rate",
|
||||
"description": "Sampling rate for logging (0.0 to 1.0, default: 1.0)",
|
||||
"required": false
|
||||
},
|
||||
"langsmith_tenant_id": {
|
||||
"type": "text",
|
||||
"ui_name": "Tenant ID",
|
||||
"description": "LangSmith tenant ID for organization-scoped API keys (required when using org-scoped keys)",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"description": "Langsmith Logging Integration"
|
||||
|
|
|
|||
|
|
@ -50,6 +50,42 @@ else:
|
|||
Langfuse = Any
|
||||
|
||||
|
||||
def _extract_cache_read_input_tokens(usage_obj) -> int:
|
||||
"""
|
||||
Extract cache_read_input_tokens from usage object.
|
||||
|
||||
Checks both:
|
||||
1. Top-level cache_read_input_tokens (Anthropic format)
|
||||
2. prompt_tokens_details.cached_tokens (Gemini, OpenAI format)
|
||||
|
||||
See: https://github.com/BerriAI/litellm/issues/18520
|
||||
|
||||
Args:
|
||||
usage_obj: Usage object from LLM response
|
||||
|
||||
Returns:
|
||||
int: Number of cached tokens read, defaults to 0
|
||||
"""
|
||||
cache_read_input_tokens = usage_obj.get("cache_read_input_tokens") or 0
|
||||
|
||||
# Check prompt_tokens_details.cached_tokens (used by Gemini and other providers)
|
||||
if hasattr(usage_obj, "prompt_tokens_details"):
|
||||
prompt_tokens_details = getattr(usage_obj, "prompt_tokens_details", None)
|
||||
if (
|
||||
prompt_tokens_details is not None
|
||||
and hasattr(prompt_tokens_details, "cached_tokens")
|
||||
):
|
||||
cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
|
||||
if (
|
||||
cached_tokens is not None
|
||||
and isinstance(cached_tokens, (int, float))
|
||||
and cached_tokens > 0
|
||||
):
|
||||
cache_read_input_tokens = cached_tokens
|
||||
|
||||
return cache_read_input_tokens
|
||||
|
||||
|
||||
class LangFuseLogger:
|
||||
# Class variables or attributes
|
||||
def __init__(
|
||||
|
|
@ -757,8 +793,8 @@ class LangFuseLogger:
|
|||
cache_creation_input_tokens = (
|
||||
_usage_obj.get("cache_creation_input_tokens") or 0
|
||||
)
|
||||
cache_read_input_tokens = (
|
||||
_usage_obj.get("cache_read_input_tokens") or 0
|
||||
cache_read_input_tokens = _extract_cache_read_input_tokens(
|
||||
_usage_obj
|
||||
)
|
||||
|
||||
usage = {
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
langsmith_project: Optional[str] = None,
|
||||
langsmith_base_url: Optional[str] = None,
|
||||
langsmith_sampling_rate: Optional[float] = None,
|
||||
langsmith_tenant_id: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.flush_lock = asyncio.Lock()
|
||||
|
|
@ -48,6 +49,7 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
langsmith_api_key=langsmith_api_key,
|
||||
langsmith_project=langsmith_project,
|
||||
langsmith_base_url=langsmith_base_url,
|
||||
langsmith_tenant_id=langsmith_tenant_id,
|
||||
)
|
||||
self.sampling_rate: float = (
|
||||
langsmith_sampling_rate
|
||||
|
|
@ -76,6 +78,7 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
langsmith_api_key: Optional[str] = None,
|
||||
langsmith_project: Optional[str] = None,
|
||||
langsmith_base_url: Optional[str] = None,
|
||||
langsmith_tenant_id: Optional[str] = None,
|
||||
) -> LangsmithCredentialsObject:
|
||||
_credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY")
|
||||
_credentials_project = (
|
||||
|
|
@ -86,11 +89,13 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
or os.getenv("LANGSMITH_BASE_URL")
|
||||
or "https://api.smith.langchain.com"
|
||||
)
|
||||
_credentials_tenant_id = langsmith_tenant_id or os.getenv("LANGSMITH_TENANT_ID")
|
||||
|
||||
return LangsmithCredentialsObject(
|
||||
LANGSMITH_API_KEY=_credentials_api_key,
|
||||
LANGSMITH_BASE_URL=_credentials_base_url,
|
||||
LANGSMITH_PROJECT=_credentials_project,
|
||||
LANGSMITH_TENANT_ID=_credentials_tenant_id,
|
||||
)
|
||||
|
||||
def _prepare_log_data(
|
||||
|
|
@ -365,8 +370,11 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
"""
|
||||
langsmith_api_base = credentials["LANGSMITH_BASE_URL"]
|
||||
langsmith_api_key = credentials["LANGSMITH_API_KEY"]
|
||||
langsmith_tenant_id = credentials.get("LANGSMITH_TENANT_ID")
|
||||
url = self._add_endpoint_to_url(langsmith_api_base, "runs/batch")
|
||||
headers = {"x-api-key": langsmith_api_key}
|
||||
if langsmith_tenant_id:
|
||||
headers["x-tenant-id"] = langsmith_tenant_id
|
||||
elements_to_log = [queue_object["data"] for queue_object in queue_objects]
|
||||
|
||||
try:
|
||||
|
|
@ -418,6 +426,7 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
api_key=credentials["LANGSMITH_API_KEY"],
|
||||
project=credentials["LANGSMITH_PROJECT"],
|
||||
base_url=credentials["LANGSMITH_BASE_URL"],
|
||||
tenant_id=credentials.get("LANGSMITH_TENANT_ID"),
|
||||
)
|
||||
|
||||
if key not in log_queue_by_credentials:
|
||||
|
|
@ -466,6 +475,9 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
langsmith_base_url=standard_callback_dynamic_params.get(
|
||||
"langsmith_base_url", None
|
||||
),
|
||||
langsmith_tenant_id=standard_callback_dynamic_params.get(
|
||||
"langsmith_tenant_id", None
|
||||
),
|
||||
)
|
||||
else:
|
||||
credentials = self.default_credentials
|
||||
|
|
@ -491,13 +503,16 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
|
||||
def get_run_by_id(self, run_id):
|
||||
langsmith_api_key = self.default_credentials["LANGSMITH_API_KEY"]
|
||||
|
||||
langsmith_api_base = self.default_credentials["LANGSMITH_BASE_URL"]
|
||||
langsmith_tenant_id = self.default_credentials.get("LANGSMITH_TENANT_ID")
|
||||
|
||||
url = f"{langsmith_api_base}/runs/{run_id}"
|
||||
headers = {"x-api-key": langsmith_api_key}
|
||||
if langsmith_tenant_id:
|
||||
headers["x-tenant-id"] = langsmith_tenant_id
|
||||
response = litellm.module_level_client.get(
|
||||
url=url,
|
||||
headers={"x-api-key": langsmith_api_key},
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
return response.json()
|
||||
|
|
|
|||
|
|
@ -196,50 +196,88 @@ class OpenTelemetry(CustomLogger):
|
|||
litellm.service_callback.append(self)
|
||||
setattr(proxy_server, "open_telemetry_logger", self)
|
||||
|
||||
def _get_or_create_provider(
|
||||
self,
|
||||
provider,
|
||||
provider_name: str,
|
||||
get_existing_provider_fn,
|
||||
sdk_provider_class,
|
||||
create_new_provider_fn,
|
||||
set_provider_fn,
|
||||
):
|
||||
"""
|
||||
Generic helper to get or create an OpenTelemetry provider (Tracer, Meter, or Logger).
|
||||
|
||||
Args:
|
||||
provider: The provider instance passed to the init function (can be None)
|
||||
provider_name: Name for logging (e.g., "TracerProvider")
|
||||
get_existing_provider_fn: Function to get the existing global provider
|
||||
sdk_provider_class: The SDK provider class to check for (e.g., TracerProvider from SDK)
|
||||
create_new_provider_fn: Function to create a new provider instance
|
||||
set_provider_fn: Function to set the provider globally
|
||||
|
||||
Returns:
|
||||
The provider to use (either existing, new, or explicitly provided)
|
||||
"""
|
||||
if provider is not None:
|
||||
# Provider explicitly provided (e.g., for testing)
|
||||
# Do NOT call set_provider_fn - the caller is responsible for managing global state
|
||||
# If they want it to be global, they've already set it before passing it to us
|
||||
verbose_logger.debug(
|
||||
"OpenTelemetry: Using provided TracerProvider: %s",
|
||||
type(provider).__name__,
|
||||
)
|
||||
return provider
|
||||
|
||||
# Check if a provider is already set globally
|
||||
try:
|
||||
existing_provider = get_existing_provider_fn()
|
||||
|
||||
# If a real SDK provider exists (set by another SDK like Langfuse), use it
|
||||
# This uses a positive check for SDK providers instead of a negative check for proxy providers
|
||||
if isinstance(existing_provider, sdk_provider_class):
|
||||
verbose_logger.debug(
|
||||
"OpenTelemetry: Using existing %s: %s",
|
||||
provider_name,
|
||||
type(existing_provider).__name__,
|
||||
)
|
||||
provider = existing_provider
|
||||
# Don't call set_provider to preserve existing context
|
||||
else:
|
||||
# Default proxy provider or unknown type, create our own
|
||||
verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name)
|
||||
provider = create_new_provider_fn()
|
||||
set_provider_fn(provider)
|
||||
except Exception as e:
|
||||
# Fallback: create a new provider if something goes wrong
|
||||
verbose_logger.debug(
|
||||
"OpenTelemetry: Exception checking existing %s, creating new one: %s",
|
||||
provider_name,
|
||||
str(e),
|
||||
)
|
||||
provider = create_new_provider_fn()
|
||||
set_provider_fn(provider)
|
||||
|
||||
return provider
|
||||
|
||||
def _init_tracing(self, tracer_provider):
|
||||
from opentelemetry import trace
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
from opentelemetry.trace import SpanKind
|
||||
|
||||
# use provided tracer or create a new one
|
||||
if tracer_provider is None:
|
||||
# Check if a TracerProvider is already set globally (e.g., by Langfuse SDK)
|
||||
try:
|
||||
from opentelemetry.trace import ProxyTracerProvider
|
||||
def create_tracer_provider():
|
||||
provider = TracerProvider(resource=_get_litellm_resource())
|
||||
provider.add_span_processor(self._get_span_processor())
|
||||
return provider
|
||||
|
||||
existing_provider = trace.get_tracer_provider()
|
||||
|
||||
# If an actual provider exists (not the default proxy), use it
|
||||
if not isinstance(existing_provider, ProxyTracerProvider):
|
||||
verbose_logger.debug(
|
||||
"OpenTelemetry: Using existing TracerProvider: %s",
|
||||
type(existing_provider).__name__,
|
||||
)
|
||||
tracer_provider = existing_provider
|
||||
# Don't call set_tracer_provider to preserve existing context
|
||||
else:
|
||||
# No real provider exists yet, create our own
|
||||
verbose_logger.debug("OpenTelemetry: Creating new TracerProvider")
|
||||
tracer_provider = TracerProvider(resource=_get_litellm_resource())
|
||||
tracer_provider.add_span_processor(self._get_span_processor())
|
||||
trace.set_tracer_provider(tracer_provider)
|
||||
except Exception as e:
|
||||
# Fallback: create a new provider if something goes wrong
|
||||
verbose_logger.debug(
|
||||
"OpenTelemetry: Exception checking existing provider, creating new one: %s",
|
||||
str(e),
|
||||
)
|
||||
tracer_provider = TracerProvider(resource=_get_litellm_resource())
|
||||
tracer_provider.add_span_processor(self._get_span_processor())
|
||||
trace.set_tracer_provider(tracer_provider)
|
||||
else:
|
||||
# Tracer provider explicitly provided (e.g., for testing)
|
||||
# Do NOT call set_tracer_provider - the caller is responsible for managing global state
|
||||
# If they want it to be global, they've already set it before passing it to us
|
||||
verbose_logger.debug(
|
||||
"OpenTelemetry: Using provided TracerProvider: %s",
|
||||
type(tracer_provider).__name__,
|
||||
)
|
||||
tracer_provider = self._get_or_create_provider(
|
||||
provider=tracer_provider,
|
||||
provider_name="TracerProvider",
|
||||
get_existing_provider_fn=trace.get_tracer_provider,
|
||||
sdk_provider_class=TracerProvider,
|
||||
create_new_provider_fn=create_tracer_provider,
|
||||
set_provider_fn=trace.set_tracer_provider,
|
||||
)
|
||||
|
||||
# Grab our tracer from the TracerProvider (not from global context)
|
||||
# This ensures we use the provided TracerProvider (e.g., for testing)
|
||||
|
|
@ -257,39 +295,24 @@ class OpenTelemetry(CustomLogger):
|
|||
return
|
||||
|
||||
from opentelemetry import metrics
|
||||
from opentelemetry.sdk.metrics import Histogram, MeterProvider
|
||||
from opentelemetry.sdk.metrics import MeterProvider
|
||||
|
||||
# Only create OTLP infrastructure if no custom meter provider is provided
|
||||
if meter_provider is None:
|
||||
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
|
||||
OTLPMetricExporter,
|
||||
)
|
||||
from opentelemetry.sdk.metrics.export import (
|
||||
AggregationTemporality,
|
||||
PeriodicExportingMetricReader,
|
||||
def create_meter_provider():
|
||||
metric_reader = self._get_metric_reader()
|
||||
return MeterProvider(
|
||||
metric_readers=[metric_reader], resource=_get_litellm_resource()
|
||||
)
|
||||
|
||||
normalized_endpoint = self._normalize_otel_endpoint(
|
||||
self.config.endpoint, "metrics"
|
||||
)
|
||||
_metric_exporter = OTLPMetricExporter(
|
||||
endpoint=normalized_endpoint,
|
||||
headers=OpenTelemetry._get_headers_dictionary(self.config.headers),
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
_metric_reader = PeriodicExportingMetricReader(
|
||||
_metric_exporter, export_interval_millis=10000
|
||||
)
|
||||
meter_provider = self._get_or_create_provider(
|
||||
provider=meter_provider,
|
||||
provider_name="MeterProvider",
|
||||
get_existing_provider_fn=metrics.get_meter_provider,
|
||||
sdk_provider_class=MeterProvider,
|
||||
create_new_provider_fn=create_meter_provider,
|
||||
set_provider_fn=metrics.set_meter_provider,
|
||||
)
|
||||
|
||||
meter_provider = MeterProvider(
|
||||
metric_readers=[_metric_reader], resource=_get_litellm_resource()
|
||||
)
|
||||
meter = meter_provider.get_meter(__name__)
|
||||
else:
|
||||
# Use the provided meter provider as-is, without creating additional OTLP infrastructure
|
||||
meter = meter_provider.get_meter(__name__)
|
||||
|
||||
metrics.set_meter_provider(meter_provider)
|
||||
meter = meter_provider.get_meter(__name__)
|
||||
|
||||
self._operation_duration_histogram = meter.create_histogram(
|
||||
name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38
|
||||
|
|
@ -327,22 +350,26 @@ class OpenTelemetry(CustomLogger):
|
|||
if not self.config.enable_events:
|
||||
return
|
||||
|
||||
from opentelemetry._logs import set_logger_provider
|
||||
from opentelemetry._logs import get_logger_provider, set_logger_provider
|
||||
from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider
|
||||
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
|
||||
|
||||
# set up log pipeline
|
||||
if logger_provider is None:
|
||||
litellm_resource = _get_litellm_resource()
|
||||
logger_provider = OTLoggerProvider(resource=litellm_resource)
|
||||
# Only add OTLP exporter if we created the logger provider ourselves
|
||||
def create_logger_provider():
|
||||
provider = OTLoggerProvider(resource=_get_litellm_resource())
|
||||
log_exporter = self._get_log_exporter()
|
||||
if log_exporter:
|
||||
logger_provider.add_log_record_processor(
|
||||
BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
|
||||
)
|
||||
provider.add_log_record_processor(
|
||||
BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
|
||||
)
|
||||
return provider
|
||||
|
||||
set_logger_provider(logger_provider)
|
||||
self._get_or_create_provider(
|
||||
provider=logger_provider,
|
||||
provider_name="LoggerProvider",
|
||||
get_existing_provider_fn=get_logger_provider,
|
||||
sdk_provider_class=OTLoggerProvider,
|
||||
create_new_provider_fn=create_logger_provider,
|
||||
set_provider_fn=set_logger_provider,
|
||||
)
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
self._handle_success(kwargs, response_obj, start_time, end_time)
|
||||
|
|
@ -944,6 +971,15 @@ class OpenTelemetry(CustomLogger):
|
|||
if not self.config.enable_events:
|
||||
return
|
||||
|
||||
# NOTE: Semantic logs (gen_ai.content.prompt/completion events) have compatibility issues
|
||||
# with OTEL SDK >= 1.39.0 due to breaking changes in PR #4676:
|
||||
# - LogRecord moved from opentelemetry.sdk._logs to opentelemetry.sdk._logs._internal
|
||||
# - LogRecord constructor no longer accepts 'resource' parameter (now inherited from LoggerProvider)
|
||||
# - LogData class was removed entirely
|
||||
# These logs work correctly in OTEL SDK < 1.39.0 but may fail in >= 1.39.0.
|
||||
# See: https://github.com/open-telemetry/opentelemetry-python/pull/4676
|
||||
# TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords
|
||||
|
||||
from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider
|
||||
from opentelemetry.sdk._logs import LogRecord as SdkLogRecord
|
||||
|
||||
|
|
@ -1807,7 +1843,8 @@ class OpenTelemetry(CustomLogger):
|
|||
)
|
||||
return self.OTEL_EXPORTER
|
||||
|
||||
if self.OTEL_EXPORTER == "console":
|
||||
otel_logs_exporter = os.getenv("OTEL_LOGS_EXPORTER")
|
||||
if self.OTEL_EXPORTER == "console" or otel_logs_exporter == "console":
|
||||
from opentelemetry.sdk._logs.export import ConsoleLogExporter
|
||||
|
||||
verbose_logger.debug(
|
||||
|
|
@ -1854,6 +1891,67 @@ class OpenTelemetry(CustomLogger):
|
|||
|
||||
return ConsoleLogExporter()
|
||||
|
||||
def _get_metric_reader(self):
|
||||
"""
|
||||
Get the appropriate metric reader based on the configuration.
|
||||
"""
|
||||
from opentelemetry.sdk.metrics import Histogram
|
||||
from opentelemetry.sdk.metrics.export import (
|
||||
AggregationTemporality,
|
||||
ConsoleMetricExporter,
|
||||
PeriodicExportingMetricReader,
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
"OpenTelemetry Logger, initializing metric reader\nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
|
||||
self.OTEL_EXPORTER,
|
||||
self.OTEL_ENDPOINT,
|
||||
self.OTEL_HEADERS,
|
||||
)
|
||||
|
||||
_split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)
|
||||
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "metrics")
|
||||
|
||||
if self.OTEL_EXPORTER == "console":
|
||||
exporter = ConsoleMetricExporter()
|
||||
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
|
||||
|
||||
elif (
|
||||
self.OTEL_EXPORTER == "otlp_http"
|
||||
or self.OTEL_EXPORTER == "http/protobuf"
|
||||
or self.OTEL_EXPORTER == "http/json"
|
||||
):
|
||||
from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
|
||||
OTLPMetricExporter,
|
||||
)
|
||||
|
||||
exporter = OTLPMetricExporter(
|
||||
endpoint=normalized_endpoint,
|
||||
headers=_split_otel_headers,
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
|
||||
|
||||
elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc":
|
||||
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
|
||||
OTLPMetricExporter,
|
||||
)
|
||||
|
||||
exporter = OTLPMetricExporter(
|
||||
endpoint=normalized_endpoint,
|
||||
headers=_split_otel_headers,
|
||||
preferred_temporality={Histogram: AggregationTemporality.DELTA},
|
||||
)
|
||||
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
|
||||
|
||||
else:
|
||||
verbose_logger.warning(
|
||||
"OpenTelemetry: Unknown metric exporter '%s', defaulting to console. Supported: console, otlp_http, otlp_grpc",
|
||||
self.OTEL_EXPORTER,
|
||||
)
|
||||
exporter = ConsoleMetricExporter()
|
||||
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
|
||||
|
||||
def _normalize_otel_endpoint(
|
||||
self, endpoint: Optional[str], signal_type: str
|
||||
) -> Optional[str]:
|
||||
|
|
|
|||
|
|
@ -2000,24 +2000,56 @@ class CustomStreamWrapper:
|
|||
)
|
||||
## Map to OpenAI Exception
|
||||
try:
|
||||
raise exception_type(
|
||||
mapped_exception = exception_type(
|
||||
model=self.model,
|
||||
custom_llm_provider=self.custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
except Exception as e:
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
except Exception as mapping_error:
|
||||
mapped_exception = mapping_error
|
||||
|
||||
raise MidStreamFallbackError(
|
||||
message=str(e),
|
||||
model=self.model,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
original_exception=e,
|
||||
generated_content=self.response_uptil_now,
|
||||
is_pre_first_chunk=not self.sent_first_chunk,
|
||||
)
|
||||
def _normalize_status_code(exc: Exception) -> Optional[int]:
|
||||
"""
|
||||
Best-effort status_code extraction.
|
||||
Uses status_code on the exception, then falls back to the response.
|
||||
"""
|
||||
try:
|
||||
code = getattr(exc, "status_code", None)
|
||||
if code is not None:
|
||||
return int(code)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
response = getattr(exc, "response", None)
|
||||
if response is not None:
|
||||
try:
|
||||
status_code = getattr(response, "status_code", None)
|
||||
if status_code is not None:
|
||||
return int(status_code)
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
mapped_status_code = _normalize_status_code(mapped_exception)
|
||||
original_status_code = _normalize_status_code(e)
|
||||
|
||||
if mapped_status_code is not None and 400 <= mapped_status_code < 500:
|
||||
raise mapped_exception
|
||||
if original_status_code is not None and 400 <= original_status_code < 500:
|
||||
raise mapped_exception
|
||||
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
raise MidStreamFallbackError(
|
||||
message=str(mapped_exception),
|
||||
model=self.model,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
original_exception=mapped_exception,
|
||||
generated_content=self.response_uptil_now,
|
||||
is_pre_first_chunk=not self.sent_first_chunk,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _strip_sse_data_from_chunk(chunk: Optional[str]) -> Optional[str]:
|
||||
|
|
|
|||
|
|
@ -242,3 +242,30 @@ class BaseResponsesAPIConfig(ABC):
|
|||
#########################################################
|
||||
########## END CANCEL RESPONSE API TRANSFORMATION #######
|
||||
#########################################################
|
||||
|
||||
#########################################################
|
||||
########## COMPACT RESPONSE API TRANSFORMATION ##########
|
||||
#########################################################
|
||||
@abstractmethod
|
||||
def transform_compact_response_api_request(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, ResponseInputParam],
|
||||
response_api_optional_request_params: Dict,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def transform_compact_response_api_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ResponsesAPIResponse:
|
||||
pass
|
||||
|
||||
#########################################################
|
||||
########## END COMPACT RESPONSE API TRANSFORMATION ######
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -91,6 +91,7 @@ from litellm.types.rerank import RerankResponse
|
|||
from litellm.types.responses.main import DeleteResponseResult
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
EmbeddingResponse,
|
||||
FileTypes,
|
||||
LiteLLMBatch,
|
||||
|
|
@ -850,7 +851,9 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
sync_httpx_client = _get_httpx_client()
|
||||
sync_httpx_client = _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
|
|
@ -896,7 +899,8 @@ class BaseLLMHTTPHandler:
|
|||
) -> EmbeddingResponse:
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider)
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
|
@ -2004,6 +2008,10 @@ class BaseLLMHTTPHandler:
|
|||
"""
|
||||
Handles responses API requests.
|
||||
When _is_async=True, returns a coroutine instead of making the call directly.
|
||||
|
||||
Keeps the pre-transform request context for streaming so post-call hooks/metadata
|
||||
(added for Responses API parity with chat) receive the original params instead of
|
||||
the provider-shaped body that caused them to be skipped before.
|
||||
"""
|
||||
|
||||
if _is_async:
|
||||
|
|
@ -2060,6 +2068,18 @@ class BaseLLMHTTPHandler:
|
|||
if extra_body:
|
||||
data.update(extra_body)
|
||||
|
||||
# Preserve the OpenAI-style request context (not sent to the provider) for streaming
|
||||
# hooks/metadata; the streaming iterator now consumes this to run deployment hooks
|
||||
# with the same info as chat, including litellm_params.
|
||||
request_context: Dict[str, Any] = {"input": input}
|
||||
try:
|
||||
request_context.update(response_api_optional_request_params)
|
||||
except Exception:
|
||||
pass
|
||||
# Needed by streaming callbacks/metadata helpers to reconstruct api_base/model_id
|
||||
# but never included in the outbound provider payload.
|
||||
request_context["litellm_params"] = dict(litellm_params)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -2097,6 +2117,8 @@ class BaseLLMHTTPHandler:
|
|||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
return SyncResponsesAPIStreamingIterator(
|
||||
|
|
@ -2106,6 +2128,8 @@ class BaseLLMHTTPHandler:
|
|||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
else:
|
||||
# For non-streaming requests
|
||||
|
|
@ -2189,6 +2213,18 @@ class BaseLLMHTTPHandler:
|
|||
if extra_body:
|
||||
data.update(extra_body)
|
||||
|
||||
# Preserve the OpenAI-style request context (not sent to the provider) for streaming
|
||||
# hooks/metadata; the streaming iterator now consumes this to run deployment hooks
|
||||
# with the same info as chat, including litellm_params.
|
||||
request_context: Dict[str, Any] = {"input": input}
|
||||
try:
|
||||
request_context.update(response_api_optional_request_params)
|
||||
except Exception:
|
||||
pass
|
||||
# Needed by streaming callbacks/metadata helpers to reconstruct api_base/model_id
|
||||
# but never included in the outbound provider payload.
|
||||
request_context["litellm_params"] = dict(litellm_params)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -2227,6 +2263,8 @@ class BaseLLMHTTPHandler:
|
|||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
# Return the streaming iterator
|
||||
|
|
@ -2237,6 +2275,8 @@ class BaseLLMHTTPHandler:
|
|||
responses_api_provider_config=responses_api_provider_config,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_context,
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
else:
|
||||
# For non-streaming, proceed as before
|
||||
|
|
@ -3526,6 +3566,174 @@ class BaseLLMHTTPHandler:
|
|||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def compact_response_api_handler(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, "ResponseInputParam"],
|
||||
responses_api_provider_config: BaseResponsesAPIConfig,
|
||||
response_api_optional_request_params: Dict,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
custom_llm_provider: Optional[str],
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
|
||||
_is_async: bool = False,
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]:
|
||||
"""
|
||||
Handler for the compact responses API.
|
||||
"""
|
||||
if _is_async:
|
||||
return self.async_compact_response_api_handler(
|
||||
model=model,
|
||||
input=input,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
response_api_optional_request_params=response_api_optional_request_params,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
sync_httpx_client = _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
headers = responses_api_provider_config.validate_environment(
|
||||
headers=extra_headers or {}, model=model, litellm_params=litellm_params
|
||||
)
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = responses_api_provider_config.get_complete_url(
|
||||
api_base=litellm_params.api_base,
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, data = responses_api_provider_config.transform_compact_response_api_request(
|
||||
model=model,
|
||||
input=input,
|
||||
response_api_optional_request_params=response_api_optional_request_params,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=url, headers=headers, json=data, timeout=timeout
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=responses_api_provider_config,
|
||||
)
|
||||
|
||||
return responses_api_provider_config.transform_compact_response_api_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_compact_response_api_handler(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, "ResponseInputParam"],
|
||||
responses_api_provider_config: BaseResponsesAPIConfig,
|
||||
response_api_optional_request_params: Dict,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
custom_llm_provider: Optional[str],
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
|
||||
_is_async: bool = False,
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
) -> ResponsesAPIResponse:
|
||||
"""
|
||||
Async version of the compact response API handler.
|
||||
"""
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
verbose_logger.debug(
|
||||
f"Creating HTTP client for compact_response with shared_session: {id(shared_session) if shared_session else None}"
|
||||
)
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
shared_session=shared_session,
|
||||
)
|
||||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
headers = responses_api_provider_config.validate_environment(
|
||||
headers=extra_headers or {}, model=model, litellm_params=litellm_params
|
||||
)
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base = responses_api_provider_config.get_complete_url(
|
||||
api_base=litellm_params.api_base,
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
url, data = responses_api_provider_config.transform_compact_response_api_request(
|
||||
model=model,
|
||||
input=input,
|
||||
response_api_optional_request_params=response_api_optional_request_params,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=url, headers=headers, json=data, timeout=timeout
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=responses_api_provider_config,
|
||||
)
|
||||
|
||||
return responses_api_provider_config.transform_compact_response_api_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def list_files(self):
|
||||
"""
|
||||
Lists all files
|
||||
|
|
@ -8288,4 +8496,4 @@ class BaseLLMHTTPHandler:
|
|||
return skills_api_provider_config.transform_delete_skill_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
23
litellm/llms/gigachat/__init__.py
Normal file
23
litellm/llms/gigachat/__init__.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
"""
|
||||
GigaChat Provider for LiteLLM
|
||||
|
||||
GigaChat is Sber AI's large language model (Russia's leading LLM).
|
||||
Supports:
|
||||
- Chat completions (sync/async)
|
||||
- Streaming (sync/async)
|
||||
- Function calling / Tools
|
||||
- Structured output via JSON schema (emulated through function calls)
|
||||
- Image input (base64 and URL)
|
||||
- Embeddings
|
||||
|
||||
API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/overview
|
||||
"""
|
||||
|
||||
from .chat.transformation import GigaChatConfig, GigaChatError
|
||||
from .embedding.transformation import GigaChatEmbeddingConfig
|
||||
|
||||
__all__ = [
|
||||
"GigaChatConfig",
|
||||
"GigaChatEmbeddingConfig",
|
||||
"GigaChatError",
|
||||
]
|
||||
241
litellm/llms/gigachat/authenticator.py
Normal file
241
litellm/llms/gigachat/authenticator.py
Normal file
|
|
@ -0,0 +1,241 @@
|
|||
"""
|
||||
GigaChat OAuth Authenticator
|
||||
|
||||
Handles OAuth 2.0 token management for GigaChat API.
|
||||
Based on official GigaChat SDK authentication flow.
|
||||
"""
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import InMemoryCache
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
HTTPHandler,
|
||||
_get_httpx_client,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
# GigaChat OAuth endpoint
|
||||
GIGACHAT_AUTH_URL = "https://ngw.devices.sberbank.ru:9443/api/v2/oauth"
|
||||
|
||||
# Default scope for personal API access
|
||||
GIGACHAT_SCOPE = "GIGACHAT_API_PERS"
|
||||
|
||||
# Token expiry buffer in milliseconds (refresh token 60s before expiry)
|
||||
TOKEN_EXPIRY_BUFFER_MS = 60000
|
||||
|
||||
# Cache for access tokens
|
||||
_token_cache = InMemoryCache()
|
||||
|
||||
|
||||
class GigaChatAuthError(BaseLLMException):
|
||||
"""GigaChat authentication error."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
def _get_credentials() -> Optional[str]:
|
||||
"""Get GigaChat credentials from environment."""
|
||||
return get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY")
|
||||
|
||||
|
||||
def _get_auth_url() -> str:
|
||||
"""Get GigaChat auth URL from environment or use default."""
|
||||
return get_secret_str("GIGACHAT_AUTH_URL") or GIGACHAT_AUTH_URL
|
||||
|
||||
|
||||
def _get_scope() -> str:
|
||||
"""Get GigaChat scope from environment or use default."""
|
||||
return get_secret_str("GIGACHAT_SCOPE") or GIGACHAT_SCOPE
|
||||
|
||||
|
||||
def _get_http_client() -> HTTPHandler:
|
||||
"""Get cached httpx client with SSL verification disabled."""
|
||||
return _get_httpx_client(params={"ssl_verify": False})
|
||||
|
||||
|
||||
def get_access_token(
|
||||
credentials: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
auth_url: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get valid access token, using cache if available.
|
||||
|
||||
Args:
|
||||
credentials: Base64-encoded credentials (client_id:client_secret)
|
||||
scope: API scope (GIGACHAT_API_PERS, GIGACHAT_API_CORP, etc.)
|
||||
auth_url: OAuth endpoint URL
|
||||
|
||||
Returns:
|
||||
Access token string
|
||||
|
||||
Raises:
|
||||
GigaChatAuthError: If authentication fails
|
||||
"""
|
||||
credentials = credentials or _get_credentials()
|
||||
if not credentials:
|
||||
raise GigaChatAuthError(
|
||||
status_code=401,
|
||||
message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.",
|
||||
)
|
||||
|
||||
scope = scope or _get_scope()
|
||||
auth_url = auth_url or _get_auth_url()
|
||||
|
||||
# Check cache
|
||||
cache_key = f"gigachat_token:{credentials[:16]}"
|
||||
cached = _token_cache.get_cache(cache_key)
|
||||
if cached:
|
||||
token, expires_at = cached
|
||||
# Check if token is still valid (with buffer)
|
||||
if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS:
|
||||
verbose_logger.debug("Using cached GigaChat access token")
|
||||
return token
|
||||
|
||||
# Request new token
|
||||
token, expires_at = _request_token_sync(credentials, scope, auth_url)
|
||||
|
||||
# Cache token
|
||||
ttl_seconds = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
|
||||
if ttl_seconds > 0:
|
||||
_token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds)
|
||||
|
||||
return token
|
||||
|
||||
|
||||
async def get_access_token_async(
|
||||
credentials: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
auth_url: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Async version of get_access_token."""
|
||||
credentials = credentials or _get_credentials()
|
||||
if not credentials:
|
||||
raise GigaChatAuthError(
|
||||
status_code=401,
|
||||
message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.",
|
||||
)
|
||||
|
||||
scope = scope or _get_scope()
|
||||
auth_url = auth_url or _get_auth_url()
|
||||
|
||||
# Check cache
|
||||
cache_key = f"gigachat_token:{credentials[:16]}"
|
||||
cached = _token_cache.get_cache(cache_key)
|
||||
if cached:
|
||||
token, expires_at = cached
|
||||
if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS:
|
||||
verbose_logger.debug("Using cached GigaChat access token")
|
||||
return token
|
||||
|
||||
# Request new token
|
||||
token, expires_at = await _request_token_async(credentials, scope, auth_url)
|
||||
|
||||
# Cache token
|
||||
ttl_seconds = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000)
|
||||
if ttl_seconds > 0:
|
||||
_token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds)
|
||||
|
||||
return token
|
||||
|
||||
|
||||
def _request_token_sync(
|
||||
credentials: str,
|
||||
scope: str,
|
||||
auth_url: str,
|
||||
) -> Tuple[str, int]:
|
||||
"""
|
||||
Request new access token from GigaChat OAuth endpoint (sync).
|
||||
|
||||
Returns:
|
||||
Tuple of (access_token, expires_at_ms)
|
||||
"""
|
||||
headers = {
|
||||
"Authorization": f"Basic {credentials}",
|
||||
"RqUID": str(uuid.uuid4()),
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
}
|
||||
data = {"scope": scope}
|
||||
|
||||
verbose_logger.debug(f"Requesting GigaChat access token from {auth_url}")
|
||||
|
||||
try:
|
||||
client = _get_http_client()
|
||||
response = client.post(auth_url, headers=headers, data=data, timeout=30)
|
||||
response.raise_for_status()
|
||||
return _parse_token_response(response)
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise GigaChatAuthError(
|
||||
status_code=e.response.status_code,
|
||||
message=f"GigaChat authentication failed: {e.response.text}",
|
||||
)
|
||||
except httpx.RequestError as e:
|
||||
raise GigaChatAuthError(
|
||||
status_code=500,
|
||||
message=f"GigaChat authentication request failed: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
async def _request_token_async(
|
||||
credentials: str,
|
||||
scope: str,
|
||||
auth_url: str,
|
||||
) -> Tuple[str, int]:
|
||||
"""Async version of _request_token_sync."""
|
||||
headers = {
|
||||
"Authorization": f"Basic {credentials}",
|
||||
"RqUID": str(uuid.uuid4()),
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
}
|
||||
data = {"scope": scope}
|
||||
|
||||
verbose_logger.debug(f"Requesting GigaChat access token from {auth_url}")
|
||||
|
||||
try:
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GIGACHAT,
|
||||
params={"ssl_verify": False},
|
||||
)
|
||||
response = await client.post(auth_url, headers=headers, data=data, timeout=30)
|
||||
response.raise_for_status()
|
||||
return _parse_token_response(response)
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise GigaChatAuthError(
|
||||
status_code=e.response.status_code,
|
||||
message=f"GigaChat authentication failed: {e.response.text}",
|
||||
)
|
||||
except httpx.RequestError as e:
|
||||
raise GigaChatAuthError(
|
||||
status_code=500,
|
||||
message=f"GigaChat authentication request failed: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
def _parse_token_response(response: httpx.Response) -> Tuple[str, int]:
|
||||
"""Parse OAuth token response."""
|
||||
data = response.json()
|
||||
|
||||
# GigaChat returns either 'tok'/'exp' or 'access_token'/'expires_at'
|
||||
access_token = data.get("tok") or data.get("access_token")
|
||||
expires_at = data.get("exp") or data.get("expires_at")
|
||||
|
||||
if not access_token:
|
||||
raise GigaChatAuthError(
|
||||
status_code=500,
|
||||
message=f"Invalid token response: {data}",
|
||||
)
|
||||
|
||||
# expires_at is in milliseconds
|
||||
if isinstance(expires_at, str):
|
||||
expires_at = int(expires_at)
|
||||
|
||||
verbose_logger.debug("GigaChat access token obtained successfully")
|
||||
return access_token, expires_at
|
||||
12
litellm/llms/gigachat/chat/__init__.py
Normal file
12
litellm/llms/gigachat/chat/__init__.py
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
"""
|
||||
GigaChat Chat Module
|
||||
"""
|
||||
|
||||
from .transformation import GigaChatConfig, GigaChatError
|
||||
from .streaming import GigaChatModelResponseIterator
|
||||
|
||||
__all__ = [
|
||||
"GigaChatConfig",
|
||||
"GigaChatError",
|
||||
"GigaChatModelResponseIterator",
|
||||
]
|
||||
134
litellm/llms/gigachat/chat/streaming.py
Normal file
134
litellm/llms/gigachat/chat/streaming.py
Normal file
|
|
@ -0,0 +1,134 @@
|
|||
"""
|
||||
GigaChat Streaming Response Handler
|
||||
"""
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
from litellm.types.llms.openai import ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk
|
||||
from litellm.types.utils import GenericStreamingChunk
|
||||
|
||||
|
||||
class GigaChatModelResponseIterator:
|
||||
"""Iterator for GigaChat streaming responses."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
streaming_response: Any,
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
):
|
||||
self.streaming_response = streaming_response
|
||||
self.response_iterator = self.streaming_response
|
||||
self.json_mode = json_mode
|
||||
|
||||
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
|
||||
"""Parse a single streaming chunk from GigaChat."""
|
||||
text = ""
|
||||
tool_use: Optional[ChatCompletionToolCallChunk] = None
|
||||
is_finished = False
|
||||
finish_reason: Optional[str] = None
|
||||
|
||||
choices = chunk.get("choices", [])
|
||||
if not choices:
|
||||
return GenericStreamingChunk(
|
||||
text="",
|
||||
tool_use=None,
|
||||
is_finished=False,
|
||||
finish_reason="",
|
||||
usage=None,
|
||||
index=0,
|
||||
)
|
||||
|
||||
choice = choices[0]
|
||||
delta = choice.get("delta", {})
|
||||
finish_reason = choice.get("finish_reason")
|
||||
|
||||
# Extract text content
|
||||
text = delta.get("content", "") or ""
|
||||
|
||||
# Handle function_call in stream
|
||||
if finish_reason == "function_call" and delta.get("function_call"):
|
||||
func_call = delta["function_call"]
|
||||
args = func_call.get("arguments", {})
|
||||
|
||||
if isinstance(args, dict):
|
||||
args = json.dumps(args, ensure_ascii=False)
|
||||
|
||||
tool_use = ChatCompletionToolCallChunk(
|
||||
id=f"call_{uuid.uuid4().hex[:24]}",
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=func_call.get("name", ""),
|
||||
arguments=args,
|
||||
),
|
||||
index=0,
|
||||
)
|
||||
finish_reason = "tool_calls"
|
||||
|
||||
if finish_reason is not None:
|
||||
is_finished = True
|
||||
|
||||
return GenericStreamingChunk(
|
||||
text=text,
|
||||
tool_use=tool_use,
|
||||
is_finished=is_finished,
|
||||
finish_reason=finish_reason or "",
|
||||
usage=None,
|
||||
index=choice.get("index", 0),
|
||||
)
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self) -> GenericStreamingChunk:
|
||||
try:
|
||||
chunk = self.response_iterator.__next__()
|
||||
if isinstance(chunk, str):
|
||||
# Parse SSE format: data: {...}
|
||||
if chunk.startswith("data: "):
|
||||
chunk = chunk[6:]
|
||||
if chunk.strip() == "[DONE]":
|
||||
raise StopIteration
|
||||
try:
|
||||
chunk = json.loads(chunk)
|
||||
except json.JSONDecodeError:
|
||||
return GenericStreamingChunk(
|
||||
text="",
|
||||
tool_use=None,
|
||||
is_finished=False,
|
||||
finish_reason="",
|
||||
usage=None,
|
||||
index=0,
|
||||
)
|
||||
return self.chunk_parser(chunk)
|
||||
except StopIteration:
|
||||
raise
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> GenericStreamingChunk:
|
||||
try:
|
||||
chunk = await self.response_iterator.__anext__()
|
||||
if isinstance(chunk, str):
|
||||
# Parse SSE format
|
||||
if chunk.startswith("data: "):
|
||||
chunk = chunk[6:]
|
||||
if chunk.strip() == "[DONE]":
|
||||
raise StopAsyncIteration
|
||||
try:
|
||||
chunk = json.loads(chunk)
|
||||
except json.JSONDecodeError:
|
||||
return GenericStreamingChunk(
|
||||
text="",
|
||||
tool_use=None,
|
||||
is_finished=False,
|
||||
finish_reason="",
|
||||
usage=None,
|
||||
index=0,
|
||||
)
|
||||
return self.chunk_parser(chunk)
|
||||
except StopAsyncIteration:
|
||||
raise
|
||||
473
litellm/llms/gigachat/chat/transformation.py
Normal file
473
litellm/llms/gigachat/chat/transformation.py
Normal file
|
|
@ -0,0 +1,473 @@
|
|||
"""
|
||||
GigaChat Chat Transformation
|
||||
|
||||
Transforms OpenAI-format requests to GigaChat format and back.
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
|
||||
from ..authenticator import get_access_token
|
||||
from ..file_handler import upload_file_sync
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
# GigaChat API endpoint
|
||||
GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1"
|
||||
|
||||
|
||||
class GigaChatError(BaseLLMException):
|
||||
"""GigaChat API error."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class GigaChatConfig(BaseConfig):
|
||||
"""
|
||||
Configuration class for GigaChat API.
|
||||
|
||||
GigaChat is Sber's (Russia's largest bank) LLM API.
|
||||
|
||||
Supported parameters:
|
||||
temperature: Sampling temperature (0-2, default 0.87)
|
||||
top_p: Nucleus sampling parameter
|
||||
max_tokens: Maximum tokens to generate
|
||||
repetition_penalty: Repetition penalty factor
|
||||
profanity_check: Enable content filtering
|
||||
stream: Enable streaming
|
||||
"""
|
||||
|
||||
temperature: Optional[float] = None
|
||||
top_p: Optional[float] = None
|
||||
max_tokens: Optional[int] = None
|
||||
repetition_penalty: Optional[float] = None
|
||||
profanity_check: Optional[bool] = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
temperature: Optional[float] = None,
|
||||
top_p: Optional[float] = None,
|
||||
max_tokens: Optional[int] = None,
|
||||
repetition_penalty: Optional[float] = None,
|
||||
profanity_check: Optional[bool] = None,
|
||||
) -> None:
|
||||
locals_ = locals().copy()
|
||||
for key, value in locals_.items():
|
||||
if key != "self" and value is not None:
|
||||
setattr(self.__class__, key, value)
|
||||
# Instance variables for current request context
|
||||
self._current_credentials: Optional[str] = None
|
||||
self._current_api_base: Optional[str] = None
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""Get complete API URL for chat completions."""
|
||||
base = api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL
|
||||
return f"{base}/chat/completions"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Set up headers with OAuth token.
|
||||
"""
|
||||
# Get access token
|
||||
credentials = api_key or get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY")
|
||||
access_token = get_access_token(credentials=credentials)
|
||||
|
||||
# Store credentials for image uploads
|
||||
self._current_credentials = credentials
|
||||
self._current_api_base = api_base
|
||||
|
||||
headers["Authorization"] = f"Bearer {access_token}"
|
||||
headers["Content-Type"] = "application/json"
|
||||
headers["Accept"] = "application/json"
|
||||
|
||||
return headers
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
"""Return list of supported OpenAI parameters."""
|
||||
return [
|
||||
"stream",
|
||||
"temperature",
|
||||
"top_p",
|
||||
"max_tokens",
|
||||
"max_completion_tokens",
|
||||
"stop",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"functions",
|
||||
"function_call",
|
||||
"response_format",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""Map OpenAI parameters to GigaChat parameters."""
|
||||
for param, value in non_default_params.items():
|
||||
if param == "stream":
|
||||
optional_params["stream"] = value
|
||||
elif param == "temperature":
|
||||
# GigaChat: temperature 0 means use top_p=0 instead
|
||||
if value == 0:
|
||||
optional_params["top_p"] = 0
|
||||
else:
|
||||
optional_params["temperature"] = value
|
||||
elif param == "top_p":
|
||||
optional_params["top_p"] = value
|
||||
elif param in ("max_tokens", "max_completion_tokens"):
|
||||
optional_params["max_tokens"] = value
|
||||
elif param == "stop":
|
||||
# GigaChat doesn't support stop sequences
|
||||
pass
|
||||
elif param == "tools":
|
||||
# Convert tools to functions format
|
||||
optional_params["functions"] = self._convert_tools_to_functions(value)
|
||||
elif param == "tool_choice":
|
||||
if isinstance(value, dict) and value.get("function"):
|
||||
optional_params["function_call"] = {"name": value["function"]["name"]}
|
||||
elif value == "auto":
|
||||
pass # Default behavior
|
||||
elif value == "required":
|
||||
# GigaChat doesn't have 'required', handled differently
|
||||
pass
|
||||
elif param == "functions":
|
||||
optional_params["functions"] = value
|
||||
elif param == "function_call":
|
||||
optional_params["function_call"] = value
|
||||
elif param == "response_format":
|
||||
# Handle structured output via function calling
|
||||
if value.get("type") == "json_schema":
|
||||
json_schema = value.get("json_schema", {})
|
||||
schema_name = json_schema.get("name", "structured_output")
|
||||
schema = json_schema.get("schema", {})
|
||||
|
||||
function_def = {
|
||||
"name": schema_name,
|
||||
"description": f"Output structured response: {schema_name}",
|
||||
"parameters": schema,
|
||||
}
|
||||
|
||||
if "functions" not in optional_params:
|
||||
optional_params["functions"] = []
|
||||
optional_params["functions"].append(function_def)
|
||||
optional_params["function_call"] = {"name": schema_name}
|
||||
optional_params["_structured_output"] = True
|
||||
|
||||
return optional_params
|
||||
|
||||
def _convert_tools_to_functions(self, tools: List[dict]) -> List[dict]:
|
||||
"""Convert OpenAI tools format to GigaChat functions format."""
|
||||
functions = []
|
||||
for tool in tools:
|
||||
if tool.get("type") == "function":
|
||||
func = tool.get("function", {})
|
||||
functions.append({
|
||||
"name": func.get("name", ""),
|
||||
"description": func.get("description", ""),
|
||||
"parameters": func.get("parameters", {}),
|
||||
})
|
||||
return functions
|
||||
|
||||
def _upload_image(self, image_url: str) -> Optional[str]:
|
||||
"""
|
||||
Upload image to GigaChat and return file_id.
|
||||
|
||||
Args:
|
||||
image_url: URL or base64 data URL of the image
|
||||
|
||||
Returns:
|
||||
file_id string or None if upload failed
|
||||
"""
|
||||
try:
|
||||
return upload_file_sync(
|
||||
image_url=image_url,
|
||||
credentials=self._current_credentials,
|
||||
api_base=self._current_api_base,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Failed to upload image: {e}")
|
||||
return None
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""Transform OpenAI request to GigaChat format."""
|
||||
# Transform messages
|
||||
giga_messages = self._transform_messages(messages)
|
||||
|
||||
# Build request
|
||||
request_data = {
|
||||
"model": model.replace("gigachat/", ""),
|
||||
"messages": giga_messages,
|
||||
}
|
||||
|
||||
# Add optional params
|
||||
for key in ["temperature", "top_p", "max_tokens", "stream",
|
||||
"repetition_penalty", "profanity_check"]:
|
||||
if key in optional_params:
|
||||
request_data[key] = optional_params[key]
|
||||
|
||||
# Add functions if present
|
||||
if "functions" in optional_params:
|
||||
request_data["functions"] = optional_params["functions"]
|
||||
if "function_call" in optional_params:
|
||||
request_data["function_call"] = optional_params["function_call"]
|
||||
|
||||
return request_data
|
||||
|
||||
def _transform_messages(self, messages: List[AllMessageValues]) -> List[dict]:
|
||||
"""Transform OpenAI messages to GigaChat format."""
|
||||
transformed = []
|
||||
|
||||
for i, msg in enumerate(messages):
|
||||
message = dict(msg)
|
||||
|
||||
# Remove unsupported fields
|
||||
message.pop("name", None)
|
||||
|
||||
# Transform roles
|
||||
role = message.get("role", "user")
|
||||
if role == "developer":
|
||||
message["role"] = "system"
|
||||
elif role == "system" and i > 0:
|
||||
# GigaChat only allows system message as first message
|
||||
message["role"] = "user"
|
||||
elif role == "tool":
|
||||
message["role"] = "function"
|
||||
content = message.get("content", "")
|
||||
if not isinstance(content, str):
|
||||
message["content"] = json.dumps(content, ensure_ascii=False)
|
||||
|
||||
# Handle None content
|
||||
if message.get("content") is None:
|
||||
message["content"] = ""
|
||||
|
||||
# Handle list content (multimodal) - extract text and images
|
||||
content = message.get("content")
|
||||
if isinstance(content, list):
|
||||
texts = []
|
||||
attachments = []
|
||||
for part in content:
|
||||
if isinstance(part, dict):
|
||||
if part.get("type") == "text":
|
||||
texts.append(part.get("text", ""))
|
||||
elif part.get("type") == "image_url":
|
||||
# Extract image URL and upload to GigaChat
|
||||
image_url = part.get("image_url", {})
|
||||
if isinstance(image_url, str):
|
||||
url = image_url
|
||||
else:
|
||||
url = image_url.get("url", "")
|
||||
if url:
|
||||
file_id = self._upload_image(url)
|
||||
if file_id:
|
||||
attachments.append(file_id)
|
||||
message["content"] = "\n".join(texts) if texts else ""
|
||||
if attachments:
|
||||
message["attachments"] = attachments
|
||||
|
||||
# Transform tool_calls to function_call
|
||||
tool_calls = message.get("tool_calls")
|
||||
if tool_calls and isinstance(tool_calls, list) and len(tool_calls) > 0:
|
||||
tool_call = tool_calls[0]
|
||||
func = tool_call.get("function", {})
|
||||
args = func.get("arguments", "{}")
|
||||
if isinstance(args, str):
|
||||
try:
|
||||
args = json.loads(args)
|
||||
except json.JSONDecodeError:
|
||||
args = {}
|
||||
message["function_call"] = {
|
||||
"name": func.get("name", ""),
|
||||
"arguments": args,
|
||||
}
|
||||
message.pop("tool_calls", None)
|
||||
|
||||
transformed.append(message)
|
||||
|
||||
# Collapse consecutive user messages
|
||||
return self._collapse_user_messages(transformed)
|
||||
|
||||
def _collapse_user_messages(self, messages: List[dict]) -> List[dict]:
|
||||
"""Collapse consecutive user messages into one."""
|
||||
collapsed: List[dict] = []
|
||||
prev_user_msg: Optional[dict] = None
|
||||
content_parts: List[str] = []
|
||||
|
||||
for msg in messages:
|
||||
if msg.get("role") == "user" and prev_user_msg is not None:
|
||||
content_parts.append(msg.get("content", ""))
|
||||
else:
|
||||
if content_parts and prev_user_msg:
|
||||
prev_user_msg["content"] = "\n".join(
|
||||
[prev_user_msg.get("content", "")] + content_parts
|
||||
)
|
||||
content_parts = []
|
||||
collapsed.append(msg)
|
||||
prev_user_msg = msg if msg.get("role") == "user" else None
|
||||
|
||||
if content_parts and prev_user_msg:
|
||||
prev_user_msg["content"] = "\n".join(
|
||||
[prev_user_msg.get("content", "")] + content_parts
|
||||
)
|
||||
|
||||
return collapsed
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
"""Transform GigaChat response to OpenAI format."""
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
except Exception:
|
||||
raise GigaChatError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"Invalid JSON response: {raw_response.text}",
|
||||
)
|
||||
|
||||
is_structured_output = optional_params.get("_structured_output", False)
|
||||
|
||||
choices = []
|
||||
for choice in response_json.get("choices", []):
|
||||
message_data = choice.get("message", {})
|
||||
finish_reason = choice.get("finish_reason", "stop")
|
||||
|
||||
# Transform function_call to tool_calls or content
|
||||
if finish_reason == "function_call" and message_data.get("function_call"):
|
||||
func_call = message_data["function_call"]
|
||||
args = func_call.get("arguments", {})
|
||||
|
||||
if is_structured_output:
|
||||
# Convert to content for structured output
|
||||
if isinstance(args, dict):
|
||||
content = json.dumps(args, ensure_ascii=False)
|
||||
else:
|
||||
content = str(args)
|
||||
message_data["content"] = content
|
||||
message_data.pop("function_call", None)
|
||||
message_data.pop("functions_state_id", None)
|
||||
finish_reason = "stop"
|
||||
else:
|
||||
# Convert to tool_calls format
|
||||
if isinstance(args, dict):
|
||||
args = json.dumps(args, ensure_ascii=False)
|
||||
message_data["tool_calls"] = [{
|
||||
"id": f"call_{uuid.uuid4().hex[:24]}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": func_call.get("name", ""),
|
||||
"arguments": args,
|
||||
}
|
||||
}]
|
||||
message_data.pop("function_call", None)
|
||||
finish_reason = "tool_calls"
|
||||
|
||||
# Clean up GigaChat-specific fields
|
||||
message_data.pop("functions_state_id", None)
|
||||
|
||||
choices.append(
|
||||
Choices(
|
||||
index=choice.get("index", 0),
|
||||
message=Message(
|
||||
role=message_data.get("role", "assistant"),
|
||||
content=message_data.get("content"),
|
||||
tool_calls=message_data.get("tool_calls"),
|
||||
),
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
)
|
||||
|
||||
# Build usage
|
||||
usage_data = response_json.get("usage", {})
|
||||
usage = Usage(
|
||||
prompt_tokens=usage_data.get("prompt_tokens", 0),
|
||||
completion_tokens=usage_data.get("completion_tokens", 0),
|
||||
total_tokens=usage_data.get("total_tokens", 0),
|
||||
)
|
||||
|
||||
model_response.id = response_json.get("id", f"chatcmpl-{uuid.uuid4().hex[:12]}")
|
||||
model_response.created = response_json.get("created", int(time.time()))
|
||||
model_response.model = model
|
||||
model_response.choices = choices # type: ignore
|
||||
setattr(model_response, "usage", usage)
|
||||
|
||||
return model_response
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Union[dict, httpx.Headers],
|
||||
) -> BaseLLMException:
|
||||
"""Return GigaChat error class."""
|
||||
return GigaChatError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def get_model_response_iterator(
|
||||
self,
|
||||
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
):
|
||||
"""Return streaming response iterator."""
|
||||
from .streaming import GigaChatModelResponseIterator
|
||||
|
||||
return GigaChatModelResponseIterator(
|
||||
streaming_response=streaming_response,
|
||||
sync_stream=sync_stream,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
7
litellm/llms/gigachat/embedding/__init__.py
Normal file
7
litellm/llms/gigachat/embedding/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
GigaChat Embedding Module
|
||||
"""
|
||||
|
||||
from .transformation import GigaChatEmbeddingConfig
|
||||
|
||||
__all__ = ["GigaChatEmbeddingConfig"]
|
||||
212
litellm/llms/gigachat/embedding/transformation.py
Normal file
212
litellm/llms/gigachat/embedding/transformation.py
Normal file
|
|
@ -0,0 +1,212 @@
|
|||
"""
|
||||
GigaChat Embedding Transformation
|
||||
|
||||
Transforms OpenAI /v1/embeddings format to GigaChat format.
|
||||
API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/rest/post-embeddings
|
||||
"""
|
||||
|
||||
import types
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm import LlmProviders
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
from ..authenticator import get_access_token
|
||||
|
||||
# GigaChat API endpoint
|
||||
GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1"
|
||||
|
||||
|
||||
class GigaChatEmbeddingError(BaseLLMException):
|
||||
"""GigaChat Embedding API error."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
|
||||
"""
|
||||
Configuration class for GigaChat Embeddings API.
|
||||
|
||||
GigaChat embeddings endpoint: POST /api/v1/embeddings
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def get_config(cls):
|
||||
return {
|
||||
k: v
|
||||
for k, v in cls.__dict__.items()
|
||||
if not k.startswith("__")
|
||||
and not isinstance(
|
||||
v,
|
||||
(
|
||||
types.FunctionType,
|
||||
types.BuiltinFunctionType,
|
||||
classmethod,
|
||||
staticmethod,
|
||||
),
|
||||
)
|
||||
and v is not None
|
||||
}
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
"""GigaChat embeddings don't support additional parameters."""
|
||||
return []
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""Map OpenAI params to GigaChat format (no special mapping needed)."""
|
||||
return optional_params
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
) -> Tuple[str, Optional[str], Optional[str]]:
|
||||
"""
|
||||
Returns provider info for GigaChat.
|
||||
|
||||
Returns:
|
||||
Tuple of (custom_llm_provider, api_base, dynamic_api_key)
|
||||
"""
|
||||
api_base = api_base or GIGACHAT_BASE_URL
|
||||
return LlmProviders.GIGACHAT.value, api_base, api_key
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""Get the complete URL for embeddings endpoint."""
|
||||
base = api_base or GIGACHAT_BASE_URL
|
||||
return f"{base}/embeddings"
|
||||
|
||||
def transform_embedding_request(
|
||||
self,
|
||||
model: str,
|
||||
input: AllEmbeddingInputValues,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform OpenAI embedding request to GigaChat format.
|
||||
|
||||
GigaChat format:
|
||||
{
|
||||
"model": "Embeddings",
|
||||
"input": ["text1", "text2", ...]
|
||||
}
|
||||
"""
|
||||
# Normalize input to list
|
||||
if isinstance(input, str):
|
||||
input_list: list = [input]
|
||||
elif isinstance(input, list):
|
||||
input_list = input
|
||||
else:
|
||||
input_list = [input]
|
||||
|
||||
# Remove gigachat/ prefix from model if present
|
||||
if model.startswith("gigachat/"):
|
||||
model = model[9:]
|
||||
|
||||
return {
|
||||
"model": model,
|
||||
"input": input_list,
|
||||
}
|
||||
|
||||
def transform_embedding_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: EmbeddingResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
Transform GigaChat embedding response to OpenAI format.
|
||||
|
||||
GigaChat returns:
|
||||
{
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "embedding": [...], "index": 0, "usage": {...}}],
|
||||
"model": "Embeddings"
|
||||
}
|
||||
"""
|
||||
response_json = raw_response.json()
|
||||
|
||||
# Log response
|
||||
logging_obj.post_call(
|
||||
input=request_data.get("input"),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
original_response=response_json,
|
||||
)
|
||||
|
||||
# Calculate total tokens from individual embeddings
|
||||
total_tokens = 0
|
||||
if "data" in response_json:
|
||||
for emb in response_json["data"]:
|
||||
if "usage" in emb and "prompt_tokens" in emb["usage"]:
|
||||
total_tokens += emb["usage"]["prompt_tokens"]
|
||||
# Remove usage from individual embeddings (not part of OpenAI format)
|
||||
if "usage" in emb:
|
||||
del emb["usage"]
|
||||
|
||||
# Set overall usage
|
||||
response_json["usage"] = {
|
||||
"prompt_tokens": total_tokens,
|
||||
"total_tokens": total_tokens,
|
||||
}
|
||||
|
||||
return EmbeddingResponse(**response_json)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Set up headers with OAuth token for GigaChat.
|
||||
"""
|
||||
# Get access token via OAuth
|
||||
access_token = get_access_token(api_key)
|
||||
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
}
|
||||
return {**default_headers, **headers}
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
"""Return GigaChat-specific error class."""
|
||||
return GigaChatEmbeddingError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
)
|
||||
211
litellm/llms/gigachat/file_handler.py
Normal file
211
litellm/llms/gigachat/file_handler.py
Normal file
|
|
@ -0,0 +1,211 @@
|
|||
"""
|
||||
GigaChat File Handler
|
||||
|
||||
Handles file uploads to GigaChat API for image processing.
|
||||
GigaChat requires files to be uploaded first, then referenced by file_id.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import re
|
||||
import uuid
|
||||
from typing import Dict, Optional, Tuple
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
_get_httpx_client,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
from .authenticator import get_access_token, get_access_token_async
|
||||
|
||||
# GigaChat API endpoint
|
||||
GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1"
|
||||
|
||||
# Simple in-memory cache for file IDs
|
||||
_file_cache: Dict[str, str] = {}
|
||||
|
||||
|
||||
def _get_url_hash(url: str) -> str:
|
||||
"""Generate hash for URL to use as cache key."""
|
||||
return hashlib.sha256(url.encode()).hexdigest()
|
||||
|
||||
|
||||
def _parse_data_url(data_url: str) -> Optional[Tuple[bytes, str, str]]:
|
||||
"""
|
||||
Parse data URL (base64 image).
|
||||
|
||||
Returns:
|
||||
Tuple of (content_bytes, content_type, extension) or None
|
||||
"""
|
||||
match = re.match(r"data:([^;]+);base64,(.+)", data_url)
|
||||
if not match:
|
||||
return None
|
||||
|
||||
content_type = match.group(1)
|
||||
base64_data = match.group(2)
|
||||
content_bytes = base64.b64decode(base64_data)
|
||||
ext = content_type.split("/")[-1].split(";")[0] or "jpg"
|
||||
|
||||
return content_bytes, content_type, ext
|
||||
|
||||
|
||||
def _download_image_sync(url: str) -> Tuple[bytes, str, str]:
|
||||
"""Download image from URL synchronously."""
|
||||
client = _get_httpx_client(params={"ssl_verify": False})
|
||||
response = client.get(url)
|
||||
response.raise_for_status()
|
||||
|
||||
content_type = response.headers.get("content-type", "image/jpeg")
|
||||
ext = content_type.split("/")[-1].split(";")[0] or "jpg"
|
||||
|
||||
return response.content, content_type, ext
|
||||
|
||||
|
||||
async def _download_image_async(url: str) -> Tuple[bytes, str, str]:
|
||||
"""Download image from URL asynchronously."""
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GIGACHAT,
|
||||
params={"ssl_verify": False},
|
||||
)
|
||||
response = await client.get(url)
|
||||
response.raise_for_status()
|
||||
|
||||
content_type = response.headers.get("content-type", "image/jpeg")
|
||||
ext = content_type.split("/")[-1].split(";")[0] or "jpg"
|
||||
|
||||
return response.content, content_type, ext
|
||||
|
||||
|
||||
def upload_file_sync(
|
||||
image_url: str,
|
||||
credentials: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Upload file to GigaChat and return file_id (sync).
|
||||
|
||||
Args:
|
||||
image_url: URL or base64 data URL of the image
|
||||
credentials: GigaChat credentials for auth
|
||||
api_base: Optional custom API base URL
|
||||
|
||||
Returns:
|
||||
file_id string or None if upload failed
|
||||
"""
|
||||
url_hash = _get_url_hash(image_url)
|
||||
|
||||
# Check cache
|
||||
if url_hash in _file_cache:
|
||||
verbose_logger.debug(f"Image found in cache: {url_hash[:16]}...")
|
||||
return _file_cache[url_hash]
|
||||
|
||||
try:
|
||||
# Get image data
|
||||
parsed = _parse_data_url(image_url)
|
||||
if parsed:
|
||||
content_bytes, content_type, ext = parsed
|
||||
verbose_logger.debug("Decoded base64 image")
|
||||
else:
|
||||
verbose_logger.debug(f"Downloading image from URL: {image_url[:80]}...")
|
||||
content_bytes, content_type, ext = _download_image_sync(image_url)
|
||||
|
||||
filename = f"{uuid.uuid4()}.{ext}"
|
||||
|
||||
# Get access token
|
||||
access_token = get_access_token(credentials)
|
||||
|
||||
# Upload to GigaChat
|
||||
base_url = api_base or GIGACHAT_BASE_URL
|
||||
upload_url = f"{base_url}/files"
|
||||
|
||||
client = _get_httpx_client(params={"ssl_verify": False})
|
||||
response = client.post(
|
||||
upload_url,
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
files={"file": (filename, content_bytes, content_type)},
|
||||
data={"purpose": "general"},
|
||||
timeout=60,
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
file_id = result.get("id")
|
||||
if file_id:
|
||||
_file_cache[url_hash] = file_id
|
||||
verbose_logger.debug(f"File uploaded successfully, file_id: {file_id}")
|
||||
|
||||
return file_id
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error uploading file to GigaChat: {e}")
|
||||
return None
|
||||
|
||||
|
||||
async def upload_file_async(
|
||||
image_url: str,
|
||||
credentials: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Upload file to GigaChat and return file_id (async).
|
||||
|
||||
Args:
|
||||
image_url: URL or base64 data URL of the image
|
||||
credentials: GigaChat credentials for auth
|
||||
api_base: Optional custom API base URL
|
||||
|
||||
Returns:
|
||||
file_id string or None if upload failed
|
||||
"""
|
||||
url_hash = _get_url_hash(image_url)
|
||||
|
||||
# Check cache
|
||||
if url_hash in _file_cache:
|
||||
verbose_logger.debug(f"Image found in cache: {url_hash[:16]}...")
|
||||
return _file_cache[url_hash]
|
||||
|
||||
try:
|
||||
# Get image data
|
||||
parsed = _parse_data_url(image_url)
|
||||
if parsed:
|
||||
content_bytes, content_type, ext = parsed
|
||||
verbose_logger.debug("Decoded base64 image")
|
||||
else:
|
||||
verbose_logger.debug(f"Downloading image from URL: {image_url[:80]}...")
|
||||
content_bytes, content_type, ext = await _download_image_async(image_url)
|
||||
|
||||
filename = f"{uuid.uuid4()}.{ext}"
|
||||
|
||||
# Get access token
|
||||
access_token = await get_access_token_async(credentials)
|
||||
|
||||
# Upload to GigaChat
|
||||
base_url = api_base or GIGACHAT_BASE_URL
|
||||
upload_url = f"{base_url}/files"
|
||||
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GIGACHAT,
|
||||
params={"ssl_verify": False},
|
||||
)
|
||||
response = await client.post(
|
||||
upload_url,
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
files={"file": (filename, content_bytes, content_type)},
|
||||
data={"purpose": "general"},
|
||||
timeout=60,
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
file_id = result.get("id")
|
||||
if file_id:
|
||||
_file_cache[url_hash] = file_id
|
||||
verbose_logger.debug(f"File uploaded successfully, file_id: {file_id}")
|
||||
|
||||
return file_id
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error uploading file to GigaChat: {e}")
|
||||
return None
|
||||
|
|
@ -500,3 +500,69 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
response._hidden_params["headers"] = raw_response_headers
|
||||
|
||||
return response
|
||||
|
||||
#########################################################
|
||||
########## COMPACT RESPONSE API TRANSFORMATION ##########
|
||||
#########################################################
|
||||
def transform_compact_response_api_request(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, ResponseInputParam],
|
||||
response_api_optional_request_params: Dict,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""
|
||||
Transform the compact response API request into a URL and data
|
||||
|
||||
OpenAI API expects the following request
|
||||
- POST /v1/responses/compact
|
||||
"""
|
||||
url = f"{api_base}/compact"
|
||||
|
||||
input = self._validate_input_param(input)
|
||||
data = dict(
|
||||
ResponsesAPIRequestParams(
|
||||
model=model, input=input, **response_api_optional_request_params
|
||||
)
|
||||
)
|
||||
|
||||
return url, data
|
||||
|
||||
def transform_compact_response_api_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ResponsesAPIResponse:
|
||||
"""
|
||||
Transform the compact response API response into a ResponsesAPIResponse
|
||||
"""
|
||||
try:
|
||||
logging_obj.post_call(
|
||||
original_response=raw_response.text,
|
||||
additional_args={"complete_input_dict": {}},
|
||||
)
|
||||
raw_response_json = raw_response.json()
|
||||
raw_response_json["created_at"] = _safe_convert_created_field(
|
||||
raw_response_json["created_at"]
|
||||
)
|
||||
except Exception:
|
||||
raise OpenAIError(
|
||||
message=raw_response.text, status_code=raw_response.status_code
|
||||
)
|
||||
raw_response_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_response_headers)
|
||||
|
||||
try:
|
||||
response = ResponsesAPIResponse(**raw_response_json)
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct"
|
||||
)
|
||||
response = ResponsesAPIResponse.model_construct(**raw_response_json)
|
||||
|
||||
response._hidden_params["additional_headers"] = processed_headers
|
||||
response._hidden_params["headers"] = raw_response_headers
|
||||
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ Handles Authentication and generating request urls for Vertex AI and Google AI S
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -168,7 +168,6 @@ class VertexBase:
|
|||
)
|
||||
|
||||
def _credentials_from_default_auth(self, scopes):
|
||||
|
||||
import google.auth as google_auth
|
||||
|
||||
return google_auth.default(scopes=scopes)
|
||||
|
|
@ -392,7 +391,7 @@ class VertexBase:
|
|||
Returns
|
||||
token, url
|
||||
"""
|
||||
version: Optional[Literal["v1beta1", "v1"]] = None
|
||||
version: Optional[Literal["v1", "v1beta1"]] = None
|
||||
if custom_llm_provider == "gemini":
|
||||
url, endpoint = _get_gemini_url(
|
||||
mode=mode,
|
||||
|
|
@ -415,7 +414,7 @@ class VertexBase:
|
|||
stream=stream,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_api_version=version,
|
||||
vertex_api_version=cast(Literal["v1", "v1beta1"], version),
|
||||
)
|
||||
|
||||
return self._check_custom_proxy(
|
||||
|
|
|
|||
|
|
@ -2141,6 +2141,49 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
|
||||
client=client,
|
||||
)
|
||||
elif custom_llm_provider == "gigachat":
|
||||
# GigaChat - Sber AI's LLM (Russia)
|
||||
api_key = (
|
||||
api_key
|
||||
or litellm.api_key
|
||||
or litellm.gigachat_key
|
||||
or get_secret("GIGACHAT_API_KEY")
|
||||
or get_secret("GIGACHAT_CREDENTIALS")
|
||||
)
|
||||
|
||||
headers = headers or litellm.headers or {}
|
||||
|
||||
## COMPLETION CALL
|
||||
try:
|
||||
response = base_llm_http_handler.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
headers=headers,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
acompletion=acompletion,
|
||||
logging_obj=logging,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
shared_session=shared_session,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
encoding=_get_encoding(),
|
||||
stream=stream,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
except Exception as e:
|
||||
## LOGGING - log the original exception returned
|
||||
logging.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=str(e),
|
||||
additional_args={"headers": headers},
|
||||
)
|
||||
raise e
|
||||
|
||||
elif custom_llm_provider == "sap":
|
||||
headers = headers or litellm.headers
|
||||
## LOAD CONFIG - if set
|
||||
|
|
@ -5224,6 +5267,28 @@ def embedding( # noqa: PLR0915
|
|||
aembedding=aembedding,
|
||||
litellm_params={},
|
||||
)
|
||||
elif custom_llm_provider == "gigachat":
|
||||
api_key = (
|
||||
api_key
|
||||
or litellm.api_key
|
||||
or litellm.gigachat_key
|
||||
or get_secret_str("GIGACHAT_CREDENTIALS")
|
||||
or get_secret_str("GIGACHAT_API_KEY")
|
||||
)
|
||||
response = base_llm_http_handler.embedding(
|
||||
model=model,
|
||||
input=input,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
logging_obj=logging,
|
||||
timeout=timeout,
|
||||
model_response=EmbeddingResponse(),
|
||||
optional_params=optional_params,
|
||||
client=client,
|
||||
aembedding=aembedding,
|
||||
litellm_params={"ssl_verify": kwargs.get("ssl_verify", None)},
|
||||
)
|
||||
else:
|
||||
raise LiteLLMUnknownProvider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
|
|
|
|||
|
|
@ -15831,6 +15831,68 @@
|
|||
"max_tokens": 8191,
|
||||
"mode": "embedding"
|
||||
},
|
||||
"gigachat/GigaChat-2-Lite": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0.0,
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true
|
||||
},
|
||||
"gigachat/GigaChat-2-Max": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0.0,
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"gigachat/GigaChat-2-Pro": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0.0,
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"gigachat/Embeddings": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 512,
|
||||
"max_tokens": 512,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 1024
|
||||
},
|
||||
"gigachat/Embeddings-2": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 512,
|
||||
"max_tokens": 512,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 1024
|
||||
},
|
||||
"gigachat/EmbeddingsGigaR": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 2560
|
||||
},
|
||||
"google.gemma-3-12b-it": {
|
||||
"input_cost_per_token": 9e-08,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
@ -32092,3 +32154,4 @@
|
|||
"mode": "chat"
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.proxy.utils import get_server_root_path
|
||||
|
||||
router = APIRouter(
|
||||
tags=["mcp"],
|
||||
|
|
@ -381,13 +382,30 @@ async def callback(code: str, state: str):
|
|||
# ------------------------------
|
||||
# Optional .well-known endpoints for MCP + OAuth discovery
|
||||
# ------------------------------
|
||||
@router.get("/.well-known/oauth-protected-resource/{mcp_server_name}/mcp")
|
||||
"""
|
||||
Per SEP-985, the client MUST:
|
||||
1. Try resource_metadata from WWW-Authenticate header (if present)
|
||||
2. Fall back to path-based well-known URI: /.well-known/oauth-protected-resource/{path}
|
||||
(
|
||||
If the resource identifier value contains a path or query component, any terminating slash (/)
|
||||
following the host component MUST be removed before inserting /.well-known/ and the well-known
|
||||
URI path suffix between the host component and the path(include root path) and/or query components.
|
||||
https://datatracker.ietf.org/doc/html/rfc9728#section-3.1)
|
||||
3. Fall back to root-based well-known URI: /.well-known/oauth-protected-resource
|
||||
"""
|
||||
@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp")
|
||||
@router.get("/.well-known/oauth-protected-resource")
|
||||
async def oauth_protected_resource_mcp(
|
||||
request: Request, mcp_server_name: Optional[str] = None
|
||||
):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
# Get the correct base URL considering X-Forwarded-* headers
|
||||
request_base_url = get_request_base_url(request)
|
||||
mcp_server: Optional[MCPServer] = None
|
||||
if mcp_server_name:
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name)
|
||||
return {
|
||||
"authorization_servers": [
|
||||
(
|
||||
|
|
@ -401,14 +419,25 @@ async def oauth_protected_resource_mcp(
|
|||
if mcp_server_name
|
||||
else f"{request_base_url}/mcp"
|
||||
), # this is what Claude will call
|
||||
"scopes_supported": mcp_server.scopes if mcp_server else [],
|
||||
}
|
||||
|
||||
|
||||
@router.get("/.well-known/oauth-authorization-server/{mcp_server_name}")
|
||||
"""
|
||||
https://datatracker.ietf.org/doc/html/rfc8414#section-3.1
|
||||
RFC 8414: Path-aware OAuth discovery
|
||||
If the issuer identifier value contains a path component, any
|
||||
terminating "/" MUST be removed before inserting "/.well-known/" and
|
||||
the well-known URI suffix between the host component and the path(include root path)
|
||||
component.
|
||||
"""
|
||||
@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}")
|
||||
@router.get("/.well-known/oauth-authorization-server")
|
||||
async def oauth_authorization_server_mcp(
|
||||
request: Request, mcp_server_name: Optional[str] = None
|
||||
):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
# Get the correct base URL considering X-Forwarded-* headers
|
||||
request_base_url = get_request_base_url(request)
|
||||
|
||||
|
|
@ -423,16 +452,21 @@ async def oauth_authorization_server_mcp(
|
|||
else f"{request_base_url}/token"
|
||||
)
|
||||
|
||||
mcp_server: Optional[MCPServer] = None
|
||||
if mcp_server_name:
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name)
|
||||
|
||||
return {
|
||||
"issuer": request_base_url, # point to your proxy
|
||||
"authorization_endpoint": authorization_endpoint,
|
||||
"token_endpoint": token_endpoint,
|
||||
"response_types_supported": ["code"],
|
||||
"grant_types_supported": ["authorization_code"],
|
||||
"scopes_supported": mcp_server.scopes if mcp_server else [],
|
||||
"grant_types_supported": ["authorization_code", "refresh_token"],
|
||||
"code_challenge_methods_supported": ["S256"],
|
||||
"token_endpoint_auth_methods_supported": ["client_secret_post"],
|
||||
# Claude expects a registration endpoint, even if we just fake it
|
||||
"registration_endpoint": f"{request_base_url}/{mcp_server_name}/register",
|
||||
"registration_endpoint": f"{request_base_url}/{mcp_server_name}/register" if mcp_server_name else f"{request_base_url}/register",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -660,14 +660,14 @@ class MCPServerManager:
|
|||
"""
|
||||
allowed_mcp_servers = await self.get_allowed_mcp_servers(user_api_key_auth)
|
||||
|
||||
list_tools_result: List[MCPTool] = []
|
||||
verbose_logger.debug("SERVER MANAGER LISTING TOOLS")
|
||||
|
||||
for server_id in allowed_mcp_servers:
|
||||
async def _fetch_server_tools(server_id: str) -> List[MCPTool]:
|
||||
"""Fetch tools from a single server with error handling."""
|
||||
server = self.get_mcp_server_by_id(server_id)
|
||||
if server is None:
|
||||
verbose_logger.warning(f"MCP Server {server_id} not found")
|
||||
continue
|
||||
return []
|
||||
|
||||
# Get server-specific auth header if available
|
||||
server_auth_header = None
|
||||
|
|
@ -685,15 +685,21 @@ class MCPServerManager:
|
|||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
)
|
||||
list_tools_result.extend(tools)
|
||||
verbose_logger.info(
|
||||
f"Successfully fetched {len(tools)} tools from server {server.name}"
|
||||
)
|
||||
return tools
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Failed to list tools from server {server.name}: {str(e)}. Continuing with other servers."
|
||||
)
|
||||
# Continue with other servers instead of failing completely
|
||||
return []
|
||||
|
||||
# Fetch tools from all servers in parallel
|
||||
tasks = [_fetch_server_tools(server_id) for server_id in allowed_mcp_servers]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
# Flatten results into single list
|
||||
list_tools_result: List[MCPTool] = [
|
||||
tool for tools in results for tool in tools
|
||||
]
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully fetched {len(list_tools_result)} tools total from all servers"
|
||||
|
|
@ -2003,6 +2009,9 @@ class MCPServerManager:
|
|||
Note: This now handles prefixed tool names
|
||||
"""
|
||||
for server in self.get_registry().values():
|
||||
if server.auth_type == MCPAuth.oauth2:
|
||||
# Skip OAuth2 servers for now as they may require user-specific tokens
|
||||
continue
|
||||
tools = await self._get_tools_from_server(server)
|
||||
for tool in tools:
|
||||
# The tool.name here is already prefixed from _get_tools_from_server
|
||||
|
|
@ -2284,14 +2293,7 @@ class MCPServerManager:
|
|||
# Check all accessible servers
|
||||
target_server_ids = allowed_server_ids
|
||||
|
||||
# Run health checks concurrently
|
||||
tasks = [self.health_check_server(server_id) for server_id in target_server_ids]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
# Filter out None results (servers that were not found)
|
||||
list_mcp_servers = [server for server in results if server is not None]
|
||||
|
||||
return list_mcp_servers
|
||||
return await self._run_health_checks(target_server_ids)
|
||||
|
||||
async def get_all_allowed_mcp_servers(
|
||||
self,
|
||||
|
|
@ -2306,8 +2308,6 @@ class MCPServerManager:
|
|||
Returns:
|
||||
List of MCP server objects without health status
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
# Get allowed server IDs
|
||||
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
|
||||
|
||||
|
|
@ -2319,40 +2319,56 @@ class MCPServerManager:
|
|||
verbose_logger.warning(f"MCP Server {server_id} not found in registry")
|
||||
continue
|
||||
|
||||
# Build LiteLLM_MCPServerTable without health check
|
||||
mcp_server_table = LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.server_name,
|
||||
alias=server.alias,
|
||||
description=(
|
||||
server.mcp_info.get("description") if server.mcp_info else None
|
||||
),
|
||||
url=server.url,
|
||||
transport=server.transport,
|
||||
auth_type=server.auth_type,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
teams=[],
|
||||
mcp_access_groups=server.access_groups or [],
|
||||
allowed_tools=server.allowed_tools or [],
|
||||
extra_headers=server.extra_headers or [],
|
||||
mcp_info=server.mcp_info,
|
||||
static_headers=server.static_headers,
|
||||
status=None, # No health check performed
|
||||
last_health_check=None, # No health check performed
|
||||
health_check_error=None,
|
||||
command=getattr(server, "command", None),
|
||||
args=getattr(server, "args", None) or [],
|
||||
env=getattr(server, "env", None) or {},
|
||||
authorization_url=server.authorization_url,
|
||||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
)
|
||||
mcp_server_table = self._build_mcp_server_table(server)
|
||||
list_mcp_servers.append(mcp_server_table)
|
||||
|
||||
return list_mcp_servers
|
||||
|
||||
def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable:
|
||||
from datetime import datetime
|
||||
|
||||
return LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.server_name,
|
||||
alias=server.alias,
|
||||
description=(
|
||||
server.mcp_info.get("description") if server.mcp_info else None
|
||||
),
|
||||
url=server.url,
|
||||
transport=server.transport,
|
||||
auth_type=server.auth_type,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
teams=[],
|
||||
mcp_access_groups=server.access_groups or [],
|
||||
allowed_tools=server.allowed_tools or [],
|
||||
extra_headers=server.extra_headers or [],
|
||||
mcp_info=server.mcp_info,
|
||||
static_headers=server.static_headers,
|
||||
status=None, # No health check performed
|
||||
last_health_check=None, # No health check performed
|
||||
health_check_error=None,
|
||||
command=getattr(server, "command", None),
|
||||
args=getattr(server, "args", None) or [],
|
||||
env=getattr(server, "env", None) or {},
|
||||
authorization_url=server.authorization_url,
|
||||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
)
|
||||
|
||||
async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]:
|
||||
"""Return all MCP servers from registry without applying access controls."""
|
||||
|
||||
registry = self.get_registry()
|
||||
if not registry:
|
||||
return []
|
||||
|
||||
servers: List[LiteLLM_MCPServerTable] = []
|
||||
for server in registry.values():
|
||||
servers.append(self._build_mcp_server_table(server))
|
||||
return servers
|
||||
|
||||
async def reload_servers_from_database(self):
|
||||
"""
|
||||
Public method to reload all MCP servers from database into registry.
|
||||
|
|
@ -2360,5 +2376,34 @@ class MCPServerManager:
|
|||
"""
|
||||
await self._add_mcp_servers_from_db_to_in_memory_registry()
|
||||
|
||||
async def get_all_mcp_servers_with_health_unfiltered(
|
||||
self, server_ids: Optional[List[str]] = None
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
"""Return health info for all servers in registry regardless of user access."""
|
||||
|
||||
registry = self.get_registry()
|
||||
if not registry:
|
||||
return []
|
||||
|
||||
if server_ids:
|
||||
target_server_ids = [sid for sid in server_ids if sid in registry]
|
||||
else:
|
||||
target_server_ids = list(registry.keys())
|
||||
|
||||
if not target_server_ids:
|
||||
return []
|
||||
|
||||
return await self._run_health_checks(target_server_ids)
|
||||
|
||||
async def _run_health_checks(
|
||||
self, target_server_ids: List[str]
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
if not target_server_ids:
|
||||
return []
|
||||
|
||||
tasks = [self.health_check_server(server_id) for server_id in target_server_ids]
|
||||
results = await asyncio.gather(*tasks)
|
||||
return [server for server in results if server is not None]
|
||||
|
||||
|
||||
global_mcp_server_manager: MCPServerManager = MCPServerManager()
|
||||
|
|
|
|||
|
|
@ -709,7 +709,8 @@ if MCP_AVAILABLE:
|
|||
|
||||
extra_headers: Optional[Dict[str, str]] = None
|
||||
if server.auth_type == MCPAuth.oauth2:
|
||||
extra_headers = oauth2_headers
|
||||
# Copy to avoid mutating the original dict (important for parallel fetching)
|
||||
extra_headers = oauth2_headers.copy() if oauth2_headers else None
|
||||
|
||||
if server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
|
|
@ -755,11 +756,10 @@ if MCP_AVAILABLE:
|
|||
# Decide whether to add prefix based on number of allowed servers
|
||||
add_prefix = not (len(allowed_mcp_servers) == 1)
|
||||
|
||||
# Get tools from each allowed server
|
||||
all_tools = []
|
||||
for server in allowed_mcp_servers:
|
||||
async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]:
|
||||
"""Fetch and filter tools from a single server with error handling."""
|
||||
if server is None:
|
||||
continue
|
||||
return []
|
||||
|
||||
server_auth_header, extra_headers = _prepare_mcp_server_headers(
|
||||
server=server,
|
||||
|
|
@ -786,16 +786,24 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
all_tools.extend(filtered_tools)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering"
|
||||
)
|
||||
return filtered_tools
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Error getting tools from server {server.name}: {str(e)}"
|
||||
)
|
||||
# Continue with other servers instead of failing completely
|
||||
return []
|
||||
|
||||
# Fetch tools from all servers in parallel
|
||||
tasks = [
|
||||
_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
# Flatten results into single list
|
||||
all_tools: List[MCPTool] = [tool for tools in results for tool in tools]
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully fetched {len(all_tools)} tools total from all MCP servers"
|
||||
|
|
|
|||
|
|
@ -1908,6 +1908,9 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase):
|
|||
}
|
||||
|
||||
|
||||
UserMCPManagementMode = Literal["restricted", "view_all"]
|
||||
|
||||
|
||||
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
||||
"""
|
||||
Documents all the fields supported by `general_settings` in config.yaml
|
||||
|
|
@ -2025,6 +2028,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).",
|
||||
)
|
||||
user_mcp_management_mode: Optional[UserMCPManagementMode] = Field(
|
||||
None,
|
||||
description="Controls how non-admin users interact with MCP servers in the dashboard. 'restricted' shows only accessible servers, 'view_all' lists every server in read-only mode.",
|
||||
)
|
||||
|
||||
|
||||
class ConfigYAML(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -319,6 +319,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"acreate_batch",
|
||||
"aretrieve_batch",
|
||||
"alist_batches",
|
||||
|
|
@ -457,6 +458,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"atext_completion",
|
||||
"aimage_edit",
|
||||
"alist_input_items",
|
||||
|
|
|
|||
|
|
@ -154,7 +154,11 @@ class PrismaManager:
|
|||
|
||||
prisma_dir = PrismaManager._get_prisma_dir()
|
||||
|
||||
return ProxyExtrasDBManager.setup_database(use_migrate=use_migrate)
|
||||
from litellm.proxy.proxy_server import redis_usage_cache
|
||||
|
||||
return ProxyExtrasDBManager.setup_database(
|
||||
use_migrate=use_migrate, redis_cache=redis_usage_cache
|
||||
)
|
||||
else:
|
||||
# Use prisma db push with increased timeout
|
||||
subprocess.run(
|
||||
|
|
|
|||
|
|
@ -118,7 +118,7 @@ class LassoGuardrail(CustomGuardrail):
|
|||
Falls back to UUID if ULID library is not available.
|
||||
"""
|
||||
if ULID_AVAILABLE and ulid is not None:
|
||||
return str(ulid.new()) # type: ignore
|
||||
return str(ulid.ULID()) # type: ignore
|
||||
else:
|
||||
verbose_proxy_logger.debug("ULID library not available, using UUID")
|
||||
return str(uuid.uuid4())
|
||||
|
|
|
|||
|
|
@ -32,8 +32,8 @@ from fastapi import (
|
|||
from fastapi.responses import JSONResponse
|
||||
|
||||
import litellm
|
||||
from litellm._uuid import uuid
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
validate_and_normalize_mcp_server_payload,
|
||||
|
|
@ -67,7 +67,6 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
build_effective_auth_contexts,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
LitellmUserRoles,
|
||||
|
|
@ -76,8 +75,10 @@ if MCP_AVAILABLE:
|
|||
SpecialMCPServerName,
|
||||
UpdateMCPServerRequest,
|
||||
UserAPIKeyAuth,
|
||||
UserMCPManagementMode,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
|
||||
from litellm.types.mcp import MCPCredentials
|
||||
|
|
@ -302,6 +303,20 @@ if MCP_AVAILABLE:
|
|||
return {"access_groups": access_groups_list}
|
||||
|
||||
## FastAPI Routes
|
||||
def _get_user_mcp_management_mode() -> UserMCPManagementMode:
|
||||
proxy_general_settings: dict = {}
|
||||
try:
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings as proxy_general_settings,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
mode = proxy_general_settings.get("user_mcp_management_mode")
|
||||
if mode == "view_all":
|
||||
return "view_all"
|
||||
return "restricted"
|
||||
|
||||
@router.get(
|
||||
"/server",
|
||||
description="Returns the mcp server list with associated teams",
|
||||
|
|
@ -319,18 +334,26 @@ if MCP_AVAILABLE:
|
|||
```
|
||||
"""
|
||||
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
user_mcp_management_mode = _get_user_mcp_management_mode()
|
||||
|
||||
aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
|
||||
for auth_context in auth_contexts:
|
||||
servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
|
||||
user_api_key_auth=auth_context
|
||||
if user_mcp_management_mode == "view_all":
|
||||
servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered()
|
||||
redacted_mcp_servers = _redact_mcp_credentials_list(servers)
|
||||
else:
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
|
||||
for auth_context in auth_contexts:
|
||||
servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
|
||||
user_api_key_auth=auth_context
|
||||
)
|
||||
for server in servers:
|
||||
if server.server_id not in aggregated_servers:
|
||||
aggregated_servers[server.server_id] = server
|
||||
|
||||
redacted_mcp_servers = _redact_mcp_credentials_list(
|
||||
aggregated_servers.values()
|
||||
)
|
||||
for server in servers:
|
||||
if server.server_id not in aggregated_servers:
|
||||
aggregated_servers[server.server_id] = server
|
||||
|
||||
redacted_mcp_servers = _redact_mcp_credentials_list(aggregated_servers.values())
|
||||
|
||||
# augment the mcp servers with public status
|
||||
if litellm.public_mcp_servers is not None:
|
||||
|
|
@ -372,6 +395,17 @@ if MCP_AVAILABLE:
|
|||
--header 'Authorization: Bearer your_api_key_here'
|
||||
```
|
||||
"""
|
||||
user_mcp_management_mode = _get_user_mcp_management_mode()
|
||||
|
||||
if user_mcp_management_mode == "view_all":
|
||||
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered(
|
||||
server_ids=server_ids
|
||||
)
|
||||
return [
|
||||
{"server_id": server.server_id, "status": server.status}
|
||||
for server in servers
|
||||
]
|
||||
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
server_status_map: Dict[
|
||||
|
|
|
|||
|
|
@ -698,6 +698,88 @@ async def get_response_input_items(
|
|||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/responses/compact",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["responses"],
|
||||
)
|
||||
@router.post(
|
||||
"/responses/compact",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["responses"],
|
||||
)
|
||||
@router.post(
|
||||
"/openai/v1/responses/compact",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["responses"],
|
||||
)
|
||||
async def compact_response(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Compact a response by running a compaction pass over a conversation.
|
||||
|
||||
Returns encrypted, opaque items that can be used to reduce context size.
|
||||
|
||||
Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/compact
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/v1/responses/compact \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "gpt-4o",
|
||||
"input": [{"role": "user", "content": "Hello"}]
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
_read_request_body,
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
select_data_generator,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
|
||||
data = await _read_request_body(request=request)
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="acompact_responses",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/responses/{response_id}/cancel",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ ROUTE_ENDPOINT_MAPPING = {
|
|||
"alist_input_items": "/responses/{response_id}/input_items",
|
||||
"aimage_edit": "/images/edits",
|
||||
"acancel_responses": "/responses/{response_id}/cancel",
|
||||
"acompact_responses": "/responses/compact",
|
||||
"aocr": "/ocr",
|
||||
"asearch": "/search",
|
||||
"avideo_generation": "/videos",
|
||||
|
|
@ -116,6 +117,7 @@ async def route_request(
|
|||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"acreate_response_reply",
|
||||
"alist_input_items",
|
||||
"_arealtime", # private function for realtime API
|
||||
|
|
|
|||
|
|
@ -1361,3 +1361,205 @@ def cancel_responses(
|
|||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def acompact_responses(
|
||||
input: Union[str, ResponseInputParam],
|
||||
model: str,
|
||||
instructions: Optional[str] = None,
|
||||
previous_response_id: Optional[str] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> ResponsesAPIResponse:
|
||||
"""
|
||||
Async version of the POST Compact Responses API
|
||||
|
||||
POST /v1/responses/compact endpoint in the responses API
|
||||
|
||||
Runs a compaction pass over a conversation, returning encrypted, opaque items.
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["acompact_responses"] = True
|
||||
|
||||
# get custom llm provider so we can use this for mapping exceptions
|
||||
if custom_llm_provider is None:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model, api_base=local_vars.get("base_url", None)
|
||||
)
|
||||
|
||||
func = partial(
|
||||
compact_responses,
|
||||
input=input,
|
||||
model=model,
|
||||
instructions=instructions,
|
||||
previous_response_id=previous_response_id,
|
||||
extra_headers=extra_headers,
|
||||
extra_query=extra_query,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
ctx = contextvars.copy_context()
|
||||
func_with_context = partial(ctx.run, func)
|
||||
init_response = await loop.run_in_executor(None, func_with_context)
|
||||
|
||||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
|
||||
# Update the responses_api_response_id with the model_id
|
||||
if isinstance(response, ResponsesAPIResponse):
|
||||
response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
|
||||
responses_api_response=response,
|
||||
litellm_metadata=kwargs.get("litellm_metadata", {}),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def compact_responses(
|
||||
input: Union[str, ResponseInputParam],
|
||||
model: str,
|
||||
instructions: Optional[str] = None,
|
||||
previous_response_id: Optional[str] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_query: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]:
|
||||
"""
|
||||
Synchronous version of the POST Compact Responses API
|
||||
|
||||
POST /v1/responses/compact endpoint in the responses API
|
||||
|
||||
Runs a compaction pass over a conversation, returning encrypted, opaque items.
|
||||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
_is_async = kwargs.pop("acompact_responses", False) is True
|
||||
|
||||
# get llm provider logic
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
(
|
||||
model,
|
||||
custom_llm_provider,
|
||||
dynamic_api_key,
|
||||
dynamic_api_base,
|
||||
) = litellm.get_llm_provider(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=litellm_params.api_base,
|
||||
api_key=litellm_params.api_key,
|
||||
)
|
||||
|
||||
if custom_llm_provider is None:
|
||||
raise ValueError("custom_llm_provider is required but passed as None")
|
||||
|
||||
# get provider config
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
raise ValueError(
|
||||
f"COMPACT responses is not supported for {custom_llm_provider}"
|
||||
)
|
||||
|
||||
local_vars.update(kwargs)
|
||||
|
||||
# Build optional params for compact endpoint
|
||||
response_api_optional_params: ResponsesAPIOptionalRequestParams = (
|
||||
ResponsesAPIRequestUtils.get_requested_response_api_optional_param(
|
||||
local_vars
|
||||
)
|
||||
)
|
||||
|
||||
# Get optional parameters for the responses API
|
||||
responses_api_request_params: Dict = (
|
||||
ResponsesAPIRequestUtils.get_optional_params_responses_api(
|
||||
model=model,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
response_api_optional_params=response_api_optional_params,
|
||||
allowed_openai_params=None,
|
||||
)
|
||||
)
|
||||
|
||||
# Pre Call logging
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model=model,
|
||||
optional_params=dict(responses_api_request_params),
|
||||
litellm_params={
|
||||
**responses_api_request_params,
|
||||
"litellm_call_id": litellm_call_id,
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
# Call the handler with _is_async flag instead of directly calling the async handler
|
||||
response = base_llm_http_handler.compact_response_api_handler(
|
||||
model=model,
|
||||
input=input,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
response_api_optional_request_params=responses_api_request_params,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout or request_timeout,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
shared_session=kwargs.get("shared_session"),
|
||||
)
|
||||
|
||||
# Update the responses_api_response_id with the model_id
|
||||
if isinstance(response, ResponsesAPIResponse):
|
||||
response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
|
||||
responses_api_response=response,
|
||||
litellm_metadata=kwargs.get("litellm_metadata", {}),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
|
@ -11,6 +12,9 @@ from litellm.litellm_core_utils.asyncify import run_async_function
|
|||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_base
|
||||
from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
|
||||
update_response_metadata,
|
||||
)
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
|
@ -22,7 +26,8 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook
|
||||
|
||||
|
||||
class BaseResponsesAPIStreamingIterator:
|
||||
|
|
@ -40,6 +45,8 @@ class BaseResponsesAPIStreamingIterator:
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
call_type: Optional[str] = None,
|
||||
):
|
||||
self.response = response
|
||||
self.model = model
|
||||
|
|
@ -47,21 +54,25 @@ class BaseResponsesAPIStreamingIterator:
|
|||
self.finished = False
|
||||
self.responses_api_provider_config = responses_api_provider_config
|
||||
self.completed_response: Optional[ResponsesAPIStreamingResponse] = None
|
||||
self.start_time = datetime.now()
|
||||
self.start_time = getattr(logging_obj, "start_time", datetime.now())
|
||||
|
||||
# set request kwargs
|
||||
# track request context for hooks
|
||||
self.litellm_metadata = litellm_metadata
|
||||
self.custom_llm_provider = custom_llm_provider
|
||||
self.request_data: Dict[str, Any] = request_data or {}
|
||||
self.call_type: Optional[str] = call_type
|
||||
|
||||
# set hidden params for response headers (e.g., x-litellm-model-id)
|
||||
# This matches ths stream wrapper in litellm/litellm_core_utils/streaming_handler.py
|
||||
# This matches the stream wrapper in litellm/litellm_core_utils/streaming_handler.py
|
||||
_api_base = get_api_base(
|
||||
model=model or "",
|
||||
optional_params=self.logging_obj.model_call_details.get(
|
||||
"litellm_params", {}
|
||||
),
|
||||
)
|
||||
_model_info: Dict = litellm_metadata.get("model_info", {}) if litellm_metadata else {}
|
||||
_model_info: Dict = (
|
||||
litellm_metadata.get("model_info", {}) if litellm_metadata else {}
|
||||
)
|
||||
self._hidden_params = {
|
||||
"model_id": _model_info.get("id", None),
|
||||
"api_base": _api_base,
|
||||
|
|
@ -102,13 +113,21 @@ class BaseResponsesAPIStreamingIterator:
|
|||
# if "response" in parsed_chunk, then encode litellm specific information like custom_llm_provider
|
||||
response_object = getattr(openai_responses_api_chunk, "response", None)
|
||||
if response_object:
|
||||
response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
|
||||
responses_api_response=response_object,
|
||||
litellm_metadata=self.litellm_metadata,
|
||||
custom_llm_provider=self.custom_llm_provider,
|
||||
response = (
|
||||
ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
|
||||
responses_api_response=response_object,
|
||||
litellm_metadata=self.litellm_metadata,
|
||||
custom_llm_provider=self.custom_llm_provider,
|
||||
)
|
||||
)
|
||||
setattr(openai_responses_api_chunk, "response", response)
|
||||
|
||||
# Allow callbacks to modify chunk before returning
|
||||
openai_responses_api_chunk = run_async_function(
|
||||
async_function=self._call_post_streaming_deployment_hook,
|
||||
chunk=openai_responses_api_chunk,
|
||||
)
|
||||
|
||||
# Store the completed response
|
||||
if (
|
||||
openai_responses_api_chunk
|
||||
|
|
@ -149,11 +168,159 @@ class BaseResponsesAPIStreamingIterator:
|
|||
except json.JSONDecodeError:
|
||||
# If we can't parse the chunk, continue
|
||||
return None
|
||||
except Exception as e:
|
||||
# Ensure failures trigger failure hooks
|
||||
self._handle_failure(e)
|
||||
raise
|
||||
|
||||
def _handle_logging_completed_response(self):
|
||||
"""Base implementation - should be overridden by subclasses"""
|
||||
pass
|
||||
|
||||
async def _call_post_streaming_deployment_hook(self, chunk):
|
||||
"""
|
||||
Allow callbacks to modify streaming chunks before returning (parity with chat).
|
||||
"""
|
||||
try:
|
||||
# Align with chat pipeline: use logging_obj model_call_details + call_type
|
||||
typed_call_type: Optional[CallTypes] = None
|
||||
if self.call_type is not None:
|
||||
try:
|
||||
typed_call_type = CallTypes(self.call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None
|
||||
if typed_call_type is None:
|
||||
try:
|
||||
typed_call_type = CallTypes(getattr(self.logging_obj, "call_type", None))
|
||||
except Exception:
|
||||
typed_call_type = None
|
||||
|
||||
request_data = self.request_data or getattr(
|
||||
self.logging_obj, "model_call_details", {}
|
||||
)
|
||||
callbacks = getattr(litellm, "callbacks", None) or []
|
||||
hooks_ran = False
|
||||
for callback in callbacks:
|
||||
if hasattr(callback, "async_post_call_streaming_deployment_hook"):
|
||||
hooks_ran = True
|
||||
result = await callback.async_post_call_streaming_deployment_hook(
|
||||
request_data=request_data,
|
||||
response_chunk=chunk,
|
||||
call_type=typed_call_type,
|
||||
)
|
||||
if result is not None:
|
||||
chunk = result
|
||||
if hooks_ran:
|
||||
setattr(chunk, "_post_streaming_hooks_ran", True)
|
||||
return chunk
|
||||
except Exception:
|
||||
return chunk
|
||||
|
||||
async def call_post_streaming_hooks_for_testing(self, chunk):
|
||||
"""
|
||||
Helper to invoke streaming deployment hooks explicitly (used in tests).
|
||||
"""
|
||||
return await self._call_post_streaming_deployment_hook(chunk)
|
||||
|
||||
def _run_post_success_hooks(self, end_time: datetime):
|
||||
"""
|
||||
Run post-call deployment hooks and update metadata similar to chat pipeline.
|
||||
"""
|
||||
if self.completed_response is None:
|
||||
return
|
||||
|
||||
request_payload: Dict[str, Any] = {}
|
||||
if isinstance(self.request_data, dict):
|
||||
request_payload.update(self.request_data)
|
||||
try:
|
||||
if hasattr(self.logging_obj, "model_call_details"):
|
||||
request_payload.update(self.logging_obj.model_call_details)
|
||||
except Exception:
|
||||
pass
|
||||
if "litellm_params" not in request_payload:
|
||||
try:
|
||||
request_payload["litellm_params"] = getattr(
|
||||
self.logging_obj, "model_call_details", {}
|
||||
).get("litellm_params", {})
|
||||
except Exception:
|
||||
request_payload["litellm_params"] = {}
|
||||
|
||||
try:
|
||||
update_response_metadata(
|
||||
result=self.completed_response,
|
||||
logging_obj=self.logging_obj,
|
||||
model=self.model,
|
||||
kwargs=request_payload,
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
except Exception:
|
||||
# Non-blocking
|
||||
pass
|
||||
|
||||
try:
|
||||
typed_call_type: Optional[CallTypes] = None
|
||||
if self.call_type is not None:
|
||||
try:
|
||||
typed_call_type = CallTypes(self.call_type)
|
||||
except ValueError:
|
||||
typed_call_type = None
|
||||
except Exception:
|
||||
typed_call_type = None
|
||||
if typed_call_type is None:
|
||||
try:
|
||||
typed_call_type = CallTypes.responses
|
||||
except Exception:
|
||||
typed_call_type = None
|
||||
|
||||
try:
|
||||
# Call synchronously; async hook will be executed via asyncio.run in a new loop
|
||||
run_async_function(
|
||||
async_function=async_post_call_success_deployment_hook,
|
||||
request_data=request_payload,
|
||||
response=self.completed_response,
|
||||
call_type=typed_call_type,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _handle_failure(self, exception: Exception):
|
||||
"""
|
||||
Trigger failure handlers before bubbling the exception.
|
||||
"""
|
||||
traceback_exception = traceback.format_exc()
|
||||
try:
|
||||
run_async_function(
|
||||
async_function=self.logging_obj.async_failure_handler,
|
||||
exception=exception,
|
||||
traceback_exception=traceback_exception,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
executor.submit(
|
||||
self.logging_obj.failure_handler,
|
||||
exception,
|
||||
traceback_exception,
|
||||
self.start_time,
|
||||
datetime.now(),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def call_post_streaming_hooks_for_testing(iterator, chunk):
|
||||
"""
|
||||
Module-level helper for tests to ensure hooks can be invoked even if the iterator is wrapped.
|
||||
"""
|
||||
hook_fn = getattr(iterator, "_call_post_streaming_deployment_hook", None)
|
||||
if hook_fn is None:
|
||||
return chunk
|
||||
return await hook_fn(chunk)
|
||||
|
||||
|
||||
class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
||||
"""
|
||||
|
|
@ -168,6 +335,8 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
call_type: Optional[str] = None,
|
||||
):
|
||||
super().__init__(
|
||||
response,
|
||||
|
|
@ -176,6 +345,8 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj,
|
||||
litellm_metadata,
|
||||
custom_llm_provider,
|
||||
request_data,
|
||||
call_type,
|
||||
)
|
||||
self.stream_iterator = response.aiter_lines()
|
||||
|
||||
|
|
@ -203,16 +374,21 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
self._handle_failure(e)
|
||||
raise e
|
||||
except Exception as e:
|
||||
self.finished = True
|
||||
self._handle_failure(e)
|
||||
raise e
|
||||
|
||||
def _handle_logging_completed_response(self):
|
||||
"""Handle logging for completed responses in async context"""
|
||||
# Create a deep copy for logging to avoid modifying the response object that will be returned to the user
|
||||
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
|
||||
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
|
||||
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
|
||||
import copy
|
||||
logging_response = copy.deepcopy(self.completed_response)
|
||||
|
||||
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_success_handler(
|
||||
result=logging_response,
|
||||
|
|
@ -229,6 +405,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
self._run_post_success_hooks(end_time=datetime.now())
|
||||
|
||||
|
||||
class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
||||
|
|
@ -244,6 +421,8 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
call_type: Optional[str] = None,
|
||||
):
|
||||
super().__init__(
|
||||
response,
|
||||
|
|
@ -252,6 +431,8 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj,
|
||||
litellm_metadata,
|
||||
custom_llm_provider,
|
||||
request_data,
|
||||
call_type,
|
||||
)
|
||||
self.stream_iterator = response.iter_lines()
|
||||
|
||||
|
|
@ -279,16 +460,21 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
self._handle_failure(e)
|
||||
raise e
|
||||
except Exception as e:
|
||||
self.finished = True
|
||||
self._handle_failure(e)
|
||||
raise e
|
||||
|
||||
def _handle_logging_completed_response(self):
|
||||
"""Handle logging for completed responses in sync context"""
|
||||
# Create a deep copy for logging to avoid modifying the response object that will be returned to the user
|
||||
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
|
||||
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
|
||||
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
|
||||
import copy
|
||||
logging_response = copy.deepcopy(self.completed_response)
|
||||
|
||||
|
||||
run_async_function(
|
||||
async_function=self.logging_obj.async_success_handler,
|
||||
result=logging_response,
|
||||
|
|
@ -304,6 +490,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
self._run_post_success_hooks(end_time=datetime.now())
|
||||
|
||||
|
||||
class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
||||
|
|
@ -324,6 +511,8 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
call_type: Optional[str] = None,
|
||||
):
|
||||
super().__init__(
|
||||
response=response,
|
||||
|
|
@ -332,6 +521,8 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
logging_obj=logging_obj,
|
||||
litellm_metadata=litellm_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_data=request_data,
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
# one-time transform
|
||||
|
|
|
|||
|
|
@ -713,6 +713,23 @@ class Router:
|
|||
self, routing_strategy: Union[RoutingStrategy, str], routing_strategy_args: dict
|
||||
):
|
||||
verbose_router_logger.info(f"Routing strategy: {routing_strategy}")
|
||||
|
||||
# Validate routing_strategy value to fail fast with helpful error
|
||||
# See: https://github.com/BerriAI/litellm/issues/11330
|
||||
# Derive valid strategies from RoutingStrategy enum + "simple-shuffle" (default, not in enum)
|
||||
valid_strategy_strings = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
|
||||
|
||||
if routing_strategy is not None:
|
||||
is_valid_string = isinstance(routing_strategy, str) and routing_strategy in valid_strategy_strings
|
||||
is_valid_enum = isinstance(routing_strategy, RoutingStrategy)
|
||||
if not is_valid_string and not is_valid_enum:
|
||||
raise ValueError(
|
||||
f"Invalid routing_strategy: '{routing_strategy}'. "
|
||||
f"Valid options: {valid_strategy_strings}. "
|
||||
f"Check 'router_settings.routing_strategy' in your config.yaml "
|
||||
f"or the 'routing_strategy' parameter if using the Router SDK directly."
|
||||
)
|
||||
|
||||
if (
|
||||
routing_strategy == RoutingStrategy.LEAST_BUSY.value
|
||||
or routing_strategy == RoutingStrategy.LEAST_BUSY
|
||||
|
|
@ -812,6 +829,9 @@ class Router:
|
|||
self.acancel_responses = self.factory_function(
|
||||
litellm.acancel_responses, call_type="acancel_responses"
|
||||
)
|
||||
self.acompact_responses = self.factory_function(
|
||||
litellm.acompact_responses, call_type="acompact_responses"
|
||||
)
|
||||
self.adelete_responses = self.factory_function(
|
||||
litellm.adelete_responses, call_type="adelete_responses"
|
||||
)
|
||||
|
|
@ -3924,6 +3944,7 @@ class Router:
|
|||
"anthropic_messages",
|
||||
"aresponses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"responses",
|
||||
"aget_responses",
|
||||
"adelete_responses",
|
||||
|
|
@ -4152,6 +4173,7 @@ class Router:
|
|||
elif call_type in (
|
||||
"aget_responses",
|
||||
"acancel_responses",
|
||||
"acompact_responses",
|
||||
"adelete_responses",
|
||||
"alist_input_items",
|
||||
):
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ class LangsmithCredentialsObject(TypedDict):
|
|||
LANGSMITH_API_KEY: Optional[str]
|
||||
LANGSMITH_PROJECT: Optional[str]
|
||||
LANGSMITH_BASE_URL: str
|
||||
LANGSMITH_TENANT_ID: Optional[str]
|
||||
|
||||
|
||||
class LangsmithQueueObject(TypedDict):
|
||||
|
|
@ -52,6 +53,7 @@ class CredentialsKey(NamedTuple):
|
|||
api_key: str
|
||||
project: str
|
||||
base_url: str
|
||||
tenant_id: Optional[str]
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
|
|||
|
|
@ -2677,6 +2677,7 @@ class StandardCallbackDynamicParams(TypedDict, total=False):
|
|||
langsmith_project: Optional[str]
|
||||
langsmith_base_url: Optional[str]
|
||||
langsmith_sampling_rate: Optional[float]
|
||||
langsmith_tenant_id: Optional[str]
|
||||
|
||||
# Humanloop dynamic params
|
||||
humanloop_api_key: Optional[str]
|
||||
|
|
@ -2946,6 +2947,7 @@ class LlmProviders(str, Enum):
|
|||
MISTRAL = "mistral"
|
||||
MILVUS = "milvus"
|
||||
GROQ = "groq"
|
||||
GIGACHAT = "gigachat"
|
||||
NVIDIA_NIM = "nvidia_nim"
|
||||
CEREBRAS = "cerebras"
|
||||
AI21_CHAT = "ai21_chat"
|
||||
|
|
|
|||
|
|
@ -7521,6 +7521,8 @@ class ProviderConfigManager:
|
|||
return litellm.CompactifAIChatConfig()
|
||||
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
|
||||
return litellm.GithubCopilotConfig()
|
||||
elif litellm.LlmProviders.GIGACHAT == provider:
|
||||
return litellm.GigaChatConfig()
|
||||
elif litellm.LlmProviders.RAGFLOW == provider:
|
||||
return litellm.RAGFlowConfig()
|
||||
elif (
|
||||
|
|
@ -7716,6 +7718,8 @@ class ProviderConfigManager:
|
|||
return litellm.CometAPIEmbeddingConfig()
|
||||
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
|
||||
return litellm.GithubCopilotEmbeddingConfig()
|
||||
elif litellm.LlmProviders.GIGACHAT == provider:
|
||||
return litellm.GigaChatEmbeddingConfig()
|
||||
elif litellm.LlmProviders.SAGEMAKER == provider:
|
||||
from litellm.llms.sagemaker.embedding.transformation import (
|
||||
SagemakerEmbeddingConfig,
|
||||
|
|
|
|||
|
|
@ -15831,6 +15831,68 @@
|
|||
"max_tokens": 8191,
|
||||
"mode": "embedding"
|
||||
},
|
||||
"gigachat/GigaChat-2-Lite": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0.0,
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true
|
||||
},
|
||||
"gigachat/GigaChat-2-Max": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0.0,
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"gigachat/GigaChat-2-Pro": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 0.0,
|
||||
"supports_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"gigachat/Embeddings": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 512,
|
||||
"max_tokens": 512,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 1024
|
||||
},
|
||||
"gigachat/Embeddings-2": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 512,
|
||||
"max_tokens": 512,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 1024
|
||||
},
|
||||
"gigachat/EmbeddingsGigaR": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "gigachat",
|
||||
"max_input_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 2560
|
||||
},
|
||||
"google.gemma-3-12b-it": {
|
||||
"input_cost_per_token": 9e-08,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
@ -32092,3 +32154,4 @@
|
|||
"mode": "chat"
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -28,7 +28,8 @@
|
|||
"list_container_files": "Supports GET /containers/{id}/files endpoint",
|
||||
"retrieve_container_file": "Supports GET /containers/{id}/files/{file_id} endpoint",
|
||||
"retrieve_container_file_content": "Supports GET /containers/{id}/files/{file_id}/content endpoint",
|
||||
"delete_container_file": "Supports DELETE /containers/{id}/files/{file_id} endpoint"
|
||||
"delete_container_file": "Supports DELETE /containers/{id}/files/{file_id} endpoint",
|
||||
"compact": "Supports /responses/compact endpoint"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
|
@ -1519,6 +1520,7 @@
|
|||
"retrieve_container_file": true,
|
||||
"retrieve_container_file_content": true,
|
||||
"delete_container_file": true,
|
||||
"compact": true,
|
||||
"a2a": true,
|
||||
"interactions": true
|
||||
}
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ google-cloud-aiplatform==1.47.0 # for vertex ai calls
|
|||
google-cloud-iam==2.19.1 # for GCP IAM Redis authentication
|
||||
google-genai==1.22.0
|
||||
anthropic[vertex]==0.54.0
|
||||
mcp==1.23.0 ; python_version >= "3.10" # for MCP server
|
||||
mcp==1.25.0 ; python_version >= "3.10" # for MCP server
|
||||
google-generativeai==0.5.0 # for vertex ai calls
|
||||
async_generator==1.10.0 # for async ollama calls
|
||||
langfuse==2.59.7 # for langfuse self-hosted logging
|
||||
|
|
|
|||
|
|
@ -61,13 +61,19 @@ async def test_bedrock_apply_guardrail_blocked():
|
|||
guardrailVersion="DRAFT",
|
||||
)
|
||||
|
||||
# Mock the make_bedrock_api_request method
|
||||
# Mock the make_bedrock_api_request method to raise an exception for blocked content
|
||||
with patch.object(
|
||||
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
|
||||
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
|
||||
) as mock_api_request:
|
||||
# Mock a blocked response from Bedrock
|
||||
mock_response = {"action": "BLOCKED", "reason": "Content violates policy"}
|
||||
mock_api_request.return_value = mock_response
|
||||
# Mock the method to raise an HTTPException as it would for blocked content
|
||||
from fastapi import HTTPException
|
||||
mock_api_request.side_effect = HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Violated guardrail policy",
|
||||
"bedrock_guardrail_response": "",
|
||||
},
|
||||
)
|
||||
|
||||
# Test the apply_guardrail method should raise an exception
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
|
|
@ -77,8 +83,9 @@ async def test_bedrock_apply_guardrail_blocked():
|
|||
input_type="request",
|
||||
)
|
||||
|
||||
assert "Content blocked by Bedrock guardrail" in str(exc_info.value)
|
||||
assert "Content violates policy" in str(exc_info.value)
|
||||
# The apply_guardrail method wraps the original exception in a generic Exception
|
||||
assert "Bedrock guardrail failed:" in str(exc_info.value)
|
||||
assert "Violated guardrail policy" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -253,7 +260,15 @@ async def test_bedrock_apply_guardrail_filters_request_messages_when_flag_enable
|
|||
with patch.object(
|
||||
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
|
||||
) as mock_api:
|
||||
mock_api.return_value = {"action": "BLOCKED", "reason": "policy"}
|
||||
# Mock the method to raise an HTTPException as it would for blocked content
|
||||
from fastapi import HTTPException
|
||||
mock_api.side_effect = HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Violated guardrail policy",
|
||||
"bedrock_guardrail_response": "policy",
|
||||
},
|
||||
)
|
||||
|
||||
with pytest.raises(Exception, match="policy") as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
|
|
@ -265,7 +280,8 @@ async def test_bedrock_apply_guardrail_filters_request_messages_when_flag_enable
|
|||
assert mock_api.called
|
||||
_, kwargs = mock_api.call_args
|
||||
assert kwargs["messages"] == [request_messages[-1]]
|
||||
assert "Content blocked by Bedrock guardrail" in str(exc_info.value)
|
||||
# The apply_guardrail method wraps the original exception in a generic Exception
|
||||
assert "Bedrock guardrail failed:" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_filters_latest_user_message_when_enabled():
|
||||
|
|
|
|||
|
|
@ -1,11 +1,14 @@
|
|||
import os
|
||||
import sys
|
||||
import time
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from litellm_proxy_extras.utils import ProxyExtrasDBManager
|
||||
from litellm_proxy_extras.utils import ProxyExtrasDBManager, MigrationLockManager
|
||||
|
||||
|
||||
|
||||
def test_custom_prisma_dir(monkeypatch):
|
||||
|
|
@ -27,101 +30,279 @@ def test_custom_prisma_dir(monkeypatch):
|
|||
assert os.path.exists(migrations_dir)
|
||||
|
||||
|
||||
class TestPermissionErrorDetection:
|
||||
"""Test cases for permission error detection in Prisma migrations"""
|
||||
class TestMigrationLockManager:
|
||||
"""Test cases for MigrationLockManager"""
|
||||
|
||||
def test_is_permission_error_postgres_42501(self):
|
||||
"""Test detection of PostgreSQL 42501 error code (insufficient privilege)"""
|
||||
error_message = "Database error code: 42501 - permission denied for table users"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is True
|
||||
def test_acquire_lock_without_redis(self):
|
||||
"""Test lock acquisition when Redis is not available"""
|
||||
lock_manager = MigrationLockManager()
|
||||
result = lock_manager.acquire_lock()
|
||||
assert result is True # Should return True when Redis is not available
|
||||
assert lock_manager.lock_acquired is True # Redis 없을 때도 lock_acquired는 True
|
||||
|
||||
def test_is_permission_error_must_be_owner(self):
|
||||
"""Test detection of 'must be owner of table' error"""
|
||||
error_message = "ERROR: must be owner of table my_table"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is True
|
||||
def test_acquire_lock_with_redis_success(self):
|
||||
"""Test successful lock acquisition with Redis"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.return_value = True
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
|
||||
def test_is_permission_error_permission_denied_schema(self):
|
||||
"""Test detection of 'permission denied for schema' error"""
|
||||
error_message = "permission denied for schema public"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is True
|
||||
result = lock_manager.acquire_lock()
|
||||
|
||||
def test_is_permission_error_permission_denied_table(self):
|
||||
"""Test detection of 'permission denied for table' error"""
|
||||
error_message = "permission denied for table my_table"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is True
|
||||
assert result is True
|
||||
assert lock_manager.lock_acquired is True
|
||||
mock_redis.set_cache.assert_called_once()
|
||||
|
||||
def test_is_permission_error_must_be_owner_schema(self):
|
||||
"""Test detection of 'must be owner of schema' error"""
|
||||
error_message = "must be owner of schema public"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is True
|
||||
def test_acquire_lock_with_redis_failure(self):
|
||||
"""Test failed lock acquisition with Redis"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.return_value = False
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
|
||||
def test_is_permission_error_case_insensitive(self):
|
||||
"""Test that permission error detection is case insensitive"""
|
||||
error_message = "PERMISSION DENIED FOR TABLE my_table"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is True
|
||||
result = lock_manager.acquire_lock()
|
||||
|
||||
def test_is_permission_error_negative(self):
|
||||
"""Test that non-permission errors are not detected as permission errors"""
|
||||
error_message = "column 'id' already exists"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is False
|
||||
assert result is False
|
||||
assert lock_manager.lock_acquired is False
|
||||
mock_redis.set_cache.assert_called_once()
|
||||
|
||||
def test_acquire_lock_with_redis_exception(self):
|
||||
"""Test lock acquisition with Redis exception"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.side_effect = Exception("Redis error")
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
|
||||
result = lock_manager.acquire_lock()
|
||||
|
||||
assert result is False
|
||||
assert lock_manager.lock_acquired is False
|
||||
|
||||
def test_wait_for_lock_release_success(self):
|
||||
"""Test successful waiting for lock release"""
|
||||
mock_redis = Mock()
|
||||
# First call returns False (lock held), second call returns True (lock acquired)
|
||||
mock_redis.set_cache.side_effect = [False, True]
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
|
||||
result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=1)
|
||||
|
||||
assert result is True
|
||||
assert lock_manager.lock_acquired is True
|
||||
assert mock_redis.set_cache.call_count == 2
|
||||
|
||||
def test_wait_for_lock_release_timeout(self):
|
||||
"""Test timeout while waiting for lock release"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.return_value = False # Lock always held
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
|
||||
result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=0.2)
|
||||
|
||||
assert result is False
|
||||
assert lock_manager.lock_acquired is False
|
||||
|
||||
def test_release_lock_not_acquired(self):
|
||||
"""Test releasing lock when not acquired"""
|
||||
mock_redis = Mock()
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
|
||||
lock_manager.release_lock()
|
||||
|
||||
mock_redis.get_cache.assert_not_called()
|
||||
mock_redis.delete_cache.assert_not_called()
|
||||
|
||||
def test_release_lock_success(self):
|
||||
"""Test successful lock release"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.get_cache.return_value = "pod_123_456"
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
lock_manager.pod_id = "pod_123_456"
|
||||
lock_manager.lock_acquired = True
|
||||
|
||||
lock_manager.release_lock()
|
||||
|
||||
mock_redis.get_cache.assert_called_once()
|
||||
mock_redis.delete_cache.assert_called_once()
|
||||
assert lock_manager.lock_acquired is False
|
||||
|
||||
def test_release_lock_wrong_owner(self):
|
||||
"""Test releasing lock when not the owner"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.get_cache.return_value = "pod_999_999" # Different pod
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
lock_manager.pod_id = "pod_123_456"
|
||||
lock_manager.lock_acquired = True
|
||||
|
||||
lock_manager.release_lock()
|
||||
|
||||
mock_redis.get_cache.assert_called_once()
|
||||
mock_redis.delete_cache.assert_not_called()
|
||||
assert lock_manager.lock_acquired is False
|
||||
|
||||
def test_context_manager(self):
|
||||
"""Test MigrationLockManager as context manager"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.return_value = True
|
||||
# Mock get_cache to return the same pod_id for successful release
|
||||
mock_redis.get_cache.return_value = "pod_123_456"
|
||||
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
lock_manager.pod_id = "pod_123_456" # Set consistent pod_id
|
||||
|
||||
with lock_manager:
|
||||
assert lock_manager.lock_acquired is True
|
||||
|
||||
# Should call release_lock when exiting context
|
||||
mock_redis.get_cache.assert_called_once()
|
||||
mock_redis.delete_cache.assert_called_once()
|
||||
|
||||
|
||||
class TestIdempotentErrorDetection:
|
||||
"""Test cases for idempotent error detection in Prisma migrations"""
|
||||
class TestProxyExtrasDBManagerMigrationLock:
|
||||
"""Test cases for ProxyExtrasDBManager with migration locking"""
|
||||
|
||||
def test_is_idempotent_error_already_exists(self):
|
||||
"""Test detection of generic 'already exists' error"""
|
||||
error_message = "object already exists"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
@patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations")
|
||||
@patch("litellm_proxy_extras.utils.subprocess.run")
|
||||
@patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
|
||||
@patch("os.chdir")
|
||||
def test_setup_database_with_redis_lock_success(
|
||||
self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations
|
||||
):
|
||||
"""Test successful database setup with Redis lock"""
|
||||
# Setup mocks
|
||||
mock_get_prisma_dir.return_value = "/test/prisma"
|
||||
mock_subprocess.return_value = Mock(stdout="Migration completed", stderr="")
|
||||
mock_resolve_migrations.return_value = None
|
||||
|
||||
def test_is_idempotent_error_column_already_exists(self):
|
||||
"""Test detection of 'column already exists' error"""
|
||||
error_message = "column 'email' already exists"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
# Mock Redis cache
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.return_value = True # Lock acquired successfully
|
||||
mock_redis.get_cache.return_value = "pod_123_456"
|
||||
|
||||
def test_is_idempotent_error_duplicate_key(self):
|
||||
"""Test detection of duplicate key violation error"""
|
||||
error_message = "duplicate key value violates unique constraint"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
# Set DATABASE_URL
|
||||
with patch.dict(
|
||||
os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
|
||||
):
|
||||
result = ProxyExtrasDBManager.setup_database(
|
||||
use_migrate=True, redis_cache=mock_redis
|
||||
)
|
||||
|
||||
def test_is_idempotent_error_relation_already_exists(self):
|
||||
"""Test detection of 'relation already exists' error"""
|
||||
error_message = "relation 'users_pkey' already exists"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
assert result is True
|
||||
# set_cache is called once in acquire_lock (__enter__ calls acquire_lock)
|
||||
assert mock_redis.set_cache.call_count == 1
|
||||
mock_subprocess.assert_called_once()
|
||||
|
||||
def test_is_idempotent_error_constraint_already_exists(self):
|
||||
"""Test detection of 'constraint already exists' error"""
|
||||
error_message = "constraint 'fk_user_id' already exists"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
@patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations")
|
||||
@patch("litellm_proxy_extras.utils.subprocess.run")
|
||||
@patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
|
||||
@patch("os.chdir")
|
||||
def test_setup_database_with_redis_lock_wait_and_skip(
|
||||
self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations
|
||||
):
|
||||
"""Test database setup when lock is held by another pod, then acquired after waiting"""
|
||||
# Setup mocks
|
||||
mock_get_prisma_dir.return_value = "/test/prisma"
|
||||
mock_resolve_migrations.return_value = None
|
||||
|
||||
def test_is_idempotent_error_case_insensitive(self):
|
||||
"""Test that idempotent error detection is case insensitive"""
|
||||
error_message = "COLUMN 'ID' ALREADY EXISTS"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
# Mock Redis cache - first call fails, second call succeeds
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.side_effect = [False, True] # First fails, then succeeds
|
||||
mock_redis.get_cache.return_value = "pod_123_456"
|
||||
|
||||
def test_is_idempotent_error_negative(self):
|
||||
"""Test that non-idempotent errors are not detected as idempotent errors"""
|
||||
error_message = "Database error code: 42501 - permission denied"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False
|
||||
# Set DATABASE_URL
|
||||
with patch.dict(
|
||||
os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
|
||||
):
|
||||
result = ProxyExtrasDBManager.setup_database(
|
||||
use_migrate=True, redis_cache=mock_redis
|
||||
)
|
||||
|
||||
assert result is True # Should return True after waiting and acquiring lock
|
||||
# set_cache is called 2 times: once in __enter__, once in wait_for_lock_release
|
||||
assert mock_redis.set_cache.call_count == 2
|
||||
# Proceed for case handling in case of migration failure
|
||||
mock_subprocess.assert_called_once()
|
||||
|
||||
class TestErrorClassificationPriority:
|
||||
"""Test cases to ensure errors are correctly classified"""
|
||||
@patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations")
|
||||
@patch("litellm_proxy_extras.utils.subprocess.run")
|
||||
@patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
|
||||
@patch("os.chdir")
|
||||
def test_setup_database_without_redis(
|
||||
self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations
|
||||
):
|
||||
"""Test database setup without Redis cache"""
|
||||
# Setup mocks
|
||||
mock_get_prisma_dir.return_value = "/test/prisma"
|
||||
mock_subprocess.return_value = Mock(stdout="Migration completed", stderr="")
|
||||
mock_resolve_migrations.return_value = None
|
||||
|
||||
def test_permission_error_not_classified_as_idempotent(self):
|
||||
"""Ensure permission errors are not mistakenly classified as idempotent"""
|
||||
error_message = "Database error code: 42501 - must be owner of table users"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is True
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False
|
||||
# Set DATABASE_URL
|
||||
with patch.dict(
|
||||
os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
|
||||
):
|
||||
result = ProxyExtrasDBManager.setup_database(
|
||||
use_migrate=True, redis_cache=None
|
||||
)
|
||||
|
||||
def test_idempotent_error_not_classified_as_permission(self):
|
||||
"""Ensure idempotent errors are not mistakenly classified as permission errors"""
|
||||
error_message = "column 'created_at' already exists"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is False
|
||||
assert result is True
|
||||
# Redis가 없을 때는 락 보호 없이 마이그레이션을 실행해야 함
|
||||
mock_subprocess.assert_called_once()
|
||||
|
||||
def test_unknown_error_classified_as_neither(self):
|
||||
"""Ensure unknown errors are classified as neither permission nor idempotent"""
|
||||
error_message = "connection timeout"
|
||||
assert ProxyExtrasDBManager._is_permission_error(error_message) is False
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False
|
||||
def test_setup_database_no_database_url(self):
|
||||
"""Test database setup without DATABASE_URL"""
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
result = ProxyExtrasDBManager.setup_database(
|
||||
use_migrate=True, redis_cache=None
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
||||
@patch("litellm_proxy_extras.utils.subprocess.run")
|
||||
@patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir")
|
||||
@patch("os.chdir")
|
||||
@patch.object(
|
||||
MigrationLockManager, "LOCK_TTL_SECONDS", 1
|
||||
) # Set short TTL for testing
|
||||
def test_setup_database_lock_timeout(
|
||||
self, mock_chdir, mock_get_prisma_dir, mock_subprocess
|
||||
):
|
||||
"""Test database setup when lock acquisition times out"""
|
||||
# Setup mocks
|
||||
mock_get_prisma_dir.return_value = "/test/prisma"
|
||||
|
||||
# Mock Redis cache - always fails to acquire lock
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.return_value = False # Always fails
|
||||
|
||||
# Set DATABASE_URL
|
||||
with patch.dict(
|
||||
os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"}
|
||||
):
|
||||
# Patch the wait_for_lock_release method to use shorter timeout
|
||||
with patch.object(
|
||||
MigrationLockManager, "wait_for_lock_release"
|
||||
) as mock_wait:
|
||||
mock_wait.return_value = False # Simulate timeout
|
||||
|
||||
result = ProxyExtrasDBManager.setup_database(
|
||||
use_migrate=True, redis_cache=mock_redis
|
||||
)
|
||||
|
||||
assert result is False # Should return False after timeout
|
||||
mock_subprocess.assert_not_called() # Should not run migration
|
||||
# Verify that wait_for_lock_release was called with default parameters
|
||||
mock_wait.assert_called_once_with()
|
||||
|
||||
def test_wait_for_lock_release_actual_timeout(self):
|
||||
"""Test actual timeout behavior of wait_for_lock_release with real timing"""
|
||||
mock_redis = Mock()
|
||||
mock_redis.set_cache.return_value = False # Always fails to acquire lock
|
||||
lock_manager = MigrationLockManager(mock_redis)
|
||||
|
||||
# Test with very short timeout to verify actual timeout behavior
|
||||
start_time = time.time()
|
||||
result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=0.5)
|
||||
end_time = time.time()
|
||||
|
||||
assert result is False # Should timeout
|
||||
assert end_time - start_time >= 0.5 # Should wait at least the max_wait time
|
||||
assert end_time - start_time < 1.0 # But not too much longer
|
||||
# Should have called set_cache multiple times during the wait
|
||||
assert mock_redis.set_cache.call_count > 1
|
||||
|
|
|
|||
|
|
@ -1814,3 +1814,49 @@ async def test_extra_body_merges_with_request_data(extra_body_mock_response_data
|
|||
assert "temperature" in request_body
|
||||
assert "custom_field" in request_body
|
||||
assert request_body["custom_field"] == "custom_value"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
async def test_openai_compact_responses_api(sync_mode):
|
||||
"""
|
||||
Test the compact_responses API for OpenAI.
|
||||
|
||||
This test verifies that the compact_responses endpoint works correctly
|
||||
for compressing conversation history.
|
||||
"""
|
||||
litellm._turn_on_debug()
|
||||
litellm.set_verbose = True
|
||||
|
||||
input_messages = [
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
{"role": "assistant", "content": "I'm doing well, thank you for asking!"},
|
||||
{"role": "user", "content": "What is the weather like today?"},
|
||||
]
|
||||
|
||||
try:
|
||||
if sync_mode:
|
||||
response = litellm.compact_responses(
|
||||
model="openai/gpt-4o",
|
||||
input=input_messages,
|
||||
instructions="Be helpful and concise",
|
||||
)
|
||||
else:
|
||||
response = await litellm.acompact_responses(
|
||||
model="openai/gpt-4o",
|
||||
input=input_messages,
|
||||
instructions="Be helpful and concise",
|
||||
)
|
||||
except litellm.InternalServerError:
|
||||
pytest.skip("Skipping test due to InternalServerError")
|
||||
except litellm.BadRequestError as e:
|
||||
# compact_responses may not be available for all models/accounts
|
||||
pytest.skip(f"Skipping test due to BadRequestError: {e}")
|
||||
|
||||
print("compact_responses response=", json.dumps(response, indent=4, default=str))
|
||||
|
||||
# Validate response structure
|
||||
assert response is not None
|
||||
assert "id" in response, "Response should have an 'id' field"
|
||||
assert "output" in response, "Response should have an 'output' field"
|
||||
assert isinstance(response["output"], list), "Output should be a list"
|
||||
|
|
|
|||
165
tests/llm_responses_api_testing/test_responses_hooks.py
Normal file
165
tests/llm_responses_api_testing/test_responses_hooks.py
Normal file
|
|
@ -0,0 +1,165 @@
|
|||
import asyncio
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.responses import streaming_iterator as streaming_module
|
||||
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
|
||||
from litellm.types.llms.openai import ResponsesAPIStreamEvents
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
class _FakeLoggingObj:
|
||||
def __init__(self):
|
||||
self.success_calls = 0
|
||||
self.async_success_calls = 0
|
||||
self.failure_calls = 0
|
||||
self.async_failure_calls = 0
|
||||
self.start_time = datetime.now()
|
||||
self.model_call_details = {"litellm_params": {}}
|
||||
|
||||
# Signature alignment with Logging handlers
|
||||
def success_handler(self, *args, **kwargs):
|
||||
self.success_calls += 1
|
||||
|
||||
async def async_success_handler(self, *args, **kwargs):
|
||||
self.async_success_calls += 1
|
||||
|
||||
def failure_handler(self, *args, **kwargs):
|
||||
self.failure_calls += 1
|
||||
|
||||
async def async_failure_handler(self, *args, **kwargs):
|
||||
self.async_failure_calls += 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_streaming_triggers_hooks(monkeypatch):
|
||||
"""
|
||||
Ensure streaming iterator fires success + post-call hooks for responses API.
|
||||
"""
|
||||
hook_calls = {"post_call": 0, "metadata": 0}
|
||||
seen = {}
|
||||
|
||||
async def fake_post_call(request_data, response, call_type):
|
||||
hook_calls["post_call"] += 1
|
||||
seen["request_data"] = request_data
|
||||
seen["call_type"] = call_type
|
||||
|
||||
def fake_update_metadata(**kwargs):
|
||||
hook_calls["metadata"] += 1
|
||||
|
||||
monkeypatch.setattr(
|
||||
streaming_module,
|
||||
"async_post_call_success_deployment_hook",
|
||||
fake_post_call,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
streaming_module,
|
||||
"update_response_metadata",
|
||||
fake_update_metadata,
|
||||
)
|
||||
|
||||
logging_obj = _FakeLoggingObj()
|
||||
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=SimpleNamespace(), # not used in this test
|
||||
logging_obj=logging_obj,
|
||||
request_data={"foo": "bar", "litellm_params": {}},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
# Simulate completed streaming event
|
||||
iterator.completed_response = SimpleNamespace(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, response=SimpleNamespace()
|
||||
)
|
||||
|
||||
iterator._handle_logging_completed_response()
|
||||
await asyncio.sleep(0.2) # allow async tasks to run
|
||||
|
||||
assert logging_obj.success_calls == 1
|
||||
assert logging_obj.async_success_calls == 1
|
||||
assert hook_calls["post_call"] == 1
|
||||
assert hook_calls["metadata"] == 1
|
||||
assert seen["request_data"]["foo"] == "bar"
|
||||
assert seen["request_data"].get("litellm_params") is not None
|
||||
assert seen["call_type"] == CallTypes.responses
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_streaming_calls_post_streaming_deployment_hook(monkeypatch):
|
||||
"""
|
||||
Ensure per-chunk streaming deployment hook can modify chunks.
|
||||
"""
|
||||
|
||||
class _HookLogger(CustomLogger):
|
||||
async def async_post_call_streaming_deployment_hook(
|
||||
self, request_data, response_chunk, call_type
|
||||
):
|
||||
response_chunk.tagged = True
|
||||
return response_chunk
|
||||
|
||||
# Set callbacks to our fake hook
|
||||
original_callbacks = litellm.callbacks
|
||||
litellm.callbacks = [_HookLogger()]
|
||||
|
||||
logging_obj = _FakeLoggingObj()
|
||||
|
||||
class _StubConfig:
|
||||
def transform_streaming_response(self, **kwargs):
|
||||
return SimpleNamespace(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None
|
||||
)
|
||||
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=_StubConfig(),
|
||||
logging_obj=logging_obj,
|
||||
request_data={"foo": "bar"},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
# Call hook helper directly to verify chunk is modified/flagged
|
||||
chunk = SimpleNamespace(type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None)
|
||||
chunk = await streaming_module.call_post_streaming_hooks_for_testing(iterator, chunk)
|
||||
assert getattr(chunk, "_post_streaming_hooks_ran", False) is True
|
||||
assert getattr(chunk, "tagged", False) is True
|
||||
|
||||
# reset callbacks
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_streaming_failure_triggers_failure_handlers():
|
||||
"""
|
||||
If transform raises, failure handlers should be called.
|
||||
"""
|
||||
|
||||
class _FailConfig:
|
||||
def transform_streaming_response(self, **kwargs):
|
||||
raise ValueError("boom")
|
||||
|
||||
logging_obj = _FakeLoggingObj()
|
||||
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="test-model",
|
||||
responses_api_provider_config=_FailConfig(),
|
||||
logging_obj=logging_obj,
|
||||
request_data={"foo": "bar"},
|
||||
call_type=CallTypes.responses.value,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
iterator._process_chunk('{"delta": "chunk"}')
|
||||
|
||||
# allow failure callbacks to run
|
||||
await asyncio.sleep(0.2)
|
||||
assert logging_obj.failure_calls >= 1
|
||||
assert logging_obj.async_failure_calls >= 1
|
||||
|
|
@ -385,7 +385,7 @@ def test_anthropic_tool_use(tool_type, tool_config, message_content):
|
|||
"computer_tool_used, prompt_caching_set, expected_beta_header",
|
||||
[
|
||||
(True, False, True),
|
||||
(False, True, True),
|
||||
(False, True, False),
|
||||
(True, True, True),
|
||||
(False, False, False),
|
||||
],
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import litellm
|
|||
from litellm.exceptions import BadRequestError
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
from litellm._version import version
|
||||
from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest
|
||||
|
||||
try:
|
||||
|
|
@ -725,6 +726,7 @@ def test_embeddings_with_sync_http_handler(monkeypatch):
|
|||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": f"litellm/{version}",
|
||||
},
|
||||
data=json.dumps(
|
||||
{
|
||||
|
|
@ -767,6 +769,7 @@ def test_embeddings_with_async_http_handler(monkeypatch):
|
|||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": f"litellm/{version}",
|
||||
},
|
||||
data=json.dumps(
|
||||
{
|
||||
|
|
@ -823,6 +826,7 @@ def test_embeddings_uses_databricks_sdk_if_api_key_and_base_not_specified(monkey
|
|||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": f"litellm/{version}",
|
||||
},
|
||||
data=json.dumps(
|
||||
{
|
||||
|
|
@ -895,6 +899,7 @@ async def test_databricks_embeddings(sync_mode, monkeypatch):
|
|||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": f"litellm/{version}",
|
||||
},
|
||||
data=json.dumps(
|
||||
{
|
||||
|
|
@ -923,6 +928,7 @@ async def test_databricks_embeddings(sync_mode, monkeypatch):
|
|||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": f"litellm/{version}",
|
||||
},
|
||||
data=json.dumps(
|
||||
{
|
||||
|
|
|
|||
349
tests/llm_translation/test_gigachat.py
Normal file
349
tests/llm_translation/test_gigachat.py
Normal file
|
|
@ -0,0 +1,349 @@
|
|||
"""
|
||||
Tests for GigaChat LiteLLM Provider
|
||||
|
||||
Tests message transformation, parameter handling, and response transformation.
|
||||
Run with: pytest tests/llm_translation/test_gigachat.py -v
|
||||
"""
|
||||
|
||||
import json
|
||||
import pytest
|
||||
from unittest.mock import Mock, MagicMock
|
||||
|
||||
|
||||
class TestGigaChatMessageTransformation:
|
||||
"""Tests for message transformation (OpenAI -> GigaChat format)"""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
|
||||
return GigaChatConfig()
|
||||
|
||||
def test_simple_user_message(self, config):
|
||||
"""Basic user message should pass through"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
result = config._transform_messages(messages)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]["role"] == "user"
|
||||
assert result[0]["content"] == "Hello"
|
||||
|
||||
def test_developer_role_to_system(self, config):
|
||||
"""Developer role should be converted to system"""
|
||||
messages = [{"role": "developer", "content": "You are helpful"}]
|
||||
result = config._transform_messages(messages)
|
||||
|
||||
assert result[0]["role"] == "system"
|
||||
|
||||
def test_system_after_first_becomes_user(self, config):
|
||||
"""System message after first position should become user"""
|
||||
messages = [
|
||||
{"role": "assistant", "content": "Response"},
|
||||
{"role": "system", "content": "Additional instruction"},
|
||||
]
|
||||
result = config._transform_messages(messages)
|
||||
|
||||
assert result[0]["role"] == "assistant"
|
||||
assert result[1]["role"] == "user" # system after first becomes user
|
||||
|
||||
def test_tool_role_to_function(self, config):
|
||||
"""Tool role should be converted to function"""
|
||||
messages = [{"role": "tool", "content": "result data"}]
|
||||
result = config._transform_messages(messages)
|
||||
|
||||
assert result[0]["role"] == "function"
|
||||
|
||||
def test_tool_calls_to_function_call(self, config):
|
||||
"""tool_calls should be converted to function_call"""
|
||||
messages = [{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{
|
||||
"id": "call_123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"city": "Moscow"}'
|
||||
}
|
||||
}]
|
||||
}]
|
||||
result = config._transform_messages(messages)
|
||||
|
||||
assert "function_call" in result[0]
|
||||
assert result[0]["function_call"]["name"] == "get_weather"
|
||||
assert result[0]["function_call"]["arguments"] == {"city": "Moscow"}
|
||||
assert "tool_calls" not in result[0]
|
||||
|
||||
def test_none_content_becomes_empty_string(self, config):
|
||||
"""None content should become empty string"""
|
||||
messages = [{"role": "assistant", "content": None}]
|
||||
result = config._transform_messages(messages)
|
||||
|
||||
assert result[0]["content"] == ""
|
||||
|
||||
def test_name_field_removed(self, config):
|
||||
"""name field should be removed (not supported by GigaChat)"""
|
||||
messages = [{"role": "user", "content": "Hi", "name": "John"}]
|
||||
result = config._transform_messages(messages)
|
||||
|
||||
assert "name" not in result[0]
|
||||
|
||||
|
||||
class TestGigaChatCollapseUserMessages:
|
||||
"""Tests for collapsing consecutive user messages"""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
|
||||
return GigaChatConfig()
|
||||
|
||||
def test_no_collapse_single_message(self, config):
|
||||
"""Single message should not be changed"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
result = config._collapse_user_messages(messages)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]["content"] == "Hello"
|
||||
|
||||
def test_collapse_consecutive_user_messages(self, config):
|
||||
"""Consecutive user messages should be collapsed"""
|
||||
messages = [
|
||||
{"role": "user", "content": "First"},
|
||||
{"role": "user", "content": "Second"},
|
||||
{"role": "user", "content": "Third"},
|
||||
]
|
||||
result = config._collapse_user_messages(messages)
|
||||
|
||||
assert len(result) == 1
|
||||
assert "First" in result[0]["content"]
|
||||
assert "Second" in result[0]["content"]
|
||||
assert "Third" in result[0]["content"]
|
||||
|
||||
def test_no_collapse_with_assistant_between(self, config):
|
||||
"""Messages with assistant between should not be collapsed"""
|
||||
messages = [
|
||||
{"role": "user", "content": "First"},
|
||||
{"role": "assistant", "content": "Response"},
|
||||
{"role": "user", "content": "Second"},
|
||||
]
|
||||
result = config._collapse_user_messages(messages)
|
||||
|
||||
assert len(result) == 3
|
||||
|
||||
|
||||
class TestGigaChatToolsTransformation:
|
||||
"""Tests for tools -> functions conversion"""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
|
||||
return GigaChatConfig()
|
||||
|
||||
def test_single_tool_conversion(self, config):
|
||||
"""Single tool should be converted correctly"""
|
||||
tools = [{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a city",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": {"type": "string"}
|
||||
}
|
||||
}
|
||||
}
|
||||
}]
|
||||
result = config._convert_tools_to_functions(tools)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]["name"] == "get_weather"
|
||||
assert result[0]["description"] == "Get weather for a city"
|
||||
|
||||
def test_multiple_tools_conversion(self, config):
|
||||
"""Multiple tools should all be converted"""
|
||||
tools = [
|
||||
{"type": "function", "function": {"name": "func1", "description": "First", "parameters": {"type": "object", "properties": {}}}},
|
||||
{"type": "function", "function": {"name": "func2", "description": "Second", "parameters": {"type": "object", "properties": {}}}},
|
||||
]
|
||||
result = config._convert_tools_to_functions(tools)
|
||||
|
||||
assert len(result) == 2
|
||||
assert result[0]["name"] == "func1"
|
||||
assert result[1]["name"] == "func2"
|
||||
|
||||
|
||||
class TestGigaChatParamsTransformation:
|
||||
"""Tests for parameter transformation"""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
|
||||
return GigaChatConfig()
|
||||
|
||||
def test_temperature_zero_becomes_top_p_zero(self, config):
|
||||
"""temperature=0 should become top_p=0"""
|
||||
params = {"temperature": 0}
|
||||
result = config.map_openai_params(
|
||||
non_default_params=params,
|
||||
optional_params={},
|
||||
model="GigaChat",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "top_p" in result
|
||||
assert result["top_p"] == 0
|
||||
assert "temperature" not in result
|
||||
|
||||
def test_temperature_nonzero_preserved(self, config):
|
||||
"""Non-zero temperature should be preserved"""
|
||||
params = {"temperature": 0.7}
|
||||
result = config.map_openai_params(
|
||||
non_default_params=params,
|
||||
optional_params={},
|
||||
model="GigaChat",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["temperature"] == 0.7
|
||||
|
||||
def test_max_completion_tokens_to_max_tokens(self, config):
|
||||
"""max_completion_tokens should become max_tokens"""
|
||||
params = {"max_completion_tokens": 100}
|
||||
result = config.map_openai_params(
|
||||
non_default_params=params,
|
||||
optional_params={},
|
||||
model="GigaChat",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["max_tokens"] == 100
|
||||
|
||||
def test_structured_output_via_json_schema(self, config):
|
||||
"""json_schema response_format should trigger structured output mode"""
|
||||
params = {
|
||||
"response_format": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "person",
|
||||
"schema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"age": {"type": "integer"}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
result = config.map_openai_params(
|
||||
non_default_params=params,
|
||||
optional_params={},
|
||||
model="GigaChat",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "_structured_output" in result
|
||||
assert result["_structured_output"] is True
|
||||
assert "function_call" in result
|
||||
assert result["function_call"]["name"] == "person"
|
||||
|
||||
|
||||
class TestGigaChatProviderRegistration:
|
||||
"""Tests for provider registration in LiteLLM"""
|
||||
|
||||
def test_gigachat_in_provider_list(self):
|
||||
"""GigaChat should be in provider list"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
assert hasattr(LlmProviders, "GIGACHAT")
|
||||
assert LlmProviders.GIGACHAT.value == "gigachat"
|
||||
|
||||
def test_gigachat_in_chat_providers(self):
|
||||
"""GigaChat should be in LITELLM_CHAT_PROVIDERS"""
|
||||
from litellm.constants import LITELLM_CHAT_PROVIDERS
|
||||
|
||||
assert "gigachat" in LITELLM_CHAT_PROVIDERS
|
||||
|
||||
def test_gigachat_key_exists(self):
|
||||
"""gigachat_key should be available"""
|
||||
import litellm
|
||||
|
||||
assert hasattr(litellm, "gigachat_key")
|
||||
|
||||
def test_gigachat_config_exists(self):
|
||||
"""GigaChatConfig should be available"""
|
||||
import litellm
|
||||
|
||||
assert hasattr(litellm, "GigaChatConfig")
|
||||
|
||||
|
||||
class TestGigaChatTransformRequest:
|
||||
"""Tests for request transformation"""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
|
||||
return GigaChatConfig()
|
||||
|
||||
def test_basic_request(self, config):
|
||||
"""Basic request should be transformed correctly"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
result = config.transform_request(
|
||||
model="gigachat/GigaChat",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["model"] == "GigaChat"
|
||||
assert len(result["messages"]) == 1
|
||||
assert result["messages"][0]["role"] == "user"
|
||||
|
||||
def test_request_with_temperature(self, config):
|
||||
"""Request with temperature should include it"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
result = config.transform_request(
|
||||
model="gigachat/GigaChat",
|
||||
messages=messages,
|
||||
optional_params={"temperature": 0.7},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["temperature"] == 0.7
|
||||
|
||||
def test_request_with_functions(self, config):
|
||||
"""Request with functions should include them"""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
functions = [{"name": "test", "description": "Test", "parameters": {}}]
|
||||
result = config.transform_request(
|
||||
model="gigachat/GigaChat",
|
||||
messages=messages,
|
||||
optional_params={"functions": functions},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert "functions" in result
|
||||
assert len(result["functions"]) == 1
|
||||
|
||||
|
||||
class TestGigaChatSupportedParams:
|
||||
"""Tests for supported parameters"""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
from litellm.llms.gigachat.chat.transformation import GigaChatConfig
|
||||
return GigaChatConfig()
|
||||
|
||||
def test_supported_params(self, config):
|
||||
"""Check supported parameters list"""
|
||||
supported = config.get_supported_openai_params("GigaChat")
|
||||
|
||||
assert "temperature" in supported
|
||||
assert "max_tokens" in supported
|
||||
assert "max_completion_tokens" in supported
|
||||
assert "tools" in supported
|
||||
assert "response_format" in supported
|
||||
assert "stream" in supported
|
||||
|
|
@ -286,7 +286,7 @@ def test_completion_claude_3_empty_response():
|
|||
},
|
||||
]
|
||||
try:
|
||||
response = litellm.completion(model="claude-3-opus-20240229", messages=messages)
|
||||
response = litellm.completion(model="claude-3-7-sonnet-20250219", messages=messages)
|
||||
print(response)
|
||||
except litellm.InternalServerError as e:
|
||||
pytest.skip(f"InternalServerError - {str(e)}")
|
||||
|
|
@ -313,7 +313,7 @@ def test_completion_claude_3():
|
|||
try:
|
||||
# test without max tokens
|
||||
response = completion(
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
model="anthropic/claude-3-7-sonnet-20250219",
|
||||
messages=messages,
|
||||
)
|
||||
# Add any assertions, here to check response args
|
||||
|
|
@ -326,7 +326,7 @@ def test_completion_claude_3():
|
|||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
["anthropic/claude-3-opus-20240229", "anthropic.claude-3-sonnet-20240229-v1:0"],
|
||||
["anthropic/claude-3-7-sonnet-20250219", "anthropic.claude-3-sonnet-20240229-v1:0"],
|
||||
)
|
||||
def test_completion_claude_3_function_call(model):
|
||||
litellm.set_verbose = True
|
||||
|
|
@ -411,7 +411,7 @@ def test_completion_claude_3_function_call(model):
|
|||
"model, api_key, api_base",
|
||||
[
|
||||
("gpt-3.5-turbo", None, None),
|
||||
("claude-3-opus-20240229", None, None),
|
||||
("claude-3-7-sonnet-20250219", None, None),
|
||||
("anthropic.claude-3-sonnet-20240229-v1:0", None, None),
|
||||
# (
|
||||
# "azure_ai/command-r-plus",
|
||||
|
|
@ -512,7 +512,7 @@ async def test_anthropic_no_content_error():
|
|||
try:
|
||||
litellm.drop_params = True
|
||||
response = await litellm.acompletion(
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
model="anthropic/claude-3-7-sonnet-20250219",
|
||||
api_key=os.getenv("ANTHROPIC_API_KEY"),
|
||||
messages=[
|
||||
{
|
||||
|
|
@ -630,7 +630,7 @@ def test_completion_claude_3_multi_turn_conversations():
|
|||
]
|
||||
try:
|
||||
response = completion(
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
model="anthropic/claude-3-7-sonnet-20250219",
|
||||
messages=messages,
|
||||
)
|
||||
print(response)
|
||||
|
|
@ -644,7 +644,7 @@ def test_completion_claude_3_stream():
|
|||
try:
|
||||
# test without max tokens
|
||||
response = completion(
|
||||
model="anthropic/claude-3-opus-20240229",
|
||||
model="anthropic/claude-3-7-sonnet-20250219",
|
||||
messages=messages,
|
||||
max_tokens=10,
|
||||
stream=True,
|
||||
|
|
@ -669,7 +669,7 @@ def encode_image(image_path):
|
|||
[
|
||||
"gpt-4o",
|
||||
"azure/gpt-4.1-mini",
|
||||
"anthropic/claude-3-opus-20240229",
|
||||
"anthropic/claude-3-7-sonnet-20250219",
|
||||
],
|
||||
) #
|
||||
def test_completion_base64(model):
|
||||
|
|
|
|||
|
|
@ -1418,7 +1418,7 @@ def test_bedrock_claude_3_streaming():
|
|||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"claude-3-opus-20240229",
|
||||
"claude-3-7-sonnet-20250219",
|
||||
"cohere.command-r-plus-v1:0", # bedrock
|
||||
"gpt-3.5-turbo",
|
||||
],
|
||||
|
|
@ -2914,7 +2914,7 @@ def test_completion_claude_3_function_call_with_streaming():
|
|||
try:
|
||||
# test without max tokens
|
||||
response = completion(
|
||||
model="claude-3-opus-20240229",
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
tool_choice="required",
|
||||
|
|
@ -2946,7 +2946,7 @@ def test_completion_claude_3_function_call_with_streaming():
|
|||
"model",
|
||||
[
|
||||
"gemini/gemini-2.5-flash-lite",
|
||||
], # "claude-3-opus-20240229"
|
||||
],
|
||||
) #
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_function_call_with_streaming(model):
|
||||
|
|
|
|||
|
|
@ -47,6 +47,19 @@ async def test_get_credentials_from_env():
|
|||
credentials = logger.get_credentials_from_env()
|
||||
assert credentials["LANGSMITH_BASE_URL"] == "https://api.smith.langchain.com"
|
||||
|
||||
# Test with tenant_id
|
||||
credentials = logger.get_credentials_from_env(
|
||||
langsmith_tenant_id="test-tenant-id"
|
||||
)
|
||||
assert credentials["LANGSMITH_TENANT_ID"] == "test-tenant-id"
|
||||
|
||||
# Test tenant_id from environment variable
|
||||
import os
|
||||
os.environ["LANGSMITH_TENANT_ID"] = "env-tenant-id"
|
||||
credentials = logger.get_credentials_from_env()
|
||||
assert credentials["LANGSMITH_TENANT_ID"] == "env-tenant-id"
|
||||
del os.environ["LANGSMITH_TENANT_ID"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_batches_by_credentials():
|
||||
|
|
@ -60,6 +73,7 @@ async def test_group_batches_by_credentials():
|
|||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -69,6 +83,7 @@ async def test_group_batches_by_credentials():
|
|||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -95,6 +110,7 @@ async def test_group_batches_by_credentials_multiple_credentials():
|
|||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -104,6 +120,7 @@ async def test_group_batches_by_credentials_multiple_credentials():
|
|||
"LANGSMITH_API_KEY": "key2", # Different API key
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -113,6 +130,7 @@ async def test_group_batches_by_credentials_multiple_credentials():
|
|||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj2", # Different project
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": None,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -127,6 +145,57 @@ async def test_group_batches_by_credentials_multiple_credentials():
|
|||
assert len(batch_group.queue_objects) == 1 # Each group should have one object
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_batches_by_credentials_with_tenant_id():
|
||||
|
||||
# Test that different tenant_ids create separate groups
|
||||
logger = LangsmithLogger(langsmith_api_key="test-key")
|
||||
|
||||
queue_obj1 = LangsmithQueueObject(
|
||||
data={"test": "data1"},
|
||||
credentials={
|
||||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": "tenant1",
|
||||
},
|
||||
)
|
||||
|
||||
queue_obj2 = LangsmithQueueObject(
|
||||
data={"test": "data2"},
|
||||
credentials={
|
||||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": "tenant2", # Different tenant_id
|
||||
},
|
||||
)
|
||||
|
||||
queue_obj3 = LangsmithQueueObject(
|
||||
data={"test": "data3"},
|
||||
credentials={
|
||||
"LANGSMITH_API_KEY": "key1",
|
||||
"LANGSMITH_PROJECT": "proj1",
|
||||
"LANGSMITH_BASE_URL": "url1",
|
||||
"LANGSMITH_TENANT_ID": "tenant1", # Same as queue_obj1
|
||||
},
|
||||
)
|
||||
|
||||
logger.log_queue = [queue_obj1, queue_obj2, queue_obj3]
|
||||
|
||||
grouped = logger._group_batches_by_credentials()
|
||||
|
||||
# Should have two groups: one for tenant1 (queue_obj1 and queue_obj3), one for tenant2 (queue_obj2)
|
||||
assert len(grouped) == 2
|
||||
for key, batch_group in grouped.items():
|
||||
assert isinstance(key, CredentialsKey)
|
||||
assert key.tenant_id in ["tenant1", "tenant2"]
|
||||
if key.tenant_id == "tenant1":
|
||||
assert len(batch_group.queue_objects) == 2
|
||||
else:
|
||||
assert len(batch_group.queue_objects) == 1
|
||||
|
||||
|
||||
# Test make_dot_order
|
||||
@pytest.mark.asyncio
|
||||
async def test_make_dot_order():
|
||||
|
|
@ -201,10 +270,43 @@ async def test_async_send_batch():
|
|||
call_args = logger.async_httpx_client.post.call_args
|
||||
assert "runs/batch" in call_args[1]["url"]
|
||||
assert "x-api-key" in call_args[1]["headers"]
|
||||
# tenant_id should not be in headers if not provided
|
||||
assert "x-tenant-id" not in call_args[1]["headers"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_langsmith_key_based_logging(mocker):
|
||||
async def test_async_send_batch_with_tenant_id():
|
||||
logger = LangsmithLogger(
|
||||
langsmith_api_key="test-key",
|
||||
langsmith_tenant_id="test-tenant-id"
|
||||
)
|
||||
|
||||
# Mock the httpx client
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
logger.async_httpx_client = AsyncMock()
|
||||
logger.async_httpx_client.post.return_value = mock_response
|
||||
|
||||
# Add test data to queue
|
||||
logger.log_queue = [
|
||||
LangsmithQueueObject(
|
||||
data={"test": "data"}, credentials=logger.default_credentials
|
||||
)
|
||||
]
|
||||
|
||||
await logger.async_send_batch()
|
||||
|
||||
# Verify the API call includes tenant_id header
|
||||
logger.async_httpx_client.post.assert_called_once()
|
||||
call_args = logger.async_httpx_client.post.call_args
|
||||
assert "runs/batch" in call_args[1]["url"]
|
||||
assert "x-api-key" in call_args[1]["headers"]
|
||||
assert "x-tenant-id" in call_args[1]["headers"]
|
||||
assert call_args[1]["headers"]["x-tenant-id"] == "test-tenant-id"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_langsmith_key_based_logging():
|
||||
"""
|
||||
In key based logging langsmith_api_key and langsmith_project are passed directly to litellm.acompletion
|
||||
"""
|
||||
|
|
@ -219,10 +321,11 @@ async def test_langsmith_key_based_logging(mocker):
|
|||
mock_response.text = ""
|
||||
mock_async_httpx_handler.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
mock_get_client = mocker.patch(
|
||||
mock_get_client = patch(
|
||||
"litellm.integrations.langsmith.get_async_httpx_client",
|
||||
return_value=mock_async_httpx_handler
|
||||
)
|
||||
mock_get_client.start()
|
||||
|
||||
litellm.set_verbose = True
|
||||
litellm.DEFAULT_FLUSH_INTERVAL_SECONDS = 1
|
||||
|
|
@ -253,6 +356,8 @@ async def test_langsmith_key_based_logging(mocker):
|
|||
|
||||
# Check headers contain the correct API key
|
||||
assert call_args[1]["headers"]["x-api-key"] == "fake_key_project2"
|
||||
# tenant_id should not be in headers if not provided
|
||||
assert "x-tenant-id" not in call_args[1]["headers"]
|
||||
|
||||
# Verify the request body contains the expected data
|
||||
request_body = call_args[1]["json"]
|
||||
|
|
@ -344,6 +449,8 @@ async def test_langsmith_key_based_logging(mocker):
|
|||
actual_body["post"][0]["session_name"]
|
||||
== expected_body["post"][0]["session_name"]
|
||||
)
|
||||
|
||||
mock_get_client.stop()
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
|
|
|||
|
|
@ -65,31 +65,6 @@ class TestOpentelemetryUnitTests(BaseLoggingCallbackTest):
|
|||
# External spans should only be closed by their creators
|
||||
parent_otel_span.end.assert_not_called()
|
||||
|
||||
def test_init_tracing_respects_existing_tracer_provider(self):
|
||||
"""
|
||||
Unit test: _init_tracing() should respect existing TracerProvider.
|
||||
|
||||
When a TracerProvider already exists (e.g., set by Langfuse SDK),
|
||||
LiteLLM should use it instead of creating a new one.
|
||||
"""
|
||||
from opentelemetry import trace
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
|
||||
# Setup: Create and set an existing TracerProvider
|
||||
tracer_provider = TracerProvider()
|
||||
trace.set_tracer_provider(tracer_provider)
|
||||
existing_provider = trace.get_tracer_provider()
|
||||
|
||||
# Act: Initialize OpenTelemetry integration (should detect existing provider)
|
||||
otel_integration = OpenTelemetry()
|
||||
|
||||
# Assert: The existing provider should still be active
|
||||
current_provider = trace.get_tracer_provider()
|
||||
assert current_provider is existing_provider, (
|
||||
"Existing TracerProvider should be respected and not overridden"
|
||||
)
|
||||
|
||||
def test_get_span_context_detects_active_span(self):
|
||||
"""
|
||||
Unit test: _get_span_context() should auto-detect active spans from global context.
|
||||
|
|
|
|||
|
|
@ -79,6 +79,73 @@ def test_routing_strategy_init(model_list):
|
|||
)
|
||||
|
||||
|
||||
def test_routing_strategy_init_invalid_strategy(model_list):
|
||||
"""Test that invalid routing_strategy raises ValueError with helpful message.
|
||||
|
||||
See: https://github.com/BerriAI/litellm/issues/11330
|
||||
Invalid strategies like 'simple' (without '-shuffle') should fail fast
|
||||
with a clear error, not silently cause 'No deployments available' errors.
|
||||
"""
|
||||
router = Router(model_list=model_list)
|
||||
|
||||
# Test common mistake: "simple" instead of "simple-shuffle"
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
router.routing_strategy_init(
|
||||
routing_strategy="simple",
|
||||
routing_strategy_args={}
|
||||
)
|
||||
|
||||
# Verify error message is helpful
|
||||
error_msg = str(exc_info.value)
|
||||
assert "Invalid routing_strategy" in error_msg
|
||||
assert "simple" in error_msg
|
||||
assert "simple-shuffle" in error_msg # Suggests the correct option
|
||||
# Verify error message tells user WHERE to fix it
|
||||
assert "config.yaml" in error_msg
|
||||
assert "router_settings.routing_strategy" in error_msg
|
||||
assert "Router SDK" in error_msg
|
||||
|
||||
# Test completely invalid strategy
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
router.routing_strategy_init(
|
||||
routing_strategy="not-a-real-strategy",
|
||||
routing_strategy_args={}
|
||||
)
|
||||
assert "Invalid routing_strategy" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_routing_strategy_init_valid_string_strategies(model_list):
|
||||
"""Test that all valid string routing strategies work without error.
|
||||
|
||||
Valid strategies are derived from RoutingStrategy enum values plus 'simple-shuffle'.
|
||||
"""
|
||||
from litellm.types.router import RoutingStrategy
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
|
||||
# All strategies from enum + simple-shuffle (default, not in enum)
|
||||
valid_strategies = ["simple-shuffle"] + [s.value for s in RoutingStrategy]
|
||||
|
||||
for strategy in valid_strategies:
|
||||
# Should not raise
|
||||
router.routing_strategy_init(
|
||||
routing_strategy=strategy, routing_strategy_args={}
|
||||
)
|
||||
|
||||
|
||||
def test_routing_strategy_init_valid_enum_strategies(model_list):
|
||||
"""Test that RoutingStrategy enum values work without error."""
|
||||
from litellm.types.router import RoutingStrategy
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
|
||||
for strategy in RoutingStrategy:
|
||||
# Should not raise when passing enum directly
|
||||
router.routing_strategy_init(
|
||||
routing_strategy=strategy, routing_strategy_args={}
|
||||
)
|
||||
|
||||
|
||||
def test_print_deployment(model_list):
|
||||
"""Test if the api key is masked correctly"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1011,39 +1011,87 @@ def test_multiple_tool_calls_in_single_choice():
|
|||
|
||||
def test_map_reasoning_effort_adds_summary_detailed():
|
||||
"""
|
||||
Test that _map_reasoning_effort adds summary="detailed" when user provides reasoning_effort as a string.
|
||||
Test that _map_reasoning_effort behavior with reasoning_auto_summary flag.
|
||||
|
||||
This ensures that when users pass reasoning_effort in the completions API for OpenAI responses/models,
|
||||
the transformation automatically includes summary="detailed" in the reasoning parameter.
|
||||
By default (flag=False), summary should NOT be added to avoid:
|
||||
1. Breaking for users without verified OpenAI orgs (400 errors)
|
||||
2. Making requests more expensive by including summary reasoning tokens
|
||||
|
||||
When flag is enabled (flag=True or env var), summary="detailed" is added.
|
||||
"""
|
||||
import os
|
||||
|
||||
import litellm
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
# Test all string effort levels
|
||||
# Test all string effort levels - DEFAULT BEHAVIOR (no summary)
|
||||
effort_levels = ["none", "low", "medium", "high", "xhigh", "minimal"]
|
||||
|
||||
for effort in effort_levels:
|
||||
result = handler._map_reasoning_effort(effort)
|
||||
# Save original flag value
|
||||
original_flag = litellm.reasoning_auto_summary
|
||||
original_env = os.environ.get("LITELLM_REASONING_AUTO_SUMMARY")
|
||||
|
||||
try:
|
||||
# Test 1: Default behavior (flag=False, no env var) - NO summary
|
||||
litellm.reasoning_auto_summary = False
|
||||
if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
|
||||
del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
|
||||
|
||||
assert result is not None, f"Result should not be None for effort={effort}"
|
||||
assert result["effort"] == effort, f"Effort should be {effort}"
|
||||
assert result["summary"] == "detailed", f"Summary should be 'detailed' for effort={effort}"
|
||||
for effort in effort_levels:
|
||||
result = handler._map_reasoning_effort(effort)
|
||||
|
||||
assert result is not None, f"Result should not be None for effort={effort}"
|
||||
assert result["effort"] == effort, f"Effort should be {effort}"
|
||||
assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}"
|
||||
|
||||
print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)")
|
||||
|
||||
print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed'")
|
||||
# Test 2: With flag enabled - summary IS added
|
||||
litellm.reasoning_auto_summary = True
|
||||
|
||||
for effort in effort_levels:
|
||||
result = handler._map_reasoning_effort(effort)
|
||||
|
||||
assert result is not None, f"Result should not be None for effort={effort}"
|
||||
assert result["effort"] == effort, f"Effort should be {effort}"
|
||||
assert result["summary"] == "detailed", f"Summary should be 'detailed' when flag is enabled for effort={effort}"
|
||||
|
||||
print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)")
|
||||
|
||||
# Test 3: With env var enabled (flag disabled) - summary IS added
|
||||
litellm.reasoning_auto_summary = False
|
||||
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true"
|
||||
|
||||
result = handler._map_reasoning_effort("high")
|
||||
assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled"
|
||||
print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly")
|
||||
|
||||
# Test 4: Dict input is passed through as-is (no modification)
|
||||
litellm.reasoning_auto_summary = False
|
||||
if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
|
||||
del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
|
||||
|
||||
dict_input = {"effort": "high", "summary": "custom_summary"}
|
||||
result_dict = handler._map_reasoning_effort(dict_input)
|
||||
assert result_dict["effort"] == "high"
|
||||
assert result_dict["summary"] == "custom_summary"
|
||||
print("✓ Dict input is passed through without modification")
|
||||
|
||||
# Test 5: None/unknown values return None
|
||||
result_unknown = handler._map_reasoning_effort("unknown_value")
|
||||
assert result_unknown is None
|
||||
print("✓ Unknown reasoning_effort values return None")
|
||||
|
||||
print("✓ All reasoning_effort behaviors work correctly with flag/env var control")
|
||||
|
||||
# Test that dict input is passed through as-is (no modification)
|
||||
dict_input = {"effort": "high", "summary": "custom_summary"}
|
||||
result_dict = handler._map_reasoning_effort(dict_input)
|
||||
assert result_dict["effort"] == "high"
|
||||
assert result_dict["summary"] == "custom_summary"
|
||||
print("✓ Dict input is passed through without modification")
|
||||
|
||||
# Test that None/unknown values return None
|
||||
result_unknown = handler._map_reasoning_effort("unknown_value")
|
||||
assert result_unknown is None
|
||||
print("✓ Unknown reasoning_effort values return None")
|
||||
|
||||
print("✓ All reasoning_effort string values correctly map to summary='detailed'")
|
||||
finally:
|
||||
# Restore original values
|
||||
litellm.reasoning_auto_summary = original_flag
|
||||
if original_env is not None:
|
||||
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = original_env
|
||||
elif "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
|
||||
del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,90 @@
|
|||
"""
|
||||
Test for Langfuse integration with Gemini cached_tokens bug
|
||||
https://github.com/BerriAI/litellm/issues/18520
|
||||
"""
|
||||
import pytest
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
|
||||
|
||||
def test_cached_tokens_extraction():
|
||||
"""
|
||||
Test that we can extract cached_tokens from prompt_tokens_details.
|
||||
This is the core logic fix for https://github.com/BerriAI/litellm/issues/18520
|
||||
"""
|
||||
# Create usage object like Gemini returns
|
||||
usage = Usage(
|
||||
prompt_tokens=20209,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
cached_tokens=20203,
|
||||
text_tokens=6,
|
||||
),
|
||||
completion_tokens=541,
|
||||
)
|
||||
|
||||
# Simulate the logic from langfuse.py lines 745-757 (after the fix)
|
||||
cache_read_input_tokens = 0 # Default value
|
||||
|
||||
# Check prompt_tokens_details.cached_tokens (the fix)
|
||||
if hasattr(usage, "prompt_tokens_details"):
|
||||
prompt_tokens_details = getattr(usage, "prompt_tokens_details", None)
|
||||
if (
|
||||
prompt_tokens_details is not None
|
||||
and hasattr(prompt_tokens_details, "cached_tokens")
|
||||
):
|
||||
cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
|
||||
if cached_tokens is not None and cached_tokens > 0:
|
||||
cache_read_input_tokens = cached_tokens
|
||||
|
||||
# Verify the fix works
|
||||
assert cache_read_input_tokens == 20203, f"Expected 20203, got {cache_read_input_tokens}"
|
||||
|
||||
|
||||
def test_cached_tokens_not_present():
|
||||
"""Test backward compatibility when cached_tokens is not present"""
|
||||
# Usage without prompt_tokens_details
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
)
|
||||
|
||||
cache_read_input_tokens = 0
|
||||
|
||||
if hasattr(usage, "prompt_tokens_details"):
|
||||
prompt_tokens_details = getattr(usage, "prompt_tokens_details", None)
|
||||
if (
|
||||
prompt_tokens_details is not None
|
||||
and hasattr(prompt_tokens_details, "cached_tokens")
|
||||
):
|
||||
cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
|
||||
if cached_tokens is not None and cached_tokens > 0:
|
||||
cache_read_input_tokens = cached_tokens
|
||||
|
||||
# Should remain 0
|
||||
assert cache_read_input_tokens == 0
|
||||
|
||||
|
||||
def test_cached_tokens_is_zero():
|
||||
"""Test when cached_tokens is explicitly 0"""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
cached_tokens=0,
|
||||
text_tokens=100,
|
||||
),
|
||||
completion_tokens=50,
|
||||
)
|
||||
|
||||
cache_read_input_tokens = 0
|
||||
|
||||
if hasattr(usage, "prompt_tokens_details"):
|
||||
prompt_tokens_details = getattr(usage, "prompt_tokens_details", None)
|
||||
if (
|
||||
prompt_tokens_details is not None
|
||||
and hasattr(prompt_tokens_details, "cached_tokens")
|
||||
):
|
||||
cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
|
||||
if cached_tokens is not None and cached_tokens > 0:
|
||||
cache_read_input_tokens = cached_tokens
|
||||
|
||||
# Should remain 0 when cached_tokens is 0
|
||||
assert cache_read_input_tokens == 0
|
||||
|
|
@ -172,6 +172,86 @@ class TestOpenTelemetryCostBreakdown(unittest.TestCase):
|
|||
assert ("gen_ai.cost.original_cost", 0.004) not in call_args_list
|
||||
|
||||
|
||||
class TestOpenTelemetryProviderInitialization(unittest.TestCase):
|
||||
"""Test suite for verifying provider initialization respects existing providers"""
|
||||
|
||||
def test_init_tracing_respects_existing_tracer_provider(self):
|
||||
"""
|
||||
Unit test: _init_tracing() should respect existing TracerProvider.
|
||||
|
||||
When a TracerProvider already exists (e.g., set by Langfuse SDK),
|
||||
LiteLLM should use it instead of creating a new one.
|
||||
"""
|
||||
from opentelemetry import trace
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
||||
# Setup: Create and set an existing TracerProvider
|
||||
tracer_provider = TracerProvider()
|
||||
trace.set_tracer_provider(tracer_provider)
|
||||
existing_provider = trace.get_tracer_provider()
|
||||
|
||||
# Act: Initialize OpenTelemetry integration (should detect existing provider)
|
||||
otel_integration = OpenTelemetry()
|
||||
|
||||
# Assert: The existing provider should still be active
|
||||
current_provider = trace.get_tracer_provider()
|
||||
assert current_provider is existing_provider, (
|
||||
"Existing TracerProvider should be respected and not overridden"
|
||||
)
|
||||
|
||||
@patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True)
|
||||
def test_init_metrics_respects_existing_meter_provider(self):
|
||||
"""
|
||||
Unit test: _init_metrics() should respect existing MeterProvider.
|
||||
|
||||
When a MeterProvider already exists (e.g., set by Langfuse SDK),
|
||||
LiteLLM should use it instead of creating a new one.
|
||||
"""
|
||||
from opentelemetry import metrics
|
||||
from opentelemetry.sdk.metrics import MeterProvider
|
||||
|
||||
# Create and set an existing MeterProvider
|
||||
meter_provider = MeterProvider()
|
||||
metrics.set_meter_provider(meter_provider)
|
||||
existing_provider = metrics.get_meter_provider()
|
||||
|
||||
# Act: Initialize OpenTelemetry integration (should detect existing provider)
|
||||
config = OpenTelemetryConfig.from_env()
|
||||
otel_integration = OpenTelemetry(config=config)
|
||||
|
||||
# Assert: The existing provider should still be active
|
||||
current_provider = metrics.get_meter_provider()
|
||||
assert current_provider is existing_provider, (
|
||||
"Existing MeterProvider should be respected and not overridden"
|
||||
)
|
||||
|
||||
@patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS": "true"}, clear=True)
|
||||
def test_init_logs_respects_existing_logger_provider(self):
|
||||
"""
|
||||
Unit test: _init_logs() should respect existing LoggerProvider.
|
||||
|
||||
When a LoggerProvider already exists (e.g., set by Langfuse SDK),
|
||||
LiteLLM should use it instead of creating a new one.
|
||||
"""
|
||||
from opentelemetry._logs import get_logger_provider, set_logger_provider
|
||||
from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider
|
||||
|
||||
# Create and set an existing LoggerProvider
|
||||
logger_provider = OTLoggerProvider()
|
||||
set_logger_provider(logger_provider)
|
||||
existing_provider = get_logger_provider()
|
||||
|
||||
# Act: Initialize OpenTelemetry integration (should detect existing provider)
|
||||
config = OpenTelemetryConfig.from_env()
|
||||
otel_integration = OpenTelemetry(config=config)
|
||||
|
||||
# Assert: The existing provider should still be active
|
||||
current_provider = get_logger_provider()
|
||||
assert current_provider is existing_provider, (
|
||||
"Existing LoggerProvider should be respected and not overridden"
|
||||
)
|
||||
|
||||
|
||||
class TestOpenTelemetry(unittest.TestCase):
|
||||
POLL_INTERVAL = 0.05
|
||||
POLL_TIMEOUT = 2.0
|
||||
|
|
@ -620,7 +700,6 @@ class TestOpenTelemetry(unittest.TestCase):
|
|||
self.assertEqual(attributes.get("extra.attr"), "extra-value")
|
||||
|
||||
|
||||
|
||||
def test_handle_success_spans_only(self):
|
||||
# make sure neither events nor metrics is on
|
||||
os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None)
|
||||
|
|
@ -687,11 +766,8 @@ class TestOpenTelemetry(unittest.TestCase):
|
|||
logs = log_exporter.get_finished_logs()
|
||||
self.assertFalse(logs, "Did not expect any logs")
|
||||
|
||||
@patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True)
|
||||
def test_handle_success_spans_and_metrics(self):
|
||||
# only metrics on
|
||||
os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None)
|
||||
os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_METRICS"] = "true"
|
||||
|
||||
# ─── build in‐memory OTEL providers/exporters ─────────────────────────────
|
||||
span_exporter = InMemorySpanExporter()
|
||||
tracer_provider = TracerProvider()
|
||||
|
|
@ -1320,6 +1396,23 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
|
|||
)
|
||||
self.assertEqual(normalized, "http://collector:4317/v1/logs")
|
||||
|
||||
def test_get_metric_reader_uses_http_exporter_for_http_protobuf(self):
|
||||
"""Test that http/protobuf protocol uses OTLPMetricExporterHTTP"""
|
||||
from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
|
||||
OTLPMetricExporter,
|
||||
)
|
||||
from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader
|
||||
|
||||
config = OpenTelemetryConfig(
|
||||
exporter="http/protobuf", endpoint="http://collector:4318"
|
||||
)
|
||||
otel = OpenTelemetry(config=config)
|
||||
|
||||
reader = otel._get_metric_reader()
|
||||
|
||||
self.assertIsInstance(reader, PeriodicExportingMetricReader)
|
||||
self.assertIsInstance(reader._exporter, OTLPMetricExporter)
|
||||
|
||||
|
||||
class TestOpenTelemetryExternalSpan(unittest.TestCase):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -691,6 +691,29 @@ async def test_streaming_completion_start_time(logging_obj: Logging):
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_streaming_bad_request_not_midstream(logging_obj: Logging):
|
||||
"""Ensure Vertex bad request errors surface as 400, not mid-stream fallbacks."""
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError
|
||||
|
||||
async def _raise_bad_request(**kwargs):
|
||||
raise VertexAIError(status_code=400, message="invalid maxOutputTokens", headers=None)
|
||||
|
||||
response = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
model="gemini-3-pro-preview",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="vertex_ai_beta",
|
||||
make_call=_raise_bad_request,
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError) as excinfo:
|
||||
await response.__anext__()
|
||||
|
||||
assert getattr(excinfo.value, "status_code", None) == 400
|
||||
assert "invalid maxOutputTokens" in str(excinfo.value)
|
||||
|
||||
|
||||
def test_streaming_handler_with_created_time_propagation(
|
||||
initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging
|
||||
):
|
||||
|
|
|
|||
|
|
@ -0,0 +1,355 @@
|
|||
"""
|
||||
Test cases for functionCall args serialization in Vertex AI Gemini.
|
||||
|
||||
This test file specifically tests the edge cases where Vertex AI might return
|
||||
functionCall args in unexpected formats that could lead to invalid JSON strings
|
||||
like: {"x":"x"}{"a":"a"}
|
||||
"""
|
||||
import json
|
||||
from typing import List, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai import HttpxPartType
|
||||
|
||||
|
||||
class TestFunctionCallArgsSerialization:
|
||||
"""Test cases for functionCall args serialization edge cases."""
|
||||
|
||||
def test_normal_dict_args(self):
|
||||
"""Test normal case: args is a dict."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": {"location": "Boston", "unit": "celsius"},
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
assert tools[0]["function"]["name"] == "get_weather"
|
||||
|
||||
# Verify arguments is a valid JSON string
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
# Should be valid JSON
|
||||
parsed = json.loads(arguments)
|
||||
assert parsed == {"location": "Boston", "unit": "celsius"}
|
||||
|
||||
def test_none_args(self):
|
||||
"""Test case: args is None."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": None,
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
# Should serialize None to "null" or empty dict
|
||||
assert isinstance(arguments, str)
|
||||
parsed = json.loads(arguments)
|
||||
# json.dumps(None) returns "null"
|
||||
assert parsed is None or parsed == {}
|
||||
|
||||
def test_args_as_string_valid_json(self):
|
||||
"""Test case: args is already a valid JSON string."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": '{"location": "Boston"}', # String, not dict
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
# If args is a string, json.dumps will double-encode it
|
||||
# This would result in: "{\"location\": \"Boston\"}"
|
||||
assert isinstance(arguments, str)
|
||||
# This is the problematic case - string gets double-encoded
|
||||
# The result would be a JSON string containing a JSON string
|
||||
parsed = json.loads(arguments)
|
||||
# If it's double-encoded, parsed would be a string, not a dict
|
||||
if isinstance(parsed, str):
|
||||
# Double-encoded case
|
||||
inner_parsed = json.loads(parsed)
|
||||
assert inner_parsed == {"location": "Boston"}
|
||||
else:
|
||||
# Normal case (shouldn't happen if args is string)
|
||||
assert parsed == {"location": "Boston"}
|
||||
|
||||
def test_args_as_string_invalid_json_concatenated(self):
|
||||
"""Test case: args is a string with concatenated JSON objects (the bug case).
|
||||
|
||||
When args is a string like '{"x":"x"}{"a":"a"}', json.dumps() will serialize it
|
||||
as a JSON string, resulting in: "{\"x\":\"x\"}{\"a\":\"a\"}"
|
||||
This is a valid JSON string (the outer quotes), but the content inside is invalid JSON.
|
||||
When you try to parse the inner content, it fails.
|
||||
"""
|
||||
# This simulates the case where Vertex might return something like:
|
||||
# args = '{"x":"x"}{"a":"a"}' # Two JSON objects concatenated
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": '{"x":"x"}{"a":"a"}', # Invalid concatenated JSON
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
|
||||
# json.dumps() on a string will escape it, so we get:
|
||||
# arguments = '"{\\"x\\":\\"x\\"}{\\"a\\":\\"a\\"}"'
|
||||
# This is a valid JSON string (the outer quotes), but the inner content is invalid
|
||||
parsed_outer = json.loads(arguments)
|
||||
assert isinstance(parsed_outer, str)
|
||||
|
||||
# The inner string is invalid JSON (two objects concatenated)
|
||||
# This is the bug: the inner content cannot be parsed as valid JSON
|
||||
with pytest.raises(json.JSONDecodeError):
|
||||
json.loads(parsed_outer)
|
||||
|
||||
# The arguments string would be: "{\"x\":\"x\"}{\"a\":\"a\"}"
|
||||
# Which when parsed gives: '{"x":"x"}{"a":"a"}' (invalid JSON)
|
||||
|
||||
def test_args_as_array(self):
|
||||
"""Test case: args is an array (unexpected but possible)."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": [{"x": "x"}, {"a": "a"}], # Array of objects
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
# Should serialize array correctly
|
||||
parsed = json.loads(arguments)
|
||||
assert parsed == [{"x": "x"}, {"a": "a"}]
|
||||
|
||||
def test_args_missing_key(self):
|
||||
"""Test case: args key is missing from functionCall.
|
||||
|
||||
This will raise a KeyError because the code directly accesses part["functionCall"]["args"]
|
||||
without checking if the key exists. This is a bug that should be fixed.
|
||||
"""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
# args key missing
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
# This should raise KeyError because args key is missing
|
||||
with pytest.raises(KeyError):
|
||||
VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
def test_multiple_function_calls(self):
|
||||
"""Test case: multiple function calls in parts."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": {"location": "Boston"},
|
||||
}
|
||||
},
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_time",
|
||||
"args": {"timezone": "EST"},
|
||||
}
|
||||
},
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 2
|
||||
assert tools[0]["function"]["name"] == "get_weather"
|
||||
assert tools[1]["function"]["name"] == "get_time"
|
||||
|
||||
# Both should have valid JSON arguments
|
||||
args1 = json.loads(tools[0]["function"]["arguments"])
|
||||
args2 = json.loads(tools[1]["function"]["arguments"])
|
||||
assert args1 == {"location": "Boston"}
|
||||
assert args2 == {"timezone": "EST"}
|
||||
|
||||
def test_args_with_vertex_protobuf_format(self):
|
||||
"""Test case: args in Vertex protobuf format with string_value, etc."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": {
|
||||
"location": {"string_value": "Boston, MA"},
|
||||
"unit": {"string_value": "celsius"},
|
||||
},
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
# Should serialize the nested structure correctly
|
||||
parsed = json.loads(arguments)
|
||||
assert "location" in parsed
|
||||
assert "unit" in parsed
|
||||
|
||||
def test_args_as_empty_dict(self):
|
||||
"""Test case: args is an empty dict."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": {},
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
parsed = json.loads(arguments)
|
||||
assert parsed == {}
|
||||
|
||||
def test_args_with_special_characters(self):
|
||||
"""Test case: args contains special characters that need escaping."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": {
|
||||
"location": 'Boston, MA "downtown"',
|
||||
"note": "Line 1\nLine 2",
|
||||
},
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
# Should handle special characters correctly
|
||||
parsed = json.loads(arguments)
|
||||
assert parsed["location"] == 'Boston, MA "downtown"'
|
||||
assert parsed["note"] == "Line 1\nLine 2"
|
||||
|
||||
def test_args_as_list_of_strings_that_look_like_json(self):
|
||||
"""Test case: args is a list containing strings that look like JSON objects."""
|
||||
# This could potentially cause issues if not handled correctly
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": ['{"x":"x"}', '{"a":"a"}'], # List of JSON strings
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
# Should serialize list correctly
|
||||
parsed = json.loads(arguments)
|
||||
assert isinstance(parsed, list)
|
||||
assert parsed == ['{"x":"x"}', '{"a":"a"}']
|
||||
|
||||
def test_args_as_dict_with_nested_structures(self):
|
||||
"""Test case: args contains nested dicts and lists."""
|
||||
parts: List[HttpxPartType] = [
|
||||
{
|
||||
"functionCall": {
|
||||
"name": "complex_function",
|
||||
"args": {
|
||||
"nested": {"key": "value"},
|
||||
"list": [1, 2, 3],
|
||||
"mixed": [{"a": 1}, {"b": 2}],
|
||||
},
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
function, tools, idx = VertexGeminiConfig._transform_parts(
|
||||
parts=parts, cumulative_tool_call_idx=0, is_function_call=False
|
||||
)
|
||||
|
||||
assert tools is not None
|
||||
assert len(tools) == 1
|
||||
arguments = tools[0]["function"]["arguments"]
|
||||
assert isinstance(arguments, str)
|
||||
parsed = json.loads(arguments)
|
||||
assert parsed["nested"] == {"key": "value"}
|
||||
assert parsed["list"] == [1, 2, 3]
|
||||
assert parsed["mixed"] == [{"a": 1}, {"b": 2}]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
|
||||
|
|
@ -10,11 +10,13 @@ from pydantic import BaseModel
|
|||
|
||||
import litellm
|
||||
from litellm import ModelResponse, completion
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai import UsageMetadata
|
||||
from litellm.types.utils import ChoiceLogprobs, Usage
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
|
||||
def test_top_logprobs():
|
||||
|
|
@ -1605,6 +1607,39 @@ def test_vertex_ai_annotation_streaming_events():
|
|||
assert "Weather information" in annotation["url_citation"]["title"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_ai_streaming_bad_request_is_not_wrapped():
|
||||
class DummyLogging:
|
||||
def __init__(self):
|
||||
self.model_call_details = {"litellm_params": {}}
|
||||
self.optional_params = {}
|
||||
self.messages = []
|
||||
self.completion_start_time = None
|
||||
self.stream_options = None
|
||||
|
||||
def failure_handler(self, *args, **kwargs):
|
||||
return None
|
||||
|
||||
async def async_failure_handler(self, *args, **kwargs):
|
||||
return None
|
||||
|
||||
async def failing_make_call(client=None, **kwargs):
|
||||
raise VertexAIError(status_code=400, message="bad input", headers={})
|
||||
|
||||
stream = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
make_call=failing_make_call,
|
||||
model="gemini-3-pro-preview",
|
||||
logging_obj=DummyLogging(),
|
||||
custom_llm_provider="vertex_ai_beta",
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
await stream.__anext__()
|
||||
|
||||
assert getattr(exc_info.value, "status_code", None) == 400
|
||||
|
||||
|
||||
def test_vertex_ai_annotation_conversion():
|
||||
"""
|
||||
Test the conversion of Vertex AI grounding metadata to OpenAI annotations.
|
||||
|
|
|
|||
|
|
@ -354,7 +354,7 @@ async def test_register_client_remote_registration_success():
|
|||
|
||||
request_payload = {
|
||||
"client_name": "Litellm Proxy",
|
||||
"grant_types": ["authorization_code"],
|
||||
"grant_types": ["authorization_code", "refresh_token"],
|
||||
"response_types": ["code"],
|
||||
"token_endpoint_auth_method": "client_secret_post",
|
||||
}
|
||||
|
|
@ -556,9 +556,33 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto():
|
|||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
oauth_protected_resource_mcp,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from fastapi import Request
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
# Clear registry
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
# Create mock OAuth2 server
|
||||
oauth2_server = MCPServer(
|
||||
server_id="test_oauth_server",
|
||||
name="test_oauth",
|
||||
server_name="test_oauth",
|
||||
alias="test_oauth",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
authorization_url="https://provider.com/oauth/authorize",
|
||||
token_url="https://provider.com/oauth/token",
|
||||
scopes=["read", "write"],
|
||||
)
|
||||
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
|
||||
|
||||
# Mock request with http base_url but X-Forwarded-Proto: https
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
|
@ -568,13 +592,14 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto():
|
|||
# Call the endpoint
|
||||
response = await oauth_protected_resource_mcp(
|
||||
request=mock_request,
|
||||
mcp_server_name="test_server",
|
||||
mcp_server_name="test_oauth",
|
||||
)
|
||||
|
||||
# Verify response uses HTTPS URLs
|
||||
assert response["authorization_servers"][0].startswith(
|
||||
"https://litellm.example.com/"
|
||||
)
|
||||
assert response["scopes_supported"] == oauth2_server.scopes
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -584,9 +609,33 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto():
|
|||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
oauth_authorization_server_mcp,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from fastapi import Request
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
# Clear registry
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
# Create mock OAuth2 server
|
||||
oauth2_server = MCPServer(
|
||||
server_id="test_oauth_server",
|
||||
name="test_oauth",
|
||||
server_name="test_oauth",
|
||||
alias="test_oauth",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
authorization_url="https://provider.com/oauth/authorize",
|
||||
token_url="https://provider.com/oauth/token",
|
||||
scopes=["read", "write"],
|
||||
)
|
||||
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
|
||||
|
||||
# Mock request with http base_url but X-Forwarded-Proto: https
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
|
@ -596,13 +645,15 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto():
|
|||
# Call the endpoint
|
||||
response = await oauth_authorization_server_mcp(
|
||||
request=mock_request,
|
||||
mcp_server_name="test_server",
|
||||
mcp_server_name="test_oauth",
|
||||
)
|
||||
|
||||
# Verify response uses HTTPS URLs
|
||||
assert response["authorization_endpoint"].startswith("https://litellm.example.com/")
|
||||
assert response["token_endpoint"].startswith("https://litellm.example.com/")
|
||||
assert response["registration_endpoint"].startswith("https://litellm.example.com/")
|
||||
assert response["grant_types_supported"] == ["authorization_code", "refresh_token"]
|
||||
assert response["scopes_supported"] == oauth2_server.scopes
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -594,7 +594,26 @@ class TestMCPServerManager:
|
|||
assert (
|
||||
server.registration_url == "https://discovered.example.com/register"
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
config = {
|
||||
"example": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"auth_type": MCPAuth.oauth2,
|
||||
"scopes": ["config"],
|
||||
"authorization_url": "https://config.example.com/auth",
|
||||
}
|
||||
}
|
||||
|
||||
await manager.load_servers_from_config(config)
|
||||
|
||||
# Initialize the tool mapping
|
||||
await manager._initialize_tool_name_to_mcp_server_name_mapping()
|
||||
assert manager.tool_name_to_mcp_server_name_mapping == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_handles_missing_server_alias(self):
|
||||
"""Test that list_tools handles servers without alias gracefully"""
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
Unit tests for Qualifire guardrail integration.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -139,76 +140,98 @@ class TestQualifireGuardrailEvaluateKwargs:
|
|||
@pytest.mark.asyncio
|
||||
async def test_evaluate_called_with_prompt_injections(self):
|
||||
"""Test that evaluate is called with prompt_injections enabled."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
|
||||
QualifireGuardrail,
|
||||
)
|
||||
# Mock the qualifire module and its types
|
||||
mock_qualifire_types = MagicMock()
|
||||
mock_llm_message = MagicMock()
|
||||
mock_llm_tool_call = MagicMock()
|
||||
mock_message_instance = MagicMock()
|
||||
mock_llm_message.return_value = mock_message_instance
|
||||
|
||||
mock_qualifire_types.LLMMessage = mock_llm_message
|
||||
mock_qualifire_types.LLMToolCall = mock_llm_tool_call
|
||||
|
||||
with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
|
||||
QualifireGuardrail,
|
||||
)
|
||||
|
||||
guardrail = QualifireGuardrail(
|
||||
api_key="test_key",
|
||||
prompt_injections=True,
|
||||
guardrail_name="test_guardrail",
|
||||
)
|
||||
guardrail = QualifireGuardrail(
|
||||
api_key="test_key",
|
||||
prompt_injections=True,
|
||||
guardrail_name="test_guardrail",
|
||||
)
|
||||
|
||||
# Mock the client
|
||||
mock_client = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.score = 100
|
||||
mock_result.status = "completed"
|
||||
mock_result.evaluationResults = []
|
||||
mock_client.evaluate.return_value = mock_result
|
||||
guardrail._client = mock_client
|
||||
# Mock the client
|
||||
mock_client = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.score = 100
|
||||
mock_result.status = "completed"
|
||||
mock_result.evaluationResults = []
|
||||
mock_client.evaluate.return_value = mock_result
|
||||
guardrail._client = mock_client
|
||||
|
||||
messages = [{"role": "user", "content": "Hello, world!"}]
|
||||
messages = [{"role": "user", "content": "Hello, world!"}]
|
||||
|
||||
await guardrail._run_qualifire_check(
|
||||
messages=messages, output=None, dynamic_params={}
|
||||
)
|
||||
await guardrail._run_qualifire_check(
|
||||
messages=messages, output=None, dynamic_params={}
|
||||
)
|
||||
|
||||
# Verify evaluate was called with correct kwargs
|
||||
mock_client.evaluate.assert_called_once()
|
||||
call_kwargs = mock_client.evaluate.call_args[1]
|
||||
assert call_kwargs["prompt_injections"] is True
|
||||
assert "messages" in call_kwargs
|
||||
# Verify evaluate was called with correct kwargs
|
||||
mock_client.evaluate.assert_called_once()
|
||||
call_kwargs = mock_client.evaluate.call_args[1]
|
||||
assert call_kwargs["prompt_injections"] is True
|
||||
assert "messages" in call_kwargs
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_called_with_multiple_checks(self):
|
||||
"""Test that evaluate is called with multiple checks enabled."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
|
||||
QualifireGuardrail,
|
||||
)
|
||||
# Mock the qualifire module and its types
|
||||
mock_qualifire_types = MagicMock()
|
||||
mock_llm_message = MagicMock()
|
||||
mock_llm_tool_call = MagicMock()
|
||||
mock_message_instance = MagicMock()
|
||||
mock_llm_message.return_value = mock_message_instance
|
||||
|
||||
mock_qualifire_types.LLMMessage = mock_llm_message
|
||||
mock_qualifire_types.LLMToolCall = mock_llm_tool_call
|
||||
|
||||
with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import (
|
||||
QualifireGuardrail,
|
||||
)
|
||||
|
||||
guardrail = QualifireGuardrail(
|
||||
api_key="test_key",
|
||||
prompt_injections=True,
|
||||
pii_check=True,
|
||||
hallucinations_check=True,
|
||||
assertions=["Output must be valid JSON"],
|
||||
guardrail_name="test_guardrail",
|
||||
)
|
||||
guardrail = QualifireGuardrail(
|
||||
api_key="test_key",
|
||||
prompt_injections=True,
|
||||
pii_check=True,
|
||||
hallucinations_check=True,
|
||||
assertions=["Output must be valid JSON"],
|
||||
guardrail_name="test_guardrail",
|
||||
)
|
||||
|
||||
# Mock the client
|
||||
mock_client = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.score = 100
|
||||
mock_result.status = "completed"
|
||||
mock_result.evaluationResults = []
|
||||
mock_client.evaluate.return_value = mock_result
|
||||
guardrail._client = mock_client
|
||||
# Mock the client
|
||||
mock_client = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.score = 100
|
||||
mock_result.status = "completed"
|
||||
mock_result.evaluationResults = []
|
||||
mock_client.evaluate.return_value = mock_result
|
||||
guardrail._client = mock_client
|
||||
|
||||
messages = [{"role": "user", "content": "Hello, world!"}]
|
||||
messages = [{"role": "user", "content": "Hello, world!"}]
|
||||
|
||||
await guardrail._run_qualifire_check(
|
||||
messages=messages, output="Test output", dynamic_params={}
|
||||
)
|
||||
await guardrail._run_qualifire_check(
|
||||
messages=messages, output="Test output", dynamic_params={}
|
||||
)
|
||||
|
||||
# Verify evaluate was called with correct kwargs
|
||||
mock_client.evaluate.assert_called_once()
|
||||
call_kwargs = mock_client.evaluate.call_args[1]
|
||||
assert call_kwargs["prompt_injections"] is True
|
||||
assert call_kwargs["pii_check"] is True
|
||||
assert call_kwargs["hallucinations_check"] is True
|
||||
assert call_kwargs["assertions"] == ["Output must be valid JSON"]
|
||||
assert call_kwargs["output"] == "Test output"
|
||||
# Verify evaluate was called with correct kwargs
|
||||
mock_client.evaluate.assert_called_once()
|
||||
call_kwargs = mock_client.evaluate.call_args[1]
|
||||
assert call_kwargs["prompt_injections"] is True
|
||||
assert call_kwargs["pii_check"] is True
|
||||
assert call_kwargs["hallucinations_check"] is True
|
||||
assert call_kwargs["assertions"] == ["Output must be valid JSON"]
|
||||
assert call_kwargs["output"] == "Test output"
|
||||
|
||||
|
||||
class TestQualifireGuardrailCheckIfFlagged:
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ class TestKeyManagementEventHooksIndependentOperations:
|
|||
mock_data = MagicMock()
|
||||
mock_data.key_alias = "test-key-alias"
|
||||
mock_data.team_id = None
|
||||
mock_data.send_invite_email = True
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump.return_value = {"key": "sk-test", "token": "test-token"}
|
||||
|
|
@ -59,6 +60,10 @@ class TestKeyManagementEventHooksIndependentOperations:
|
|||
KeyManagementEventHooks,
|
||||
"_store_virtual_key_in_secret_manager",
|
||||
side_effect=mock_store_secret,
|
||||
), patch.object(
|
||||
KeyManagementEventHooks,
|
||||
"_is_email_sending_enabled",
|
||||
return_value=True,
|
||||
), patch(
|
||||
"litellm.store_audit_logs", False
|
||||
), patch(
|
||||
|
|
@ -96,6 +101,7 @@ class TestKeyManagementEventHooksIndependentOperations:
|
|||
mock_data = MagicMock()
|
||||
mock_data.key_alias = "test-key-alias"
|
||||
mock_data.team_id = None
|
||||
mock_data.send_invite_email = True
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump.return_value = {"key": "sk-test", "token": "test-token"}
|
||||
|
|
@ -115,6 +121,10 @@ class TestKeyManagementEventHooksIndependentOperations:
|
|||
KeyManagementEventHooks,
|
||||
"_store_virtual_key_in_secret_manager",
|
||||
side_effect=mock_store_secret_raises,
|
||||
), patch.object(
|
||||
KeyManagementEventHooks,
|
||||
"_is_email_sending_enabled",
|
||||
return_value=True,
|
||||
), patch(
|
||||
"litellm.store_audit_logs", False
|
||||
), patch(
|
||||
|
|
|
|||
|
|
@ -231,6 +231,40 @@ class TestListMCPServers:
|
|||
assert server.url == "https://mcp.deepwiki.com/mcp"
|
||||
assert server.transport == "http"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_mcp_servers_view_all_mode(self):
|
||||
"""Users should see all MCP servers when view_all mode is enabled."""
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
mock_servers = [
|
||||
generate_mock_mcp_server_db_record(server_id="server-1", alias="One"),
|
||||
generate_mock_mcp_server_db_record(server_id="server-2", alias="Two"),
|
||||
]
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_all_mcp_servers_unfiltered = AsyncMock(
|
||||
return_value=mock_servers
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
|
||||
return_value="view_all",
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_all_mcp_servers,
|
||||
)
|
||||
|
||||
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
|
||||
|
||||
assert len(result) == 2
|
||||
assert {server.server_id for server in result} == {"server-1", "server-2"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_mcp_servers_combined_config_and_db(self):
|
||||
"""
|
||||
|
|
@ -1096,6 +1130,51 @@ class TestHealthCheckServers:
|
|||
assert result[0]["server_id"] == "server-1"
|
||||
assert result[0]["status"] == "healthy"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_view_all_mode(self):
|
||||
"""view_all mode should return health info for all MCP servers."""
|
||||
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
health_check_servers,
|
||||
)
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
health_result_one = generate_mock_mcp_server_db_record(
|
||||
server_id="server-1", alias="One"
|
||||
)
|
||||
health_result_one.status = "healthy"
|
||||
|
||||
health_result_two = generate_mock_mcp_server_db_record(
|
||||
server_id="server-2", alias="Two"
|
||||
)
|
||||
health_result_two.status = "unhealthy"
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_all_mcp_servers_with_health_unfiltered = AsyncMock(
|
||||
return_value=[health_result_one, health_result_two]
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
|
||||
return_value="view_all",
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
result = await health_check_servers(
|
||||
server_ids=None,
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert len(result) == 2
|
||||
assert result[0]["server_id"] == "server-1"
|
||||
assert result[0]["status"] == "healthy"
|
||||
assert result[1]["server_id"] == "server-2"
|
||||
assert result[1]["status"] == "unhealthy"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_unauthorized_servers(self):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -742,18 +742,16 @@ class TestProxySettingEndpoints:
|
|||
):
|
||||
"""Test updating UI settings with an allowlisted field"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
class MockUser:
|
||||
def __init__(self, user_role):
|
||||
self.user_role = user_role
|
||||
|
||||
async def mock_admin_auth():
|
||||
return MockUser(LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.user_api_key_auth",
|
||||
mock_admin_auth,
|
||||
# Override the FastAPI dependency with a proper mock
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
user_id="test-user-123",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
|
||||
|
|
@ -761,7 +759,11 @@ class TestProxySettingEndpoints:
|
|||
|
||||
payload = {"disable_model_add_for_internal_users": True}
|
||||
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
try:
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
finally:
|
||||
# Clean up the dependency override
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
|
|
@ -780,18 +782,16 @@ class TestProxySettingEndpoints:
|
|||
):
|
||||
"""Test non-allowlisted UI settings are ignored on update"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
class MockUser:
|
||||
def __init__(self, user_role):
|
||||
self.user_role = user_role
|
||||
|
||||
async def mock_admin_auth():
|
||||
return MockUser(LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.user_api_key_auth",
|
||||
mock_admin_auth,
|
||||
# Override the FastAPI dependency with a proper mock
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
user_id="test-user-123",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
|
||||
|
|
@ -802,7 +802,11 @@ class TestProxySettingEndpoints:
|
|||
"unsupported_flag": True,
|
||||
}
|
||||
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
try:
|
||||
response = client.patch("/update/ui_settings", json=payload)
|
||||
finally:
|
||||
# Clean up the dependency override
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { getSSOSettings } from "@/components/networking";
|
||||
import { useQuery, UseQueryResult } from "@tanstack/react-query";
|
||||
import { createQueryKeys } from "../common/queryKeysFactory";
|
||||
import { getSSOSettings } from "@/components/networking";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
|
||||
export interface SSOFieldSchema {
|
||||
description: string;
|
||||
|
|
@ -27,13 +27,15 @@ export interface SSOSettingsValues {
|
|||
proxy_base_url: string | null;
|
||||
user_email: string | null;
|
||||
ui_access_mode: string | null;
|
||||
role_mappings: {
|
||||
provider: string;
|
||||
group_claim: string;
|
||||
default_role: "internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer";
|
||||
roles: {
|
||||
[key: string]: string[];
|
||||
};
|
||||
role_mappings: RoleMappings;
|
||||
}
|
||||
|
||||
export interface RoleMappings {
|
||||
provider: string;
|
||||
group_claim: string;
|
||||
default_role: "internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer";
|
||||
roles: {
|
||||
[key: string]: string[];
|
||||
};
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ import { useQueryClient } from "@tanstack/react-query";
|
|||
import { Col, Grid, Icon, Tab, TabGroup, TabList, TabPanel, TabPanels, Text } from "@tremor/react";
|
||||
import type { UploadProps } from "antd";
|
||||
import { Form, Typography } from "antd";
|
||||
import { PlusCircleOutlined } from "@ant-design/icons";
|
||||
import React, { useEffect, useMemo, useState } from "react";
|
||||
import AddModelTab from "../../../components/add_model/add_model_tab";
|
||||
import HealthCheckComponent from "../../../components/model_dashboard/HealthCheckComponent";
|
||||
|
|
@ -274,6 +275,30 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({ premiumUser, te
|
|||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Missing Provider Banner */}
|
||||
<div className="mb-4 px-4 py-3 bg-blue-50 rounded-lg border border-blue-100 flex items-center gap-4">
|
||||
<div className="flex-shrink-0 w-10 h-10 bg-white rounded-full flex items-center justify-center border border-blue-200">
|
||||
<PlusCircleOutlined style={{ fontSize: '18px', color: '#6366f1' }} />
|
||||
</div>
|
||||
<div className="flex-1 min-w-0">
|
||||
<h4 className="text-gray-900 font-semibold text-sm m-0">Missing a provider?</h4>
|
||||
<p className="text-gray-500 text-xs m-0 mt-0.5">
|
||||
The LiteLLM engineering team is constantly adding support for new LLM models, providers, endpoints. If you don't see the one you need, let us know and we'll prioritize it.
|
||||
</p>
|
||||
</div>
|
||||
<a
|
||||
href="https://models.litellm.ai/?request=true"
|
||||
target="_blank"
|
||||
rel="noopener noreferrer"
|
||||
className="flex-shrink-0 inline-flex items-center gap-2 px-4 py-2 bg-[#6366f1] hover:bg-[#5558e3] text-white text-sm font-medium rounded-lg transition-colors"
|
||||
>
|
||||
Request Provider
|
||||
<svg xmlns="http://www.w3.org/2000/svg" className="h-4 w-4" fill="none" viewBox="0 0 24 24" stroke="currentColor" strokeWidth={2}>
|
||||
<path strokeLinecap="round" strokeLinejoin="round" d="M10 6H6a2 2 0 00-2 2v10a2 2 0 002 2h10a2 2 0 002-2v-4M14 4h6m0 0v6m0-6L10 14" />
|
||||
</svg>
|
||||
</a>
|
||||
</div>
|
||||
{selectedModelId && !isLoading ? (
|
||||
<ModelInfoView
|
||||
modelId={selectedModelId}
|
||||
|
|
|
|||
|
|
@ -189,11 +189,16 @@ const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormS
|
|||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings}
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.use_role_mappings !== currentValues.use_role_mappings ||
|
||||
prevValues.sso_provider !== currentValues.sso_provider
|
||||
}
|
||||
>
|
||||
{({ getFieldValue }) => {
|
||||
const useRoleMappings = getFieldValue("use_role_mappings");
|
||||
return useRoleMappings ? (
|
||||
const provider = getFieldValue("sso_provider");
|
||||
const supportsRoleMappings = provider === "okta" || provider === "generic";
|
||||
return useRoleMappings && supportsRoleMappings ? (
|
||||
<Form.Item
|
||||
label="Group Claim"
|
||||
name="group_claim"
|
||||
|
|
@ -207,11 +212,16 @@ const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormS
|
|||
|
||||
<Form.Item
|
||||
noStyle
|
||||
shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings}
|
||||
shouldUpdate={(prevValues, currentValues) =>
|
||||
prevValues.use_role_mappings !== currentValues.use_role_mappings ||
|
||||
prevValues.sso_provider !== currentValues.sso_provider
|
||||
}
|
||||
>
|
||||
{({ getFieldValue }) => {
|
||||
const useRoleMappings = getFieldValue("use_role_mappings");
|
||||
return useRoleMappings ? (
|
||||
const provider = getFieldValue("sso_provider");
|
||||
const supportsRoleMappings = provider === "okta" || provider === "generic";
|
||||
return useRoleMappings && supportsRoleMappings ? (
|
||||
<>
|
||||
<Form.Item label="Default Role" name="default_role" initialValue="Internal User">
|
||||
<Select>
|
||||
|
|
|
|||
|
|
@ -1,20 +1,60 @@
|
|||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { render, screen } from "@testing-library/react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import DeleteSSOSettingsModal from "./DeleteSSOSettingsModal";
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/sso/useSSOSettings", () => ({
|
||||
useSSOSettings: vi.fn(() => ({
|
||||
data: {
|
||||
values: {
|
||||
google_client_id: "test-client-id",
|
||||
},
|
||||
},
|
||||
})),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/sso/useEditSSOSettings", () => ({
|
||||
useEditSSOSettings: vi.fn(() => ({
|
||||
mutateAsync: vi.fn(),
|
||||
isPending: false,
|
||||
})),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
||||
default: vi.fn(() => ({
|
||||
accessToken: "test-token",
|
||||
userId: "test-user-id",
|
||||
userRole: "proxy_admin",
|
||||
})),
|
||||
}));
|
||||
|
||||
const createQueryClient = () =>
|
||||
new QueryClient({
|
||||
defaultOptions: {
|
||||
queries: {
|
||||
retry: false,
|
||||
gcTime: 0,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
describe("DeleteSSOSettingsModal", () => {
|
||||
it("should render", () => {
|
||||
const onCancel = vi.fn();
|
||||
const onSuccess = vi.fn();
|
||||
const queryClient = createQueryClient();
|
||||
|
||||
render(
|
||||
<DeleteSSOSettingsModal isVisible={true} onCancel={onCancel} onSuccess={onSuccess} accessToken="test-token" />,
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<DeleteSSOSettingsModal isVisible={true} onCancel={onCancel} onSuccess={onSuccess} />
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("Confirm Clear SSO Settings")).toBeInTheDocument();
|
||||
expect(
|
||||
screen.getByText("Are you sure you want to clear all SSO settings? This action cannot be undone."),
|
||||
screen.getByText(
|
||||
"Are you sure you want to clear all SSO settings? Users will no longer be able to login using SSO after this change.",
|
||||
),
|
||||
).toBeInTheDocument();
|
||||
expect(screen.getByText("Users will no longer be able to login using SSO after this change.")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,79 +1,66 @@
|
|||
import { Modal } from "antd";
|
||||
import { useEditSSOSettings } from "@/app/(dashboard)/hooks/sso/useEditSSOSettings";
|
||||
import { useSSOSettings } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
import React from "react";
|
||||
import DeleteResourceModal from "../../../../common_components/DeleteResourceModal";
|
||||
import NotificationsManager from "../../../../molecules/notifications_manager";
|
||||
import { updateSSOSettings } from "../../../../networking";
|
||||
import { parseErrorMessage } from "../../../../shared/errorUtils";
|
||||
import { detectSSOProvider } from "../utils";
|
||||
|
||||
interface DeleteSSOSettingsModalProps {
|
||||
isVisible: boolean;
|
||||
onCancel: () => void;
|
||||
onSuccess: () => void;
|
||||
accessToken: string | null;
|
||||
}
|
||||
|
||||
const DeleteSSOSettingsModal: React.FC<DeleteSSOSettingsModalProps> = ({
|
||||
isVisible,
|
||||
onCancel,
|
||||
onSuccess,
|
||||
accessToken,
|
||||
}) => {
|
||||
const DeleteSSOSettingsModal: React.FC<DeleteSSOSettingsModalProps> = ({ isVisible, onCancel, onSuccess }) => {
|
||||
const { data: ssoSettings } = useSSOSettings();
|
||||
const { mutateAsync: editSSOSettings, isPending: isEditingSSOSettings } = useEditSSOSettings();
|
||||
|
||||
// Handle clearing SSO settings
|
||||
const handleClearSSO = async () => {
|
||||
if (!accessToken) {
|
||||
NotificationsManager.fromBackend("No access token available");
|
||||
return;
|
||||
}
|
||||
const clearSettings = {
|
||||
google_client_id: null,
|
||||
google_client_secret: null,
|
||||
microsoft_client_id: null,
|
||||
microsoft_client_secret: null,
|
||||
microsoft_tenant: null,
|
||||
generic_client_id: null,
|
||||
generic_client_secret: null,
|
||||
generic_authorization_endpoint: null,
|
||||
generic_token_endpoint: null,
|
||||
generic_userinfo_endpoint: null,
|
||||
proxy_base_url: null,
|
||||
user_email: null,
|
||||
sso_provider: null,
|
||||
role_mappings: null,
|
||||
};
|
||||
|
||||
try {
|
||||
// Clear all SSO settings
|
||||
const clearSettings = {
|
||||
google_client_id: null,
|
||||
google_client_secret: null,
|
||||
microsoft_client_id: null,
|
||||
microsoft_client_secret: null,
|
||||
microsoft_tenant: null,
|
||||
generic_client_id: null,
|
||||
generic_client_secret: null,
|
||||
generic_authorization_endpoint: null,
|
||||
generic_token_endpoint: null,
|
||||
generic_userinfo_endpoint: null,
|
||||
proxy_base_url: null,
|
||||
user_email: null,
|
||||
sso_provider: null,
|
||||
};
|
||||
|
||||
await updateSSOSettings(accessToken, clearSettings);
|
||||
|
||||
NotificationsManager.success("SSO settings cleared successfully");
|
||||
|
||||
// Close modal and trigger success callback
|
||||
onCancel();
|
||||
onSuccess();
|
||||
} catch (error) {
|
||||
console.error("Failed to clear SSO settings:", error);
|
||||
NotificationsManager.fromBackend("Failed to clear SSO settings: " + parseErrorMessage(error));
|
||||
}
|
||||
await editSSOSettings(clearSettings, {
|
||||
onSuccess: () => {
|
||||
NotificationsManager.success("SSO settings cleared successfully");
|
||||
onCancel();
|
||||
onSuccess();
|
||||
},
|
||||
onError: (error) => {
|
||||
NotificationsManager.fromBackend("Failed to clear SSO settings: " + parseErrorMessage(error));
|
||||
},
|
||||
});
|
||||
};
|
||||
|
||||
return (
|
||||
<Modal
|
||||
<DeleteResourceModal
|
||||
isOpen={isVisible}
|
||||
title="Confirm Clear SSO Settings"
|
||||
visible={isVisible}
|
||||
onOk={handleClearSSO}
|
||||
alertMessage="This action cannot be undone."
|
||||
message="Are you sure you want to clear all SSO settings? Users will no longer be able to login using SSO after this change."
|
||||
resourceInformationTitle="SSO Settings"
|
||||
resourceInformation={[
|
||||
{ label: "Provider", value: (ssoSettings?.values && detectSSOProvider(ssoSettings?.values)) || "Generic" },
|
||||
]}
|
||||
onCancel={onCancel}
|
||||
okText="Yes, Clear"
|
||||
cancelText="Cancel"
|
||||
okButtonProps={{
|
||||
danger: true,
|
||||
style: {
|
||||
backgroundColor: "#dc2626",
|
||||
borderColor: "#dc2626",
|
||||
},
|
||||
}}
|
||||
>
|
||||
<p>Are you sure you want to clear all SSO settings? This action cannot be undone.</p>
|
||||
<p>Users will no longer be able to login using SSO after this change.</p>
|
||||
</Modal>
|
||||
onOk={handleClearSSO}
|
||||
confirmLoading={isEditingSSOSettings}
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,92 @@
|
|||
import type { RoleMappings as RoleMappingsType } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
import { screen } from "@testing-library/react";
|
||||
import { describe, expect, it } from "vitest";
|
||||
import { renderWithProviders } from "../../../../../tests/test-utils";
|
||||
import RoleMappings from "./RoleMappings";
|
||||
|
||||
describe("RoleMappings", () => {
|
||||
it("should render successfully", () => {
|
||||
const roleMappings: RoleMappingsType = {
|
||||
provider: "generic",
|
||||
group_claim: "groups",
|
||||
default_role: "internal_user",
|
||||
roles: {
|
||||
proxy_admin: ["admin-group"],
|
||||
proxy_admin_viewer: [],
|
||||
internal_user: ["user-group"],
|
||||
internal_user_viewer: [],
|
||||
},
|
||||
};
|
||||
|
||||
renderWithProviders(<RoleMappings roleMappings={roleMappings} />);
|
||||
|
||||
expect(screen.getByText("Role Mappings")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should return null when roleMappings is undefined", () => {
|
||||
const { container } = renderWithProviders(<RoleMappings roleMappings={undefined} />);
|
||||
|
||||
expect(container.firstChild).toBeNull();
|
||||
});
|
||||
|
||||
it("should display Group Claim and Default Role with correct values and display names", () => {
|
||||
const testCases: Array<{ role: RoleMappingsType["default_role"]; displayName: string; groupClaim: string }> = [
|
||||
{ role: "internal_user_viewer", displayName: "Internal Viewer", groupClaim: "custom-groups-1" },
|
||||
{ role: "internal_user", displayName: "Internal User", groupClaim: "custom-groups-2" },
|
||||
{ role: "proxy_admin_viewer", displayName: "Proxy Admin Viewer", groupClaim: "custom-groups-3" },
|
||||
{ role: "proxy_admin", displayName: "Proxy Admin", groupClaim: "custom-groups-4" },
|
||||
];
|
||||
|
||||
testCases.forEach(({ role, displayName, groupClaim }) => {
|
||||
const roleMappings: RoleMappingsType = {
|
||||
provider: "generic",
|
||||
group_claim: groupClaim,
|
||||
default_role: role,
|
||||
roles: {
|
||||
proxy_admin: [],
|
||||
proxy_admin_viewer: [],
|
||||
internal_user: [],
|
||||
internal_user_viewer: [],
|
||||
},
|
||||
};
|
||||
|
||||
const { unmount } = renderWithProviders(<RoleMappings roleMappings={roleMappings} />);
|
||||
|
||||
expect(screen.getByText("Group Claim")).toBeInTheDocument();
|
||||
expect(screen.getByText(groupClaim)).toBeInTheDocument();
|
||||
expect(screen.getByText("Default Role")).toBeInTheDocument();
|
||||
const displayNameElements = screen.getAllByText(displayName);
|
||||
expect(displayNameElements.length).toBeGreaterThan(0);
|
||||
unmount();
|
||||
});
|
||||
});
|
||||
|
||||
it("should display table with roles, groups as Tags when mapped, and 'No groups mapped' when empty", () => {
|
||||
const roleMappings: RoleMappingsType = {
|
||||
provider: "generic",
|
||||
group_claim: "groups",
|
||||
default_role: "internal_user",
|
||||
roles: {
|
||||
proxy_admin: ["admin-group-1", "admin-group-2", "admin-group-3"],
|
||||
proxy_admin_viewer: ["viewer-group"],
|
||||
internal_user: ["user-group"],
|
||||
internal_user_viewer: [],
|
||||
},
|
||||
};
|
||||
|
||||
renderWithProviders(<RoleMappings roleMappings={roleMappings} />);
|
||||
|
||||
expect(screen.getByText("Role")).toBeInTheDocument();
|
||||
expect(screen.getByText("Mapped Groups")).toBeInTheDocument();
|
||||
expect(screen.getAllByText("Proxy Admin").length).toBeGreaterThan(0);
|
||||
expect(screen.getAllByText("Proxy Admin Viewer").length).toBeGreaterThan(0);
|
||||
expect(screen.getAllByText("Internal User").length).toBeGreaterThan(0);
|
||||
expect(screen.getAllByText("Internal Viewer").length).toBeGreaterThan(0);
|
||||
expect(screen.getByText("admin-group-1")).toBeInTheDocument();
|
||||
expect(screen.getByText("admin-group-2")).toBeInTheDocument();
|
||||
expect(screen.getByText("admin-group-3")).toBeInTheDocument();
|
||||
expect(screen.getByText("viewer-group")).toBeInTheDocument();
|
||||
expect(screen.getByText("user-group")).toBeInTheDocument();
|
||||
expect(screen.getByText("No groups mapped")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,74 @@
|
|||
import type { RoleMappings as RoleMappingsType } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
import { Card, Divider, Table, Tag, Typography } from "antd";
|
||||
import { Users } from "lucide-react";
|
||||
import { defaultRoleDisplayNames } from "./constants";
|
||||
const { Title, Text } = Typography;
|
||||
|
||||
export default function RoleMappings({ roleMappings }: { roleMappings: RoleMappingsType | undefined }) {
|
||||
if (!roleMappings) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const roleMappingsColumns = [
|
||||
{
|
||||
title: "Role",
|
||||
dataIndex: "role",
|
||||
key: "role",
|
||||
render: (text: string) => <Text strong>{defaultRoleDisplayNames[text]}</Text>,
|
||||
},
|
||||
{
|
||||
title: "Mapped Groups",
|
||||
dataIndex: "groups",
|
||||
key: "groups",
|
||||
render: (groups: string[]) => (
|
||||
<>
|
||||
{groups.length > 0 ? (
|
||||
groups.map((group, index) => (
|
||||
<Tag key={index} color="blue">
|
||||
{group}
|
||||
</Tag>
|
||||
))
|
||||
) : (
|
||||
<Text className="text-gray-400 italic">No groups mapped</Text>
|
||||
)}
|
||||
</>
|
||||
),
|
||||
},
|
||||
];
|
||||
return (
|
||||
<Card>
|
||||
<div className="flex items-center gap-3">
|
||||
<Users className="w-6 h-6 text-gray-400 mb-2" />
|
||||
<Title level={3}>Role Mappings</Title>
|
||||
</div>
|
||||
<div className="space-y-8">
|
||||
<div className="grid grid-cols-2 gap-4">
|
||||
<div>
|
||||
<Title level={5}>Group Claim</Title>
|
||||
<div>
|
||||
<Text code>{roleMappings.group_claim}</Text>
|
||||
</div>
|
||||
</div>
|
||||
<div>
|
||||
<Title level={5}>Default Role</Title>
|
||||
<div>
|
||||
<Text strong>{defaultRoleDisplayNames[roleMappings.default_role]}</Text>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<Divider />
|
||||
<Table
|
||||
columns={roleMappingsColumns}
|
||||
dataSource={Object.entries(roleMappings.roles).map(([role, groups]) => ({
|
||||
role,
|
||||
groups,
|
||||
}))}
|
||||
pagination={false}
|
||||
bordered
|
||||
size="small"
|
||||
className="w-full"
|
||||
/>
|
||||
</div>
|
||||
</Card>
|
||||
);
|
||||
}
|
||||
|
|
@ -1,23 +1,23 @@
|
|||
"use client";
|
||||
|
||||
import { useSSOSettings, type SSOSettingsValues } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { Button, Card, Descriptions, Space, Typography } from "antd";
|
||||
import { Edit, Shield, Trash2 } from "lucide-react";
|
||||
import { useState } from "react";
|
||||
import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./constants";
|
||||
import AddSSOSettingsModal from "./Modals/AddSSOSettingsModal";
|
||||
import DeleteSSOSettingsModal from "./Modals/DeleteSSOSettingsModal";
|
||||
import EditSSOSettingsModal from "./Modals/EditSSOSettingsModal";
|
||||
import RedactableField from "./RedactableField";
|
||||
import RoleMappings from "./RoleMappings";
|
||||
import SSOSettingsEmptyPlaceholder from "./SSOSettingsEmptyPlaceholder";
|
||||
import SSOSettingsLoadingSkeleton from "./SSOSettingsLoadingSkeleton";
|
||||
import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./constants";
|
||||
import { detectSSOProvider } from "./utils";
|
||||
|
||||
const { Title, Text } = Typography;
|
||||
|
||||
export default function SSOSettings() {
|
||||
const { data: ssoSettings, refetch, isLoading } = useSSOSettings();
|
||||
const { accessToken } = useAuthorized();
|
||||
const [isDeleteModalVisible, setIsDeleteModalVisible] = useState(false);
|
||||
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
|
||||
const [isEditModalVisible, setIsEditModalVisible] = useState(false);
|
||||
|
|
@ -26,24 +26,8 @@ export default function SSOSettings() {
|
|||
Boolean(ssoSettings?.values.microsoft_client_id) ||
|
||||
Boolean(ssoSettings?.values.generic_client_id);
|
||||
|
||||
// Determine the SSO provider based on the configuration
|
||||
const detectSSOProvider = (values: SSOSettingsValues): string | null => {
|
||||
if (values.google_client_id) return "google";
|
||||
if (values.microsoft_client_id) return "microsoft";
|
||||
if (values.generic_client_id) {
|
||||
// Check if it looks like Okta/Auth0 based on endpoints
|
||||
if (
|
||||
values.generic_authorization_endpoint?.includes("okta") ||
|
||||
values.generic_authorization_endpoint?.includes("auth0")
|
||||
) {
|
||||
return "okta";
|
||||
}
|
||||
return "generic";
|
||||
}
|
||||
return null;
|
||||
};
|
||||
|
||||
const selectedProvider = ssoSettings?.values ? detectSSOProvider(ssoSettings.values) : null;
|
||||
const isRoleMappingsEnabled = Boolean(ssoSettings?.values.role_mappings);
|
||||
|
||||
const renderEndpointValue = (value?: string | null) => (
|
||||
<Text className="font-mono text-gray-600 text-sm" copyable={!!value}>
|
||||
|
|
@ -185,46 +169,52 @@ export default function SSOSettings() {
|
|||
{isLoading ? (
|
||||
<SSOSettingsLoadingSkeleton />
|
||||
) : (
|
||||
<Card>
|
||||
<Space direction="vertical" size="large" className="w-full">
|
||||
{/* Header Section */}
|
||||
<div className="flex items-center justify-between">
|
||||
<div className="flex items-center gap-3">
|
||||
<Shield className="w-6 h-6 text-gray-400" />
|
||||
<div>
|
||||
<Title level={3}>SSO Configuration</Title>
|
||||
<Text type="secondary">Manage Single Sign-On authentication settings</Text>
|
||||
<Space direction="vertical" size="large" className="w-full">
|
||||
<Card>
|
||||
<Space direction="vertical" size="large" className="w-full">
|
||||
{/* Header Section */}
|
||||
<div className="flex items-center justify-between">
|
||||
<div className="flex items-center gap-3">
|
||||
<Shield className="w-6 h-6 text-gray-400" />
|
||||
<div>
|
||||
<Title level={3}>SSO Configuration</Title>
|
||||
<Text type="secondary">Manage Single Sign-On authentication settings</Text>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="flex items-center gap-3">
|
||||
{isSSOConfigured && (
|
||||
<>
|
||||
<Button icon={<Edit className="w-4 h-4" />} onClick={() => setIsEditModalVisible(true)}>
|
||||
Edit SSO Settings
|
||||
</Button>
|
||||
<Button
|
||||
danger
|
||||
icon={<Trash2 className="w-4 h-4" />}
|
||||
onClick={() => setIsDeleteModalVisible(true)}
|
||||
>
|
||||
Delete SSO Settings
|
||||
</Button>
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="flex items-center gap-3">
|
||||
{isSSOConfigured && (
|
||||
<>
|
||||
<Button icon={<Edit className="w-4 h-4" />} onClick={() => setIsEditModalVisible(true)}>
|
||||
Edit SSO Settings
|
||||
</Button>
|
||||
<Button danger icon={<Trash2 className="w-4 h-4" />} onClick={() => setIsDeleteModalVisible(true)}>
|
||||
Delete SSO Settings
|
||||
</Button>
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{isSSOConfigured ? (
|
||||
renderSSOSettings()
|
||||
) : (
|
||||
<SSOSettingsEmptyPlaceholder onAdd={() => setIsAddModalVisible(true)} />
|
||||
)}
|
||||
</Space>
|
||||
</Card>
|
||||
{isSSOConfigured ? (
|
||||
renderSSOSettings()
|
||||
) : (
|
||||
<SSOSettingsEmptyPlaceholder onAdd={() => setIsAddModalVisible(true)} />
|
||||
)}
|
||||
</Space>
|
||||
</Card>
|
||||
{isRoleMappingsEnabled && <RoleMappings roleMappings={ssoSettings?.values.role_mappings} />}
|
||||
</Space>
|
||||
)}
|
||||
|
||||
<DeleteSSOSettingsModal
|
||||
isVisible={isDeleteModalVisible}
|
||||
onCancel={() => setIsDeleteModalVisible(false)}
|
||||
onSuccess={() => refetch()}
|
||||
accessToken={accessToken}
|
||||
/>
|
||||
|
||||
<AddSSOSettingsModal
|
||||
|
|
|
|||
|
|
@ -13,3 +13,10 @@ export const ssoProviderDisplayNames: Record<string, string> = {
|
|||
okta: "Okta / Auth0 SSO",
|
||||
generic: "Generic SSO",
|
||||
};
|
||||
|
||||
export const defaultRoleDisplayNames: Record<string, string> = {
|
||||
internal_user_viewer: "Internal Viewer",
|
||||
internal_user: "Internal User",
|
||||
proxy_admin_viewer: "Proxy Admin Viewer",
|
||||
proxy_admin: "Proxy Admin",
|
||||
};
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
other_field: "value",
|
||||
};
|
||||
|
||||
|
|
@ -83,6 +84,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -100,6 +102,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -121,6 +124,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -142,6 +146,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -160,6 +165,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -174,6 +180,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -186,6 +193,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "internal_user",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -198,6 +206,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin_viewer",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -210,6 +219,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "proxy_admin",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -222,6 +232,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
default_role: "unknown_role",
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
@ -233,6 +244,7 @@ describe("processSSOSettingsPayload", () => {
|
|||
const formValues = {
|
||||
group_claim: "groups",
|
||||
use_role_mappings: true,
|
||||
sso_provider: "generic",
|
||||
};
|
||||
|
||||
const result = processSSOSettingsPayload(formValues);
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
import { SSOSettingsValues } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
|
||||
/**
|
||||
* Processes SSO settings form values and transforms them into the payload format expected by the API
|
||||
* Handles role mappings transformation and field extraction
|
||||
|
|
@ -18,7 +20,7 @@ export const processSSOSettingsPayload = (formValues: Record<string, any>): Reco
|
|||
...rest,
|
||||
};
|
||||
|
||||
// Add role mappings if use_role_mappings is checked
|
||||
// Add role mappings only if use_role_mappings is checked AND provider supports role mappings
|
||||
if (use_role_mappings) {
|
||||
// Helper function to split comma-separated string into array
|
||||
const splitTeams = (teams: string | undefined): string[] => {
|
||||
|
|
@ -52,3 +54,20 @@ export const processSSOSettingsPayload = (formValues: Record<string, any>): Reco
|
|||
|
||||
return payload;
|
||||
};
|
||||
|
||||
// Determine the SSO provider based on the configuration
|
||||
export const detectSSOProvider = (values: SSOSettingsValues): string | null => {
|
||||
if (values.google_client_id) return "google";
|
||||
if (values.microsoft_client_id) return "microsoft";
|
||||
if (values.generic_client_id) {
|
||||
// Check if it looks like Okta/Auth0 based on endpoints
|
||||
if (
|
||||
values.generic_authorization_endpoint?.includes("okta") ||
|
||||
values.generic_authorization_endpoint?.includes("auth0")
|
||||
) {
|
||||
return "okta";
|
||||
}
|
||||
return "generic";
|
||||
}
|
||||
return null;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -136,7 +136,7 @@ export const useMcpOAuthFlow = ({
|
|||
if (!hasPreconfiguredCredentials) {
|
||||
const registration = await registerMcpOAuthClient(accessToken, serverId, {
|
||||
client_name: temporaryPayload.alias || temporaryPayload.server_name || serverId,
|
||||
grant_types: ["authorization_code"],
|
||||
grant_types: ["authorization_code", "refresh_token"],
|
||||
response_types: ["code"],
|
||||
token_endpoint_auth_method:
|
||||
temporaryPayload.credentials && temporaryPayload.credentials.client_secret ? "client_secret_post" : "none",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue