mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin' into litellm_user_promethus_metrics
This commit is contained in:
commit
02eec19a1f
105 changed files with 4592 additions and 446 deletions
174
.github/workflows/label-component.yml
vendored
174
.github/workflows/label-component.yml
vendored
|
|
@ -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]
|
||||
});
|
||||
|
|
|
|||
11
Dockerfile
11
Dockerfile
|
|
@ -20,7 +20,8 @@ RUN python -m pip install build
|
|||
COPY . .
|
||||
|
||||
# Build Admin UI
|
||||
RUN chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
# Convert Windows line endings to Unix and make executable
|
||||
RUN sed -i 's/\r$//' docker/build_admin_ui.sh && chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
|
||||
# Build the package
|
||||
RUN rm -rf dist/* && python -m build
|
||||
|
|
@ -65,12 +66,14 @@ RUN find /usr/lib -type f -path "*/tornado/test/*" -delete && \
|
|||
find /usr/lib -type d -path "*/tornado/test" -delete
|
||||
|
||||
# Install semantic_router and aurelio-sdk using script
|
||||
RUN chmod +x docker/install_auto_router.sh && ./docker/install_auto_router.sh
|
||||
# Convert Windows line endings to Unix and make executable
|
||||
RUN sed -i 's/\r$//' docker/install_auto_router.sh && chmod +x docker/install_auto_router.sh && ./docker/install_auto_router.sh
|
||||
|
||||
# Generate prisma client
|
||||
RUN prisma generate
|
||||
RUN chmod +x docker/entrypoint.sh
|
||||
RUN chmod +x docker/prod_entrypoint.sh
|
||||
# Convert Windows line endings to Unix for entrypoint scripts
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
|
||||
RUN sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,8 @@ WORKDIR /app
|
|||
COPY config.yaml .
|
||||
|
||||
# Make sure your docker/entrypoint.sh is executable
|
||||
RUN chmod +x docker/entrypoint.sh
|
||||
# Convert Windows line endings to Unix
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
|
||||
|
||||
# Expose the necessary port
|
||||
EXPOSE 4000/tcp
|
||||
|
|
|
|||
|
|
@ -46,8 +46,9 @@ COPY --from=builder /wheels/ /wheels/
|
|||
# Install the built wheel using pip; again using a wildcard if it's the only file
|
||||
RUN pip install *.whl /wheels/* --no-index --find-links=/wheels/ && rm -f *.whl && rm -rf /wheels
|
||||
|
||||
RUN chmod +x docker/entrypoint.sh
|
||||
RUN chmod +x docker/prod_entrypoint.sh
|
||||
# Convert Windows line endings to Unix for entrypoint scripts
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
|
||||
RUN sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
|
|
|
|||
|
|
@ -32,8 +32,9 @@ RUN rm -rf /app/litellm/proxy/_experimental/out/* && \
|
|||
WORKDIR /app
|
||||
|
||||
# Make sure your docker/entrypoint.sh is executable
|
||||
RUN chmod +x docker/entrypoint.sh
|
||||
RUN chmod +x docker/prod_entrypoint.sh
|
||||
# Convert Windows line endings to Unix for entrypoint scripts
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
|
||||
RUN sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
|
||||
|
||||
# Expose the necessary port
|
||||
EXPOSE 4000/tcp
|
||||
|
|
|
|||
|
|
@ -27,7 +27,8 @@ RUN python -m pip install build
|
|||
COPY . .
|
||||
|
||||
# Build Admin UI
|
||||
RUN chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
# Convert Windows line endings to Unix and make executable
|
||||
RUN sed -i 's/\r$//' docker/build_admin_ui.sh && chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
|
||||
# Build the package
|
||||
RUN rm -rf dist/* && python -m build
|
||||
|
|
@ -63,20 +64,23 @@ COPY --from=builder /wheels/ /wheels/
|
|||
RUN pip install *.whl /wheels/* --no-index --find-links=/wheels/ && rm -f *.whl && rm -rf /wheels
|
||||
|
||||
# Install semantic_router and aurelio-sdk using script
|
||||
RUN chmod +x docker/install_auto_router.sh && ./docker/install_auto_router.sh
|
||||
# Convert Windows line endings to Unix and make executable
|
||||
RUN sed -i 's/\r$//' docker/install_auto_router.sh && chmod +x docker/install_auto_router.sh && ./docker/install_auto_router.sh
|
||||
|
||||
# ensure pyjwt is used, not jwt
|
||||
RUN pip uninstall jwt -y
|
||||
RUN pip uninstall PyJWT -y
|
||||
RUN pip install PyJWT==2.9.0 --no-cache-dir
|
||||
|
||||
# Build Admin UI
|
||||
RUN chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
# Build Admin UI (runtime stage)
|
||||
# Convert Windows line endings to Unix and make executable
|
||||
RUN sed -i 's/\r$//' docker/build_admin_ui.sh && chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
|
||||
# Generate prisma client
|
||||
RUN prisma generate
|
||||
RUN chmod +x docker/entrypoint.sh
|
||||
RUN chmod +x docker/prod_entrypoint.sh
|
||||
# Convert Windows line endings to Unix for entrypoint scripts
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
|
||||
RUN sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
RUN apk add --no-cache supervisor
|
||||
|
|
|
|||
|
|
@ -40,7 +40,8 @@ COPY enterprise/ ./enterprise/
|
|||
COPY docker/ ./docker/
|
||||
|
||||
# Build Admin UI once
|
||||
RUN chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
# Convert Windows line endings to Unix and make executable
|
||||
RUN sed -i 's/\r$//' docker/build_admin_ui.sh && chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
|
||||
|
||||
# Build the package
|
||||
RUN rm -rf dist/* && python -m build
|
||||
|
|
@ -79,8 +80,12 @@ RUN pip install --no-cache-dir *.whl /wheels/* --no-index --find-links=/wheels/
|
|||
rm -rf /wheels
|
||||
|
||||
# Generate prisma client and set permissions
|
||||
# Convert Windows line endings to Unix for entrypoint scripts
|
||||
RUN prisma generate && \
|
||||
chmod +x docker/entrypoint.sh docker/prod_entrypoint.sh
|
||||
sed -i 's/\r$//' docker/entrypoint.sh && \
|
||||
sed -i 's/\r$//' docker/prod_entrypoint.sh && \
|
||||
chmod +x docker/entrypoint.sh && \
|
||||
chmod +x docker/prod_entrypoint.sh
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
|
|
|
|||
|
|
@ -144,7 +144,10 @@ RUN pip install --no-index --find-links=/wheels/ -r requirements.txt && \
|
|||
fi
|
||||
|
||||
# Permissions, cleanup, and Prisma prep
|
||||
RUN chmod +x docker/entrypoint.sh docker/prod_entrypoint.sh && \
|
||||
# Convert Windows line endings to Unix for entrypoint scripts
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && \
|
||||
sed -i 's/\r$//' docker/prod_entrypoint.sh && \
|
||||
chmod +x docker/entrypoint.sh docker/prod_entrypoint.sh && \
|
||||
mkdir -p /nonexistent /.npm /var/lib/litellm/assets /var/lib/litellm/ui && \
|
||||
chown -R nobody:nogroup /app /var/lib/litellm/ui /var/lib/litellm/assets /nonexistent /.npm && \
|
||||
pip uninstall jwt -y || true && \
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) |
|
||||
194
docs/my-website/docs/providers/manus.md
Normal file
194
docs/my-website/docs/providers/manus.md
Normal 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)
|
||||
|
|
@ -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) |
|
||||
|
|
|
|||
|
|
@ -576,10 +576,31 @@ custom_tokenizer:
|
|||
|
||||
```yaml
|
||||
general_settings:
|
||||
database_connection_pool_limit: 10 # sets connection pool for prisma client to postgres db (default: 10, recommended: 10-20)
|
||||
database_connection_pool_limit: 10 # sets connection pool per worker for prisma client to postgres db (default: 10, recommended: 10-20)
|
||||
database_connection_timeout: 60 # sets a 60s timeout for any connection call to the db
|
||||
```
|
||||
|
||||
**How to calculate the right value:**
|
||||
|
||||
The connection limit is applied **per worker process**, not per instance. This means if you have multiple workers, each worker will create its own connection pool.
|
||||
|
||||
**Formula:**
|
||||
```
|
||||
database_connection_pool_limit = MAX_DB_CONNECTIONS ÷ (number_of_instances × number_of_workers_per_instance)
|
||||
```
|
||||
|
||||
**Example:**
|
||||
- Your database allows a maximum of **100 connections**
|
||||
- You're running **1 instance** of LiteLLM
|
||||
- Each instance has **8 workers** (set via `--num_workers 8`)
|
||||
|
||||
Calculation: `100 ÷ (1 × 8) = 12.5`
|
||||
|
||||
Since you shouldn't use 12.5, round down to **10** to leave a safety buffer. This means:
|
||||
- Each of the 8 workers will have a connection pool limit of 10
|
||||
- Total maximum connections: 8 workers × 10 connections = 80 connections
|
||||
- This stays safely under your database's 100 connection limit
|
||||
|
||||
## Extras
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -19,7 +19,11 @@ general_settings:
|
|||
master_key: sk-1234 # enter your own master key, ensure it starts with 'sk-'
|
||||
alerting: ["slack"] # Setup slack alerting - get alerts on LLM exceptions, Budget Alerts, Slow LLM Responses
|
||||
proxy_batch_write_at: 60 # Batch write spend updates every 60s
|
||||
database_connection_pool_limit: 10 # limit the number of database connections to = MAX Number of DB Connections/Number of instances of litellm proxy (Around 10-20 is good number)
|
||||
database_connection_pool_limit: 10 # connection pool limit per worker process. Total connections = limit × workers × instances. Calculate: MAX_DB_CONNECTIONS / (instances × workers). Default: 10.
|
||||
|
||||
:::warning
|
||||
**Multiple instances:** If running multiple LiteLLM instances (e.g., Kubernetes pods), remember each instance multiplies your total connections. Example: 3 instances × 4 workers × 10 connections = 120 total connections.
|
||||
:::
|
||||
|
||||
# OPTIONAL Best Practices
|
||||
disable_error_logs: True # turn off writing LLM Exceptions to DB
|
||||
|
|
|
|||
|
|
@ -710,6 +710,7 @@ const sidebars = {
|
|||
"providers/llamafile",
|
||||
"providers/llamagate",
|
||||
"providers/lm_studio",
|
||||
"providers/manus",
|
||||
"providers/meta_llama",
|
||||
"providers/milvus_vector_stores",
|
||||
"providers/mistral",
|
||||
|
|
|
|||
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20.tar.gz
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -0,0 +1,6 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN "router_settings" JSONB DEFAULT '{}';
|
||||
|
||||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "router_settings" JSONB DEFAULT '{}';
|
||||
|
||||
|
|
@ -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"));
|
||||
|
|
@ -124,6 +124,7 @@ model LiteLLM_TeamTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
model_spend Json @default("{}")
|
||||
model_max_budget Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
team_member_permissions String[] @default([])
|
||||
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
|
|
@ -225,6 +226,7 @@ model LiteLLM_VerificationToken {
|
|||
models String[]
|
||||
aliases Json @default("{}")
|
||||
config Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
permissions Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.19"
|
||||
version = "0.4.20"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
|
|
@ -22,7 +22,7 @@ requires = ["poetry-core"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.19"
|
||||
version = "0.4.20"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -909,6 +909,7 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[
|
|||
"twelvelabs",
|
||||
"openai",
|
||||
"stability",
|
||||
"moonshot",
|
||||
]
|
||||
|
||||
BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from typing import (
|
|||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
|
|
@ -815,6 +816,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", {})
|
||||
|
|
@ -1230,11 +1236,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()
|
||||
|
||||
|
|
@ -1248,7 +1260,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(
|
||||
|
|
@ -1268,6 +1279,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,
|
||||
|
|
@ -1293,6 +1437,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,
|
||||
|
|
@ -1307,8 +1459,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,
|
||||
|
|
@ -1346,6 +1498,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,
|
||||
|
|
@ -1401,6 +1558,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(
|
||||
|
|
@ -1415,9 +1581,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"
|
||||
],
|
||||
|
|
@ -1482,6 +1646,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", {})
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -2137,6 +2137,14 @@ def anthropic_messages_pt( # noqa: PLR0915
|
|||
assistant_content.append(
|
||||
cast(AnthropicMessagesTextParam, _cached_message)
|
||||
)
|
||||
# handle server_tool_use blocks (tool search, web search, etc.)
|
||||
# Pass through as-is since these are Anthropic-native content types
|
||||
elif m.get("type", "") == "server_tool_use":
|
||||
assistant_content.append(m) # type: ignore
|
||||
# handle tool_search_tool_result blocks
|
||||
# Pass through as-is since these are Anthropic-native content types
|
||||
elif m.get("type", "") == "tool_search_tool_result":
|
||||
assistant_content.append(m) # type: ignore
|
||||
elif (
|
||||
"content" in assistant_content_block
|
||||
and isinstance(assistant_content_block["content"], str)
|
||||
|
|
@ -3168,6 +3176,11 @@ def _convert_to_bedrock_tool_call_invoke(
|
|||
id = tool["id"]
|
||||
name = tool["function"].get("name", "")
|
||||
arguments = tool["function"].get("arguments", "")
|
||||
arguments_dict = json.loads(arguments) if arguments else {}
|
||||
# Ensure arguments_dict is always a dict (Bedrock requires toolUse.input to be an object)
|
||||
# When some providers return arguments: '""' (JSON-encoded empty string), json.loads returns ""
|
||||
if not isinstance(arguments_dict, dict):
|
||||
arguments_dict = {}
|
||||
if not arguments or not arguments.strip():
|
||||
arguments_dict = {}
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
|
|
|
|||
2
litellm/llms/manus/__init__.py
Normal file
2
litellm/llms/manus/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
# Manus provider implementation
|
||||
|
||||
2
litellm/llms/manus/responses/__init__.py
Normal file
2
litellm/llms/manus/responses/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
# Manus Responses API implementation
|
||||
|
||||
308
litellm/llms/manus/responses/transformation.py
Normal file
308
litellm/llms/manus/responses/transformation.py
Normal 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
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -863,6 +863,7 @@ class KeyRequestBase(GenerateRequestBase):
|
|||
tpm_limit_type: Optional[
|
||||
Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"]
|
||||
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
|
||||
router_settings: Optional[UpdateRouterConfig] = None
|
||||
|
||||
|
||||
class LiteLLMKeyType(str, enum.Enum):
|
||||
|
|
@ -918,6 +919,7 @@ class GenerateKeyResponse(KeyRequestBase):
|
|||
"config",
|
||||
"permissions",
|
||||
"model_max_budget",
|
||||
"router_settings",
|
||||
]
|
||||
for field in dict_fields:
|
||||
value = values.get(field)
|
||||
|
|
@ -1460,6 +1462,7 @@ class TeamBase(LiteLLMPydanticObjectBase):
|
|||
|
||||
models: list = []
|
||||
blocked: bool = False
|
||||
router_settings: Optional[dict] = None
|
||||
|
||||
|
||||
class NewTeamRequest(TeamBase):
|
||||
|
|
@ -1542,6 +1545,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
model_rpm_limit: Optional[Dict[str, int]] = None
|
||||
model_tpm_limit: Optional[Dict[str, int]] = None
|
||||
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
|
||||
router_settings: Optional[dict] = None
|
||||
|
||||
|
||||
class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase):
|
||||
|
|
@ -1684,6 +1688,7 @@ class LiteLLM_TeamTable(TeamBase):
|
|||
"permissions",
|
||||
"model_max_budget",
|
||||
"model_aliases",
|
||||
"router_settings",
|
||||
]
|
||||
|
||||
if isinstance(values, BaseModel):
|
||||
|
|
|
|||
214
litellm/proxy/common_utils/performance_utils.md
Normal file
214
litellm/proxy/common_utils/performance_utils.md
Normal file
|
|
@ -0,0 +1,214 @@
|
|||
# Performance Utilities Documentation
|
||||
|
||||
This module provides performance monitoring and profiling functionality for LiteLLM proxy server using `cProfile` and `line_profiler`.
|
||||
|
||||
## Table of Contents
|
||||
|
||||
- [Line Profiler Usage](#line-profiler-usage)
|
||||
- [Example 1: Wrapping a function directly](#example-1-wrapping-a-function-directly)
|
||||
- [Example 2: Wrapping a module function dynamically](#example-2-wrapping-a-module-function-dynamically)
|
||||
- [Example 3: Manual stats collection](#example-3-manual-stats-collection)
|
||||
- [Example 4: Analyzing the profile output](#example-4-analyzing-the-profile-output)
|
||||
- [Example 5: Using in a decorator pattern](#example-5-using-in-a-decorator-pattern)
|
||||
- [cProfile Usage](#cprofile-usage)
|
||||
- [Installation](#installation)
|
||||
- [Notes](#notes)
|
||||
|
||||
## Line Profiler Usage
|
||||
|
||||
### Example 1: Wrapping a function directly
|
||||
|
||||
This is how it's used in `litellm/utils.py` to profile `wrapper_async`:
|
||||
|
||||
```python
|
||||
from litellm.proxy.common_utils.performance_utils import (
|
||||
register_shutdown_handler,
|
||||
wrap_function_directly,
|
||||
)
|
||||
|
||||
def client(original_function):
|
||||
@wraps(original_function)
|
||||
async def wrapper_async(*args, **kwargs):
|
||||
# ... function implementation ...
|
||||
pass
|
||||
|
||||
# Wrap the function with line_profiler
|
||||
wrapper_async = wrap_function_directly(wrapper_async)
|
||||
|
||||
# Register shutdown handler to collect stats on server shutdown
|
||||
register_shutdown_handler(output_file="wrapper_async_line_profile.lprof")
|
||||
|
||||
return wrapper_async
|
||||
```
|
||||
|
||||
### Example 2: Wrapping a module function dynamically
|
||||
|
||||
```python
|
||||
import my_module
|
||||
from litellm.proxy.common_utils.performance_utils import (
|
||||
wrap_function_with_line_profiler,
|
||||
register_shutdown_handler,
|
||||
)
|
||||
|
||||
# Wrap a function in a module
|
||||
wrap_function_with_line_profiler(my_module, "expensive_function")
|
||||
|
||||
# Register shutdown handler
|
||||
register_shutdown_handler(output_file="my_profile.lprof")
|
||||
|
||||
# Now all calls to my_module.expensive_function will be profiled
|
||||
my_module.expensive_function()
|
||||
```
|
||||
|
||||
### Example 3: Manual stats collection
|
||||
|
||||
```python
|
||||
from litellm.proxy.common_utils.performance_utils import (
|
||||
wrap_function_directly,
|
||||
collect_line_profiler_stats,
|
||||
)
|
||||
|
||||
def my_function():
|
||||
# ... implementation ...
|
||||
pass
|
||||
|
||||
# Wrap the function
|
||||
my_function = wrap_function_directly(my_function)
|
||||
|
||||
# Run your code
|
||||
my_function()
|
||||
|
||||
# Collect stats manually (instead of waiting for shutdown)
|
||||
collect_line_profiler_stats(output_file="manual_profile.lprof")
|
||||
```
|
||||
|
||||
### Example 4: Analyzing the profile output
|
||||
|
||||
After running your code, analyze the `.lprof` file:
|
||||
|
||||
```bash
|
||||
# View the profile
|
||||
python -m line_profiler wrapper_async_line_profile.lprof
|
||||
|
||||
# Save to text file
|
||||
python -m line_profiler wrapper_async_line_profile.lprof > profile_report.txt
|
||||
```
|
||||
|
||||
The output shows:
|
||||
- **Line #**: Line number in the source file
|
||||
- **Hits**: Number of times the line was executed
|
||||
- **Time**: Total time spent on that line (in microseconds)
|
||||
- **Per Hit**: Average time per execution
|
||||
- **% Time**: Percentage of total function time
|
||||
- **Line Contents**: The actual source code
|
||||
|
||||
Example output:
|
||||
```
|
||||
Timer unit: 1e-06 s
|
||||
|
||||
Total time: 3.73697 s
|
||||
File: litellm/utils.py
|
||||
Function: client.<locals>.wrapper_async at line 1657
|
||||
|
||||
Line # Hits Time Per Hit % Time Line Contents
|
||||
==============================================================
|
||||
1657 @wraps(original_function)
|
||||
1658 async def wrapper_async(*args, **kwargs):
|
||||
1659 2005 7577.1 3.8 0.2 print_args_passed_to_litellm(...)
|
||||
1763 2005 1351909.0 674.3 36.2 result = await original_function(*args, **kwargs)
|
||||
1846 4010 1543688.1 385.0 41.3 update_response_metadata(...)
|
||||
```
|
||||
|
||||
### Example 5: Using in a decorator pattern
|
||||
|
||||
```python
|
||||
from litellm.proxy.common_utils.performance_utils import (
|
||||
wrap_function_directly,
|
||||
register_shutdown_handler,
|
||||
)
|
||||
|
||||
def profile_decorator(func):
|
||||
# Wrap the function
|
||||
profiled_func = wrap_function_directly(func)
|
||||
|
||||
# Register shutdown handler (only once)
|
||||
if not hasattr(profile_decorator, '_registered'):
|
||||
register_shutdown_handler(output_file="decorated_functions.lprof")
|
||||
profile_decorator._registered = True
|
||||
|
||||
return profiled_func
|
||||
|
||||
@profile_decorator
|
||||
async def my_async_function():
|
||||
# This function will be profiled
|
||||
pass
|
||||
```
|
||||
|
||||
## cProfile Usage
|
||||
|
||||
### Example: Using the profile_endpoint decorator
|
||||
|
||||
```python
|
||||
from litellm.proxy.common_utils.performance_utils import profile_endpoint
|
||||
|
||||
@profile_endpoint(sampling_rate=0.1) # Profile 10% of requests
|
||||
async def my_endpoint():
|
||||
# ... implementation ...
|
||||
pass
|
||||
```
|
||||
|
||||
The `sampling_rate` parameter controls what percentage of requests are profiled:
|
||||
- `1.0`: Profile all requests (100%)
|
||||
- `0.1`: Profile 1 in 10 requests (10%)
|
||||
- `0.0`: Profile no requests (0%)
|
||||
|
||||
## Installation
|
||||
|
||||
`line_profiler` must be installed to use the line profiling functionality:
|
||||
|
||||
```bash
|
||||
pip install line_profiler
|
||||
```
|
||||
|
||||
On Windows with Python 3.14+, you may need to install Microsoft Visual C++ Build Tools to compile `line_profiler` from source.
|
||||
|
||||
## Notes
|
||||
|
||||
- The profiler aggregates stats by source code location, so multiple instances of the same function (e.g., closures) will be profiled together
|
||||
- Stats are automatically collected on server shutdown via `atexit` handler when using `register_shutdown_handler()`
|
||||
- You can also manually collect stats using `collect_line_profiler_stats()`
|
||||
- The line profiler will fail with an `ImportError` if `line_profiler` is not installed (as configured in `litellm/utils.py`)
|
||||
|
||||
## API Reference
|
||||
|
||||
### `wrap_function_directly(func: Callable) -> Callable`
|
||||
|
||||
Wrap a function directly with line_profiler. This is the recommended way to profile functions, especially closures or functions created dynamically.
|
||||
|
||||
**Raises:**
|
||||
- `ImportError`: If line_profiler is not available
|
||||
- `RuntimeError`: If line_profiler cannot be enabled or function cannot be wrapped
|
||||
|
||||
### `wrap_function_with_line_profiler(module: Any, function_name: str) -> bool`
|
||||
|
||||
Dynamically wrap a function in a module with line_profiler.
|
||||
|
||||
**Returns:** `True` if wrapping was successful, `False` otherwise
|
||||
|
||||
### `collect_line_profiler_stats(output_file: Optional[str] = None) -> None`
|
||||
|
||||
Collect and save line_profiler statistics. If `output_file` is provided, saves to file. Otherwise, prints to stdout.
|
||||
|
||||
### `register_shutdown_handler(output_file: Optional[str] = None) -> None`
|
||||
|
||||
Register an `atexit` handler that will automatically save profiling statistics when the Python process exits. Safe to call multiple times (only registers once).
|
||||
|
||||
**Default output file:** `line_profile_stats.lprof` if not specified
|
||||
|
||||
### `profile_endpoint(sampling_rate: float = 1.0)`
|
||||
|
||||
Decorator to sample endpoint hits and save to a profile file using cProfile.
|
||||
|
||||
**Args:**
|
||||
- `sampling_rate`: Rate of requests to profile (0.0 to 1.0)
|
||||
|
||||
|
|
@ -2,14 +2,19 @@
|
|||
Performance utilities for LiteLLM proxy server.
|
||||
|
||||
This module provides performance monitoring and profiling functionality for endpoint
|
||||
performance analysis using cProfile with configurable sampling rates.
|
||||
performance analysis using cProfile with configurable sampling rates, and line_profiler
|
||||
for line-by-line profiling.
|
||||
|
||||
See performance_utils.md for detailed usage examples and documentation.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import atexit
|
||||
import cProfile
|
||||
import functools
|
||||
import threading
|
||||
from pathlib import Path as PathLib
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
|
|
@ -20,6 +25,11 @@ _last_profile_file_path = None
|
|||
_sample_counter = 0
|
||||
_sample_counter_lock = threading.Lock()
|
||||
|
||||
# Global line_profiler state
|
||||
_line_profiler: Optional[Any] = None
|
||||
_line_profiler_lock = threading.Lock()
|
||||
_wrapped_functions: dict[str, Callable] = {} # Store original functions
|
||||
|
||||
|
||||
def _should_sample(profile_sampling_rate: float) -> bool:
|
||||
"""Determine if current request should be sampled based on sampling rate."""
|
||||
|
|
@ -123,3 +133,156 @@ def profile_endpoint(sampling_rate: float = 1.0):
|
|||
raise
|
||||
return sync_wrapper
|
||||
return decorator
|
||||
|
||||
|
||||
def enable_line_profiler() -> None:
|
||||
"""Enable line_profiler for dynamic function wrapping.
|
||||
|
||||
Raises:
|
||||
ImportError: If line_profiler is not available
|
||||
"""
|
||||
global _line_profiler
|
||||
from line_profiler import LineProfiler # Will raise ImportError if not available
|
||||
|
||||
with _line_profiler_lock:
|
||||
if _line_profiler is None:
|
||||
_line_profiler = LineProfiler()
|
||||
verbose_proxy_logger.info("Line profiler enabled")
|
||||
|
||||
|
||||
def wrap_function_with_line_profiler(module: Any, function_name: str) -> bool:
|
||||
"""Dynamically wrap a function with line_profiler.
|
||||
|
||||
Args:
|
||||
module: The module containing the function
|
||||
function_name: Name of the function to wrap
|
||||
|
||||
Returns:
|
||||
True if wrapping was successful, False otherwise
|
||||
"""
|
||||
try:
|
||||
enable_line_profiler() # May raise ImportError if not available
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
if _line_profiler is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
original_function = getattr(module, function_name, None)
|
||||
if original_function is None:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Function {function_name} not found in module {module.__name__}"
|
||||
)
|
||||
return False
|
||||
|
||||
# Store original function if not already wrapped
|
||||
if function_name not in _wrapped_functions:
|
||||
_wrapped_functions[function_name] = original_function
|
||||
|
||||
# Wrap with line_profiler
|
||||
profiled_function = _line_profiler(original_function)
|
||||
setattr(module, function_name, profiled_function)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Wrapped {module.__name__}.{function_name} with line_profiler"
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"Error wrapping {function_name} with line_profiler: {e}"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def wrap_function_directly(func: Callable) -> Callable:
|
||||
"""Wrap a function directly with line_profiler.
|
||||
|
||||
This is the recommended way to profile functions, especially closures or
|
||||
functions created dynamically (like wrapper_async in litellm/utils.py).
|
||||
|
||||
Args:
|
||||
func: The function to wrap
|
||||
|
||||
Returns:
|
||||
The wrapped function that will be profiled when called
|
||||
|
||||
Raises:
|
||||
ImportError: If line_profiler is not available
|
||||
RuntimeError: If line_profiler cannot be enabled or function cannot be wrapped
|
||||
"""
|
||||
import warnings
|
||||
|
||||
enable_line_profiler() # Will raise ImportError if not available
|
||||
|
||||
if _line_profiler is None:
|
||||
raise RuntimeError("Line profiler was not initialized")
|
||||
|
||||
# Suppress warnings about __wrapped__ - we intentionally want to profile the wrapper
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings('ignore', message='.*__wrapped__.*', category=UserWarning)
|
||||
# Add function to line_profiler and wrap it
|
||||
_line_profiler.add_function(func)
|
||||
profiled_function = _line_profiler(func)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Wrapped function {func.__name__} with line_profiler"
|
||||
)
|
||||
return profiled_function
|
||||
|
||||
|
||||
def collect_line_profiler_stats(output_file: Optional[str] = None) -> None:
|
||||
"""Collect and save line_profiler statistics.
|
||||
|
||||
This can be called manually to collect stats at any time, or it's
|
||||
automatically called on shutdown if register_shutdown_handler() was used.
|
||||
|
||||
Args:
|
||||
output_file: Optional path to save stats. If None, prints to stdout.
|
||||
"""
|
||||
global _line_profiler
|
||||
|
||||
with _line_profiler_lock:
|
||||
if _line_profiler is None:
|
||||
verbose_proxy_logger.debug("Line profiler not enabled, nothing to collect")
|
||||
return
|
||||
|
||||
try:
|
||||
if output_file:
|
||||
# Save to file
|
||||
output_path = PathLib(output_file)
|
||||
_line_profiler.dump_stats(str(output_path))
|
||||
verbose_proxy_logger.info(
|
||||
f"Line profiler stats saved to {output_path}"
|
||||
)
|
||||
else:
|
||||
# Print to stdout
|
||||
from io import StringIO
|
||||
|
||||
stream = StringIO()
|
||||
_line_profiler.print_stats(stream=stream)
|
||||
stats_output = stream.getvalue()
|
||||
verbose_proxy_logger.info("Line profiler stats:\n" + stats_output)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error collecting line profiler stats: {e}")
|
||||
|
||||
|
||||
def register_shutdown_handler(output_file: Optional[str] = None) -> None:
|
||||
"""Register a shutdown handler to collect line_profiler stats.
|
||||
|
||||
This registers an atexit handler that will automatically save profiling
|
||||
statistics when the Python process exits. Safe to call multiple times
|
||||
(only registers once).
|
||||
|
||||
Args:
|
||||
output_file: Optional path to save stats on shutdown.
|
||||
Defaults to 'line_profile_stats.lprof'
|
||||
"""
|
||||
if output_file is None:
|
||||
output_file = "line_profile_stats.lprof"
|
||||
|
||||
def shutdown_handler():
|
||||
collect_line_profiler_stats(output_file=output_file)
|
||||
|
||||
atexit.register(shutdown_handler)
|
||||
verbose_proxy_logger.debug(f"Registered line_profiler shutdown handler for {output_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
|
||||
|
|
|
|||
|
|
@ -14,9 +14,10 @@ import copy
|
|||
import json
|
||||
import secrets
|
||||
import traceback
|
||||
import yaml
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Literal, Optional, Tuple, cast
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
|
||||
|
||||
|
|
@ -1033,7 +1034,7 @@ async def generate_key_fn(
|
|||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated (regenerated)
|
||||
- rotation_interval: Optional[str] - How often to auto-rotate this key (e.g., '30s', '30m', '30h', '30d'). Required if auto_rotate=True.
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
Examples:
|
||||
|
||||
|
|
@ -1388,6 +1389,10 @@ async def prepare_key_update_data(
|
|||
if "model_max_budget" in non_default_values:
|
||||
validate_model_max_budget(non_default_values["model_max_budget"])
|
||||
|
||||
# Serialize router_settings to JSON if present
|
||||
if "router_settings" in non_default_values and non_default_values["router_settings"] is not None:
|
||||
non_default_values["router_settings"] = safe_dumps(non_default_values["router_settings"])
|
||||
|
||||
non_default_values = prepare_metadata_fields(
|
||||
data=data, non_default_values=non_default_values, existing_metadata=_metadata
|
||||
)
|
||||
|
|
@ -1489,7 +1494,8 @@ async def update_key_fn(
|
|||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated
|
||||
- rotation_interval: Optional[str] - How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/key/update' \
|
||||
|
|
@ -2080,6 +2086,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None,
|
||||
auto_rotate: Optional[bool] = None,
|
||||
rotation_interval: Optional[str] = None,
|
||||
router_settings: Optional[dict] = None,
|
||||
):
|
||||
from litellm.proxy.proxy_server import premium_user, prisma_client
|
||||
|
||||
|
|
@ -2114,6 +2121,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
aliases_json = json.dumps(aliases)
|
||||
config_json = json.dumps(config)
|
||||
permissions_json = json.dumps(permissions)
|
||||
router_settings_json = safe_dumps(router_settings) if router_settings is not None else safe_dumps({})
|
||||
|
||||
# Add model_rpm_limit and model_tpm_limit to metadata
|
||||
if model_rpm_limit is not None:
|
||||
|
|
@ -2189,6 +2197,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
"updated_by": updated_by,
|
||||
"allowed_routes": allowed_routes or [],
|
||||
"object_permission_id": object_permission_id,
|
||||
"router_settings": router_settings_json,
|
||||
}
|
||||
|
||||
# Add rotation fields if auto_rotate is enabled
|
||||
|
|
@ -2225,6 +2234,13 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
saved_token["model_max_budget"] = json.loads(
|
||||
saved_token["model_max_budget"]
|
||||
)
|
||||
router_settings = cast(Optional[dict], saved_token.get("router_settings"))
|
||||
if router_settings is not None and isinstance(router_settings, str):
|
||||
try:
|
||||
saved_token["router_settings"] = yaml.safe_load(router_settings)
|
||||
except yaml.YAMLError:
|
||||
# If it's not valid JSON/YAML, keep as is or set to empty dict
|
||||
saved_token["router_settings"] = {}
|
||||
|
||||
if saved_token.get("expires", None) is not None and isinstance(
|
||||
saved_token["expires"], datetime
|
||||
|
|
@ -2269,6 +2285,15 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
)
|
||||
key_data["created_at"] = getattr(create_key_response, "created_at", None)
|
||||
key_data["updated_at"] = getattr(create_key_response, "updated_at", None)
|
||||
|
||||
# Deserialize router_settings from JSON string to dict for response
|
||||
router_settings_value = key_data.get("router_settings")
|
||||
if router_settings_value is not None and isinstance(router_settings_value, str):
|
||||
try:
|
||||
key_data["router_settings"] = yaml.safe_load(router_settings_value)
|
||||
except yaml.YAMLError:
|
||||
# If it's not valid JSON/YAML, keep as is or set to empty dict
|
||||
key_data["router_settings"] = {}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"litellm.proxy.proxy_server.generate_key_helper_fn(): Exception occured - {}".format(
|
||||
|
|
|
|||
|
|
@ -100,7 +100,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
TeamMemberAddResult,
|
||||
UpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
|
|
@ -696,8 +696,7 @@ async def new_team( # noqa: PLR0915
|
|||
- allowed_passthrough_routes: Optional[List[str]] - List of allowed pass through routes for the team.
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
- secret_manager_settings: Optional[dict] - Secret manager settings for the team. [Docs](https://docs.litellm.ai/docs/secret_managers/overview)
|
||||
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
Returns:
|
||||
- team_id: (str) Unique team id - used for tracking spend across multiple keys for same team id.
|
||||
|
|
@ -911,6 +910,12 @@ async def new_team( # noqa: PLR0915
|
|||
complete_team_data.members_with_roles = []
|
||||
|
||||
complete_team_data_dict = complete_team_data.model_dump(exclude_none=True)
|
||||
|
||||
# Serialize router_settings to JSON (matching key creation pattern)
|
||||
router_settings_value = getattr(data, "router_settings", None)
|
||||
router_settings_json = safe_dumps(router_settings_value) if router_settings_value is not None else safe_dumps({})
|
||||
complete_team_data_dict["router_settings"] = router_settings_json
|
||||
|
||||
complete_team_data_dict = prisma_client.jsonify_team_object(
|
||||
db_data=complete_team_data_dict
|
||||
)
|
||||
|
|
@ -1234,7 +1239,7 @@ async def update_team( # noqa: PLR0915
|
|||
Example - update team TPM Limit
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
- secret_manager_settings: Optional[dict] - Secret manager settings for the team. [Docs](https://docs.litellm.ai/docs/secret_managers/overview)
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
```
|
||||
curl --location 'http://0.0.0.0:4000/team/update' \
|
||||
|
|
@ -1396,6 +1401,10 @@ async def update_team( # noqa: PLR0915
|
|||
if _model_id is not None:
|
||||
updated_kv["model_id"] = _model_id
|
||||
|
||||
# Serialize router_settings to JSON if present (matching key update pattern)
|
||||
if "router_settings" in updated_kv and updated_kv["router_settings"] is not None:
|
||||
updated_kv["router_settings"] = safe_dumps(updated_kv["router_settings"])
|
||||
|
||||
updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv)
|
||||
team_row: Optional[LiteLLM_TeamTable] = (
|
||||
await prisma_client.db.litellm_teamtable.update(
|
||||
|
|
|
|||
|
|
@ -3402,8 +3402,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:
|
||||
|
|
@ -3412,6 +3412,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:
|
||||
|
|
@ -9762,6 +9765,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 = []
|
||||
"""
|
||||
[
|
||||
|
|
|
|||
|
|
@ -124,6 +124,7 @@ model LiteLLM_TeamTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
model_spend Json @default("{}")
|
||||
model_max_budget Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
team_member_permissions String[] @default([])
|
||||
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
|
|
@ -225,6 +226,7 @@ model LiteLLM_VerificationToken {
|
|||
models String[]
|
||||
aliases Json @default("{}")
|
||||
config Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
permissions Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -59,7 +59,7 @@ websockets = {version = "^15.0.1", optional = true}
|
|||
boto3 = {version = "1.36.0", optional = true}
|
||||
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
|
||||
mcp = {version = "^1.21.2", optional = true, python = ">=3.10"}
|
||||
litellm-proxy-extras = {version = "0.4.19", optional = true}
|
||||
litellm-proxy-extras = {version = "0.4.20", optional = true}
|
||||
rich = {version = "13.7.1", optional = true}
|
||||
litellm-enterprise = {version = "0.1.27", optional = true}
|
||||
diskcache = {version = "^5.6.1", optional = true}
|
||||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ sentry_sdk==2.21.0 # for sentry error handling
|
|||
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
|
||||
cryptography==44.0.1
|
||||
tzdata==2025.1 # IANA time zone database
|
||||
litellm-proxy-extras==0.4.19 # for proxy extras - e.g. prisma migrations
|
||||
litellm-proxy-extras==0.4.20 # for proxy extras - e.g. prisma migrations
|
||||
llm-sandbox==0.3.31 # for skill execution in sandbox
|
||||
### LITELLM PACKAGE DEPENDENCIES
|
||||
python-dotenv==1.0.1 # for env
|
||||
|
|
|
|||
|
|
@ -124,6 +124,7 @@ model LiteLLM_TeamTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
model_spend Json @default("{}")
|
||||
model_max_budget Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
team_member_permissions String[] @default([])
|
||||
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
|
|
@ -225,6 +226,7 @@ model LiteLLM_VerificationToken {
|
|||
models String[]
|
||||
aliases Json @default("{}")
|
||||
config Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
permissions Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
115
tests/llm_responses_api_testing/test_manus_responses_api.py
Normal file
115
tests/llm_responses_api_testing/test_manus_responses_api.py
Normal 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
|
||||
|
||||
290
tests/llm_translation/test_bedrock_moonshot.py
Normal file
290
tests/llm_translation/test_bedrock_moonshot.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -1137,3 +1137,94 @@ def test_bedrock_create_bedrock_block_different_document_formats():
|
|||
assert f"DocumentPDFmessages_" in block["document"]["name"]
|
||||
assert block["document"]["name"].endswith(f"_{format_type}")
|
||||
assert block["document"]["format"] == format_type
|
||||
|
||||
|
||||
def test_anthropic_messages_pt_server_tool_use_passthrough():
|
||||
"""
|
||||
Test that anthropic_messages_pt passes through server_tool_use and
|
||||
tool_search_tool_result blocks in assistant message content.
|
||||
|
||||
These are Anthropic-native content types used for tool search functionality
|
||||
that need to be preserved when reconstructing multi-turn conversations.
|
||||
|
||||
Fixes: https://github.com/BerriAI/litellm/issues/XXXXX
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import anthropic_messages_pt
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "I need help with time information."
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01ABC123",
|
||||
"name": "tool_search_tool_regex",
|
||||
"input": {"query": ".*time.*"}
|
||||
},
|
||||
{
|
||||
"type": "tool_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_01ABC123",
|
||||
"content": {
|
||||
"type": "tool_search_tool_search_result",
|
||||
"tool_references": [
|
||||
{"type": "tool_reference", "tool_name": "get_time"}
|
||||
]
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "I found the time tool. How can I help you?"
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What's the time in New York?"
|
||||
},
|
||||
]
|
||||
|
||||
result = anthropic_messages_pt(
|
||||
messages=messages,
|
||||
model="claude-sonnet-4-5-20250929",
|
||||
llm_provider="anthropic",
|
||||
)
|
||||
|
||||
# Verify we have 3 messages (user, assistant, user)
|
||||
assert len(result) == 3
|
||||
|
||||
# Verify the assistant message content
|
||||
assistant_msg = result[1]
|
||||
assert assistant_msg["role"] == "assistant"
|
||||
assert isinstance(assistant_msg["content"], list)
|
||||
|
||||
# Find the different content block types
|
||||
content_types = [block.get("type") for block in assistant_msg["content"]]
|
||||
|
||||
# Verify server_tool_use block is preserved
|
||||
assert "server_tool_use" in content_types
|
||||
server_tool_use_block = next(
|
||||
b for b in assistant_msg["content"] if b.get("type") == "server_tool_use"
|
||||
)
|
||||
assert server_tool_use_block["id"] == "srvtoolu_01ABC123"
|
||||
assert server_tool_use_block["name"] == "tool_search_tool_regex"
|
||||
assert server_tool_use_block["input"] == {"query": ".*time.*"}
|
||||
|
||||
# Verify tool_search_tool_result block is preserved
|
||||
assert "tool_search_tool_result" in content_types
|
||||
tool_result_block = next(
|
||||
b for b in assistant_msg["content"] if b.get("type") == "tool_search_tool_result"
|
||||
)
|
||||
assert tool_result_block["tool_use_id"] == "srvtoolu_01ABC123"
|
||||
assert tool_result_block["content"]["type"] == "tool_search_tool_search_result"
|
||||
assert tool_result_block["content"]["tool_references"][0]["tool_name"] == "get_time"
|
||||
|
||||
# Verify text block is also preserved
|
||||
assert "text" in content_types
|
||||
text_block = next(
|
||||
b for b in assistant_msg["content"] if b.get("type") == "text"
|
||||
)
|
||||
assert text_block["text"] == "I found the time tool. How can I help you?"
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
|
|
|||
2
tests/test_litellm/llms/manus/__init__.py
Normal file
2
tests/test_litellm/llms/manus/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
# Manus provider tests
|
||||
|
||||
2
tests/test_litellm/llms/manus/responses/__init__.py
Normal file
2
tests/test_litellm/llms/manus/responses/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
# Manus Responses API tests
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
150
tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py
Normal file
150
tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py
Normal 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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import os
|
|||
import sys
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -3642,3 +3643,152 @@ async def test_update_key_negative_max_budget():
|
|||
# Should not raise any errors at model level
|
||||
request = UpdateKeyRequest(key="test-key", max_budget=-5.0)
|
||||
assert request.max_budget == -5.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_key_with_router_settings(monkeypatch):
|
||||
"""
|
||||
Test that /key/generate correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when saving to database
|
||||
3. Storing router_settings in the key record
|
||||
"""
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.jsonify_object = lambda data: data
|
||||
|
||||
# Mock prisma_client.insert_data for both user and key tables
|
||||
async def _insert_data_side_effect(*args, **kwargs):
|
||||
table_name = kwargs.get("table_name")
|
||||
if table_name == "user":
|
||||
return MagicMock(models=[], spend=0)
|
||||
elif table_name == "key":
|
||||
return MagicMock(
|
||||
token="hashed_token_router",
|
||||
litellm_budget_table=None,
|
||||
object_permission=None,
|
||||
)
|
||||
return MagicMock()
|
||||
|
||||
mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect)
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken = MagicMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
generate_key_fn,
|
||||
)
|
||||
|
||||
# Test router_settings with sample data
|
||||
# Using valid UpdateRouterConfig fields (retry_policy is not a valid field,
|
||||
# but model_group_retry_policy is, which also tests nested dict serialization)
|
||||
router_settings_data = {
|
||||
"routing_strategy": "usage-based",
|
||||
"num_retries": 3,
|
||||
"model_group_retry_policy": {"max_retries": 5},
|
||||
}
|
||||
|
||||
request_data = GenerateKeyRequest(
|
||||
models=["gpt-4"],
|
||||
router_settings=router_settings_data,
|
||||
)
|
||||
|
||||
await generate_key_fn(
|
||||
data=request_data,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="user-router-1",
|
||||
),
|
||||
)
|
||||
|
||||
# Verify key insertion was called
|
||||
assert mock_prisma_client.insert_data.call_count >= 1
|
||||
key_insert_calls = [
|
||||
call.kwargs
|
||||
for call in mock_prisma_client.insert_data.call_args_list
|
||||
if call.kwargs.get("table_name") == "key"
|
||||
]
|
||||
assert len(key_insert_calls) >= 1
|
||||
key_data = key_insert_calls[0]["data"]
|
||||
|
||||
# Verify router_settings is present
|
||||
assert "router_settings" in key_data
|
||||
|
||||
# router_settings should be present in the data passed to insert_data
|
||||
# The code uses safe_dumps to serialize router_settings, so it will be a JSON string
|
||||
router_settings_value = key_data["router_settings"]
|
||||
|
||||
# Get the actual settings value for comparison
|
||||
# The code uses safe_dumps to serialize and yaml.safe_load to deserialize
|
||||
if isinstance(router_settings_value, str):
|
||||
# If it's a JSON string (from safe_dumps), deserialize it using json.loads
|
||||
# (safe_dumps produces JSON, and json.loads is the correct way to deserialize it)
|
||||
actual_settings = json.loads(router_settings_value)
|
||||
elif isinstance(router_settings_value, dict):
|
||||
# If it's still a dict, use it directly
|
||||
actual_settings = router_settings_value
|
||||
else:
|
||||
raise AssertionError(
|
||||
f"router_settings should be str or dict, got {type(router_settings_value)}"
|
||||
)
|
||||
|
||||
# Verify router_settings matches input (regardless of serialization state)
|
||||
assert actual_settings == router_settings_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_with_router_settings(monkeypatch):
|
||||
"""
|
||||
Test that /key/update correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when updating database
|
||||
3. Updating router_settings in the key record
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
prepare_key_update_data,
|
||||
)
|
||||
|
||||
# Mock existing key
|
||||
existing_key = LiteLLM_VerificationToken(
|
||||
token="test-token-router",
|
||||
key_alias="test-key",
|
||||
models=["gpt-3.5-turbo"],
|
||||
user_id="test-user",
|
||||
team_id=None,
|
||||
auto_rotate=False,
|
||||
rotation_interval=None,
|
||||
metadata={},
|
||||
)
|
||||
|
||||
# Test updating router_settings
|
||||
router_settings_data = {
|
||||
"routing_strategy": "latency-based",
|
||||
"num_retries": 2,
|
||||
}
|
||||
|
||||
update_request = UpdateKeyRequest(
|
||||
key="test-token-router", router_settings=router_settings_data
|
||||
)
|
||||
|
||||
result = await prepare_key_update_data(
|
||||
data=update_request, existing_key_row=existing_key
|
||||
)
|
||||
|
||||
# Verify router_settings is serialized to JSON string
|
||||
assert "router_settings" in result
|
||||
assert isinstance(result["router_settings"], str)
|
||||
|
||||
# Verify router_settings can be deserialized and matches input
|
||||
deserialized_settings = json.loads(result["router_settings"])
|
||||
assert deserialized_settings == router_settings_data
|
||||
|
|
|
|||
|
|
@ -4393,3 +4393,162 @@ async def test_new_team_positive_budgets_accepted():
|
|||
)
|
||||
assert request.max_budget == 100.0
|
||||
assert request.team_member_budget == 50.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
Test that /team/new correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when saving to database
|
||||
3. Storing router_settings in the team record
|
||||
"""
|
||||
# Configure mocked prisma client
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.get_data = AsyncMock(return_value=None)
|
||||
mock_db_client.update_data = AsyncMock(return_value=MagicMock())
|
||||
mock_db_client.db = MagicMock()
|
||||
|
||||
# Mock model table creation
|
||||
mock_db_client.db.litellm_modeltable = MagicMock()
|
||||
mock_db_client.db.litellm_modeltable.create = AsyncMock(
|
||||
return_value=MagicMock(id="model123")
|
||||
)
|
||||
|
||||
# Capture team table creation
|
||||
team_create_result = MagicMock(
|
||||
team_id="team-router-456",
|
||||
)
|
||||
team_create_result.model_dump.return_value = {
|
||||
"team_id": "team-router-456",
|
||||
}
|
||||
mock_team_create = AsyncMock(return_value=team_create_result)
|
||||
mock_team_count = AsyncMock(return_value=0)
|
||||
mock_db_client.db.litellm_teamtable = MagicMock()
|
||||
mock_db_client.db.litellm_teamtable.create = mock_team_create
|
||||
mock_db_client.db.litellm_teamtable.count = mock_team_count
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(
|
||||
return_value=team_create_result
|
||||
)
|
||||
|
||||
# Mock user table
|
||||
mock_db_client.db.litellm_usertable = MagicMock()
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import NewTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
|
||||
# Test router_settings with sample data
|
||||
router_settings_data = {
|
||||
"routing_strategy": "usage-based",
|
||||
"num_retries": 3,
|
||||
"retry_policy": {"max_retries": 5},
|
||||
}
|
||||
|
||||
# Build request with router_settings
|
||||
team_request = NewTeamRequest(
|
||||
team_alias="my-team-router",
|
||||
router_settings=router_settings_data,
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
# Execute the endpoint function
|
||||
await new_team(
|
||||
data=team_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
# Verify team creation was called
|
||||
assert mock_team_create.call_count == 1
|
||||
created_team_kwargs = mock_team_create.call_args.kwargs
|
||||
team_data = created_team_kwargs["data"]
|
||||
|
||||
# Verify router_settings is serialized to JSON string
|
||||
assert "router_settings" in team_data
|
||||
assert isinstance(team_data["router_settings"], str)
|
||||
|
||||
# Verify router_settings can be deserialized and matches input
|
||||
deserialized_settings = json.loads(team_data["router_settings"])
|
||||
assert deserialized_settings == router_settings_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
Test that /team/update correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when updating database
|
||||
3. Updating router_settings in the team record
|
||||
"""
|
||||
# Configure mocked prisma client
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.db = MagicMock()
|
||||
|
||||
# Mock existing team row
|
||||
existing_team_mock = MagicMock()
|
||||
existing_team_mock.team_id = "team-router-update-789"
|
||||
existing_team_mock.organization_id = None
|
||||
existing_team_mock.models = []
|
||||
existing_team_mock.members_with_roles = []
|
||||
existing_team_mock.model_dump.return_value = {
|
||||
"team_id": "team-router-update-789",
|
||||
"organization_id": None,
|
||||
"models": [],
|
||||
"members_with_roles": [],
|
||||
}
|
||||
|
||||
# Mock team table find_unique and update
|
||||
updated_team_result = MagicMock(
|
||||
team_id="team-router-update-789",
|
||||
)
|
||||
updated_team_result.model_dump.return_value = {
|
||||
"team_id": "team-router-update-789",
|
||||
}
|
||||
mock_team_find_unique = AsyncMock(return_value=existing_team_mock)
|
||||
mock_team_update = AsyncMock(return_value=updated_team_result)
|
||||
mock_db_client.db.litellm_teamtable = MagicMock()
|
||||
mock_db_client.db.litellm_teamtable.find_unique = mock_team_find_unique
|
||||
mock_db_client.db.litellm_teamtable.update = mock_team_update
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import UpdateTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import update_team
|
||||
|
||||
# Test router_settings with updated data
|
||||
router_settings_data = {
|
||||
"routing_strategy": "latency-based",
|
||||
"num_retries": 2,
|
||||
}
|
||||
|
||||
# Build update request with router_settings
|
||||
team_update_request = UpdateTeamRequest(
|
||||
team_id="team-router-update-789",
|
||||
router_settings=router_settings_data,
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
# Execute the endpoint function
|
||||
await update_team(
|
||||
data=team_update_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
# Verify team update was called
|
||||
assert mock_team_update.call_count == 1
|
||||
updated_team_kwargs = mock_team_update.call_args.kwargs
|
||||
team_data = updated_team_kwargs["data"]
|
||||
|
||||
# Verify router_settings is serialized to JSON string
|
||||
assert "router_settings" in team_data
|
||||
assert isinstance(team_data["router_settings"], str)
|
||||
|
||||
# Verify router_settings can be deserialized and matches input
|
||||
deserialized_settings = json.loads(team_data["router_settings"])
|
||||
assert deserialized_settings == router_settings_data
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -3036,3 +3036,91 @@ def test_get_image_root_case_uses_current_dir(monkeypatch):
|
|||
|
||||
# Verify FileResponse was called
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}`;
|
||||
|
|
|
|||
|
|
@ -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> = {};
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue