Merge remote-tracking branch 'origin' into litellm_deleted_keys_team

This commit is contained in:
yuneng-jiang 2026-01-16 09:55:52 -08:00
commit 33ff58b70a
111 changed files with 7483 additions and 1921 deletions

View file

@ -144,8 +144,8 @@ jobs:
pip install "google-generativeai==0.3.2"
pip install "google-cloud-aiplatform==1.43.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install lunary==0.2.5
pip install "azure-identity==1.16.1"
@ -260,8 +260,8 @@ jobs:
pip install "google-generativeai==0.3.2"
pip install "google-cloud-aiplatform==1.43.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install lunary==0.2.5
pip install "azure-identity==1.16.1"
@ -367,8 +367,8 @@ jobs:
pip install "google-generativeai==0.3.2"
pip install "google-cloud-aiplatform==1.43.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install lunary==0.2.5
pip install "azure-identity==1.16.1"
@ -637,8 +637,8 @@ jobs:
pip install "google-generativeai==0.3.2"
pip install "google-cloud-aiplatform==1.43.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install "langfuse>=2.0.0"
pip install "logfire==0.29.0"
@ -759,8 +759,8 @@ jobs:
pip install "google-cloud-aiplatform==1.43.0"
pip install "google-genai==1.22.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install lunary==0.2.5
pip install "azure-identity==1.16.1"
@ -865,8 +865,8 @@ jobs:
pip install "google-cloud-aiplatform==1.43.0"
pip install "google-genai==1.22.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install lunary==0.2.5
pip install "azure-identity==1.16.1"
@ -972,8 +972,8 @@ jobs:
pip install "google-cloud-aiplatform==1.43.0"
pip install "google-genai==1.22.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install lunary==0.2.5
pip install "azure-identity==1.16.1"
@ -1198,7 +1198,7 @@ jobs:
pip install "pytest-asyncio==0.21.1"
pip install "respx==0.22.0"
pip install "pydantic==2.10.2"
pip install "boto3==1.36.0"
pip install "boto3==1.40.61"
# Run pytest and generate JUnit XML report
- run:
name: Run tests
@ -1879,7 +1879,7 @@ jobs:
pip install aiohttp
pip install openai
pip install click
pip install "boto3==1.36.0"
pip install "boto3==1.40.61"
pip install jinja2
pip install "tokenizers==0.20.0"
pip install "uvloop==0.21.0"
@ -2176,8 +2176,8 @@ jobs:
pip install "google-generativeai==0.3.2"
pip install "google-cloud-aiplatform==1.43.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install "langfuse>=2.0.0"
pip install "logfire==0.29.0"
@ -2316,8 +2316,8 @@ jobs:
pip install "google-generativeai==0.3.2"
pip install "google-cloud-aiplatform==1.43.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install "langchain_mcp_adapters==0.0.5"
pip install "langfuse>=2.0.0"
@ -2462,8 +2462,8 @@ jobs:
pip install "google-generativeai==0.3.2"
pip install "google-cloud-aiplatform==1.43.0"
pip install pyarrow
pip install "boto3==1.36.0"
pip install "aioboto3==13.4.0"
pip install "boto3==1.40.61"
pip install "aioboto3==15.5.0"
pip install langchain
pip install "langfuse>=2.0.0"
pip install "logfire==0.29.0"
@ -3118,7 +3118,7 @@ jobs:
pip install "pytest==7.3.1"
pip install "pytest-mock==3.12.0"
pip install "pytest-asyncio==0.21.1"
pip install "boto3==1.36.0"
pip install "boto3==1.40.61"
pip install "mypy==1.18.2"
pip install pyarrow
pip install numpydoc

View file

@ -1,3 +1,3 @@
ignore:
- vulnerability: CVE-2019-1010022
reason: no fixed glibc package is available yet in the Wolfi repositories, so this is ignored temporarily until an upstream release exists
- vulnerability: CVE-2026-22184
reason: no fixed zlib package is available yet in the Wolfi repositories, so this is ignored temporarily until an upstream release exists

View file

@ -129,11 +129,14 @@ run_grype_scans() {
"CVE-2025-13836" # Python 3.13 HTTP response reading OOM/DoS - no fix available in base image
"CVE-2025-12084" # Python 3.13 xml.dom.minidom quadratic algorithm - no fix available in base image
"CVE-2025-60876" # BusyBox wget HTTP request splitting - no fix available in Chainguard Wolfi base image
"CVE-2026-0861" # Wolfi glibc still flagged even on 2.42-r5; upstream patched build unavailable yet
"CVE-2010-4756" # glibc glob DoS - awaiting patched Wolfi glibc build
"CVE-2019-1010022" # glibc stack guard bypass - awaiting patched Wolfi glibc build
"CVE-2019-1010023" # glibc ldd remap issue - awaiting patched Wolfi glibc build
"CVE-2019-1010024" # glibc ASLR mitigation bypass - awaiting patched Wolfi glibc build
"CVE-2019-1010025" # glibc pthread heap address leak - awaiting patched Wolfi glibc build
"CVE-2026-22184" # zlib untgz buffer overflow - untgz unused + no fixed Wolfi build yet
"GHSA-58pv-8j8x-9vj2" # jaraco.context path traversal - setuptools vendored only (v5.3.0), not used in application code (using v6.1.0+)
)
# Build JSON array of allowlisted CVE IDs for jq

View file

@ -0,0 +1,468 @@
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 works with all LLM providers that support tool calling:
- ✅ Anthropic (Claude)
- ✅ OpenAI (GPT-4, GPT-3.5)
- ✅ AWS Bedrock (Claude, Titan)
- ✅ Google Vertex AI (Claude, Gemini)
- ✅ Azure OpenAI
- ✅ And all other providers with tool calling support
## 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)

View file

@ -12,100 +12,340 @@ LiteLLM supports SAP Generative AI Hub's Orchestration Service.
| Supported Endpoints | `/chat/completions`, `/embeddings` |
| API Reference | [SAP AI Core Documentation](https://help.sap.com/docs/sap-ai-core) |
## Prerequisites
Before you begin, ensure you have:
1. **SAP BTP Account** with access to SAP AI Core
2. **AI Core Service Instance** provisioned in your subaccount
3. **Service Key** created for your AI Core instance (this contains your credentials)
4. **Resource Group** with deployed AI models (check with your SAP administrator)
:::tip Where to Find Your Credentials
Your credentials come from the **Service Key** you create in SAP BTP Cockpit:
1. Navigate to your **Subaccount** → **Instances and Subscriptions**
2. Find your **AI Core** instance and click on it
3. Go to **Service Keys** and create one (or use existing)
4. The JSON contains all values needed below
The service key JSON looks like this:
```json
{
"clientid": "sb-abc123...",
"clientsecret": "xyz789...",
"url": "https://myinstance.authentication.eu10.hana.ondemand.com",
"serviceurls": {
"AI_API_URL": "https://api.ai.prod.eu-central-1.aws.ml.hana.ondemand.com"
}
}
```
:::info Resource Group
The resource group is typically configured separately in your AI Core deployment, not in the service key itself. You can set it via the `AICORE_RESOURCE_GROUP` environment variable (defaults to "default").
:::
## Quick Start
### Step 1: Install LiteLLM
```bash
pip install litellm
```
### Step 2: Set Your Credentials
Choose **one** of these authentication methods:
<Tabs>
<TabItem value="service-key" label="Service Key JSON (Recommended)">
The simplest approach - paste your entire service key as a single environment variable. The service key must be wrapped in a `credentials` object:
```bash
export AICORE_SERVICE_KEY='{
"credentials": {
"clientid": "your-client-id",
"clientsecret": "your-client-secret",
"url": "https://<your-instance>.authentication.sap.hana.ondemand.com",
"serviceurls": {
"AI_API_URL": "https://api.ai.<your-region>.aws.ml.hana.ondemand.com"
}
}
}'
export AICORE_RESOURCE_GROUP="default"
```
</TabItem>
<TabItem value="individual" label="Individual Variables">
Alternatively, instead of using the service key above, you could set each credential separately:
```bash
export AICORE_AUTH_URL="https://<your-instance>.authentication.sap.hana.ondemand.com/oauth/token"
export AICORE_CLIENT_ID="your-client-id"
export AICORE_CLIENT_SECRET="your-client-secret"
export AICORE_RESOURCE_GROUP="default"
export AICORE_BASE_URL="https://api.ai.<your-region>.aws.ml.hana.ondemand.com/v2"
```
</TabItem>
</Tabs>
### Step 3: Make Your First Request
```python title="test_sap.py"
from litellm import completion
response = completion(
model="sap/gpt-4o",
messages=[{"role": "user", "content": "Hello from LiteLLM!"}]
)
print(response.choices[0].message.content)
```
Run it:
```bash
python test_sap.py
```
**Expected output:**
```text
Hello! How can I assist you today?
```
### Step 4: Verify Your Setup (Optional)
Test that everything is working with this diagnostic script:
```python title="verify_sap_setup.py"
import os
import litellm
# Enable debug logging to see what's happening
import os
os.environ["LITELLM_LOG"] = "DEBUG"
# Either use AICORE_SERVICE_KEY (contains all credentials including resourcegroup)
# OR use individual variables (all required together)
individual_vars = ["AICORE_AUTH_URL", "AICORE_CLIENT_ID", "AICORE_CLIENT_SECRET", "AICORE_BASE_URL", "AICORE_RESOURCE_GROUP"]
print("=== SAP Gen AI Hub Setup Verification ===\n")
# Check for service key method
if os.environ.get("AICORE_SERVICE_KEY"):
print("✓ Using AICORE_SERVICE_KEY authentication (includes resource group)")
else:
# Check individual variables
missing = [v for v in individual_vars if not os.environ.get(v)]
if missing:
print(f"✗ Missing environment variables: {missing}")
else:
print("✓ Using individual variable authentication")
print(f"✓ Resource group: {os.environ.get('AICORE_RESOURCE_GROUP')}")
# Test API connection
print("\n=== Testing API Connection ===\n")
try:
response = litellm.completion(
model="sap/gpt-4o",
messages=[{"role": "user", "content": "Say 'Connection successful!' and nothing else."}],
max_tokens=20
)
print(f"✓ API Response: {response.choices[0].message.content}")
print("\n🎉 Setup complete! You're ready to use SAP Gen AI Hub with LiteLLM.")
except Exception as e:
print(f"✗ API Error: {e}")
print("\nTroubleshooting tips:")
print(" 1. Verify your service key credentials are correct")
print(" 2. Check that 'gpt-4o' is deployed in your resource group")
print(" 3. Ensure your SAP AI Core instance is running")
```
Run the verification:
```bash
python verify_sap_setup.py
```
**Expected output on success:**
```text
=== SAP Gen AI Hub Setup Verification ===
✓ Using AICORE_SERVICE_KEY authentication
✓ Resource group: default
=== Testing API Connection ===
✓ API Response: Connection successful!
🎉 Setup complete! You're ready to use SAP Gen AI Hub with LiteLLM.
```
## Authentication
SAP Generative AI Hub uses service key authentication. You can provide credentials via:
SAP Generative AI Hub uses OAuth2 service keys for authentication. See [Quick Start](#quick-start) for setup instructions.
1. **Environment variable** - Set `AICORE_SERVICE_KEY` with your service key JSON
2. **Direct parameter** - Pass `api_key` with the service key JSON string
### Environment Variables Reference
```python showLineNumbers title="Environment Variable"
import os
os.environ["AICORE_SERVICE_KEY"] = '{"clientid": "...", "clientsecret": "...", ...}'
| Variable | Required | Description |
|----------|----------|-------------|
| `AICORE_SERVICE_KEY` | Yes* | Complete service key JSON (recommended method) |
| `AICORE_RESOURCE_GROUP` | Yes | Your AI Core resource group name |
| `AICORE_AUTH_URL` | Yes* | OAuth token URL (alternative to service key) |
| `AICORE_CLIENT_ID` | Yes* | OAuth client ID (alternative to service key) |
| `AICORE_CLIENT_SECRET` | Yes* | OAuth client secret (alternative to service key) |
| `AICORE_BASE_URL` | Yes* | AI Core API base URL (alternative to service key) |
*Choose either `AICORE_SERVICE_KEY` OR the individual variables (`AICORE_AUTH_URL`, `AICORE_CLIENT_ID`, `AICORE_CLIENT_SECRET`, `AICORE_BASE_URL`).
## Model Naming Conventions
Understanding model naming is crucial for using SAP Gen AI Hub correctly. The naming pattern differs depending on whether you're using the SDK directly or through the proxy.
### Direct SDK Usage
When calling LiteLLM's SDK directly, you **must** include the `sap/` prefix in the model name:
```python
# Correct - includes sap/ prefix
model="sap/gpt-4o"
model="sap/anthropic--claude-4.5-sonnet"
model="sap/gemini-2.5-pro"
# Incorrect - missing prefix
model="gpt-4o" # ❌ Won't work
```
3. **Environment variables** - Set the following list of credentials in .env file
<pre>
AICORE_AUTH_URL = "https://* * * .authentication.sap.hana.ondemand.com/oauth/token",
AICORE_CLIENT_ID = " *** ",
AICORE_CLIENT_SECRET = " *** ",
AICORE_RESOURCE_GROUP = " *** ",
AICORE_BASE_URL = "https://api.ai.***.cfapps.sap.hana.ondemand.com/v2"
</pre>
## Usage - LiteLLM Python SDK
```python showLineNumbers title="SAP Chat Completion"
from litellm import completion
import os
### Proxy Usage
os.environ["AICORE_SERVICE_KEY"] = '{"clientid": "...", "clientsecret": "...", ...}'
When using the LiteLLM Proxy, you use the **friendly `model_name`** defined in your configuration. The proxy automatically handles the `sap/` prefix routing.
response = completion(
model="sap/gpt-4",
messages=[{"role": "user", "content": "Hello from LiteLLM"}]
```yaml
# In config.yaml, define the mapping
model_list:
- model_name: gpt-4o # ← Use this name in client requests
litellm_params:
model: sap/gpt-4o # ← Proxy handles the sap/ prefix
```
```python
# Client request - no sap/ prefix needed
client.chat.completions.create(
model="gpt-4o", # ✓ Correct for proxy usage
messages=[...]
)
print(response)
```
```python showLineNumbers title="SAP Chat Completion - Streaming"
### Anthropic Models Special Syntax
Anthropic models use a double-dash (`--`) prefix convention:
| Provider | Model Example | LiteLLM Format |
|----------|---------------|----------------|
| OpenAI | GPT-4o | `sap/gpt-4o` |
| Anthropic | Claude 4.5 Sonnet | `sap/anthropic--claude-4.5-sonnet` |
| Google | Gemini 2.5 Pro | `sap/gemini-2.5-pro` |
| Mistral | Mistral Large | `sap/mistral-large` |
### Quick Reference Table
| Usage Type | Model Format | Example |
|------------|--------------|---------|
| Direct SDK | `sap/<model-name>` | `sap/gpt-4o` |
| Direct SDK (Anthropic) | `sap/anthropic--<model>` | `sap/anthropic--claude-4.5-sonnet` |
| Proxy Client | `<friendly-name>` | `gpt-4o` or `claude-sonnet` |
## Using the Python SDK
The LiteLLM Python SDK automatically detects your authentication method. Simply set your environment variables and make requests.
```python showLineNumbers title="Basic Completion"
from litellm import completion
import os
os.environ["AICORE_SERVICE_KEY"] = '{"clientid": "...", "clientsecret": "...", ...}'
# Assumes AICORE_AUTH_URL, AICORE_CLIENT_ID, etc. are set
response = completion(
model="sap/gpt-4",
messages=[{"role": "user", "content": "Hello from LiteLLM"}],
stream=True
model="sap/anthropic--claude-4.5-sonnet",
messages=[{"role": "user", "content": "Explain quantum computing"}]
)
for chunk in response:
print(chunk.choices[0].delta.content or "", end="")
print(response.choices[0].message.content)
```
```python showLineNumbers title="SAP Embedding"
from litellm import embedding
import os
Both authentication methods (individual variables or service key JSON) work automatically - no code changes required.
os.environ["AICORE_SERVICE_KEY"] = '{"clientid": "...", "clientsecret": "...", ...}'
## Using the Proxy Server
result = embedding(
model="sap/text-embedding-3-small",
input="Answer to the ultimate question of life, the universe, and everything is 42")
print(result.data[0])
```
The LiteLLM Proxy provides a unified OpenAI-compatible API for your SAP models.
## Usage - LiteLLM Proxy
### Configuration
Add to your LiteLLM Proxy config:
Create a `config.yaml` file in your project directory with your model mappings and credentials:
```yaml showLineNumbers title="config.yaml"
model_list:
- model_name: "sap/*"
# OpenAI models
- model_name: gpt-5
litellm_params:
model: "sap/*"
model: sap/gpt-5
general_settings:
master_key: your-proxy-api-key
# Anthropic models (note the double-dash)
- model_name: claude-sonnet
litellm_params:
model: sap/anthropic--claude-4.5-sonnet
- model_name: claude-opus
litellm_params:
model: sap/anthropic--claude-4.5-opus
# Embeddings
- model_name: text-embedding-3-small
litellm_params:
model: sap/text-embedding-3-small
litellm_settings:
drop_params: true
set_verbose: false
request_timeout: 600
num_retries: 2
forward_client_headers_to_llm_api: ["anthropic-version"]
general_settings:
master_key: "sk-1234" # Enter here your desired master key starting with 'sk-'.
# UI Admin is not required but helpful including the management of keys for your team(s). If you are using a database, these parameters are required:
database_url: "Enter you database URL."
UI_USERNAME: "Your desired UI admin account name"
UI_PASSWORD: "Your desired and strong pwd"
# Authentication
environment_variables:
AICORE_SERVICE_KEY: '{"clientid": "...", "clientsecret": "...", ...}'
AICORE_SERVICE_KEY: '{"credentials": {"clientid": "...", "clientsecret": "...", "url": "...", "serviceurls": {"AI_API_URL": "..."}}}'
AICORE_RESOURCE_GROUP: "default"
```
Start the proxy:
### Starting the Proxy
```bash showLineNumbers title="Start Proxy"
litellm --config config.yaml
```
The proxy will start on `http://localhost:4000` by default.
### Making Requests
<Tabs>
<TabItem value="curl" label="cURL">
```bash showLineNumbers title="Test Request"
curl http://localhost:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer your-proxy-api-key" \
-H "Authorization: Bearer sk-1234" \
-d '{
"model": "sap/gpt-4",
"model": "gpt-4o",
"messages": [{"role": "user", "content": "Hello"}]
}'
```
@ -118,11 +358,11 @@ from openai import OpenAI
client = OpenAI(
base_url="http://localhost:4000",
api_key="your-proxy-api-key"
api_key="sk-1234"
)
response = client.chat.completions.create(
model="sap/gpt-4",
model="gpt-4o",
messages=[{"role": "user", "content": "Hello"}]
)
print(response.choices[0].message.content)
@ -134,12 +374,14 @@ print(response.choices[0].message.content)
```python showLineNumbers title="LiteLLM SDK"
import os
import litellm
os.environ["LITELLM_PROXY_API_KEY"] = "your-proxy-api-key"
litellm.use_litellm_proxy = True # it is important to set this parameter
os.environ["LITELLM_PROXY_API_KEY"] = "sk-1234"
litellm.use_litellm_proxy = True
response = litellm.completion(
model="sap/gpt-4o",
messages=[{ "content": "Hello, how are you?","role": "user"}],
api_base="http://your-proxy-api-base"
model="claude-sonnet",
messages=[{"content": "Hello, how are you?", "role": "user"}],
api_base="http://localhost:4000"
)
print(response)
@ -148,15 +390,170 @@ print(response)
</TabItem>
</Tabs>
## Supported Parameters
## Features
| Parameter | Description |
|-----------|-------------|
| `temperature` | Controls randomness |
| `max_tokens` | Maximum tokens in response |
| `top_p` | Nucleus sampling |
| `tools` | Function calling tools |
| `tool_choice` | Tool selection behavior |
| `response_format` | Output format (json_object, json_schema) |
| `stream` | Enable streaming |
### Streaming Responses
Stream responses in real-time for better user experience:
```python showLineNumbers title="Streaming Chat Completion"
from litellm import completion
response = completion(
model="sap/gpt-4o",
messages=[{"role": "user", "content": "Count from 1 to 10"}],
stream=True
)
for chunk in response:
if chunk.choices[0].delta.content:
print(chunk.choices[0].delta.content, end="", flush=True)
```
### Structured Output
#### JSON Schema (Recommended)
Use JSON Schema for structured output with strict validation:
```python showLineNumbers title="JSON Schema Response"
from litellm import completion
response = completion(
model="sap/gpt-4o",
messages=[{
"role": "user",
"content": "Generate info about Tokyo"
}],
response_format={
"type": "json_schema",
"json_schema": {
"name": "city_info",
"schema": {
"type": "object",
"properties": {
"name": {"type": "string"},
"population": {"type": "number"},
"country": {"type": "string"}
},
"required": ["name", "population", "country"],
"additionalProperties": False
},
"strict": True
}
}
)
print(response.choices[0].message.content)
# Output: {"name":"Tokyo","population":37000000,"country":"Japan"}
```
#### JSON Object Format
For flexible JSON output without schema validation:
```python showLineNumbers title="JSON Object Response"
from litellm import completion
response = completion(
model="sap/gpt-4o",
messages=[{
"role": "user",
"content": "Generate a person object in JSON format with name and age"
}],
response_format={"type": "json_object"}
)
print(response.choices[0].message.content)
```
:::note SAP Platform Requirement
When using `json_object` type, SAP's orchestration service requires the word "json" to appear in your prompt. This ensures explicit intent for JSON formatting. For schema-validated output without this requirement, use `json_schema` instead (recommended).
:::
### Multi-turn Conversations
Maintain conversation context across multiple turns:
```python showLineNumbers title="Multi-turn Conversation"
from litellm import completion
response = completion(
model="sap/gpt-4o",
messages=[
{"role": "user", "content": "My name is Alice"},
{"role": "assistant", "content": "Hello Alice! Nice to meet you."},
{"role": "user", "content": "What is my name?"}
]
)
print(response.choices[0].message.content)
# Output: Your name is Alice.
```
### Embeddings
Generate vector embeddings for semantic search and retrieval:
```python showLineNumbers title="Create Embeddings"
from litellm import embedding
response = embedding(
model="sap/text-embedding-3-small",
input=["Hello world", "Machine learning is fascinating"]
)
print(response.data[0]["embedding"]) # Vector representation
```
## Reference
### Supported Parameters
| Parameter | Type | Description |
|-----------|------|-------------|
| `model` | string | Model identifier (with `sap/` prefix for SDK) |
| `messages` | array | Conversation messages |
| `temperature` | float | Controls randomness (0-2) |
| `max_tokens` | integer | Maximum tokens in response |
| `top_p` | float | Nucleus sampling threshold |
| `stream` | boolean | Enable streaming responses |
| `response_format` | object | Output format (`json_object`, `json_schema`) |
| `tools` | array | Function calling tool definitions |
| `tool_choice` | string/object | Tool selection behavior |
### Supported Models
For the complete and up-to-date list of available models provided by SAP Gen AI Hub, please refer to the [SAP AI Core Generative AI Hub documentation](https://help.sap.com/docs/sap-ai-core/sap-ai-core-service-guide/models-and-scenarios-in-generative-ai-hub).
:::info Model Availability
Model availability varies by SAP deployment region and your subscription. Contact your SAP administrator to confirm which models are available in your environment.
:::
### Troubleshooting
**Authentication Errors**
If you receive authentication errors:
1. Verify all required environment variables are set correctly
2. Check that your service key hasn't expired
3. Confirm your resource group has access to the desired models
4. Ensure the `AICORE_AUTH_URL` and `AICORE_BASE_URL` match your SAP region
**Model Not Found**
If a model returns "not found":
1. Verify the model is available in your SAP deployment
2. Check you're using the correct model name format (`sap/` prefix for SDK)
3. Confirm your resource group has access to that specific model
4. For Anthropic models, ensure you're using the `anthropic--` double-dash prefix
**Rate Limiting**
SAP Gen AI Hub enforces rate limits based on your subscription. If you hit limits:
1. Implement exponential backoff retry logic
2. Consider using the proxy's built-in rate limiting features
3. Contact your SAP administrator to review quota allocations

View file

@ -416,7 +416,6 @@ response = image_edit(
image=open("original_image.png", "rb"),
mask=open("mask_image.png", "rb"),
prompt="Add flowers in the masked area",
size="1024x1024",
)
print(response)
```

View file

@ -9,7 +9,6 @@ LiteLLM provides flexible cost tracking and pricing customization for all LLM pr
- **Custom Pricing** - Override default model costs or set pricing for custom models
- **Cost Per Token** - Track costs based on input/output tokens (most common)
- **Cost Per Second** - Track costs based on runtime (e.g., Sagemaker)
- **Zero-Cost Models** - Bypass budget checks for free/on-premises models by setting costs to 0
- **[Provider Discounts](./provider_discounts.md)** - Apply percentage-based discounts to specific providers
- **[Provider Margins](./provider_margins.md)** - Add fees/margins to LLM costs for internal billing
- **Base Model Mapping** - Ensure accurate cost tracking for Azure deployments
@ -107,51 +106,6 @@ There are other keys you can use to specify costs for different scenarios and mo
These keys evolve based on how new models handle multimodality. The latest version can be found at [https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json).
## Zero-Cost Models (Bypass Budget Checks)
**Use Case**: You have on-premises or free models that should be accessible even when users exceed their budget limits.
**Solution** ✅: Set both `input_cost_per_token` and `output_cost_per_token` to `0` (explicitly) to bypass all budget checks for that model.
:::info
When a model is configured with zero cost, LiteLLM will automatically skip ALL budget checks (user, team, team member, end-user, organization, and global proxy budget) for requests to that model.
**Important**: Both costs must be **explicitly set to 0**. If costs are `null` or undefined, the model will be treated as having cost and budget checks will apply.
:::
### Configuration Example
```yaml
model_list:
# On-premises model - free to use
- model_name: on-prem-llama
litellm_params:
model: ollama/llama3
api_base: http://localhost:11434
model_info:
input_cost_per_token: 0 # 👈 Explicitly set to 0
output_cost_per_token: 0 # 👈 Explicitly set to 0
# Paid cloud model - budget checks apply
- model_name: gpt-4
litellm_params:
model: gpt-4
api_key: os.environ/OPENAI_API_KEY
# No model_info - uses default pricing from cost map
```
### Behavior
With the above configuration:
- **User over budget** → Can still use `on-prem-llama` ✅, but blocked from `gpt-4`
- **Team over budget** → Can still use `on-prem-llama` ✅, but blocked from `gpt-4`
- **End-user over budget** → Can still use `on-prem-llama` ✅, but blocked from `gpt-4`
This ensures your free/on-premises models remain accessible regardless of budget constraints, while paid models are still properly governed.
## Set 'base_model' for Cost Tracking (e.g. Azure deployments)
**Problem**: Azure returns `gpt-4` in the response when `azure/gpt-4-1106-preview` is used. This leads to inaccurate cost tracking

View file

@ -22,19 +22,22 @@ Customer Usage enables you to track spend and usage for individual customers (en
## How to Track Spend
Track customer spend by including a `user` field in your API requests. The customer ID will be automatically tracked and associated with all spend from that request.
Track customer spend by including a `user` field in your API requests or by passing a customer ID header. The customer ID will be automatically tracked and associated with all spend from that request.
### Example using cURL
<Tabs>
<TabItem value="body" label="Request Body" default>
### Using Request Body
Make a `/chat/completions` call with the `user` field containing your customer ID:
```bash showLineNumbers title="Track spend with customer ID"
```bash showLineNumbers title="Track spend 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' \ # 👈 YOUR PROXY KEY
--header 'Authorization: Bearer sk-1234' \
--data '{
"model": "gpt-3.5-turbo",
"user": "customer-123", # 👈 CUSTOMER ID
"user": "customer-123",
"messages": [
{
"role": "user",
@ -44,7 +47,49 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
}'
```
The customer ID (`customer-123`) will be automatically upserted into the database with the new spend. If the customer ID already exists, spend will be incremented.
</TabItem>
<TabItem value="header" label="Request Header">
### Using Request Headers
You can also pass the customer ID via HTTP headers. This is useful for tools that support custom headers but don't allow modifying the request body (like Claude Code with `ANTHROPIC_CUSTOM_HEADERS`).
LiteLLM automatically recognizes these standard headers (no configuration required):
- `x-litellm-customer-id`
- `x-litellm-end-user-id`
```bash showLineNumbers title="Track spend 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' \
--header 'x-litellm-customer-id: customer-123' \
--data '{
"model": "gpt-3.5-turbo",
"messages": [
{
"role": "user",
"content": "What is the capital of France?"
}
]
}'
```
#### Using with Claude Code
Claude Code supports custom headers via the `ANTHROPIC_CUSTOM_HEADERS` environment variable. Set it to pass your customer ID:
```bash title="Configure Claude Code with customer tracking"
export ANTHROPIC_BASE_URL="http://0.0.0.0:4000/v1/messages"
export ANTHROPIC_API_KEY="sk-1234"
export ANTHROPIC_CUSTOM_HEADERS="x-litellm-customer-id: my-customer-id"
```
Now all requests from Claude Code will automatically track spend under `my-customer-id`.
</TabItem>
</Tabs>
The customer ID will be automatically upserted into the database with the new spend. If the customer ID already exists, spend will be incremented.
### Example using OpenWebUI

View file

@ -0,0 +1,273 @@
# [New] Fallback Management Endpoints
Dedicated endpoints for managing model fallbacks separately from the general configuration.
## Overview
These endpoints allow you to configure, retrieve, and delete fallback models without modifying the entire proxy configuration. This provides a cleaner and safer way to manage fallbacks compared to using the `/config/update` endpoint.
## Prerequisites
- Database storage must be enabled: Set `STORE_MODEL_IN_DB=True` in your environment
- Models must exist in the router before configuring fallbacks
## Endpoints
### POST /fallback
Create or update fallbacks for a specific model.
**Request Body:**
```json
{
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4", "claude-3-haiku"],
"fallback_type": "general"
}
```
**Parameters:**
- `model` (string, required): The primary model name to configure fallbacks for
- `fallback_models` (array of strings, required): List of fallback model names in priority order
- `fallback_type` (string, optional): Type of fallback. Options:
- `"general"` (default): Standard fallbacks for any error
- `"context_window"`: Fallbacks for context window exceeded errors
- `"content_policy"`: Fallbacks for content policy violations
**Response:**
```json
{
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4", "claude-3-haiku"],
"fallback_type": "general",
"message": "Fallback configuration created successfully"
}
```
**Example using cURL:**
```bash
curl -X POST "http://localhost:4000/fallback" \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4", "claude-3-haiku"],
"fallback_type": "general"
}'
```
**Example using Python:**
```python
import requests
response = requests.post(
"http://localhost:4000/fallback",
headers={
"Authorization": "Bearer sk-1234",
"Content-Type": "application/json"
},
json={
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4", "claude-3-haiku"],
"fallback_type": "general"
}
)
print(response.json())
```
### GET /fallback/{model}
Get fallback configuration for a specific model.
**Parameters:**
- `model` (path parameter, required): The model name to get fallbacks for
- `fallback_type` (query parameter, optional): Type of fallback to retrieve (default: "general")
**Response:**
```json
{
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4", "claude-3-haiku"],
"fallback_type": "general"
}
```
**Example using cURL:**
```bash
curl -X GET "http://localhost:4000/fallback/gpt-3.5-turbo?fallback_type=general" \
-H "Authorization: Bearer sk-1234"
```
**Example using Python:**
```python
import requests
response = requests.get(
"http://localhost:4000/fallback/gpt-3.5-turbo",
headers={"Authorization": "Bearer sk-1234"},
params={"fallback_type": "general"}
)
print(response.json())
```
### DELETE /fallback/{model}
Delete fallback configuration for a specific model.
**Parameters:**
- `model` (path parameter, required): The model name to delete fallbacks for
- `fallback_type` (query parameter, optional): Type of fallback to delete (default: "general")
**Response:**
```json
{
"model": "gpt-3.5-turbo",
"fallback_type": "general",
"message": "Fallback configuration deleted successfully"
}
```
**Example using cURL:**
```bash
curl -X DELETE "http://localhost:4000/fallback/gpt-3.5-turbo?fallback_type=general" \
-H "Authorization: Bearer sk-1234"
```
**Example using Python:**
```python
import requests
response = requests.delete(
"http://localhost:4000/fallback/gpt-3.5-turbo",
headers={"Authorization": "Bearer sk-1234"},
params={"fallback_type": "general"}
)
print(response.json())
```
### Test fallback
</TabItem>
<TabItem value="proxy" label="PROXY">
```bash
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-d '{
"model": "gpt-3.5-turbo",
"messages": [
{
"role": "user",
"content": "ping"
}
],
"mock_testing_fallbacks": true
}
'
```
</TabItem>
</Tabs>
## Validation
The endpoints perform the following validations:
1. **Model Existence**: Verifies that the primary model exists in the router
2. **Fallback Model Existence**: Ensures all fallback models exist in the router
3. **No Self-Fallback**: Prevents a model from being its own fallback
4. **No Duplicates**: Ensures no duplicate models in the fallback list
5. **Database Enabled**: Requires `STORE_MODEL_IN_DB=True` to be set
## Error Responses
### 400 Bad Request
```json
{
"detail": {
"error": "Invalid fallback models: ['non-existent-model']",
"available_models": ["gpt-3.5-turbo", "gpt-4", "claude-3-haiku"]
}
}
```
### 404 Not Found
```json
{
"detail": {
"error": "Model 'gpt-3.5-turbo' not found in router",
"available_models": ["gpt-4", "claude-3-haiku"]
}
}
```
### 500 Internal Server Error
```json
{
"detail": {
"error": "Router not initialized"
}
}
```
## Fallback Types Explained
### General Fallbacks
Used for any type of error that occurs during model invocation. This is the most common type of fallback.
**Use Case:** When a model is unavailable, rate-limited, or returns an error.
```json
{
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4", "claude-3-haiku"],
"fallback_type": "general"
}
```
### Context Window Fallbacks
Specifically triggered when a context window exceeded error occurs.
**Use Case:** When the input is too long for the primary model, fallback to a model with a larger context window.
```json
{
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4-32k", "claude-3-opus"],
"fallback_type": "context_window"
}
```
### Content Policy Fallbacks
Specifically triggered when content policy violations occur.
**Use Case:** When the primary model rejects content due to safety filters, fallback to a model with different content policies.
```json
{
"model": "gpt-4",
"fallback_models": ["claude-3-haiku"],
"fallback_type": "content_policy"
}
```
## Benefits Over /config/update
1. **Safety**: Only modifies fallback configuration, won't accidentally change other settings
2. **Simplicity**: Focused API with clear validation messages
3. **Granularity**: Manage fallbacks per model and per type
4. **Validation**: Comprehensive checks ensure configuration is valid before applying
5. **Clarity**: Clear error messages with available models listed
## Notes
- Fallbacks are triggered after the configured number of retries fails
- Fallbacks are attempted in the order specified in `fallback_models`
- The maximum number of fallbacks attempted is controlled by the router's `max_fallbacks` setting
- Changes take effect immediately and are persisted to the database

View file

@ -30,6 +30,9 @@ general_settings:
# Optional: set how frequently cleanup should run - default is daily
maximum_spend_logs_retention_interval: "1d" # Run cleanup daily
# Optional: set exact time for cleanup (Cron syntax)
maximum_spend_logs_cleanup_cron: "0 4 * * *" # Run at 04:00 AM daily
litellm_settings:
cache: true
cache_params:
@ -51,6 +54,15 @@ How long logs should be kept before deletion. Supported formats:
How often the cleanup job should run. Uses the same format as above. If not set, cleanup will run every 24 hours if and only if `maximum_spend_logs_retention_period` is set.
#### `maximum_spend_logs_cleanup_cron` (optional)
Schedule the cleanup using standard cron syntax. This takes precedence over `maximum_spend_logs_retention_interval`.
Examples:
- `"0 4 * * *"` Run at 04:00 AM daily
- `"0 0 * * 0"` Run at midnight every Sunday
- `"*/30 * * * *"` Run every 30 minutes
## How it works
### Step 1. Lock Acquisition (Optional with Redis)

View file

@ -0,0 +1,99 @@
# Claude Code - Granular Cost Tracking
Track Claude Code usage by customer or tags using LiteLLM proxy. This enables granular cost attribution for billing, budgeting, and analytics.
## How It Works
Claude Code supports custom headers via `ANTHROPIC_CUSTOM_HEADERS`. LiteLLM automatically tracks requests with specific headers for cost attribution.
## Tracking Options
Choose how you want to attribute costs:
| Track By | Header | Use Case |
|----------|--------|----------|
| Customer | `x-litellm-customer-id` | Bill customers, per-user budgets |
| Tags | `x-litellm-tags` | Project tracking, cost centers, environments |
## Environment Variables
| Variable | Description | Example |
|----------|-------------|---------|
| `ANTHROPIC_BASE_URL` | LiteLLM proxy URL | `http://localhost:4000` |
| `ANTHROPIC_API_KEY` | LiteLLM API key | `sk-1234` |
| `ANTHROPIC_CUSTOM_HEADERS` | Custom headers (`header-name: value` format) | See examples below |
## Option 1: Track by Customer
Use this to attribute costs to specific customers or end-users.
```bash
export ANTHROPIC_BASE_URL=http://localhost:4000
export ANTHROPIC_API_KEY=sk-1234
export ANTHROPIC_CUSTOM_HEADERS="x-litellm-customer-id: claude-ishaan-local"
```
## Option 2: Track by Tags
Use this to attribute costs to projects, cost centers, or environments. Pass comma-separated tags.
```bash
export ANTHROPIC_BASE_URL=http://localhost:4000
export ANTHROPIC_API_KEY=sk-1234
export ANTHROPIC_CUSTOM_HEADERS="x-litellm-tags: project:acme,env:prod,team:backend"
```
## Quick Start
### 1. Set Environment Variables
```bash
export ANTHROPIC_BASE_URL=http://localhost:4000
export ANTHROPIC_API_KEY=sk-1234
export ANTHROPIC_CUSTOM_HEADERS="x-litellm-customer-id: claude-ishaan-local"
```
### 2. Use Claude Code
```bash
claude
```
All requests will now be tracked under the customer ID `claude-ishaan-local`.
![](https://colony-recorder.s3.amazonaws.com/files/2026-01-16/8f45872e-2d00-4d01-bf3d-4d6ae11d1396/ascreenshot_d2a745b8da4f4a56aaf2cac02871ef53_text_export.jpeg)
![](https://colony-recorder.s3.amazonaws.com/files/2026-01-16/dd41eae3-2592-4bc9-a8d2-d6d02614cd2d/ascreenshot_43ec9ee48ad946cca49732f007e786fc_text_export.jpeg)
![](https://colony-recorder.s3.amazonaws.com/files/2026-01-16/0c30309e-7117-4999-a3df-d22a2d5629c1/ascreenshot_d76a48c53b9a4fad8f6727baf4aa6a9c_text_export.jpeg)
### 3. View Usage in LiteLLM UI
Navigate to the **Logs** tab in the LiteLLM UI.
![](https://colony-recorder.s3.amazonaws.com/files/2026-01-16/ff774392-69f5-483e-83e2-fb749c94ee90/ascreenshot_d264fc04c9ee47edb047f61b6eb8c4d7_text_export.jpeg)
Click on a request to see details.
![](https://colony-recorder.s3.amazonaws.com/files/2026-01-16/5f71589b-5fdd-4759-9b6e-e6874be0eb21/ascreenshot_92dd86dadccb4764b1169c29c10dfe65_text_export.jpeg)
Filter by customer ID to see all requests for that customer.
![](https://colony-recorder.s3.amazonaws.com/files/2026-01-16/dd1c8aba-e75b-4714-9eee-c785e9db99af/ascreenshot_36aaec0fe12f4189b64f704a551e6729_text_export.jpeg)
## Supported Headers
| Header | Description |
|--------|-------------|
| `x-litellm-customer-id` | Track by customer/end-user ID |
| `x-litellm-end-user-id` | Alternative customer ID header |
| `x-litellm-tags` | Comma-separated tags for cost attribution |
## Related
- [Claude Code Quickstart](./claude_responses_api.md)
- [Customer Budgets](../proxy/customers.md)
- [Tag Budgets](../proxy/tag_budgets.md)
- [Track Usage for Coding Tools](./cost_tracking_coding.md)

View file

@ -121,6 +121,7 @@ const sidebars = {
label: "Claude Code",
items: [
"tutorials/claude_responses_api",
"tutorials/claude_code_customer_tracking",
"tutorials/claude_mcp",
"tutorials/claude_non_anthropic_models",
]
@ -821,6 +822,7 @@ const sidebars = {
"completion/knowledgebase",
"guides/code_interpreter",
"completion/message_trimming",
"completion/message_sanitization",
"completion/model_alias",
"completion/mock_requests",
"completion/predict_outputs",
@ -857,6 +859,7 @@ const sidebars = {
"proxy/load_balancing",
"proxy/provider_budget_routing",
"proxy/reliability",
"proxy/fallback_management",
"proxy/tag_routing",
"proxy/timeout",
"wildcard_routing"

View file

@ -9,7 +9,7 @@ warnings.filterwarnings("ignore", message=".*conflict with protected namespace.*
warnings.filterwarnings(
"ignore", message=".*Accessing the.*attribute on the instance is deprecated.*"
)
### INIT VARIABLES ########################
### INIT VARIABLES #########################
import threading
import os
from typing import (

View file

@ -1073,6 +1073,13 @@ LITELLM_TRUNCATED_PAYLOAD_FIELD = "litellm_truncated"
########################### LiteLLM Proxy Specific Constants ###########################
########################################################################################
# Standard headers that are always checked for customer/end-user ID (no configuration required)
# These headers work out-of-the-box for tools like Claude Code that support custom headers
STANDARD_CUSTOMER_ID_HEADERS = [
"x-litellm-customer-id",
"x-litellm-end-user-id",
]
MAX_SPENDLOG_ROWS_TO_QUERY = int(
os.getenv("MAX_SPENDLOG_ROWS_TO_QUERY", 1_000_000)
) # if spendLogs has more than 1M rows, do not query the DB

View file

@ -952,7 +952,8 @@ def completion_cost( # noqa: PLR0915
)
potential_model_names = [selected_model, _get_response_model(completion_response)]
if model is not None:
potential_model_names.append(model)
for idx, model in enumerate(potential_model_names):
try:

View file

@ -52,7 +52,7 @@ def _get_cached_end_user_id_for_cost_tracking():
class PrometheusLogger(CustomLogger):
# Class variables or attributes
def __init__(
def __init__( # noqa: PLR0915
self,
**kwargs,
):

View file

@ -4338,6 +4338,38 @@ class StandardLoggingPayloadSetup:
return messages
@staticmethod
def merge_litellm_metadata(litellm_params: dict) -> dict:
"""
Merge both litellm_metadata and metadata from litellm_params.
litellm_metadata contains model-related fields, metadata contains user API key fields.
We need both for complete standard logging payload.
Args:
litellm_params: Dictionary containing metadata and litellm_metadata
Returns:
dict: Merged metadata with user API key fields taking precedence
"""
merged_metadata: dict = {}
# Start with metadata (user API key fields) - but skip non-serializable objects
if litellm_params.get("metadata") and isinstance(litellm_params.get("metadata"), dict):
for key, value in litellm_params["metadata"].items():
# Skip non-serializable objects like UserAPIKeyAuth
if key == "user_api_key_auth":
continue
merged_metadata[key] = value
# Then merge litellm_metadata (model-related fields) - this will NOT overwrite existing keys
if litellm_params.get("litellm_metadata") and isinstance(litellm_params.get("litellm_metadata"), dict):
for key, value in litellm_params["litellm_metadata"].items():
if key not in merged_metadata: # Don't overwrite existing keys from metadata
merged_metadata[key] = value
return merged_metadata
@staticmethod
def get_standard_logging_metadata(
metadata: Optional[Dict[str, Any]],
@ -4456,7 +4488,7 @@ class StandardLoggingPayloadSetup:
@staticmethod
def get_usage_from_response_obj(
response_obj: Optional[Union[dict, BaseModel]], combined_usage_object: Optional[Usage] = None
response_obj: Optional[dict], combined_usage_object: Optional[Usage] = None
) -> Usage:
## BASE CASE ##
if combined_usage_object is not None:
@ -4468,32 +4500,27 @@ class StandardLoggingPayloadSetup:
total_tokens=0,
)
usage = _safe_extract_usage_from_obj(response_obj)
if usage is None:
usage = response_obj.get("usage", None) or {}
if usage is None or (
not isinstance(usage, dict) and not isinstance(usage, Usage)
):
return Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=0,
)
if isinstance(usage, Usage):
elif isinstance(usage, Usage):
return usage
transformed_usage = _try_transform_response_api_usage(usage)
if transformed_usage is not None:
return transformed_usage
if isinstance(usage, dict):
created_usage = _try_create_usage_from_dict(usage)
if created_usage is not None:
return created_usage
return Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=0,
)
elif isinstance(usage, dict):
if ResponseAPILoggingUtils._is_response_api_usage(usage):
return (
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
)
return Usage(**usage)
raise ValueError(f"usage is required, got={usage} of type {type(usage)}")
@staticmethod
def get_model_cost_information(
@ -4534,18 +4561,13 @@ class StandardLoggingPayloadSetup:
@staticmethod
def get_final_response_obj(
response_obj: Union[dict, BaseModel], init_response_obj: Union[Any, BaseModel, dict], kwargs: dict
response_obj: dict, init_response_obj: Union[Any, BaseModel, dict], kwargs: dict
) -> Optional[Union[dict, str, list]]:
"""
Get final response object after redacting the message input/output from logging
"""
if response_obj:
if isinstance(response_obj, BaseModel):
final_response_obj: Optional[Union[dict, str, list]] = _safe_model_dump(
response_obj, default={}
)
else:
final_response_obj = response_obj
final_response_obj: Optional[Union[dict, str, list]] = response_obj
elif isinstance(init_response_obj, list) or isinstance(init_response_obj, str):
final_response_obj = init_response_obj
else:
@ -4559,7 +4581,7 @@ class StandardLoggingPayloadSetup:
if modified_final_response_obj is not None and isinstance(
modified_final_response_obj, BaseModel
):
final_response_obj = _safe_model_dump(modified_final_response_obj, default={})
final_response_obj = modified_final_response_obj.model_dump()
else:
final_response_obj = modified_final_response_obj
@ -4830,125 +4852,6 @@ class StandardLoggingPayloadSetup:
return request_tags
def _safe_model_dump(
obj: BaseModel, default: Optional[Union[dict, str, list]] = None
) -> Union[dict, str, list]:
"""
Safely call model_dump() on a BaseModel with fallback strategies.
Args:
obj: BaseModel instance to dump
default: Default value to return if all strategies fail
Returns:
Dict representation of the BaseModel, or fallback value
"""
if default is None:
default = {}
try:
return obj.model_dump()
except (AttributeError, TypeError) as e:
verbose_logger.debug(
f"Error calling model_dump() on BaseModel: {e}, type: {type(obj)}"
)
try:
if hasattr(obj, "__dict__"):
return obj.__dict__
else:
return str(obj)
except Exception:
return default
def _safe_get_attribute(
obj: Union[dict, BaseModel, Any], attr_name: str, default: Any = None
) -> Any:
"""
Safely get an attribute from a dict or BaseModel object.
Args:
obj: Object to get attribute from (dict, BaseModel, or any object)
attr_name: Name of the attribute to get
default: Default value to return if attribute doesn't exist
Returns:
Attribute value or default
"""
try:
if isinstance(obj, dict):
return obj.get(attr_name, default)
else:
return getattr(obj, attr_name, default)
except (AttributeError, TypeError) as e:
verbose_logger.debug(
f"Error getting attribute '{attr_name}' from object: {e}, type: {type(obj)}"
)
return default
def _safe_extract_usage_from_obj(
response_obj: Union[dict, BaseModel, Any]
) -> Optional[Union[dict, Usage, Any]]:
"""
Safely extract usage from response_obj (dict or BaseModel).
Args:
response_obj: Response object (dict, BaseModel, or any object)
Returns:
Usage object, dict, or None
"""
return _safe_get_attribute(response_obj, "usage", None)
def _try_transform_response_api_usage(usage: Any) -> Optional[Usage]:
"""
Try to transform ResponseAPIUsage to Usage object.
Args:
usage: Usage object (dict, ResponseAPIUsage, or other)
Returns:
Transformed Usage object, or None if transformation fails
"""
try:
if ResponseAPILoggingUtils._is_response_api_usage(usage):
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
except (AttributeError, TypeError, KeyError) as e:
verbose_logger.debug(
f"Error checking/transforming ResponseAPIUsage: {e}, type: {type(usage)}"
)
return None
def _try_create_usage_from_dict(usage: dict) -> Optional[Usage]:
"""
Try to create Usage object from dict.
Args:
usage: Dict containing usage information
Returns:
Usage object, or None if creation fails
"""
try:
return Usage(**usage)
except (TypeError, ValueError) as e:
# Avoid logging full dict contents, which may include sensitive data
try:
usage_keys = list(usage.keys())
except Exception:
usage_keys = None
verbose_logger.debug(
"Error creating Usage from dict: %s, usage keys: %s, usage type: %s",
e,
usage_keys,
type(usage),
)
return None
def _get_status_fields(
status: StandardLoggingPayloadStatus,
guardrail_information: Optional[List[dict]],
@ -4998,21 +4901,17 @@ def _get_status_fields(
def _extract_response_obj_and_hidden_params(
init_response_obj: Union[Any, BaseModel, dict],
original_exception: Optional[Exception],
) -> Tuple[Union[dict, BaseModel], Optional[dict]]:
) -> Tuple[dict, Optional[dict]]:
"""Extract response_obj and hidden_params from init_response_obj."""
hidden_params: Optional[dict] = None
if init_response_obj is None:
response_obj: Union[dict, BaseModel] = {}
response_obj = {}
elif isinstance(init_response_obj, BaseModel):
response_obj = init_response_obj
hidden_params = _safe_get_attribute(init_response_obj, "_hidden_params", None)
response_obj = init_response_obj.model_dump()
hidden_params = getattr(init_response_obj, "_hidden_params", None)
elif isinstance(init_response_obj, dict):
response_obj = init_response_obj
else:
verbose_logger.debug(
f"Unknown init_response_obj type: {type(init_response_obj)}, defaulting to empty dict"
)
response_obj = {}
if original_exception is not None and hidden_params is None:
@ -5059,11 +4958,8 @@ def get_standard_logging_object_payload(
litellm_params = kwargs.get("litellm_params", {}) or {}
proxy_server_request = litellm_params.get("proxy_server_request") or {}
metadata: dict = (
litellm_params.get("litellm_metadata")
or litellm_params.get("metadata", None)
or {}
)
# Merge both litellm_metadata and metadata to get complete metadata
metadata: dict = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
completion_start_time = kwargs.get("completion_start_time", end_time)
call_type = kwargs.get("call_type")
@ -5075,10 +4971,7 @@ def get_standard_logging_object_payload(
),
)
# Preserve falsy values (0, "", False) if they exist in response_obj
id = _safe_get_attribute(response_obj, "id", None)
if id is None:
id = kwargs.get("litellm_call_id")
id = response_obj.get("id", kwargs.get("litellm_call_id"))
_model_id = metadata.get("model_info", {}).get("id", "")
_model_group = metadata.get("model_group", "")

View file

@ -45,7 +45,6 @@ from .common_utils import (
infer_content_type_from_url_and_content,
is_non_content_values_set,
parse_tool_call_arguments,
unpack_defs,
)
from .image_handling import convert_url_to_base64
@ -1463,56 +1462,6 @@ def convert_to_gemini_tool_call_invoke(
)
def _clean_refs_for_gemini(obj: Any) -> None:
"""
Recursively clean $defs, $ref, and definitions from a dict for Gemini compatibility.
Gemini rejects:
- $defs sections (even after $ref has been inlined)
- Any remaining $ref (circular refs, external URLs)
This function:
1. Removes all $defs/definitions keys
2. Replaces any remaining $ref with a placeholder object
"""
if isinstance(obj, dict):
# Remove $defs and definitions at this level
obj.pop("$defs", None)
obj.pop("definitions", None)
# Check for and handle remaining $ref (circular or external)
if "$ref" in obj:
ref_value = obj.pop("$ref")
# Replace with a generic object type as placeholder
obj["type"] = "object"
obj["description"] = f"(schema reference: {ref_value})"
# Recurse into values
for value in obj.values():
_clean_refs_for_gemini(value)
elif isinstance(obj, list):
for item in obj:
_clean_refs_for_gemini(item)
def _prepare_response_for_gemini(response_data: dict) -> dict:
"""
Prepare a tool response dict for Gemini by inlining $ref and removing $defs.
Gemini rejects JSON schemas with $defs/$ref in function_response content.
This function applies unpack_defs to inline references, then cleans up
any remaining $defs sections and unresolved $refs (circular or external).
Returns a new dict (does not mutate the input).
"""
import copy
result = copy.deepcopy(response_data)
unpack_defs(result, {})
_clean_refs_for_gemini(result)
return result
def convert_to_gemini_tool_call_result(
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
last_message_with_tool_calls: Optional[dict],
@ -1621,11 +1570,6 @@ def convert_to_gemini_tool_call_result(
# Not valid JSON, wrap in content field
response_data = {"content": content_str}
# Gemini rejects JSON schemas with $defs/$ref in function_response content.
# Inline $refs and clean up for Gemini compatibility.
if isinstance(response_data, dict):
response_data = _prepare_response_for_gemini(response_data)
# We can't determine from openai message format whether it's a successful or
# error call result so default to the successful result template
_function_response = VertexFunctionResponse(
@ -2045,6 +1989,223 @@ 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 = 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(
current_message: AllMessageValues,
messages: List[AllMessageValues],
current_index: int,
) -> List[AllMessageValues]:
"""
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 list containing the assistant message followed by any dummy tool results needed
"""
result_messages: List[AllMessageValues] = []
tool_calls = current_message.get("tool_calls")
if not tool_calls or len(tool_calls) == 0:
return [current_message]
# Collect all tool_call_ids from this assistant message
expected_tool_call_ids = set()
for tool_call in 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)
found_tool_call_ids = set()
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:
found_tool_call_ids.add(tool_call_id)
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)
for tool_call_id in missing_tool_call_ids:
tool_name = "unknown_tool"
for tool_call in 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 result_messages
return [current_message]
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 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 = _add_missing_tool_results(current_message, messages, i)
# If dummy tool results were added, extend sanitized_messages and continue
if len(result_messages) > 1:
sanitized_messages.extend(result_messages)
i += 1
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,
@ -2064,6 +2225,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.
@ -3289,17 +3453,21 @@ def _convert_to_bedrock_tool_call_invoke(
id = tool["id"]
name = tool["function"].get("name", "")
arguments = tool["function"].get("arguments", "")
arguments_dict = json.loads(arguments) if arguments else {}
# Ensure arguments_dict is always a dict (Bedrock requires toolUse.input to be an object)
# When some providers return arguments: '""' (JSON-encoded empty string), json.loads returns ""
if not isinstance(arguments_dict, dict):
arguments_dict = {}
if not arguments or not arguments.strip():
arguments_dict = {}
arguments_input = {}
else:
arguments_dict = json.loads(arguments)
# Try to parse the arguments JSON
try:
arguments_input = json.loads(arguments)
except json.JSONDecodeError as e:
verbose_logger.warning(
f"Malformed JSON in tool call arguments for tool '{name}': {str(e)}. "
f"Storing as raw string to allow conversation to continue."
)
arguments_input = arguments
bedrock_tool = BedrockToolUseBlock(
input=arguments_dict, name=name, toolUseId=id
input=arguments_input, name=name, toolUseId=id
)
bedrock_content_block = BedrockContentBlock(toolUse=bedrock_tool)
_parts_list.append(bedrock_content_block)

View file

@ -132,7 +132,7 @@ class ChunkProcessor:
)
return response
def get_combined_tool_content(
def get_combined_tool_content( # noqa: PLR0915
self, tool_call_chunks: List[Dict[str, Any]]
) -> List[ChatCompletionMessageToolCall]:
tool_calls_list: List[ChatCompletionMessageToolCall] = []
@ -147,10 +147,26 @@ class ChunkProcessor:
tool_calls = delta.get("tool_calls", [])
for tool_call in tool_calls:
if not tool_call or not hasattr(tool_call, "function"):
# Handle both dict and object formats
if not tool_call:
continue
# Check if tool_call has function (either as attribute or dict key)
has_function = False
if isinstance(tool_call, dict):
has_function = "function" in tool_call and tool_call["function"] is not None
else:
has_function = hasattr(tool_call, "function") and tool_call.function is not None
if not has_function:
continue
index = getattr(tool_call, "index", 0)
# Get index (handle both dict and object)
if isinstance(tool_call, dict):
index = tool_call.get("index", 0)
else:
index = getattr(tool_call, "index", 0)
if index not in tool_call_map:
tool_call_map[index] = {
"id": None,
@ -160,30 +176,56 @@ class ChunkProcessor:
"provider_specific_fields": None,
}
if hasattr(tool_call, "id") and tool_call.id:
tool_call_map[index]["id"] = tool_call.id
if hasattr(tool_call, "type") and tool_call.type:
tool_call_map[index]["type"] = tool_call.type
if hasattr(tool_call, "function"):
if (
hasattr(tool_call.function, "name")
and tool_call.function.name
):
tool_call_map[index]["name"] = tool_call.function.name
if (
hasattr(tool_call.function, "arguments")
and tool_call.function.arguments
):
tool_call_map[index]["arguments"].append(
tool_call.function.arguments
)
# Extract id, type, and function data (handle both dict and object)
if isinstance(tool_call, dict):
if tool_call.get("id"):
tool_call_map[index]["id"] = tool_call["id"]
if tool_call.get("type"):
tool_call_map[index]["type"] = tool_call["type"]
function = tool_call.get("function", {})
if isinstance(function, dict):
if function.get("name"):
tool_call_map[index]["name"] = function["name"]
if function.get("arguments"):
tool_call_map[index]["arguments"].append(function["arguments"])
else:
# function is an object
if hasattr(function, "name") and function.name:
tool_call_map[index]["name"] = function.name
if hasattr(function, "arguments") and function.arguments:
tool_call_map[index]["arguments"].append(function.arguments)
else:
# tool_call is an object
if hasattr(tool_call, "id") and tool_call.id:
tool_call_map[index]["id"] = tool_call.id
if hasattr(tool_call, "type") and tool_call.type:
tool_call_map[index]["type"] = tool_call.type
if hasattr(tool_call, "function"):
if (
hasattr(tool_call.function, "name")
and tool_call.function.name
):
tool_call_map[index]["name"] = tool_call.function.name
if (
hasattr(tool_call.function, "arguments")
and tool_call.function.arguments
):
tool_call_map[index]["arguments"].append(
tool_call.function.arguments
)
# Preserve provider_specific_fields from streaming chunks
provider_fields = None
if hasattr(tool_call, "provider_specific_fields") and tool_call.provider_specific_fields:
provider_fields = tool_call.provider_specific_fields
elif hasattr(tool_call, "function") and hasattr(tool_call.function, "provider_specific_fields") and tool_call.function.provider_specific_fields:
provider_fields = tool_call.function.provider_specific_fields
if isinstance(tool_call, dict):
provider_fields = tool_call.get("provider_specific_fields")
if not provider_fields and isinstance(tool_call.get("function"), dict):
provider_fields = tool_call["function"].get("provider_specific_fields")
else:
if hasattr(tool_call, "provider_specific_fields") and tool_call.provider_specific_fields:
provider_fields = tool_call.provider_specific_fields
elif hasattr(tool_call, "function") and hasattr(tool_call.function, "provider_specific_fields") and tool_call.function.provider_specific_fields:
provider_fields = tool_call.function.provider_specific_fields
if provider_fields:
# Merge provider_specific_fields if multiple chunks have them
@ -222,6 +264,7 @@ class ChunkProcessor:
return tool_calls_list
def get_combined_function_call_content(
self, function_call_chunks: List[Dict[str, Any]]
) -> FunctionCall:

View file

@ -2,7 +2,8 @@ from typing import Any, AsyncIterator, Dict, List, Optional, Tuple
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj, verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import verbose_logger
from litellm.llms.base_llm.anthropic_messages.transformation import (
BaseAnthropicMessagesConfig,
)
@ -13,9 +14,10 @@ from litellm.types.llms.anthropic import (
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.llms.anthropic_tool_search import get_tool_search_beta_header
from litellm.types.router import GenericLiteLLMParams
from ...common_utils import AnthropicError
from ...common_utils import AnthropicError, AnthropicModelInfo
DEFAULT_ANTHROPIC_API_BASE = "https://api.anthropic.com"
DEFAULT_ANTHROPIC_API_VERSION = "2023-06-01"
@ -75,9 +77,9 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
if "content-type" not in headers:
headers["content-type"] = "application/json"
headers = self._update_headers_with_optional_anthropic_beta(
headers = self._update_headers_with_anthropic_beta(
headers=headers,
context_management=optional_params.get("context_management"),
optional_params=optional_params,
)
return headers, api_base
@ -153,16 +155,44 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
)
@staticmethod
def _update_headers_with_optional_anthropic_beta(
headers: dict, context_management: Optional[Dict]
def _update_headers_with_anthropic_beta(
headers: dict,
optional_params: dict,
custom_llm_provider: str = "anthropic",
) -> dict:
if context_management is None:
return headers
"""
Auto-inject anthropic-beta headers based on features used.
Handles:
- context_management: adds 'context-management-2025-06-27'
- tool_search: adds provider-specific tool search header
Args:
headers: Request headers dict
optional_params: Optional parameters including tools, context_management
custom_llm_provider: Provider name for looking up correct tool search header
"""
beta_values: set = set()
# Get existing beta headers if any
existing_beta = headers.get("anthropic-beta")
beta_value = ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
if existing_beta is None:
headers["anthropic-beta"] = beta_value
elif beta_value not in [beta.strip() for beta in existing_beta.split(",")]:
headers["anthropic-beta"] = f"{existing_beta}, {beta_value}"
if existing_beta:
beta_values.update(b.strip() for b in existing_beta.split(","))
# Check for context management
if optional_params.get("context_management") is not None:
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value)
# Check for tool search tools
tools = optional_params.get("tools")
if tools:
anthropic_model_info = AnthropicModelInfo()
if anthropic_model_info.is_tool_search_used(tools):
# Use provider-specific tool search header
tool_search_header = get_tool_search_beta_header(custom_llm_provider)
beta_values.add(tool_search_header)
if beta_values:
headers["anthropic-beta"] = ",".join(sorted(beta_values))
return headers

View file

@ -664,8 +664,29 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
**data, timeout=timeout
)
headers = dict(raw_response.headers)
response = raw_response.parse()
# Convert json.JSONDecodeError to AzureOpenAIError for two critical reasons:
#
# 1. ROUTER BEHAVIOR: The router relies on exception.status_code to determine cooldown logic:
# - JSONDecodeError has no status_code → router skips cooldown evaluation
# - AzureOpenAIError has status_code → router properly evaluates for cooldown
#
# 2. CONNECTION CLEANUP: When response.parse() throws JSONDecodeError, the response
# body may not be fully consumed, preventing httpx from properly returning the
# connection to the pool. By catching the exception and accessing raw_response.status_code,
# we trigger httpx's internal cleanup logic. Without this:
# - parse() fails → JSONDecodeError bubbles up → httpx never knows response was acknowledged → connection leak
# This completely eliminates "Unclosed connection" warnings during high load.
try:
response = raw_response.parse()
except json.JSONDecodeError as json_error:
raise AzureOpenAIError(
status_code=raw_response.status_code or 500,
message=f"Failed to parse raw Azure embedding response: {str(json_error)}"
) from json_error
stringified_response = response.model_dump()
## LOGGING
logging_obj.post_call(
input=input,

View file

@ -62,10 +62,10 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig):
if "content-type" not in headers:
headers["content-type"] = "application/json"
# Update headers with optional anthropic beta features
headers = self._update_headers_with_optional_anthropic_beta(
# Update headers with anthropic beta features (context management, tool search, etc.)
headers = self._update_headers_with_anthropic_beta(
headers=headers,
context_management=optional_params.get("context_management"),
optional_params=optional_params,
)
return headers, api_base

View file

@ -87,7 +87,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig):
def transform_image_edit_request(
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
@ -99,6 +99,9 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig):
FLUX 2 uses the same endpoint for generation and editing,
with the image passed as base64 in the JSON body.
"""
if prompt is None:
raise ValueError("FLUX 2 image edit requires a prompt.")
image_b64 = self._convert_image_to_base64(image)
# Build request body with required params

View file

@ -92,7 +92,7 @@ class BaseImageEditConfig(ABC):
def transform_image_edit_request(
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,

View file

@ -1395,9 +1395,16 @@ class AmazonConverseConfig(BaseConfig):
response_tool_name = get_bedrock_tool_name(
response_tool_name=_response_tool_name
)
tool_input = content["toolUse"]["input"]
if isinstance(tool_input, str):
arguments_str = tool_input
else:
# Otherwise, serialize it to JSON
arguments_str = json.dumps(tool_input)
_function_chunk = ChatCompletionToolCallFunctionChunk(
name=response_tool_name,
arguments=json.dumps(content["toolUse"]["input"]),
arguments=arguments_str,
)
_tool_response_chunk = ChatCompletionToolCallChunk(

View file

@ -425,6 +425,15 @@ def strip_bedrock_routing_prefix(model: str) -> str:
return model
def strip_bedrock_throughput_suffix(model: str) -> str:
""" Strip throughput tier suffixes from Bedrock model names. """
import re
# Pattern matches model:version:throughput where throughput is like 51k, 18k, etc.
# Keep the model:version part, strip the :throughput suffix
return re.sub(r"(:\d+):\d+k$", r"\1", model)
def get_bedrock_base_model(model: str) -> str:
"""
Get the base model from the given model name.
@ -432,9 +441,11 @@ def get_bedrock_base_model(model: str) -> str:
Handle model names like:
- "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"
"""
model = strip_bedrock_routing_prefix(model)
model = extract_model_name_from_bedrock_arn(model)
model = strip_bedrock_throughput_suffix(model)
potential_region = model.split(".", 1)[0]
alt_potential_region = model.split("/", 1)[0]

View file

@ -261,7 +261,7 @@ class BedrockImageEdit(BaseAWSLLM):
"""
config_class = self.get_config_class(model=model)
config_instance = config_class()
request_body = config_instance.transform_image_edit_request(
request_body, _ = config_instance.transform_image_edit_request(
model=model,
prompt=prompt,
image=image[0] if image else None,

View file

@ -21,18 +21,18 @@ Supported models:
API Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters.html
"""
import json
import base64
import json
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple
import httpx
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.types.images.main import ImageEditOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
from litellm.types.llms.stability import (
OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import FileTypes, ImageObject, ImageResponse
from litellm.utils import get_model_info
@ -153,7 +153,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig):
def transform_image_edit_request(
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
@ -164,6 +164,9 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig):
Returns the request body dict that will be JSON-encoded by the handler.
"""
if prompt is None:
raise ValueError("Bedrock Stability image edit requires a prompt.")
# Build Bedrock Stability request
data: Dict[str, Any] = {
"prompt": prompt,

View file

@ -129,6 +129,37 @@ class AmazonAnthropicClaudeMessagesConfig(
if isinstance(cache_control, dict) and "ttl" in cache_control:
cache_control.pop("ttl", None)
def _get_tool_search_beta_header_for_bedrock(
self,
model: str,
tool_search_used: bool,
programmatic_tool_calling_used: bool,
input_examples_used: bool,
beta_set: set,
) -> None:
"""
Adjust tool search beta header for Bedrock.
Bedrock requires a different beta header for tool search on Opus 4 models
when tool search is used without programmatic tool calling or input examples.
Note: On Amazon Bedrock, server-side tool search is only supported on Claude Opus 4
with the `tool-search-tool-2025-10-19` beta header.
Ref: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool
Args:
model: The model name
tool_search_used: Whether tool search is used
programmatic_tool_calling_used: Whether programmatic tool calling is used
input_examples_used: Whether input examples are used
beta_set: The set of beta headers to modify in-place
"""
if tool_search_used and not (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():
beta_set.add("tool-search-tool-2025-10-19")
def transform_anthropic_messages_request(
self,
model: str,
@ -189,13 +220,13 @@ class AmazonAnthropicClaudeMessagesConfig(
)
beta_set.update(auto_betas)
if (
tool_search_used
and not (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():
beta_set.add("tool-search-tool-2025-10-19")
self._get_tool_search_beta_header_for_bedrock(
model=model,
tool_search_used=tool_search_used,
programmatic_tool_calling_used=programmatic_tool_calling_used,
input_examples_used=input_examples_used,
beta_set=beta_set,
)
if beta_set:
anthropic_messages_request["anthropic_beta"] = list(beta_set)

View file

@ -245,7 +245,6 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
allow_redirects=False,
auto_decompress=False,
timeout=ClientTimeout(
total=timeout.get("read"),
sock_connect=timeout.get("connect"),
sock_read=timeout.get("read"),
connect=timeout.get("pool"),

View file

@ -4453,7 +4453,7 @@ class BaseLLMHTTPHandler:
self,
model: str,
image: Any,
prompt: str,
prompt: Optional[str],
image_edit_provider_config: BaseImageEditConfig,
image_edit_optional_request_params: Dict,
custom_llm_provider: str,
@ -4572,7 +4572,7 @@ class BaseLLMHTTPHandler:
self,
model: str,
image: FileTypes,
prompt: str,
prompt: Optional[str],
image_edit_provider_config: BaseImageEditConfig,
image_edit_optional_request_params: Dict,
custom_llm_provider: str,

View file

@ -201,7 +201,7 @@ class CustomLLM(BaseLLM):
self,
model: str,
image: Any,
prompt: str,
prompt: Optional[str],
model_response: ImageResponse,
api_key: Optional[str],
api_base: Optional[str],
@ -216,7 +216,7 @@ class CustomLLM(BaseLLM):
self,
model: str,
image: Any,
prompt: str,
prompt: Optional[str],
model_response: ImageResponse,
api_key: Optional[str],
api_base: Optional[str],

View file

@ -80,7 +80,7 @@ class GeminiImageEditConfig(BaseImageEditConfig):
def transform_image_edit_request( # type: ignore[override]
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict[str, Any],
litellm_params: GenericLiteLLMParams,
@ -90,6 +90,9 @@ class GeminiImageEditConfig(BaseImageEditConfig):
if not inline_parts:
raise ValueError("Gemini image edit requires at least one image.")
if prompt is None:
raise ValueError("Gemini image edit requires a prompt.")
contents = [
{
"parts": inline_parts + [{"text": prompt}],

View file

@ -1,5 +1,5 @@
from io import BufferedReader
from typing import TYPE_CHECKING, Any, Dict, List, Tuple, cast
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
from httpx._types import RequestFiles
@ -30,7 +30,7 @@ class DallE2ImageEditConfig(OpenAIImageEditConfig):
def transform_image_edit_request(
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
@ -41,6 +41,9 @@ class DallE2ImageEditConfig(OpenAIImageEditConfig):
DALL-E-2 only accepts a single image with field name "image" (not "image[]").
"""
if prompt is None:
raise ValueError("DALL-E-2 image edit requires a prompt.")
request = ImageEditRequestParams(
model=model,
image=image,

View file

@ -79,7 +79,7 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
def transform_image_edit_request(
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
@ -91,6 +91,9 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
Handles multipart/form-data for images. Uses "image[]" field name
to support multiple images (e.g., for gpt-image-1).
"""
if prompt is None:
raise ValueError("OpenAI image edit requires a prompt.")
request = ImageEditRequestParams(
model=model,
image=image,

View file

@ -101,7 +101,7 @@ class RecraftImageEditConfig(BaseImageEditConfig):
def transform_image_edit_request(
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
@ -114,6 +114,9 @@ class RecraftImageEditConfig(BaseImageEditConfig):
https://www.recraft.ai/docs#image-to-image
"""
if prompt is None:
raise ValueError("Recraft image edit requires a prompt.")
request_body: RecraftImageEditRequestParams = RecraftImageEditRequestParams(
model=model,
prompt=prompt,
@ -124,7 +127,7 @@ class RecraftImageEditConfig(BaseImageEditConfig):
#########################################################
# Reuse OpenAI logic: Separate images as `files` and send other parameters as `data`
#########################################################
files_list = self._get_image_files_for_request(image=image)
files_list = self._get_image_files_for_request(image=image) if image is not None else []
data_without_images = {k: v for k, v in request_dict.items() if k != "image"}
return data_without_images, files_list
@ -132,7 +135,7 @@ class RecraftImageEditConfig(BaseImageEditConfig):
def _get_image_files_for_request(
self,
image: FileTypes,
image: Optional[FileTypes],
) -> List[Tuple[str, Any]]:
files_list: List[Tuple[str, Any]] = []

View file

@ -14,11 +14,11 @@ from httpx._types import RequestFiles
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.images.main import ImageEditOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
from litellm.types.llms.stability import (
OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO,
STABILITY_EDIT_ENDPOINTS,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import FileTypes, ImageObject, ImageResponse
from litellm.utils import get_model_info
@ -170,7 +170,7 @@ class StabilityImageEditConfig(BaseImageEditConfig):
def transform_image_edit_request(
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
@ -186,9 +186,12 @@ class StabilityImageEditConfig(BaseImageEditConfig):
# Populate multipart form-data as separate text fields (data) and files.
# Stability expects prompt/output_format/etc. as normal form fields, not file parts.
data: Dict[str, Any] = {
"prompt": prompt,
"output_format": "png", # Default to PNG
}
# Add prompt only if provided (some Stability endpoints don't require it)
if prompt is not None:
data["prompt"] = prompt
# Handle image parameter - could be a single file or list
image_file = image[0] if isinstance(image, list) else image # type: ignore
files: Dict[str, Any] = {"image": image_file}

View file

@ -665,11 +665,11 @@ def add_object_type(schema):
if "required" in schema and schema["required"] is None:
schema.pop("required", None)
# Gemini doesn't accept empty properties for object types
# If properties is empty, remove it and the type field
# If properties is empty, remove it but keep type as object
if not properties:
schema.pop("properties", None)
schema.pop("type", None)
schema.pop("required", None)
schema["type"] = "object"
else:
schema["type"] = "object"
for name, value in properties.items():
@ -776,6 +776,16 @@ def get_vertex_location_from_url(url: str) -> Optional[str]:
return match.group(1) if match else None
def get_vertex_model_id_from_url(url: str) -> Optional[str]:
"""
Get the vertex model id from the url
`https://${LOCATION}-aiplatform.googleapis.com/v1/projects/${PROJECT_ID}/locations/${LOCATION}/publishers/google/models/${MODEL_ID}:streamGenerateContent`
"""
match = re.search(r"/models/([^/:]+)", url)
return match.group(1) if match else None
def replace_project_and_location_in_route(
requested_route: str, vertex_project: str, vertex_location: str
) -> str:
@ -825,6 +835,15 @@ def construct_target_url(
if "cachedContent" in requested_route:
vertex_version = "v1beta1"
# Check if the requested route starts with a version
# e.g. /v1beta1/publishers/google/models/gemini-3-pro-preview:streamGenerateContent
if requested_route.startswith("/v1/"):
vertex_version = "v1"
requested_route = requested_route.replace("/v1/", "/", 1)
elif requested_route.startswith("/v1beta1/"):
vertex_version = "v1beta1"
requested_route = requested_route.replace("/v1beta1/", "/", 1)
base_requested_route = "{}/projects/{}/locations/{}".format(
vertex_version, vertex_project, vertex_location
)

View file

@ -68,6 +68,8 @@ def _convert_detail_to_media_resolution_enum(
) -> Optional[Dict[str, str]]:
if detail == "low":
return {"level": "MEDIA_RESOLUTION_LOW"}
elif detail == "medium":
return {"level": "MEDIA_RESOLUTION_MEDIUM"}
elif detail == "high":
return {"level": "MEDIA_RESOLUTION_HIGH"}
return None

View file

@ -151,7 +151,7 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM):
def transform_image_edit_request( # type: ignore[override]
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict[str, Any],
litellm_params: GenericLiteLLMParams,
@ -161,6 +161,9 @@ class VertexAIGeminiImageEditConfig(BaseImageEditConfig, VertexLLM):
if not inline_parts:
raise ValueError("Vertex AI Gemini image edit requires at least one image.")
if prompt is None:
raise ValueError("Vertex AI Gemini image edit requires a prompt.")
# Correct format for Vertex AI Gemini image editing
contents = {
"role": "USER",

View file

@ -143,17 +143,22 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
def transform_image_edit_request( # type: ignore[override]
self,
model: str,
prompt: str,
prompt: Optional[str],
image: FileTypes,
image_edit_optional_request_params: Dict[str, Any],
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[Dict[str, Any], Optional[RequestFiles]]:
# Prepare reference images in the correct Imagen format
if image is None:
raise ValueError("Vertex AI Imagen image edit requires at least one reference image.")
reference_images = self._prepare_reference_images(image, image_edit_optional_request_params)
if not reference_images:
raise ValueError("Vertex AI Imagen image edit requires at least one reference image.")
if prompt is None:
raise ValueError("Vertex AI Imagen image edit requires a prompt.")
# Correct Imagen instances format
instances = [
{

View file

@ -1,11 +1,16 @@
from typing import Any, Dict, List, Optional, Tuple
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
from litellm.types.llms.anthropic import (
ANTHROPIC_BETA_HEADER_VALUES,
ANTHROPIC_HOSTED_TOOLS,
)
from litellm.types.llms.anthropic_tool_search import get_tool_search_beta_header
from litellm.types.llms.vertex_ai import VertexPartnerProvider
from litellm.types.router import GenericLiteLLMParams
from litellm.types.llms.anthropic import ANTHROPIC_BETA_HEADER_VALUES, ANTHROPIC_HOSTED_TOOLS
from ....vertex_llm_base import VertexBase
@ -51,13 +56,28 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert
headers["content-type"] = "application/json"
# Add web search beta header for Vertex AI only if not already set
if "anthropic-beta" not in headers:
tools = optional_params.get("tools", [])
for tool in tools:
if isinstance(tool, dict) and tool.get("type", "").startswith(ANTHROPIC_HOSTED_TOOLS.WEB_SEARCH.value):
headers["anthropic-beta"] = ANTHROPIC_BETA_HEADER_VALUES.WEB_SEARCH_2025_03_05.value
break
# Add beta headers for Vertex AI
tools = optional_params.get("tools", [])
beta_values: set[str] = set()
# Get existing beta headers if any
existing_beta = headers.get("anthropic-beta")
if existing_beta:
beta_values.update(b.strip() for b in existing_beta.split(","))
# Check for web search tool
for tool in tools:
if isinstance(tool, dict) and tool.get("type", "").startswith(ANTHROPIC_HOSTED_TOOLS.WEB_SEARCH.value):
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.WEB_SEARCH_2025_03_05.value)
break
# Check for tool search tools - Vertex AI uses different beta header
anthropic_model_info = AnthropicModelInfo()
if anthropic_model_info.is_tool_search_used(tools):
beta_values.add(get_tool_search_beta_header("vertex_ai"))
if beta_values:
headers["anthropic-beta"] = ",".join(beta_values)
return headers, api_base

View file

@ -69,6 +69,9 @@ class VertexAIAnthropicConfig(AnthropicConfig):
data.pop("model", None) # vertex anthropic doesn't accept 'model' parameter
# VertexAI doesn't support output_format parameter, remove it if present
data.pop("output_format", None)
tools = optional_params.get("tools")
tool_search_used = self.is_tool_search_used(tools)
auto_betas = self.get_anthropic_beta_list(
@ -89,6 +92,37 @@ class VertexAIAnthropicConfig(AnthropicConfig):
return data
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Override parent method to ensure VertexAI always uses tool-based structured outputs.
VertexAI doesn't support the output_format parameter, so we force all models
to use the tool-based approach for structured outputs.
"""
# Temporarily override model name to force tool-based approach
# This ensures Claude Sonnet 4.5 uses tools instead of output_format
original_model = model
if "response_format" in non_default_params:
model = "claude-3-sonnet-20240229" # Use a model that will use tool-based approach
# Call parent method with potentially modified model name
optional_params = super().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=drop_params,
)
# Restore original model name for any other processing
model = original_model
return optional_params
def transform_response(
self,
model: str,

View file

@ -3634,6 +3634,37 @@
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5.2-codex": {
"cache_read_input_token_cost": 1.75e-07,
"input_cost_per_token": 1.75e-06,
"litellm_provider": "azure",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 1.4e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5.2-pro": {
"input_cost_per_token": 2.1e-05,
"litellm_provider": "azure",
@ -10170,6 +10201,48 @@
"mode": "completion",
"output_cost_per_token": 5e-07
},
"deepseek-v3-2-251201": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 98304,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"glm-4-7-251222": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 204800,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"kimi-k2-thinking-251104": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 229376,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"doubao-embedding": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
@ -25526,13 +25599,13 @@
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
"output_cost_per_image": 0.4
"output_cost_per_image": 0.40
},
"stability.stable-creative-upscale-v1:0": {
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
"output_cost_per_image": 0.6
"output_cost_per_image": 0.60
},
"stability.stable-fast-upscale-v1:0": {
"litellm_provider": "bedrock",
@ -28782,13 +28855,13 @@
"supports_web_search": true
},
"vertex_ai/zai-org/glm-4.7-maas": {
"input_cost_per_token": 3e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "vertex_ai-zai_models",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"output_cost_per_token": 2.2e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -33930,4 +34003,4 @@
"litellm_provider": "llamagate",
"mode": "embedding"
}
}
}

View file

@ -74,75 +74,6 @@ db_cache_expiry = DEFAULT_IN_MEMORY_TTL # refresh every 5s
all_routes = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
def _is_model_cost_zero(
model: Optional[Union[str, List[str]]], llm_router: Optional[Router]
) -> bool:
"""
Check if a model has zero cost (no configured pricing).
Uses the router's get_model_group_info method to get pricing information.
Args:
model: The model name or list of model names
llm_router: The LiteLLM router instance
Returns:
bool: True if all costs for the model are zero, False otherwise
"""
if model is None or llm_router is None:
return False
# Handle list of models
model_list = [model] if isinstance(model, str) else model
for model_name in model_list:
try:
# Use router's get_model_group_info method directly for better reliability
model_group_info = llm_router.get_model_group_info(model_group=model_name)
if model_group_info is None:
# Model not found or no pricing info available
# Conservative approach: assume it has cost
verbose_proxy_logger.debug(
f"No model group info found for {model_name}, assuming it has cost"
)
return False
# Check costs for this model
# Only allow bypass if BOTH costs are explicitly set to 0 (not None)
input_cost = model_group_info.input_cost_per_token
output_cost = model_group_info.output_cost_per_token
# If costs are not explicitly configured (None), assume it has cost
if input_cost is None or output_cost is None:
verbose_proxy_logger.debug(
f"Model {model_name} has undefined cost (input: {input_cost}, output: {output_cost}), assuming it has cost"
)
return False
# If either cost is non-zero, return False
if input_cost > 0 or output_cost > 0:
verbose_proxy_logger.debug(
f"Model {model_name} has non-zero cost (input: {input_cost}, output: {output_cost})"
)
return False
# This model has zero cost explicitly configured
verbose_proxy_logger.debug(
f"Model {model_name} has zero cost explicitly configured (input: {input_cost}, output: {output_cost})"
)
except Exception as e:
# If we can't determine the cost, assume it has cost (conservative approach)
verbose_proxy_logger.debug(
f"Error checking cost for model {model_name}: {str(e)}, assuming it has cost"
)
return False
# All models checked have zero cost
return True
async def common_checks(
request_body: dict,
team_object: Optional[LiteLLM_TeamTable],
@ -155,7 +86,6 @@ async def common_checks(
proxy_logging_obj: ProxyLogging,
valid_token: Optional[UserAPIKeyAuth],
request: Request,
skip_budget_checks: bool = False,
) -> bool:
"""
Common checks across jwt + key-based auth.
@ -207,66 +137,64 @@ async def common_checks(
user_object=user_object,
)
# If this is a free model, skip all budget checks
if not skip_budget_checks:
# 3. If team is in budget
await _team_max_budget_check(
team_object=team_object,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
# 3. If team is in budget
await _team_max_budget_check(
team_object=team_object,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
# 3.1. If organization is in budget
await _organization_max_budget_check(
valid_token=valid_token,
team_object=team_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# 3.1. If organization is in budget
await _organization_max_budget_check(
valid_token=valid_token,
team_object=team_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await _tag_max_budget_check(
request_body=request_body,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
await _tag_max_budget_check(
request_body=request_body,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
# 4. If user is in budget
## 4.1 check personal budget, if personal key
if (
(team_object is None or team_object.team_id is None)
and user_object is not None
and user_object.max_budget is not None
):
user_budget = user_object.max_budget
if user_budget < user_object.spend:
raise litellm.BudgetExceededError(
current_cost=user_object.spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}",
)
# 4. If user is in budget
## 4.1 check personal budget, if personal key
if (
(team_object is None or team_object.team_id is None)
and user_object is not None
and user_object.max_budget is not None
):
user_budget = user_object.max_budget
if user_budget < user_object.spend:
raise litellm.BudgetExceededError(
current_cost=user_object.spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}",
)
## 4.2 check team member budget, if team key
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
## 4.2 check team member budget, if team key
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
if end_user_object is not None and end_user_object.litellm_budget_table is not None:
end_user_budget = end_user_object.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_object.spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_object.spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
)
# 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
if end_user_object is not None and end_user_object.litellm_budget_table is not None:
end_user_budget = end_user_object.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_object.spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_object.spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
)
# 6. [OPTIONAL] If 'enforce_user_param' enabled - did developer pass in 'user' param for openai endpoints
if (
@ -309,7 +237,6 @@ async def common_checks(
# 7. [OPTIONAL] If 'litellm.max_budget' is set (>0), is proxy under budget
if (
litellm.max_budget > 0
and not skip_budget_checks
and global_proxy_spend is not None
# only run global budget checks for OpenAI routes
# Reason - the Admin UI should continue working if the proxy crosses it's global budget

View file

@ -7,6 +7,7 @@ from fastapi import HTTPException, Request, status
from litellm import Router, provider_list
from litellm._logging import verbose_proxy_logger
from litellm.constants import STANDARD_CUSTOMER_ID_HEADERS
from litellm.proxy._types import *
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS
@ -561,6 +562,32 @@ def get_customer_user_header_from_mapping(user_id_mapping) -> Optional[str]:
return header_name
return None
def _get_customer_id_from_standard_headers(
request_headers: Optional[dict],
) -> Optional[str]:
"""
Check standard customer ID headers for a customer/end-user ID.
This enables tools like Claude Code to pass customer IDs via ANTHROPIC_CUSTOM_HEADERS.
No configuration required - these headers are always checked.
Args:
request_headers: The request headers dict
Returns:
The customer ID if found in standard headers, None otherwise
"""
if request_headers is None:
return None
for standard_header in STANDARD_CUSTOMER_ID_HEADERS:
for header_name, header_value in request_headers.items():
if header_name.lower() == standard_header.lower():
user_id_str = str(header_value) if header_value is not None else ""
if user_id_str.strip():
return user_id_str
return None
def get_end_user_id_from_request_body(
request_body: dict, request_headers: Optional[dict] = None
@ -569,7 +596,12 @@ def get_end_user_id_from_request_body(
# and to ensure it's fetched at runtime.
from litellm.proxy.proxy_server import general_settings
# Check 1 : Follow the user header mappings feature, if not found, then check for deprecated user_header_name (only if request_headers is provided)
# Check 1: Standard customer ID headers (always checked, no configuration required)
customer_id = _get_customer_id_from_standard_headers(request_headers=request_headers)
if customer_id is not None:
return customer_id
# Check 2: Follow the user header mappings feature, if not found, then check for deprecated user_header_name (only if request_headers is provided)
# User query: "system not respecting user_header_name property"
# This implies the key in general_settings is 'user_header_name'.
if request_headers is not None:
@ -602,19 +634,19 @@ def get_end_user_id_from_request_body(
if user_id_str.strip():
return user_id_str
# Check 2: 'user' field in request_body (commonly OpenAI)
# Check 3: 'user' field in request_body (commonly OpenAI)
if "user" in request_body and request_body["user"] is not None:
user_from_body_user_field = request_body["user"]
return str(user_from_body_user_field)
# Check 3: 'litellm_metadata.user' in request_body (commonly Anthropic)
# Check 4: 'litellm_metadata.user' in request_body (commonly Anthropic)
litellm_metadata = request_body.get("litellm_metadata")
if isinstance(litellm_metadata, dict):
user_from_litellm_metadata = litellm_metadata.get("user")
if user_from_litellm_metadata is not None:
return str(user_from_litellm_metadata)
# Check 4: 'metadata.user_id' in request_body (another common pattern)
# Check 5: 'metadata.user_id' in request_body (another common pattern)
metadata_dict = request_body.get("metadata")
if isinstance(metadata_dict, dict):
user_id_from_metadata_field = metadata_dict.get("user_id")

View file

@ -586,21 +586,6 @@ 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
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
skip_budget_checks = _is_model_cost_zero(
model=model, llm_router=llm_router
)
if skip_budget_checks:
verbose_proxy_logger.info(
f"Skipping all budget checks for zero-cost model: {model}"
)
# run through common checks
_ = await common_checks(
request=request,
@ -614,7 +599,6 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
skip_budget_checks=skip_budget_checks,
)
# return UserAPIKeyAuth object
@ -1006,22 +990,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
user_obj = None
# Check 2a. Check if model has zero cost - if so, skip all budget checks
model = get_model_from_request(request_data, route)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
skip_budget_checks = _is_model_cost_zero(
model=model, llm_router=llm_router
)
if skip_budget_checks:
verbose_proxy_logger.info(
f"Skipping all budget checks for zero-cost model: {model}"
)
# Check 3. Check if user is in their team budget
if not skip_budget_checks and valid_token.team_member_spend is not None:
if valid_token.team_member_spend is not None:
if prisma_client is not None:
_cache_key = f"{valid_token.team_id}_{valid_token.user_id}"
@ -1085,47 +1055,46 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
param=abbreviate_api_key(api_key=api_key),
)
if not skip_budget_checks:
# Check 4. Token Spend is under budget
if RouteChecks.is_llm_api_route(route=route):
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 5. Max Budget Alert Check
await _virtual_key_max_budget_alert_check(
# Check 4. Token Spend is under budget
if RouteChecks.is_llm_api_route(route=route):
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 6. Soft Budget Check
await _virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
# Check 5. Max Budget Alert Check
await _virtual_key_max_budget_alert_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 6. Soft Budget Check
await _virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 5. Token Model Spend is under Model budget
max_budget_per_model = valid_token.model_max_budget
current_model = request_data.get("model", None)
if (
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and prisma_client is not None
and current_model is not None
and valid_token.token is not None
):
## GET THE SPEND FOR THIS MODEL
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
# Check 5. Token Model Spend is under Model budget
max_budget_per_model = valid_token.model_max_budget
current_model = request_data.get("model", None)
if (
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and prisma_client is not None
and current_model is not None
and valid_token.token is not None
):
## GET THE SPEND FOR THIS MODEL
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
# Check 6: Additional Common Checks across jwt + key auth
if valid_token.team_id is not None:
_team_obj: Optional[LiteLLM_TeamTable] = LiteLLM_TeamTable(
@ -1193,7 +1162,6 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
skip_budget_checks=skip_budget_checks,
)
# Token passed all checks
if valid_token is None:

View file

@ -49,7 +49,9 @@ if TYPE_CHECKING:
ProxyConfig = _ProxyConfig
else:
ProxyConfig = Any
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.litellm_pre_call_utils import (
add_litellm_data_to_request,
)
from litellm.types.utils import ModelResponse, ModelResponseStream, Usage

View file

@ -50,8 +50,32 @@ from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter impor
ContentFilterDetection,
PatternDetection,
)
from .patterns import PATTERN_EXTRA_CONFIG, get_compiled_pattern
from .patterns import get_compiled_pattern
MAX_KEYWORD_VALUE_GAP_WORDS = 1
GAP_WORD_TOKENIZER = re.compile(r"\b\w+\b")
WORD_NUMBER_MAP = {
"zero": "0",
"oh": "0",
"one": "1",
"two": "2",
"three": "3",
"four": "4",
"five": "5",
"six": "6",
"seven": "7",
"eight": "8",
"nine": "9",
}
WORD_NUMBER_TOKEN_REGEX = "|".join(WORD_NUMBER_MAP.keys())
WORD_NUMBER_SEQUENCE_PATTERN = re.compile(
rf"(?<![A-Za-z])(?:{WORD_NUMBER_TOKEN_REGEX})(?:[\s\-]+(?:{WORD_NUMBER_TOKEN_REGEX}))+(?![A-Za-z])",
re.IGNORECASE,
)
WORD_NUMBER_TOKEN_FINDER = re.compile(rf"(?:{WORD_NUMBER_TOKEN_REGEX})", re.IGNORECASE)
# Helper data structure for category-based detection
@ -144,9 +168,9 @@ class ContentFilterGuardrail(CustomGuardrail):
self.image_model = image_model
# Store loaded categories
self.loaded_categories: Dict[str, CategoryConfig] = {}
self.category_keywords: Dict[str, Tuple[str, str, ContentFilterAction]] = (
{}
) # keyword -> (category, severity, action)
self.category_keywords: Dict[
str, Tuple[str, str, ContentFilterAction]
] = {} # keyword -> (category, severity, action)
# Load categories if provided
if categories:
@ -170,7 +194,7 @@ class ContentFilterGuardrail(CustomGuardrail):
normalized_blocked_words.append(word)
# Compile regex patterns
self.compiled_patterns: List[Tuple[Pattern, str, ContentFilterAction]] = []
self.compiled_patterns: List[Dict[str, Any]] = []
for pattern_config in normalized_patterns:
self._add_pattern(pattern_config)
@ -323,11 +347,13 @@ class ContentFilterGuardrail(CustomGuardrail):
pattern_config: ContentFilterPattern configuration
"""
try:
extra_config: Dict[str, Any] = {}
if pattern_config.pattern_type == "prebuilt":
if not pattern_config.pattern_name:
raise ValueError("pattern_name is required for prebuilt patterns")
compiled = get_compiled_pattern(pattern_config.pattern_name)
pattern_name = pattern_config.pattern_name
extra_config = PATTERN_EXTRA_CONFIG.get(pattern_name, {}) or {}
elif pattern_config.pattern_type == "regex":
if not pattern_config.pattern:
raise ValueError("pattern is required for regex patterns")
@ -336,8 +362,20 @@ class ContentFilterGuardrail(CustomGuardrail):
else:
raise ValueError(f"Unknown pattern_type: {pattern_config.pattern_type}")
keyword_regex: Optional[Pattern] = None
if extra_config.get("keyword_pattern"):
keyword_regex = re.compile(
extra_config["keyword_pattern"], re.IGNORECASE
)
self.compiled_patterns.append(
(compiled, pattern_name, pattern_config.action)
{
"regex": compiled,
"pattern_name": pattern_name,
"action": pattern_config.action,
"keyword_regex": keyword_regex,
"allow_word_numbers": bool(extra_config.get("allow_word_numbers")),
}
)
verbose_proxy_logger.debug(
f"Added pattern: {pattern_name} with action {pattern_config.action}"
@ -395,6 +433,130 @@ class ContentFilterGuardrail(CustomGuardrail):
except Exception as e:
raise Exception(f"Error loading blocked words file {file_path}: {str(e)}")
def _find_pattern_spans(
self, text: str, pattern_entry: Dict[str, Any]
) -> List[Tuple[int, int]]:
"""Return all match spans for a pattern, applying contextual rules if required."""
regex: Pattern = pattern_entry["regex"]
keyword_regex: Optional[Pattern] = pattern_entry.get("keyword_regex")
allow_word_numbers: bool = pattern_entry.get("allow_word_numbers", False)
keyword_matches: Optional[List[re.Match]] = None
if keyword_regex is not None:
keyword_matches = list(keyword_regex.finditer(text))
if not keyword_matches:
return []
match_spans: List[Tuple[int, int]] = []
for match in regex.finditer(text):
if keyword_matches is not None and not self._match_near_keyword(
match.start(), match.end(), keyword_matches, text
):
continue
match_spans.append((match.start(), match.end()))
if allow_word_numbers:
for word_match in WORD_NUMBER_SEQUENCE_PATTERN.finditer(text):
digits = self._convert_word_number_sequence(word_match.group())
if not digits:
continue
if not regex.fullmatch(digits):
continue
if keyword_matches is not None and not self._match_near_keyword(
word_match.start(), word_match.end(), keyword_matches, text
):
continue
match_spans.append((word_match.start(), word_match.end()))
return self._merge_spans(match_spans)
def _match_near_keyword(
self,
value_start: int,
value_end: int,
keyword_matches: List[re.Match],
text: str,
) -> bool:
"""Check if a value is separated from a keyword by an allowed gap."""
for keyword_match in keyword_matches:
keyword_start = keyword_match.start()
keyword_end = keyword_match.end()
if value_start >= keyword_end:
gap_text = text[keyword_end:value_start]
elif keyword_start >= value_end:
gap_text = text[value_end:keyword_start]
else:
return True # overlapping
if self._gap_text_allowed(gap_text):
return True
return False
def _gap_text_allowed(self, gap_text: str) -> bool:
"""Return True if the gap between keyword and value meets word-count rules."""
if not gap_text.strip():
return True
if any(char.isdigit() for char in gap_text):
return False
words = GAP_WORD_TOKENIZER.findall(gap_text)
return len(words) <= MAX_KEYWORD_VALUE_GAP_WORDS
def _merge_spans(self, spans: List[Tuple[int, int]]) -> List[Tuple[int, int]]:
"""Merge overlapping spans to avoid double-masking."""
if not spans:
return []
spans.sort(key=lambda item: item[0])
merged: List[Tuple[int, int]] = [spans[0]]
for start, end in spans[1:]:
last_start, last_end = merged[-1]
if start <= last_end:
merged[-1] = (last_start, max(last_end, end))
else:
merged.append((start, end))
return merged
def _mask_spans(
self, text: str, spans: List[Tuple[int, int]], redaction: str
) -> str:
"""Apply masking for the provided spans using the given redaction tag."""
if not spans:
return text
result_parts: List[str] = []
previous_end = 0
for start, end in spans:
result_parts.append(text[previous_end:start])
result_parts.append(redaction)
previous_end = end
result_parts.append(text[previous_end:])
return "".join(result_parts)
def _convert_word_number_sequence(self, sequence: str) -> Optional[str]:
"""Convert a spelled-out digit sequence (e.g., 'One-Two') into digits."""
tokens = WORD_NUMBER_TOKEN_FINDER.findall(sequence)
if not tokens:
return None
digits: List[str] = []
for token in tokens:
digit = WORD_NUMBER_MAP.get(token.lower())
if digit is None:
return None
digits.append(digit)
return "".join(digits) if digits else None
def _check_patterns(
self, text: str
) -> Optional[Tuple[str, str, ContentFilterAction]]:
@ -407,10 +569,13 @@ class ContentFilterGuardrail(CustomGuardrail):
Returns:
Tuple of (matched_text, pattern_name, action) if match found, None otherwise
"""
for compiled_pattern, pattern_name, action in self.compiled_patterns:
match = compiled_pattern.search(text)
if match:
matched_text = match.group(0)
for pattern_entry in self.compiled_patterns:
spans = self._find_pattern_spans(text, pattern_entry)
if spans:
start, end = spans[0]
matched_text = text[start:end]
pattern_name = pattern_entry["pattern_name"]
action = pattern_entry["action"]
verbose_proxy_logger.debug(
f"Pattern '{pattern_name}' matched: {matched_text[:20]}..."
)
@ -582,11 +747,13 @@ class ContentFilterGuardrail(CustomGuardrail):
)
# Check regex patterns - process ALL patterns, not just first match
for compiled_pattern, pattern_name, action in self.compiled_patterns:
match = compiled_pattern.search(text)
if not match:
for pattern_entry in self.compiled_patterns:
spans = self._find_pattern_spans(text, pattern_entry)
if not spans:
continue
pattern_name = pattern_entry["pattern_name"]
action = pattern_entry["action"]
if detections is not None:
# Don't log matched_text to avoid exposing sensitive content (emails, credit cards, etc.)
pattern_detection: PatternDetection = {
@ -604,11 +771,10 @@ class ContentFilterGuardrail(CustomGuardrail):
detail={"error": error_msg, "pattern": pattern_name},
)
elif action == ContentFilterAction.MASK:
# Replace ALL matches of this pattern with redaction tag
redaction_tag = self.pattern_redaction_format.format(
pattern_name=pattern_name.upper()
)
text = compiled_pattern.sub(redaction_tag, text)
text = self._mask_spans(text, spans, redaction_tag)
verbose_proxy_logger.info(
f"Masked all {pattern_name} matches in content"
)
@ -924,19 +1090,28 @@ class ContentFilterGuardrail(CustomGuardrail):
if pattern_match:
matched_text, pattern_name, action = pattern_match
if action == ContentFilterAction.BLOCK:
error_msg = f"Content blocked: {pattern_name} pattern detected"
error_msg = (
f"Content blocked: {pattern_name} pattern detected"
)
verbose_proxy_logger.warning(error_msg)
raise HTTPException(
status_code=403,
detail={"error": error_msg, "pattern": pattern_name},
detail={
"error": error_msg,
"pattern": pattern_name,
},
)
# Check blocked words
blocked_word_match = self._check_blocked_words(accumulated_content)
blocked_word_match = self._check_blocked_words(
accumulated_content
)
if blocked_word_match:
keyword, action, description = blocked_word_match
if action == ContentFilterAction.BLOCK:
error_msg = f"Content blocked: keyword '{keyword}' detected"
error_msg = (
f"Content blocked: keyword '{keyword}' detected"
)
if description:
error_msg += f" ({description})"
verbose_proxy_logger.warning(error_msg)

View file

@ -120,11 +120,11 @@
"description": "Detects URLs (http/https)"
},
{
"name": "passport_us",
"display_name": "Passport (US)",
"pattern": "\\b[0-9]{9}\\b",
"category": "PII Patterns",
"description": "US passport numbers (9 digits)"
"name": "passport_us",
"display_name": "Passport (US)",
"pattern": "\\b[0-9]{9}\\b",
"category": "PII Patterns",
"description": "US passport numbers (9 digits)"
},
{
"name": "passport_uk",
@ -203,7 +203,6 @@
"category": "Protected Class - Fair Lending",
"description": "Detects race, ethnicity and national origin terms - protected under ECOA and Fair Housing Act"
},
{
"name": "religion",
"display_name": "Religion & Creed (Protected Class)",
@ -236,7 +235,7 @@
"name": "military_status",
"display_name": "Military Status (Protected Class)",
"pattern": "\\b(veteran|military|armed\\s+forces|army|navy|air\\s+force|marine(s|\\s+corps)?|coast\\s+guard|national\\s+guard|reserve(s|ist)?|active\\s+duty|deployment|deployed|enlisted|commissioned|honorable\\s+discharge|dishonorable\\s+discharge|VA\\s+benefits|GI\\s+bill|military\\s+service|service\\s+member|servicemember|SCRA|MLA|military\\s+lending)\\b",
"category": "Protected Class - Fair Lending",
"category": "Protected Class - Fair Lending",
"description": "Detects military status terms - protected under SCRA and MLA"
},
{
@ -245,7 +244,7 @@
"pattern": "\\b(welfare|public\\s+assistance|food\\s+stamps|SNAP|WIC|TANF|medicaid|section\\s+8|housing\\s+voucher|subsidized\\s+housing|public\\s+housing|government\\s+benefits|social\\s+services|unemployment\\s+(benefits|insurance)|UI\\s+benefits|EBT|benefit\\s+recipient)\\b",
"category": "Protected Class - Fair Lending",
"description": "Detects public assistance terms - protected under ECOA"
} ,
},
{
"name": "weapons_firearms",
"display_name": "Weapons & Firearms",
@ -313,10 +312,12 @@
{
"name": "nl_bsn_contextual",
"display_name": "BSN (Dutch Citizen Service Number)",
"pattern": "\\b(?:BSN|B\\.S\\.N\\.|burgerservicenummer|burger\\s*service\\s*nummer|sofi\\s*nummer|sofinummer|persoonsnummer|identificatienummer|citizen\\s*service\\s*number)[:\\s]*[0-9]{9}\\b|\\b[0-9]{9}\\b(?=\\s*(?:BSN|burgerservicenummer|sofinummer))",
"pattern": "\\b[0-9]{9}\\b",
"category": "PII Patterns",
"action": "MASK",
"description": "Detects Dutch BSN numbers with contextual keywords"
"description": "Detects Dutch BSN numbers with contextual keywords",
"keyword_pattern": "(?:\\b(?:BSN|B\\.S\\.N\\.|burgerservicenummer|burger\\s*service\\s*nummer|sofi\\s*nummer|sofinummer|persoonsnummer|identificatienummer|citizen\\s*service\\s*number)\\b|8\\s*5\\s*\\|\\\\\\|)",
"allow_word_numbers": true
},
{
"name": "br_cpf",
@ -369,5 +370,3 @@
}
]
}

View file

@ -9,7 +9,7 @@ import json
import os
import re
from enum import Enum
from typing import Dict, List, Pattern
from typing import Any, Dict, List, Pattern
def _load_patterns_from_json() -> Dict:
@ -41,6 +41,26 @@ PREBUILT_PATTERNS: Dict[str, str] = {
}
# Capture any extra configuration declared per pattern (e.g., contextual keywords)
KNOWN_PATTERN_KEYS = {
"name",
"display_name",
"pattern",
"category",
"action",
"description",
}
PATTERN_EXTRA_CONFIG: Dict[str, Dict[str, Any]] = {}
for pattern_data in _PATTERNS_DATA["patterns"]:
extra_config = {
key: value
for key, value in pattern_data.items()
if key not in KNOWN_PATTERN_KEYS
}
PATTERN_EXTRA_CONFIG[pattern_data["name"]] = extra_config
def get_compiled_pattern(pattern_name: str) -> Pattern:
"""
Get a compiled regex pattern by name.

View file

@ -114,25 +114,25 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
) -> Optional[str]:
"""
Get priority from user_api_key_dict.
Checks team metadata first (takes precedence), then falls back to key metadata.
Args:
user_api_key_dict: User authentication info
Returns:
Priority string if found, None otherwise
"""
priority: Optional[str] = None
# Check team metadata first (takes precedence)
if user_api_key_dict.team_metadata is not None:
priority = user_api_key_dict.team_metadata.get("priority", None)
# Fall back to key metadata
if priority is None:
priority = user_api_key_dict.metadata.get("priority", None)
return priority
def _normalize_priority_weights(
@ -299,10 +299,13 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
"""
descriptors: List[RateLimitDescriptor] = []
if litellm.priority_reservation is None:
return descriptors
# Get model group info
model_group_info: Optional[ModelGroupInfo] = (
self.llm_router.get_model_group_info(model_group=model)
)
model_group_info: Optional[
ModelGroupInfo
] = self.llm_router.get_model_group_info(model_group=model)
if model_group_info is None:
return descriptors
@ -577,9 +580,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
)
# Get model configuration
model_group_info: Optional[ModelGroupInfo] = (
self.llm_router.get_model_group_info(model_group=model)
)
model_group_info: Optional[
ModelGroupInfo
] = self.llm_router.get_model_group_info(model_group=model)
if model_group_info is None:
verbose_proxy_logger.debug(
f"No model group info for {model}, allowing request"
@ -703,7 +706,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
# Get priority from user_api_key_auth_metadata in standard_logging_metadata
# This is where user_api_key_dict.metadata is stored during pre-call
user_api_key_auth_metadata = standard_logging_metadata.get("user_api_key_auth_metadata") or {}
user_api_key_auth_metadata = (
standard_logging_metadata.get("user_api_key_auth_metadata") or {}
)
key_priority: Optional[str] = user_api_key_auth_metadata.get("priority")
# Get total tokens from response
@ -775,7 +780,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
# Only log 'priority' if it's known safe; otherwise, redact.
SAFE_PRIORITIES = {"low", "medium", "high", "default"}
logged_priority = key_priority if key_priority in SAFE_PRIORITIES else "REDACTED"
logged_priority = (
key_priority if key_priority in SAFE_PRIORITIES else "REDACTED"
)
verbose_proxy_logger.debug(
f"[Dynamic Rate Limiter] Incremented tokens by {total_tokens} for "
f"model={model_group}, priority={logged_priority}"

View file

@ -1236,7 +1236,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return pipeline_operations
def _get_total_tokens_from_usage(
self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"]
self, usage: Optional[Any], rate_limit_type: Literal["output", "input", "total"]
) -> int:
"""
Get total tokens from response usage for rate limiting.

View file

@ -846,7 +846,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915
# Add headers to metadata for guardrails to access (fixes #17477)
# Guardrails use metadata["headers"] to access request headers (e.g., User-Agent)
if _metadata_variable_name in data and isinstance(data[_metadata_variable_name], dict):
if _metadata_variable_name in data and isinstance(
data[_metadata_variable_name], dict
):
data[_metadata_variable_name]["headers"] = _headers
# check for forwardable headers
@ -1314,6 +1316,9 @@ def move_guardrails_to_metadata(
- If guardrails set on API Key metadata then sets guardrails on request metadata
- If guardrails not set on API key, then checks request metadata
Note: We copy (not pop) guardrails from data to metadata to ensure deployment-level
guardrails merged by the router remain in kwargs for async_pre_call_deployment_hook.
"""
# Check key-level guardrails
_add_guardrails_from_key_or_team_metadata(
@ -1326,15 +1331,25 @@ def move_guardrails_to_metadata(
#########################################################################################
# User's might send "guardrails" in the request body, we need to add them to the request metadata.
# Since downstream logic requires "guardrails" to be in the request metadata
#
# IMPORTANT: We copy instead of pop to preserve guardrails in kwargs for
# async_pre_call_deployment_hook (custom_guardrail.py:290) which checks kwargs.get("guardrails").
# This is the event-based approach for deployment-level guardrails.
#########################################################################################
if "guardrails" in data:
request_body_guardrails = data.pop("guardrails")
request_body_guardrails = data.get("guardrails")
if request_body_guardrails is None:
return
if "guardrails" in data[_metadata_variable_name] and isinstance(
data[_metadata_variable_name]["guardrails"], list
):
data[_metadata_variable_name]["guardrails"].extend(request_body_guardrails)
# Merge unique guardrails
existing = data[_metadata_variable_name]["guardrails"]
for g in request_body_guardrails:
if g not in existing:
existing.append(g)
else:
data[_metadata_variable_name]["guardrails"] = request_body_guardrails
data[_metadata_variable_name]["guardrails"] = list(request_body_guardrails)
#########################################################################################
if "guardrail_config" in data:

View file

@ -343,7 +343,7 @@ def _build_where_conditions(
start_date: str,
end_date: str,
model: Optional[str],
api_key: Optional[Union[str, List[str]]],
api_key: Optional[str],
exclude_entity_ids: Optional[List[str]] = None,
) -> Dict[str, Any]:
"""Build prisma where clause for daily activity queries."""
@ -357,10 +357,7 @@ def _build_where_conditions(
if model:
where_conditions["model"] = model
if api_key:
if isinstance(api_key, list):
where_conditions["api_key"] = {"in": api_key}
else:
where_conditions["api_key"] = api_key
where_conditions["api_key"] = api_key
if entity_id is not None:
if isinstance(entity_id, list):
@ -448,7 +445,7 @@ async def get_daily_activity(
start_date: Optional[str],
end_date: Optional[str],
model: Optional[str],
api_key: Optional[Union[str, List[str]]],
api_key: Optional[str],
page: int,
page_size: int,
exclude_entity_ids: Optional[List[str]] = None,

View file

@ -0,0 +1,375 @@
"""
FALLBACK MANAGEMENT ENDPOINTS
Dedicated endpoints for managing model fallbacks separately from general config.
POST /fallback - Create or update fallbacks for a specific model
GET /fallback/{model} - Get fallbacks for a specific model
DELETE /fallback/{model} - Delete fallbacks for a specific model
"""
# pyright: reportMissingImports=false
import json
from typing import TYPE_CHECKING, Dict, List, Literal
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.model_checks import get_all_fallbacks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
if TYPE_CHECKING:
from fastapi import APIRouter, Depends, HTTPException, status
else:
try:
from fastapi import APIRouter, Depends, HTTPException, status
except ImportError:
# fastapi is only required for proxy, not for SDK usage
pass
from litellm.types.management_endpoints.router_settings_endpoints import (
FallbackCreateRequest,
FallbackDeleteResponse,
FallbackGetResponse,
FallbackResponse,
)
router = APIRouter()
@router.post(
"/fallback",
tags=["Fallback Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=FallbackResponse,
status_code=status.HTTP_200_OK,
)
async def create_fallback(
data: FallbackCreateRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Create or update fallbacks for a specific model.
This endpoint allows you to configure fallback models separately from the general config.
Fallbacks are triggered when a model call fails after retries.
**Example Request:**
```json
{
"model": "gpt-3.5-turbo",
"fallback_models": ["gpt-4", "claude-3-haiku"],
"fallback_type": "general"
}
```
**Fallback Types:**
- `general`: Standard fallbacks for any error (default)
- `context_window`: Fallbacks specifically for context window exceeded errors
- `content_policy`: Fallbacks specifically for content policy violations
"""
from litellm.proxy.proxy_server import (
llm_router,
prisma_client,
proxy_config,
store_model_in_db,
)
try:
# Validate that we have a router
if llm_router is None:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": "Router not initialized"},
)
# Validate that the model exists in the router
model_names = llm_router.model_names
if data.model not in model_names:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={
"error": f"Model '{data.model}' not found in router",
"available_models": list(model_names),
},
)
# Validate that all fallback models exist in the router
invalid_fallback_models = [
m for m in data.fallback_models if m not in model_names
]
if invalid_fallback_models:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": f"Invalid fallback models: {invalid_fallback_models}",
"available_models": list(model_names),
},
)
# Check if fallback model is the same as the primary model
if data.model in data.fallback_models:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": f"Model '{data.model}' cannot be its own fallback"
},
)
# Check if we need to store in DB
if store_model_in_db is not True or prisma_client is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "Database storage not enabled. Set 'STORE_MODEL_IN_DB=True' in your environment to use this feature."
},
)
# Load existing config
config = await proxy_config.get_config()
router_settings = config.get("router_settings", {})
# Get the appropriate fallback list based on type
fallback_key = "fallbacks"
if data.fallback_type == "context_window":
fallback_key = "context_window_fallbacks"
elif data.fallback_type == "content_policy":
fallback_key = "content_policy_fallbacks"
# Get existing fallbacks
existing_fallbacks: List[Dict[str, List[str]]] = router_settings.get(
fallback_key, []
)
# Update or add the fallback configuration
fallback_updated = False
for i, fallback_dict in enumerate(existing_fallbacks):
if data.model in fallback_dict:
# Update existing fallback
existing_fallbacks[i] = {data.model: data.fallback_models}
fallback_updated = True
break
if not fallback_updated:
# Add new fallback
existing_fallbacks.append({data.model: data.fallback_models})
# Update router settings
router_settings[fallback_key] = existing_fallbacks
# Save to database - convert router_settings to JSON string
router_settings_json = json.dumps(router_settings)
await prisma_client.db.litellm_config.upsert(
where={"param_name": "router_settings"},
data={
"create": {
"param_name": "router_settings",
"param_value": router_settings_json,
},
"update": {
"param_value": router_settings_json
},
},
)
# Update the in-memory router configuration
setattr(llm_router, fallback_key, existing_fallbacks)
verbose_proxy_logger.info(
f"Fallback configured: {data.model} -> {data.fallback_models} (type: {data.fallback_type})"
)
return FallbackResponse(
model=data.model,
fallback_models=data.fallback_models,
fallback_type=data.fallback_type,
message=f"Fallback configuration {'updated' if fallback_updated else 'created'} successfully",
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.error(f"Error creating fallback: {str(e)}", exc_info=True)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Failed to create fallback: {str(e)}"},
)
@router.get(
"/fallback/{model}",
tags=["Fallback Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=FallbackGetResponse,
)
async def get_fallback(
model: str,
fallback_type: Literal["general", "context_window", "content_policy"] = "general",
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get fallback configuration for a specific model.
**Parameters:**
- `model`: The model name to get fallbacks for
- `fallback_type`: Type of fallback to retrieve (query parameter)
**Example:**
```
GET /fallback/gpt-3.5-turbo?fallback_type=general
```
"""
from litellm.proxy.proxy_server import llm_router
try:
if llm_router is None:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": "Router not initialized"},
)
# Get fallbacks using the existing utility function
fallback_models = get_all_fallbacks(
model=model, llm_router=llm_router, fallback_type=fallback_type
)
if not fallback_models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={
"error": f"No {fallback_type} fallbacks configured for model '{model}'"
},
)
return FallbackGetResponse(
model=model,
fallback_models=fallback_models,
fallback_type=fallback_type,
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.error(f"Error getting fallback: {str(e)}", exc_info=True)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Failed to get fallback: {str(e)}"},
)
@router.delete(
"/fallback/{model}",
tags=["Fallback Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=FallbackDeleteResponse,
)
async def delete_fallback(
model: str,
fallback_type: Literal["general", "context_window", "content_policy"] = "general",
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Delete fallback configuration for a specific model.
**Parameters:**
- `model`: The model name to delete fallbacks for
- `fallback_type`: Type of fallback to delete (query parameter)
**Example:**
```
DELETE /fallback/gpt-3.5-turbo?fallback_type=general
```
"""
from litellm.proxy.proxy_server import (
llm_router,
prisma_client,
proxy_config,
store_model_in_db,
)
try:
if llm_router is None:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": "Router not initialized"},
)
if store_model_in_db is not True or prisma_client is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "Database storage not enabled. Set 'STORE_MODEL_IN_DB=True' in your environment to use this feature."
},
)
# Load existing config
config = await proxy_config.get_config()
router_settings = config.get("router_settings", {})
# Get the appropriate fallback list based on type
fallback_key = "fallbacks"
if fallback_type == "context_window":
fallback_key = "context_window_fallbacks"
elif fallback_type == "content_policy":
fallback_key = "content_policy_fallbacks"
# Get existing fallbacks
existing_fallbacks: List[Dict[str, List[str]]] = router_settings.get(
fallback_key, []
)
# Find and remove the fallback configuration
fallback_found = False
updated_fallbacks = []
for fallback_dict in existing_fallbacks:
if model not in fallback_dict:
updated_fallbacks.append(fallback_dict)
else:
fallback_found = True
if not fallback_found:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={
"error": f"No {fallback_type} fallbacks configured for model '{model}'"
},
)
# Update router settings
router_settings[fallback_key] = updated_fallbacks
# Save to database - convert router_settings to JSON string
router_settings_json = json.dumps(router_settings)
await prisma_client.db.litellm_config.upsert(
where={"param_name": "router_settings"},
data={
"create": {
"param_name": "router_settings",
"param_value": router_settings_json,
},
"update": {
"param_value": router_settings_json
},
},
)
# Update the in-memory router configuration
setattr(llm_router, fallback_key, updated_fallbacks)
verbose_proxy_logger.info(
f"Fallback deleted: {model} (type: {fallback_type})"
)
return FallbackDeleteResponse(
model=model,
fallback_type=fallback_type,
message="Fallback configuration deleted successfully",
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.error(f"Error deleting fallback: {str(e)}", exc_info=True)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Failed to delete fallback: {str(e)}"},
)

View file

@ -3715,7 +3715,7 @@ async def get_team_daily_activity(
},
)
## Fetch team aliases and check team admin status
## Fetch team aliases
where_condition = {}
if team_ids_list:
where_condition["team_id"] = {"in": list(team_ids_list)}
@ -3726,36 +3726,6 @@ async def get_team_daily_activity(
t.team_id: {"team_alias": t.team_alias} for t in team_aliases
}
# Check if user is team admin for any requested teams
# If not, filter by user's API keys
user_api_keys: Optional[List[str]] = None
if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases:
# Check if user is team admin for any of the teams
is_team_admin_for_any = False
for team_alias in team_aliases:
team_obj = LiteLLM_TeamTable(**team_alias.model_dump())
if _is_user_team_admin(
user_api_key_dict=user_api_key_dict, team_obj=team_obj
):
is_team_admin_for_any = True
break
# If user is not a team admin for any team, filter by their API keys
if not is_team_admin_for_any:
# Get all API keys for this user
user_keys = await prisma_client.db.litellm_verificationtoken.find_many(
where={"user_id": user_api_key_dict.user_id}
)
user_api_keys = [key.token for key in user_keys if key.token]
# If user has no API keys, return empty result
if not user_api_keys:
user_api_keys = [""] # Use empty string to ensure no matches
# If api_key parameter is provided, use it; otherwise use user_api_keys if set
final_api_key_filter: Optional[Union[str, List[str]]] = api_key
if final_api_key_filter is None and user_api_keys is not None:
final_api_key_filter = user_api_keys
return await get_daily_activity(
prisma_client=prisma_client,
table_name="litellm_dailyteamspend",
@ -3766,7 +3736,7 @@ async def get_team_daily_activity(
start_date=start_date,
end_date=end_date,
model=model,
api_key=final_api_key_filter,
api_key=api_key,
page=page,
page_size=page_size,
)

View file

@ -1554,6 +1554,7 @@ async def _base_vertex_proxy_route(
from litellm.llms.vertex_ai.common_utils import (
construct_target_url,
get_vertex_location_from_url,
get_vertex_model_id_from_url,
get_vertex_project_id_from_url,
)
@ -1583,6 +1584,25 @@ async def _base_vertex_proxy_route(
vertex_location=vertex_location,
)
if vertex_project is None or vertex_location is None:
# Check if model is in router config
model_id = get_vertex_model_id_from_url(endpoint)
if model_id:
from litellm.proxy.proxy_server import llm_router
if llm_router:
try:
# Use the dedicated pass-through deployment selection method to automatically filter use_in_pass_through=True
deployment = llm_router.get_available_deployment_for_pass_through(model=model_id)
if deployment:
litellm_params = deployment.get("litellm_params", {})
vertex_project = litellm_params.get("vertex_project")
vertex_location = litellm_params.get("vertex_location")
except Exception as e:
verbose_proxy_logger.debug(
f"Error getting available deployment for model {model_id}: {e}"
)
vertex_credentials = passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
location=vertex_location,

View file

@ -26,3 +26,5 @@ if exit_code != 0:
verbose_proxy_logger.error(
f"'prisma generate' stderr: {result.stderr}"
) # Log stderr
sys.exit(exit_code)

View file

@ -187,6 +187,7 @@ class ProxyInitializationHelpers:
ssl_certfile_path: str,
ssl_keyfile_path: str,
max_requests_before_restart: Optional[int] = None,
keepalive_timeout: Optional[int] = None,
):
"""
Run litellm with `gunicorn`
@ -267,6 +268,10 @@ class ProxyInitializationHelpers:
"access_log_format": '%(h)s %(l)s %(u)s %(t)s "%(r)s" %(s)s %(b)s',
}
# Optional: set keepalive timeout if specified by user
if keepalive_timeout is not None:
gunicorn_options["keepalive"] = keepalive_timeout
# Optional: recycle workers after N requests to mitigate memory growth
if max_requests_before_restart is not None:
gunicorn_options["max_requests"] = max_requests_before_restart
@ -489,7 +494,7 @@ class ProxyInitializationHelpers:
"--keepalive_timeout",
default=None,
type=int,
help="Set the uvicorn keepalive timeout in seconds (uvicorn timeout_keep_alive parameter)",
help="Set the keepalive timeout in seconds. For Uvicorn: timeout_keep_alive parameter. For Gunicorn: keepalive parameter. Default: Uvicorn uses ~75s, Gunicorn uses 90s",
envvar="KEEPALIVE_TIMEOUT",
)
@click.option(
@ -859,6 +864,7 @@ def run_server( # noqa: PLR0915
ssl_certfile_path=ssl_certfile_path,
ssl_keyfile_path=ssl_keyfile_path,
max_requests_before_restart=max_requests_before_restart,
keepalive_timeout=keepalive_timeout,
)
elif run_hypercorn is True:
ProxyInitializationHelpers._init_hypercorn_server(

View file

@ -37,6 +37,12 @@ model_list:
model_info:
litellm_provider: bedrock_converse
mode: chat
- model_name: azure-claude-opus-4-5
litellm_params:
model: azure_ai/claude-opus-4-5
api_base: https://krish-mh44t553-eastus2.services.ai.azure.com
api_key: os.environ/AZURE_ANTHROPIC_API_KEY
general_settings:
store_prompts_in_spend_logs: true

View file

@ -297,10 +297,15 @@ from litellm.proxy.management_endpoints.cost_tracking_settings import (
from litellm.proxy.management_endpoints.customer_endpoints import (
router as customer_router,
)
from litellm.proxy.management_endpoints.fallback_management_endpoints import (
router as fallback_management_router,
)
from litellm.proxy.management_endpoints.internal_user_endpoints import (
router as internal_user_router,
)
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
from litellm.proxy.management_endpoints.internal_user_endpoints import (
user_update,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
delete_verification_tokens,
duration_in_seconds,
@ -354,7 +359,9 @@ from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
from litellm.proxy.openai_files_endpoints.files_endpoints import (
router as openai_files_router,
)
from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config
from litellm.proxy.openai_files_endpoints.files_endpoints import (
set_files_config,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
passthrough_endpoint_router,
)
@ -449,7 +456,9 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
LiteLLM_UpperboundKeyGenerateParams,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import DeploymentTypedDict
from litellm.types.router import (
DeploymentTypedDict,
)
from litellm.types.router import ModelInfo as RouterModelInfo
from litellm.types.router import (
RouterGeneralSettings,
@ -3253,20 +3262,23 @@ class ProxyConfig:
) -> Optional[dict]:
"""
Get router_settings in priority order: Key > Team > Global
Returns:
dict: Combined router_settings, or None if no settings found
"""
if prisma_client is None:
return None
import json
import yaml
# 1. Try key-level router_settings
if user_api_key_dict is not None:
# Check if router_settings is available on the key object
key_router_settings_value = getattr(user_api_key_dict, "router_settings", None)
key_router_settings_value = getattr(
user_api_key_dict, "router_settings", None
)
if key_router_settings_value is not None:
key_router_settings = None
if isinstance(key_router_settings_value, str):
@ -3279,11 +3291,15 @@ class ProxyConfig:
pass
elif isinstance(key_router_settings_value, dict):
key_router_settings = key_router_settings_value
# If key has router_settings (non-empty dict), use it
if key_router_settings is not None and isinstance(key_router_settings, dict) and key_router_settings:
if (
key_router_settings is not None
and isinstance(key_router_settings, dict)
and key_router_settings
):
return key_router_settings
# 2. Try team-level router_settings
if user_api_key_dict is not None and user_api_key_dict.team_id is not None:
try:
@ -3291,37 +3307,51 @@ class ProxyConfig:
where={"team_id": user_api_key_dict.team_id}
)
if team_obj is not None:
team_router_settings_value = getattr(team_obj, "router_settings", None)
team_router_settings_value = getattr(
team_obj, "router_settings", None
)
if team_router_settings_value is not None:
team_router_settings = None
if isinstance(team_router_settings_value, str):
try:
team_router_settings = yaml.safe_load(team_router_settings_value)
team_router_settings = yaml.safe_load(
team_router_settings_value
)
except (yaml.YAMLError, json.JSONDecodeError):
try:
team_router_settings = json.loads(team_router_settings_value)
team_router_settings = json.loads(
team_router_settings_value
)
except json.JSONDecodeError:
pass
elif isinstance(team_router_settings_value, dict):
team_router_settings = team_router_settings_value
# If team has router_settings (non-empty dict), use it
if team_router_settings is not None and isinstance(team_router_settings, dict) and team_router_settings:
if (
team_router_settings is not None
and isinstance(team_router_settings, dict)
and team_router_settings
):
return team_router_settings
except Exception:
# If team lookup fails, continue to global settings
pass
# 3. Try global router_settings
try:
db_router_settings = await prisma_client.db.litellm_config.find_first(
where={"param_name": "router_settings"}
)
if db_router_settings is not None and isinstance(db_router_settings.param_value, dict) and db_router_settings.param_value:
if (
db_router_settings is not None
and isinstance(db_router_settings.param_value, dict)
and db_router_settings.param_value
):
return db_router_settings.param_value
except Exception:
pass
return None
async def _add_router_settings_from_db_config(
@ -4688,27 +4718,48 @@ class ProxyStartupEvent:
### SPEND LOG CLEANUP ###
if general_settings.get("maximum_spend_logs_retention_period") is not None:
spend_log_cleanup = SpendLogCleanup()
# Get the interval from config or default to 1 day
retention_interval = general_settings.get(
"maximum_spend_logs_retention_interval", "1d"
)
try:
interval_seconds = duration_in_seconds(retention_interval)
scheduler.add_job(
spend_log_cleanup.cleanup_old_spend_logs,
"interval",
seconds=interval_seconds
+ random.randint(0, 60), # Add small random offset
# REMOVED jitter parameter - major cause of memory leak
args=[prisma_client],
id="spend_log_cleanup_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
except ValueError:
verbose_proxy_logger.error(
"Invalid maximum_spend_logs_retention_interval value"
cleanup_cron = general_settings.get("maximum_spend_logs_cleanup_cron")
if cleanup_cron:
from apscheduler.triggers.cron import CronTrigger
try:
cron_trigger = CronTrigger.from_crontab(cleanup_cron)
scheduler.add_job(
spend_log_cleanup.cleanup_old_spend_logs,
cron_trigger,
args=[prisma_client],
id="spend_log_cleanup_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
verbose_proxy_logger.info(
f"Spend log cleanup scheduled with cron: {cleanup_cron}"
)
except ValueError:
verbose_proxy_logger.error(
f"Invalid maximum_spend_logs_cleanup_cron value: {cleanup_cron}"
)
else:
# Interval-based scheduling (existing behavior)
retention_interval = general_settings.get(
"maximum_spend_logs_retention_interval", "1d"
)
try:
interval_seconds = duration_in_seconds(retention_interval)
scheduler.add_job(
spend_log_cleanup.cleanup_old_spend_logs,
"interval",
seconds=interval_seconds + random.randint(0, 60),
args=[prisma_client],
id="spend_log_cleanup_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
except ValueError:
verbose_proxy_logger.error(
"Invalid maximum_spend_logs_retention_interval value"
)
### CHECK BATCH COST ###
if llm_router is not None:
try:
@ -9922,7 +9973,9 @@ async def get_config(): # noqa: PLR0915
_success_callbacks = normalize_callback(_success_callbacks)
_failure_callbacks = normalize_callback(_failure_callbacks)
_success_and_failure_callbacks = normalize_callback(_success_and_failure_callbacks)
_success_and_failure_callbacks = normalize_callback(
_success_and_failure_callbacks
)
_data_to_return = []
"""
@ -10475,6 +10528,7 @@ app.include_router(model_access_group_management_router)
app.include_router(tag_management_router)
app.include_router(cost_tracking_settings_router)
app.include_router(router_settings_router)
app.include_router(fallback_management_router)
app.include_router(cache_settings_router)
app.include_router(user_agent_analytics_router)
app.include_router(enterprise_router)

View file

@ -256,7 +256,9 @@ async def video_status(
# Resolve model_name from model_id if available
# This allows the router to automatically inject litellm_params from the model config
if model_id_from_decoded and llm_router:
resolved_model = llm_router.resolve_model_name_from_model_id(model_id_from_decoded)
resolved_model = llm_router.resolve_model_name_from_model_id(
model_id_from_decoded, custom_llm_provider=provider_from_id
)
if resolved_model:
data["model"] = resolved_model
@ -354,7 +356,9 @@ async def video_content(
# Resolve model_name from model_id if available
# This allows the router to automatically inject litellm_params from the model config
if model_id_from_decoded and llm_router:
resolved_model = llm_router.resolve_model_name_from_model_id(model_id_from_decoded)
resolved_model = llm_router.resolve_model_name_from_model_id(
model_id_from_decoded, custom_llm_provider=provider_from_id
)
if resolved_model:
data["model"] = resolved_model
# Process request using ProxyBaseLLMRequestProcessing
@ -466,7 +470,9 @@ async def video_remix(
# Resolve model_name from model_id if available
# This allows the router to automatically inject litellm_params from the model config
if model_id_from_decoded and llm_router:
resolved_model = llm_router.resolve_model_name_from_model_id(model_id_from_decoded)
resolved_model = llm_router.resolve_model_name_from_model_id(
model_id_from_decoded, custom_llm_provider=provider_from_id
)
if resolved_model:
data["model"] = resolved_model

View file

@ -6971,7 +6971,7 @@ class Router:
return candidate_id in self.model_id_to_deployment_index_map
def resolve_model_name_from_model_id(
self, model_id: Optional[str]
self, model_id: Optional[str], custom_llm_provider: Optional[str] = None
) -> Optional[str]:
"""
Resolve model_name from model_id.
@ -6981,12 +6981,15 @@ class Router:
Strategy:
1. First, check if model_id directly matches a model_name or deployment ID
2. If not, search through router's model_list to find a match by litellm_params.model
3. Return the model_name if found, None otherwise
2. If custom_llm_provider is provided, check with provider prefix
3. Search through router's model_list to find a match by litellm_params.model
4. If custom_llm_provider is provided, try to find a wildcard pattern match
5. Return the model_name if found, None otherwise
Args:
model_id: The model_id extracted from decoded video_id
(could be model_name or litellm_params.model value)
custom_llm_provider: The provider name (e.g., "vertex_ai") for wildcard matching
Returns:
model_name if found, None otherwise. If None, the request will fall through
@ -6999,15 +7002,26 @@ class Router:
if model_id in self.model_names or self.has_model_id(model_id):
return model_id
# Strategy 2: Search through router's model_list to find by litellm_params.model
# Strategy 2: Check with provider prefix (e.g., "vertex_ai/veo-3.0-generate-preview")
if custom_llm_provider:
full_model_name = f"{custom_llm_provider}/{model_id}"
if full_model_name in self.model_names or self.has_model_id(full_model_name):
return full_model_name
# Strategy 3: Search through router's model_list to find by litellm_params.model
all_models = self.get_model_list(model_name=None)
if not all_models:
return None
# First pass: exact matches (non-wildcard)
for deployment in all_models:
litellm_params = deployment.get("litellm_params", {})
actual_model = litellm_params.get("model")
# Skip wildcard patterns in first pass
if actual_model and actual_model.endswith("/*"):
continue
# Match by exact match or by checking if actual_model ends with /model_id or :model_id
# e.g., model_id="veo-2.0-generate-001" matches actual_model="vertex_ai/veo-2.0-generate-001"
matches = (
@ -7021,6 +7035,19 @@ class Router:
if model_name:
return model_name
# Strategy 4: Wildcard patterns using PatternMatchRouter
# For video status/content, we need to match model_id like "veo-3.0-generate-preview"
# to wildcard patterns like "vertex_ai/*"
if custom_llm_provider:
full_model_name = f"{custom_llm_provider}/{model_id}"
pattern_deployments = self.pattern_router.route(full_model_name)
if pattern_deployments:
# Return the first matching wildcard model_name
for pattern_deployment in pattern_deployments:
matched_model_name = pattern_deployment.get("model_name")
if matched_model_name:
return matched_model_name
# No match found
return None
@ -8032,6 +8059,154 @@ class Router:
)
raise e
async def async_get_available_deployment_for_pass_through(
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, str]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
):
"""
Async version of get_available_deployment_for_pass_through
Only returns deployments configured with use_in_pass_through=True
"""
try:
parent_otel_span = _get_parent_otel_span_from_kwargs(request_kwargs)
# 1. Execute pre-routing hook
pre_routing_hook_response = await self.async_pre_routing_hook(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
)
if pre_routing_hook_response is not None:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
# 2. Get healthy deployments
healthy_deployments = await self.async_get_healthy_deployments(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
)
# 3. If specific deployment returned, verify if it supports pass-through
if isinstance(healthy_deployments, dict):
litellm_params = healthy_deployments.get("litellm_params", {})
if litellm_params.get("use_in_pass_through"):
return healthy_deployments
else:
raise litellm.BadRequestError(
message=f"Deployment {healthy_deployments.get('model_info', {}).get('id')} does not support pass-through endpoint (use_in_pass_through=False)",
model=model,
llm_provider="",
)
# 4. Filter deployments that support pass-through
pass_through_deployments = self._filter_pass_through_deployments(
healthy_deployments=healthy_deployments
)
if len(pass_through_deployments) == 0:
raise litellm.BadRequestError(
message=f"Model {model} has no deployments configured with use_in_pass_through=True. Please add use_in_pass_through: true to the deployment configuration",
model=model,
llm_provider="",
)
# 5. Apply load balancing strategy
start_time = time.perf_counter()
if (
self.routing_strategy == "usage-based-routing-v2"
and self.lowesttpm_logger_v2 is not None
):
deployment = (
await self.lowesttpm_logger_v2.async_get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
)
)
elif (
self.routing_strategy == "latency-based-routing"
and self.lowestlatency_logger is not None
):
deployment = (
await self.lowestlatency_logger.async_get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
)
elif self.routing_strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
healthy_deployments=pass_through_deployments,
model=model,
)
elif (
self.routing_strategy == "least-busy"
and self.leastbusy_logger is not None
):
deployment = (
await self.leastbusy_logger.async_get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
)
)
else:
deployment = None
if deployment is None:
exception = await async_raise_no_deployment_exception(
litellm_router_instance=self,
model=model,
parent_otel_span=parent_otel_span,
)
raise exception
verbose_router_logger.info(
f"async_get_available_deployment_for_pass_through model: {model}, selected deployment: {self.print_deployment(deployment)}"
)
end_time = time.perf_counter()
_duration = end_time - start_time
asyncio.create_task(
self.service_logger_obj.async_service_success_hook(
service=ServiceTypes.ROUTER,
duration=_duration,
call_type="<routing_strategy>.async_get_available_deployments",
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
)
)
return deployment
except Exception as e:
traceback_exception = traceback.format_exc()
if request_kwargs is not None:
logging_obj = request_kwargs.get("litellm_logging_obj", None)
if logging_obj is not None:
threading.Thread(
target=logging_obj.failure_handler,
args=(e, traceback_exception),
).start()
asyncio.create_task(
logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
)
raise e
async def async_pre_routing_hook(
self,
model: str,
@ -8184,6 +8359,169 @@ class Router:
)
return deployment
def get_available_deployment_for_pass_through(
self,
model: str,
messages: Optional[List[Dict[str, str]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
request_kwargs: Optional[Dict] = None,
):
"""
Returns deployments available for pass-through endpoints (based on load balancing strategy)
Similar to get_available_deployment, but only returns deployments with use_in_pass_through=True
Args:
model: Model name
messages: Optional list of messages
input: Optional input data
specific_deployment: Whether to find a specific deployment
request_kwargs: Optional request parameters
Returns:
Dict: Selected deployment configuration
Raises:
BadRequestError: If no deployment is configured with use_in_pass_through=True
RouterRateLimitError: If no pass-through deployments are available
"""
# 1. Perform common checks to get healthy deployments list
model, healthy_deployments = self._common_checks_available_deployment(
model=model,
messages=messages,
input=input,
specific_deployment=specific_deployment,
)
# 2. If the returned is a specific deployment (Dict), verify and return directly
if isinstance(healthy_deployments, dict):
litellm_params = healthy_deployments.get("litellm_params", {})
if litellm_params.get("use_in_pass_through"):
return healthy_deployments
else:
# Specific deployment does not support pass-through
raise litellm.BadRequestError(
message=f"Deployment {healthy_deployments.get('model_info', {}).get('id')} does not support pass-through endpoint (use_in_pass_through=False)",
model=model,
llm_provider="",
)
# 3. Filter deployments that support pass-through
pass_through_deployments = self._filter_pass_through_deployments(
healthy_deployments=healthy_deployments
)
if len(pass_through_deployments) == 0:
# No deployments support pass-through
raise litellm.BadRequestError(
message=f"Model {model} has no deployment configured with use_in_pass_through=True. Please add use_in_pass_through: true in the deployment configuration",
model=model,
llm_provider="",
)
# 4. Apply cooldown filtering
parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(
request_kwargs
)
cooldown_deployments = _get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
)
pass_through_deployments = self._filter_cooldown_deployments(
healthy_deployments=pass_through_deployments,
cooldown_deployments=cooldown_deployments,
)
# 5. Apply pre-call checks (if enabled)
if self.enable_pre_call_checks and messages is not None:
pass_through_deployments = self._pre_call_checks(
model=model,
healthy_deployments=pass_through_deployments,
messages=messages,
request_kwargs=request_kwargs,
)
if len(pass_through_deployments) == 0:
model_ids = self.get_model_ids(model_name=model)
_cooldown_time = self.cooldown_cache.get_min_cooldown(
model_ids=model_ids, parent_otel_span=parent_otel_span
)
_cooldown_list = _get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
)
raise RouterRateLimitError(
model=model,
cooldown_time=_cooldown_time,
enable_pre_call_checks=self.enable_pre_call_checks,
cooldown_list=_cooldown_list,
)
# 6. Apply load balancing strategy
if self.routing_strategy == "least-busy" and self.leastbusy_logger is not None:
deployment = self.leastbusy_logger.get_available_deployments(
model_group=model, healthy_deployments=pass_through_deployments # type: ignore
)
elif self.routing_strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
healthy_deployments=pass_through_deployments,
model=model,
)
elif (
self.routing_strategy == "latency-based-routing"
and self.lowestlatency_logger is not None
):
deployment = self.lowestlatency_logger.get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
request_kwargs=request_kwargs,
)
elif (
self.routing_strategy == "usage-based-routing"
and self.lowesttpm_logger is not None
):
deployment = self.lowesttpm_logger.get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
)
elif (
self.routing_strategy == "usage-based-routing-v2"
and self.lowesttpm_logger_v2 is not None
):
deployment = self.lowesttpm_logger_v2.get_available_deployments(
model_group=model,
healthy_deployments=pass_through_deployments, # type: ignore
messages=messages,
input=input,
)
else:
deployment = None
if deployment is None:
verbose_router_logger.info(
f"get_available_deployment_for_pass_through model: {model}, no available deployments"
)
model_ids = self.get_model_ids(model_name=model)
_cooldown_time = self.cooldown_cache.get_min_cooldown(
model_ids=model_ids, parent_otel_span=parent_otel_span
)
_cooldown_list = _get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
)
raise RouterRateLimitError(
model=model,
cooldown_time=_cooldown_time,
enable_pre_call_checks=self.enable_pre_call_checks,
cooldown_list=_cooldown_list,
)
verbose_router_logger.info(
f"get_available_deployment_for_pass_through model: {model}, selected deployment: {self.print_deployment(deployment)}"
)
return deployment
def _filter_cooldown_deployments(
self, healthy_deployments: List[Dict], cooldown_deployments: List[str]
) -> List[Dict]:
@ -8206,6 +8544,34 @@ class Router:
if deployment["model_info"]["id"] not in cooldown_set
]
def _filter_pass_through_deployments(
self, healthy_deployments: List[Dict]
) -> List[Dict]:
"""
Filter out deployments configured with use_in_pass_through=True
Args:
healthy_deployments: List of healthy deployments
Returns:
List[Dict]: Only includes a list of deployments that support pass-through
"""
verbose_router_logger.debug(
f"Filter pass-through deployments from {len(healthy_deployments)} healthy deployments"
)
pass_through_deployments = [
deployment
for deployment in healthy_deployments
if deployment.get("litellm_params", {}).get("use_in_pass_through", False)
]
verbose_router_logger.debug(
f"Found {len(pass_through_deployments)} deployments with pass-through enabled"
)
return pass_through_deployments
def _track_deployment_metrics(
self, deployment, parent_otel_span: Optional[Span], response=None
):

View file

@ -636,8 +636,10 @@ class ANTHROPIC_BETA_HEADER_VALUES(str, Enum):
ADVANCED_TOOL_USE_2025_11_20 = "advanced-tool-use-2025-11-20"
# Tool search beta header constant
# Tool search beta header constant (for Anthropic direct API and Microsoft Foundry)
ANTHROPIC_TOOL_SEARCH_BETA_HEADER = "advanced-tool-use-2025-11-20"
# Effort beta header constant
ANTHROPIC_EFFORT_BETA_HEADER = "effort-2025-11-24"

View file

@ -0,0 +1,36 @@
"""
Tool Search Beta Header Configuration
Reference: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool
"""
from typing import Dict
from litellm.types.utils import LlmProviders
# Tool search beta header values
TOOL_SEARCH_BETA_HEADER_ANTHROPIC = "advanced-tool-use-2025-11-20"
TOOL_SEARCH_BETA_HEADER_VERTEX = "tool-search-tool-2025-10-19"
TOOL_SEARCH_BETA_HEADER_BEDROCK = "tool-search-tool-2025-10-19"
# Mapping of custom_llm_provider -> tool search beta header
TOOL_SEARCH_BETA_HEADER_BY_PROVIDER: Dict[str, str] = {
LlmProviders.ANTHROPIC.value: TOOL_SEARCH_BETA_HEADER_ANTHROPIC,
LlmProviders.AZURE.value: TOOL_SEARCH_BETA_HEADER_ANTHROPIC,
LlmProviders.AZURE_AI.value: TOOL_SEARCH_BETA_HEADER_ANTHROPIC,
LlmProviders.VERTEX_AI.value: TOOL_SEARCH_BETA_HEADER_VERTEX,
LlmProviders.VERTEX_AI_BETA.value: TOOL_SEARCH_BETA_HEADER_VERTEX,
LlmProviders.BEDROCK.value: TOOL_SEARCH_BETA_HEADER_BEDROCK,
}
def get_tool_search_beta_header(custom_llm_provider: str) -> str:
"""
Get the tool search beta header for a given provider.
"""
return TOOL_SEARCH_BETA_HEADER_BY_PROVIDER.get(
custom_llm_provider,
TOOL_SEARCH_BETA_HEADER_ANTHROPIC
)

View file

@ -62,7 +62,7 @@ class ToolResultBlock(TypedDict, total=False):
class ToolUseBlock(TypedDict):
input: dict
input: Any # Per boto3 spec: document type can be dict, list, int, float, str, bool, or None
name: str
toolUseId: str

View file

@ -2,9 +2,70 @@
Types and field definitions for router settings management endpoints
"""
from typing import Any, Dict, List, Optional
from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel
from pydantic import BaseModel, Field, field_validator
# Fallback Management Types
class FallbackCreateRequest(BaseModel):
"""Request model for creating/updating fallbacks"""
model: str = Field(
description="The model name to configure fallbacks for (e.g., 'gpt-3.5-turbo')"
)
fallback_models: List[str] = Field(
description="List of fallback model names in order of priority",
min_length=1,
)
fallback_type: Literal["general", "context_window", "content_policy"] = Field(
default="general",
description="Type of fallback: 'general' (default), 'context_window', or 'content_policy'",
)
@field_validator("fallback_models")
@classmethod
def validate_fallback_models(cls, v: List[str]) -> List[str]:
if not v:
raise ValueError("fallback_models must contain at least one model")
if len(v) != len(set(v)):
raise ValueError("fallback_models must not contain duplicates")
return v
@field_validator("model")
@classmethod
def validate_model(cls, v: str) -> str:
if not v or not v.strip():
raise ValueError("model must be a non-empty string")
return v.strip()
class FallbackResponse(BaseModel):
"""Response model for fallback operations"""
model: str = Field(description="The model name")
fallback_models: List[str] = Field(description="List of fallback model names")
fallback_type: str = Field(description="Type of fallback")
message: str = Field(description="Success message")
class FallbackGetResponse(BaseModel):
"""Response model for getting fallbacks"""
model: str = Field(description="The model name")
fallback_models: List[str] = Field(description="List of fallback model names")
fallback_type: str = Field(description="Type of fallback")
class FallbackDeleteResponse(BaseModel):
"""Response model for deleting fallbacks"""
model: str = Field(description="The model name")
fallback_type: str = Field(description="Type of fallback")
message: str = Field(description="Success message")
# Router Settings Types
class RouterSettingsField(BaseModel):

View file

@ -3634,6 +3634,37 @@
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5.2-codex": {
"cache_read_input_token_cost": 1.75e-07,
"input_cost_per_token": 1.75e-06,
"litellm_provider": "azure",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 1.4e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5.2-pro": {
"input_cost_per_token": 2.1e-05,
"litellm_provider": "azure",
@ -10170,6 +10201,48 @@
"mode": "completion",
"output_cost_per_token": 5e-07
},
"deepseek-v3-2-251201": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 98304,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"glm-4-7-251222": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 204800,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"kimi-k2-thinking-251104": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 229376,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"doubao-embedding": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
@ -25526,13 +25599,13 @@
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
"output_cost_per_image": 0.4
"output_cost_per_image": 0.40
},
"stability.stable-creative-upscale-v1:0": {
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"mode": "image_edit",
"output_cost_per_image": 0.6
"output_cost_per_image": 0.60
},
"stability.stable-fast-upscale-v1:0": {
"litellm_provider": "bedrock",
@ -28782,13 +28855,13 @@
"supports_web_search": true
},
"vertex_ai/zai-org/glm-4.7-maas": {
"input_cost_per_token": 3e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "vertex_ai-zai_models",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"output_cost_per_token": 2.2e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supports_function_calling": true,
"supports_reasoning": true,
@ -33930,4 +34003,4 @@
"litellm_provider": "llamagate",
"mode": "embedding"
}
}
}

38
poetry.lock generated
View file

@ -1,4 +1,4 @@
# This file is automatically @generated by Poetry 2.2.0 and should not be changed by hand.
# This file is automatically @generated by Poetry 2.2.1 and should not be changed by hand.
[[package]]
name = "aiofiles"
@ -525,36 +525,36 @@ files = [
[[package]]
name = "boto3"
version = "1.36.0"
version = "1.40.61"
description = "The AWS SDK for Python"
optional = true
python-versions = ">=3.8"
python-versions = ">=3.9"
groups = ["main"]
markers = "extra == \"proxy\""
files = [
{file = "boto3-1.36.0-py3-none-any.whl", hash = "sha256:d0ca7a58ce25701a52232cc8df9d87854824f1f2964b929305722ebc7959d5a9"},
{file = "boto3-1.36.0.tar.gz", hash = "sha256:159898f51c2997a12541c0e02d6e5a8fe2993ddb307b9478fd9a339f98b57e00"},
{file = "boto3-1.40.61-py3-none-any.whl", hash = "sha256:6b9c57b2a922b5d8c17766e29ed792586a818098efe84def27c8f582b33f898c"},
{file = "boto3-1.40.61.tar.gz", hash = "sha256:d6c56277251adf6c2bdd25249feae625abe4966831676689ff23b4694dea5b12"},
]
[package.dependencies]
botocore = ">=1.36.0,<1.37.0"
botocore = ">=1.40.61,<1.41.0"
jmespath = ">=0.7.1,<2.0.0"
s3transfer = ">=0.11.0,<0.12.0"
s3transfer = ">=0.14.0,<0.15.0"
[package.extras]
crt = ["botocore[crt] (>=1.21.0,<2.0a0)"]
[[package]]
name = "botocore"
version = "1.36.26"
version = "1.40.76"
description = "Low-level, data-driven core of boto 3."
optional = true
python-versions = ">=3.8"
python-versions = ">=3.9"
groups = ["main"]
markers = "extra == \"proxy\""
files = [
{file = "botocore-1.36.26-py3-none-any.whl", hash = "sha256:4e3f19913887a58502e71ef8d696fe7eaa54de7813ff73390cd5883f837dfa6e"},
{file = "botocore-1.36.26.tar.gz", hash = "sha256:4a63bcef7ecf6146fd3a61dc4f9b33b7473b49bdaf1770e9aaca6eee0c9eab62"},
{file = "botocore-1.40.76-py3-none-any.whl", hash = "sha256:fe425d386e48ac64c81cbb4a7181688d813df2e2b4c78b95ebe833c9e868c6f4"},
{file = "botocore-1.40.76.tar.gz", hash = "sha256:2b16024d68b29b973005adfb5039adfe9099ebe772d40a90ca89f2e165c495dc"},
]
[package.dependencies]
@ -566,7 +566,7 @@ urllib3 = [
]
[package.extras]
crt = ["awscrt (==0.23.8)"]
crt = ["awscrt (==0.28.4)"]
[[package]]
name = "cachetools"
@ -6255,22 +6255,22 @@ files = [
[[package]]
name = "s3transfer"
version = "0.11.3"
version = "0.14.0"
description = "An Amazon S3 Transfer Manager"
optional = true
python-versions = ">=3.8"
python-versions = ">=3.9"
groups = ["main"]
markers = "extra == \"proxy\""
files = [
{file = "s3transfer-0.11.3-py3-none-any.whl", hash = "sha256:ca855bdeb885174b5ffa95b9913622459d4ad8e331fc98eb01e6d5eb6a30655d"},
{file = "s3transfer-0.11.3.tar.gz", hash = "sha256:edae4977e3a122445660c7c114bba949f9d191bae3b34a096f18a1c8c354527a"},
{file = "s3transfer-0.14.0-py3-none-any.whl", hash = "sha256:ea3b790c7077558ed1f02a3072fb3cb992bbbd253392f4b6e9e8976941c7d456"},
{file = "s3transfer-0.14.0.tar.gz", hash = "sha256:eff12264e7c8b4985074ccce27a3b38a485bb7f7422cc8046fee9be4983e4125"},
]
[package.dependencies]
botocore = ">=1.36.0,<2.0a.0"
botocore = ">=1.37.4,<2.0a.0"
[package.extras]
crt = ["botocore[crt] (>=1.36.0,<2.0a.0)"]
crt = ["botocore[crt] (>=1.37.4,<2.0a.0)"]
[[package]]
name = "scikit-learn"
@ -7981,4 +7981,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "ea62b77c662ab9fc486e421c576f0868bcde16d62a24703ee1f4916a0465ffb2"
content-hash = "f391c702cf58ef2ba7641acdc3ae13d7c8e672faede68c0a624bd2ba0fb46b12"

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.80.16"
version = "1.80.17"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -56,7 +56,7 @@ google-cloud-iam = {version = "^2.19.1", optional = true}
resend = {version = ">=0.8.0", optional = true}
pynacl = {version = "^1.5.0", optional = true}
websockets = {version = "^15.0.1", optional = true}
boto3 = {version = "1.36.0", optional = true}
boto3 = {version = "1.40.61", optional = true}
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
mcp = {version = "^1.21.2", optional = true, python = ">=3.10"}
litellm-proxy-extras = {version = "0.4.21", optional = true}
@ -167,7 +167,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.80.16"
version = "1.80.17"
version_files = [
"pyproject.toml:^version"
]

View file

@ -10,7 +10,7 @@ uvicorn==0.31.1 # server dep
gunicorn==23.0.0 # server dep
fastuuid==0.13.5 # for uuid4
uvloop==0.21.0 # uvicorn dep, gives us much better performance under load
boto3==1.36.0 # aws bedrock/sagemaker calls
boto3==1.40.61 # aws bedrock/sagemaker calls
redis==5.2.1 # redis caching
prisma==0.11.0 # for db
nodejs-wheel-binaries==24.12.0 ## required by prisma for migrations, prevents runtime download (updated from nodejs-bin for security fixes)
@ -33,6 +33,7 @@ fastapi-sso==0.19.0 # admin UI, SSO
pyjwt[crypto]==2.10.1 ; python_version >= "3.9"
python-multipart==0.0.18 # admin UI
Pillow==11.0.0
jaraco.context>=6.1.0
azure-ai-contentsafety==1.0.0 # for azure content safety
azure-identity==1.16.1 ; python_version >= "3.9" # for azure content safety
azure-keyvault==4.2.0 # for azure KMS integration
@ -62,7 +63,7 @@ click==8.1.7 # for proxy cli
rich==13.7.1 # for litellm proxy cli
jinja2==3.1.6 # for prompt templates
aiohttp==3.13.3 # for network calls
aioboto3==13.4.0 # for async sagemaker calls
aioboto3==15.5.0 # for async sagemaker calls
tenacity==8.5.0 # for retrying requests, when litellm.num_retries set
pydantic>=2.11,<3 # proxy + openai req. + mcp
jsonschema>=4.23.0,<5.0.0 # validating json schema - aligned with openapi-core + mcp

View file

@ -139,4 +139,4 @@ fastuuid: >=0.13.0 # BSD-3-Clause license
llm-sandbox: >=0.3.31 # MIT License - https://github.com/vndee/llm-sandbox
nodejs-wheel-binaries: >=24.12.0 # MIT license manually verified
grpcio: >=1.69.0 # Apache License 2.0
jaraco.context: >=6.1.0 # Unknown license

File diff suppressed because one or more lines are too long

View file

@ -12,6 +12,7 @@ from litellm.llms.bedrock.common_utils import (
get_bedrock_base_model,
get_bedrock_cross_region_inference_regions,
strip_bedrock_routing_prefix,
strip_bedrock_throughput_suffix,
)
from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter
@ -46,6 +47,21 @@ class TestStripBedrockRoutingPrefix:
)
class TestStripBedrockThroughputSuffix:
"""Tests for strip_bedrock_throughput_suffix function."""
@pytest.mark.parametrize("input_model,expected", [
("anthropic.claude-3-5-sonnet-20241022-v2:0:51k", "anthropic.claude-3-5-sonnet-20241022-v2:0"),
("anthropic.claude-3-5-sonnet-20241022-v2:0:18k", "anthropic.claude-3-5-sonnet-20241022-v2:0"),
("model:1:51k", "model:1"),
("model:123:18k", "model:123"),
("anthropic.claude-3-5-sonnet-20241022-v2:0", "anthropic.claude-3-5-sonnet-20241022-v2:0"),
("anthropic.claude-3-sonnet", "anthropic.claude-3-sonnet"),
])
def test_strip_throughput_suffix(self, input_model, expected):
assert strip_bedrock_throughput_suffix(input_model) == expected
class TestExtractModelNameFromBedrockArn:
"""Tests for extract_model_name_from_bedrock_arn function."""
@ -118,6 +134,16 @@ class TestGetBedrockBaseModel:
== "anthropic.claude-3-sonnet-20240229-v1:0"
)
@pytest.mark.parametrize("input_model,expected", [
("anthropic.claude-3-5-sonnet-20241022-v2:0:51k", "anthropic.claude-3-5-sonnet-20241022-v2:0"),
("anthropic.claude-3-5-sonnet-20241022-v2:0:18k", "anthropic.claude-3-5-sonnet-20241022-v2:0"),
("bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0:51k", "anthropic.claude-3-5-sonnet-20241022-v2:0"),
("us.anthropic.claude-3-5-sonnet-20241022-v2:0:51k", "anthropic.claude-3-5-sonnet-20241022-v2:0"),
])
def test_strips_throughput_suffix(self, input_model, expected):
"""Test that throughput tier suffixes like :51k are stripped. Issue #19113."""
assert get_bedrock_base_model(input_model) == expected
class TestBedrockModelInfoWrappers:
"""Tests that BedrockModelInfo methods correctly wrap standalone functions."""

View file

@ -3954,3 +3954,157 @@ def test_bedrock_openai_error_handling():
assert exc_info.value.status_code == 422
print("✓ Error handling works correctly")
def test_bedrock_malformed_tool_json_handling():
"""
Test that Bedrock handles malformed JSON in tool call arguments gracefully.
This test covers the issue where:
1. LLM generates malformed JSON in tool call arguments
2. Subsequent requests with conversation history should not crash
3. The toolUse.input field should handle any JSON value type per boto3 spec
Related issue: https://github.com/BerriAI/litellm/issues/[issue_number]
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
_convert_to_bedrock_tool_call_invoke,
)
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
from litellm.types.llms.bedrock import ContentBlock
# Test 1: Malformed JSON in tool call arguments
malformed_tool_calls = [
{
"id": "call_123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "Paris", "invalid_json', # Malformed JSON
},
}
]
# Should not raise an exception, but store as raw string
result = _convert_to_bedrock_tool_call_invoke(malformed_tool_calls)
assert len(result) == 1
assert result[0]["toolUse"]["name"] == "get_weather"
# The malformed JSON should be stored as a string
assert isinstance(result[0]["toolUse"]["input"], str)
assert result[0]["toolUse"]["input"] == '{"location": "Paris", "invalid_json'
print("✓ Malformed JSON stored as raw string")
# Test 2: Valid JSON should still work normally
valid_tool_calls = [
{
"id": "call_456",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "London"}',
},
}
]
result = _convert_to_bedrock_tool_call_invoke(valid_tool_calls)
assert len(result) == 1
assert result[0]["toolUse"]["name"] == "get_weather"
assert isinstance(result[0]["toolUse"]["input"], dict)
assert result[0]["toolUse"]["input"] == {"location": "London"}
print("✓ Valid JSON parsed correctly")
# Test 3: Empty arguments should create empty dict
empty_tool_calls = [
{
"id": "call_789",
"type": "function",
"function": {
"name": "no_args_function",
"arguments": "",
},
}
]
result = _convert_to_bedrock_tool_call_invoke(empty_tool_calls)
assert len(result) == 1
assert result[0]["toolUse"]["input"] == {}
print("✓ Empty arguments handled correctly")
# Test 4: Bedrock to OpenAI conversion handles string input
converse_config = AmazonConverseConfig()
content_blocks = [
ContentBlock(
toolUse={
"name": "get_weather",
"toolUseId": "call_123",
"input": '{"location": "Paris", "invalid_json', # String input (malformed)
}
)
]
content_str, tools, reasoning = converse_config._translate_message_content(
content_blocks
)
assert len(tools) == 1
assert tools[0]["function"]["name"] == "get_weather"
# Should return the string as-is
assert tools[0]["function"]["arguments"] == '{"location": "Paris", "invalid_json'
print("✓ Bedrock to OpenAI conversion handles string input")
# Test 5: Bedrock to OpenAI conversion handles dict input
content_blocks_dict = [
ContentBlock(
toolUse={
"name": "get_weather",
"toolUseId": "call_456",
"input": {"location": "London"}, # Dict input (normal case)
}
)
]
content_str, tools, reasoning = converse_config._translate_message_content(
content_blocks_dict
)
assert len(tools) == 1
assert tools[0]["function"]["name"] == "get_weather"
# Should serialize dict to JSON string
assert tools[0]["function"]["arguments"] == '{"location": "London"}'
print("✓ Bedrock to OpenAI conversion handles dict input")
# Test 6: Round-trip conversion with malformed JSON
# Test that we can convert OpenAI -> Bedrock -> OpenAI with malformed JSON
malformed_tool_calls_roundtrip = [
{
"id": "call_999",
"type": "function",
"function": {
"name": "test_function",
"arguments": '{"key": "value", "broken', # Malformed
},
}
]
# Step 1: OpenAI to Bedrock (should store as string)
bedrock_blocks = _convert_to_bedrock_tool_call_invoke(malformed_tool_calls_roundtrip)
assert isinstance(bedrock_blocks[0]["toolUse"]["input"], str)
# Step 2: Bedrock back to OpenAI (should preserve the string)
content_blocks_roundtrip = [
ContentBlock(
toolUse={
"name": bedrock_blocks[0]["toolUse"]["name"],
"toolUseId": bedrock_blocks[0]["toolUse"]["toolUseId"],
"input": bedrock_blocks[0]["toolUse"]["input"],
}
)
]
content_str, tools_roundtrip, reasoning = converse_config._translate_message_content(
content_blocks_roundtrip
)
# Should preserve the malformed JSON string through the round trip
assert tools_roundtrip[0]["function"]["arguments"] == '{"key": "value", "broken'
print("✓ Round-trip conversion preserves malformed JSON")
print("✓ All malformed JSON handling tests passed")

View file

@ -592,3 +592,205 @@ async def test_weighted_selection_router_async(rpm_list, tpm_list):
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")
def test_get_available_deployment_for_pass_through():
"""
Test get_available_deployment_for_pass_through function
- Tests that only deployments with use_in_pass_through=True are returned
- Tests that BadRequestError is raised when no pass-through deployments exist
"""
try:
litellm.set_verbose = False
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
"use_in_pass_through": True,
},
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION"),
"use_in_pass_through": False,
},
},
]
router = Router(
model_list=model_list,
)
# Test that only pass-through deployment is returned
selected_model = router.get_available_deployment_for_pass_through(
"gpt-3.5-turbo"
)
assert selected_model["litellm_params"]["model"] == "gpt-3.5-turbo"
assert selected_model["litellm_params"]["use_in_pass_through"] is True
router.reset()
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")
def test_get_available_deployment_for_pass_through_no_deployments():
"""
Test get_available_deployment_for_pass_through raises BadRequestError
when no deployments have use_in_pass_through=True
"""
try:
litellm.set_verbose = False
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
"use_in_pass_through": False,
},
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION"),
"use_in_pass_through": False,
},
},
]
router = Router(
model_list=model_list,
)
# Test that BadRequestError is raised when no pass-through deployments exist
try:
router.get_available_deployment_for_pass_through("gpt-3.5-turbo")
pytest.fail(
"Expected BadRequestError when no pass-through deployments exist"
)
except litellm.BadRequestError as e:
assert "use_in_pass_through=True" in str(e)
router.reset()
except Exception as e:
if isinstance(e, litellm.BadRequestError):
pass # Expected error
else:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")
@pytest.mark.asyncio
async def test_async_get_available_deployment_for_pass_through():
"""
Test async_get_available_deployment_for_pass_through function
- Tests that only deployments with use_in_pass_through=True are returned
- Tests async version works correctly
"""
try:
litellm.set_verbose = False
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
"use_in_pass_through": True,
},
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION"),
"use_in_pass_through": False,
},
},
]
router = Router(
model_list=model_list,
)
# Test that only pass-through deployment is returned
selected_model = await router.async_get_available_deployment_for_pass_through(
model="gpt-3.5-turbo", request_kwargs={}
)
assert selected_model["litellm_params"]["model"] == "gpt-3.5-turbo"
assert selected_model["litellm_params"]["use_in_pass_through"] is True
router.reset()
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")
def test_filter_pass_through_deployments():
"""
Test _filter_pass_through_deployments function
- Tests that it correctly filters deployments with use_in_pass_through=True
"""
try:
litellm.set_verbose = False
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
"use_in_pass_through": True,
},
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION"),
"use_in_pass_through": False,
},
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-35-turbo",
"api_key": os.getenv("AZURE_API_KEY"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_version": os.getenv("AZURE_API_VERSION"),
"use_in_pass_through": True,
},
},
]
router = Router(
model_list=model_list,
)
# Get all healthy deployments
healthy_deployments = router.get_model_list()
# Filter pass-through deployments
pass_through_deployments = router._filter_pass_through_deployments(
healthy_deployments
)
# Should only have 2 deployments with use_in_pass_through=True
assert len(pass_through_deployments) == 2
# Verify all returned deployments have use_in_pass_through=True
for deployment in pass_through_deployments:
assert deployment["litellm_params"]["use_in_pass_through"] is True
router.reset()
except Exception as e:
traceback.print_exc()
pytest.fail(f"Error occurred: {e}")

View file

@ -703,3 +703,188 @@ def test_cost_breakdown_missing_in_standard_logging_payload():
assert payload["response_cost"] == 0.0001
print("✅ Cost breakdown missing test passed!")
def test_merge_litellm_metadata_basic():
"""
Test that merge_litellm_metadata correctly merges metadata and litellm_metadata.
User API key fields (from metadata) should take precedence over model-related fields (from litellm_metadata).
"""
litellm_params = {
"metadata": {
"user_api_key": "test-key-123",
"user_api_key_user_id": "user-456",
"user_api_key_team_id": "team-789",
},
"litellm_metadata": {
"model_group": "gpt-4-group",
"model_info": {"id": "model-123"},
"tags": ["tag1", "tag2"],
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# Check that user API key fields are present
assert result["user_api_key"] == "test-key-123"
assert result["user_api_key_user_id"] == "user-456"
assert result["user_api_key_team_id"] == "team-789"
# Check that model-related fields are present
assert result["model_group"] == "gpt-4-group"
assert result["model_info"] == {"id": "model-123"}
assert result["tags"] == ["tag1", "tag2"]
def test_merge_litellm_metadata_precedence():
"""
Test that metadata fields take precedence over litellm_metadata when there are conflicts.
"""
litellm_params = {
"metadata": {
"tags": ["user-tag1", "user-tag2"],
"custom_field": "from_metadata",
},
"litellm_metadata": {
"tags": ["model-tag1", "model-tag2"], # This should NOT overwrite
"custom_field": "from_litellm_metadata", # This should NOT overwrite
"model_group": "gpt-4-group", # This should be included
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# metadata values should take precedence
assert result["tags"] == ["user-tag1", "user-tag2"]
assert result["custom_field"] == "from_metadata"
# litellm_metadata values should only be included if not in metadata
assert result["model_group"] == "gpt-4-group"
def test_merge_litellm_metadata_skip_non_serializable():
"""
Test that non-serializable objects like UserAPIKeyAuth are skipped.
"""
from litellm.proxy._types import UserAPIKeyAuth
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="test-team",
)
litellm_params = {
"metadata": {
"user_api_key": "test-key-123",
"user_api_key_auth": user_api_key_auth, # This should be skipped
"safe_field": "safe_value",
},
"litellm_metadata": {
"model_group": "gpt-4-group",
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# user_api_key_auth should be skipped
assert "user_api_key_auth" not in result
# Other fields should be present
assert result["user_api_key"] == "test-key-123"
assert result["safe_field"] == "safe_value"
assert result["model_group"] == "gpt-4-group"
def test_merge_litellm_metadata_empty_params():
"""
Test that merge_litellm_metadata handles empty or missing metadata gracefully.
"""
# Test with empty litellm_params
result = StandardLoggingPayloadSetup.merge_litellm_metadata({})
assert result == {}
# Test with only metadata
litellm_params = {
"metadata": {
"user_api_key": "test-key",
}
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
assert result == {"user_api_key": "test-key"}
# Test with only litellm_metadata
litellm_params = {
"litellm_metadata": {
"model_group": "gpt-4-group",
}
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
assert result == {"model_group": "gpt-4-group"}
# Test with None values
litellm_params = {
"metadata": None,
"litellm_metadata": None,
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
assert result == {}
def test_merge_litellm_metadata_bedrock_passthrough_scenario():
"""
Test merge_litellm_metadata in a Bedrock passthrough scenario where both
user API key metadata and model metadata need to be merged.
This is the specific scenario that was fixed - bedrock passthrough requests
should include complete user authentication metadata in logging.
"""
litellm_params = {
"metadata": {
# User API key fields from authentication
"user_api_key": "sk-bedrock-test-key-123",
"user_api_key_hash": "hashed-key-123",
"user_api_key_user_id": "bedrock-user-456",
"user_api_key_team_id": "bedrock-team-789",
"user_api_key_org_id": "bedrock-org-101",
"user_api_key_alias": "bedrock-key-alias",
"user_api_key_team_alias": "bedrock-team-alias",
"user_api_key_end_user_id": "end-user-123",
"user_api_key_request_route": "/bedrock/model/invoke",
},
"litellm_metadata": {
# Model-related fields from Bedrock configuration
"model_group": "bedrock-claude-group",
"model_info": {
"id": "anthropic.claude-3-sonnet",
"mode": "chat",
},
"aws_region_name": "us-east-1",
"tags": ["production", "bedrock"],
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# Verify all user API key fields are present
assert result["user_api_key"] == "sk-bedrock-test-key-123"
assert result["user_api_key_hash"] == "hashed-key-123"
assert result["user_api_key_user_id"] == "bedrock-user-456"
assert result["user_api_key_team_id"] == "bedrock-team-789"
assert result["user_api_key_org_id"] == "bedrock-org-101"
assert result["user_api_key_alias"] == "bedrock-key-alias"
assert result["user_api_key_team_alias"] == "bedrock-team-alias"
assert result["user_api_key_end_user_id"] == "end-user-123"
assert result["user_api_key_request_route"] == "/bedrock/model/invoke"
# Verify all model-related fields are present
assert result["model_group"] == "bedrock-claude-group"
assert result["model_info"] == {
"id": "anthropic.claude-3-sonnet",
"mode": "chat",
}
assert result["aws_region_name"] == "us-east-1"
assert result["tags"] == ["production", "bedrock"]
# Verify total number of fields (9 user fields + 4 model fields = 13)
assert len(result) == 13

View file

@ -0,0 +1,294 @@
"""
Base test class for Anthropic Messages API tool search E2E tests.
Tests that tool search works correctly via litellm.anthropic.messages interface
by making actual API calls and validating that tool search discovers deferred tools.
Reference: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool
"""
import json
import os
import sys
from abc import ABC, abstractmethod
from typing import Any, Dict, List
sys.path.insert(0, os.path.abspath("../../.."))
import pytest
import litellm
# Sample tools for tool search testing
def get_deferred_tools() -> List[Dict[str, Any]]:
"""
Returns a list of tools with defer_loading: true.
These tools should only be discovered via tool search.
"""
return [
{
"name": "get_weather",
"description": "Get the current weather for a location",
"input_schema": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA"
}
},
"required": ["location"]
},
"defer_loading": True
},
{
"name": "get_stock_price",
"description": "Get the current stock price for a ticker symbol",
"input_schema": {
"type": "object",
"properties": {
"ticker": {
"type": "string",
"description": "The stock ticker symbol, e.g. AAPL"
}
},
"required": ["ticker"]
},
"defer_loading": True
},
{
"name": "search_web",
"description": "Search the web for information",
"input_schema": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "The search query"
}
},
"required": ["query"]
},
"defer_loading": True
},
]
def get_tool_search_tool_regex() -> Dict[str, Any]:
"""Returns the tool search tool using regex variant."""
return {
"type": "tool_search_tool_regex_20251119",
"name": "tool_search_tool_regex"
}
def get_tool_search_tool_bm25() -> Dict[str, Any]:
"""Returns the tool search tool using BM25 variant."""
return {
"type": "tool_search_tool_bm25_20251119",
"name": "tool_search_tool_bm25"
}
class BaseAnthropicMessagesToolSearchTest(ABC):
"""
Base test class for tool search E2E tests across different providers.
Subclasses must implement:
- get_model(): Returns the model string to use for tests
Tests pass the anthropic-beta header via extra_headers to validate
that the header is correctly forwarded to downstream providers.
"""
@abstractmethod
def get_model(self) -> str:
"""
Returns the model string to use for tests.
Examples:
- "anthropic/claude-sonnet-4-20250514"
- "vertex_ai/claude-sonnet-4@20250514"
- "bedrock/invoke/anthropic.claude-sonnet-4-20250514-v1:0"
"""
pass
def get_extra_headers(self) -> Dict[str, str]:
"""
Returns extra headers to pass with the request.
Includes the anthropic-beta header for tool search.
This is what claude code forwards, simulate the same behavior here.
"""
return {"anthropic-beta": "advanced-tool-use-2025-11-20"}
def get_tools_with_tool_search(self) -> List[Dict[str, Any]]:
"""
Returns tools list with tool search tool and deferred tools.
"""
return [get_tool_search_tool_regex()] + get_deferred_tools()
@pytest.mark.asyncio
async def test_tool_search_basic_request(self):
"""
E2E test: Basic tool search request should succeed.
This validates that the tool search beta header is being passed via
extra_headers and forwarded correctly to the downstream provider.
"""
litellm._turn_on_debug()
tools = self.get_tools_with_tool_search()
messages = [
{
"role": "user",
"content": "What's the weather in San Francisco?"
}
]
response = await litellm.anthropic.messages.acreate(
model=self.get_model(),
messages=messages,
tools=tools,
max_tokens=1024,
extra_headers=self.get_extra_headers(),
)
print(f"Response: {json.dumps(response, indent=2, default=str)}")
# Validate response structure
assert "content" in response, "Response should contain content"
assert "usage" in response, "Response should contain usage"
# The model should either respond with text or use a tool
content = response.get("content", [])
assert len(content) > 0, "Response should have content"
@pytest.mark.asyncio
async def test_tool_search_discovers_tool(self):
"""
E2E test: Tool search should discover and use a deferred tool.
This validates that when the user asks about weather, the model
discovers the get_weather tool via tool search and attempts to use it.
"""
litellm._turn_on_debug()
tools = self.get_tools_with_tool_search()
messages = [
{
"role": "user",
"content": "I need to know the current weather in New York City. Please use the appropriate tool."
}
]
response = await litellm.anthropic.messages.acreate(
model=self.get_model(),
messages=messages,
tools=tools,
max_tokens=1024,
extra_headers=self.get_extra_headers(),
)
print(f"Response: {json.dumps(response, indent=2, default=str)}")
content = response.get("content", [])
# Check if the model used tool_use (either tool_search or get_weather)
tool_uses = [block for block in content if block.get("type") == "tool_use"]
print(f"Tool uses: {json.dumps(tool_uses, indent=2, default=str)}")
# The model should attempt to use tools when asked about weather
# It might use tool_search first, or directly use get_weather if discovered
if response.get("stop_reason") == "tool_use":
assert len(tool_uses) > 0, "Expected tool_use blocks when stop_reason is tool_use"
@pytest.mark.asyncio
async def test_tool_search_streaming(self):
"""
E2E test: Tool search should work with streaming responses.
"""
litellm._turn_on_debug()
tools = self.get_tools_with_tool_search()
messages = [
{
"role": "user",
"content": "What's the weather like in Tokyo?"
}
]
response = await litellm.anthropic.messages.acreate(
model=self.get_model(),
messages=messages,
tools=tools,
max_tokens=1024,
stream=True,
extra_headers=self.get_extra_headers(),
)
# Collect all chunks
chunks = []
async for chunk in response:
if isinstance(chunk, bytes):
chunk_str = chunk.decode("utf-8")
for line in chunk_str.split("\n"):
if line.startswith("data: "):
try:
json_data = json.loads(line[6:])
chunks.append(json_data)
print(f"Chunk: {json.dumps(json_data, indent=2, default=str)}")
except json.JSONDecodeError:
pass
elif isinstance(chunk, dict):
chunks.append(chunk)
print(f"Chunk: {json.dumps(chunk, indent=2, default=str)}")
# Should have received chunks
assert len(chunks) > 0, "Expected to receive streaming chunks"
# Should have message_start
message_starts = [c for c in chunks if c.get("type") == "message_start"]
assert len(message_starts) > 0, "Expected message_start in streaming response"
@pytest.mark.asyncio
async def test_tool_search_with_multiple_deferred_tools(self):
"""
E2E test: Tool search should work with multiple deferred tools.
This validates that the model can discover the appropriate tool
from a larger catalog of deferred tools.
"""
litellm._turn_on_debug()
tools = self.get_tools_with_tool_search()
messages = [
{
"role": "user",
"content": "What's the stock price of Apple (AAPL)?"
}
]
response = await litellm.anthropic.messages.acreate(
model=self.get_model(),
messages=messages,
tools=tools,
max_tokens=1024,
extra_headers=self.get_extra_headers(),
)
print(f"Response: {json.dumps(response, indent=2, default=str)}")
# Validate response
assert "content" in response, "Response should contain content"
content = response.get("content", [])
tool_uses = [block for block in content if block.get("type") == "tool_use"]
# If the model decides to use a tool, it should be related to stocks
if tool_uses:
tool_names = [t.get("name") for t in tool_uses]
print(f"Tools used: {tool_names}")

View file

@ -0,0 +1,83 @@
"""
E2E Test suite for Anthropic Messages API tool search across different providers.
Tests that tool search works correctly via litellm.anthropic.messages interface
by making actual API calls.
Supported providers:
- Anthropic API: advanced-tool-use-2025-11-20
- Azure Anthropic: advanced-tool-use-2025-11-20
- Vertex AI: tool-search-tool-2025-10-19
- Bedrock Invoke: tool-search-tool-2025-10-19
Reference: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool
"""
import os
import sys
sys.path.insert(0, os.path.abspath("../../.."))
import pytest
from base_anthropic_messages_tool_search_test import (
BaseAnthropicMessagesToolSearchTest,
)
class TestAnthropicAPIToolSearch(BaseAnthropicMessagesToolSearchTest):
"""
E2E tests for tool search with Anthropic API directly.
Uses the anthropic/ prefix which routes through the native
Anthropic Messages API.
Beta header: advanced-tool-use-2025-11-20
Note: Tool search is only supported on Claude Opus 4.5 and Claude Sonnet 4.5.
"""
def get_model(self) -> str:
return "anthropic/claude-sonnet-4-5-20250929"
# class TestAzureAnthropicToolSearch(BaseAnthropicMessagesToolSearchTest):
# """
# E2E tests for tool search with Azure Anthropic (Microsoft Foundry).
# Uses the azure/ prefix which routes through Azure's Anthropic endpoint.
# Beta header: advanced-tool-use-2025-11-20
# """
# def get_model(self) -> str:
# return "azure/claude-sonnet-4-20250514"
# class TestVertexAIToolSearch(BaseAnthropicMessagesToolSearchTest):
# """
# E2E tests for tool search with Vertex AI.
# Uses the vertex_ai/ prefix which routes through Google Cloud's
# Vertex AI Anthropic partner models.
# Beta header: tool-search-tool-2025-10-19
# """
# def get_model(self) -> str:
# return "vertex_ai/claude-sonnet-4@20250514"
class TestBedrockInvokeToolSearch(BaseAnthropicMessagesToolSearchTest):
"""
E2E tests for tool search with Bedrock Invoke API.
Uses the bedrock/invoke/ prefix which routes through the native
Anthropic Messages API format on Bedrock.
Beta header: advanced-tool-use-2025-11-20 (passed via extra_headers)
Note: Tool search on Bedrock is only supported on Claude Opus 4.5.
"""
def get_model(self) -> str:
return "bedrock/invoke/us.anthropic.claude-opus-4-5-20251101-v1:0"

View file

@ -56,6 +56,12 @@ def test_routes_on_litellm_proxy():
# realtime routes - /realtime?model=gpt-4o
if "realtime" in route:
assert "/realtime" in _all_routes
# wildcard patterns like /containers/* - check that base path exists
elif RouteChecks._is_wildcard_pattern(pattern=route):
# For wildcard patterns, check that the base path (without * and trailing /) exists
base_path = route[:-1].rstrip("/") # Remove the trailing * and any trailing /
# Check if base path exists (e.g., /containers or /v1/containers)
assert base_path in _all_routes, f"Wildcard pattern {route} requires base path {base_path} to exist"
else:
assert route in _all_routes

View file

@ -1,590 +0,0 @@
"""
Tests for zero-cost model budget bypass functionality.
When a user exceeds their budget, the system should still allow requests
to models with zero cost (e.g., on-premises models).
"""
import asyncio
from typing import Optional
from unittest.mock import MagicMock, patch
import pytest
import litellm
from litellm.caching.caching import DualCache
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_UserTable,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import (
_check_team_member_budget,
_is_model_cost_zero,
_team_max_budget_check,
common_checks,
)
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
@pytest.fixture
def mock_router_with_zero_cost_model():
"""Create a mock router with a zero-cost model."""
router = Router(
model_list=[
{
"model_name": "on-prem-model",
"litellm_params": {
"model": "ollama/llama2",
"api_base": "http://localhost:11434",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
},
"model_info": {
"id": "on-prem-model-id",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
},
},
{
"model_name": "cloud-model",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "sk-test",
},
"model_info": {
"id": "cloud-model-id",
},
},
]
)
return router
@pytest.fixture
def mock_router_with_paid_model():
"""Create a mock router with only paid models."""
router = Router(
model_list=[
{
"model_name": "cloud-model",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "sk-test",
},
"model_info": {
"id": "cloud-model-id",
},
}
]
)
return router
@pytest.fixture
def mock_proxy_logging():
"""Create a mock ProxyLogging instance."""
proxy_logging = ProxyLogging(user_api_key_cache=None)
async def mock_budget_alerts(*args, **kwargs):
pass
proxy_logging.budget_alerts = mock_budget_alerts
return proxy_logging
class TestIsModelCostZero:
"""Tests for _is_model_cost_zero helper function."""
def test_zero_cost_model_in_router(self, mock_router_with_zero_cost_model):
"""Test that a zero-cost model in router is correctly identified."""
result = _is_model_cost_zero(
model="on-prem-model", llm_router=mock_router_with_zero_cost_model
)
assert result is True
def test_paid_model_in_router(self, mock_router_with_zero_cost_model):
"""Test that a paid model is correctly identified as non-zero cost."""
with patch("litellm.get_model_info") as mock_get_model_info:
# Mock the return value for gpt-3.5-turbo
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
result = _is_model_cost_zero(
model="cloud-model", llm_router=mock_router_with_zero_cost_model
)
assert result is False
def test_none_model(self, mock_router_with_zero_cost_model):
"""Test that None model returns False."""
result = _is_model_cost_zero(
model=None, llm_router=mock_router_with_zero_cost_model
)
assert result is False
def test_none_router(self):
"""Test that None router returns False."""
result = _is_model_cost_zero(model="some-model", llm_router=None)
assert result is False
def test_list_of_zero_cost_models(self, mock_router_with_zero_cost_model):
"""Test that a list of zero-cost models returns True."""
result = _is_model_cost_zero(
model=["on-prem-model"], llm_router=mock_router_with_zero_cost_model
)
assert result is True
def test_mixed_cost_models(self, mock_router_with_zero_cost_model):
"""Test that a list with mixed cost models returns False."""
with patch("litellm.get_model_info") as mock_get_model_info:
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
result = _is_model_cost_zero(
model=["on-prem-model", "cloud-model"],
llm_router=mock_router_with_zero_cost_model,
)
assert result is False
class TestUserBudgetBypass:
"""Tests for user budget bypass with zero-cost models."""
@pytest.mark.asyncio
async def test_user_over_budget_with_zero_cost_model_allowed(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that user over budget can still use zero-cost models."""
user_object = LiteLLM_UserTable(
user_id="test-user",
spend=100.0,
max_budget=50.0,
)
request_body = {"model": "on-prem-model"}
# Should not raise BudgetExceededError
result = await common_checks(
request_body=request_body,
team_object=None,
user_object=user_object,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=UserAPIKeyAuth(
token="test-token",
user_id="test-user",
),
request=MagicMock(),
skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
)
assert result is True
@pytest.mark.asyncio
async def test_user_over_budget_with_paid_model_blocked(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that user over budget cannot use paid models."""
user_object = LiteLLM_UserTable(
user_id="test-user",
spend=100.0,
max_budget=50.0,
)
request_body = {"model": "cloud-model"}
with patch("litellm.get_model_info") as mock_get_model_info:
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body=request_body,
team_object=None,
user_object=user_object,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=UserAPIKeyAuth(
token="test-token",
user_id="test-user",
),
request=MagicMock(),
)
assert exc_info.value.current_cost == 100.0
assert exc_info.value.max_budget == 50.0
assert "test-user" in str(exc_info.value)
class TestEndUserBudgetBypass:
"""Tests for end user budget bypass with zero-cost models."""
@pytest.mark.asyncio
async def test_end_user_over_budget_with_zero_cost_model_allowed(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that end user over budget can still use zero-cost models."""
end_user_budget = LiteLLM_BudgetTable(max_budget=20.0)
end_user_object = LiteLLM_EndUserTable(
user_id="end-user-123",
spend=50.0,
litellm_budget_table=end_user_budget,
blocked=False,
)
request_body = {"model": "on-prem-model", "user": "end-user-123"}
# In the real flow, skip_budget_checks would be set to True for zero-cost models
result = await common_checks(
request_body=request_body,
team_object=None,
user_object=None,
end_user_object=end_user_object,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=UserAPIKeyAuth(
token="test-token",
),
request=MagicMock(),
skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
)
assert result is True
@pytest.mark.asyncio
async def test_end_user_over_budget_with_paid_model_blocked(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that end user over budget cannot use paid models."""
end_user_budget = LiteLLM_BudgetTable(max_budget=20.0)
end_user_object = LiteLLM_EndUserTable(
user_id="end-user-123",
spend=50.0,
litellm_budget_table=end_user_budget,
blocked=False,
)
request_body = {"model": "cloud-model", "user": "end-user-123"}
with patch("litellm.get_model_info") as mock_get_model_info:
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body=request_body,
team_object=None,
user_object=None,
end_user_object=end_user_object,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=UserAPIKeyAuth(
token="test-token",
),
request=MagicMock(),
)
assert exc_info.value.current_cost == 50.0
assert exc_info.value.max_budget == 20.0
assert "end-user-123" in str(exc_info.value)
class TestTeamBudgetBypass:
"""Tests for team budget bypass with zero-cost models."""
@pytest.mark.asyncio
async def test_team_over_budget_with_zero_cost_model_allowed(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that team over budget can still use zero-cost models."""
team_object = LiteLLM_TeamTable(
team_id="test-team",
spend=150.0,
max_budget=100.0,
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="test-team",
)
request_body = {"model": "on-prem-model"}
# In the real flow, skip_budget_checks would be set to True for zero-cost models
result = await common_checks(
request_body=request_body,
team_object=team_object,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=valid_token,
request=MagicMock(),
skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
)
assert result is True
@pytest.mark.asyncio
async def test_team_over_budget_with_paid_model_blocked(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that team over budget cannot use paid models."""
team_object = LiteLLM_TeamTable(
team_id="test-team",
spend=150.0,
max_budget=100.0,
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="test-team",
)
request_body = {"model": "cloud-model"}
with patch("litellm.get_model_info") as mock_get_model_info:
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body=request_body,
team_object=team_object,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=valid_token,
request=MagicMock(),
)
assert exc_info.value.current_cost == 150.0
assert exc_info.value.max_budget == 100.0
assert "test-team" in str(exc_info.value)
class TestTeamMemberBudgetBypass:
"""Tests for team member budget bypass with zero-cost models."""
@pytest.mark.asyncio
async def test_team_member_over_budget_with_zero_cost_model_allowed(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that team member over budget can still use zero-cost models."""
team_object = LiteLLM_TeamTable(
team_id="test-team",
)
user_object = LiteLLM_UserTable(
user_id="test-user",
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user",
team_id="test-team",
)
member_budget = LiteLLM_BudgetTable(max_budget=30.0)
team_membership = LiteLLM_TeamMembership(
user_id="test-user",
team_id="test-team",
spend=60.0,
litellm_budget_table=member_budget,
)
request_body = {"model": "on-prem-model"}
# Mock get_team_membership
with patch(
"litellm.proxy.auth.auth_checks.get_team_membership"
) as mock_get_membership:
mock_get_membership.return_value = team_membership
# In the real flow, skip_budget_checks would be set to True for zero-cost models
result = await common_checks(
request_body=request_body,
team_object=team_object,
user_object=user_object,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=valid_token,
request=MagicMock(),
skip_budget_checks=True, # This is set by user_api_key_auth for zero-cost models
)
assert result is True
@pytest.mark.asyncio
async def test_team_member_over_budget_with_paid_model_blocked(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that team member over budget cannot use paid models."""
team_object = LiteLLM_TeamTable(
team_id="test-team",
)
user_object = LiteLLM_UserTable(
user_id="test-user",
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user",
team_id="test-team",
)
member_budget = LiteLLM_BudgetTable(max_budget=30.0)
team_membership = LiteLLM_TeamMembership(
user_id="test-user",
team_id="test-team",
spend=60.0,
litellm_budget_table=member_budget,
)
request_body = {"model": "cloud-model"}
with patch(
"litellm.proxy.auth.auth_checks.get_team_membership"
) as mock_get_membership:
mock_get_membership.return_value = team_membership
with patch("litellm.get_model_info") as mock_get_model_info:
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body=request_body,
team_object=team_object,
user_object=user_object,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=valid_token,
request=MagicMock(),
)
assert exc_info.value.current_cost == 60.0
assert exc_info.value.max_budget == 30.0
assert "test-user" in str(exc_info.value)
assert "test-team" in str(exc_info.value)
class TestEdgeCases:
"""Tests for edge cases and error handling."""
def test_model_not_in_router(self, mock_router_with_zero_cost_model):
"""Test behavior when model is not found in router."""
with patch("litellm.get_model_info") as mock_get_model_info:
# Simulate model not found
mock_get_model_info.side_effect = Exception("Model not found")
result = _is_model_cost_zero(
model="nonexistent-model", llm_router=mock_router_with_zero_cost_model
)
# Should return False (conservative approach)
assert result is False
@pytest.mark.asyncio
async def test_user_under_budget_with_paid_model_allowed(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that user under budget can use paid models normally."""
user_object = LiteLLM_UserTable(
user_id="test-user",
spend=30.0,
max_budget=100.0,
)
request_body = {"model": "cloud-model"}
with patch("litellm.get_model_info") as mock_get_model_info:
mock_get_model_info.return_value = {
"input_cost_per_token": 0.0000015,
"output_cost_per_token": 0.000002,
}
# Should not raise BudgetExceededError
result = await common_checks(
request_body=request_body,
team_object=None,
user_object=user_object,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=UserAPIKeyAuth(
token="test-token",
user_id="test-user",
),
request=MagicMock(),
)
assert result is True
@pytest.mark.asyncio
async def test_user_under_budget_with_zero_cost_model_allowed(
self, mock_router_with_zero_cost_model, mock_proxy_logging
):
"""Test that user under budget can use zero-cost models normally."""
user_object = LiteLLM_UserTable(
user_id="test-user",
spend=30.0,
max_budget=100.0,
)
request_body = {"model": "on-prem-model"}
# Should not raise BudgetExceededError
result = await common_checks(
request_body=request_body,
team_object=None,
user_object=user_object,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=mock_router_with_zero_cost_model,
proxy_logging_obj=mock_proxy_logging,
valid_token=UserAPIKeyAuth(
token="test-token",
user_id="test-user",
),
request=MagicMock(),
)
assert result is True

View file

@ -422,7 +422,7 @@ def test_streaming_tool_calls_transformation():
ChatCompletionDeltaToolCall,
Delta,
Function,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
@ -454,7 +454,7 @@ def test_streaming_tool_calls_transformation():
delta=mock_delta
)
mock_response = ModelResponse(
mock_response = ModelResponseStream(
id="test-streaming",
choices=[mock_choice],
created=1234567890,
@ -493,7 +493,7 @@ def test_streaming_partial_tool_calls_accumulation():
ChatCompletionDeltaToolCall,
Delta,
Function,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
@ -543,7 +543,7 @@ def test_streaming_partial_tool_calls_accumulation():
delta=mock_delta
)
mock_response = ModelResponse(
mock_response = ModelResponseStream(
id="test-streaming",
choices=[mock_choice],
created=1234567890,
@ -595,7 +595,7 @@ def test_streaming_multiple_partial_tool_calls():
ChatCompletionDeltaToolCall,
Delta,
Function,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
@ -642,7 +642,7 @@ def test_streaming_multiple_partial_tool_calls():
delta=mock_delta
)
mock_response = ModelResponse(
mock_response = ModelResponseStream(
id="test-streaming",
choices=[mock_choice],
created=1234567890,

View file

@ -0,0 +1,380 @@
"""
Test message sanitization for Anthropic API when modify_params=True
Tests three cases:
A. Missing tool_result for tool_use (orphaned tool calls)
B. Orphaned tool_result without matching tool_use
C. Empty text content
"""
import pytest
import sys
import os
# Add the parent directory to the path so we can import litellm
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")))
import litellm
from litellm.litellm_core_utils.prompt_templates.factory import (
sanitize_messages_for_tool_calling,
anthropic_messages_pt,
)
class TestMessageSanitization:
"""Test message sanitization for tool calling scenarios"""
def setup_method(self):
"""Setup for each test"""
# Save original modify_params value
self.original_modify_params = litellm.modify_params
litellm.modify_params = True
def teardown_method(self):
"""Cleanup after each test"""
# Restore original modify_params value
litellm.modify_params = self.original_modify_params
def test_case_a_orphaned_tool_call_single(self):
"""
Test Case A: Assistant message with tool_calls but no tool result
Should add a dummy tool result message
"""
messages = [
{
"role": "user",
"content": "What is the weather in Nashik?"
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "toolu_01Kus2cC3ydjBW7UK4GJqBP4",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "Nashik, India"}'
}
}
]
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
# Should have 3 messages: user, assistant, and dummy tool result
assert len(sanitized) == 3
assert sanitized[0]["role"] == "user"
assert sanitized[1]["role"] == "assistant"
assert sanitized[2]["role"] == "tool"
assert sanitized[2]["tool_call_id"] == "toolu_01Kus2cC3ydjBW7UK4GJqBP4"
assert "skipped" in sanitized[2]["content"].lower() or "interrupted" in sanitized[2]["content"].lower()
assert "get_weather" in sanitized[2]["content"]
def test_case_a_orphaned_tool_call_multiple(self):
"""
Test Case A: Assistant message with multiple tool_calls, some missing results
"""
messages = [
{
"role": "user",
"content": "Get weather for Nashik and Mumbai"
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "Nashik"}'
}
},
{
"id": "call_2",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "Mumbai"}'
}
}
]
},
{
"role": "tool",
"tool_call_id": "call_1",
"content": "Weather in Nashik: 25°C"
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
# Should have 4 messages: user, assistant, tool result for call_1, dummy for call_2
assert len(sanitized) == 4
assert sanitized[0]["role"] == "user"
assert sanitized[1]["role"] == "assistant"
assert sanitized[2]["tool_call_id"] == "call_2" # Dummy added first
assert sanitized[3]["tool_call_id"] == "call_1" # Original tool result
def test_case_b_orphaned_tool_result(self):
"""
Test Case B: Tool result without matching tool_call in previous assistant message
Should remove the orphaned tool result
"""
messages = [
{
"role": "user",
"content": "Hello"
},
{
"role": "assistant",
"content": "Hi there!"
},
{
"role": "tool",
"tool_call_id": "nonexistent_id",
"content": "Some result"
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
# Should have only 2 messages, orphaned tool result removed
assert len(sanitized) == 2
assert sanitized[0]["role"] == "user"
assert sanitized[1]["role"] == "assistant"
def test_case_b_valid_tool_result_preserved(self):
"""
Test Case B: Valid tool result with matching tool_call should be preserved
"""
messages = [
{
"role": "user",
"content": "What's the weather?"
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "Boston"}'
}
}
]
},
{
"role": "tool",
"tool_call_id": "call_123",
"content": "Weather: 20°C"
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
# All messages should be preserved
assert len(sanitized) == 3
assert sanitized[2]["role"] == "tool"
assert sanitized[2]["tool_call_id"] == "call_123"
def test_case_c_empty_text_content_user(self):
"""
Test Case C: Empty text content in user message
Should replace with placeholder
"""
messages = [
{
"role": "user",
"content": ""
},
{
"role": "assistant",
"content": "Hello!"
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
assert len(sanitized) == 2
assert sanitized[0]["role"] == "user"
assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]"
def test_case_c_whitespace_only_content(self):
"""
Test Case C: Whitespace-only content
Should replace with placeholder
"""
messages = [
{
"role": "user",
"content": " \n \t "
},
{
"role": "assistant",
"content": " "
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
assert len(sanitized) == 2
assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]"
assert sanitized[1]["content"] == "[System: Empty message content sanitised to satisfy protocol]"
def test_case_c_valid_content_preserved(self):
"""
Test Case C: Valid non-empty content should be preserved
"""
messages = [
{
"role": "user",
"content": "Hello"
},
{
"role": "assistant",
"content": "Hi there!"
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
assert len(sanitized) == 2
assert sanitized[0]["content"] == "Hello"
assert sanitized[1]["content"] == "Hi there!"
def test_combined_cases(self):
"""
Test combination of multiple cases
"""
messages = [
{
"role": "user",
"content": "Get weather"
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}'
}
}
]
},
# Missing tool result for call_1
{
"role": "user",
"content": "" # Empty content
},
{
"role": "assistant",
"content": "Response"
},
{
"role": "tool",
"tool_call_id": "orphaned_id", # Orphaned tool result
"content": "Some data"
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
# Should have: user, assistant, dummy tool result, user (sanitized), assistant
# Orphaned tool result should be removed
assert len(sanitized) == 5
assert sanitized[0]["role"] == "user"
assert sanitized[1]["role"] == "assistant"
assert sanitized[2]["role"] == "tool"
assert sanitized[2]["tool_call_id"] == "call_1" # Dummy added
assert sanitized[3]["role"] == "user"
assert sanitized[3]["content"] == "[System: Empty message content sanitised to satisfy protocol]"
assert sanitized[4]["role"] == "assistant"
def test_modify_params_false_no_sanitization(self):
"""
Test that sanitization is skipped when modify_params=False
"""
litellm.modify_params = False
messages = [
{
"role": "user",
"content": ""
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{}'
}
}
]
}
]
sanitized = sanitize_messages_for_tool_calling(messages)
# Messages should be unchanged
assert len(sanitized) == 2
assert sanitized[0]["content"] == ""
assert len(sanitized[1].get("tool_calls", [])) == 1
def test_anthropic_messages_pt_integration(self):
"""
Test that sanitization is integrated into anthropic_messages_pt
"""
litellm.modify_params = True
messages = [
{
"role": "user",
"content": "What is the weather in Nashik?"
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "toolu_01Kus2cC3ydjBW7UK4GJqBP4",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "Nashik, India"}'
}
}
]
}
]
# This should not raise an error and should add dummy tool result
result = anthropic_messages_pt(
messages=messages,
model="claude-sonnet-4-5",
llm_provider="anthropic"
)
# Should have at least 2 messages (user and assistant)
# The tool result will be merged into user content
assert len(result) >= 2
assert result[0]["role"] == "user"
assert result[1]["role"] == "assistant"
if __name__ == "__main__":
pytest.main([__file__, "-v"])

View file

@ -440,6 +440,7 @@ def test_select_azure_base_url_called(setup_mocks):
"asearch",
"avector_store_create",
"avector_store_search",
"acreate_skill",
]
],
)

View file

@ -333,15 +333,18 @@ def _make_mock_response(should_fail=False, fail_count={"count": 0}):
@pytest.mark.asyncio
async def test_handle_async_request_total_timeout_triggers():
async def test_handle_async_request_sock_read_timeout_triggers():
"""
Ensure that LiteLLMAiohttpTransport raises httpx.TimeoutException
when the total timeout duration elapses.
when the sock_read timeout duration elapses (individual read operation timeout).
This is the correct behavior for stream_timeout - it should timeout on slow reads,
not on the total duration of the stream.
"""
import asyncio
from aiohttp import web
async def slow_handler(request):
# Sleep longer than the sock_read timeout
await asyncio.sleep(0.3)
return web.Response(text="ok")
@ -361,11 +364,12 @@ async def test_handle_async_request_total_timeout_triggers():
request = httpx.Request("GET", f"http://127.0.0.1:{port}/")
# Set a short sock_read timeout - this should trigger
# Note: total timeout is NOT set, allowing long-running streams
request.extensions["timeout"] = {
"connect": 0.1,
"read": 0.1,
"pool": 0.1,
"total": 0.1,
"connect": 5.0,
"read": 0.1, # Short timeout for individual reads
"pool": 5.0,
}
try:
@ -376,6 +380,77 @@ async def test_handle_async_request_total_timeout_triggers():
await runner.cleanup()
@pytest.mark.asyncio
async def test_handle_async_request_streaming_does_not_timeout_on_total_duration():
"""
Ensure that LiteLLMAiohttpTransport does NOT timeout on long-running
streaming responses as long as individual chunks arrive within the sock_read timeout.
This is the fix for issue #19184 - stream_timeout should only control the timeout
for individual chunks, not the total stream duration.
"""
import asyncio
from aiohttp import web
async def streaming_handler(request):
# Simulate a streaming response that takes longer than a single timeout
# but each chunk arrives quickly
response = web.StreamResponse()
await response.prepare(request)
# Send 5 chunks over 0.5 seconds total (0.1s between chunks)
for i in range(5):
await asyncio.sleep(0.05) # Less than sock_read timeout
await response.write(f"chunk{i}\n".encode())
await response.write_eof()
return response
app = web.Application()
app.router.add_get("/stream", streaming_handler)
runner = web.AppRunner(app)
await runner.setup()
site = web.TCPSite(runner, "127.0.0.1", 0)
await site.start()
port = site._server.sockets[0].getsockname()[1]
def factory():
return aiohttp.ClientSession()
transport = LiteLLMAiohttpTransport(client=factory) # type: ignore
request = httpx.Request("GET", f"http://127.0.0.1:{port}/stream")
# Set sock_read timeout that's longer than individual chunk delays
# but shorter than total stream duration
# Total duration: ~0.25s, sock_read timeout: 0.15s per chunk
# This should NOT timeout because each chunk arrives within 0.15s
request.extensions["timeout"] = {
"connect": 5.0,
"read": 0.15, # Timeout for individual reads
"pool": 5.0,
# Note: total is NOT set - this is the fix!
}
try:
# This should succeed without timing out
response = await transport.handle_async_request(request)
assert response.status_code == 200
# Read the streaming response
chunks = []
async for chunk in response.aiter_bytes():
chunks.append(chunk)
# Verify we got all chunks
full_response = b"".join(chunks).decode()
assert "chunk0" in full_response
assert "chunk4" in full_response
finally:
await transport.aclose()
await runner.cleanup()
def _make_mock_session(closed=False):
"""Helper to create a mock aiohttp session"""

View file

@ -0,0 +1,16 @@
"""Test for Gemini schema handling with empty properties."""
import os
import sys
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.llms.vertex_ai.common_utils import add_object_type
def test_add_object_type_empty_properties_keeps_type():
"""Gemini requires type: object even when properties is empty."""
schema = {"properties": {}, "type": "object"}
add_object_type(schema)
assert schema.get("type") == "object"
assert "properties" not in schema

View file

@ -1,7 +1,6 @@
import os
import sys
from typing import Any, Dict
from unittest.mock import MagicMock, call, patch
from unittest.mock import patch
import pytest
@ -11,7 +10,6 @@ sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.vertex_ai.common_utils import (
_get_vertex_url,
convert_anyof_null_to_nullable,
@ -798,9 +796,54 @@ def test_fix_enum_empty_strings():
assert "mobile" in enum_values
assert "tablet" in enum_values
# 3. Other properties preserved
assert input_schema["properties"]["user_agent_type"]["type"] == "string"
assert input_schema["properties"]["user_agent_type"]["description"] == "Device type for user agent"
def test_get_vertex_model_id_from_url():
"""Test get_vertex_model_id_from_url with various URLs"""
from litellm.llms.vertex_ai.common_utils import get_vertex_model_id_from_url
# Test with valid URL
url = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent"
model_id = get_vertex_model_id_from_url(url)
assert model_id == "gemini-pro"
# Test with invalid URL
url = "https://invalid-url.com"
model_id = get_vertex_model_id_from_url(url)
assert model_id is None
def test_construct_target_url_with_version_prefix():
"""Test construct_target_url with version prefixes"""
from litellm.llms.vertex_ai.common_utils import construct_target_url
# Test with /v1/ prefix
url = "/v1/publishers/google/models/gemini-pro:streamGenerateContent"
vertex_project = "test-project"
vertex_location = "us-central1"
base_url = "https://us-central1-aiplatform.googleapis.com"
target_url = construct_target_url(
base_url=base_url,
requested_route=url,
vertex_project=vertex_project,
vertex_location=vertex_location,
)
expected_url = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent"
assert str(target_url) == expected_url
# Test with /v1beta1/ prefix
url = "/v1beta1/publishers/google/models/gemini-pro:streamGenerateContent"
target_url = construct_target_url(
base_url=base_url,
requested_route=url,
vertex_project=vertex_project,
vertex_location=vertex_location,
)
expected_url = "https://us-central1-aiplatform.googleapis.com/v1beta1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent"
assert str(target_url) == expected_url
def test_fix_enum_types():
@ -862,7 +905,7 @@ def test_fix_enum_types():
"truncateMode": {
"enum": ["auto", "none", "start", "end"], # Kept - string type
"type": "string",
"description": "How to truncate content"
"description": "How to truncate content",
},
"maxLength": { # enum removed
"type": "integer",
@ -1254,8 +1297,8 @@ def test_build_vertex_schema_empty_properties():
# Verify empty properties was removed
assert "properties" not in go_back_schema, "Empty properties should be removed"
# Verify type was also removed (since object without properties is invalid in Gemini)
assert "type" not in go_back_schema, "Type should be removed when properties is empty"
# Verify type is kept as object (Gemini requires type: object even without properties)
assert go_back_schema.get("type") == "object", "Type should be kept as object when properties is empty"
# Verify required was also removed
assert "required" not in go_back_schema, "Required should be removed when properties is empty"

View file

@ -115,3 +115,147 @@ def test_vertex_ai_anthropic_structured_output_header_not_added():
"Non-Vertex request SHOULD have anthropic-beta header for structured output"
assert result_non_vertex["anthropic-beta"] == "structured-outputs-2025-11-13", \
f"Expected 'structured-outputs-2025-11-13', got: {result_non_vertex.get('anthropic-beta')}"
def test_vertex_ai_claude_sonnet_4_5_structured_output_fix():
"""
Test fix for issue #18625: Claude Sonnet 4.5 on VertexAI should use tool-based
structured outputs instead of output_format parameter.
This test verifies that:
1. Claude Sonnet 4.5 uses tool-based structured outputs on VertexAI
2. output_format parameter is removed from the final request
3. The fix prevents "Extra inputs are not permitted" error
"""
config = VertexAIAnthropicConfig()
# Test data matching the issue report
response_format = {
"type": "json_schema",
"json_schema": {
"name": "questions",
"strict": True,
"schema": {
"type": "object",
"properties": {
"question": {
"type": "string"
},
"response": {
"type": "string"
}
},
"required": ["question", "response"],
"additionalProperties": False
}
}
}
messages = [
{"role": "user", "content": "Generate a question and answer about AI."}
]
# Test parameters that would trigger the issue
non_default_params = {
"response_format": response_format,
"max_tokens": 1000,
}
# Test 1: Verify map_openai_params forces tool-based approach for Claude Sonnet 4.5
optional_params = {}
result_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="claude-3-5-sonnet-20241022", # Claude Sonnet 4.5 model
drop_params=False,
)
# Should have tools and tool_choice (tool-based approach)
assert "tools" in result_params, "Tools should be present for structured output"
assert "tool_choice" in result_params, "Tool choice should be present for structured output"
assert "json_mode" in result_params, "JSON mode should be enabled"
# Verify the tool is the response format tool
tools = result_params["tools"]
assert len(tools) == 1, "Should have exactly one tool for response format"
assert tools[0]["name"] == "json_tool_call", "Tool should be named json_tool_call"
# Test 2: Verify transform_request removes output_format parameter
# Simulate what would happen if parent class added output_format
test_data = {
"model": "claude-3-5-sonnet-20241022",
"messages": messages,
"max_tokens": 1000,
"tools": tools,
"tool_choice": result_params["tool_choice"],
"output_format": { # This would be added by parent class for Sonnet 4.5
"type": "json_schema",
"schema": response_format["json_schema"]["schema"]
}
}
# Mock the parent transform_request to return data with output_format
original_transform = config.__class__.__bases__[0].transform_request
def mock_transform_request(self, model, messages, optional_params, litellm_params, headers):
# Return test data that includes output_format
return test_data.copy()
# Temporarily replace parent method
config.__class__.__bases__[0].transform_request = mock_transform_request
try:
final_data = config.transform_request(
model="claude-3-5-sonnet-20241022",
messages=messages,
optional_params=result_params,
litellm_params={},
headers={},
)
# Verify that output_format was removed (fixes the "Extra inputs are not permitted" error)
assert "output_format" not in final_data, "output_format should be removed for VertexAI"
assert "model" not in final_data, "model should be removed for VertexAI"
assert "tools" in final_data, "tools should still be present"
assert "tool_choice" in final_data, "tool_choice should still be present"
finally:
# Restore original method
config.__class__.__bases__[0].transform_request = original_transform
def test_vertex_ai_anthropic_other_models_still_use_tools():
"""
Test that other Anthropic models (non-Sonnet 4.5) on VertexAI also use tool-based
structured outputs, ensuring consistency across all models.
"""
config = VertexAIAnthropicConfig()
response_format = {
"type": "json_schema",
"json_schema": {
"name": "test_schema",
"schema": {
"type": "object",
"properties": {
"result": {"type": "string"}
}
}
}
}
# Test with Claude 3 Sonnet (not 4.5)
non_default_params = {"response_format": response_format}
optional_params = {}
result_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="claude-3-sonnet-20240229",
drop_params=False,
)
# Should still use tool-based approach
assert "tools" in result_params, "Claude 3 Sonnet should also use tool-based structured output"
assert "tool_choice" in result_params, "Tool choice should be present"
assert "json_mode" in result_params, "JSON mode should be enabled"

View file

@ -1,9 +1,13 @@
"""
Unit tests for auth_utils functions related to rate limiting.
Unit tests for auth_utils functions related to rate limiting and customer ID extraction.
"""
from unittest.mock import patch
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
_get_customer_id_from_standard_headers,
get_end_user_id_from_request_body,
get_key_model_rpm_limit,
get_key_model_tpm_limit,
)
@ -129,3 +133,56 @@ class TestGetKeyModelTpmLimit:
)
result = get_key_model_tpm_limit(user_api_key_dict)
assert result == {"gpt-4": 10000}
class TestGetCustomerIdFromStandardHeaders:
"""Tests for _get_customer_id_from_standard_headers helper function."""
def test_should_return_customer_id_from_x_litellm_customer_id_header(self):
"""Should extract customer ID from x-litellm-customer-id header."""
headers = {"x-litellm-customer-id": "customer-123"}
result = _get_customer_id_from_standard_headers(request_headers=headers)
assert result == "customer-123"
def test_should_return_customer_id_from_x_litellm_end_user_id_header(self):
"""Should extract customer ID from x-litellm-end-user-id header."""
headers = {"x-litellm-end-user-id": "end-user-456"}
result = _get_customer_id_from_standard_headers(request_headers=headers)
assert result == "end-user-456"
def test_should_return_none_when_headers_is_none(self):
"""Should return None when headers is None."""
result = _get_customer_id_from_standard_headers(request_headers=None)
assert result is None
def test_should_return_none_when_no_standard_headers_present(self):
"""Should return None when no standard customer ID headers are present."""
headers = {"x-other-header": "some-value"}
result = _get_customer_id_from_standard_headers(request_headers=headers)
assert result is None
class TestGetEndUserIdFromRequestBodyWithStandardHeaders:
"""Tests for get_end_user_id_from_request_body with standard customer ID headers."""
def test_should_prioritize_standard_header_over_body_user(self):
"""Standard customer ID header should take precedence over body user field."""
headers = {"x-litellm-customer-id": "header-customer"}
request_body = {"user": "body-user"}
with patch("litellm.proxy.proxy_server.general_settings", {}):
result = get_end_user_id_from_request_body(
request_body=request_body, request_headers=headers
)
assert result == "header-customer"
def test_should_fall_back_to_body_when_no_standard_header(self):
"""Should fall back to body user when no standard headers are present."""
headers = {"x-other-header": "value"}
request_body = {"user": "body-user"}
with patch("litellm.proxy.proxy_server.general_settings", {}):
result = get_end_user_id_from_request_body(
request_body=request_body, request_headers=headers
)
assert result == "body-user"

View file

@ -14,7 +14,6 @@ sys.path.insert(
from fastapi import HTTPException
import litellm
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)

View file

@ -20,7 +20,6 @@ from litellm.proxy._types import (
LiteLLM_OrganizationTable,
LiteLLM_OrganizationTableWithMembers,
LiteLLM_TeamTable,
LiteLLM_UserTable,
LitellmUserRoles,
Member,
ProxyErrorTypes,
@ -4826,187 +4825,6 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth):
assert deserialized_settings == router_settings_data
@pytest.mark.asyncio
async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys(
mock_db_client,
):
"""
Test that non-team-admin users only see their own spend (filtered by their API keys)
when calling /team/daily/activity endpoint.
"""
from litellm.proxy.management_endpoints.team_endpoints import (
get_team_daily_activity,
)
# Create a non-admin user
user_id = "test_user_123"
team_id = "test_team_456"
user_api_key_dict = UserAPIKeyAuth(
user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER
)
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id=user_id,
teams=[team_id],
max_budget=1000.0,
spend=0.0,
user_email="test@example.com",
user_role="internal_user",
)
# Mock team with user as non-admin member
mock_team_member = Member(user_id=user_id, role="user")
mock_team = MagicMock(spec=LiteLLM_TeamTable)
mock_team.team_id = team_id
mock_team.team_alias = "Test Team"
mock_team.members_with_roles = [mock_team_member]
mock_team.model_dump.return_value = {
"team_id": team_id,
"team_alias": "Test Team",
"members_with_roles": [{"user_id": user_id, "role": "user"}],
}
# Mock user's API keys
user_api_key_1 = MagicMock()
user_api_key_1.token = "user_key_1"
user_api_key_2 = MagicMock()
user_api_key_2.token = "user_key_2"
# Setup mocks
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(
return_value=[mock_team]
)
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[user_api_key_1, user_api_key_2]
)
# Mock get_user_object
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
new_callable=AsyncMock,
) as mock_get_user_object:
mock_get_user_object.return_value = mock_user_info
# Mock get_daily_activity to capture the api_key parameter
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity:
mock_get_daily_activity.return_value = MagicMock()
# Call the endpoint
await get_team_daily_activity(
team_ids=team_id,
start_date="2024-01-01",
end_date="2024-01-02",
model=None,
api_key=None,
page=1,
page_size=10,
exclude_team_ids=None,
user_api_key_dict=user_api_key_dict,
)
# Verify get_daily_activity was called with user's API keys as filter
mock_get_daily_activity.assert_called_once()
call_kwargs = mock_get_daily_activity.call_args[1]
assert call_kwargs["api_key"] == ["user_key_1", "user_key_2"]
assert call_kwargs["entity_id"] == [team_id]
# Verify user's API keys were fetched
mock_db_client.db.litellm_verificationtoken.find_many.assert_called_once()
api_key_call_kwargs = (
mock_db_client.db.litellm_verificationtoken.find_many.call_args[1]
)
assert api_key_call_kwargs["where"] == {"user_id": user_id}
@pytest.mark.asyncio
async def test_get_team_daily_activity_team_admin_sees_all_spend(mock_db_client):
"""
Test that team admin users see all team spend (no API key filtering)
when calling /team/daily/activity endpoint.
"""
from litellm.proxy.management_endpoints.team_endpoints import (
get_team_daily_activity,
)
# Create a team admin user
user_id = "test_admin_123"
team_id = "test_team_456"
user_api_key_dict = UserAPIKeyAuth(
user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER
)
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id=user_id,
teams=[team_id],
max_budget=1000.0,
spend=0.0,
user_email="admin@example.com",
user_role="internal_user",
)
# Mock team with user as admin member
mock_team_member = Member(user_id=user_id, role="admin")
mock_team = MagicMock(spec=LiteLLM_TeamTable)
mock_team.team_id = team_id
mock_team.team_alias = "Test Team"
mock_team.members_with_roles = [mock_team_member]
mock_team.model_dump.return_value = {
"team_id": team_id,
"team_alias": "Test Team",
"members_with_roles": [{"user_id": user_id, "role": "admin"}],
}
# Setup mocks
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(
return_value=[mock_team]
)
# Mock get_user_object
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
new_callable=AsyncMock,
) as mock_get_user_object:
mock_get_user_object.return_value = mock_user_info
# Mock get_daily_activity to capture the api_key parameter
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity:
mock_get_daily_activity.return_value = MagicMock()
# Call the endpoint
await get_team_daily_activity(
team_ids=team_id,
start_date="2024-01-01",
end_date="2024-01-02",
model=None,
api_key=None,
page=1,
page_size=10,
exclude_team_ids=None,
user_api_key_dict=user_api_key_dict,
)
# Verify get_daily_activity was called WITHOUT API key filtering
mock_get_daily_activity.assert_called_once()
call_kwargs = mock_get_daily_activity.call_args[1]
assert call_kwargs["api_key"] is None
assert call_kwargs["entity_id"] == [team_id]
# Verify user's API keys were NOT fetched (since they're admin)
if hasattr(
mock_db_client.db.litellm_verificationtoken, "find_many"
) and mock_db_client.db.litellm_verificationtoken.find_many.called:
# If it was called, that's unexpected for admin users
assert False, "API keys should not be fetched for team admin users"
@pytest.mark.asyncio
async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth):
"""
@ -5083,184 +4901,3 @@ async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth)
# Verify router_settings can be deserialized and matches input
deserialized_settings = json.loads(team_data["router_settings"])
assert deserialized_settings == router_settings_data
@pytest.mark.asyncio
async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys(
mock_db_client,
):
"""
Test that non-team-admin users only see their own spend (filtered by their API keys)
when calling /team/daily/activity endpoint.
"""
from litellm.proxy.management_endpoints.team_endpoints import (
get_team_daily_activity,
)
# Create a non-admin user
user_id = "test_user_123"
team_id = "test_team_456"
user_api_key_dict = UserAPIKeyAuth(
user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER
)
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id=user_id,
teams=[team_id],
max_budget=1000.0,
spend=0.0,
user_email="test@example.com",
user_role="internal_user",
)
# Mock team with user as non-admin member
mock_team_member = Member(user_id=user_id, role="user")
mock_team = MagicMock(spec=LiteLLM_TeamTable)
mock_team.team_id = team_id
mock_team.team_alias = "Test Team"
mock_team.members_with_roles = [mock_team_member]
mock_team.model_dump.return_value = {
"team_id": team_id,
"team_alias": "Test Team",
"members_with_roles": [{"user_id": user_id, "role": "user"}],
}
# Mock user's API keys
user_api_key_1 = MagicMock()
user_api_key_1.token = "user_key_1"
user_api_key_2 = MagicMock()
user_api_key_2.token = "user_key_2"
# Setup mocks
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(
return_value=[mock_team]
)
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[user_api_key_1, user_api_key_2]
)
# Mock get_user_object
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
new_callable=AsyncMock,
) as mock_get_user_object:
mock_get_user_object.return_value = mock_user_info
# Mock get_daily_activity to capture the api_key parameter
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity:
mock_get_daily_activity.return_value = MagicMock()
# Call the endpoint
await get_team_daily_activity(
team_ids=team_id,
start_date="2024-01-01",
end_date="2024-01-02",
model=None,
api_key=None,
page=1,
page_size=10,
exclude_team_ids=None,
user_api_key_dict=user_api_key_dict,
)
# Verify get_daily_activity was called with user's API keys as filter
mock_get_daily_activity.assert_called_once()
call_kwargs = mock_get_daily_activity.call_args[1]
assert call_kwargs["api_key"] == ["user_key_1", "user_key_2"]
assert call_kwargs["entity_id"] == [team_id]
# Verify user's API keys were fetched
mock_db_client.db.litellm_verificationtoken.find_many.assert_called_once()
api_key_call_kwargs = (
mock_db_client.db.litellm_verificationtoken.find_many.call_args[1]
)
assert api_key_call_kwargs["where"] == {"user_id": user_id}
@pytest.mark.asyncio
async def test_get_team_daily_activity_team_admin_sees_all_spend(mock_db_client):
"""
Test that team admin users see all team spend (no API key filtering)
when calling /team/daily/activity endpoint.
"""
from litellm.proxy.management_endpoints.team_endpoints import (
get_team_daily_activity,
)
# Create a team admin user
user_id = "test_admin_123"
team_id = "test_team_456"
user_api_key_dict = UserAPIKeyAuth(
user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER
)
# Mock user info
mock_user_info = LiteLLM_UserTable(
user_id=user_id,
teams=[team_id],
max_budget=1000.0,
spend=0.0,
user_email="admin@example.com",
user_role="internal_user",
)
# Mock team with user as admin member
mock_team_member = Member(user_id=user_id, role="admin")
mock_team = MagicMock(spec=LiteLLM_TeamTable)
mock_team.team_id = team_id
mock_team.team_alias = "Test Team"
mock_team.members_with_roles = [mock_team_member]
mock_team.model_dump.return_value = {
"team_id": team_id,
"team_alias": "Test Team",
"members_with_roles": [{"user_id": user_id, "role": "admin"}],
}
# Setup mocks
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(
return_value=[mock_team]
)
# Mock get_user_object
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
new_callable=AsyncMock,
) as mock_get_user_object:
mock_get_user_object.return_value = mock_user_info
# Mock get_daily_activity to capture the api_key parameter
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity",
new_callable=AsyncMock,
) as mock_get_daily_activity:
mock_get_daily_activity.return_value = MagicMock()
# Call the endpoint
await get_team_daily_activity(
team_ids=team_id,
start_date="2024-01-01",
end_date="2024-01-02",
model=None,
api_key=None,
page=1,
page_size=10,
exclude_team_ids=None,
user_api_key_dict=user_api_key_dict,
)
# Verify get_daily_activity was called WITHOUT API key filtering
mock_get_daily_activity.assert_called_once()
call_kwargs = mock_get_daily_activity.call_args[1]
assert call_kwargs["api_key"] is None
assert call_kwargs["entity_id"] == [team_id]
# Verify user's API keys were NOT fetched (since they're admin)
if hasattr(
mock_db_client.db.litellm_verificationtoken, "find_many"
) and mock_db_client.db.litellm_verificationtoken.find_many.called:
# If it was called, that's unexpected for admin users
assert False, "API keys should not be fetched for team admin users"

View file

@ -0,0 +1,222 @@
import pytest
from unittest.mock import MagicMock, AsyncMock, patch
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import _base_vertex_proxy_route
from litellm.types.router import DeploymentTypedDict
@pytest.mark.asyncio
async def test_vertex_passthrough_load_balancing():
"""
Test that _base_vertex_proxy_route uses llm_router.get_available_deployment_for_pass_through
instead of get_model_list to ensure load balancing works with pass-through filtering.
"""
# Setup mocks
mock_request = MagicMock()
mock_response = MagicMock()
mock_handler = MagicMock()
# Mock the router
mock_router = MagicMock()
mock_deployment = {
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "test-project-lb",
"vertex_location": "us-central1-lb",
"use_in_pass_through": True
}
}
mock_router.get_available_deployment_for_pass_through.return_value = mock_deployment
# Mock get_vertex_model_id_from_url to return a model ID
with patch("litellm.llms.vertex_ai.common_utils.get_vertex_model_id_from_url", return_value="gemini-pro"), \
patch("litellm.proxy.proxy_server.llm_router", mock_router), \
patch("litellm.llms.vertex_ai.common_utils.get_vertex_project_id_from_url", return_value=None), \
patch("litellm.llms.vertex_ai.common_utils.get_vertex_location_from_url", return_value=None), \
patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router") as mock_pt_router, \
patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._prepare_vertex_auth_headers", new_callable=AsyncMock) as mock_prep_headers, \
patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route") as mock_create_route, \
patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth", new_callable=AsyncMock) as mock_auth:
# Setup additional mocks to avoid side effects
mock_pt_router.get_vertex_credentials.return_value = MagicMock()
mock_prep_headers.return_value = ({}, "https://test.url", False, "test-project-lb", "us-central1-lb")
mock_endpoint_func = AsyncMock()
mock_create_route.return_value = mock_endpoint_func
mock_auth.return_value = {}
# Execute
await _base_vertex_proxy_route(
endpoint="https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-pro:streamGenerateContent",
request=mock_request,
fastapi_response=mock_response,
get_vertex_pass_through_handler=mock_handler
)
# Verify
# 1. Check that get_available_deployment_for_pass_through was called with the correct model ID
mock_router.get_available_deployment_for_pass_through.assert_called_once_with(model="gemini-pro")
# 2. Check that get_model_list was NOT called (this ensures we aren't doing the old logic)
mock_router.get_model_list.assert_not_called()
# 3. Verify that the project and location from the deployment were used (passed to _prepare_vertex_auth_headers)
# The args are: request, vertex_credentials, router_credentials, vertex_project, vertex_location, ...
# We check the 4th and 5th args (index 3 and 4)
call_args = mock_prep_headers.call_args
assert call_args[1]['vertex_project'] == "test-project-lb"
assert call_args[1]['vertex_location'] == "us-central1-lb"
def test_get_available_deployment_for_pass_through_filters_correctly():
"""
Test that get_available_deployment_for_pass_through filters deployments correctly
"""
from litellm.router import Router
# Configure router with both pass-through and non-pass-through deployments
model_list = [
{
"model_name": "gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "project-1",
"vertex_location": "us-central1",
"use_in_pass_through": True, # Supports pass-through
}
},
{
"model_name": "gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "project-2",
"vertex_location": "us-west1",
"use_in_pass_through": False, # Does not support pass-through
}
},
{
"model_name": "gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "project-3",
"vertex_location": "us-east1",
# use_in_pass_through not set (defaults to False)
}
},
]
router = Router(model_list=model_list, routing_strategy="simple-shuffle")
# Test: Should only return project-1 (use_in_pass_through=True)
deployment = router.get_available_deployment_for_pass_through(model="gemini-pro")
assert deployment is not None
assert deployment["litellm_params"]["vertex_project"] == "project-1"
assert deployment["litellm_params"]["use_in_pass_through"] is True
def test_get_available_deployment_for_pass_through_no_deployments():
"""
Test that correct error is thrown when there are no pass-through deployments
"""
import litellm
from litellm.router import Router
model_list = [
{
"model_name": "gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "project-1",
"vertex_location": "us-central1",
"use_in_pass_through": False, # Does not support pass-through
}
}
]
router = Router(model_list=model_list)
# Should throw BadRequestError
with pytest.raises(litellm.BadRequestError) as exc_info:
router.get_available_deployment_for_pass_through(model="gemini-pro")
assert "use_in_pass_through=True" in str(exc_info.value)
def test_get_available_deployment_for_pass_through_load_balancing():
"""
Test load balancing for pass-through deployments
"""
from litellm.router import Router
model_list = [
{
"model_name": "gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "project-1",
"vertex_location": "us-central1",
"use_in_pass_through": True,
"rpm": 100,
}
},
{
"model_name": "gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "project-2",
"vertex_location": "us-west1",
"use_in_pass_through": True,
"rpm": 200, # Higher RPM should be selected more frequently
}
},
]
router = Router(
model_list=model_list,
routing_strategy="simple-shuffle"
)
# Call multiple times and track selected deployments
selections = {"project-1": 0, "project-2": 0}
for _ in range(100):
deployment = router.get_available_deployment_for_pass_through(model="gemini-pro")
project = deployment["litellm_params"]["vertex_project"]
selections[project] += 1
# Due to rpm weight, project-2 should be selected more times
assert selections["project-2"] > selections["project-1"]
@pytest.mark.asyncio
async def test_async_get_available_deployment_for_pass_through():
"""
Test the async version of get_available_deployment_for_pass_through
"""
from litellm.router import Router
model_list = [
{
"model_name": "gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "project-1",
"vertex_location": "us-central1",
"use_in_pass_through": True,
}
}
]
router = Router(
model_list=model_list,
routing_strategy="simple-shuffle"
)
deployment = await router.async_get_available_deployment_for_pass_through(
model="gemini-pro",
request_kwargs={}
)
assert deployment is not None
assert deployment["litellm_params"]["use_in_pass_through"] is True

View file

@ -0,0 +1,494 @@
"""
Tests for fallback management endpoints
Tests:
1. Create fallback configuration
2. Get fallback configuration
3. Delete fallback configuration
4. Validation tests (invalid models, duplicate fallbacks, etc.)
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.proxy.management_endpoints.fallback_management_endpoints import (
FallbackCreateRequest,
create_fallback,
delete_fallback,
get_fallback,
)
class TestFallbackCreateRequest:
"""Test the FallbackCreateRequest validation"""
def test_valid_request(self):
"""Test valid fallback request"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4", "claude-3-haiku"],
fallback_type="general",
)
assert request.model == "gpt-3.5-turbo"
assert request.fallback_models == ["gpt-4", "claude-3-haiku"]
assert request.fallback_type == "general"
def test_default_fallback_type(self):
"""Test default fallback type is 'general'"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4"],
)
assert request.fallback_type == "general"
def test_empty_fallback_models(self):
"""Test that empty fallback_models raises validation error"""
with pytest.raises(ValueError, match="at least 1 item"):
FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=[],
)
def test_duplicate_fallback_models(self):
"""Test that duplicate fallback models raise validation error"""
with pytest.raises(ValueError, match="fallback_models must not contain duplicates"):
FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4", "gpt-4"],
)
def test_empty_model_name(self):
"""Test that empty model name raises validation error"""
with pytest.raises(ValueError, match="model must be a non-empty string"):
FallbackCreateRequest(
model="",
fallback_models=["gpt-4"],
)
def test_whitespace_model_name(self):
"""Test that whitespace-only model name raises validation error"""
with pytest.raises(ValueError, match="model must be a non-empty string"):
FallbackCreateRequest(
model=" ",
fallback_models=["gpt-4"],
)
def test_model_name_trimmed(self):
"""Test that model name is trimmed"""
request = FallbackCreateRequest(
model=" gpt-3.5-turbo ",
fallback_models=["gpt-4"],
)
assert request.model == "gpt-3.5-turbo"
def test_context_window_fallback_type(self):
"""Test context_window fallback type"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4-32k"],
fallback_type="context_window",
)
assert request.fallback_type == "context_window"
def test_content_policy_fallback_type(self):
"""Test content_policy fallback type"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4"],
fallback_type="content_policy",
)
assert request.fallback_type == "content_policy"
@pytest.mark.asyncio
class TestCreateFallback:
"""Test the create_fallback endpoint"""
@pytest.fixture
def mock_router(self):
"""Create a mock router"""
router = MagicMock()
router.model_names = {"gpt-3.5-turbo", "gpt-4", "claude-3-haiku"}
router.fallbacks = []
router.context_window_fallbacks = []
router.content_policy_fallbacks = []
return router
@pytest.fixture
def mock_prisma_client(self):
"""Create a mock prisma client"""
client = MagicMock()
client.db.litellm_config.upsert = AsyncMock()
client.jsonify_object = lambda x: x
return client
@pytest.fixture
def mock_proxy_config(self):
"""Create a mock proxy config"""
config = MagicMock()
config.get_config = AsyncMock(return_value={"router_settings": {}})
return config
@pytest.fixture
def mock_user_api_key_dict(self):
"""Create a mock user API key dict"""
return MagicMock()
async def test_create_fallback_success(
self, mock_router, mock_prisma_client, mock_proxy_config, mock_user_api_key_dict
):
"""Test successful fallback creation"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4", "claude-3-haiku"],
fallback_type="general",
)
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router,
), patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
), patch(
"litellm.proxy.proxy_server.proxy_config",
mock_proxy_config,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
True,
):
response = await create_fallback(request, mock_user_api_key_dict)
assert response.model == "gpt-3.5-turbo"
assert response.fallback_models == ["gpt-4", "claude-3-haiku"]
assert response.fallback_type == "general"
assert "created" in response.message.lower() or "updated" in response.message.lower()
# Verify database was updated
mock_prisma_client.db.litellm_config.upsert.assert_called_once()
async def test_create_fallback_router_not_initialized(
self, mock_prisma_client, mock_proxy_config, mock_user_api_key_dict
):
"""Test error when router is not initialized"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4"],
)
with patch(
"litellm.proxy.proxy_server.llm_router",
None,
), pytest.raises(HTTPException) as exc_info:
await create_fallback(request, mock_user_api_key_dict)
assert exc_info.value.status_code == 500
assert "Router not initialized" in str(exc_info.value.detail)
async def test_create_fallback_model_not_found(
self, mock_router, mock_prisma_client, mock_proxy_config, mock_user_api_key_dict
):
"""Test error when model is not found in router"""
request = FallbackCreateRequest(
model="invalid-model",
fallback_models=["gpt-4"],
)
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router,
), patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
True,
), pytest.raises(HTTPException) as exc_info:
await create_fallback(request, mock_user_api_key_dict)
assert exc_info.value.status_code == 404
assert "not found in router" in str(exc_info.value.detail)
async def test_create_fallback_invalid_fallback_model(
self, mock_router, mock_prisma_client, mock_proxy_config, mock_user_api_key_dict
):
"""Test error when fallback model is not found in router"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["invalid-fallback-model"],
)
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router,
), patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
True,
), pytest.raises(HTTPException) as exc_info:
await create_fallback(request, mock_user_api_key_dict)
assert exc_info.value.status_code == 400
assert "Invalid fallback models" in str(exc_info.value.detail)
async def test_create_fallback_model_is_own_fallback(
self, mock_router, mock_prisma_client, mock_proxy_config, mock_user_api_key_dict
):
"""Test error when model is its own fallback"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-3.5-turbo", "gpt-4"],
)
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router,
), patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
True,
), pytest.raises(HTTPException) as exc_info:
await create_fallback(request, mock_user_api_key_dict)
assert exc_info.value.status_code == 400
assert "cannot be its own fallback" in str(exc_info.value.detail)
async def test_create_fallback_db_not_enabled(
self, mock_router, mock_user_api_key_dict
):
"""Test error when database storage is not enabled"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4"],
)
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
False,
), pytest.raises(HTTPException) as exc_info:
await create_fallback(request, mock_user_api_key_dict)
assert exc_info.value.status_code == 400
assert "Database storage not enabled" in str(exc_info.value.detail)
async def test_create_fallback_context_window_type(
self, mock_router, mock_prisma_client, mock_proxy_config, mock_user_api_key_dict
):
"""Test creating context_window fallback"""
request = FallbackCreateRequest(
model="gpt-3.5-turbo",
fallback_models=["gpt-4"],
fallback_type="context_window",
)
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router,
), patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
), patch(
"litellm.proxy.proxy_server.proxy_config",
mock_proxy_config,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
True,
):
response = await create_fallback(request, mock_user_api_key_dict)
assert response.fallback_type == "context_window"
# Verify the correct attribute was updated
assert hasattr(mock_router, "context_window_fallbacks")
@pytest.mark.asyncio
class TestGetFallback:
"""Test the get_fallback endpoint"""
@pytest.fixture
def mock_router_with_fallbacks(self):
"""Create a mock router with fallbacks configured"""
router = MagicMock()
router.fallbacks = [{"gpt-3.5-turbo": ["gpt-4", "claude-3-haiku"]}]
router.context_window_fallbacks = []
router.content_policy_fallbacks = []
return router
@pytest.fixture
def mock_user_api_key_dict(self):
"""Create a mock user API key dict"""
return MagicMock()
async def test_get_fallback_success(
self, mock_router_with_fallbacks, mock_user_api_key_dict
):
"""Test successful fallback retrieval"""
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router_with_fallbacks,
):
response = await get_fallback(
"gpt-3.5-turbo", "general", mock_user_api_key_dict
)
assert response.model == "gpt-3.5-turbo"
assert response.fallback_models == ["gpt-4", "claude-3-haiku"]
assert response.fallback_type == "general"
async def test_get_fallback_not_found(
self, mock_router_with_fallbacks, mock_user_api_key_dict
):
"""Test error when fallback is not found"""
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router_with_fallbacks,
), pytest.raises(HTTPException) as exc_info:
await get_fallback("gpt-4", "general", mock_user_api_key_dict)
assert exc_info.value.status_code == 404
assert "No general fallbacks configured" in str(exc_info.value.detail)
async def test_get_fallback_router_not_initialized(self, mock_user_api_key_dict):
"""Test error when router is not initialized"""
with patch(
"litellm.proxy.proxy_server.llm_router",
None,
), pytest.raises(HTTPException) as exc_info:
await get_fallback("gpt-3.5-turbo", "general", mock_user_api_key_dict)
assert exc_info.value.status_code == 500
assert "Router not initialized" in str(exc_info.value.detail)
@pytest.mark.asyncio
class TestDeleteFallback:
"""Test the delete_fallback endpoint"""
@pytest.fixture
def mock_router_with_fallbacks(self):
"""Create a mock router with fallbacks configured"""
router = MagicMock()
router.fallbacks = [{"gpt-3.5-turbo": ["gpt-4", "claude-3-haiku"]}]
router.context_window_fallbacks = []
router.content_policy_fallbacks = []
return router
@pytest.fixture
def mock_prisma_client(self):
"""Create a mock prisma client"""
client = MagicMock()
client.db.litellm_config.upsert = AsyncMock()
client.jsonify_object = lambda x: x
return client
@pytest.fixture
def mock_proxy_config(self):
"""Create a mock proxy config"""
config = MagicMock()
config.get_config = AsyncMock(
return_value={
"router_settings": {
"fallbacks": [{"gpt-3.5-turbo": ["gpt-4", "claude-3-haiku"]}]
}
}
)
return config
@pytest.fixture
def mock_user_api_key_dict(self):
"""Create a mock user API key dict"""
return MagicMock()
async def test_delete_fallback_success(
self,
mock_router_with_fallbacks,
mock_prisma_client,
mock_proxy_config,
mock_user_api_key_dict,
):
"""Test successful fallback deletion"""
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router_with_fallbacks,
), patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
), patch(
"litellm.proxy.proxy_server.proxy_config",
mock_proxy_config,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
True,
):
response = await delete_fallback(
"gpt-3.5-turbo", "general", mock_user_api_key_dict
)
assert response.model == "gpt-3.5-turbo"
assert response.fallback_type == "general"
assert "deleted" in response.message.lower()
# Verify database was updated
mock_prisma_client.db.litellm_config.upsert.assert_called_once()
async def test_delete_fallback_not_found(
self,
mock_router_with_fallbacks,
mock_prisma_client,
mock_proxy_config,
mock_user_api_key_dict,
):
"""Test error when fallback to delete is not found"""
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router_with_fallbacks,
), patch(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,
), patch(
"litellm.proxy.proxy_server.proxy_config",
mock_proxy_config,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
True,
), pytest.raises(HTTPException) as exc_info:
await delete_fallback("gpt-4", "general", mock_user_api_key_dict)
assert exc_info.value.status_code == 404
assert "No general fallbacks configured" in str(exc_info.value.detail)
async def test_delete_fallback_router_not_initialized(self, mock_user_api_key_dict):
"""Test error when router is not initialized"""
with patch(
"litellm.proxy.proxy_server.llm_router",
None,
), pytest.raises(HTTPException) as exc_info:
await delete_fallback("gpt-3.5-turbo", "general", mock_user_api_key_dict)
assert exc_info.value.status_code == 500
assert "Router not initialized" in str(exc_info.value.detail)
async def test_delete_fallback_db_not_enabled(
self, mock_router_with_fallbacks, mock_user_api_key_dict
):
"""Test error when database storage is not enabled"""
with patch(
"litellm.proxy.proxy_server.llm_router",
mock_router_with_fallbacks,
), patch(
"litellm.proxy.proxy_server.store_model_in_db",
False,
), pytest.raises(HTTPException) as exc_info:
await delete_fallback("gpt-3.5-turbo", "general", mock_user_api_key_dict)
assert exc_info.value.status_code == 400
assert "Database storage not enabled" in str(exc_info.value.detail)

View file

@ -483,6 +483,75 @@ class TestProxyInitializationHelpers:
# Verify that uvicorn.run was called again
mock_uvicorn_run.assert_called_once()
@patch("litellm.proxy.proxy_cli.ProxyInitializationHelpers._run_gunicorn_server")
@patch("builtins.print")
def test_gunicorn_keepalive_timeout_flag(self, mock_print, mock_gunicorn):
"""Test that the keepalive_timeout flag is properly passed to Gunicorn"""
from click.testing import CliRunner
from litellm.proxy.proxy_cli import run_server
runner = CliRunner()
mock_app = MagicMock()
mock_proxy_config = MagicMock()
mock_key_mgmt = MagicMock()
mock_save_worker_config = MagicMock()
with patch.dict(
"sys.modules",
{
"proxy_server": MagicMock(
app=mock_app,
ProxyConfig=mock_proxy_config,
KeyManagementSettings=mock_key_mgmt,
save_worker_config=mock_save_worker_config,
)
},
):
result = runner.invoke(
run_server, ["--local", "--run_gunicorn", "--keepalive_timeout", "120"]
)
assert result.exit_code == 0
# Verify _run_gunicorn_server was called with keepalive_timeout
mock_gunicorn.assert_called_once()
call_kwargs = mock_gunicorn.call_args.kwargs
assert call_kwargs["keepalive_timeout"] == 120
@patch("litellm.proxy.proxy_cli.ProxyInitializationHelpers._run_gunicorn_server")
@patch("builtins.print")
def test_gunicorn_keepalive_default(self, mock_print, mock_gunicorn):
"""Test that Gunicorn uses default 90s when keepalive_timeout not specified"""
from click.testing import CliRunner
from litellm.proxy.proxy_cli import run_server
runner = CliRunner()
mock_app = MagicMock()
mock_proxy_config = MagicMock()
mock_key_mgmt = MagicMock()
mock_save_worker_config = MagicMock()
with patch.dict(
"sys.modules",
{
"proxy_server": MagicMock(
app=mock_app,
ProxyConfig=mock_proxy_config,
KeyManagementSettings=mock_key_mgmt,
save_worker_config=mock_save_worker_config,
)
},
):
result = runner.invoke(run_server, ["--local", "--run_gunicorn"])
assert result.exit_code == 0
# Verify default behavior (keepalive_timeout is None, Gunicorn will use 90)
call_kwargs = mock_gunicorn.call_args.kwargs
assert call_kwargs.get("keepalive_timeout") is None
class TestHealthAppFactory:
"""Test cases for the health app factory module"""

View file

@ -10,6 +10,114 @@ import pytest
from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup
def test_spend_log_cleanup_cron_scheduling():
"""Test that cron expressions are correctly parsed for spend log cleanup scheduling"""
from apscheduler.triggers.cron import CronTrigger
# Valid cron expressions
cron_expr = "0 4 * * *" # 4:00 AM daily
trigger = CronTrigger.from_crontab(cron_expr)
assert trigger is not None
# Every minute (useful for testing)
trigger_minute = CronTrigger.from_crontab("*/1 * * * *")
assert trigger_minute is not None
# Specific day and hour
trigger_weekly = CronTrigger.from_crontab("0 3 * * 0") # 3 AM every Sunday
assert trigger_weekly is not None
# Invalid cron expression should raise ValueError
with pytest.raises(ValueError):
CronTrigger.from_crontab("invalid cron")
with pytest.raises(ValueError):
CronTrigger.from_crontab("60 25 * * *") # Invalid minute and hour
def test_spend_log_cleanup_cron_scheduler_integration():
"""
Integration test: Verify the proxy_server scheduler logic correctly adds
cron-based cleanup job when maximum_spend_logs_cleanup_cron is configured.
This tests the logic in proxy_server.py lines 4671-4717 without requiring
a real database connection.
"""
from unittest.mock import MagicMock
from apscheduler.triggers.cron import CronTrigger
# Mock scheduler
mock_scheduler = MagicMock()
mock_prisma_client = MagicMock()
mock_cleanup_instance = MagicMock()
# Test Case 1: Cron-based scheduling
general_settings_cron = {
"maximum_spend_logs_retention_period": "7d",
"maximum_spend_logs_cleanup_cron": "0 4 * * *", # 4 AM daily
}
cleanup_cron = general_settings_cron.get("maximum_spend_logs_cleanup_cron")
assert cleanup_cron is not None
# Simulate the scheduler logic from proxy_server.py
cron_trigger = CronTrigger.from_crontab(cleanup_cron)
mock_scheduler.add_job(
mock_cleanup_instance.cleanup_old_spend_logs,
cron_trigger,
args=[mock_prisma_client],
id="spend_log_cleanup_job",
replace_existing=True,
misfire_grace_time=3600,
)
# Verify scheduler was called correctly
mock_scheduler.add_job.assert_called_once()
call_args = mock_scheduler.add_job.call_args
# Verify the trigger is a CronTrigger
assert isinstance(call_args[0][1], CronTrigger)
# Verify job ID
assert call_args[1]["id"] == "spend_log_cleanup_job"
assert call_args[1]["replace_existing"] is True
# Test Case 2: Interval-based scheduling (fallback)
mock_scheduler.reset_mock()
general_settings_interval = {
"maximum_spend_logs_retention_period": "7d",
# No cron, so it should fall back to interval
}
cleanup_cron_fallback = general_settings_interval.get(
"maximum_spend_logs_cleanup_cron"
)
assert cleanup_cron_fallback is None # No cron configured
# Simulate interval-based scheduling fallback
retention_interval = general_settings_interval.get(
"maximum_spend_logs_retention_interval", "1d"
)
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
interval_seconds = duration_in_seconds(retention_interval)
mock_scheduler.add_job(
mock_cleanup_instance.cleanup_old_spend_logs,
"interval",
seconds=interval_seconds,
args=[mock_prisma_client],
id="spend_log_cleanup_job",
replace_existing=True,
)
# Verify interval scheduling was called
mock_scheduler.add_job.assert_called_once()
interval_call_args = mock_scheduler.add_job.call_args
assert interval_call_args[0][1] == "interval"
assert interval_call_args[1]["seconds"] == 86400 # 1 day in seconds
@pytest.mark.asyncio
async def test_should_delete_spend_logs():
# Test case 1: No retention set

View file

@ -2066,3 +2066,190 @@ async def test_aguardrail():
assert result["result"] == "success"
assert result["selected_guardrail"]["id"] == "guardrail-1"
def test_resolve_model_name_from_model_id_wildcard_pattern():
"""
Test that resolve_model_name_from_model_id correctly resolves model names
for wildcard patterns using PatternMatchRouter.
This is critical for video status/content endpoints where model_id extracted
from video_id (e.g., "veo-3.0-generate-preview") needs to match wildcard
patterns like "vertex_ai/*" to inject credentials from the model config.
"""
# Set up router with wildcard pattern
router = litellm.Router(
model_list=[
{
"model_name": "vertex_ai/*",
"litellm_params": {
"model": "vertex_ai/*",
"vertex_project": "test-project",
"vertex_location": "us-central1",
},
},
{
"model_name": "specific-model",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "specific-project",
"vertex_location": "us-east1",
},
},
],
)
# Test Case 1: Wildcard pattern matching with custom_llm_provider
# This simulates video_id like "vertex_ai:veo-3.0-generate-preview:..."
result = router.resolve_model_name_from_model_id(
model_id="veo-3.0-generate-preview",
custom_llm_provider="vertex_ai",
)
assert result == "vertex_ai/*", f"Expected 'vertex_ai/*', got '{result}'"
# Test Case 2: Different model name should also match wildcard
result = router.resolve_model_name_from_model_id(
model_id="gemini-2.0-flash",
custom_llm_provider="vertex_ai",
)
assert result == "vertex_ai/*", f"Expected 'vertex_ai/*', got '{result}'"
# Test Case 3: Without custom_llm_provider, should not match wildcard
result = router.resolve_model_name_from_model_id(
model_id="veo-3.0-generate-preview",
custom_llm_provider=None,
)
assert result is None, f"Expected None without provider, got '{result}'"
# Test Case 4: Exact model_name match should take precedence
result = router.resolve_model_name_from_model_id(
model_id="specific-model",
custom_llm_provider="vertex_ai",
)
assert result == "specific-model", f"Expected 'specific-model', got '{result}'"
def test_resolve_model_name_from_model_id_exact_match():
"""
Test that resolve_model_name_from_model_id correctly resolves exact model names.
"""
router = litellm.Router(
model_list=[
{
"model_name": "my-gpt-model",
"litellm_params": {
"model": "azure/gpt-4",
"api_key": "test-key",
},
},
{
"model_name": "veo-model",
"litellm_params": {
"model": "vertex_ai/veo-2.0-generate-001",
"vertex_project": "test-project",
},
},
],
)
# Test Case 1: Direct model_name match
result = router.resolve_model_name_from_model_id(model_id="my-gpt-model")
assert result == "my-gpt-model", f"Expected 'my-gpt-model', got '{result}'"
# Test Case 2: Match by litellm_params.model suffix
result = router.resolve_model_name_from_model_id(model_id="veo-2.0-generate-001")
assert result == "veo-model", f"Expected 'veo-model', got '{result}'"
# Test Case 3: Non-existent model should return None
result = router.resolve_model_name_from_model_id(model_id="non-existent-model")
assert result is None, f"Expected None, got '{result}'"
def test_resolve_model_name_from_model_id_provider_prefix():
"""
Test that resolve_model_name_from_model_id handles provider prefix correctly.
"""
router = litellm.Router(
model_list=[
{
"model_name": "vertex_ai/gemini-pro",
"litellm_params": {
"model": "vertex_ai/gemini-pro",
"vertex_project": "test-project",
},
},
],
)
# Test Case 1: Full model name with provider prefix as model_name
result = router.resolve_model_name_from_model_id(
model_id="vertex_ai/gemini-pro",
custom_llm_provider=None,
)
assert result == "vertex_ai/gemini-pro", f"Expected 'vertex_ai/gemini-pro', got '{result}'"
# Test Case 2: Model ID with provider prefix constructed from custom_llm_provider
result = router.resolve_model_name_from_model_id(
model_id="gemini-pro",
custom_llm_provider="vertex_ai",
)
assert result == "vertex_ai/gemini-pro", f"Expected 'vertex_ai/gemini-pro', got '{result}'"
def test_resolve_model_name_from_model_id_multiple_wildcards():
"""
Test that resolve_model_name_from_model_id works with multiple wildcard patterns.
"""
router = litellm.Router(
model_list=[
{
"model_name": "vertex_ai/*",
"litellm_params": {
"model": "vertex_ai/*",
"vertex_project": "vertex-project",
},
},
{
"model_name": "openai/*",
"litellm_params": {
"model": "openai/*",
"api_key": "openai-key",
},
},
{
"model_name": "anthropic/*",
"litellm_params": {
"model": "anthropic/*",
"api_key": "anthropic-key",
},
},
],
)
# Test Case 1: Match vertex_ai wildcard
result = router.resolve_model_name_from_model_id(
model_id="veo-3.0-generate-preview",
custom_llm_provider="vertex_ai",
)
assert result == "vertex_ai/*", f"Expected 'vertex_ai/*', got '{result}'"
# Test Case 2: Match openai wildcard
result = router.resolve_model_name_from_model_id(
model_id="gpt-4o",
custom_llm_provider="openai",
)
assert result == "openai/*", f"Expected 'openai/*', got '{result}'"
# Test Case 3: Match anthropic wildcard
result = router.resolve_model_name_from_model_id(
model_id="claude-3-opus",
custom_llm_provider="anthropic",
)
assert result == "anthropic/*", f"Expected 'anthropic/*', got '{result}'"
# Test Case 4: Non-matching provider should return None
result = router.resolve_model_name_from_model_id(
model_id="some-model",
custom_llm_provider="bedrock",
)
assert result is None, f"Expected None for non-matching provider, got '{result}'"

