mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
chore: resolve merge conflict with main
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
commit
063238a789
183 changed files with 16247 additions and 2032 deletions
|
|
@ -21,9 +21,7 @@ commands:
|
|||
- run:
|
||||
name: "Install local version of litellm-enterprise"
|
||||
command: |
|
||||
cd enterprise
|
||||
python -m pip install -e .
|
||||
cd ..
|
||||
pip install --force-reinstall --no-deps -e enterprise/
|
||||
setup_litellm_test_deps:
|
||||
steps:
|
||||
- checkout
|
||||
|
|
|
|||
36
.claude/settings.json
Normal file
36
.claude/settings.json
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
{
|
||||
"permissions": {
|
||||
"allow": [
|
||||
"Bash(git show:*)",
|
||||
"Bash(git worktree add:*)",
|
||||
"Read(//Users/krrishdholakia/Documents/litellm/**)",
|
||||
"Read(//Users/krrishdholakia/Documents/litellm-claude-code-guardrails/litellm/types/**)",
|
||||
"Read(//Users/krrishdholakia/Documents/litellm-claude-code-guardrails/**)",
|
||||
"Read(//Users/krrishdholakia/Documents/litellm-claude-code-guardrails/litellm/**)",
|
||||
"Bash(python:*)",
|
||||
"Bash(python -c \"\nimport sys; sys.path.insert\\(0, ''.''\\)\nfrom litellm.proxy.guardrails.guardrail_hooks.claude_code.guardrail import ClaudeCodeGuardrail, HOSTED_TOOL_PREFIXES\nprint\\(''HOSTED_TOOL_PREFIXES:'', HOSTED_TOOL_PREFIXES\\)\nprint\\(''ClaudeCodeGuardrail imported OK''\\)\n\")",
|
||||
"Read(//Users/krrishdholakia/Documents/litellm-mcp-jwt-groups/litellm/proxy/**)",
|
||||
"Read(//Users/krrishdholakia/Documents/litellm-mcp-jwt-groups/**)",
|
||||
"Bash(poetry run pytest:*)",
|
||||
"Bash(git add:*)",
|
||||
"Bash(git commit:*)",
|
||||
"Bash(poetry run python:*)",
|
||||
"Bash(poetry run pip:*)",
|
||||
"Bash(git reset:*)",
|
||||
"Bash(git cherry-pick:*)",
|
||||
"Bash(git checkout:*)",
|
||||
"Read(//Users/krrishdholakia/Documents/litellm/litellm/proxy/guardrails/guardrail_hooks/**)",
|
||||
"Read(//Users/krrishdholakia/Documents/**)",
|
||||
"Bash(git -C /Users/krrishdholakia/Documents/litellm-mcp-user-permissions worktree list)",
|
||||
"Bash(ls:*)"
|
||||
],
|
||||
"additionalDirectories": [
|
||||
"/Users/krrishdholakia/Documents/litellm-mcp-group-plan/plan",
|
||||
"/Users/krrishdholakia/Documents/litellm-claude-code-guardrails/litellm/proxy/guardrails/guardrail_hooks/claude_code",
|
||||
"/Users/krrishdholakia/Documents/litellm-claude-code-guardrails/litellm/types",
|
||||
"/Users/krrishdholakia/Documents/litellm-claude-code-guardrails",
|
||||
"/Users/krrishdholakia/Documents/litellm-mcp-jwt-groups/litellm/proxy",
|
||||
"/Users/krrishdholakia/Documents/litellm-mcp-jwt-groups/tests/test_litellm/proxy/auth"
|
||||
]
|
||||
}
|
||||
}
|
||||
6
.github/workflows/test-litellm-matrix.yml
vendored
6
.github/workflows/test-litellm-matrix.yml
vendored
|
|
@ -100,7 +100,11 @@ jobs:
|
|||
|
||||
- name: Setup litellm-enterprise
|
||||
run: |
|
||||
cd enterprise && poetry run pip install -e . && cd ..
|
||||
poetry run pip install --force-reinstall --no-deps -e enterprise/
|
||||
|
||||
- name: Generate Prisma client
|
||||
run: |
|
||||
poetry run prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Run tests - ${{ matrix.test-group.name }}
|
||||
run: |
|
||||
|
|
|
|||
4
.github/workflows/test-litellm.yml
vendored
4
.github/workflows/test-litellm.yml
vendored
|
|
@ -42,9 +42,7 @@ jobs:
|
|||
poetry run pip install "openapi-core"
|
||||
- name: Setup litellm-enterprise as local package
|
||||
run: |
|
||||
cd enterprise
|
||||
poetry run pip install -e .
|
||||
cd ..
|
||||
poetry run pip install --force-reinstall --no-deps -e enterprise/
|
||||
- name: Run tests
|
||||
run: |
|
||||
poetry run pytest tests/test_litellm --tb=short -vv --maxfail=10 -n 4 --durations=50
|
||||
|
|
|
|||
4
.github/workflows/test-mcp.yml
vendored
4
.github/workflows/test-mcp.yml
vendored
|
|
@ -40,9 +40,7 @@ jobs:
|
|||
|
||||
- name: Setup litellm-enterprise as local package
|
||||
run: |
|
||||
cd enterprise
|
||||
python -m pip install -e .
|
||||
cd ..
|
||||
poetry run pip install --force-reinstall --no-deps -e enterprise/
|
||||
|
||||
- name: Run MCP tests
|
||||
run: |
|
||||
|
|
|
|||
2
.github/workflows/test_server_root_path.yml
vendored
2
.github/workflows/test_server_root_path.yml
vendored
|
|
@ -26,7 +26,7 @@ jobs:
|
|||
uses: docker/build-push-action@v5
|
||||
with:
|
||||
context: .
|
||||
file: ./docker/Dockerfile.database
|
||||
file: ./docker/Dockerfile.non_root
|
||||
tags: litellm-test:${{ github.sha }}
|
||||
load: true
|
||||
cache-from: type=gha
|
||||
|
|
|
|||
|
|
@ -24,6 +24,8 @@ hide_table_of_contents: false
|
|||
**Severity:** High
|
||||
**Status:** Resolved
|
||||
|
||||
> **Note:** This fix will be available starting from `v1.81.13-nightly` or higher of LiteLLM.
|
||||
|
||||
## Summary
|
||||
|
||||
Claude Code began sending unsupported Anthropic beta headers to non-Anthropic providers (Bedrock, Azure AI, Vertex AI), causing `invalid beta flag` errors. LiteLLM was forwarding all beta headers without provider-specific validation. Users experienced request failures when routing Claude Code requests through LiteLLM to these providers.
|
||||
|
|
|
|||
117
docs/my-website/blog/vllm_embeddings_incident/index.md
Normal file
117
docs/my-website/blog/vllm_embeddings_incident/index.md
Normal file
|
|
@ -0,0 +1,117 @@
|
|||
---
|
||||
slug: vllm-embeddings-incident
|
||||
title: "Incident Report: vLLM Embeddings Broken by encoding_format Parameter"
|
||||
date: 2026-02-18T10:00:00
|
||||
authors:
|
||||
- name: Sameer Kankute
|
||||
title: SWE @ LiteLLM (LLM Translation)
|
||||
url: https://www.linkedin.com/in/sameer-kankute/
|
||||
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
|
||||
- name: Krrish Dholakia
|
||||
title: "CEO, LiteLLM"
|
||||
url: https://www.linkedin.com/in/krish-d/
|
||||
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
|
||||
- name: Ishaan Jaff
|
||||
title: "CTO, LiteLLM"
|
||||
url: https://www.linkedin.com/in/reffajnaahsi/
|
||||
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
|
||||
tags: [incident-report, embeddings, vllm]
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
**Date:** Feb 16, 2026
|
||||
**Duration:** ~3 hours
|
||||
**Severity:** High (for vLLM embedding users)
|
||||
**Status:** Resolved
|
||||
|
||||
## Summary
|
||||
|
||||
A commit ([`dbcae4a`](https://github.com/BerriAI/litellm/commit/dbcae4aca5836770d0e9cd43abab0333c3d61ab2)) intended to fix OpenAI SDK behavior broke vLLM embeddings by explicitly passing `encoding_format=None` in API requests. vLLM rejects this with error: `"unknown variant \`\`, expected float or base64"`.
|
||||
|
||||
- **vLLM embedding calls:** Complete failure - all requests rejected
|
||||
- **Other providers:** No impact - OpenAI and other providers functioned normally
|
||||
- **Other vLLM functionality:** No impact - only embeddings were affected
|
||||
|
||||
{/* truncate */}
|
||||
|
||||
---
|
||||
|
||||
## Background
|
||||
|
||||
The `encoding_format` parameter for embeddings specifies whether vectors should be returned as `float` arrays or `base64` encoded strings. Different providers have different expectations:
|
||||
|
||||
- **OpenAI SDK:** If `encoding_format` is omitted, the SDK adds a default value of `"float"`
|
||||
- **vLLM:** Strictly validates `encoding_format` - only accepts `"float"`, `"base64"`, or complete omission. Rejects `None` or empty string values.
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
A["1. User calls litellm.embedding()
|
||||
litellm/main.py"] --> B["2. Transform request for provider
|
||||
litellm/llms/openai_like/embedding/handler.py"]
|
||||
B --> C["3. Send request to vLLM endpoint"]
|
||||
C -->|"encoding_format omitted"| D["4a. ✅ vLLM processes request"]
|
||||
C -->|"encoding_format='float' or 'base64'"| D
|
||||
C -->|"encoding_format=None or ''"| E["4b. ❌ vLLM rejects with error:
|
||||
'unknown variant, expected float or base64'"]
|
||||
|
||||
style D fill:#d4edda,stroke:#28a745
|
||||
style E fill:#f8d7da,stroke:#dc3545
|
||||
style B fill:#fff3cd,stroke:#ffc107
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Root cause
|
||||
|
||||
A well-intentioned fix for OpenAI SDK behavior inadvertently broke vLLM embeddings:
|
||||
|
||||
**The Breaking Change ([`dbcae4a`](https://github.com/BerriAI/litellm/commit/dbcae4aca5836770d0e9cd43abab0333c3d61ab2)):**
|
||||
|
||||
In `litellm/main.py`, the code was changed to explicitly set `encoding_format=None` instead of omitting it:
|
||||
|
||||
```python
|
||||
# Added in dbcae4a
|
||||
if encoding_format is not None:
|
||||
optional_params["encoding_format"] = encoding_format
|
||||
else:
|
||||
# Omitting causes openai sdk to add default value of "float"
|
||||
optional_params["encoding_format"] = None
|
||||
```
|
||||
|
||||
This fix worked correctly for OpenAI - explicitly passing `None` prevented the SDK from adding its default value. However, vLLM's strict parameter validation rejected `None` values, causing all embedding requests to fail.
|
||||
|
||||
---
|
||||
|
||||
## The Fix
|
||||
|
||||
Fix deployed ([`55348dd`](https://github.com/BerriAI/litellm/commit/55348dd9c51b5b028f676d25ad023b8f052fc071)). The solution filters out `None` and empty string values from `optional_params` before sending requests to OpenAI-like providers (including vLLM).
|
||||
|
||||
**In `litellm/llms/openai_like/embedding/handler.py`:**
|
||||
|
||||
```python
|
||||
# Before (broken)
|
||||
data = {"model": model, "input": input, **optional_params}
|
||||
|
||||
# After (fixed)
|
||||
filtered_optional_params = {k: v for k, v in optional_params.items() if v not in (None, '')}
|
||||
data = {"model": model, "input": input, **filtered_optional_params}
|
||||
```
|
||||
|
||||
This ensures:
|
||||
- Valid values (`"float"`, `"base64"`) are preserved and sent
|
||||
- `None` and empty string values are filtered out (parameter omitted entirely)
|
||||
- OpenAI SDK no longer adds defaults because liteLLM handles the parameter upstream
|
||||
|
||||
---
|
||||
|
||||
## Remediation
|
||||
|
||||
| # | Action | Status | Code |
|
||||
|---|---|---|---|
|
||||
| 1 | Filter `None` and empty string values in OpenAI-like embedding handler | ✅ Done | [`handler.py#L108`](https://github.com/BerriAI/litellm/blob/main/litellm/llms/openai_like/embedding/handler.py#L108) |
|
||||
| 2 | Unit tests for parameter filtering (None, empty string, valid values) | ✅ Done | [`test_openai_like_embedding.py`](https://github.com/BerriAI/litellm/blob/main/tests/test_litellm/llms/openai_like/embedding/test_openai_like_embedding.py) |
|
||||
| 3 | Transformation tests for hosted_vllm embedding config | ✅ Done | [`test_hosted_vllm_embedding_transformation.py`](https://github.com/BerriAI/litellm/blob/main/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py) |
|
||||
| 4 | E2E tests with actual vLLM endpoint | ✅ Done | [`test_hosted_vllm_embedding_e2e.py`](https://github.com/BerriAI/litellm/blob/main/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_e2e.py) |
|
||||
| 5 | Validate JSON payload structure matches vLLM expectations | ✅ Done | Tests verify exact JSON sent to endpoint |
|
||||
|
||||
---
|
||||
|
|
@ -237,6 +237,7 @@ litellm_settings:
|
|||
mode: pre_call # or post_call, during_call
|
||||
api_base: https://your-guardrail-api.com
|
||||
api_key: os.environ/YOUR_GUARDRAIL_API_KEY # optional
|
||||
unreachable_fallback: fail_closed # default: fail_closed. Set to fail_open to proceed if the guardrail endpoint is unreachable (network errors, or HTTP 502/503/504 from an upstream proxy/LB).
|
||||
additional_provider_specific_params:
|
||||
# your custom parameters
|
||||
threshold: 0.8
|
||||
|
|
|
|||
465
docs/my-website/docs/completion/message_sanitization.md
Normal file
465
docs/my-website/docs/completion/message_sanitization.md
Normal file
|
|
@ -0,0 +1,465 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Message Sanitization for Tool Calling for anthropic models
|
||||
|
||||
**Automatically fix common message formatting issues when using tool calling with `modify_params=True`**
|
||||
|
||||
LiteLLM can automatically sanitize messages to handle common issues that occur during tool calling workflows, especially when using OpenAI-compatible clients with providers that have strict message format requirements (like Anthropic Claude).
|
||||
|
||||
## Overview
|
||||
|
||||
When `litellm.modify_params = True` is enabled, LiteLLM automatically sanitizes messages to fix three common issues:
|
||||
|
||||
1. **Orphaned Tool Calls** - Assistant messages with tool_calls but missing tool results
|
||||
2. **Orphaned Tool Results** - Tool messages that reference non-existent tool_call_ids
|
||||
3. **Empty Message Content** - Messages with empty or whitespace-only text content
|
||||
|
||||
This ensures your tool calling workflows work seamlessly across different LLM providers without manual message validation.
|
||||
|
||||
## Why Message Sanitization?
|
||||
|
||||
Different LLM providers have varying requirements for message formats, especially during tool calling:
|
||||
|
||||
- **Anthropic Claude** requires every tool_call to have a corresponding tool result
|
||||
- Some providers reject messages with empty content
|
||||
- OpenAI-compatible clients may not always maintain perfect message consistency
|
||||
|
||||
Without sanitization, these issues cause API errors that interrupt your workflows. With `modify_params=True`, LiteLLM handles these edge cases automatically.
|
||||
|
||||
## Quick Start
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# Enable automatic message sanitization
|
||||
litellm.modify_params = True
|
||||
|
||||
# This will work even if messages have formatting issues
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=[
|
||||
{"role": "user", "content": "What's the weather in Boston?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_123",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"city": "Boston"}'}
|
||||
}
|
||||
]
|
||||
# Missing tool result - LiteLLM will add a dummy result automatically
|
||||
},
|
||||
{"role": "user", "content": "Thanks!"}
|
||||
],
|
||||
tools=[{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a city",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"]
|
||||
}
|
||||
}
|
||||
}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
modify_params: true # Enable automatic message sanitization
|
||||
|
||||
model_list:
|
||||
- model_name: claude-3-5-sonnet
|
||||
litellm_params:
|
||||
model: anthropic/claude-3-5-sonnet-20241022
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Sanitization Cases
|
||||
|
||||
### Case A: Orphaned Tool Calls (Missing Tool Results)
|
||||
|
||||
**Problem:** An assistant message contains `tool_calls`, but no corresponding tool result messages follow.
|
||||
|
||||
**Solution:** LiteLLM automatically adds dummy tool result messages for any missing tool results.
|
||||
|
||||
**Example:**
|
||||
|
||||
```python
|
||||
import litellm
|
||||
litellm.modify_params = True
|
||||
|
||||
# Messages with orphaned tool calls
|
||||
messages = [
|
||||
{"role": "user", "content": "Search for Python tutorials"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {"name": "web_search", "arguments": '{"query": "Python tutorials"}'}
|
||||
}
|
||||
]
|
||||
},
|
||||
# Missing tool result here!
|
||||
{"role": "user", "content": "What about JavaScript?"}
|
||||
]
|
||||
|
||||
# LiteLLM automatically adds:
|
||||
# {
|
||||
# "role": "tool",
|
||||
# "tool_call_id": "call_abc123",
|
||||
# "content": "[System: Tool execution skipped/interrupted by user. No result provided for tool 'web_search'.]"
|
||||
# }
|
||||
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=messages,
|
||||
tools=[...]
|
||||
)
|
||||
```
|
||||
|
||||
**When this happens:**
|
||||
- User interrupts tool execution
|
||||
- Client loses tool results due to network issues
|
||||
- Conversation flow changes before tool completes
|
||||
- Multi-turn conversations where tools are optional
|
||||
|
||||
### Case B: Orphaned Tool Results (Invalid tool_call_id)
|
||||
|
||||
**Problem:** A tool message references a `tool_call_id` that doesn't exist in any previous assistant message.
|
||||
|
||||
**Solution:** LiteLLM automatically removes these orphaned tool result messages.
|
||||
|
||||
**Example:**
|
||||
|
||||
```python
|
||||
import litellm
|
||||
litellm.modify_params = True
|
||||
|
||||
# Messages with orphaned tool result
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi! How can I help?"},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_nonexistent", # This tool_call_id doesn't exist!
|
||||
"content": "Some result"
|
||||
}
|
||||
]
|
||||
|
||||
# LiteLLM automatically removes the orphaned tool message
|
||||
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=messages
|
||||
)
|
||||
```
|
||||
|
||||
**When this happens:**
|
||||
- Message history is manually edited
|
||||
- Tool results are duplicated or mismatched
|
||||
- Conversation state is restored incorrectly
|
||||
- Messages are merged from different conversations
|
||||
|
||||
### Case C: Empty Message Content
|
||||
|
||||
**Problem:** User or assistant messages have empty or whitespace-only content.
|
||||
|
||||
**Solution:** LiteLLM replaces empty content with a system placeholder message.
|
||||
|
||||
**Example:**
|
||||
|
||||
```python
|
||||
import litellm
|
||||
litellm.modify_params = True
|
||||
|
||||
# Messages with empty content
|
||||
messages = [
|
||||
{"role": "user", "content": ""}, # Empty content
|
||||
{"role": "assistant", "content": " "}, # Whitespace only
|
||||
]
|
||||
|
||||
# LiteLLM automatically replaces with:
|
||||
# {"role": "user", "content": "[System: Empty message content sanitised to satisfy protocol]"}
|
||||
# {"role": "assistant", "content": "[System: Empty message content sanitised to satisfy protocol]"}
|
||||
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=messages
|
||||
)
|
||||
```
|
||||
|
||||
**When this happens:**
|
||||
- UI sends empty messages
|
||||
- Content is stripped during preprocessing
|
||||
- Placeholder messages in conversation history
|
||||
- Edge cases in message construction
|
||||
|
||||
## Configuration
|
||||
|
||||
### Enable Globally
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# Enable for all completion calls
|
||||
litellm.modify_params = True
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
modify_params: true
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="env" label="Environment Variable">
|
||||
|
||||
```bash
|
||||
export LITELLM_MODIFY_PARAMS=True
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Enable Per-Request
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# Enable only for specific requests
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=messages,
|
||||
modify_params=True # Override global setting
|
||||
)
|
||||
```
|
||||
|
||||
## Supported Providers
|
||||
|
||||
Message sanitization currently works with:
|
||||
|
||||
- ✅ Anthropic (Claude)
|
||||
|
||||
**Note:** While the sanitization logic is provider-agnostic, it is currently only applied in the Anthropic message transformation pipeline. Support for additional providers may be added in future releases.
|
||||
|
||||
## Implementation Details
|
||||
|
||||
### How It Works
|
||||
|
||||
The message sanitization process runs **before** messages are converted to provider-specific formats:
|
||||
|
||||
1. **Input:** OpenAI-format messages with potential issues
|
||||
2. **Sanitization:** Three helper functions process the messages:
|
||||
- `_sanitize_empty_text_content()` - Fixes empty content
|
||||
- `_add_missing_tool_results()` - Adds dummy tool results
|
||||
- `_is_orphaned_tool_result()` - Identifies orphaned results
|
||||
3. **Output:** Clean, provider-compatible messages
|
||||
|
||||
### Code Reference
|
||||
|
||||
The sanitization logic is implemented in:
|
||||
- `litellm/litellm_core_utils/prompt_templates/factory.py`
|
||||
- Function: `sanitize_messages_for_tool_calling()`
|
||||
|
||||
### Logging
|
||||
|
||||
When sanitization occurs, LiteLLM logs debug messages:
|
||||
|
||||
```python
|
||||
import litellm
|
||||
litellm.set_verbose = True # Enable debug logging
|
||||
|
||||
# You'll see logs like:
|
||||
# "_add_missing_tool_results: Found 1 orphaned tool calls. Adding dummy tool results."
|
||||
# "_is_orphaned_tool_result: Found orphaned tool result with tool_call_id=call_123"
|
||||
# "_sanitize_empty_text_content: Replaced empty text content in user message"
|
||||
```
|
||||
|
||||
## Best Practices
|
||||
|
||||
### 1. Enable for Production Workflows
|
||||
|
||||
```python
|
||||
# Recommended for production
|
||||
litellm.modify_params = True
|
||||
|
||||
# Ensures robust handling of edge cases
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=messages,
|
||||
tools=tools
|
||||
)
|
||||
```
|
||||
|
||||
### 2. Preserve Tool Results When Possible
|
||||
|
||||
While sanitization handles missing tool results, it's better to provide actual results:
|
||||
|
||||
```python
|
||||
# Good: Provide actual tool results
|
||||
messages = [
|
||||
{"role": "user", "content": "Search for Python"},
|
||||
{"role": "assistant", "tool_calls": [...]},
|
||||
{"role": "tool", "tool_call_id": "call_123", "content": "Actual search results"}
|
||||
]
|
||||
|
||||
# Fallback: Sanitization adds dummy result if missing
|
||||
messages = [
|
||||
{"role": "user", "content": "Search for Python"},
|
||||
{"role": "assistant", "tool_calls": [...]},
|
||||
# Missing tool result - sanitization adds dummy
|
||||
]
|
||||
```
|
||||
|
||||
### 3. Monitor Sanitization Events
|
||||
|
||||
Use logging to track when sanitization occurs:
|
||||
|
||||
```python
|
||||
import litellm
|
||||
import logging
|
||||
|
||||
# Enable debug logging
|
||||
litellm.set_verbose = True
|
||||
logging.basicConfig(level=logging.DEBUG)
|
||||
|
||||
# Track sanitization events in your application
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=messages
|
||||
)
|
||||
```
|
||||
|
||||
### 4. Test Edge Cases
|
||||
|
||||
Ensure your application handles sanitized messages correctly:
|
||||
|
||||
```python
|
||||
import litellm
|
||||
litellm.modify_params = True
|
||||
|
||||
# Test orphaned tool calls
|
||||
test_messages = [
|
||||
{"role": "user", "content": "Test"},
|
||||
{"role": "assistant", "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "test", "arguments": "{}"}}]},
|
||||
{"role": "user", "content": "Continue"} # No tool result
|
||||
]
|
||||
|
||||
response = litellm.completion(
|
||||
model="anthropic/claude-3-5-sonnet-20241022",
|
||||
messages=test_messages,
|
||||
tools=[...]
|
||||
)
|
||||
|
||||
# Verify the response handles the dummy tool result appropriately
|
||||
```
|
||||
|
||||
## Related Features
|
||||
|
||||
- **[Drop Params](./drop_params.md)** - Drop unsupported parameters for specific providers
|
||||
- **[Message Trimming](./message_trimming.md)** - Trim messages to fit token limits
|
||||
- **[Function Calling](./function_call.md)** - Complete guide to tool/function calling
|
||||
- **[Reasoning Content](../reasoning_content.md)** - Extended thinking with tool calling
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Sanitization Not Working
|
||||
|
||||
**Issue:** Messages still cause errors despite `modify_params=True`
|
||||
|
||||
**Solution:**
|
||||
1. Verify `modify_params` is enabled:
|
||||
```python
|
||||
import litellm
|
||||
print(litellm.modify_params) # Should be True
|
||||
```
|
||||
|
||||
2. Check if the issue is provider-specific:
|
||||
```python
|
||||
litellm.set_verbose = True # Enable debug logging
|
||||
```
|
||||
|
||||
3. Ensure you're using a recent version of LiteLLM:
|
||||
```bash
|
||||
pip install --upgrade litellm
|
||||
```
|
||||
|
||||
### Unexpected Dummy Tool Results
|
||||
|
||||
**Issue:** Dummy tool results appear when you expect actual results
|
||||
|
||||
**Cause:** Tool result messages are missing or have incorrect `tool_call_id`
|
||||
|
||||
**Solution:**
|
||||
1. Verify tool result messages have correct `tool_call_id`:
|
||||
```python
|
||||
# Correct
|
||||
{"role": "tool", "tool_call_id": "call_123", "content": "result"}
|
||||
|
||||
# Incorrect - will be treated as orphaned
|
||||
{"role": "tool", "tool_call_id": "wrong_id", "content": "result"}
|
||||
```
|
||||
|
||||
2. Ensure tool results immediately follow assistant messages with tool_calls
|
||||
|
||||
### Performance Impact
|
||||
|
||||
**Issue:** Concerned about performance overhead
|
||||
|
||||
**Details:** Message sanitization has minimal performance impact:
|
||||
- Runs in O(n) time where n = number of messages
|
||||
- Only processes messages when `modify_params=True`
|
||||
- Typically adds < 1ms to request processing time
|
||||
|
||||
## FAQ
|
||||
|
||||
**Q: Does sanitization modify my original messages?**
|
||||
|
||||
A: No, sanitization creates a new list of messages. Your original messages remain unchanged.
|
||||
|
||||
**Q: Can I disable specific sanitization cases?**
|
||||
|
||||
A: Currently, all three cases are handled together when `modify_params=True`. To disable sanitization entirely, set `modify_params=False`.
|
||||
|
||||
**Q: What happens to the dummy tool results?**
|
||||
|
||||
A: Dummy tool results are sent to the LLM provider along with other messages. The model sees them as regular tool results with informative error messages.
|
||||
|
||||
**Q: Does this work with streaming?**
|
||||
|
||||
A: Yes, message sanitization works with both streaming and non-streaming requests.
|
||||
|
||||
**Q: Is this related to `drop_params`?**
|
||||
|
||||
A: No, they're separate features:
|
||||
- `modify_params` - Modifies/fixes message content and structure
|
||||
- `drop_params` - Removes unsupported API parameters
|
||||
|
||||
Both can be enabled simultaneously.
|
||||
|
||||
## See Also
|
||||
|
||||
- [Reasoning Content with Tool Calling](../reasoning_content.md)
|
||||
- [Function Calling Guide](./function_call.md)
|
||||
- [Bedrock Provider Documentation](../providers/bedrock.md)
|
||||
- [Anthropic Provider Documentation](../providers/anthropic.md)
|
||||
|
|
@ -808,6 +808,68 @@ If your stdio MCP server needs per-request credentials, you can map HTTP headers
|
|||
|
||||
In this example, when a client makes a request with the `X-GITHUB_PERSONAL_ACCESS_TOKEN` header, the proxy forwards that value into the stdio process as the `GITHUB_PERSONAL_ACCESS_TOKEN` environment variable.
|
||||
|
||||
## Control MCP Access for End Users
|
||||
|
||||
Control which MCP servers end users of your AI application can access (e.g. users of an internal chat UI). Pass the customer ID in the `x-litellm-end-user-id` header to:
|
||||
- Enforce object permissions (limit which MCP servers they can access)
|
||||
- Apply customer-specific budgets
|
||||
- Track spend per customer
|
||||
|
||||
**FastMCP Client Example:**
|
||||
|
||||
```python title="Track customer spend with x-litellm-end-user-id" showLineNumbers
|
||||
from fastmcp import Client
|
||||
import asyncio
|
||||
|
||||
# MCP client configuration with customer tracking
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"github": {
|
||||
"url": "http://localhost:4000/github_mcp/mcp",
|
||||
"headers": {
|
||||
"x-litellm-api-key": "Bearer sk-1234",
|
||||
"x-litellm-end-user-id": "customer_123", # 👈 CUSTOMER ID
|
||||
"Authorization": "Bearer gho_token"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
client = Client(config)
|
||||
|
||||
async def main():
|
||||
async with client:
|
||||
# All MCP calls will be tracked under customer_123
|
||||
tools = await client.list_tools()
|
||||
result = await client.call_tool(tools[0].name, {})
|
||||
print(f"Tool result: {result}")
|
||||
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
**Cursor IDE Example:**
|
||||
|
||||
```json title="Cursor config with customer tracking" showLineNumbers
|
||||
{
|
||||
"mcpServers": {
|
||||
"GitHub": {
|
||||
"url": "http://localhost:4000/github_mcp/mcp",
|
||||
"headers": {
|
||||
"x-litellm-api-key": "Bearer $LITELLM_API_KEY",
|
||||
"x-litellm-end-user-id": "customer_123"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**What happens:**
|
||||
- Customer-specific object permissions are enforced (only allowed MCP servers are accessible)
|
||||
- Customer budgets are applied
|
||||
- All tool calls are tracked under `customer_123`
|
||||
|
||||
[Learn more about customer management →](./proxy/customers)
|
||||
|
||||
## Using your MCP with client side credentials
|
||||
|
||||
Use this if you want to pass a client side authentication token to LiteLLM to then pass to your MCP to auth to your MCP.
|
||||
|
|
|
|||
|
|
@ -253,3 +253,12 @@ LiteLLM supports customizing the following Datadog environment variables
|
|||
\* **Required when using Direct API** (default): `DD_API_KEY` and `DD_SITE` are required
|
||||
\* **Optional when using DataDog Agent**: Set `LITELLM_DD_AGENT_HOST` to use agent mode; `DD_API_KEY` and `DD_SITE` are not required for **Datadog Logs**. (**Note: `DD_API_KEY` IS REQUIRED for Datadog LLM Observability**)
|
||||
|
||||
## Automatic Tags
|
||||
|
||||
LiteLLM automatically adds the following tags to your Datadog logs and metrics if the information is available in the request:
|
||||
|
||||
| Tag | Description | Source |
|
||||
|-----|-------------|--------|
|
||||
| `team` | The team alias or ID associated with the API Key | `user_api_key_team_alias`, `team_alias`, `user_api_key_team_id`, or `team_id` in metadata |
|
||||
| `request_tag` | Custom tags passed in the request | `request_tags` in logging payload |
|
||||
|
||||
|
|
|
|||
52
docs/my-website/docs/providers/watsonx/rerank.md
Normal file
52
docs/my-website/docs/providers/watsonx/rerank.md
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
# watsonx.ai Rerank
|
||||
|
||||
## Overview
|
||||
|
||||
| Property | Details |
|
||||
|----------|--------------------------------------------------------------------------|
|
||||
| Description | watsonx.ai rerank integration |
|
||||
| Provider Route on LiteLLM | `watsonx/` |
|
||||
| Supported Operations | `/ml/v1/text/rerank` |
|
||||
| Link to Provider Doc | [IBM WatsonX.ai ↗](https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank) |
|
||||
|
||||
## Quick Start
|
||||
|
||||
### **LiteLLM SDK**
|
||||
|
||||
```python
|
||||
import os
|
||||
from litellm import rerank
|
||||
|
||||
os.environ["WATSONX_APIKEY"] = "YOUR_WATSONX_APIKEY"
|
||||
os.environ["WATSONX_API_BASE"] = "YOUR_WATSONX_API_BASE"
|
||||
os.environ["WATSONX_PROJECT_ID"] = "YOUR_WATSONX_PROJECT_ID"
|
||||
|
||||
query="Best programming language for beginners?"
|
||||
documents=[
|
||||
"Python is great for beginners due to simple syntax.",
|
||||
"JavaScript runs in browsers and is versatile.",
|
||||
"Rust has a steep learning curve but is very safe.",
|
||||
]
|
||||
|
||||
response = rerank(
|
||||
model="watsonx/cross-encoder/ms-marco-minilm-l-12-v2",
|
||||
query=query,
|
||||
documents=documents,
|
||||
top_n=2,
|
||||
return_documents=True,
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
### **LiteLLM Proxy**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: cross-encoder/ms-marco-minilm-l-12-v2
|
||||
litellm_params:
|
||||
model: watsonx/cross-encoder/ms-marco-minilm-l-12-v2
|
||||
api_key: os.environ/WATSONX_APIKEY
|
||||
api_base: os.environ/WATSONX_API_BASE
|
||||
project_id: os.environ/WATSONX_PROJECT_ID
|
||||
```
|
||||
|
|
@ -358,7 +358,8 @@ router_settings:
|
|||
| redis_url | str | URL for Redis server. **Known performance issue with Redis URL.** |
|
||||
| cache_responses | boolean | Flag to enable caching LLM Responses, if cache set under `router_settings`. If true, caches responses. Defaults to False. |
|
||||
| router_general_settings | RouterGeneralSettings | [SDK-Only] Router general settings - contains optimizations like 'async_only_mode'. [Docs](../routing.md#router-general-settings) |
|
||||
| optional_pre_call_checks | List[str] | List of pre-call checks to add to the router. Currently supported: 'router_budget_limiting', 'prompt_caching' |
|
||||
| optional_pre_call_checks | List[str] | List of pre-call checks to add to the router. Supported: `router_budget_limiting`, `prompt_caching`, `responses_api_deployment_check`, `deployment_affinity`, `forward_client_headers_by_model_group` |
|
||||
| deployment_affinity_ttl_seconds | int | TTL (seconds) for user-key → deployment affinity mapping when `deployment_affinity` is enabled (configured at Router init / proxy startup). Defaults to `3600` (1 hour). |
|
||||
| ignore_invalid_deployments | boolean | If true, ignores invalid deployments. Default for proxy is True - to prevent invalid models from blocking other models from being loaded. |
|
||||
| search_tools | List[SearchToolTypedDict] | List of search tool configurations for Search API integration. Each tool specifies a search_tool_name and litellm_params with search_provider, api_key, api_base, etc. [Further Docs](../search.md) |
|
||||
| guardrail_list | List[GuardrailTypedDict] | List of guardrail configurations for guardrail load balancing. Enables load balancing across multiple guardrail deployments with the same guardrail_name. [Further Docs](./guardrails/guardrail_load_balancing.md) |
|
||||
|
|
@ -493,6 +494,7 @@ router_settings:
|
|||
| DATABASE_USER | Username for database connection
|
||||
| DATABASE_USERNAME | Alias for database user
|
||||
| DATABRICKS_API_BASE | Base URL for Databricks API
|
||||
| DATABRICKS_API_KEY | API key (Personal Access Token) for Databricks API authentication
|
||||
| DATABRICKS_CLIENT_ID | Client ID for Databricks OAuth M2M authentication (Service Principal application ID)
|
||||
| DATABRICKS_CLIENT_SECRET | Client secret for Databricks OAuth M2M authentication
|
||||
| DATABRICKS_USER_AGENT | Custom user agent string for Databricks API requests. Used for partner telemetry attribution
|
||||
|
|
@ -540,7 +542,7 @@ router_settings:
|
|||
| DEFAULT_IMAGE_WIDTH | Default width for images. Default is 300
|
||||
| DEFAULT_IN_MEMORY_TTL | Default time-to-live for in-memory cache in seconds. Default is 5
|
||||
| DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL | Default time-to-live in seconds for management objects (User, Team, Key, Organization) in memory cache. Default is 60 seconds.
|
||||
| DEFAULT_MAX_LRU_CACHE_SIZE | Default maximum size for LRU cache. Default is 16
|
||||
| DEFAULT_MAX_LRU_CACHE_SIZE | Default maximum size for LRU cache. Default is 64
|
||||
| DEFAULT_MAX_RECURSE_DEPTH | Default maximum recursion depth. Default is 100
|
||||
| DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER | Default maximum recursion depth for sensitive data masker. Default is 10
|
||||
| DEFAULT_MAX_RETRIES | Default maximum retry attempts. Default is 2
|
||||
|
|
|
|||
|
|
@ -2,29 +2,98 @@ import Image from '@theme/IdealImage';
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Customers / End-User Budgets
|
||||
# Customers / End-Users
|
||||
|
||||
Track spend, set budgets for your customers.
|
||||
Track spend, set budgets and permissions for your customers.
|
||||
|
||||
## Tracking Customer Spend
|
||||
## Tracking Customer Spend + Permissions
|
||||
|
||||
### 1. Make LLM API call w/ Customer ID
|
||||
|
||||
Make a /chat/completions call, pass 'user' - First call Works
|
||||
LiteLLM checks for a customer/end-user ID in the following order (first match wins):
|
||||
|
||||
```bash showLineNumbers title="Make request with customer ID"
|
||||
| Priority | Method | Where | Notes |
|
||||
|----------|--------|-------|-------|
|
||||
| 1 | `x-litellm-customer-id` header | Request headers | Standard header, always checked |
|
||||
| 2 | `x-litellm-end-user-id` header | Request headers | Standard header, always checked |
|
||||
| 3 | Custom header via `user_header_mappings` | Request headers | Configured in `general_settings` |
|
||||
| 4 | Custom header via `user_header_name` | Request headers | Deprecated — use `user_header_mappings` |
|
||||
| 5 | `user` field | Request body | Standard OpenAI field |
|
||||
| 6 | `litellm_metadata.user` field | Request body | Anthropic-style metadata |
|
||||
| 7 | `metadata.user_id` field | Request body | Generic metadata pattern |
|
||||
| 8 | `safety_identifier` field | Request body | Responses API |
|
||||
|
||||
**Option 1: Standard headers** (recommended — no request body modification needed)
|
||||
|
||||
```bash showLineNumbers title="Make request with customer ID in header"
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-1234' \ # 👈 YOUR PROXY KEY
|
||||
--data ' {
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'x-litellm-end-user-id: ishaan3' \
|
||||
--data '{
|
||||
"model": "azure-gpt-3.5",
|
||||
"user": "ishaan3", # 👈 CUSTOMER ID
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what time is it"
|
||||
}
|
||||
]
|
||||
"messages": [{"role": "user", "content": "what time is it"}]
|
||||
}'
|
||||
```
|
||||
|
||||
Both `x-litellm-customer-id` and `x-litellm-end-user-id` are supported and always checked without any configuration.
|
||||
|
||||
**Option 2: `user` field in request body** (OpenAI-compatible)
|
||||
|
||||
```bash showLineNumbers title="Make request with customer ID in body"
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--data '{
|
||||
"model": "azure-gpt-3.5",
|
||||
"user": "ishaan3",
|
||||
"messages": [{"role": "user", "content": "what time is it"}]
|
||||
}'
|
||||
```
|
||||
|
||||
**Option 3: Custom header via `user_header_mappings`** (configurable)
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
general_settings:
|
||||
user_header_mappings:
|
||||
- header_name: "x-my-app-user-id"
|
||||
litellm_user_role: "customer"
|
||||
```
|
||||
|
||||
```bash showLineNumbers title="Make request with custom header"
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'x-my-app-user-id: ishaan3' \
|
||||
--data '{
|
||||
"model": "azure-gpt-3.5",
|
||||
"messages": [{"role": "user", "content": "what time is it"}]
|
||||
}'
|
||||
```
|
||||
|
||||
**Option 4: `litellm_metadata.user`** (Anthropic-style)
|
||||
|
||||
```bash showLineNumbers title="Make request with litellm_metadata.user"
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--data '{
|
||||
"model": "claude-3-5-sonnet",
|
||||
"messages": [{"role": "user", "content": "what time is it"}],
|
||||
"litellm_metadata": {"user": "ishaan3"}
|
||||
}'
|
||||
```
|
||||
|
||||
**Option 5: `metadata.user_id`**
|
||||
|
||||
```bash showLineNumbers title="Make request with metadata.user_id"
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--data '{
|
||||
"model": "azure-gpt-3.5",
|
||||
"messages": [{"role": "user", "content": "what time is it"}],
|
||||
"metadata": {"user_id": "ishaan3"}
|
||||
}'
|
||||
```
|
||||
|
||||
|
|
@ -123,7 +192,171 @@ Expected Response
|
|||
</Tabs>
|
||||
|
||||
|
||||
## Setting Customer Budgets
|
||||
## Setting Customer Object Permissions
|
||||
|
||||
Control which resources (MCP servers, vector stores, agents) a customer can access.
|
||||
|
||||
### What are Object Permissions?
|
||||
|
||||
Object permissions allow you to restrict customer access to specific:
|
||||
- **MCP Servers**: Limit which MCP servers the customer can call
|
||||
- **MCP Access Groups**: Assign customers to predefined groups of MCP servers
|
||||
- **MCP Tool Permissions**: Granular control over which tools within an MCP server the customer can use
|
||||
- **Vector Stores**: Control which vector stores the customer can query
|
||||
- **Agents**: Restrict which agents the customer can interact with
|
||||
- **Agent Access Groups**: Assign customers to predefined groups of agents
|
||||
|
||||
### Creating a Customer with Object Permissions
|
||||
|
||||
```bash showLineNumbers title="Create customer with object permissions"
|
||||
curl -L -X POST 'http://localhost:4000/customer/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"user_id": "user_1",
|
||||
"object_permission": {
|
||||
"mcp_servers": ["server_1", "server_2"],
|
||||
"mcp_access_groups": ["public_group"],
|
||||
"mcp_tool_permissions": {
|
||||
"server_1": ["tool_a", "tool_b"]
|
||||
},
|
||||
"vector_stores": ["vector_store_1"],
|
||||
"agents": ["agent_1"],
|
||||
"agent_access_groups": ["basic_agents"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**Parameters:**
|
||||
- `mcp_servers` (Optional[List[str]]): List of allowed MCP server IDs
|
||||
- `mcp_access_groups` (Optional[List[str]]): List of MCP access group names
|
||||
- `mcp_tool_permissions` (Optional[Dict[str, List[str]]]): Map of server ID to allowed tool names
|
||||
- `vector_stores` (Optional[List[str]]): List of allowed vector store IDs
|
||||
- `agents` (Optional[List[str]]): List of allowed agent IDs
|
||||
- `agent_access_groups` (Optional[List[str]]): List of agent access group names
|
||||
|
||||
**Note:** If `object_permission` is `null` or `{}`, the customer has no object-level restrictions.
|
||||
|
||||
### Updating Customer Object Permissions
|
||||
|
||||
You can update object permissions for existing customers:
|
||||
|
||||
```bash showLineNumbers title="Update customer object permissions"
|
||||
curl -L -X POST 'http://localhost:4000/customer/update' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"user_id": "user_1",
|
||||
"object_permission": {
|
||||
"mcp_servers": ["server_3"],
|
||||
"vector_stores": ["vector_store_2", "vector_store_3"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### Viewing Customer Object Permissions
|
||||
|
||||
When you query customer info, object permissions are included in the response:
|
||||
|
||||
```bash showLineNumbers title="Get customer info with object permissions"
|
||||
curl -X GET 'http://0.0.0.0:4000/customer/info?end_user_id=user_1' \
|
||||
-H 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
|
||||
**Response:**
|
||||
```json showLineNumbers title="Response with object permissions"
|
||||
{
|
||||
"user_id": "user_1",
|
||||
"blocked": false,
|
||||
"alias": "John Doe",
|
||||
"spend": 0.0,
|
||||
"object_permission": {
|
||||
"object_permission_id": "perm_abc123",
|
||||
"mcp_servers": ["server_1", "server_2"],
|
||||
"mcp_access_groups": ["public_group"],
|
||||
"mcp_tool_permissions": {
|
||||
"server_1": ["tool_a", "tool_b"]
|
||||
},
|
||||
"vector_stores": ["vector_store_1"],
|
||||
"agents": ["agent_1"],
|
||||
"agent_access_groups": ["basic_agents"]
|
||||
},
|
||||
"litellm_budget_table": null
|
||||
}
|
||||
```
|
||||
|
||||
### Use Cases
|
||||
|
||||
**1. Tiered Access Control**
|
||||
Create different permission tiers for your customers:
|
||||
|
||||
```bash showLineNumbers title="Free tier customer"
|
||||
# Free tier - limited access
|
||||
curl -L -X POST 'http://localhost:4000/customer/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"user_id": "free_user",
|
||||
"budget_id": "free_tier",
|
||||
"object_permission": {
|
||||
"mcp_access_groups": ["public_group"],
|
||||
"agent_access_groups": ["basic_agents"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
```bash showLineNumbers title="Premium tier customer"
|
||||
# Premium tier - full access
|
||||
curl -L -X POST 'http://localhost:4000/customer/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"user_id": "premium_user",
|
||||
"budget_id": "premium_tier",
|
||||
"object_permission": {
|
||||
"mcp_servers": ["server_1", "server_2", "server_3"],
|
||||
"vector_stores": ["vector_store_1", "vector_store_2"],
|
||||
"agents": ["agent_1", "agent_2", "agent_3"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**2. Department-Specific Access**
|
||||
Restrict customers to resources relevant to their department:
|
||||
|
||||
```bash showLineNumbers title="Sales team customer"
|
||||
curl -L -X POST 'http://localhost:4000/customer/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"user_id": "sales_user",
|
||||
"object_permission": {
|
||||
"mcp_servers": ["crm_server", "email_server"],
|
||||
"agents": ["sales_assistant"],
|
||||
"vector_stores": ["sales_knowledge_base"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**3. Tool-Level Restrictions**
|
||||
Grant access to specific tools within an MCP server:
|
||||
|
||||
```bash showLineNumbers title="Limited tool access"
|
||||
curl -L -X POST 'http://localhost:4000/customer/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"user_id": "restricted_user",
|
||||
"object_permission": {
|
||||
"mcp_servers": ["database_server"],
|
||||
"mcp_tool_permissions": {
|
||||
"database_server": ["read_only_query", "get_table_schema"]
|
||||
}
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
## Setting Customer Budgets
|
||||
|
||||
Set customer budgets (e.g. monthly budgets, tpm/rpm limits) on LiteLLM Proxy
|
||||
|
||||
|
|
|
|||
|
|
@ -20,6 +20,10 @@ By default, LiteLLM does not forward client headers to LLM provider APIs. Howeve
|
|||
|
||||
`x-litellm-spend-logs-metadata`: Optional[str]: JSON string containing custom metadata to include in spend logs. Example: `{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion"}`. [Learn More](../proxy/enterprise#tracking-spend-with-custom-metadata)
|
||||
|
||||
`x-litellm-customer-id`: Optional[str]: Standard header for passing a customer/end-user ID. Always checked without any configuration. [Learn More](./customers)
|
||||
|
||||
`x-litellm-end-user-id`: Optional[str]: Standard header for passing a customer/end-user ID. Always checked without any configuration. [Learn More](./customers)
|
||||
|
||||
## Anthropic Headers
|
||||
|
||||
`anthropic-version` Optional[str]: The version of the Anthropic API to use.
|
||||
|
|
|
|||
|
|
@ -8,15 +8,15 @@ LiteLLM Follows the [cohere api request / response for the rerank api](https://c
|
|||
|
||||
## Overview
|
||||
|
||||
| Feature | Supported | Notes |
|
||||
|---------|-----------|-------|
|
||||
| Cost Tracking | ✅ | Works with all supported models |
|
||||
| Logging | ✅ | Works across all integrations |
|
||||
| End-user Tracking | ✅ | |
|
||||
| Fallbacks | ✅ | Works between supported models |
|
||||
| Loadbalancing | ✅ | Works between supported models |
|
||||
| Guardrails | ✅ | Applies to input query only (not documents) |
|
||||
| Supported Providers | Cohere, Together AI, Azure AI, DeepInfra, Nvidia NIM, Infinity, Fireworks AI, Voyage AI | |
|
||||
| Feature | Supported | Notes |
|
||||
|---------|-----------------------------------------------------------------------------------------------------|-------|
|
||||
| Cost Tracking | ✅ | Works with all supported models |
|
||||
| Logging | ✅ | Works across all integrations |
|
||||
| End-user Tracking | ✅ | |
|
||||
| Fallbacks | ✅ | Works between supported models |
|
||||
| Loadbalancing | ✅ | Works between supported models |
|
||||
| Guardrails | ✅ | Applies to input query only (not documents) |
|
||||
| Supported Providers | Cohere, Together AI, Azure AI, DeepInfra, Nvidia NIM, Infinity, Fireworks AI, Voyage AI, watsonx.ai | |
|
||||
|
||||
## **LiteLLM Python SDK Usage**
|
||||
### Quick Start
|
||||
|
|
@ -123,17 +123,18 @@ curl http://0.0.0.0:4000/rerank \
|
|||
|
||||
#### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/)
|
||||
|
||||
| Provider | Link to Usage |
|
||||
|-------------|--------------------|
|
||||
| Cohere (v1 + v2 clients) | [Usage](#quick-start) |
|
||||
| Together AI| [Usage](../docs/providers/togetherai) |
|
||||
| Azure AI| [Usage](../docs/providers/azure_ai#rerank-endpoint) |
|
||||
| Jina AI| [Usage](../docs/providers/jina_ai) |
|
||||
| AWS Bedrock| [Usage](../docs/providers/bedrock#rerank-api) |
|
||||
| HuggingFace| [Usage](../docs/providers/huggingface_rerank) |
|
||||
| Infinity| [Usage](../docs/providers/infinity) |
|
||||
| vLLM| [Usage](../docs/providers/vllm#rerank-endpoint) |
|
||||
| DeepInfra| [Usage](../docs/providers/deepinfra#rerank-endpoint) |
|
||||
| Vertex AI| [Usage](../docs/providers/vertex#rerank-api) |
|
||||
| Fireworks AI| [Usage](../docs/providers/fireworks_ai#rerank-endpoint) |
|
||||
| Voyage AI| [Usage](../docs/providers/voyage#rerank) |
|
||||
| Provider | Link to Usage |
|
||||
|--------------------------|------------------------------------------------------|
|
||||
| Cohere (v1 + v2 clients) | [Usage](#quick-start) |
|
||||
| Together AI | [Usage](../docs/providers/togetherai) |
|
||||
| Azure AI | [Usage](../docs/providers/azure_ai#rerank-endpoint) |
|
||||
| Jina AI | [Usage](../docs/providers/jina_ai) |
|
||||
| AWS Bedrock | [Usage](../docs/providers/bedrock#rerank-api) |
|
||||
| HuggingFace | [Usage](../docs/providers/huggingface_rerank) |
|
||||
| Infinity | [Usage](../docs/providers/infinity) |
|
||||
| vLLM | [Usage](../docs/providers/vllm#rerank-endpoint) |
|
||||
| DeepInfra | [Usage](../docs/providers/deepinfra#rerank-endpoint) |
|
||||
| Vertex AI | [Usage](../docs/providers/vertex#rerank-api) |
|
||||
| Fireworks AI | [Usage](../docs/providers/fireworks_ai#rerank-endpoint) |
|
||||
| Voyage AI | [Usage](../docs/providers/voyage#rerank) |
|
||||
| IBM watsonx.ai | [Usage](../docs/providers/watsonx/rerank) |
|
||||
|
|
@ -884,7 +884,12 @@ router = litellm.Router(
|
|||
},
|
||||
},
|
||||
],
|
||||
optional_pre_call_checks=["responses_api_deployment_check"],
|
||||
# `responses_api_deployment_check` ensures Requests with `previous_response_id`
|
||||
# are routed to the same deployment. `deployment_affinity` adds sticky sessions
|
||||
# for requests without `previous_response_id` (useful for implicit caching).
|
||||
optional_pre_call_checks=["responses_api_deployment_check", "deployment_affinity"],
|
||||
# Optional (default is 3600 seconds / 1 hour)
|
||||
deployment_affinity_ttl_seconds=3600,
|
||||
)
|
||||
|
||||
# Initial request
|
||||
|
|
@ -911,7 +916,16 @@ follow_up = await router.aresponses(
|
|||
|
||||
#### 1. Setup session continuity on proxy config.yaml
|
||||
|
||||
To enable session continuity for Responses API in your LiteLLM proxy, set `optional_pre_call_checks: ["responses_api_deployment_check"]` in your proxy config.yaml.
|
||||
To enable session continuity for Responses API in your LiteLLM proxy, set `optional_pre_call_checks` in your proxy config.yaml.
|
||||
|
||||
- `responses_api_deployment_check`: high priority routing when `previous_response_id` is provided
|
||||
- `deployment_affinity`: sticky sessions based on user key (applies even without `previous_response_id`)
|
||||
|
||||
Notes:
|
||||
- User-key affinity is keyed on `metadata.user_api_key_hash` (the API key hash). The OpenAI `user` request parameter is an end-user identifier and is intentionally not used for deployment affinity.
|
||||
- `user_api_key_hash` is already SHA-256, and is used as-is (no double hashing).
|
||||
- Affinity is scoped by a stable model identifier (the model-map key, e.g. `model_map_information.model_map_key`) so model aliases map to the same stickiness bucket.
|
||||
- The mapping TTL is controlled by `deployment_affinity_ttl_seconds` (configured on Router init / proxy startup).
|
||||
|
||||
```yaml showLineNumbers title="config.yaml with Session Continuity"
|
||||
model_list:
|
||||
|
|
@ -929,7 +943,11 @@ model_list:
|
|||
api_base: https://endpoint2.openai.azure.com
|
||||
|
||||
router_settings:
|
||||
optional_pre_call_checks: ["responses_api_deployment_check"]
|
||||
optional_pre_call_checks:
|
||||
- responses_api_deployment_check
|
||||
- deployment_affinity
|
||||
# Optional (default is 3600 seconds / 1 hour)
|
||||
deployment_affinity_ttl_seconds: 3600
|
||||
```
|
||||
|
||||
#### 2. Use the OpenAI Python SDK to make requests to LiteLLM Proxy
|
||||
|
|
@ -1356,8 +1374,3 @@ Response:
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -276,6 +276,7 @@ The response follows Perplexity's search format with the following structure:
|
|||
| Firecrawl | `FIRECRAWL_API_KEY` | `firecrawl` |
|
||||
| SearXNG | `SEARXNG_API_BASE` (required) | `searxng` |
|
||||
| Linkup | `LINKUP_API_KEY` | `linkup` |
|
||||
| DuckDuckGo | `DUCKDUCKGO_API_BASE` | `duckduckgo` |
|
||||
|
||||
See the individual provider documentation for detailed setup instructions and provider-specific parameters.
|
||||
|
||||
|
|
|
|||
|
|
@ -944,6 +944,7 @@ const sidebars = {
|
|||
"providers/anthropic_tool_search",
|
||||
"guides/code_interpreter",
|
||||
"completion/message_trimming",
|
||||
"completion/message_sanitization",
|
||||
"completion/model_alias",
|
||||
"completion/mock_requests",
|
||||
"completion/predict_outputs",
|
||||
|
|
|
|||
|
|
@ -1051,6 +1051,168 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
"""Handled in files_endpoints.py"""
|
||||
return []
|
||||
|
||||
def _is_batch_polling_enabled(self) -> bool:
|
||||
"""
|
||||
Check if batch cost tracking is actually enabled and running.
|
||||
Returns:
|
||||
bool: True if batch cost tracking is active, False otherwise
|
||||
"""
|
||||
try:
|
||||
# Import here to avoid circular dependencies
|
||||
import litellm.proxy.proxy_server as proxy_server_module
|
||||
|
||||
# Check if the scheduler has the batch cost checking job registered
|
||||
scheduler = getattr(proxy_server_module, 'scheduler', None)
|
||||
if scheduler is None:
|
||||
return False
|
||||
|
||||
# Check if the check_batch_cost_job exists in the scheduler
|
||||
try:
|
||||
job = scheduler.get_job('check_batch_cost_job')
|
||||
if job is not None:
|
||||
return True
|
||||
except Exception:
|
||||
# Job not found or scheduler doesn't support get_job
|
||||
pass
|
||||
|
||||
return False
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Error checking batch polling configuration: {e}. Assuming disabled."
|
||||
)
|
||||
return False
|
||||
|
||||
async def _get_batches_referencing_file(
|
||||
self, file_id: str
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Find batches in non-terminal states that reference this file.
|
||||
|
||||
Non-terminal states: validating, in_progress, finalizing
|
||||
Terminal states: completed, complete, failed, expired, cancelled
|
||||
|
||||
Args:
|
||||
file_id: The unified file ID to check
|
||||
|
||||
Returns:
|
||||
List of batch objects referencing this file in non-terminal state
|
||||
(max 10 for error message display)
|
||||
"""
|
||||
# Prepare list of file IDs to check (both unified and provider IDs)
|
||||
file_ids_to_check = [file_id]
|
||||
|
||||
# Get model-specific file IDs for this unified file ID if it's a managed file
|
||||
try:
|
||||
model_file_id_mapping = await self.get_model_file_id_mapping(
|
||||
[file_id], litellm_parent_otel_span=None
|
||||
)
|
||||
|
||||
if model_file_id_mapping and file_id in model_file_id_mapping:
|
||||
# Add all provider file IDs for this unified file
|
||||
provider_file_ids = list(model_file_id_mapping[file_id].values())
|
||||
file_ids_to_check.extend(provider_file_ids)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Could not get model file ID mapping for {file_id}: {e}. "
|
||||
f"Will only check unified file ID."
|
||||
)
|
||||
MAX_MATCHES_TO_RETURN = 10
|
||||
|
||||
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
|
||||
where={
|
||||
"file_purpose": "batch",
|
||||
"status": {"in": ["validating", "in_progress", "finalizing"]},
|
||||
},
|
||||
take=MAX_MATCHES_TO_RETURN,
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
referencing_batches = []
|
||||
for batch in batches:
|
||||
try:
|
||||
# Parse the batch file_object to check for file references
|
||||
batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object
|
||||
|
||||
# Extract file IDs from batch
|
||||
# Batches typically reference the unified file ID in input_file_id
|
||||
# Output and error files are generated by the provider
|
||||
input_file_id = batch_data.get("input_file_id")
|
||||
output_file_id = batch_data.get("output_file_id")
|
||||
error_file_id = batch_data.get("error_file_id")
|
||||
|
||||
referenced_file_ids = [fid for fid in [input_file_id, output_file_id, error_file_id] if fid]
|
||||
|
||||
# Check if any referenced file ID matches the file we're trying to delete
|
||||
if any(ref_id in file_ids_to_check for ref_id in referenced_file_ids):
|
||||
referencing_batches.append({
|
||||
"batch_id": batch.unified_object_id,
|
||||
"status": batch.status,
|
||||
"created_at": batch.created_at,
|
||||
})
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Error parsing batch object {batch.unified_object_id}: {e}"
|
||||
)
|
||||
continue
|
||||
|
||||
return referencing_batches
|
||||
|
||||
async def _check_file_deletion_allowed(self, file_id: str) -> None:
|
||||
"""
|
||||
Check if file deletion should be blocked due to batch references.
|
||||
|
||||
Blocks deletion if:
|
||||
1. File is referenced by any batch in non-terminal state, AND
|
||||
2. Batch polling is configured (user wants cost tracking)
|
||||
|
||||
Args:
|
||||
file_id: The unified file ID to check
|
||||
|
||||
Raises:
|
||||
HTTPException: If file deletion should be blocked
|
||||
"""
|
||||
# Check if batch polling is enabled
|
||||
if not self._is_batch_polling_enabled():
|
||||
# Batch polling not configured, allow deletion
|
||||
return
|
||||
|
||||
# Check if file is referenced by any non-terminal batches
|
||||
referencing_batches = await self._get_batches_referencing_file(file_id)
|
||||
|
||||
if referencing_batches:
|
||||
# File is referenced by non-terminal batches and polling is enabled
|
||||
MAX_BATCHES_IN_ERROR = 5 # Limit batches shown in error message for readability
|
||||
|
||||
# Show up to MAX_BATCHES_IN_ERROR in the error message
|
||||
batches_to_show = referencing_batches[:MAX_BATCHES_IN_ERROR]
|
||||
batch_statuses = [f"{b['batch_id']}: {b['status']}" for b in batches_to_show]
|
||||
|
||||
# Determine the count message
|
||||
count_message = f"{len(referencing_batches)}"
|
||||
if len(referencing_batches) >= 10: # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file
|
||||
count_message = "10+"
|
||||
|
||||
error_message = (
|
||||
f"Cannot delete file {file_id}. "
|
||||
f"The file is referenced by {count_message} batch(es) in non-terminal state"
|
||||
)
|
||||
|
||||
# Add specific batch details if not too many
|
||||
if len(referencing_batches) <= MAX_BATCHES_IN_ERROR:
|
||||
error_message += f": {', '.join(batch_statuses)}. "
|
||||
else:
|
||||
error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. "
|
||||
|
||||
error_message += (
|
||||
f"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
|
||||
f"Alternatively, wait for all batches to complete processing."
|
||||
)
|
||||
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=error_message,
|
||||
)
|
||||
|
||||
async def afile_delete(
|
||||
self,
|
||||
file_id: str,
|
||||
|
|
@ -1059,6 +1221,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
**data: Dict,
|
||||
) -> OpenAIFileObject:
|
||||
|
||||
# Check if file deletion should be blocked due to batch references
|
||||
await self._check_file_deletion_allowed(file_id)
|
||||
|
||||
# file_id = convert_b64_uid_to_unified_uid(file_id)
|
||||
model_file_id_mapping = await self.get_model_file_id_mapping(
|
||||
[file_id], litellm_parent_otel_span
|
||||
|
|
|
|||
|
|
@ -1,10 +1,13 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ManagedVectorStoresTable" ADD COLUMN "team_id" TEXT,
|
||||
ADD COLUMN "user_id" TEXT;
|
||||
ALTER TABLE "LiteLLM_ManagedVectorStoresTable"
|
||||
ADD COLUMN IF NOT EXISTS "team_id" TEXT,
|
||||
ADD COLUMN IF NOT EXISTS "user_id" TEXT;
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_ManagedVectorStoresTable_team_id_idx" ON "LiteLLM_ManagedVectorStoresTable"("team_id");
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_ManagedVectorStoresTable_team_id_idx"
|
||||
ON "LiteLLM_ManagedVectorStoresTable"("team_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_ManagedVectorStoresTable_user_id_idx" ON "LiteLLM_ManagedVectorStoresTable"("user_id");
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_ManagedVectorStoresTable_user_id_idx"
|
||||
ON "LiteLLM_ManagedVectorStoresTable"("user_id");
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,6 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_EndUserTable" ADD COLUMN "object_permission_id" TEXT;
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_EndUserTable" ADD CONSTRAINT "LiteLLM_EndUserTable_object_permission_id_fkey" FOREIGN KEY ("object_permission_id") REFERENCES "LiteLLM_ObjectPermissionTable"("object_permission_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
|
||||
|
|
@ -233,6 +233,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
verification_tokens LiteLLM_VerificationToken[]
|
||||
organizations LiteLLM_OrganizationTable[]
|
||||
users LiteLLM_UserTable[]
|
||||
end_users LiteLLM_EndUserTable[]
|
||||
}
|
||||
|
||||
// Holds the MCP server configuration
|
||||
|
|
@ -403,7 +404,9 @@ model LiteLLM_EndUserTable {
|
|||
allowed_model_region String? // require all user requests to use models in this specific region
|
||||
default_model String? // use along with 'allowed_model_region'. if no available model in region, default to this model.
|
||||
budget_id String?
|
||||
object_permission_id String?
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
blocked Boolean @default(false)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1355,6 +1355,7 @@ if TYPE_CHECKING:
|
|||
from .llms.vertex_ai.rerank.transformation import VertexAIRerankConfig as VertexAIRerankConfig
|
||||
from .llms.fireworks_ai.rerank.transformation import FireworksAIRerankConfig as FireworksAIRerankConfig
|
||||
from .llms.voyage.rerank.transformation import VoyageRerankConfig as VoyageRerankConfig
|
||||
from .llms.watsonx.rerank.transformation import IBMWatsonXRerankConfig as IBMWatsonXRerankConfig
|
||||
from .llms.clarifai.chat.transformation import ClarifaiConfig as ClarifaiConfig
|
||||
from .llms.ai21.chat.transformation import AI21ChatConfig as AI21ChatConfig
|
||||
from .llms.meta_llama.chat.transformation import LlamaAPIConfig as LlamaAPIConfig
|
||||
|
|
|
|||
|
|
@ -155,6 +155,7 @@ LLM_CONFIG_NAMES = (
|
|||
"VertexAIRerankConfig",
|
||||
"FireworksAIRerankConfig",
|
||||
"VoyageRerankConfig",
|
||||
"IBMWatsonXRerankConfig",
|
||||
"ClarifaiConfig",
|
||||
"AI21ChatConfig",
|
||||
"LlamaAPIConfig",
|
||||
|
|
@ -672,6 +673,7 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
"FireworksAIRerankConfig",
|
||||
),
|
||||
"VoyageRerankConfig": (".llms.voyage.rerank.transformation", "VoyageRerankConfig"),
|
||||
"IBMWatsonXRerankConfig": (".llms.watsonx.rerank.transformation", "IBMWatsonXRerankConfig"),
|
||||
"ClarifaiConfig": (".llms.clarifai.chat.transformation", "ClarifaiConfig"),
|
||||
"AI21ChatConfig": (".llms.ai21.chat.transformation", "AI21ChatConfig"),
|
||||
"LlamaAPIConfig": (".llms.meta_llama.chat.transformation", "LlamaAPIConfig"),
|
||||
|
|
|
|||
|
|
@ -67,7 +67,7 @@
|
|||
"compact-2026-01-12": null,
|
||||
"computer-use-2025-01-24": "computer-use-2025-01-24",
|
||||
"computer-use-2025-11-24": "computer-use-2025-11-24",
|
||||
"context-1m-2025-08-07": null,
|
||||
"context-1m-2025-08-07": "context-1m-2025-08-07",
|
||||
"context-management-2025-06-27": "context-management-2025-06-27",
|
||||
"effort-2025-11-24": null,
|
||||
"fast-mode-2026-02-01": null,
|
||||
|
|
@ -148,5 +148,35 @@
|
|||
"tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19",
|
||||
"web-fetch-2025-09-10": null,
|
||||
"web-search-2025-03-05": "web-search-2025-03-05"
|
||||
},
|
||||
"databricks": {
|
||||
"advanced-tool-use-2025-11-20": "advanced-tool-use-2025-11-20",
|
||||
"bash_20241022": null,
|
||||
"bash_20250124": null,
|
||||
"code-execution-2025-08-25": "code-execution-2025-08-25",
|
||||
"compact-2026-01-12": "compact-2026-01-12",
|
||||
"computer-use-2025-01-24": "computer-use-2025-01-24",
|
||||
"computer-use-2025-11-24": "computer-use-2025-11-24",
|
||||
"context-1m-2025-08-07": "context-1m-2025-08-07",
|
||||
"context-management-2025-06-27": "context-management-2025-06-27",
|
||||
"effort-2025-11-24": "effort-2025-11-24",
|
||||
"fast-mode-2026-02-01": "fast-mode-2026-02-01",
|
||||
"files-api-2025-04-14": "files-api-2025-04-14",
|
||||
"structured-output-2024-03-01": null,
|
||||
"fine-grained-tool-streaming-2025-05-14": "fine-grained-tool-streaming-2025-05-14",
|
||||
"interleaved-thinking-2025-05-14": "interleaved-thinking-2025-05-14",
|
||||
"mcp-client-2025-11-20": "mcp-client-2025-11-20",
|
||||
"mcp-client-2025-04-04": "mcp-client-2025-04-04",
|
||||
"mcp-servers-2025-12-04": null,
|
||||
"oauth-2025-04-20": "oauth-2025-04-20",
|
||||
"output-128k-2025-02-19": "output-128k-2025-02-19",
|
||||
"prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05",
|
||||
"skills-2025-10-02": "skills-2025-10-02",
|
||||
"structured-outputs-2025-11-13": "structured-outputs-2025-11-13",
|
||||
"text_editor_20241022": null,
|
||||
"text_editor_20250124": null,
|
||||
"token-efficient-tools-2025-02-19": "token-efficient-tools-2025-02-19",
|
||||
"web-fetch-2025-09-10": "web-fetch-2025-09-10",
|
||||
"web-search-2025-03-05": "web-search-2025-03-05"
|
||||
}
|
||||
}
|
||||
|
|
@ -62,9 +62,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
def __init__(self):
|
||||
pass
|
||||
|
||||
def _handle_raw_dict_response_item(
|
||||
self, item: Dict[str, Any], index: int
|
||||
) -> Tuple[Optional[Any], int]:
|
||||
def _handle_raw_dict_response_item(self, item: Dict[str, Any], index: int) -> Tuple[Optional[Any], int]:
|
||||
"""
|
||||
Handle raw dict response items from Responses API (e.g., GPT-5 Codex format).
|
||||
|
||||
|
|
@ -107,13 +105,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if item_type == "function_call":
|
||||
# Extract provider_specific_fields if present and pass through as-is
|
||||
provider_specific_fields = item.get("provider_specific_fields")
|
||||
if provider_specific_fields and not isinstance(
|
||||
provider_specific_fields, dict
|
||||
):
|
||||
if provider_specific_fields and not isinstance(provider_specific_fields, dict):
|
||||
provider_specific_fields = (
|
||||
dict(provider_specific_fields)
|
||||
if hasattr(provider_specific_fields, "__dict__")
|
||||
else {}
|
||||
dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {}
|
||||
)
|
||||
|
||||
tool_call_dict = {
|
||||
|
|
@ -129,9 +123,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if provider_specific_fields:
|
||||
tool_call_dict["provider_specific_fields"] = provider_specific_fields
|
||||
# Also add to function's provider_specific_fields for consistency
|
||||
tool_call_dict["function"][
|
||||
"provider_specific_fields"
|
||||
] = provider_specific_fields
|
||||
tool_call_dict["function"]["provider_specific_fields"] = provider_specific_fields
|
||||
|
||||
msg = Message(
|
||||
content=None,
|
||||
|
|
@ -169,7 +161,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"type": "message",
|
||||
"role": role,
|
||||
"content": self._convert_content_to_responses_format(
|
||||
content, role # type: ignore
|
||||
content,
|
||||
role, # type: ignore
|
||||
),
|
||||
}
|
||||
)
|
||||
|
|
@ -186,7 +179,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
elif isinstance(content, list):
|
||||
# Transform list content to Responses API format
|
||||
tool_output = self._convert_content_to_responses_format(
|
||||
content, "user" # Use "user" role to get input_* types
|
||||
content,
|
||||
"user", # Use "user" role to get input_* types
|
||||
)
|
||||
else:
|
||||
# Fallback: convert unexpected types to input_text
|
||||
|
|
@ -219,9 +213,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
{
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"content": self._convert_content_to_responses_format(
|
||||
content, cast(str, role)
|
||||
),
|
||||
"content": self._convert_content_to_responses_format(content, cast(str, role)),
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -344,9 +336,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
previous_response_id = optional_params.get("previous_response_id")
|
||||
if previous_response_id:
|
||||
# Use the existing session handler for responses API
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Warning ignoring previous response ID: {previous_response_id}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Warning ignoring previous response ID: {previous_response_id}")
|
||||
|
||||
# Convert back to responses API format for the actual request
|
||||
|
||||
|
|
@ -368,9 +358,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"client": client,
|
||||
}
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Final request model={api_model}, input_items={len(input_items)}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Final request model={api_model}, input_items={len(input_items)}")
|
||||
|
||||
self._merge_responses_api_request_into_request_data(
|
||||
request_data, responses_api_request, instructions
|
||||
|
|
@ -450,9 +438,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=item,
|
||||
index=tool_call_index,
|
||||
tool_call_dict = (
|
||||
LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=item,
|
||||
index=tool_call_index,
|
||||
)
|
||||
)
|
||||
accumulated_tool_calls.append(tool_call_dict)
|
||||
tool_call_index += 1
|
||||
|
|
@ -472,9 +462,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
tool_calls=accumulated_tool_calls,
|
||||
reasoning_content=reasoning_content,
|
||||
)
|
||||
choices.append(
|
||||
Choices(message=msg, finish_reason="tool_calls", index=index)
|
||||
)
|
||||
choices.append(Choices(message=msg, finish_reason="tool_calls", index=index))
|
||||
reasoning_content = None
|
||||
|
||||
return choices
|
||||
|
|
@ -510,17 +498,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
)
|
||||
|
||||
if len(choices) == 0:
|
||||
if (
|
||||
raw_response.incomplete_details is not None
|
||||
and raw_response.incomplete_details.reason is not None
|
||||
):
|
||||
raise ValueError(
|
||||
f"{model} unable to complete request: {raw_response.incomplete_details.reason}"
|
||||
)
|
||||
if raw_response.incomplete_details is not None and raw_response.incomplete_details.reason is not None:
|
||||
raise ValueError(f"{model} unable to complete request: {raw_response.incomplete_details.reason}")
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown items in responses API response: {raw_response.output}"
|
||||
)
|
||||
raise ValueError(f"Unknown items in responses API response: {raw_response.output}")
|
||||
|
||||
setattr(model_response, "choices", choices)
|
||||
|
||||
|
|
@ -529,11 +510,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
setattr(
|
||||
model_response,
|
||||
"usage",
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
raw_response.usage
|
||||
),
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_response.usage),
|
||||
)
|
||||
|
||||
|
||||
# Preserve hidden params from the ResponsesAPIResponse, especially the headers
|
||||
# which contain important provider information like x-request-id
|
||||
raw_response_hidden_params = getattr(raw_response, "_hidden_params", {})
|
||||
|
|
@ -550,24 +529,18 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
model_response._hidden_params[key] = merged_headers
|
||||
else:
|
||||
model_response._hidden_params[key] = value
|
||||
|
||||
|
||||
return model_response
|
||||
|
||||
def get_model_response_iterator(
|
||||
self,
|
||||
streaming_response: Union[
|
||||
Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel"
|
||||
],
|
||||
streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel"],
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
) -> BaseModelResponseIterator:
|
||||
return OpenAiResponsesToChatCompletionStreamIterator(
|
||||
streaming_response, sync_stream, json_mode
|
||||
)
|
||||
return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode)
|
||||
|
||||
def _convert_content_str_to_input_text(
|
||||
self, content: str, role: str
|
||||
) -> Dict[str, Any]:
|
||||
def _convert_content_str_to_input_text(self, content: str, role: str) -> Dict[str, Any]:
|
||||
if role == "user" or role == "system" or role == "tool":
|
||||
return {"type": "input_text", "text": content}
|
||||
else:
|
||||
|
|
@ -594,9 +567,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if actual_image_url is None:
|
||||
raise ValueError(f"Invalid image URL: {content_image_url}")
|
||||
|
||||
image_param = ResponseInputImageParam(
|
||||
image_url=actual_image_url, detail="auto", type="input_image"
|
||||
)
|
||||
image_param = ResponseInputImageParam(image_url=actual_image_url, detail="auto", type="input_image")
|
||||
|
||||
if detail:
|
||||
image_param["detail"] = detail
|
||||
|
|
@ -605,31 +576,29 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
def _convert_content_to_responses_format(
|
||||
self,
|
||||
content: Union[
|
||||
str,
|
||||
Iterable[
|
||||
Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock"]
|
||||
],
|
||||
content: Optional[
|
||||
Union[
|
||||
str,
|
||||
Iterable[Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock"]],
|
||||
]
|
||||
],
|
||||
role: str,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Convert chat completion content to responses API format"""
|
||||
from litellm.types.llms.openai import ChatCompletionImageObject
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Converting content to responses format - input type: {type(content)}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Converting content to responses format - input type: {type(content)}")
|
||||
|
||||
if isinstance(content, str):
|
||||
if content is None:
|
||||
return [self._convert_content_str_to_input_text("", role)]
|
||||
elif isinstance(content, str):
|
||||
result = [self._convert_content_str_to_input_text(content, role)]
|
||||
verbose_logger.debug(f"Chat provider: String content -> {result}")
|
||||
return result
|
||||
elif isinstance(content, list):
|
||||
result = []
|
||||
for i, item in enumerate(content):
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Processing content item {i}: {type(item)} = {item}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Processing content item {i}: {type(item)} = {item}")
|
||||
if isinstance(item, str):
|
||||
converted = self._convert_content_str_to_input_text(item, role)
|
||||
result.append(converted)
|
||||
|
|
@ -638,9 +607,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
# Handle multimodal content
|
||||
original_type = item.get("type")
|
||||
if original_type == "text":
|
||||
converted = self._convert_content_str_to_input_text(
|
||||
item.get("text", ""), role
|
||||
)
|
||||
converted = self._convert_content_str_to_input_text(item.get("text", ""), role)
|
||||
result.append(converted)
|
||||
verbose_logger.debug(f"Chat provider: text -> {converted}")
|
||||
elif original_type == "image_url":
|
||||
|
|
@ -652,18 +619,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
),
|
||||
)
|
||||
result.append(converted)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: image_url -> {converted}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: image_url -> {converted}")
|
||||
else:
|
||||
# Try to map other types to responses API format
|
||||
item_type = original_type or "input_text"
|
||||
if item_type == "image":
|
||||
converted = {"type": "input_image", **item}
|
||||
result.append(converted)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: image -> {converted}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: image -> {converted}")
|
||||
elif item_type in [
|
||||
"input_text",
|
||||
"input_image",
|
||||
|
|
@ -675,18 +638,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
]:
|
||||
# Already in responses API format
|
||||
result.append(item)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: passthrough -> {item}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: passthrough -> {item}")
|
||||
else:
|
||||
# Default to input_text for unknown types
|
||||
converted = self._convert_content_str_to_input_text(
|
||||
str(item.get("text", item)), role
|
||||
)
|
||||
converted = self._convert_content_str_to_input_text(str(item.get("text", item)), role)
|
||||
result.append(converted)
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: unknown({original_type}) -> {converted}"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: unknown({original_type}) -> {converted}")
|
||||
verbose_logger.debug(f"Chat provider: Final converted content: {result}")
|
||||
return result
|
||||
else:
|
||||
|
|
@ -694,17 +651,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
verbose_logger.debug(f"Chat provider: Other content type -> {result}")
|
||||
return result
|
||||
|
||||
def _convert_tools_to_responses_format(
|
||||
self, tools: List[Dict[str, Any]]
|
||||
) -> List["ALL_RESPONSES_API_TOOL_PARAMS"]:
|
||||
def _convert_tools_to_responses_format(self, tools: List[Dict[str, Any]]) -> List["ALL_RESPONSES_API_TOOL_PARAMS"]:
|
||||
"""Convert chat completion tools to responses API tools format"""
|
||||
responses_tools: List["ALL_RESPONSES_API_TOOL_PARAMS"] = []
|
||||
for tool in tools:
|
||||
# convert function tool from chat completion to responses API format
|
||||
if tool.get("type") == "function":
|
||||
function_tool = cast(
|
||||
ChatCompletionToolParamFunctionChunk, tool.get("function")
|
||||
)
|
||||
function_tool = cast(ChatCompletionToolParamFunctionChunk, tool.get("function"))
|
||||
responses_tools.append(
|
||||
FunctionToolParam(
|
||||
name=function_tool["name"],
|
||||
|
|
@ -730,9 +683,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if not extra_body:
|
||||
return optional_params
|
||||
|
||||
supported_responses_api_params = set(
|
||||
ResponsesAPIOptionalRequestParams.__annotations__.keys()
|
||||
)
|
||||
supported_responses_api_params = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
|
||||
# Also include params we handle specially
|
||||
supported_responses_api_params.update(
|
||||
{
|
||||
|
|
@ -750,9 +701,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
return optional_params
|
||||
|
||||
def _map_reasoning_effort(
|
||||
self, reasoning_effort: Union[str, Dict[str, Any]]
|
||||
) -> Optional[Reasoning]:
|
||||
def _map_reasoning_effort(self, reasoning_effort: Union[str, Dict[str, Any]]) -> Optional[Reasoning]:
|
||||
# If dict is passed, convert it directly to Reasoning object
|
||||
if isinstance(reasoning_effort, dict):
|
||||
return Reasoning(**reasoning_effort) # type: ignore[typeddict-item]
|
||||
|
|
@ -760,8 +709,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
# 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"
|
||||
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
|
||||
|
|
@ -772,11 +720,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
elif reasoning_effort == "xhigh":
|
||||
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") if auto_summary_enabled else Reasoning(effort="medium")
|
||||
return (
|
||||
Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium")
|
||||
)
|
||||
elif reasoning_effort == "low":
|
||||
return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low")
|
||||
elif reasoning_effort == "minimal":
|
||||
return Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal")
|
||||
return (
|
||||
Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal")
|
||||
)
|
||||
return None
|
||||
|
||||
def _add_web_search_tool(
|
||||
|
|
@ -855,7 +807,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
return {"format": {"type": "text"}}
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _convert_annotations_to_chat_format(
|
||||
annotations: Optional[List[Any]],
|
||||
|
|
@ -908,9 +860,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
|
||||
class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
||||
def __init__(
|
||||
self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False
|
||||
):
|
||||
def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False):
|
||||
super().__init__(streaming_response, sync_stream, json_mode)
|
||||
|
||||
def _handle_string_chunk(
|
||||
|
|
@ -923,9 +873,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
|
||||
if not str_line or str_line.startswith("event:"):
|
||||
# ignore.
|
||||
return GenericStreamingChunk(
|
||||
text="", tool_use=None, is_finished=False, finish_reason="", usage=None
|
||||
)
|
||||
return GenericStreamingChunk(text="", tool_use=None, is_finished=False, finish_reason="", usage=None)
|
||||
index = str_line.find("data:")
|
||||
if index != -1:
|
||||
str_line = str_line[index + 5 :]
|
||||
|
|
@ -988,13 +936,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
if output_item.get("type") == "function_call":
|
||||
# Extract provider_specific_fields if present
|
||||
provider_specific_fields = output_item.get("provider_specific_fields")
|
||||
if provider_specific_fields and not isinstance(
|
||||
provider_specific_fields, dict
|
||||
):
|
||||
if provider_specific_fields and not isinstance(provider_specific_fields, dict):
|
||||
provider_specific_fields = (
|
||||
dict(provider_specific_fields)
|
||||
if hasattr(provider_specific_fields, "__dict__")
|
||||
else {}
|
||||
dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {}
|
||||
)
|
||||
|
||||
function_chunk = ChatCompletionToolCallFunctionChunk(
|
||||
|
|
@ -1003,9 +947,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
)
|
||||
|
||||
if provider_specific_fields:
|
||||
function_chunk["provider_specific_fields"] = (
|
||||
provider_specific_fields
|
||||
)
|
||||
function_chunk["provider_specific_fields"] = provider_specific_fields
|
||||
|
||||
tool_call_chunk = ChatCompletionToolCallChunk(
|
||||
id=output_item.get("call_id"),
|
||||
|
|
@ -1040,9 +982,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
id=None,
|
||||
index=0,
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=None, arguments=content_part
|
||||
),
|
||||
function=ChatCompletionToolCallFunctionChunk(name=None, arguments=content_part),
|
||||
)
|
||||
]
|
||||
),
|
||||
|
|
@ -1051,22 +991,16 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
]
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Chat provider: Invalid function argument delta {parsed_chunk}"
|
||||
)
|
||||
raise ValueError(f"Chat provider: Invalid function argument delta {parsed_chunk}")
|
||||
elif event_type == "response.output_item.done":
|
||||
# New output item added
|
||||
output_item = parsed_chunk.get("item", {})
|
||||
if output_item.get("type") == "function_call":
|
||||
# Extract provider_specific_fields if present
|
||||
provider_specific_fields = output_item.get("provider_specific_fields")
|
||||
if provider_specific_fields and not isinstance(
|
||||
provider_specific_fields, dict
|
||||
):
|
||||
if provider_specific_fields and not isinstance(provider_specific_fields, dict):
|
||||
provider_specific_fields = (
|
||||
dict(provider_specific_fields)
|
||||
if hasattr(provider_specific_fields, "__dict__")
|
||||
else {}
|
||||
dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {}
|
||||
)
|
||||
|
||||
function_chunk = ChatCompletionToolCallFunctionChunk(
|
||||
|
|
@ -1076,9 +1010,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
|
||||
# Add provider_specific_fields to function if present
|
||||
if provider_specific_fields:
|
||||
function_chunk["provider_specific_fields"] = (
|
||||
provider_specific_fields
|
||||
)
|
||||
function_chunk["provider_specific_fields"] = provider_specific_fields
|
||||
|
||||
tool_call_chunk = ChatCompletionToolCallChunk(
|
||||
id=output_item.get("call_id"),
|
||||
|
|
@ -1142,21 +1074,31 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
elif event_type == "response.completed":
|
||||
# Response is fully complete - now we can signal is_finished=True
|
||||
# This ensures we don't prematurely end the stream before tool_calls arrive
|
||||
|
||||
# Check if response contains function_call items in output
|
||||
# to determine correct finish_reason
|
||||
response_data = parsed_chunk.get("response", {})
|
||||
output_items = response_data.get("output", []) if response_data else []
|
||||
|
||||
has_function_calls = any(
|
||||
item.get("type") == "function_call" for item in output_items if isinstance(item, dict)
|
||||
)
|
||||
|
||||
finish_reason = "tool_calls" if has_function_calls else "stop"
|
||||
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=""),
|
||||
finish_reason="stop",
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
pass
|
||||
# For any unhandled event types, create a minimal valid chunk or skip
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: Unhandled event type '{event_type}', creating empty chunk"
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: Unhandled event type '{event_type}', creating empty chunk")
|
||||
|
||||
# Return a minimal valid chunk for unknown events
|
||||
return ModelResponseStream(
|
||||
|
|
@ -1179,9 +1121,5 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
Returns:
|
||||
ModelResponseStream: OpenAI-formatted streaming chunk
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"Chat provider: transform_streaming_response called with chunk: {chunk}"
|
||||
)
|
||||
return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
|
||||
chunk
|
||||
)
|
||||
verbose_logger.debug(f"Chat provider: transform_streaming_response called with chunk: {chunk}")
|
||||
return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
|
||||
|
|
|
|||
|
|
@ -287,7 +287,9 @@ MIN_NON_ZERO_TEMPERATURE = float(os.getenv("MIN_NON_ZERO_TEMPERATURE", 0.0001))
|
|||
REPEATED_STREAMING_CHUNK_LIMIT = int(
|
||||
os.getenv("REPEATED_STREAMING_CHUNK_LIMIT", 100)
|
||||
) # catch if model starts looping the same chunk while streaming. Uses high default to prevent false positives.
|
||||
DEFAULT_MAX_LRU_CACHE_SIZE = int(os.getenv("DEFAULT_MAX_LRU_CACHE_SIZE", 16))
|
||||
# Shared maxsize for functools.lru_cache usage across hot paths.
|
||||
# Defaulted to 64 to avoid cache thrash in multi-model production workloads.
|
||||
DEFAULT_MAX_LRU_CACHE_SIZE = int(os.getenv("DEFAULT_MAX_LRU_CACHE_SIZE", 64))
|
||||
_REALTIME_BODY_CACHE_SIZE = 1000 # Keep realtime helper caches bounded; workloads rarely exceed 1k models/intents
|
||||
INITIAL_RETRY_DELAY = float(os.getenv("INITIAL_RETRY_DELAY", 0.5))
|
||||
MAX_RETRY_DELAY = float(os.getenv("MAX_RETRY_DELAY", 8.0))
|
||||
|
|
@ -576,6 +578,11 @@ OPENAI_CHAT_COMPLETION_PARAMS = [
|
|||
"thinking",
|
||||
"web_search_options",
|
||||
"service_tier",
|
||||
"store",
|
||||
"prompt_cache_key",
|
||||
"prompt_cache_retention",
|
||||
"safety_identifier",
|
||||
"verbosity",
|
||||
]
|
||||
|
||||
OPENAI_TRANSCRIPTION_PARAMS = [
|
||||
|
|
|
|||
|
|
@ -448,7 +448,9 @@ def cost_per_token( # noqa: PLR0915
|
|||
elif custom_llm_provider == "anthropic":
|
||||
return anthropic_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "bedrock":
|
||||
return bedrock_cost_per_token(model=model, usage=usage_block)
|
||||
return bedrock_cost_per_token(
|
||||
model=model, usage=usage_block, service_tier=service_tier
|
||||
)
|
||||
elif custom_llm_provider == "openai":
|
||||
return openai_cost_per_token(
|
||||
model=model, usage=usage_block, service_tier=service_tier
|
||||
|
|
@ -2146,4 +2148,3 @@ def handle_realtime_stream_cost_calculation(
|
|||
|
||||
return total_cost
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -93,7 +93,9 @@ class DatadogCostManagementLogger(CustomBatchLogger):
|
|||
Aggregates costs by Provider, Model, and Date.
|
||||
Returns a list of DatadogFOCUSCostEntry.
|
||||
"""
|
||||
aggregator: Dict[Tuple[str, str, str, Tuple[Tuple[str, str], ...]], DatadogFOCUSCostEntry] = {}
|
||||
aggregator: Dict[
|
||||
Tuple[str, str, str, Tuple[Tuple[str, str], ...]], DatadogFOCUSCostEntry
|
||||
] = {}
|
||||
|
||||
for log in logs:
|
||||
try:
|
||||
|
|
@ -167,10 +169,20 @@ class DatadogCostManagementLogger(CustomBatchLogger):
|
|||
metadata = log.get("metadata", {})
|
||||
if metadata:
|
||||
# Add user info
|
||||
if "user_api_key_alias" in metadata:
|
||||
# Add user info
|
||||
if metadata.get("user_api_key_alias"):
|
||||
tags["user"] = str(metadata["user_api_key_alias"])
|
||||
if "user_api_key_team_alias" in metadata:
|
||||
tags["team"] = str(metadata["user_api_key_team_alias"])
|
||||
|
||||
# Add Team Tag
|
||||
team_tag = (
|
||||
metadata.get("user_api_key_team_alias")
|
||||
or metadata.get("team_alias") # type: ignore
|
||||
or metadata.get("user_api_key_team_id")
|
||||
or metadata.get("team_id") # type: ignore
|
||||
)
|
||||
|
||||
if team_tag:
|
||||
tags["team"] = str(team_tag)
|
||||
# model_group is not in StandardLoggingMetadata TypedDict, so we need to access it via dict.get()
|
||||
model_group = metadata.get("model_group") # type: ignore[misc]
|
||||
if model_group:
|
||||
|
|
|
|||
|
|
@ -55,4 +55,15 @@ def get_datadog_tags(
|
|||
request_tags = standard_logging_object.get("request_tags", []) or []
|
||||
tags.extend(f"request_tag:{tag}" for tag in request_tags)
|
||||
|
||||
# Add Team Tag
|
||||
metadata = standard_logging_object.get("metadata", {}) or {}
|
||||
team_tag = (
|
||||
metadata.get("user_api_key_team_alias")
|
||||
or metadata.get("team_alias")
|
||||
or metadata.get("user_api_key_team_id")
|
||||
or metadata.get("team_id")
|
||||
)
|
||||
if team_tag:
|
||||
tags.append(f"team:{team_tag}")
|
||||
|
||||
return ",".join(tags)
|
||||
|
|
|
|||
|
|
@ -22,6 +22,10 @@ from typing import (
|
|||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_litellm_metadata_from_kwargs,
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_DeletedVerificationToken,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
@ -1055,16 +1059,16 @@ class PrometheusLogger(CustomLogger):
|
|||
enum_values=enum_values,
|
||||
)
|
||||
|
||||
if (
|
||||
standard_logging_payload["stream"] is True
|
||||
): # log successful streaming requests from logging event hook.
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
|
||||
# increment litellm_proxy_total_requests_metric for all successful requests
|
||||
# (both streaming and non-streaming) in this single location to prevent
|
||||
# double-counting that occurs when async_post_call_success_hook also increments
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
|
||||
|
||||
def _increment_token_metrics(
|
||||
self,
|
||||
|
|
@ -1086,13 +1090,6 @@ class PrometheusLogger(CustomLogger):
|
|||
):
|
||||
_tags = standard_logging_payload["request_tags"]
|
||||
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_total_tokens_metric"
|
||||
|
|
@ -1655,49 +1652,12 @@ class PrometheusLogger(CustomLogger):
|
|||
):
|
||||
"""
|
||||
Proxy level tracking - triggered when the proxy responds with a success response to the client
|
||||
|
||||
Note: litellm_proxy_total_requests_metric is NOT incremented here to avoid
|
||||
double-counting. It is incremented in async_log_success_event which fires
|
||||
for all successful requests (both streaming and non-streaming).
|
||||
"""
|
||||
try:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
StandardLoggingPayloadSetup,
|
||||
)
|
||||
|
||||
if self._should_skip_metrics_for_invalid_key(
|
||||
user_api_key_dict=user_api_key_dict
|
||||
):
|
||||
return
|
||||
|
||||
_metadata = data.get("metadata", {}) or {}
|
||||
enum_values = UserAPIKeyLabelValues(
|
||||
end_user=user_api_key_dict.end_user_id,
|
||||
hashed_api_key=user_api_key_dict.api_key,
|
||||
api_key_alias=user_api_key_dict.key_alias,
|
||||
requested_model=data.get("model", ""),
|
||||
team=user_api_key_dict.team_id,
|
||||
team_alias=user_api_key_dict.team_alias,
|
||||
user=user_api_key_dict.user_id,
|
||||
user_email=user_api_key_dict.user_email,
|
||||
status_code="200",
|
||||
route=user_api_key_dict.request_route,
|
||||
tags=StandardLoggingPayloadSetup._get_request_tags(
|
||||
litellm_params=data,
|
||||
proxy_server_request=data.get("proxy_server_request", {}),
|
||||
),
|
||||
client_ip=_metadata.get("requester_ip_address"),
|
||||
user_agent=_metadata.get("user_agent"),
|
||||
)
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"prometheus Layer Error(): Exception occured - {}".format(str(e))
|
||||
)
|
||||
pass
|
||||
pass
|
||||
|
||||
def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
|
||||
"""Get value from dict or Pydantic model."""
|
||||
|
|
@ -2004,7 +1964,7 @@ class PrometheusLogger(CustomLogger):
|
|||
|
||||
api_base = standard_logging_payload["api_base"]
|
||||
_litellm_params = request_kwargs.get("litellm_params", {}) or {}
|
||||
_metadata = _litellm_params.get("metadata", {})
|
||||
_metadata = get_litellm_metadata_from_kwargs(request_kwargs)
|
||||
litellm_model_name = request_kwargs.get("model", None)
|
||||
llm_provider = _litellm_params.get("custom_llm_provider", None)
|
||||
_model_info = _metadata.get("model_info") or {}
|
||||
|
|
@ -2220,7 +2180,8 @@ class PrometheusLogger(CustomLogger):
|
|||
original_model_group,
|
||||
kwargs,
|
||||
)
|
||||
_metadata = kwargs.get("metadata", {})
|
||||
_metadata_key = get_metadata_variable_name_from_kwargs(kwargs)
|
||||
_metadata = kwargs.get(_metadata_key) or {}
|
||||
standard_metadata: StandardLoggingMetadata = (
|
||||
StandardLoggingPayloadSetup.get_standard_logging_metadata(
|
||||
metadata=_metadata
|
||||
|
|
@ -2265,7 +2226,8 @@ class PrometheusLogger(CustomLogger):
|
|||
kwargs,
|
||||
)
|
||||
_new_model = kwargs.get("model")
|
||||
_metadata = kwargs.get("metadata", {})
|
||||
_metadata_key = get_metadata_variable_name_from_kwargs(kwargs)
|
||||
_metadata = kwargs.get(_metadata_key) or {}
|
||||
_tags = cast(List[str], kwargs.get("tags") or [])
|
||||
standard_metadata: StandardLoggingMetadata = (
|
||||
StandardLoggingPayloadSetup.get_standard_logging_metadata(
|
||||
|
|
|
|||
|
|
@ -1335,7 +1335,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
|
||||
# Store additional costs if provided (free-form dict for extensibility)
|
||||
if additional_costs and isinstance(additional_costs, dict) and len(additional_costs) > 0:
|
||||
if (
|
||||
additional_costs
|
||||
and isinstance(additional_costs, dict)
|
||||
and len(additional_costs) > 0
|
||||
):
|
||||
self.cost_breakdown["additional_costs"] = additional_costs
|
||||
|
||||
# Store discount information if provided
|
||||
|
|
@ -4519,13 +4523,19 @@ class StandardLoggingPayloadSetup:
|
|||
requester_custom_headers=None,
|
||||
cold_storage_object_key=None,
|
||||
user_api_key_auth_metadata=None,
|
||||
team_alias=None,
|
||||
team_id=None,
|
||||
)
|
||||
if isinstance(metadata, dict):
|
||||
for key in metadata.keys() & _STANDARD_LOGGING_METADATA_KEYS:
|
||||
clean_metadata[key] = metadata[key] # type: ignore
|
||||
|
||||
user_api_key = metadata.get("user_api_key")
|
||||
if user_api_key and isinstance(user_api_key, str) and is_valid_sha256_hash(user_api_key):
|
||||
if (
|
||||
user_api_key
|
||||
and isinstance(user_api_key, str)
|
||||
and is_valid_sha256_hash(user_api_key)
|
||||
):
|
||||
clean_metadata["user_api_key_hash"] = user_api_key
|
||||
_potential_requester_metadata = metadata.get(
|
||||
"metadata", None
|
||||
|
|
@ -5279,6 +5289,8 @@ def get_standard_logging_metadata(
|
|||
user_api_key_request_route=None,
|
||||
cold_storage_object_key=None,
|
||||
user_api_key_auth_metadata=None,
|
||||
team_alias=None,
|
||||
team_id=None,
|
||||
)
|
||||
if isinstance(metadata, dict):
|
||||
# Update the clean_metadata with values from input metadata that match StandardLoggingMetadata fields
|
||||
|
|
|
|||
|
|
@ -546,7 +546,11 @@ def convert_to_model_response_object( # noqa: PLR0915
|
|||
message = litellm.Message(content=json_mode_content_str)
|
||||
finish_reason = "stop"
|
||||
if message is None:
|
||||
provider_specific_fields = {}
|
||||
# Preserve provider_specific_fields if already present
|
||||
# in the response (e.g. from proxy passthrough)
|
||||
provider_specific_fields = dict(
|
||||
choice["message"].get("provider_specific_fields", None) or {}
|
||||
)
|
||||
message_keys = Message.model_fields.keys()
|
||||
for field in choice["message"].keys():
|
||||
if field not in message_keys:
|
||||
|
|
|
|||
|
|
@ -2018,6 +2018,235 @@ def anthropic_process_openai_file_message(
|
|||
)
|
||||
|
||||
|
||||
def _sanitize_empty_text_content(
|
||||
message: AllMessageValues,
|
||||
) -> AllMessageValues:
|
||||
"""
|
||||
Case C: Sanitize empty text content
|
||||
- Replace empty or whitespace-only text content with a placeholder message.
|
||||
|
||||
Returns:
|
||||
The message with sanitized content if needed, otherwise the original message
|
||||
"""
|
||||
if message.get("role") in ["user", "assistant"]:
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
if not content or not content.strip():
|
||||
message = cast(AllMessageValues, dict(message)) # Make a copy
|
||||
message["content"] = "[System: Empty message content sanitised to satisfy protocol]"
|
||||
verbose_logger.debug(
|
||||
f"_sanitize_empty_text_content: Replaced empty text content in {message.get('role')} message"
|
||||
)
|
||||
return message
|
||||
|
||||
|
||||
def _add_missing_tool_results( # noqa: PLR0915
|
||||
current_message: AllMessageValues,
|
||||
messages: List[AllMessageValues],
|
||||
current_index: int,
|
||||
) -> Tuple[List[AllMessageValues], int]:
|
||||
"""
|
||||
Case A: Missing tool_result for tool_use (orphaned tool calls)
|
||||
- If an assistant message has tool_calls but no corresponding tool result follows,
|
||||
add a dummy tool result message indicating the user did not provide the result.
|
||||
|
||||
Returns:
|
||||
A tuple of:
|
||||
- List containing the assistant message, followed by existing tool results,
|
||||
followed by any dummy tool results needed
|
||||
- Number of original messages consumed (to adjust iteration index)
|
||||
"""
|
||||
result_messages: List[AllMessageValues] = []
|
||||
tool_calls = current_message.get("tool_calls")
|
||||
|
||||
if not tool_calls or len(cast(list, tool_calls)) == 0:
|
||||
return ([current_message], 0)
|
||||
|
||||
# Collect all tool_call_ids from this assistant message
|
||||
expected_tool_call_ids = set()
|
||||
for tool_call in cast(list, tool_calls):
|
||||
tool_call_id = None
|
||||
if isinstance(tool_call, dict):
|
||||
tool_call_id = tool_call.get("id")
|
||||
else:
|
||||
tool_call_id = getattr(tool_call, "id", None)
|
||||
if tool_call_id:
|
||||
expected_tool_call_ids.add(tool_call_id)
|
||||
|
||||
# Collect actual tool result messages that follow this assistant message
|
||||
found_tool_call_ids = set()
|
||||
actual_tool_results: List[AllMessageValues] = []
|
||||
j = current_index + 1
|
||||
|
||||
while j < len(messages):
|
||||
next_msg = messages[j]
|
||||
next_role = next_msg.get("role")
|
||||
|
||||
if next_role == "assistant":
|
||||
break
|
||||
|
||||
if next_role in ["tool", "function"]:
|
||||
tool_call_id = next_msg.get("tool_call_id")
|
||||
if tool_call_id and tool_call_id in expected_tool_call_ids:
|
||||
found_tool_call_ids.add(tool_call_id)
|
||||
actual_tool_results.append(next_msg)
|
||||
|
||||
j += 1
|
||||
|
||||
# Find missing tool results
|
||||
missing_tool_call_ids = expected_tool_call_ids - found_tool_call_ids
|
||||
|
||||
if missing_tool_call_ids:
|
||||
verbose_logger.debug(
|
||||
f"_add_missing_tool_results: Found {len(missing_tool_call_ids)} orphaned tool calls. Adding dummy tool results."
|
||||
)
|
||||
|
||||
result_messages.append(current_message)
|
||||
|
||||
# Add existing tool results FIRST
|
||||
result_messages.extend(actual_tool_results)
|
||||
|
||||
# Then add dummy tool results for missing ones
|
||||
for tool_call_id in missing_tool_call_ids:
|
||||
tool_name = "unknown_tool"
|
||||
for tool_call in cast(list, tool_calls):
|
||||
tc_id = None
|
||||
if isinstance(tool_call, dict):
|
||||
tc_id = tool_call.get("id")
|
||||
else:
|
||||
tc_id = getattr(tool_call, "id", None)
|
||||
|
||||
if tc_id == tool_call_id:
|
||||
if isinstance(tool_call, dict):
|
||||
function = tool_call.get("function", {})
|
||||
if isinstance(function, dict):
|
||||
tool_name = function.get("name", "unknown_tool")
|
||||
else:
|
||||
tool_name = getattr(function, "name", "unknown_tool")
|
||||
else:
|
||||
function = getattr(tool_call, "function", None)
|
||||
if function:
|
||||
tool_name = getattr(function, "name", "unknown_tool")
|
||||
break
|
||||
|
||||
dummy_tool_result: ChatCompletionToolMessage = {
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_call_id,
|
||||
"content": f"[System: Tool execution skipped/interrupted by user. No result provided for tool '{tool_name}'.]",
|
||||
}
|
||||
result_messages.append(dummy_tool_result)
|
||||
|
||||
# Return the messages and the number of original messages to skip
|
||||
return (result_messages, len(actual_tool_results))
|
||||
|
||||
return ([current_message], 0)
|
||||
|
||||
|
||||
def _is_orphaned_tool_result(
|
||||
current_message: AllMessageValues,
|
||||
sanitized_messages: List[AllMessageValues],
|
||||
) -> bool:
|
||||
"""
|
||||
Case B: Orphaned tool_result (unexpected result)
|
||||
- Check if a tool message references a tool_call_id that doesn't exist in the previous
|
||||
assistant message.
|
||||
|
||||
Returns:
|
||||
True if this is an orphaned tool result that should be removed, False otherwise
|
||||
"""
|
||||
if current_message.get("role") not in ["tool", "function"]:
|
||||
return False
|
||||
|
||||
tool_call_id = current_message.get("tool_call_id")
|
||||
|
||||
if not tool_call_id:
|
||||
return False
|
||||
|
||||
# Look back to find the most recent assistant message with tool_calls
|
||||
found_matching_tool_call = False
|
||||
|
||||
for j in range(len(sanitized_messages) - 1, -1, -1):
|
||||
prev_msg = sanitized_messages[j]
|
||||
if prev_msg.get("role") == "assistant":
|
||||
tool_calls = prev_msg.get("tool_calls")
|
||||
if tool_calls:
|
||||
for tool_call in cast(list, tool_calls):
|
||||
tc_id = None
|
||||
if isinstance(tool_call, dict):
|
||||
tc_id = tool_call.get("id")
|
||||
else:
|
||||
tc_id = getattr(tool_call, "id", None)
|
||||
|
||||
if tc_id == tool_call_id:
|
||||
found_matching_tool_call = True
|
||||
break
|
||||
|
||||
break
|
||||
|
||||
if not found_matching_tool_call:
|
||||
verbose_logger.debug(
|
||||
"_is_orphaned_tool_result: Found orphaned tool result with redacted tool_call_id"
|
||||
)
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def sanitize_messages_for_tool_calling(
|
||||
messages: List[AllMessageValues],
|
||||
) -> List[AllMessageValues]:
|
||||
"""
|
||||
Sanitize messages for tool calling to handle common issues when modify_params=True:
|
||||
|
||||
Case A: Missing tool_result for tool_use (orphaned tool calls)
|
||||
- If an assistant message has tool_calls but no corresponding tool result follows,
|
||||
add a dummy tool result message indicating the user did not provide the result.
|
||||
|
||||
Case B: Orphaned tool_result (unexpected result)
|
||||
- If a tool message references a tool_call_id that doesn't exist in the previous
|
||||
assistant message, remove that tool message.
|
||||
|
||||
Case C: Empty text content
|
||||
- Replace empty or whitespace-only text content with a placeholder message.
|
||||
|
||||
This function operates on OpenAI format messages before they are converted to
|
||||
provider-specific formats.
|
||||
"""
|
||||
if not litellm.modify_params:
|
||||
return messages
|
||||
|
||||
sanitized_messages: List[AllMessageValues] = []
|
||||
i = 0
|
||||
|
||||
while i < len(messages):
|
||||
current_message = messages[i]
|
||||
|
||||
# Case C: Sanitize empty text content
|
||||
current_message = _sanitize_empty_text_content(current_message)
|
||||
|
||||
# Case A: Check if assistant message has tool_calls without following tool results
|
||||
if current_message.get("role") == "assistant":
|
||||
result_messages, messages_consumed = _add_missing_tool_results(current_message, messages, i)
|
||||
|
||||
# If dummy tool results were added, extend sanitized_messages and skip consumed messages
|
||||
if len(result_messages) > 1:
|
||||
sanitized_messages.extend(result_messages)
|
||||
# Skip the assistant message and any actual tool results that were included
|
||||
i += 1 + messages_consumed
|
||||
continue
|
||||
|
||||
# Case B: Check for orphaned tool results
|
||||
if _is_orphaned_tool_result(current_message, sanitized_messages):
|
||||
i += 1
|
||||
continue # Skip this orphaned tool result
|
||||
|
||||
# Add the message to sanitized list
|
||||
sanitized_messages.append(current_message)
|
||||
i += 1
|
||||
|
||||
return sanitized_messages
|
||||
|
||||
|
||||
def anthropic_messages_pt( # noqa: PLR0915
|
||||
messages: List[AllMessageValues],
|
||||
model: str,
|
||||
|
|
@ -2037,6 +2266,9 @@ def anthropic_messages_pt( # noqa: PLR0915
|
|||
5. System messages are a separate param to the Messages API
|
||||
6. Ensure we only accept role, content. (message.name is not supported)
|
||||
"""
|
||||
# Sanitize messages for tool calling issues when modify_params=True
|
||||
messages = sanitize_messages_for_tool_calling(messages)
|
||||
|
||||
# add role=tool support to allow function call result/error submission
|
||||
user_message_types = {"user", "tool", "function"}
|
||||
# reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, merge them.
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import collections.abc
|
||||
import datetime
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
import traceback
|
||||
|
|
@ -435,7 +436,7 @@ class CustomStreamWrapper:
|
|||
|
||||
def handle_openai_chat_completion_chunk(self, chunk):
|
||||
try:
|
||||
print_verbose(f"\nRaw OpenAI Chunk\n{chunk}\n")
|
||||
|
||||
str_line = chunk
|
||||
text = ""
|
||||
is_finished = False
|
||||
|
|
@ -485,7 +486,7 @@ class CustomStreamWrapper:
|
|||
|
||||
def handle_azure_text_completion_chunk(self, chunk):
|
||||
try:
|
||||
print_verbose(f"\nRaw OpenAI Chunk\n{chunk}\n")
|
||||
|
||||
text = ""
|
||||
is_finished = False
|
||||
finish_reason = None
|
||||
|
|
@ -506,7 +507,7 @@ class CustomStreamWrapper:
|
|||
|
||||
def handle_openai_text_completion_chunk(self, chunk):
|
||||
try:
|
||||
print_verbose(f"\nRaw OpenAI Chunk\n{chunk}\n")
|
||||
|
||||
text = ""
|
||||
is_finished = False
|
||||
finish_reason = None
|
||||
|
|
@ -870,9 +871,6 @@ class CustomStreamWrapper:
|
|||
preserve_upstream_non_openai_attributes,
|
||||
)
|
||||
|
||||
print_verbose(
|
||||
f"completion_obj: {completion_obj}, model_response.choices[0]: {model_response.choices[0]}, response_obj: {response_obj}"
|
||||
)
|
||||
is_chunk_non_empty = self.is_chunk_non_empty(
|
||||
completion_obj, model_response, response_obj
|
||||
)
|
||||
|
|
@ -899,11 +897,9 @@ class CustomStreamWrapper:
|
|||
choice_json.pop(
|
||||
"finish_reason", None
|
||||
) # for mistral etc. which return a value in their last chunk (not-openai compatible).
|
||||
print_verbose(f"choice_json: {choice_json}")
|
||||
choices.append(StreamingChoices(**choice_json))
|
||||
except Exception:
|
||||
choices.append(StreamingChoices())
|
||||
print_verbose(f"choices in streaming: {choices}")
|
||||
setattr(model_response, "choices", choices)
|
||||
else:
|
||||
return
|
||||
|
|
@ -921,9 +917,11 @@ class CustomStreamWrapper:
|
|||
)
|
||||
|
||||
model_response = self.strip_role_from_delta(model_response)
|
||||
verbose_logger.debug(
|
||||
f"model_response.choices[0].delta inside is_chunk_non_empty: {model_response.choices[0].delta}"
|
||||
)
|
||||
if verbose_logger.isEnabledFor(logging.DEBUG):
|
||||
verbose_logger.debug(
|
||||
"model_response.choices[0].delta: %s",
|
||||
model_response.choices[0].delta,
|
||||
)
|
||||
else:
|
||||
## else
|
||||
completion_obj["content"] = model_response_str
|
||||
|
|
@ -1370,9 +1368,6 @@ class CustomStreamWrapper:
|
|||
)
|
||||
|
||||
model_response.model = self.model
|
||||
print_verbose(
|
||||
f"model_response finish reason 3: {self.received_finish_reason}; response_obj={response_obj}"
|
||||
)
|
||||
## FUNCTION CALL PARSING
|
||||
original_chunk = (
|
||||
response_obj.get("original_chunk") if response_obj is not None else None
|
||||
|
|
@ -1432,7 +1427,6 @@ class CustomStreamWrapper:
|
|||
):
|
||||
t.function.arguments = ""
|
||||
_json_delta = delta.model_dump()
|
||||
print_verbose(f"_json_delta: {_json_delta}")
|
||||
if "role" not in _json_delta or _json_delta["role"] is None:
|
||||
_json_delta[
|
||||
"role"
|
||||
|
|
@ -1466,11 +1460,7 @@ class CustomStreamWrapper:
|
|||
if original_chunk.choices[0].delta is None
|
||||
else dict(original_chunk.choices[0].delta)
|
||||
)
|
||||
print_verbose(f"original delta: {delta}")
|
||||
model_response.choices[0].delta = Delta(**delta)
|
||||
print_verbose(
|
||||
f"new delta: {model_response.choices[0].delta}"
|
||||
)
|
||||
except Exception:
|
||||
model_response.choices[0].delta = Delta()
|
||||
else:
|
||||
|
|
@ -1480,11 +1470,6 @@ class CustomStreamWrapper:
|
|||
):
|
||||
return model_response
|
||||
return
|
||||
print_verbose(
|
||||
f"model_response.choices[0].delta: {model_response.choices[0].delta}; completion_obj: {completion_obj}"
|
||||
)
|
||||
print_verbose(f"self.sent_first_chunk: {self.sent_first_chunk}")
|
||||
|
||||
## CHECK FOR TOOL USE
|
||||
|
||||
if "tool_calls" in completion_obj and len(completion_obj["tool_calls"]) > 0:
|
||||
|
|
@ -1915,18 +1900,9 @@ class CustomStreamWrapper:
|
|||
and len(chunk.parts) == 0
|
||||
):
|
||||
continue
|
||||
# chunk_creator() does logging/stream chunk building. We need to let it know its being called in_async_func, so we don't double add chunks.
|
||||
# __anext__ also calls async_success_handler, which does logging
|
||||
verbose_logger.debug(
|
||||
f"PROCESSED ASYNC CHUNK PRE CHUNK CREATOR: {chunk}"
|
||||
)
|
||||
|
||||
processed_chunk: Optional[ModelResponseStream] = self.chunk_creator(
|
||||
chunk=chunk
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"PROCESSED ASYNC CHUNK POST CHUNK CREATOR: {processed_chunk}"
|
||||
)
|
||||
if processed_chunk is None:
|
||||
continue
|
||||
|
||||
|
|
@ -1943,31 +1919,33 @@ class CustomStreamWrapper:
|
|||
self.rules.post_call_rules(
|
||||
input=self.response_uptil_now, model=self.model
|
||||
)
|
||||
self.chunks.append(processed_chunk)
|
||||
|
||||
# Store a shallow copy so usage stripping below
|
||||
# does not mutate the stored chunk.
|
||||
self.chunks.append(processed_chunk.model_copy())
|
||||
|
||||
# Add mcp_list_tools to first chunk if present
|
||||
if not self.sent_first_chunk:
|
||||
processed_chunk = self._add_mcp_list_tools_to_first_chunk(processed_chunk)
|
||||
self.sent_first_chunk = True
|
||||
if hasattr(
|
||||
processed_chunk, "usage"
|
||||
): # remove usage from chunk, only send on final chunk
|
||||
# Convert the object to a dictionary
|
||||
if (
|
||||
hasattr(processed_chunk, "usage")
|
||||
and getattr(processed_chunk, "usage", None) is not None
|
||||
):
|
||||
# Strip usage from the outgoing chunk so it's not sent twice
|
||||
# (once in the chunk, once in _hidden_params).
|
||||
# Create a new object without usage, matching sync behavior.
|
||||
# The copy in self.chunks retains usage for calculate_total_usage().
|
||||
obj_dict = processed_chunk.model_dump()
|
||||
|
||||
# Remove an attribute (e.g., 'attr2')
|
||||
if "usage" in obj_dict:
|
||||
del obj_dict["usage"]
|
||||
|
||||
# Create a new object without the removed attribute
|
||||
processed_chunk = self.model_response_creator(chunk=obj_dict)
|
||||
processed_chunk = self.model_response_creator(
|
||||
chunk=obj_dict, hidden_params=processed_chunk._hidden_params
|
||||
)
|
||||
is_empty = is_model_response_stream_empty(
|
||||
model_response=cast(ModelResponseStream, processed_chunk)
|
||||
)
|
||||
|
||||
if is_empty:
|
||||
continue
|
||||
print_verbose(f"final returned processed chunk: {processed_chunk}")
|
||||
|
||||
# add usage as hidden param
|
||||
if self.sent_last_chunk is True and self.stream_options is None:
|
||||
|
|
@ -1982,7 +1960,7 @@ class CustomStreamWrapper:
|
|||
)
|
||||
)
|
||||
# Add MCP metadata to final chunk if present (after hooks)
|
||||
processed_chunk = self._add_mcp_metadata_to_final_chunk(processed_chunk)
|
||||
processed_chunk = self._add_mcp_metadata_to_final_chunk(processed_chunk) # type: ignore[reportArgumentType]
|
||||
|
||||
return processed_chunk
|
||||
raise StopAsyncIteration
|
||||
|
|
@ -1996,13 +1974,9 @@ class CustomStreamWrapper:
|
|||
else:
|
||||
chunk = next(self.completion_stream)
|
||||
if chunk is not None and chunk != b"":
|
||||
print_verbose(f"PROCESSED CHUNK PRE CHUNK CREATOR: {chunk}")
|
||||
processed_chunk: Optional[
|
||||
ModelResponseStream
|
||||
] = self.chunk_creator(chunk=chunk)
|
||||
print_verbose(
|
||||
f"PROCESSED CHUNK POST CHUNK CREATOR: {processed_chunk}"
|
||||
)
|
||||
if processed_chunk is None:
|
||||
continue
|
||||
|
||||
|
|
@ -2193,7 +2167,7 @@ def calculate_total_usage(chunks: List[ModelResponse]) -> Usage:
|
|||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
for chunk in chunks:
|
||||
if "usage" in chunk:
|
||||
if "usage" in chunk and chunk["usage"] is not None:
|
||||
if "prompt_tokens" in chunk["usage"]:
|
||||
prompt_tokens = chunk["usage"].get("prompt_tokens", 0) or 0
|
||||
if "completion_tokens" in chunk["usage"]:
|
||||
|
|
|
|||
|
|
@ -124,6 +124,9 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
guardrailed_tools = guardrailed_inputs.get("tools")
|
||||
if guardrailed_tools is not None:
|
||||
data["tools"] = guardrailed_tools
|
||||
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
|
|
@ -194,7 +197,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
openai_tools = self.adapter.translate_anthropic_tools_to_openai(
|
||||
tools=cast(List[AllAnthropicToolsValues], tools)
|
||||
)
|
||||
tools_to_check.extend(openai_tools)
|
||||
tools_to_check.extend(openai_tools) # type: ignore
|
||||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -172,8 +172,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
@staticmethod
|
||||
def _is_claude_opus_4_6(model: str) -> bool:
|
||||
"""Check if the model is Claude Opus 4.5."""
|
||||
return "opus-4-6" in model.lower() or "opus_4_6" in model.lower()
|
||||
"""Check if the model is Claude Opus 4.5 or Sonnet 4.6."""
|
||||
return "opus-4-6" in model.lower() or "opus_4_6" in model.lower() or "sonnet-4-6" in model.lower() or "sonnet_4_6" in model.lower() or "sonnet-4.6" in model.lower()
|
||||
|
||||
def get_supported_openai_params(self, model: str):
|
||||
params = [
|
||||
|
|
@ -881,6 +881,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"opus-4-5",
|
||||
"opus-4.6",
|
||||
"opus-4-6",
|
||||
"sonnet-4.6",
|
||||
"sonnet-4-6",
|
||||
"sonnet_4.6",
|
||||
"sonnet_4_6",
|
||||
}
|
||||
):
|
||||
_output_format = (
|
||||
|
|
|
|||
|
|
@ -22,6 +22,15 @@ from litellm.types.llms.anthropic import (
|
|||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
def is_anthropic_oauth_key(value: Optional[str]) -> bool:
|
||||
"""Check if a value contains an Anthropic OAuth token (sk-ant-oat*)."""
|
||||
if value is None:
|
||||
return False
|
||||
# Handle both raw token and "Bearer <token>" format
|
||||
if value.startswith("Bearer "):
|
||||
value = value[7:]
|
||||
return value.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX)
|
||||
|
||||
def optionally_handle_anthropic_oauth(
|
||||
headers: dict, api_key: Optional[str]
|
||||
) -> tuple[dict, Optional[str]]:
|
||||
|
|
|
|||
|
|
@ -299,6 +299,26 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
"""
|
||||
return ["messages", "metadata", "system", "tool_choice", "tools", "thinking", "output_format"]
|
||||
|
||||
def _is_web_search_tool(self, tool: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
Check if a tool is an Anthropic web search tool.
|
||||
|
||||
Anthropic web search tools have:
|
||||
- type starting with "web_search" (e.g., "web_search_20260209")
|
||||
- name = "web_search"
|
||||
|
||||
Args:
|
||||
tool: Tool definition dict
|
||||
|
||||
Returns:
|
||||
True if this is a web search tool
|
||||
"""
|
||||
tool_type = tool.get("type", "")
|
||||
tool_name = tool.get("name", "")
|
||||
return (
|
||||
isinstance(tool_type, str) and tool_type.startswith("web_search")
|
||||
) or tool_name == "web_search"
|
||||
|
||||
def translate_anthropic_messages_to_openai( # noqa: PLR0915
|
||||
self,
|
||||
messages: List[
|
||||
|
|
@ -872,10 +892,25 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
if "tools" in anthropic_message_request:
|
||||
tools = anthropic_message_request["tools"]
|
||||
if tools:
|
||||
new_kwargs["tools"], tool_name_mapping = self.translate_anthropic_tools_to_openai(
|
||||
tools=cast(List[AllAnthropicToolsValues], tools),
|
||||
model=new_kwargs.get("model"),
|
||||
)
|
||||
# Separate web search tools from regular tools
|
||||
web_search_tools = []
|
||||
regular_tools = []
|
||||
for tool in tools:
|
||||
if self._is_web_search_tool(cast(Dict[str, Any], tool)):
|
||||
web_search_tools.append(tool)
|
||||
else:
|
||||
regular_tools.append(tool)
|
||||
|
||||
# If web search tools are present, add web_search_options parameter
|
||||
if web_search_tools:
|
||||
new_kwargs["web_search_options"] = {} # type: ignore
|
||||
|
||||
# Only translate regular tools (non-web-search)
|
||||
if regular_tools:
|
||||
new_kwargs["tools"], tool_name_mapping = self.translate_anthropic_tools_to_openai(
|
||||
tools=cast(List[AllAnthropicToolsValues], regular_tools),
|
||||
model=new_kwargs.get("model"),
|
||||
)
|
||||
|
||||
## CONVERT THINKING
|
||||
if "thinking" in anthropic_message_request:
|
||||
|
|
|
|||
|
|
@ -384,6 +384,14 @@ class BaseAWSLLM:
|
|||
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
|
||||
model_id, spec="moonshot"
|
||||
)
|
||||
elif "nova-2/" in model_id:
|
||||
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
|
||||
model_id, spec="nova-2"
|
||||
)
|
||||
elif "nova/" in model_id:
|
||||
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
|
||||
model_id, spec="nova"
|
||||
)
|
||||
return model_id
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -272,7 +272,18 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
if unencoded_model_id is not None:
|
||||
modelId = self.encode_model_id(model_id=unencoded_model_id)
|
||||
else:
|
||||
modelId = self.encode_model_id(model_id=model)
|
||||
# Strip nova spec prefixes before encoding model ID for API URL
|
||||
_model_for_id = model
|
||||
_stripped = _model_for_id
|
||||
for rp in ["bedrock/converse/", "bedrock/", "converse/"]:
|
||||
if _stripped.startswith(rp):
|
||||
_stripped = _stripped[len(rp):]
|
||||
break
|
||||
for _nova_prefix in ["nova-2/", "nova/"]:
|
||||
if _stripped.startswith(_nova_prefix):
|
||||
_model_for_id = _model_for_id.replace(_nova_prefix, "", 1)
|
||||
break
|
||||
modelId = self.encode_model_id(model_id=_model_for_id)
|
||||
|
||||
fake_stream = litellm.AmazonConverseConfig().should_fake_stream(
|
||||
fake_stream=fake_stream,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Translating between OpenAI's `/chat/completion` format and Amazon's `/converse`
|
|||
"""
|
||||
|
||||
import copy
|
||||
import json
|
||||
import time
|
||||
import types
|
||||
from typing import List, Literal, Optional, Tuple, Union, cast, overload
|
||||
|
|
@ -85,9 +86,37 @@ BEDROCK_COMPUTER_USE_TOOLS = [
|
|||
UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS = [
|
||||
"advanced-tool-use", # Bedrock Converse doesn't support advanced-tool-use beta headers
|
||||
"prompt-caching", # Prompt caching not supported in Converse API
|
||||
"compact-2026-01-12", # The compact beta feature is not currently supported on the Converse and ConverseStream APIs
|
||||
"compact-2026-01-12", # The compact beta feature is not currently supported on the Converse and ConverseStream APIs
|
||||
]
|
||||
|
||||
# Models that support Bedrock's native structured outputs API (outputConfig.textFormat)
|
||||
# Uses substring matching against the Bedrock model ID
|
||||
# Ref: https://docs.aws.amazon.com/bedrock/latest/userguide/structured-output.html
|
||||
BEDROCK_NATIVE_STRUCTURED_OUTPUT_MODELS = {
|
||||
# Anthropic Claude 4.5+
|
||||
"claude-haiku-4-5",
|
||||
"claude-sonnet-4-5",
|
||||
"claude-opus-4-5",
|
||||
"claude-opus-4-6",
|
||||
# Qwen3
|
||||
"qwen3",
|
||||
# DeepSeek
|
||||
"deepseek-v3.1",
|
||||
# Gemma 3
|
||||
"gemma-3",
|
||||
# MiniMax
|
||||
"minimax-m2",
|
||||
# Mistral (magistral-small excluded: broken constrained decoding on Bedrock)
|
||||
"ministral",
|
||||
"mistral-large-3",
|
||||
"voxtral",
|
||||
# Moonshot
|
||||
"kimi-k2",
|
||||
# NVIDIA
|
||||
"nemotron-nano",
|
||||
# OpenAI (gpt-oss excluded: broken constrained decoding, works via tool-call fallback)
|
||||
}
|
||||
|
||||
|
||||
class AmazonConverseConfig(BaseConfig):
|
||||
"""
|
||||
|
|
@ -270,45 +299,56 @@ class AmazonConverseConfig(BaseConfig):
|
|||
llm_provider="bedrock",
|
||||
)
|
||||
|
||||
def _is_nova_lite_2_model(self, model: str) -> bool:
|
||||
def _is_nova_2_model(self, model: str) -> bool:
|
||||
"""
|
||||
Check if the model is a Nova Lite 2 model that supports reasoningConfig.
|
||||
Check if the model is a Nova 2 model that supports reasoningConfig.
|
||||
|
||||
Nova Lite 2 models use a different reasoning configuration structure compared to
|
||||
Nova 2 models use a different reasoning configuration structure compared to
|
||||
Anthropic's thinking parameter and GPT-OSS's reasoning_effort parameter.
|
||||
|
||||
Supported models:
|
||||
- amazon.nova-2-lite-v1:0
|
||||
- amazon.nova-2-pro-preview-20251202-v1:0
|
||||
- us.amazon.nova-2-lite-v1:0
|
||||
- eu.amazon.nova-2-lite-v1:0
|
||||
- apac.amazon.nova-2-lite-v1:0
|
||||
- (and other regional variants)
|
||||
|
||||
Args:
|
||||
model: The model identifier
|
||||
|
||||
Returns:
|
||||
True if the model is a Nova Lite 2 model, False otherwise
|
||||
True if the model is a Nova 2 model, False otherwise
|
||||
|
||||
Examples:
|
||||
>>> config = AmazonConverseConfig()
|
||||
>>> config._is_nova_lite_2_model("amazon.nova-2-lite-v1:0")
|
||||
>>> config._is_nova_2_model("amazon.nova-2-lite-v1:0")
|
||||
True
|
||||
>>> config._is_nova_lite_2_model("us.amazon.nova-2-lite-v1:0")
|
||||
>>> config._is_nova_2_model("us.amazon.nova-2-lite-v1:0")
|
||||
True
|
||||
>>> config._is_nova_lite_2_model("amazon.nova-pro-1-5-v1:0")
|
||||
>>> config._is_nova_2_model("us.amazon.nova-2-pro-preview-20251202-v1:0")
|
||||
True
|
||||
>>> config._is_nova_2_model("amazon.nova-pro-1-5-v1:0")
|
||||
False
|
||||
>>> config._is_nova_lite_2_model("amazon.nova-pro-v1:0")
|
||||
>>> config._is_nova_2_model("amazon.nova-pro-v1:0")
|
||||
False
|
||||
"""
|
||||
# Remove regional prefix if present (us., eu., apac.)
|
||||
# Remove provider routing prefix if present (bedrock/converse/, bedrock/, converse/)
|
||||
model_without_region = model
|
||||
for prefix in ["us.", "eu.", "apac."]:
|
||||
if model.startswith(prefix):
|
||||
model_without_region = model[len(prefix) :]
|
||||
for routing_prefix in ["bedrock/converse/", "bedrock/", "converse/"]:
|
||||
if model_without_region.startswith(routing_prefix):
|
||||
model_without_region = model_without_region[len(routing_prefix) :]
|
||||
break
|
||||
|
||||
# Check if the model is specifically Nova Lite 2
|
||||
return "nova-2-lite" in model_without_region
|
||||
# Remove regional prefix if present (us., eu., apac.)
|
||||
for prefix in ["us.", "eu.", "apac."]:
|
||||
if model_without_region.startswith(prefix):
|
||||
model_without_region = model_without_region[len(prefix) :]
|
||||
break
|
||||
|
||||
# Check if the model is a Nova 2 model (matches nova-2-lite, nova-2-pro, etc.)
|
||||
# Also check for nova-2/ spec prefix for imported models
|
||||
return model_without_region.startswith("amazon.nova-2-") or model_without_region.startswith("nova-2/")
|
||||
|
||||
def _map_web_search_options(
|
||||
self, web_search_options: dict, model: str
|
||||
|
|
@ -396,7 +436,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
Different model families handle reasoning effort differently:
|
||||
- GPT-OSS models: Keep reasoning_effort as-is (passed to additionalModelRequestFields)
|
||||
- Nova Lite 2 models: Transform to reasoningConfig structure
|
||||
- Nova 2 models: Transform to reasoningConfig structure
|
||||
- Other models (Anthropic, etc.): Convert to thinking parameter
|
||||
|
||||
Args:
|
||||
|
|
@ -425,8 +465,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
# GPT-OSS models: keep reasoning_effort as-is
|
||||
# It will be passed through to additionalModelRequestFields
|
||||
optional_params["reasoning_effort"] = reasoning_effort
|
||||
elif self._is_nova_lite_2_model(model):
|
||||
# Nova Lite 2 models: transform to reasoningConfig
|
||||
elif self._is_nova_2_model(model):
|
||||
# Nova 2 models: transform to reasoningConfig
|
||||
reasoning_config = self._transform_reasoning_effort_to_reasoning_config(
|
||||
reasoning_effort
|
||||
)
|
||||
|
|
@ -480,6 +520,9 @@ class AmazonConverseConfig(BaseConfig):
|
|||
supported_params.append("tool_choice")
|
||||
supported_params.append("thinking")
|
||||
supported_params.append("reasoning_effort")
|
||||
# For nova imported models, also add web_search_options
|
||||
if "nova" in model.lower():
|
||||
supported_params.append("web_search_options")
|
||||
return supported_params
|
||||
|
||||
## Filter out 'cross-region' from model name
|
||||
|
|
@ -514,8 +557,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
if "gpt-oss" in model:
|
||||
supported_params.append("reasoning_effort")
|
||||
elif self._is_nova_lite_2_model(model):
|
||||
# Nova Lite 2 models support reasoning_effort (transformed to reasoningConfig)
|
||||
elif self._is_nova_2_model(model):
|
||||
# Nova 2 models support reasoning_effort (transformed to reasoningConfig)
|
||||
# These models use a different reasoning structure than Anthropic's thinking parameter
|
||||
supported_params.append("reasoning_effort")
|
||||
elif (
|
||||
|
|
@ -714,6 +757,100 @@ class AmazonConverseConfig(BaseConfig):
|
|||
)
|
||||
return _tool
|
||||
|
||||
@staticmethod
|
||||
def _supports_native_structured_outputs(model: str) -> bool:
|
||||
"""Check if the Bedrock model supports native structured outputs (outputConfig.textFormat)."""
|
||||
return any(
|
||||
substring in model
|
||||
for substring in BEDROCK_NATIVE_STRUCTURED_OUTPUT_MODELS
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _add_additional_properties_to_schema(schema: dict) -> dict:
|
||||
"""
|
||||
Recursively ensure all object types in a JSON schema have
|
||||
``"additionalProperties": false``.
|
||||
|
||||
Bedrock's native structured-outputs API requires this field to be
|
||||
explicitly set on every object node, otherwise it returns a
|
||||
validation error.
|
||||
"""
|
||||
if not isinstance(schema, dict):
|
||||
return schema
|
||||
|
||||
result = dict(schema)
|
||||
|
||||
if result.get("type") == "object" and "additionalProperties" not in result:
|
||||
result["additionalProperties"] = False
|
||||
|
||||
# Recurse into nested schemas
|
||||
if "properties" in result and isinstance(result["properties"], dict):
|
||||
result["properties"] = {
|
||||
k: AmazonConverseConfig._add_additional_properties_to_schema(v)
|
||||
for k, v in result["properties"].items()
|
||||
}
|
||||
if "items" in result and isinstance(result["items"], dict):
|
||||
result["items"] = AmazonConverseConfig._add_additional_properties_to_schema(
|
||||
result["items"]
|
||||
)
|
||||
for defs_key in ("$defs", "definitions"):
|
||||
if defs_key in result and isinstance(result[defs_key], dict):
|
||||
result[defs_key] = {
|
||||
k: AmazonConverseConfig._add_additional_properties_to_schema(v)
|
||||
for k, v in result[defs_key].items()
|
||||
}
|
||||
for key in ("anyOf", "allOf", "oneOf"):
|
||||
if key in result and isinstance(result[key], list):
|
||||
result[key] = [
|
||||
AmazonConverseConfig._add_additional_properties_to_schema(item)
|
||||
for item in result[key]
|
||||
]
|
||||
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _create_output_config_for_response_format(
|
||||
json_schema: Optional[dict] = None,
|
||||
name: Optional[str] = None,
|
||||
description: Optional[str] = None,
|
||||
) -> "OutputConfigBlock":
|
||||
"""
|
||||
Build an outputConfig block for Bedrock's native structured outputs API.
|
||||
|
||||
The Converse API expects:
|
||||
{
|
||||
"outputConfig": {
|
||||
"textFormat": {
|
||||
"type": "json_schema",
|
||||
"structure": {
|
||||
"jsonSchema": {
|
||||
"schema": "<json-string>",
|
||||
"name": "optional",
|
||||
"description": "optional"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
"""
|
||||
if json_schema is not None:
|
||||
json_schema = AmazonConverseConfig._add_additional_properties_to_schema(
|
||||
json_schema
|
||||
)
|
||||
schema_str = json.dumps(json_schema) if json_schema is not None else "{}"
|
||||
json_schema_def: JsonSchemaDefinition = {"schema": schema_str}
|
||||
if name is not None:
|
||||
json_schema_def["name"] = name
|
||||
if description is not None:
|
||||
json_schema_def["description"] = description
|
||||
|
||||
return OutputConfigBlock(
|
||||
textFormat=OutputFormat(
|
||||
type="json_schema",
|
||||
structure=OutputFormatStructure(jsonSchema=json_schema_def),
|
||||
)
|
||||
)
|
||||
|
||||
def _apply_tool_call_transformation(
|
||||
self,
|
||||
tools: List[OpenAIChatCompletionToolParam],
|
||||
|
|
@ -806,8 +943,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
)
|
||||
|
||||
# Only update thinking tokens for non-GPT-OSS models and non-Nova-Lite-2 models
|
||||
# Nova Lite 2 handles token budgeting differently through reasoningConfig
|
||||
if "gpt-oss" not in model and not self._is_nova_lite_2_model(model):
|
||||
# Nova 2 handles token budgeting differently through reasoningConfig
|
||||
if "gpt-oss" not in model and not self._is_nova_2_model(model):
|
||||
self.update_optional_params_with_thinking_tokens(
|
||||
non_default_params=non_default_params, optional_params=optional_params
|
||||
)
|
||||
|
|
@ -843,45 +980,53 @@ class AmazonConverseConfig(BaseConfig):
|
|||
return optional_params
|
||||
|
||||
json_schema: Optional[dict] = None
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
if "response_schema" in value:
|
||||
json_schema = value["response_schema"]
|
||||
elif "json_schema" in value:
|
||||
json_schema = value["json_schema"]["schema"]
|
||||
name = value["json_schema"].get("name")
|
||||
description = value["json_schema"].get("description")
|
||||
|
||||
if "type" in value and value["type"] == "text":
|
||||
return optional_params
|
||||
|
||||
"""
|
||||
Follow similar approach to anthropic - translate to a single tool call.
|
||||
|
||||
When using tools in this way: - https://docs.anthropic.com/en/docs/build-with-claude/tool-use#json-mode
|
||||
- You usually want to provide a single tool
|
||||
- You should set tool_choice (see Forcing tool use) to instruct the model to explicitly use that tool
|
||||
- Remember that the model will pass the input to the tool, so the name of the tool and description should be from the model’s perspective.
|
||||
"""
|
||||
_tool = self._create_json_tool_call_for_response_format(
|
||||
json_schema=json_schema,
|
||||
description=description,
|
||||
)
|
||||
optional_params = self._add_tools_to_optional_params(
|
||||
optional_params=optional_params, tools=[_tool]
|
||||
)
|
||||
|
||||
if (
|
||||
litellm.utils.supports_tool_choice(
|
||||
model=model, custom_llm_provider=self.custom_llm_provider
|
||||
if self._supports_native_structured_outputs(model) and json_schema is not None:
|
||||
# Use Bedrock's native structured outputs API (outputConfig.textFormat)
|
||||
# No synthetic tool injection, no fake_stream needed.
|
||||
# Requires an explicit schema — json_object with no schema falls through
|
||||
# to the tool-call path below.
|
||||
output_config = self._create_output_config_for_response_format(
|
||||
json_schema=json_schema,
|
||||
name=name,
|
||||
description=description,
|
||||
)
|
||||
and not is_thinking_enabled
|
||||
):
|
||||
optional_params["tool_choice"] = ToolChoiceValuesBlock(
|
||||
tool=SpecificToolChoiceBlock(name=RESPONSE_FORMAT_TOOL_NAME)
|
||||
optional_params["outputConfig"] = output_config
|
||||
else:
|
||||
# Fallback: translate to a synthetic tool call
|
||||
# https://docs.anthropic.com/en/docs/build-with-claude/tool-use#json-mode
|
||||
_tool = self._create_json_tool_call_for_response_format(
|
||||
json_schema=json_schema,
|
||||
description=description,
|
||||
)
|
||||
optional_params = self._add_tools_to_optional_params(
|
||||
optional_params=optional_params, tools=[_tool]
|
||||
)
|
||||
|
||||
if (
|
||||
litellm.utils.supports_tool_choice(
|
||||
model=model, custom_llm_provider=self.custom_llm_provider
|
||||
)
|
||||
and not is_thinking_enabled
|
||||
):
|
||||
optional_params["tool_choice"] = ToolChoiceValuesBlock(
|
||||
tool=SpecificToolChoiceBlock(name=RESPONSE_FORMAT_TOOL_NAME)
|
||||
)
|
||||
if non_default_params.get("stream", False) is True:
|
||||
optional_params["fake_stream"] = True
|
||||
|
||||
optional_params["json_mode"] = True
|
||||
if non_default_params.get("stream", False) is True:
|
||||
optional_params["fake_stream"] = True
|
||||
|
||||
return optional_params
|
||||
|
||||
def update_optional_params_with_thinking_tokens(
|
||||
|
|
@ -1024,7 +1169,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
def _prepare_request_params(
|
||||
self, optional_params: dict, model: str
|
||||
) -> Tuple[dict, dict, dict]:
|
||||
) -> Tuple[dict, dict, dict, Optional[OutputConfigBlock]]:
|
||||
"""Prepare and separate request parameters."""
|
||||
# Filter out exception objects before deepcopy to prevent deepcopy failures
|
||||
# Exceptions should not be stored in optional_params (this is a defensive fix)
|
||||
|
|
@ -1047,6 +1192,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
if request_metadata is not None:
|
||||
self._validate_request_metadata(request_metadata)
|
||||
|
||||
output_config: Optional[OutputConfigBlock] = inference_params.pop("outputConfig", None)
|
||||
|
||||
# keep supported params in 'inference_params', and set all model-specific params in 'additional_request_params'
|
||||
additional_request_params = {
|
||||
k: v for k, v in inference_params.items() if k not in total_supported_params
|
||||
|
|
@ -1071,7 +1218,12 @@ class AmazonConverseConfig(BaseConfig):
|
|||
additional_request_params
|
||||
)
|
||||
|
||||
return inference_params, additional_request_params, request_metadata
|
||||
return (
|
||||
inference_params,
|
||||
additional_request_params,
|
||||
request_metadata,
|
||||
output_config,
|
||||
)
|
||||
|
||||
def _process_tools_and_beta(
|
||||
self,
|
||||
|
|
@ -1125,22 +1277,44 @@ class AmazonConverseConfig(BaseConfig):
|
|||
# "computer-use-2025-01-24" for Claude Sonnet 4.5, Haiku 4.5, Opus 4.1, Sonnet 4, Opus 4, and Sonnet 3.7
|
||||
# "computer-use-2024-10-22" for older models
|
||||
model_lower = model.lower()
|
||||
if "opus-4.6" in model_lower or "opus_4.6" in model_lower or "opus-4-6" in model_lower or "opus_4_6" in model_lower:
|
||||
if "opus-4.6" in model_lower or "opus_4.6" in model_lower or "opus-4-6" in model_lower or "opus_4_6" in model_lower or "sonnet-4.6" in model_lower or "sonnet_4.6" in model_lower or "sonnet-4-6" in model_lower or "sonnet_4_6" in model_lower:
|
||||
computer_use_header = "computer-use-2025-11-24"
|
||||
elif "opus-4.5" in model_lower or "opus_4.5" in model_lower or "opus-4-5" in model_lower or "opus_4_5" in model_lower:
|
||||
elif (
|
||||
"opus-4.5" in model_lower
|
||||
or "opus_4.5" in model_lower
|
||||
or "opus-4-5" in model_lower
|
||||
or "opus_4_5" in model_lower
|
||||
):
|
||||
computer_use_header = "computer-use-2025-11-24"
|
||||
elif any(pattern in model_lower for pattern in [
|
||||
"sonnet-4.5", "sonnet_4.5", "sonnet-4-5", "sonnet_4_5",
|
||||
"haiku-4.5", "haiku_4.5", "haiku-4-5", "haiku_4_5",
|
||||
"opus-4.1", "opus_4.1", "opus-4-1", "opus_4_1",
|
||||
"sonnet-4", "sonnet_4",
|
||||
"opus-4", "opus_4",
|
||||
"sonnet-3.7", "sonnet_3.7", "sonnet-3-7", "sonnet_3_7"
|
||||
]):
|
||||
elif any(
|
||||
pattern in model_lower
|
||||
for pattern in [
|
||||
"sonnet-4.5",
|
||||
"sonnet_4.5",
|
||||
"sonnet-4-5",
|
||||
"sonnet_4_5",
|
||||
"haiku-4.5",
|
||||
"haiku_4.5",
|
||||
"haiku-4-5",
|
||||
"haiku_4_5",
|
||||
"opus-4.1",
|
||||
"opus_4.1",
|
||||
"opus-4-1",
|
||||
"opus_4_1",
|
||||
"sonnet-4",
|
||||
"sonnet_4",
|
||||
"opus-4",
|
||||
"opus_4",
|
||||
"sonnet-3.7",
|
||||
"sonnet_3.7",
|
||||
"sonnet-3-7",
|
||||
"sonnet_3_7",
|
||||
]
|
||||
):
|
||||
computer_use_header = "computer-use-2025-01-24"
|
||||
else:
|
||||
computer_use_header = "computer-use-2024-10-22"
|
||||
|
||||
|
||||
anthropic_beta_list.append(computer_use_header)
|
||||
# Transform computer use tools to proper Bedrock format
|
||||
transformed_computer_tools = self._transform_computer_use_tools(
|
||||
|
|
@ -1214,6 +1388,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
inference_params,
|
||||
additional_request_params,
|
||||
request_metadata,
|
||||
output_config,
|
||||
) = self._prepare_request_params(optional_params, model)
|
||||
|
||||
original_tools = inference_params.pop("tools", [])
|
||||
|
|
@ -1256,6 +1431,9 @@ class AmazonConverseConfig(BaseConfig):
|
|||
if request_metadata is not None:
|
||||
data["requestMetadata"] = request_metadata
|
||||
|
||||
if output_config is not None:
|
||||
data["outputConfig"] = output_config
|
||||
|
||||
return data
|
||||
|
||||
async def _async_transform_request(
|
||||
|
|
@ -1504,9 +1682,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
return message, returned_finish_reason
|
||||
|
||||
def _translate_message_content(
|
||||
self, content_blocks: List[ContentBlock]
|
||||
) -> Tuple[
|
||||
def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[
|
||||
str,
|
||||
List[ChatCompletionToolCallChunk],
|
||||
Optional[List[BedrockConverseReasoningContentBlock]],
|
||||
|
|
@ -1523,9 +1699,9 @@ class AmazonConverseConfig(BaseConfig):
|
|||
"""
|
||||
content_str = ""
|
||||
tools: List[ChatCompletionToolCallChunk] = []
|
||||
reasoningContentBlocks: Optional[
|
||||
List[BedrockConverseReasoningContentBlock]
|
||||
] = None
|
||||
reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
|
||||
None
|
||||
)
|
||||
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
|
||||
for idx, content in enumerate(content_blocks):
|
||||
"""
|
||||
|
|
@ -1652,9 +1828,9 @@ class AmazonConverseConfig(BaseConfig):
|
|||
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
|
||||
content_str = ""
|
||||
tools: List[ChatCompletionToolCallChunk] = []
|
||||
reasoningContentBlocks: Optional[
|
||||
List[BedrockConverseReasoningContentBlock]
|
||||
] = None
|
||||
reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
|
||||
None
|
||||
)
|
||||
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
|
||||
|
||||
if message is not None:
|
||||
|
|
@ -1673,17 +1849,17 @@ class AmazonConverseConfig(BaseConfig):
|
|||
provider_specific_fields["citationsContent"] = citationsContentBlocks
|
||||
|
||||
if provider_specific_fields:
|
||||
chat_completion_message[
|
||||
"provider_specific_fields"
|
||||
] = provider_specific_fields
|
||||
chat_completion_message["provider_specific_fields"] = (
|
||||
provider_specific_fields
|
||||
)
|
||||
|
||||
if reasoningContentBlocks is not None:
|
||||
chat_completion_message[
|
||||
"reasoning_content"
|
||||
] = self._transform_reasoning_content(reasoningContentBlocks)
|
||||
chat_completion_message[
|
||||
"thinking_blocks"
|
||||
] = self._transform_thinking_blocks(reasoningContentBlocks)
|
||||
chat_completion_message["reasoning_content"] = (
|
||||
self._transform_reasoning_content(reasoningContentBlocks)
|
||||
)
|
||||
chat_completion_message["thinking_blocks"] = (
|
||||
self._transform_thinking_blocks(reasoningContentBlocks)
|
||||
)
|
||||
chat_completion_message["content"] = content_str
|
||||
if (
|
||||
json_mode is True
|
||||
|
|
@ -1696,8 +1872,6 @@ class AmazonConverseConfig(BaseConfig):
|
|||
)
|
||||
json_mode_content_str: Optional[str] = tools[0]["function"].get("arguments")
|
||||
if json_mode_content_str is not None:
|
||||
import json
|
||||
|
||||
# Bedrock returns the response wrapped in a "properties" object
|
||||
# We need to extract the actual content from this wrapper
|
||||
try:
|
||||
|
|
@ -1716,7 +1890,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
pass
|
||||
|
||||
chat_completion_message["content"] = json_mode_content_str
|
||||
else:
|
||||
elif tools:
|
||||
chat_completion_message["tool_calls"] = tools
|
||||
|
||||
## CALCULATING USAGE - bedrock returns usage in the headers
|
||||
|
|
|
|||
|
|
@ -404,7 +404,7 @@ def extract_model_name_from_bedrock_arn(model: str) -> str:
|
|||
|
||||
def strip_bedrock_routing_prefix(model: str) -> str:
|
||||
"""Strip LiteLLM routing prefixes from model name."""
|
||||
for prefix in ["bedrock/", "converse/", "invoke/", "openai/"]:
|
||||
for prefix in ["bedrock/", "converse/", "invoke/", "openai/", "nova-2/", "nova/"]:
|
||||
if model.startswith(prefix):
|
||||
model = model.split("/", 1)[1]
|
||||
return model
|
||||
|
|
@ -427,7 +427,20 @@ def get_bedrock_base_model(model: str) -> str:
|
|||
- "us.meta.llama3-2-11b-instruct-v1:0" -> "meta.llama3-2-11b-instruct-v1"
|
||||
- "bedrock/converse/model" -> "model"
|
||||
- "anthropic.claude-3-5-sonnet-20241022-v2:0:51k" -> "anthropic.claude-3-5-sonnet-20241022-v2:0"
|
||||
- "bedrock/nova-2/arn:aws:..." -> "amazon.nova-2-custom"
|
||||
- "bedrock/nova/arn:aws:..." -> "amazon.nova-custom"
|
||||
"""
|
||||
# Detect nova spec prefixes before stripping them
|
||||
stripped = model
|
||||
for rp in ["bedrock/converse/", "bedrock/", "converse/"]:
|
||||
if stripped.startswith(rp):
|
||||
stripped = stripped[len(rp):]
|
||||
break
|
||||
if stripped.startswith("nova-2/"):
|
||||
return "amazon.nova-2-custom"
|
||||
elif stripped.startswith("nova/"):
|
||||
return "amazon.nova-custom"
|
||||
|
||||
model = strip_bedrock_routing_prefix(model)
|
||||
model = extract_model_name_from_bedrock_arn(model)
|
||||
model = strip_bedrock_throughput_suffix(model)
|
||||
|
|
@ -465,6 +478,14 @@ def is_claude_4_5_on_bedrock(model: str) -> bool:
|
|||
"opus_4.5",
|
||||
"opus-4-5",
|
||||
"opus_4_5",
|
||||
"sonnet-4.6",
|
||||
"sonnet_4.6",
|
||||
"sonnet-4-6",
|
||||
"sonnet_4_6",
|
||||
"opus-4.6",
|
||||
"opus_4.6",
|
||||
"opus-4-6",
|
||||
"opus_4_6",
|
||||
]
|
||||
return any(pattern in model_lower for pattern in claude_4_5_patterns)
|
||||
|
||||
|
|
@ -594,6 +615,11 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
if prefix in model:
|
||||
return route_type
|
||||
|
||||
# Check for nova spec prefixes (nova/ and nova-2/)
|
||||
_model_after_bedrock = model.replace("bedrock/", "", 1)
|
||||
if _model_after_bedrock.startswith("nova-2/") or _model_after_bedrock.startswith("nova/"):
|
||||
return "converse"
|
||||
|
||||
base_model = BedrockModelInfo.get_base_model(model)
|
||||
alt_model = BedrockModelInfo.get_non_litellm_routing_model_name(model=model)
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Helper util for handling bedrock-specific cost calculation
|
|||
- e.g.: prompt caching
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Tuple
|
||||
from typing import TYPE_CHECKING, Optional, Tuple
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
|
||||
|
|
@ -11,12 +11,17 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]:
|
||||
def cost_per_token(
|
||||
model: str, usage: "Usage", service_tier: Optional[str] = None
|
||||
) -> Tuple[float, float]:
|
||||
"""
|
||||
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
|
||||
|
||||
Follows the same logic as Anthropic's cost per token calculation.
|
||||
"""
|
||||
return generic_cost_per_token(
|
||||
model=model, usage=usage, custom_llm_provider="bedrock"
|
||||
)
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider="bedrock",
|
||||
service_tier=service_tier,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -180,6 +180,14 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
"opus_4", # Opus 4
|
||||
"sonnet-4",
|
||||
"sonnet_4", # Sonnet 4
|
||||
"sonnet-4.6",
|
||||
"sonnet_4.6",
|
||||
"sonnet-4-6",
|
||||
"sonnet_4_6",
|
||||
"opus-4.6",
|
||||
"opus_4.6",
|
||||
"opus-4-6",
|
||||
"opus_4_6",
|
||||
]
|
||||
|
||||
return any(pattern in model_lower for pattern in supported_patterns)
|
||||
|
|
@ -251,6 +259,11 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
"opus_4.6",
|
||||
"opus-4-6",
|
||||
"opus_4_6",
|
||||
#sonnet 4.6
|
||||
"sonnet-4.6",
|
||||
"sonnet_4.6",
|
||||
"sonnet-4-6",
|
||||
"sonnet_4_6",
|
||||
]
|
||||
|
||||
return any(pattern in model_lower for pattern in supported_patterns)
|
||||
|
|
@ -285,7 +298,7 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
programmatic_tool_calling_used or input_examples_used
|
||||
):
|
||||
beta_set.discard(ANTHROPIC_TOOL_SEARCH_BETA_HEADER)
|
||||
if "opus-4" in model.lower() or "opus_4" in model.lower():
|
||||
if self._supports_tool_search_on_bedrock(model):
|
||||
beta_set.add("tool-search-tool-2025-10-19")
|
||||
|
||||
def _convert_output_format_to_inline_schema(
|
||||
|
|
@ -420,10 +433,8 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
beta_set=beta_set,
|
||||
)
|
||||
|
||||
# --- Custom logic: if tool-search-tool-2025-10-19 is present, add tool-examples-2025-10-29 ---
|
||||
if "tool-search-tool-2025-10-19" in beta_set:
|
||||
beta_set.add("tool-examples-2025-10-29")
|
||||
# ------------------------------------------------------------------------------
|
||||
|
||||
if beta_set:
|
||||
anthropic_messages_request["anthropic_beta"] = list(beta_set)
|
||||
|
|
|
|||
6
litellm/llms/duckduckgo/search/__init__.py
Normal file
6
litellm/llms/duckduckgo/search/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
"""
|
||||
DuckDuckGo Search API module.
|
||||
"""
|
||||
from litellm.llms.duckduckgo.search.transformation import DuckDuckGoSearchConfig
|
||||
|
||||
__all__ = ["DuckDuckGoSearchConfig"]
|
||||
252
litellm/llms/duckduckgo/search/transformation.py
Normal file
252
litellm/llms/duckduckgo/search/transformation.py
Normal file
|
|
@ -0,0 +1,252 @@
|
|||
"""
|
||||
Calls DuckDuckGo's Instant Answer API to search the web.
|
||||
|
||||
DuckDuckGo API Reference: https://duckduckgo.com/api
|
||||
"""
|
||||
from typing import Dict, List, Literal, Optional, TypedDict, Union
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.search.transformation import (
|
||||
BaseSearchConfig,
|
||||
SearchResponse,
|
||||
SearchResult,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class _DuckDuckGoSearchRequestRequired(TypedDict):
|
||||
"""Required fields for DuckDuckGo Search API request."""
|
||||
q: str # Required - search query
|
||||
|
||||
|
||||
class DuckDuckGoSearchRequest(_DuckDuckGoSearchRequestRequired, total=False):
|
||||
"""
|
||||
DuckDuckGo Instant Answer API request format.
|
||||
Based on: https://duckduckgo.com/api
|
||||
"""
|
||||
format: str # Optional - output format ('json', 'xml'), default 'json'
|
||||
pretty: int # Optional - pretty print (0 or 1), default 1
|
||||
no_redirect: int # Optional - skip HTTP redirects (0 or 1), default 0
|
||||
no_html: int # Optional - remove HTML from text (0 or 1), default 0
|
||||
skip_disambig: int # Optional - skip disambiguation results (0 or 1), default 0
|
||||
|
||||
|
||||
class DuckDuckGoSearchConfig(BaseSearchConfig):
|
||||
DUCKDUCKGO_API_BASE = "https://api.duckduckgo.com"
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "DuckDuckGo"
|
||||
|
||||
def get_http_method(self) -> Literal["GET", "POST"]:
|
||||
"""
|
||||
Get HTTP method for search requests.
|
||||
DuckDuckGo Instant Answer API uses GET requests.
|
||||
|
||||
Returns:
|
||||
HTTP method 'GET'
|
||||
"""
|
||||
return "GET"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
"""
|
||||
Validate environment and return headers.
|
||||
DuckDuckGo Instant Answer API does not require authentication.
|
||||
"""
|
||||
# DuckDuckGo API is free and doesn't require API key
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
optional_params: dict,
|
||||
data: Optional[Union[Dict, List[Dict]]] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for Search endpoint.
|
||||
DuckDuckGo uses query parameters, so we construct the URL with the query.
|
||||
"""
|
||||
api_base = api_base or get_secret_str("DUCKDUCKGO_API_BASE") or self.DUCKDUCKGO_API_BASE
|
||||
|
||||
# Build query parameters from the transformed request body
|
||||
if data and isinstance(data, dict) and "_duckduckgo_params" in data:
|
||||
params = data["_duckduckgo_params"]
|
||||
query_string = urlencode(params, doseq=True)
|
||||
return f"{api_base}/?{query_string}"
|
||||
|
||||
return api_base
|
||||
|
||||
|
||||
def transform_search_request(
|
||||
self,
|
||||
query: Union[str, List[str]],
|
||||
optional_params: dict,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
"""
|
||||
Transform Search request to DuckDuckGo API format.
|
||||
|
||||
Args:
|
||||
query: Search query (string or list of strings). DuckDuckGo only supports single string queries.
|
||||
optional_params: Optional parameters for the request
|
||||
- max_results: Maximum number of search results (DuckDuckGo API doesn't directly support this, used for filtering)
|
||||
- format: Output format ('json', 'xml')
|
||||
- pretty: Pretty print (0 or 1)
|
||||
- no_redirect: Skip HTTP redirects (0 or 1)
|
||||
- no_html: Remove HTML from text (0 or 1)
|
||||
- skip_disambig: Skip disambiguation results (0 or 1)
|
||||
|
||||
Returns:
|
||||
Dict with typed request data following DuckDuckGoSearchRequest spec
|
||||
"""
|
||||
if isinstance(query, list):
|
||||
# DuckDuckGo only supports single string queries
|
||||
query = " ".join(query)
|
||||
|
||||
request_data: DuckDuckGoSearchRequest = {
|
||||
"q": query,
|
||||
"format": "json", # Always use JSON format
|
||||
}
|
||||
|
||||
# Convert to dict before dynamic key assignments
|
||||
result_data = dict(request_data)
|
||||
|
||||
if "max_results" in optional_params:
|
||||
result_data["_max_results"] = optional_params["max_results"]
|
||||
|
||||
# Pass through DuckDuckGo-specific parameters
|
||||
ddg_params = ["pretty", "no_redirect", "no_html", "skip_disambig"]
|
||||
for param in ddg_params:
|
||||
if param in optional_params:
|
||||
result_data[param] = optional_params[param]
|
||||
|
||||
return {
|
||||
"_duckduckgo_params": result_data,
|
||||
}
|
||||
|
||||
def transform_search_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
**kwargs,
|
||||
) -> SearchResponse:
|
||||
"""
|
||||
Transform DuckDuckGo API response to LiteLLM unified SearchResponse format.
|
||||
|
||||
DuckDuckGo → LiteLLM mappings:
|
||||
- RelatedTopics[].Text → SearchResult.title + snippet
|
||||
- RelatedTopics[].FirstURL → SearchResult.url
|
||||
- RelatedTopics[].Text → SearchResult.snippet
|
||||
- No date/last_updated fields in DuckDuckGo response (set to None)
|
||||
|
||||
Args:
|
||||
raw_response: Raw httpx response from DuckDuckGo API
|
||||
logging_obj: Logging object for tracking
|
||||
|
||||
Returns:
|
||||
SearchResponse with standardized format
|
||||
"""
|
||||
response_json = raw_response.json()
|
||||
|
||||
# Extract max_results from the request URL params
|
||||
query_params = raw_response.request.url.params if raw_response.request else {}
|
||||
max_results = None
|
||||
if "_max_results" in query_params:
|
||||
try:
|
||||
max_results = int(query_params["_max_results"])
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
# Transform results to SearchResult objects
|
||||
results = []
|
||||
|
||||
# DuckDuckGo can return results in different fields
|
||||
# Priority: Abstract > Answer > RelatedTopics
|
||||
|
||||
# Check if there's an Abstract with URL
|
||||
if response_json.get("AbstractURL") and response_json.get("AbstractText"):
|
||||
abstract_result = SearchResult(
|
||||
title=response_json.get("Heading", ""),
|
||||
url=response_json.get("AbstractURL", ""),
|
||||
snippet=response_json.get("AbstractText", ""),
|
||||
date=None,
|
||||
last_updated=None,
|
||||
)
|
||||
results.append(abstract_result)
|
||||
|
||||
# Process RelatedTopics
|
||||
related_topics = response_json.get("RelatedTopics", [])
|
||||
for topic in related_topics:
|
||||
# Stop if we've reached max_results
|
||||
if max_results is not None and len(results) >= max_results:
|
||||
break
|
||||
|
||||
if isinstance(topic, dict):
|
||||
# Check if it's a direct result
|
||||
if "FirstURL" in topic and "Text" in topic:
|
||||
text = topic.get("Text", "")
|
||||
url = topic.get("FirstURL", "")
|
||||
|
||||
# Try to split title and snippet
|
||||
if " - " in text:
|
||||
parts = text.split(" - ", 1)
|
||||
title = parts[0]
|
||||
snippet = parts[1] if len(parts) > 1 else text
|
||||
else:
|
||||
title = text[:50] + "..." if len(text) > 50 else text
|
||||
snippet = text
|
||||
|
||||
search_result = SearchResult(
|
||||
title=title,
|
||||
url=url,
|
||||
snippet=snippet,
|
||||
date=None,
|
||||
last_updated=None,
|
||||
)
|
||||
results.append(search_result)
|
||||
|
||||
# Check if it contains nested topics
|
||||
elif "Topics" in topic:
|
||||
nested_topics = topic.get("Topics", [])
|
||||
for nested_topic in nested_topics:
|
||||
# Stop if we've reached max_results
|
||||
if max_results is not None and len(results) >= max_results:
|
||||
break
|
||||
|
||||
if "FirstURL" in nested_topic and "Text" in nested_topic:
|
||||
text = nested_topic.get("Text", "")
|
||||
url = nested_topic.get("FirstURL", "")
|
||||
|
||||
# Try to split title and snippet
|
||||
if " - " in text:
|
||||
parts = text.split(" - ", 1)
|
||||
title = parts[0]
|
||||
snippet = parts[1] if len(parts) > 1 else text
|
||||
else:
|
||||
title = text[:50] + "..." if len(text) > 50 else text
|
||||
snippet = text
|
||||
|
||||
search_result = SearchResult(
|
||||
title=title,
|
||||
url=url,
|
||||
snippet=snippet,
|
||||
date=None,
|
||||
last_updated=None,
|
||||
)
|
||||
results.append(search_result)
|
||||
|
||||
return SearchResponse(
|
||||
results=results,
|
||||
object="search",
|
||||
)
|
||||
|
|
@ -770,14 +770,36 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
|
||||
|
||||
class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
|
||||
def _map_reasoning_to_reasoning_content(self, choices: list) -> list:
|
||||
"""
|
||||
Map 'reasoning' field to 'reasoning_content' field in delta.
|
||||
|
||||
Some OpenAI-compatible providers (e.g., GLM-5, hosted_vllm) return
|
||||
delta.reasoning, but LiteLLM expects delta.reasoning_content.
|
||||
|
||||
Args:
|
||||
choices: List of choice objects from the streaming chunk
|
||||
|
||||
Returns:
|
||||
List of choices with reasoning field mapped to reasoning_content
|
||||
"""
|
||||
for choice in choices:
|
||||
delta = choice.get("delta", {})
|
||||
if "reasoning" in delta:
|
||||
delta["reasoning_content"] = delta.pop("reasoning")
|
||||
return choices
|
||||
|
||||
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
|
||||
try:
|
||||
choices = chunk.get("choices", [])
|
||||
choices = self._map_reasoning_to_reasoning_content(choices)
|
||||
|
||||
kwargs = {
|
||||
"id": chunk["id"],
|
||||
"object": "chat.completion.chunk",
|
||||
"created": chunk.get("created"),
|
||||
"model": chunk.get("model"),
|
||||
"choices": chunk.get("choices", []),
|
||||
"choices": choices,
|
||||
}
|
||||
if "usage" in chunk and chunk["usage"] is not None:
|
||||
kwargs["usage"] = chunk["usage"]
|
||||
|
|
|
|||
|
|
@ -107,6 +107,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
guardrailed_tool_calls = guardrailed_inputs.get("tool_calls", [])
|
||||
guardrailed_tools = guardrailed_inputs.get("tools")
|
||||
if guardrailed_tools is not None:
|
||||
data["tools"] = guardrailed_tools
|
||||
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
if guardrailed_texts and texts_to_check:
|
||||
|
|
|
|||
|
|
@ -96,10 +96,11 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
# Handle simple string input
|
||||
if isinstance(input_data, str):
|
||||
inputs = GenericGuardrailAPIInputs(texts=[input_data])
|
||||
original_tools: List[Dict[str, Any]] = []
|
||||
|
||||
# Extract and transform tools if present
|
||||
|
||||
if "tools" in data and data["tools"]:
|
||||
original_tools = list(data["tools"])
|
||||
self._extract_and_transform_tools(data["tools"], tools_to_check)
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
|
|
@ -118,6 +119,9 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data
|
||||
self._apply_guardrailed_tools_to_data(
|
||||
data, original_tools, guardrailed_inputs.get("tools")
|
||||
)
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed string input")
|
||||
return data
|
||||
|
||||
|
|
@ -128,8 +132,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
# Track (message_index, content_index) for each text
|
||||
# content_index is None for string content, int for list content
|
||||
original_tools_list: List[Dict[str, Any]] = list(data.get("tools") or [])
|
||||
|
||||
# Step 1: Extract all text content, images, and tools
|
||||
for msg_idx, message in enumerate(input_data):
|
||||
|
|
@ -166,6 +169,11 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
self._apply_guardrailed_tools_to_data(
|
||||
data,
|
||||
original_tools_list,
|
||||
guardrailed_inputs.get("tools"),
|
||||
)
|
||||
|
||||
# Step 3: Map guardrail responses back to original input structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
|
|
@ -203,6 +211,53 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
cast(List[ChatCompletionToolParam], transformed_tools)
|
||||
)
|
||||
|
||||
def _remap_tools_to_responses_api_format(
|
||||
self, guardrailed_tools: List[Any]
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Remap guardrail-returned tools (Chat Completion format) back to
|
||||
Responses API request tool format.
|
||||
"""
|
||||
return LiteLLMCompletionResponsesConfig.transform_chat_completion_tool_params_to_responses_api_tools(
|
||||
guardrailed_tools # type: ignore
|
||||
)
|
||||
|
||||
def _merge_tools_after_guardrail(
|
||||
self,
|
||||
original_tools: List[Dict[str, Any]],
|
||||
remapped: List[Dict[str, Any]],
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Merge remapped guardrailed tools with original tools that were not sent
|
||||
to the guardrail (e.g. web_search, web_search_preview), preserving order.
|
||||
"""
|
||||
if not original_tools:
|
||||
return remapped
|
||||
result: List[Dict[str, Any]] = []
|
||||
j = 0
|
||||
for tool in original_tools:
|
||||
if isinstance(tool, dict) and tool.get("type") in (
|
||||
"web_search",
|
||||
"web_search_preview",
|
||||
):
|
||||
result.append(tool)
|
||||
else:
|
||||
if j < len(remapped):
|
||||
result.append(remapped[j])
|
||||
j += 1
|
||||
return result
|
||||
|
||||
def _apply_guardrailed_tools_to_data(
|
||||
self,
|
||||
data: dict,
|
||||
original_tools: List[Dict[str, Any]],
|
||||
guardrailed_tools: Optional[List[Any]],
|
||||
) -> None:
|
||||
"""Remap guardrailed tools to Responses API format and merge with original, then set data['tools']."""
|
||||
if guardrailed_tools is not None:
|
||||
remapped = self._remap_tools_to_responses_api_format(guardrailed_tools)
|
||||
data["tools"] = self._merge_tools_after_guardrail(original_tools, remapped)
|
||||
|
||||
def _extract_input_text_and_images(
|
||||
self,
|
||||
message: Any, # Can be Dict[str, Any] or ResponseInputParam
|
||||
|
|
@ -407,7 +462,10 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
List[ChatCompletionToolCallChunk], tool_calls
|
||||
)
|
||||
# Include model information if available
|
||||
if hasattr(model_response_stream, "model") and model_response_stream.model:
|
||||
if (
|
||||
hasattr(model_response_stream, "model")
|
||||
and model_response_stream.model
|
||||
):
|
||||
inputs["model"] = model_response_stream.model
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
|
|
@ -448,7 +506,9 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
)
|
||||
return responses_so_far
|
||||
else:
|
||||
verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping output guardrail - model response has no choices"
|
||||
)
|
||||
# model_response_stream = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(final_chunk)
|
||||
# tool_calls = model_response_stream.choices[0].tool_calls
|
||||
# convert openai response to model response
|
||||
|
|
@ -456,7 +516,11 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
inputs = GenericGuardrailAPIInputs(texts=[string_so_far])
|
||||
# Try to get model from the final chunk if available
|
||||
if isinstance(final_chunk, dict):
|
||||
response_model = final_chunk.get("response", {}).get("model") if isinstance(final_chunk.get("response"), dict) else None
|
||||
response_model = (
|
||||
final_chunk.get("response", {}).get("model")
|
||||
if isinstance(final_chunk.get("response"), dict)
|
||||
else None
|
||||
)
|
||||
if response_model:
|
||||
inputs["model"] = response_model
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
|
|
@ -591,8 +655,8 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
content = generic_response_output_item.content
|
||||
except Exception:
|
||||
# Try to extract content directly from output_item if validation fails
|
||||
if hasattr(output_item, "content") and output_item.content:
|
||||
content = output_item.content
|
||||
if hasattr(output_item, "content") and output_item.content: # type: ignore
|
||||
content = output_item.content # type: ignore
|
||||
else:
|
||||
return
|
||||
elif isinstance(output_item, dict):
|
||||
|
|
@ -669,10 +733,10 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
if isinstance(content_item, OutputText):
|
||||
content_item.text = guardrail_response
|
||||
# Update the original response output
|
||||
if hasattr(output_item, "content") and output_item.content:
|
||||
original_content = output_item.content[content_idx]
|
||||
if hasattr(output_item, "content") and output_item.content: # type: ignore
|
||||
original_content = output_item.content[content_idx] # type: ignore
|
||||
if hasattr(original_content, "text"):
|
||||
original_content.text = guardrail_response
|
||||
original_content.text = guardrail_response # type: ignore
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(output_item, dict):
|
||||
|
|
|
|||
|
|
@ -1072,7 +1072,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
elif param == "modalities" and isinstance(value, list):
|
||||
response_modalities = self.map_response_modalities(value)
|
||||
optional_params["responseModalities"] = response_modalities
|
||||
elif param == "web_search_options" and value and isinstance(value, dict):
|
||||
elif param == "web_search_options" and isinstance(value, dict):
|
||||
_tools = self._map_web_search_options(value)
|
||||
optional_params = self._add_tools_to_optional_params(
|
||||
optional_params, [_tools]
|
||||
|
|
|
|||
0
litellm/llms/watsonx/__init__.py
Normal file
0
litellm/llms/watsonx/__init__.py
Normal file
0
litellm/llms/watsonx/chat/__init__.py
Normal file
0
litellm/llms/watsonx/chat/__init__.py
Normal file
0
litellm/llms/watsonx/completion/__init__.py
Normal file
0
litellm/llms/watsonx/completion/__init__.py
Normal file
0
litellm/llms/watsonx/embed/__init__.py
Normal file
0
litellm/llms/watsonx/embed/__init__.py
Normal file
0
litellm/llms/watsonx/rerank/__init__.py
Normal file
0
litellm/llms/watsonx/rerank/__init__.py
Normal file
204
litellm/llms/watsonx/rerank/transformation.py
Normal file
204
litellm/llms/watsonx/rerank/transformation.py
Normal file
|
|
@ -0,0 +1,204 @@
|
|||
"""
|
||||
Transformation logic for IBM watsonx.ai's /ml/v1/text/rerank endpoint.
|
||||
|
||||
Docs - https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.watsonx import (
|
||||
WatsonXAIEndpoint,
|
||||
)
|
||||
from litellm.types.rerank import (
|
||||
RerankResponse,
|
||||
RerankResponseMeta,
|
||||
RerankTokens,
|
||||
)
|
||||
|
||||
from ..common_utils import IBMWatsonXMixin, _generate_watsonx_token, _get_api_params
|
||||
|
||||
|
||||
class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
|
||||
"""
|
||||
IBM watsonx.ai Rerank API configuration
|
||||
"""
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
model: str,
|
||||
optional_params: Optional[dict] = None,
|
||||
) -> str:
|
||||
base_url = self._get_base_url(api_base=api_base)
|
||||
endpoint = WatsonXAIEndpoint.RERANK.value
|
||||
|
||||
url = base_url.rstrip("/") + endpoint
|
||||
|
||||
params = optional_params or {}
|
||||
|
||||
complete_url = self._add_api_version_to_url(url=url, api_version=(params.get("api_version", None)))
|
||||
return complete_url
|
||||
|
||||
def get_supported_cohere_rerank_params(self, model: str) -> list:
|
||||
return [
|
||||
"query",
|
||||
"documents",
|
||||
"top_n",
|
||||
"return_documents",
|
||||
"max_tokens_per_doc",
|
||||
]
|
||||
|
||||
def validate_environment( # type: ignore[override]
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
optional_params: Optional[dict] = None,
|
||||
) -> Dict:
|
||||
optional_params = optional_params or {}
|
||||
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
|
||||
if "Authorization" in headers:
|
||||
return {**default_headers, **headers}
|
||||
token = cast(
|
||||
Optional[str],
|
||||
optional_params.pop("token", None) or get_secret_str("WATSONX_TOKEN"),
|
||||
)
|
||||
zen_api_key = cast(
|
||||
Optional[str],
|
||||
optional_params.pop("zen_api_key", None) or get_secret_str("WATSONX_ZENAPIKEY"),
|
||||
)
|
||||
if token:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
elif zen_api_key:
|
||||
headers["Authorization"] = f"ZenApiKey {zen_api_key}"
|
||||
else:
|
||||
token = _generate_watsonx_token(api_key=api_key, token=token)
|
||||
# build auth headers
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
return {**default_headers, **headers}
|
||||
|
||||
def map_cohere_rerank_params(
|
||||
self,
|
||||
non_default_params: Optional[dict],
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
query: str,
|
||||
documents: List[Union[str, Dict[str, Any]]],
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
top_n: Optional[int] = None,
|
||||
rank_fields: Optional[List[str]] = None,
|
||||
return_documents: Optional[bool] = True,
|
||||
max_chunks_per_doc: Optional[int] = None,
|
||||
max_tokens_per_doc: Optional[int] = None,
|
||||
) -> Dict:
|
||||
"""
|
||||
Map Cohere rerank params to IBM watsonx.ai rerank params
|
||||
"""
|
||||
optional_rerank_params = {}
|
||||
if non_default_params is not None:
|
||||
for k, v in non_default_params.items():
|
||||
if k == "query" and v is not None:
|
||||
optional_rerank_params["query"] = v
|
||||
elif k == "documents" and v is not None:
|
||||
optional_rerank_params["inputs"] = [
|
||||
{"text": el} if isinstance(el, str) else el for el in v
|
||||
]
|
||||
elif k == "top_n" and v is not None:
|
||||
optional_rerank_params.setdefault("parameters", {}).setdefault("return_options", {})["top_n"] = v
|
||||
elif k == "return_documents" and v is not None and isinstance(v, bool):
|
||||
optional_rerank_params.setdefault("parameters", {}).setdefault("return_options", {})["inputs"] = v
|
||||
elif k == "max_tokens_per_doc" and v is not None:
|
||||
optional_rerank_params.setdefault("parameters", {})["truncate_input_tokens"] = v
|
||||
|
||||
# IBM watsonx.ai require one of below parameters
|
||||
elif k == "project_id" and v is not None:
|
||||
optional_rerank_params["project_id"] = v
|
||||
elif k == "space_id" and v is not None:
|
||||
optional_rerank_params["space_id"] = v
|
||||
|
||||
return dict(optional_rerank_params)
|
||||
|
||||
def transform_rerank_request(
|
||||
self,
|
||||
model: str,
|
||||
optional_rerank_params: Dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform request to IBM watsonx.ai rerank format
|
||||
"""
|
||||
watsonx_api_params = _get_api_params(params=optional_rerank_params, model=model)
|
||||
watsonx_auth_payload = self._prepare_payload(
|
||||
model=model,
|
||||
api_params=watsonx_api_params,
|
||||
)
|
||||
|
||||
return optional_rerank_params | watsonx_auth_payload
|
||||
|
||||
def transform_rerank_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: RerankResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str] = None,
|
||||
request_data: dict = {},
|
||||
optional_params: dict = {},
|
||||
litellm_params: dict = {},
|
||||
) -> RerankResponse:
|
||||
"""
|
||||
Transform IBM watsonx.ai rerank response to LiteLLM RerankResponse format
|
||||
"""
|
||||
try:
|
||||
raw_response_json = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Failed to parse response: {str(e)}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
_results: Optional[List[dict]] = raw_response_json.get("results")
|
||||
if _results is None:
|
||||
raise ValueError(f"No results found in the response={raw_response_json}")
|
||||
|
||||
transformed_results = []
|
||||
|
||||
for result in _results:
|
||||
transformed_result: Dict[str, Any] = {
|
||||
"index": result["index"],
|
||||
"relevance_score": result["score"],
|
||||
}
|
||||
|
||||
if "input" in result:
|
||||
if isinstance(result["input"], str):
|
||||
transformed_result["document"] = {"text": result["input"]}
|
||||
else:
|
||||
transformed_result["document"] = result["input"]
|
||||
|
||||
transformed_results.append(transformed_result)
|
||||
|
||||
response_id = raw_response_json.get("id") or raw_response_json.get("model_id") or str(uuid.uuid4())
|
||||
|
||||
# Extract usage information
|
||||
_tokens = RerankTokens(
|
||||
input_tokens=raw_response_json.get("input_token_count", 0),
|
||||
)
|
||||
rerank_meta = RerankResponseMeta(tokens=_tokens)
|
||||
|
||||
return RerankResponse(
|
||||
id=response_id,
|
||||
results=transformed_results, # type: ignore
|
||||
meta=rerank_meta,
|
||||
)
|
||||
|
|
@ -8294,6 +8294,37 @@
|
|||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 346
|
||||
},
|
||||
"us/claude-sonnet-4-6": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6.6e-06,
|
||||
"litellm_provider": "anthropic",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 64000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 2.475e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 346,
|
||||
"inference_geo": "us"
|
||||
},
|
||||
"claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
|
|
@ -22465,6 +22496,20 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/devstral-small-latest": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-07,
|
||||
"source": "https://docs.mistral.ai/models/devstral-small-2-25-12",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/labs-devstral-small-2512": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "mistral",
|
||||
|
|
@ -22479,6 +22524,34 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/devstral-latest": {
|
||||
"input_cost_per_token": 4e-07,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-06,
|
||||
"source": "https://mistral.ai/news/devstral-2-vibe-cli",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/devstral-medium-latest": {
|
||||
"input_cost_per_token": 4e-07,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-06,
|
||||
"source": "https://mistral.ai/news/devstral-2-vibe-cli",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/devstral-2512": {
|
||||
"input_cost_per_token": 4e-07,
|
||||
"litellm_provider": "mistral",
|
||||
|
|
@ -37270,5 +37343,13 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
}
|
||||
},
|
||||
"duckduckgo/search": {
|
||||
"litellm_provider": "duckduckgo",
|
||||
"mode": "search",
|
||||
"input_cost_per_query": 0.0,
|
||||
"metadata": {
|
||||
"notes": "DuckDuckGo Instant Answer API is free and does not require an API key."
|
||||
}
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -336,15 +336,21 @@ class MCPRequestHandler:
|
|||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get list of allowed MCP servers for the given user/key based on permissions
|
||||
Get list of allowed MCP servers for the given user/key based on permissions.
|
||||
|
||||
Permission hierarchy (all rules are intersections):
|
||||
1. Get allowed servers from key permissions
|
||||
2. Get allowed servers from team permissions
|
||||
3. Get allowed servers from end_user permissions
|
||||
4. Final result = intersection of key/team AND end_user (if end_user has permissions set)
|
||||
|
||||
Returns:
|
||||
List[str]: List of allowed MCP servers by server id
|
||||
"""
|
||||
from typing import List
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
try:
|
||||
allowed_mcp_servers: List[str] = []
|
||||
# Get allowed servers from key and team
|
||||
allowed_mcp_servers_for_key = (
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth
|
||||
|
|
@ -357,8 +363,9 @@ class MCPRequestHandler:
|
|||
)
|
||||
|
||||
#########################################################
|
||||
# If team has mcp_servers, handle inheritance and intersection logic
|
||||
# Calculate key/team allowed servers using inheritance and intersection logic
|
||||
#########################################################
|
||||
allowed_mcp_servers: List[str] = []
|
||||
if len(allowed_mcp_servers_for_team) > 0:
|
||||
if len(allowed_mcp_servers_for_key) > 0:
|
||||
# Key has its own MCP permissions - use intersection with team permissions
|
||||
|
|
@ -371,6 +378,40 @@ class MCPRequestHandler:
|
|||
else:
|
||||
allowed_mcp_servers = allowed_mcp_servers_for_key
|
||||
|
||||
#########################################################
|
||||
# Check end_user permissions if end_user_id is set
|
||||
#########################################################
|
||||
if user_api_key_auth and user_api_key_auth.end_user_id:
|
||||
allowed_mcp_servers_for_end_user = (
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_end_user(
|
||||
user_api_key_auth
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# If end_user has explicit MCP server permissions, apply intersection
|
||||
if len(allowed_mcp_servers_for_end_user) > 0:
|
||||
verbose_logger.debug(
|
||||
f"End user {user_api_key_auth.end_user_id} has explicit MCP permissions: {allowed_mcp_servers_for_end_user}"
|
||||
)
|
||||
|
||||
# Always apply intersection: key/team AND end_user
|
||||
# This ensures end_user can only access servers that both they AND their key/team are authorized for
|
||||
filtered_servers = []
|
||||
for _mcp_server in allowed_mcp_servers:
|
||||
if _mcp_server in allowed_mcp_servers_for_end_user:
|
||||
filtered_servers.append(_mcp_server)
|
||||
allowed_mcp_servers = filtered_servers
|
||||
verbose_logger.debug(
|
||||
f"Applied end_user intersection filter. Final allowed servers: {allowed_mcp_servers}"
|
||||
)
|
||||
# If flag is enabled but end_user has no permissions, block all access
|
||||
elif general_settings.get("require_end_user_mcp_access_defined", False):
|
||||
verbose_logger.debug(
|
||||
f"require_end_user_mcp_access_defined=True and end_user {user_api_key_auth.end_user_id} has no MCP permissions - blocking MCP access"
|
||||
)
|
||||
return []
|
||||
|
||||
return list(set(allowed_mcp_servers))
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
|
||||
|
|
@ -614,6 +655,66 @@ class MCPRequestHandler:
|
|||
)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_end_user(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get allowed MCP servers for an end user.
|
||||
|
||||
Returns the MCP servers from the end_user's object_permission.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_end_user_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not user_api_key_auth or not user_api_key_auth.end_user_id:
|
||||
return []
|
||||
|
||||
if prisma_client is None:
|
||||
|
||||
verbose_logger.debug("prisma_client is None")
|
||||
return []
|
||||
|
||||
try:
|
||||
# Use optimized get_end_user_object function with caching
|
||||
end_user_obj = await get_end_user_object(
|
||||
end_user_id=user_api_key_auth.end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
route="/mcp",
|
||||
)
|
||||
|
||||
|
||||
if end_user_obj is None or end_user_obj.object_permission is None:
|
||||
return []
|
||||
|
||||
# Get direct MCP servers
|
||||
direct_mcp_servers = end_user_obj.object_permission.mcp_servers or []
|
||||
|
||||
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = (
|
||||
await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
end_user_obj.object_permission.mcp_access_groups or []
|
||||
)
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
all_servers = direct_mcp_servers + access_group_servers
|
||||
return list(set(all_servers))
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Failed to get allowed MCP servers for end_user: {str(e)}"
|
||||
)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _get_config_server_ids_for_access_groups(
|
||||
config_mcp_servers, access_groups: List[str]
|
||||
|
|
@ -691,8 +792,6 @@ class MCPRequestHandler:
|
|||
"""
|
||||
Get list of MCP access groups for the given user/key based on permissions
|
||||
"""
|
||||
from typing import List
|
||||
|
||||
access_groups: List[str] = []
|
||||
access_groups_for_key = await MCPRequestHandler._get_mcp_access_groups_for_key(
|
||||
user_api_key_auth
|
||||
|
|
|
|||
|
|
@ -43,9 +43,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
|||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
LiteLLMProxyRequestSetup,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
|
||||
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
|
||||
|
|
@ -795,6 +793,7 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
)
|
||||
|
||||
|
||||
return allowed_mcp_servers
|
||||
|
||||
|
|
@ -938,9 +937,6 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
)
|
||||
|
||||
# Decide whether to add prefix based on number of allowed servers
|
||||
add_prefix = not (len(allowed_mcp_servers) == 1)
|
||||
|
||||
async def _fetch_and_filter_server_tools(
|
||||
server: MCPServer,
|
||||
) -> List[MCPTool]:
|
||||
|
|
@ -961,7 +957,7 @@ if MCP_AVAILABLE:
|
|||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
add_prefix=True, # Always add server prefix
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
|
@ -1079,8 +1075,6 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
)
|
||||
|
||||
# Decide whether to add prefix based on number of allowed servers
|
||||
add_prefix = not (len(allowed_mcp_servers) == 1)
|
||||
|
||||
# Get prompts from each allowed server
|
||||
all_prompts = []
|
||||
|
|
@ -1101,7 +1095,7 @@ if MCP_AVAILABLE:
|
|||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
add_prefix=True, # Always add server prefix
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
|
|
@ -1140,7 +1134,6 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
)
|
||||
|
||||
add_prefix = not (len(allowed_mcp_servers) == 1)
|
||||
|
||||
all_resources: List[Resource] = []
|
||||
for server in allowed_mcp_servers:
|
||||
|
|
@ -1160,7 +1153,7 @@ if MCP_AVAILABLE:
|
|||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
add_prefix=True, # Always add server prefix
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
all_resources.extend(resources)
|
||||
|
|
@ -1197,7 +1190,6 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
)
|
||||
|
||||
add_prefix = not (len(allowed_mcp_servers) == 1)
|
||||
|
||||
all_resource_templates: List[ResourceTemplate] = []
|
||||
for server in allowed_mcp_servers:
|
||||
|
|
@ -1218,7 +1210,7 @@ if MCP_AVAILABLE:
|
|||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
add_prefix=True, # Always add server prefix
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
)
|
||||
|
|
@ -1676,14 +1668,9 @@ if MCP_AVAILABLE:
|
|||
detail="User not allowed to get this prompt.",
|
||||
)
|
||||
|
||||
# Decide whether to add prefix based on number of allowed servers
|
||||
add_prefix = not (len(allowed_mcp_servers) == 1)
|
||||
|
||||
if add_prefix:
|
||||
original_prompt_name, server_name = split_server_prefix_from_name(name)
|
||||
else:
|
||||
original_prompt_name = name
|
||||
server_name = allowed_mcp_servers[0].name
|
||||
# Extract server name from prefixed prompt name
|
||||
original_prompt_name, server_name = split_server_prefix_from_name(name)
|
||||
|
||||
server = next((s for s in allowed_mcp_servers if s.name == server_name), None)
|
||||
if server is None:
|
||||
|
|
|
|||
|
|
@ -13,26 +13,25 @@ model_list:
|
|||
- model_name: gpt-4.1-mini
|
||||
litellm_params:
|
||||
model: openai/gpt-4.1-mini
|
||||
|
||||
|
||||
# guardrails:
|
||||
# - guardrail_name: generic-guardrail
|
||||
# litellm_params:
|
||||
# guardrail: generic_guardrail_api
|
||||
# mode: ["pre_call"]
|
||||
# headers:
|
||||
# Authorization: Bearer mock-bedrock-token-12345
|
||||
# api_base: http://localhost:8080
|
||||
# default_on: true
|
||||
|
||||
prompts:
|
||||
- prompt_id: "simple_prompt"
|
||||
- model_name: gpt-5-mini
|
||||
litellm_params:
|
||||
prompt_integration: "generic_prompt_management"
|
||||
provider_specific_query_params:
|
||||
project_name: litellm
|
||||
slug: hello-world-prompt-2bac
|
||||
api_base: http://localhost:8080
|
||||
api_key: os.environ/BRAINTRUST_API_KEY
|
||||
ignore_prompt_manager_model: true
|
||||
ignore_prompt_manager_optional_params: true
|
||||
model: openai/gpt-5-mini
|
||||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: mcp-user-permissions
|
||||
litellm_params:
|
||||
guardrail: mcp_end_user_permission
|
||||
mode: pre_call
|
||||
default_on: true
|
||||
|
||||
mcp_servers:
|
||||
my_http_server:
|
||||
url: "http://0.0.0.0:8001/mcp"
|
||||
transport: "http"
|
||||
description: "My custom MCP server"
|
||||
available_on_public_internet: true
|
||||
|
||||
general_settings:
|
||||
store_model_in_db: true
|
||||
store_prompts_in_spend_logs: true
|
||||
|
|
|
|||
|
|
@ -1409,12 +1409,13 @@ class NewCustomerRequest(BudgetNewRequest):
|
|||
blocked: bool = False # allow/disallow requests for this end-user
|
||||
budget_id: Optional[str] = None # give either a budget_id or max_budget
|
||||
spend: Optional[float] = None
|
||||
allowed_model_region: Optional[
|
||||
AllowedModelRegion
|
||||
] = None # require all user requests to use models in this specific region
|
||||
default_model: Optional[
|
||||
str
|
||||
] = None # if no equivalent model in allowed region - default all requests to this model
|
||||
allowed_model_region: Optional[AllowedModelRegion] = (
|
||||
None # require all user requests to use models in this specific region
|
||||
)
|
||||
default_model: Optional[str] = (
|
||||
None # if no equivalent model in allowed region - default all requests to this model
|
||||
)
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
@ -1436,12 +1437,13 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
|
|||
blocked: bool = False # allow/disallow requests for this end-user
|
||||
max_budget: Optional[float] = None
|
||||
budget_id: Optional[str] = None # give either a budget_id or max_budget
|
||||
allowed_model_region: Optional[
|
||||
AllowedModelRegion
|
||||
] = None # require all user requests to use models in this specific region
|
||||
default_model: Optional[
|
||||
str
|
||||
] = None # if no equivalent model in allowed region - default all requests to this model
|
||||
allowed_model_region: Optional[AllowedModelRegion] = (
|
||||
None # require all user requests to use models in this specific region
|
||||
)
|
||||
default_model: Optional[str] = (
|
||||
None # if no equivalent model in allowed region - default all requests to this model
|
||||
)
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
|
||||
|
||||
|
||||
class DeleteCustomerRequest(LiteLLMPydanticObjectBase):
|
||||
|
|
@ -2125,6 +2127,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For headers are only trusted from these IPs.",
|
||||
)
|
||||
store_model_in_db: Optional[bool] = Field(
|
||||
None,
|
||||
description="If True, models and config are stored in and loaded from the database. Default is False.",
|
||||
)
|
||||
|
||||
|
||||
class ConfigYAML(LiteLLMPydanticObjectBase):
|
||||
|
|
@ -2297,6 +2303,7 @@ class UserAPIKeyAuth(
|
|||
user_max_budget: Optional[float] = None
|
||||
request_route: Optional[str] = None
|
||||
user: Optional[Any] = None # Expanded user object when expand=user is used
|
||||
end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
|
|
@ -2531,6 +2538,8 @@ class LiteLLM_EndUserTable(LiteLLMPydanticObjectBase):
|
|||
allowed_model_region: Optional[AllowedModelRegion] = None
|
||||
default_model: Optional[str] = None
|
||||
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
|
||||
object_permission_id: Optional[str] = None
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -11,7 +11,8 @@ Run checks for:
|
|||
import asyncio
|
||||
import re
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
|
||||
from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union,
|
||||
cast)
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -20,41 +21,27 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.dual_cache import LimitedSizeOrderedDict
|
||||
from litellm.constants import (
|
||||
CLI_JWT_EXPIRATION_HOURS,
|
||||
CLI_JWT_TOKEN_NAME,
|
||||
DEFAULT_ACCESS_GROUP_CACHE_TTL,
|
||||
DEFAULT_IN_MEMORY_TTL,
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE,
|
||||
)
|
||||
from litellm.constants import (CLI_JWT_EXPIRATION_HOURS, CLI_JWT_TOKEN_NAME,
|
||||
DEFAULT_ACCESS_GROUP_CACHE_TTL,
|
||||
DEFAULT_IN_MEMORY_TTL,
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE)
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.proxy._types import (
|
||||
RBAC_ROLES,
|
||||
CallInfo,
|
||||
LiteLLM_AccessGroupTable,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_EndUserTable,
|
||||
Litellm_EntityType,
|
||||
LiteLLM_JWTAuth,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_OrganizationMembershipTable,
|
||||
LiteLLM_OrganizationTable,
|
||||
LiteLLM_TagTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
LiteLLMRoutes,
|
||||
LitellmUserRoles,
|
||||
NewTeamRequest,
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
RoleBasedPermissions,
|
||||
SpecialModelNames,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy._types import (RBAC_ROLES, CallInfo,
|
||||
LiteLLM_AccessGroupTable,
|
||||
LiteLLM_BudgetTable, LiteLLM_EndUserTable,
|
||||
Litellm_EntityType, LiteLLM_JWTAuth,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_OrganizationMembershipTable,
|
||||
LiteLLM_OrganizationTable, LiteLLM_TagTable,
|
||||
LiteLLM_TeamMembership, LiteLLM_TeamTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable, LiteLLMRoutes,
|
||||
LitellmUserRoles, NewTeamRequest,
|
||||
ProxyErrorTypes, ProxyException,
|
||||
RoleBasedPermissions, SpecialModelNames,
|
||||
UserAPIKeyAuth)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
|
||||
|
|
@ -366,7 +353,8 @@ async def common_checks(
|
|||
_request_metadata: dict = request_body.get("metadata", {}) or {}
|
||||
if _request_metadata.get("guardrails"):
|
||||
# check if team allowed to modify guardrails
|
||||
from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails
|
||||
from litellm.proxy.guardrails.guardrail_helpers import \
|
||||
can_modify_guardrails
|
||||
|
||||
can_modify: bool = can_modify_guardrails(team_object)
|
||||
if can_modify is False:
|
||||
|
|
@ -792,7 +780,7 @@ async def get_end_user_object(
|
|||
try:
|
||||
response = await prisma_client.db.litellm_endusertable.find_unique(
|
||||
where={"user_id": end_user_id},
|
||||
include={"litellm_budget_table": True},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
||||
if response is None:
|
||||
|
|
@ -1812,9 +1800,8 @@ class ExperimentalUIJWTToken:
|
|||
def get_experimental_ui_login_jwt_auth_token(user_info: LiteLLM_UserTable) -> str:
|
||||
from datetime import timedelta
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import \
|
||||
encrypt_value_helper
|
||||
|
||||
if user_info.user_role is None:
|
||||
raise Exception("User role is required for experimental UI login")
|
||||
|
|
@ -1860,9 +1847,8 @@ class ExperimentalUIJWTToken:
|
|||
"""
|
||||
from datetime import timedelta
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import \
|
||||
encrypt_value_helper
|
||||
|
||||
if user_info.user_role is None:
|
||||
raise Exception("User role is required for CLI JWT login")
|
||||
|
|
@ -1901,9 +1887,8 @@ class ExperimentalUIJWTToken:
|
|||
import json
|
||||
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
decrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import \
|
||||
decrypt_value_helper
|
||||
|
||||
decrypted_token = decrypt_value_helper(
|
||||
hashed_token, key="ui_hash_key", exception_type="debug"
|
||||
|
|
@ -2150,11 +2135,11 @@ async def _get_resources_from_access_groups(
|
|||
|
||||
# Lazy import to avoid circular imports
|
||||
if prisma_client is None or user_api_key_cache is None:
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client as _prisma_client,
|
||||
proxy_logging_obj as _proxy_logging_obj,
|
||||
user_api_key_cache as _user_api_key_cache,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client as _prisma_client
|
||||
from litellm.proxy.proxy_server import \
|
||||
proxy_logging_obj as _proxy_logging_obj
|
||||
from litellm.proxy.proxy_server import \
|
||||
user_api_key_cache as _user_api_key_cache
|
||||
|
||||
prisma_client = prisma_client or _prisma_client
|
||||
user_api_key_cache = user_api_key_cache or _user_api_key_cache
|
||||
|
|
@ -2936,7 +2921,8 @@ async def _tag_max_budget_check(
|
|||
BudgetExceededError if any tag is over its max budget.
|
||||
Triggers a budget alert if any tag is over its max budget.
|
||||
"""
|
||||
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
|
||||
from litellm.proxy.common_utils.http_parsing_utils import \
|
||||
get_tags_from_request_body
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -736,6 +736,16 @@ def get_end_user_id_from_request_body(
|
|||
user_id_from_metadata_field = metadata_dict.get("user_id")
|
||||
if user_id_from_metadata_field is not None:
|
||||
return str(user_id_from_metadata_field)
|
||||
|
||||
|
||||
# Check 6: 'safety_identifier' in request body (OpenAI Responses API parameter)
|
||||
# SECURITY NOTE: safety_identifier can be set by any caller in the request body.
|
||||
# Only use this for end-user identification in trusted environments where you control
|
||||
# the calling application. For untrusted callers, prefer using headers or server-side
|
||||
# middleware to set the end_user_id to prevent impersonation.
|
||||
if request_body.get("safety_identifier") is not None:
|
||||
user_from_body_user_field = request_body["safety_identifier"]
|
||||
return str(user_from_body_user_field)
|
||||
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -223,12 +223,14 @@ def get_known_models_from_wildcard(
|
|||
except ValueError: # safely fail
|
||||
return []
|
||||
|
||||
if litellm_params is None: # need litellm params to extract litellm model name
|
||||
return []
|
||||
|
||||
try:
|
||||
provider = litellm_params.model.split("/", 1)[0]
|
||||
except ValueError:
|
||||
# Use provider from litellm_params when available, otherwise from wildcard prefix
|
||||
# (e.g., "openai" from "openai/*" - needed for BYOK where wildcard isn't in router)
|
||||
if litellm_params is not None:
|
||||
try:
|
||||
provider = litellm_params.model.split("/", 1)[0]
|
||||
except ValueError:
|
||||
provider = wildcard_provider_prefix
|
||||
else:
|
||||
provider = wildcard_provider_prefix
|
||||
|
||||
# get all known provider models
|
||||
|
|
@ -282,7 +284,7 @@ def _get_wildcard_models(
|
|||
## get litellm params from model
|
||||
if llm_router is not None:
|
||||
model_list = llm_router.get_model_list(model_name=model)
|
||||
if model_list is not None:
|
||||
if model_list:
|
||||
for router_model in model_list:
|
||||
wildcard_models = get_known_models_from_wildcard(
|
||||
wildcard_model=model,
|
||||
|
|
@ -291,11 +293,22 @@ def _get_wildcard_models(
|
|||
),
|
||||
)
|
||||
all_wildcard_models.extend(wildcard_models)
|
||||
else:
|
||||
# Router has no deployment for this wildcard (e.g., BYOK team models)
|
||||
# Fall back to expanding from known provider models
|
||||
wildcard_models = get_known_models_from_wildcard(
|
||||
wildcard_model=model, litellm_params=None
|
||||
)
|
||||
if wildcard_models:
|
||||
models_to_remove.add(model)
|
||||
all_wildcard_models.extend(wildcard_models)
|
||||
else:
|
||||
# get all known provider models
|
||||
wildcard_models = get_known_models_from_wildcard(wildcard_model=model)
|
||||
wildcard_models = get_known_models_from_wildcard(
|
||||
wildcard_model=model, litellm_params=None
|
||||
)
|
||||
|
||||
if wildcard_models is not None:
|
||||
if wildcard_models:
|
||||
models_to_remove.add(model)
|
||||
all_wildcard_models.extend(wildcard_models)
|
||||
|
||||
|
|
|
|||
|
|
@ -307,6 +307,9 @@ class RouteChecks:
|
|||
):
|
||||
return True
|
||||
|
||||
if route in LiteLLMRoutes.litellm_native_routes.value:
|
||||
return True
|
||||
|
||||
# fuzzy match routes like "/v1/threads/thread_49EIN5QF32s4mH20M7GFKdlZ"
|
||||
# Check for routes with placeholders or wildcard patterns
|
||||
for openai_route in LiteLLMRoutes.openai_routes.value:
|
||||
|
|
|
|||
|
|
@ -644,7 +644,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
if team_object is not None
|
||||
else None,
|
||||
)
|
||||
|
||||
|
||||
# Check if model has zero cost - if so, skip all budget checks
|
||||
model = get_model_from_request(request_data, route)
|
||||
skip_budget_checks = False
|
||||
|
|
@ -831,6 +831,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
valid_token=valid_token, end_user_params=end_user_params
|
||||
)
|
||||
valid_token.parent_otel_span = parent_otel_span
|
||||
if _end_user_object is not None:
|
||||
valid_token.end_user_object_permission = _end_user_object.object_permission
|
||||
|
||||
return valid_token
|
||||
|
||||
|
|
@ -1277,6 +1279,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
|
||||
if _end_user_object is not None:
|
||||
valid_token_dict.update(end_user_params)
|
||||
valid_token_dict["end_user_object_permission"] = (
|
||||
_end_user_object.object_permission
|
||||
)
|
||||
|
||||
# check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions
|
||||
# sso/login, ui/login, /key functions and /user functions
|
||||
|
|
|
|||
|
|
@ -18,6 +18,9 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
additional_provider_specific_params=getattr(
|
||||
litellm_params, "additional_provider_specific_params", {}
|
||||
),
|
||||
unreachable_fallback=getattr(
|
||||
litellm_params, "unreachable_fallback", "fail_closed"
|
||||
),
|
||||
guardrail_name=guardrail.get("guardrail_name", ""),
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ litellm_settings:
|
|||
mode: pre_call # Options: pre_call, post_call, during_call, [pre_call, post_call]
|
||||
api_key: os.environ/GENERIC_GUARDRAIL_API_KEY # Optional if using Bearer auth
|
||||
api_base: http://localhost:8080 # Required. Endpoint /beta/litellm_basic_guardrail_api is automatically appended
|
||||
unreachable_fallback: fail_closed # Options: fail_closed (default, raise), fail_open (proceed if endpoint unreachable or upstream returns 502/503/504)
|
||||
default_on: false # Set to true to apply to all requests by default
|
||||
additional_provider_specific_params:
|
||||
# Any additional parameters your guardrail API needs
|
||||
|
|
|
|||
|
|
@ -9,9 +9,11 @@ import fnmatch
|
|||
import os
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
from litellm.exceptions import GuardrailRaisedException, Timeout
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
|
|
@ -34,17 +36,19 @@ if TYPE_CHECKING:
|
|||
GUARDRAIL_NAME = "generic_guardrail_api"
|
||||
|
||||
# Headers whose values are forwarded as-is (case-insensitive). Glob patterns supported (e.g. x-stainless-*, x-litellm*).
|
||||
_HEADER_VALUE_ALLOWLIST = frozenset({
|
||||
"host",
|
||||
"accept-encoding",
|
||||
"connection",
|
||||
"accept",
|
||||
"content-type",
|
||||
"user-agent",
|
||||
"x-stainless-*",
|
||||
"x-litellm-*",
|
||||
"content-length",
|
||||
})
|
||||
_HEADER_VALUE_ALLOWLIST = frozenset(
|
||||
{
|
||||
"host",
|
||||
"accept-encoding",
|
||||
"connection",
|
||||
"accept",
|
||||
"content-type",
|
||||
"user-agent",
|
||||
"x-stainless-*",
|
||||
"x-litellm-*",
|
||||
"content-length",
|
||||
}
|
||||
)
|
||||
|
||||
# Placeholder for headers that exist but are not on the allowlist (we don't expose their value).
|
||||
_HEADER_PRESENT_PLACEHOLDER = "[present]"
|
||||
|
|
@ -166,6 +170,7 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
api_base: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
additional_provider_specific_params: Optional[Dict[str, Any]] = None,
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
|
||||
**kwargs,
|
||||
):
|
||||
self.async_handler = get_async_httpx_client(
|
||||
|
|
@ -196,6 +201,10 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
additional_provider_specific_params or {}
|
||||
)
|
||||
|
||||
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
|
||||
unreachable_fallback
|
||||
)
|
||||
|
||||
# Set supported event hooks
|
||||
if "supported_event_hooks" not in kwargs:
|
||||
kwargs["supported_event_hooks"] = [
|
||||
|
|
@ -259,6 +268,54 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
|
||||
return result_metadata
|
||||
|
||||
def _fail_open_passthrough(
|
||||
self,
|
||||
*,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"],
|
||||
error: Exception,
|
||||
http_status_code: Optional[int] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
status_suffix = f" http_status_code={http_status_code}" if http_status_code else ""
|
||||
verbose_proxy_logger.critical(
|
||||
"Generic Guardrail API unreachable (fail-open). Proceeding without guardrail.%s "
|
||||
"guardrail_name=%s api_base=%s input_type=%s litellm_call_id=%s litellm_trace_id=%s",
|
||||
status_suffix,
|
||||
getattr(self, "guardrail_name", None),
|
||||
getattr(self, "api_base", None),
|
||||
input_type,
|
||||
getattr(logging_obj, "litellm_call_id", None) if logging_obj else None,
|
||||
getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None,
|
||||
exc_info=error,
|
||||
)
|
||||
# Keep flow going - treat as action=NONE (no modifications)
|
||||
return_inputs: GenericGuardrailAPIInputs = {}
|
||||
return_inputs.update(inputs)
|
||||
return return_inputs
|
||||
|
||||
def _build_guardrail_return_inputs(
|
||||
self,
|
||||
*,
|
||||
texts: list,
|
||||
images: Any,
|
||||
tools: Any,
|
||||
guardrail_response: GenericGuardrailAPIResponse,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
# Action is NONE or no modifications needed
|
||||
return_inputs = GenericGuardrailAPIInputs(texts=texts)
|
||||
if guardrail_response.texts:
|
||||
return_inputs["texts"] = guardrail_response.texts
|
||||
if guardrail_response.images:
|
||||
return_inputs["images"] = guardrail_response.images
|
||||
elif images:
|
||||
return_inputs["images"] = images
|
||||
if guardrail_response.tools:
|
||||
return_inputs["tools"] = guardrail_response.tools
|
||||
elif tools:
|
||||
return_inputs["tools"] = tools
|
||||
return return_inputs
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
@ -313,7 +370,9 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
|
||||
# Extract user API key metadata
|
||||
user_metadata = self._extract_user_api_key_metadata(request_data)
|
||||
inbound_headers = _extract_inbound_headers(request_data=request_data, logging_obj=logging_obj)
|
||||
inbound_headers = _extract_inbound_headers(
|
||||
request_data=request_data, logging_obj=logging_obj
|
||||
)
|
||||
|
||||
# Create request payload
|
||||
guardrail_request = GenericGuardrailAPIRequest(
|
||||
|
|
@ -370,23 +429,64 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
should_wrap_with_default_message=False,
|
||||
)
|
||||
|
||||
# Action is NONE or no modifications needed
|
||||
return_inputs = GenericGuardrailAPIInputs(texts=texts)
|
||||
if guardrail_response.texts:
|
||||
return_inputs["texts"] = guardrail_response.texts
|
||||
if guardrail_response.images:
|
||||
return_inputs["images"] = guardrail_response.images
|
||||
elif images:
|
||||
return_inputs["images"] = images
|
||||
if guardrail_response.tools:
|
||||
return_inputs["tools"] = guardrail_response.tools
|
||||
elif tools:
|
||||
return_inputs["tools"] = tools
|
||||
return return_inputs
|
||||
return self._build_guardrail_return_inputs(
|
||||
texts=texts,
|
||||
images=images,
|
||||
tools=tools,
|
||||
guardrail_response=guardrail_response,
|
||||
)
|
||||
|
||||
except GuardrailRaisedException:
|
||||
# Re-raise guardrail exceptions as-is
|
||||
raise
|
||||
except Timeout as e:
|
||||
# AsyncHTTPHandler wraps httpx.TimeoutException into litellm.Timeout
|
||||
if self.unreachable_fallback == "fail_open":
|
||||
return self._fail_open_passthrough(
|
||||
inputs=inputs,
|
||||
input_type=input_type,
|
||||
logging_obj=logging_obj,
|
||||
error=e,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.error(
|
||||
"Generic Guardrail API: failed to make request: %s", str(e)
|
||||
)
|
||||
raise Exception(f"Generic Guardrail API failed: {str(e)}")
|
||||
except httpx.HTTPStatusError as e:
|
||||
# Common reverse-proxy/LB failures can present as HTTP errors even when the backend is unreachable.
|
||||
status_code = getattr(getattr(e, "response", None), "status_code", None)
|
||||
if self.unreachable_fallback == "fail_open" and status_code in (
|
||||
502,
|
||||
503,
|
||||
504,
|
||||
):
|
||||
return self._fail_open_passthrough(
|
||||
inputs=inputs,
|
||||
input_type=input_type,
|
||||
logging_obj=logging_obj,
|
||||
error=e,
|
||||
http_status_code=status_code,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.error(
|
||||
"Generic Guardrail API: failed to make request: %s", str(e)
|
||||
)
|
||||
raise Exception(f"Generic Guardrail API failed: {str(e)}")
|
||||
except httpx.RequestError as e:
|
||||
# Guardrail endpoint is unreachable (DNS/connect/timeout/etc)
|
||||
if self.unreachable_fallback == "fail_open":
|
||||
return self._fail_open_passthrough(
|
||||
inputs=inputs,
|
||||
input_type=input_type,
|
||||
logging_obj=logging_obj,
|
||||
error=e,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.error(
|
||||
"Generic Guardrail API: failed to make request: %s", str(e)
|
||||
)
|
||||
raise Exception(f"Generic Guardrail API failed: {str(e)}")
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"Generic Guardrail API: failed to make request: %s", str(e)
|
||||
|
|
|
|||
|
|
@ -31,11 +31,15 @@ from litellm import Router
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import GuardrailTracingDetail, ModelResponseStream
|
||||
from litellm.types.utils import (
|
||||
GenericGuardrailAPIInputs,
|
||||
GuardrailStatus,
|
||||
GuardrailTracingDetail,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus
|
||||
|
||||
from litellm.types.guardrails import (
|
||||
BlockedWord,
|
||||
|
|
@ -201,12 +205,6 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
# Load categories if provided
|
||||
if categories:
|
||||
self._load_categories(categories)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
"ContentFilterGuardrail has no content categories configured. "
|
||||
"Toxic/abuse and other category-based keyword filtering will not run. "
|
||||
"Add categories (e.g. harm_toxic_abuse) in the guardrail config to enable them."
|
||||
)
|
||||
|
||||
# Normalize inputs: convert dicts to Pydantic models for consistent handling
|
||||
normalized_patterns: List[ContentFilterPattern] = []
|
||||
|
|
@ -1546,8 +1544,6 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
Raises:
|
||||
HTTPException: If sensitive content is detected and action is BLOCK
|
||||
"""
|
||||
from litellm.types.utils import GuardrailStatus
|
||||
|
||||
start_time = datetime.now()
|
||||
detections: List[ContentFilterDetection] = []
|
||||
masked_entity_count: Dict[str, int] = {}
|
||||
|
|
@ -1693,4 +1689,4 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
LitellmContentFilterGuardrailConfigModel,
|
||||
)
|
||||
|
||||
return LitellmContentFilterGuardrailConfigModel
|
||||
return LitellmContentFilterGuardrailConfigModel
|
||||
|
|
@ -452,6 +452,47 @@
|
|||
"pattern": "\\b\\d{1,6}\\s+[A-Za-z0-9][A-Za-z0-9\\s.'-]*\\s+(?:Street|St|Avenue|Ave|Road|Rd|Boulevard|Blvd|Drive|Dr|Lane|Ln|Way|Court|Ct|Place|Pl|Circle|Cir)\\b",
|
||||
"category": "PII Patterns",
|
||||
"description": "Detects street addresses (number + street name + street type)"
|
||||
},
|
||||
{
|
||||
"name": "airline_pnr",
|
||||
"display_name": "Airline PNR / Booking Reference",
|
||||
"pattern": "\\b[A-Z]{6}\\b",
|
||||
"category": "Aviation PII Patterns",
|
||||
"description": "Detects airline PNR / booking references (6 uppercase alpha characters) when near booking context",
|
||||
"keyword_pattern": "\\b(?:PNR|booking\\s*(?:reference|ref|code|number|confirmation)|reservation\\s*(?:code|number|ref)|record\\s*locator|confirmation\\s*(?:code|number)|itinerary\\s*(?:number|ref))\\b",
|
||||
"allow_word_numbers": false
|
||||
},
|
||||
{
|
||||
"name": "skywards_number",
|
||||
"display_name": "Emirates Skywards / Frequent Flyer Number",
|
||||
"pattern": "\\b(?:EK\\s?)?\\d{9,10}\\b",
|
||||
"category": "Aviation PII Patterns",
|
||||
"description": "Detects Emirates Skywards frequent flyer numbers (9-10 digits, optional EK prefix) when near loyalty/frequent flyer context",
|
||||
"keyword_pattern": "\\b(?:[Ss]kywards|frequent\\s*flyer|FF\\s*(?:number|no|#)|loyalty\\s*(?:number|id|member)|member\\s*(?:number|id|#)|miles\\s*(?:account|number)|tier\\s*(?:number|status))\\b",
|
||||
"allow_word_numbers": false
|
||||
},
|
||||
{
|
||||
"name": "uae_emirates_id",
|
||||
"display_name": "UAE Emirates ID",
|
||||
"pattern": "\\b784-?\\d{4}-?\\d{7}-?\\d\\b",
|
||||
"category": "UAE PII Patterns",
|
||||
"description": "Detects UAE Emirates ID numbers (784-YYYY-NNNNNNN-C format, 15 digits starting with 784)"
|
||||
},
|
||||
{
|
||||
"name": "uae_phone",
|
||||
"display_name": "Phone Number (UAE)",
|
||||
"pattern": "(?<!\\d)(?:\\+971|00971|0)\\s?(?:2|3|4|6|7|9|50|52|54|55|56|58)\\s?\\d{3}\\s?\\d{4}\\b",
|
||||
"category": "UAE PII Patterns",
|
||||
"description": "Detects UAE phone numbers (+971/00971/0 prefix, landline and mobile)"
|
||||
},
|
||||
{
|
||||
"name": "flight_number",
|
||||
"display_name": "Airline Flight Number",
|
||||
"pattern": "\\b(?:EK|QF|BA|AA|LH|AF|SQ|CX|QR|EY|FZ|WY|GF|KU|SV|MS|RJ|ME|PK|AI|UL|WS|9W|6E)\\s?\\d{1,4}\\b",
|
||||
"category": "Aviation PII Patterns",
|
||||
"description": "Detects airline flight numbers (major IATA 2-letter codes + 1-4 digit flight number) when near flight context",
|
||||
"keyword_pattern": "\\b(?:flight|departure|arrival|gate|boarding|schedule|operate|route|aircraft|plane|outbound|inbound|leg|sector|flying)\\b",
|
||||
"allow_word_numbers": false
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -0,0 +1,89 @@
|
|||
# Airline Brand Protection
|
||||
# Prevents AI from generating fake incident reports, unauthorized statements, or negative PR
|
||||
# NOTE: The identifier_words use {{brand_name}} placeholder - this file is used as a reference.
|
||||
# The actual brand-specific keywords are injected at template instantiation time via blocked_words.
|
||||
category_name: "airline_brand_protection"
|
||||
description: "Prevents AI from generating fake incident reports, unauthorized airline statements, or reputation-damaging content"
|
||||
default_action: "BLOCK"
|
||||
|
||||
identifier_words:
|
||||
# Generic airline context (brand-specific terms injected at runtime)
|
||||
- "airline"
|
||||
- "carrier"
|
||||
- "airways"
|
||||
- "flights"
|
||||
|
||||
additional_block_words:
|
||||
# Fake incident reports
|
||||
- "crashed"
|
||||
- "crash landing"
|
||||
- "emergency landing"
|
||||
- "engine failure"
|
||||
- "engine fire"
|
||||
- "depressurization"
|
||||
- "hull breach"
|
||||
- "mayday"
|
||||
- "mid-air collision"
|
||||
- "near miss"
|
||||
- "fatality"
|
||||
- "fatalities"
|
||||
- "dead passengers"
|
||||
- "casualties"
|
||||
# Unauthorized statements
|
||||
- "official statement"
|
||||
- "spokesperson said"
|
||||
- "press release"
|
||||
- "confirms"
|
||||
- "announces"
|
||||
- "announced today"
|
||||
- "breaking news"
|
||||
- "just in"
|
||||
- "BREAKING"
|
||||
- "statement from"
|
||||
# PR / reputation attacks
|
||||
- "scandal"
|
||||
- "cover up"
|
||||
- "coverup"
|
||||
- "lawsuit"
|
||||
- "class action"
|
||||
- "mistreatment"
|
||||
- "safety violation"
|
||||
- "whistleblower"
|
||||
- "banned"
|
||||
- "grounded fleet"
|
||||
- "recall"
|
||||
- "boycott"
|
||||
- "worst airline"
|
||||
- "unsafe"
|
||||
- "dangerous"
|
||||
- "negligence"
|
||||
|
||||
always_block_keywords:
|
||||
- keyword: "plane crash"
|
||||
severity: "high"
|
||||
- keyword: "flight crashed"
|
||||
severity: "high"
|
||||
- keyword: "mid-air collision"
|
||||
severity: "high"
|
||||
- keyword: "passengers dead"
|
||||
severity: "high"
|
||||
- keyword: "confirms fatalities"
|
||||
severity: "high"
|
||||
- keyword: "safety scandal"
|
||||
severity: "high"
|
||||
- keyword: "fleet grounded"
|
||||
severity: "high"
|
||||
|
||||
exceptions:
|
||||
- "how safe is"
|
||||
- "safety record"
|
||||
- "safety rating"
|
||||
- "what is"
|
||||
- "explain"
|
||||
- "historical"
|
||||
- "aviation history"
|
||||
- "customer review"
|
||||
- "customer service"
|
||||
- "feedback"
|
||||
- "improve"
|
||||
- "complaint resolution"
|
||||
|
|
@ -0,0 +1,120 @@
|
|||
# Aviation Safety & Security Topic Restrictions
|
||||
# Prevents AI from generating content about aircraft vulnerabilities, security bypass, cockpit access
|
||||
category_name: "aviation_safety_topics"
|
||||
description: "Blocks AI from generating content about aviation security vulnerabilities, bypass procedures, and restricted operations data"
|
||||
default_action: "BLOCK"
|
||||
|
||||
identifier_words:
|
||||
- "aircraft"
|
||||
- "airplane"
|
||||
- "plane"
|
||||
- "aviation"
|
||||
- "airline"
|
||||
- "airport"
|
||||
- "cockpit"
|
||||
- "flight deck"
|
||||
- "airside"
|
||||
- "tarmac"
|
||||
- "runway"
|
||||
- "cargo hold"
|
||||
- "fuselage"
|
||||
- "avionics"
|
||||
- "transponder"
|
||||
- "ATC"
|
||||
- "air traffic control"
|
||||
- "ACARS"
|
||||
- "flight management system"
|
||||
- "FMS"
|
||||
- "autopilot"
|
||||
- "A380"
|
||||
- "B777"
|
||||
- "777"
|
||||
- "787"
|
||||
|
||||
additional_block_words:
|
||||
# Security bypass
|
||||
- "bypass security"
|
||||
- "bypass screening"
|
||||
- "defeat security"
|
||||
- "circumvent security"
|
||||
- "evade detection"
|
||||
- "avoid screening"
|
||||
- "smuggle"
|
||||
- "smuggling"
|
||||
- "conceal weapon"
|
||||
- "hidden weapon"
|
||||
# Cockpit access
|
||||
- "cockpit door code"
|
||||
- "cockpit access code"
|
||||
- "override cockpit"
|
||||
- "force cockpit door"
|
||||
- "break into cockpit"
|
||||
- "cockpit intrusion"
|
||||
- "cockpit breach"
|
||||
# Aircraft vulnerabilities
|
||||
- "vulnerability"
|
||||
- "exploit"
|
||||
- "disable"
|
||||
- "jamming"
|
||||
- "spoofing"
|
||||
- "hijack"
|
||||
- "hijacking"
|
||||
- "take over"
|
||||
- "commandeer"
|
||||
- "crash intentionally"
|
||||
- "bring down"
|
||||
- "shoot down"
|
||||
- "sabotage"
|
||||
- "tamper"
|
||||
- "interfere with"
|
||||
# System exploitation
|
||||
- "hack"
|
||||
- "hack into"
|
||||
- "inject malware"
|
||||
- "remote access"
|
||||
- "unauthorized access"
|
||||
|
||||
always_block_keywords:
|
||||
- keyword: "hijack a plane"
|
||||
severity: "high"
|
||||
- keyword: "hijack an aircraft"
|
||||
severity: "high"
|
||||
- keyword: "cockpit door override code"
|
||||
severity: "high"
|
||||
- keyword: "bypass airport security"
|
||||
severity: "high"
|
||||
- keyword: "smuggle weapons onto aircraft"
|
||||
severity: "high"
|
||||
- keyword: "disable aircraft transponder"
|
||||
severity: "high"
|
||||
- keyword: "bring down an aircraft"
|
||||
severity: "high"
|
||||
- keyword: "crash a plane into"
|
||||
severity: "high"
|
||||
- keyword: "aircraft bomb"
|
||||
severity: "high"
|
||||
- keyword: "ACARS vulnerability"
|
||||
severity: "high"
|
||||
- keyword: "flight management system hack"
|
||||
severity: "high"
|
||||
- keyword: "spoof ADS-B"
|
||||
severity: "high"
|
||||
|
||||
exceptions:
|
||||
- "aviation safety"
|
||||
- "safety report"
|
||||
- "safety investigation"
|
||||
- "safety improvement"
|
||||
- "NTSB"
|
||||
- "accident investigation"
|
||||
- "safety training"
|
||||
- "security awareness"
|
||||
- "safety management system"
|
||||
- "IATA safety"
|
||||
- "ICAO"
|
||||
- "how to report"
|
||||
- "incident report"
|
||||
- "safety culture"
|
||||
- "what is"
|
||||
- "explain"
|
||||
- "research on"
|
||||
|
|
@ -0,0 +1,71 @@
|
|||
# UAE Anti-Discrimination Compliance
|
||||
# UAE Federal Decree-Law No. 2/2015 on combating discrimination and hatred
|
||||
category_name: "uae_anti_discrimination"
|
||||
description: "Blocks content violating UAE Federal anti-discrimination laws (Decree-Law No. 2/2015)"
|
||||
default_action: "BLOCK"
|
||||
|
||||
identifier_words:
|
||||
- "Emirati"
|
||||
- "Emiratis"
|
||||
- "expatriate"
|
||||
- "expat"
|
||||
- "migrant worker"
|
||||
- "laborer"
|
||||
- "domestic worker"
|
||||
- "nationality"
|
||||
- "religion"
|
||||
- "race"
|
||||
- "ethnicity"
|
||||
- "caste"
|
||||
- "sect"
|
||||
|
||||
additional_block_words:
|
||||
- "inferior"
|
||||
- "superior"
|
||||
- "subhuman"
|
||||
- "uncivilized"
|
||||
- "backward"
|
||||
- "primitive"
|
||||
- "servant class"
|
||||
- "slave"
|
||||
- "slavery"
|
||||
- "second class"
|
||||
- "second-class citizen"
|
||||
- "don't belong"
|
||||
- "go back to"
|
||||
- "deport"
|
||||
- "hate"
|
||||
- "hatred"
|
||||
- "incite"
|
||||
- "incitement"
|
||||
- "sedition"
|
||||
- "sectarian"
|
||||
|
||||
always_block_keywords:
|
||||
- keyword: "racial superiority"
|
||||
severity: "high"
|
||||
- keyword: "ethnic cleansing"
|
||||
severity: "high"
|
||||
- keyword: "sectarian violence"
|
||||
severity: "high"
|
||||
- keyword: "incite hatred"
|
||||
severity: "high"
|
||||
- keyword: "incite discrimination"
|
||||
severity: "high"
|
||||
- keyword: "religious hatred"
|
||||
severity: "high"
|
||||
|
||||
exceptions:
|
||||
- "anti-discrimination"
|
||||
- "equal rights"
|
||||
- "diversity and inclusion"
|
||||
- "tolerance"
|
||||
- "coexistence"
|
||||
- "UAE Year of Tolerance"
|
||||
- "human rights"
|
||||
- "what is"
|
||||
- "explain"
|
||||
- "history of"
|
||||
- "education"
|
||||
- "awareness"
|
||||
- "compliance training"
|
||||
|
|
@ -0,0 +1,72 @@
|
|||
# UAE Cultural Sensitivity
|
||||
# Respects UAE cultural norms, royal family references, and religious content policies
|
||||
category_name: "uae_cultural_sensitivity"
|
||||
description: "Ensures AI output respects UAE cultural norms, royal family, and religious content policies"
|
||||
default_action: "BLOCK"
|
||||
|
||||
identifier_words:
|
||||
- "UAE"
|
||||
- "United Arab Emirates"
|
||||
- "Dubai"
|
||||
- "Abu Dhabi"
|
||||
- "Sharjah"
|
||||
- "Ajman"
|
||||
- "Fujairah"
|
||||
- "Ras Al Khaimah"
|
||||
- "Umm Al Quwain"
|
||||
- "Emirati"
|
||||
- "Emiratis"
|
||||
|
||||
additional_block_words:
|
||||
# Royal family disrespect
|
||||
- "mock"
|
||||
- "ridicule"
|
||||
- "insult"
|
||||
- "satirize"
|
||||
- "caricature"
|
||||
- "parody"
|
||||
- "defame"
|
||||
- "slander"
|
||||
- "disrespect"
|
||||
# Cultural sensitivity
|
||||
- "alcohol promotion"
|
||||
- "gambling promotion"
|
||||
- "drug use"
|
||||
- "immodest"
|
||||
- "indecent"
|
||||
|
||||
always_block_keywords:
|
||||
- keyword: "insult the ruler"
|
||||
severity: "high"
|
||||
- keyword: "mock the sheikh"
|
||||
severity: "high"
|
||||
- keyword: "ridicule the royal family"
|
||||
severity: "high"
|
||||
- keyword: "defame Mohammed bin"
|
||||
severity: "high"
|
||||
- keyword: "blasphemy against Islam"
|
||||
severity: "high"
|
||||
- keyword: "insult Islam"
|
||||
severity: "high"
|
||||
- keyword: "mock the Prophet"
|
||||
severity: "high"
|
||||
- keyword: "anti-Islamic"
|
||||
severity: "high"
|
||||
- keyword: "desecrate the Quran"
|
||||
severity: "high"
|
||||
|
||||
exceptions:
|
||||
- "what is"
|
||||
- "explain"
|
||||
- "cultural guide"
|
||||
- "travel guide"
|
||||
- "etiquette"
|
||||
- "customs"
|
||||
- "tradition"
|
||||
- "how to respect"
|
||||
- "cultural awareness"
|
||||
- "history of"
|
||||
- "learn about"
|
||||
- "understand"
|
||||
- "policy on"
|
||||
- "regulation"
|
||||
|
|
@ -0,0 +1,35 @@
|
|||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
from .mcp_end_user_permission import MCPEndUserPermissionGuardrail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
|
||||
import litellm
|
||||
|
||||
# Default to always-on. Only disable if the user explicitly sets default_on: false.
|
||||
# We check the raw guardrail dict because LitellmParams normalizes None → False,
|
||||
# making it impossible to distinguish "not set" from "explicitly false" via litellm_params.
|
||||
_raw_default_on = guardrail.get("litellm_params", {}).get("default_on")
|
||||
_default_on = False if _raw_default_on is False else True
|
||||
|
||||
_callback = MCPEndUserPermissionGuardrail(
|
||||
guardrail_name=guardrail.get("guardrail_name", ""),
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=_default_on,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_callback)
|
||||
return _callback
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.MCP_END_USER_PERMISSION.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.MCP_END_USER_PERMISSION.value: MCPEndUserPermissionGuardrail,
|
||||
}
|
||||
|
|
@ -0,0 +1,262 @@
|
|||
"""
|
||||
MCP End User Permission Guardrail Hook
|
||||
|
||||
Enforces end user permissions for MCP server access via apply_guardrail:
|
||||
- input_type="request" → filter tools the end user cannot access
|
||||
|
||||
Permission logic:
|
||||
- No end_user_id → allow all (key/team-level permissions apply)
|
||||
- end_user_id, no mcp_servers → allow all (default)
|
||||
- end_user_id + mcp_servers → allow only those servers
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Literal, Optional, Type
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
GUARDRAIL_NAME = "mcp_end_user_permission"
|
||||
|
||||
|
||||
class MCPEndUserPermissionGuardrail(CustomGuardrail):
|
||||
"""
|
||||
Guardrail that enforces end user permissions for MCP server access.
|
||||
|
||||
Runs on input only (pre-call). Filters tools in the request that the
|
||||
end user is not permitted to call based on their object_permission.
|
||||
|
||||
end_user_object_permission is populated on UserAPIKeyAuth during auth.
|
||||
The guardrail resolves it via a cached get_end_user_object lookup —
|
||||
no extra DB round-trip when the cache is warm.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
if "supported_event_hooks" not in kwargs:
|
||||
kwargs["supported_event_hooks"] = [
|
||||
GuardrailEventHooks.pre_call,
|
||||
]
|
||||
super().__init__(**kwargs)
|
||||
verbose_proxy_logger.debug("MCP End User Permission Guardrail initialized")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# apply_guardrail — filters MCP tools on the request side only
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"] = "request",
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""
|
||||
Filters MCP tools the end user cannot access based on their
|
||||
object_permission.mcp_servers / mcp_access_groups settings.
|
||||
"""
|
||||
object_permission = await self._resolve_end_user_object_permission(request_data)
|
||||
return await self._check_request_tools(inputs, object_permission)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Private — request-side tool filtering
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _check_request_tools(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionTable],
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
tools = inputs.get("tools")
|
||||
if not tools:
|
||||
return inputs
|
||||
|
||||
allowed_mcp_servers = (
|
||||
await self._get_allowed_mcp_servers_from_object_permission(
|
||||
object_permission
|
||||
)
|
||||
)
|
||||
if allowed_mcp_servers is None:
|
||||
return inputs # No restrictions → pass through unchanged
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"MCP guardrail: end user restricted to MCP servers: {allowed_mcp_servers}"
|
||||
)
|
||||
|
||||
filtered_tools = []
|
||||
removed_tools = []
|
||||
|
||||
for tool in tools:
|
||||
tool_name = self._get_tool_name_from_definition(tool)
|
||||
server_name = (
|
||||
self._extract_mcp_server_name(tool_name) if tool_name else None
|
||||
)
|
||||
|
||||
if server_name is None:
|
||||
# Not an MCP tool (no prefix) or unrecognised format → keep
|
||||
filtered_tools.append(tool)
|
||||
elif server_name in allowed_mcp_servers:
|
||||
filtered_tools.append(tool)
|
||||
else:
|
||||
removed_tools.append(tool_name)
|
||||
verbose_proxy_logger.warning(
|
||||
f"MCP guardrail: removing tool '{tool_name}' "
|
||||
f"(server: '{server_name}') — not in end user's allowed servers"
|
||||
)
|
||||
|
||||
if removed_tools:
|
||||
verbose_proxy_logger.debug(
|
||||
f"MCP guardrail: removed {len(removed_tools)} unauthorized MCP tool(s): {removed_tools}"
|
||||
)
|
||||
inputs["tools"] = filtered_tools
|
||||
|
||||
return inputs
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Private — end user permission resolution
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_end_user_object_permission(
|
||||
request_data: dict,
|
||||
) -> Optional[LiteLLM_ObjectPermissionTable]:
|
||||
"""
|
||||
Resolve the end user's object_permission via the cached auth lookup.
|
||||
|
||||
Uses get_end_user_object (same path as auth) so no extra DB round-trip
|
||||
when the cache is warm.
|
||||
"""
|
||||
end_user_id = MCPEndUserPermissionGuardrail._get_end_user_id_from_request_data(
|
||||
request_data
|
||||
)
|
||||
if not end_user_id:
|
||||
return None
|
||||
|
||||
end_user_object = await MCPEndUserPermissionGuardrail._fetch_end_user_object(
|
||||
end_user_id
|
||||
)
|
||||
return (
|
||||
end_user_object.object_permission if end_user_object is not None else None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_end_user_id_from_request_data(request_data: dict) -> Optional[str]:
|
||||
return request_data.get("user_api_key_end_user_id") or request_data.get(
|
||||
"litellm_metadata", {}
|
||||
).get("user_api_key_end_user_id")
|
||||
|
||||
@staticmethod
|
||||
async def _fetch_end_user_object(end_user_id: str): # type: ignore[return]
|
||||
"""
|
||||
Fetch end user object via the same cached path used during auth.
|
||||
No extra DB round-trip when the cache is warm.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_end_user_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
return await get_end_user_object(
|
||||
end_user_id=end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
route="/mcp",
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
f"MCP guardrail: failed to fetch end_user_object for '{end_user_id}': {e}"
|
||||
)
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Private — permission derivation
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_from_object_permission(
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionTable],
|
||||
) -> Optional[List[str]]:
|
||||
"""
|
||||
Returns:
|
||||
None — no restrictions configured, allow all MCP servers
|
||||
list — restrict to exactly these server names
|
||||
"""
|
||||
if object_permission is None:
|
||||
return None
|
||||
|
||||
direct_mcp_servers = object_permission.mcp_servers or []
|
||||
mcp_access_groups = object_permission.mcp_access_groups or []
|
||||
|
||||
if not direct_mcp_servers and not mcp_access_groups:
|
||||
return None # Both empty → no restrictions
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
|
||||
access_group_servers = (
|
||||
await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
mcp_access_groups
|
||||
)
|
||||
)
|
||||
|
||||
return list(set(direct_mcp_servers + access_group_servers))
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Config model — exposes this guardrail in the UI
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.mcp_end_user_permission import (
|
||||
MCPEndUserPermissionGuardrailConfigModel,
|
||||
)
|
||||
|
||||
return MCPEndUserPermissionGuardrailConfigModel
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Private — tool name extraction
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _extract_mcp_server_name(tool_name: str) -> Optional[str]:
|
||||
"""
|
||||
Split "github-create_issue" → "github".
|
||||
Returns None if the tool name has no '-' prefix (not an MCP tool).
|
||||
"""
|
||||
if not tool_name or "-" not in tool_name:
|
||||
return None
|
||||
return tool_name.split("-", 1)[0]
|
||||
|
||||
@staticmethod
|
||||
def _get_tool_name_from_definition(tool: Any) -> Optional[str]:
|
||||
"""
|
||||
Extract tool name from a definition dict.
|
||||
|
||||
OpenAI format: {"type": "function", "function": {"name": "..."}}
|
||||
Anthropic format: {"name": "...", "input_schema": {...}}
|
||||
"""
|
||||
if not isinstance(tool, dict):
|
||||
return None
|
||||
function_def = tool.get("function")
|
||||
if isinstance(function_def, dict):
|
||||
name = function_def.get("name")
|
||||
if name:
|
||||
return name
|
||||
return tool.get("name")
|
||||
|
|
@ -85,6 +85,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
|
||||
verbose_proxy_logger.debug("Running UnifiedLLMGuardrails pre-call hook")
|
||||
|
||||
guardrail_to_apply: CustomGuardrail = data.pop("guardrail_to_apply", None)
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ def init_guardrails_v2(
|
|||
if initialized_guardrail:
|
||||
guardrail_list.append(initialized_guardrail)
|
||||
|
||||
verbose_proxy_logger.debug(f"\nGuardrail List:{guardrail_list}\n")
|
||||
# verbose_proxy_logger.debug(f"\nGuardrail List:{guardrail_list}\n")
|
||||
|
||||
# Populate router's guardrail_list for load balancing support
|
||||
_populate_router_guardrail_list(guardrail_list=guardrail_list)
|
||||
|
|
|
|||
|
|
@ -239,6 +239,8 @@ def clean_headers(
|
|||
"""
|
||||
Removes litellm api key from headers
|
||||
"""
|
||||
from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key
|
||||
|
||||
clean_headers = {}
|
||||
litellm_key_lower = (
|
||||
litellm_key_header_name.lower() if litellm_key_header_name is not None else None
|
||||
|
|
@ -246,8 +248,13 @@ def clean_headers(
|
|||
|
||||
for header, value in headers.items():
|
||||
header_lower = header.lower()
|
||||
# Preserve Authorization header if it contains Anthropic OAuth token (sk-ant-oat*)
|
||||
# This allows OAuth tokens to be forwarded to Anthropic-compatible providers
|
||||
# via add_provider_specific_headers_to_request()
|
||||
if header_lower == "authorization" and is_anthropic_oauth_key(value):
|
||||
clean_headers[header] = value
|
||||
# Check if header should be excluded: either in special headers cache or matches custom litellm key
|
||||
if header_lower not in _SPECIAL_HEADERS_CACHE and (
|
||||
elif header_lower not in _SPECIAL_HEADERS_CACHE and (
|
||||
litellm_key_lower is None or header_lower != litellm_key_lower
|
||||
):
|
||||
clean_headers[header] = value
|
||||
|
|
@ -872,7 +879,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
general_settings, user_api_key_dict, _headers
|
||||
)
|
||||
|
||||
# Parse user info from headers
|
||||
# Parse user info from headers (fallback to general_settings.user_header_name)
|
||||
user = LiteLLMProxyRequestSetup.get_user_from_headers(_headers, general_settings)
|
||||
if user is not None:
|
||||
if user_api_key_dict.end_user_id is None:
|
||||
|
|
@ -1533,9 +1540,7 @@ def _match_and_track_policies(
|
|||
add_policy_sources_to_metadata,
|
||||
add_policy_to_applied_policies_header,
|
||||
)
|
||||
from litellm.proxy.policy_engine.attachment_registry import (
|
||||
get_attachment_registry,
|
||||
)
|
||||
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
|
||||
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
|
||||
|
||||
# Get matching policies via attachments (with match reasons for attribution)
|
||||
|
|
@ -1670,9 +1675,7 @@ def add_guardrails_from_policy_engine(
|
|||
user_api_key_dict: The user's API key authentication info
|
||||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
get_tags_from_request_body,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
from litellm.types.proxy.policy_engine import PolicyMatchContext
|
||||
|
||||
|
|
@ -1717,6 +1720,8 @@ def add_provider_specific_headers_to_request(
|
|||
data: dict,
|
||||
headers: dict,
|
||||
):
|
||||
from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key
|
||||
|
||||
anthropic_headers = {}
|
||||
# boolean to indicate if a header was added
|
||||
added_header = False
|
||||
|
|
@ -1726,6 +1731,14 @@ def add_provider_specific_headers_to_request(
|
|||
anthropic_headers[header] = header_value
|
||||
added_header = True
|
||||
|
||||
# Check for Authorization header with Anthropic OAuth token (sk-ant-oat*)
|
||||
# This needs to be handled via provider-specific headers to ensure it only
|
||||
# goes to Anthropic-compatible providers, not all providers in the router
|
||||
for header, value in headers.items():
|
||||
if header.lower() == "authorization" and is_anthropic_oauth_key(value):
|
||||
anthropic_headers[header] = value
|
||||
added_header = True
|
||||
break
|
||||
if added_header is True:
|
||||
# Anthropic headers work across multiple providers
|
||||
# Store as comma-separated list so retrieval can match any of them
|
||||
|
|
|
|||
|
|
@ -19,11 +19,13 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import \
|
||||
get_daily_activity
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_set_object_permission, handle_update_object_permission_common)
|
||||
from litellm.proxy.utils import handle_exception_on_proxy
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import \
|
||||
SpendAnalyticsPaginatedResponse
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
|
@ -107,9 +109,8 @@ async def unblock_user(data: BlockUsers):
|
|||
```
|
||||
"""
|
||||
try:
|
||||
from enterprise.enterprise_hooks.blocked_user_list import (
|
||||
_ENTERPRISE_BlockedUserList,
|
||||
)
|
||||
from enterprise.enterprise_hooks.blocked_user_list import \
|
||||
_ENTERPRISE_BlockedUserList
|
||||
except ImportError:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -164,6 +165,38 @@ def new_budget_request(data: NewCustomerRequest) -> Optional[BudgetNewRequest]:
|
|||
return None
|
||||
|
||||
|
||||
async def _handle_customer_object_permission_update(
|
||||
non_default_values: dict,
|
||||
end_user_table_data_typed: Optional[LiteLLM_EndUserTable],
|
||||
update_end_user_table_data: dict,
|
||||
prisma_client,
|
||||
) -> None:
|
||||
"""
|
||||
Handle object permission updates for customer endpoints.
|
||||
|
||||
Updates the update_end_user_table_data dict in place with the new object_permission_id.
|
||||
|
||||
Args:
|
||||
non_default_values: Dictionary containing the update values including object_permission
|
||||
end_user_table_data_typed: Existing end user table data
|
||||
update_end_user_table_data: Dictionary to update with new object_permission_id
|
||||
prisma_client: Prisma database client
|
||||
"""
|
||||
if "object_permission" in non_default_values:
|
||||
existing_object_permission_id = (
|
||||
end_user_table_data_typed.object_permission_id
|
||||
if end_user_table_data_typed is not None
|
||||
else None
|
||||
)
|
||||
object_permission_id = await handle_update_object_permission_common(
|
||||
data_json=non_default_values,
|
||||
existing_object_permission_id=existing_object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if object_permission_id is not None:
|
||||
update_end_user_table_data["object_permission_id"] = object_permission_id
|
||||
|
||||
|
||||
@router.post(
|
||||
"/end_user/new",
|
||||
tags=["Customer Management"],
|
||||
|
|
@ -200,6 +233,16 @@ async def new_end_user(
|
|||
- soft_budget: Optional[float] - [Not Implemented Yet] Get alerts when customer crosses given budget, doesn't block requests.
|
||||
- spend: Optional[float] - Specify initial spend for a given customer.
|
||||
- budget_reset_at: Optional[str] - Specify the date and time when the budget should be reset.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - Customer-specific object permissions to control access to resources.
|
||||
Supported fields:
|
||||
* mcp_servers: List[str] - List of allowed MCP server IDs
|
||||
* mcp_access_groups: List[str] - List of MCP access group names
|
||||
* mcp_tool_permissions: Dict[str, List[str]] - Map of server ID to allowed tool names (e.g., {"server_1": ["tool_a", "tool_b"]})
|
||||
* vector_stores: List[str] - List of allowed vector store IDs
|
||||
* agents: List[str] - List of allowed agent IDs
|
||||
* agent_access_groups: List[str] - List of agent access group names
|
||||
Example: {"mcp_servers": ["server_1", "server_2"], "vector_stores": ["vector_store_1"], "agents": ["agent_1"]}
|
||||
IF null or {} then no object-level restrictions apply.
|
||||
|
||||
|
||||
- Allow specifying allowed regions
|
||||
|
|
@ -214,9 +257,22 @@ async def new_end_user(
|
|||
"user_id" : "ishaan-jaff-3",
|
||||
"allowed_region": "eu",
|
||||
"budget_id": "free_tier",
|
||||
"default_model": "azure/gpt-3.5-turbo-eu" <- all calls from this user, use this model?
|
||||
"default_model": "azure/gpt-3.5-turbo-eu"
|
||||
}'
|
||||
|
||||
# With object permissions
|
||||
curl -L -X POST 'http://localhost:4000/customer/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"user_id": "user_1",
|
||||
"object_permission": {
|
||||
"mcp_servers": ["server_1"],
|
||||
"mcp_access_groups": ["public_group"],
|
||||
"vector_stores": ["vector_store_1"]
|
||||
}
|
||||
}'
|
||||
|
||||
# return end-user object
|
||||
```
|
||||
|
||||
|
|
@ -233,11 +289,8 @@ async def new_end_user(
|
|||
- end-user object
|
||||
- currently allowed models
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
llm_router,
|
||||
prisma_client,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (litellm_proxy_admin_name,
|
||||
llm_router, prisma_client)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -289,13 +342,34 @@ async def new_end_user(
|
|||
if k not in BudgetNewRequest.model_fields.keys():
|
||||
new_end_user_obj[k] = v
|
||||
|
||||
## Handle Object Permission - MCP Servers, Vector Stores etc.
|
||||
new_end_user_obj = await _set_object_permission(
|
||||
data_json=new_end_user_obj,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
# Ensure object_permission is not in the data being sent to create
|
||||
# It should have been converted to object_permission_id by _set_object_permission
|
||||
if "object_permission" in new_end_user_obj:
|
||||
verbose_proxy_logger.warning(
|
||||
f"object_permission still in new_end_user_obj after _set_object_permission: {new_end_user_obj.get('object_permission')}"
|
||||
)
|
||||
new_end_user_obj.pop("object_permission", None)
|
||||
|
||||
## WRITE TO DB ##
|
||||
end_user_record = await prisma_client.db.litellm_endusertable.create(
|
||||
data=new_end_user_obj, # type: ignore
|
||||
include={"litellm_budget_table": True},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
||||
return end_user_record
|
||||
# Convert to dict and clean up recursive fields
|
||||
response_dict = end_user_record.model_dump()
|
||||
if response_dict.get("object_permission"):
|
||||
# Remove reverse relations from object_permission
|
||||
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
|
||||
response_dict["object_permission"].pop(field, None)
|
||||
|
||||
return response_dict
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.customer_endpoints.new_end_user(): Exception occured - {}".format(
|
||||
|
|
@ -351,7 +425,7 @@ async def end_user_info(
|
|||
)
|
||||
|
||||
user_info = await prisma_client.db.litellm_endusertable.find_first(
|
||||
where={"user_id": end_user_id}, include={"litellm_budget_table": True}
|
||||
where={"user_id": end_user_id}, include={"litellm_budget_table": True, "object_permission": True}
|
||||
)
|
||||
|
||||
if user_info is None:
|
||||
|
|
@ -361,7 +435,15 @@ async def end_user_info(
|
|||
code=404,
|
||||
param="end_user_id",
|
||||
)
|
||||
return user_info.model_dump(exclude_none=True)
|
||||
|
||||
# Convert to dict and clean up recursive fields
|
||||
response_dict = user_info.model_dump(exclude_none=True)
|
||||
if response_dict.get("object_permission"):
|
||||
# Remove reverse relations from object_permission
|
||||
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
|
||||
response_dict["object_permission"].pop(field, None)
|
||||
|
||||
return response_dict
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
|
|
@ -401,6 +483,16 @@ async def update_end_user(
|
|||
- default_model: Optional[str] = (
|
||||
None # if no equivalent model in allowed region - default all requests to this model
|
||||
)
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - Customer-specific object permissions to control access to resources.
|
||||
Supported fields:
|
||||
* mcp_servers: List[str] - List of allowed MCP server IDs
|
||||
* mcp_access_groups: List[str] - List of MCP access group names
|
||||
* mcp_tool_permissions: Dict[str, List[str]] - Map of server ID to allowed tool names
|
||||
* vector_stores: List[str] - List of allowed vector store IDs
|
||||
* agents: List[str] - List of allowed agent IDs
|
||||
* agent_access_groups: List[str] - List of agent access group names
|
||||
Example: {"mcp_servers": ["server_1"], "vector_stores": ["vector_store_1"]}
|
||||
IF null or {} then no object-level restrictions apply.
|
||||
|
||||
Example curl:
|
||||
```
|
||||
|
|
@ -412,11 +504,24 @@ async def update_end_user(
|
|||
"budget_id": "paid_tier"
|
||||
}'
|
||||
|
||||
See below for all params
|
||||
# Updating object permissions
|
||||
curl -L -X POST 'http://localhost:4000/customer/update' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"user_id": "user_1",
|
||||
"object_permission": {
|
||||
"mcp_servers": ["server_3"],
|
||||
"vector_stores": ["vector_store_2", "vector_store_3"]
|
||||
}
|
||||
}'
|
||||
|
||||
See below for all params
|
||||
```
|
||||
"""
|
||||
|
||||
from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client
|
||||
from litellm.proxy.proxy_server import (litellm_proxy_admin_name,
|
||||
prisma_client)
|
||||
|
||||
try:
|
||||
data_json: dict = data.json()
|
||||
|
|
@ -467,6 +572,14 @@ async def update_end_user(
|
|||
elif k in LiteLLM_EndUserTable.model_fields.keys():
|
||||
update_end_user_table_data[k] = v
|
||||
|
||||
## Handle object permission updates (MCP servers, vector stores, etc.)
|
||||
await _handle_customer_object_permission_update(
|
||||
non_default_values=non_default_values,
|
||||
end_user_table_data_typed=end_user_table_data_typed,
|
||||
update_end_user_table_data=update_end_user_table_data,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
## Check if we need to create a new budget (only if budget fields are provided, not just budget_id) ##
|
||||
if budget_table_data:
|
||||
if end_user_budget_table is None:
|
||||
|
|
@ -498,11 +611,20 @@ async def update_end_user(
|
|||
|
||||
## Update user table, with update params + new budget id (if set) ##
|
||||
verbose_proxy_logger.debug("/customer/update: Received data = %s", data)
|
||||
|
||||
# Ensure object_permission is not in the update data
|
||||
# It should have been converted to object_permission_id by handle_update_object_permission_common
|
||||
if "object_permission" in update_end_user_table_data:
|
||||
verbose_proxy_logger.warning(
|
||||
f"object_permission still in update_end_user_table_data: {update_end_user_table_data.get('object_permission')}"
|
||||
)
|
||||
update_end_user_table_data.pop("object_permission", None)
|
||||
|
||||
if data.user_id is not None and len(data.user_id) > 0:
|
||||
update_end_user_table_data["user_id"] = data.user_id # type: ignore
|
||||
verbose_proxy_logger.debug("In update customer, user_id condition block.")
|
||||
response = await prisma_client.db.litellm_endusertable.update(
|
||||
where={"user_id": data.user_id}, data=update_end_user_table_data, include={"litellm_budget_table": True} # type: ignore
|
||||
where={"user_id": data.user_id}, data=update_end_user_table_data, include={"litellm_budget_table": True, "object_permission": True} # type: ignore
|
||||
)
|
||||
if response is None:
|
||||
raise ValueError(
|
||||
|
|
@ -511,7 +633,15 @@ async def update_end_user(
|
|||
verbose_proxy_logger.debug(
|
||||
f"received response from updating prisma client. response={response}"
|
||||
)
|
||||
return response
|
||||
|
||||
# Convert to dict and clean up recursive fields
|
||||
response_dict = response.model_dump()
|
||||
if response_dict.get("object_permission"):
|
||||
# Remove reverse relations from object_permission
|
||||
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
|
||||
response_dict["object_permission"].pop(field, None)
|
||||
|
||||
return response_dict
|
||||
else:
|
||||
raise ValueError(f"user_id is required, passed user_id = {data.user_id}")
|
||||
|
||||
|
|
@ -663,12 +793,17 @@ async def list_end_user(
|
|||
)
|
||||
|
||||
response = await prisma_client.db.litellm_endusertable.find_many(
|
||||
include={"litellm_budget_table": True}
|
||||
include={"litellm_budget_table": True, "object_permission": True}
|
||||
)
|
||||
|
||||
returned_response: List[LiteLLM_EndUserTable] = []
|
||||
for item in response:
|
||||
returned_response.append(LiteLLM_EndUserTable(**item.model_dump()))
|
||||
item_dict = item.model_dump()
|
||||
# Remove reverse relations from object_permission
|
||||
if item_dict.get("object_permission"):
|
||||
for field in ["teams", "verification_tokens", "organizations", "users", "end_users"]:
|
||||
item_dict["object_permission"].pop(field, None)
|
||||
returned_response.append(LiteLLM_EndUserTable(**item_dict))
|
||||
return returned_response
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -706,9 +841,7 @@ async def get_customer_daily_activity(
|
|||
"""
|
||||
Get daily activity for specific organizations or all accessible organizations.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from datetime import datetime, timedelta, timezone
|
|||
from typing import Any, Dict, List, Literal, Optional, Tuple, cast
|
||||
|
||||
import fastapi
|
||||
import prisma
|
||||
import yaml
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
|
||||
|
||||
|
|
@ -74,7 +75,6 @@ from litellm.proxy.utils import (
|
|||
_hash_token_if_needed,
|
||||
handle_exception_on_proxy,
|
||||
is_valid_api_key,
|
||||
jsonify_object,
|
||||
)
|
||||
from litellm.router import Router
|
||||
from litellm.secret_managers.main import get_secret
|
||||
|
|
@ -3052,7 +3052,7 @@ async def delete_key_aliases(
|
|||
)
|
||||
|
||||
|
||||
async def _rotate_master_key(
|
||||
async def _rotate_master_key( # noqa: PLR0915
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
current_master_key: str,
|
||||
|
|
@ -3095,13 +3095,17 @@ async def _rotate_master_key(
|
|||
should_create_model_in_db=False,
|
||||
)
|
||||
if new_model:
|
||||
new_models.append(jsonify_object(new_model.model_dump()))
|
||||
_dumped = new_model.model_dump(exclude_none=True)
|
||||
_dumped["litellm_params"] = prisma.Json(_dumped["litellm_params"]) # type: ignore[attr-defined]
|
||||
_dumped["model_info"] = prisma.Json(_dumped["model_info"]) # type: ignore[attr-defined]
|
||||
new_models.append(_dumped)
|
||||
verbose_proxy_logger.debug("Resetting proxy model table")
|
||||
await prisma_client.db.litellm_proxymodeltable.delete_many()
|
||||
verbose_proxy_logger.debug("Creating %s models", len(new_models))
|
||||
await prisma_client.db.litellm_proxymodeltable.create_many(
|
||||
data=new_models,
|
||||
)
|
||||
async with prisma_client.db.tx() as tx:
|
||||
await tx.litellm_proxymodeltable.delete_many()
|
||||
verbose_proxy_logger.debug("Creating %s models", len(new_models))
|
||||
await tx.litellm_proxymodeltable.create_many(
|
||||
data=new_models,
|
||||
)
|
||||
# 3. process config table
|
||||
try:
|
||||
config = await prisma_client.db.litellm_config.find_many()
|
||||
|
|
@ -3127,15 +3131,20 @@ async def _rotate_master_key(
|
|||
if encrypted_env_vars:
|
||||
await prisma_client.db.litellm_config.update(
|
||||
where={"param_name": "environment_variables"},
|
||||
data={"param_value": jsonify_object(encrypted_env_vars)},
|
||||
data={"param_value": prisma.Json(encrypted_env_vars)}, # type: ignore[attr-defined]
|
||||
)
|
||||
|
||||
# 4. process MCP server table
|
||||
await rotate_mcp_server_credentials_master_key(
|
||||
prisma_client=prisma_client,
|
||||
touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
||||
new_master_key=new_master_key,
|
||||
)
|
||||
try:
|
||||
await rotate_mcp_server_credentials_master_key(
|
||||
prisma_client=prisma_client,
|
||||
touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
||||
new_master_key=new_master_key,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to rotate MCP server credentials: %s", str(e)
|
||||
)
|
||||
|
||||
# 5. process credentials table
|
||||
try:
|
||||
|
|
@ -3153,13 +3162,19 @@ async def _rotate_master_key(
|
|||
updated_patch=decrypted_cred,
|
||||
new_encryption_key=new_master_key,
|
||||
)
|
||||
credential_object_jsonified = jsonify_object(
|
||||
encrypted_cred.model_dump()
|
||||
)
|
||||
_cred_data = encrypted_cred.model_dump(exclude_none=True)
|
||||
if "credential_values" in _cred_data:
|
||||
_cred_data["credential_values"] = prisma.Json( # type: ignore[attr-defined]
|
||||
_cred_data["credential_values"]
|
||||
)
|
||||
if "credential_info" in _cred_data:
|
||||
_cred_data["credential_info"] = prisma.Json( # type: ignore[attr-defined]
|
||||
_cred_data["credential_info"]
|
||||
)
|
||||
await prisma_client.db.litellm_credentialstable.update(
|
||||
where={"credential_name": cred.credential_name},
|
||||
data={
|
||||
**credential_object_jsonified,
|
||||
**_cred_data,
|
||||
"updated_by": user_api_key_dict.user_id,
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -11,13 +11,19 @@ All /policy management endpoints
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, List, Literal, Optional, TypedDict, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
|
||||
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
PolicyGuardrailsResponse,
|
||||
PolicyInfoResponse,
|
||||
|
|
@ -29,10 +35,186 @@ from litellm.types.proxy.policy_engine import (
|
|||
PolicyValidateRequest,
|
||||
PolicyValidationResponse,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class GuardrailApplyError(Exception):
|
||||
"""
|
||||
Raised when a guardrail's apply_guardrail fails during apply_policies.
|
||||
|
||||
Consumers (e.g. Compliance UI) can use guardrail_name and message to show
|
||||
which guardrail triggered and the error reason.
|
||||
"""
|
||||
|
||||
def __init__(self, guardrail_name: str, message: str) -> None:
|
||||
self.guardrail_name = guardrail_name
|
||||
self.message = message
|
||||
super().__init__(f"Guardrail '{guardrail_name}' failed: {message}")
|
||||
|
||||
|
||||
class GuardrailErrorEntry(TypedDict):
|
||||
"""One guardrail failure for ApplyPoliciesResult.guardrail_errors."""
|
||||
|
||||
guardrail_name: str
|
||||
message: str
|
||||
|
||||
|
||||
class ApplyPoliciesResult(TypedDict):
|
||||
"""Result of apply_policies: inputs plus any guardrail failures."""
|
||||
|
||||
inputs: GenericGuardrailAPIInputs
|
||||
guardrail_errors: List[GuardrailErrorEntry]
|
||||
|
||||
|
||||
async def apply_policies(
|
||||
policy_names: Optional[list[str]],
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
proxy_logging_obj: "LiteLLMLoggingObj",
|
||||
guardrail_names: Optional[list[str]] = None,
|
||||
) -> ApplyPoliciesResult:
|
||||
"""
|
||||
Apply guardrails to inputs from policy names and/or a direct list of guardrail names.
|
||||
|
||||
Runs all guardrails in order; if one fails, the error is recorded and execution
|
||||
continues so that all inputs can complete testing and all guardrail failures are
|
||||
collected. No exception is raised; failures are returned in guardrail_errors.
|
||||
|
||||
Guardrails can be specified in two ways (both can be used together; names are merged):
|
||||
- policy_names: resolve guardrails from the policy registry (with inheritance).
|
||||
- guardrail_names: use this list of guardrail names directly (no policy registry needed).
|
||||
|
||||
Returns:
|
||||
ApplyPoliciesResult with "inputs" (final GenericGuardrailAPIInputs) and
|
||||
"guardrail_errors" (list of {"guardrail_name", "message"} for each failure).
|
||||
"""
|
||||
guardrail_errors: List[GuardrailErrorEntry] = []
|
||||
|
||||
guardrail_name_set: set[str] = set()
|
||||
|
||||
if guardrail_names:
|
||||
guardrail_name_set.update(guardrail_names)
|
||||
|
||||
if policy_names:
|
||||
registry = get_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
verbose_proxy_logger.debug(
|
||||
"apply_policies: policy engine not initialized, skipping policy-resolved guardrails"
|
||||
)
|
||||
else:
|
||||
policies = registry.get_all_policies()
|
||||
for policy_name in policy_names:
|
||||
resolved = PolicyResolver.resolve_policy_guardrails(
|
||||
policy_name=policy_name,
|
||||
policies=policies,
|
||||
context=None,
|
||||
)
|
||||
guardrail_name_set.update(resolved.guardrails)
|
||||
|
||||
if not guardrail_name_set:
|
||||
return {"inputs": inputs, "guardrail_errors": guardrail_errors}
|
||||
|
||||
guardrail_registry = GuardrailRegistry()
|
||||
current_inputs = cast(GenericGuardrailAPIInputs, dict(inputs))
|
||||
|
||||
for guardrail_name in sorted(guardrail_name_set):
|
||||
callback = guardrail_registry.get_initialized_guardrail_callback(
|
||||
guardrail_name=guardrail_name
|
||||
)
|
||||
if callback is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"apply_policies: guardrail '%s' not found, skipping",
|
||||
guardrail_name,
|
||||
)
|
||||
continue
|
||||
if not isinstance(callback, CustomGuardrail):
|
||||
continue
|
||||
if "apply_guardrail" not in type(callback).__dict__:
|
||||
verbose_proxy_logger.debug(
|
||||
"apply_policies: guardrail '%s' has no apply_guardrail, skipping",
|
||||
guardrail_name,
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
current_inputs = await callback.apply_guardrail(
|
||||
inputs=current_inputs,
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
error_reason = str(e)
|
||||
verbose_proxy_logger.debug(
|
||||
"apply_policies: guardrail '%s' failed: %s",
|
||||
guardrail_name,
|
||||
error_reason,
|
||||
)
|
||||
guardrail_errors.append(
|
||||
GuardrailErrorEntry(
|
||||
guardrail_name=guardrail_name,
|
||||
message=error_reason,
|
||||
)
|
||||
)
|
||||
# Continue to next guardrail; current_inputs unchanged for this failure
|
||||
|
||||
return {"inputs": current_inputs, "guardrail_errors": guardrail_errors}
|
||||
|
||||
|
||||
class TestPoliciesAndGuardrailsRequest(BaseModel):
|
||||
"""Request body for POST /utils/test_policies_and_guardrails."""
|
||||
|
||||
policy_names: Optional[List[str]] = Field(default=None, description="Policy names to resolve guardrails from")
|
||||
guardrail_names: Optional[List[str]] = Field(default=None, description="Guardrail names to apply directly")
|
||||
inputs: dict = Field(description="GenericGuardrailAPIInputs, e.g. { \"texts\": [\"...\"] }")
|
||||
request_data: dict = Field(default_factory=dict, description="Request context (model, user_id, etc.)")
|
||||
input_type: Literal["request", "response"] = Field(default="request", description="Whether inputs are request or response")
|
||||
|
||||
|
||||
@router.post(
|
||||
"/utils/test_policies_and_guardrails",
|
||||
tags=["utils"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def test_policies_and_guardrails(
|
||||
request: Request,
|
||||
data: TestPoliciesAndGuardrailsRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Apply policies and/or guardrails to inputs (for compliance UI testing).
|
||||
|
||||
Runs all guardrails in order; failures are collected and returned in guardrail_errors.
|
||||
Returns inputs (possibly modified) and any guardrail errors so the UI can show which
|
||||
guardrails failed and why.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
from litellm.proxy.utils import handle_exception_on_proxy
|
||||
|
||||
try:
|
||||
inputs_typed = cast(GenericGuardrailAPIInputs, data.inputs)
|
||||
logging_obj = cast(LiteLLMLoggingObj, proxy_logging_obj)
|
||||
result = await apply_policies(
|
||||
policy_names=data.policy_names,
|
||||
inputs=inputs_typed,
|
||||
request_data=data.request_data,
|
||||
input_type=data.input_type,
|
||||
proxy_logging_obj=logging_obj,
|
||||
guardrail_names=data.guardrail_names,
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policy/validate",
|
||||
tags=["policy management"],
|
||||
|
|
@ -263,7 +445,9 @@ async def test_policy_matching(
|
|||
)
|
||||
|
||||
|
||||
POLICY_TEMPLATES_GITHUB_URL = "https://raw.githubusercontent.com/BerriAI/litellm/main/policy_templates.json"
|
||||
POLICY_TEMPLATES_GITHUB_URL = (
|
||||
"https://raw.githubusercontent.com/BerriAI/litellm/main/policy_templates.json"
|
||||
)
|
||||
|
||||
|
||||
def _load_policy_templates_from_local_backup() -> list:
|
||||
|
|
@ -322,3 +506,137 @@ async def get_policy_templates(
|
|||
)
|
||||
|
||||
return _load_policy_templates_from_local_backup()
|
||||
|
||||
|
||||
class EnrichTemplateRequest(BaseModel):
|
||||
template_id: str
|
||||
parameters: dict
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policy/templates/enrich",
|
||||
tags=["policy management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def enrich_policy_template(
|
||||
data: EnrichTemplateRequest,
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> dict:
|
||||
"""
|
||||
Enrich a policy template with LLM-discovered data (e.g. competitor names).
|
||||
|
||||
Calls an onboarded LLM to discover competitors for the given brand name,
|
||||
then returns enriched guardrailDefinitions with the discovered data populated.
|
||||
"""
|
||||
templates = _load_policy_templates_from_local_backup()
|
||||
template = next((t for t in templates if t.get("id") == data.template_id), None)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
|
||||
|
||||
llm_enrichment = template.get("llm_enrichment")
|
||||
if llm_enrichment is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Template does not support LLM enrichment",
|
||||
)
|
||||
|
||||
brand_name = data.parameters.get(llm_enrichment["parameter"], "")
|
||||
if not brand_name:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Parameter '{llm_enrichment['parameter']}' is required",
|
||||
)
|
||||
|
||||
prompt = llm_enrichment["prompt"].replace(
|
||||
"{{" + llm_enrichment["parameter"] + "}}", brand_name
|
||||
)
|
||||
|
||||
competitors = await _discover_competitors_via_llm(prompt)
|
||||
|
||||
enriched_definitions = _build_competitor_guardrail_definitions(
|
||||
template.get("guardrailDefinitions", []),
|
||||
competitors,
|
||||
brand_name,
|
||||
)
|
||||
|
||||
return {"guardrailDefinitions": enriched_definitions, "competitors": competitors}
|
||||
|
||||
|
||||
async def _discover_competitors_via_llm(prompt: str) -> list:
|
||||
"""Call an onboarded LLM to discover competitor names."""
|
||||
import litellm
|
||||
|
||||
try:
|
||||
response = await litellm.acompletion(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
temperature=0.3,
|
||||
)
|
||||
raw = response.choices[0].message.content or "" # type: ignore
|
||||
competitors = [
|
||||
line.strip().strip(".-) ").strip()
|
||||
for line in raw.strip().split("\n")
|
||||
if line.strip() and len(line.strip()) > 1
|
||||
]
|
||||
return competitors[:15]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("LLM competitor discovery failed: %s", e)
|
||||
return []
|
||||
|
||||
|
||||
def _build_competitor_guardrail_definitions(
|
||||
definitions: list,
|
||||
competitors: list,
|
||||
brand_name: str,
|
||||
) -> list:
|
||||
"""Build enriched guardrailDefinitions with competitor names populated."""
|
||||
import copy
|
||||
|
||||
enriched = copy.deepcopy(definitions)
|
||||
|
||||
output_blocked = [
|
||||
{"keyword": comp, "action": "BLOCK", "description": f"Competitor: {comp}"}
|
||||
for comp in competitors
|
||||
]
|
||||
|
||||
recommendation_blocked = []
|
||||
for comp in competitors:
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"try {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"use {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"switch to {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"consider {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
|
||||
comparison_blocked = []
|
||||
for comp in competitors:
|
||||
comparison_blocked.append(
|
||||
{"keyword": f"{comp} is better", "action": "BLOCK", "description": "Unfavorable comparison"}
|
||||
)
|
||||
comparison_blocked.append(
|
||||
{"keyword": f"better than {brand_name}", "action": "BLOCK", "description": "Unfavorable comparison"}
|
||||
)
|
||||
comparison_blocked.append(
|
||||
{"keyword": f"{brand_name} is worse", "action": "BLOCK", "description": "Unfavorable comparison"}
|
||||
)
|
||||
|
||||
blocked_words_map = {
|
||||
"competitor-output-blocker": output_blocked,
|
||||
"competitor-recommendation-filter": recommendation_blocked,
|
||||
"competitor-comparison-filter": comparison_blocked,
|
||||
}
|
||||
|
||||
for defn in enriched:
|
||||
guardrail_name = defn.get("guardrail_name", "")
|
||||
if guardrail_name in blocked_words_map:
|
||||
defn["litellm_params"]["blocked_words"] = blocked_words_map[guardrail_name]
|
||||
|
||||
return enriched
|
||||
|
|
|
|||
|
|
@ -1407,50 +1407,22 @@ async def insert_sso_user(
|
|||
if user_defined_values is None:
|
||||
raise ValueError("user_defined_values is None")
|
||||
|
||||
# Check if role_mappings is configured in SSO settings
|
||||
role_mappings_configured = False
|
||||
try:
|
||||
from litellm.proxy.utils import get_prisma_client_or_throw
|
||||
|
||||
prisma_client = get_prisma_client_or_throw(
|
||||
"Prisma client is None, connect a database to your proxy"
|
||||
)
|
||||
|
||||
# Get SSO config from dedicated table
|
||||
sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique(
|
||||
where={"id": "sso_config"}
|
||||
)
|
||||
|
||||
if sso_db_record and sso_db_record.sso_settings:
|
||||
sso_settings_dict = dict(sso_db_record.sso_settings)
|
||||
role_mappings_data = sso_settings_dict.get("role_mappings")
|
||||
role_mappings_configured = role_mappings_data is not None
|
||||
generic_user_role_mappings = os.getenv("GENERIC_USER_ROLE_MAPPINGS", None)
|
||||
if generic_user_role_mappings is not None:
|
||||
role_mappings_configured = True
|
||||
|
||||
except Exception as e:
|
||||
# If we can't check role_mappings, continue with existing logic
|
||||
verbose_proxy_logger.debug(
|
||||
f"Could not check role_mappings configuration: {e}. Using default behavior."
|
||||
)
|
||||
|
||||
# Apply default_internal_user_params
|
||||
if litellm.default_internal_user_params:
|
||||
# If role_mappings is configured and user_role is already set from SSO, preserve it
|
||||
if (
|
||||
role_mappings_configured
|
||||
and user_defined_values.get("user_role") is not None
|
||||
):
|
||||
# Preserve the SSO-extracted role if it's a valid LiteLLM role,
|
||||
# regardless of how it was determined (role_mappings, Microsoft app_roles,
|
||||
# GENERIC_USER_ROLE_ATTRIBUTE, custom SSO handler, etc.)
|
||||
sso_role = user_defined_values.get("user_role")
|
||||
if _should_use_role_from_sso_response(sso_role):
|
||||
# Preserve the SSO-extracted role, but apply other defaults
|
||||
preserved_role = user_defined_values.get("user_role")
|
||||
preserved_role = sso_role
|
||||
user_defined_values.update(litellm.default_internal_user_params) # type: ignore
|
||||
user_defined_values["user_role"] = preserved_role # Restore preserved role
|
||||
verbose_proxy_logger.debug(
|
||||
f"Preserved SSO-extracted role '{preserved_role}' (role_mappings configured)"
|
||||
f"Preserved SSO-extracted role '{preserved_role}'"
|
||||
)
|
||||
else:
|
||||
# Default behavior: update all values including role
|
||||
# SSO didn't provide a valid role, apply all defaults including role
|
||||
user_defined_values.update(litellm.default_internal_user_params) # type: ignore
|
||||
|
||||
# Set budget for internal users
|
||||
|
|
|
|||
|
|
@ -1,16 +1,21 @@
|
|||
"""
|
||||
Prometheus Auth Middleware
|
||||
|
||||
Pure ASGI middleware — avoids Starlette's BaseHTTPMiddleware which wraps
|
||||
streaming responses with receive_or_disconnect per chunk, blocking the
|
||||
event loop and causing severe throughput degradation under concurrent
|
||||
streaming load.
|
||||
"""
|
||||
from fastapi import Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import SpecialHeaders
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
|
||||
class PrometheusAuthMiddleware(BaseHTTPMiddleware):
|
||||
class PrometheusAuthMiddleware:
|
||||
"""
|
||||
Middleware to authenticate requests to the metrics endpoint
|
||||
|
||||
|
|
@ -24,8 +29,15 @@ class PrometheusAuthMiddleware(BaseHTTPMiddleware):
|
|||
```
|
||||
"""
|
||||
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
# Check if this is a request to the metrics endpoint
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] not in ("http", "websocket"):
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
request = Request(scope, receive)
|
||||
|
||||
if self._is_prometheus_metrics_endpoint(request):
|
||||
if self._should_run_auth_on_metrics_endpoint() is True:
|
||||
|
|
@ -38,15 +50,14 @@ class PrometheusAuthMiddleware(BaseHTTPMiddleware):
|
|||
or "",
|
||||
)
|
||||
except Exception as e:
|
||||
return JSONResponse(
|
||||
response = JSONResponse(
|
||||
status_code=401,
|
||||
content=f"Unauthorized access to metrics endpoint: {getattr(e, 'message', str(e))}",
|
||||
)
|
||||
await response(scope, receive, send)
|
||||
return
|
||||
|
||||
# Process the request and get the response
|
||||
response = await call_next(request)
|
||||
|
||||
return response
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
@staticmethod
|
||||
def _is_prometheus_metrics_endpoint(request: Request):
|
||||
|
|
|
|||
|
|
@ -2808,6 +2808,7 @@ class ProxyConfig:
|
|||
store_model_in_db = general_settings.get("store_model_in_db", False)
|
||||
if store_model_in_db is None:
|
||||
store_model_in_db = False
|
||||
general_settings["store_model_in_db"] = store_model_in_db
|
||||
### CUSTOM API KEY AUTH ###
|
||||
## pass filepath
|
||||
custom_auth = general_settings.get("custom_auth", None)
|
||||
|
|
@ -3845,7 +3846,7 @@ class ProxyConfig:
|
|||
"""
|
||||
Pull from DB, read general settings value
|
||||
"""
|
||||
global general_settings
|
||||
global general_settings, store_model_in_db
|
||||
if db_general_settings is None:
|
||||
return
|
||||
_general_settings = dict(db_general_settings)
|
||||
|
|
@ -3897,6 +3898,19 @@ class ProxyConfig:
|
|||
# For other types, convert to bool
|
||||
general_settings["store_prompts_in_spend_logs"] = bool(value)
|
||||
|
||||
## STORE MODEL IN DB ##
|
||||
if "store_model_in_db" in _general_settings:
|
||||
value = _general_settings["store_model_in_db"]
|
||||
if value is None:
|
||||
pass # Don't change store_model_in_db to None; keep current value
|
||||
elif isinstance(value, bool):
|
||||
store_model_in_db = value
|
||||
elif isinstance(value, str):
|
||||
store_model_in_db = value.lower() == "true"
|
||||
else:
|
||||
store_model_in_db = bool(value)
|
||||
general_settings["store_model_in_db"] = store_model_in_db
|
||||
|
||||
## MAXIMUM SPEND LOGS RETENTION PERIOD ##
|
||||
if "maximum_spend_logs_retention_period" in _general_settings:
|
||||
old_value = general_settings.get("maximum_spend_logs_retention_period")
|
||||
|
|
@ -5074,10 +5088,6 @@ async def async_data_generator(
|
|||
response=response,
|
||||
request_data=request_data,
|
||||
):
|
||||
verbose_proxy_logger.debug(
|
||||
"async_data_generator: received streaming chunk - {}".format(chunk)
|
||||
)
|
||||
|
||||
### CALL HOOKS ### - modify outgoing data
|
||||
chunk = await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -5433,6 +5443,31 @@ class ProxyStartupEvent:
|
|||
get_secret_bool("STORE_MODEL_IN_DB", store_model_in_db) or store_model_in_db
|
||||
)
|
||||
|
||||
# If store_model_in_db is still False, check DB for override.
|
||||
# This breaks the chicken-and-egg where DB has store_model_in_db=True
|
||||
# but YAML config has False.
|
||||
if store_model_in_db is not True and prisma_client is not None:
|
||||
try:
|
||||
_db_gs_record = await prisma_client.db.litellm_config.find_first(
|
||||
where={"param_name": "general_settings"}
|
||||
)
|
||||
if _db_gs_record is not None and isinstance(
|
||||
_db_gs_record.param_value, dict
|
||||
):
|
||||
_db_val = _db_gs_record.param_value.get("store_model_in_db")
|
||||
if _db_val is True or (
|
||||
isinstance(_db_val, str)
|
||||
and _db_val.lower() == "true"
|
||||
):
|
||||
store_model_in_db = True
|
||||
verbose_proxy_logger.info(
|
||||
"store_model_in_db=True loaded from DB, overriding config/env"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to check DB for store_model_in_db: %s", str(e)
|
||||
)
|
||||
|
||||
if store_model_in_db is True:
|
||||
# MEMORY LEAK FIX: Increase interval from 10s to 30s minimum
|
||||
# Frequent polling was causing excessive memory allocations
|
||||
|
|
@ -11321,6 +11356,7 @@ async def get_config_list(
|
|||
"max_request_size_mb": {"type": "Integer"},
|
||||
"max_response_size_mb": {"type": "Integer"},
|
||||
"pass_through_endpoints": {"type": "PydanticModel"},
|
||||
"store_model_in_db": {"type": "Boolean"},
|
||||
"store_prompts_in_spend_logs": {"type": "Boolean"},
|
||||
"maximum_spend_logs_retention_period": {"type": "String"},
|
||||
"mcp_internal_ip_ranges": {"type": "List"},
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import base64
|
|||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
import litellm
|
||||
|
|
@ -357,6 +357,19 @@ async def rag_ingest(
|
|||
# Parse request
|
||||
ingest_options, file_data, file_url, file_id = await parse_rag_ingest_request(request)
|
||||
|
||||
# INTERNAL_USER_VIEW_ONLY can ingest to existing vector stores only
|
||||
if (
|
||||
user_api_key_dict.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value
|
||||
and not ingest_options.get("vector_store", {}).get("vector_store_id")
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={
|
||||
"error": "internal_user_viewer role can only ingest files to an existing vector store. "
|
||||
"Provide 'vector_store_id' in ingest_options.vector_store."
|
||||
},
|
||||
)
|
||||
|
||||
# Add litellm data
|
||||
request_data: Dict[str, Any] = {}
|
||||
request_data = await add_litellm_data_to_request(
|
||||
|
|
|
|||
|
|
@ -233,6 +233,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
verification_tokens LiteLLM_VerificationToken[]
|
||||
organizations LiteLLM_OrganizationTable[]
|
||||
users LiteLLM_UserTable[]
|
||||
end_users LiteLLM_EndUserTable[]
|
||||
}
|
||||
|
||||
// Holds the MCP server configuration
|
||||
|
|
@ -403,7 +404,9 @@ model LiteLLM_EndUserTable {
|
|||
allowed_model_region String? // require all user requests to use models in this specific region
|
||||
default_model String? // use along with 'allowed_model_region'. if no available model in region, default to this model.
|
||||
budget_id String?
|
||||
object_permission_id String?
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
blocked Boolean @default(false)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -456,18 +456,26 @@ class ProxyLogging:
|
|||
def _init_litellm_callbacks(self, llm_router: Optional[Router] = None):
|
||||
self._add_proxy_hooks(llm_router)
|
||||
litellm.logging_callback_manager.add_litellm_callback(self.service_logging_obj) # type: ignore
|
||||
for callback in litellm.callbacks:
|
||||
|
||||
# Track string callbacks and their initialized instances so we can
|
||||
# replace them in-place, preventing duplicates (string + instance) in
|
||||
# litellm.callbacks which caused double-counting of metrics.
|
||||
string_callbacks_to_replace: Dict[int, CustomLogger] = {}
|
||||
|
||||
for idx, callback in enumerate(litellm.callbacks):
|
||||
if isinstance(callback, str):
|
||||
callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class( # type: ignore
|
||||
initialized_callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class(
|
||||
cast(_custom_logger_compatible_callbacks_literal, callback),
|
||||
internal_usage_cache=self.internal_usage_cache.dual_cache,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
if callback is None:
|
||||
continue
|
||||
if initialized_callback is not None:
|
||||
string_callbacks_to_replace[idx] = initialized_callback
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(callback)
|
||||
# Replace string entries in litellm.callbacks with initialized instances
|
||||
for idx, initialized_callback in string_callbacks_to_replace.items():
|
||||
litellm.callbacks[idx] = initialized_callback
|
||||
|
||||
async def update_request_status(
|
||||
self, litellm_call_id: str, status: Literal["success", "fail"]
|
||||
|
|
@ -1313,6 +1321,7 @@ class ProxyLogging:
|
|||
metadata = data.get("metadata", data.get("litellm_metadata", {})) or {}
|
||||
pipeline_managed: set = metadata.get("_pipeline_managed_guardrails", set())
|
||||
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
start_time = time.time()
|
||||
_callback = None
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
|
|||
from litellm.llms.bedrock.rerank.handler import BedrockRerankHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.together_ai.rerank.handler import TogetherAIRerank
|
||||
from litellm.llms.watsonx.common_utils import IBMWatsonXMixin
|
||||
from litellm.rerank_api.rerank_utils import get_optional_rerank_params
|
||||
from litellm.secret_managers.main import get_secret, get_secret_str
|
||||
from litellm.types.rerank import RerankResponse
|
||||
|
|
@ -29,7 +30,7 @@ async def arerank(
|
|||
model: str,
|
||||
query: str,
|
||||
documents: List[Union[str, Dict[str, Any]]],
|
||||
custom_llm_provider: Optional[Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage"]] = None,
|
||||
custom_llm_provider: Optional[Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage", "watsonx"]] = None,
|
||||
top_n: Optional[int] = None,
|
||||
rank_fields: Optional[List[str]] = None,
|
||||
return_documents: Optional[bool] = None,
|
||||
|
|
@ -85,6 +86,7 @@ def rerank( # noqa: PLR0915
|
|||
"deepinfra",
|
||||
"fireworks_ai",
|
||||
"voyage",
|
||||
"watsonx",
|
||||
]
|
||||
] = None,
|
||||
top_n: Optional[int] = None,
|
||||
|
|
@ -478,6 +480,31 @@ def rerank( # noqa: PLR0915
|
|||
or get_secret_str("VOYAGE_API_BASE")
|
||||
)
|
||||
|
||||
response = base_llm_http_handler.rerank(
|
||||
model=model,
|
||||
custom_llm_provider=_custom_llm_provider,
|
||||
provider_config=rerank_provider_config,
|
||||
optional_rerank_params=optional_rerank_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
timeout=optional_params.timeout,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
_is_async=_is_async,
|
||||
headers=headers or litellm.headers or {},
|
||||
client=client,
|
||||
model_response=model_response,
|
||||
)
|
||||
elif _custom_llm_provider == litellm.LlmProviders.WATSONX:
|
||||
credentials = IBMWatsonXMixin.get_watsonx_credentials(
|
||||
optional_params=dict(optional_params), api_key=dynamic_api_key, api_base=dynamic_api_base
|
||||
)
|
||||
|
||||
api_key = credentials["api_key"]
|
||||
api_base = credentials["api_base"]
|
||||
|
||||
if credentials.get("token") is not None:
|
||||
optional_rerank_params["token"] = credentials["token"]
|
||||
|
||||
response = base_llm_http_handler.rerank(
|
||||
model=model,
|
||||
custom_llm_provider=_custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from collections.abc import Sequence
|
|||
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
|
||||
|
||||
from openai.types.responses import ResponseFunctionToolCall
|
||||
from openai.types.responses.response_create_params import ResponseInputParam
|
||||
from openai.types.responses.tool_param import FunctionToolParam
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
|
|
@ -32,7 +33,6 @@ from litellm.types.llms.openai import (
|
|||
OpenAIWebSearchUserLocation,
|
||||
OutputTokensDetails,
|
||||
ResponseAPIUsage,
|
||||
ResponseInputParam,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStatus,
|
||||
|
|
@ -738,9 +738,25 @@ class LiteLLMCompletionResponsesConfig:
|
|||
|
||||
@staticmethod
|
||||
def _ensure_tool_results_have_corresponding_tool_calls(
|
||||
messages: List[Union[AllMessageValues, GenericChatCompletionMessage, ChatCompletionResponseMessage]],
|
||||
messages: Sequence[
|
||||
Union[
|
||||
AllMessageValues,
|
||||
GenericChatCompletionMessage,
|
||||
ChatCompletionResponseMessage,
|
||||
ChatCompletionMessageToolCall,
|
||||
Message,
|
||||
]
|
||||
],
|
||||
tools: Optional[List[Any]] = None,
|
||||
) -> List[Union[AllMessageValues, GenericChatCompletionMessage, ChatCompletionResponseMessage]]:
|
||||
) -> List[
|
||||
Union[
|
||||
AllMessageValues,
|
||||
GenericChatCompletionMessage,
|
||||
ChatCompletionResponseMessage,
|
||||
ChatCompletionMessageToolCall,
|
||||
Message,
|
||||
]
|
||||
]:
|
||||
"""
|
||||
Ensure that tool_result messages have corresponding tool_calls in the previous assistant message.
|
||||
|
||||
|
|
@ -755,11 +771,19 @@ class LiteLLMCompletionResponsesConfig:
|
|||
List of messages with tool_calls added to assistant messages when needed
|
||||
"""
|
||||
if not messages:
|
||||
return messages
|
||||
|
||||
# Create a deep copy to avoid modifying the original
|
||||
return list(messages)
|
||||
|
||||
# Create a deep copy to avoid modifying the original (use list() so we can mutate and return List)
|
||||
import copy
|
||||
fixed_messages = copy.deepcopy(messages)
|
||||
fixed_messages: List[
|
||||
Union[
|
||||
AllMessageValues,
|
||||
GenericChatCompletionMessage,
|
||||
ChatCompletionResponseMessage,
|
||||
ChatCompletionMessageToolCall,
|
||||
Message,
|
||||
]
|
||||
] = list(copy.deepcopy(messages))
|
||||
messages_to_remove = []
|
||||
|
||||
# Count non-tool messages to avoid removing all messages
|
||||
|
|
@ -1306,6 +1330,50 @@ class LiteLLMCompletionResponsesConfig:
|
|||
chat_completion_tools.append(cast(Union[ChatCompletionToolParam, OpenAIMcpServerTool], tool))
|
||||
return chat_completion_tools, web_search_options
|
||||
|
||||
@staticmethod
|
||||
def transform_chat_completion_tool_params_to_responses_api_tools(
|
||||
chat_completion_tools: Optional[
|
||||
List[Union[ChatCompletionToolParam, OpenAIMcpServerTool]]
|
||||
],
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Transform Chat Completion tool params (e.g. from guardrail output) back to
|
||||
Responses API request tool format. Inverse of
|
||||
transform_responses_api_tools_to_chat_completion_tools for the tools list.
|
||||
"""
|
||||
if chat_completion_tools is None or not chat_completion_tools:
|
||||
return []
|
||||
result: List[Dict[str, Any]] = []
|
||||
for tool in chat_completion_tools:
|
||||
if not isinstance(tool, dict):
|
||||
result.append(tool) # type: ignore
|
||||
continue
|
||||
if tool.get("type") == "function":
|
||||
fn = tool.get("function") or {}
|
||||
parameters = dict(fn.get("parameters", {}) or {})
|
||||
if not parameters or "type" not in parameters:
|
||||
parameters["type"] = "object"
|
||||
responses_tool: Dict[str, Any] = {
|
||||
"type": "function",
|
||||
"name": fn.get("name") or "",
|
||||
"description": fn.get("description") or "",
|
||||
"parameters": parameters,
|
||||
"strict": fn.get("strict", False) or False,
|
||||
}
|
||||
if tool.get("cache_control") is not None:
|
||||
responses_tool["cache_control"] = tool.get("cache_control")
|
||||
if tool.get("defer_loading") is not None:
|
||||
responses_tool["defer_loading"] = tool.get("defer_loading")
|
||||
if tool.get("allowed_callers") is not None:
|
||||
responses_tool["allowed_callers"] = tool.get("allowed_callers")
|
||||
if tool.get("input_examples") is not None:
|
||||
responses_tool["input_examples"] = tool.get("input_examples")
|
||||
result.append(responses_tool)
|
||||
else:
|
||||
# mcp or other: pass through unchanged
|
||||
result.append(dict(tool))
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def transform_chat_completion_tools_to_responses_tools(
|
||||
chat_completion_response: ModelResponse,
|
||||
|
|
|
|||
|
|
@ -600,8 +600,12 @@ def responses(
|
|||
# Update input and tools with provider-specific file IDs if managed files are used
|
||||
#########################################################
|
||||
model_file_id_mapping = kwargs.get("model_file_id_mapping")
|
||||
model_info_id = kwargs.get("model_info", {}).get("id") if isinstance(kwargs.get("model_info"), dict) else None
|
||||
|
||||
model_info_id = (
|
||||
kwargs.get("model_info", {}).get("id")
|
||||
if isinstance(kwargs.get("model_info"), dict)
|
||||
else None
|
||||
)
|
||||
|
||||
input = cast(
|
||||
Union[str, ResponseInputParam],
|
||||
update_responses_input_with_model_file_ids(
|
||||
|
|
@ -611,7 +615,7 @@ def responses(
|
|||
),
|
||||
)
|
||||
local_vars["input"] = input
|
||||
|
||||
|
||||
# Update tools with provider-specific file IDs if needed
|
||||
if tools:
|
||||
tools = cast(
|
||||
|
|
@ -696,7 +700,10 @@ def responses(
|
|||
)
|
||||
)
|
||||
|
||||
# Pre Call logging
|
||||
# Pre Call logging - preserve metadata for custom callbacks
|
||||
# When called from completion bridge (codex models), metadata is in litellm_metadata
|
||||
metadata_for_callbacks = metadata or kwargs.get("litellm_metadata") or {}
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model=model,
|
||||
user=user,
|
||||
|
|
@ -705,7 +712,7 @@ def responses(
|
|||
**responses_api_request_params,
|
||||
"aresponses": _is_async,
|
||||
"litellm_call_id": litellm_call_id,
|
||||
"metadata": metadata,
|
||||
"metadata": metadata_for_callbacks,
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -113,12 +113,12 @@ from litellm.router_utils.handle_error import (
|
|||
from litellm.router_utils.pre_call_checks.model_rate_limit_check import (
|
||||
ModelRateLimitingCheck,
|
||||
)
|
||||
from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
|
||||
DeploymentAffinityCheck,
|
||||
)
|
||||
from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import (
|
||||
PromptCachingDeploymentCheck,
|
||||
)
|
||||
from litellm.router_utils.pre_call_checks.responses_api_deployment_check import (
|
||||
ResponsesApiDeploymentCheck,
|
||||
)
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
|
||||
increment_deployment_failures_for_current_minute,
|
||||
increment_deployment_successes_for_current_minute,
|
||||
|
|
@ -293,6 +293,7 @@ class Router:
|
|||
router_general_settings: Optional[
|
||||
RouterGeneralSettings
|
||||
] = RouterGeneralSettings(),
|
||||
deployment_affinity_ttl_seconds: int = 3600,
|
||||
ignore_invalid_deployments: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
|
|
@ -326,6 +327,7 @@ class Router:
|
|||
routing_strategy_args (dict): Additional args for latency-based routing. Defaults to {}.
|
||||
alerting_config (AlertingConfig): Slack alerting configuration. Defaults to None.
|
||||
provider_budget_config (ProviderBudgetConfig): Provider budget configuration. Use this to set llm_provider budget limits. example $100/day to OpenAI, $100/day to Azure, etc. Defaults to None.
|
||||
deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600.
|
||||
ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error.
|
||||
Returns:
|
||||
Router: An instance of the litellm.Router class.
|
||||
|
|
@ -604,6 +606,7 @@ class Router:
|
|||
litellm.failure_callback = [self.deployment_callback_on_failure]
|
||||
self.routing_strategy_args = routing_strategy_args
|
||||
self.provider_budget_config = provider_budget_config
|
||||
self.deployment_affinity_ttl_seconds = deployment_affinity_ttl_seconds
|
||||
self.router_budget_logger: Optional[RouterBudgetLimiting] = None
|
||||
if RouterBudgetLimiting.should_init_router_budget_limiter(
|
||||
model_list=model_list, provider_budget_config=self.provider_budget_config
|
||||
|
|
@ -1184,26 +1187,78 @@ class Router:
|
|||
def add_optional_pre_call_checks(
|
||||
self, optional_pre_call_checks: Optional[OptionalPreCallChecks]
|
||||
):
|
||||
if optional_pre_call_checks is not None:
|
||||
for pre_call_check in optional_pre_call_checks:
|
||||
_callback: Optional[CustomLogger] = None
|
||||
if pre_call_check == "prompt_caching":
|
||||
_callback = PromptCachingDeploymentCheck(cache=self.cache)
|
||||
elif pre_call_check == "router_budget_limiting":
|
||||
_callback = RouterBudgetLimiting(
|
||||
dual_cache=self.cache,
|
||||
provider_budget_config=self.provider_budget_config,
|
||||
model_list=self.model_list,
|
||||
)
|
||||
elif pre_call_check == "responses_api_deployment_check":
|
||||
_callback = ResponsesApiDeploymentCheck()
|
||||
elif pre_call_check == "enforce_model_rate_limits":
|
||||
_callback = ModelRateLimitingCheck(dual_cache=self.cache)
|
||||
if _callback is not None:
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
self.optional_callbacks.append(_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_callback)
|
||||
if optional_pre_call_checks is None:
|
||||
return
|
||||
|
||||
# ---------------------------------------------------------------------
|
||||
# Unified deployment affinity (session stickiness)
|
||||
# ---------------------------------------------------------------------
|
||||
enable_user_key_affinity = "deployment_affinity" in optional_pre_call_checks
|
||||
enable_responses_api_affinity = (
|
||||
"responses_api_deployment_check" in optional_pre_call_checks
|
||||
)
|
||||
if enable_user_key_affinity or enable_responses_api_affinity:
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
|
||||
existing_affinity_callback: Optional[DeploymentAffinityCheck] = None
|
||||
for cb in self.optional_callbacks:
|
||||
if isinstance(cb, DeploymentAffinityCheck):
|
||||
existing_affinity_callback = cb
|
||||
break
|
||||
|
||||
if existing_affinity_callback is not None:
|
||||
existing_affinity_callback.enable_user_key_affinity = (
|
||||
existing_affinity_callback.enable_user_key_affinity
|
||||
or enable_user_key_affinity
|
||||
)
|
||||
existing_affinity_callback.enable_responses_api_affinity = (
|
||||
existing_affinity_callback.enable_responses_api_affinity
|
||||
or enable_responses_api_affinity
|
||||
)
|
||||
existing_affinity_callback.ttl_seconds = (
|
||||
self.deployment_affinity_ttl_seconds
|
||||
)
|
||||
else:
|
||||
affinity_callback = DeploymentAffinityCheck(
|
||||
cache=self.cache,
|
||||
ttl_seconds=self.deployment_affinity_ttl_seconds,
|
||||
enable_user_key_affinity=enable_user_key_affinity,
|
||||
enable_responses_api_affinity=enable_responses_api_affinity,
|
||||
)
|
||||
self.optional_callbacks.append(affinity_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(
|
||||
affinity_callback
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------
|
||||
# Remaining optional pre-call checks
|
||||
# ---------------------------------------------------------------------
|
||||
for pre_call_check in optional_pre_call_checks:
|
||||
_callback: Optional[CustomLogger] = None
|
||||
if pre_call_check in (
|
||||
"deployment_affinity",
|
||||
"responses_api_deployment_check",
|
||||
):
|
||||
continue
|
||||
if pre_call_check == "prompt_caching":
|
||||
_callback = PromptCachingDeploymentCheck(cache=self.cache)
|
||||
elif pre_call_check == "router_budget_limiting":
|
||||
_callback = RouterBudgetLimiting(
|
||||
dual_cache=self.cache,
|
||||
provider_budget_config=self.provider_budget_config,
|
||||
model_list=self.model_list,
|
||||
)
|
||||
elif pre_call_check == "enforce_model_rate_limits":
|
||||
_callback = ModelRateLimitingCheck(dual_cache=self.cache)
|
||||
|
||||
if _callback is None:
|
||||
continue
|
||||
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
self.optional_callbacks.append(_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_callback)
|
||||
|
||||
def print_deployment(self, deployment: dict):
|
||||
"""
|
||||
|
|
@ -7600,9 +7655,16 @@ class Router:
|
|||
Used by `.get_model_list` to get model list from model alias.
|
||||
"""
|
||||
returned_models: List[DeploymentTypedDict] = []
|
||||
for model_alias, model_value in self.model_group_alias.items():
|
||||
if model_name is not None and model_alias != model_name:
|
||||
continue
|
||||
|
||||
if model_name is not None:
|
||||
# Fast path: direct dict lookup avoids scanning all aliases for non-alias model names.
|
||||
if model_name not in self.model_group_alias:
|
||||
return returned_models
|
||||
alias_items = [(model_name, self.model_group_alias[model_name])]
|
||||
else:
|
||||
alias_items = list(self.model_group_alias.items())
|
||||
|
||||
for model_alias, model_value in alias_items:
|
||||
if isinstance(model_value, str):
|
||||
_router_model_name: str = model_value
|
||||
elif isinstance(model_value, dict):
|
||||
|
|
@ -9099,4 +9161,3 @@ class Router:
|
|||
litellm._async_failure_callback = []
|
||||
self.retry_policy = None
|
||||
self.flush_cache()
|
||||
|
||||
|
|
|
|||
|
|
@ -335,13 +335,14 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
):
|
||||
lowest_tpm = float("inf")
|
||||
potential_deployments = [] # if multiple deployments have the same low value
|
||||
deployment_lookup = {
|
||||
deployment.get("model_info", {}).get("id"): deployment
|
||||
for deployment in healthy_deployments
|
||||
}
|
||||
for item, item_tpm in all_deployments.items():
|
||||
## get the item from model list
|
||||
_deployment = None
|
||||
item = item.split(":")[0]
|
||||
for m in healthy_deployments:
|
||||
if item == m["model_info"]["id"]:
|
||||
_deployment = m
|
||||
_deployment = deployment_lookup.get(item)
|
||||
if _deployment is None:
|
||||
continue # skip to next one
|
||||
elif item_tpm is None:
|
||||
|
|
|
|||
|
|
@ -58,7 +58,7 @@ def filter_team_based_models(
|
|||
request_team_id = metadata.get("user_api_key_team_id") or litellm_metadata.get(
|
||||
"user_api_key_team_id"
|
||||
)
|
||||
ids_to_remove = []
|
||||
ids_to_remove = set()
|
||||
if isinstance(healthy_deployments, dict):
|
||||
return healthy_deployments
|
||||
for deployment in healthy_deployments:
|
||||
|
|
@ -67,7 +67,7 @@ def filter_team_based_models(
|
|||
if model_team_id is None:
|
||||
continue
|
||||
if model_team_id != request_team_id:
|
||||
ids_to_remove.append(deployment.get("model_info", {}).get("id"))
|
||||
ids_to_remove.add(_model_info.get("id"))
|
||||
|
||||
return [
|
||||
deployment
|
||||
|
|
@ -125,4 +125,3 @@ def filter_web_search_deployments(
|
|||
if len(healthy_deployments) > 0 and len(final_deployments) == 0:
|
||||
verbose_logger.warning("No deployments support web search for request")
|
||||
return final_deployments
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,396 @@
|
|||
"""
|
||||
Unified deployment affinity (session stickiness) for the Router.
|
||||
|
||||
Features (independently enable-able):
|
||||
1. Responses API continuity: when a `previous_response_id` is provided, route to the
|
||||
deployment that generated the original response (highest priority).
|
||||
2. API-key affinity: map an API key hash -> deployment id for a TTL and re-use that
|
||||
deployment for subsequent requests to the same router deployment model name
|
||||
(alias-safe, aligns to `model_map_information.model_map_key`).
|
||||
|
||||
This is designed to support "implicit prompt caching" scenarios (no explicit cache_control),
|
||||
where routing to a consistent deployment is still beneficial.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
class DeploymentAffinityCacheValue(TypedDict):
|
||||
model_id: str
|
||||
|
||||
|
||||
class DeploymentAffinityCheck(CustomLogger):
|
||||
"""
|
||||
Router deployment affinity callback.
|
||||
|
||||
NOTE: This is a Router-only callback intended to be wired through
|
||||
`Router(optional_pre_call_checks=[...])`.
|
||||
"""
|
||||
|
||||
CACHE_KEY_PREFIX = "deployment_affinity:v1"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cache: DualCache,
|
||||
ttl_seconds: int,
|
||||
enable_user_key_affinity: bool,
|
||||
enable_responses_api_affinity: bool,
|
||||
):
|
||||
super().__init__()
|
||||
self.cache = cache
|
||||
self.ttl_seconds = ttl_seconds
|
||||
self.enable_user_key_affinity = enable_user_key_affinity
|
||||
self.enable_responses_api_affinity = enable_responses_api_affinity
|
||||
|
||||
@staticmethod
|
||||
def _looks_like_sha256_hex(value: str) -> bool:
|
||||
if len(value) != 64:
|
||||
return False
|
||||
try:
|
||||
int(value, 16)
|
||||
except ValueError:
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _hash_user_key(user_key: str) -> str:
|
||||
"""
|
||||
Hash user identifiers before storing them in cache keys.
|
||||
|
||||
This avoids putting raw API keys / user identifiers into Redis keys (and therefore
|
||||
into logs/metrics), while keeping the cache key stable and a fixed length.
|
||||
"""
|
||||
# If the proxy already provides a stable SHA-256 (e.g. `metadata.user_api_key_hash`),
|
||||
# keep it as-is to avoid double-hashing and to make correlation/debugging possible.
|
||||
if DeploymentAffinityCheck._looks_like_sha256_hex(user_key):
|
||||
return user_key.lower()
|
||||
|
||||
return hashlib.sha256(user_key.encode("utf-8")).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _get_model_map_key_from_litellm_model_name(litellm_model_name: str) -> Optional[str]:
|
||||
"""
|
||||
Best-effort derivation of a stable "model map key" for affinity scoping.
|
||||
|
||||
The intent is to align with `standard_logging_payload.model_map_information.model_map_key`,
|
||||
which is typically the base model identifier (stable across deployments/endpoints).
|
||||
|
||||
Notes:
|
||||
- When the model name is in "provider/model" format, the provider prefix is stripped.
|
||||
- For Azure, the string after "azure/" is commonly an *Azure deployment name*, which may
|
||||
differ across instances. If `base_model` is not explicitly set, we skip deriving a
|
||||
model-map key from the model string to avoid generating unstable keys.
|
||||
"""
|
||||
if not litellm_model_name:
|
||||
return None
|
||||
|
||||
if "/" not in litellm_model_name:
|
||||
return litellm_model_name
|
||||
|
||||
provider_prefix, remainder = litellm_model_name.split("/", 1)
|
||||
if provider_prefix == "azure":
|
||||
return None
|
||||
|
||||
return remainder
|
||||
|
||||
@staticmethod
|
||||
def _get_model_map_key_from_deployment(deployment: dict) -> Optional[str]:
|
||||
"""
|
||||
Derive a stable model-map key from a router deployment dict.
|
||||
|
||||
Primary source: `deployment.model_name` (Router's canonical group name after
|
||||
alias resolution). This is stable across provider-specific deployments (e.g.,
|
||||
Azure/Vertex/Bedrock for the same logical model) and aligns with
|
||||
`model_map_information.model_map_key` in standard logging.
|
||||
|
||||
Prefer `base_model` when available (important for Azure), otherwise fall back to
|
||||
parsing `litellm_params.model`.
|
||||
"""
|
||||
model_name = deployment.get("model_name")
|
||||
if isinstance(model_name, str) and model_name:
|
||||
return model_name
|
||||
|
||||
model_info = deployment.get("model_info")
|
||||
if isinstance(model_info, dict):
|
||||
base_model = model_info.get("base_model")
|
||||
if isinstance(base_model, str) and base_model:
|
||||
return base_model
|
||||
|
||||
litellm_params = deployment.get("litellm_params")
|
||||
if isinstance(litellm_params, dict):
|
||||
base_model = litellm_params.get("base_model")
|
||||
if isinstance(base_model, str) and base_model:
|
||||
return base_model
|
||||
litellm_model_name = litellm_params.get("model")
|
||||
if isinstance(litellm_model_name, str) and litellm_model_name:
|
||||
return DeploymentAffinityCheck._get_model_map_key_from_litellm_model_name(
|
||||
litellm_model_name
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _get_stable_model_map_key_from_deployments(
|
||||
healthy_deployments: List[dict],
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Only use model-map key scoping when it is stable across the deployment set.
|
||||
|
||||
This prevents accidentally keying on per-deployment identifiers like Azure deployment
|
||||
names (when `base_model` is not configured).
|
||||
"""
|
||||
if not healthy_deployments:
|
||||
return None
|
||||
|
||||
keys: List[str] = []
|
||||
for deployment in healthy_deployments:
|
||||
key = DeploymentAffinityCheck._get_model_map_key_from_deployment(deployment)
|
||||
if key is None:
|
||||
return None
|
||||
keys.append(key)
|
||||
|
||||
unique_keys = set(keys)
|
||||
if len(unique_keys) != 1:
|
||||
return None
|
||||
return keys[0]
|
||||
|
||||
@staticmethod
|
||||
def _shorten_for_logs(value: str, keep: int = 8) -> str:
|
||||
if len(value) <= keep:
|
||||
return value
|
||||
return f"{value[:keep]}..."
|
||||
|
||||
@classmethod
|
||||
def get_affinity_cache_key(cls, model_group: str, user_key: str) -> str:
|
||||
hashed_user_key = cls._hash_user_key(user_key=user_key)
|
||||
return f"{cls.CACHE_KEY_PREFIX}:{model_group}:{hashed_user_key}"
|
||||
|
||||
@staticmethod
|
||||
def _get_user_key_from_metadata_dict(metadata: dict) -> Optional[str]:
|
||||
# NOTE: affinity is keyed on the *API key hash* provided by the proxy (not the
|
||||
# OpenAI `user` parameter, which is an end-user identifier).
|
||||
user_key = metadata.get("user_api_key_hash")
|
||||
if user_key is None:
|
||||
return None
|
||||
return str(user_key)
|
||||
|
||||
@staticmethod
|
||||
def _iter_metadata_dicts(request_kwargs: dict) -> List[dict]:
|
||||
"""
|
||||
Return all metadata dicts available on the request.
|
||||
|
||||
Depending on the endpoint, Router may populate `metadata` or `litellm_metadata`.
|
||||
Users may also send one or both, so we check both (rather than using `or`).
|
||||
"""
|
||||
metadata_dicts: List[dict] = []
|
||||
for key in ("litellm_metadata", "metadata"):
|
||||
md = request_kwargs.get(key)
|
||||
if isinstance(md, dict):
|
||||
metadata_dicts.append(md)
|
||||
return metadata_dicts
|
||||
|
||||
@staticmethod
|
||||
def _get_user_key_from_request_kwargs(request_kwargs: dict) -> Optional[str]:
|
||||
"""
|
||||
Extract a stable affinity key from request kwargs.
|
||||
|
||||
Source (proxy): `metadata.user_api_key_hash`
|
||||
|
||||
Note: the OpenAI `user` parameter is an end-user identifier and is intentionally
|
||||
not used for deployment affinity.
|
||||
"""
|
||||
# Check metadata dicts (Proxy usage)
|
||||
for metadata in DeploymentAffinityCheck._iter_metadata_dicts(request_kwargs):
|
||||
user_key = DeploymentAffinityCheck._get_user_key_from_metadata_dict(
|
||||
metadata=metadata
|
||||
)
|
||||
if user_key is not None:
|
||||
return user_key
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _find_deployment_by_model_id(healthy_deployments: List[dict], model_id: str) -> Optional[dict]:
|
||||
for deployment in healthy_deployments:
|
||||
model_info = deployment.get("model_info")
|
||||
if not isinstance(model_info, dict):
|
||||
continue
|
||||
deployment_model_id = model_info.get("id")
|
||||
if deployment_model_id is not None and str(deployment_model_id) == str(model_id):
|
||||
return deployment
|
||||
return None
|
||||
|
||||
async def async_filter_deployments(
|
||||
self,
|
||||
model: str,
|
||||
healthy_deployments: List,
|
||||
messages: Optional[List[AllMessageValues]],
|
||||
request_kwargs: Optional[dict] = None,
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
) -> List[dict]:
|
||||
"""
|
||||
Optionally filter healthy deployments based on:
|
||||
1. `previous_response_id` (Responses API continuity) [highest priority]
|
||||
2. cached API-key deployment affinity
|
||||
"""
|
||||
request_kwargs = request_kwargs or {}
|
||||
typed_healthy_deployments = cast(List[dict], healthy_deployments)
|
||||
|
||||
# 1) Responses API continuity (high priority)
|
||||
if self.enable_responses_api_affinity:
|
||||
previous_response_id = request_kwargs.get("previous_response_id")
|
||||
if previous_response_id is not None:
|
||||
responses_model_id = ResponsesAPIRequestUtils.get_model_id_from_response_id(str(previous_response_id))
|
||||
if responses_model_id is not None:
|
||||
deployment = self._find_deployment_by_model_id(
|
||||
healthy_deployments=typed_healthy_deployments,
|
||||
model_id=responses_model_id,
|
||||
)
|
||||
if deployment is not None:
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: previous_response_id pinning -> deployment=%s",
|
||||
responses_model_id,
|
||||
)
|
||||
return [deployment]
|
||||
|
||||
# 2) User key -> deployment affinity
|
||||
if not self.enable_user_key_affinity:
|
||||
return typed_healthy_deployments
|
||||
|
||||
user_key = self._get_user_key_from_request_kwargs(request_kwargs=request_kwargs)
|
||||
if user_key is None:
|
||||
return typed_healthy_deployments
|
||||
|
||||
stable_model_map_key = self._get_stable_model_map_key_from_deployments(
|
||||
healthy_deployments=typed_healthy_deployments
|
||||
)
|
||||
if stable_model_map_key is None:
|
||||
return typed_healthy_deployments
|
||||
|
||||
cache_key = self.get_affinity_cache_key(
|
||||
model_group=stable_model_map_key, user_key=user_key
|
||||
)
|
||||
cache_result = await self.cache.async_get_cache(key=cache_key)
|
||||
|
||||
model_id: Optional[str] = None
|
||||
if isinstance(cache_result, dict):
|
||||
model_id = cast(Optional[str], cache_result.get("model_id"))
|
||||
elif isinstance(cache_result, str):
|
||||
# Backwards / safety: allow raw string values.
|
||||
model_id = cache_result
|
||||
|
||||
if not model_id:
|
||||
return typed_healthy_deployments
|
||||
|
||||
deployment = self._find_deployment_by_model_id(
|
||||
healthy_deployments=typed_healthy_deployments,
|
||||
model_id=model_id,
|
||||
)
|
||||
if deployment is None:
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: pinned deployment=%s not found in healthy_deployments",
|
||||
model_id,
|
||||
)
|
||||
return typed_healthy_deployments
|
||||
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: api-key affinity hit -> deployment=%s user_key=%s",
|
||||
model_id,
|
||||
self._shorten_for_logs(user_key),
|
||||
)
|
||||
return [deployment]
|
||||
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Persist/update the API-key -> deployment mapping for this request.
|
||||
|
||||
Why pre-call?
|
||||
- LiteLLM runs async success callbacks via a background logging worker for performance.
|
||||
- We want affinity to be immediately available for subsequent requests.
|
||||
"""
|
||||
if not self.enable_user_key_affinity:
|
||||
return None
|
||||
|
||||
user_key = self._get_user_key_from_request_kwargs(request_kwargs=kwargs)
|
||||
if user_key is None:
|
||||
return None
|
||||
|
||||
metadata_dicts = self._iter_metadata_dicts(kwargs)
|
||||
|
||||
model_info = kwargs.get("model_info")
|
||||
if not isinstance(model_info, dict):
|
||||
model_info = None
|
||||
|
||||
if model_info is None:
|
||||
for metadata in metadata_dicts:
|
||||
maybe_model_info = metadata.get("model_info")
|
||||
if isinstance(maybe_model_info, dict):
|
||||
model_info = maybe_model_info
|
||||
break
|
||||
|
||||
if model_info is None:
|
||||
# Router sets `model_info` after selecting a deployment. If it's missing, this is
|
||||
# likely a non-router call or a call path that doesn't support affinity.
|
||||
return None
|
||||
|
||||
model_id = model_info.get("id")
|
||||
if not model_id:
|
||||
verbose_router_logger.warning(
|
||||
"DeploymentAffinityCheck: model_id missing; skipping affinity cache update."
|
||||
)
|
||||
return None
|
||||
|
||||
# Scope affinity by the Router deployment model name (alias-safe, consistent across
|
||||
# heterogeneous providers, and matches standard logging's `model_map_key`).
|
||||
deployment_model_name: Optional[str] = None
|
||||
for metadata in metadata_dicts:
|
||||
maybe_deployment_model_name = metadata.get("deployment_model_name")
|
||||
if isinstance(maybe_deployment_model_name, str) and maybe_deployment_model_name:
|
||||
deployment_model_name = maybe_deployment_model_name
|
||||
break
|
||||
|
||||
if not deployment_model_name:
|
||||
verbose_router_logger.warning(
|
||||
"DeploymentAffinityCheck: deployment_model_name missing; skipping affinity cache update. model_id=%s",
|
||||
model_id,
|
||||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
cache_key = self.get_affinity_cache_key(
|
||||
model_group=deployment_model_name, user_key=user_key
|
||||
)
|
||||
await self.cache.async_set_cache(
|
||||
cache_key,
|
||||
DeploymentAffinityCacheValue(model_id=str(model_id)),
|
||||
ttl=self.ttl_seconds,
|
||||
)
|
||||
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: set affinity mapping model_map_key=%s deployment=%s ttl=%s user_key=%s",
|
||||
deployment_model_name,
|
||||
model_id,
|
||||
self.ttl_seconds,
|
||||
self._shorten_for_logs(user_key),
|
||||
)
|
||||
except Exception as e:
|
||||
# Non-blocking: affinity is a best-effort optimization.
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: failed to set affinity cache. model_map_key=%s error=%s",
|
||||
deployment_model_name,
|
||||
e,
|
||||
)
|
||||
|
||||
return None
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue