Merge branch 'main' into akshoop/fastuuid-dep-make-optional

This commit is contained in:
Alex Shoop 2025-09-24 01:45:38 +09:00 • committed by GitHub
commit c659b7e587
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
83 changed files with 4199 additions and 636 deletions

48
.github/workflows/test-mcp.yml vendored Normal file
View file

@ -0,0 +1,48 @@
name: LiteLLM MCP Tests (folder - tests/mcp_tests)
on:
pull_request:
branches: [ main ]
jobs:
test:
runs-on: ubuntu-latest
timeout-minutes: 25
steps:
- uses: actions/checkout@v4
- name: Thank You Message
run: |
echo "### 🙏 Thank you for contributing to LiteLLM!" >> $GITHUB_STEP_SUMMARY
echo "Your PR is being tested now. We appreciate your help in making LiteLLM better!" >> $GITHUB_STEP_SUMMARY
- name: Set up Python
uses: actions/setup-python@v4
with:
python-version: '3.12'
- name: Install Poetry
uses: snok/install-poetry@v1
- name: Install dependencies
run: |
poetry install --with dev,proxy-dev --extras "proxy semantic-router"
poetry run pip install "pytest==7.3.1"
poetry run pip install "pytest-retry==1.6.3"
poetry run pip install "pytest-cov==5.0.0"
poetry run pip install "pytest-asyncio==0.21.1"
poetry run pip install "respx==0.22.0"
poetry run pip install "pydantic==2.10.2"
poetry run pip install "mcp==1.10.1"
poetry run pip install pytest-xdist
- name: Setup litellm-enterprise as local package
run: |
cd enterprise
python -m pip install -e .
cd ..
- name: Run MCP tests
run: |
poetry run pytest tests/mcp_tests -x -vv -n 4 --cov=litellm --cov-report=xml --durations=5

View file

@ -0,0 +1,62 @@
#!/usr/bin/env python3
"""
Example: Using CLI token with LiteLLM SDK
This example shows how to use the CLI authentication token
in your Python scripts after running `litellm-proxy login`.
"""
from textwrap import indent
import litellm
LITELLM_BASE_URL = "http://localhost:4000/"
def main():
"""Using CLI token with LiteLLM SDK"""
print("🚀 Using CLI Token with LiteLLM SDK")
print("=" * 40)
#litellm._turn_on_debug()
# Get the CLI token
api_key = litellm.get_litellm_gateway_api_key()
if not api_key:
print("❌ No CLI token found. Please run 'litellm-proxy login' first.")
return
print("✅ Found CLI token.")
available_models = litellm.get_valid_models(
check_provider_endpoint=True,
custom_llm_provider="litellm_proxy",
api_key=api_key,
api_base=LITELLM_BASE_URL
)
print("✅ Available models:")
if available_models:
for i, model in enumerate(available_models, 1):
print(f" {i:2d}. {model}")
else:
print(" No models available")
# Use with LiteLLM
try:
response = litellm.completion(
model="litellm_proxy/gemini/gemini-2.5-flash",
messages=[{"role": "user", "content": "Hello from CLI token!"}],
api_key=api_key,
base_url=LITELLM_BASE_URL
)
print(f"✅ LLM Response: {response.model_dump_json(indent=4)}")
except Exception as e:
print(f"❌ Error: {e}")
if __name__ == "__main__":
main()
print("\n💡 Tips:")
print("1. Run 'litellm-proxy login' to authenticate first")
print("2. Replace 'https://your-proxy.com' with your actual proxy URL")
print("3. The token is stored locally at ~/.litellm/token.json")

View file

@ -423,7 +423,7 @@ model_list:
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-D '{
-d '{
"model": "llama-3-8b-instruct",
"messages": [
{
@ -431,8 +431,9 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
"content": "What'\''s the weather like in Boston today?"
}
],
"adapater_id": "my-special-adapter-id" # 👈 PROVIDER-SPECIFIC PARAM
}'
"adapater_id": "my-special-adapter-id"
}'
```
## Provider-Specific Metadata Parameters
@ -482,5 +483,4 @@ response = litellm.completion(
```
</TabItem>
</Tabs>
```
</Tabs>

View file

@ -26,6 +26,7 @@ response = completion(
print(response.usage)
```
> **Note:** LiteLLM supports endpoint bridging—if a model does not natively support a requested endpoint, LiteLLM will automatically route the call to the correct supported endpoint (such as bridging `/chat/completions` to `/responses` or vice versa) based on the model's `mode`set in `model_prices_and_context_window`.
## Streaming Usage

View file

@ -1,6 +1,11 @@
import Image from '@theme/IdealImage';
# Enterprise
:::info
✨ SSO is free for up to 5 users. After that, an enterprise license is required. [Get Started with Enterprise here](https://www.litellm.ai/enterprise)
:::
For companies that need SSO, user management and professional support for LiteLLM Proxy
:::info

View file

@ -13,6 +13,8 @@ This is an Enterprise only endpoint [Get Started with Enterprise here](https://c
| Feature | Supported | Notes |
|-------|-------|-------|
| Supported Providers | OpenAI, Azure OpenAI, Vertex AI | - |
#### ⚡️See an exhaustive list of supported models and providers at [models.litellm.ai](https://models.litellm.ai/)
| Cost Tracking | 🟡 | [Let us know if you need this](https://github.com/BerriAI/litellm/issues) |
| Logging | ✅ | Works across all logging integrations |

View file

@ -32,7 +32,8 @@ Next Steps 👉 [Call all supported models - e.g. Claude-2, Llama2-70b, etc.](./
More details 👉
- [Completion() function details](./completion/)
- [All supported models / providers on LiteLLM](./providers/)
- [Overview of supported models / providers on LiteLLM](./providers/)
- [Search all models / providers](https://models.litellm.ai/)
- [Build your own OpenAI proxy](https://github.com/BerriAI/liteLLM-proxy/tree/main)
## streaming

View file

@ -18,6 +18,9 @@ LiteLLM provides image editing functionality that maps to OpenAI's `/images/edit
| Supported LiteLLM Proxy Versions | 1.71.1+ | |
| Supported LLM providers | **OpenAI** | Currently only `openai` is supported |
#### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/)
## Usage
### LiteLLM Python SDK

View file

@ -279,6 +279,8 @@ print(f"response: {response}")
## Supported Providers
#### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/)
| Provider | Documentation Link |
|----------|-------------------|
| OpenAI | [OpenAI Image Generation →](./providers/openai) |

View file

@ -524,6 +524,15 @@ try:
except OpenAIError as e:
print(e)
```
### See How LiteLLM Transforms Your Requests
Want to understand how LiteLLM parses and normalizes your LLM API requests? Use the `/utils/transform_request` endpoint to see exactly how your request is transformed internally.
You can try it out now directly on our Demo App!
Go to the [LiteLLM API docs for transform_request](https://litellm-api.up.railway.app/#/llm%20utils/transform_request_utils_transform_request_post)
LiteLLM will show you the normalized, provider-agnostic version of your request. This is useful for debugging, learning, and understanding how LiteLLM handles different providers and options.
### Logging Observability - Log LLM Input/Output ([Docs](https://docs.litellm.ai/docs/observability/callbacks))
LiteLLM exposes pre defined callbacks to send data to Lunary, MLflow, Langfuse, Helicone, Promptlayer, Traceloop, Slack

View file

@ -130,6 +130,8 @@ Here's the exact json output and type you can expect from all moderation calls:
## **Supported Providers**
#### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/)
| Provider |
|-------------|
| OpenAI |

View file

@ -5,13 +5,15 @@
liteLLM provides `input_callbacks`, `success_callbacks` and `failure_callbacks`, making it easy for you to send data to a particular provider depending on the status of your responses.
:::tip
**New to LiteLLM Callbacks?** Check out our comprehensive [Callback Management Guide](./callback_management.md) to understand when to use different callback hooks like `async_log_success_event` vs `async_post_call_success_hook`.
**New to LiteLLM Callbacks?**
- For proxy/server logging and observability, see the [Proxy Logging Guide](https://docs.litellm.ai/docs/proxy/logging).
- To write your own callback logic, see the [Custom Callbacks Guide](https://docs.litellm.ai/docs/observability/custom_callback).
:::
liteLLM supports:
- [Custom Callback Functions](https://docs.litellm.ai/docs/observability/custom_callback)
- [Callback Management Guide](./callback_management.md) - **Comprehensive guide for choosing the right hooks**
### Supported Callback Integrations
- [Lunary](https://lunary.ai/docs)
- [Langfuse](https://langfuse.com/docs)
- [LangSmith](https://www.langchain.com/langsmith)
@ -21,9 +23,20 @@ liteLLM supports:
- [Sentry](https://docs.sentry.io/platforms/python/)
- [PostHog](https://posthog.com/docs/libraries/python)
- [Slack](https://slack.dev/bolt-python/concepts)
- [Arize](https://docs.arize.com/)
- [PromptLayer](https://docs.promptlayer.com/)
This is **not** an extensive list. Please check the dropdown for all logging integrations.
### Related Cookbooks
Try out our cookbooks for code snippets and interactive demos:
- [Langfuse Callback Example (Colab)](https://colab.research.google.com/github/BerriAI/litellm/blob/main/cookbook/logging_observability/LiteLLM_Langfuse.ipynb)
- [Lunary Callback Example (Colab)](https://colab.research.google.com/github/BerriAI/litellm/blob/main/cookbook/logging_observability/LiteLLM_Lunary.ipynb)
- [Arize Callback Example (Colab)](https://colab.research.google.com/github/BerriAI/litellm/blob/main/cookbook/logging_observability/LiteLLM_Arize.ipynb)
- [Proxy + Langfuse Callback Example (Colab)](https://colab.research.google.com/github/BerriAI/litellm/blob/main/cookbook/logging_observability/LiteLLM_Proxy_Langfuse.ipynb)
- [PromptLayer Callback Example (Colab)](https://colab.research.google.com/github/BerriAI/litellm/blob/main/cookbook/LiteLLM_PromptLayer.ipynb)
### Quick Start
```python

View file

@ -67,6 +67,23 @@ asyncio.run(completion())
- `async_post_call_success_hook` - Access user data + modify responses
- `async_pre_call_hook` - Modify requests before sending
### Example: Modifying the Response in async_post_call_success_hook
You can use `async_post_call_success_hook` to add custom headers or metadata to the response before it is returned to the client. For example:
```python
async def async_post_call_success_hook(data, user_api_key_dict, response):
# Add a custom header to the response
additional_headers = getattr(response, "_hidden_params", {}).get("additional_headers", {}) or {}
additional_headers["x-litellm-custom-header"] = "my-value"
if not hasattr(response, "_hidden_params"):
response._hidden_params = {}
response._hidden_params["additional_headers"] = additional_headers
return response
```
This allows you to inject custom metadata or headers into the response for downstream consumers. You can use this pattern to pass information to clients, proxies, or observability tools.
## Callback Functions
If you just want to log on a specific event (e.g. on input) - you can use callback functions.

View file

@ -0,0 +1,260 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Azure AI Image Editing
Azure AI provides powerful image editing capabilities using FLUX models from Black Forest Labs to modify existing images based on text descriptions.
## Overview
| Property | Details |
|----------|---------|
| Description | Azure AI Image Editing uses FLUX models to modify existing images based on text prompts. |
| Provider Route on LiteLLM | `azure_ai/` |
| Provider Doc | [Azure AI FLUX Models ↗](https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/black-forest-labs-flux-1-kontext-pro-and-flux1-1-pro-now-available-in-azure-ai-f/4434659) |
| Supported Operations | [`/images/edits`](#image-editing) |
## Setup
### API Key & Base URL & API Version
```python showLineNumbers
# Set your Azure AI API credentials
import os
os.environ["AZURE_AI_API_KEY"] = "your-api-key-here"
os.environ["AZURE_AI_API_BASE"] = "your-azure-ai-endpoint" # e.g., https://your-endpoint.eastus2.inference.ai.azure.com/
os.environ["AZURE_AI_API_VERSION"] = "2025-04-01-preview" # Example API version
```
Get your API key and endpoint from [Azure AI Studio](https://ai.azure.com/).
## Supported Models
| Model Name | Description | Cost per Image |
|------------|-------------|----------------|
| `azure_ai/FLUX.1-Kontext-pro` | FLUX 1 Kontext Pro model with enhanced context understanding for editing | $0.04 |
## Image Editing
### Usage - LiteLLM Python SDK
<Tabs>
<TabItem value="basic-edit" label="Basic Usage">
```python showLineNumbers title="Basic Image Editing"
import os
import base64
from pathlib import Path
import litellm
# Set your API credentials
os.environ["AZURE_AI_API_KEY"] = "your-api-key-here"
os.environ["AZURE_AI_API_BASE"] = "your-azure-ai-endpoint"
os.environ["AZURE_AI_API_VERSION"] = "2025-04-01-preview"
# Edit an image with a prompt
response = litellm.image_edit(
model="azure_ai/FLUX.1-Kontext-pro",
image=open("path/to/your/image.png", "rb"),
prompt="Add a winter theme with snow and cold colors",
api_base=os.environ["AZURE_AI_API_BASE"],
api_key=os.environ["AZURE_AI_API_KEY"],
api_version=os.environ["AZURE_AI_API_VERSION"]
)
img_base64 = response.data[0].get("b64_json")
img_bytes = base64.b64decode(img_base64)
path = Path("edited_image.png")
path.write_bytes(img_bytes)
```
</TabItem>
<TabItem value="async-edit" label="Async Usage">
```python showLineNumbers title="Async Image Editing"
import os
import base64
from pathlib import Path
import litellm
import asyncio
# Set your API credentials
os.environ["AZURE_AI_API_KEY"] = "your-api-key-here"
os.environ["AZURE_AI_API_BASE"] = "your-azure-ai-endpoint"
os.environ["AZURE_AI_API_VERSION"] = "2025-04-01-preview"
async def edit_image():
# Edit image asynchronously
response = await litellm.aimage_edit(
model="azure_ai/FLUX.1-Kontext-pro",
image=open("path/to/your/image.png", "rb"),
prompt="Make this image look like a watercolor painting",
api_base=os.environ["AZURE_AI_API_BASE"],
api_key=os.environ["AZURE_AI_API_KEY"],
api_version=os.environ["AZURE_AI_API_VERSION"]
)
img_base64 = response.data[0].get("b64_json")
img_bytes = base64.b64decode(img_base64)
path = Path("async_edited_image.png")
path.write_bytes(img_bytes)
# Run the async function
asyncio.run(edit_image())
```
</TabItem>
<TabItem value="advanced-edit" label="Advanced Parameters">
```python showLineNumbers title="Advanced Image Editing with Parameters"
import os
import base64
from pathlib import Path
import litellm
# Set your API credentials
os.environ["AZURE_AI_API_KEY"] = "your-api-key-here"
os.environ["AZURE_AI_API_BASE"] = "your-azure-ai-endpoint"
os.environ["AZURE_AI_API_VERSION"] = "2025-04-01-preview"
# Edit image with additional parameters
response = litellm.image_edit(
model="azure_ai/FLUX.1-Kontext-pro",
image=open("path/to/your/image.png", "rb"),
prompt="Add magical elements like floating crystals and mystical lighting",
api_base=os.environ["AZURE_AI_API_BASE"],
api_key=os.environ["AZURE_AI_API_KEY"],
api_version=os.environ["AZURE_AI_API_VERSION"],
n=1
)
img_base64 = response.data[0].get("b64_json")
img_bytes = base64.b64decode(img_base64)
path = Path("advanced_edited_image.png")
path.write_bytes(img_bytes)
```
</TabItem>
</Tabs>
### Usage - LiteLLM Proxy Server
#### 1. Configure your config.yaml
```yaml showLineNumbers title="Azure AI Image Editing Configuration"
model_list:
- model_name: azure-flux-kontext-edit
litellm_params:
model: azure_ai/FLUX.1-Kontext-pro
api_key: os.environ/AZURE_AI_API_KEY
api_base: os.environ/AZURE_AI_API_BASE
api_version: os.environ/AZURE_AI_API_VERSION
model_info:
mode: image_edit
general_settings:
master_key: sk-1234
```
#### 2. Start LiteLLM Proxy Server
```bash showLineNumbers title="Start LiteLLM Proxy Server"
litellm --config /path/to/config.yaml
# RUNNING on http://0.0.0.0:4000
```
#### 3. Make image editing requests with OpenAI Python SDK
<Tabs>
<TabItem value="openai-edit-sdk" label="OpenAI SDK">
```python showLineNumbers title="Azure AI Image Editing via Proxy - OpenAI SDK"
from openai import OpenAI
# Initialize client with your proxy URL
client = OpenAI(
base_url="http://localhost:4000", # Your proxy URL
api_key="sk-1234" # Your proxy API key
)
# Edit image with FLUX Kontext Pro
response = client.images.edit(
model="azure-flux-kontext-edit",
image=open("path/to/your/image.png", "rb"),
prompt="Transform this image into a beautiful oil painting style",
)
img_base64 = response.data[0].b64_json
img_bytes = base64.b64decode(img_base64)
path = Path("proxy_edited_image.png")
path.write_bytes(img_bytes)
```
</TabItem>
<TabItem value="litellm-edit-sdk" label="LiteLLM SDK">
```python showLineNumbers title="Azure AI Image Editing via Proxy - LiteLLM SDK"
import litellm
# Edit image through proxy
response = litellm.image_edit(
model="litellm_proxy/azure-flux-kontext-edit",
image=open("path/to/your/image.png", "rb"),
prompt="Add a mystical forest background with magical creatures",
api_base="http://localhost:4000",
api_key="sk-1234"
)
img_base64 = response.data[0].b64_json
img_bytes = base64.b64decode(img_base64)
path = Path("proxy_edited_image.png")
path.write_bytes(img_bytes)
```
</TabItem>
<TabItem value="curl-edit" label="cURL">
```bash showLineNumbers title="Azure AI Image Editing via Proxy - cURL"
curl --location 'http://localhost:4000/v1/images/edits' \
--header 'Authorization: Bearer sk-1234' \
--form 'model="azure-flux-kontext-edit"' \
--form 'prompt="Convert this image to a vintage sepia tone with old-fashioned effects"' \
--form 'image=@"path/to/your/image.png"'
```
</TabItem>
</Tabs>
## Supported Parameters
Azure AI Image Editing supports the following OpenAI-compatible parameters:
| Parameter | Type | Description | Default | Example |
|-----------|------|-------------|---------|---------|
| `image` | file | The image file to edit | Required | File object or binary data |
| `prompt` | string | Text description of the desired changes | Required | `"Add snow and winter elements"` |
| `model` | string | The FLUX model to use for editing | Required | `"azure_ai/FLUX.1-Kontext-pro"` |
| `n` | integer | Number of edited images to generate (You can specify only 1) | `1` | `1` |
| `api_base` | string | Your Azure AI endpoint URL | Required | `"https://your-endpoint.eastus2.inference.ai.azure.com/"` |
| `api_key` | string | Your Azure AI API key | Required | Environment variable or direct value |
| `api_version` | string | API version for Azure AI | Required | `"2025-04-01-preview"` |
## Getting Started
1. Create an account at [Azure AI Studio](https://ai.azure.com/)
2. Deploy a FLUX model in your Azure AI Studio workspace
3. Get your API key and endpoint from the deployment details
4. Set your `AZURE_AI_API_KEY`, `AZURE_AI_API_BASE` and `AZURE_AI_API_VERSION` environment variables
5. Prepare your source image
6. Use `litellm.image_edit()` to modify your images with text instructions
## Additional Resources
- [Azure AI Studio Documentation](https://docs.microsoft.com/en-us/azure/ai-services/)
- [FLUX Models Announcement](https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/black-forest-labs-flux-1-kontext-pro-and-flux1-1-pro-now-available-in-azure-ai-f/4434659)

View file

@ -2340,6 +2340,39 @@ response = completion(
Make the bedrock completion call
---
### Required AWS IAM Policy for AssumeRole
To use `aws_role_name` (STS AssumeRole) with LiteLLM, your IAM user or role **must** have permission to call `sts:AssumeRole` on the target role. If you see an error like:
```
An error occurred (AccessDenied) when calling the AssumeRole operation: User: arn:aws:sts::...:assumed-role/litellm-ecs-task-role/... is not authorized to perform: sts:AssumeRole on resource: arn:aws:iam::...:role/Enterprise/BedrockCrossAccountConsumer
```
This means the IAM identity running LiteLLM does **not** have permission to assume the target role. You must update your IAM policy to allow this action.
#### Example IAM Policy
Replace `<TARGET_ROLE_ARN>` with the ARN of the role you want to assume (e.g., `arn:aws:iam::123456789012:role/Enterprise/BedrockCrossAccountConsumer`).
```json
{
"Version": "2012-10-17",
"Statement": [
{
"Effect": "Allow",
"Action": "sts:AssumeRole",
"Resource": "<TARGET_ROLE_ARN>"
}
]
}
```
**Note:** The target role itself must also trust the calling IAM identity (via its trust policy) for AssumeRole to succeed. See [AWS AssumeRole docs](https://docs.aws.amazon.com/IAM/latest/UserGuide/id_roles_use_switch-role-api.html) for more details.
---
<Tabs>
<TabItem value="sdk" label="SDK">

View file

@ -45,7 +45,7 @@ vertex_credentials_json = json.dumps(vertex_credentials)
## COMPLETION CALL
response = completion(
model="vertex_ai/gemini-pro",
model="vertex_ai/gemini-2.5-pro",
messages=[{ "content": "Hello, how are you?","role": "user"}],
vertex_credentials=vertex_credentials_json
)
@ -69,7 +69,7 @@ vertex_credentials_json = json.dumps(vertex_credentials)
response = completion(
model="vertex_ai/gemini-pro",
model="vertex_ai/gemini-2.5-pro",
messages=[{"content": "You are a good bot.","role": "system"}, {"content": "Hello, how are you?","role": "user"}],
vertex_credentials=vertex_credentials_json
)
@ -189,14 +189,26 @@ print(json.loads(completion.choices[0].message.content))
1. Add model to config.yaml
```yaml
model_list:
- model_name: gemini-pro
- model_name: gemini-2.5-pro
litellm_params:
model: vertex_ai/gemini-1.5-pro
vertex_project: "project-id"
vertex_location: "us-central1"
vertex_credentials: "/path/to/service_account.json" # [OPTIONAL] Do this OR `!gcloud auth application-default login` - run this to add vertex credentials to your env
```
or
```yaml
model_list:
- model_name: gemini-pro
litellm_params:
model: vertex_ai/gemini-1.5-pro
litellm_credential_name: vertex-global
vertex_project: project-name-here
vertex_location: global
base_model: gemini
model_info:
provider: Vertex
```
2. Start Proxy
```
@ -210,7 +222,7 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-D '{
"model": "gemini-pro",
"model": "gemini-2.5-pro",
"messages": [
{"role": "user", "content": "List 5 popular cookie recipes."}
],
@ -262,7 +274,7 @@ except JSONSchemaValidationError as e:
1. Add model to config.yaml
```yaml
model_list:
- model_name: gemini-pro
- model_name: gemini-2.5-pro
litellm_params:
model: vertex_ai/gemini-1.5-pro
vertex_project: "project-id"
@ -283,7 +295,7 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-D '{
"model": "gemini-pro",
"model": "gemini-2.5-pro",
"messages": [
{"role": "user", "content": "List 5 popular cookie recipes."}
],
@ -391,7 +403,7 @@ client = OpenAI(
)
response = client.chat.completions.create(
model="gemini-pro",
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "Who won the world cup?"}],
tools=[{"googleSearch": {}}],
)
@ -406,7 +418,7 @@ curl http://localhost:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-1234" \
-d '{
"model": "gemini-pro",
"model": "gemini-2.5-pro",
"messages": [
{"role": "user", "content": "Who won the world cup?"}
],
@ -527,7 +539,7 @@ client = OpenAI(
)
response = client.chat.completions.create(
model="gemini-pro",
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "Who won the world cup?"}],
tools=[{"enterpriseWebSearch": {}}],
)
@ -542,7 +554,7 @@ curl http://localhost:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-1234" \
-d '{
"model": "gemini-pro",
"model": "gemini-2.5-pro",
"messages": [
{"role": "user", "content": "Who won the world cup?"}
],
@ -835,7 +847,7 @@ import litellm
litellm.vertex_project = "hardy-device-38811" # Your Project ID
litellm.vertex_location = "us-central1" # proj location
response = litellm.completion(model="gemini-pro", messages=[{"role": "user", "content": "write code for saying hi from LiteLLM"}])
response = litellm.completion(model="gemini-2.5-pro", messages=[{"role": "user", "content": "write code for saying hi from LiteLLM"}])
```
## Usage with LiteLLM Proxy Server
@ -876,9 +888,9 @@ Here's how to use Vertex AI with the LiteLLM Proxy Server
vertex_location: "us-central1" # proj location
model_list:
-model_name: team1-gemini-pro
-model_name: team1-gemini-2.5-pro
litellm_params:
model: gemini-pro
model: gemini-2.5-pro
```
</TabItem>
@ -905,7 +917,7 @@ Here's how to use Vertex AI with the LiteLLM Proxy Server
)
response = client.chat.completions.create(
model="team1-gemini-pro",
model="team1-gemini-2.5-pro",
messages = [
{
"role": "user",
@ -925,7 +937,7 @@ Here's how to use Vertex AI with the LiteLLM Proxy Server
--header 'Authorization: Bearer sk-1234' \
--header 'Content-Type: application/json' \
--data '{
"model": "team1-gemini-pro",
"model": "team1-gemini-2.5-pro",
"messages": [
{
"role": "user",
@ -975,7 +987,7 @@ vertex_credentials_json = json.dumps(vertex_credentials)
response = completion(
model="vertex_ai/gemini-pro",
model="vertex_ai/gemini-2.5-pro",
messages=[{"content": "You are a good bot.","role": "system"}, {"content": "Hello, how are you?","role": "user"}],
vertex_credentials=vertex_credentials_json,
vertex_project="my-special-project",
@ -1039,7 +1051,7 @@ In certain use-cases you may need to make calls to the models and pass [safety s
```python
response = completion(
model="vertex_ai/gemini-pro",
model="vertex_ai/gemini-2.5-pro",
messages=[{"role": "user", "content": "write code for saying hi from LiteLLM"}]
safety_settings=[
{
@ -1153,7 +1165,7 @@ litellm.vertex_ai_safety_settings = [
},
]
response = completion(
model="vertex_ai/gemini-pro",
model="vertex_ai/gemini-2.5-pro",
messages=[{"role": "user", "content": "write code for saying hi from LiteLLM"}]
)
```
@ -1212,7 +1224,7 @@ litellm.vertex_location = "us-central1 # Your Location
## Gemini Pro
| Model Name | Function Call |
|------------------|--------------------------------------|
| gemini-pro | `completion('gemini-pro', messages)`, `completion('vertex_ai/gemini-pro', messages)` |
| gemini-2.5-pro | `completion('gemini-2.5-pro', messages)`, `completion('vertex_ai/gemini-2.5-pro', messages)` |
## Fine-tuned Models
@ -1307,7 +1319,7 @@ curl --location 'https://0.0.0.0:4000/v1/chat/completions' \
## Gemini Pro Vision
| Model Name | Function Call |
|------------------|--------------------------------------|
| gemini-pro-vision | `completion('gemini-pro-vision', messages)`, `completion('vertex_ai/gemini-pro-vision', messages)`|
| gemini-2.5-pro-vision | `completion('gemini-2.5-pro-vision', messages)`, `completion('vertex_ai/gemini-2.5-pro-vision', messages)`|
## Gemini 1.5 Pro (and Vision)
| Model Name | Function Call |
@ -1321,7 +1333,7 @@ curl --location 'https://0.0.0.0:4000/v1/chat/completions' \
#### Using Gemini Pro Vision
Call `gemini-pro-vision` in the same input/output format as OpenAI [`gpt-4-vision`](https://docs.litellm.ai/docs/providers/openai#openai-vision-models)
Call `gemini-2.5-pro-vision` in the same input/output format as OpenAI [`gpt-4-vision`](https://docs.litellm.ai/docs/providers/openai#openai-vision-models)
LiteLLM Supports the following image types passed in `url`
- Images with Cloud Storage URIs - gs://cloud-samples-data/generative-ai/image/boats.jpeg
@ -1339,7 +1351,7 @@ LiteLLM Supports the following image types passed in `url`
import litellm
response = litellm.completion(
model = "vertex_ai/gemini-pro-vision",
model = "vertex_ai/gemini-2.5-pro-vision",
messages=[
{
"role": "user",
@ -1377,7 +1389,7 @@ image_path = "cached_logo.jpg"
# Getting the base64 string
base64_image = encode_image(image_path)
response = litellm.completion(
model="vertex_ai/gemini-pro-vision",
model="vertex_ai/gemini-2.5-pro-vision",
messages=[
{
"role": "user",
@ -1433,7 +1445,7 @@ tools = [
messages = [{"role": "user", "content": "What's the weather like in Boston today?"}]
response = completion(
model="vertex_ai/gemini-pro-vision",
model="vertex_ai/gemini-2.5-pro-vision",
messages=messages,
tools=tools,
)

View file

@ -958,6 +958,19 @@ curl http://localhost:4000/v1/chat/completions \
</Tabs>
## Redis max_connections
You can set the `max_connections` parameter in your `cache_params` for Redis. This is passed directly to the Redis client and controls the maximum number of simultaneous connections in the pool. If you see errors like `No connection available`, try increasing this value:
```yaml
litellm_settings:
cache: true
cache_params:
type: redis
max_connections: 100
```
## Supported `cache_params` on proxy config.yaml
```yaml
@ -966,6 +979,7 @@ cache_params:
ttl: Optional[float]
default_in_memory_ttl: Optional[float]
default_in_redis_ttl: Optional[float]
max_connections: Optional[Int]
# Type of cache (options: "local", "redis", "s3")
type: s3

View file

@ -50,6 +50,7 @@ litellm_settings:
port: 6379 # The port number for the Redis cache. Required if type is "redis".
password: "your_password" # The password for the Redis cache. Required if type is "redis".
namespace: "litellm.caching.caching" # namespace for redis cache
max_connections: 100 # [OPTIONAL] Set Maximum number of Redis connections. Passed directly to redis-py.
# Optional - Redis Cluster Settings
redis_startup_nodes: [{"host": "127.0.0.1", "port": "7001"}]

View file

@ -1,9 +1,7 @@
# ✨ Event Hooks for SSO Login
:::info
✨ This is an Enterprise only feature [Get Started with Enterprise here](https://www.litellm.ai/enterprise)
✨ SSO is free for up to 5 users. After that, an enterprise license is required. [Get Started with Enterprise here](https://www.litellm.ai/enterprise)
:::
## Overview

View file

@ -84,3 +84,29 @@ LiteLLM emits the following prometheus metrics to monitor the health/status of t
| `litellm_in_memory_spend_update_queue_size` | In-memory aggregate spend values for keys, users, teams, team members, etc.| In-Memory |
| `litellm_redis_spend_update_queue_size` | Redis aggregate spend values for keys, users, teams, etc. | Redis |
## Troubleshooting: Redis Connection Errors
You may see errors like:
```
LiteLLM Redis Caching: async async_increment() - Got exception from REDIS No connection available., Writing value=21
LiteLLM Redis Caching: async set_cache_pipeline() - Got exception from REDIS No connection available., Writing value=None
```
This means all available Redis connections are in use, and LiteLLM cannot obtain a new connection from the pool. This can happen under high load or with many concurrent proxy requests.
**Solution:**
- Increase the `max_connections` parameter in your Redis config section in `proxy_config.yaml` to allow more simultaneous connections. For example:
```yaml
litellm_settings:
cache: True
cache_params:
type: redis
max_connections: 100 # Increase as needed for your traffic
```
Adjust this value based on your expected concurrency and Redis server capacity.

View file

@ -4,6 +4,10 @@ import TabItem from '@theme/TabItem';
# Bedrock Guardrails
:::tip ⚡️
If you haven't set up or authenticated your Bedrock provider yet, see the [Bedrock Provider Setup & Authentication Guide](../../providers/bedrock.md).
:::
LiteLLM supports Bedrock guardrails via the [Bedrock ApplyGuardrail API](https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ApplyGuardrail.html).
## Quick Start

View file

@ -172,6 +172,9 @@ router_settings:
redis_host: <your redis host>
redis_password: <your redis password>
redis_port: 1992
cache_params:
type: redis
max_connections: 100 # maximum Redis connections in the pool; tune based on expected concurrency/load
```
## Router settings on config - routing_strategy, model_group_alias

View file

@ -227,7 +227,7 @@ export PROXY_LOGOUT_URL="https://www.google.com"
<Image img={require('../../img/ui_logout.png')} style={{ width: '400px', height: 'auto' }} />
### Set max budget for internal users
### Set default max budget for internal users
Automatically apply budget per internal user when they sign up. By default the table will be checked every 10 minutes, for users to reset. To modify this, [see this](./users.md#reset-budgets)
@ -239,6 +239,10 @@ litellm_settings:
This sets a max budget of $10 USD for internal users when they sign up.
You can also manage these settings visually in the UI:
<Image img={require('../../img/default_user_settings_admin_ui.png')} style={{ width: '700px', height: 'auto' }} />
This budget only applies to personal keys created by that user - seen under `Default Team` on the UI.
<Image img={require('../../img/max_budget_for_internal_users.png')} style={{ width: '500px', height: 'auto' }} />

View file

@ -27,7 +27,7 @@ Email us @ krrish@berri.ai
## Supported Models for LiteLLM Key
These are the models that currently work with the "sk-litellm-.." keys.
For a complete list of models/providers that you can call with LiteLLM, [check out our provider list](./providers/)
For a complete list of models/providers that you can call with LiteLLM, [check out our provider list](./providers/) or check out [models.litellm.ai](https://models.litellm.ai/)
* OpenAI models - [OpenAI docs](./providers/openai.md)
* gpt-4

View file

@ -109,6 +109,8 @@ curl http://0.0.0.0:4000/rerank \
## **Supported Providers**
#### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/)
| Provider | Link to Usage |
|-------------|--------------------|
| Cohere (v1 + v2 clients) | [Usage](#quick-start) |

View file

@ -3,8 +3,11 @@ import TabItem from '@theme/TabItem';
# /responses [Beta]
LiteLLM provides a BETA endpoint in the spec of [OpenAI's `/responses` API](https://platform.openai.com/docs/api-reference/responses)
Requests to /chat/completions may be bridged here automatically when the provider lacks support for that endpoint. The model’s default `mode` determines how bridging works.(see `model_prices_and_context_window`)
| Feature | Supported | Notes |
|---------|-----------|--------|
| Cost Tracking | ✅ | Works with all supported models |
@ -78,6 +81,43 @@ print(retrieved_response)
# retrieved_response = await litellm.aget_responses(response_id=response_id)
```
#### CANCEL a Response
You can cancel an in-progress response (if supported by the provider):
```python showLineNumbers title="Cancel Response by ID"
import litellm
# First, create a response
response = litellm.responses(
model="openai/o1-pro",
input="Tell me a three sentence bedtime story about a unicorn.",
max_output_tokens=100
)
# Get the response ID
response_id = response.id
# Cancel the response by ID
cancel_response = litellm.cancel_responses(
response_id=response_id
)
print(cancel_response)
# For async usage
# cancel_response = await litellm.acancel_responses(response_id=response_id)
```
**REST API:**
```bash
curl -X POST http://localhost:4000/v1/responses/response_id/cancel \
-H "Authorization: Bearer sk-1234"
```
This will attempt to cancel the in-progress response with the given ID.
**Note:** Not all providers support response cancellation. If unsupported, an error will be raised.
#### DELETE a Response
```python showLineNumbers title="Delete Response by ID"
import litellm

Binary file not shown.

After

Width:  |  Height:  |  Size: 234 KiB

View file

@ -57,31 +57,31 @@ const sidebars = {
type: "category",
label: "Alerting & Monitoring",
items: [
"proxy/prometheus",
"proxy/alerting",
"proxy/pagerduty"
].sort()
"proxy/pagerduty",
"proxy/prometheus"
]
},
{
type: "category",
label: "[Beta] Prompt Management",
items: [
"proxy/prompt_management",
"proxy/custom_prompt_management",
"proxy/native_litellm_prompt",
"proxy/custom_prompt_management"
].sort()
"proxy/prompt_management"
]
},
{
type: "category",
label: "AI Tools (OpenWebUI, Claude Code, etc.)",
items: [
"tutorials/openweb_ui",
"tutorials/openai_codex",
"tutorials/litellm_gemini_cli",
"tutorials/litellm_qwen_code_cli",
"tutorials/github_copilot_integration",
"tutorials/claude_responses_api",
"tutorials/cost_tracking_coding",
"tutorials/github_copilot_integration",
"tutorials/litellm_gemini_cli",
"tutorials/litellm_qwen_code_cli",
"tutorials/openai_codex",
"tutorials/openweb_ui"
]
},
@ -111,41 +111,63 @@ const sidebars = {
label: "Setup & Deployment",
items: [
"proxy/quick_start",
"proxy/deploy",
"proxy/prod",
"proxy/cli",
"proxy/release_cycle",
"proxy/model_management",
"proxy/health",
"proxy/debugging",
"proxy/deploy",
"proxy/health",
"proxy/master_key_rotations",
"proxy/model_management",
"proxy/prod",
"proxy/release_cycle",
],
},
"proxy/demo",
{
type: "category",
label: "Admin UI",
items: [
"proxy/admin_ui_sso",
"proxy/custom_root_ui",
"proxy/custom_sso",
"proxy/model_hub",
"proxy/public_teams",
"proxy/self_serve",
"proxy/ui",
"proxy/ui/bulk_edit_users",
"proxy/ui_credentials",
"tutorials/scim_litellm",
{
type: "category",
label: "UI Logs",
items: [
"proxy/ui_logs",
"proxy/ui_logs_sessions"
]
}
],
},
{
type: "category",
label: "Architecture",
items: ["proxy/architecture", "proxy/control_plane_and_data_plane", "proxy/db_info", "proxy/db_deadlocks", "router_architecture", "proxy/user_management_heirarchy", "proxy/jwt_auth_arch", "proxy/image_handling", "proxy/spend_logs_deletion"],
items: [
"proxy/architecture",
"proxy/control_plane_and_data_plane",
"proxy/db_deadlocks",
"proxy/db_info",
"proxy/image_handling",
"proxy/jwt_auth_arch",
"proxy/spend_logs_deletion",
"proxy/user_management_heirarchy",
"router_architecture"
],
},
{
type: "link",
label: "All Endpoints (Swagger)",
href: "https://litellm-api.up.railway.app/",
},
"proxy/enterprise",
"proxy/management_cli",
{
type: "category",
label: "Making LLM Requests",
items: [
"proxy/user_keys",
"proxy/clientside_auth",
"proxy/request_headers",
"proxy/response_headers",
"proxy/forward_client_headers",
"proxy/model_discovery",
],
},
"proxy/enterprise",
"proxy/management_cli",
{
type: "category",
label: "Authentication",
@ -163,45 +185,25 @@ const sidebars = {
},
{
type: "category",
label: "Model Access",
label: "Budgets + Rate Limits",
items: [
"proxy/model_access",
"proxy/team_model_add"
]
},
{
type: "category",
label: "Admin UI",
items: [
"proxy/ui",
"proxy/admin_ui_sso",
"proxy/custom_root_ui",
"proxy/model_hub",
"proxy/self_serve",
"proxy/public_teams",
"tutorials/scim_litellm",
"proxy/custom_sso",
"proxy/ui_credentials",
"proxy/ui/bulk_edit_users",
{
type: "category",
label: "UI Logs",
items: [
"proxy/ui_logs",
"proxy/ui_logs_sessions"
]
}
"proxy/customers",
"proxy/dynamic_rate_limit",
"proxy/rate_limit_tiers",
"proxy/team_budgets",
"proxy/temporary_budget_increase",
"proxy/users"
],
},
"proxy/caching",
{
type: "category",
label: "Spend Tracking",
items: ["proxy/cost_tracking", "proxy/custom_pricing", "proxy/billing",],
},
{
type: "category",
label: "Budgets + Rate Limits",
items: ["proxy/users", "proxy/temporary_budget_increase", "proxy/rate_limit_tiers", "proxy/team_budgets", "proxy/dynamic_rate_limit", "proxy/customers"],
label: "Create Custom Plugins",
description: "Modify requests, responses, and more",
items: [
"proxy/call_hooks",
"proxy/rules",
]
},
{
type: "link",
@ -212,13 +214,32 @@ const sidebars = {
type: "category",
label: "Logging, Alerting, Metrics",
items: [
"proxy/dynamic_logging",
"proxy/logging",
"proxy/logging_spec",
"proxy/team_logging",
"proxy/dynamic_logging"
"proxy/team_logging"
],
},
{
type: "category",
label: "Making LLM Requests",
items: [
"proxy/user_keys",
"proxy/clientside_auth",
"proxy/request_headers",
"proxy/response_headers",
"proxy/forward_client_headers",
"proxy/model_discovery",
],
},
{
type: "category",
label: "Model Access",
items: [
"proxy/model_access",
"proxy/team_model_add"
]
},
{
type: "category",
label: "Secret Managers",
@ -229,14 +250,13 @@ const sidebars = {
},
{
type: "category",
label: "Create Custom Plugins",
description: "Modify requests, responses, and more",
label: "Spend Tracking",
items: [
"proxy/call_hooks",
"proxy/rules",
]
"proxy/billing",
"proxy/cost_tracking",
"proxy/custom_pricing"
],
},
"proxy/caching",
]
},
{
@ -250,6 +270,23 @@ const sidebars = {
slug: "/supported_endpoints",
},
items: [
"assistants",
{
type: "category",
label: "/audio",
items: [
"audio_transcription",
"text_to_speech",
]
},
{
type: "category",
label: "/batches",
items: [
"batches",
"proxy/managed_batches",
]
},
{
type: "category",
label: "/chat/completions",
@ -266,57 +303,8 @@ const sidebars = {
"completion/http_handler_config",
],
},
"response_api",
"text_completion",
"embedding/supported_embedding",
"anthropic_unified",
"mcp",
"generateContent",
{
type: "category",
label: "/images",
items: [
"image_generation",
"image_edits",
"image_variations",
]
},
{
type: "category",
label: "/audio",
"items": [
"audio_transcription",
"text_to_speech",
]
},
{
type: "category",
label: "/vector_stores",
items: [
"vector_stores/search",
]
},
{
type: "category",
label: "Pass-through Endpoints (Anthropic SDK, etc.)",
items: [
"pass_through/intro",
"pass_through/vertex_ai",
"pass_through/google_ai_studio",
"pass_through/cohere",
"pass_through/vllm",
"pass_through/mistral",
"pass_through/openai_passthrough",
"pass_through/anthropic_completion",
"pass_through/bedrock",
"pass_through/assembly_ai",
"pass_through/langfuse",
"proxy/pass_through",
],
},
"rerank",
"assistants",
{
type: "category",
label: "/files",
@ -325,15 +313,6 @@ const sidebars = {
"proxy/litellm_managed_files",
],
},
{
type: "category",
label: "/batches",
items: [
"batches",
"proxy/managed_batches",
]
},
"realtime",
{
type: "category",
label: "/fine_tuning",
@ -342,8 +321,48 @@ const sidebars = {
"proxy/managed_finetuning",
]
},
"generateContent",
"apply_guardrail",
{
type: "category",
label: "/images",
items: [
"image_edits",
"image_generation",
"image_variations",
]
},
"mcp",
"moderation",
"apply_guardrail",
{
type: "category",
label: "Pass-through Endpoints (Anthropic SDK, etc.)",
items: [
"pass_through/intro",
"pass_through/anthropic_completion",
"pass_through/assembly_ai",
"pass_through/bedrock",
"pass_through/cohere",
"pass_through/google_ai_studio",
"pass_through/langfuse",
"pass_through/mistral",
"pass_through/openai_passthrough",
"pass_through/vertex_ai",
"pass_through/vllm",
"proxy/pass_through"
]
},
"realtime",
"rerank",
"response_api",
"anthropic_unified",
{
type: "category",
label: "/vector_stores",
items: [
"vector_stores/search",
]
},
],
},
{
@ -383,6 +402,7 @@ const sidebars = {
items: [
"providers/azure_ai",
"providers/azure_ai_img",
"providers/azure_ai_img_edit",
]
},
{
@ -499,33 +519,32 @@ const sidebars = {
type: "category",
label: "Guides",
items: [
"exception_mapping",
"completion/audio",
"completion/batching",
"completion/computer_use",
"completion/document_understanding",
"completion/drop_params",
"completion/function_call",
"completion/image_generation_chat",
"completion/json_mode",
"completion/knowledgebase",
"completion/message_trimming",
"completion/model_alias",
"completion/mock_requests",
"completion/predict_outputs",
"completion/prefix",
"completion/prompt_caching",
"completion/prompt_formatting",
"completion/reliable_completions",
"completion/stream",
"completion/provider_specific_params",
"completion/vision",
"completion/web_search",
"exception_mapping",
"guides/finetuned_models",
"guides/security_settings",
"completion/audio",
"completion/image_generation_chat",
"completion/web_search",
"completion/document_understanding",
"completion/vision",
"completion/json_mode",
"reasoning_content",
"completion/computer_use",
"completion/prompt_caching",
"completion/predict_outputs",
"completion/knowledgebase",
"completion/prefix",
"completion/drop_params",
"completion/prompt_formatting",
"completion/stream",
"completion/message_trimming",
"completion/function_call",
"completion/model_alias",
"completion/batching",
"completion/mock_requests",
"completion/reliable_completions",
"proxy/veo_video_generation",
"reasoning_content"
]
},
@ -538,25 +557,37 @@ const sidebars = {
description: "Learn how to load balance, route, and set fallbacks for your LLM requests",
slug: "/routing-load-balancing",
},
items: ["routing", "scheduler", "proxy/load_balancing", "proxy/reliability", "proxy/timeout", "proxy/auto_routing", "proxy/tag_routing", "proxy/provider_budget_routing", "wildcard_routing"],
items: [
"routing",
"scheduler",
"proxy/auto_routing",
"proxy/load_balancing",
"proxy/provider_budget_routing",
"proxy/reliability",
"proxy/tag_routing",
"proxy/timeout",
"wildcard_routing"
],
},
{
type: "category",
label: "LiteLLM Python SDK",
items: [
"set_keys",
"budget_manager",
"caching/all_caches",
"completion/token_usage",
"sdk/headers",
"sdk_custom_pricing",
"embedding/async_embedding",
"embedding/moderation",
"budget_manager",
"caching/all_caches",
"migration",
"sdk_custom_pricing",
{
type: "category",
label: "LangChain, LlamaIndex, Instructor Integration",
items: ["langchain/langchain", "tutorials/instructor"],
},
}
],
},

View file

@ -2328,7 +2328,6 @@ def get_custom_labels_from_tags(tags: List[str]) -> Dict[str, str]:
"tag_Service_web_app_v1": "false",
}
"""
import re
from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
from litellm.types.integrations.prometheus import _sanitize_prometheus_label_name

View file

@ -2,7 +2,6 @@
Enterprise internal user management endpoints
"""
import os
from fastapi import APIRouter, Depends, HTTPException

View file

@ -11,7 +11,7 @@ All /vector_store management endpoints
import copy
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi import APIRouter, Depends, HTTPException
import litellm
from litellm._logging import verbose_proxy_logger

View file

@ -60,6 +60,7 @@ from litellm.constants import (
empower_models,
together_ai_models,
baseten_models,
WANDB_MODELS,
REPEATED_STREAMING_CHUNK_LIMIT,
request_timeout,
open_ai_embedding_models,
@ -242,6 +243,7 @@ novita_api_key: Optional[str] = None
snowflake_key: Optional[str] = None
gradient_ai_api_key: Optional[str] = None
nebius_key: Optional[str] = None
wandb_key: Optional[str] = None
heroku_key: Optional[str] = None
cometapi_key: Optional[str] = None
ovhcloud_key: Optional[str] = None
@ -524,6 +526,7 @@ cometapi_models: Set = set()
oci_models: Set = set()
vercel_ai_gateway_models: Set = set()
volcengine_models: Set = set()
wandb_models: Set = set(WANDB_MODELS)
ovhcloud_models: Set = set()
ovhcloud_embedding_models: Set = set()
@ -740,6 +743,8 @@ def add_known_models():
oci_models.add(key)
elif value.get("litellm_provider") == "volcengine":
volcengine_models.add(key)
elif value.get("litellm_provider") == "wandb":
wandb_models.add(key)
elif value.get("litellm_provider") == "ovhcloud":
ovhcloud_models.add(key)
elif value.get("litellm_provider") == "ovhcloud-embedding-models":
@ -838,6 +843,7 @@ model_list = list(
| heroku_models
| vercel_ai_gateway_models
| volcengine_models
| wandb_models
| ovhcloud_models
)
@ -920,6 +926,7 @@ models_by_provider: dict = {
"cometapi": cometapi_models,
"oci": oci_models,
"volcengine": volcengine_models,
"wandb": wandb_models,
"ovhcloud": ovhcloud_models | ovhcloud_embedding_models,
}
@ -1259,6 +1266,7 @@ from .llms.watsonx.chat.transformation import IBMWatsonXChatConfig
from .llms.watsonx.embed.transformation import IBMWatsonXEmbeddingConfig
from .llms.github_copilot.chat.transformation import GithubCopilotConfig
from .llms.nebius.chat.transformation import NebiusConfig
from .llms.wandb.chat.transformation import WandbConfig
from .llms.dashscope.chat.transformation import DashScopeChatConfig
from .llms.moonshot.chat.transformation import MoonshotChatConfig
from .llms.v0.chat.transformation import V0ChatConfig
@ -1335,5 +1343,8 @@ disable_hf_tokenizer_download: Optional[bool] = (
)
global_disable_no_log_param: bool = False
### CLI UTILITIES ###
from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key
### PASSTHROUGH ###
from .passthrough import allm_passthrough_route, llm_passthrough_route

View file

@ -313,6 +313,7 @@ LITELLM_CHAT_PROVIDERS = [
"morph",
"lambda_ai",
"vercel_ai_gateway",
"wandb",
"ovhcloud",
]
@ -448,6 +449,7 @@ openai_compatible_endpoints: List = [
"https://api.lambda.ai/v1",
"https://api.hyperbolic.xyz/v1",
"https://ai-gateway.vercel.sh/v1",
"https://api.inference.wandb.ai/v1",
]
@ -492,6 +494,7 @@ openai_compatible_providers: List = [
"hyperbolic",
"vercel_ai_gateway",
"aiml",
"wandb",
]
openai_text_completion_compatible_providers: List = (
[ # providers that support `/v1/completions`
@ -507,6 +510,7 @@ openai_text_completion_compatible_providers: List = (
"v0",
"lambda_ai",
"hyperbolic",
"wandb",
]
)
_openai_like_providers: List = [
@ -757,6 +761,38 @@ nebius_embedding_models: set = set(
]
)
WANDB_MODELS: set = set(
[
# openai models
"openai/gpt-oss-120b",
"openai/gpt-oss-20b",
# zai-org models
"zai-org/GLM-4.5",
# Qwen models
"Qwen/Qwen3-235B-A22B-Instruct-2507",
"Qwen/Qwen3-Coder-480B-A35B-Instruct",
"Qwen/Qwen3-235B-A22B-Thinking-2507",
# moonshotai
"moonshotai/Kimi-K2-Instruct",
# meta models
"meta-llama/Llama-3.1-8B-Instruct",
"meta-llama/Llama-3.3-70B-Instruct",
"meta-llama/Llama-4-Scout-17B-16E-Instruct",
# deepseek-ai
"deepseek-ai/DeepSeek-V3.1",
"deepseek-ai/DeepSeek-R1-0528",
"deepseek-ai/DeepSeek-V3-0324",
# microsoft
"microsoft/Phi-4-mini-instruct",
]
)
BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[
"cohere",
"anthropic",
@ -947,6 +983,7 @@ HEALTH_CHECK_TIMEOUT_SECONDS = int(
os.getenv("HEALTH_CHECK_TIMEOUT_SECONDS", 60)
) # 60 seconds
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME = "litellm-internal-health-check"
LITTELM_CLI_SERVICE_ACCOUNT_NAME = "litellm-cli"
UI_SESSION_TOKEN_TEAM_ID = "litellm-dashboard"
LITELLM_PROXY_ADMIN_NAME = "default_user_id"

View file

@ -148,6 +148,8 @@ def cost_per_token( # noqa: PLR0915
### CALL TYPE ###
call_type: CallTypesLiteral = "completion",
audio_transcription_file_duration: float = 0.0, # for audio transcription calls - the file time in seconds
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
) -> Tuple[float, float]: # type: ignore
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -278,6 +280,7 @@ def cost_per_token( # noqa: PLR0915
model=model_without_prefix,
usage=usage_block,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
)
return prompt_cost, completion_cost
@ -327,7 +330,7 @@ def cost_per_token( # noqa: PLR0915
elif custom_llm_provider == "bedrock":
return bedrock_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "openai":
return openai_cost_per_token(model=model, usage=usage_block)
return openai_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
elif custom_llm_provider == "databricks":
return databricks_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "fireworks_ai":
@ -606,6 +609,8 @@ def completion_cost( # noqa: PLR0915
litellm_model_name: Optional[str] = None,
router_model_id: Optional[str] = None,
litellm_logging_obj: Optional[LitellmLoggingObject] = None,
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
) -> float:
"""
Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm.
@ -658,6 +663,10 @@ def completion_cost( # noqa: PLR0915
completion_response=completion_response
)
rerank_billed_units: Optional[RerankBilledUnits] = None
# Extract service_tier from optional_params if not provided directly
if service_tier is None and optional_params is not None:
service_tier = optional_params.get("service_tier")
selected_model = _select_model_name_for_cost_calc(
model=model,
@ -909,6 +918,7 @@ def completion_cost( # noqa: PLR0915
call_type=cast(CallTypesLiteral, call_type),
audio_transcription_file_duration=audio_transcription_file_duration,
rerank_billed_units=rerank_billed_units,
service_tier=service_tier,
)
_final_cost = (
prompt_tokens_cost_usd_dollar + completion_tokens_cost_usd_dollar
@ -1003,6 +1013,8 @@ def response_cost_calculator(
litellm_model_name: Optional[str] = None,
router_model_id: Optional[str] = None,
litellm_logging_obj: Optional[LitellmLoggingObject] = None,
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
) -> float:
"""
Returns
@ -1036,6 +1048,7 @@ def response_cost_calculator(
litellm_model_name=litellm_model_name,
router_model_id=router_model_id,
litellm_logging_obj=litellm_logging_obj,
service_tier=service_tier,
)
return response_cost
except Exception as e:

View file

@ -0,0 +1,58 @@
"""
CLI Token Utilities
SDK-level utilities for reading CLI authentication tokens.
This module has no dependencies on proxy code and can be safely imported at the SDK level.
"""
import json
import os
from pathlib import Path
from typing import Optional
def get_cli_token_file_path() -> str:
"""Get the path to the CLI token file"""
home_dir = Path.home()
config_dir = home_dir / ".litellm"
return str(config_dir / "token.json")
def load_cli_token() -> Optional[dict]:
"""Load CLI token data from file"""
token_file = get_cli_token_file_path()
if not os.path.exists(token_file):
return None
try:
with open(token_file, 'r') as f:
return json.load(f)
except (json.JSONDecodeError, IOError):
return None
def get_litellm_gateway_api_key() -> Optional[str]:
"""
Get the stored CLI API key for use with LiteLLM SDK.
This function reads the token file created by `litellm-proxy login`
and returns the API key for use in Python scripts.
Returns:
str: The API key if found, None otherwise
Example:
>>> import litellm
>>> api_key = litellm.get_litellm_gateway_api_key()
>>> if api_key:
>>> response = litellm.completion(
>>> model="gpt-3.5-turbo",
>>> messages=[{"role": "user", "content": "Hello"}],
>>> api_key=api_key,
>>> base_url="https://your-proxy.com/v1"
>>> )
"""
token_data = load_cli_token()
if token_data and 'key' in token_data:
return token_data['key']
return None

View file

@ -252,6 +252,9 @@ def get_llm_provider( # noqa: PLR0915
elif endpoint == "https://ai-gateway.vercel.sh/v1":
custom_llm_provider = "vercel_ai_gateway"
dynamic_api_key = get_secret_str("VERCEL_AI_GATEWAY_API_KEY")
elif endpoint == "https://api.inference.wandb.ai/v1":
custom_llm_provider = "wandb"
dynamic_api_key = get_secret_str("WANDB_API_KEY")
if api_base is not None and not isinstance(api_base, str):
raise Exception(
@ -773,6 +776,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
) = litellm.AIMLChatConfig()._get_openai_compatible_provider_info(
api_base, api_key
)
elif custom_llm_provider == "wandb":
api_base = (
api_base
or get_secret("WANDB_API_BASE")
or "https://api.inference.wandb.ai/v1"
) # type: ignore
dynamic_api_key = api_key or get_secret_str("WANDB_API_KEY")
if api_base is not None and not isinstance(api_base, str):
raise Exception("api base needs to be a string. api_base={}".format(api_base))

View file

@ -149,6 +149,9 @@ def get_supported_openai_params( # noqa: PLR0915
elif custom_llm_provider == "nebius":
if request_type == "chat_completion":
return litellm.NebiusConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "wandb":
if request_type == "chat_completion":
return litellm.WandbConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "replicate":
return litellm.ReplicateConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "huggingface":

View file

@ -1228,6 +1228,7 @@ class Logging(LiteLLMLoggingBaseClass):
"standard_built_in_tools_params": self.standard_built_in_tools_params,
"router_model_id": router_model_id,
"litellm_logging_obj": self,
"service_tier": self.optional_params.get("service_tier") if self.optional_params else None,
}
except Exception as e: # error creating kwargs for cost calculation
debug_info = StandardLoggingModelCostFailureDebugInformation(

View file

@ -12,6 +12,7 @@ from litellm.types.utils import (
ModelInfo,
PassthroughCallTypes,
Usage,
ServiceTier,
)
from litellm.utils import get_model_info
@ -114,8 +115,30 @@ def _generic_cost_per_character(
return prompt_cost, completion_cost
def _get_service_tier_cost_key(base_key: str, service_tier: Optional[str]) -> str:
"""
Get the appropriate cost key based on service tier.
Args:
base_key: The base cost key (e.g., "input_cost_per_token")
service_tier: The service tier ("flex", "priority", or None for standard)
Returns:
str: The cost key to use (e.g., "input_cost_per_token_flex" or "input_cost_per_token")
"""
if service_tier is None:
return base_key
# Only use service tier specific keys for "flex" and "priority"
if service_tier.lower() in [ServiceTier.FLEX.value, ServiceTier.PRIORITY.value]:
return f"{base_key}_{service_tier.lower()}"
# For any other service tier, use standard pricing
return base_key
def _get_token_base_cost(
model_info: ModelInfo, usage: Usage
model_info: ModelInfo, usage: Usage, service_tier: Optional[str] = None
) -> Tuple[float, float, float, float, float]:
"""
Return prompt cost, completion cost, and cache costs for a given model and usage.
@ -126,21 +149,27 @@ def _get_token_base_cost(
Returns:
Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost)
"""
# Get service tier aware cost keys
input_cost_key = _get_service_tier_cost_key("input_cost_per_token", service_tier)
output_cost_key = _get_service_tier_cost_key("output_cost_per_token", service_tier)
cache_creation_cost_key = _get_service_tier_cost_key("cache_creation_input_token_cost", service_tier)
cache_read_cost_key = _get_service_tier_cost_key("cache_read_input_token_cost", service_tier)
prompt_base_cost = cast(
float, _get_cost_per_unit(model_info, "input_cost_per_token")
float, _get_cost_per_unit(model_info, input_cost_key)
)
completion_base_cost = cast(
float, _get_cost_per_unit(model_info, "output_cost_per_token")
float, _get_cost_per_unit(model_info, output_cost_key)
)
cache_creation_cost = cast(
float, _get_cost_per_unit(model_info, "cache_creation_input_token_cost")
float, _get_cost_per_unit(model_info, cache_creation_cost_key)
)
cache_creation_cost_above_1hr = cast(
float,
_get_cost_per_unit(model_info, "cache_creation_input_token_cost_above_1hr"),
)
cache_read_cost = cast(
float, _get_cost_per_unit(model_info, "cache_read_input_token_cost")
float, _get_cost_per_unit(model_info, cache_read_cost_key)
)
## CHECK IF ABOVE THRESHOLD
@ -249,6 +278,29 @@ def _get_cost_per_unit(
verbose_logger.exception(
f"litellm.litellm_core_utils.llm_cost_calc.utils.py::calculate_cost_per_component(): Exception occured - {cost_per_unit}\nDefaulting to 0.0"
)
# If the service tier key doesn't exist or is None, try to fall back to the standard key
if cost_per_unit is None:
# Check if any service tier suffix exists in the cost key using ServiceTier enum
for service_tier in ServiceTier:
suffix = f"_{service_tier.value}"
if suffix in cost_key:
# Extract the base key by removing the matched suffix
base_key = cost_key.replace(suffix, '')
fallback_cost = model_info.get(base_key)
if isinstance(fallback_cost, float):
return fallback_cost
if isinstance(fallback_cost, int):
return float(fallback_cost)
if isinstance(fallback_cost, str):
try:
return float(fallback_cost)
except ValueError:
verbose_logger.exception(
f"litellm.litellm_core_utils.llm_cost_calc.utils.py::_get_cost_per_unit(): Exception occured - {fallback_cost}\nDefaulting to 0.0"
)
break # Only try the first matching suffix
return default_value
@ -443,7 +495,7 @@ def _calculate_input_cost(
def generic_cost_per_token(
model: str, usage: Usage, custom_llm_provider: str
model: str, usage: Usage, custom_llm_provider: str, service_tier: Optional[str] = None
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -495,7 +547,7 @@ def generic_cost_per_token(
cache_creation_cost,
cache_creation_cost_above_1hr,
cache_read_cost,
) = _get_token_base_cost(model_info=model_info, usage=usage)
) = _get_token_base_cost(model_info=model_info, usage=usage, service_tier=service_tier)
prompt_cost = _calculate_input_cost(
prompt_tokens_details=prompt_tokens_details,

View file

@ -0,0 +1,15 @@
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from .transformation import AzureFoundryFluxImageEditConfig
__all__ = ["AzureFoundryFluxImageEditConfig"]
def get_azure_ai_image_edit_config(model: str) -> BaseImageEditConfig:
model = model.lower()
model = model.replace("-", "")
model = model.replace("_", "")
if model == "" or "flux" in model: # empty model is flux
return AzureFoundryFluxImageEditConfig()
else:
raise ValueError(f"Model {model} is not supported for Azure AI image editing.")

View file

@ -0,0 +1,99 @@
from typing import Optional
import httpx
import litellm
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
from litellm.secret_managers.main import get_secret_str
from litellm.utils import _add_path_to_api_base
class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig):
"""
Azure AI Foundry FLUX image edit config
Supports FLUX models including FLUX-1-kontext-pro for image editing.
Azure AI Foundry FLUX models handle image editing through the /images/edits endpoint,
same as standard Azure OpenAI models. The request format uses multipart/form-data
with image files and prompt.
"""
def validate_environment(
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
) -> dict:
"""
Validate Azure AI Foundry environment and set up authentication
Uses Api-Key header format
"""
api_key = AzureFoundryModelInfo.get_api_key(api_key)
if not api_key:
raise ValueError(
f"Azure AI API key is required for model {model}. Set AZURE_AI_API_KEY environment variable or pass api_key parameter."
)
headers.update(
{
"Api-Key": api_key, # Azure AI Foundry uses Api-Key header format
}
)
return headers
def get_complete_url(
self,
model: str,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Constructs a complete URL for Azure AI Foundry image edits API request.
Azure AI Foundry FLUX models handle image editing through the /images/edits
endpoint.
Args:
- model: Model name (deployment name for Azure AI Foundry)
- api_base: Base URL for Azure AI endpoint
- litellm_params: Additional parameters including api_version
Returns:
- Complete URL for the image edits endpoint
"""
api_base = AzureFoundryModelInfo.get_api_base(api_base)
if api_base is None:
raise ValueError(
"Azure AI API base is required. Set AZURE_AI_API_BASE environment variable or pass api_base parameter."
)
api_version = (litellm_params.get("api_version") or litellm.api_version
or get_secret_str("AZURE_AI_API_VERSION")
)
if api_version is None:
# API version is mandatory for Azure AI Foundry
raise ValueError(
"Azure API version is required. Set AZURE_AI_API_VERSION environment variable or pass api_version parameter."
)
# Add the path to the base URL using the model as deployment name
# Azure AI Foundry FLUX models use /images/edits for editing
if "/openai/deployments/" in api_base:
new_url = _add_path_to_api_base(
api_base=api_base,
ending_path="/images/edits",
)
else:
new_url = _add_path_to_api_base(
api_base=api_base,
ending_path=f"/openai/deployments/{model}/images/edits",
)
# Use the new query_params dictionary
final_url = httpx.URL(new_url).copy_with(params={"api-version": api_version})
return str(final_url)

View file

@ -200,8 +200,8 @@ class AmazonConverseConfig(BaseConfig):
llm_provider="bedrock",
)
key_pattern = re.compile(r'^[a-zA-Z0-9\s:_@$#=/+,.-]{1,256}$')
value_pattern = re.compile(r'^[a-zA-Z0-9\s:_@$#=/+,.-]{0,256}$')
key_pattern = re.compile(r"^[a-zA-Z0-9\s:_@$#=/+,.-]{1,256}$")
value_pattern = re.compile(r"^[a-zA-Z0-9\s:_@$#=/+,.-]{0,256}$")
for key, value in metadata.items():
if not isinstance(key, str):
@ -762,7 +762,9 @@ class AmazonConverseConfig(BaseConfig):
return {}
def _prepare_request_params(self, optional_params: dict, model: str) -> tuple[dict, dict, dict]:
def _prepare_request_params(
self, optional_params: dict, model: str
) -> Tuple[dict, dict, dict]:
"""Prepare and separate request parameters."""
inference_params = copy.deepcopy(optional_params)
supported_converse_params = list(
@ -797,7 +799,13 @@ class AmazonConverseConfig(BaseConfig):
return inference_params, additional_request_params, request_metadata
def _process_tools_and_beta(self, original_tools: list, model: str, headers: Optional[dict], additional_request_params: dict) -> tuple[List[ToolBlock], list]:
def _process_tools_and_beta(
self,
original_tools: list,
model: str,
headers: Optional[dict],
additional_request_params: dict,
) -> Tuple[List[ToolBlock], list]:
"""Process tools and collect anthropic_beta values."""
bedrock_tools: List[ToolBlock] = []
@ -871,12 +879,16 @@ class AmazonConverseConfig(BaseConfig):
)
# Prepare and separate parameters
inference_params, additional_request_params, request_metadata = self._prepare_request_params(optional_params, model)
inference_params, additional_request_params, request_metadata = (
self._prepare_request_params(optional_params, model)
)
original_tools = inference_params.pop("tools", [])
# Process tools and collect beta values
bedrock_tools, anthropic_beta_list = self._process_tools_and_beta(original_tools, model, headers, additional_request_params)
bedrock_tools, anthropic_beta_list = self._process_tools_and_beta(
original_tools, model, headers, additional_request_params
)
bedrock_tool_config: Optional[ToolConfigBlock] = None
if len(bedrock_tools) > 0:
@ -1157,9 +1169,7 @@ class AmazonConverseConfig(BaseConfig):
return message, returned_finish_reason
def _translate_message_content(
self, content_blocks: List[ContentBlock]
) -> Tuple[
def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[
str,
List[ChatCompletionToolCallChunk],
Optional[List[BedrockConverseReasoningContentBlock]],
@ -1174,9 +1184,9 @@ class AmazonConverseConfig(BaseConfig):
"""
content_str = ""
tools: List[ChatCompletionToolCallChunk] = []
reasoningContentBlocks: Optional[
List[BedrockConverseReasoningContentBlock]
] = None
reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
None
)
for idx, content in enumerate(content_blocks):
"""
- Content is either a tool response or text
@ -1297,9 +1307,9 @@ class AmazonConverseConfig(BaseConfig):
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
content_str = ""
tools: List[ChatCompletionToolCallChunk] = []
reasoningContentBlocks: Optional[
List[BedrockConverseReasoningContentBlock]
] = None
reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
None
)
if message is not None:
(
@ -1312,12 +1322,12 @@ class AmazonConverseConfig(BaseConfig):
chat_completion_message["provider_specific_fields"] = {
"reasoningContentBlocks": reasoningContentBlocks,
}
chat_completion_message[
"reasoning_content"
] = self._transform_reasoning_content(reasoningContentBlocks)
chat_completion_message[
"thinking_blocks"
] = self._transform_thinking_blocks(reasoningContentBlocks)
chat_completion_message["reasoning_content"] = (
self._transform_reasoning_content(reasoningContentBlocks)
)
chat_completion_message["thinking_blocks"] = (
self._transform_thinking_blocks(reasoningContentBlocks)
)
chat_completion_message["content"] = content_str
if (
json_mode is True

View file

@ -18,7 +18,7 @@ def cost_router(call_type: CallTypes) -> Literal["cost_per_token", "cost_per_sec
return "cost_per_token"
def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
def cost_per_token(model: str, usage: Usage, service_tier: Optional[str] = None) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -31,7 +31,7 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
"""
## CALCULATE INPUT COST
return generic_cost_per_token(
model=model, usage=usage, custom_llm_provider="openai"
model=model, usage=usage, custom_llm_provider="openai", service_tier=service_tier
)
# ### Non-cached text tokens
# non_cached_text_tokens = usage.prompt_tokens

View file

@ -11,7 +11,21 @@ from litellm.utils import _add_path_to_api_base
class VLLMError(BaseLLMException):
pass
def __init__(
self,
status_code: int,
message: str,
request: Optional[httpx.Request] = None,
response: Optional[httpx.Response] = None,
headers: Optional[Union[httpx.Headers, dict]] = None,
):
super().__init__(
status_code=status_code,
message=message,
request=request,
response=response,
headers=headers,
)
class VLLMModelInfo(BaseLLMModelInfo):
@ -25,7 +39,8 @@ class VLLMModelInfo(BaseLLMModelInfo):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
"""Google AI Studio sends api key in query params"""
if api_key is not None:
headers["x-api-key"] = api_key
return headers
@staticmethod
@ -53,7 +68,7 @@ class VLLMModelInfo(BaseLLMModelInfo):
endpoint = "/v1/models"
if api_base is None or api_key is None:
raise ValueError(
"GEMINI_API_BASE or GEMINI_API_KEY is not set. Please set the environment variable, to query Gemini's `/models` endpoint."
"VLLM_API_BASE or VLLM_API_KEY is not set. Please set the environment variable, to query VLLM's `/models` endpoint."
)
url = _add_path_to_api_base(api_base, endpoint)

View file

View file

View file

@ -0,0 +1,27 @@
"""
Wandb Chat Completions API - Transformation
This is OpenAI compatible - no translation needed / occurs
"""
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
class WandbConfig(OpenAIGPTConfig):
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
map max_completion_tokens param to max_tokens
"""
supported_openai_params = self.get_supported_openai_params(model=model)
for param, value in non_default_params.items():
if param == "max_completion_tokens":
optional_params["max_tokens"] = value
elif param in supported_openai_params:
optional_params[param] = value
return optional_params

View file

@ -1981,6 +1981,7 @@ def completion( # type: ignore # noqa: PLR0915
or custom_llm_provider == "openai"
or custom_llm_provider == "together_ai"
or custom_llm_provider == "nebius"
or custom_llm_provider == "wandb"
or custom_llm_provider in litellm.openai_compatible_providers
or "ft:gpt-3.5-turbo" in model # finetune gpt-3.5-turbo
): # allow user to make an openai call with a custom base
@ -4400,6 +4401,27 @@ def embedding( # noqa: PLR0915
or "api.studio.nebius.ai/v1"
)
response = openai_chat_completions.embedding(
model=model,
input=input,
api_base=api_base,
api_key=api_key,
logging_obj=logging,
timeout=timeout,
model_response=EmbeddingResponse(),
optional_params=optional_params,
client=client,
aembedding=aembedding,
)
elif custom_llm_provider == "wandb":
api_key = api_key or litellm.api_key or get_secret_str("WANDB_API_KEY")
api_base = (
api_base
or litellm.api_base
or get_secret_str("WANDB_API_BASE")
or "https://api.inference.wandb.ai/v1"
)
response = openai_chat_completions.embedding(
model=model,
input=input,

View file

@ -11534,8 +11534,10 @@
},
"gpt-4.1": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_priority": 8.75e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_batches": 1e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 1047576,
"max_output_tokens": 32768,
@ -11543,6 +11545,7 @@
"mode": "chat",
"output_cost_per_token": 8e-06,
"output_cost_per_token_batches": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -11600,8 +11603,10 @@
},
"gpt-4.1-mini": {
"cache_read_input_token_cost": 1e-07,
"cache_read_input_token_cost_priority": 1.75e-07,
"input_cost_per_token": 4e-07,
"input_cost_per_token_batches": 2e-07,
"input_cost_per_token_priority": 7e-07,
"litellm_provider": "openai",
"max_input_tokens": 1047576,
"max_output_tokens": 32768,
@ -11609,6 +11614,7 @@
"mode": "chat",
"output_cost_per_token": 1.6e-06,
"output_cost_per_token_batches": 8e-07,
"output_cost_per_token_priority": 2.8e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -11666,8 +11672,10 @@
},
"gpt-4.1-nano": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_priority": 5e-08,
"input_cost_per_token": 1e-07,
"input_cost_per_token_batches": 5e-08,
"input_cost_per_token_priority": 2e-07,
"litellm_provider": "openai",
"max_input_tokens": 1047576,
"max_output_tokens": 32768,
@ -11675,6 +11683,7 @@
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_batches": 2e-07,
"output_cost_per_token_priority": 8e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -11773,8 +11782,10 @@
},
"gpt-4o": {
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_priority": 2.125e-06,
"input_cost_per_token": 2.5e-06,
"input_cost_per_token_batches": 1.25e-06,
"input_cost_per_token_priority": 4.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
@ -11782,6 +11793,7 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"output_cost_per_token_priority": 1.7e-05,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -11794,6 +11806,7 @@
"gpt-4o-2024-05-13": {
"input_cost_per_token": 5e-06,
"input_cost_per_token_batches": 2.5e-06,
"input_cost_per_token_priority": 8.75e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 4096,
@ -11801,6 +11814,7 @@
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_priority": 2.625e-05,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -11919,8 +11933,10 @@
},
"gpt-4o-mini": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_priority": 1.25e-07,
"input_cost_per_token": 1.5e-07,
"input_cost_per_token_batches": 7.5e-08,
"input_cost_per_token_priority": 2.5e-07,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
@ -11928,6 +11944,7 @@
"mode": "chat",
"output_cost_per_token": 6e-07,
"output_cost_per_token_batches": 3e-07,
"output_cost_per_token_priority": 1e-06,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -12243,13 +12260,19 @@
},
"gpt-5": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_flex": 6.25e-08,
"cache_read_input_token_cost_priority": 2.5e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_flex": 6.25e-07,
"input_cost_per_token_priority": 2.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12275,13 +12298,19 @@
},
"gpt-5-2025-08-07": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_flex": 6.25e-08,
"cache_read_input_token_cost_priority": 2.5e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_flex": 6.25e-07,
"input_cost_per_token_priority": 2.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12371,13 +12400,19 @@
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,
"cache_read_input_token_cost_priority": 4.5e-08,
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_flex": 1.25e-07,
"input_cost_per_token_priority": 4.5e-07,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12403,13 +12438,19 @@
},
"gpt-5-mini-2025-08-07": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,
"cache_read_input_token_cost_priority": 4.5e-08,
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_flex": 1.25e-07,
"input_cost_per_token_priority": 4.5e-07,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12435,13 +12476,16 @@
},
"gpt-5-nano": {
"cache_read_input_token_cost": 5e-09,
"cache_read_input_token_cost_flex": 2.5e-09,
"input_cost_per_token": 5e-08,
"input_cost_per_token_flex": 2.5e-08,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12467,13 +12511,16 @@
},
"gpt-5-nano-2025-08-07": {
"cache_read_input_token_cost": 5e-09,
"cache_read_input_token_cost_flex": 2.5e-09,
"input_cost_per_token": 5e-08,
"input_cost_per_token_flex": 2.5e-08,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -15177,13 +15224,19 @@
},
"o3": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_flex": 2.5e-07,
"cache_read_input_token_cost_priority": 8.75e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_flex": 1e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 200000,
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
"output_cost_per_token": 8e-06,
"output_cost_per_token_flex": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
"supported_endpoints": [
"/v1/responses",
"/v1/chat/completions",
@ -15399,13 +15452,19 @@
},
"o4-mini": {
"cache_read_input_token_cost": 2.75e-07,
"cache_read_input_token_cost_flex": 1.38e-07,
"cache_read_input_token_cost_priority": 5e-07,
"input_cost_per_token": 1.1e-06,
"input_cost_per_token_flex": 5.5e-07,
"input_cost_per_token_priority": 2e-06,
"litellm_provider": "openai",
"max_input_tokens": 200000,
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"output_cost_per_token_flex": 2.2e-06,
"output_cost_per_token_priority": 8e-06,
"supports_function_calling": true,
"supports_parallel_function_calling": false,
"supports_pdf_input": true,
@ -16900,6 +16959,20 @@
"supports_tool_choice": true,
"supports_web_search": true
},
"openrouter/x-ai/grok-4-fast:free": {
"input_cost_per_token": 0,
"litellm_provider": "openrouter",
"max_input_tokens": 2000000,
"max_output_tokens": 30000,
"max_tokens": 2000000,
"mode": "chat",
"output_cost_per_token": 0,
"source": "https://openrouter.ai/x-ai/grok-4-fast:free",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_web_search": false
},
"ovhcloud/DeepSeek-R1-Distill-Llama-70B": {
"input_cost_per_token": 6.7e-07,
"litellm_provider": "ovhcloud",
@ -20943,6 +21016,132 @@
"mode": "embedding",
"output_cost_per_token": 0.0
},
"wandb/openai/gpt-oss-120b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.015,
"output_cost_per_token": 0.06,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/openai/gpt-oss-20b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.005,
"output_cost_per_token": 0.02,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/zai-org/GLM-4.5": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.055,
"output_cost_per_token": 0.2,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/Qwen/Qwen3-235B-A22B-Instruct-2507": {
"max_tokens": 262144,
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"input_cost_per_token": 0.01,
"output_cost_per_token": 0.01,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/Qwen/Qwen3-Coder-480B-A35B-Instruct": {
"max_tokens": 262144,
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"input_cost_per_token": 0.1,
"output_cost_per_token": 0.15,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/Qwen/Qwen3-235B-A22B-Thinking-2507": {
"max_tokens": 262144,
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"input_cost_per_token": 0.01,
"output_cost_per_token": 0.01,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/moonshotai/Kimi-K2-Instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.135,
"output_cost_per_token": 0.4,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/meta-llama/Llama-3.1-8B-Instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.022,
"output_cost_per_token": 0.022,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/deepseek-ai/DeepSeek-V3.1": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.055,
"output_cost_per_token": 0.165,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/deepseek-ai/DeepSeek-R1-0528": {
"max_tokens": 161000,
"max_input_tokens": 161000,
"max_output_tokens": 161000,
"input_cost_per_token": 0.135,
"output_cost_per_token": 0.54,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/deepseek-ai/DeepSeek-V3-0324": {
"max_tokens": 161000,
"max_input_tokens": 161000,
"max_output_tokens": 161000,
"input_cost_per_token": 0.114,
"output_cost_per_token": 0.275,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/meta-llama/Llama-3.3-70B-Instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.071,
"output_cost_per_token": 0.071,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/meta-llama/Llama-4-Scout-17B-16E-Instruct": {
"max_tokens": 64000,
"max_input_tokens": 64000,
"max_output_tokens": 64000,
"input_cost_per_token": 0.017,
"output_cost_per_token": 0.066,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/microsoft/Phi-4-mini-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.008,
"output_cost_per_token": 0.035,
"litellm_provider": "wandb",
"mode": "chat"
},
"watsonx/ibm/granite-3-8b-instruct": {
"input_cost_per_token": 0.0002,
"litellm_provider": "watsonx",
@ -21337,4 +21536,4 @@
"supports_vision": true,
"supports_web_search": true
}
}
}

View file

@ -121,15 +121,13 @@ class MCPServerManager:
for server_name, server_config in mcp_servers_config.items():
validate_mcp_server_name(server_name)
_mcp_info: Dict[str, Any] = server_config.get("mcp_info", None) or {}
# Convert Dict[str, Any] to MCPInfo properly
mcp_info: MCPInfo = {
"server_name": _mcp_info.get("server_name", server_name),
"description": _mcp_info.get(
"description", server_config.get("description", None)
),
"logo_url": _mcp_info.get("logo_url", None),
"mcp_server_cost_info": _mcp_info.get("mcp_server_cost_info", None),
}
# Preserve all custom fields from config while setting defaults for core fields
mcp_info: MCPInfo = _mcp_info.copy()
# Set default values for core fields if not present
if "server_name" not in mcp_info:
mcp_info["server_name"] = server_name
if "description" not in mcp_info and server_config.get("description"):
mcp_info["description"] = server_config.get("description")
# Use alias for name if present, else server_name
alias = server_config.get("alias", None)
@ -243,6 +241,14 @@ class MCPServerManager:
name_for_prefix = (
mcp_server.alias or mcp_server.server_name or mcp_server.server_id
)
# Preserve all custom fields from database while setting defaults for core fields
mcp_info: MCPInfo = _mcp_info.copy()
# Set default values for core fields if not present
if "server_name" not in mcp_info:
mcp_info["server_name"] = mcp_server.server_name or mcp_server.server_id
if "description" not in mcp_info and mcp_server.description:
mcp_info["description"] = mcp_server.description
new_server = MCPServer(
server_id=mcp_server.server_id,
name=name_for_prefix,
@ -251,11 +257,7 @@ class MCPServerManager:
url=mcp_server.url,
transport=cast(MCPTransportType, mcp_server.transport),
auth_type=cast(MCPAuthType, mcp_server.auth_type),
mcp_info=MCPInfo(
server_name=mcp_server.server_name or mcp_server.server_id,
description=mcp_server.description,
mcp_server_cost_info=_mcp_info.get("mcp_server_cost_info", None),
),
mcp_info=mcp_info,
# Stdio-specific fields
command=getattr(mcp_server, "command", None),
args=getattr(mcp_server, "args", None) or [],
@ -419,6 +421,7 @@ class MCPServerManager:
self,
server: MCPServer,
mcp_auth_header: Optional[str] = None,
add_prefix: bool = True,
) -> List[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
@ -443,9 +446,11 @@ class MCPServerManager:
tools = await self._fetch_tools_with_timeout(client, server.name)
prefixed_tools = self._create_prefixed_tools(tools, server)
prefixed_or_original_tools = self._create_prefixed_tools(
tools, server, add_prefix=add_prefix
)
return prefixed_tools
return prefixed_or_original_tools
except Exception as e:
verbose_logger.warning(
@ -514,7 +519,7 @@ class MCPServerManager:
return []
def _create_prefixed_tools(
self, tools: List[MCPTool], server: MCPServer
self, tools: List[MCPTool], server: MCPServer, add_prefix: bool = True
) -> List[MCPTool]:
"""
Create prefixed tools and update tool mapping.
@ -532,14 +537,16 @@ class MCPServerManager:
for tool in tools:
prefixed_name = add_server_prefix_to_tool_name(tool.name, prefix)
prefixed_tool = MCPTool(
name=prefixed_name,
name_to_use = prefixed_name if add_prefix else tool.name
tool_obj = MCPTool(
name=name_to_use,
description=tool.description,
inputSchema=tool.inputSchema,
)
prefixed_tools.append(prefixed_tool)
prefixed_tools.append(tool_obj)
# Update tool to server mapping with both original and prefixed names
# Update tool to server mapping for resolution (support both forms)
self.tool_name_to_mcp_server_name_mapping[tool.name] = prefix
self.tool_name_to_mcp_server_name_mapping[prefixed_name] = prefix

View file

@ -73,6 +73,7 @@ if MCP_AVAILABLE:
tools = await global_mcp_server_manager._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
add_prefix=False,
)
return _create_tool_response_objects(tools, server.mcp_info)

View file

@ -384,6 +384,9 @@ if MCP_AVAILABLE:
allowed_mcp_servers=allowed_mcp_servers,
)
# Decide whether to add prefix based on number of allowed servers
add_prefix = not (len(allowed_mcp_servers) == 1)
# Get tools from each allowed server
all_tools = []
for server_id in allowed_mcp_servers:
@ -406,6 +409,7 @@ if MCP_AVAILABLE:
tools = await global_mcp_server_manager._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
add_prefix=add_prefix,
)
all_tools.extend(tools)
verbose_logger.debug(
@ -637,27 +641,35 @@ if MCP_AVAILABLE:
# Server names can contain slashes (e.g., "custom_solutions/user_123")
mcp_path_match = re.match(r"^/mcp/([^?#]+)(?:\?.*)?(?:#.*)?$", path)
if mcp_path_match:
servers_and_path = mcp_path_match.group(1)
if servers_and_path:
# Check if it contains commas (comma-separated servers)
if ',' in servers_and_path:
# For comma-separated, look for a path at the end
# Common patterns: /tools, /chat/completions, etc.
path_match = re.search(r'/([^/,]+(?:/[^/,]+)*)$', servers_and_path)
if path_match:
# Path found at the end, remove it from servers
path_part = '/' + path_match.group(1)
servers_part = servers_and_path[:-len(path_part)]
mcp_servers_from_path = [s.strip() for s in servers_part.split(',') if s.strip()]
else:
# No path, just comma-separated servers
mcp_servers_from_path = [s.strip() for s in servers_and_path.split(',') if s.strip()]
mcp_servers_str = mcp_path_match.group(1)
optional_path = mcp_path_match.group(2)
if mcp_servers_str:
# First, try to split by comma for comma-separated lists
if "," in mcp_servers_str:
# For comma-separated lists, we need to handle the case where the last item
# might include the path (e.g., "zapier,group1/tools" -> ["zapier", "group1/tools"])
parts = [s.strip() for s in mcp_servers_str.split(",") if s.strip()]
# If there's an optional path AND the last part contains a slash that matches the optional path,
# remove the path portion from the last server name
if optional_path and len(parts) > 0 and "/" in parts[-1]:
last_part = parts[-1]
# Check if the last part ends with the optional path
if optional_path and last_part.endswith(
optional_path.lstrip("/")
):
# Remove the path portion from the last server name
parts[-1] = last_part[: -len(optional_path.lstrip("/"))]
mcp_servers_from_path = parts
else:
# Single server case - use regex approach for server/path separation
# This handles cases like "custom_solutions/user_123/chat/completions"
# where we want to extract "custom_solutions/user_123" as the server name
single_server_match = re.match(r"^([^/]+(?:/[^/]+)?)(?:/.*)?$", servers_and_path)
# For single server, it might be just a name or contain slashes
# We need to determine where the server name ends and the path begins
# This is tricky - let's use the original logic but handle comma cases differently
single_server_match = re.match(
r"^([^/]+(?:/[^/]+)?)(?:/.*)?$", mcp_servers_str
)
if single_server_match:
server_name = single_server_match.group(1)
mcp_servers_from_path = [server_name]

View file

@ -1915,6 +1915,22 @@ class UserAPIKeyAuth(
key_alias=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
team_alias=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
)
@classmethod
def get_litellm_cli_user_api_key_auth(cls) -> "UserAPIKeyAuth":
"""
Returns a `UserAPIKeyAuth` object for the litellm internal health check service account.
This is used to track number of requests/spend for health check calls.
"""
from litellm.constants import LITTELM_CLI_SERVICE_ACCOUNT_NAME
return cls(
api_key=LITTELM_CLI_SERVICE_ACCOUNT_NAME,
team_id=LITTELM_CLI_SERVICE_ACCOUNT_NAME,
key_alias=LITTELM_CLI_SERVICE_ACCOUNT_NAME,
team_alias=LITTELM_CLI_SERVICE_ACCOUNT_NAME,
)
class UserInfoResponse(LiteLLMPydanticObjectBase):

View file

@ -1,11 +1,15 @@
import json
import os
import sys
import time
import webbrowser
from pathlib import Path
from typing import Any, Dict, Optional
from typing import Any, Dict, List, Optional
import click
import requests
from rich.console import Console
from rich.table import Table
# Token storage utilities
@ -44,10 +48,256 @@ def clear_token() -> None:
def get_stored_api_key() -> Optional[str]:
"""Get the stored API key from token file"""
token_data = load_token()
if token_data and 'key' in token_data:
return token_data['key']
return None
# Use the SDK-level utility
from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key
return get_litellm_gateway_api_key()
# Team selection utilities
def display_teams_table(teams: List[Dict[str, Any]]) -> None:
"""Display teams in a formatted table"""
console = Console()
if not teams:
console.print("❌ No teams found for your user.")
return
table = Table(title="Available Teams")
table.add_column("Index", style="cyan", no_wrap=True)
table.add_column("Team Alias", style="magenta")
table.add_column("Team ID", style="green")
table.add_column("Models", style="yellow")
table.add_column("Max Budget", style="blue")
for i, team in enumerate(teams):
team_alias = team.get("team_alias") or "N/A"
team_id = team.get("team_id", "N/A")
models = team.get("models", [])
max_budget = team.get("max_budget")
# Format models list
if models:
if len(models) > 3:
models_str = ", ".join(models[:3]) + f" (+{len(models) - 3} more)"
else:
models_str = ", ".join(models)
else:
models_str = "All models"
# Format budget
budget_str = f"${max_budget}" if max_budget else "Unlimited"
table.add_row(
str(i + 1),
team_alias,
team_id,
models_str,
budget_str
)
console.print(table)
def get_user_teams(base_url: str, api_key: str, user_id: str) -> List[Dict[str, Any]]:
"""Fetch teams for the current user"""
from litellm.proxy.client import Client
client = Client(base_url=base_url, api_key=api_key)
try:
response = client.teams.list_v2(user_id=user_id)
# Extract just the teams array from the paginated response
if isinstance(response, dict) and 'teams' in response:
return response['teams']
else:
# Fallback in case the response structure is different
return response if isinstance(response, list) else []
except Exception as e:
click.echo(f"❌ Error fetching teams: {e}")
return []
def get_key_input():
"""Get a single key input from the user (cross-platform)"""
try:
if sys.platform == 'win32':
import msvcrt
key = msvcrt.getch()
if key == b'\xe0': # Arrow keys on Windows
key = msvcrt.getch()
if key == b'H': # Up arrow
return 'up'
elif key == b'P': # Down arrow
return 'down'
elif key == b'\r': # Enter key
return 'enter'
elif key == b'\x1b': # Escape key
return 'escape'
elif key == b'q':
return 'quit'
return None
else:
import termios
import tty
fd = sys.stdin.fileno()
old_settings = termios.tcgetattr(fd)
try:
tty.setraw(sys.stdin.fileno())
key = sys.stdin.read(1)
if key == '\x1b': # Escape sequence
key += sys.stdin.read(2)
if key == '\x1b[A': # Up arrow
return 'up'
elif key == '\x1b[B': # Down arrow
return 'down'
elif key == '\x1b': # Just escape
return 'escape'
elif key == '\r' or key == '\n': # Enter key
return 'enter'
elif key == 'q':
return 'quit'
return None
finally:
termios.tcsetattr(fd, termios.TCSADRAIN, old_settings)
except ImportError:
# Fallback to simple input if termios/msvcrt not available
return None
def display_interactive_team_selection(teams: List[Dict[str, Any]], selected_index: int = 0) -> None:
"""Display teams with one highlighted for selection"""
console = Console()
# Clear the screen using Rich's method
console.clear()
console.print("🎯 Select a Team (Use ↑↓ arrows, Enter to select, 'q' to skip):\n")
for i, team in enumerate(teams):
team_alias = team.get("team_alias") or "N/A"
team_id = team.get("team_id", "N/A")
models = team.get("models", [])
max_budget = team.get("max_budget")
# Format models list
if models:
if len(models) > 3:
models_str = ", ".join(models[:3]) + f" (+{len(models) - 3} more)"
else:
models_str = ", ".join(models)
else:
models_str = "All models"
# Format budget
budget_str = f"${max_budget}" if max_budget else "Unlimited"
# Highlight the selected item
if i == selected_index:
console.print(f"➤ [bold cyan]{team_alias}[/bold cyan] ({team_id})")
console.print(f" Models: [yellow]{models_str}[/yellow]")
console.print(f" Budget: [blue]{budget_str}[/blue]\n")
else:
console.print(f" [dim]{team_alias}[/dim] ({team_id})")
console.print(f" Models: [dim]{models_str}[/dim]")
console.print(f" Budget: [dim]{budget_str}[/dim]\n")
def prompt_team_selection(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]:
"""Interactive team selection with arrow keys"""
if not teams:
return None
selected_index = 0
try:
# Check if we can use interactive mode
if not sys.stdin.isatty():
# Fallback to simple selection for non-interactive environments
return prompt_team_selection_fallback(teams)
while True:
display_interactive_team_selection(teams, selected_index)
key = get_key_input()
if key == 'up':
selected_index = (selected_index - 1) % len(teams)
elif key == 'down':
selected_index = (selected_index + 1) % len(teams)
elif key == 'enter':
selected_team = teams[selected_index]
# Clear screen and show selection
console = Console()
console.clear()
click.echo(f"✅ Selected team: {selected_team.get('team_alias', 'N/A')} ({selected_team.get('team_id')})")
return selected_team
elif key == 'quit' or key == 'escape':
# Clear screen
console = Console()
console.clear()
click.echo("ℹ️ Team selection skipped.")
return None
elif key is None:
# If we can't get key input, fall back to simple selection
return prompt_team_selection_fallback(teams)
except KeyboardInterrupt:
console = Console()
console.clear()
click.echo("\n❌ Team selection cancelled.")
return None
except Exception:
# If interactive mode fails, fall back to simple selection
return prompt_team_selection_fallback(teams)
def prompt_team_selection_fallback(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]:
"""Fallback team selection for non-interactive environments"""
if not teams:
return None
while True:
try:
choice = click.prompt(
"\nSelect a team by entering the index number (or 'skip' to continue without a team)",
type=str
).strip()
if choice.lower() == 'skip':
return None
index = int(choice) - 1
if 0 <= index < len(teams):
selected_team = teams[index]
click.echo(f"\n✅ Selected team: {selected_team.get('team_alias', 'N/A')} ({selected_team.get('team_id')})")
return selected_team
else:
click.echo(f"❌ Invalid selection. Please enter a number between 1 and {len(teams)}")
except ValueError:
click.echo("❌ Invalid input. Please enter a number or 'skip'")
except KeyboardInterrupt:
click.echo("\n❌ Team selection cancelled.")
return None
def update_key_with_team(base_url: str, api_key: str, team_id: str) -> bool:
"""Update the API key to be associated with the selected team"""
from litellm.proxy.client import Client
client = Client(base_url=base_url, api_key=api_key)
try:
result = client.keys.update(key=api_key, team_id=team_id)
click.echo(f"✅ Successfully assigned key to team: {team_id}")
return True
except requests.exceptions.HTTPError as e:
# Bubble up the response text for detailed error info
error_msg = e.response.text if e.response else str(e)
click.echo(f"❌ Error updating key with team: {error_msg}")
return False
except Exception as e:
click.echo(f"❌ Error updating key with team: {e}")
return False
# Polling-based authentication - no local server needed
@ -57,13 +307,14 @@ def login(ctx: click.Context):
"""Login to LiteLLM proxy using SSO authentication"""
import uuid
import requests
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
from litellm.proxy.client.cli.interface import show_commands
base_url = ctx.obj["base_url"]
# Check if we have an existing key to regenerate
existing_key = get_stored_api_key()
# Generate unique key ID for this login session
key_id = f"sk-{str(uuid.uuid4())}"
@ -71,6 +322,10 @@ def login(ctx: click.Context):
# Construct SSO login URL with CLI source and pre-generated key
sso_url = f"{base_url}/sso/key/generate?source={LITELLM_CLI_SOURCE_IDENTIFIER}&key={key_id}"
# If we have an existing key, include it so the server can regenerate it
if existing_key:
sso_url += f"&existing_key={existing_key}"
click.echo(f"Opening browser to: {sso_url}")
click.echo("Please complete the SSO authentication in your browser...")
click.echo(f"Session ID: {key_id}")
@ -109,6 +364,36 @@ def login(ctx: click.Context):
click.echo(f"API Key: {api_key[:20]}...")
click.echo("You can now use the CLI without specifying --api-key")
# Fetch and display user's teams
click.echo("\n" + "="*60)
click.echo("📋 Fetching your teams...")
teams = get_user_teams(
base_url=base_url,
api_key=api_key,
user_id=data.get("user_id"),
)
if teams:
# Prompt for team selection (will display teams interactively)
selected_team = prompt_team_selection(teams)
if selected_team:
team_id = selected_team.get('team_id')
if team_id:
click.echo(f"\n🔄 Assigning your key to team: {selected_team.get('team_alias', team_id)}")
success = update_key_with_team(base_url, api_key, team_id)
if success:
click.echo(f"✅ Your CLI key is now associated with team: {selected_team.get('team_alias', team_id)}")
click.echo(f"🎯 You can now access models: {', '.join(selected_team.get('models', ['All models']))}")
else:
click.echo("⚠️ Key assignment failed, but you can still use the CLI")
else:
click.echo("ℹ️ Continuing without team assignment. You can assign a team later using the CLI.")
else:
click.echo("ℹ️ No teams found. You can create or join teams using the web interface.")
# Show available commands after successful login
click.echo("\n" + "="*60)
show_commands()
@ -164,5 +449,8 @@ def whoami():
if age_hours > 24:
click.echo("⚠️ Warning: Token is more than 24 hours old and may have expired.")
# Export functions for use by other CLI commands
__all__ = ['login', 'logout', 'whoami', 'prompt_team_selection']
# Export individual commands instead of grouping them
# login, logout, and whoami will be added as top-level commands

View file

@ -0,0 +1,179 @@
"""Team management commands for LiteLLM CLI."""
from typing import Any, Dict, List, Optional
import click
import requests
from rich.console import Console
from rich.table import Table
from litellm.proxy.client import Client
@click.group()
def teams():
"""Manage teams and team assignments"""
pass
def display_teams_table(teams: List[Dict[str, Any]]) -> None:
"""Display teams in a formatted table"""
console = Console()
if not teams:
console.print("❌ No teams found for your user.")
return
table = Table(title="Available Teams")
table.add_column("Index", style="cyan", no_wrap=True)
table.add_column("Team Alias", style="magenta")
table.add_column("Team ID", style="green")
table.add_column("Models", style="yellow")
table.add_column("Max Budget", style="blue")
table.add_column("Role", style="red")
for i, team in enumerate(teams):
team_alias = team.get("team_alias") or "N/A"
team_id = team.get("team_id", "N/A")
models = team.get("models", [])
max_budget = team.get("max_budget")
# Format models list
if models:
if len(models) > 3:
models_str = ", ".join(models[:3]) + f" (+{len(models) - 3} more)"
else:
models_str = ", ".join(models)
else:
models_str = "All models"
# Format budget
budget_str = f"${max_budget}" if max_budget else "Unlimited"
# Try to determine role (this might vary based on API response structure)
role = "Member" # Default role
if isinstance(team, dict) and 'members_with_roles' in team and team['members_with_roles']:
# This would need to be implemented based on actual API response structure
pass
table.add_row(
str(i + 1),
team_alias,
team_id,
models_str,
budget_str,
role
)
console.print(table)
@teams.command()
@click.pass_context
def list(ctx: click.Context):
"""List teams that you belong to"""
client = Client(ctx.obj["base_url"], ctx.obj["api_key"])
try:
# Use list() for simpler response structure (returns array directly)
teams = client.teams.list()
display_teams_table(teams)
except requests.exceptions.HTTPError as e:
click.echo(f"Error: HTTP {e.response.status_code}", err=True)
try:
error_body = e.response.json()
click.echo(f"Details: {error_body.get('detail', 'Unknown error')}", err=True)
except:
click.echo(e.response.text, err=True)
raise click.Abort()
except Exception as e:
click.echo(f"Error: {str(e)}", err=True)
raise click.Abort()
@teams.command()
@click.pass_context
def available(ctx: click.Context):
"""List teams that are available to join"""
client = Client(ctx.obj["base_url"], ctx.obj["api_key"])
try:
teams = client.teams.get_available()
if teams:
console = Console()
console.print("\n🎯 Available Teams to Join:")
display_teams_table(teams)
else:
click.echo("ℹ️ No available teams to join.")
except requests.exceptions.HTTPError as e:
click.echo(f"Error: HTTP {e.response.status_code}", err=True)
try:
error_body = e.response.json()
click.echo(f"Details: {error_body.get('detail', 'Unknown error')}", err=True)
except:
click.echo(e.response.text, err=True)
raise click.Abort()
except Exception as e:
click.echo(f"Error: {str(e)}", err=True)
raise click.Abort()
@teams.command()
@click.option("--team-id", type=str, help="Team ID to assign the key to")
@click.pass_context
def assign_key(ctx: click.Context, team_id: Optional[str]):
"""Assign your current CLI key to a team"""
client = Client(ctx.obj["base_url"], ctx.obj["api_key"])
api_key = ctx.obj["api_key"]
if not api_key:
click.echo("❌ No API key found. Please login first using 'litellm login'")
raise click.Abort()
try:
# If no team_id provided, show teams and let user select
if not team_id:
teams = client.teams.list()
if not teams:
click.echo("❌ No teams found for your user.")
return
# Use interactive selection from auth module
from .auth import prompt_team_selection
selected_team = prompt_team_selection(teams)
if selected_team:
team_id = selected_team.get('team_id')
else:
click.echo("❌ Operation cancelled.")
return
# Update the key with the selected team
if team_id:
click.echo(f"\n🔄 Assigning your key to team: {team_id}")
result = client.keys.update(key=api_key, team_id=team_id)
click.echo(f"✅ Successfully assigned key to team: {team_id}")
# Show team details if available
teams = client.teams.list()
for team in teams:
if team.get('team_id') == team_id:
models = team.get('models', [])
if models:
click.echo(f"🎯 You can now access models: {', '.join(models)}")
else:
click.echo("🎯 You can now access all available models")
break
except requests.exceptions.HTTPError as e:
click.echo(f"Error: HTTP {e.response.status_code}", err=True)
try:
error_body = e.response.json()
click.echo(f"Details: {error_body.get('detail', 'Unknown error')}", err=True)
except:
click.echo(e.response.text, err=True)
raise click.Abort()
except Exception as e:
click.echo(f"Error: {str(e)}", err=True)
raise click.Abort()

View file

@ -87,6 +87,7 @@ def show_commands():
("chat", "Interactive chat with models"),
("http", "Make HTTP requests to the proxy"),
("keys", "Manage API keys"),
("teams", "Manage teams and team assignments"),
("users", "Manage users"),
("version", "Show version information"),
("help", "Show this help message"),

View file

@ -15,6 +15,7 @@ from .commands.keys import keys
# local imports
from .commands.models import models
from .commands.teams import teams
from .commands.users import users
from .interface import interactive_shell
@ -98,6 +99,8 @@ cli.add_command(chat)
cli.add_command(http)
# Add the keys command group
cli.add_command(keys)
# Add the teams command group
cli.add_command(teams)
# Add the users command group
cli.add_command(users)

View file

@ -1,11 +1,12 @@
from typing import Optional
from .http_client import HTTPClient
from .models import ModelsManagementClient
from .model_groups import ModelGroupsManagementClient
from .chat import ChatClient
from .keys import KeysManagementClient
from .credentials import CredentialsManagementClient
from .http_client import HTTPClient
from .keys import KeysManagementClient
from .model_groups import ModelGroupsManagementClient
from .models import ModelsManagementClient
from .teams import TeamsManagementClient
class Client:
@ -36,3 +37,4 @@ class Client:
self.chat = ChatClient(base_url=self._base_url, api_key=self._api_key)
self.keys = KeysManagementClient(base_url=self._base_url, api_key=self._api_key)
self.credentials = CredentialsManagementClient(base_url=self._base_url, api_key=self._api_key)
self.teams = TeamsManagementClient(base_url=self._base_url, api_key=self._api_key)

View file

@ -1,5 +1,7 @@
from typing import Any, Dict, List, Optional, Union
import requests
from typing import Dict, Any, Optional, Union, List
from .exceptions import UnauthorizedError
@ -221,6 +223,66 @@ class KeysManagementClient:
raise UnauthorizedError(e)
raise
def update(
self,
key: str,
models: Optional[List[str]] = None,
aliases: Optional[Dict[str, str]] = None,
spend: Optional[float] = None,
duration: Optional[str] = None,
key_alias: Optional[str] = None,
team_id: Optional[str] = None,
user_id: Optional[str] = None,
) -> Union[Dict[str, Any], requests.Request]:
"""
Update an existing API key's parameters.
Args:
models: Optional[List[str]] = None,
aliases: Optional[Dict[str, str]] = None,
spend: Optional[float] = None,
duration: Optional[str] = None,
key_alias: Optional[str] = None,
team_id: Optional[str] = None,
user_id: Optional[str] = None,
Returns:
Union[Dict[str, Any], requests.Request]: Either the response from the server or
a prepared request object if return_request is True
Raises:
UnauthorizedError: If the request fails with a 401 status code
requests.exceptions.RequestException: If the request fails with any other error
"""
url = f"{self._base_url}/key/update"
data: Dict[str, Any] = {"key": key}
if key_alias is not None:
data["key_alias"] = key_alias
if user_id is not None:
data["user_id"] = user_id
if team_id is not None:
data["team_id"] = team_id
if models is not None:
data["models"] = models
if spend is not None:
data["spend"] = spend
if duration is not None:
data["duration"] = duration
if aliases is not None:
data["aliases"] = aliases
request = requests.Request("POST", url, headers=self._get_headers(), json=data)
session = requests.Session()
try:
response = session.send(request.prepare())
response.raise_for_status()
return response.json()
except requests.exceptions.HTTPError as e:
if e.response.status_code == 401:
raise UnauthorizedError(e)
raise
def info(self, key: str, return_request: bool = False) -> Union[Dict[str, Any], requests.Request]:
"""
Get information about API keys.

View file

@ -0,0 +1,146 @@
"""Teams management client for LiteLLM proxy."""
from typing import Any, Dict, List, Optional
import requests
from .exceptions import UnauthorizedError
class TeamsManagementClient:
"""Client for managing teams in LiteLLM proxy."""
def __init__(self, base_url: str, api_key: Optional[str] = None):
"""
Initialize the TeamsManagementClient.
Args:
base_url (str): The base URL of the LiteLLM proxy server (e.g., "http://localhost:4000")
api_key (Optional[str]): API key for authentication. If provided, it will be sent as a Bearer token.
"""
self._base_url = base_url.rstrip("/") # Remove trailing slash if present
self._api_key = api_key
def _get_headers(self) -> Dict[str, str]:
"""
Get the headers for API requests, including authorization if api_key is set.
Returns:
Dict[str, str]: Headers to use for API requests
"""
headers = {"Content-Type": "application/json"}
if self._api_key:
headers["Authorization"] = f"Bearer {self._api_key}"
return headers
def list(
self,
user_id: Optional[str] = None,
organization_id: Optional[str] = None,
) -> List[Dict[str, Any]]:
"""
List teams that the user belongs to.
Args:
user_id (Optional[str]): Only return teams which this user belongs to
organization_id (Optional[str]): Only return teams which belong to this organization
Returns:
List[Dict[str, Any]]: List of team objects
Raises:
requests.exceptions.HTTPError: If the request fails
UnauthorizedError: If authentication fails
"""
url = f"{self._base_url}/team/list"
params = {}
if user_id:
params["user_id"] = user_id
if organization_id:
params["organization_id"] = organization_id
response = requests.get(url, headers=self._get_headers(), params=params)
if response.status_code == 401:
raise UnauthorizedError("Authentication failed. Check your API key.")
response.raise_for_status()
return response.json()
def list_v2(
self,
user_id: Optional[str] = None,
organization_id: Optional[str] = None,
team_id: Optional[str] = None,
team_alias: Optional[str] = None,
page: int = 1,
page_size: int = 10,
sort_by: Optional[str] = None,
sort_order: str = "asc",
) -> Dict[str, Any]:
"""
Get a paginated list of teams with filtering and sorting options.
Args:
user_id (Optional[str]): Only return teams which this user belongs to
organization_id (Optional[str]): Only return teams which belong to this organization
team_id (Optional[str]): Filter teams by exact team_id match
team_alias (Optional[str]): Filter teams by partial team_alias match
page (int): Page number for pagination
page_size (int): Number of teams per page
sort_by (Optional[str]): Column to sort by (e.g. 'team_id', 'team_alias', 'created_at')
sort_order (str): Sort order ('asc' or 'desc')
Returns:
Dict[str, Any]: Paginated response containing teams and pagination info
Raises:
requests.exceptions.HTTPError: If the request fails
UnauthorizedError: If authentication fails
"""
url = f"{self._base_url}/v2/team/list"
params = {
"page": page,
"page_size": page_size,
"sort_order": sort_order,
}
if user_id:
params["user_id"] = user_id
if organization_id:
params["organization_id"] = organization_id
if team_id:
params["team_id"] = team_id
if team_alias:
params["team_alias"] = team_alias
if sort_by:
params["sort_by"] = sort_by
response = requests.get(url, headers=self._get_headers(), params=params)
if response.status_code == 401:
raise UnauthorizedError("Authentication failed. Check your API key.")
response.raise_for_status()
return response.json()
def get_available(self) -> List[Dict[str, Any]]:
"""
Get list of available teams that the user can join.
Returns:
List[Dict[str, Any]]: List of available team objects
Raises:
requests.exceptions.HTTPError: If the request fails
UnauthorizedError: If authentication fails
"""
url = f"{self._base_url}/team/available"
response = requests.get(url, headers=self._get_headers())
if response.status_code == 401:
raise UnauthorizedError("Authentication failed. Check your API key.")
response.raise_for_status()
return response.json()

View file

@ -233,14 +233,16 @@ async def get_request_body(request: Request) -> Dict[str, Any]:
"""
Read the request body and parse it as JSON.
"""
if request.headers.get("content-type") == "application/json":
return await _read_request_body(request)
elif (
request.headers.get("content-type") == "multipart/form-data"
or request.headers.get("content-type") == "application/x-www-form-urlencoded"
):
return await get_form_data(request)
else:
raise ValueError(
f"Unsupported content type: {request.headers.get('content-type')}"
)
if request.method == "POST":
if request.headers.get("content-type", "") == "application/json":
return await _read_request_body(request)
elif (
"multipart/form-data" in request.headers.get("content-type", "")
or "application/x-www-form-urlencoded" in request.headers.get("content-type", "")
):
return await get_form_data(request)
else:
raise ValueError(
f"Unsupported content type: {request.headers.get('content-type')}"
)
return {}

View file

@ -88,11 +88,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
def _create_sanitize_request(
self, content: str, source: Literal["user_prompt", "model_response"]
) -> dict:
"""Create request body for Model Armor API with correct camelCase field names."""
"""Create request body for Model Armor API."""
if source == "user_prompt":
return {"userPromptData": {"text": content}}
return {"user_prompt_data": {"text": content}}
else:
return {"modelResponseData": {"text": content}}
return {"model_response_data": {"text": content}}
def _extract_content_from_response(
self, response: Union[Any, ModelResponse]
@ -119,16 +119,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
async def make_model_armor_request(
self,
content: Optional[str] = None,
source: Literal["user_prompt", "model_response"] = "user_prompt",
content: str,
source: Literal["user_prompt", "model_response"],
request_data: Optional[dict] = None,
file_bytes: Optional[bytes] = None,
file_type: Optional[str] = None,
) -> dict:
"""
Make request to Model Armor API. Supports both text and file prompt sanitization.
If file_bytes and file_type are provided, file prompt sanitization is performed.
"""
"""Make request to Model Armor API."""
# Get access token using VertexBase auth
access_token, resolved_project_id = await self._ensure_access_token_async(
credentials=self.credentials,
@ -148,14 +143,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
url = f"{endpoint}/v1/projects/{self.project_id}/locations/{self.location}/templates/{self.template_id}:sanitizeModelResponse"
# Create request body
if file_bytes is not None and file_type is not None:
body = self.sanitize_file_prompt(file_bytes, file_type, source)
elif content is not None:
body = self._create_sanitize_request(content, source)
else:
raise ValueError(
"Either content or file_bytes and file_type must be provided."
)
body = self._create_sanitize_request(content, source)
# Set headers
headers = {
@ -201,110 +189,57 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
return await json_response
return json_response
def sanitize_file_prompt(
self, file_bytes: bytes, file_type: str, source: str = "user_prompt"
) -> dict:
"""
Helper to build the request body for file prompt sanitization for Model Armor.
file_type should be one of: PLAINTEXT_UTF8, PDF, WORD_DOCUMENT, EXCEL_DOCUMENT, POWERPOINT_DOCUMENT, TXT, CSV
Returns the request body dict.
"""
import base64
base64_data = base64.b64encode(file_bytes).decode("utf-8")
if source == "user_prompt":
return {
"userPromptData": {
"byteItem": {"byteDataType": file_type, "byteData": base64_data}
}
}
else:
return {
"modelResponseData": {
"byteItem": {"byteDataType": file_type, "byteData": base64_data}
}
}
def _should_block_content(self, armor_response: dict) -> bool:
"""Check if Model Armor response indicates content should be blocked, including both inspectResult and deidentifyResult."""
"""Check if Model Armor response indicates content should be blocked."""
# Check the sanitizationResult from Model Armor API
sanitization_result = armor_response.get("sanitizationResult", {})
filter_results = sanitization_result.get("filterResults", {})
# filterResults can be a dict (named keys) or a list (array of filter result dicts)
filter_result_items = []
if isinstance(filter_results, dict):
filter_result_items = [filter_results]
elif isinstance(filter_results, list):
filter_result_items = filter_results
# Check blocking filters (these should cause the request to be blocked)
# RAI (Responsible AI) filters
rai_results = filter_results.get("rai", {}).get("raiFilterResult", {})
if rai_results.get("matchState") == "MATCH_FOUND":
return True
# Prompt injection and jailbreak filters
pi_jailbreak = filter_results.get("piAndJailbreakFilterResult", {})
if pi_jailbreak.get("matchState") == "MATCH_FOUND":
return True
# Malicious URI filters
malicious_uri = filter_results.get("maliciousUriFilterResult", {})
if malicious_uri.get("matchState") == "MATCH_FOUND":
return True
# CSAM filters
csam = filter_results.get("csamFilterFilterResult", {})
if csam.get("matchState") == "MATCH_FOUND":
return True
# Virus scan filters
virus_scan = filter_results.get("virusScanFilterResult", {})
if virus_scan.get("matchState") == "MATCH_FOUND":
return True
for filt in filter_result_items:
# Check RAI, PI/Jailbreak, Malicious URI, CSAM, Virus scan as before
if filt.get("raiFilterResult", {}).get("matchState") == "MATCH_FOUND":
return True
if (
filt.get("piAndJailbreakFilterResult", {}).get("matchState")
== "MATCH_FOUND"
):
return True
if (
filt.get("maliciousUriFilterResult", {}).get("matchState")
== "MATCH_FOUND"
):
return True
if (
filt.get("csamFilterFilterResult", {}).get("matchState")
== "MATCH_FOUND"
):
return True
if filt.get("virusScanFilterResult", {}).get("matchState") == "MATCH_FOUND":
return True
# Check sdpFilterResult for both inspectResult and deidentifyResult
sdp = filt.get("sdpFilterResult")
if sdp:
if sdp.get("inspectResult", {}).get("matchState") == "MATCH_FOUND":
return True
if sdp.get("deidentifyResult", {}).get("matchState") == "MATCH_FOUND":
return True
# Fallback dict code removed; all cases handled above
return False
def _get_sanitized_content(self, armor_response: dict) -> Optional[str]:
"""
Get the sanitized content from a Model Armor response, if available.
Looks for sanitized text in deidentifyResult, and falls back to root-level fields if not found.
"""
result = armor_response.get("sanitizationResult", {})
filter_results = result.get("filterResults", {})
"""Extract sanitized content from Model Armor response."""
# Model Armor returns sanitized content in the sanitizationResult
sanitization_result = armor_response.get("sanitizationResult", {})
# filterResults can be a dict (single filter) or a list (multiple filters)
filters = (
[filter_results]
if isinstance(filter_results, dict)
else filter_results
if isinstance(filter_results, list)
else []
)
# Check for sdp structure (for deidentification)
filter_results = sanitization_result.get("filterResults", {})
sdp = filter_results.get("sdp", {}).get("sdpFilterResult")
# Prefer sanitized text from deidentifyResult if present
for filter_entry in filters:
sdp = filter_entry.get("sdpFilterResult")
if sdp:
deid = sdp.get("deidentifyResult", {})
sanitized = deid.get("data", {}).get("text", "")
# If Model Armor found something and returned a sanitized version, use it
if deid.get("matchState") == "MATCH_FOUND" and sanitized:
return sanitized
if sdp is not None:
# Model Armor returns sanitized text under deidentifyResult in sdp
deidentify_result = sdp.get("deidentifyResult", {})
sanitized_text = deidentify_result.get("data", {}).get("text", "")
if deidentify_result.get("matchState") == "MATCH_FOUND" and sanitized_text:
return sanitized_text
# If no deidentifyResult, optionally check for inspectResult (rare, but could have findings)
for filter_entry in filters:
sdp = filter_entry.get("sdpFilterResult")
if sdp:
inspect = sdp.get("inspectResult", {})
# If Model Armor flagged something but didn't sanitize, return None
if inspect.get("matchState") == "MATCH_FOUND":
return None
# Fallback: if Model Armor put sanitized text at the root, use it
# Fallback to checking root level
return armor_response.get("sanitizedText") or armor_response.get("text")
def _process_response(

View file

@ -67,6 +67,7 @@ from litellm.proxy.utils import (
)
from litellm.secret_managers.main import get_secret_bool, str_to_bool
from litellm.types.proxy.management_endpoints.ui_sso import *
from litellm.types.proxy.ui_sso import ParsedOpenIDResult
if TYPE_CHECKING:
from fastapi_sso.sso.base import OpenID
@ -114,7 +115,7 @@ def process_sso_jwt_access_token(
@router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False)
async def google_login(
request: Request, source: Optional[str] = None, key: Optional[str] = None
request: Request, source: Optional[str] = None, key: Optional[str] = None, existing_key: Optional[str] = None
): # noqa: PLR0915
"""
Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env
@ -173,12 +174,14 @@ async def google_login(
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=request,
sso_callback_route="sso/callback",
existing_key=existing_key,
)
# Store CLI key in state for OAuth flow
cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state(
source=source,
key=key,
existing_key=existing_key,
)
# check if user defined a custom auth sso sign in handler, if yes, use it
@ -586,13 +589,6 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
# Check if this is a CLI login (state starts with our CLI prefix)
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
# Extract the key ID from the state
key_id = state.split(":", 1)[1]
verbose_proxy_logger.info(f"CLI SSO callback detected for key: {key_id}")
return await cli_sso_callback(request, key=key_id)
from litellm.proxy._types import LiteLLM_JWTAuth
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.proxy_server import (
@ -668,6 +664,17 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
status_code=401,
detail="Result not returned by SSO provider.",
)
if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
# Extract the key ID from the state
key_id = state.split(":", 1)[1]
# Get existing_key from query parameters if provided
existing_key = request.query_params.get("existing_key")
verbose_proxy_logger.info(f"CLI SSO callback detected for key: {key_id}, existing_key: {existing_key}")
return await cli_sso_callback(request=request, key=key_id, existing_key=existing_key, result=result)
return await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=result,
@ -678,13 +685,64 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
)
async def cli_sso_callback(request: Request, key: Optional[str] = None):
"""CLI SSO callback - generates the key with pre-specified ID"""
verbose_proxy_logger.info(f"CLI SSO callback for key: {key}")
async def _regenerate_cli_key(existing_key: str, new_key: str, user_id: Optional[str] = None) -> None:
"""Regenerate an existing CLI key with a new token"""
from litellm.proxy._types import RegenerateKeyRequest, UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
regenerate_key_fn,
)
verbose_proxy_logger.info(f"Regenerating existing CLI key: {existing_key}")
admin_user_dict = UserAPIKeyAuth.get_litellm_cli_user_api_key_auth()
regenerate_request = RegenerateKeyRequest(
key=existing_key,
new_key=new_key,
duration="24hr",
user_id=user_id,
)
await regenerate_key_fn(
key=existing_key,
data=regenerate_request,
user_api_key_dict=admin_user_dict
)
verbose_proxy_logger.info(f"Regenerated CLI key: {new_key}")
async def _create_new_cli_key(
key: str,
user_id: Optional[str] = None,
) -> None:
"""Create a new CLI key"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_helper_fn,
)
verbose_proxy_logger.info("Creating new CLI key")
await generate_key_helper_fn(
request_type="key",
duration="24hr",
key_max_budget=litellm.max_ui_session_budget,
aliases={},
config={},
spend=0,
user_id=user_id,
team_id="litellm-cli",
table_name="key",
token=key,
)
verbose_proxy_logger.info(f"Created new CLI key: {key}")
async def cli_sso_callback(request: Request, key: Optional[str] = None, existing_key: Optional[str] = None, result: Optional[Union[OpenID, dict]] = None):
"""CLI SSO callback - regenerates existing CLI key or creates new one"""
verbose_proxy_logger.info(f"CLI SSO callback for key: {key}, existing_key: {existing_key}")
from litellm.proxy.proxy_server import prisma_client
if not key or not key.startswith("sk-"):
@ -697,22 +755,22 @@ async def cli_sso_callback(request: Request, key: Optional[str] = None):
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result(result=result)
verbose_proxy_logger.debug(f"parsed_openid_result: {parsed_openid_result}")
# Generate a simple key for CLI usage with the pre-specified key ID
try:
await generate_key_helper_fn(
request_type="key",
duration="24hr",
key_max_budget=litellm.max_ui_session_budget,
aliases={},
config={},
spend=0,
team_id="litellm-cli",
table_name="key",
token=key, # Use the pre-specified key ID
)
verbose_proxy_logger.info(f"Generated CLI key: {key}")
if existing_key:
await _regenerate_cli_key(
existing_key=existing_key,
new_key=key,
user_id=parsed_openid_result.get("user_id"),
)
else:
await _create_new_cli_key(
key=key,
user_id=parsed_openid_result.get("user_id"),
)
# Return success page
from fastapi.responses import HTMLResponse
@ -725,13 +783,14 @@ async def cli_sso_callback(request: Request, key: Optional[str] = None):
return HTMLResponse(content=html_content, status_code=200)
except Exception as e:
verbose_proxy_logger.error(f"Error generating CLI key: {e}")
raise HTTPException(status_code=500, detail=f"Failed to generate key: {str(e)}")
verbose_proxy_logger.error(f"Error with CLI key: {e}")
raise HTTPException(status_code=500, detail=f"Failed to process CLI key: {str(e)}")
@router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False)
async def cli_poll_key(key_id: str):
"""CLI polling endpoint - checks if key exists in DB"""
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.proxy_server import prisma_client
if not key_id.startswith("sk-"):
@ -751,10 +810,11 @@ async def cli_poll_key(key_id: str):
key_obj = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
)
key_obj: LiteLLM_VerificationToken = cast(LiteLLM_VerificationToken, key_obj)
if key_obj:
verbose_proxy_logger.info(f"CLI key found: {key_id}")
return {"status": "ready", "key": key_id}
return {"status": "ready", "key": key_id, "user_id": key_obj.user_id}
else:
return {"status": "pending"}
@ -993,20 +1053,51 @@ class SSOAuthenticationHandler:
# or a cryptographicly signed state that we can verify stateless
# For simplification we are using a static state, this is not perfect but some
# SSO providers do not allow stateless verification
redirect_params = {}
state = os.getenv("GENERIC_CLIENT_STATE", None)
if state:
redirect_params["state"] = state
elif "okta" in generic_authorization_endpoint:
redirect_params["state"] = (
uuid.uuid4().hex
) # set state param for okta - required
redirect_params = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=state,
generic_authorization_endpoint=generic_authorization_endpoint
)
return await generic_sso.get_login_redirect(**redirect_params) # type: ignore
raise ValueError(
"Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso"
)
@staticmethod
def _get_generic_sso_redirect_params(
state: Optional[str] = None,
generic_authorization_endpoint: Optional[str] = None
) -> dict:
"""
Get redirect parameters for Generic SSO with proper state priority handling.
Priority order:
1. CLI state (if provided)
2. GENERIC_CLIENT_STATE environment variable
3. Generated UUID for Okta (if Okta endpoint detected)
Args:
state: Optional state parameter (e.g., CLI state)
generic_authorization_endpoint: Authorization endpoint URL
Returns:
dict: Redirect parameters for SSO login
"""
redirect_params = {}
if state:
# CLI state takes priority
# the litellm proxy cli sends the "state" parameter to the proxy server for auth. We should maintain the state parameter for the cli if it is provided
redirect_params["state"] = state
else:
generic_client_state = os.getenv("GENERIC_CLIENT_STATE", None)
if generic_client_state:
redirect_params["state"] = generic_client_state
elif generic_authorization_endpoint and "okta" in generic_authorization_endpoint:
redirect_params["state"] = uuid.uuid4().hex # set state param for okta - required
return redirect_params
@staticmethod
def should_use_sso_handler(
google_client_id: Optional[str] = None,
@ -1025,6 +1116,7 @@ class SSOAuthenticationHandler:
def get_redirect_url_for_sso(
request: Request,
sso_callback_route: str,
existing_key: Optional[str] = None,
) -> str:
"""
Get the redirect URL for SSO
@ -1036,6 +1128,11 @@ class SSOAuthenticationHandler:
redirect_url += sso_callback_route
else:
redirect_url += "/" + sso_callback_route
# Append existing_key as query parameter if provided
if existing_key:
redirect_url += f"?existing_key={existing_key}"
return redirect_url
@staticmethod
@ -1218,7 +1315,7 @@ class SSOAuthenticationHandler:
return team_request
@staticmethod
def _get_cli_state(source: Optional[str], key: Optional[str]) -> Optional[str]:
def _get_cli_state(source: Optional[str], key: Optional[str], existing_key: Optional[str] = None) -> Optional[str]:
"""
Checks the request 'source' if a cli state token was passed in
@ -1229,45 +1326,25 @@ class SSOAuthenticationHandler:
LITELLM_CLI_SOURCE_IDENTIFIER,
)
return (
f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
if source == LITELLM_CLI_SOURCE_IDENTIFIER and key
else None
)
if source == LITELLM_CLI_SOURCE_IDENTIFIER and key:
# Just use the key - existing_key will be passed separately via query params
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
else:
return None
@staticmethod
async def get_redirect_response_from_openid( # noqa: PLR0915
result: Union[OpenID, dict, CustomOpenID],
request: Request,
received_response: Optional[dict] = None,
def _get_user_email_and_id_from_result(
result: Optional[Union[OpenID, dict]],
generic_client_id: Optional[str] = None,
ui_access_mode: Optional[Dict] = None,
) -> RedirectResponse:
import jwt
from litellm.proxy.proxy_server import (
general_settings,
generate_key_helper_fn,
master_key,
premium_user,
proxy_logging_obj,
user_api_key_cache,
user_custom_sso,
)
from litellm.proxy.utils import get_prisma_client_or_throw
from litellm.types.proxy.ui_sso import ReturnedUITokenObject
prisma_client = get_prisma_client_or_throw(
"Prisma client is None, connect a database to your proxy"
)
# User is Authe'd in - generate key for the UI to access Proxy
verbose_proxy_logger.info(f"SSO callback result: {result}")
) -> ParsedOpenIDResult:
"""
Gets the user email and id from the OpenID result after validating the email domain
"""
user_email: Optional[str] = getattr(result, "email", None)
user_id: Optional[str] = (
getattr(result, "id", None) if result is not None else None
)
user_role: Optional[str] = None
if user_email is not None and os.getenv("ALLOWED_EMAIL_DOMAINS") is not None:
email_domain = user_email.split("@")[1]
@ -1298,6 +1375,46 @@ class SSOAuthenticationHandler:
if user_email is not None and (user_id is None or len(user_id) == 0):
user_id = user_email
return ParsedOpenIDResult(
user_email=user_email,
user_id=user_id,
user_role=user_role,
)
@staticmethod
async def get_redirect_response_from_openid( # noqa: PLR0915
result: Union[OpenID, dict, CustomOpenID],
request: Request,
received_response: Optional[dict] = None,
generic_client_id: Optional[str] = None,
ui_access_mode: Optional[Dict] = None,
) -> RedirectResponse:
import jwt
from litellm.proxy.proxy_server import (
general_settings,
generate_key_helper_fn,
master_key,
premium_user,
proxy_logging_obj,
user_api_key_cache,
user_custom_sso,
)
from litellm.proxy.utils import get_prisma_client_or_throw
from litellm.types.proxy.ui_sso import ReturnedUITokenObject
prisma_client = get_prisma_client_or_throw(
"Prisma client is None, connect a database to your proxy"
)
# User is Authe'd in - generate key for the UI to access Proxy
parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result(result=result, generic_client_id=generic_client_id)
user_email = parsed_openid_result.get("user_email")
user_id = parsed_openid_result.get("user_id")
user_role = parsed_openid_result.get("user_role")
verbose_proxy_logger.info(f"SSO callback result: {result}")
user_info = None
user_id_models: List = []

View file

@ -108,7 +108,16 @@ async def llm_passthrough_factory_proxy_route(
# Construct the full target URL using httpx
base_url = httpx.URL(base_target_url)
updated_url = base_url.copy_with(path=encoded_endpoint)
# Join paths correctly by removing trailing/leading slashes as needed
if not base_url.path or base_url.path == "/":
# If base URL has no path, just use the new path
updated_url = base_url.copy_with(path=encoded_endpoint)
else:
# Otherwise, combine the paths
base_path = base_url.path.rstrip("/")
clean_path = encoded_endpoint.lstrip("/")
full_path = f"{base_path}/{clean_path}"
updated_url = base_url.copy_with(path=full_path)
# Add or update query parameters
provider_api_key = passthrough_endpoint_router.get_credentials(
@ -130,7 +139,11 @@ async def llm_passthrough_factory_proxy_route(
is_streaming_request = False
# anthropic is streaming when 'stream' = True is in the body
if request.method == "POST":
_request_body = await request.json()
if "multipart/form-data" not in request.headers.get("content-type", ""):
_request_body = await request.json()
else:
_request_body = await get_form_data(request)
if _request_body.get("stream"):
is_streaming_request = True

View file

@ -11,3 +11,15 @@ model_list:
aws_batch_role_arn: arn:aws:iam::888602223428:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV
model_info:
mode: batch
- model_name: anthropic/*
litellm_params:
model: anthropic/*
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: openai/*
litellm_params:
model: openai/*
api_key: os.environ/OPENAI_API_KEY
- model_name: gemini/*
litellm_params:
model: gemini/*
api_key: os.environ/GEMINI_API_KEY

View file

@ -1,4 +1,4 @@
from typing import Dict, List, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, ConfigDict
from typing_extensions import TypedDict
@ -7,11 +7,8 @@ from litellm.proxy._types import MCPAuthType, MCPTransportType
from litellm.types.mcp import MCPServerCostInfo
class MCPInfo(TypedDict, total=False):
server_name: str
description: Optional[str]
logo_url: Optional[str]
mcp_server_cost_info: Optional[MCPServerCostInfo]
# MCPInfo now allows arbitrary additional fields for custom metadata
MCPInfo = Dict[str, Any]
class MCPServer(BaseModel):

View file

@ -17,3 +17,12 @@ class ReturnedUITokenObject(TypedDict):
auth_header_name: str
disabled_non_admin_personal_key_creation: bool
server_root_path: str # e.g. `/litellm`
class ParsedOpenIDResult(TypedDict, total=False):
"""
Parsed OpenID result
"""
user_email: Optional[str]
user_id: Optional[str]
user_role: Optional[str]

View file

@ -122,9 +122,13 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
max_input_tokens: Required[Optional[int]]
max_output_tokens: Required[Optional[int]]
input_cost_per_token: Required[float]
input_cost_per_token_flex: Optional[float] # OpenAI flex service tier pricing
input_cost_per_token_priority: Optional[float] # OpenAI priority service tier pricing
cache_creation_input_token_cost: Optional[float]
cache_creation_input_token_cost_above_1hr: Optional[float]
cache_read_input_token_cost: Optional[float]
cache_read_input_token_cost_flex: Optional[float] # OpenAI flex service tier pricing
cache_read_input_token_cost_priority: Optional[float] # OpenAI priority service tier pricing
input_cost_per_character: Optional[float] # only for vertex ai models
input_cost_per_audio_token: Optional[float]
input_cost_per_token_above_128k_tokens: Optional[float] # only for vertex ai models
@ -142,6 +146,8 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
input_cost_per_token_batches: Optional[float]
output_cost_per_token_batches: Optional[float]
output_cost_per_token: Required[float]
output_cost_per_token_flex: Optional[float] # OpenAI flex service tier pricing
output_cost_per_token_priority: Optional[float] # OpenAI priority service tier pricing
output_cost_per_character: Optional[float] # only for vertex ai models
output_cost_per_audio_token: Optional[float]
output_cost_per_token_above_128k_tokens: Optional[
@ -2397,6 +2403,7 @@ class LlmProviders(str, Enum):
AUTO_ROUTER = "auto_router"
VERCEL_AI_GATEWAY = "vercel_ai_gateway"
DOTPROMPT = "dotprompt"
WANDB = "wandb"
OVHCLOUD = "ovhcloud"
@ -2583,6 +2590,12 @@ class SpecialEnums(Enum):
LITELLM_MANAGED_GENERIC_RESPONSE_COMPLETE_STR = "litellm_proxy;model_id:{};generic_response_id:{}" # generic implementation of 'managed batches' - used for finetuning and any future work.
class ServiceTier(Enum):
"""Enum for service tier types used in cost calculations."""
FLEX = "flex"
PRIORITY = "priority"
LLMResponseTypes = Union[
ModelResponse,
EmbeddingResponse,

View file

@ -3275,6 +3275,7 @@ def pre_process_optional_params(
and custom_llm_provider != "openrouter"
and custom_llm_provider != "vercel_ai_gateway"
and custom_llm_provider != "nebius"
and custom_llm_provider != "wandb"
and custom_llm_provider not in litellm.openai_compatible_providers
):
if custom_llm_provider == "ollama":
@ -4446,6 +4447,9 @@ def get_api_key(llm_provider: str, dynamic_api_key: Optional[str]):
# nebius
elif llm_provider == "nebius":
api_key = api_key or litellm.nebius_key or get_secret("NEBIUS_API_KEY")
# wandb
elif llm_provider == "wandb":
api_key = api_key or litellm.wandb_key or get_secret("WANDB_API_KEY")
return api_key
@ -4874,12 +4878,16 @@ def _get_model_info_helper( # noqa: PLR0915
max_input_tokens=_model_info.get("max_input_tokens", None),
max_output_tokens=_model_info.get("max_output_tokens", None),
input_cost_per_token=_input_cost_per_token,
input_cost_per_token_flex=_model_info.get("input_cost_per_token_flex", None),
input_cost_per_token_priority=_model_info.get("input_cost_per_token_priority", None),
cache_creation_input_token_cost=_model_info.get(
"cache_creation_input_token_cost", None
),
cache_read_input_token_cost=_model_info.get(
"cache_read_input_token_cost", None
),
cache_read_input_token_cost_flex=_model_info.get("cache_read_input_token_cost_flex", None),
cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None),
cache_creation_input_token_cost_above_1hr=_model_info.get(
"cache_creation_input_token_cost_above_1hr", None
),
@ -4904,6 +4912,8 @@ def _get_model_info_helper( # noqa: PLR0915
"output_cost_per_token_batches"
),
output_cost_per_token=_output_cost_per_token,
output_cost_per_token_flex=_model_info.get("output_cost_per_token_flex", None),
output_cost_per_token_priority=_model_info.get("output_cost_per_token_priority", None),
output_cost_per_audio_token=_model_info.get(
"output_cost_per_audio_token", None
),
@ -5530,6 +5540,11 @@ def validate_environment( # noqa: PLR0915
keys_in_environment = True
else:
missing_keys.append("NEBIUS_API_KEY")
elif custom_llm_provider == "wandb":
if "WANDB_API_KEY" in os.environ:
keys_in_environment = True
else:
missing_keys.append("WANDB_API_KEY")
elif custom_llm_provider == "dashscope":
if "DASHSCOPE_API_KEY" in os.environ:
keys_in_environment = True
@ -5644,6 +5659,11 @@ def validate_environment( # noqa: PLR0915
keys_in_environment = True
else:
missing_keys.append("NEBIUS_API_KEY")
elif model in litellm.wandb_models:
if "WANDB_API_KEY" in os.environ:
keys_in_environment = True
else:
missing_keys.append("WANDB_API_KEY")
def filter_missing_keys(keys: List[str], exclude_pattern: str) -> List[str]:
"""Filter out keys that contain the exclude_pattern (case insensitive)."""
@ -6408,6 +6428,8 @@ def get_valid_models(
check_provider_endpoint: Optional[bool] = None,
custom_llm_provider: Optional[str] = None,
litellm_params: Optional[LiteLLM_Params] = None,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> List[str]:
"""
Returns a list of valid LLMs based on the set environment variables
@ -6415,11 +6437,24 @@ def get_valid_models(
Args:
check_provider_endpoint: If True, will check the provider's endpoint for valid models.
custom_llm_provider: If provided, will only check the provider's endpoint for valid models.
api_key: If provided, will use the API key to get valid models.
api_base: If provided, will use the API base to get valid models.
Returns:
A list of valid LLMs
"""
try:
################################
# init litellm_params
#################################
if litellm_params is None:
litellm_params = LiteLLM_Params(model="")
if api_key is not None:
litellm_params.api_key = api_key
if api_base is not None:
litellm_params.api_base = api_base
#################################
check_provider_endpoint = (
check_provider_endpoint or litellm.check_provider_endpoint
)
@ -7046,6 +7081,8 @@ class ProviderConfigManager:
return litellm.NovitaConfig()
elif litellm.LlmProviders.NEBIUS == provider:
return litellm.NebiusConfig()
elif litellm.LlmProviders.WANDB == provider:
return litellm.WandbConfig()
elif litellm.LlmProviders.DASHSCOPE == provider:
return litellm.DashScopeChatConfig()
elif litellm.LlmProviders.MOONSHOT == provider:
@ -7501,6 +7538,12 @@ class ProviderConfigManager:
)
return RecraftImageEditConfig()
elif LlmProviders.AZURE_AI == provider:
from litellm.llms.azure_ai.image_edit import (
get_azure_ai_image_edit_config,
)
return get_azure_ai_image_edit_config(model)
elif LlmProviders.LITELLM_PROXY == provider:
from litellm.llms.litellm_proxy.image_edit.transformation import (
LiteLLMProxyImageEditConfig,

View file

@ -11534,8 +11534,10 @@
},
"gpt-4.1": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_priority": 8.75e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_batches": 1e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 1047576,
"max_output_tokens": 32768,
@ -11543,6 +11545,7 @@
"mode": "chat",
"output_cost_per_token": 8e-06,
"output_cost_per_token_batches": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -11600,8 +11603,10 @@
},
"gpt-4.1-mini": {
"cache_read_input_token_cost": 1e-07,
"cache_read_input_token_cost_priority": 1.75e-07,
"input_cost_per_token": 4e-07,
"input_cost_per_token_batches": 2e-07,
"input_cost_per_token_priority": 7e-07,
"litellm_provider": "openai",
"max_input_tokens": 1047576,
"max_output_tokens": 32768,
@ -11609,6 +11614,7 @@
"mode": "chat",
"output_cost_per_token": 1.6e-06,
"output_cost_per_token_batches": 8e-07,
"output_cost_per_token_priority": 2.8e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -11666,8 +11672,10 @@
},
"gpt-4.1-nano": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_priority": 5e-08,
"input_cost_per_token": 1e-07,
"input_cost_per_token_batches": 5e-08,
"input_cost_per_token_priority": 2e-07,
"litellm_provider": "openai",
"max_input_tokens": 1047576,
"max_output_tokens": 32768,
@ -11675,6 +11683,7 @@
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_batches": 2e-07,
"output_cost_per_token_priority": 8e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -11773,8 +11782,10 @@
},
"gpt-4o": {
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_priority": 2.125e-06,
"input_cost_per_token": 2.5e-06,
"input_cost_per_token_batches": 1.25e-06,
"input_cost_per_token_priority": 4.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
@ -11782,6 +11793,7 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"output_cost_per_token_priority": 1.7e-05,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -11794,6 +11806,7 @@
"gpt-4o-2024-05-13": {
"input_cost_per_token": 5e-06,
"input_cost_per_token_batches": 2.5e-06,
"input_cost_per_token_priority": 8.75e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 4096,
@ -11801,6 +11814,7 @@
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"output_cost_per_token_priority": 2.625e-05,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -11919,8 +11933,10 @@
},
"gpt-4o-mini": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_priority": 1.25e-07,
"input_cost_per_token": 1.5e-07,
"input_cost_per_token_batches": 7.5e-08,
"input_cost_per_token_priority": 2.5e-07,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 16384,
@ -11928,6 +11944,7 @@
"mode": "chat",
"output_cost_per_token": 6e-07,
"output_cost_per_token_batches": 3e-07,
"output_cost_per_token_priority": 1e-06,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -12243,13 +12260,19 @@
},
"gpt-5": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_flex": 6.25e-08,
"cache_read_input_token_cost_priority": 2.5e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_flex": 6.25e-07,
"input_cost_per_token_priority": 2.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12275,13 +12298,19 @@
},
"gpt-5-2025-08-07": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_flex": 6.25e-08,
"cache_read_input_token_cost_priority": 2.5e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_flex": 6.25e-07,
"input_cost_per_token_priority": 2.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12303,6 +12332,7 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true
},
"gpt-5-chat": {
@ -12371,13 +12401,19 @@
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,
"cache_read_input_token_cost_priority": 4.5e-08,
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_flex": 1.25e-07,
"input_cost_per_token_priority": 4.5e-07,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12403,13 +12439,19 @@
},
"gpt-5-mini-2025-08-07": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,
"cache_read_input_token_cost_priority": 4.5e-08,
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_flex": 1.25e-07,
"input_cost_per_token_priority": 4.5e-07,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12435,13 +12477,17 @@
},
"gpt-5-nano": {
"cache_read_input_token_cost": 5e-09,
"cache_read_input_token_cost_flex": 2.5e-09,
"input_cost_per_token": 5e-08,
"input_cost_per_token_flex": 2.5e-08,
"input_cost_per_token_priority": 2.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -12467,13 +12513,16 @@
},
"gpt-5-nano-2025-08-07": {
"cache_read_input_token_cost": 5e-09,
"cache_read_input_token_cost_flex": 2.5e-09,
"input_cost_per_token": 5e-08,
"input_cost_per_token_flex": 2.5e-08,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -15177,13 +15226,19 @@
},
"o3": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_flex": 2.5e-07,
"cache_read_input_token_cost_priority": 8.75e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_flex": 1e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 200000,
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
"output_cost_per_token": 8e-06,
"output_cost_per_token_flex": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
"supported_endpoints": [
"/v1/responses",
"/v1/chat/completions",
@ -15399,13 +15454,19 @@
},
"o4-mini": {
"cache_read_input_token_cost": 2.75e-07,
"cache_read_input_token_cost_flex": 1.375e-07,
"cache_read_input_token_cost_priority": 5e-07,
"input_cost_per_token": 1.1e-06,
"input_cost_per_token_flex": 5.5e-07,
"input_cost_per_token_priority": 2e-06,
"litellm_provider": "openai",
"max_input_tokens": 200000,
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"output_cost_per_token_flex": 2.2e-06,
"output_cost_per_token_priority": 8e-06,
"supports_function_calling": true,
"supports_parallel_function_calling": false,
"supports_pdf_input": true,
@ -16900,6 +16961,20 @@
"supports_tool_choice": true,
"supports_web_search": true
},
"openrouter/x-ai/grok-4-fast:free": {
"input_cost_per_token": 0,
"litellm_provider": "openrouter",
"max_input_tokens": 2000000,
"max_output_tokens": 30000,
"max_tokens": 2000000,
"mode": "chat",
"output_cost_per_token": 0,
"source": "https://openrouter.ai/x-ai/grok-4-fast:free",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_web_search": false
},
"ovhcloud/DeepSeek-R1-Distill-Llama-70B": {
"input_cost_per_token": 6.7e-07,
"litellm_provider": "ovhcloud",
@ -20943,6 +21018,132 @@
"mode": "embedding",
"output_cost_per_token": 0.0
},
"wandb/openai/gpt-oss-120b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.015,
"output_cost_per_token": 0.06,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/openai/gpt-oss-20b": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.005,
"output_cost_per_token": 0.02,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/zai-org/GLM-4.5": {
"max_tokens": 131072,
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"input_cost_per_token": 0.055,
"output_cost_per_token": 0.2,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/Qwen/Qwen3-235B-A22B-Instruct-2507": {
"max_tokens": 262144,
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"input_cost_per_token": 0.01,
"output_cost_per_token": 0.01,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/Qwen/Qwen3-Coder-480B-A35B-Instruct": {
"max_tokens": 262144,
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"input_cost_per_token": 0.1,
"output_cost_per_token": 0.15,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/Qwen/Qwen3-235B-A22B-Thinking-2507": {
"max_tokens": 262144,
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"input_cost_per_token": 0.01,
"output_cost_per_token": 0.01,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/moonshotai/Kimi-K2-Instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.135,
"output_cost_per_token": 0.4,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/meta-llama/Llama-3.1-8B-Instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.022,
"output_cost_per_token": 0.022,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/deepseek-ai/DeepSeek-V3.1": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.055,
"output_cost_per_token": 0.165,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/deepseek-ai/DeepSeek-R1-0528": {
"max_tokens": 161000,
"max_input_tokens": 161000,
"max_output_tokens": 161000,
"input_cost_per_token": 0.135,
"output_cost_per_token": 0.54,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/deepseek-ai/DeepSeek-V3-0324": {
"max_tokens": 161000,
"max_input_tokens": 161000,
"max_output_tokens": 161000,
"input_cost_per_token": 0.114,
"output_cost_per_token": 0.275,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/meta-llama/Llama-3.3-70B-Instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.071,
"output_cost_per_token": 0.071,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/meta-llama/Llama-4-Scout-17B-16E-Instruct": {
"max_tokens": 64000,
"max_input_tokens": 64000,
"max_output_tokens": 64000,
"input_cost_per_token": 0.017,
"output_cost_per_token": 0.066,
"litellm_provider": "wandb",
"mode": "chat"
},
"wandb/microsoft/Phi-4-mini-instruct": {
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"input_cost_per_token": 0.008,
"output_cost_per_token": 0.035,
"litellm_provider": "wandb",
"mode": "chat"
},
"watsonx/ibm/granite-3-8b-instruct": {
"input_cost_per_token": 0.0002,
"litellm_provider": "watsonx",
@ -21337,4 +21538,4 @@
"supports_vision": true,
"supports_web_search": true
}
}
}

View file

@ -12,6 +12,11 @@
- For installation and configuration, see: [Self-hosting guided](https://docs.litellm.ai/docs/proxy/deploy)
- **Telemetry** We run no telemetry when you self host LiteLLM
:::info
✨ SSO is free for up to 5 users. After that, an enterprise license is required. [Get Started with Enterprise here](https://www.litellm.ai/enterprise)
:::
### LiteLLM Cloud
- We encrypt all data stored using your `LITELLM_MASTER_KEY` and in transit using TLS.

View file

@ -463,3 +463,166 @@ def test_calculate_cache_writing_cost():
)
assert result_zero == 0.0
def test_service_tier_flex_pricing():
"""Test that flex service tier uses correct pricing (approximately 50% of standard)."""
# Set up environment for local model cost map
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
# Test with gpt-5-nano which has flex pricing
model = "gpt-5-nano"
custom_llm_provider = "openai"
# Create usage object
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500
)
# Test standard pricing
std_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier=None
)
std_total = std_cost[0] + std_cost[1]
# Test flex pricing
flex_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier="flex"
)
flex_total = flex_cost[0] + flex_cost[1]
# Verify flex is approximately 50% of standard
assert std_total > 0, "Standard cost should be greater than 0"
assert flex_total > 0, "Flex cost should be greater than 0"
flex_ratio = flex_total / std_total
assert 0.45 <= flex_ratio <= 0.55, f"Flex pricing should be ~50% of standard, got {flex_ratio:.2f}"
# Verify specific costs match expected values
# gpt-5-nano flex: input=2.5e-08, output=2e-07
expected_flex_prompt = 1000 * 2.5e-08 # 0.000025
expected_flex_completion = 500 * 2e-07 # 0.0001
expected_flex_total = expected_flex_prompt + expected_flex_completion
assert abs(flex_cost[0] - expected_flex_prompt) < 1e-10, f"Flex prompt cost mismatch: {flex_cost[0]} vs {expected_flex_prompt}"
assert abs(flex_cost[1] - expected_flex_completion) < 1e-10, f"Flex completion cost mismatch: {flex_cost[1]} vs {expected_flex_completion}"
assert abs(flex_total - expected_flex_total) < 1e-10, f"Flex total cost mismatch: {flex_total} vs {expected_flex_total}"
def test_service_tier_default_pricing():
"""Test that when no service tier is provided, standard pricing is used."""
# Set up environment for local model cost map
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
# Test with gpt-5-nano
model = "gpt-5-nano"
custom_llm_provider = "openai"
# Create usage object
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500
)
# Test with no service tier (should use standard)
default_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier=None
)
# Test with explicit standard service tier
standard_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier="standard"
)
# Both should be identical
assert abs(default_cost[0] - standard_cost[0]) < 1e-10, "Default and standard prompt costs should be identical"
assert abs(default_cost[1] - standard_cost[1]) < 1e-10, "Default and standard completion costs should be identical"
# Verify specific costs match expected standard values
# gpt-5-nano standard: input=5e-08, output=4e-07
expected_standard_prompt = 1000 * 5e-08 # 0.00005
expected_standard_completion = 500 * 4e-07 # 0.0002
expected_standard_total = expected_standard_prompt + expected_standard_completion
assert abs(default_cost[0] - expected_standard_prompt) < 1e-10, f"Standard prompt cost mismatch: {default_cost[0]} vs {expected_standard_prompt}"
assert abs(default_cost[1] - expected_standard_completion) < 1e-10, f"Standard completion cost mismatch: {default_cost[1]} vs {expected_standard_completion}"
def test_service_tier_fallback_pricing():
"""Test that when service tier is provided but model doesn't have those keys, it falls back to standard pricing."""
# Set up environment for local model cost map
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
# Test with gpt-4 which doesn't have flex pricing keys
model = "gpt-4"
custom_llm_provider = "openai"
# Create usage object
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500
)
# Test standard pricing
std_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier=None
)
std_total = std_cost[0] + std_cost[1]
# Test flex pricing (should fall back to standard since gpt-4 doesn't have flex keys)
flex_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier="flex"
)
flex_total = flex_cost[0] + flex_cost[1]
# Test priority pricing (should fall back to standard since gpt-4 doesn't have priority keys)
priority_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier="priority"
)
priority_total = priority_cost[0] + priority_cost[1]
# All should be identical (fallback to standard)
assert abs(std_total - flex_total) < 1e-10, f"Standard and flex costs should be identical (fallback): {std_total} vs {flex_total}"
assert abs(std_total - priority_total) < 1e-10, f"Standard and priority costs should be identical (fallback): {std_total} vs {priority_total}"
# Verify costs are reasonable (not zero)
assert std_total > 0, "Standard cost should be greater than 0"
assert flex_total > 0, "Flex cost should be greater than 0 (fallback)"
assert priority_total > 0, "Priority cost should be greater than 0 (fallback)"
# Verify specific costs match expected gpt-4 values
# gpt-4 standard: input=3e-05, output=6e-05
expected_standard_prompt = 1000 * 3e-05 # 0.03
expected_standard_completion = 500 * 6e-05 # 0.03
expected_standard_total = expected_standard_prompt + expected_standard_completion
assert abs(std_cost[0] - expected_standard_prompt) < 1e-10, f"Standard prompt cost mismatch: {std_cost[0]} vs {expected_standard_prompt}"
assert abs(std_cost[1] - expected_standard_completion) < 1e-10, f"Standard completion cost mismatch: {std_cost[1]} vs {expected_standard_completion}"

View file

@ -0,0 +1,69 @@
"""
Unit tests for CLI token utilities
"""
import json
import os
import tempfile
from pathlib import Path
from unittest.mock import mock_open, patch
import pytest
from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key
class TestCLITokenUtils:
"""Test CLI token utility functions"""
def test_get_litellm_gateway_api_key_success(self):
"""Test getting CLI API key when token file exists and is valid"""
token_data = {
'key': 'sk-test-cli-key-123',
'user_id': 'test-user',
'user_email': 'test@example.com',
'timestamp': 1234567890
}
with patch('os.path.exists', return_value=True), \
patch('builtins.open', mock_open(read_data=json.dumps(token_data))), \
patch('litellm.litellm_core_utils.cli_token_utils.get_cli_token_file_path', return_value='/test/.litellm/token.json'):
result = get_litellm_gateway_api_key()
assert result == 'sk-test-cli-key-123'
def test_get_litellm_gateway_api_key_no_file(self):
"""Test getting CLI API key when token file doesn't exist"""
with patch('os.path.exists', return_value=False), \
patch('litellm.litellm_core_utils.cli_token_utils.get_cli_token_file_path', return_value='/test/.litellm/token.json'):
result = get_litellm_gateway_api_key()
assert result is None
def test_get_litellm_gateway_api_key_invalid_json(self):
"""Test getting CLI API key when token file has invalid JSON"""
with patch('os.path.exists', return_value=True), \
patch('builtins.open', mock_open(read_data='invalid json')), \
patch('litellm.litellm_core_utils.cli_token_utils.get_cli_token_file_path', return_value='/test/.litellm/token.json'):
result = get_litellm_gateway_api_key()
assert result is None
def test_get_litellm_gateway_api_key_no_key_field(self):
"""Test getting CLI API key when token file exists but has no key field"""
token_data = {
'user_id': 'test-user',
'user_email': 'test@example.com'
# Missing 'key' field
}
with patch('os.path.exists', return_value=True), \
patch('builtins.open', mock_open(read_data=json.dumps(token_data))), \
patch('litellm.litellm_core_utils.cli_token_utils.get_cli_token_file_path', return_value='/test/.litellm/token.json'):
result = get_litellm_gateway_api_key()
assert result is None

View file

@ -0,0 +1,32 @@
import os
import sys
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
from litellm.llms.azure_ai.image_edit.transformation import AzureFoundryFluxImageEditConfig
def test_azure_ai_validate_environment():
"""Test Azure AI environment validation"""
config = AzureFoundryFluxImageEditConfig()
headers = {}
config.validate_environment(headers, "FLUX.1-Kontext-pro", api_key="test-key")
assert "Api-Key" in headers
assert headers["Api-Key"] == "test-key"
def test_azure_ai_url_generation():
"""Test Azure AI URL generation"""
config = AzureFoundryFluxImageEditConfig()
api_base = "https://test-endpoint.eastus2.inference.ai.azure.com"
complete_url = config.get_complete_url(
model="FLUX.1-Kontext-pro",
api_base=api_base,
litellm_params={"api_version": "2025-04-01-preview"}
)
expected_url = f"{api_base}/openai/deployments/FLUX.1-Kontext-pro/images/edits?api-version=2025-04-01-preview"
assert complete_url == expected_url

View file

@ -0,0 +1,146 @@
"""
Unit tests for WandB Inference configuration.
These tests validate the WandbInferenceConfig class which extends OpenAIGPTConfig.
Nebius AI Studio is an OpenAI-compatible provider with minor customizations.
"""
import os
import sys
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
import pytest
import litellm
from litellm import completion
from litellm.llms.wandb.chat.transformation import WandbConfig
class TestWandbConfig:
"""Test class for WandB Inference functionality"""
def test_default_api_base(self):
"""Test that default API base is used when none is provided"""
config = WandbConfig()
headers = {}
api_key = "fake-wandb-key"
# Call validate_environment without specifying api_base
result = config.validate_environment(
headers=headers,
model="wandb/openai/gpt-oss-20b",
messages=[{"role": "user", "content": "Hey"}],
optional_params={},
litellm_params={},
api_key=api_key,
api_base=None, # Not providing api_base
)
# Verify headers are still set correctly
assert result["Authorization"] == f"Bearer {api_key}"
assert result["Content-Type"] == "application/json"
# We can't directly test the api_base value here since validate_environment
# only returns the headers, but we can verify it doesn't raise an exception
# which would happen if api_base handling was incorrect
@pytest.mark.respx()
def test_wandb_completion_mock(self, respx_mock):
"""
Mock test for WandB Inference completion using the model format from docs.
This test mocks the actual HTTP request to test the integration properly.
"""
litellm.disable_aiohttp_transport = (
True # since this uses respx, we need to set use_aiohttp_transport to False
)
# Set up environment variables for the test
api_key = "fake-wandb-key"
api_base = "https://api.inference.wandb.ai/v1"
model = "wandb/openai/gpt-oss-20b"
model_name = "gpt-oss-20b" # The actual model name without provider prefix
# Mock the HTTP request to the WandB Inference API
respx_mock.post(f"{api_base}/chat/completions").respond(
json={
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": model_name,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": '```python\nprint("Hey from LiteLLM!")\n```\n\nThis simple Python code prints a greeting message from LiteLLM.',
},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 9,
"completion_tokens": 12,
"total_tokens": 21,
},
},
status_code=200,
)
# Make the actual API call through LiteLLM
response = completion(
model=model,
messages=[
{"role": "user", "content": "write code for saying hey from LiteLLM"}
],
api_key=api_key,
api_base=api_base,
)
# Verify response structure
assert response is not None
# If response is a streaming wrapper, extract the first chunk for assertions
# This handles both streaming and non-streaming responses
# For streaming, response is typically an iterator yielding (event, data) tuples
if hasattr(response, "__iter__") and not hasattr(response, "choices"):
# Streaming response: get the first chunk
first_chunk = next(iter(response))
# first_chunk is likely a tuple: (event, data)
# Try to extract the data part
if isinstance(first_chunk, tuple) and len(first_chunk) == 2:
data = first_chunk[1]
else:
data = first_chunk
# The data object should have .choices[0] with .delta or .message
choices = getattr(data, "choices", None)
assert choices is not None
assert len(choices) > 0
choice = choices[0]
# For streaming, content may be in .delta or .message
content = None
if hasattr(choice, "delta") and hasattr(choice.delta, "content"):
content = choice.delta.content
elif hasattr(choice, "message") and hasattr(choice.message, "content"):
content = choice.message.content
assert content is not None
assert "```python" in content
assert "Hey from LiteLLM" in content
else:
# Non-streaming response
choices = getattr(response, "choices", None)
assert choices is not None
assert len(choices) > 0
choice = choices[0]
message = getattr(choice, "message", None)
assert message is not None
content = getattr(message, "content", None)
assert content is not None
# Check for specific content in the response
assert "```python" in content
assert "Hey from LiteLLM" in content

View file

@ -0,0 +1,212 @@
"""
Test suite for MCP server custom fields functionality.
Tests that mcp_info can accept arbitrary custom fields in addition to predefined ones.
"""
import pytest
import sys
import os
from unittest.mock import Mock, patch
from typing import Dict, Any
# Add the path to find the modules
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adjust the path as needed
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.types.mcp import MCPAuth
from litellm.proxy._types import LiteLLM_MCPServerTable
class TestMCPCustomFields:
"""Test custom fields functionality in MCP server configuration."""
def test_custom_fields_preserved_from_config(self):
"""Test that custom fields in mcp_info are preserved when loading from config."""
manager = MCPServerManager()
# Mock config with custom fields
mock_config = {
"test_server": {
"url": "http://localhost:3000",
"transport": "http",
"auth_type": "bearer_token",
"authentication_token": "test-token",
"mcp_info": {
"server_name": "Test Server",
"description": "A test server",
"custom_field_1": "custom_value_1",
"custom_field_2": {"nested": "value"},
"custom_field_3": ["list", "values"],
"priority": 10,
"tags": ["production", "api"]
}
}
}
# Load servers from config
manager.load_servers_from_config(mock_config)
# Get the loaded server
servers = list(manager.config_mcp_servers.values())
assert len(servers) == 1
server = servers[0]
mcp_info = server.mcp_info
# Verify standard fields are preserved
assert mcp_info["server_name"] == "Test Server"
assert mcp_info["description"] == "A test server"
# Verify custom fields are preserved
assert mcp_info["custom_field_1"] == "custom_value_1"
assert mcp_info["custom_field_2"] == {"nested": "value"}
assert mcp_info["custom_field_3"] == ["list", "values"]
assert mcp_info["priority"] == 10
assert mcp_info["tags"] == ["production", "api"]
def test_custom_fields_preserved_from_database(self):
"""Test that custom fields in mcp_info are preserved when adding from database."""
manager = MCPServerManager()
# Mock database record with custom fields
mock_server = Mock(spec=LiteLLM_MCPServerTable)
mock_server.server_id = "test-server-id"
mock_server.server_name = "Test Server"
mock_server.description = "A test server"
mock_server.url = "http://localhost:3000"
mock_server.transport = "http"
mock_server.auth_type = MCPAuth.bearer_token
mock_server.alias = None
mock_server.mcp_info = {
"server_name": "Test Server",
"description": "A test server",
"custom_db_field": "database_value",
"metadata": {"source": "database"},
"version": "1.0.0"
}
mock_server.command = None
mock_server.args = None
mock_server.env = None
mock_server.mcp_access_groups = None
# Add server to manager
manager.add_update_server(mock_server)
# Get the added server
server = manager.get_mcp_server_by_id("test-server-id")
assert server is not None
mcp_info = server.mcp_info
# Verify standard fields are preserved
assert mcp_info["server_name"] == "Test Server"
assert mcp_info["description"] == "A test server"
# Verify custom fields are preserved
assert mcp_info["custom_db_field"] == "database_value"
assert mcp_info["metadata"] == {"source": "database"}
assert mcp_info["version"] == "1.0.0"
def test_empty_mcp_info_handled_gracefully(self):
"""Test that empty or missing mcp_info is handled gracefully."""
manager = MCPServerManager()
# Config with empty mcp_info
mock_config = {
"test_server": {
"url": "http://localhost:3000",
"transport": "http",
"mcp_info": {}
}
}
manager.load_servers_from_config(mock_config)
servers = list(manager.config_mcp_servers.values())
assert len(servers) == 1
server = servers[0]
mcp_info = server.mcp_info
# Should have default server_name
assert mcp_info["server_name"] == "test_server"
def test_missing_mcp_info_creates_defaults(self):
"""Test that missing mcp_info creates appropriate defaults."""
manager = MCPServerManager()
# Config without mcp_info
mock_config = {
"test_server": {
"url": "http://localhost:3000",
"transport": "http",
"description": "Server description"
}
}
manager.load_servers_from_config(mock_config)
servers = list(manager.config_mcp_servers.values())
assert len(servers) == 1
server = servers[0]
mcp_info = server.mcp_info
# Should have default server_name and description from config
assert mcp_info["server_name"] == "test_server"
assert mcp_info["description"] == "Server description"
def test_config_description_fallback(self):
"""Test that description from config level is used as fallback."""
manager = MCPServerManager()
# Config with description at server level but not in mcp_info
mock_config = {
"test_server": {
"url": "http://localhost:3000",
"transport": "http",
"description": "Config level description",
"mcp_info": {
"custom_field": "custom_value"
}
}
}
manager.load_servers_from_config(mock_config)
servers = list(manager.config_mcp_servers.values())
server = servers[0]
mcp_info = server.mcp_info
# Should use config level description as fallback
assert mcp_info["description"] == "Config level description"
assert mcp_info["custom_field"] == "custom_value"
def test_mcp_info_description_takes_precedence(self):
"""Test that description in mcp_info takes precedence over config level."""
manager = MCPServerManager()
# Config with description at both levels
mock_config = {
"test_server": {
"url": "http://localhost:3000",
"transport": "http",
"description": "Config level description",
"mcp_info": {
"description": "MCP info description",
"custom_field": "custom_value"
}
}
}
manager.load_servers_from_config(mock_config)
servers = list(manager.config_mcp_servers.values())
server = servers[0]
mcp_info = server.mcp_info
# Should use mcp_info description, not config level
assert mcp_info["description"] == "MCP info description"
assert mcp_info["custom_field"] == "custom_value"

View file

@ -102,7 +102,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
working_server if server_id == "working_server" else failing_server
)
async def mock_get_tools_from_server(server, mcp_auth_header=None):
async def mock_get_tools_from_server(server, mcp_auth_header=None, add_prefix=True):
if server.name == "working_server":
# Working server returns tools
tool1 = MagicMock()
@ -184,7 +184,7 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
failing_server1 if server_id == "failing_server1" else failing_server2
)
async def mock_get_tools_from_server(server, mcp_auth_header=None):
async def mock_get_tools_from_server(server, mcp_auth_header=None, add_prefix=True):
# All servers fail
raise Exception(f"Server {server.name} connection failed")
@ -448,3 +448,121 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name():
assert (
called_servers[0].server_id == specific_server.server_id
), "Should have contacted the specific server alias, not the group."
@pytest.mark.asyncio
async def test_list_tools_single_server_unprefixed_names():
"""When only one MCP server is allowed, list tools should return unprefixed names."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
set_auth_context,
)
except ImportError:
pytest.skip("MCP server not available")
# Mock user auth
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user")
set_auth_context(user_api_key_auth)
# One allowed server
server = MagicMock()
server.server_id = "server1"
server.name = "Zapier MCP"
server.alias = "zapier"
# Mock manager: allow just one server and return a tool based on add_prefix flag
mock_manager = MagicMock()
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = (
lambda server_id: server if server_id == "server1" else None
)
async def mock_get_tools_from_server(
server, mcp_auth_header=None, add_prefix=False
):
tool = MagicMock()
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
tool.description = "desc"
tool.inputSchema = {}
return [tool]
mock_manager._get_tools_from_server = mock_get_tools_from_server
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
mock_manager,
):
tools = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
)
# Should be unprefixed since only one server is allowed
assert len(tools) == 1
assert tools[0].name == "toolA"
@pytest.mark.asyncio
async def test_list_tools_multiple_servers_prefixed_names():
"""When multiple MCP servers are allowed, list tools should return prefixed names."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
set_auth_context,
)
except ImportError:
pytest.skip("MCP server not available")
# Mock user auth
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user")
set_auth_context(user_api_key_auth)
# Two allowed servers
server1 = MagicMock()
server1.server_id = "server1"
server1.name = "Zapier MCP"
server1.alias = "zapier"
server2 = MagicMock()
server2.server_id = "server2"
server2.name = "Jira MCP"
server2.alias = "jira"
# Mock manager
mock_manager = MagicMock()
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1", "server2"]
)
mock_manager.get_mcp_server_by_id = (
lambda server_id: server1 if server_id == "server1" else server2
)
async def mock_get_tools_from_server(
server, mcp_auth_header=None, add_prefix=True
):
tool = MagicMock()
# When multiple servers, add_prefix should be True -> prefixed names
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
tool.description = "desc"
tool.inputSchema = {}
return [tool]
mock_manager._get_tools_from_server = mock_get_tools_from_server
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
mock_manager,
):
tools = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
)
# Should be prefixed since multiple servers are allowed
names = sorted([t.name for t in tools])
assert names == ["jira-toolA", "zapier-toolA"]

View file

@ -420,6 +420,112 @@ class TestMCPServerManager:
assert result["status"] == "healthy"
assert result["tools_count"] == 1
@pytest.mark.asyncio
async def test_get_tools_from_server_add_prefix(self):
"""Verify _get_tools_from_server respects add_prefix True/False."""
manager = MCPServerManager()
# Create a minimal server with alias used as prefix
server = MCPServer(
server_id="zapier",
name="zapier",
transport=MCPTransport.http,
)
# Mock client creation and fetching tools
manager._create_mcp_client = MagicMock(return_value=object())
# Tools returned upstream (unprefixed from provider)
upstream_tool = MagicMock()
upstream_tool.name = "send_email"
upstream_tool.description = "Send an email"
upstream_tool.inputSchema = {}
manager._fetch_tools_with_timeout = AsyncMock(return_value=[upstream_tool])
# Case 1: add_prefix=True (default for multi-server) -> expect prefixed
tools_prefixed = await manager._get_tools_from_server(server, add_prefix=True)
assert len(tools_prefixed) == 1
assert tools_prefixed[0].name == "zapier-send_email"
# Case 2: add_prefix=False (single-server) -> expect unprefixed
tools_unprefixed = await manager._get_tools_from_server(
server, add_prefix=False
)
assert len(tools_unprefixed) == 1
assert tools_unprefixed[0].name == "send_email"
def test_create_prefixed_tools_updates_mapping_for_both_forms(self):
"""_create_prefixed_tools should populate mapping for prefixed and original names even when not adding prefix in output."""
manager = MCPServerManager()
server = MCPServer(
server_id="jira",
name="jira",
transport=MCPTransport.http,
)
# Input tools as would come from upstream
t1 = MagicMock()
t1.name = "create_issue"
t1.description = ""
t1.inputSchema = {}
t2 = MagicMock()
t2.name = "close_issue"
t2.description = ""
t2.inputSchema = {}
# Do not add prefix in returned objects
out_tools = manager._create_prefixed_tools([t1, t2], server, add_prefix=False)
# Returned names should be unprefixed
names = sorted([t.name for t in out_tools])
assert names == ["close_issue", "create_issue"]
# Mapping should include both original and prefixed names -> resolves calls either way
assert manager.tool_name_to_mcp_server_name_mapping["create_issue"] == "jira"
assert (
manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira"
)
assert manager.tool_name_to_mcp_server_name_mapping["close_issue"] == "jira"
assert (
manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira"
)
def test_get_mcp_server_from_tool_name_with_prefixed_and_unprefixed(self):
"""After mapping is populated, manager resolves both prefixed and unprefixed tool names to the same server."""
manager = MCPServerManager()
server = MCPServer(
server_id="zapier",
name="zapier",
server_name="zapier",
transport=MCPTransport.http,
)
# Register server so resolution can find it
manager.registry = {server.server_id: server}
# Populate mapping (add_prefix value doesn't matter for mapping population)
base_tool = MagicMock()
base_tool.name = "create_zap"
base_tool.description = ""
base_tool.inputSchema = {}
_ = manager._create_prefixed_tools([base_tool], server, add_prefix=False)
# Unprefixed resolution
resolved_server_unpref = manager._get_mcp_server_from_tool_name("create_zap")
print(resolved_server_unpref)
assert resolved_server_unpref is not None
assert resolved_server_unpref.server_id == server.server_id
# Prefixed resolution
resolved_server_pref = manager._get_mcp_server_from_tool_name(
"zapier-create_zap"
)
assert resolved_server_pref is not None
assert resolved_server_pref.server_id == server.server_id
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -435,3 +435,105 @@ class TestWhoamiCommand:
assert "✅ Authenticated" in result.output
# Should calculate age based on timestamp=0
assert "Token age:" in result.output
class TestCLIKeyRegenerationFlow:
"""Test the end-to-end CLI key regeneration flow from CLI perspective"""
def setup_method(self):
"""Setup for each test"""
self.runner = CliRunner()
def test_login_with_existing_key_regeneration_flow(self):
"""Test complete login flow when user has existing key - should regenerate it"""
mock_context = Mock()
mock_context.obj = {"base_url": "https://test.example.com"}
# Mock existing stored key
existing_key = "sk-existing-key-123"
# Mock successful regeneration response
mock_response = Mock()
mock_response.status_code = 200
mock_response.json.return_value = {
"status": "ready",
"key": "sk-regenerated-key-456" # New regenerated key
}
with patch('webbrowser.open') as mock_browser, \
patch('requests.get', return_value=mock_response) as mock_get, \
patch('litellm.proxy.client.cli.commands.auth.get_stored_api_key', return_value=existing_key) as mock_get_stored, \
patch('litellm.proxy.client.cli.commands.auth.save_token') as mock_save, \
patch('litellm.proxy.client.cli.interface.show_commands') as mock_show_commands, \
patch('uuid.uuid4', return_value='new-session-uuid-789'):
result = self.runner.invoke(login, obj=mock_context.obj)
assert result.exit_code == 0
assert "✅ Login successful!" in result.output
assert "API Key: sk-regenerated-key-456" in result.output
# Verify existing key was retrieved
mock_get_stored.assert_called_once()
# Verify browser was opened with correct URL including existing key
mock_browser.assert_called_once()
call_args = mock_browser.call_args[0][0]
assert "https://test.example.com/sso/key/generate" in call_args
assert "source=litellm-cli" in call_args
assert "key=sk-new-session-uuid-789" in call_args
assert f"existing_key={existing_key}" in call_args
# Verify polling was done with correct session key
mock_get.assert_called()
poll_url = mock_get.call_args[0][0]
assert "sk-new-session-uuid-789" in poll_url
# Verify regenerated key was saved
mock_save.assert_called_once()
saved_data = mock_save.call_args[0][0]
assert saved_data['key'] == 'sk-regenerated-key-456'
assert saved_data['user_id'] == 'cli-user'
mock_show_commands.assert_called_once()
def test_login_without_existing_key_creation_flow(self):
"""Test complete login flow when user has no existing key - should create new one"""
mock_context = Mock()
mock_context.obj = {"base_url": "https://test.example.com"}
# Mock no existing key
mock_response = Mock()
mock_response.status_code = 200
mock_response.json.return_value = {
"status": "ready",
"key": "sk-new-created-key-789"
}
with patch('webbrowser.open') as mock_browser, \
patch('requests.get', return_value=mock_response), \
patch('litellm.proxy.client.cli.commands.auth.get_stored_api_key', return_value=None) as mock_get_stored, \
patch('litellm.proxy.client.cli.commands.auth.save_token') as mock_save, \
patch('litellm.proxy.client.cli.interface.show_commands'), \
patch('uuid.uuid4', return_value='new-session-uuid-999'):
result = self.runner.invoke(login, obj=mock_context.obj)
assert result.exit_code == 0
assert "✅ Login successful!" in result.output
# Verify existing key check was done
mock_get_stored.assert_called_once()
# Verify browser was opened with correct URL WITHOUT existing key
mock_browser.assert_called_once()
call_args = mock_browser.call_args[0][0]
assert "https://test.example.com/sso/key/generate" in call_args
assert "source=litellm-cli" in call_args
assert "key=sk-new-session-uuid-999" in call_args
assert "existing_key=" not in call_args # Should not include existing_key param
# Verify new key was saved
mock_save.assert_called_once()
saved_data = mock_save.call_args[0][0]
assert saved_data['key'] == 'sk-new-created-key-789'

View file

@ -3,6 +3,7 @@ Test priority-based rate limiting for dynamic_rate_limiter_v3.
Core tests to validate that priority weights are respected (0.9/0.1) instead of equal splitting (0.5/0.5).
"""
import asyncio
import os
import sys
@ -25,22 +26,25 @@ from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
async def test_priority_weight_allocation():
"""
Test that priority weights are correctly applied instead of equal splitting.
With priority_reservation = {"high": 0.9, "low": 0.1}:
- High priority should get 90% of TPM (900 out of 1000)
- Low priority should get 10% of TPM (100 out of 1000)
This validates the core fix where before it would split 50/50.
"""
# Set up environment for premium feature
os.environ["LITELLM_LICENSE"] = "test-license-key"
# Set up priority reservations
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
dual_cache = DualCache()
handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache)
model = "test-model"
total_tpm = 1000
llm_router = Router(
model_list=[
{
@ -55,72 +59,75 @@ async def test_priority_weight_allocation():
]
)
handler.update_variables(llm_router=llm_router)
# Test high priority allocation
high_priority_user = UserAPIKeyAuth()
high_priority_user.metadata = {"priority": "high"}
high_descriptors = handler._create_priority_based_descriptors(
model=model,
user_api_key_dict=high_priority_user,
priority="high",
)
assert len(high_descriptors) == 1
high_descriptor = high_descriptors[0]
expected_high_tpm = int(total_tpm * 0.9) # 900
actual_high_tpm = high_descriptor["rate_limit"]["tokens_per_unit"]
assert actual_high_tpm == expected_high_tpm, (
f"High priority should get {expected_high_tpm} TPM (90%), got {actual_high_tpm}"
)
assert (
actual_high_tpm == expected_high_tpm
), f"High priority should get {expected_high_tpm} TPM (90%), got {actual_high_tpm}"
assert high_descriptor["value"] == f"{model}:high"
# Test low priority allocation
low_priority_user = UserAPIKeyAuth()
low_priority_user.metadata = {"priority": "low"}
low_descriptors = handler._create_priority_based_descriptors(
model=model,
user_api_key_dict=low_priority_user,
priority="low",
)
assert len(low_descriptors) == 1
low_descriptor = low_descriptors[0]
expected_low_tpm = int(total_tpm * 0.1) # 100
actual_low_tpm = low_descriptor["rate_limit"]["tokens_per_unit"]
assert actual_low_tpm == expected_low_tpm, (
f"Low priority should get {expected_low_tpm} TPM (10%), got {actual_low_tpm}"
)
assert (
actual_low_tpm == expected_low_tpm
), f"Low priority should get {expected_low_tpm} TPM (10%), got {actual_low_tpm}"
assert low_descriptor["value"] == f"{model}:low"
# Verify the ratio is 9:1, not 1:1 (equal splitting)
ratio = actual_high_tpm / actual_low_tpm
expected_ratio = 9.0
assert abs(ratio - expected_ratio) < 0.1, (
f"High:Low ratio should be {expected_ratio}:1, got {ratio}:1"
)
assert (
abs(ratio - expected_ratio) < 0.1
), f"High:Low ratio should be {expected_ratio}:1, got {ratio}:1"
@pytest.mark.asyncio
async def test_concurrent_priority_requests():
"""
Test the core issue: 5 concurrent requests with different priorities should get
Test the core issue: 5 concurrent requests with different priorities should get
proper allocation based on priority weights, not equal splitting.
This tests the exact scenario mentioned: priorities 0.9 and 0.1 should be 0.9/0.1, not 0.5/0.5.
"""
# Set up environment for premium feature
os.environ["LITELLM_LICENSE"] = "test-license-key"
# Set up the exact scenario from the issue
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
dual_cache = DualCache()
handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache)
model = "test-model"
total_tpm = 1000
llm_router = Router(
model_list=[
{
@ -135,23 +142,23 @@ async def test_concurrent_priority_requests():
]
)
handler.update_variables(llm_router=llm_router)
# Create 5 concurrent users - 3 high priority, 2 low priority
high_priority_users = []
low_priority_users = []
for i in range(3): # 3 high priority users
user = UserAPIKeyAuth()
user.metadata = {"priority": "high"}
user.user_id = f"high_user_{i}"
high_priority_users.append(user)
for i in range(2): # 2 low priority users
for i in range(2): # 2 low priority users
user = UserAPIKeyAuth()
user.metadata = {"priority": "low"}
user.user_id = f"low_user_{i}"
low_priority_users.append(user)
# Test all high priority users get the same allocation (not divided)
for user in high_priority_users:
descriptors = handler._create_priority_based_descriptors(
@ -159,7 +166,7 @@ async def test_concurrent_priority_requests():
user_api_key_dict=user,
priority="high",
)
assert len(descriptors) == 1
descriptor = descriptors[0]
# Each high priority user should get 900 TPM, not divided by 3
@ -168,7 +175,7 @@ async def test_concurrent_priority_requests():
f"got {descriptor['rate_limit']['tokens_per_unit']}"
)
assert descriptor["value"] == f"{model}:high"
# Test all low priority users get the same allocation (not divided)
for user in low_priority_users:
descriptors = handler._create_priority_based_descriptors(
@ -176,7 +183,7 @@ async def test_concurrent_priority_requests():
user_api_key_dict=user,
priority="low",
)
assert len(descriptors) == 1
descriptor = descriptors[0]
# Each low priority user should get 100 TPM, not divided by 2
@ -191,21 +198,24 @@ async def test_concurrent_priority_requests():
async def test_100_concurrent_priority_requests():
"""
Stress test: 100 concurrent requests with mixed priorities over 10 seconds.
This validates that the priority system works correctly under high load:
- 70 high priority requests (should get 900 TPM each)
- 30 low priority requests (should get 100 TPM each)
- Spread across 10 seconds to simulate real-world load
"""
# Set up environment for premium feature
os.environ["LITELLM_LICENSE"] = "test-license-key"
# Set up priority reservations
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
dual_cache = DualCache()
handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache)
model = "stress-test-model"
total_tpm = 1000
llm_router = Router(
model_list=[
{
@ -221,108 +231,130 @@ async def test_100_concurrent_priority_requests():
]
)
handler.update_variables(llm_router=llm_router)
# Create 100 users: 70 high priority, 30 low priority
all_users = []
# 70 high priority users
for i in range(70):
user = UserAPIKeyAuth()
user.metadata = {"priority": "high"}
user.user_id = f"high_stress_user_{i}"
all_users.append((user, "high", 900, 450)) # expected TPM, expected RPM
# 30 low priority users
for i in range(30):
user = UserAPIKeyAuth()
user.metadata = {"priority": "low"}
user.user_id = f"low_stress_user_{i}"
all_users.append((user, "low", 100, 50)) # expected TPM, expected RPM
async def test_user_descriptors(user_data):
"""Test descriptor creation for a single user."""
user, priority, expected_tpm, expected_rpm = user_data
descriptors = handler._create_priority_based_descriptors(
model=model,
user_api_key_dict=user,
priority=priority,
)
assert len(descriptors) == 1, f"User {user.user_id} should have exactly 1 descriptor"
assert (
len(descriptors) == 1
), f"User {user.user_id} should have exactly 1 descriptor"
descriptor = descriptors[0]
# Validate TPM allocation
actual_tpm = descriptor["rate_limit"]["tokens_per_unit"]
assert actual_tpm == expected_tpm, (
f"User {user.user_id} ({priority}) should get {expected_tpm} TPM, got {actual_tpm}"
)
assert (
actual_tpm == expected_tpm
), f"User {user.user_id} ({priority}) should get {expected_tpm} TPM, got {actual_tpm}"
# Validate RPM allocation
actual_rpm = descriptor["rate_limit"]["requests_per_unit"]
assert actual_rpm == expected_rpm, (
f"User {user.user_id} ({priority}) should get {expected_rpm} RPM, got {actual_rpm}"
)
assert (
actual_rpm == expected_rpm
), f"User {user.user_id} ({priority}) should get {expected_rpm} RPM, got {actual_rpm}"
# Validate descriptor key
assert descriptor["value"] == f"{model}:{priority}"
assert descriptor["key"] == "priority_model"
return {
"user_id": user.user_id,
"priority": priority,
"tpm": actual_tpm,
"rpm": actual_rpm,
"success": True
"success": True,
}
# Run all 100 requests concurrently to simulate high load
start_time = time.time()
# Split into batches to simulate requests over 10 seconds
batch_size = 10 # 10 requests per batch
batches = [all_users[i:i + batch_size] for i in range(0, len(all_users), batch_size)]
batches = [
all_users[i : i + batch_size] for i in range(0, len(all_users), batch_size)
]
all_results = []
for batch_idx, batch in enumerate(batches):
# Process each batch concurrently
batch_tasks = [test_user_descriptors(user_data) for user_data in batch]
batch_results = await asyncio.gather(*batch_tasks, return_exceptions=True)
all_results.extend(batch_results)
# Add small delay between batches to spread over ~10 seconds
if batch_idx < len(batches) - 1: # Don't sleep after last batch
await asyncio.sleep(1.0) # 1 second between batches
end_time = time.time()
total_duration = end_time - start_time
# Validate that the test ran over approximately 10 seconds
assert total_duration >= 9.0, f"Test should take ~10 seconds, took {total_duration:.2f}s"
assert (
total_duration >= 9.0
), f"Test should take ~10 seconds, took {total_duration:.2f}s"
assert total_duration <= 15.0, f"Test took too long: {total_duration:.2f}s"
# Validate all requests were successful
successful_results = [r for r in all_results if isinstance(r, dict) and r.get("success")]
assert len(successful_results) == 100, f"Expected 100 successful results, got {len(successful_results)}"
successful_results = [
r for r in all_results if isinstance(r, dict) and r.get("success")
]
assert (
len(successful_results) == 100
), f"Expected 100 successful results, got {len(successful_results)}"
# Validate priority distribution
high_priority_results = [r for r in successful_results if r["priority"] == "high"]
low_priority_results = [r for r in successful_results if r["priority"] == "low"]
assert len(high_priority_results) == 70, f"Expected 70 high priority results, got {len(high_priority_results)}"
assert len(low_priority_results) == 30, f"Expected 30 low priority results, got {len(low_priority_results)}"
assert (
len(high_priority_results) == 70
), f"Expected 70 high priority results, got {len(high_priority_results)}"
assert (
len(low_priority_results) == 30
), f"Expected 30 low priority results, got {len(low_priority_results)}"
# Validate all high priority users got correct allocation
for result in high_priority_results:
assert result["tpm"] == 900, f"High priority user {result['user_id']} got {result['tpm']} TPM, expected 900"
assert result["rpm"] == 450, f"High priority user {result['user_id']} got {result['rpm']} RPM, expected 450"
assert (
result["tpm"] == 900
), f"High priority user {result['user_id']} got {result['tpm']} TPM, expected 900"
assert (
result["rpm"] == 450
), f"High priority user {result['user_id']} got {result['rpm']} RPM, expected 450"
# Validate all low priority users got correct allocation
for result in low_priority_results:
assert result["tpm"] == 100, f"Low priority user {result['user_id']} got {result['tpm']} TPM, expected 100"
assert result["rpm"] == 50, f"Low priority user {result['user_id']} got {result['rpm']} RPM, expected 50"
assert (
result["tpm"] == 100
), f"Low priority user {result['user_id']} got {result['tpm']} TPM, expected 100"
assert (
result["rpm"] == 50
), f"Low priority user {result['user_id']} got {result['rpm']} RPM, expected 50"
print(f"✅ Successfully processed 100 concurrent requests in {total_duration:.2f}s")
print(f" - 70 high priority users: 900 TPM, 450 RPM each")
print(f" - 30 low priority users: 100 TPM, 50 RPM each")
@ -333,17 +365,20 @@ async def test_100_concurrent_priority_requests():
async def test_concurrent_pre_call_hooks_stress():
"""
Stress test: 50 concurrent pre-call hooks with priority enforcement.
This tests the actual rate limiting logic under concurrent load.
"""
# Set up environment for premium feature
os.environ["LITELLM_LICENSE"] = "test-license-key"
litellm.priority_reservation = {"premium": 0.8, "standard": 0.2}
dual_cache = DualCache()
handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache)
model = "pre-call-stress-model"
total_tpm = 2000
llm_router = Router(
model_list=[
{
@ -358,71 +393,80 @@ async def test_concurrent_pre_call_hooks_stress():
]
)
handler.update_variables(llm_router=llm_router)
# Mock the v3 limiter to simulate different scenarios
successful_requests = []
rate_limited_requests = []
async def mock_should_rate_limit(descriptors, parent_otel_span=None):
"""Mock rate limiter that allows premium users, limits some standard users."""
descriptor = descriptors[0]
priority = descriptor["value"].split(":")[-1]
if priority == "premium":
# Allow all premium requests
return {
"overall_code": "OK",
"statuses": [{
"code": "OK",
"descriptor_key": descriptor["value"],
"rate_limit_type": "tokens_per_unit",
"limit_remaining": 1000
}]
"statuses": [
{
"code": "OK",
"descriptor_key": descriptor["value"],
"rate_limit_type": "tokens_per_unit",
"limit_remaining": 1000,
}
],
}
else:
# Rate limit some standard requests (simulate load)
import random
if random.random() < 0.3: # 30% of standard requests get rate limited
return {
"overall_code": "OVER_LIMIT",
"statuses": [{
"code": "OVER_LIMIT",
"descriptor_key": descriptor["value"],
"rate_limit_type": "tokens_per_unit",
"limit_remaining": 0
}]
"statuses": [
{
"code": "OVER_LIMIT",
"descriptor_key": descriptor["value"],
"rate_limit_type": "tokens_per_unit",
"limit_remaining": 0,
}
],
}
else:
return {
"overall_code": "OK",
"statuses": [{
"code": "OK",
"descriptor_key": descriptor["value"],
"rate_limit_type": "tokens_per_unit",
"limit_remaining": 100
}]
"statuses": [
{
"code": "OK",
"descriptor_key": descriptor["value"],
"rate_limit_type": "tokens_per_unit",
"limit_remaining": 100,
}
],
}
# Create 50 users: 30 premium, 20 standard
users = []
for i in range(30):
user = UserAPIKeyAuth()
user.metadata = {"priority": "premium"}
user.user_id = f"premium_hook_user_{i}"
users.append((user, "premium"))
for i in range(20):
user = UserAPIKeyAuth()
user.metadata = {"priority": "standard"}
user.user_id = f"standard_hook_user_{i}"
users.append((user, "standard"))
async def make_request(user_data):
"""Make a pre-call hook request."""
user, priority = user_data
with patch.object(handler.v3_limiter, 'should_rate_limit', side_effect=mock_should_rate_limit):
with patch.object(
handler.v3_limiter, "should_rate_limit", side_effect=mock_should_rate_limit
):
try:
result = await handler.async_pre_call_hook(
user_api_key_dict=user,
@ -430,53 +474,79 @@ async def test_concurrent_pre_call_hooks_stress():
data={"model": model},
call_type="completion",
)
# If no exception, request was allowed
successful_requests.append({
successful_requests.append(
{"user_id": user.user_id, "priority": priority, "result": "allowed"}
)
return {
"status": "success",
"user_id": user.user_id,
"priority": priority,
"result": "allowed"
})
return {"status": "success", "user_id": user.user_id, "priority": priority}
}
except Exception as e:
# Request was rate limited
rate_limited_requests.append({
rate_limited_requests.append(
{"user_id": user.user_id, "priority": priority, "error": str(e)}
)
return {
"status": "rate_limited",
"user_id": user.user_id,
"priority": priority,
"error": str(e)
})
return {"status": "rate_limited", "user_id": user.user_id, "priority": priority}
}
# Run all 50 requests concurrently
start_time = time.time()
tasks = [make_request(user_data) for user_data in users]
results = await asyncio.gather(*tasks, return_exceptions=True)
end_time = time.time()
# Analyze results
successful_count = len([r for r in results if isinstance(r, dict) and r["status"] == "success"])
rate_limited_count = len([r for r in results if isinstance(r, dict) and r["status"] == "rate_limited"])
successful_count = len(
[r for r in results if isinstance(r, dict) and r["status"] == "success"]
)
rate_limited_count = len(
[r for r in results if isinstance(r, dict) and r["status"] == "rate_limited"]
)
# Validate that premium users were mostly successful (priority worked)
premium_results = [r for r in results if isinstance(r, dict) and r["priority"] == "premium"]
premium_results = [
r for r in results if isinstance(r, dict) and r["priority"] == "premium"
]
premium_success = len([r for r in premium_results if r["status"] == "success"])
standard_results = [r for r in results if isinstance(r, dict) and r["priority"] == "standard"]
standard_results = [
r for r in results if isinstance(r, dict) and r["priority"] == "standard"
]
standard_success = len([r for r in standard_results if r["status"] == "success"])
# Premium users should have higher success rate due to priority
premium_success_rate = premium_success / len(premium_results) if premium_results else 0
standard_success_rate = standard_success / len(standard_results) if standard_results else 0
assert premium_success_rate >= 0.9, f"Premium success rate should be >= 90%, got {premium_success_rate:.2%}"
assert standard_success_rate >= 0.5, f"Standard success rate should be >= 50%, got {standard_success_rate:.2%}"
assert premium_success_rate > standard_success_rate, "Premium should have higher success rate than standard"
premium_success_rate = (
premium_success / len(premium_results) if premium_results else 0
)
standard_success_rate = (
standard_success / len(standard_results) if standard_results else 0
)
assert (
premium_success_rate >= 0.9
), f"Premium success rate should be >= 90%, got {premium_success_rate:.2%}"
assert (
standard_success_rate >= 0.5
), f"Standard success rate should be >= 50%, got {standard_success_rate:.2%}"
assert (
premium_success_rate > standard_success_rate
), "Premium should have higher success rate than standard"
total_duration = end_time - start_time
print(f"✅ Processed 50 concurrent pre-call hooks in {total_duration:.2f}s")
print(f" - Premium users: {premium_success}/{len(premium_results)} success ({premium_success_rate:.1%})")
print(f" - Standard users: {standard_success}/{len(standard_results)} success ({standard_success_rate:.1%})")
print(
f" - Premium users: {premium_success}/{len(premium_results)} success ({premium_success_rate:.1%})"
)
print(
f" - Standard users: {standard_success}/{len(standard_results)} success ({standard_success_rate:.1%})"
)
print(f" - Total successful: {successful_count}/50 ({successful_count/50:.1%})")
print(f" - Priority system working: Premium > Standard success rates")

View file

@ -1247,6 +1247,156 @@ class TestCustomUISSO:
assert result.status_code == 303
class TestCLIKeyRegenerationFlow:
"""Test the end-to-end CLI key regeneration flow"""
@pytest.mark.asyncio
async def test_cli_sso_callback_regenerate_existing_key(self):
"""Test CLI SSO callback regenerating an existing key"""
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
# Mock request
mock_request = MagicMock(spec=Request)
# Test data
existing_key = "sk-existing-key-123"
new_key = "sk-new-key-456"
# Mock the regenerate helper function
with patch("litellm.proxy.management_endpoints.ui_sso._regenerate_cli_key") as mock_regenerate, \
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \
patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="<html>Success</html>"):
# Act
result = await cli_sso_callback(
request=mock_request,
key=new_key,
existing_key=existing_key
)
# Assert
mock_regenerate.assert_called_once_with(existing_key, new_key)
assert result.status_code == 200
assert "Success" in result.body.decode()
@pytest.mark.asyncio
async def test_cli_sso_callback_create_new_key(self):
"""Test CLI SSO callback creating a new key when no existing key provided"""
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
# Mock request
mock_request = MagicMock(spec=Request)
# Test data
new_key = "sk-new-key-789"
# Mock the create helper function
with patch("litellm.proxy.management_endpoints.ui_sso._create_new_cli_key") as mock_create, \
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \
patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="<html>Success</html>"):
# Act
result = await cli_sso_callback(
request=mock_request,
key=new_key,
existing_key=None
)
# Assert
mock_create.assert_called_once_with(new_key)
assert result.status_code == 200
assert "Success" in result.body.decode()
@pytest.mark.asyncio
async def test_auth_callback_routes_to_cli_with_existing_key(self):
"""Test that auth_callback properly routes CLI requests and preserves existing_key parameter"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
from litellm.proxy.management_endpoints.ui_sso import auth_callback
# Mock request with existing_key query parameter
mock_request = MagicMock(spec=Request)
mock_request.query_params.get.return_value = "sk-existing-cli-key-123"
# CLI state
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-session-key-456"
# Mock the CLI callback
with patch("litellm.proxy.management_endpoints.ui_sso.cli_sso_callback") as mock_cli_callback:
mock_cli_callback.return_value = MagicMock()
# Act
await auth_callback(request=mock_request, state=cli_state)
# Assert
mock_cli_callback.assert_called_once_with(
mock_request,
key="sk-new-session-key-456",
existing_key="sk-existing-cli-key-123"
)
def test_get_redirect_url_preserves_existing_key(self):
"""Test that redirect URL generation preserves existing_key parameter"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock request
mock_request = MagicMock()
mock_request.base_url = "https://test.litellm.ai/"
with patch("litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"):
# Test with existing_key
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=mock_request,
sso_callback_route="sso/callback",
existing_key="sk-existing-123"
)
assert "https://test.litellm.ai/sso/callback?existing_key=sk-existing-123" == redirect_url
def test_get_redirect_url_without_existing_key(self):
"""Test that redirect URL generation works without existing_key parameter"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock request
mock_request = MagicMock()
mock_request.base_url = "https://test.litellm.ai/"
with patch("litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"):
# Test without existing_key
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=mock_request,
sso_callback_route="sso/callback"
)
assert "https://test.litellm.ai/sso/callback" == redirect_url
@pytest.mark.asyncio
async def test_cli_sso_callback_regenerate_vs_create_flow(self):
"""Test CLI SSO callback calls regenerate_key_fn when existing_key provided, generate_key_helper_fn when not"""
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
mock_request = MagicMock(spec=Request)
with patch("litellm.proxy.management_endpoints.key_management_endpoints.regenerate_key_fn") as mock_regenerate, \
patch("litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn") as mock_generate, \
patch("litellm.proxy._types.UserAPIKeyAuth.get_litellm_cli_user_api_key_auth"), \
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \
patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="<html>Success</html>"):
# Test regeneration path
await cli_sso_callback(mock_request, key="sk-new-123", existing_key="sk-existing-456")
mock_regenerate.assert_called_once()
mock_generate.assert_not_called()
# Reset mocks
mock_regenerate.reset_mock()
mock_generate.reset_mock()
# Test creation path
await cli_sso_callback(mock_request, key="sk-new-789", existing_key=None)
mock_regenerate.assert_not_called()
mock_generate.assert_called_once()
class TestProcessSSOJWTAccessToken:
"""Test the process_sso_jwt_access_token helper function"""

View file

@ -19,6 +19,8 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
BaseOpenAIPassThroughHandler,
RouteChecks,
create_pass_through_route,
llm_passthrough_factory_proxy_route,
vllm_proxy_route,
vertex_discovery_proxy_route,
vertex_proxy_route,
bedrock_llm_proxy_route,
@ -914,3 +916,119 @@ class TestBedrockLLMProxyRoute:
# For regular models, model should be just the model ID
assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0"
assert result == "success"
class TestLLMPassthroughFactoryProxyRoute:
@pytest.mark.asyncio
async def test_llm_passthrough_factory_proxy_route_success(self):
from litellm.types.utils import LlmProviders
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.json = AsyncMock(return_value={"stream": False})
mock_fastapi_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock()
with patch(
"litellm.utils.ProviderConfigManager.get_provider_model_info"
) as mock_get_provider, patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials"
) as mock_get_creds, patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route"
) as mock_create_route:
mock_provider_config = MagicMock()
mock_provider_config.get_api_base.return_value = "https://example.com/v1"
mock_provider_config.validate_environment.return_value = {
"x-api-key": "dummy"
}
mock_get_provider.return_value = mock_provider_config
mock_get_creds.return_value = "dummy"
mock_endpoint_func = AsyncMock(return_value="success")
mock_create_route.return_value = mock_endpoint_func
result = await llm_passthrough_factory_proxy_route(
custom_llm_provider=LlmProviders.VLLM,
endpoint="/chat/completions",
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
)
assert result == "success"
mock_get_provider.assert_called_once_with(
provider=litellm.LlmProviders(LlmProviders.VLLM), model=None
)
mock_get_creds.assert_called_once_with(
custom_llm_provider=LlmProviders.VLLM, region_name=None
)
mock_create_route.assert_called_once_with(
endpoint="/chat/completions",
target="https://example.com/v1/chat/completions",
custom_headers={"x-api-key": "dummy"},
)
mock_endpoint_func.assert_awaited_once()
class TestVLLMProxyRoute:
@pytest.mark.asyncio
@patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
return_value={"model": "router-model", "stream": False},
)
@patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
return_value=True,
)
@patch("litellm.proxy.proxy_server.llm_router")
async def test_vllm_proxy_route_with_router_model(
self, mock_llm_router, mock_is_router, mock_get_body
):
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.headers = {"content-type": "application/json"}
mock_request.query_params = {}
mock_fastapi_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock()
mock_llm_router.allm_passthrough_route = AsyncMock(
return_value=httpx.Response(200, json={"response": "success"})
)
await vllm_proxy_route(
endpoint="/chat/completions",
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
)
mock_is_router.assert_called_once()
mock_llm_router.allm_passthrough_route.assert_awaited_once()
@pytest.mark.asyncio
@patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
return_value={"model": "other-model"},
)
@patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
return_value=False,
)
@patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.llm_passthrough_factory_proxy_route"
)
async def test_vllm_proxy_route_fallback_to_factory(
self, mock_factory_route, mock_is_router, mock_get_body
):
mock_request = MagicMock(spec=Request)
mock_fastapi_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock()
mock_factory_route.return_value = "factory_success"
result = await vllm_proxy_route(
endpoint="/chat/completions",
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
)
assert result == "factory_success"
mock_factory_route.assert_awaited_once()

View file

@ -2409,3 +2409,53 @@ def test_model_info_for_vertex_ai_deepseek_model():
assert model_info["input_cost_per_token"] is not None
assert model_info["output_cost_per_token"] is not None
print("vertex deepseek model info", model_info)
class TestGetValidModelsWithCLI:
"""Test get_valid_models function as used in CLI token usage"""
def test_get_valid_models_with_cli_pattern(self):
"""Test get_valid_models with litellm_proxy provider and CLI token pattern"""
# Mock the HTTP request that get_valid_models makes to the proxy
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"data": [
{"id": "gpt-3.5-turbo", "object": "model"},
{"id": "gpt-4", "object": "model"},
{"id": "litellm_proxy/gemini/gemini-2.5-flash", "object": "model"},
{"id": "claude-3-sonnet", "object": "model"}
]
}
with patch('requests.get', return_value=mock_response) as mock_get:
# Test the exact pattern used in cli_token_usage.py
result = litellm.get_valid_models(
check_provider_endpoint=True,
custom_llm_provider="litellm_proxy",
api_key="sk-test-cli-key-123",
api_base="http://localhost:4000/"
)
# Verify the function returns a list of model names
assert isinstance(result, list)
assert len(result) == 4
assert "gpt-3.5-turbo" in result
assert "gpt-4" in result
assert "litellm_proxy/gemini/gemini-2.5-flash" in result
assert "claude-3-sonnet" in result
# Verify the HTTP request was made with correct parameters
mock_get.assert_called_once()
call_args = mock_get.call_args
# Check that the request was made to the correct endpoint
assert "http://localhost:4000/" in call_args[0][0]
assert "/v1/models" in call_args[0][0]
# Check that the API key was included in headers
assert "headers" in call_args.kwargs
headers = call_args.kwargs["headers"]
assert "Authorization" in headers
assert "Bearer sk-test-cli-key-123" == headers["Authorization"]