View file

@ -1,9 +1,23 @@
import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import { modelInfoCall, modelHubCall } from "@/components/networking";
import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking";
import useAuthorized from "../useAuthorized";
export interface ProxyModel {
id: string;
object: string;
created: number;
owned_by: string;
}
export interface AllProxyModelsResponse {
data: ProxyModel[];
}
const modelKeys = createQueryKeys("models");
const modelHubKeys = createQueryKeys("modelHub");
const allProxyModelsKeys = createQueryKeys("allProxyModels");
const selectedTeamModelsKeys = createQueryKeys("selectedTeamModels");
export const useModelsInfo = () => {
const { accessToken, userId, userRole } = useAuthorized();
@ -27,3 +41,21 @@ export const useModelHub = () => {
enabled: Boolean(accessToken),
});
};
export const useAllProxyModels = () => {
const { accessToken, userId, userRole } = useAuthorized();
return useQuery<AllProxyModelsResponse>({
queryKey: allProxyModelsKeys.list({}),
queryFn: async () => await modelAvailableCall(accessToken!, userId!, userRole!, true),
enabled: Boolean(accessToken && userId && userRole),
});
};
export const useSelectedTeamModels = (teamID: string | null) => {
const { accessToken, userId, userRole } = useAuthorized();
return useQuery<AllProxyModelsResponse>({
queryKey: selectedTeamModelsKeys.list({}),
queryFn: async () => await modelAvailableCall(accessToken!, userId!, userRole!, true, teamID!),
enabled: Boolean(accessToken && userId && userRole && teamID),
});
};

View file

@ -1,10 +1,9 @@
import { useQuery, UseQueryResult } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import { organizationListCall, Organization } from "@/components/networking";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { Organization, organizationInfoCall, organizationListCall } from "@/components/networking";
import { useQuery, useQueryClient, UseQueryResult } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
const organizationKeys = createQueryKeys("organizations");
export const useOrganizations = (): UseQueryResult<Organization[]> => {
const { accessToken, userId, userRole } = useAuthorized();
return useQuery<Organization[]>({
@ -13,3 +12,28 @@ export const useOrganizations = (): UseQueryResult<Organization[]> => {
enabled: Boolean(accessToken && userId && userRole),
});
};
export const useOrganization = (organizationID?: string) => {
const queryClient = useQueryClient();
const { accessToken } = useAuthorized();
return useQuery<Organization>({
queryKey: organizationKeys.detail(organizationID!),
enabled: Boolean(accessToken && organizationID),
queryFn: async () => {
if (!accessToken || !organizationID) {
throw new Error("Missing auth or teamId");
}
return organizationInfoCall(accessToken, organizationID);
},
initialData: () => {
if (!organizationID) return undefined;
const organizations = queryClient.getQueryData<Organization[]>(organizationKeys.list({}));
return organizations?.find((organization: Organization) => organization.organization_id === organizationID);
},
});
};

Some files were not shown because too many files have changed in this diff Show more