diff --git a/.github/workflows/test-mcp.yml b/.github/workflows/test-mcp.yml new file mode 100644 index 00000000000..2da6980951a --- /dev/null +++ b/.github/workflows/test-mcp.yml @@ -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 diff --git a/cookbook/litellm_proxy_server/cli_token_usage.py b/cookbook/litellm_proxy_server/cli_token_usage.py new file mode 100644 index 00000000000..6ee5555695e --- /dev/null +++ b/cookbook/litellm_proxy_server/cli_token_usage.py @@ -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") diff --git a/docs/my-website/docs/completion/provider_specific_params.md b/docs/my-website/docs/completion/provider_specific_params.md index 772ca13e293..250b410c9c4 100644 --- a/docs/my-website/docs/completion/provider_specific_params.md +++ b/docs/my-website/docs/completion/provider_specific_params.md @@ -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( ``` - -``` \ No newline at end of file + \ No newline at end of file diff --git a/docs/my-website/docs/completion/usage.md b/docs/my-website/docs/completion/usage.md index 2a9eab941ea..c388e5bfee1 100644 --- a/docs/my-website/docs/completion/usage.md +++ b/docs/my-website/docs/completion/usage.md @@ -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 diff --git a/docs/my-website/docs/enterprise.md b/docs/my-website/docs/enterprise.md index 9101d8e3751..cc3466fc103 100644 --- a/docs/my-website/docs/enterprise.md +++ b/docs/my-website/docs/enterprise.md @@ -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 diff --git a/docs/my-website/docs/fine_tuning.md b/docs/my-website/docs/fine_tuning.md index f9a9297e062..f3f955cb01d 100644 --- a/docs/my-website/docs/fine_tuning.md +++ b/docs/my-website/docs/fine_tuning.md @@ -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 | diff --git a/docs/my-website/docs/getting_started.md b/docs/my-website/docs/getting_started.md index 15ee00a7273..6b2c1fd531e 100644 --- a/docs/my-website/docs/getting_started.md +++ b/docs/my-website/docs/getting_started.md @@ -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 diff --git a/docs/my-website/docs/image_edits.md b/docs/my-website/docs/image_edits.md index 246e1c70f0e..84dddd5e4ad 100644 --- a/docs/my-website/docs/image_edits.md +++ b/docs/my-website/docs/image_edits.md @@ -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 diff --git a/docs/my-website/docs/image_generation.md b/docs/my-website/docs/image_generation.md index 7e7ff9922d6..8cd5803aa6c 100644 --- a/docs/my-website/docs/image_generation.md +++ b/docs/my-website/docs/image_generation.md @@ -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) | diff --git a/docs/my-website/docs/index.md b/docs/my-website/docs/index.md index 3f5e1b479c3..11d2963b7a3 100644 --- a/docs/my-website/docs/index.md +++ b/docs/my-website/docs/index.md @@ -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 diff --git a/docs/my-website/docs/moderation.md b/docs/my-website/docs/moderation.md index 95fe8b2856d..f9c2810bc8a 100644 --- a/docs/my-website/docs/moderation.md +++ b/docs/my-website/docs/moderation.md @@ -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 | diff --git a/docs/my-website/docs/observability/callbacks.md b/docs/my-website/docs/observability/callbacks.md index 040d83697d3..b752bdc2764 100644 --- a/docs/my-website/docs/observability/callbacks.md +++ b/docs/my-website/docs/observability/callbacks.md @@ -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 diff --git a/docs/my-website/docs/observability/custom_callback.md b/docs/my-website/docs/observability/custom_callback.md index c206c23d0f4..cfe97ca42c0 100644 --- a/docs/my-website/docs/observability/custom_callback.md +++ b/docs/my-website/docs/observability/custom_callback.md @@ -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. diff --git a/docs/my-website/docs/providers/azure_ai_img_edit.md b/docs/my-website/docs/providers/azure_ai_img_edit.md new file mode 100644 index 00000000000..0d5408f0af4 --- /dev/null +++ b/docs/my-website/docs/providers/azure_ai_img_edit.md @@ -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 + + + + +```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) +``` + + + + + +```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()) +``` + + + + + +```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) +``` + + + + +### 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 + + + + +```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) +``` + + + + + +```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) +``` + + + + + +```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"' +``` + + + + +## 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) \ No newline at end of file diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 86e9ac5e3e6..fe996099145 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -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 `` 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": "" + } + ] +} +``` + +**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. + +--- + diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md index cb90b7434e7..260cc55c2e9 100644 --- a/docs/my-website/docs/providers/vertex.md +++ b/docs/my-website/docs/providers/vertex.md @@ -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 ``` @@ -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, ) diff --git a/docs/my-website/docs/proxy/caching.md b/docs/my-website/docs/proxy/caching.md index 1fb7385f689..617609cf08a 100644 --- a/docs/my-website/docs/proxy/caching.md +++ b/docs/my-website/docs/proxy/caching.md @@ -958,6 +958,19 @@ curl http://localhost:4000/v1/chat/completions \ + +## 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 diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 974e95a07bd..f70701886b5 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -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"}] diff --git a/docs/my-website/docs/proxy/custom_sso.md b/docs/my-website/docs/proxy/custom_sso.md index 8e869a11393..bbd7f41bee1 100644 --- a/docs/my-website/docs/proxy/custom_sso.md +++ b/docs/my-website/docs/proxy/custom_sso.md @@ -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 diff --git a/docs/my-website/docs/proxy/db_deadlocks.md b/docs/my-website/docs/proxy/db_deadlocks.md index 0eee928fa64..ef9d31d6232 100644 --- a/docs/my-website/docs/proxy/db_deadlocks.md +++ b/docs/my-website/docs/proxy/db_deadlocks.md @@ -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. + diff --git a/docs/my-website/docs/proxy/guardrails/bedrock.md b/docs/my-website/docs/proxy/guardrails/bedrock.md index 6725acf1f25..4a1a0a246f8 100644 --- a/docs/my-website/docs/proxy/guardrails/bedrock.md +++ b/docs/my-website/docs/proxy/guardrails/bedrock.md @@ -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 diff --git a/docs/my-website/docs/proxy/load_balancing.md b/docs/my-website/docs/proxy/load_balancing.md index bcbc4e93651..54c917bbbca 100644 --- a/docs/my-website/docs/proxy/load_balancing.md +++ b/docs/my-website/docs/proxy/load_balancing.md @@ -172,6 +172,9 @@ router_settings: redis_host: 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 diff --git a/docs/my-website/docs/proxy/self_serve.md b/docs/my-website/docs/proxy/self_serve.md index dff55a8ac04..b54344c1d05 100644 --- a/docs/my-website/docs/proxy/self_serve.md +++ b/docs/my-website/docs/proxy/self_serve.md @@ -227,7 +227,7 @@ export PROXY_LOGOUT_URL="https://www.google.com" -### 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: + + + This budget only applies to personal keys created by that user - seen under `Default Team` on the UI. diff --git a/docs/my-website/docs/proxy_api.md b/docs/my-website/docs/proxy_api.md index 89bfacbe19f..7612645fb54 100644 --- a/docs/my-website/docs/proxy_api.md +++ b/docs/my-website/docs/proxy_api.md @@ -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 diff --git a/docs/my-website/docs/rerank.md b/docs/my-website/docs/rerank.md index c57eacbb224..cad64718384 100644 --- a/docs/my-website/docs/rerank.md +++ b/docs/my-website/docs/rerank.md @@ -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) | diff --git a/docs/my-website/docs/response_api.md b/docs/my-website/docs/response_api.md index 94d7c73be05..b03e4f8be92 100644 --- a/docs/my-website/docs/response_api.md +++ b/docs/my-website/docs/response_api.md @@ -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 diff --git a/docs/my-website/img/default_user_settings_admin_ui.png b/docs/my-website/img/default_user_settings_admin_ui.png new file mode 100644 index 00000000000..5910154cd51 Binary files /dev/null and b/docs/my-website/img/default_user_settings_admin_ui.png differ diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index fe6dbc27290..fc14c47a157 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -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"], - }, + } ], }, diff --git a/enterprise/litellm_enterprise/integrations/prometheus.py b/enterprise/litellm_enterprise/integrations/prometheus.py index 4a1e8d75435..4451d76bed0 100644 --- a/enterprise/litellm_enterprise/integrations/prometheus.py +++ b/enterprise/litellm_enterprise/integrations/prometheus.py @@ -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 diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py index e60b4d69905..2f53f9e9281 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py @@ -2,7 +2,6 @@ Enterprise internal user management endpoints """ -import os from fastapi import APIRouter, Depends, HTTPException diff --git a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py index 43bdfa3844f..bb4b546b8d3 100644 --- a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py +++ b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py @@ -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 diff --git a/litellm/__init__.py b/litellm/__init__.py index 038787f5cce..7a68aa3a8d6 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 diff --git a/litellm/constants.py b/litellm/constants.py index 9b44613b855..005eb2bb6d0 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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" diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 5d8f5faadf1..36a562b3574 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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: diff --git a/litellm/litellm_core_utils/cli_token_utils.py b/litellm/litellm_core_utils/cli_token_utils.py new file mode 100644 index 00000000000..2aedb1c19d2 --- /dev/null +++ b/litellm/litellm_core_utils/cli_token_utils.py @@ -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 diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index 414ccb7ab83..69c996d8139 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -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)) diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index d77f53bd798..06e650f938d 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -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": diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 2aaeed40bea..64986970d00 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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( diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 60a31198415..626a3f3625f 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -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, diff --git a/litellm/llms/azure_ai/image_edit/__init__.py b/litellm/llms/azure_ai/image_edit/__init__.py new file mode 100644 index 00000000000..e0e57bec403 --- /dev/null +++ b/litellm/llms/azure_ai/image_edit/__init__.py @@ -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.") diff --git a/litellm/llms/azure_ai/image_edit/transformation.py b/litellm/llms/azure_ai/image_edit/transformation.py new file mode 100644 index 00000000000..47f612912ce --- /dev/null +++ b/litellm/llms/azure_ai/image_edit/transformation.py @@ -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) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index e8094444330..a61cfa39e47 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -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 diff --git a/litellm/llms/openai/cost_calculation.py b/litellm/llms/openai/cost_calculation.py index 304c444e37a..229f75f2657 100644 --- a/litellm/llms/openai/cost_calculation.py +++ b/litellm/llms/openai/cost_calculation.py @@ -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 diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index 8dca3e1de25..e2ed0daafe4 100644 --- a/litellm/llms/vllm/common_utils.py +++ b/litellm/llms/vllm/common_utils.py @@ -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) diff --git a/litellm/llms/wandb/__init__.py b/litellm/llms/wandb/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/wandb/chat/__init__.py b/litellm/llms/wandb/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/wandb/chat/transformation.py b/litellm/llms/wandb/chat/transformation.py new file mode 100644 index 00000000000..1cb2ab492bc --- /dev/null +++ b/litellm/llms/wandb/chat/transformation.py @@ -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 diff --git a/litellm/main.py b/litellm/main.py index 908c364220c..5493c7e34e3 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 43ea3af320f..578523abff4 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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 } -} \ No newline at end of file +} diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d0eadb36ba3..1e7840d95b3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 2a9174717d1..399b79b4f7c 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2cf91c84dcc..8558229cf0b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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] diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b0a4e71e23a..9d5298c30f8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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): diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 7d89d39ed70..9be74268059 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -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 \ No newline at end of file diff --git a/litellm/proxy/client/cli/commands/teams.py b/litellm/proxy/client/cli/commands/teams.py new file mode 100644 index 00000000000..d4e05c890dc --- /dev/null +++ b/litellm/proxy/client/cli/commands/teams.py @@ -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() diff --git a/litellm/proxy/client/cli/interface.py b/litellm/proxy/client/cli/interface.py index d3f3a24eb45..78aeb442351 100644 --- a/litellm/proxy/client/cli/interface.py +++ b/litellm/proxy/client/cli/interface.py @@ -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"), diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index fb4a37c3a17..eab9b31482a 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -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) diff --git a/litellm/proxy/client/client.py b/litellm/proxy/client/client.py index 93e877d1563..8ed9a4ff89f 100644 --- a/litellm/proxy/client/client.py +++ b/litellm/proxy/client/client.py @@ -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) diff --git a/litellm/proxy/client/keys.py b/litellm/proxy/client/keys.py index 8c62eb8e4d0..fc307648a2f 100644 --- a/litellm/proxy/client/keys.py +++ b/litellm/proxy/client/keys.py @@ -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. diff --git a/litellm/proxy/client/teams.py b/litellm/proxy/client/teams.py new file mode 100644 index 00000000000..61ddbe6adae --- /dev/null +++ b/litellm/proxy/client/teams.py @@ -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() diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 74da9992631..fd84eeda081 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -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 {} \ No newline at end of file diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 31d7d70d5f3..480be4a651b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -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( diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index e6fdcef2805..3828f7318c1 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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 = [] diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 56ee599325a..5e171af5252 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -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 diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index fdd2e51dfde..addf2443d36 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -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 diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index eb1eb3250ba..bd09ef9c199 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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): diff --git a/litellm/types/proxy/ui_sso.py b/litellm/types/proxy/ui_sso.py index 0c036274c79..04523e88b1b 100644 --- a/litellm/types/proxy/ui_sso.py +++ b/litellm/types/proxy/ui_sso.py @@ -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] \ No newline at end of file diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 2ab8f844615..ac4a6b116ec 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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, diff --git a/litellm/utils.py b/litellm/utils.py index d4df9b206b5..0721b023d2b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 43ea3af320f..5ca0bf08ba9 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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 } -} \ No newline at end of file +} diff --git a/security.md b/security.md index d126dabcc67..2da073661c5 100644 --- a/security.md +++ b/security.md @@ -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. diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index be3133ca37e..0c7a6b2942b 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -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}" diff --git a/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py b/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py new file mode 100644 index 00000000000..fe78b6ecfe3 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py @@ -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 diff --git a/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py new file mode 100644 index 00000000000..249d19eceb3 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py b/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py new file mode 100644 index 00000000000..ef7bb0e44f0 --- /dev/null +++ b/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_custom_fields.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_custom_fields.py new file mode 100644 index 00000000000..d2fa7853e78 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_custom_fields.py @@ -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" \ No newline at end of file diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 8fa9964b1fc..754487d11e7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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"] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 3237de37636..e3bb085d5f4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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__]) diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index 611fde7767e..5ef96e4f9af 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -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' diff --git a/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py index a7b5e4f1b1d..05e0d4287a3 100644 --- a/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py @@ -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") diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index d72e9f7caa5..80aebc98497 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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="Success"): + + # 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="Success"): + + # 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="Success"): + + # 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""" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 1702a317a29..702ae4bd42f 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -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() diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index fc7088f4b62..618a2902540 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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"]