mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'main' into akshoop/fastuuid-dep-make-optional
This commit is contained in:
commit
c659b7e587
83 changed files with 4199 additions and 636 deletions
48
.github/workflows/test-mcp.yml
vendored
Normal file
48
.github/workflows/test-mcp.yml
vendored
Normal 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
|
||||
62
cookbook/litellm_proxy_server/cli_token_usage.py
Normal file
62
cookbook/litellm_proxy_server/cli_token_usage.py
Normal 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")
|
||||
|
|
@ -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>
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
260
docs/my-website/docs/providers/azure_ai_img_edit.md
Normal file
260
docs/my-website/docs/providers/azure_ai_img_edit.md
Normal 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)
|
||||
|
|
@ -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">
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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' }} />
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
BIN
docs/my-website/img/default_user_settings_admin_ui.png
Normal file
BIN
docs/my-website/img/default_user_settings_admin_ui.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 234 KiB |
|
|
@ -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"],
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
Enterprise internal user management endpoints
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
58
litellm/litellm_core_utils/cli_token_utils.py
Normal file
58
litellm/litellm_core_utils/cli_token_utils.py
Normal 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
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
15
litellm/llms/azure_ai/image_edit/__init__.py
Normal file
15
litellm/llms/azure_ai/image_edit/__init__.py
Normal 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.")
|
||||
99
litellm/llms/azure_ai/image_edit/transformation.py
Normal file
99
litellm/llms/azure_ai/image_edit/transformation.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
0
litellm/llms/wandb/__init__.py
Normal file
0
litellm/llms/wandb/__init__.py
Normal file
0
litellm/llms/wandb/chat/__init__.py
Normal file
0
litellm/llms/wandb/chat/__init__.py
Normal file
27
litellm/llms/wandb/chat/transformation.py
Normal file
27
litellm/llms/wandb/chat/transformation.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
179
litellm/proxy/client/cli/commands/teams.py
Normal file
179
litellm/proxy/client/cli/commands/teams.py
Normal 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()
|
||||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
146
litellm/proxy/client/teams.py
Normal file
146
litellm/proxy/client/teams.py
Normal 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()
|
||||
|
|
@ -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 {}
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
146
tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py
Normal file
146
tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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'
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue