Merge remote-tracking branch 'origin' into litellm_key_team_routing_3

This commit is contained in:
yuneng-jiang 2026-01-08 10:39:12 -08:00
commit 1b9c7deec6
79 changed files with 3676 additions and 414 deletions

View file

@ -11,134 +11,72 @@ jobs:
permissions:
issues: write
steps:
- name: Add SDK label
if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nSDK (litellm Python package)')
- name: Add component labels
uses: actions/github-script@v7
with:
github-token: ${{ secrets.GITHUB_TOKEN }}
script: |
const labelName = 'sdk';
try {
await github.rest.issues.getLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName
});
} catch (error) {
if (error.status === 404) {
await github.rest.issues.createLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName,
color: '0E7C86',
description: 'Issues related to the litellm Python SDK'
});
} else {
throw error;
}
}
await github.rest.issues.addLabels({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: context.issue.number,
labels: [labelName]
});
const body = context.payload.issue.body;
if (!body) return;
- name: Add Proxy label
if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nProxy')
uses: actions/github-script@v7
with:
github-token: ${{ secrets.GITHUB_TOKEN }}
script: |
const labelName = 'proxy';
try {
await github.rest.issues.getLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName
});
} catch (error) {
if (error.status === 404) {
await github.rest.issues.createLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName,
color: '5319E7',
description: 'Issues related to the LiteLLM Proxy'
});
} else {
throw error;
// Define component mappings with regex patterns that handle flexible whitespace
const components = [
{
pattern: /What part of LiteLLM is this about\?\s*SDK \(litellm Python package\)/,
label: 'sdk',
color: '0E7C86',
description: 'Issues related to the litellm Python SDK'
},
{
pattern: /What part of LiteLLM is this about\?\s*Proxy/,
label: 'proxy',
color: '5319E7',
description: 'Issues related to the LiteLLM Proxy'
},
{
pattern: /What part of LiteLLM is this about\?\s*UI Dashboard/,
label: 'ui-dashboard',
color: 'D876E3',
description: 'Issues related to the LiteLLM UI Dashboard'
},
{
pattern: /What part of LiteLLM is this about\?\s*Docs/,
label: 'docs',
color: 'FBCA04',
description: 'Issues related to LiteLLM documentation'
}
}
await github.rest.issues.addLabels({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: context.issue.number,
labels: [labelName]
});
];
- name: Add UI Dashboard label
if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nUI Dashboard')
uses: actions/github-script@v7
with:
github-token: ${{ secrets.GITHUB_TOKEN }}
script: |
const labelName = 'ui-dashboard';
try {
await github.rest.issues.getLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName
});
} catch (error) {
if (error.status === 404) {
await github.rest.issues.createLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName,
color: 'D876E3',
description: 'Issues related to the LiteLLM UI Dashboard'
});
} else {
throw error;
}
}
await github.rest.issues.addLabels({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: context.issue.number,
labels: [labelName]
});
// Find matching component
for (const component of components) {
if (component.pattern.test(body)) {
// Ensure label exists
try {
await github.rest.issues.getLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: component.label
});
} catch (error) {
if (error.status === 404) {
await github.rest.issues.createLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: component.label,
color: component.color,
description: component.description
});
}
}
- name: Add Docs label
if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nDocs')
uses: actions/github-script@v7
with:
github-token: ${{ secrets.GITHUB_TOKEN }}
script: |
const labelName = 'docs';
try {
await github.rest.issues.getLabel({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName
});
} catch (error) {
if (error.status === 404) {
await github.rest.issues.createLabel({
// Add label to issue
await github.rest.issues.addLabels({
owner: context.repo.owner,
repo: context.repo.repo,
name: labelName,
color: 'FBCA04',
description: 'Issues related to LiteLLM documentation'
issue_number: context.issue.number,
labels: [component.label]
});
} else {
throw error;
break;
}
}
await github.rest.issues.addLabels({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: context.issue.number,
labels: [labelName]
});

View file

@ -1692,9 +1692,9 @@ Assistant:
```
## Usage - PDF
## Usage - PDF
Pass base64 encoded PDF files to Anthropic models using the `image_url` field.
Pass base64 encoded PDF files to Anthropic models using the `file` content type with a `file_data` field.
<Tabs>
<TabItem value="sdk" label="SDK">

View file

@ -7,7 +7,7 @@ ALL Bedrock models (Anthropic, Meta, Deepseek, Mistral, Amazon, etc.) are Suppor
| Property | Details |
|-------|-------|
| Description | Amazon Bedrock is a fully managed service that offers a choice of high-performing foundation models (FMs). |
| Provider Route on LiteLLM | `bedrock/`, [`bedrock/converse/`](#set-converse--invoke-route), [`bedrock/invoke/`](#set-invoke-route), [`bedrock/converse_like/`](#calling-via-internal-proxy), [`bedrock/llama/`](#deepseek-not-r1), [`bedrock/deepseek_r1/`](#deepseek-r1), [`bedrock/qwen3/`](#qwen3-imported-models), [`bedrock/qwen2/`](./bedrock_imported.md#qwen2-imported-models), [`bedrock/openai/`](./bedrock_imported.md#openai-compatible-imported-models-qwen-25-vl-etc) |
| Provider Route on LiteLLM | `bedrock/`, [`bedrock/converse/`](#set-converse--invoke-route), [`bedrock/invoke/`](#set-invoke-route), [`bedrock/converse_like/`](#calling-via-internal-proxy), [`bedrock/llama/`](#deepseek-not-r1), [`bedrock/deepseek_r1/`](#deepseek-r1), [`bedrock/qwen3/`](#qwen3-imported-models), [`bedrock/qwen2/`](./bedrock_imported.md#qwen2-imported-models), [`bedrock/openai/`](./bedrock_imported.md#openai-compatible-imported-models-qwen-25-vl-etc), [`bedrock/moonshot`](./bedrock_imported.md#moonshot-kimi-k2-thinking) |
| Provider Doc | [Amazon Bedrock ↗](https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html) |
| Supported OpenAI Endpoints | `/chat/completions`, `/completions`, `/embeddings`, `/images/generations` |
| Rerank Endpoint | `/rerank` |
@ -1941,6 +1941,7 @@ Here's an example of using a bedrock model with LiteLLM. For a complete list, re
| Mixtral 8x7B Instruct | `completion(model='bedrock/mistral.mixtral-8x7b-instruct-v0:1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
| TwelveLabs Pegasus 1.2 (US) | `completion(model='bedrock/us.twelvelabs.pegasus-1-2-v1:0', messages=messages, mediaSource={...})` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
| TwelveLabs Pegasus 1.2 (EU) | `completion(model='bedrock/eu.twelvelabs.pegasus-1-2-v1:0', messages=messages, mediaSource={...})` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
| Moonshot Kimi K2 Thinking | `completion(model='bedrock/moonshot.kimi-k2-thinking', messages=messages)` or `completion(model='bedrock/invoke/moonshot.kimi-k2-thinking', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
## Bedrock Embedding

View file

@ -431,4 +431,180 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
"max_tokens": 300,
"temperature": 0.5
}'
```
```
### Moonshot Kimi K2 Thinking
Moonshot AI's Kimi K2 Thinking model is now available on Amazon Bedrock. This model features advanced reasoning capabilities with automatic reasoning content extraction.
| Property | Details |
|----------|---------|
| Provider Route | `bedrock/moonshot.kimi-k2-thinking`, `bedrock/invoke/moonshot.kimi-k2-thinking` |
| Provider Documentation | [AWS Bedrock Moonshot Announcement ↗](https://aws.amazon.com/about-aws/whats-new/2025/12/amazon-bedrock-fully-managed-open-weight-models/) |
| Supported Parameters | `temperature`, `max_tokens`, `top_p`, `stream`, `tools`, `tool_choice` |
| Special Features | Reasoning content extraction, Tool calling |
#### Supported Features
- **Reasoning Content Extraction**: Automatically extracts `<reasoning>` tags and returns them as `reasoning_content` (similar to OpenAI's o1 models)
- **Tool Calling**: Full support for function/tool calling with tool responses
- **Streaming**: Both streaming and non-streaming responses
- **System Messages**: System message support
#### Basic Usage
<Tabs>
<TabItem value="sdk" label="SDK">
```python title="Moonshot Kimi K2 SDK Usage" showLineNumbers
from litellm import completion
import os
os.environ["AWS_ACCESS_KEY_ID"] = "your-aws-access-key"
os.environ["AWS_SECRET_ACCESS_KEY"] = "your-aws-secret-key"
os.environ["AWS_REGION_NAME"] = "us-west-2" # or your preferred region
# Basic completion
response = completion(
model="bedrock/moonshot.kimi-k2-thinking", # or bedrock/invoke/moonshot.kimi-k2-thinking
messages=[
{"role": "user", "content": "What is 2+2? Think step by step."}
],
temperature=0.7,
max_tokens=200
)
print(response.choices[0].message.content)
# Access reasoning content if present
if response.choices[0].message.reasoning_content:
print("Reasoning:", response.choices[0].message.reasoning_content)
```
</TabItem>
<TabItem value="proxy" label="Proxy">
**1. Add to config**
```yaml title="config.yaml" showLineNumbers
model_list:
- model_name: kimi-k2
litellm_params:
model: bedrock/moonshot.kimi-k2-thinking
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-west-2
```
**2. Start proxy**
```bash title="Start LiteLLM Proxy" showLineNumbers
litellm --config /path/to/config.yaml
# RUNNING at http://0.0.0.0:4000
```
**3. Test it!**
```bash title="Test Kimi K2 via Proxy" showLineNumbers
curl --location 'http://0.0.0.0:4000/chat/completions' \
--header 'Authorization: Bearer sk-1234' \
--header 'Content-Type: application/json' \
--data '{
"model": "kimi-k2",
"messages": [
{
"role": "user",
"content": "What is 2+2? Think step by step."
}
],
"temperature": 0.7,
"max_tokens": 200
}'
```
</TabItem>
</Tabs>
#### Tool Calling Example
```python title="Kimi K2 with Tool Calling" showLineNumbers
from litellm import completion
import os
os.environ["AWS_ACCESS_KEY_ID"] = "your-aws-access-key"
os.environ["AWS_SECRET_ACCESS_KEY"] = "your-aws-secret-key"
os.environ["AWS_REGION_NAME"] = "us-west-2"
# Tool calling example
response = completion(
model="bedrock/moonshot.kimi-k2-thinking",
messages=[
{"role": "user", "content": "What's the weather in Tokyo?"}
],
tools=[
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather in a location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city name"
}
},
"required": ["location"]
}
}
}
]
)
if response.choices[0].message.tool_calls:
tool_call = response.choices[0].message.tool_calls[0]
print(f"Tool called: {tool_call.function.name}")
print(f"Arguments: {tool_call.function.arguments}")
```
#### Streaming Example
```python title="Kimi K2 Streaming" showLineNumbers
from litellm import completion
import os
os.environ["AWS_ACCESS_KEY_ID"] = "your-aws-access-key"
os.environ["AWS_SECRET_ACCESS_KEY"] = "your-aws-secret-key"
os.environ["AWS_REGION_NAME"] = "us-west-2"
response = completion(
model="bedrock/moonshot.kimi-k2-thinking",
messages=[
{"role": "user", "content": "Explain quantum computing in simple terms."}
],
stream=True,
temperature=0.7
)
for chunk in response:
if chunk.choices[0].delta.content:
print(chunk.choices[0].delta.content, end="")
# Check for reasoning content in streaming
if hasattr(chunk.choices[0].delta, 'reasoning_content') and chunk.choices[0].delta.reasoning_content:
print(f"\n[Reasoning: {chunk.choices[0].delta.reasoning_content}]")
```
#### Supported Parameters
| Parameter | Type | Description | Supported |
|-----------|------|-------------|-----------|
| `temperature` | float (0-1) | Controls randomness in output | ✅ |
| `max_tokens` | integer | Maximum tokens to generate | ✅ |
| `top_p` | float | Nucleus sampling parameter | ✅ |
| `stream` | boolean | Enable streaming responses | ✅ |
| `tools` | array | Tool/function definitions | ✅ |
| `tool_choice` | string/object | Tool choice specification | ✅ |
| `stop` | array | Stop sequences | ❌ (Not supported on Bedrock) |

View file

@ -0,0 +1,194 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Manus
Use Manus AI agents through LiteLLM's OpenAI-compatible Responses API.
| Property | Details |
|----------|---------|
| Description | Manus is an AI agent platform for complex reasoning tasks, document analysis, and multi-step workflows with asynchronous task execution. |
| Provider Route on LiteLLM | `manus/{agent_profile}` |
| Supported Operations | `/responses` (Responses API) |
| Provider Doc | [Manus API ↗](https://open.manus.im/docs/openai-compatibility) |
## Model Format
```shell
manus/{agent_profile}
```
**Examples:**
- `manus/manus-1.6` - General purpose agent
- `manus/manus-1.6-lite` - Lightweight agent for simple tasks
- `manus/manus-1.6-max` - Advanced agent for complex analysis
## LiteLLM Python SDK
```python showLineNumbers title="Basic Usage"
import litellm
import os
import time
# Set API key
os.environ["MANUS_API_KEY"] = "your-manus-api-key"
# Create task
response = litellm.responses(
model="manus/manus-1.6",
input="What's the capital of France?",
)
print(f"Task ID: {response.id}")
print(f"Status: {response.status}") # "running"
# Poll until complete
task_id = response.id
while response.status == "running":
time.sleep(5)
response = litellm.get_response(
response_id=task_id,
custom_llm_provider="manus",
)
print(f"Status: {response.status}")
# Get results
if response.status == "completed":
for message in response.output:
if message.role == "assistant":
print(message.content[0].text)
```
## LiteLLM AI Gateway
### Setup
```yaml showLineNumbers title="config.yaml"
model_list:
- model_name: manus-agent
litellm_params:
model: manus/manus-1.6
api_key: os.environ/MANUS_API_KEY
```
```bash title="Start Proxy"
litellm --config config.yaml
```
### Usage
<Tabs>
<TabItem value="curl" label="cURL">
```bash showLineNumbers title="Create Task"
# Create task
curl -X POST http://localhost:4000/responses \
-H "Authorization: Bearer your-proxy-key" \
-H "Content-Type: application/json" \
-d '{
"model": "manus-agent",
"input": "What is the capital of France?"
}'
# Response
{
"id": "task_abc123",
"status": "running",
"metadata": {
"task_url": "https://manus.im/app/task_abc123"
}
}
```
```bash showLineNumbers title="Poll for Completion"
# Check status (repeat until status is "completed")
curl http://localhost:4000/responses/task_abc123 \
-H "Authorization: Bearer your-proxy-key"
# When completed
{
"id": "task_abc123",
"status": "completed",
"output": [
{
"role": "user",
"content": [{"text": "What is the capital of France?"}]
},
{
"role": "assistant",
"content": [{"text": "The capital of France is Paris."}]
}
]
}
```
</TabItem>
<TabItem value="openai" label="OpenAI SDK">
```python showLineNumbers title="Create Task and Poll"
import openai
import time
client = openai.OpenAI(
base_url="http://localhost:4000",
api_key="your-proxy-key"
)
# Create task
response = client.responses.create(
model="manus-agent",
input="What is the capital of France?"
)
print(f"Task ID: {response.id}")
print(f"Status: {response.status}") # "running"
# Poll until complete
task_id = response.id
while response.status == "running":
time.sleep(5)
response = client.responses.retrieve(response_id=task_id)
print(f"Status: {response.status}")
# Get results
if response.status == "completed":
for message in response.output:
if message.role == "assistant":
print(message.content[0].text)
```
</TabItem>
</Tabs>
## How It Works
Manus operates as an **asynchronous agent API**:
1. **Create Task**: When you call `litellm.responses()`, Manus creates a task and returns immediately with `status: "running"`
2. **Task Executes**: The agent works on your request in the background
3. **Poll for Completion**: You must repeatedly call `litellm.get_response()` or `client.responses.retrieve()` until the status changes to `"completed"`
4. **Get Results**: Once completed, the `output` field contains the full conversation
**Task Statuses:**
- `running` - Agent is actively working
- `pending` - Agent is waiting for input
- `completed` - Task finished successfully
- `error` - Task failed
:::tip Production Usage
For production applications, use [webhooks](https://open.manus.im/docs/webhooks) instead of polling to get notified when tasks complete.
:::
## Supported Parameters
| Parameter | Supported | Notes |
|-----------|-----------|-------|
| `input` | ✅ | Text, images, or structured content |
| `stream` | ✅ | Fake streaming (task runs async) |
| `max_output_tokens` | ✅ | Limits response length |
| `previous_response_id` | ✅ | For multi-turn conversations |
## Related Documentation
- [LiteLLM Responses API](/docs/response_api)
- [Manus OpenAI Compatibility](https://open.manus.im/docs/openai-compatibility)

View file

@ -146,6 +146,7 @@ router_settings:
cooldown_time: 30 # (in seconds) how long to cooldown model if fails/min > allowed_fails
disable_cooldowns: True # bool - Disable cooldowns for all models
enable_tag_filtering: True # bool - Use tag based routing for requests
tag_filtering_match_any: True # bool - Tag matching behavior (only when enable_tag_filtering=true). `true`: match if deployment has ANY requested tag; `false`: match only if deployment has ALL requested tags
retry_policy: { # Dict[str, int]: retry policy for different types of exceptions
"AuthenticationErrorRetries": 3,
"TimeoutErrorRetries": 3,
@ -293,6 +294,7 @@ router_settings:
cooldown_time: 30 # (in seconds) how long to cooldown model if fails/min > allowed_fails
disable_cooldowns: True # bool - Disable cooldowns for all models
enable_tag_filtering: True # bool - Use tag based routing for requests
tag_filtering_match_any: True # bool - Tag matching behavior (only when enable_tag_filtering=true). `true`: match if deployment has ANY requested tag; `false`: match only if deployment has ALL requested tags
retry_policy: { # Dict[str, int]: retry policy for different types of exceptions
"AuthenticationErrorRetries": 3,
"TimeoutErrorRetries": 3,
@ -322,6 +324,7 @@ router_settings:
| content_policy_fallbacks | array of objects | Specifies fallback models for content policy violations. [More information here](reliability) |
| fallbacks | array of objects | Specifies fallback models for all types of errors. [More information here](reliability) |
| enable_tag_filtering | boolean | If true, uses tag based routing for requests [Tag Based Routing](tag_routing) |
| tag_filtering_match_any | boolean | Tag matching behavior (only when enable_tag_filtering=true). `true`: match if deployment has ANY requested tag; `false`: match only if deployment has ALL requested tags |
| cooldown_time | integer | The duration (in seconds) to cooldown a model if it exceeds the allowed failures. |
| disable_cooldowns | boolean | If true, disables cooldowns for all models. [More information here](reliability) |
| retry_policy | object | Specifies the number of retries for different types of exceptions. [More information here](reliability) |

View file

@ -710,6 +710,7 @@ const sidebars = {
"providers/llamafile",
"providers/llamagate",
"providers/lm_studio",
"providers/manus",
"providers/meta_llama",
"providers/milvus_vector_stores",
"providers/mistral",

View file

@ -0,0 +1,9 @@
-- CreateIndex
-- Fixes performance issue in _check_duplicate_user_email function
-- by enabling fast case-insensitive email lookups.
--
-- Without this index, queries with mode: "insensitive" cause full table scans.
-- With this index, PostgreSQL can use an Index Scan for O(log n) performance.
--
-- Related: GitHub Issue #18411
CREATE INDEX "LiteLLM_UserTable_user_email_lower_idx" ON "LiteLLM_UserTable"(LOWER("user_email"));

View file

@ -486,6 +486,7 @@ vertex_mistral_models: Set = set()
vertex_openai_models: Set = set()
vertex_minimax_models: Set = set()
vertex_moonshot_models: Set = set()
vertex_zai_models: Set = set()
ai21_models: Set = set()
ai21_chat_models: Set = set()
nlp_cloud_models: Set = set()
@ -664,6 +665,9 @@ def add_known_models():
elif value.get("litellm_provider") == "vertex_ai-moonshot_models":
key = key.replace("vertex_ai/", "")
vertex_moonshot_models.add(key)
elif value.get("litellm_provider") == "vertex_ai-zai_models":
key = key.replace("vertex_ai/", "")
vertex_zai_models.add(key)
elif value.get("litellm_provider") == "ai21":
if value.get("mode") == "chat":
ai21_chat_models.add(key)
@ -950,7 +954,8 @@ models_by_provider: dict = {
| vertex_language_models
| vertex_deepseek_models
| vertex_minimax_models
| vertex_moonshot_models,
| vertex_moonshot_models
| vertex_zai_models,
"ai21": ai21_models,
"bedrock": bedrock_models | bedrock_converse_models,
"petals": petals_models,
@ -1338,6 +1343,7 @@ if TYPE_CHECKING:
from .llms.bedrock.chat.invoke_transformations.amazon_llama_transformation import AmazonLlamaConfig as AmazonLlamaConfig
from .llms.bedrock.chat.invoke_transformations.amazon_deepseek_transformation import AmazonDeepSeekR1Config as AmazonDeepSeekR1Config
from .llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation import AmazonMistralConfig as AmazonMistralConfig
from .llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import AmazonMoonshotConfig as AmazonMoonshotConfig
from .llms.bedrock.chat.invoke_transformations.amazon_titan_transformation import AmazonTitanConfig as AmazonTitanConfig
from .llms.bedrock.chat.invoke_transformations.amazon_twelvelabs_pegasus_transformation import AmazonTwelveLabsPegasusConfig as AmazonTwelveLabsPegasusConfig
from .llms.bedrock.chat.invoke_transformations.base_invoke_transformation import AmazonInvokeConfig as AmazonInvokeConfig
@ -1367,6 +1373,7 @@ if TYPE_CHECKING:
from .llms.azure.responses.o_series_transformation import AzureOpenAIOSeriesResponsesAPIConfig as AzureOpenAIOSeriesResponsesAPIConfig
from .llms.xai.responses.transformation import XAIResponsesAPIConfig as XAIResponsesAPIConfig
from .llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig as LiteLLMProxyResponsesAPIConfig
from .llms.manus.responses.transformation import ManusResponsesAPIConfig as ManusResponsesAPIConfig
from .llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig
from .llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig as OpenAIOSeriesConfig, OpenAIOSeriesConfig as OpenAIO1Config
from .llms.anthropic.skills.transformation import AnthropicSkillsConfig as AnthropicSkillsConfig

View file

@ -165,6 +165,7 @@ LLM_CONFIG_NAMES = (
"AmazonLlamaConfig",
"AmazonDeepSeekR1Config",
"AmazonMistralConfig",
"AmazonMoonshotConfig",
"AmazonTitanConfig",
"AmazonTwelveLabsPegasusConfig",
"AmazonInvokeConfig",
@ -252,6 +253,7 @@ LLM_CONFIG_NAMES = (
"IBMWatsonXAudioTranscriptionConfig",
"GithubCopilotConfig",
"GithubCopilotResponsesAPIConfig",
"ManusResponsesAPIConfig",
"GithubCopilotEmbeddingConfig",
"NebiusConfig",
"WandbConfig",
@ -556,6 +558,7 @@ _LLM_CONFIGS_IMPORT_MAP = {
"AmazonLlamaConfig": (".llms.bedrock.chat.invoke_transformations.amazon_llama_transformation", "AmazonLlamaConfig"),
"AmazonDeepSeekR1Config": (".llms.bedrock.chat.invoke_transformations.amazon_deepseek_transformation", "AmazonDeepSeekR1Config"),
"AmazonMistralConfig": (".llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation", "AmazonMistralConfig"),
"AmazonMoonshotConfig": (".llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation", "AmazonMoonshotConfig"),
"AmazonTitanConfig": (".llms.bedrock.chat.invoke_transformations.amazon_titan_transformation", "AmazonTitanConfig"),
"AmazonTwelveLabsPegasusConfig": (".llms.bedrock.chat.invoke_transformations.amazon_twelvelabs_pegasus_transformation", "AmazonTwelveLabsPegasusConfig"),
"AmazonInvokeConfig": (".llms.bedrock.chat.invoke_transformations.base_invoke_transformation", "AmazonInvokeConfig"),
@ -588,6 +591,7 @@ _LLM_CONFIGS_IMPORT_MAP = {
"AzureOpenAIOSeriesResponsesAPIConfig": (".llms.azure.responses.o_series_transformation", "AzureOpenAIOSeriesResponsesAPIConfig"),
"XAIResponsesAPIConfig": (".llms.xai.responses.transformation", "XAIResponsesAPIConfig"),
"LiteLLMProxyResponsesAPIConfig": (".llms.litellm_proxy.responses.transformation", "LiteLLMProxyResponsesAPIConfig"),
"ManusResponsesAPIConfig": (".llms.manus.responses.transformation", "ManusResponsesAPIConfig"),
"GoogleAIStudioInteractionsConfig": (".llms.gemini.interactions.transformation", "GoogleAIStudioInteractionsConfig"),
"OpenAIOSeriesConfig": (".llms.openai.chat.o_series_transformation", "OpenAIOSeriesConfig"),
"AnthropicSkillsConfig": (".llms.anthropic.skills.transformation", "AnthropicSkillsConfig"),

View file

@ -31,6 +31,7 @@ from litellm.llms.base_llm.bridges.completion_transformation import (
CompletionTransformationBridge,
)
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
ChatCompletionToolParamFunctionChunk,
Reasoning,
ResponsesAPIOptionalRequestParams,
@ -90,9 +91,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
content_type = content_item.get("type")
if content_type == "output_text":
response_text = content_item.get("text", "")
# Extract annotations from content if present
annotations = LiteLLMResponsesTransformationHandler._convert_annotations_to_chat_format(
content_item.get("annotations", None)
)
msg = Message(
role=item.get("role", "assistant"),
content=response_text if response_text else "",
annotations=annotations,
)
choice = Choices(message=msg, finish_reason="stop", index=index)
return choice, index + 1
@ -364,10 +370,16 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
elif isinstance(item, ResponseOutputMessage):
for content in item.content:
response_text = getattr(content, "text", "")
# Extract annotations from content if present
raw_annotations = getattr(content, "annotations", None)
annotations = LiteLLMResponsesTransformationHandler._convert_annotations_to_chat_format(
raw_annotations
)
msg = Message(
role=item.role,
content=response_text if response_text else "",
reasoning_content=reasoning_content,
annotations=annotations,
)
choices.append(
@ -763,6 +775,42 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return {"format": {"type": "text"}}
return None
@staticmethod
def _convert_annotations_to_chat_format(
annotations: Optional[List[Any]],
) -> Optional[List["ChatCompletionAnnotation"]]:
"""
Convert annotations from Responses API to Chat Completions format.
Annotations are already in compatible format between both APIs,
so we just need to convert Pydantic models to dicts.
"""
if not annotations:
return None
result: List[ChatCompletionAnnotation] = []
for annotation in annotations:
try:
# Convert Pydantic models to dicts (handles both v1 and v2)
if hasattr(annotation, "model_dump"):
annotation_dict = annotation.model_dump()
elif hasattr(annotation, "dict"):
annotation_dict = annotation.dict()
elif isinstance(annotation, dict):
annotation_dict = annotation
else:
# Skip unsupported annotation types
verbose_logger.debug(f"Skipping unsupported annotation type: {type(annotation)}")
continue
result.append(annotation_dict) # type: ignore
except Exception as e:
# Skip malformed annotations
verbose_logger.debug(f"Skipping malformed annotation: {annotation}, error: {e}")
continue
return result if result else None
def _map_responses_status_to_finish_reason(self, status: Optional[str]) -> str:
"""Map responses API status to chat completion finish_reason"""

View file

@ -909,6 +909,7 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[
"twelvelabs",
"openai",
"stability",
"moonshot",
]
BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[

View file

@ -14,6 +14,7 @@ from typing import (
Literal,
Optional,
Tuple,
Union,
cast,
)
@ -791,6 +792,11 @@ class PrometheusLogger(CustomLogger):
f"standard_logging_object is required, got={standard_logging_payload}"
)
if self._should_skip_metrics_for_invalid_key(
kwargs=kwargs, standard_logging_payload=standard_logging_payload
):
return
model = kwargs.get("model", "")
litellm_params = kwargs.get("litellm_params", {}) or {}
_metadata = litellm_params.get("metadata", {})
@ -1189,11 +1195,17 @@ class PrometheusLogger(CustomLogger):
f"prometheus Logging - Enters failure logging function for kwargs {kwargs}"
)
# unpack kwargs
model = kwargs.get("model", "")
standard_logging_payload: StandardLoggingPayload = kwargs.get(
"standard_logging_object", {}
)
if self._should_skip_metrics_for_invalid_key(
kwargs=kwargs, standard_logging_payload=standard_logging_payload
):
return
model = kwargs.get("model", "")
litellm_params = kwargs.get("litellm_params", {}) or {}
get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking()
@ -1207,7 +1219,6 @@ class PrometheusLogger(CustomLogger):
user_api_team_alias = standard_logging_payload["metadata"][
"user_api_key_team_alias"
]
kwargs.get("exception", None)
try:
self.litellm_llm_api_failed_requests_metric.labels(
@ -1227,6 +1238,139 @@ class PrometheusLogger(CustomLogger):
pass
pass
def _extract_status_code(
self,
kwargs: Optional[dict] = None,
enum_values: Optional[Any] = None,
exception: Optional[Exception] = None,
) -> Optional[int]:
"""
Extract HTTP status code from various input formats for validation.
This is a centralized helper to extract status code from different
callback function signatures. Handles both ProxyException (uses 'code')
and standard exceptions (uses 'status_code').
Args:
kwargs: Dictionary potentially containing 'exception' key
enum_values: Object with 'status_code' attribute
exception: Exception object to extract status code from directly
Returns:
Status code as integer if found, None otherwise
"""
status_code = None
# Try from enum_values first (most common in our callbacks)
if enum_values and hasattr(enum_values, "status_code") and enum_values.status_code:
try:
status_code = int(enum_values.status_code)
except (ValueError, TypeError):
pass
if not status_code and exception:
# ProxyException uses 'code' attribute, other exceptions may use 'status_code'
status_code = getattr(exception, "status_code", None) or getattr(exception, "code", None)
if status_code is not None:
try:
status_code = int(status_code)
except (ValueError, TypeError):
status_code = None
if not status_code and kwargs:
exception_in_kwargs = kwargs.get("exception")
if exception_in_kwargs:
status_code = getattr(exception_in_kwargs, "status_code", None) or getattr(exception_in_kwargs, "code", None)
if status_code is not None:
try:
status_code = int(status_code)
except (ValueError, TypeError):
status_code = None
return status_code
def _is_invalid_api_key_request(
self,
status_code: Optional[int],
exception: Optional[Exception] = None,
) -> bool:
"""
Determine if a request has an invalid API key based on status code and exception.
This method prevents invalid authentication attempts from being recorded in
Prometheus metrics. A 401 status code is the definitive indicator of authentication
failure. Additionally, we check exception messages for authentication error patterns
to catch cases where the exception hasn't been converted to a ProxyException yet.
Args:
status_code: HTTP status code (401 indicates authentication error)
exception: Exception object to check for auth-related error messages
Returns:
True if the request has an invalid API key and metrics should be skipped,
False otherwise
"""
if status_code == 401:
return True
# Handle cases where AssertionError is raised before conversion to ProxyException
if exception is not None:
exception_str = str(exception).lower()
auth_error_patterns = [
"virtual key expected",
"expected to start with 'sk-'",
"authentication error",
"invalid api key",
"api key not valid",
]
if any(pattern in exception_str for pattern in auth_error_patterns):
return True
return False
def _should_skip_metrics_for_invalid_key(
self,
kwargs: Optional[dict] = None,
user_api_key_dict: Optional[Any] = None,
enum_values: Optional[Any] = None,
standard_logging_payload: Optional[Union[dict, StandardLoggingPayload]] = None,
exception: Optional[Exception] = None,
) -> bool:
"""
Determine if Prometheus metrics should be skipped for invalid API key requests.
This is a centralized validation method that extracts status code and exception
information from various callback function signatures and determines if the request
represents an invalid API key attempt that should be filtered from metrics.
Args:
kwargs: Dictionary potentially containing exception and other data
user_api_key_dict: User API key authentication object (currently unused)
enum_values: Object with status_code attribute
standard_logging_payload: Standard logging payload dictionary
exception: Exception object to check directly
Returns:
True if metrics should be skipped (invalid key detected), False otherwise
"""
status_code = self._extract_status_code(
kwargs=kwargs,
enum_values=enum_values,
exception=exception,
)
if exception is None and kwargs:
exception = kwargs.get("exception")
if self._is_invalid_api_key_request(status_code, exception=exception):
verbose_logger.debug(
"Skipping Prometheus metrics for invalid API key request: "
f"status_code={status_code}, exception={type(exception).__name__ if exception else None}"
)
return True
return False
async def async_post_call_failure_hook(
self,
request_data: dict,
@ -1252,6 +1396,14 @@ class PrometheusLogger(CustomLogger):
StandardLoggingPayloadSetup,
)
if self._should_skip_metrics_for_invalid_key(
user_api_key_dict=user_api_key_dict,
exception=original_exception,
):
return
status_code = self._extract_status_code(exception=original_exception)
try:
_tags = StandardLoggingPayloadSetup._get_request_tags(
litellm_params=request_data,
@ -1266,8 +1418,8 @@ class PrometheusLogger(CustomLogger):
team=user_api_key_dict.team_id,
team_alias=user_api_key_dict.team_alias,
requested_model=request_data.get("model", ""),
status_code=str(getattr(original_exception, "status_code", None)),
exception_status=str(getattr(original_exception, "status_code", None)),
status_code=str(status_code),
exception_status=str(status_code),
exception_class=self._get_exception_class_name(original_exception),
tags=_tags,
route=user_api_key_dict.request_route,
@ -1305,6 +1457,11 @@ class PrometheusLogger(CustomLogger):
StandardLoggingPayloadSetup,
)
if self._should_skip_metrics_for_invalid_key(
user_api_key_dict=user_api_key_dict
):
return
enum_values = UserAPIKeyLabelValues(
end_user=user_api_key_dict.end_user_id,
hashed_api_key=user_api_key_dict.api_key,
@ -1360,6 +1517,15 @@ class PrometheusLogger(CustomLogger):
exception = request_kwargs.get("exception", None)
llm_provider = _litellm_params.get("custom_llm_provider", None)
if self._should_skip_metrics_for_invalid_key(
kwargs=request_kwargs,
standard_logging_payload=standard_logging_payload,
):
return
hashed_api_key = standard_logging_payload.get("metadata", {}).get(
"user_api_key_hash"
)
# Create enum_values for the label factory (always create for use in different metrics)
enum_values = UserAPIKeyLabelValues(
@ -1374,9 +1540,7 @@ class PrometheusLogger(CustomLogger):
self._get_exception_class_name(exception) if exception else None
),
requested_model=model_group,
hashed_api_key=standard_logging_payload["metadata"][
"user_api_key_hash"
],
hashed_api_key=hashed_api_key,
api_key_alias=standard_logging_payload["metadata"][
"user_api_key_alias"
],
@ -1441,6 +1605,14 @@ class PrometheusLogger(CustomLogger):
if standard_logging_payload is None:
return
# Skip recording metrics for invalid API key requests
if self._should_skip_metrics_for_invalid_key(
kwargs=request_kwargs,
enum_values=enum_values,
standard_logging_payload=standard_logging_payload,
):
return
api_base = standard_logging_payload["api_base"]
_litellm_params = request_kwargs.get("litellm_params", {}) or {}
_metadata = _litellm_params.get("metadata", {})

View file

@ -913,6 +913,14 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
or "http://localhost:2024"
)
dynamic_api_key = api_key or get_secret_str("LANGGRAPH_API_KEY")
elif custom_llm_provider == "manus":
# Manus is OpenAI compatible for responses API
api_base = (
api_base
or get_secret_str("MANUS_API_BASE")
or "https://api.manus.im"
)
dynamic_api_key = api_key or get_secret_str("MANUS_API_KEY")
if api_base is not None and not isinstance(api_base, str):
raise Exception("api base needs to be a string. api_base={}".format(api_base))

View file

@ -4800,7 +4800,7 @@ class StandardLoggingPayloadSetup:
"""
Extract additional header tags for spend tracking based on config.
"""
extra_headers: List[str] = litellm.extra_spend_tag_headers or []
extra_headers: List[str] = getattr(litellm, "extra_spend_tag_headers", None) or []
if not extra_headers:
return None

View file

@ -6,6 +6,7 @@ import io
import mimetypes
import re
from os import PathLike
from pathlib import Path
from typing import (
TYPE_CHECKING,
Any,
@ -533,6 +534,12 @@ def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
# Convert content to bytes
if isinstance(file_content, (str, PathLike)):
# If it's a path, open and read the file
# Extract filename from path if not already set
if filename is None:
if isinstance(file_content, PathLike):
filename = Path(file_content).name
else:
filename = Path(str(file_content)).name
with open(file_content, "rb") as f:
content = f.read()
elif isinstance(file_content, io.IOBase):
@ -550,11 +557,11 @@ def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
# Use provided content type or guess based on filename
if not content_type:
content_type = (
mimetypes.guess_type(filename)[0]
if filename
else "application/octet-stream"
)
if filename:
guessed_type = mimetypes.guess_type(filename)[0]
content_type = guessed_type if guessed_type else "application/octet-stream"
else:
content_type = "application/octet-stream"
return ExtractedFileData(
filename=filename,

View file

@ -1,4 +1,5 @@
from typing import Any, Dict, Optional, Set
from collections.abc import Mapping
from typing import Any, Dict, List, Optional, Set
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
@ -17,6 +18,7 @@ class SensitiveDataMasker:
"key",
"token",
"auth",
"authorization",
"credential",
"access",
"private",
@ -42,22 +44,52 @@ class SensitiveDataMasker:
else:
return f"{value_str[:self.visible_prefix]}{self.mask_char * masked_length}{value_str[-self.visible_suffix:]}"
def is_sensitive_key(self, key: str, excluded_keys: Optional[Set[str]] = None) -> bool:
def is_sensitive_key(
self, key: str, excluded_keys: Optional[Set[str]] = None
) -> bool:
# Check if key is in excluded_keys first (exact match)
if excluded_keys and key in excluded_keys:
return False
key_lower = str(key).lower()
# Split on underscores and check if any segment matches the pattern
# Split on underscores/hyphens and check if any segment matches the pattern
# This avoids false positives like "max_tokens" matching "token"
# but still catches "api_key", "access_token", etc.
key_segments = key_lower.replace('-', '_').split('_')
result = any(
pattern in key_segments
for pattern in self.sensitive_patterns
)
key_segments = key_lower.replace("-", "_").split("_")
result = any(pattern in key_segments for pattern in self.sensitive_patterns)
return result
def _mask_sequence(
self,
values: List[Any],
depth: int,
max_depth: int,
excluded_keys: Optional[Set[str]],
key_is_sensitive: bool,
) -> List[Any]:
masked_items: List[Any] = []
if depth >= max_depth:
return values
for item in values:
if isinstance(item, Mapping):
masked_items.append(
self.mask_dict(dict(item), depth + 1, max_depth, excluded_keys)
)
elif isinstance(item, list):
masked_items.append(
self._mask_sequence(
item, depth + 1, max_depth, excluded_keys, key_is_sensitive
)
)
elif key_is_sensitive and isinstance(item, str):
masked_items.append(self._mask_value(item))
else:
masked_items.append(
item if isinstance(item, (int, float, bool, str, list)) else str(item)
)
return masked_items
def mask_dict(
self,
data: Dict[str, Any],
@ -71,11 +103,20 @@ class SensitiveDataMasker:
masked_data: Dict[str, Any] = {}
for k, v in data.items():
try:
if isinstance(v, dict):
masked_data[k] = self.mask_dict(v, depth + 1, max_depth, excluded_keys)
key_is_sensitive = self.is_sensitive_key(k, excluded_keys)
if isinstance(v, Mapping):
masked_data[k] = self.mask_dict(
dict(v), depth + 1, max_depth, excluded_keys
)
elif isinstance(v, list):
masked_data[k] = self._mask_sequence(
v, depth + 1, max_depth, excluded_keys, key_is_sensitive
)
elif hasattr(v, "__dict__") and not isinstance(v, type):
masked_data[k] = self.mask_dict(vars(v), depth + 1, max_depth, excluded_keys)
elif self.is_sensitive_key(k, excluded_keys):
masked_data[k] = self.mask_dict(
vars(v), depth + 1, max_depth, excluded_keys
)
elif key_is_sensitive:
str_value = str(v) if v is not None else ""
masked_data[k] = self._mask_value(str_value)
else:

View file

@ -1265,14 +1265,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
cache_creation_tokens=cache_creation_input_tokens,
cache_creation_token_details=cache_creation_token_details,
)
completion_token_details = (
CompletionTokensDetailsWrapper(
reasoning_tokens=token_counter(
text=reasoning_content, count_response_tokens=True
)
)
# Always populate completion_token_details, not just when there's reasoning_content
reasoning_tokens = (
token_counter(text=reasoning_content, count_response_tokens=True)
if reasoning_content
else None
else 0
)
completion_token_details = CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens if reasoning_tokens > 0 else None,
text_tokens=completion_tokens - reasoning_tokens if reasoning_tokens > 0 else completion_tokens,
)
total_tokens = prompt_tokens + completion_tokens

View file

@ -369,6 +369,10 @@ class BaseAWSLLM:
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
model_id, spec="stability"
)
elif provider == "moonshot" and "moonshot/" in model_id:
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
model_id, spec="moonshot"
)
return model_id
@staticmethod

View file

@ -0,0 +1,256 @@
"""
Transformation for Bedrock Moonshot AI (Kimi K2) models.
Supports the Kimi K2 Thinking model available on Amazon Bedrock.
Model format: bedrock/moonshot.kimi-k2-thinking-v1:0
Reference: https://aws.amazon.com/about-aws/whats-new/2025/12/amazon-bedrock-fully-managed-open-weight-models/
"""
from typing import TYPE_CHECKING, Any, List, Optional, Union
import re
import httpx
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
)
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.moonshot.chat.transformation import MoonshotChatConfig
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.types.utils import ModelResponse
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig):
"""
Configuration for Bedrock Moonshot AI (Kimi K2) models.
Reference:
https://aws.amazon.com/about-aws/whats-new/2025/12/amazon-bedrock-fully-managed-open-weight-models/
https://platform.moonshot.ai/docs/api/chat
Supported Params for the Amazon / Moonshot models:
- `max_tokens` (integer) max tokens
- `temperature` (float) temperature for model (0-1 for Moonshot)
- `top_p` (float) top p for model
- `stream` (bool) whether to stream responses
- `tools` (list) tool definitions (supported on kimi-k2-thinking)
- `tool_choice` (str|dict) tool choice specification (supported on kimi-k2-thinking)
NOT Supported on Bedrock:
- `stop` sequences (Bedrock doesn't support stopSequences field for this model)
Note: The kimi-k2-thinking model DOES support tool calls, unlike kimi-thinking-preview.
"""
def __init__(self, **kwargs):
AmazonInvokeConfig.__init__(self, **kwargs)
MoonshotChatConfig.__init__(self, **kwargs)
@property
def custom_llm_provider(self) -> Optional[str]:
return "bedrock"
def _get_model_id(self, model: str) -> str:
"""
Extract the actual model ID from the LiteLLM model name.
Removes routing prefixes like:
- bedrock/invoke/moonshot.kimi-k2-thinking -> moonshot.kimi-k2-thinking
- invoke/moonshot.kimi-k2-thinking -> moonshot.kimi-k2-thinking
- moonshot.kimi-k2-thinking -> moonshot.kimi-k2-thinking
"""
# Remove bedrock/ prefix if present
if model.startswith("bedrock/"):
model = model[8:]
# Remove invoke/ prefix if present
if model.startswith("invoke/"):
model = model[7:]
# Remove any provider prefix (e.g., moonshot/)
if "/" in model and not model.startswith("arn:"):
parts = model.split("/", 1)
if len(parts) == 2:
model = parts[1]
return model
def get_supported_openai_params(self, model: str) -> List[str]:
"""
Get the supported OpenAI params for Moonshot AI models on Bedrock.
Bedrock-specific limitations:
- stopSequences field is not supported on Bedrock (unlike native Moonshot API)
- functions parameter is not supported (use tools instead)
- tool_choice doesn't support "required" value
Note: kimi-k2-thinking DOES support tool calls (unlike kimi-thinking-preview)
The parent MoonshotChatConfig class handles the kimi-thinking-preview exclusion.
"""
excluded_params: List[str] = ["functions", "stop"] # Bedrock doesn't support stopSequences
base_openai_params = super(MoonshotChatConfig, self).get_supported_openai_params(model=model)
final_params: List[str] = []
for param in base_openai_params:
if param not in excluded_params:
final_params.append(param)
return final_params
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI parameters to Moonshot AI parameters for Bedrock.
Handles Moonshot AI specific limitations:
- tool_choice doesn't support "required" value
- Temperature <0.3 limitation for n>1
- Temperature range is [0, 1] (not [0, 2] like OpenAI)
"""
return MoonshotChatConfig.map_openai_params(
self,
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=drop_params,
)
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the request for Bedrock Moonshot AI models.
Uses the Moonshot transformation logic which handles:
- Converting content lists to strings (Moonshot doesn't support list format)
- Adding tool_choice="required" message if needed
- Temperature and parameter validation
"""
# Filter out AWS credentials using the existing method from BaseAWSLLM
self._get_boto_credentials_from_optional_params(optional_params, model)
# Strip routing prefixes to get the actual model ID
clean_model_id = self._get_model_id(model)
# Use Moonshot's transform_request which handles message transformation
# and tool_choice="required" workaround
return MoonshotChatConfig.transform_request(
self,
model=clean_model_id,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
def _extract_reasoning_from_content(self, content: str) -> tuple[Optional[str], str]:
"""
Extract reasoning content from <reasoning> tags in the response.
Moonshot AI's Kimi K2 Thinking model returns reasoning in <reasoning> tags.
This method extracts that content and returns it separately.
Args:
content: The full content string from the API response
Returns:
tuple: (reasoning_content, main_content)
"""
if not content:
return None, content
# Match <reasoning>...</reasoning> tags
reasoning_match = re.match(
r"<reasoning>(.*?)</reasoning>\s*(.*)",
content,
re.DOTALL
)
if reasoning_match:
reasoning_content = reasoning_match.group(1).strip()
main_content = reasoning_match.group(2).strip()
return reasoning_content, main_content
return None, content
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: "ModelResponse",
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> "ModelResponse":
"""
Transform the response from Bedrock Moonshot AI models.
Moonshot AI uses OpenAI-compatible response format, but returns reasoning
content in <reasoning> tags. This method:
1. Calls parent class transformation
2. Extracts reasoning content from <reasoning> tags
3. Sets reasoning_content on the message object
"""
# First, get the standard transformation
model_response = MoonshotChatConfig.transform_response(
self,
model=model,
raw_response=raw_response,
model_response=model_response,
logging_obj=logging_obj,
request_data=request_data,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=encoding,
api_key=api_key,
json_mode=json_mode,
)
# Extract reasoning content from <reasoning> tags
if model_response.choices and len(model_response.choices) > 0:
for choice in model_response.choices:
# Only process Choices (not StreamingChoices) which have message attribute
if isinstance(choice, Choices) and choice.message and choice.message.content:
reasoning_content, main_content = self._extract_reasoning_from_content(
choice.message.content
)
if reasoning_content:
# Set the reasoning_content field
choice.message.reasoning_content = reasoning_content
# Update the main content without reasoning tags
choice.message.content = main_content
return model_response
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BedrockError:
"""Return the appropriate error class for Bedrock."""
return BedrockError(status_code=status_code, message=error_message)

View file

@ -629,6 +629,8 @@ def get_bedrock_chat_config(model: str):
return litellm.AmazonCohereConfig()
elif bedrock_invoke_provider == "mistral":
return litellm.AmazonMistralConfig()
elif bedrock_invoke_provider == "moonshot":
return litellm.AmazonMoonshotConfig()
elif bedrock_invoke_provider == "deepseek_r1":
return litellm.AmazonDeepSeekR1Config()
elif bedrock_invoke_provider == "nova":

View file

@ -34,11 +34,12 @@ class BedrockPassthroughConfig(
litellm_params: dict,
) -> Tuple["URL", str]:
optional_params = litellm_params.copy()
model_id = optional_params.get("model_id", None)
aws_region_name = self._get_aws_region_name(
optional_params=optional_params,
model=model,
model_id=None,
model_id=model_id,
)
aws_bedrock_runtime_endpoint = optional_params.get("aws_bedrock_runtime_endpoint")
@ -49,6 +50,12 @@ class BedrockPassthroughConfig(
endpoint_type="runtime",
)
# If model_id is provided (e.g., Application Inference Profile ARN), use it in the endpoint
# instead of the translated model name
if model_id is not None:
# Replace the model name in the endpoint with the model_id
import re
endpoint = re.sub(r'model/[^/]+/', f'model/{model_id}/', endpoint)
return self.format_url(endpoint, endpoint_url, request_query_params or {}), endpoint_url
def sign_request(

View file

@ -1,9 +1,11 @@
from typing import Optional, Tuple, Union
import json
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, cast, overload
import litellm
from litellm.constants import MIN_NON_ZERO_TEMPERATURE
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
class DeepInfraConfig(OpenAIGPTConfig):
@ -117,6 +119,79 @@ class DeepInfraConfig(OpenAIGPTConfig):
optional_params[param] = value
return optional_params
def _transform_tool_message_content(self, messages: List[AllMessageValues]) -> List[AllMessageValues]:
"""
Transform tool message content from array to string format for DeepInfra compatibility.
DeepInfra requires tool message content to be a string, not an array.
This method converts tool message content from array format to string format.
Example transformation:
- Input: {"role": "tool", "content": [{"type": "text", "text": "20"}]}
- Output: {"role": "tool", "content": "20"}
Or if content is complex:
- Input: {"role": "tool", "content": [{"type": "text", "text": "result"}]}
- Output: {"role": "tool", "content": "[{\"type\": \"text\", \"text\": \"result\"}]"}
"""
for message in messages:
if message.get("role") == "tool":
content = message.get("content")
# If content is a list/array, convert it to string
if isinstance(content, list):
# Check if it's a simple single text item
if (
len(content) == 1
and isinstance(content[0], dict)
and content[0].get("type") == "text"
and "text" in content[0]
):
# Extract just the text value for simple cases
message["content"] = content[0]["text"]
else:
# For complex content, serialize the entire array as JSON string
message["content"] = json.dumps(content)
return messages
@overload
def _transform_messages(
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
) -> Coroutine[Any, Any, List[AllMessageValues]]:
...
@overload
def _transform_messages(
self, messages: List[AllMessageValues], model: str, is_async: Literal[False] = False
) -> List[AllMessageValues]:
...
def _transform_messages(
self, messages: List[AllMessageValues], model: str, is_async: bool = False
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
"""
Transform messages for DeepInfra compatibility.
Handles both sync and async transformations.
"""
if is_async:
# For async case, create an async function that awaits parent and applies our transformation
async def _async_transform():
# Call parent with is_async=True (literal) for async case
parent_result = super(DeepInfraConfig, self)._transform_messages(
messages=messages, model=model, is_async=cast(Literal[True], True)
)
transformed_messages = await parent_result
return self._transform_tool_message_content(transformed_messages)
return _async_transform()
else:
# Call parent with is_async=False (literal) for sync case
parent_result = super()._transform_messages(
messages=messages, model=model, is_async=cast(Literal[False], False)
)
# For sync case, parent_result is already the transformed messages
return self._transform_tool_message_content(parent_result)
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:

View file

@ -0,0 +1,2 @@
# Manus provider implementation

View file

@ -0,0 +1,2 @@
# Manus Responses API implementation

View file

@ -0,0 +1,308 @@
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
_safe_convert_created_field,
)
from litellm.llms.openai.common_utils import OpenAIError
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseInputParam,
ResponsesAPIResponse,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
MANUS_API_BASE = "https://api.manus.im"
class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
Configuration for Manus API's Responses API.
Manus API is OpenAI-compatible but has some differences:
- API key passed via `API_KEY` header (not `Authorization: Bearer`)
- Model format: `manus/{agent_profile}` (e.g., `manus/manus-1.6`)
- Requires `extra_body` with `task_mode: "agent"` and `agent_profile`
Reference: https://open.manus.im/docs/openai-compatibility
"""
@property
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.MANUS
def should_fake_stream(
self,
model: Optional[str],
stream: Optional[bool],
custom_llm_provider: Optional[str] = None,
) -> bool:
"""
Manus API doesn't support real-time streaming.
It returns a task that runs asynchronously.
We fake streaming by converting the response into streaming events.
"""
return stream is True
def _extract_agent_profile(self, model: str) -> str:
"""
Extract agent profile from model name.
Model format: `manus/{agent_profile}`
Examples: `manus/manus-1.6`, `manus/manus-1.6-lite`, `manus/manus-1.6-max`
Returns:
str: The agent profile (e.g., "manus-1.6")
"""
if "/" in model:
return model.split("/", 1)[1]
# If no slash, assume the model name itself is the agent profile
return model
def validate_environment(
self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]
) -> dict:
"""
Validate environment and set up headers for Manus API.
Manus uses `API_KEY` header instead of `Authorization: Bearer`.
"""
litellm_params = litellm_params or GenericLiteLLMParams()
api_key = (
litellm_params.api_key
or litellm.api_key
or get_secret_str("MANUS_API_KEY")
)
if not api_key:
raise ValueError(
"Manus API key is required. Set MANUS_API_KEY environment variable or pass api_key parameter."
)
# Manus uses API_KEY header, not Authorization: Bearer
headers.update(
{
"API_KEY": api_key,
}
)
return headers
def get_complete_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Get the complete URL for Manus Responses API endpoint.
Returns:
str: The full URL for the Manus /v1/responses endpoint
"""
api_base = (
api_base
or litellm.api_base
or get_secret_str("MANUS_API_BASE")
or MANUS_API_BASE
)
# Remove trailing slashes
api_base = api_base.rstrip("/")
# Manus API uses /v1/responses endpoint (OpenAI-compatible)
if api_base.endswith("/v1"):
return f"{api_base}/responses"
return f"{api_base}/v1/responses"
def transform_responses_api_request(
self,
model: str,
input: Union[str, ResponseInputParam],
response_api_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Dict:
"""
Transform the request for Manus API.
Manus requires:
- `task_mode: "agent"` in the request body
- `agent_profile` extracted from model name in the request body
"""
# First, get the base OpenAI request
base_request = super().transform_responses_api_request(
model=model,
input=input,
response_api_optional_request_params=response_api_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
# Extract agent profile from model name
agent_profile = self._extract_agent_profile(model=model)
# Add Manus-specific parameters directly to the request body
# These will be sent as part of the request
base_request["task_mode"] = "agent"
base_request["agent_profile"] = agent_profile
# Merge any existing extra_body into the request
extra_body = response_api_optional_request_params.get("extra_body", {}) or {}
if extra_body:
base_request.update(extra_body)
# Avoid logging potentially sensitive agent_profile value
verbose_logger.debug("Manus: Using task_mode=agent")
return base_request
def transform_response_api_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIResponse:
"""
Transform Manus API response to OpenAI-compatible format.
Manus uses camelCase (createdAt) instead of snake_case (created_at).
"""
try:
logging_obj.post_call(
original_response=raw_response.text,
additional_args={"complete_input_dict": {}},
)
raw_response_json = raw_response.json()
# Manus uses camelCase "createdAt" instead of snake_case "created_at"
if "createdAt" in raw_response_json and "created_at" not in raw_response_json:
raw_response_json["created_at"] = _safe_convert_created_field(
raw_response_json["createdAt"]
)
# Ensure created_at is set
if "created_at" in raw_response_json:
raw_response_json["created_at"] = _safe_convert_created_field(
raw_response_json["created_at"]
)
except Exception:
raise OpenAIError(
message=raw_response.text, status_code=raw_response.status_code
)
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
# Ensure reasoning is an empty dict if not present, OpenAI SDK does not allow None
if "reasoning" not in raw_response_json or raw_response_json.get("reasoning") is None:
raw_response_json["reasoning"] = {}
if "text" not in raw_response_json or raw_response_json.get("text") is None:
raw_response_json["text"] = {}
if "output" not in raw_response_json or raw_response_json.get("output") is None:
raw_response_json["output"] = []
# Ensure usage is present with default values if not provided
if "usage" not in raw_response_json or raw_response_json.get("usage") is None:
raw_response_json["usage"] = ResponseAPIUsage(
input_tokens=0,
output_tokens=0,
total_tokens=0,
)
try:
response = ResponsesAPIResponse(**raw_response_json)
except Exception:
verbose_logger.debug(
f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct"
)
response = ResponsesAPIResponse.model_construct(**raw_response_json)
# Store processed headers in additional_headers so they get returned to the client
response._hidden_params["additional_headers"] = processed_headers
response._hidden_params["headers"] = raw_response_headers
return response
def transform_get_response_api_request(
self,
response_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, Dict]:
"""
Transform the get response API request into a URL and data.
Manus API follows OpenAI-compatible format:
- GET /v1/responses/{response_id}
Reference: https://open.manus.im/docs/openai-compatibility
"""
url = f"{api_base}/{response_id}"
data: Dict = {}
return url, data
def transform_get_response_api_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIResponse:
"""
Transform Manus API GET response to OpenAI-compatible format.
Manus uses camelCase (createdAt) instead of snake_case (created_at).
Same transformation as transform_response_api_response.
"""
try:
logging_obj.post_call(
original_response=raw_response.text,
additional_args={"complete_input_dict": {}},
)
raw_response_json = raw_response.json()
# Manus uses camelCase "createdAt" instead of snake_case "created_at"
if "createdAt" in raw_response_json and "created_at" not in raw_response_json:
raw_response_json["created_at"] = _safe_convert_created_field(
raw_response_json["createdAt"]
)
# Ensure created_at is set
if "created_at" in raw_response_json:
raw_response_json["created_at"] = _safe_convert_created_field(
raw_response_json["created_at"]
)
except Exception:
raise OpenAIError(
message=raw_response.text, status_code=raw_response.status_code
)
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
try:
response = ResponsesAPIResponse(**raw_response_json)
except Exception:
verbose_logger.debug(
f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct"
)
response = ResponsesAPIResponse.model_construct(**raw_response_json)
# Store processed headers in additional_headers so they get returned to the client
response._hidden_params["additional_headers"] = processed_headers
response._hidden_params["headers"] = raw_response_headers
return response

View file

@ -771,9 +771,9 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
return ModelResponseStream(
id=chunk["id"],
object="chat.completion.chunk",
created=chunk["created"],
model=chunk["model"],
choices=chunk["choices"],
created=chunk.get("created"),
model=chunk.get("model"),
choices=chunk.get("choices", []),
)
except Exception as e:
raise e

View file

@ -115,7 +115,7 @@ def _process_gemini_image(
and (image_type := format or _get_image_mime_type_from_url(image_url))
is not None
):
file_data = FileDataType(file_uri=image_url, mime_type=image_type)
file_data = FileDataType(mime_type=image_type, file_uri=image_url)
part = {"file_data": file_data}
if media_resolution_enum is not None and model is not None:

View file

@ -40,6 +40,7 @@ class PartnerModelPrefixes(str, Enum):
GPT_OSS_PREFIX = "openai/gpt-oss-"
MINIMAX_PREFIX = "minimaxai/"
MOONSHOT_PREFIX = "moonshotai/"
ZAI_PREFIX = "zai-org/"
class VertexAIPartnerModels(VertexBase):
@ -66,6 +67,7 @@ class VertexAIPartnerModels(VertexBase):
or model.startswith(PartnerModelPrefixes.GPT_OSS_PREFIX)
or model.startswith(PartnerModelPrefixes.MINIMAX_PREFIX)
or model.startswith(PartnerModelPrefixes.MOONSHOT_PREFIX)
or model.startswith(PartnerModelPrefixes.ZAI_PREFIX)
):
return True
return False
@ -79,6 +81,7 @@ class VertexAIPartnerModels(VertexBase):
PartnerModelPrefixes.GPT_OSS_PREFIX,
PartnerModelPrefixes.MINIMAX_PREFIX,
PartnerModelPrefixes.MOONSHOT_PREFIX,
PartnerModelPrefixes.ZAI_PREFIX,
]
if any(provider in model for provider in OPENAI_LIKE_VERTEX_PROVIDERS):
return True

View file

@ -28345,6 +28345,19 @@
"supports_tool_choice": true,
"supports_web_search": true
},
"vertex_ai/zai-org/glm-4.7-maas": {
"input_cost_per_token": 3e-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,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"vertex_ai/mistral-medium-3": {
"input_cost_per_token": 4e-07,
"litellm_provider": "vertex_ai-mistral_models",

View file

@ -216,6 +216,11 @@ def llm_passthrough_route(
)
litellm_params_dict = get_litellm_params(**kwargs)
# Add model_id to litellm_params if present in kwargs (for Bedrock Application Inference Profiles)
if "model_id" in kwargs:
litellm_params_dict["model_id"] = kwargs["model_id"]
litellm_logging_obj.update_environment_variables(
model=model,
litellm_params=litellm_params_dict,

View file

@ -218,10 +218,10 @@ if MCP_AVAILABLE:
from fastapi import HTTPException
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.proxy.proxy_server import add_litellm_data_to_request, proxy_config
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy.proxy_server import add_litellm_data_to_request, proxy_config
try:
data = await request.json()
@ -252,7 +252,12 @@ if MCP_AVAILABLE:
if mcp_server_auth_headers:
data["mcp_server_auth_headers"] = mcp_server_auth_headers
data["raw_headers"] = raw_headers_from_request
# Extract user_api_key_auth from metadata and add to top level
# call_mcp_tool expects user_api_key_auth as a top-level parameter
if "metadata" in data and "user_api_key_auth" in data["metadata"]:
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
result = await call_mcp_tool(**data)
return result
except BlockedPiiEntityError as e:

View file

@ -336,8 +336,8 @@ class SkillsInjectionHook(CustomLogger):
)
# Check if code execution is enabled for this request
litellm_metadata = request_data.get("litellm_metadata", {})
metadata = request_data.get("metadata", {})
litellm_metadata = request_data.get("litellm_metadata") or {}
metadata = request_data.get("metadata") or {}
code_exec_enabled = (
litellm_metadata.get("_litellm_code_execution_enabled") or

View file

@ -3480,8 +3480,8 @@ class ProxyConfig:
def _deep_merge_dicts(dst: dict, src: dict) -> None:
"""
Deep-merge src into dst, skipping None values from src.
On conflicts, src (DB) wins.
Deep-merge src into dst, skipping None values and empty lists from src.
On conflicts, src (DB) wins, but empty lists are treated as "no value" and don't overwrite.
"""
stack = [(dst, src)]
while stack:
@ -3490,6 +3490,9 @@ class ProxyConfig:
if v is None:
# Preserve existing config when DB value is None (matches prior behavior)
continue
# Skip empty lists - treat them as "no value" to preserve file config
if isinstance(v, list) and len(v) == 0:
continue
if isinstance(v, dict) and isinstance(d.get(k), dict):
stack.append((d[k], v))
else:
@ -9840,6 +9843,18 @@ async def get_config(): # noqa: PLR0915
_failure_callbacks = _litellm_settings.get("failure_callback", [])
_success_and_failure_callbacks = _litellm_settings.get("callbacks", [])
# Normalize string callbacks to lists
def normalize_callback(callback):
if isinstance(callback, str):
return [callback]
elif callback is None:
return []
return callback
_success_callbacks = normalize_callback(_success_callbacks)
_failure_callbacks = normalize_callback(_failure_callbacks)
_success_and_failure_callbacks = normalize_callback(_success_and_failure_callbacks)
_data_to_return = []
"""
[

View file

@ -1195,7 +1195,7 @@ class ProxyLogging:
and _callback.__class__.async_pre_call_hook
!= CustomLogger.async_pre_call_hook
):
if call_type == "mcp_call" and user_api_key_dict is None:
if call_type == "call_mcp_tool" and user_api_key_dict is None:
continue
response = await _callback.async_pre_call_hook(

View file

@ -577,7 +577,13 @@ def responses(
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
)
# Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True)
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
if dynamic_api_base is not None:
litellm_params.api_base = dynamic_api_base
#########################################################
# Update input with provider-specific file IDs if managed files are used
#########################################################
@ -1483,6 +1489,12 @@ def compact_responses(
api_key=litellm_params.api_key,
)
# Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True)
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
if dynamic_api_base is not None:
litellm_params.api_base = dynamic_api_base
if custom_llm_provider is None:
raise ValueError("custom_llm_provider is required but passed as None")

View file

@ -255,6 +255,7 @@ class Router:
] = {},
enable_pre_call_checks: bool = False,
enable_tag_filtering: bool = False,
tag_filtering_match_any: bool = True,
retry_after: int = 0, # min time to wait before retrying a failed request
retry_policy: Optional[
Union[RetryPolicy, dict]
@ -363,6 +364,7 @@ class Router:
self.debug_level = debug_level
self.enable_pre_call_checks = enable_pre_call_checks
self.enable_tag_filtering = enable_tag_filtering
self.tag_filtering_match_any = tag_filtering_match_any
from litellm._service_logger import ServiceLogging
self.service_logger_obj: ServiceLogging = ServiceLogging()

View file

@ -20,17 +20,28 @@ else:
def is_valid_deployment_tag(
deployment_tags: List[str], request_tags: List[str]
deployment_tags: List[str], request_tags: List[str], match_any: bool = True
) -> bool:
"""
Check if a tag is valid
Check if a tag is valid, the matching can be either any or all based on `match_any` flag
"""
if not request_tags:
return False
if any(tag in deployment_tags for tag in request_tags):
dep_set = set(deployment_tags)
req_set = set(request_tags)
if match_any:
is_valid_deployment = bool(dep_set & req_set)
else:
is_valid_deployment = req_set.issubset(dep_set)
if is_valid_deployment:
verbose_logger.debug(
"adding deployment with tags: %s, request tags: %s",
"adding deployment with tags: %s, request tags: %s for match_any=%s",
deployment_tags,
request_tags,
match_any,
)
return True
return False
@ -68,6 +79,7 @@ async def get_deployments_for_tag(
if metadata_variable_name in request_kwargs:
metadata = request_kwargs[metadata_variable_name]
request_tags = metadata.get("tags")
match_any = llm_router_instance.tag_filtering_match_any
new_healthy_deployments = []
default_deployments = []
@ -76,7 +88,6 @@ async def get_deployments_for_tag(
"get_deployments_for_tag routing: router_keys: %s", request_tags
)
# example this can be router_keys=["free", "custom"]
# get all deployments that have a superset of these router keys
for deployment in healthy_deployments:
deployment_litellm_params = deployment.get("litellm_params")
deployment_tags = deployment_litellm_params.get("tags")
@ -90,7 +101,7 @@ async def get_deployments_for_tag(
if deployment_tags is None:
continue
if is_valid_deployment_tag(deployment_tags, request_tags):
if is_valid_deployment_tag(deployment_tags, request_tags, match_any):
new_healthy_deployments.append(deployment)
if "default" in deployment_tags:

View file

@ -184,6 +184,14 @@ ROUTER_SETTINGS_FIELDS: List[RouterSettingsField] = [
field_default=False,
ui_field_name="Enable Tag Filtering",
link="https://docs.litellm.ai/docs/proxy/tag_routing",
),
RouterSettingsField(
field_name="tag_filtering_match_any",
field_type="Boolean",
field_value=None,
field_description="Match any tag instead of all tags for tag-based routing",
field_default=True,
ui_field_name="Tag Filtering Match Any",
),
RouterSettingsField(
field_name="disable_cooldowns",

View file

@ -3016,6 +3016,7 @@ class LlmProviders(str, Enum):
AUTO_ROUTER = "auto_router"
VERCEL_AI_GATEWAY = "vercel_ai_gateway"
DOTPROMPT = "dotprompt"
MANUS = "manus"
WANDB = "wandb"
OVHCLOUD = "ovhcloud"
LEMONADE = "lemonade"
@ -3028,6 +3029,7 @@ class LlmProviders(str, Enum):
NANOGPT = "nano-gpt"
POE = "poe"
CHUTES = "chutes"
XIAOMI_MIMO = "xiaomi_mimo"

View file

@ -7871,6 +7871,8 @@ class ProviderConfigManager:
return litellm.GithubCopilotResponsesAPIConfig()
elif litellm.LlmProviders.LITELLM_PROXY == provider:
return litellm.LiteLLMProxyResponsesAPIConfig()
elif litellm.LlmProviders.MANUS == provider:
return litellm.ManusResponsesAPIConfig()
return None
@staticmethod

View file

@ -7800,7 +7800,7 @@
"litellm_provider": "deepseek",
"max_input_tokens": 131072,
"max_output_tokens": 8192,
"max_tokens": 131072,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.7e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
@ -7854,7 +7854,7 @@
"litellm_provider": "dashscope",
"max_input_tokens": 997952,
"max_output_tokens": 32768,
"max_tokens": 1000000,
"max_tokens": 32768,
"mode": "chat",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
@ -8579,7 +8579,7 @@
"litellm_provider": "databricks",
"max_input_tokens": 128000,
"max_output_tokens": 32000,
"max_tokens": 128000,
"max_tokens": 32000,
"metadata": {
"notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation."
},
@ -28345,6 +28345,19 @@
"supports_tool_choice": true,
"supports_web_search": true
},
"vertex_ai/zai-org/glm-4.7-maas": {
"input_cost_per_token": 3e-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,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"vertex_ai/mistral-medium-3": {
"input_cost_per_token": 4e-07,
"litellm_provider": "vertex_ai-mistral_models",

View file

@ -2304,6 +2304,24 @@
"messages": true,
"responses": true
}
},
"manus": {
"display_name": "Manus (`manus`)",
"url": "https://docs.litellm.ai/docs/providers/manus",
"endpoints": {
"chat_completions": true,
"messages": true,
"responses": true,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": true,
"interactions": true
}
}
},
"endpoints": {

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.80.12"
version = "1.80.13"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -167,7 +167,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.80.12"
version = "1.80.13"
version_files = [
"pyproject.toml:^version"
]

View file

@ -154,8 +154,8 @@ class TestAzureAIFlux2ImageEdit(BaseLLMImageEditTest):
return {
"model": "azure_ai/flux.2-pro",
"image": SINGLE_TEST_IMAGE,
"api_base": os.getenv("AZURE_AI_API_BASE", "https://litellm-ci-cd-prod.services.ai.azure.com"),
"api_key": os.getenv("AZURE_AI_API_KEY"),
"api_base": "https://litellm-ci-cd-prod.services.ai.azure.com",
"api_key": os.getenv("AZURE_API_KEY"),
"api_version": "preview",
}

View file

@ -54,9 +54,10 @@ def validate_responses_api_response(response, final_chunk: bool = False):
assert "created_at" in response and isinstance(
response["created_at"], int
), "Response should have an integer 'created_at' field"
assert "output" in response and isinstance(
response["output"], list
), "Response should have a list 'output' field"
if response.get("status") == "completed":
assert "output" in response and isinstance(
response["output"], list
), "Response should have a list 'output' field"
# Optional fields with their expected types
optional_fields = {
@ -91,7 +92,7 @@ def validate_responses_api_response(response, final_chunk: bool = False):
), f"Field '{field}' should be of type {expected_type}, but got {type(response[field])}"
# Check if output has at least one item
if final_chunk is True:
if final_chunk is True and response.get("status") == "completed":
assert (
len(response["output"]) > 0
), "Response 'output' field should have at least one item"
@ -170,48 +171,57 @@ class BaseResponsesAPITest(ABC):
elif event.type == "response.completed":
response_completed_event = event
# assert the delta chunks content had len(collected_content_string) > 0
# this content is typically rendered on chat ui's
assert len(collected_content_string) > 0
# assert the response completed event is not None
assert response_completed_event is not None
# assert the response completed event has a response
assert response_completed_event.response is not None
# assert the response completed event includes the usage
assert response_completed_event.response.usage is not None
# For async agent APIs (like Manus), the response may be in 'running' state
# without content yet - this is valid behavior
response_status = response_completed_event.response.status
if response_status in ["running", "pending"]:
# Running/pending state is acceptable - task started successfully
print(f"Response is in '{response_status}' state - async agent API behavior")
assert response_completed_event.response.id is not None
else:
# For completed responses, validate content and usage
# assert the delta chunks content had len(collected_content_string) > 0
# this content is typically rendered on chat ui's
assert len(collected_content_string) > 0
# basic test assert the usage seems reasonable
print(
"response_completed_event.response.usage=",
response_completed_event.response.usage,
)
assert (
response_completed_event.response.usage.input_tokens > 0
and response_completed_event.response.usage.input_tokens < 100
)
assert (
response_completed_event.response.usage.output_tokens > 0
and response_completed_event.response.usage.output_tokens < 2000
)
assert (
response_completed_event.response.usage.total_tokens > 0
and response_completed_event.response.usage.total_tokens < 2000
)
# assert the response completed event includes the usage
assert response_completed_event.response.usage is not None
# total tokens should be the sum of input and output tokens
assert (
response_completed_event.response.usage.total_tokens
== response_completed_event.response.usage.input_tokens
+ response_completed_event.response.usage.output_tokens
)
# basic test assert the usage seems reasonable
print(
"response_completed_event.response.usage=",
response_completed_event.response.usage,
)
assert (
response_completed_event.response.usage.input_tokens > 0
and response_completed_event.response.usage.input_tokens < 100
)
assert (
response_completed_event.response.usage.output_tokens > 0
and response_completed_event.response.usage.output_tokens < 2000
)
assert (
response_completed_event.response.usage.total_tokens > 0
and response_completed_event.response.usage.total_tokens < 2000
)
# assert the response completed event includes cost when include_cost_in_streaming_usage is True
assert hasattr(response_completed_event.response.usage, "cost"), "Cost should be included in streaming responses API usage object"
assert response_completed_event.response.usage.cost > 0, "Cost should be greater than 0"
print(f"Cost found in streaming response: {response_completed_event.response.usage.cost}")
# total tokens should be the sum of input and output tokens
assert (
response_completed_event.response.usage.total_tokens
== response_completed_event.response.usage.input_tokens
+ response_completed_event.response.usage.output_tokens
)
# assert the response completed event includes cost when include_cost_in_streaming_usage is True
assert hasattr(response_completed_event.response.usage, "cost"), "Cost should be included in streaming responses API usage object"
assert response_completed_event.response.usage.cost > 0, "Cost should be greater than 0"
print(f"Cost found in streaming response: {response_completed_event.response.usage.cost}")
# Reset the setting
litellm.include_cost_in_streaming_usage = False
@ -450,7 +460,13 @@ class BaseResponsesAPITest(ABC):
# Additional assertions specific to tool calls
assert response is not None
assert "output" in response
assert len(response["output"]) > 0
# For async agent APIs (like Manus), the response may be in 'running' state
# without output yet - this is valid behavior
if response.get("status") in ["running", "pending"]:
print(f"Response is in '{response.get('status')}' state - async agent API behavior")
assert response.get("id") is not None
else:
assert len(response["output"]) > 0
@pytest.mark.asyncio
async def test_responses_api_multi_turn_with_reasoning_and_structured_output(self):

View file

@ -0,0 +1,115 @@
import os
import sys
import pytest
import asyncio
from typing import Optional
from unittest.mock import patch, AsyncMock
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.integrations.custom_logger import CustomLogger
import json
from litellm.types.utils import StandardLoggingPayload
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponseAPIUsage,
IncompleteDetails,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from base_responses_api import BaseResponsesAPITest
# class TestManusResponsesAPITest(BaseResponsesAPITest):
# def get_base_completion_call_args(self):
# return {
# "model": "manus/manus-1.6",
# "api_key": os.getenv("MANUS_API_KEY"),
# }
# @pytest.mark.parametrize("sync_mode", [True, False])
# @pytest.mark.asyncio
# async def test_basic_openai_responses_delete_endpoint(self, sync_mode):
# pytest.skip("DELETE responses is not supported for Manus")
# @pytest.mark.parametrize("sync_mode", [True, False])
# @pytest.mark.asyncio
# async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode):
# pytest.skip("DELETE responses is not supported for Manus")
# # GET responses is now supported for Manus
# @pytest.mark.parametrize("sync_mode", [True, False])
# @pytest.mark.asyncio
# async def test_basic_openai_responses_get_endpoint(self, sync_mode):
# pytest.skip("GET responses is not supported for Manus")
# @pytest.mark.parametrize("sync_mode", [True, False])
# @pytest.mark.asyncio
# async def test_basic_openai_responses_cancel_endpoint(self, sync_mode):
# pytest.skip("CANCEL responses is not supported for Manus")
# @pytest.mark.parametrize("sync_mode", [True, False])
# @pytest.mark.asyncio
# async def test_cancel_responses_invalid_response_id(self, sync_mode):
# pytest.skip("CANCEL responses is not supported for Manus")
# @pytest.mark.asyncio
# async def test_multiturn_responses_api(self):
# pytest.skip("Multiturn responses is not supported for Manus")
# @pytest.mark.asyncio
# async def test_manus_responses_api_with_agent_profile():
# """
# Test that Manus API correctly extracts agent profile from model name
# and includes task_mode and agent_profile in the request.
# """
# litellm._turn_on_debug()
# response = await litellm.aresponses(
# model="manus/manus-1.6",
# input="What's the color of the sky?",
# api_key=os.getenv("MANUS_API_KEY"),
# max_output_tokens=50,
# )
# print("Manus response=", json.dumps(response, indent=4, default=str))
# # Validate response structure
# assert isinstance(response, ResponsesAPIResponse), "Response should be ResponsesAPIResponse"
# assert response.id is not None, "Response should have an ID"
# assert response.status in ["running", "completed", "pending"], f"Status should be valid, got {response.status}"
# # Check that metadata includes Manus-specific fields
# if response.metadata:
# assert "task_id" in response.metadata or "task_url" in response.metadata, (
# "Manus response should include task_id or task_url in metadata"
# )
# @pytest.mark.asyncio
# async def test_manus_responses_api_different_agent_profiles():
# """
# Test that different agent profiles work correctly.
# """
# litellm._turn_on_debug()
# # Test with different agent profile variants
# agent_profiles = ["manus-1.6", "manus-1.6-lite", "manus-1.6-max"]
# for profile in agent_profiles:
# try:
# response = await litellm.aresponses(
# model=f"manus/{profile}",
# input="Hello",
# api_key=os.getenv("MANUS_API_KEY"),
# max_output_tokens=20,
# )
# assert response.id is not None, f"Response for {profile} should have an ID"
# print(f"✓ {profile} works: {response.id}")
# except Exception as e:
# # Some profiles might not be available, that's okay
# print(f"⚠ {profile} not available: {e}")
# pass

View file

@ -0,0 +1,290 @@
"""
Tests for Bedrock Moonshot (Kimi K2) integration.
This test suite verifies:
1. Basic completion functionality
2. Streaming responses
3. System message support
4. Temperature parameter handling
5. Reasoning content extraction from <reasoning> tags
6. Tool calling support (including tool response handling)
7. Parameter validation (e.g., stop sequences not supported)
"""
from base_llm_unit_tests import BaseLLMChatTest
import pytest
import sys
import os
import json
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.llms.bedrock.common_utils import get_bedrock_chat_config
class TestBedrockMoonshotInvoke(BaseLLMChatTest):
"""
Test suite for Bedrock Moonshot via invoke route.
Inherits all standard LLM tests from BaseLLMChatTest.
"""
def get_base_completion_call_args(self) -> dict:
litellm._turn_on_debug()
return {
"model": "bedrock/invoke/moonshot.kimi-k2-thinking",
}
def test_tool_call_no_arguments(self, tool_call_no_arguments):
"""Test that tool calls with no arguments is translated correctly."""
pass
class TestBedrockMoonshotBasic:
"""Unit tests for Bedrock Moonshot configuration and transformations."""
def test_provider_detection_invoke(self):
"""Test that Bedrock Moonshot invoke models are correctly detected."""
config = get_bedrock_chat_config("bedrock/invoke/moonshot.kimi-k2-thinking")
assert config is not None
assert config.__class__.__name__ == "AmazonMoonshotConfig"
def test_provider_detection_converse(self):
"""Test that Bedrock Moonshot converse models are correctly detected."""
config = get_bedrock_chat_config("bedrock/moonshot.kimi-k2-thinking")
assert config is not None
def test_config_initialization(self):
"""Test that AmazonMoonshotConfig initializes correctly."""
config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking")
assert config is not None
assert config.custom_llm_provider == "bedrock"
def test_supported_params(self):
"""Test that supported OpenAI params are correctly defined."""
config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking")
supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking")
# Should support these params
assert "temperature" in supported_params
assert "max_tokens" in supported_params
assert "top_p" in supported_params
assert "stream" in supported_params
assert "tools" in supported_params
assert "tool_choice" in supported_params
# Should NOT support stop sequences on Bedrock
assert "stop" not in supported_params
# Should NOT support functions (use tools instead)
assert "functions" not in supported_params
def test_transform_request_strips_model_prefix(self):
"""Test that model ID prefixes are correctly stripped in transform_request."""
from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import (
AmazonMoonshotConfig,
)
config = AmazonMoonshotConfig()
messages = [{"role": "user", "content": "Hello"}]
# Test that bedrock/invoke/ prefix is stripped
transformed = config.transform_request(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=messages,
optional_params={},
litellm_params={},
headers={}
)
# The model ID in the request body should be stripped
assert transformed["model"] == "moonshot.kimi-k2-thinking"
class TestBedrockMoonshotReasoningContent:
"""Tests for reasoning content extraction."""
def test_reasoning_content_extraction(self):
"""Test that reasoning content is extracted from <reasoning> tags."""
from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import (
AmazonMoonshotConfig,
)
config = AmazonMoonshotConfig()
# Test with reasoning tags
content_with_reasoning = "<reasoning>This is my thought process</reasoning>This is the answer"
reasoning, content = config._extract_reasoning_from_content(content_with_reasoning)
assert reasoning == "This is my thought process"
assert content == "This is the answer"
assert "<reasoning>" not in content
# Test without reasoning tags
content_without_reasoning = "This is just a regular answer"
reasoning, content = config._extract_reasoning_from_content(content_without_reasoning)
assert reasoning is None
assert content == "This is just a regular answer"
class TestBedrockMoonshotToolCalling:
"""Unit tests for tool calling functionality."""
def test_tool_calling_supported(self):
"""Test that tool calling is supported for Kimi K2 Thinking model."""
config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking")
supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking")
# Kimi K2 Thinking DOES support tool calls (unlike kimi-thinking-preview)
assert "tools" in supported_params
assert "tool_choice" in supported_params
def test_tool_call_request_format(self):
"""Test that tool call requests are formatted correctly."""
from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import (
AmazonMoonshotConfig,
)
config = AmazonMoonshotConfig()
messages = [
{"role": "user", "content": "What's the weather in San Francisco?"}
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather",
"parameters": {
"type": "object",
"properties": {
"location": {"type": "string"}
},
"required": ["location"]
}
}
}
]
}
transformed = config.transform_request(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={}
)
# Verify model ID is stripped
assert transformed["model"] == "moonshot.kimi-k2-thinking"
# Verify tools are included
assert "tools" in transformed
assert len(transformed["tools"]) == 1
assert transformed["tools"][0]["function"]["name"] == "get_weather"
def test_tool_response_message_format(self):
"""Test that tool response messages are formatted correctly."""
# This tests the proper format for sending tool responses back
tool_response_message = {
"role": "tool",
"tool_call_id": "call_123",
"content": json.dumps({"temperature": 72, "condition": "sunny"})
}
# Verify the message structure
assert tool_response_message["role"] == "tool"
assert "tool_call_id" in tool_response_message
assert "content" in tool_response_message
class TestBedrockMoonshotParameterValidation:
"""Tests for parameter validation and edge cases."""
def test_stop_sequences_not_supported(self):
"""Test that stop sequences are correctly excluded from supported params."""
config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking")
supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking")
# Bedrock Moonshot doesn't support stopSequences field
assert "stop" not in supported_params
def test_temperature_range(self):
"""Test that temperature parameter is handled correctly."""
# Moonshot models support temperature 0-1
# This is handled by the parent MoonshotChatConfig class
config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking")
# Verify config exists and can handle temperature
assert config is not None
supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking")
assert "temperature" in supported_params
class TestBedrockMoonshotTransformations:
"""Tests for request/response transformations."""
def test_transform_request_basic(self):
"""Test basic request transformation."""
from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import (
AmazonMoonshotConfig,
)
config = AmazonMoonshotConfig()
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"}
]
optional_params = {
"temperature": 0.7,
"max_tokens": 100
}
transformed = config.transform_request(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={}
)
# Verify model ID is stripped
assert transformed["model"] == "moonshot.kimi-k2-thinking"
# Verify messages are included
assert "messages" in transformed
assert len(transformed["messages"]) >= 1
# Verify optional params are included
assert transformed["temperature"] == 0.7
assert transformed["max_tokens"] == 100
def test_transform_request_with_system_message(self):
"""Test request transformation with system message."""
from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import (
AmazonMoonshotConfig,
)
config = AmazonMoonshotConfig()
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"}
]
transformed = config.transform_request(
model="moonshot.kimi-k2-thinking",
messages=messages,
optional_params={},
litellm_params={},
headers={}
)
# System messages should be supported
assert "messages" in transformed

View file

@ -1095,3 +1095,186 @@ def test_map_reasoning_effort_adds_summary_detailed():
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = original_env
elif "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
def test_transform_response_preserves_annotations():
"""
Test that annotations from Responses API are preserved when transforming to Chat Completions format.
This is a regression test for the bug where annotations (like url_citation) were being
dropped during the transformation from ResponsesAPIResponse to ModelResponse.
The fix ensures annotations are extracted from ResponseOutputText content items and
passed through to the Message object in the Chat Completions response.
"""
from unittest.mock import Mock
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
from litellm.types.llms.openai import (
InputTokensDetails,
OutputTokensDetails,
ResponseAPIUsage,
ResponsesAPIResponse,
)
from litellm.types.utils import ModelResponse, Usage
handler = LiteLLMResponsesTransformationHandler()
# Create annotations similar to what OpenAI Responses API returns
annotations = [
{
"type": "url_citation",
"start_index": 0,
"end_index": 100,
"title": "Example Article",
"url": "https://example.com/article",
},
{
"type": "url_citation",
"start_index": 101,
"end_index": 200,
"title": "Another Source",
"url": "https://example.com/source",
},
]
# Create output text with annotations
output_text = ResponseOutputText(
annotations=annotations,
text="Here is some information with citations.",
type="output_text",
logprobs=[],
)
# Create output message
output_message = ResponseOutputMessage(
id="msg_test123",
content=[output_text],
role="assistant",
status="completed",
type="message",
)
# Create usage information
usage = ResponseAPIUsage(
input_tokens=10,
input_tokens_details=InputTokensDetails(
audio_tokens=None, cached_tokens=0, text_tokens=None
),
output_tokens=20,
output_tokens_details=OutputTokensDetails(
reasoning_tokens=0, text_tokens=None
),
total_tokens=30,
cost=None,
)
# Create the full ResponsesAPIResponse
raw_response = ResponsesAPIResponse(
id="resp_test123",
created_at=1234567890,
error=None,
incomplete_details=None,
instructions=None,
metadata={},
model="gpt-5.1",
object="response",
output=[output_message],
parallel_tool_calls=True,
temperature=1.0,
tool_choice="auto",
tools=[],
top_p=1.0,
max_output_tokens=None,
previous_response_id=None,
reasoning=None,
status="completed",
text={"format": {"type": "text"}, "verbosity": "medium"},
truncation="disabled",
usage=usage,
user=None,
store=True,
background=False,
billing={"payer": "openai"},
max_tool_calls=None,
prompt_cache_key=None,
safety_identifier=None,
service_tier="default",
top_logprobs=0,
)
# Create empty model_response
model_response = ModelResponse(
id="chatcmpl-test123",
created=1234567890,
model=None,
object="chat.completion",
system_fingerprint=None,
choices=[],
usage=Usage(completion_tokens=0, prompt_tokens=0, total_tokens=0),
)
# Create mock objects for required parameters
logging_obj = Mock()
messages = [{"role": "user", "content": "Tell me about AI"}]
request_data = {"model": "gpt-5.1"}
optional_params = {}
litellm_params = {"acompletion": False, "api_key": None}
encoding = Mock()
# Call transform_response
result = handler.transform_response(
model="gpt-5.1",
raw_response=raw_response,
model_response=model_response,
logging_obj=logging_obj,
request_data=request_data,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=encoding,
api_key=None,
json_mode=None,
)
# Assertions
assert result.model == "gpt-5.1"
assert len(result.choices) == 1
# Check the choice
choice = result.choices[0]
assert choice.finish_reason == "stop"
assert choice.index == 0
assert choice.message.role == "assistant"
assert choice.message.content == "Here is some information with citations."
# Check that annotations are preserved
assert hasattr(choice.message, "annotations"), "Message should have annotations attribute"
assert choice.message.annotations is not None, "Annotations should not be None"
assert len(choice.message.annotations) == 2, f"Expected 2 annotations, got {len(choice.message.annotations)}"
# Verify annotation content
annotation1 = choice.message.annotations[0]
assert annotation1["type"] == "url_citation"
assert annotation1["title"] == "Example Article"
assert annotation1["url"] == "https://example.com/article"
assert annotation1["start_index"] == 0
assert annotation1["end_index"] == 100
annotation2 = choice.message.annotations[1]
assert annotation2["type"] == "url_citation"
assert annotation2["title"] == "Another Source"
assert annotation2["url"] == "https://example.com/source"
assert annotation2["start_index"] == 101
assert annotation2["end_index"] == 200
# Check usage
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 20
assert result.usage.total_tokens == 30
print("✓ Annotations from Responses API are correctly preserved in Chat Completions format")

View file

@ -0,0 +1,161 @@
"""
Unit tests for Prometheus invalid API key request filtering.
Tests functionality that prevents invalid API key requests (401 status codes)
from being recorded in Prometheus metrics.
"""
import os
import sys
from unittest.mock import Mock, patch
import pytest
from prometheus_client import REGISTRY
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy._types import UserAPIKeyAuth
@pytest.fixture(scope="function")
def prometheus_logger():
"""Create a PrometheusLogger instance for testing."""
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
REGISTRY.unregister(collector)
return PrometheusLogger()
class ExceptionWithCode:
"""Exception-like object with 'code' attribute (ProxyException pattern)."""
def __init__(self, code):
self.code = code
class ExceptionWithStatusCode:
"""Exception-like object with 'status_code' attribute."""
def __init__(self, status_code):
self.status_code = status_code
class TestExtractStatusCode:
"""Test status code extraction from various sources."""
@pytest.mark.parametrize("exception_class,code_value,expected", [
(ExceptionWithCode, "401", 401),
(ExceptionWithStatusCode, 401, 401),
])
def test_extract_from_exception(self, prometheus_logger, exception_class, code_value, expected):
exception = exception_class(code_value)
assert prometheus_logger._extract_status_code(exception=exception) == expected
def test_extract_from_kwargs(self, prometheus_logger):
exception = ExceptionWithCode("401")
assert prometheus_logger._extract_status_code(kwargs={"exception": exception}) == 401
def test_extract_from_enum_values(self, prometheus_logger):
enum_values = Mock(status_code="401")
assert prometheus_logger._extract_status_code(enum_values=enum_values) == 401
class TestInvalidAPIKeyDetection:
"""Test invalid API key request detection logic."""
@pytest.mark.parametrize("status_code,expected", [
(401, True),
(200, False),
(500, False),
(None, False),
])
def test_status_code_detection(self, prometheus_logger, status_code, expected):
assert prometheus_logger._is_invalid_api_key_request(status_code=status_code) == expected
def test_auth_error_message_detection(self, prometheus_logger):
exception = AssertionError("LiteLLM Virtual Key expected. Received=invalid-key-12345, expected to start with 'sk-'.")
assert prometheus_logger._is_invalid_api_key_request(status_code=None, exception=exception) is True
def test_non_auth_exception_not_detected(self, prometheus_logger):
exception = ValueError("Some other error")
assert prometheus_logger._is_invalid_api_key_request(status_code=None, exception=exception) is False
class TestSkipMetricsValidation:
"""Test high-level validation method that orchestrates detection and extraction."""
def test_skip_for_401_exception(self, prometheus_logger):
"""Test full flow: extraction -> detection -> skip decision."""
exception = ExceptionWithCode("401")
assert prometheus_logger._should_skip_metrics_for_invalid_key(exception=exception) is True
def test_skip_for_auth_error_message(self, prometheus_logger):
"""Test full flow: exception message -> detection -> skip decision."""
exception = AssertionError("expected to start with 'sk-'")
assert prometheus_logger._should_skip_metrics_for_invalid_key(exception=exception) is True
def test_no_skip_for_valid_request(self, prometheus_logger):
assert prometheus_logger._should_skip_metrics_for_invalid_key() is False
class TestAsyncHooks:
"""Test async hook methods skip metrics for invalid API keys."""
@pytest.fixture
def mock_user_api_key(self):
"""Create a mock UserAPIKeyAuth object."""
user_key = Mock(spec=UserAPIKeyAuth)
user_key.api_key = "test-key"
user_key.end_user_id = None
user_key.user_id = None
user_key.user_email = None
user_key.key_alias = None
user_key.team_id = None
user_key.team_alias = None
user_key.request_route = "/test"
return user_key
@pytest.mark.asyncio
async def test_post_call_failure_hook_skips_401(self, prometheus_logger, mock_user_api_key):
exception = ExceptionWithCode("401")
exception.__class__.__name__ = "ProxyException"
with patch.object(prometheus_logger, 'litellm_proxy_failed_requests_metric') as mock_failed, \
patch.object(prometheus_logger, 'litellm_proxy_total_requests_metric') as mock_total:
await prometheus_logger.async_post_call_failure_hook(
request_data={"model": "test-model"},
original_exception=exception,
user_api_key_dict=mock_user_api_key
)
mock_failed.labels.assert_not_called()
mock_total.labels.assert_not_called()
@pytest.mark.asyncio
async def test_log_failure_event_skips_401(self, prometheus_logger):
exception = ExceptionWithCode("401")
kwargs = {
"model": "test-model",
"standard_logging_object": {
"metadata": {
"user_api_key_hash": "test-key",
"user_api_key_user_id": "test-user",
},
"model_group": "test-model",
},
"exception": exception,
"litellm_params": {},
}
with patch.object(prometheus_logger, 'litellm_llm_api_failed_requests_metric') as mock_failed, \
patch.object(prometheus_logger, 'set_llm_deployment_failure_metrics') as mock_deployment:
await prometheus_logger.async_log_failure_event(
kwargs=kwargs,
response_obj=None,
start_time=None,
end_time=None
)
mock_failed.labels.assert_not_called()
mock_deployment.assert_not_called()

View file

@ -75,3 +75,47 @@ def test_excluded_keys_exact_match():
assert masked["api_key"] == "sk-1234567890abcdef" # Should NOT be masked
assert masked["access_token"] != "token-12345" # Should still be masked
assert "*" in masked["access_token"]
def test_extra_headers_are_masked_recursively():
"""
Ensure nested dictionaries (like extra_headers) are masked.
"""
masker = SensitiveDataMasker()
data = {
"litellm_params": {
"model": "openai/gpt-4",
"extra_headers": {
"rits_api_key": "sk-secret-12345-very-sensitive",
"Authorization": "Bearer token123",
},
}
}
masked = masker.mask_dict(data)
extra_headers = masked["litellm_params"]["extra_headers"]
assert extra_headers["rits_api_key"] != "sk-secret-12345-very-sensitive"
assert "*" in extra_headers["rits_api_key"]
assert extra_headers["Authorization"] != "Bearer token123"
assert "*" in extra_headers["Authorization"]
def test_lists_with_sensitive_keys_are_masked():
"""
Lists belonging to sensitive keys should have their values masked.
"""
masker = SensitiveDataMasker()
data = {
"api_key": ["sk-123", "sk-456"],
"tags": ["prod", "test"],
}
masked = masker.mask_dict(data)
# sensitive key list entries should be masked
assert masked["api_key"][0] != "sk-123"
assert "*" in masked["api_key"][0]
# non-sensitive list should remain unchanged
assert masked["tags"] == ["prod", "test"]

View file

@ -1754,3 +1754,61 @@ def test_transform_request_respects_user_max_tokens():
)
assert result["max_tokens"] == 1000
def test_calculate_usage_completion_tokens_details_always_populated():
"""
Test that completion_tokens_details is always populated in Usage object,
not just when there's reasoning_content.
Fixes: https://github.com/BerriAI/litellm/issues/18772
Bug: completion_tokens_details was None for regular Claude responses without reasoning
"""
config = AnthropicConfig()
# Test without reasoning_content - completion_tokens_details should still be populated
usage_object = {
"input_tokens": 37,
"output_tokens": 248,
}
usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None)
# completion_tokens_details should NOT be None
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens is None
assert usage.completion_tokens_details.text_tokens == 248
assert usage.completion_tokens == 248
assert usage.prompt_tokens == 37
assert usage.total_tokens == 285
def test_calculate_usage_completion_tokens_details_with_reasoning():
"""
Test that completion_tokens_details correctly splits text_tokens and reasoning_tokens
when reasoning_content is present.
Fixes: https://github.com/BerriAI/litellm/issues/18772
"""
config = AnthropicConfig()
# Test with reasoning_content - should split tokens correctly
usage_object = {
"input_tokens": 100,
"output_tokens": 500,
}
# Simulating reasoning content that would count as ~50 tokens
reasoning_content = "Let me think about this step by step. " * 10 # Roughly 50 tokens
usage = config.calculate_usage(
usage_object=usage_object,
reasoning_content=reasoning_content
)
# completion_tokens_details should be populated with both reasoning and text tokens
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens is not None
assert usage.completion_tokens_details.reasoning_tokens > 0
# text_tokens should be total minus reasoning
expected_text_tokens = 500 - usage.completion_tokens_details.reasoning_tokens
assert usage.completion_tokens_details.text_tokens == expected_text_tokens
assert usage.completion_tokens == 500

View file

@ -175,3 +175,132 @@ def test_format_url_handles_trailing_slash_normalization():
assert str(result_with_slash) == "http://proxy.com/bedrockproxy/model/test/invoke"
def test_bedrock_passthrough_with_application_inference_profile():
"""
Test get_complete_url with Application Inference Profile ARN as model_id.
This test verifies the fix for GitHub issue #18761 where Bedrock passthrough
was not working with Application Inference Profiles. The model_id (ARN) should
replace the translated model name in the endpoint URL.
"""
config = BedrockPassthroughConfig()
model = "anthropic.claude-sonnet-4-20250514-v1:0"
model_id = "arn:aws:bedrock:eu-west-1:123456789:application-inference-profile/abcdefgh1234"
endpoint = f"model/{model}/invoke"
with patch.object(config, '_get_aws_region_name', return_value="eu-west-1"), \
patch.object(config, 'get_runtime_endpoint', return_value=(
"https://bedrock-runtime.eu-west-1.amazonaws.com",
"https://bedrock-runtime.eu-west-1.amazonaws.com"
)):
url, api_base = config.get_complete_url(
api_base=None,
api_key=None,
model=model,
endpoint=endpoint,
request_query_params=None,
litellm_params={"model_id": model_id, "aws_region_name": "eu-west-1"}
)
# Verify that the URL contains the model_id (ARN) instead of the model name
url_str = str(url)
assert model_id in url_str, f"Expected model_id ARN in URL, but got: {url_str}"
assert model not in url_str, f"Model name should be replaced by model_id, but got: {url_str}"
assert "/invoke" in url_str, "Expected /invoke action in URL"
# Verify the complete URL structure
expected_url = f"https://bedrock-runtime.eu-west-1.amazonaws.com/model/{model_id}/invoke"
assert url_str == expected_url, f"Expected {expected_url}, but got: {url_str}"
def test_bedrock_passthrough_with_inference_profile_converse_endpoint():
"""Test Application Inference Profile with converse endpoint"""
config = BedrockPassthroughConfig()
model = "anthropic.claude-sonnet-4-20250514-v1:0"
model_id = "arn:aws:bedrock:us-east-1:123456789:application-inference-profile/xyz123"
endpoint = f"model/{model}/converse"
with patch.object(config, '_get_aws_region_name', return_value="us-east-1"), \
patch.object(config, 'get_runtime_endpoint', return_value=(
"https://bedrock-runtime.us-east-1.amazonaws.com",
"https://bedrock-runtime.us-east-1.amazonaws.com"
)):
url, api_base = config.get_complete_url(
api_base=None,
api_key=None,
model=model,
endpoint=endpoint,
request_query_params=None,
litellm_params={"model_id": model_id}
)
url_str = str(url)
assert model_id in url_str
assert "/converse" in url_str
assert model not in url_str
def test_bedrock_passthrough_without_model_id_backward_compatibility():
"""
Test that passthrough still works without model_id (backward compatibility).
When model_id is not provided, the system should use the model name as before.
"""
config = BedrockPassthroughConfig()
model = "anthropic.claude-3-sonnet"
endpoint = f"model/{model}/invoke"
with patch.object(config, '_get_aws_region_name', return_value="us-east-1"), \
patch.object(config, 'get_runtime_endpoint', return_value=(
"https://bedrock-runtime.us-east-1.amazonaws.com",
"https://bedrock-runtime.us-east-1.amazonaws.com"
)):
url, api_base = config.get_complete_url(
api_base=None,
api_key=None,
model=model,
endpoint=endpoint,
request_query_params=None,
litellm_params={} # No model_id provided
)
# Verify that the URL contains the model name (not replaced)
url_str = str(url)
assert model in url_str, f"Expected model name in URL when model_id not provided, but got: {url_str}"
expected_url = f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{model}/invoke"
assert url_str == expected_url
def test_bedrock_passthrough_region_extraction_from_inference_profile_arn():
"""Test that AWS region is correctly extracted from Application Inference Profile ARN"""
config = BedrockPassthroughConfig()
model = "anthropic.claude-sonnet-4-20250514-v1:0"
# ARN contains us-west-2 region
model_id = "arn:aws:bedrock:us-west-2:123456789:application-inference-profile/test123"
endpoint = f"model/{model}/invoke"
# Don't provide aws_region_name in litellm_params to test ARN extraction
with patch.object(config, 'get_runtime_endpoint', return_value=(
"https://bedrock-runtime.us-west-2.amazonaws.com",
"https://bedrock-runtime.us-west-2.amazonaws.com"
)):
url, api_base = config.get_complete_url(
api_base=None,
api_key=None,
model=model,
endpoint=endpoint,
request_query_params=None,
litellm_params={"model_id": model_id} # Region should be extracted from ARN
)
# Verify that the region from ARN is used in the base URL
assert "us-west-2" in api_base, f"Expected region 'us-west-2' from ARN in base URL, but got: {api_base}"

View file

@ -24,3 +24,194 @@ def test_deepseek_supported_openai_params():
supported_openai_params = DeepInfraConfig().get_supported_openai_params(model="deepinfra/deepseek-ai/DeepSeek-V3.1")
print(supported_openai_params)
assert "reasoning_effort" in supported_openai_params
def test_deepinfra_tool_message_content_transformation():
"""
Test that DeepInfra transforms tool message content from array to string.
This fixes the issue where LibreChat sends tool messages with content as an array:
{"role": "tool", "content": [{"type": "text", "text": "20"}]}
DeepInfra requires content to be a string, so we transform it to:
{"role": "tool", "content": "20"}
Related to issue #13982
"""
from litellm.llms.deepinfra.chat.transformation import DeepInfraConfig
config = DeepInfraConfig()
# Test case 1: Simple single text item in array (common case from LibreChat)
messages_with_array_content = [
{
"role": "user",
"content": "Calculate 10 + 10"
},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "calculator",
"arguments": '{"input": "10 + 10"}'
}
}
]
},
{
"role": "tool",
"tool_call_id": "call_123",
"name": "calculator",
"content": [{"type": "text", "text": "20"}] # Array format from LibreChat
}
]
transformed_messages = config._transform_messages(
messages=messages_with_array_content,
model="deepinfra/Qwen/Qwen3-235B-A22B"
)
# Verify the tool message content was converted to string
tool_message = transformed_messages[2]
assert tool_message["role"] == "tool"
assert isinstance(tool_message["content"], str)
assert tool_message["content"] == "20"
print(f"✓ Test case 1 passed: {tool_message['content']}")
# Test case 2: Complex content array (multiple items)
messages_with_complex_content = [
{
"role": "user",
"content": "Test"
},
{
"role": "assistant",
"tool_calls": [
{
"id": "call_456",
"type": "function",
"function": {"name": "test", "arguments": "{}"}
}
]
},
{
"role": "tool",
"tool_call_id": "call_456",
"content": [
{"type": "text", "text": "Result 1"},
{"type": "text", "text": "Result 2"}
]
}
]
transformed_messages_complex = config._transform_messages(
messages=messages_with_complex_content,
model="deepinfra/Qwen/Qwen3-235B-A22B"
)
tool_message_complex = transformed_messages_complex[2]
assert tool_message_complex["role"] == "tool"
assert isinstance(tool_message_complex["content"], str)
# For complex content, it should be JSON stringified
parsed_content = json.loads(tool_message_complex["content"])
assert len(parsed_content) == 2
assert parsed_content[0]["text"] == "Result 1"
print(f"✓ Test case 2 passed: {tool_message_complex['content']}")
# Test case 3: Tool message with string content (should remain unchanged)
messages_with_string_content = [
{
"role": "user",
"content": "Test"
},
{
"role": "assistant",
"tool_calls": [
{
"id": "call_789",
"type": "function",
"function": {"name": "test", "arguments": "{}"}
}
]
},
{
"role": "tool",
"tool_call_id": "call_789",
"content": "Simple string result" # Already a string
}
]
transformed_messages_string = config._transform_messages(
messages=messages_with_string_content,
model="deepinfra/Qwen/Qwen3-235B-A22B"
)
tool_message_string = transformed_messages_string[2]
assert tool_message_string["role"] == "tool"
assert isinstance(tool_message_string["content"], str)
assert tool_message_string["content"] == "Simple string result"
print(f"✓ Test case 3 passed: {tool_message_string['content']}")
print("\n✅ All DeepInfra tool message transformation tests passed!")
@pytest.mark.asyncio
async def test_deepinfra_tool_message_content_transformation_async():
"""
Test that DeepInfra transforms tool message content from array to string in async mode.
This ensures the async path works correctly when is_async=True.
Related to issue #13982
"""
from litellm.llms.deepinfra.chat.transformation import DeepInfraConfig
config = DeepInfraConfig()
# Test async transformation with tool message containing array content
messages_with_array_content = [
{
"role": "user",
"content": "Calculate 10 + 10"
},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "calculator",
"arguments": '{"input": "10 + 10"}'
}
}
]
},
{
"role": "tool",
"tool_call_id": "call_123",
"name": "calculator",
"content": [{"type": "text", "text": "20"}] # Array format from LibreChat
}
]
# Call with is_async=True
transformed_messages = await config._transform_messages(
messages=messages_with_array_content,
model="deepinfra/Qwen/Qwen3-235B-A22B",
is_async=True
)
# Verify the tool message content was converted to string
tool_message = transformed_messages[2]
assert tool_message["role"] == "tool"
assert isinstance(tool_message["content"], str)
assert tool_message["content"] == "20"
print(f"✓ Async test passed: {tool_message['content']}")
print("\n✅ DeepInfra async tool message transformation test passed!")

View file

@ -0,0 +1,2 @@
# Manus provider tests

View file

@ -0,0 +1,2 @@
# Manus Responses API tests

View file

@ -0,0 +1,60 @@
"""
Tests for Manus Responses API transformation
Tests the ManusResponsesAPIConfig class that handles Manus-specific
transformations for the Responses API.
Source: litellm/llms/manus/responses/transformation.py
"""
import os
import sys
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.manus.responses.transformation import ManusResponsesAPIConfig
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
def test_extract_agent_profile():
"""Test that agent profile is correctly extracted from model name"""
config = ManusResponsesAPIConfig()
assert config._extract_agent_profile("manus/manus-1.6") == "manus-1.6"
assert config._extract_agent_profile("manus/manus-1.6-lite") == "manus-1.6-lite"
assert config._extract_agent_profile("manus/manus-1.6-max") == "manus-1.6-max"
def test_transform_responses_api_request_adds_manus_params():
"""Test that transform_responses_api_request adds task_mode and agent_profile"""
config = ManusResponsesAPIConfig()
input_param = [
{
"role": "user",
"content": [
{
"type": "input_text",
"text": "What's the color of the sky?",
}
],
}
]
optional_params = ResponsesAPIOptionalRequestParams()
litellm_params = GenericLiteLLMParams()
headers = {}
result = config.transform_responses_api_request(
model="manus/manus-1.6",
input=input_param,
response_api_optional_request_params=dict(optional_params),
litellm_params=litellm_params,
headers=headers,
)
assert result["task_mode"] == "agent"
assert result["agent_profile"] == "manus-1.6"
assert "input" in result
assert "model" in result

View file

@ -0,0 +1,150 @@
"""
Tests for Xiaomi MiMo provider configuration and integration.
Related to issue #18794
"""
import os
import sys
from unittest.mock import MagicMock, patch
try:
import pytest
except ImportError:
pytest = None
# Add workspace to path
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
sys.path.insert(0, workspace_path)
import litellm
class TestXiaomiMiMoProviderConfig:
"""Test Xiaomi MiMo provider configuration"""
def test_xiaomi_mimo_in_provider_list(self):
"""Test that xiaomi_mimo is in the provider list (fixes #18794)"""
from litellm import LlmProviders
# Verify xiaomi_mimo is in the enum
assert hasattr(LlmProviders, 'XIAOMI_MIMO')
assert LlmProviders.XIAOMI_MIMO.value == 'xiaomi_mimo'
# Verify it's in the provider list
assert 'xiaomi_mimo' in litellm.provider_list
def test_xiaomi_mimo_json_config_exists(self):
"""Test that xiaomi_mimo is configured in providers.json"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
# Verify xiaomi_mimo is loaded
assert JSONProviderRegistry.exists("xiaomi_mimo")
# Get xiaomi_mimo config
xiaomi_mimo = JSONProviderRegistry.get("xiaomi_mimo")
assert xiaomi_mimo is not None
assert xiaomi_mimo.base_url == "https://api.xiaomimimo.com/v1"
assert xiaomi_mimo.api_key_env == "XIAOMI_MIMO_API_KEY"
assert xiaomi_mimo.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_xiaomi_mimo_provider_resolution(self):
"""Test that provider resolution finds xiaomi_mimo"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="xiaomi_mimo/mimo-v2-flash",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "mimo-v2-flash"
assert provider == "xiaomi_mimo"
assert api_base == "https://api.xiaomimimo.com/v1"
def test_xiaomi_mimo_router_config(self):
"""Test that xiaomi_mimo can be used in Router configuration (fixes #18794)"""
from litellm import Router
# This should not raise "Unsupported provider - xiaomi_mimo"
router = Router(
model_list=[
{
"model_name": "mimo-v2-flash",
"litellm_params": {
"model": "xiaomi_mimo/mimo-v2-flash",
"api_key": "test-key",
},
}
]
)
# Verify the deployment was created successfully
assert len(router.model_list) == 1
assert router.model_list[0]["model_name"] == "mimo-v2-flash"
class TestXiaomiMiMoIntegration:
"""Integration tests for Xiaomi MiMo provider"""
def test_xiaomi_mimo_completion_basic(self):
"""Test basic completion call to Xiaomi MiMo"""
# Skip test if API key not set in environment
if not os.environ.get("XIAOMI_MIMO_API_KEY"):
if pytest:
pytest.skip("XIAOMI_MIMO_API_KEY not set")
return
try:
response = litellm.completion(
model="xiaomi_mimo/mimo-v2-flash",
messages=[{"role": "user", "content": "Say 'test successful' and nothing else"}],
max_tokens=10,
)
# Verify response structure
assert response is not None
assert hasattr(response, "choices")
assert len(response.choices) > 0
assert hasattr(response.choices[0], "message")
assert hasattr(response.choices[0].message, "content")
assert response.choices[0].message.content is not None
# Check that we got a response
content = response.choices[0].message.content.lower()
assert len(content) > 0
print(f"✓ Xiaomi MiMo completion successful: {response.choices[0].message.content}")
except Exception as e:
if pytest:
pytest.fail(f"Xiaomi MiMo completion failed: {str(e)}")
else:
raise
if __name__ == "__main__":
# Run basic tests
print("Testing Xiaomi MiMo Provider...")
test_config = TestXiaomiMiMoProviderConfig()
print("\n1. Testing provider in list...")
test_config.test_xiaomi_mimo_in_provider_list()
print(" ✓ xiaomi_mimo in provider list")
print("\n2. Testing JSON config...")
test_config.test_xiaomi_mimo_json_config_exists()
print(" ✓ xiaomi_mimo JSON config loaded")
print("\n3. Testing provider resolution...")
test_config.test_xiaomi_mimo_provider_resolution()
print(" ✓ Provider resolution works")
print("\n4. Testing router configuration...")
test_config.test_xiaomi_mimo_router_config()
print(" ✓ Router configuration works (issue #18794 fixed)")
print("\n" + "="*50)
print("✓ All configuration tests passed!")
print("="*50)

View file

@ -721,3 +721,189 @@ def test_convert_tool_response_text_only():
# Check inline_data does NOT exist (no image provided)
assert "inline_data" not in result
def test_file_data_field_order():
"""
Test that file_data fields are in the correct order (mime_type before file_uri).
The Gemini API is sensitive to field order in the file_data object.
This test verifies that mime_type comes before file_uri in both:
1. Dictionary key order
2. JSON serialization
Related issue: Gemini API returns 400 INVALID_ARGUMENT when fields are in wrong order.
"""
import json
from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_image
# Test with HTTPS URL and explicit format (audio file)
file_url = "https://generativelanguage.googleapis.com/v1beta/files/test123"
format = "audio/mpeg"
result = _process_gemini_image(image_url=file_url, format=format)
# Verify the result has file_data
assert "file_data" in result
file_data = result["file_data"]
# Verify both fields are present
assert "mime_type" in file_data
assert "file_uri" in file_data
assert file_data["mime_type"] == "audio/mpeg"
assert file_data["file_uri"] == file_url
# Verify field order by checking dictionary keys
# In Python 3.7+, dict maintains insertion order
file_data_keys = list(file_data.keys())
assert file_data_keys.index("mime_type") < file_data_keys.index("file_uri"), \
"mime_type must come before file_uri in the file_data dict"
# Also verify by serializing to JSON string
json_str = json.dumps(file_data)
mime_type_pos = json_str.find('"mime_type"')
file_uri_pos = json_str.find('"file_uri"')
assert mime_type_pos < file_uri_pos, \
"mime_type must appear before file_uri in JSON serialization"
def test_file_data_field_order_gcs_urls():
"""Test that GCS URLs also maintain correct field order."""
import json
from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_image
# Test with GCS URL
gcs_url = "gs://bucket/audio.mp3"
result = _process_gemini_image(image_url=gcs_url)
# Verify the result has file_data
assert "file_data" in result
file_data = result["file_data"]
# Verify both fields are present
assert "mime_type" in file_data
assert "file_uri" in file_data
# Verify field order
file_data_keys = list(file_data.keys())
assert file_data_keys.index("mime_type") < file_data_keys.index("file_uri"), \
"mime_type must come before file_uri in the file_data dict"
def test_extract_file_data_with_path_object():
"""
Test that filename is correctly extracted from Path objects for MIME type detection.
When uploading files using Path objects (e.g., Path("speech.mp3")), the filename
must be extracted to enable proper MIME type detection. Without this, files get
uploaded with 'application/octet-stream' instead of the correct MIME type.
Related issue: Files uploaded with wrong MIME type cause Gemini API to reject
requests where the specified format doesn't match the uploaded file's MIME type.
"""
from pathlib import Path
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
import tempfile
import os
# Create a temporary MP3 file
with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as tmp:
tmp.write(b"fake mp3 content")
tmp_path = tmp.name
try:
# Test with Path object
path_obj = Path(tmp_path)
extracted = extract_file_data(path_obj)
# Verify filename was extracted
assert extracted["filename"] is not None
assert extracted["filename"].endswith(".mp3")
# Verify MIME type was correctly detected
assert extracted["content_type"] == "audio/mpeg", \
f"Expected 'audio/mpeg' but got '{extracted['content_type']}'"
# Verify content was read
assert extracted["content"] == b"fake mp3 content"
finally:
# Clean up temporary file
os.unlink(tmp_path)
def test_extract_file_data_with_string_path():
"""Test that filename is correctly extracted from string paths."""
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
import tempfile
import os
# Create a temporary WAV file
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
tmp.write(b"fake wav content")
tmp_path = tmp.name
try:
# Test with string path
extracted = extract_file_data(tmp_path)
# Verify filename was extracted
assert extracted["filename"] is not None
assert extracted["filename"].endswith(".wav")
# Verify MIME type was correctly detected (can be audio/wav or audio/x-wav depending on system)
assert extracted["content_type"] in ["audio/wav", "audio/x-wav"], \
f"Expected 'audio/wav' or 'audio/x-wav' but got '{extracted['content_type']}'"
# Verify content was read
assert extracted["content"] == b"fake wav content"
finally:
# Clean up temporary file
os.unlink(tmp_path)
def test_extract_file_data_with_tuple_format():
"""Test that tuple format (with explicit content_type) still works correctly."""
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
# Test with tuple format: (filename, content, content_type)
filename = "test_audio.mp3"
content = b"test audio content"
content_type = "audio/mpeg"
extracted = extract_file_data((filename, content, content_type))
# Verify all fields are correct
assert extracted["filename"] == filename
assert extracted["content"] == content
assert extracted["content_type"] == content_type
def test_extract_file_data_fallback_to_octet_stream():
"""Test that unknown file types fall back to application/octet-stream."""
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
import tempfile
import os
# Create a temporary file with unknown extension
with tempfile.NamedTemporaryFile(suffix=".xyz123", delete=False) as tmp:
tmp.write(b"unknown content")
tmp_path = tmp.name
try:
# Test with unknown file type
extracted = extract_file_data(tmp_path)
# Verify filename was extracted
assert extracted["filename"] is not None
assert extracted["filename"].endswith(".xyz123")
# Verify MIME type falls back to octet-stream
assert extracted["content_type"] == "application/octet-stream", \
f"Expected 'application/octet-stream' for unknown type, got '{extracted['content_type']}'"
finally:
# Clean up temporary file
os.unlink(tmp_path)

View file

@ -1175,6 +1175,30 @@ def test_vertex_ai_moonshot_uses_openai_handler():
)
def test_vertex_ai_zai_uses_openai_handler():
"""
Ensure ZAI partner models re-use the OpenAI-format handler.
"""
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
VertexAIPartnerModels,
)
assert VertexAIPartnerModels.should_use_openai_handler(
"zai-org/glm-4.7-maas"
)
def test_vertex_ai_zai_is_partner_model():
"""
Ensure ZAI models are detected as Vertex AI partner models.
"""
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
VertexAIPartnerModels,
)
assert VertexAIPartnerModels.is_vertex_partner_model("zai-org/glm-4.7-maas")
def test_build_vertex_schema_empty_properties():
"""
Test _build_vertex_schema handles empty properties objects correctly.

View file

@ -1242,7 +1242,7 @@ class TestSpendLogsPayload:
"model": "claude-3-7-sonnet-20250219",
"user": "",
"team_id": "",
"metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}',
"metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}',
"cache_key": "Cache OFF",
"spend": 0.01383,
"total_tokens": 2598,
@ -1334,7 +1334,7 @@ class TestSpendLogsPayload:
"model": "claude-3-7-sonnet-20250219",
"user": "",
"team_id": "",
"metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}',
"metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}',
"cache_key": "Cache OFF",
"spend": 0.01383,
"total_tokens": 2598,

View file

@ -3038,6 +3038,94 @@ def test_get_image_root_case_uses_current_dir(monkeypatch):
assert mock_file_response.called, "FileResponse should be called"
def test_get_config_normalizes_string_callbacks(monkeypatch):
"""
Test that /get/config/callbacks normalizes string callbacks to lists.
"""
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
config_data = {
"litellm_settings": {
"success_callback": "langfuse",
"failure_callback": None,
"callbacks": ["prometheus", "datadog"],
},
"general_settings": {},
"environment_variables": {},
}
mock_router = MagicMock()
mock_router.get_settings.return_value = {}
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
monkeypatch.setattr(
proxy_config, "get_config", AsyncMock(return_value=config_data)
)
original_overrides = app.dependency_overrides.copy()
app.dependency_overrides[user_api_key_auth] = lambda: MagicMock()
client = TestClient(app)
try:
response = client.get("/get/config/callbacks")
finally:
app.dependency_overrides = original_overrides
assert response.status_code == 200
callbacks = response.json()["callbacks"]
success_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "success"]
failure_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "failure"]
success_and_failure_callbacks = [
cb["name"] for cb in callbacks if cb.get("type") == "success_and_failure"
]
assert "langfuse" in success_callbacks
assert len(failure_callbacks) == 0
assert "prometheus" in success_and_failure_callbacks
assert "datadog" in success_and_failure_callbacks
def test_deep_merge_dicts_skips_none_and_empty_lists(monkeypatch):
"""
Test that _update_config_fields deep merge skips None values and empty lists.
"""
from litellm.proxy.proxy_server import ProxyConfig
proxy_config = ProxyConfig()
current_config = {
"general_settings": {
"max_parallel_requests": 10,
"allowed_models": ["gpt-3.5-turbo", "gpt-4"],
"nested": {
"key1": "value1",
"key2": "value2",
},
}
}
db_param_value = {
"max_parallel_requests": None,
"allowed_models": [],
"new_key": "new_value",
"nested": {
"key1": "updated_value1",
"key3": "value3",
},
}
result = proxy_config._update_config_fields(
current_config, "general_settings", db_param_value
)
assert result["general_settings"]["max_parallel_requests"] == 10
assert result["general_settings"]["allowed_models"] == ["gpt-3.5-turbo", "gpt-4"]
assert result["general_settings"]["new_key"] == "new_value"
assert result["general_settings"]["nested"]["key1"] == "updated_value1"
assert result["general_settings"]["nested"]["key2"] == "value2"
assert result["general_settings"]["nested"]["key3"] == "value3"
@pytest.mark.asyncio
async def test_get_hierarchical_router_settings():
"""

View file

@ -34,7 +34,7 @@ class TestTextFormatConversion:
Test that when text_format parameter is passed to litellm.aresponses,
it gets converted to text parameter in the raw API call to OpenAI.
"""
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
class TestResponse(BaseModel):
"""Test Pydantic model for structured output"""
@ -42,20 +42,8 @@ class TestTextFormatConversion:
answer: str
confidence: float
class MockResponse:
"""Mock response class for testing"""
def __init__(self, json_data, status_code):
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = {}
def json(self):
return self._json_data
# Mock response from OpenAI
mock_response = {
mock_response_data = {
"id": "resp_123",
"object": "response",
"created_at": 1741476542,
@ -101,13 +89,74 @@ class TestTextFormatConversion:
base_completion_call_args = self.get_base_completion_call_args()
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
# Configure the mock to return our response
mock_post.return_value = MockResponse(mock_response, 200)
# Mock the response_api_handler function to capture the request
captured_request = {}
def mock_handler(
model,
input,
responses_api_provider_config,
response_api_optional_request_params,
custom_llm_provider,
litellm_params,
logging_obj,
extra_headers=None,
extra_body=None,
timeout=None,
client=None,
fake_stream=False,
litellm_metadata=None,
shared_session=None,
_is_async=False,
):
# Capture the request parameters
captured_request["model"] = model
captured_request["input"] = input
captured_request["params"] = response_api_optional_request_params
# Return a mock ResponsesAPIResponse wrapped in a coroutine if async
async def async_response():
return ResponsesAPIResponse(
id="resp_123",
object="response",
created_at=1741476542,
status="completed",
model="gpt-4o",
output=mock_response_data["output"],
usage=ResponseAPIUsage(
input_tokens=10,
output_tokens=20,
total_tokens=30,
),
text=mock_response_data.get("text"),
error=None,
incomplete_details=None,
)
if _is_async:
return async_response()
else:
return ResponsesAPIResponse(
id="resp_123",
object="response",
created_at=1741476542,
status="completed",
model="gpt-4o",
output=mock_response_data["output"],
usage=ResponseAPIUsage(
input_tokens=10,
output_tokens=20,
total_tokens=30,
),
text=mock_response_data.get("text"),
error=None,
incomplete_details=None,
)
with patch(
"litellm.responses.main.base_llm_http_handler.response_api_handler",
new=mock_handler,
):
litellm._turn_on_debug()
litellm.set_verbose = True
@ -118,21 +167,19 @@ class TestTextFormatConversion:
**base_completion_call_args,
)
# Verify the request was made correctly
mock_post.assert_called_once()
request_body = mock_post.call_args.kwargs["json"]
print("Request body:", json.dumps(request_body, indent=4))
# Verify the captured request
print("Captured request:", json.dumps(captured_request, indent=4, default=str))
# Validate that text_format was converted to text parameter
assert (
"text" in request_body
), "text parameter should be present in request body"
"text" in captured_request["params"]
), "text parameter should be present in request params"
assert (
"text_format" not in request_body
), "text_format should not be in request body"
"text_format" not in captured_request["params"]
), "text_format should not be in request params"
# Validate the text parameter structure
text_param = request_body["text"]
text_param = captured_request["params"]["text"]
assert "format" in text_param, "text parameter should have format field"
assert (
text_param["format"]["type"] == "json_schema"
@ -156,7 +203,7 @@ class TestTextFormatConversion:
), "schema should have confidence property"
# Validate other request parameters
assert request_body["input"] == "What is the capital of France?"
assert captured_request["input"] == "What is the capital of France?"
# Validate the response
print("Response:", json.dumps(response, indent=4, default=str))

View file

@ -313,17 +313,31 @@ async def test_error_from_tag_routing():
def test_tag_routing_with_list_of_tags():
"""
Test that the router can handle a list of tags
Test that the router can handle a list of tags with match_any behavior
"""
from litellm.router_strategy.tag_based_routing import is_valid_deployment_tag
assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA"])
assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamB"])
assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamC"])
assert is_valid_deployment_tag(["teamA"], ["teamA", "teamB"])
assert not is_valid_deployment_tag(["teamA", "teamB"], ["teamC"])
assert not is_valid_deployment_tag(["teamA", "teamB"], [])
assert not is_valid_deployment_tag(["default"], ["teamA"])
def test_tag_routing_with_list_of_tags_match_all():
"""
Test that the router can handle a list of tags with match_all behavior
"""
from litellm.router_strategy.tag_based_routing import is_valid_deployment_tag
assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA"], match_any=False)
assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamB"], match_any=False)
assert not is_valid_deployment_tag(["teamA", "teamB", "teamC"], ["teamA", "teamD"], match_any=False)
assert not is_valid_deployment_tag(["teamA"], ["teamA", "teamB"], match_any=False)
assert not is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamC"], match_any=False)
assert not is_valid_deployment_tag(["teamA", "teamB"], [], match_any=False)
assert not is_valid_deployment_tag(["default"], ["teamA"], match_any=False)
@pytest.mark.asyncio()
async def test_router_free_paid_tier_with_responses_api():

View file

@ -42,34 +42,45 @@ from litellm._lazy_imports import (
def _clear_names_from_globals(names: tuple):
"""Clear all names from litellm globals."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in names:
if name in litellm.__dict__:
del litellm.__dict__[name]
if name in litellm_globals:
del litellm_globals[name]
def _clear_names_from_utils_globals(names: tuple):
"""Clear all names from litellm.utils globals."""
# Get the actual globals dict, not a copy
utils_globals = sys.modules["litellm.utils"].__dict__
for name in names:
if name in litellm.utils.__dict__:
del litellm.utils.__dict__[name]
if name in utils_globals:
del utils_globals[name]
def _verify_only_requested_name_imported(name: str, all_names: tuple):
"""Verify that only the requested name is in globals, not the others."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for other_name in all_names:
if other_name != name:
assert other_name not in litellm.__dict__, f"{other_name} should not be imported when importing {name}"
assert other_name not in litellm_globals, f"{other_name} should not be imported when importing {name}"
def _verify_only_requested_name_imported_in_utils(name: str, all_names: tuple):
"""Verify that only the requested name is in utils globals, not the others."""
# Get the actual globals dict, not a copy
utils_globals = sys.modules["litellm.utils"].__dict__
for other_name in all_names:
if other_name != name:
assert other_name not in litellm.utils.__dict__, f"{other_name} should not be imported when importing {name}"
assert other_name not in utils_globals, f"{other_name} should not be imported when importing {name}"
def test_cost_calculator_lazy_imports():
"""Test that all cost calculator functions can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
# Test each name individually - only that name should be imported
for name in COST_CALCULATOR_NAMES:
# Clear all names before importing just one
@ -78,7 +89,7 @@ def test_cost_calculator_lazy_imports():
func = _lazy_import_cost_calculator(name)
assert func is not None
assert callable(func)
assert name in litellm.__dict__
assert name in litellm_globals
# Verify only the requested name is in globals, not the others
_verify_only_requested_name_imported(name, COST_CALCULATOR_NAMES)
@ -86,6 +97,9 @@ def test_cost_calculator_lazy_imports():
def test_litellm_logging_lazy_imports():
"""Test that all litellm_logging items can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
# Test each name individually - only that name should be imported
for name in LITELLM_LOGGING_NAMES:
# Clear all names before importing just one
@ -93,7 +107,7 @@ def test_litellm_logging_lazy_imports():
item = _lazy_import_litellm_logging(name)
assert item is not None
assert name in litellm.__dict__
assert name in litellm_globals
# Verify only the requested name is in globals, not the others
_verify_only_requested_name_imported(name, LITELLM_LOGGING_NAMES)
@ -101,6 +115,9 @@ def test_litellm_logging_lazy_imports():
def test_utils_lazy_imports():
"""Test that all utils functions can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
# Test each name individually - only that name should be imported
for name in UTILS_NAMES:
# Clear all names before importing just one
@ -108,7 +125,7 @@ def test_utils_lazy_imports():
attr = _lazy_import_utils(name)
assert attr is not None
assert name in litellm.__dict__
assert name in litellm_globals
# Verify only the requested name is in globals, not the others
_verify_only_requested_name_imported(name, UTILS_NAMES)
@ -116,6 +133,9 @@ def test_utils_lazy_imports():
def test_caching_lazy_imports():
"""Test that all caching classes can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
# Test each name individually - only that name should be imported
for name in CACHING_NAMES:
# Clear all names before importing just one
@ -123,7 +143,7 @@ def test_caching_lazy_imports():
cls = _lazy_import_caching(name)
assert cls is not None
assert name in litellm.__dict__
assert name in litellm_globals
# Verify only the requested name is in globals, not the others
_verify_only_requested_name_imported(name, CACHING_NAMES)
@ -131,71 +151,89 @@ def test_caching_lazy_imports():
def test_token_counter_lazy_imports():
"""Test that token counter utilities can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in TOKEN_COUNTER_NAMES:
_clear_names_from_globals(TOKEN_COUNTER_NAMES)
func = _lazy_import_token_counter(name)
assert func is not None
assert name in litellm.__dict__
assert name in litellm_globals
_verify_only_requested_name_imported(name, TOKEN_COUNTER_NAMES)
def test_bedrock_types_lazy_imports():
"""Test that Bedrock type aliases can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in BEDROCK_TYPES_NAMES:
_clear_names_from_globals(BEDROCK_TYPES_NAMES)
alias = _lazy_import_bedrock_types(name)
assert alias is not None
assert name in litellm.__dict__
assert name in litellm_globals
_verify_only_requested_name_imported(name, BEDROCK_TYPES_NAMES)
def test_types_utils_lazy_imports():
"""Test that common types.utils symbols can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in TYPES_UTILS_NAMES:
_clear_names_from_globals(TYPES_UTILS_NAMES)
obj = _lazy_import_types_utils(name)
assert obj is not None
assert name in litellm.__dict__
assert name in litellm_globals
_verify_only_requested_name_imported(name, TYPES_UTILS_NAMES)
def test_llm_client_cache_lazy_imports():
"""Test that LLM client cache class and singleton can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in LLM_CLIENT_CACHE_NAMES:
_clear_names_from_globals(LLM_CLIENT_CACHE_NAMES)
obj = _lazy_import_llm_client_cache(name)
assert obj is not None
assert name in litellm.__dict__
assert name in litellm_globals
_verify_only_requested_name_imported(name, LLM_CLIENT_CACHE_NAMES)
def test_http_handler_lazy_imports():
"""Test that HTTP handler singletons can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in HTTP_HANDLER_NAMES:
_clear_names_from_globals(HTTP_HANDLER_NAMES)
handler = _lazy_import_http_handlers(name)
assert handler is not None
assert name in litellm.__dict__
assert name in litellm_globals
_verify_only_requested_name_imported(name, HTTP_HANDLER_NAMES)
def test_dotprompt_lazy_imports():
"""Test that dotprompt globals can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in DOTPROMPT_NAMES:
_clear_names_from_globals(DOTPROMPT_NAMES)
obj = _lazy_import_dotprompt(name)
assert name in litellm.__dict__
assert name in litellm_globals
# Only the setter must be callable; others may be None by default
if name == "set_global_prompt_directory":
@ -245,12 +283,15 @@ def test_unknown_attribute_raises_error():
def test_llm_config_lazy_imports():
"""Test that LLM config classes can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in LLM_CONFIG_NAMES:
_clear_names_from_globals(LLM_CONFIG_NAMES)
obj = _lazy_import_llm_configs(name)
assert obj is not None
assert name in litellm.__dict__
assert name in litellm_globals
# Config classes should be classes/types
assert isinstance(obj, type), f"{name} should be a class"
@ -259,12 +300,15 @@ def test_llm_config_lazy_imports():
def test_types_lazy_imports():
"""Test that type classes can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in TYPES_NAMES:
_clear_names_from_globals(TYPES_NAMES)
obj = _lazy_import_types(name)
assert obj is not None
assert name in litellm.__dict__
assert name in litellm_globals
# Type classes should be classes/types
assert isinstance(obj, type), f"{name} should be a class"
@ -273,25 +317,31 @@ def test_types_lazy_imports():
def test_llm_provider_logic_lazy_imports():
"""Test that LLM provider logic functions can be lazy imported."""
# Get the actual globals dict, not a copy
litellm_globals = sys.modules["litellm"].__dict__
for name in LLM_PROVIDER_LOGIC_NAMES:
_clear_names_from_globals(LLM_PROVIDER_LOGIC_NAMES)
func = _lazy_import_llm_provider_logic(name)
assert func is not None
assert callable(func)
assert name in litellm.__dict__
assert name in litellm_globals
_verify_only_requested_name_imported(name, LLM_PROVIDER_LOGIC_NAMES)
def test_utils_module_lazy_imports():
"""Test that utils module attributes can be lazy imported."""
# Get the actual globals dict, not a copy
utils_globals = sys.modules["litellm.utils"].__dict__
for name in UTILS_MODULE_NAMES:
_clear_names_from_utils_globals(UTILS_MODULE_NAMES)
obj = _lazy_import_utils_module(name)
assert obj is not None
assert name in litellm.utils.__dict__
assert name in utils_globals
_verify_only_requested_name_imported_in_utils(name, UTILS_MODULE_NAMES)

View file

@ -42,8 +42,11 @@ class TestIsEncryptedResponseId:
def test_is_encrypted_response_id_valid(self, responses_id_security):
"""Test that a properly encrypted response ID is identified correctly"""
with patch(
"litellm.proxy.hooks.responses_id_security.decrypt_value_helper"
# Patch at the module level where it's imported
import litellm.proxy.hooks.responses_id_security as responses_module
with patch.object(
responses_module, "decrypt_value_helper"
) as mock_decrypt:
mock_decrypt.return_value = f"{SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value}response_id:resp_123;user_id:user-456"
@ -56,8 +59,11 @@ class TestIsEncryptedResponseId:
def test_is_encrypted_response_id_invalid(self, responses_id_security):
"""Test that an unencrypted response ID returns False"""
with patch(
"litellm.proxy.hooks.responses_id_security.decrypt_value_helper"
# Patch at the module level where it's imported
import litellm.proxy.hooks.responses_id_security as responses_module
with patch.object(
responses_module, "decrypt_value_helper"
) as mock_decrypt:
mock_decrypt.return_value = None
@ -71,8 +77,11 @@ class TestDecryptResponseId:
def test_decrypt_response_id_valid(self, responses_id_security):
"""Test decrypting a valid encrypted response ID"""
with patch(
"litellm.proxy.hooks.responses_id_security.decrypt_value_helper"
# Patch at the module level where it's imported
import litellm.proxy.hooks.responses_id_security as responses_module
with patch.object(
responses_module, "decrypt_value_helper"
) as mock_decrypt:
mock_decrypt.return_value = f"{SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value}response_id:resp_original_123;user_id:user-456;team_id:team-789"
@ -86,8 +95,11 @@ class TestDecryptResponseId:
def test_decrypt_response_id_no_encryption(self, responses_id_security):
"""Test decrypting a non-encrypted response ID"""
with patch(
"litellm.proxy.hooks.responses_id_security.decrypt_value_helper"
# Patch at the module level where it's imported
import litellm.proxy.hooks.responses_id_security as responses_module
with patch.object(
responses_module, "decrypt_value_helper"
) as mock_decrypt:
mock_decrypt.return_value = None

View file

@ -2,7 +2,7 @@ import asyncio
import json
import os
import sys
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -18,6 +18,7 @@ from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.gemini.videos.transformation import GeminiVideoConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.types.videos.main import VideoObject, VideoResponse
from litellm.videos import main as videos_main
from litellm.videos.main import (
avideo_generation,
avideo_status,
@ -31,32 +32,29 @@ class TestVideoGeneration:
def test_video_generation_basic(self):
"""Test basic video generation functionality."""
# Mock the video generation response
mock_response = VideoObject(
id="video_123",
object="video",
status="queued",
created_at=1712697600,
# Use mock_response parameter for reliable testing
response = video_generation(
prompt="Show them running around the room",
model="sora-2",
seconds="8",
size="720x1280",
seconds="8"
mock_response={
"id": "video_123",
"object": "video",
"status": "queued",
"created_at": 1712697600,
"model": "sora-2",
"size": "720x1280",
"seconds": "8"
}
)
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
mock_handler.video_generation_handler.return_value = mock_response
response = video_generation(
prompt="Show them running around the room",
model="sora-2",
seconds="8",
size="720x1280"
)
assert isinstance(response, VideoObject)
assert response.id == "video_123"
assert response.model == "sora-2"
assert response.size == "720x1280"
assert response.seconds == "8"
assert isinstance(response, VideoObject)
assert response.id == "video_123"
assert response.status == "queued"
assert response.model == "sora-2"
assert response.size == "720x1280"
assert response.seconds == "8"
def test_video_generation_with_mock_response(self):
"""Test video generation with mock response."""
@ -97,26 +95,27 @@ class TestVideoGeneration:
progress=50
)
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
mock_handler.video_generation_handler.return_value = mock_response
import asyncio
async def test_async():
response = await avideo_generation(
prompt="A cat playing with a ball",
model="sora-2",
seconds="5",
size="720x1280"
)
return response
response = asyncio.run(test_async())
assert isinstance(response, VideoObject)
assert response.id == "video_async_123"
assert response.status == "processing"
assert response.progress == 50
# Mock the async_video_generation_handler to return the mock_response
async_mock = AsyncMock(return_value=mock_response)
with patch.object(videos_main.base_llm_http_handler, 'async_video_generation_handler', async_mock):
with patch.object(videos_main.base_llm_http_handler, 'video_generation_handler', side_effect=lambda **kwargs: async_mock(**kwargs)):
import asyncio
async def test_async():
response = await avideo_generation(
prompt="A cat playing with a ball",
model="sora-2",
seconds="5",
size="720x1280"
)
return response
response = asyncio.run(test_async())
assert isinstance(response, VideoObject)
assert response.id == "video_async_123"
assert response.status == "processing"
assert response.progress == 50
def test_video_generation_parameter_validation(self):
"""Test video generation parameter validation."""
@ -132,9 +131,7 @@ class TestVideoGeneration:
def test_video_generation_error_handling(self):
"""Test video generation error handling."""
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
mock_handler.video_generation_handler.side_effect = Exception("API Error")
with patch.object(videos_main.base_llm_http_handler, 'video_generation_handler', side_effect=Exception("API Error")):
with pytest.raises(Exception):
video_generation(
prompt="Test video",
@ -443,32 +440,28 @@ class TestVideoGeneration:
def test_video_status_basic(self):
"""Test basic video status functionality."""
# Mock the video status response
mock_response = VideoObject(
id="video_123",
object="video",
status="completed",
created_at=1712697600,
completed_at=1712697660,
# Use mock_response parameter for reliable testing
response = video_status(
video_id="video_123",
model="sora-2",
progress=100,
size="720x1280",
seconds="8"
mock_response={
"id": "video_123",
"object": "video",
"status": "completed",
"created_at": 1712697600,
"completed_at": 1712697660,
"model": "sora-2",
"progress": 100,
"size": "720x1280",
"seconds": "8"
}
)
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
mock_handler.video_status_handler.return_value = mock_response
response = video_status(
video_id="video_123",
model="sora-2"
)
assert isinstance(response, VideoObject)
assert response.id == "video_123"
assert response.status == "completed"
assert response.progress == 100
assert response.model == "sora-2"
assert isinstance(response, VideoObject)
assert response.id == "video_123"
assert response.status == "completed"
assert response.progress == 100
assert response.model == "sora-2"
def test_video_status_with_mock_response(self):
"""Test video status with mock response."""
@ -506,24 +499,25 @@ class TestVideoGeneration:
progress=0
)
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
mock_handler.video_status_handler.return_value = mock_response
import asyncio
async def test_async():
response = await avideo_status(
video_id="video_async_123",
model="sora-2"
)
return response
response = asyncio.run(test_async())
assert isinstance(response, VideoObject)
assert response.id == "video_async_123"
assert response.status == "queued"
assert response.progress == 0
# Mock the async_video_status_handler to return the mock_response
async_mock = AsyncMock(return_value=mock_response)
with patch.object(videos_main.base_llm_http_handler, 'async_video_status_handler', async_mock):
with patch.object(videos_main.base_llm_http_handler, 'video_status_handler', side_effect=lambda **kwargs: async_mock(**kwargs)):
import asyncio
async def test_async():
response = await avideo_status(
video_id="video_async_123",
model="sora-2"
)
return response
response = asyncio.run(test_async())
assert isinstance(response, VideoObject)
assert response.id == "video_async_123"
assert response.status == "queued"
assert response.progress == 0
def test_video_status_parameter_validation(self):
"""Test video status parameter validation."""
@ -539,9 +533,7 @@ class TestVideoGeneration:
def test_video_status_error_handling(self):
"""Test video status error handling."""
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
mock_handler.video_status_handler.side_effect = Exception("API Error")
with patch.object(videos_main.base_llm_http_handler, 'video_status_handler', side_effect=Exception("API Error")):
with pytest.raises(Exception):
video_status(
video_id="test_video_id",
@ -672,33 +664,30 @@ class TestVideoGeneration:
def test_video_status_async_inside_async_function(self):
"""Test that sync video_status works inside async functions (no asyncio.run issues)."""
mock_response = VideoObject(
id="video_sync_in_async",
object="video",
status="completed",
created_at=1712697600,
model="sora-2",
progress=100
)
import asyncio
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
mock_handler.video_status_handler.return_value = mock_response
import asyncio
async def test_sync_in_async():
# This should work without asyncio.run() issues
response = video_status(
video_id="video_sync_in_async",
model="sora-2"
)
return response
response = asyncio.run(test_sync_in_async())
assert isinstance(response, VideoObject)
assert response.id == "video_sync_in_async"
assert response.status == "completed"
async def test_sync_in_async():
# This should work without asyncio.run() issues
# Use mock_response parameter for reliable testing
response = video_status(
video_id="video_sync_in_async",
model="sora-2",
mock_response={
"id": "video_sync_in_async",
"object": "video",
"status": "completed",
"created_at": 1712697600,
"model": "sora-2",
"progress": 100
}
)
return response
response = asyncio.run(test_sync_in_async())
assert isinstance(response, VideoObject)
assert response.id == "video_sync_in_async"
assert response.status == "completed"
def test_video_status_url_construction(self):
"""Test video status URL construction."""

View file

@ -130,6 +130,9 @@ const ChatUI: React.FC<ChatUIProps> = ({
return disabledPersonalKeyCreation ? "custom" : "session";
});
const [apiKey, setApiKey] = useState<string>(() => sessionStorage.getItem("apiKey") || "");
const [customProxyBaseUrl, setCustomProxyBaseUrl] = useState<string>(
() => sessionStorage.getItem("customProxyBaseUrl") || ""
);
const [inputMessage, setInputMessage] = useState("");
const [chatHistory, setChatHistory] = useState<MessageType[]>(() => {
try {
@ -392,7 +395,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
const loadAgents = async () => {
try {
const agents = await fetchAvailableAgents(userApiKey);
const agents = await fetchAvailableAgents(userApiKey, customProxyBaseUrl || undefined);
setAgentInfo(agents);
// Clear selection if current agent not in list
if (selectedAgent && !agents.some((a) => a.agent_name === selectedAgent)) {
@ -404,7 +407,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
};
loadAgents();
}, [accessToken, apiKeySource, apiKey, endpointType]);
}, [accessToken, apiKeySource, apiKey, endpointType, customProxyBaseUrl, selectedAgent]);
useEffect(() => {
// Scroll to the bottom of the chat whenever chatHistory updates
@ -900,6 +903,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
useAdvancedParams ? temperature : undefined,
useAdvancedParams ? maxTokens : undefined,
updateTotalLatency,
customProxyBaseUrl || undefined,
mcpServers,
mcpServerToolRestrictions,
);
@ -912,6 +916,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
effectiveApiKey,
selectedTags,
signal,
customProxyBaseUrl || undefined,
);
} else if (endpointType === EndpointType.SPEECH) {
// For audio speech
@ -923,6 +928,9 @@ const ChatUI: React.FC<ChatUIProps> = ({
effectiveApiKey,
selectedTags,
signal,
undefined, // responseFormat
undefined, // speed
customProxyBaseUrl || undefined,
);
} else if (endpointType === EndpointType.IMAGE_EDITS) {
// For image edits
@ -935,6 +943,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
effectiveApiKey,
selectedTags,
signal,
customProxyBaseUrl || undefined,
);
}
} else if (endpointType === EndpointType.RESPONSES) {
@ -973,6 +982,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
handleMCPEvent, // Pass MCP event handler
codeInterpreter.enabled, // Enable Code Interpreter tool
codeInterpreter.setResult, // Handle code interpreter output
customProxyBaseUrl || undefined,
mcpServers,
mcpServerToolRestrictions,
);
@ -997,6 +1007,8 @@ const ChatUI: React.FC<ChatUIProps> = ({
traceId,
selectedVectorStores.length > 0 ? selectedVectorStores : undefined,
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
selectedMCPTools, // Pass the selected tools array
customProxyBaseUrl || undefined,
);
} else if (endpointType === EndpointType.EMBEDDINGS) {
await makeOpenAIEmbeddingsRequest(
@ -1005,6 +1017,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
selectedModel,
effectiveApiKey,
selectedTags,
customProxyBaseUrl || undefined,
);
} else if (endpointType === EndpointType.TRANSCRIPTION) {
// For audio transcriptions
@ -1016,6 +1029,11 @@ const ChatUI: React.FC<ChatUIProps> = ({
effectiveApiKey,
selectedTags,
signal,
undefined, // language
undefined, // prompt
undefined, // responseFormat
undefined, // temperature
customProxyBaseUrl || undefined,
);
}
}
@ -1032,6 +1050,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
updateTimingData,
updateTotalLatency,
updateA2AMetadata,
customProxyBaseUrl || undefined,
);
}
} catch (error) {
@ -1156,6 +1175,42 @@ const ChatUI: React.FC<ChatUIProps> = ({
)}
</div>
<div>
<div className="flex items-center justify-between mb-2">
<Text className="font-medium text-gray-700 flex items-center">
<SettingOutlined className="mr-2" /> Custom Proxy Base URL
</Text>
{customProxyBaseUrl && (
<Button
type="link"
size="small"
icon={<ClearOutlined />}
onClick={() => {
setCustomProxyBaseUrl("");
sessionStorage.removeItem("customProxyBaseUrl");
}}
className="text-gray-500 hover:text-gray-700"
>
Clear
</Button>
)}
</div>
<TextInput
placeholder="Optional: Enter custom proxy URL (e.g., http://localhost:5000)"
onValueChange={(value) => {
setCustomProxyBaseUrl(value);
sessionStorage.setItem("customProxyBaseUrl", value);
}}
value={customProxyBaseUrl}
icon={ApiOutlined}
/>
{customProxyBaseUrl && (
<Text className="text-xs text-gray-500 mt-1">
API calls will be sent to: {customProxyBaseUrl}
</Text>
)}
</div>
<div>
<Text className="font-medium block mb-2 text-gray-700 flex items-center">
<ApiOutlined className="mr-2" /> Endpoint Type

View file

@ -106,6 +106,9 @@ export default function CompareUI({ accessToken, disabledPersonalKeyCreation }:
);
const [customApiKey, setCustomApiKey] = useState("");
const [debouncedCustomApiKey, setDebouncedCustomApiKey] = useState("");
const [customProxyBaseUrl] = useState<string>(
() => sessionStorage.getItem("customProxyBaseUrl") || ""
);
useEffect(() => {
const timer = setTimeout(() => {
setDebouncedCustomApiKey(customApiKey);
@ -171,7 +174,7 @@ export default function CompareUI({ accessToken, disabledPersonalKeyCreation }:
}
setIsLoadingAgents(true);
try {
const agents = await fetchAvailableAgents(effectiveApiKey);
const agents = await fetchAvailableAgents(effectiveApiKey, customProxyBaseUrl || undefined);
if (!active) return;
setAgentOptions(agents);
} catch (error) {
@ -598,6 +601,8 @@ export default function CompareUI({ accessToken, disabledPersonalKeyCreation }:
undefined,
(time) => updateTimingDataForComparison(prepared.id, time),
(latency) => updateTotalLatencyForComparison(prepared.id, latency),
undefined, // onA2AMetadata
customProxyBaseUrl || undefined,
)
: makeOpenAIChatCompletionRequest(
prepared.apiChatHistory,
@ -618,6 +623,7 @@ export default function CompareUI({ accessToken, disabledPersonalKeyCreation }:
useAdvancedParams ? prepared.temperature : undefined,
useAdvancedParams ? prepared.maxTokens : undefined,
(latency) => updateTotalLatencyForComparison(prepared.id, latency),
customProxyBaseUrl || undefined,
);
requestPromise

View file

@ -113,8 +113,9 @@ export const makeA2ASendMessageRequest = async (
onTimingData?: (timeToFirstToken: number) => void,
onTotalLatency?: (totalLatency: number) => void,
onA2AMetadata?: (metadata: A2ATaskMetadata) => void,
customBaseUrl?: string,
): Promise<void> => {
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
const url = proxyBaseUrl
? `${proxyBaseUrl}/a2a/${agentId}/message/send`
: `/a2a/${agentId}/message/send`;
@ -242,8 +243,9 @@ export const makeA2AStreamMessageRequest = async (
onTimingData?: (timeToFirstToken: number) => void,
onTotalLatency?: (totalLatency: number) => void,
onA2AMetadata?: (metadata: A2ATaskMetadata) => void,
customBaseUrl?: string,
): Promise<void> => {
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
const url = proxyBaseUrl
? `${proxyBaseUrl}/a2a/${agentId}`
: `/a2a/${agentId}`;

View file

@ -17,6 +17,8 @@ export async function makeAnthropicMessagesRequest(
traceId?: string,
vector_store_ids?: string[],
guardrails?: string[],
selectedMCPTools?: string[],
customBaseUrl?: string,
) {
if (!accessToken) {
throw new Error("Virtual Key is required");
@ -27,7 +29,7 @@ export async function makeAnthropicMessagesRequest(
console.log = function () {};
}
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
// Prepare headers with tags and trace ID
const headers: Record<string, string> = {};

View file

@ -13,6 +13,7 @@ export async function makeOpenAIAudioSpeechRequest(
signal?: AbortSignal,
responseFormat?: string,
speed?: number,
customBaseUrl?: string,
) {
// base url should be the current base_url
const isLocal = process.env.NODE_ENV === "development";
@ -20,7 +21,7 @@ export async function makeOpenAIAudioSpeechRequest(
console.log = function () {};
}
console.log("isLocal:", isLocal);
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
const client = new openai.OpenAI({
apiKey: accessToken,
baseURL: proxyBaseUrl,

View file

@ -13,6 +13,7 @@ export async function makeOpenAIAudioTranscriptionRequest(
prompt?: string,
responseFormat?: string,
temperature?: number,
customBaseUrl?: string,
) {
// base url should be the current base_url
const isLocal = process.env.NODE_ENV === "development";
@ -20,7 +21,7 @@ export async function makeOpenAIAudioTranscriptionRequest(
console.log = function () {};
}
console.log("isLocal:", isLocal);
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
const client = new openai.OpenAI({
apiKey: accessToken,

View file

@ -24,6 +24,7 @@ export async function makeOpenAIChatCompletionRequest(
temperature?: number,
max_tokens?: number,
onTotalLatency?: (latency: number) => void,
customBaseUrl?: string,
mcpServers?: MCPServer[],
mcpServerToolRestrictions?: Record<string, string[]>,
) {
@ -33,7 +34,7 @@ export async function makeOpenAIChatCompletionRequest(
console.log = function () {};
}
console.log("isLocal:", isLocal);
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
// Prepare headers with tags and trace ID
const headers: Record<string, string> = {};
if (tags && tags.length > 0) {

View file

@ -7,6 +7,7 @@ export async function makeOpenAIEmbeddingsRequest(
selectedModel: string,
accessToken: string,
tags?: string[],
customBaseUrl?: string,
) {
if (!accessToken) {
throw new Error("Virtual Key is required");
@ -18,7 +19,7 @@ export async function makeOpenAIEmbeddingsRequest(
console.log = function () {};
}
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
// Prepare headers with tags and trace ID
const headers: Record<string, string> = {};
if (tags && tags.length > 0) {

View file

@ -16,9 +16,12 @@ export interface Agent {
/**
* Fetches available A2A agents from /v1/agents endpoint.
*/
export const fetchAvailableAgents = async (accessToken: string): Promise<Agent[]> => {
export const fetchAvailableAgents = async (
accessToken: string,
customBaseUrl?: string,
): Promise<Agent[]> => {
try {
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/agents` : `/v1/agents`;
const response = await fetch(url, {

View file

@ -10,6 +10,7 @@ export async function makeOpenAIImageEditsRequest(
accessToken: string,
tags?: string[],
signal?: AbortSignal,
customBaseUrl?: string,
) {
// base url should be the current base_url
const isLocal = process.env.NODE_ENV === "development";
@ -17,7 +18,7 @@ export async function makeOpenAIImageEditsRequest(
console.log = function () {};
}
console.log("isLocal:", isLocal);
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
const client = new openai.OpenAI({
apiKey: accessToken,

View file

@ -9,6 +9,7 @@ export async function makeOpenAIImageGenerationRequest(
accessToken: string,
tags?: string[],
signal?: AbortSignal,
customBaseUrl?: string,
) {
// base url should be the current base_url
const isLocal = process.env.NODE_ENV === "development";
@ -16,7 +17,7 @@ export async function makeOpenAIImageGenerationRequest(
console.log = function () {};
}
console.log("isLocal:", isLocal);
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
const client = new openai.OpenAI({
apiKey: accessToken,
baseURL: proxyBaseUrl,

View file

@ -33,6 +33,7 @@ export async function makeOpenAIResponsesRequest(
onMCPEvent?: (event: MCPEvent) => void,
codeInterpreterEnabled?: boolean,
onCodeInterpreterResult?: (result: CodeInterpreterResult) => void,
customBaseUrl?: string,
mcpServers?: MCPServer[],
mcpServerToolRestrictions?: Record<string, string[]>,
) {
@ -50,7 +51,7 @@ export async function makeOpenAIResponsesRequest(
console.log = function () {};
}
const proxyBaseUrl = getProxyBaseUrl();
const proxyBaseUrl = customBaseUrl || getProxyBaseUrl();
// Prepare headers with tags and trace ID
const headers: Record<string, string> = {};
if (tags && tags.length > 0) {