mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Merge remote-tracking branch 'origin' into litellm_deleted_keys_team
This commit is contained in:
commit
33ff58b70a
111 changed files with 7483 additions and 1921 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
468
docs/my-website/docs/completion/message_sanitization.md
Normal file
468
docs/my-website/docs/completion/message_sanitization.md
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
```
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
273
docs/my-website/docs/proxy/fallback_management.md
Normal file
273
docs/my-website/docs/proxy/fallback_management.md
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
### 3. View Usage in LiteLLM UI
|
||||
|
||||
Navigate to the **Logs** tab in the LiteLLM UI.
|
||||
|
||||

|
||||
|
||||
Click on a request to see details.
|
||||
|
||||

|
||||
|
||||
Filter by customer ID to see all requests for that customer.
|
||||
|
||||

|
||||
|
||||
## 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)
|
||||
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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}],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]] = []
|
||||
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 @@
|
|||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)}"},
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -26,3 +26,5 @@ if exit_code != 0:
|
|||
verbose_proxy_logger.error(
|
||||
f"'prisma generate' stderr: {result.stderr}"
|
||||
) # Log stderr
|
||||
|
||||
sys.exit(exit_code)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
36
litellm/types/llms/anthropic_tool_search.py
Normal file
36
litellm/types/llms/anthropic_tool_search.py
Normal 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
|
||||
)
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
38
poetry.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
||||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
380
tests/test_litellm/llms/anthropic/test_message_sanitization.py
Normal file
380
tests/test_litellm/llms/anthropic/test_message_sanitization.py
Normal 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"])
|
||||
|
|
@ -440,6 +440,7 @@ def test_select_azure_base_url_called(setup_mocks):
|
|||
"asearch",
|
||||
"avector_store_create",
|
||||
"avector_store_search",
|
||||
"acreate_skill",
|
||||
]
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
494
tests/test_litellm/proxy/test_fallback_management_endpoints.py
Normal file
494
tests/test_litellm/proxy/test_fallback_management_endpoints.py
Normal 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)
|
||||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}'"
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue