diff --git a/.circleci/config.yml b/.circleci/config.yml index 476f138b1d4..ebbd986bf2a 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -2038,6 +2038,39 @@ jobs: - run: python ./tests/code_coverage_tests/memory_test.py - run: helm lint ./deploy/charts/litellm-helm + memory_leak_tests: + docker: + - image: cimg/python:3.11 + auth: + username: ${DOCKERHUB_USERNAME} + password: ${DOCKERHUB_PASSWORD} + working_directory: ~/project + resource_class: large + steps: + - setup_litellm_test_deps + - run: + name: Install Memory Test Dependencies + command: | + pip install "psutil>=5.9.0" + pip install "fastapi>=0.100.0" + pip install "httpx>=0.24.0" + pip install "uvicorn>=0.23.0" + - run: + name: Run Linear Memory Growth Tests + command: | + echo "Running memory leak tests individually to avoid baseline drift..." + echo "Running test_memory_baseline_1k..." + python -m pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_1k -v -s --tb=short + echo "Running test_memory_baseline_2k..." + python -m pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_2k -v -s --tb=short + echo "Running test_memory_baseline_4k..." + python -m pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_4k -v -s --tb=short + echo "Running test_memory_baseline_10k..." + python -m pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_10k -v -s --tb=short + echo "Running test_memory_baseline_30k..." + python -m pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_30k -v -s --tb=short + no_output_timeout: 60m + db_migration_disable_update_check: machine: image: ubuntu-2204:2023.10.1 @@ -3557,12 +3590,34 @@ jobs: name: Install Playwright Browsers command: | npx playwright install + - run: + name: Install Neon CLI + command: | + npm i -g neonctl + - run: + name: Create Neon branch + command: | + export EXPIRES_AT=$(date -u -d "+3 hours" +"%Y-%m-%dT%H:%M:%SZ") + echo "Expires at: $EXPIRES_AT" + neon branches create \ + --project-id $NEON_PROJECT_ID \ + --name preview/commit-${CIRCLE_SHA1:0:7} \ + --expires-at $EXPIRES_AT \ + --parent br-fancy-paper-ad1olsb3 \ + --api-key $NEON_API_KEY || true - run: name: Run Docker container command: | + E2E_UI_TEST_DATABASE_URL=$(neon connection-string \ + --project-id $NEON_PROJECT_ID \ + --api-key $NEON_API_KEY \ + --branch preview/commit-${CIRCLE_SHA1:0:7} \ + --database-name yuneng-trial-db \ + --role neondb_owner) + echo $E2E_UI_TEST_DATABASE_URL docker run -d \ -p 4000:4000 \ - -e DATABASE_URL=$SMALL_DATABASE_URL \ + -e DATABASE_URL=$E2E_UI_TEST_DATABASE_URL \ -e LITELLM_MASTER_KEY="sk-1234" \ -e OPENAI_API_KEY=$OPENAI_API_KEY \ -e UI_USERNAME="admin" \ @@ -3765,6 +3820,12 @@ workflows: only: - main - /litellm_.*/ + - memory_leak_tests: + filters: + branches: + only: + - main + - /litellm_.*/ - ui_build: filters: branches: @@ -3792,6 +3853,7 @@ workflows: - main - /litellm_.*/ - e2e_ui_testing: + context: e2e_ui_tests requires: - ui_build - build_docker_database_image diff --git a/.github/workflows/label-component.yml b/.github/workflows/label-component.yml index 9a547c162a6..76b8316790c 100644 --- a/.github/workflows/label-component.yml +++ b/.github/workflows/label-component.yml @@ -11,134 +11,72 @@ jobs: permissions: issues: write steps: - - name: Add SDK label - if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nSDK (litellm Python package)') + - name: Add component labels uses: actions/github-script@v7 with: github-token: ${{ secrets.GITHUB_TOKEN }} script: | - const labelName = 'sdk'; - try { - await github.rest.issues.getLabel({ - owner: context.repo.owner, - repo: context.repo.repo, - name: labelName - }); - } catch (error) { - if (error.status === 404) { - await github.rest.issues.createLabel({ - owner: context.repo.owner, - repo: context.repo.repo, - name: labelName, - color: '0E7C86', - description: 'Issues related to the litellm Python SDK' - }); - } else { - throw error; - } - } - await github.rest.issues.addLabels({ - owner: context.repo.owner, - repo: context.repo.repo, - issue_number: context.issue.number, - labels: [labelName] - }); + const body = context.payload.issue.body; + if (!body) return; - - name: Add Proxy label - if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nProxy') - uses: actions/github-script@v7 - with: - github-token: ${{ secrets.GITHUB_TOKEN }} - script: | - const labelName = 'proxy'; - try { - await github.rest.issues.getLabel({ - owner: context.repo.owner, - repo: context.repo.repo, - name: labelName - }); - } catch (error) { - if (error.status === 404) { - await github.rest.issues.createLabel({ - owner: context.repo.owner, - repo: context.repo.repo, - name: labelName, - color: '5319E7', - description: 'Issues related to the LiteLLM Proxy' - }); - } else { - throw error; + // Define component mappings with regex patterns that handle flexible whitespace + const components = [ + { + pattern: /What part of LiteLLM is this about\?\s*SDK \(litellm Python package\)/, + label: 'sdk', + color: '0E7C86', + description: 'Issues related to the litellm Python SDK' + }, + { + pattern: /What part of LiteLLM is this about\?\s*Proxy/, + label: 'proxy', + color: '5319E7', + description: 'Issues related to the LiteLLM Proxy' + }, + { + pattern: /What part of LiteLLM is this about\?\s*UI Dashboard/, + label: 'ui-dashboard', + color: 'D876E3', + description: 'Issues related to the LiteLLM UI Dashboard' + }, + { + pattern: /What part of LiteLLM is this about\?\s*Docs/, + label: 'docs', + color: 'FBCA04', + description: 'Issues related to LiteLLM documentation' } - } - await github.rest.issues.addLabels({ - owner: context.repo.owner, - repo: context.repo.repo, - issue_number: context.issue.number, - labels: [labelName] - }); + ]; - - name: Add UI Dashboard label - if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nUI Dashboard') - uses: actions/github-script@v7 - with: - github-token: ${{ secrets.GITHUB_TOKEN }} - script: | - const labelName = 'ui-dashboard'; - try { - await github.rest.issues.getLabel({ - owner: context.repo.owner, - repo: context.repo.repo, - name: labelName - }); - } catch (error) { - if (error.status === 404) { - await github.rest.issues.createLabel({ - owner: context.repo.owner, - repo: context.repo.repo, - name: labelName, - color: 'D876E3', - description: 'Issues related to the LiteLLM UI Dashboard' - }); - } else { - throw error; - } - } - await github.rest.issues.addLabels({ - owner: context.repo.owner, - repo: context.repo.repo, - issue_number: context.issue.number, - labels: [labelName] - }); + // Find matching component + for (const component of components) { + if (component.pattern.test(body)) { + // Ensure label exists + try { + await github.rest.issues.getLabel({ + owner: context.repo.owner, + repo: context.repo.repo, + name: component.label + }); + } catch (error) { + if (error.status === 404) { + await github.rest.issues.createLabel({ + owner: context.repo.owner, + repo: context.repo.repo, + name: component.label, + color: component.color, + description: component.description + }); + } + } - - name: Add Docs label - if: contains(github.event.issue.body, 'What part of LiteLLM is this about?\n\nDocs') - uses: actions/github-script@v7 - with: - github-token: ${{ secrets.GITHUB_TOKEN }} - script: | - const labelName = 'docs'; - try { - await github.rest.issues.getLabel({ - owner: context.repo.owner, - repo: context.repo.repo, - name: labelName - }); - } catch (error) { - if (error.status === 404) { - await github.rest.issues.createLabel({ + // Add label to issue + await github.rest.issues.addLabels({ owner: context.repo.owner, repo: context.repo.repo, - name: labelName, - color: 'FBCA04', - description: 'Issues related to LiteLLM documentation' + issue_number: context.issue.number, + labels: [component.label] }); - } else { - throw error; + + break; } } - await github.rest.issues.addLabels({ - owner: context.repo.owner, - repo: context.repo.repo, - issue_number: context.issue.number, - labels: [labelName] - }); diff --git a/.github/workflows/publish-migrations.yml b/.github/workflows/publish-migrations.yml index 8e5a67bcf85..a5187cb2f55 100644 --- a/.github/workflows/publish-migrations.yml +++ b/.github/workflows/publish-migrations.yml @@ -13,6 +13,7 @@ on: jobs: publish-migrations: + if: github.repository == 'BerriAI/litellm' runs-on: ubuntu-latest services: postgres: diff --git a/README.md b/README.md index a020bd80898..75a23faa5c1 100644 --- a/README.md +++ b/README.md @@ -262,6 +262,7 @@ Support for more providers. Missing a provider or LLM Platform, raise a [feature | Provider | `/chat/completions` | `/messages` | `/responses` | `/embeddings` | `/image/generations` | `/audio/transcriptions` | `/audio/speech` | `/moderations` | `/batches` | `/rerank` | |-------------------------------------------------------------------------------------|---------------------|-------------|--------------|---------------|----------------------|-------------------------|-----------------|----------------|-----------|-----------| +| [Abliteration (`abliteration`)](https://docs.litellm.ai/docs/providers/abliteration) | ✅ | | | | | | | | | | | [AI/ML API (`aiml`)](https://docs.litellm.ai/docs/providers/aiml) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | | | [AI21 (`ai21`)](https://docs.litellm.ai/docs/providers/ai21) | ✅ | ✅ | ✅ | | | | | | | | | [AI21 Chat (`ai21_chat`)](https://docs.litellm.ai/docs/providers/ai21) | ✅ | ✅ | ✅ | | | | | | | | @@ -455,4 +456,3 @@ All these checks must pass before your PR can be merged. - diff --git a/docs/my-website/docs/mcp_control.md b/docs/my-website/docs/mcp_control.md index a7d66a6b7fc..96c71ef9278 100644 --- a/docs/my-website/docs/mcp_control.md +++ b/docs/my-website/docs/mcp_control.md @@ -649,3 +649,16 @@ general_settings: ``` This is useful when you want discoverability for MCP offerings without granting additional execution privileges. + + +## Publish MCP Registry + +If you want other systems—for example external agent frameworks such as MCP-capable IDEs running outside your network—to automatically discover the MCP servers hosted on LiteLLM, you can expose a Model Context Protocol Registry endpoint. This registry lists the built-in LiteLLM MCP server and every server you have configured, using the [official MCP Registry spec](https://github.com/modelcontextprotocol/registry). + +1. Set `enable_mcp_registry: true` under `general_settings` in your proxy config (or DB settings) and restart the proxy. +2. LiteLLM will serve the registry at `GET /v1/mcp/registry.json`. +3. Each entry points to either `/mcp` (built-in server) or `/{mcp_server_name}/mcp` for your custom servers, so clients can connect directly using the advertised Streamable HTTP URL. + +:::note Permissions still apply +The registry only advertises server URLs. Actual access control is still enforced by LiteLLM when the client connects to `/mcp` or `/{server}/mcp`, so publishing the registry does not bypass per-key permissions. +::: diff --git a/docs/my-website/docs/observability/focus.md b/docs/my-website/docs/observability/focus.md new file mode 100644 index 00000000000..c282f4a220c --- /dev/null +++ b/docs/my-website/docs/observability/focus.md @@ -0,0 +1,93 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Focus Export (Experimental) + +:::caution Experimental feature +Focus Format export is under active development and currently considered experimental. +Interfaces, schema mappings, and configuration options may change as we iterate based on user feedback. +Please treat this integration as a preview and report any issues or suggestions to help us stabilize and improve the workflow. +::: + +LiteLLM can emit usage data in the [FinOps FOCUS format](https://focus.finops.org/focus-specification/v1-2/) and push artifacts (for example Parquet files) to destinations such as Amazon S3. This enables downstream cost-analysis tooling to ingest a standardised dataset directly from LiteLLM. + +LiteLLM currently conforms to the FinOps FOCUS v1.2 specification when emitting this dataset. + +## Overview + +| Property | Details | +|----------|---------| +| Destination | Export LiteLLM usage data in FOCUS format to managed storage (currently S3) | +| Callback name | `focus` | +| Supported operations | Automatic scheduled export | +| Data format | FOCUS Normalised Dataset (Parquet) | + +## Environment Variables + +### Common settings + +| Variable | Required | Description | +|----------|----------|-------------| +| `FOCUS_PROVIDER` | No | Destination provider (defaults to `s3`). | +| `FOCUS_FORMAT` | No | Output format (currently only `parquet`). | +| `FOCUS_FREQUENCY` | No | Export cadence. Prefer `hourly` or `daily` for production; `interval` is intended for short test loops. Defaults to `hourly`. | +| `FOCUS_CRON_OFFSET` | No | Minute offset used for hourly/daily cron triggers. Defaults to `5`. | +| `FOCUS_INTERVAL_SECONDS` | No | Interval (seconds) when `FOCUS_FREQUENCY="interval"`. | +| `FOCUS_PREFIX` | No | Object key prefix/folder. Defaults to `focus_exports`. | + +### S3 destination + +| Variable | Required | Description | +|----------|----------|-------------| +| `FOCUS_S3_BUCKET_NAME` | Yes | Destination bucket for exported files. | +| `FOCUS_S3_REGION_NAME` | No | AWS region for the bucket. | +| `FOCUS_S3_ENDPOINT_URL` | No | Custom endpoint (useful for S3-compatible storage). | +| `FOCUS_S3_ACCESS_KEY` | Yes | AWS access key for uploads. | +| `FOCUS_S3_SECRET_KEY` | Yes | AWS secret key for uploads. | +| `FOCUS_S3_SESSION_TOKEN` | No | AWS session token if using temporary credentials. | + +## Setup via Config + +### Configure environment variables + +```bash +export FOCUS_PROVIDER="s3" +export FOCUS_PREFIX="focus_exports" + +# S3 example +export FOCUS_S3_BUCKET_NAME="my-litellm-focus-bucket" +export FOCUS_S3_REGION_NAME="us-east-1" +export FOCUS_S3_ACCESS_KEY="AKIA..." +export FOCUS_S3_SECRET_KEY="..." +``` + +### Update LiteLLM config + +```yaml +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: sk-your-key + +litellm_settings: + callbacks: ["focus"] +``` + +### Start the proxy + +```bash +litellm --config /path/to/config.yaml +``` + +During boot LiteLLM registers the Focus logger and a background job that runs according to the configured frequency. + +## Planned Enhancements +- Add "Setup on UI" flow alongside the current configuration-based setup. +- Add GCS / Azure Blob to the Destination options. +- Support CSV output alongside Parquet. + +## Related Links + +- [Focus](https://focus.finops.org/) + diff --git a/docs/my-website/docs/observability/qualifire_integration.md b/docs/my-website/docs/observability/qualifire_integration.md new file mode 100644 index 00000000000..cf866f467bf --- /dev/null +++ b/docs/my-website/docs/observability/qualifire_integration.md @@ -0,0 +1,122 @@ +import Image from '@theme/IdealImage'; + +# Qualifire - LLM Evaluation, Guardrails & Observability + +[Qualifire](https://qualifire.ai/) provides real-time Agentic evaluations, guardrails and observability for production AI applications. + +**Key Features:** + +- **Evaluation** - Systematically assess AI behavior to detect hallucinations, jailbreaks, policy breaches, and other vulnerabilities +- **Guardrails** - Real-time interventions to prevent risks like brand damage, data leaks, and compliance breaches +- **Observability** - Complete tracing and logging for RAG pipelines, chatbots, and AI agents +- **Prompt Management** - Centralized prompt management with versioning and no-code studio + +:::tip + +Looking for Qualifire Guardrails? Check out the [Qualifire Guardrails Integration](../proxy/guardrails/qualifire.md) for real-time content moderation, prompt injection detection, PII checks, and more. + +::: + +## Pre-Requisites + +1. Create an account on [Qualifire](https://app.qualifire.ai/) +2. Get your API key and webhook URL from the Qualifire dashboard + +```bash +pip install litellm +``` + +## Quick Start + +Use just 2 lines of code to instantly log your responses **across all providers** with Qualifire. + +```python +litellm.callbacks = ["qualifire_eval"] +``` + +```python +import litellm +import os + +# Set Qualifire credentials +os.environ["QUALIFIRE_API_KEY"] = "your-qualifire-api-key" +os.environ["QUALIFIRE_WEBHOOK_URL"] = "https://your-qualifire-webhook-url" + +# LLM API Keys +os.environ['OPENAI_API_KEY'] = "your-openai-api-key" + +# Set qualifire_eval as a callback & LiteLLM will send the data to Qualifire +litellm.callbacks = ["qualifire_eval"] + +# OpenAI call +response = litellm.completion( + model="gpt-5", + messages=[ + {"role": "user", "content": "Hi 👋 - i'm openai"} + ] +) +``` + +## Using with LiteLLM Proxy + +1. Setup config.yaml + +```yaml +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + +litellm_settings: + callbacks: ["qualifire_eval"] + +general_settings: + master_key: "sk-1234" + +environment_variables: + QUALIFIRE_API_KEY: "your-qualifire-api-key" + QUALIFIRE_WEBHOOK_URL: "https://app.qualifire.ai/api/v1/webhooks/evaluations" +``` + +2. Start the proxy + +```bash +litellm --config config.yaml +``` + +3. Test it! + +```bash +curl -X POST 'http://0.0.0.0:4000/chat/completions' \ +-H 'Content-Type: application/json' \ +-H 'Authorization: Bearer sk-1234' \ +-d '{ "model": "gpt-4o", "messages": [{"role": "user", "content": "Hi 👋 - i'm openai"}]}' +``` + +## Environment Variables + +| Variable | Description | +| ----------------------- | ------------------------------------------------------ | +| `QUALIFIRE_API_KEY` | Your Qualifire API key for authentication | +| `QUALIFIRE_WEBHOOK_URL` | The Qualifire webhook endpoint URL from your dashboard | + +## What Gets Logged? + +The [LiteLLM Standard Logging Payload](https://docs.litellm.ai/docs/proxy/logging_spec) is sent to your Qualifire endpoint on each successful LLM API call. + +This includes: + +- Request messages and parameters +- Response content and metadata +- Token usage statistics +- Latency metrics +- Model information +- Cost data + +Once data is in Qualifire, you can: + +- Run evaluations to detect hallucinations, toxicity, and policy violations +- Set up guardrails to block or modify responses in real-time +- View traces across your entire AI pipeline +- Track performance and quality metrics over time diff --git a/docs/my-website/docs/providers/abliteration.md b/docs/my-website/docs/providers/abliteration.md new file mode 100644 index 00000000000..a0fc7f39310 --- /dev/null +++ b/docs/my-website/docs/providers/abliteration.md @@ -0,0 +1,109 @@ +# Abliteration + +## Overview + +| Property | Details | +|-------|-------| +| Description | Abliteration provides an OpenAI-compatible `/chat/completions` endpoint. | +| Provider Route on LiteLLM | `abliteration/` | +| Link to Provider Doc | [Abliteration](https://abliteration.ai) | +| Base URL | `https://api.abliteration.ai/v1` | +| Supported Operations | [`/chat/completions`](#sample-usage) | + +
+ +## Required Variables + +```python showLineNumbers title="Environment Variables" +os.environ["ABLITERATION_API_KEY"] = "" # your Abliteration API key +``` + +## Sample Usage + +```python showLineNumbers title="Abliteration Completion" +import os +from litellm import completion + +os.environ["ABLITERATION_API_KEY"] = "" + +response = completion( + model="abliteration/abliterated-model", + messages=[{"role": "user", "content": "Hello from LiteLLM"}], +) + +print(response) +``` + +## Sample Usage - Streaming + +```python showLineNumbers title="Abliteration Streaming Completion" +import os +from litellm import completion + +os.environ["ABLITERATION_API_KEY"] = "" + +response = completion( + model="abliteration/abliterated-model", + messages=[{"role": "user", "content": "Stream a short reply"}], + stream=True, +) + +for chunk in response: + print(chunk) +``` + +## Usage with LiteLLM Proxy Server + +1. Add the model to your proxy config: + +```yaml showLineNumbers title="config.yaml" +model_list: + - model_name: abliteration-chat + litellm_params: + model: abliteration/abliterated-model + api_key: os.environ/ABLITERATION_API_KEY +``` + +2. Start the proxy: + +```bash +litellm --config /path/to/config.yaml +``` + +## Direct API Usage (Bearer Token) + +Use the environment variable as a Bearer token against the OpenAI-compatible endpoint: +`https://api.abliteration.ai/v1/chat/completions`. + +```bash showLineNumbers title="cURL" +export ABLITERATION_API_KEY="" +curl https://api.abliteration.ai/v1/chat/completions \ + -H "Authorization: Bearer ${ABLITERATION_API_KEY}" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "abliterated-model", + "messages": [{"role": "user", "content": "Hello from Abliteration"}] + }' +``` + +```python showLineNumbers title="Python (requests)" +import os +import requests + +api_key = os.environ["ABLITERATION_API_KEY"] + +response = requests.post( + "https://api.abliteration.ai/v1/chat/completions", + headers={ + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + }, + json={ + "model": "abliterated-model", + "messages": [{"role": "user", "content": "Hello from Abliteration"}], + }, + timeout=60, +) + +print(response.json()) +``` diff --git a/docs/my-website/docs/providers/anthropic.md b/docs/my-website/docs/providers/anthropic.md index cae8657f1a0..446d663c5ac 100644 --- a/docs/my-website/docs/providers/anthropic.md +++ b/docs/my-website/docs/providers/anthropic.md @@ -1692,9 +1692,9 @@ Assistant: ``` -## Usage - PDF +## Usage - PDF -Pass base64 encoded PDF files to Anthropic models using the `image_url` field. +Pass base64 encoded PDF files to Anthropic models using the `file` content type with a `file_data` field. diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index f1eed4b4d52..5b247707696 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -7,7 +7,7 @@ ALL Bedrock models (Anthropic, Meta, Deepseek, Mistral, Amazon, etc.) are Suppor | Property | Details | |-------|-------| | Description | Amazon Bedrock is a fully managed service that offers a choice of high-performing foundation models (FMs). | -| Provider Route on LiteLLM | `bedrock/`, [`bedrock/converse/`](#set-converse--invoke-route), [`bedrock/invoke/`](#set-invoke-route), [`bedrock/converse_like/`](#calling-via-internal-proxy), [`bedrock/llama/`](#deepseek-not-r1), [`bedrock/deepseek_r1/`](#deepseek-r1), [`bedrock/qwen3/`](#qwen3-imported-models), [`bedrock/qwen2/`](./bedrock_imported.md#qwen2-imported-models), [`bedrock/openai/`](./bedrock_imported.md#openai-compatible-imported-models-qwen-25-vl-etc) | +| Provider Route on LiteLLM | `bedrock/`, [`bedrock/converse/`](#set-converse--invoke-route), [`bedrock/invoke/`](#set-invoke-route), [`bedrock/converse_like/`](#calling-via-internal-proxy), [`bedrock/llama/`](#deepseek-not-r1), [`bedrock/deepseek_r1/`](#deepseek-r1), [`bedrock/qwen3/`](#qwen3-imported-models), [`bedrock/qwen2/`](./bedrock_imported.md#qwen2-imported-models), [`bedrock/openai/`](./bedrock_imported.md#openai-compatible-imported-models-qwen-25-vl-etc), [`bedrock/moonshot`](./bedrock_imported.md#moonshot-kimi-k2-thinking) | | Provider Doc | [Amazon Bedrock ↗](https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html) | | Supported OpenAI Endpoints | `/chat/completions`, `/completions`, `/embeddings`, `/images/generations` | | Rerank Endpoint | `/rerank` | @@ -1941,6 +1941,7 @@ Here's an example of using a bedrock model with LiteLLM. For a complete list, re | Mixtral 8x7B Instruct | `completion(model='bedrock/mistral.mixtral-8x7b-instruct-v0:1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | | TwelveLabs Pegasus 1.2 (US) | `completion(model='bedrock/us.twelvelabs.pegasus-1-2-v1:0', messages=messages, mediaSource={...})` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | | TwelveLabs Pegasus 1.2 (EU) | `completion(model='bedrock/eu.twelvelabs.pegasus-1-2-v1:0', messages=messages, mediaSource={...})` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | +| Moonshot Kimi K2 Thinking | `completion(model='bedrock/moonshot.kimi-k2-thinking', messages=messages)` or `completion(model='bedrock/invoke/moonshot.kimi-k2-thinking', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | ## Bedrock Embedding diff --git a/docs/my-website/docs/providers/bedrock_imported.md b/docs/my-website/docs/providers/bedrock_imported.md index 0784f716925..709736e6109 100644 --- a/docs/my-website/docs/providers/bedrock_imported.md +++ b/docs/my-website/docs/providers/bedrock_imported.md @@ -431,4 +431,180 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \ "max_tokens": 300, "temperature": 0.5 }' -``` \ No newline at end of file +``` + +### Moonshot Kimi K2 Thinking + +Moonshot AI's Kimi K2 Thinking model is now available on Amazon Bedrock. This model features advanced reasoning capabilities with automatic reasoning content extraction. + +| Property | Details | +|----------|---------| +| Provider Route | `bedrock/moonshot.kimi-k2-thinking`, `bedrock/invoke/moonshot.kimi-k2-thinking` | +| Provider Documentation | [AWS Bedrock Moonshot Announcement ↗](https://aws.amazon.com/about-aws/whats-new/2025/12/amazon-bedrock-fully-managed-open-weight-models/) | +| Supported Parameters | `temperature`, `max_tokens`, `top_p`, `stream`, `tools`, `tool_choice` | +| Special Features | Reasoning content extraction, Tool calling | + +#### Supported Features + +- **Reasoning Content Extraction**: Automatically extracts `` tags and returns them as `reasoning_content` (similar to OpenAI's o1 models) +- **Tool Calling**: Full support for function/tool calling with tool responses +- **Streaming**: Both streaming and non-streaming responses +- **System Messages**: System message support + +#### Basic Usage + + + + +```python title="Moonshot Kimi K2 SDK Usage" showLineNumbers +from litellm import completion +import os + +os.environ["AWS_ACCESS_KEY_ID"] = "your-aws-access-key" +os.environ["AWS_SECRET_ACCESS_KEY"] = "your-aws-secret-key" +os.environ["AWS_REGION_NAME"] = "us-west-2" # or your preferred region + +# Basic completion +response = completion( + model="bedrock/moonshot.kimi-k2-thinking", # or bedrock/invoke/moonshot.kimi-k2-thinking + messages=[ + {"role": "user", "content": "What is 2+2? Think step by step."} + ], + temperature=0.7, + max_tokens=200 +) + +print(response.choices[0].message.content) + +# Access reasoning content if present +if response.choices[0].message.reasoning_content: + print("Reasoning:", response.choices[0].message.reasoning_content) +``` + + + + +**1. Add to config** + +```yaml title="config.yaml" showLineNumbers +model_list: + - model_name: kimi-k2 + litellm_params: + model: bedrock/moonshot.kimi-k2-thinking + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + aws_region_name: us-west-2 +``` + +**2. Start proxy** + +```bash title="Start LiteLLM Proxy" showLineNumbers +litellm --config /path/to/config.yaml + +# RUNNING at http://0.0.0.0:4000 +``` + +**3. Test it!** + +```bash title="Test Kimi K2 via Proxy" showLineNumbers +curl --location 'http://0.0.0.0:4000/chat/completions' \ + --header 'Authorization: Bearer sk-1234' \ + --header 'Content-Type: application/json' \ + --data '{ + "model": "kimi-k2", + "messages": [ + { + "role": "user", + "content": "What is 2+2? Think step by step." + } + ], + "temperature": 0.7, + "max_tokens": 200 + }' +``` + + + + +#### Tool Calling Example + +```python title="Kimi K2 with Tool Calling" showLineNumbers +from litellm import completion +import os + +os.environ["AWS_ACCESS_KEY_ID"] = "your-aws-access-key" +os.environ["AWS_SECRET_ACCESS_KEY"] = "your-aws-secret-key" +os.environ["AWS_REGION_NAME"] = "us-west-2" + +# Tool calling example +response = completion( + model="bedrock/moonshot.kimi-k2-thinking", + messages=[ + {"role": "user", "content": "What's the weather in Tokyo?"} + ], + tools=[ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the current weather in a location", + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "The city name" + } + }, + "required": ["location"] + } + } + } + ] +) + +if response.choices[0].message.tool_calls: + tool_call = response.choices[0].message.tool_calls[0] + print(f"Tool called: {tool_call.function.name}") + print(f"Arguments: {tool_call.function.arguments}") +``` + +#### Streaming Example + +```python title="Kimi K2 Streaming" showLineNumbers +from litellm import completion +import os + +os.environ["AWS_ACCESS_KEY_ID"] = "your-aws-access-key" +os.environ["AWS_SECRET_ACCESS_KEY"] = "your-aws-secret-key" +os.environ["AWS_REGION_NAME"] = "us-west-2" + +response = completion( + model="bedrock/moonshot.kimi-k2-thinking", + messages=[ + {"role": "user", "content": "Explain quantum computing in simple terms."} + ], + stream=True, + temperature=0.7 +) + +for chunk in response: + if chunk.choices[0].delta.content: + print(chunk.choices[0].delta.content, end="") + + # Check for reasoning content in streaming + if hasattr(chunk.choices[0].delta, 'reasoning_content') and chunk.choices[0].delta.reasoning_content: + print(f"\n[Reasoning: {chunk.choices[0].delta.reasoning_content}]") +``` + +#### Supported Parameters + +| Parameter | Type | Description | Supported | +|-----------|------|-------------|-----------| +| `temperature` | float (0-1) | Controls randomness in output | ✅ | +| `max_tokens` | integer | Maximum tokens to generate | ✅ | +| `top_p` | float | Nucleus sampling parameter | ✅ | +| `stream` | boolean | Enable streaming responses | ✅ | +| `tools` | array | Tool/function definitions | ✅ | +| `tool_choice` | string/object | Tool choice specification | ✅ | +| `stop` | array | Stop sequences | ❌ (Not supported on Bedrock) | \ No newline at end of file diff --git a/docs/my-website/docs/providers/manus.md b/docs/my-website/docs/providers/manus.md new file mode 100644 index 00000000000..2981ec6d247 --- /dev/null +++ b/docs/my-website/docs/providers/manus.md @@ -0,0 +1,194 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Manus + +Use Manus AI agents through LiteLLM's OpenAI-compatible Responses API. + +| Property | Details | +|----------|---------| +| Description | Manus is an AI agent platform for complex reasoning tasks, document analysis, and multi-step workflows with asynchronous task execution. | +| Provider Route on LiteLLM | `manus/{agent_profile}` | +| Supported Operations | `/responses` (Responses API) | +| Provider Doc | [Manus API ↗](https://open.manus.im/docs/openai-compatibility) | + +## Model Format + +```shell +manus/{agent_profile} +``` + +**Examples:** +- `manus/manus-1.6` - General purpose agent +- `manus/manus-1.6-lite` - Lightweight agent for simple tasks +- `manus/manus-1.6-max` - Advanced agent for complex analysis + +## LiteLLM Python SDK + +```python showLineNumbers title="Basic Usage" +import litellm +import os +import time + +# Set API key +os.environ["MANUS_API_KEY"] = "your-manus-api-key" + +# Create task +response = litellm.responses( + model="manus/manus-1.6", + input="What's the capital of France?", +) + +print(f"Task ID: {response.id}") +print(f"Status: {response.status}") # "running" + +# Poll until complete +task_id = response.id +while response.status == "running": + time.sleep(5) + response = litellm.get_response( + response_id=task_id, + custom_llm_provider="manus", + ) + print(f"Status: {response.status}") + +# Get results +if response.status == "completed": + for message in response.output: + if message.role == "assistant": + print(message.content[0].text) +``` + +## LiteLLM AI Gateway + +### Setup + +```yaml showLineNumbers title="config.yaml" +model_list: + - model_name: manus-agent + litellm_params: + model: manus/manus-1.6 + api_key: os.environ/MANUS_API_KEY +``` + +```bash title="Start Proxy" +litellm --config config.yaml +``` + +### Usage + + + + +```bash showLineNumbers title="Create Task" +# Create task +curl -X POST http://localhost:4000/responses \ + -H "Authorization: Bearer your-proxy-key" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "manus-agent", + "input": "What is the capital of France?" + }' + +# Response +{ + "id": "task_abc123", + "status": "running", + "metadata": { + "task_url": "https://manus.im/app/task_abc123" + } +} +``` + +```bash showLineNumbers title="Poll for Completion" +# Check status (repeat until status is "completed") +curl http://localhost:4000/responses/task_abc123 \ + -H "Authorization: Bearer your-proxy-key" + +# When completed +{ + "id": "task_abc123", + "status": "completed", + "output": [ + { + "role": "user", + "content": [{"text": "What is the capital of France?"}] + }, + { + "role": "assistant", + "content": [{"text": "The capital of France is Paris."}] + } + ] +} +``` + + + + +```python showLineNumbers title="Create Task and Poll" +import openai +import time + +client = openai.OpenAI( + base_url="http://localhost:4000", + api_key="your-proxy-key" +) + +# Create task +response = client.responses.create( + model="manus-agent", + input="What is the capital of France?" +) + +print(f"Task ID: {response.id}") +print(f"Status: {response.status}") # "running" + +# Poll until complete +task_id = response.id +while response.status == "running": + time.sleep(5) + response = client.responses.retrieve(response_id=task_id) + print(f"Status: {response.status}") + +# Get results +if response.status == "completed": + for message in response.output: + if message.role == "assistant": + print(message.content[0].text) +``` + + + + +## How It Works + +Manus operates as an **asynchronous agent API**: + +1. **Create Task**: When you call `litellm.responses()`, Manus creates a task and returns immediately with `status: "running"` +2. **Task Executes**: The agent works on your request in the background +3. **Poll for Completion**: You must repeatedly call `litellm.get_response()` or `client.responses.retrieve()` until the status changes to `"completed"` +4. **Get Results**: Once completed, the `output` field contains the full conversation + +**Task Statuses:** +- `running` - Agent is actively working +- `pending` - Agent is waiting for input +- `completed` - Task finished successfully +- `error` - Task failed + +:::tip Production Usage +For production applications, use [webhooks](https://open.manus.im/docs/webhooks) instead of polling to get notified when tasks complete. +::: + +## Supported Parameters + +| Parameter | Supported | Notes | +|-----------|-----------|-------| +| `input` | ✅ | Text, images, or structured content | +| `stream` | ✅ | Fake streaming (task runs async) | +| `max_output_tokens` | ✅ | Limits response length | +| `previous_response_id` | ✅ | For multi-turn conversations | + +## Related Documentation + +- [LiteLLM Responses API](/docs/response_api) +- [Manus OpenAI Compatibility](https://open.manus.im/docs/openai-compatibility) diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md index 33ebf535d29..f46608aa57c 100644 --- a/docs/my-website/docs/providers/vertex.md +++ b/docs/my-website/docs/providers/vertex.md @@ -35,6 +35,8 @@ import json # !gcloud auth application-default login - run this to add vertex credentials to your env ## OR ## file_path = 'path/to/vertex_ai_service_account.json' +## OR ## +export VERTEXAI_API_KEY="your-api-key" # Load the JSON file with open(file_path, 'r') as file: @@ -47,7 +49,7 @@ vertex_credentials_json = json.dumps(vertex_credentials) response = completion( model="vertex_ai/gemini-2.5-pro", messages=[{ "content": "Hello, how are you?","role": "user"}], - vertex_credentials=vertex_credentials_json + vertex_credentials=vertex_credentials_json # Can remove this is added VERTEXAI_API_KEY in env ) ``` @@ -1329,15 +1331,41 @@ Here's how to use Vertex AI with the LiteLLM Proxy Server ## Authentication - vertex_project, vertex_location, etc. +LiteLLM supports two authentication methods for Vertex AI: + +1. **API Key Authentication** (Recommended for getting started) +2. **Service Account Credentials** (Recommended for production) + Set your vertex credentials via: - dynamic params OR - env vars +### **Authentication Method 1: -### **Dynamic Params** +The simplest way to authenticate with Vertex AI. You can set: +- `api_key` (str) - Your Vertex AI API key -You can set: +**Environment Variables:** +```bash +export VERTEXAI_API_KEY="your-api-key" +``` + +**Or pass as parameters:** +```python +from litellm import completion + +response = completion( + model="vertex_ai/gemini-2.0-flash-exp", + messages=[{"role": "user", "content": "Hello!"}], + api_key="your-vertex-api-key", + +) +``` + +### **Authentication Method 2: Service Account Credentials** + +For production environments with fine-grained access control. You can set: - `vertex_credentials` (str) - can be a json string or filepath to your vertex ai service account.json - `vertex_location` (str) - place where vertex model is deployed (us-central1, asia-southeast1, etc.). Some models support the global location, please see [Vertex AI documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#supported_models) - `vertex_project` Optional[str] - use if vertex project different from the one in vertex_credentials @@ -1392,7 +1420,16 @@ model_list: ### **Environment Variables** -You can set: +#### For API Key Authentication: + +- `VERTEXAI_API_KEY` or `VERTEX_API_KEY` - Your Vertex AI API key + +```bash +export VERTEXAI_API_KEY="your-vertex-api-key" +``` + +#### For Service Account Authentication: + - `GOOGLE_APPLICATION_CREDENTIALS` - store the filepath for your service_account.json in here (used by vertex sdk directly). - VERTEXAI_LOCATION - place where vertex model is deployed (us-central1, asia-southeast1, etc.) - VERTEXAI_PROJECT - Optional[str] - use if vertex project different from the one in vertex_credentials diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index dfc0efd37ad..42494c1e5f5 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -146,6 +146,7 @@ router_settings: cooldown_time: 30 # (in seconds) how long to cooldown model if fails/min > allowed_fails disable_cooldowns: True # bool - Disable cooldowns for all models enable_tag_filtering: True # bool - Use tag based routing for requests + tag_filtering_match_any: True # bool - Tag matching behavior (only when enable_tag_filtering=true). `true`: match if deployment has ANY requested tag; `false`: match only if deployment has ALL requested tags retry_policy: { # Dict[str, int]: retry policy for different types of exceptions "AuthenticationErrorRetries": 3, "TimeoutErrorRetries": 3, @@ -293,6 +294,7 @@ router_settings: cooldown_time: 30 # (in seconds) how long to cooldown model if fails/min > allowed_fails disable_cooldowns: True # bool - Disable cooldowns for all models enable_tag_filtering: True # bool - Use tag based routing for requests + tag_filtering_match_any: True # bool - Tag matching behavior (only when enable_tag_filtering=true). `true`: match if deployment has ANY requested tag; `false`: match only if deployment has ALL requested tags retry_policy: { # Dict[str, int]: retry policy for different types of exceptions "AuthenticationErrorRetries": 3, "TimeoutErrorRetries": 3, @@ -322,6 +324,7 @@ router_settings: | content_policy_fallbacks | array of objects | Specifies fallback models for content policy violations. [More information here](reliability) | | fallbacks | array of objects | Specifies fallback models for all types of errors. [More information here](reliability) | | enable_tag_filtering | boolean | If true, uses tag based routing for requests [Tag Based Routing](tag_routing) | +| tag_filtering_match_any | boolean | Tag matching behavior (only when enable_tag_filtering=true). `true`: match if deployment has ANY requested tag; `false`: match only if deployment has ALL requested tags | | cooldown_time | integer | The duration (in seconds) to cooldown a model if it exceeds the allowed failures. | | disable_cooldowns | boolean | If true, disables cooldowns for all models. [More information here](reliability) | | retry_policy | object | Specifies the number of retries for different types of exceptions. [More information here](reliability) | @@ -578,6 +581,18 @@ router_settings: | FIREWORKS_AI_56_B_MOE | Size parameter for Fireworks AI 56B MOE model. Default is 56 | FIREWORKS_AI_80_B | Size parameter for Fireworks AI 80B model. Default is 80 | FIREWORKS_AI_176_B_MOE | Size parameter for Fireworks AI 176B MOE model. Default is 176 +| FOCUS_PROVIDER | Destination provider for Focus exports (e.g., `s3`). Defaults to `s3`. +| FOCUS_FORMAT | Output format for Focus exports. Defaults to `parquet`. +| FOCUS_FREQUENCY | Frequency for scheduled Focus exports (`hourly`, `daily`, or `interval`). Defaults to `hourly`. +| FOCUS_CRON_OFFSET | Minute offset used when scheduling hourly/daily Focus exports. Defaults to `5` minutes. +| FOCUS_INTERVAL_SECONDS | Interval (in seconds) for Focus exports when `frequency` is `interval`. +| FOCUS_PREFIX | Object key prefix (or folder) used when uploading Focus export files. Defaults to `focus_exports`. +| FOCUS_S3_BUCKET_NAME | S3 bucket to upload Focus export files when using the S3 destination. +| FOCUS_S3_REGION_NAME | AWS region for the Focus export S3 bucket. +| FOCUS_S3_ENDPOINT_URL | Custom endpoint for the Focus export S3 client (optional; useful for S3-compatible storage). +| FOCUS_S3_ACCESS_KEY | AWS access key ID used by the Focus export S3 client. +| FOCUS_S3_SECRET_KEY | AWS secret access key used by the Focus export S3 client. +| FOCUS_S3_SESSION_TOKEN | AWS session token used by the Focus export S3 client (optional). | FUNCTION_DEFINITION_TOKEN_COUNT | Token count for function definitions. Default is 9 | GALILEO_BASE_URL | Base URL for Galileo platform | GALILEO_PASSWORD | Password for Galileo authentication diff --git a/docs/my-website/docs/proxy/guardrails/qualifire.md b/docs/my-website/docs/proxy/guardrails/qualifire.md index 66961c92d9d..850af37e47f 100644 --- a/docs/my-website/docs/proxy/guardrails/qualifire.md +++ b/docs/my-website/docs/proxy/guardrails/qualifire.md @@ -8,13 +8,7 @@ Use [Qualifire](https://qualifire.ai) to evaluate LLM outputs for quality, safet ## Quick Start -### 1. Install the Qualifire SDK - -```bash -pip install qualifire -``` - -### 2. Define Guardrails on your LiteLLM config.yaml +### 1. Define Guardrails on your LiteLLM config.yaml Define your guardrails under the `guardrails` section: @@ -61,13 +55,13 @@ guardrails: - `post_call` Run **after** LLM call, on **input & output** - `during_call` Run **during** LLM call, on **input**. Same as `pre_call` but runs in parallel as LLM call. Response not returned until guardrail check completes -### 3. Start LiteLLM Gateway +### 2. Start LiteLLM Gateway ```shell litellm --config config.yaml --detailed_debug ``` -### 4. Test request +### 3. Test request **[Langchain, OpenAI SDK Usage Examples](../proxy/user_keys#request-format)** @@ -142,7 +136,7 @@ guardrails: evaluation_id: eval_abc123 # Your evaluation ID from Qualifire dashboard ``` -When `evaluation_id` is provided, LiteLLM will use `invoke_evaluation()` instead of `evaluate()`, running the pre-configured evaluation from your dashboard. +When `evaluation_id` is provided, LiteLLM will use the invoke evaluation API endpoint instead of the evaluate endpoint, running the pre-configured evaluation from your dashboard. ## Available Checks @@ -213,19 +207,19 @@ guardrails: ### Parameter Reference -| Parameter | Type | Default | Description | -| ------------------------------ | ----------- | --------------------------- | -------------------------------------------------------- | -| `api_key` | `str` | `QUALIFIRE_API_KEY` env var | Your Qualifire API key | -| `api_base` | `str` | `None` | Custom API base URL (optional) | -| `evaluation_id` | `str` | `None` | Pre-configured evaluation ID from Qualifire dashboard | -| `prompt_injections` | `bool` | `true` (if no other checks) | Enable prompt injection detection | -| `hallucinations_check` | `bool` | `None` | Enable hallucination detection | -| `grounding_check` | `bool` | `None` | Enable grounding verification | -| `pii_check` | `bool` | `None` | Enable PII detection | -| `content_moderation_check` | `bool` | `None` | Enable content moderation | -| `tool_selection_quality_check` | `bool` | `None` | Enable tool selection quality check | -| `assertions` | `List[str]` | `None` | Custom assertions to validate | -| `on_flagged` | `str` | `"block"` | Action when content is flagged: `"block"` or `"monitor"` | +| Parameter | Type | Default | Description | +| ------------------------------ | ----------- | ---------------------------- | -------------------------------------------------------- | +| `api_key` | `str` | `QUALIFIRE_API_KEY` env var | Your Qualifire API key | +| `api_base` | `str` | `https://proxy.qualifire.ai` | Custom API base URL (optional) | +| `evaluation_id` | `str` | `None` | Pre-configured evaluation ID from Qualifire dashboard | +| `prompt_injections` | `bool` | `true` (if no other checks) | Enable prompt injection detection | +| `hallucinations_check` | `bool` | `None` | Enable hallucination detection | +| `grounding_check` | `bool` | `None` | Enable grounding verification | +| `pii_check` | `bool` | `None` | Enable PII detection | +| `content_moderation_check` | `bool` | `None` | Enable content moderation | +| `tool_selection_quality_check` | `bool` | `None` | Enable tool selection quality check | +| `assertions` | `List[str]` | `None` | Custom assertions to validate | +| `on_flagged` | `str` | `"block"` | Action when content is flagged: `"block"` or `"monitor"` | ### Default Behavior @@ -261,4 +255,3 @@ This evaluates whether the LLM selected the appropriate tools and provided corre - [Qualifire Documentation](https://docs.qualifire.ai) - [Qualifire Dashboard](https://app.qualifire.ai) -- [Qualifire Python SDK](https://github.com/qualifire-dev/qualifire-python-sdk) diff --git a/docs/my-website/docs/proxy/logging.md b/docs/my-website/docs/proxy/logging.md index 5fe8f17d7b0..a27b6dcf083 100644 --- a/docs/my-website/docs/proxy/logging.md +++ b/docs/my-website/docs/proxy/logging.md @@ -67,7 +67,7 @@ Set `litellm.turn_off_message_logging=True` This will prevent the messages and r -**1. Setup config.yaml ** +**1. Setup config.yaml** ```yaml model_list: - model_name: gpt-3.5-turbo diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 482d855082e..2e8ea07ab75 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -55,6 +55,7 @@ const sidebars = { "proxy/guardrails/test_playground", "proxy/guardrails/litellm_content_filter", ...[ + "proxy/guardrails/qualifire", "proxy/guardrails/aim_security", "proxy/guardrails/onyx_security", "proxy/guardrails/aporia_api", @@ -653,12 +654,13 @@ const sidebars = { "providers/bedrock_writer", "providers/bedrock_batches", "providers/aws_polly", - "providers/bedrock_vector_store", - ] - }, - "providers/litellm_proxy", - "providers/ai21", - "providers/aiml", + "providers/bedrock_vector_store", + ] + }, + "providers/litellm_proxy", + "providers/abliteration", + "providers/ai21", + "providers/aiml", "providers/aleph_alpha", "providers/amazon_nova", "providers/anyscale", @@ -710,6 +712,7 @@ const sidebars = { "providers/llamafile", "providers/llamagate", "providers/lm_studio", + "providers/manus", "providers/meta_llama", "providers/milvus_vector_stores", "providers/mistral", diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260108_add_user_email_lower_idx/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260108_add_user_email_lower_idx/migration.sql new file mode 100644 index 00000000000..add80b39e7f --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260108_add_user_email_lower_idx/migration.sql @@ -0,0 +1,9 @@ +-- CreateIndex +-- Fixes performance issue in _check_duplicate_user_email function +-- by enabling fast case-insensitive email lookups. +-- +-- Without this index, queries with mode: "insensitive" cause full table scans. +-- With this index, PostgreSQL can use an Index Scan for O(log n) performance. +-- +-- Related: GitHub Issue #18411 +CREATE INDEX "LiteLLM_UserTable_user_email_lower_idx" ON "LiteLLM_UserTable"(LOWER("user_email")); diff --git a/litellm/__init__.py b/litellm/__init__.py index 1bc690e561f..4c43ae60aa6 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -134,6 +134,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "bitbucket", "gitlab", "cloudzero", + "focus", "posthog", "levo", ] @@ -486,6 +487,7 @@ vertex_mistral_models: Set = set() vertex_openai_models: Set = set() vertex_minimax_models: Set = set() vertex_moonshot_models: Set = set() +vertex_zai_models: Set = set() ai21_models: Set = set() ai21_chat_models: Set = set() nlp_cloud_models: Set = set() @@ -664,6 +666,9 @@ def add_known_models(): elif value.get("litellm_provider") == "vertex_ai-moonshot_models": key = key.replace("vertex_ai/", "") vertex_moonshot_models.add(key) + elif value.get("litellm_provider") == "vertex_ai-zai_models": + key = key.replace("vertex_ai/", "") + vertex_zai_models.add(key) elif value.get("litellm_provider") == "ai21": if value.get("mode") == "chat": ai21_chat_models.add(key) @@ -950,7 +955,8 @@ models_by_provider: dict = { | vertex_language_models | vertex_deepseek_models | vertex_minimax_models - | vertex_moonshot_models, + | vertex_moonshot_models + | vertex_zai_models, "ai21": ai21_models, "bedrock": bedrock_models | bedrock_converse_models, "petals": petals_models, @@ -1338,6 +1344,7 @@ if TYPE_CHECKING: from .llms.bedrock.chat.invoke_transformations.amazon_llama_transformation import AmazonLlamaConfig as AmazonLlamaConfig from .llms.bedrock.chat.invoke_transformations.amazon_deepseek_transformation import AmazonDeepSeekR1Config as AmazonDeepSeekR1Config from .llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation import AmazonMistralConfig as AmazonMistralConfig + from .llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import AmazonMoonshotConfig as AmazonMoonshotConfig from .llms.bedrock.chat.invoke_transformations.amazon_titan_transformation import AmazonTitanConfig as AmazonTitanConfig from .llms.bedrock.chat.invoke_transformations.amazon_twelvelabs_pegasus_transformation import AmazonTwelveLabsPegasusConfig as AmazonTwelveLabsPegasusConfig from .llms.bedrock.chat.invoke_transformations.base_invoke_transformation import AmazonInvokeConfig as AmazonInvokeConfig @@ -1367,6 +1374,7 @@ if TYPE_CHECKING: from .llms.azure.responses.o_series_transformation import AzureOpenAIOSeriesResponsesAPIConfig as AzureOpenAIOSeriesResponsesAPIConfig from .llms.xai.responses.transformation import XAIResponsesAPIConfig as XAIResponsesAPIConfig from .llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig as LiteLLMProxyResponsesAPIConfig + from .llms.manus.responses.transformation import ManusResponsesAPIConfig as ManusResponsesAPIConfig from .llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig from .llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig as OpenAIOSeriesConfig, OpenAIOSeriesConfig as OpenAIO1Config from .llms.anthropic.skills.transformation import AnthropicSkillsConfig as AnthropicSkillsConfig diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 26133ebc222..f37c4dc6d04 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -165,6 +165,7 @@ LLM_CONFIG_NAMES = ( "AmazonLlamaConfig", "AmazonDeepSeekR1Config", "AmazonMistralConfig", + "AmazonMoonshotConfig", "AmazonTitanConfig", "AmazonTwelveLabsPegasusConfig", "AmazonInvokeConfig", @@ -252,6 +253,7 @@ LLM_CONFIG_NAMES = ( "IBMWatsonXAudioTranscriptionConfig", "GithubCopilotConfig", "GithubCopilotResponsesAPIConfig", + "ManusResponsesAPIConfig", "GithubCopilotEmbeddingConfig", "NebiusConfig", "WandbConfig", @@ -556,6 +558,7 @@ _LLM_CONFIGS_IMPORT_MAP = { "AmazonLlamaConfig": (".llms.bedrock.chat.invoke_transformations.amazon_llama_transformation", "AmazonLlamaConfig"), "AmazonDeepSeekR1Config": (".llms.bedrock.chat.invoke_transformations.amazon_deepseek_transformation", "AmazonDeepSeekR1Config"), "AmazonMistralConfig": (".llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation", "AmazonMistralConfig"), + "AmazonMoonshotConfig": (".llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation", "AmazonMoonshotConfig"), "AmazonTitanConfig": (".llms.bedrock.chat.invoke_transformations.amazon_titan_transformation", "AmazonTitanConfig"), "AmazonTwelveLabsPegasusConfig": (".llms.bedrock.chat.invoke_transformations.amazon_twelvelabs_pegasus_transformation", "AmazonTwelveLabsPegasusConfig"), "AmazonInvokeConfig": (".llms.bedrock.chat.invoke_transformations.base_invoke_transformation", "AmazonInvokeConfig"), @@ -588,6 +591,7 @@ _LLM_CONFIGS_IMPORT_MAP = { "AzureOpenAIOSeriesResponsesAPIConfig": (".llms.azure.responses.o_series_transformation", "AzureOpenAIOSeriesResponsesAPIConfig"), "XAIResponsesAPIConfig": (".llms.xai.responses.transformation", "XAIResponsesAPIConfig"), "LiteLLMProxyResponsesAPIConfig": (".llms.litellm_proxy.responses.transformation", "LiteLLMProxyResponsesAPIConfig"), + "ManusResponsesAPIConfig": (".llms.manus.responses.transformation", "ManusResponsesAPIConfig"), "GoogleAIStudioInteractionsConfig": (".llms.gemini.interactions.transformation", "GoogleAIStudioInteractionsConfig"), "OpenAIOSeriesConfig": (".llms.openai.chat.o_series_transformation", "OpenAIOSeriesConfig"), "AnthropicSkillsConfig": (".llms.anthropic.skills.transformation", "AnthropicSkillsConfig"), diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index a89efc4e82b..af8185aa215 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -31,6 +31,7 @@ from litellm.llms.base_llm.bridges.completion_transformation import ( CompletionTransformationBridge, ) from litellm.types.llms.openai import ( + ChatCompletionAnnotation, ChatCompletionToolParamFunctionChunk, Reasoning, ResponsesAPIOptionalRequestParams, @@ -90,9 +91,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): content_type = content_item.get("type") if content_type == "output_text": response_text = content_item.get("text", "") + # Extract annotations from content if present + annotations = LiteLLMResponsesTransformationHandler._convert_annotations_to_chat_format( + content_item.get("annotations", None) + ) msg = Message( role=item.get("role", "assistant"), content=response_text if response_text else "", + annotations=annotations, ) choice = Choices(message=msg, finish_reason="stop", index=index) return choice, index + 1 @@ -364,10 +370,16 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): elif isinstance(item, ResponseOutputMessage): for content in item.content: response_text = getattr(content, "text", "") + # Extract annotations from content if present + raw_annotations = getattr(content, "annotations", None) + annotations = LiteLLMResponsesTransformationHandler._convert_annotations_to_chat_format( + raw_annotations + ) msg = Message( role=item.role, content=response_text if response_text else "", reasoning_content=reasoning_content, + annotations=annotations, ) choices.append( @@ -763,6 +775,42 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return {"format": {"type": "text"}} return None + + @staticmethod + def _convert_annotations_to_chat_format( + annotations: Optional[List[Any]], + ) -> Optional[List["ChatCompletionAnnotation"]]: + """ + Convert annotations from Responses API to Chat Completions format. + + Annotations are already in compatible format between both APIs, + so we just need to convert Pydantic models to dicts. + """ + if not annotations: + return None + + result: List[ChatCompletionAnnotation] = [] + for annotation in annotations: + try: + # Convert Pydantic models to dicts (handles both v1 and v2) + if hasattr(annotation, "model_dump"): + annotation_dict = annotation.model_dump() + elif hasattr(annotation, "dict"): + annotation_dict = annotation.dict() + elif isinstance(annotation, dict): + annotation_dict = annotation + else: + # Skip unsupported annotation types + verbose_logger.debug(f"Skipping unsupported annotation type: {type(annotation)}") + continue + + result.append(annotation_dict) # type: ignore + except Exception as e: + # Skip malformed annotations + verbose_logger.debug(f"Skipping malformed annotation: {annotation}, error: {e}") + continue + + return result if result else None def _map_responses_status_to_finish_reason(self, status: Optional[str]) -> str: """Map responses API status to chat completion finish_reason""" diff --git a/litellm/constants.py b/litellm/constants.py index db9d0114118..e24186d567e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -909,6 +909,7 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[ "twelvelabs", "openai", "stability", + "moonshot", ] BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[ diff --git a/litellm/google_genai/main.py b/litellm/google_genai/main.py index b7523ef8c16..1dc805a6b54 100644 --- a/litellm/google_genai/main.py +++ b/litellm/google_genai/main.py @@ -130,6 +130,9 @@ class GenerateContentHelper: api_key=litellm_params.api_key, ) + if litellm_params.custom_llm_provider is None: + litellm_params.custom_llm_provider = custom_llm_provider + # get provider config generate_content_provider_config: Optional[ BaseGoogleGenAIGenerateContentConfig @@ -407,6 +410,9 @@ async def agenerate_content_stream( # Check if we should use the adapter (when provider config is None) if setup_result.generate_content_provider_config is None: + if "stream" in kwargs: + kwargs.pop("stream", None) + # Use the adapter to convert to completion format return ( await GenerateContentToCompletionHandler.async_generate_content_handler( @@ -490,6 +496,9 @@ def generate_content_stream( # Check if we should use the adapter (when provider config is None) if setup_result.generate_content_provider_config is None: + if "stream" in kwargs: + kwargs.pop("stream", None) + # Use the adapter to convert to completion format return GenerateContentToCompletionHandler.generate_content_handler( model=model, diff --git a/litellm/integrations/braintrust_logging.py b/litellm/integrations/braintrust_logging.py index 364fa3f5def..585de510e8b 100644 --- a/litellm/integrations/braintrust_logging.py +++ b/litellm/integrations/braintrust_logging.py @@ -225,10 +225,13 @@ class BraintrustLogger(CustomLogger): "id": litellm_call_id, "input": prompt["messages"], "metadata": standard_logging_object, - "tags": tags, "span_attributes": {"name": span_name, "type": "llm"}, } - + + # Braintrust cannot specify 'tags' for non-root spans + if dynamic_metadata.get("root_span_id") is None: + request_data["tags"] = tags + # Only add those that are not None (or falsy) for key, value in span_attributes.items(): if value: @@ -351,14 +354,37 @@ class BraintrustLogger(CustomLogger): # Allow metadata override for span name span_name = dynamic_metadata.get("span_name", "Chat Completion") + # Span parents is a special case + span_parents = dynamic_metadata.get("span_parents") + + # Convert comma-separated string to list if present + if span_parents: + span_parents = [s.strip() for s in span_parents.split(",") if s.strip()] + + # Add optional span attributes only if present + span_attributes = { + "span_id": dynamic_metadata.get("span_id"), + "root_span_id": dynamic_metadata.get("root_span_id"), + "span_parents": span_parents, + } + request_data = { "id": litellm_call_id, "input": prompt["messages"], "output": output, "metadata": standard_logging_object, - "tags": tags, "span_attributes": {"name": span_name, "type": "llm"}, } + + # Braintrust cannot specify 'tags' for non-root spans + if dynamic_metadata.get("root_span_id") is None: + request_data["tags"] = tags + + # Only add those that are not None (or falsy) + for key, value in span_attributes.items(): + if value: + request_data[key] = value + if choices is not None: request_data["output"] = [choice.dict() for choice in choices] else: @@ -367,9 +393,6 @@ class BraintrustLogger(CustomLogger): if metrics is not None: request_data["metrics"] = metrics - if metrics is not None: - request_data["metrics"] = metrics - try: await self.global_braintrust_http_handler.post( url=f"{self.api_base}/project_logs/{project_id}/insert", diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 2128b55bf83..71929398103 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -19,7 +19,7 @@ """Database connection and data extraction for LiteLLM.""" from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any, Optional, List import polars as pl @@ -46,19 +46,9 @@ class LiteLLMDatabase: """Retrieve usage data from LiteLLM daily user spend table.""" client = self._ensure_prisma_client() - # Build WHERE clause for time filtering - where_conditions = [] - if start_time_utc: - where_conditions.append(f"dus.updated_at >= '{start_time_utc.isoformat()}'") - if end_time_utc: - where_conditions.append(f"dus.updated_at <= '{end_time_utc.isoformat()}'") - - where_clause = "" - if where_conditions: - where_clause = "WHERE " + " AND ".join(where_conditions) - - # Query to get user spend data with team information - query = f""" + # Query to get user spend data with team information. Use parameter binding to + # avoid SQL injection from user-supplied timestamps or limits. + query = """ SELECT dus.id, dus.date, @@ -85,163 +75,27 @@ class LiteLLMDatabase: LEFT JOIN "LiteLLM_VerificationToken" vt ON dus.api_key = vt.token LEFT JOIN "LiteLLM_TeamTable" tt ON vt.team_id = tt.team_id LEFT JOIN "LiteLLM_UserTable" ut ON dus.user_id = ut.user_id - {where_clause} + WHERE ($1::timestamptz IS NULL OR dus.updated_at >= $1::timestamptz) + AND ($2::timestamptz IS NULL OR dus.updated_at <= $2::timestamptz) ORDER BY dus.date DESC, dus.created_at DESC """ - if limit: - query += f" LIMIT {limit}" + params: List[Any] = [ + start_time_utc, + end_time_utc, + ] + + if limit is not None: + try: + params.append(int(limit)) + except (TypeError, ValueError): + raise ValueError("limit must be an integer") + query += " LIMIT $3" try: - db_response = await client.db.query_raw(query) + db_response = await client.db.query_raw(query, *params) # Convert the response to polars DataFrame with full schema inference # This prevents schema mismatch errors when data types vary across rows return pl.DataFrame(db_response, infer_schema_length=None) except Exception as e: raise Exception(f"Error retrieving usage data: {str(e)}") - - async def get_table_info(self) -> Dict[str, Any]: - """Get information about the daily user spend table.""" - client = self._ensure_prisma_client() - - try: - # Get row count from user spend table - user_count = await self._get_table_row_count("LiteLLM_DailyUserSpend") - - # Get column structure from user spend table - query = """ - SELECT column_name, data_type, is_nullable - FROM information_schema.columns - WHERE table_name = 'LiteLLM_DailyUserSpend' - ORDER BY ordinal_position; - """ - columns_response = await client.db.query_raw(query) - - return { - "columns": columns_response, - "row_count": user_count, - "table_name": "LiteLLM_DailyUserSpend", - } - except Exception as e: - raise Exception(f"Error getting table info: {str(e)}") - - async def _get_table_row_count(self, table_name: str) -> int: - """Get row count from specified table.""" - client = self._ensure_prisma_client() - - try: - query = f'SELECT COUNT(*) as count FROM "{table_name}"' - response = await client.db.query_raw(query) - - if response and len(response) > 0: - return response[0].get("count", 0) - return 0 - except Exception: - return 0 - - async def discover_all_tables(self) -> Dict[str, Any]: - """Discover all tables in the LiteLLM database and their schemas.""" - client = self._ensure_prisma_client() - - try: - # Get all LiteLLM tables - litellm_tables_query = """ - SELECT table_name - FROM information_schema.tables - WHERE table_schema = 'public' - AND table_name LIKE 'LiteLLM_%' - ORDER BY table_name; - """ - tables_response = await client.db.query_raw(litellm_tables_query) - table_names = [row["table_name"] for row in tables_response] - - # Get detailed schema for each table - tables_info = {} - for table_name in table_names: - # Get column information - columns_query = """ - SELECT - column_name, - data_type, - is_nullable, - column_default, - character_maximum_length, - numeric_precision, - numeric_scale, - ordinal_position - FROM information_schema.columns - WHERE table_name = $1 - AND table_schema = 'public' - ORDER BY ordinal_position; - """ - columns_response = await client.db.query_raw(columns_query, table_name) - - # Get primary key information - pk_query = """ - SELECT a.attname - FROM pg_index i - JOIN pg_attribute a ON a.attrelid = i.indrelid AND a.attnum = ANY(i.indkey) - WHERE i.indrelid = $1::regclass AND i.indisprimary; - """ - pk_response = await client.db.query_raw(pk_query, f'"{table_name}"') - primary_keys = ( - [row["attname"] for row in pk_response] if pk_response else [] - ) - - # Get foreign key information - fk_query = """ - SELECT - tc.constraint_name, - kcu.column_name, - ccu.table_name AS foreign_table_name, - ccu.column_name AS foreign_column_name - FROM information_schema.table_constraints AS tc - JOIN information_schema.key_column_usage AS kcu - ON tc.constraint_name = kcu.constraint_name - JOIN information_schema.constraint_column_usage AS ccu - ON ccu.constraint_name = tc.constraint_name - WHERE tc.constraint_type = 'FOREIGN KEY' - AND tc.table_name = $1; - """ - fk_response = await client.db.query_raw(fk_query, table_name) - foreign_keys = fk_response if fk_response else [] - - # Get indexes - indexes_query = """ - SELECT - i.relname AS index_name, - array_agg(a.attname ORDER BY a.attnum) AS column_names, - ix.indisunique AS is_unique - FROM pg_class t - JOIN pg_index ix ON t.oid = ix.indrelid - JOIN pg_class i ON i.oid = ix.indexrelid - JOIN pg_attribute a ON a.attrelid = t.oid AND a.attnum = ANY(ix.indkey) - WHERE t.relname = $1 - AND t.relkind = 'r' - GROUP BY i.relname, ix.indisunique - ORDER BY i.relname; - """ - indexes_response = await client.db.query_raw(indexes_query, table_name) - indexes = indexes_response if indexes_response else [] - - # Get row count - try: - row_count = await self._get_table_row_count(table_name) - except Exception: - row_count = 0 - - tables_info[table_name] = { - "columns": columns_response, - "primary_keys": primary_keys, - "foreign_keys": foreign_keys, - "indexes": indexes, - "row_count": row_count, - } - - return { - "tables": tables_info, - "table_count": len(table_names), - "table_names": table_names, - } - except Exception as e: - raise Exception(f"Error discovering tables: {str(e)}") diff --git a/litellm/integrations/focus/__init__.py b/litellm/integrations/focus/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/integrations/focus/database.py b/litellm/integrations/focus/database.py new file mode 100644 index 00000000000..298254670eb --- /dev/null +++ b/litellm/integrations/focus/database.py @@ -0,0 +1,113 @@ +"""Database access helpers for Focus export.""" + +from __future__ import annotations + +from datetime import datetime +from typing import Any, Dict, Optional + +import polars as pl + + +class FocusLiteLLMDatabase: + """Retrieves LiteLLM usage data for Focus export workflows.""" + + def _ensure_prisma_client(self): + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise RuntimeError( + "Database not connected. Connect a database to your proxy - " + "https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" + ) + return prisma_client + + async def get_usage_data( + self, + *, + limit: Optional[int] = None, + start_time_utc: Optional[datetime] = None, + end_time_utc: Optional[datetime] = None, + ) -> pl.DataFrame: + """Return usage data for the requested window.""" + client = self._ensure_prisma_client() + + where_clauses: list[str] = [] + query_params: list[Any] = [] + placeholder_index = 1 + if start_time_utc: + where_clauses.append(f"dus.updated_at >= ${placeholder_index}::timestamptz") + query_params.append(start_time_utc) + placeholder_index += 1 + if end_time_utc: + where_clauses.append(f"dus.updated_at <= ${placeholder_index}::timestamptz") + query_params.append(end_time_utc) + placeholder_index += 1 + + where_clause = "" + if where_clauses: + where_clause = "WHERE " + " AND ".join(where_clauses) + + limit_clause = "" + if limit is not None: + try: + limit_value = int(limit) + except (TypeError, ValueError) as exc: # pragma: no cover - defensive guard + raise ValueError("limit must be an integer") from exc + if limit_value < 0: + raise ValueError("limit must be non-negative") + limit_clause = f" LIMIT ${placeholder_index}" + query_params.append(limit_value) + + query = f""" + SELECT + dus.id, + dus.date, + dus.user_id, + dus.api_key, + dus.model, + dus.model_group, + dus.custom_llm_provider, + dus.prompt_tokens, + dus.completion_tokens, + dus.spend, + dus.api_requests, + dus.successful_requests, + dus.failed_requests, + dus.cache_creation_input_tokens, + dus.cache_read_input_tokens, + dus.created_at, + dus.updated_at, + vt.team_id, + vt.key_alias as api_key_alias, + tt.team_alias, + ut.user_email as user_email + FROM "LiteLLM_DailyUserSpend" dus + LEFT JOIN "LiteLLM_VerificationToken" vt ON dus.api_key = vt.token + LEFT JOIN "LiteLLM_TeamTable" tt ON vt.team_id = tt.team_id + LEFT JOIN "LiteLLM_UserTable" ut ON dus.user_id = ut.user_id + {where_clause} + ORDER BY dus.date DESC, dus.created_at DESC + {limit_clause} + """ + + try: + db_response = await client.db.query_raw(query, *query_params) + return pl.DataFrame(db_response, infer_schema_length=None) + except Exception as exc: + raise RuntimeError(f"Error retrieving usage data: {exc}") from exc + + async def get_table_info(self) -> Dict[str, Any]: + """Return metadata about the spend table for diagnostics.""" + client = self._ensure_prisma_client() + + info_query = """ + SELECT column_name, data_type, is_nullable + FROM information_schema.columns + WHERE table_name = 'LiteLLM_DailyUserSpend' + ORDER BY ordinal_position; + """ + try: + columns_response = await client.db.query_raw(info_query) + return {"columns": columns_response, "table_name": "LiteLLM_DailyUserSpend"} + except Exception as exc: + raise RuntimeError(f"Error getting table info: {exc}") from exc diff --git a/litellm/integrations/focus/destinations/__init__.py b/litellm/integrations/focus/destinations/__init__.py new file mode 100644 index 00000000000..233f1da0c9b --- /dev/null +++ b/litellm/integrations/focus/destinations/__init__.py @@ -0,0 +1,12 @@ +"""Destination implementations for Focus export.""" + +from .base import FocusDestination, FocusTimeWindow +from .factory import FocusDestinationFactory +from .s3_destination import FocusS3Destination + +__all__ = [ + "FocusDestination", + "FocusDestinationFactory", + "FocusTimeWindow", + "FocusS3Destination", +] diff --git a/litellm/integrations/focus/destinations/base.py b/litellm/integrations/focus/destinations/base.py new file mode 100644 index 00000000000..8042a7e23b9 --- /dev/null +++ b/litellm/integrations/focus/destinations/base.py @@ -0,0 +1,30 @@ +"""Abstract destination interfaces for Focus export.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from typing import Protocol + + +@dataclass(frozen=True) +class FocusTimeWindow: + """Represents the span of data exported in a single batch.""" + + start_time: datetime + end_time: datetime + frequency: str + + +class FocusDestination(Protocol): + """Protocol for anything that can receive Focus export files.""" + + async def deliver( + self, + *, + content: bytes, + time_window: FocusTimeWindow, + filename: str, + ) -> None: + """Persist the serialized export for the provided time window.""" + ... diff --git a/litellm/integrations/focus/destinations/factory.py b/litellm/integrations/focus/destinations/factory.py new file mode 100644 index 00000000000..cb7696a11de --- /dev/null +++ b/litellm/integrations/focus/destinations/factory.py @@ -0,0 +1,59 @@ +"""Factory helpers for Focus export destinations.""" + +from __future__ import annotations + +import os +from typing import Any, Dict, Optional + +from .base import FocusDestination +from .s3_destination import FocusS3Destination + + +class FocusDestinationFactory: + """Builds destination instances based on provider/config settings.""" + + @staticmethod + def create( + *, + provider: str, + prefix: str, + config: Optional[Dict[str, Any]] = None, + ) -> FocusDestination: + """Return a destination implementation for the requested provider.""" + provider_lower = provider.lower() + normalized_config = FocusDestinationFactory._resolve_config( + provider=provider_lower, overrides=config or {} + ) + if provider_lower == "s3": + return FocusS3Destination(prefix=prefix, config=normalized_config) + raise NotImplementedError( + f"Provider '{provider}' not supported for Focus export" + ) + + @staticmethod + def _resolve_config( + *, + provider: str, + overrides: Dict[str, Any], + ) -> Dict[str, Any]: + if provider == "s3": + resolved = { + "bucket_name": overrides.get("bucket_name") + or os.getenv("FOCUS_S3_BUCKET_NAME"), + "region_name": overrides.get("region_name") + or os.getenv("FOCUS_S3_REGION_NAME"), + "endpoint_url": overrides.get("endpoint_url") + or os.getenv("FOCUS_S3_ENDPOINT_URL"), + "aws_access_key_id": overrides.get("aws_access_key_id") + or os.getenv("FOCUS_S3_ACCESS_KEY"), + "aws_secret_access_key": overrides.get("aws_secret_access_key") + or os.getenv("FOCUS_S3_SECRET_KEY"), + "aws_session_token": overrides.get("aws_session_token") + or os.getenv("FOCUS_S3_SESSION_TOKEN"), + } + if not resolved.get("bucket_name"): + raise ValueError("FOCUS_S3_BUCKET_NAME must be provided for S3 exports") + return {k: v for k, v in resolved.items() if v is not None} + raise NotImplementedError( + f"Provider '{provider}' not supported for Focus export configuration" + ) diff --git a/litellm/integrations/focus/destinations/s3_destination.py b/litellm/integrations/focus/destinations/s3_destination.py new file mode 100644 index 00000000000..c6d5554b438 --- /dev/null +++ b/litellm/integrations/focus/destinations/s3_destination.py @@ -0,0 +1,74 @@ +"""S3 destination implementation for Focus export.""" + +from __future__ import annotations + +import asyncio +from datetime import timezone +from typing import Any, Optional + +import boto3 + +from .base import FocusDestination, FocusTimeWindow + + +class FocusS3Destination(FocusDestination): + """Handles uploading serialized exports to S3 buckets.""" + + def __init__( + self, + *, + prefix: str, + config: Optional[dict[str, Any]] = None, + ) -> None: + config = config or {} + bucket_name = config.get("bucket_name") + if not bucket_name: + raise ValueError("bucket_name must be provided for S3 destination") + self.bucket_name = bucket_name + self.prefix = prefix.rstrip("/") + self.config = config + + async def deliver( + self, + *, + content: bytes, + time_window: FocusTimeWindow, + filename: str, + ) -> None: + object_key = self._build_object_key(time_window=time_window, filename=filename) + await asyncio.to_thread(self._upload, content, object_key) + + def _build_object_key(self, *, time_window: FocusTimeWindow, filename: str) -> str: + start_utc = time_window.start_time.astimezone(timezone.utc) + date_component = f"date={start_utc.strftime('%Y-%m-%d')}" + parts = [self.prefix, date_component] + if time_window.frequency == "hourly": + parts.append(f"hour={start_utc.strftime('%H')}") + key_prefix = "/".join(filter(None, parts)) + return f"{key_prefix}/{filename}" if key_prefix else filename + + def _upload(self, content: bytes, object_key: str) -> None: + client_kwargs: dict[str, Any] = {} + region_name = self.config.get("region_name") + if region_name: + client_kwargs["region_name"] = region_name + endpoint_url = self.config.get("endpoint_url") + if endpoint_url: + client_kwargs["endpoint_url"] = endpoint_url + + session_kwargs: dict[str, Any] = {} + for key in ( + "aws_access_key_id", + "aws_secret_access_key", + "aws_session_token", + ): + if self.config.get(key): + session_kwargs[key] = self.config[key] + + s3_client = boto3.client("s3", **client_kwargs, **session_kwargs) + s3_client.put_object( + Bucket=self.bucket_name, + Key=object_key, + Body=content, + ContentType="application/octet-stream", + ) diff --git a/litellm/integrations/focus/export_engine.py b/litellm/integrations/focus/export_engine.py new file mode 100644 index 00000000000..22ebce2a168 --- /dev/null +++ b/litellm/integrations/focus/export_engine.py @@ -0,0 +1,124 @@ +"""Core export engine for Focus integrations (heavy dependencies).""" + +from __future__ import annotations + +from typing import Any, Dict, Optional + +import polars as pl + +from litellm._logging import verbose_logger + +from .database import FocusLiteLLMDatabase +from .destinations import FocusDestinationFactory, FocusTimeWindow +from .serializers import FocusParquetSerializer, FocusSerializer +from .transformer import FocusTransformer + + +class FocusExportEngine: + """Engine that fetches, normalizes, and uploads Focus exports.""" + + def __init__( + self, + *, + provider: str, + export_format: str, + prefix: str, + destination_config: Optional[dict[str, Any]] = None, + ) -> None: + self.provider = provider + self.export_format = export_format + self.prefix = prefix + self._destination = FocusDestinationFactory.create( + provider=self.provider, + prefix=self.prefix, + config=destination_config, + ) + self._serializer = self._init_serializer() + self._transformer = FocusTransformer() + self._database = FocusLiteLLMDatabase() + + def _init_serializer(self) -> FocusSerializer: + if self.export_format != "parquet": + raise NotImplementedError("Only parquet export supported currently") + return FocusParquetSerializer() + + async def dry_run_export_usage_data(self, limit: Optional[int]) -> Dict[str, Any]: + data = await self._database.get_usage_data(limit=limit) + normalized = self._transformer.transform(data) + + usage_sample = data.head(min(50, len(data))).to_dicts() + normalized_sample = normalized.head(min(50, len(normalized))).to_dicts() + + summary = { + "total_records": len(normalized), + "total_spend": self._sum_column(normalized, "spend"), + "total_tokens": self._sum_column(normalized, "total_tokens"), + "unique_teams": self._count_unique(normalized, "team_id"), + "unique_models": self._count_unique(normalized, "model"), + } + + return { + "usage_data": usage_sample, + "normalized_data": normalized_sample, + "summary": summary, + } + + async def export_window( + self, + *, + window: FocusTimeWindow, + limit: Optional[int], + ) -> None: + data = await self._database.get_usage_data( + limit=limit, + start_time_utc=window.start_time, + end_time_utc=window.end_time, + ) + if data.is_empty(): + verbose_logger.debug("Focus export: no usage data for window %s", window) + return + + normalized = self._transformer.transform(data) + if normalized.is_empty(): + verbose_logger.debug( + "Focus export: normalized data empty for window %s", window + ) + return + + await self._serialize_and_upload(normalized, window) + + async def _serialize_and_upload( + self, frame: pl.DataFrame, window: FocusTimeWindow + ) -> None: + payload = self._serializer.serialize(frame) + if not payload: + verbose_logger.debug("Focus export: serializer returned empty payload") + return + await self._destination.deliver( + content=payload, + time_window=window, + filename=self._build_filename(), + ) + + def _build_filename(self) -> str: + if not self._serializer.extension: + raise ValueError("Serializer must declare a file extension") + return f"usage.{self._serializer.extension}" + + @staticmethod + def _sum_column(frame: pl.DataFrame, column: str) -> float: + if frame.is_empty() or column not in frame.columns: + return 0.0 + value = frame.select(pl.col(column).sum().alias("sum")).row(0)[0] + if value is None: + return 0.0 + return float(value) + + @staticmethod + def _count_unique(frame: pl.DataFrame, column: str) -> int: + if frame.is_empty() or column not in frame.columns: + return 0 + value = frame.select(pl.col(column).n_unique().alias("unique")).row(0)[0] + if value is None: + return 0 + return int(value) diff --git a/litellm/integrations/focus/focus_logger.py b/litellm/integrations/focus/focus_logger.py new file mode 100644 index 00000000000..ade1cf861b1 --- /dev/null +++ b/litellm/integrations/focus/focus_logger.py @@ -0,0 +1,211 @@ +"""Focus export logger orchestrating DB pull/transform/upload.""" + +from __future__ import annotations + +import os +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast + +import litellm +from litellm._logging import verbose_logger +from litellm.integrations.custom_logger import CustomLogger + +from .destinations import FocusTimeWindow + +if TYPE_CHECKING: + from apscheduler.schedulers.asyncio import AsyncIOScheduler + from .export_engine import FocusExportEngine +else: + AsyncIOScheduler = Any + +FOCUS_USAGE_DATA_JOB_NAME = "focus_export_usage_data" +DEFAULT_DRY_RUN_LIMIT = 500 + + +class FocusLogger(CustomLogger): + """Coordinates Focus export jobs across transformer/serializer/destination layers.""" + + def __init__( + self, + *, + provider: Optional[str] = None, + export_format: Optional[str] = None, + frequency: Optional[str] = None, + cron_offset_minute: Optional[int] = None, + interval_seconds: Optional[int] = None, + prefix: Optional[str] = None, + destination_config: Optional[dict[str, Any]] = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self.provider = (provider or os.getenv("FOCUS_PROVIDER") or "s3").lower() + self.export_format = ( + export_format or os.getenv("FOCUS_FORMAT") or "parquet" + ).lower() + self.frequency = (frequency or os.getenv("FOCUS_FREQUENCY") or "hourly").lower() + self.cron_offset_minute = ( + cron_offset_minute + if cron_offset_minute is not None + else int(os.getenv("FOCUS_CRON_OFFSET", "5")) + ) + raw_interval = ( + interval_seconds + if interval_seconds is not None + else os.getenv("FOCUS_INTERVAL_SECONDS") + ) + self.interval_seconds = int(raw_interval) if raw_interval is not None else None + env_prefix = os.getenv("FOCUS_PREFIX") + self.prefix: str = ( + prefix if prefix is not None else (env_prefix if env_prefix else "focus_exports") + ) + + self._destination_config = destination_config + self._engine: Optional["FocusExportEngine"] = None + + def _ensure_engine(self) -> "FocusExportEngine": + """Instantiate the heavy export engine lazily.""" + if self._engine is None: + from .export_engine import FocusExportEngine + + self._engine = FocusExportEngine( + provider=self.provider, + export_format=self.export_format, + prefix=self.prefix, + destination_config=self._destination_config, + ) + return self._engine + + async def export_usage_data( + self, + *, + limit: Optional[int] = None, + start_time_utc: Optional[datetime] = None, + end_time_utc: Optional[datetime] = None, + ) -> None: + """Public hook to trigger export immediately.""" + if bool(start_time_utc) ^ bool(end_time_utc): + raise ValueError( + "start_time_utc and end_time_utc must be provided together" + ) + + if start_time_utc and end_time_utc: + window = FocusTimeWindow( + start_time=start_time_utc, + end_time=end_time_utc, + frequency=self.frequency, + ) + else: + window = self._compute_time_window(datetime.now(timezone.utc)) + await self._export_window(window=window, limit=limit) + + async def dry_run_export_usage_data( + self, limit: Optional[int] = DEFAULT_DRY_RUN_LIMIT + ) -> dict[str, Any]: + """Return transformed data without uploading.""" + engine = self._ensure_engine() + return await engine.dry_run_export_usage_data(limit=limit) + + async def initialize_focus_export_job(self) -> None: + """Entry point for scheduler jobs to run export cycle with locking.""" + from litellm.proxy.proxy_server import proxy_logging_obj + + pod_lock_manager = None + if proxy_logging_obj is not None: + writer = getattr(proxy_logging_obj, "db_spend_update_writer", None) + if writer is not None: + pod_lock_manager = getattr(writer, "pod_lock_manager", None) + + if pod_lock_manager and pod_lock_manager.redis_cache: + acquired = await pod_lock_manager.acquire_lock( + cronjob_id=FOCUS_USAGE_DATA_JOB_NAME + ) + if not acquired: + verbose_logger.debug("Focus export: unable to acquire pod lock") + return + try: + await self._run_scheduled_export() + finally: + await pod_lock_manager.release_lock( + cronjob_id=FOCUS_USAGE_DATA_JOB_NAME + ) + else: + await self._run_scheduled_export() + + @staticmethod + async def init_focus_export_background_job( + scheduler: AsyncIOScheduler, + ) -> None: + """Register the export cron/interval job with the provided scheduler.""" + + focus_loggers: List[ + CustomLogger + ] = litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=FocusLogger + ) + if not focus_loggers: + verbose_logger.debug( + "No Focus export logger registered; skipping scheduler" + ) + return + + focus_logger = cast(FocusLogger, focus_loggers[0]) + trigger_kwargs = focus_logger._build_scheduler_trigger() + scheduler.add_job( + focus_logger.initialize_focus_export_job, + **trigger_kwargs, + ) + + def _build_scheduler_trigger(self) -> Dict[str, Any]: + """Return scheduler configuration for the selected frequency.""" + if self.frequency == "interval": + seconds = self.interval_seconds or 60 + return {"trigger": "interval", "seconds": seconds} + + if self.frequency == "hourly": + minute = max(0, min(59, self.cron_offset_minute)) + return {"trigger": "cron", "minute": minute, "second": 0} + + if self.frequency == "daily": + total_minutes = max(0, self.cron_offset_minute) + hour = min(23, total_minutes // 60) + minute = min(59, total_minutes % 60) + return {"trigger": "cron", "hour": hour, "minute": minute, "second": 0} + + raise ValueError(f"Unsupported frequency: {self.frequency}") + + async def _run_scheduled_export(self) -> None: + """Execute the scheduled export for the configured window.""" + window = self._compute_time_window(datetime.now(timezone.utc)) + await self._export_window(window=window, limit=None) + + async def _export_window( + self, + *, + window: FocusTimeWindow, + limit: Optional[int], + ) -> None: + engine = self._ensure_engine() + await engine.export_window(window=window, limit=limit) + + def _compute_time_window(self, now: datetime) -> FocusTimeWindow: + """Derive the time window to export based on configured frequency.""" + now_utc = now.astimezone(timezone.utc) + if self.frequency == "hourly": + end_time = now_utc.replace(minute=0, second=0, microsecond=0) + start_time = end_time - timedelta(hours=1) + elif self.frequency == "daily": + end_time = now_utc.replace(hour=0, minute=0, second=0, microsecond=0) + start_time = end_time - timedelta(days=1) + elif self.frequency == "interval": + interval = timedelta(seconds=self.interval_seconds or 60) + end_time = now_utc + start_time = end_time - interval + else: + raise ValueError(f"Unsupported frequency: {self.frequency}") + return FocusTimeWindow( + start_time=start_time, + end_time=end_time, + frequency=self.frequency, + ) + +__all__ = ["FocusLogger"] diff --git a/litellm/integrations/focus/schema.py b/litellm/integrations/focus/schema.py new file mode 100644 index 00000000000..ac2f33dad0a --- /dev/null +++ b/litellm/integrations/focus/schema.py @@ -0,0 +1,50 @@ +"""Schema definitions for Focus export data.""" + +from __future__ import annotations + +import polars as pl + +# see: https://focus.finops.org/focus-specification/v1-2/ +FOCUS_NORMALIZED_SCHEMA = pl.Schema( + [ + ("BilledCost", pl.Decimal(18, 6)), + ("BillingAccountId", pl.String), + ("BillingAccountName", pl.String), + ("BillingCurrency", pl.String), + ("BillingPeriodStart", pl.Datetime(time_unit="us")), + ("BillingPeriodEnd", pl.Datetime(time_unit="us")), + ("ChargeCategory", pl.String), + ("ChargeClass", pl.String), + ("ChargeDescription", pl.String), + ("ChargeFrequency", pl.String), + ("ChargePeriodStart", pl.Datetime(time_unit="us")), + ("ChargePeriodEnd", pl.Datetime(time_unit="us")), + ("ConsumedQuantity", pl.Decimal(18, 6)), + ("ConsumedUnit", pl.String), + ("ContractedCost", pl.Decimal(18, 6)), + ("ContractedUnitPrice", pl.Decimal(18, 6)), + ("EffectiveCost", pl.Decimal(18, 6)), + ("InvoiceIssuerName", pl.String), + ("ListCost", pl.Decimal(18, 6)), + ("ListUnitPrice", pl.Decimal(18, 6)), + ("PricingCategory", pl.String), + ("PricingQuantity", pl.Decimal(18, 6)), + ("PricingUnit", pl.String), + ("ProviderName", pl.String), + ("PublisherName", pl.String), + ("RegionId", pl.String), + ("RegionName", pl.String), + ("ResourceId", pl.String), + ("ResourceName", pl.String), + ("ResourceType", pl.String), + ("ServiceCategory", pl.String), + ("ServiceSubcategory", pl.String), + ("ServiceName", pl.String), + ("SubAccountId", pl.String), + ("SubAccountName", pl.String), + ("SubAccountType", pl.String), + ("Tags", pl.Object), + ] +) + +__all__ = ["FOCUS_NORMALIZED_SCHEMA"] diff --git a/litellm/integrations/focus/serializers/__init__.py b/litellm/integrations/focus/serializers/__init__.py new file mode 100644 index 00000000000..18187bf73e5 --- /dev/null +++ b/litellm/integrations/focus/serializers/__init__.py @@ -0,0 +1,6 @@ +"""Serializer package exports for Focus integration.""" + +from .base import FocusSerializer +from .parquet import FocusParquetSerializer + +__all__ = ["FocusSerializer", "FocusParquetSerializer"] diff --git a/litellm/integrations/focus/serializers/base.py b/litellm/integrations/focus/serializers/base.py new file mode 100644 index 00000000000..6da080dae81 --- /dev/null +++ b/litellm/integrations/focus/serializers/base.py @@ -0,0 +1,18 @@ +"""Serializer abstractions for Focus export.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod + +import polars as pl + + +class FocusSerializer(ABC): + """Base serializer turning Focus frames into bytes.""" + + extension: str = "" + + @abstractmethod + def serialize(self, frame: pl.DataFrame) -> bytes: + """Convert the normalized Focus frame into the chosen format.""" + raise NotImplementedError diff --git a/litellm/integrations/focus/serializers/parquet.py b/litellm/integrations/focus/serializers/parquet.py new file mode 100644 index 00000000000..6b3dde5903d --- /dev/null +++ b/litellm/integrations/focus/serializers/parquet.py @@ -0,0 +1,22 @@ +"""Parquet serializer for Focus export.""" + +from __future__ import annotations + +import io + +import polars as pl + +from .base import FocusSerializer + + +class FocusParquetSerializer(FocusSerializer): + """Serialize normalized Focus frames to Parquet bytes.""" + + extension = "parquet" + + def serialize(self, frame: pl.DataFrame) -> bytes: + """Encode the provided frame as a parquet payload.""" + target = frame if not frame.is_empty() else pl.DataFrame(schema=frame.schema) + buffer = io.BytesIO() + target.write_parquet(buffer, compression="snappy") + return buffer.getvalue() diff --git a/litellm/integrations/focus/transformer.py b/litellm/integrations/focus/transformer.py new file mode 100644 index 00000000000..cac12b7be14 --- /dev/null +++ b/litellm/integrations/focus/transformer.py @@ -0,0 +1,90 @@ +"""Focus export data transformer.""" + +from __future__ import annotations + +from datetime import timedelta + +import polars as pl + +from .schema import FOCUS_NORMALIZED_SCHEMA + + +class FocusTransformer: + """Transforms LiteLLM DB rows into Focus-compatible schema.""" + + schema = FOCUS_NORMALIZED_SCHEMA + + def transform(self, frame: pl.DataFrame) -> pl.DataFrame: + """Return a normalized frame expected by downstream serializers.""" + if frame.is_empty(): + return pl.DataFrame(schema=self.schema) + + # derive period start/end from usage date + frame = frame.with_columns( + pl.col("date") + .cast(pl.Utf8) + .str.strptime(pl.Datetime(time_unit="us"), format="%Y-%m-%d", strict=False) + .alias("usage_date"), + ) + frame = frame.with_columns( + pl.col("usage_date").alias("ChargePeriodStart"), + (pl.col("usage_date") + timedelta(days=1)).alias("ChargePeriodEnd"), + ) + + def fmt(col): + return col.dt.strftime("%Y-%m-%dT%H:%M:%SZ") + + DEC = pl.Decimal(18, 6) + + def dec(col): + return col.cast(DEC) + + none_str = pl.lit(None, dtype=pl.Utf8) + none_dec = pl.lit(None, dtype=pl.Decimal(18, 6)) + + return frame.select( + dec(pl.col("spend").fill_null(0.0)).alias("BilledCost"), + pl.col("api_key").cast(pl.String).alias("BillingAccountId"), + pl.col("api_key_alias").cast(pl.String).alias("BillingAccountName"), + pl.lit("API Key").alias("BillingAccountType"), + pl.lit("USD").alias("BillingCurrency"), + fmt(pl.col("ChargePeriodEnd")).alias("BillingPeriodEnd"), + fmt(pl.col("ChargePeriodStart")).alias("BillingPeriodStart"), + pl.lit("Usage").alias("ChargeCategory"), + none_str.alias("ChargeClass"), + pl.col("model").cast(pl.String).alias("ChargeDescription"), + pl.lit("Usage-Based").alias("ChargeFrequency"), + fmt(pl.col("ChargePeriodEnd")).alias("ChargePeriodEnd"), + fmt(pl.col("ChargePeriodStart")).alias("ChargePeriodStart"), + dec(pl.lit(1.0)).alias("ConsumedQuantity"), + pl.lit("Requests").alias("ConsumedUnit"), + dec(pl.col("spend").fill_null(0.0)).alias("ContractedCost"), + none_str.alias("ContractedUnitPrice"), + dec(pl.col("spend").fill_null(0.0)).alias("EffectiveCost"), + pl.col("custom_llm_provider").cast(pl.String).alias("InvoiceIssuerName"), + none_str.alias("InvoiceId"), + dec(pl.col("spend").fill_null(0.0)).alias("ListCost"), + none_dec.alias("ListUnitPrice"), + none_str.alias("AvailabilityZone"), + pl.lit("USD").alias("PricingCurrency"), + none_str.alias("PricingCategory"), + dec(pl.lit(1.0)).alias("PricingQuantity"), + none_dec.alias("PricingCurrencyContractedUnitPrice"), + dec(pl.col("spend").fill_null(0.0)).alias("PricingCurrencyEffectiveCost"), + none_dec.alias("PricingCurrencyListUnitPrice"), + pl.lit("Requests").alias("PricingUnit"), + pl.col("custom_llm_provider").cast(pl.String).alias("ProviderName"), + pl.col("custom_llm_provider").cast(pl.String).alias("PublisherName"), + none_str.alias("RegionId"), + none_str.alias("RegionName"), + pl.col("model").cast(pl.String).alias("ResourceId"), + pl.col("model").cast(pl.String).alias("ResourceName"), + pl.col("model").cast(pl.String).alias("ResourceType"), + pl.lit("AI and Machine Learning").alias("ServiceCategory"), + pl.lit("Generative AI").alias("ServiceSubcategory"), + pl.col("model_group").cast(pl.String).alias("ServiceName"), + pl.col("team_id").cast(pl.String).alias("SubAccountId"), + pl.col("team_alias").cast(pl.String).alias("SubAccountName"), + none_str.alias("SubAccountType"), + none_str.alias("Tags"), + ) diff --git a/litellm/integrations/generic_api/generic_api_compatible_callbacks.json b/litellm/integrations/generic_api/generic_api_compatible_callbacks.json index 12dc4ae643c..13fe79ae671 100644 --- a/litellm/integrations/generic_api/generic_api_compatible_callbacks.json +++ b/litellm/integrations/generic_api/generic_api_compatible_callbacks.json @@ -1,28 +1,37 @@ { - "sample_callback": { - "event_types": ["llm_api_success", "llm_api_failure"], - "endpoint": "{{environment_variables.SAMPLE_CALLBACK_URL}}", - "headers": { - "Content-Type": "application/json", - "Authorization": "Bearer {{environment_variables.SAMPLE_CALLBACK_API_KEY}}" - }, - "environment_variables": ["SAMPLE_CALLBACK_URL", "SAMPLE_CALLBACK_API_KEY"] + "sample_callback": { + "event_types": ["llm_api_success", "llm_api_failure"], + "endpoint": "{{environment_variables.SAMPLE_CALLBACK_URL}}", + "headers": { + "Content-Type": "application/json", + "Authorization": "Bearer {{environment_variables.SAMPLE_CALLBACK_API_KEY}}" }, - "rubrik": { - "event_types": ["llm_api_success"], - "endpoint": "{{environment_variables.RUBRIK_WEBHOOK_URL}}", - "headers": { - "Content-Type": "application/json", - "Authorization": "Bearer {{environment_variables.RUBRIK_API_KEY}}" - }, - "environment_variables": ["RUBRIK_API_KEY", "RUBRIK_WEBHOOK_URL"] + "environment_variables": ["SAMPLE_CALLBACK_URL", "SAMPLE_CALLBACK_API_KEY"] + }, + "rubrik": { + "event_types": ["llm_api_success"], + "endpoint": "{{environment_variables.RUBRIK_WEBHOOK_URL}}", + "headers": { + "Content-Type": "application/json", + "Authorization": "Bearer {{environment_variables.RUBRIK_API_KEY}}" }, - "sumologic": { - "endpoint": "{{environment_variables.SUMOLOGIC_WEBHOOK_URL}}", - "headers": { - "Content-Type": "application/json" - }, - "environment_variables": ["SUMOLOGIC_WEBHOOK_URL"], - "log_format": "ndjson" - } -} \ No newline at end of file + "environment_variables": ["RUBRIK_API_KEY", "RUBRIK_WEBHOOK_URL"] + }, + "sumologic": { + "endpoint": "{{environment_variables.SUMOLOGIC_WEBHOOK_URL}}", + "headers": { + "Content-Type": "application/json" + }, + "environment_variables": ["SUMOLOGIC_WEBHOOK_URL"], + "log_format": "ndjson" + }, + "qualifire_eval": { + "event_types": ["llm_api_success"], + "endpoint": "{{environment_variables.QUALIFIRE_WEBHOOK_URL}}", + "headers": { + "Content-Type": "application/json", + "X-Qualifire-API-Key": "{{environment_variables.QUALIFIRE_API_KEY}}" + }, + "environment_variables": ["QUALIFIRE_API_KEY", "QUALIFIRE_WEBHOOK_URL"] + } +} diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index c01f7481277..c32a7b75c51 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -14,6 +14,7 @@ from typing import ( Literal, Optional, Tuple, + Union, cast, ) @@ -44,6 +45,7 @@ def _get_cached_end_user_id_for_cost_tracking(): global _get_end_user_id_for_cost_tracking if _get_end_user_id_for_cost_tracking is None: from litellm.utils import get_end_user_id_for_cost_tracking + _get_end_user_id_for_cost_tracking = get_end_user_id_for_cost_tracking return _get_end_user_id_for_cost_tracking @@ -237,6 +239,36 @@ class PrometheusLogger(CustomLogger): ), buckets=LATENCY_BUCKETS, ) + + # Request queue time metric + self.litellm_request_queue_time_metric = self._histogram_factory( + "litellm_request_queue_time_seconds", + "Time spent in request queue before processing starts (seconds)", + labelnames=self.get_labels_for_metric( + "litellm_request_queue_time_seconds" + ), + buckets=LATENCY_BUCKETS, + ) + + # Guardrail metrics + self.litellm_guardrail_latency_metric = self._histogram_factory( + "litellm_guardrail_latency_seconds", + "Latency (seconds) for guardrail execution", + labelnames=["guardrail_name", "status", "error_type", "hook_type"], + buckets=LATENCY_BUCKETS, + ) + + self.litellm_guardrail_errors_total = self._counter_factory( + "litellm_guardrail_errors_total", + "Total number of errors encountered during guardrail execution", + labelnames=["guardrail_name", "error_type", "hook_type"], + ) + + self.litellm_guardrail_requests_total = self._counter_factory( + "litellm_guardrail_requests_total", + "Total number of guardrail invocations", + labelnames=["guardrail_name", "status", "hook_type"], + ) # llm api provider budget metrics self.litellm_provider_remaining_budget_metric = self._gauge_factory( "litellm_provider_remaining_budget_metric", @@ -329,6 +361,25 @@ class PrometheusLogger(CustomLogger): labelnames=self.get_labels_for_metric("litellm_requests_metric"), ) + # Cache metrics + self.litellm_cache_hits_metric = self._counter_factory( + name="litellm_cache_hits_metric", + documentation="Total number of LiteLLM cache hits", + labelnames=self.get_labels_for_metric("litellm_cache_hits_metric"), + ) + + self.litellm_cache_misses_metric = self._counter_factory( + name="litellm_cache_misses_metric", + documentation="Total number of LiteLLM cache misses", + labelnames=self.get_labels_for_metric("litellm_cache_misses_metric"), + ) + + self.litellm_cached_tokens_metric = self._counter_factory( + name="litellm_cached_tokens_metric", + documentation="Total tokens served from LiteLLM cache", + labelnames=self.get_labels_for_metric("litellm_cached_tokens_metric"), + ) + except Exception as e: print_verbose(f"Got exception on init prometheus client {str(e)}") raise e @@ -791,11 +842,16 @@ class PrometheusLogger(CustomLogger): f"standard_logging_object is required, got={standard_logging_payload}" ) + if self._should_skip_metrics_for_invalid_key( + kwargs=kwargs, standard_logging_payload=standard_logging_payload + ): + return + model = kwargs.get("model", "") litellm_params = kwargs.get("litellm_params", {}) or {} _metadata = litellm_params.get("metadata", {}) get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking() - + end_user_id = get_end_user_id_for_cost_tracking( litellm_params, service_type="prometheus" ) @@ -815,7 +871,7 @@ class PrometheusLogger(CustomLogger): user_api_key_auth_metadata: Optional[dict] = standard_logging_payload[ "metadata" ].get("user_api_key_auth_metadata") - + # Include top-level metadata fields (excluding nested dictionaries) # This allows accessing fields like requester_ip_address from top-level metadata top_level_metadata = standard_logging_payload.get("metadata", {}) @@ -826,7 +882,7 @@ class PrometheusLogger(CustomLogger): for k, v in top_level_metadata.items() if not isinstance(v, dict) # Exclude nested dicts to avoid conflicts } - + combined_metadata: Dict[str, Any] = { **top_level_fields, # Include top-level fields first **(_requester_metadata if _requester_metadata else {}), @@ -945,6 +1001,12 @@ class PrometheusLogger(CustomLogger): kwargs, start_time, end_time, enum_values, output_tokens ) + # cache metrics + self._increment_cache_metrics( + standard_logging_payload=standard_logging_payload, # type: ignore + enum_values=enum_values, + ) + if ( standard_logging_payload["stream"] is True ): # log successful streaming requests from logging event hook. @@ -1014,6 +1076,54 @@ class PrometheusLogger(CustomLogger): standard_logging_payload["completion_tokens"] ) + def _increment_cache_metrics( + self, + standard_logging_payload: StandardLoggingPayload, + enum_values: UserAPIKeyLabelValues, + ): + """ + Increment cache-related Prometheus metrics based on cache hit/miss status. + + Args: + standard_logging_payload: Contains cache_hit field (True/False/None) + enum_values: Label values for Prometheus metrics + """ + cache_hit = standard_logging_payload.get("cache_hit") + + # Only track if cache_hit has a definite value (True or False) + if cache_hit is None: + return + + if cache_hit is True: + # Increment cache hits counter + _labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_cache_hits_metric" + ), + enum_values=enum_values, + ) + self.litellm_cache_hits_metric.labels(**_labels).inc() + + # Increment cached tokens counter + total_tokens = standard_logging_payload.get("total_tokens", 0) + if total_tokens > 0: + _labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_cached_tokens_metric" + ), + enum_values=enum_values, + ) + self.litellm_cached_tokens_metric.labels(**_labels).inc(total_tokens) + else: + # cache_hit is False - increment cache misses counter + _labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_cache_misses_metric" + ), + enum_values=enum_values, + ) + self.litellm_cache_misses_metric.labels(**_labels).inc() + async def _increment_remaining_budget_metrics( self, user_api_team: Optional[str], @@ -1182,6 +1292,22 @@ class PrometheusLogger(CustomLogger): total_time_seconds ) + # request queue time (time from arrival to processing start) + _litellm_params = kwargs.get("litellm_params", {}) or {} + queue_time_seconds = _litellm_params.get("metadata", {}).get( + "queue_time_seconds" + ) + if queue_time_seconds is not None and queue_time_seconds >= 0: + _labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_request_queue_time_seconds" + ), + enum_values=enum_values, + ) + self.litellm_request_queue_time_metric.labels(**_labels).observe( + queue_time_seconds + ) + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): from litellm.types.utils import StandardLoggingPayload @@ -1189,14 +1315,20 @@ class PrometheusLogger(CustomLogger): f"prometheus Logging - Enters failure logging function for kwargs {kwargs}" ) - # unpack kwargs - model = kwargs.get("model", "") standard_logging_payload: StandardLoggingPayload = kwargs.get( "standard_logging_object", {} ) + + if self._should_skip_metrics_for_invalid_key( + kwargs=kwargs, standard_logging_payload=standard_logging_payload + ): + return + + model = kwargs.get("model", "") + litellm_params = kwargs.get("litellm_params", {}) or {} get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking() - + end_user_id = get_end_user_id_for_cost_tracking( litellm_params, service_type="prometheus" ) @@ -1207,7 +1339,6 @@ class PrometheusLogger(CustomLogger): user_api_team_alias = standard_logging_payload["metadata"][ "user_api_key_team_alias" ] - kwargs.get("exception", None) try: self.litellm_llm_api_failed_requests_metric.labels( @@ -1227,6 +1358,139 @@ class PrometheusLogger(CustomLogger): pass pass + def _extract_status_code( + self, + kwargs: Optional[dict] = None, + enum_values: Optional[Any] = None, + exception: Optional[Exception] = None, + ) -> Optional[int]: + """ + Extract HTTP status code from various input formats for validation. + + This is a centralized helper to extract status code from different + callback function signatures. Handles both ProxyException (uses 'code') + and standard exceptions (uses 'status_code'). + + Args: + kwargs: Dictionary potentially containing 'exception' key + enum_values: Object with 'status_code' attribute + exception: Exception object to extract status code from directly + + Returns: + Status code as integer if found, None otherwise + """ + status_code = None + + # Try from enum_values first (most common in our callbacks) + if enum_values and hasattr(enum_values, "status_code") and enum_values.status_code: + try: + status_code = int(enum_values.status_code) + except (ValueError, TypeError): + pass + + if not status_code and exception: + # ProxyException uses 'code' attribute, other exceptions may use 'status_code' + status_code = getattr(exception, "status_code", None) or getattr(exception, "code", None) + if status_code is not None: + try: + status_code = int(status_code) + except (ValueError, TypeError): + status_code = None + + if not status_code and kwargs: + exception_in_kwargs = kwargs.get("exception") + if exception_in_kwargs: + status_code = getattr(exception_in_kwargs, "status_code", None) or getattr(exception_in_kwargs, "code", None) + if status_code is not None: + try: + status_code = int(status_code) + except (ValueError, TypeError): + status_code = None + + return status_code + + def _is_invalid_api_key_request( + self, + status_code: Optional[int], + exception: Optional[Exception] = None, + ) -> bool: + """ + Determine if a request has an invalid API key based on status code and exception. + + This method prevents invalid authentication attempts from being recorded in + Prometheus metrics. A 401 status code is the definitive indicator of authentication + failure. Additionally, we check exception messages for authentication error patterns + to catch cases where the exception hasn't been converted to a ProxyException yet. + + Args: + status_code: HTTP status code (401 indicates authentication error) + exception: Exception object to check for auth-related error messages + + Returns: + True if the request has an invalid API key and metrics should be skipped, + False otherwise + """ + if status_code == 401: + return True + + # Handle cases where AssertionError is raised before conversion to ProxyException + if exception is not None: + exception_str = str(exception).lower() + auth_error_patterns = [ + "virtual key expected", + "expected to start with 'sk-'", + "authentication error", + "invalid api key", + "api key not valid", + ] + if any(pattern in exception_str for pattern in auth_error_patterns): + return True + + return False + + def _should_skip_metrics_for_invalid_key( + self, + kwargs: Optional[dict] = None, + user_api_key_dict: Optional[Any] = None, + enum_values: Optional[Any] = None, + standard_logging_payload: Optional[Union[dict, StandardLoggingPayload]] = None, + exception: Optional[Exception] = None, + ) -> bool: + """ + Determine if Prometheus metrics should be skipped for invalid API key requests. + + This is a centralized validation method that extracts status code and exception + information from various callback function signatures and determines if the request + represents an invalid API key attempt that should be filtered from metrics. + + Args: + kwargs: Dictionary potentially containing exception and other data + user_api_key_dict: User API key authentication object (currently unused) + enum_values: Object with status_code attribute + standard_logging_payload: Standard logging payload dictionary + exception: Exception object to check directly + + Returns: + True if metrics should be skipped (invalid key detected), False otherwise + """ + status_code = self._extract_status_code( + kwargs=kwargs, + enum_values=enum_values, + exception=exception, + ) + + if exception is None and kwargs: + exception = kwargs.get("exception") + + if self._is_invalid_api_key_request(status_code, exception=exception): + verbose_logger.debug( + "Skipping Prometheus metrics for invalid API key request: " + f"status_code={status_code}, exception={type(exception).__name__ if exception else None}" + ) + return True + + return False + async def async_post_call_failure_hook( self, request_data: dict, @@ -1252,6 +1516,14 @@ class PrometheusLogger(CustomLogger): StandardLoggingPayloadSetup, ) + if self._should_skip_metrics_for_invalid_key( + user_api_key_dict=user_api_key_dict, + exception=original_exception, + ): + return + + status_code = self._extract_status_code(exception=original_exception) + try: _tags = StandardLoggingPayloadSetup._get_request_tags( litellm_params=request_data, @@ -1266,8 +1538,8 @@ class PrometheusLogger(CustomLogger): team=user_api_key_dict.team_id, team_alias=user_api_key_dict.team_alias, requested_model=request_data.get("model", ""), - status_code=str(getattr(original_exception, "status_code", None)), - exception_status=str(getattr(original_exception, "status_code", None)), + status_code=str(status_code), + exception_status=str(status_code), exception_class=self._get_exception_class_name(original_exception), tags=_tags, route=user_api_key_dict.request_route, @@ -1305,6 +1577,11 @@ class PrometheusLogger(CustomLogger): StandardLoggingPayloadSetup, ) + if self._should_skip_metrics_for_invalid_key( + user_api_key_dict=user_api_key_dict + ): + return + enum_values = UserAPIKeyLabelValues( end_user=user_api_key_dict.end_user_id, hashed_api_key=user_api_key_dict.api_key, @@ -1360,6 +1637,15 @@ class PrometheusLogger(CustomLogger): exception = request_kwargs.get("exception", None) llm_provider = _litellm_params.get("custom_llm_provider", None) + + if self._should_skip_metrics_for_invalid_key( + kwargs=request_kwargs, + standard_logging_payload=standard_logging_payload, + ): + return + hashed_api_key = standard_logging_payload.get("metadata", {}).get( + "user_api_key_hash" + ) # Create enum_values for the label factory (always create for use in different metrics) enum_values = UserAPIKeyLabelValues( @@ -1374,9 +1660,7 @@ class PrometheusLogger(CustomLogger): self._get_exception_class_name(exception) if exception else None ), requested_model=model_group, - hashed_api_key=standard_logging_payload["metadata"][ - "user_api_key_hash" - ], + hashed_api_key=hashed_api_key, api_key_alias=standard_logging_payload["metadata"][ "user_api_key_alias" ], @@ -1398,7 +1682,6 @@ class PrometheusLogger(CustomLogger): api_provider=llm_provider or "", ) if exception is not None: - _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric( metric_name="litellm_deployment_failure_responses" @@ -1431,16 +1714,23 @@ class PrometheusLogger(CustomLogger): enum_values: UserAPIKeyLabelValues, output_tokens: float = 1.0, ): - try: verbose_logger.debug("setting remaining tokens requests metric") - standard_logging_payload: Optional[StandardLoggingPayload] = ( - request_kwargs.get("standard_logging_object") - ) + standard_logging_payload: Optional[ + StandardLoggingPayload + ] = request_kwargs.get("standard_logging_object") if standard_logging_payload is None: return + # Skip recording metrics for invalid API key requests + if self._should_skip_metrics_for_invalid_key( + kwargs=request_kwargs, + enum_values=enum_values, + standard_logging_payload=standard_logging_payload, + ): + return + api_base = standard_logging_payload["api_base"] _litellm_params = request_kwargs.get("litellm_params", {}) or {} _metadata = _litellm_params.get("metadata", {}) @@ -1571,6 +1861,50 @@ class PrometheusLogger(CustomLogger): ) return + def _record_guardrail_metrics( + self, + guardrail_name: str, + latency_seconds: float, + status: str, + error_type: Optional[str], + hook_type: str, + ): + """ + Record guardrail metrics for prometheus. + + Args: + guardrail_name: Name of the guardrail + latency_seconds: Execution latency in seconds + status: "success" or "error" + error_type: Type of error if any, None otherwise + hook_type: "pre_call", "during_call", or "post_call" + """ + try: + # Record latency + self.litellm_guardrail_latency_metric.labels( + guardrail_name=guardrail_name, + status=status, + error_type=error_type or "none", + hook_type=hook_type, + ).observe(latency_seconds) + + # Record request count + self.litellm_guardrail_requests_total.labels( + guardrail_name=guardrail_name, + status=status, + hook_type=hook_type, + ).inc() + + # Record error count if there was an error + if status == "error" and error_type: + self.litellm_guardrail_errors_total.labels( + guardrail_name=guardrail_name, + error_type=error_type, + hook_type=hook_type, + ).inc() + except Exception as e: + verbose_logger.debug(f"Error recording guardrail metrics: {str(e)}") + @staticmethod def _get_exception_class_name(exception: Exception) -> str: exception_class_name = "" @@ -2208,10 +2542,10 @@ class PrometheusLogger(CustomLogger): from litellm.constants import PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES from litellm.integrations.custom_logger import CustomLogger - prometheus_loggers: List[CustomLogger] = ( - litellm.logging_callback_manager.get_custom_loggers_for_type( - callback_type=PrometheusLogger - ) + prometheus_loggers: List[ + CustomLogger + ] = litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=PrometheusLogger ) # we need to get the initialized prometheus logger instance(s) and call logger.initialize_remaining_budget_metrics() on them verbose_logger.debug("found %s prometheus loggers", len(prometheus_loggers)) @@ -2283,7 +2617,7 @@ def prometheus_label_factory( if UserAPIKeyLabelNames.END_USER.value in filtered_labels: get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking() - + filtered_labels["end_user"] = get_end_user_id_for_cost_tracking( litellm_params={"user_api_key_end_user_id": enum_values.end_user}, service_type="prometheus", diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 47cbcb8aec9..a3c25ab65e9 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -18,6 +18,7 @@ from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLog from litellm.integrations.bitbucket import BitBucketPromptManager from litellm.integrations.braintrust_logging import BraintrustLogger from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger +from litellm.integrations.focus.focus_logger import FocusLogger from litellm.integrations.datadog.datadog import DataDogLogger from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger from litellm.integrations.deepeval import DeepEvalLogger @@ -93,6 +94,7 @@ class CustomLoggerRegistry: "bitbucket": BitBucketPromptManager, "gitlab": GitLabPromptManager, "cloudzero": CloudZeroLogger, + "focus": FocusLogger, "posthog": PostHogLogger, } diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index b753e9fa8b5..21d69177336 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -913,6 +913,14 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915 or "http://localhost:2024" ) dynamic_api_key = api_key or get_secret_str("LANGGRAPH_API_KEY") + elif custom_llm_provider == "manus": + # Manus is OpenAI compatible for responses API + api_base = ( + api_base + or get_secret_str("MANUS_API_BASE") + or "https://api.manus.im" + ) + dynamic_api_key = api_key or get_secret_str("MANUS_API_KEY") if api_base is not None and not isinstance(api_base, str): raise Exception("api base needs to be a string. api_base={}".format(api_base)) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 5448fe7c771..e0a799d8e5c 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3756,6 +3756,15 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 cloudzero_logger = CloudZeroLogger() _in_memory_loggers.append(cloudzero_logger) return cloudzero_logger # type: ignore + elif logging_integration == "focus": + from litellm.integrations.focus.focus_logger import FocusLogger + + for callback in _in_memory_loggers: + if isinstance(callback, FocusLogger): + return callback # type: ignore + focus_logger = FocusLogger() + _in_memory_loggers.append(focus_logger) + return focus_logger # type: ignore elif logging_integration == "deepeval": for callback in _in_memory_loggers: if isinstance(callback, DeepEvalLogger): @@ -4076,6 +4085,12 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 for callback in _in_memory_loggers: if isinstance(callback, CloudZeroLogger): return callback + elif logging_integration == "focus": + from litellm.integrations.focus.focus_logger import FocusLogger + + for callback in _in_memory_loggers: + if isinstance(callback, FocusLogger): + return callback elif logging_integration == "deepeval": for callback in _in_memory_loggers: if isinstance(callback, DeepEvalLogger): @@ -4800,7 +4815,7 @@ class StandardLoggingPayloadSetup: """ Extract additional header tags for spend tracking based on config. """ - extra_headers: List[str] = litellm.extra_spend_tag_headers or [] + extra_headers: List[str] = getattr(litellm, "extra_spend_tag_headers", None) or [] if not extra_headers: return None @@ -4824,9 +4839,9 @@ class StandardLoggingPayloadSetup: metadata = litellm_params.get("metadata") or {} litellm_metadata = litellm_params.get("litellm_metadata") or {} if metadata.get("tags", []): - request_tags = metadata.get("tags", []) + request_tags = metadata.get("tags", []).copy() elif litellm_metadata.get("tags", []): - request_tags = litellm_metadata.get("tags", []) + request_tags = litellm_metadata.get("tags", []).copy() else: request_tags = [] user_agent_tags = StandardLoggingPayloadSetup._get_user_agent_tags( diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index b100b9b516b..2f8568db704 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -6,6 +6,7 @@ import io import mimetypes import re from os import PathLike +from pathlib import Path from typing import ( TYPE_CHECKING, Any, @@ -533,6 +534,12 @@ def extract_file_data(file_data: FileTypes) -> ExtractedFileData: # Convert content to bytes if isinstance(file_content, (str, PathLike)): # If it's a path, open and read the file + # Extract filename from path if not already set + if filename is None: + if isinstance(file_content, PathLike): + filename = Path(file_content).name + else: + filename = Path(str(file_content)).name with open(file_content, "rb") as f: content = f.read() elif isinstance(file_content, io.IOBase): @@ -550,11 +557,11 @@ def extract_file_data(file_data: FileTypes) -> ExtractedFileData: # Use provided content type or guess based on filename if not content_type: - content_type = ( - mimetypes.guess_type(filename)[0] - if filename - else "application/octet-stream" - ) + if filename: + guessed_type = mimetypes.guess_type(filename)[0] + content_type = guessed_type if guessed_type else "application/octet-stream" + else: + content_type = "application/octet-stream" return ExtractedFileData( filename=filename, diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 0c331e43038..d8e82199272 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1645,9 +1645,12 @@ def convert_to_anthropic_tool_result( ) elif content["type"] == "image_url": format = content["image_url"].get("format") if isinstance(content["image_url"], dict) else None - anthropic_content_list.append( - create_anthropic_image_param(content["image_url"], format=format) + _anthropic_image_param = create_anthropic_image_param(content["image_url"], format=format) + _anthropic_image_param = add_cache_control_to_content( + anthropic_content_element=_anthropic_image_param, + original_content_element=content, ) + anthropic_content_list.append(_anthropic_image_param) anthropic_content = anthropic_content_list anthropic_tool_result: Optional[AnthropicMessagesToolResultParam] = None diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index 206810943ca..8b6ae744637 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -1,4 +1,5 @@ -from typing import Any, Dict, Optional, Set +from collections.abc import Mapping +from typing import Any, Dict, List, Optional, Set from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER @@ -17,6 +18,7 @@ class SensitiveDataMasker: "key", "token", "auth", + "authorization", "credential", "access", "private", @@ -42,22 +44,52 @@ class SensitiveDataMasker: else: return f"{value_str[:self.visible_prefix]}{self.mask_char * masked_length}{value_str[-self.visible_suffix:]}" - def is_sensitive_key(self, key: str, excluded_keys: Optional[Set[str]] = None) -> bool: + def is_sensitive_key( + self, key: str, excluded_keys: Optional[Set[str]] = None + ) -> bool: # Check if key is in excluded_keys first (exact match) if excluded_keys and key in excluded_keys: return False - + key_lower = str(key).lower() - # Split on underscores and check if any segment matches the pattern + # Split on underscores/hyphens and check if any segment matches the pattern # This avoids false positives like "max_tokens" matching "token" # but still catches "api_key", "access_token", etc. - key_segments = key_lower.replace('-', '_').split('_') - result = any( - pattern in key_segments - for pattern in self.sensitive_patterns - ) + key_segments = key_lower.replace("-", "_").split("_") + result = any(pattern in key_segments for pattern in self.sensitive_patterns) return result + def _mask_sequence( + self, + values: List[Any], + depth: int, + max_depth: int, + excluded_keys: Optional[Set[str]], + key_is_sensitive: bool, + ) -> List[Any]: + masked_items: List[Any] = [] + if depth >= max_depth: + return values + + for item in values: + if isinstance(item, Mapping): + masked_items.append( + self.mask_dict(dict(item), depth + 1, max_depth, excluded_keys) + ) + elif isinstance(item, list): + masked_items.append( + self._mask_sequence( + item, depth + 1, max_depth, excluded_keys, key_is_sensitive + ) + ) + elif key_is_sensitive and isinstance(item, str): + masked_items.append(self._mask_value(item)) + else: + masked_items.append( + item if isinstance(item, (int, float, bool, str, list)) else str(item) + ) + return masked_items + def mask_dict( self, data: Dict[str, Any], @@ -71,11 +103,20 @@ class SensitiveDataMasker: masked_data: Dict[str, Any] = {} for k, v in data.items(): try: - if isinstance(v, dict): - masked_data[k] = self.mask_dict(v, depth + 1, max_depth, excluded_keys) + key_is_sensitive = self.is_sensitive_key(k, excluded_keys) + if isinstance(v, Mapping): + masked_data[k] = self.mask_dict( + dict(v), depth + 1, max_depth, excluded_keys + ) + elif isinstance(v, list): + masked_data[k] = self._mask_sequence( + v, depth + 1, max_depth, excluded_keys, key_is_sensitive + ) elif hasattr(v, "__dict__") and not isinstance(v, type): - masked_data[k] = self.mask_dict(vars(v), depth + 1, max_depth, excluded_keys) - elif self.is_sensitive_key(k, excluded_keys): + masked_data[k] = self.mask_dict( + vars(v), depth + 1, max_depth, excluded_keys + ) + elif key_is_sensitive: str_value = str(v) if v is not None else "" masked_data[k] = self._mask_value(str_value) else: diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index c71edcdc2d1..57391c152cb 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1265,14 +1265,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): cache_creation_tokens=cache_creation_input_tokens, cache_creation_token_details=cache_creation_token_details, ) - completion_token_details = ( - CompletionTokensDetailsWrapper( - reasoning_tokens=token_counter( - text=reasoning_content, count_response_tokens=True - ) - ) + # Always populate completion_token_details, not just when there's reasoning_content + reasoning_tokens = ( + token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content - else None + else 0 + ) + completion_token_details = CompletionTokensDetailsWrapper( + reasoning_tokens=reasoning_tokens if reasoning_tokens > 0 else None, + text_tokens=completion_tokens - reasoning_tokens if reasoning_tokens > 0 else completion_tokens, ) total_tokens = prompt_tokens + completion_tokens diff --git a/litellm/llms/azure/chat/gpt_5_transformation.py b/litellm/llms/azure/chat/gpt_5_transformation.py index 87f81d117f0..506b7fdfe5e 100644 --- a/litellm/llms/azure/chat/gpt_5_transformation.py +++ b/litellm/llms/azure/chat/gpt_5_transformation.py @@ -25,7 +25,24 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): return "gpt-5" in model or "gpt5_series" in model def get_supported_openai_params(self, model: str) -> List[str]: - return OpenAIGPT5Config.get_supported_openai_params(self, model=model) + """Get supported parameters for Azure OpenAI GPT-5 models. + + Azure OpenAI GPT-5.2 models support logprobs, unlike OpenAI's GPT-5. + This overrides the parent class to add logprobs support back for gpt-5.2. + + Reference: + - Tested with Azure OpenAI GPT-5.2 (api-version: 2025-01-01-preview) + - Azure returns logprobs successfully despite Microsoft's general + documentation stating reasoning models don't support it. + """ + params = OpenAIGPT5Config.get_supported_openai_params(self, model=model) + + # Only gpt-5.2 has been verified to support logprobs on Azure + if self.is_model_gpt_5_2_model(model): + azure_supported_params = ["logprobs", "top_logprobs"] + params.extend(azure_supported_params) + + return params def map_openai_params( self, diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 71d21001cc3..18e9deb53b0 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -74,6 +74,41 @@ class BaseAWSLLM: "aws_external_id", ] + def _get_ssl_verify(self): + """ + Get SSL verification setting for boto3 clients. + + This ensures that custom CA certificates are properly used for all AWS API calls, + including STS and Bedrock services. + + Returns: + Union[bool, str]: SSL verification setting - False to disable, True to enable, + or a string path to a CA bundle file + """ + import litellm + from litellm.secret_managers.main import str_to_bool + + # Check environment variable first (highest priority) + ssl_verify = os.getenv("SSL_VERIFY", litellm.ssl_verify) + + # Convert string "False"/"True" to boolean + if isinstance(ssl_verify, str): + # Check if it's a file path + if os.path.exists(ssl_verify): + return ssl_verify + # Otherwise try to convert to boolean + ssl_verify_bool = str_to_bool(ssl_verify) + if ssl_verify_bool is not None: + ssl_verify = ssl_verify_bool + + # Check SSL_CERT_FILE environment variable for custom CA bundle + if ssl_verify is True or ssl_verify == "True": + ssl_cert_file = os.getenv("SSL_CERT_FILE") + if ssl_cert_file and os.path.exists(ssl_cert_file): + return ssl_cert_file + + return ssl_verify + def get_cache_key(self, credential_args: Dict[str, Optional[str]]) -> str: """ Generate a unique cache key based on the credential arguments. @@ -314,6 +349,12 @@ class BaseAWSLLM: if model.startswith("invoke/"): model = model.replace("invoke/", "", 1) + # Special case: Check for "nova" in model name first (before "amazon") + # This handles amazon.nova-* models which would otherwise match "amazon" (Titan) + if "nova" in model.lower(): + if "nova" in get_args(BEDROCK_INVOKE_PROVIDERS_LITERAL): + return cast(BEDROCK_INVOKE_PROVIDERS_LITERAL, "nova") + _split_model = model.split(".")[0] if _split_model in get_args(BEDROCK_INVOKE_PROVIDERS_LITERAL): return cast(BEDROCK_INVOKE_PROVIDERS_LITERAL, _split_model) @@ -323,13 +364,9 @@ class BaseAWSLLM: if provider is not None: return provider - # check if provider == "nova" - if "nova" in model: - return "nova" - else: - for provider in get_args(BEDROCK_INVOKE_PROVIDERS_LITERAL): - if provider in model: - return provider + for provider in get_args(BEDROCK_INVOKE_PROVIDERS_LITERAL): + if provider in model: + return provider return None @staticmethod @@ -364,11 +401,15 @@ class BaseAWSLLM: elif provider == "qwen3" and "qwen3/" in model_id: model_id = BaseAWSLLM._get_model_id_from_model_with_spec( model_id, spec="qwen3" - ) + ) elif provider == "stability" and "stability/" in model_id: model_id = BaseAWSLLM._get_model_id_from_model_with_spec( model_id, spec="stability" ) + elif provider == "moonshot" and "moonshot/" in model_id: + model_id = BaseAWSLLM._get_model_id_from_model_with_spec( + model_id, spec="moonshot" + ) return model_id @staticmethod @@ -412,7 +453,7 @@ class BaseAWSLLM: if "nova" in model.lower(): if "nova" in get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL): return cast(BEDROCK_EMBEDDING_PROVIDERS_LITERAL, "nova") - + # Handle regional models like us.twelvelabs.marengo-embed-2-7-v1:0 if "." in model: parts = model.split(".") @@ -563,6 +604,7 @@ class BaseAWSLLM: "sts", region_name=aws_region_name, endpoint_url=sts_endpoint, + verify=self._get_ssl_verify(), ) # https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html @@ -619,7 +661,7 @@ class BaseAWSLLM: # Create an STS client without credentials with tracer.trace("boto3.client(sts) for manual IRSA"): - sts_client = boto3.client("sts", region_name=region) + sts_client = boto3.client("sts", region_name=region, verify=self._get_ssl_verify()) # Manually assume the IRSA role with the session name verbose_logger.debug( @@ -642,6 +684,7 @@ class BaseAWSLLM: aws_access_key_id=irsa_creds["AccessKeyId"], aws_secret_access_key=irsa_creds["SecretAccessKey"], aws_session_token=irsa_creds["SessionToken"], + verify=self._get_ssl_verify(), ) # Get current caller identity for debugging @@ -680,7 +723,7 @@ class BaseAWSLLM: verbose_logger.debug("Same account role assumption, using automatic IRSA") with tracer.trace("boto3.client(sts) with automatic IRSA"): - sts_client = boto3.client("sts", region_name=region) + sts_client = boto3.client("sts", region_name=region, verify=self._get_ssl_verify()) # Get current caller identity for debugging try: @@ -803,7 +846,7 @@ class BaseAWSLLM: # This allows the web identity token to work automatically if aws_access_key_id is None and aws_secret_access_key is None: with tracer.trace("boto3.client(sts)"): - sts_client = boto3.client("sts") + sts_client = boto3.client("sts", verify=self._get_ssl_verify()) else: with tracer.trace("boto3.client(sts)"): sts_client = boto3.client( @@ -811,6 +854,7 @@ class BaseAWSLLM: aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, aws_session_token=aws_session_token, + verify=self._get_ssl_verify(), ) assume_role_params = { @@ -958,7 +1002,9 @@ class BaseAWSLLM: return endpoint_url, proxy_endpoint_url def _select_default_endpoint_url( - self, endpoint_type: Optional[Literal["runtime", "agent", "agentcore"]], aws_region_name: str + self, + endpoint_type: Optional[Literal["runtime", "agent", "agentcore"]], + aws_region_name: str, ) -> str: """ Select the default endpoint url based on the endpoint type diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_moonshot_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_moonshot_transformation.py new file mode 100644 index 00000000000..e53410760dd --- /dev/null +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_moonshot_transformation.py @@ -0,0 +1,256 @@ +""" +Transformation for Bedrock Moonshot AI (Kimi K2) models. + +Supports the Kimi K2 Thinking model available on Amazon Bedrock. +Model format: bedrock/moonshot.kimi-k2-thinking-v1:0 + +Reference: https://aws.amazon.com/about-aws/whats-new/2025/12/amazon-bedrock-fully-managed-open-weight-models/ +""" + +from typing import TYPE_CHECKING, Any, List, Optional, Union +import re + +import httpx + +from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( + AmazonInvokeConfig, +) +from litellm.llms.bedrock.common_utils import BedrockError +from litellm.llms.moonshot.chat.transformation import MoonshotChatConfig +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import Choices + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + from litellm.types.utils import ModelResponse + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig): + """ + Configuration for Bedrock Moonshot AI (Kimi K2) models. + + Reference: + https://aws.amazon.com/about-aws/whats-new/2025/12/amazon-bedrock-fully-managed-open-weight-models/ + https://platform.moonshot.ai/docs/api/chat + + Supported Params for the Amazon / Moonshot models: + - `max_tokens` (integer) max tokens + - `temperature` (float) temperature for model (0-1 for Moonshot) + - `top_p` (float) top p for model + - `stream` (bool) whether to stream responses + - `tools` (list) tool definitions (supported on kimi-k2-thinking) + - `tool_choice` (str|dict) tool choice specification (supported on kimi-k2-thinking) + + NOT Supported on Bedrock: + - `stop` sequences (Bedrock doesn't support stopSequences field for this model) + + Note: The kimi-k2-thinking model DOES support tool calls, unlike kimi-thinking-preview. + """ + + def __init__(self, **kwargs): + AmazonInvokeConfig.__init__(self, **kwargs) + MoonshotChatConfig.__init__(self, **kwargs) + + @property + def custom_llm_provider(self) -> Optional[str]: + return "bedrock" + + def _get_model_id(self, model: str) -> str: + """ + Extract the actual model ID from the LiteLLM model name. + + Removes routing prefixes like: + - bedrock/invoke/moonshot.kimi-k2-thinking -> moonshot.kimi-k2-thinking + - invoke/moonshot.kimi-k2-thinking -> moonshot.kimi-k2-thinking + - moonshot.kimi-k2-thinking -> moonshot.kimi-k2-thinking + """ + # Remove bedrock/ prefix if present + if model.startswith("bedrock/"): + model = model[8:] + + # Remove invoke/ prefix if present + if model.startswith("invoke/"): + model = model[7:] + + # Remove any provider prefix (e.g., moonshot/) + if "/" in model and not model.startswith("arn:"): + parts = model.split("/", 1) + if len(parts) == 2: + model = parts[1] + + return model + + def get_supported_openai_params(self, model: str) -> List[str]: + """ + Get the supported OpenAI params for Moonshot AI models on Bedrock. + + Bedrock-specific limitations: + - stopSequences field is not supported on Bedrock (unlike native Moonshot API) + - functions parameter is not supported (use tools instead) + - tool_choice doesn't support "required" value + + Note: kimi-k2-thinking DOES support tool calls (unlike kimi-thinking-preview) + The parent MoonshotChatConfig class handles the kimi-thinking-preview exclusion. + """ + excluded_params: List[str] = ["functions", "stop"] # Bedrock doesn't support stopSequences + + base_openai_params = super(MoonshotChatConfig, self).get_supported_openai_params(model=model) + final_params: List[str] = [] + for param in base_openai_params: + if param not in excluded_params: + final_params.append(param) + + return final_params + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to Moonshot AI parameters for Bedrock. + + Handles Moonshot AI specific limitations: + - tool_choice doesn't support "required" value + - Temperature <0.3 limitation for n>1 + - Temperature range is [0, 1] (not [0, 2] like OpenAI) + """ + return MoonshotChatConfig.map_openai_params( + self, + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=drop_params, + ) + + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Transform the request for Bedrock Moonshot AI models. + + Uses the Moonshot transformation logic which handles: + - Converting content lists to strings (Moonshot doesn't support list format) + - Adding tool_choice="required" message if needed + - Temperature and parameter validation + + """ + # Filter out AWS credentials using the existing method from BaseAWSLLM + self._get_boto_credentials_from_optional_params(optional_params, model) + + # Strip routing prefixes to get the actual model ID + clean_model_id = self._get_model_id(model) + + # Use Moonshot's transform_request which handles message transformation + # and tool_choice="required" workaround + return MoonshotChatConfig.transform_request( + self, + model=clean_model_id, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + ) + + def _extract_reasoning_from_content(self, content: str) -> tuple[Optional[str], str]: + """ + Extract reasoning content from tags in the response. + + Moonshot AI's Kimi K2 Thinking model returns reasoning in tags. + This method extracts that content and returns it separately. + + Args: + content: The full content string from the API response + + Returns: + tuple: (reasoning_content, main_content) + """ + if not content: + return None, content + + # Match ... tags + reasoning_match = re.match( + r"(.*?)\s*(.*)", + content, + re.DOTALL + ) + + if reasoning_match: + reasoning_content = reasoning_match.group(1).strip() + main_content = reasoning_match.group(2).strip() + return reasoning_content, main_content + + return None, content + + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: "ModelResponse", + logging_obj: LiteLLMLoggingObj, + request_data: dict, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + encoding: Any, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, + ) -> "ModelResponse": + """ + Transform the response from Bedrock Moonshot AI models. + + Moonshot AI uses OpenAI-compatible response format, but returns reasoning + content in tags. This method: + 1. Calls parent class transformation + 2. Extracts reasoning content from tags + 3. Sets reasoning_content on the message object + """ + # First, get the standard transformation + model_response = MoonshotChatConfig.transform_response( + self, + model=model, + raw_response=raw_response, + model_response=model_response, + logging_obj=logging_obj, + request_data=request_data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + api_key=api_key, + json_mode=json_mode, + ) + + # Extract reasoning content from tags + if model_response.choices and len(model_response.choices) > 0: + for choice in model_response.choices: + # Only process Choices (not StreamingChoices) which have message attribute + if isinstance(choice, Choices) and choice.message and choice.message.content: + reasoning_content, main_content = self._extract_reasoning_from_content( + choice.message.content + ) + + if reasoning_content: + # Set the reasoning_content field + choice.message.reasoning_content = reasoning_content + # Update the main content without reasoning tags + choice.message.content = main_content + + return model_response + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BedrockError: + """Return the appropriate error class for Bedrock.""" + return BedrockError(status_code=status_code, message=error_message) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index c602b71fe05..cf8aee6954b 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -524,6 +524,12 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): if model.startswith("invoke/"): model = model.replace("invoke/", "", 1) + # Special case: Check for "nova" in model name first (before "amazon") + # This handles amazon.nova-* models which would otherwise match "amazon" (Titan) + if "nova" in model.lower(): + if "nova" in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL): + return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, "nova") + _split_model = model.split(".")[0] if _split_model in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL): return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, _split_model) @@ -533,10 +539,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): if provider is not None: return provider - # check if provider == "nova" - if "nova" in model: - return "nova" - for provider in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL): if provider in model: return provider diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 21a78c30343..d8e0a05dc01 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -15,7 +15,7 @@ import litellm from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, ) -from litellm.llms.base_llm.base_utils import BaseLLMModelInfo +from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.secret_managers.main import get_secret @@ -178,7 +178,26 @@ def init_bedrock_client( ) = params_to_check # SSL certificates (a.k.a CA bundle) used to verify the identity of requested hosts. + # Use the same logic as BaseAWSLLM._get_ssl_verify() for consistency + from litellm.secret_managers.main import str_to_bool ssl_verify = os.getenv("SSL_VERIFY", litellm.ssl_verify) + + # Convert string "False"/"True" to boolean + if isinstance(ssl_verify, str): + # Check if it's a file path + if os.path.exists(ssl_verify): + pass # Keep the file path + else: + # Otherwise try to convert to boolean + ssl_verify_bool = str_to_bool(ssl_verify) + if ssl_verify_bool is not None: + ssl_verify = ssl_verify_bool + + # Check SSL_CERT_FILE environment variable for custom CA bundle + if ssl_verify is True or ssl_verify == "True": + ssl_cert_file = os.getenv("SSL_CERT_FILE") + if ssl_cert_file and os.path.exists(ssl_cert_file): + ssl_verify = ssl_cert_file ### SET REGION NAME if region_name: @@ -229,7 +248,7 @@ def init_bedrock_client( status_code=401, ) - sts_client = boto3.client("sts") + sts_client = boto3.client("sts", verify=ssl_verify) # https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html # https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html @@ -359,6 +378,70 @@ def get_bedrock_tool_name(response_tool_name: str) -> str: return response_tool_name +# Cache the global regions list at module level +_BEDROCK_GLOBAL_REGIONS: Optional[List[str]] = None + + +def _get_all_bedrock_regions() -> List[str]: + """Get all Bedrock regions, cached at module level.""" + global _BEDROCK_GLOBAL_REGIONS + if _BEDROCK_GLOBAL_REGIONS is None: + _BEDROCK_GLOBAL_REGIONS = AmazonBedrockGlobalConfig().get_all_regions() + return _BEDROCK_GLOBAL_REGIONS + + +def get_bedrock_cross_region_inference_regions() -> List[str]: + """Abbreviations of regions AWS Bedrock supports for cross region inference.""" + return ["global", "us", "eu", "apac", "jp", "au", "us-gov"] + + +def extract_model_name_from_bedrock_arn(model: str) -> str: + """ + Extract the model name from an AWS Bedrock ARN. + Returns the string after the last '/' if 'arn' is in the input string. + """ + if "arn" in model.lower(): + return model.split("/")[-1] + return model + + +def strip_bedrock_routing_prefix(model: str) -> str: + """Strip LiteLLM routing prefixes from model name.""" + for prefix in ["bedrock/", "converse/", "invoke/", "openai/"]: + if model.startswith(prefix): + model = model.split("/", 1)[1] + return model + + +def get_bedrock_base_model(model: str) -> str: + """ + Get the base model from the given model name. + + Handle model names like: + - "us.meta.llama3-2-11b-instruct-v1:0" -> "meta.llama3-2-11b-instruct-v1" + - "bedrock/converse/model" -> "model" + """ + model = strip_bedrock_routing_prefix(model) + model = extract_model_name_from_bedrock_arn(model) + + potential_region = model.split(".", 1)[0] + alt_potential_region = model.split("/", 1)[0] + + if potential_region in get_bedrock_cross_region_inference_regions(): + return model.split(".", 1)[1] + elif ( + alt_potential_region in _get_all_bedrock_regions() + and len(model.split("/", 1)) > 1 + ): + return model.split("/", 1)[1] + + return model + + +# Import after standalone functions to avoid circular imports +from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter + + class BedrockModelInfo(BaseLLMModelInfo): global_config = AmazonBedrockGlobalConfig() all_global_regions = global_config.get_all_regions() @@ -394,76 +477,34 @@ class BedrockModelInfo(BaseLLMModelInfo): ) -> List[str]: return [] - @staticmethod - def extract_model_name_from_arn(model: str) -> str: + def get_token_counter(self) -> Optional[BaseTokenCounter]: """ - Extract the model name from an AWS Bedrock ARN. - Returns the string after the last '/' if 'arn' is in the input string. - - Args: - arn (str): The ARN string to parse + Factory method to create a Bedrock token counter. Returns: - str: The extracted model name if 'arn' is in the string, - otherwise returns the original string + BedrockTokenCounter instance for this provider. """ - if "arn" in model.lower(): - return model.split("/")[-1] - return model + return BedrockTokenCounter() + + @staticmethod + def extract_model_name_from_arn(model: str) -> str: + """Wrapper for standalone function. See extract_model_name_from_bedrock_arn().""" + return extract_model_name_from_bedrock_arn(model) @staticmethod def get_non_litellm_routing_model_name(model: str) -> str: - if model.startswith("bedrock/"): - model = model.split("/", 1)[1] - - if model.startswith("converse/"): - model = model.split("/", 1)[1] - - if model.startswith("invoke/"): - model = model.split("/", 1)[1] - - if model.startswith("openai/"): - model = model.split("/", 1)[1] - - return model + """Wrapper for standalone function. See strip_bedrock_routing_prefix().""" + return strip_bedrock_routing_prefix(model) @staticmethod def get_base_model(model: str) -> str: - """ - Get the base model from the given model name. - - Handle model names like - "us.meta.llama3-2-11b-instruct-v1:0" -> "meta.llama3-2-11b-instruct-v1" - AND "meta.llama3-2-11b-instruct-v1:0" -> "meta.llama3-2-11b-instruct-v1" - """ - - model = BedrockModelInfo.get_non_litellm_routing_model_name(model=model) - model = BedrockModelInfo.extract_model_name_from_arn(model) - - potential_region = model.split(".", 1)[0] - - alt_potential_region = model.split("/", 1)[ - 0 - ] # in model cost map we store regional information like `/us-west-2/bedrock-model` - - if ( - potential_region - in BedrockModelInfo._supported_cross_region_inference_region() - ): - return model.split(".", 1)[1] - elif ( - alt_potential_region in BedrockModelInfo.all_global_regions - and len(model.split("/", 1)) > 1 - ): - return model.split("/", 1)[1] - - return model + """Wrapper for standalone function. See get_bedrock_base_model().""" + return get_bedrock_base_model(model) @staticmethod def _supported_cross_region_inference_region() -> List[str]: - """ - Abbreviations of regions AWS Bedrock supports for cross region inference - """ - return ["global", "us", "eu", "apac", "jp", "au", "us-gov"] + """Wrapper for standalone function. See get_bedrock_cross_region_inference_regions().""" + return get_bedrock_cross_region_inference_regions() @staticmethod def get_bedrock_route( @@ -629,6 +670,8 @@ def get_bedrock_chat_config(model: str): return litellm.AmazonCohereConfig() elif bedrock_invoke_provider == "mistral": return litellm.AmazonMistralConfig() + elif bedrock_invoke_provider == "moonshot": + return litellm.AmazonMoonshotConfig() elif bedrock_invoke_provider == "deepseek_r1": return litellm.AmazonDeepSeekR1Config() elif bedrock_invoke_provider == "nova": diff --git a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py new file mode 100644 index 00000000000..b680bd046ef --- /dev/null +++ b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py @@ -0,0 +1,87 @@ +""" +Bedrock Token Counter implementation using the CountTokens API. +""" + +from typing import Any, Dict, List, Optional + +from litellm._logging import verbose_logger +from litellm.llms.base_llm.base_utils import BaseTokenCounter +from litellm.llms.bedrock.common_utils import get_bedrock_base_model +from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler +from litellm.types.utils import LlmProviders, TokenCountResponse + + +class BedrockTokenCounter(BaseTokenCounter): + """Token counter implementation for AWS Bedrock provider using the CountTokens API.""" + + def should_use_token_counting_api( + self, + custom_llm_provider: Optional[str] = None, + ) -> bool: + """ + Returns True if we should use the Bedrock CountTokens API for token counting. + """ + return custom_llm_provider == LlmProviders.BEDROCK.value + + async def count_tokens( + self, + model_to_use: str, + messages: Optional[List[Dict[str, Any]]], + contents: Optional[List[Dict[str, Any]]], + deployment: Optional[Dict[str, Any]] = None, + request_model: str = "", + ) -> Optional[TokenCountResponse]: + """ + Count tokens using AWS Bedrock's CountTokens API. + + This method calls the existing BedrockCountTokensHandler to make an API call + to Bedrock's token counting endpoint, bypassing the local tiktoken-based counting. + + Args: + model_to_use: The model identifier + messages: The messages to count tokens for + contents: Alternative content format (not used for Bedrock) + deployment: Deployment configuration containing litellm_params + request_model: The original request model name + + Returns: + TokenCountResponse with token count, or None if counting fails + """ + if not messages: + return None + + deployment = deployment or {} + litellm_params = deployment.get("litellm_params", {}) + + # Build request data in the format expected by BedrockCountTokensHandler + request_data = { + "model": model_to_use, + "messages": messages, + } + + # Get the resolved model (strip prefixes like bedrock/, converse/, etc.) + resolved_model = get_bedrock_base_model(model_to_use) + + try: + handler = BedrockCountTokensHandler() + result = await handler.handle_count_tokens_request( + request_data=request_data, + litellm_params=litellm_params, + resolved_model=resolved_model, + ) + + # Transform response to TokenCountResponse + if result is not None: + return TokenCountResponse( + total_tokens=result.get("input_tokens", 0), + request_model=request_model, + model_used=model_to_use, + tokenizer_type="bedrock_api", + original_response=result, + ) + except Exception as e: + verbose_logger.warning( + f"Error calling Bedrock CountTokens API: {e}, falling back to default tokenizer" + ) + + return None diff --git a/litellm/llms/bedrock/count_tokens/handler.py b/litellm/llms/bedrock/count_tokens/handler.py index d4355c0c360..e8366165b65 100644 --- a/litellm/llms/bedrock/count_tokens/handler.py +++ b/litellm/llms/bedrock/count_tokens/handler.py @@ -6,10 +6,9 @@ Simplified handler leveraging existing LiteLLM Bedrock infrastructure. from typing import Any, Dict -from fastapi import HTTPException - import litellm from litellm._logging import verbose_logger +from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.bedrock.count_tokens.transformation import BedrockCountTokensConfig from litellm.llms.custom_httpx.http_handler import get_async_httpx_client @@ -70,6 +69,8 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): verbose_logger.debug(f"Making request to: {endpoint_url}") # Use existing _sign_request method from BaseAWSLLM + # Extract api_key for bearer token auth if provided + api_key = litellm_params.get("api_key", None) headers = {"Content-Type": "application/json"} signed_headers, signed_body = self._sign_request( service_name="bedrock", @@ -78,6 +79,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): request_data=bedrock_request, api_base=endpoint_url, model=resolved_model, + api_key=api_key, ) async_client = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK) @@ -94,9 +96,9 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): if response.status_code != 200: error_text = response.text verbose_logger.error(f"AWS Bedrock error: {error_text}") - raise HTTPException( - status_code=400, - detail={"error": f"AWS Bedrock error: {error_text}"}, + raise BedrockError( + status_code=response.status_code, + message=f"AWS Bedrock error: {error_text}", ) bedrock_response = response.json() @@ -112,12 +114,12 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): return final_response - except HTTPException: - # Re-raise HTTP exceptions as-is + except BedrockError: + # Re-raise Bedrock exceptions as-is raise except Exception as e: verbose_logger.error(f"Error in CountTokens handler: {str(e)}") - raise HTTPException( + raise BedrockError( status_code=500, - detail={"error": f"CountTokens processing error: {str(e)}"}, + message=f"CountTokens processing error: {str(e)}", ) diff --git a/litellm/llms/bedrock/count_tokens/transformation.py b/litellm/llms/bedrock/count_tokens/transformation.py index d46ed3aa452..b313cc9df3c 100644 --- a/litellm/llms/bedrock/count_tokens/transformation.py +++ b/litellm/llms/bedrock/count_tokens/transformation.py @@ -8,7 +8,7 @@ to AWS Bedrock's CountTokens API format and vice versa. from typing import Any, Dict, List from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM -from litellm.llms.bedrock.common_utils import BedrockModelInfo +from litellm.llms.bedrock.common_utils import get_bedrock_base_model class BedrockCountTokensConfig(BaseAWSLLM): @@ -141,7 +141,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): Complete endpoint URL for CountTokens API """ # Use existing LiteLLM function to get the base model ID (removes region prefix) - model_id = BedrockModelInfo.get_base_model(model) + model_id = get_bedrock_base_model(model) # Remove bedrock/ prefix if present if model_id.startswith("bedrock/"): diff --git a/litellm/llms/bedrock/files/handler.py b/litellm/llms/bedrock/files/handler.py index d6177e090d5..0350271dc44 100644 --- a/litellm/llms/bedrock/files/handler.py +++ b/litellm/llms/bedrock/files/handler.py @@ -142,6 +142,7 @@ class BedrockFilesHandler(BaseAWSLLM): aws_secret_access_key=credentials.secret_key, aws_session_token=credentials.token, region_name=aws_region_name, + verify=self._get_ssl_verify(), ) # Download file from S3 diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index 5791bfb8013..568fe941716 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -34,11 +34,12 @@ class BedrockPassthroughConfig( litellm_params: dict, ) -> Tuple["URL", str]: optional_params = litellm_params.copy() + model_id = optional_params.get("model_id", None) aws_region_name = self._get_aws_region_name( optional_params=optional_params, model=model, - model_id=None, + model_id=model_id, ) aws_bedrock_runtime_endpoint = optional_params.get("aws_bedrock_runtime_endpoint") @@ -49,6 +50,12 @@ class BedrockPassthroughConfig( endpoint_type="runtime", ) + # If model_id is provided (e.g., Application Inference Profile ARN), use it in the endpoint + # instead of the translated model name + if model_id is not None: + # Replace the model name in the endpoint with the model_id + import re + endpoint = re.sub(r'model/[^/]+/', f'model/{model_id}/', endpoint) return self.format_url(endpoint, endpoint_url, request_query_params or {}), endpoint_url def sign_request( diff --git a/litellm/llms/deepinfra/chat/transformation.py b/litellm/llms/deepinfra/chat/transformation.py index 09cdabcdd82..5198260a24b 100644 --- a/litellm/llms/deepinfra/chat/transformation.py +++ b/litellm/llms/deepinfra/chat/transformation.py @@ -1,9 +1,11 @@ -from typing import Optional, Tuple, Union +import json +from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, cast, overload import litellm from litellm.constants import MIN_NON_ZERO_TEMPERATURE from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllMessageValues class DeepInfraConfig(OpenAIGPTConfig): @@ -117,6 +119,79 @@ class DeepInfraConfig(OpenAIGPTConfig): optional_params[param] = value return optional_params + def _transform_tool_message_content(self, messages: List[AllMessageValues]) -> List[AllMessageValues]: + """ + Transform tool message content from array to string format for DeepInfra compatibility. + + DeepInfra requires tool message content to be a string, not an array. + This method converts tool message content from array format to string format. + + Example transformation: + - Input: {"role": "tool", "content": [{"type": "text", "text": "20"}]} + - Output: {"role": "tool", "content": "20"} + + Or if content is complex: + - Input: {"role": "tool", "content": [{"type": "text", "text": "result"}]} + - Output: {"role": "tool", "content": "[{\"type\": \"text\", \"text\": \"result\"}]"} + """ + for message in messages: + if message.get("role") == "tool": + content = message.get("content") + + # If content is a list/array, convert it to string + if isinstance(content, list): + # Check if it's a simple single text item + if ( + len(content) == 1 + and isinstance(content[0], dict) + and content[0].get("type") == "text" + and "text" in content[0] + ): + # Extract just the text value for simple cases + message["content"] = content[0]["text"] + else: + # For complex content, serialize the entire array as JSON string + message["content"] = json.dumps(content) + + return messages + + @overload + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: Literal[False] = False + ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + """ + Transform messages for DeepInfra compatibility. + Handles both sync and async transformations. + """ + if is_async: + # For async case, create an async function that awaits parent and applies our transformation + async def _async_transform(): + # Call parent with is_async=True (literal) for async case + parent_result = super(DeepInfraConfig, self)._transform_messages( + messages=messages, model=model, is_async=cast(Literal[True], True) + ) + transformed_messages = await parent_result + return self._transform_tool_message_content(transformed_messages) + return _async_transform() + else: + # Call parent with is_async=False (literal) for sync case + parent_result = super()._transform_messages( + messages=messages, model=model, is_async=cast(Literal[False], False) + ) + # For sync case, parent_result is already the transformed messages + return self._transform_tool_message_content(parent_result) + def _get_openai_compatible_provider_info( self, api_base: Optional[str], api_key: Optional[str] ) -> Tuple[Optional[str], Optional[str]]: diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index e53829d3329..30c5b4f17c5 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -150,6 +150,15 @@ def get_api_key_from_env() -> Optional[str]: return get_secret_str("GOOGLE_API_KEY") or get_secret_str("GEMINI_API_KEY") +def get_vertex_api_key_from_env() -> Optional[str]: + """ + Get API key from environment for Vertex AI. + Checks VERTEXAI_API_KEY and VERTEX_API_KEY environment variables. + This allows using Vertex AI with API keys instead of service account credentials. + """ + return get_secret_str("VERTEXAI_API_KEY") or get_secret_str("VERTEX_API_KEY") + + class GoogleAIStudioTokenCounter(BaseTokenCounter): """Token counter implementation for Google AI Studio provider.""" def should_use_token_counting_api( diff --git a/litellm/llms/manus/__init__.py b/litellm/llms/manus/__init__.py new file mode 100644 index 00000000000..81eef025461 --- /dev/null +++ b/litellm/llms/manus/__init__.py @@ -0,0 +1,2 @@ +# Manus provider implementation + diff --git a/litellm/llms/manus/responses/__init__.py b/litellm/llms/manus/responses/__init__.py new file mode 100644 index 00000000000..e8cabc54266 --- /dev/null +++ b/litellm/llms/manus/responses/__init__.py @@ -0,0 +1,2 @@ +# Manus Responses API implementation + diff --git a/litellm/llms/manus/responses/transformation.py b/litellm/llms/manus/responses/transformation.py new file mode 100644 index 00000000000..7a72f23dd56 --- /dev/null +++ b/litellm/llms/manus/responses/transformation.py @@ -0,0 +1,308 @@ +from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union + +import httpx + +import litellm +from litellm._logging import verbose_logger +from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( + _safe_convert_created_field, +) +from litellm.llms.openai.common_utils import OpenAIError +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + ResponseAPIUsage, + ResponseInputParam, + ResponsesAPIResponse, +) +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + +MANUS_API_BASE = "https://api.manus.im" + + +class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): + """ + Configuration for Manus API's Responses API. + + Manus API is OpenAI-compatible but has some differences: + - API key passed via `API_KEY` header (not `Authorization: Bearer`) + - Model format: `manus/{agent_profile}` (e.g., `manus/manus-1.6`) + - Requires `extra_body` with `task_mode: "agent"` and `agent_profile` + + Reference: https://open.manus.im/docs/openai-compatibility + """ + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.MANUS + + def should_fake_stream( + self, + model: Optional[str], + stream: Optional[bool], + custom_llm_provider: Optional[str] = None, + ) -> bool: + """ + Manus API doesn't support real-time streaming. + It returns a task that runs asynchronously. + We fake streaming by converting the response into streaming events. + """ + return stream is True + + def _extract_agent_profile(self, model: str) -> str: + """ + Extract agent profile from model name. + + Model format: `manus/{agent_profile}` + Examples: `manus/manus-1.6`, `manus/manus-1.6-lite`, `manus/manus-1.6-max` + + Returns: + str: The agent profile (e.g., "manus-1.6") + """ + if "/" in model: + return model.split("/", 1)[1] + # If no slash, assume the model name itself is the agent profile + return model + + def validate_environment( + self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] + ) -> dict: + """ + Validate environment and set up headers for Manus API. + + Manus uses `API_KEY` header instead of `Authorization: Bearer`. + """ + litellm_params = litellm_params or GenericLiteLLMParams() + api_key = ( + litellm_params.api_key + or litellm.api_key + or get_secret_str("MANUS_API_KEY") + ) + + if not api_key: + raise ValueError( + "Manus API key is required. Set MANUS_API_KEY environment variable or pass api_key parameter." + ) + + # Manus uses API_KEY header, not Authorization: Bearer + headers.update( + { + "API_KEY": api_key, + } + ) + return headers + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Get the complete URL for Manus Responses API endpoint. + + Returns: + str: The full URL for the Manus /v1/responses endpoint + """ + api_base = ( + api_base + or litellm.api_base + or get_secret_str("MANUS_API_BASE") + or MANUS_API_BASE + ) + + # Remove trailing slashes + api_base = api_base.rstrip("/") + + # Manus API uses /v1/responses endpoint (OpenAI-compatible) + if api_base.endswith("/v1"): + return f"{api_base}/responses" + return f"{api_base}/v1/responses" + + def transform_responses_api_request( + self, + model: str, + input: Union[str, ResponseInputParam], + response_api_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + """ + Transform the request for Manus API. + + Manus requires: + - `task_mode: "agent"` in the request body + - `agent_profile` extracted from model name in the request body + """ + # First, get the base OpenAI request + base_request = super().transform_responses_api_request( + model=model, + input=input, + response_api_optional_request_params=response_api_optional_request_params, + litellm_params=litellm_params, + headers=headers, + ) + + # Extract agent profile from model name + agent_profile = self._extract_agent_profile(model=model) + + # Add Manus-specific parameters directly to the request body + # These will be sent as part of the request + base_request["task_mode"] = "agent" + base_request["agent_profile"] = agent_profile + + # Merge any existing extra_body into the request + extra_body = response_api_optional_request_params.get("extra_body", {}) or {} + if extra_body: + base_request.update(extra_body) + + # Avoid logging potentially sensitive agent_profile value + verbose_logger.debug("Manus: Using task_mode=agent") + + return base_request + + def transform_response_api_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + """ + Transform Manus API response to OpenAI-compatible format. + + Manus uses camelCase (createdAt) instead of snake_case (created_at). + """ + try: + logging_obj.post_call( + original_response=raw_response.text, + additional_args={"complete_input_dict": {}}, + ) + raw_response_json = raw_response.json() + + # Manus uses camelCase "createdAt" instead of snake_case "created_at" + if "createdAt" in raw_response_json and "created_at" not in raw_response_json: + raw_response_json["created_at"] = _safe_convert_created_field( + raw_response_json["createdAt"] + ) + + # Ensure created_at is set + if "created_at" in raw_response_json: + raw_response_json["created_at"] = _safe_convert_created_field( + raw_response_json["created_at"] + ) + except Exception: + raise OpenAIError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + # Ensure reasoning is an empty dict if not present, OpenAI SDK does not allow None + if "reasoning" not in raw_response_json or raw_response_json.get("reasoning") is None: + raw_response_json["reasoning"] = {} + + if "text" not in raw_response_json or raw_response_json.get("text") is None: + raw_response_json["text"] = {} + + if "output" not in raw_response_json or raw_response_json.get("output") is None: + raw_response_json["output"] = [] + + # Ensure usage is present with default values if not provided + if "usage" not in raw_response_json or raw_response_json.get("usage") is None: + raw_response_json["usage"] = ResponseAPIUsage( + input_tokens=0, + output_tokens=0, + total_tokens=0, + ) + + try: + response = ResponsesAPIResponse(**raw_response_json) + except Exception: + verbose_logger.debug( + f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct" + ) + response = ResponsesAPIResponse.model_construct(**raw_response_json) + + # Store processed headers in additional_headers so they get returned to the client + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + def transform_get_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + """ + Transform the get response API request into a URL and data. + + Manus API follows OpenAI-compatible format: + - GET /v1/responses/{response_id} + + Reference: https://open.manus.im/docs/openai-compatibility + """ + url = f"{api_base}/{response_id}" + data: Dict = {} + return url, data + + def transform_get_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + """ + Transform Manus API GET response to OpenAI-compatible format. + + Manus uses camelCase (createdAt) instead of snake_case (created_at). + Same transformation as transform_response_api_response. + """ + try: + logging_obj.post_call( + original_response=raw_response.text, + additional_args={"complete_input_dict": {}}, + ) + raw_response_json = raw_response.json() + + # Manus uses camelCase "createdAt" instead of snake_case "created_at" + if "createdAt" in raw_response_json and "created_at" not in raw_response_json: + raw_response_json["created_at"] = _safe_convert_created_field( + raw_response_json["createdAt"] + ) + + # Ensure created_at is set + if "created_at" in raw_response_json: + raw_response_json["created_at"] = _safe_convert_created_field( + raw_response_json["created_at"] + ) + except Exception: + raise OpenAIError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + try: + response = ResponsesAPIResponse(**raw_response_json) + except Exception: + verbose_logger.debug( + f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct" + ) + response = ResponsesAPIResponse.model_construct(**raw_response_json) + + # Store processed headers in additional_headers so they get returned to the client + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 034ccae94ad..04a10bd7fbe 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -771,9 +771,9 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator): return ModelResponseStream( id=chunk["id"], object="chat.completion.chunk", - created=chunk["created"], - model=chunk["model"], - choices=chunk["choices"], + created=chunk.get("created"), + model=chunk.get("model"), + choices=chunk.get("choices", []), ) except Exception as e: raise e diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index 206aee1359d..bda3684a8a8 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -61,6 +61,10 @@ "max_completion_tokens": "max_tokens" } }, + "abliteration": { + "base_url": "https://api.abliteration.ai/v1", + "api_key_env": "ABLITERATION_API_KEY" + }, "llamagate": { "base_url": "https://api.llamagate.dev/v1", "api_key_env": "LLAMAGATE_API_KEY", diff --git a/litellm/llms/openrouter/embedding/transformation.py b/litellm/llms/openrouter/embedding/transformation.py new file mode 100644 index 00000000000..d1d0e911d16 --- /dev/null +++ b/litellm/llms/openrouter/embedding/transformation.py @@ -0,0 +1,182 @@ +""" +OpenRouter Embedding API Configuration. + +This module provides the configuration for OpenRouter's Embedding API. +OpenRouter is OpenAI-compatible and supports embeddings via the /v1/embeddings endpoint. + +Docs: https://openrouter.ai/docs +""" +from typing import TYPE_CHECKING, Any, Optional + +import httpx + +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig +from litellm.types.llms.openai import AllEmbeddingInputValues +from litellm.types.utils import EmbeddingResponse +from litellm.utils import convert_to_model_response_object + +from ..common_utils import OpenRouterException + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class OpenrouterEmbeddingConfig(BaseEmbeddingConfig): + """ + Configuration for OpenRouter's Embedding API. + + Reference: https://openrouter.ai/docs + """ + + def validate_environment( + self, + headers: dict, + model: str, + messages: list, + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + """ + Validate environment and set up headers for OpenRouter API. + + OpenRouter requires: + - Authorization header with Bearer token + - HTTP-Referer header (site URL) + - X-Title header (app name) + """ + from litellm import get_secret + + # Get OpenRouter-specific headers + openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai" + openrouter_app_name = get_secret("OR_APP_NAME") or "liteLLM" + + openrouter_headers = { + "HTTP-Referer": openrouter_site_url, + "X-Title": openrouter_app_name, + "Content-Type": "application/json", + } + + # Add Authorization header if api_key is provided + if api_key: + openrouter_headers["Authorization"] = f"Bearer {api_key}" + + # Merge with existing headers (user's extra_headers take priority) + merged_headers = {**openrouter_headers, **headers} + + return merged_headers + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + """ + Get the complete URL for OpenRouter Embedding API endpoint. + """ + # api_base is already set to https://openrouter.ai/api/v1 in main.py + # Remove trailing slashes + if api_base: + api_base = api_base.rstrip("/") + else: + api_base = "https://openrouter.ai/api/v1" + + # Return the embeddings endpoint + return f"{api_base}/embeddings" + + def transform_embedding_request( + self, + model: str, + input: AllEmbeddingInputValues, + optional_params: dict, + headers: dict, + ) -> dict: + """ + Transform embedding request to OpenRouter format (OpenAI-compatible). + """ + # Ensure input is a list + if isinstance(input, str): + input = [input] + + # OpenRouter expects the full model name (e.g., google/gemini-embedding-001) + # Strip 'openrouter/' prefix if present + if model.startswith("openrouter/"): + model = model.replace("openrouter/", "", 1) + + return { + "model": model, + "input": input, + **optional_params, + } + + def transform_embedding_response( + self, + model: str, + raw_response: httpx.Response, + model_response: EmbeddingResponse, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + request_data: dict, + optional_params: dict, + litellm_params: dict, + ) -> EmbeddingResponse: + """ + Transform embedding response from OpenRouter format (OpenAI-compatible). + """ + logging_obj.post_call(original_response=raw_response.text) + + # OpenRouter returns standard OpenAI-compatible embedding response + response_json = raw_response.json() + + return convert_to_model_response_object( + response_object=response_json, + model_response_object=model_response, + response_type="embedding", + ) + + def get_supported_openai_params(self, model: str) -> list: + """ + Get list of supported OpenAI parameters for OpenRouter embeddings. + """ + return [ + "timeout", + "dimensions", + "encoding_format", + "user", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to OpenRouter format. + """ + for param, value in non_default_params.items(): + if param in self.get_supported_openai_params(model): + optional_params[param] = value + return optional_params + + def get_error_class( + self, error_message: str, status_code: int, headers: Any + ) -> Any: + """ + Get the error class for OpenRouter errors. + """ + return OpenRouterException( + message=error_message, + status_code=status_code, + headers=headers, + ) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 2bbdfa17cde..22042f7d641 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -115,7 +115,7 @@ def _process_gemini_image( and (image_type := format or _get_image_mime_type_from_url(image_url)) is not None ): - file_data = FileDataType(file_uri=image_url, mime_type=image_type) + file_data = FileDataType(mime_type=image_type, file_uri=image_url) part = {"file_data": file_data} if media_resolution_enum is not None and model is not None: diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index ba1788a217f..91100cf7d7b 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -480,20 +480,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): or tool_name == VertexToolName.CODE_EXECUTION.value ): # code_execution maintained for backwards compatibility code_execution = self.get_tool_value(tool, "codeExecution") - elif tool_name and tool_name == VertexToolName.GOOGLE_SEARCH.value: - googleSearch = self.get_tool_value( - tool, VertexToolName.GOOGLE_SEARCH.value - ) - elif ( - tool_name and tool_name == VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value + elif tool_name and ( + tool_name == VertexToolName.GOOGLE_SEARCH.value + or tool_name == "google_search" ): - googleSearchRetrieval = self.get_tool_value( - tool, VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value - ) - elif tool_name and tool_name == VertexToolName.ENTERPRISE_WEB_SEARCH.value: - enterpriseWebSearch = self.get_tool_value( - tool, VertexToolName.ENTERPRISE_WEB_SEARCH.value - ) + googleSearch = self.get_tool_value(tool, tool_name) + elif tool_name and ( + tool_name == VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value + or tool_name == "google_search_retrieval" + ): + googleSearchRetrieval = self.get_tool_value(tool, tool_name) + elif tool_name and ( + tool_name == VertexToolName.ENTERPRISE_WEB_SEARCH.value + or tool_name == "enterprise_web_search" + ): + enterpriseWebSearch = self.get_tool_value(tool, tool_name) elif tool_name and ( tool_name == VertexToolName.URL_CONTEXT.value or tool_name == "urlContext" @@ -1811,6 +1812,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): functions: Optional[ChatCompletionToolCallFunctionChunk] = None thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None reasoning_content: Optional[str] = None + thought_signatures: Optional[Any] = None for idx, candidate in enumerate(_candidates): if "content" not in candidate: diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py index 712a06dece1..123d925f7c1 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py @@ -40,6 +40,7 @@ class PartnerModelPrefixes(str, Enum): GPT_OSS_PREFIX = "openai/gpt-oss-" MINIMAX_PREFIX = "minimaxai/" MOONSHOT_PREFIX = "moonshotai/" + ZAI_PREFIX = "zai-org/" class VertexAIPartnerModels(VertexBase): @@ -66,6 +67,7 @@ class VertexAIPartnerModels(VertexBase): or model.startswith(PartnerModelPrefixes.GPT_OSS_PREFIX) or model.startswith(PartnerModelPrefixes.MINIMAX_PREFIX) or model.startswith(PartnerModelPrefixes.MOONSHOT_PREFIX) + or model.startswith(PartnerModelPrefixes.ZAI_PREFIX) ): return True return False @@ -79,6 +81,7 @@ class VertexAIPartnerModels(VertexBase): PartnerModelPrefixes.GPT_OSS_PREFIX, PartnerModelPrefixes.MINIMAX_PREFIX, PartnerModelPrefixes.MOONSHOT_PREFIX, + PartnerModelPrefixes.ZAI_PREFIX, ] if any(provider in model for provider in OPENAI_LIKE_VERTEX_PROVIDERS): return True diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 826f151df35..a3606ff9deb 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -388,6 +388,10 @@ class VertexBase: Internal function. Returns the token and url for the call. Handles logic if it's google ai studio vs. vertex ai. + + For Vertex AI: + - If gemini_api_key is provided, use API key authentication (x-goog-api-key header) + - Otherwise, use service account credentials (OAuth2 Bearer token) Returns token, url @@ -400,7 +404,7 @@ class VertexBase: stream=stream, gemini_api_key=gemini_api_key, ) - auth_header = None # this field is not used for gemin + auth_header = None # this field is not used for gemini else: vertex_location = self.get_vertex_region( vertex_region=vertex_location, @@ -409,14 +413,32 @@ class VertexBase: ### SET RUNTIME ENDPOINT ### version = "v1beta1" if should_use_v1beta1_features is True else "v1" - url, endpoint = _get_vertex_url( - mode=mode, - model=model, - stream=stream, - vertex_project=vertex_project, - vertex_location=vertex_location, - vertex_api_version=version, - ) + + # Check if using API key authentication for Vertex AI + if gemini_api_key and not vertex_credentials: + # When using API key with Vertex AI, use the Google AI Studio endpoint + # This is because Vertex AI API keys work with generativelanguage.googleapis.com + verbose_logger.debug( + f"Using Vertex AI API key authentication for model: {model} - routing to Google AI Studio endpoint" + ) + url, endpoint = _get_gemini_url( + mode=mode, + model=model, + stream=stream, + gemini_api_key=gemini_api_key, + ) + # API key is already included in the URL by _get_gemini_url + auth_header = None + else: + # Use OAuth2 Bearer token authentication (traditional Vertex AI) + url, endpoint = _get_vertex_url( + mode=mode, + model=model, + stream=stream, + vertex_project=vertex_project, + vertex_location=vertex_location, + vertex_api_version=version, + ) return self._check_custom_proxy( api_base=api_base, diff --git a/litellm/llms/watsonx/audio_transcription/transformation.py b/litellm/llms/watsonx/audio_transcription/transformation.py index 186d858321a..c7e6a77b96f 100644 --- a/litellm/llms/watsonx/audio_transcription/transformation.py +++ b/litellm/llms/watsonx/audio_transcription/transformation.py @@ -7,13 +7,14 @@ WatsonX follows the OpenAI spec for audio transcription. from typing import Any, Dict, List, Optional import litellm +from httpx import Response from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.types.llms.openai import ( AllMessageValues, OpenAIAudioTranscriptionOptionalParams, ) from litellm.types.llms.watsonx import WatsonXAudioTranscriptionRequestBody -from litellm.types.utils import FileTypes +from litellm.types.utils import FileTypes, TranscriptionResponse from ...base_llm.audio_transcription.transformation import ( AudioTranscriptionRequestData, @@ -156,3 +157,48 @@ class IBMWatsonXAudioTranscriptionConfig( url = f"{url}?version={api_version}" return url + + def transform_audio_transcription_response( + self, + raw_response: Response, + ) -> TranscriptionResponse: + """ + Transform the audio transcription response from WatsonX. + + WatsonX may include a 'model' field in the response, which needs to be + removed before creating the TranscriptionResponse object. + """ + try: + raw_response_json = raw_response.json() + except Exception as e: + raise ValueError( + f"Error transforming response to json: {str(e)}\nResponse: {raw_response.text}" + ) + + # Extract only valid fields for TranscriptionResponse.__init__() + # TranscriptionResponse only accepts 'text' and 'usage' in __init__() + text = raw_response_json.get("text") + usage = raw_response_json.get("usage") + + # Create response with only valid fields + response_kwargs = {} + if text is not None: + response_kwargs["text"] = text + if usage is not None: + response_kwargs["usage"] = usage + + if not response_kwargs: + raise ValueError( + "Invalid response format. Received response does not match the expected format. Got: ", + raw_response_json, + ) + + response = TranscriptionResponse(**response_kwargs) + + # Add other fields using dictionary-style assignment (like duration, task, etc.) + # Skip fields that TranscriptionResponse doesn't accept in __init__() + for key, value in raw_response_json.items(): + if key not in ["text", "usage", "model"]: # text/usage already set, model should be excluded + response[key] = value + + return response diff --git a/litellm/main.py b/litellm/main.py index e8a8b504d96..10e3bcac04b 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -189,7 +189,7 @@ from .llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from .llms.custom_llm import CustomLLM, custom_chat_llm_router from .llms.databricks.embed.handler import DatabricksEmbeddingHandler from .llms.deprecated_providers import aleph_alpha, palm -from .llms.gemini.common_utils import get_api_key_from_env +from .llms.gemini.common_utils import get_api_key_from_env, get_vertex_api_key_from_env from .llms.groq.chat.handler import GroqChatCompletion from .llms.heroku.chat.transformation import HerokuChatConfig from .llms.huggingface.embedding.handler import HuggingFaceEmbedding @@ -3230,6 +3230,12 @@ def completion( # type: ignore # noqa: PLR0915 or get_secret("VERTEXAI_CREDENTIALS") ) + vertex_api_key = ( + api_key + or get_vertex_api_key_from_env() + or litellm.api_key + ) + api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE") new_params = safe_deep_copy(optional_params or {}) @@ -3271,7 +3277,7 @@ def completion( # type: ignore # noqa: PLR0915 vertex_location=vertex_ai_location, vertex_project=vertex_ai_project, vertex_credentials=vertex_credentials, - gemini_api_key=None, + gemini_api_key=vertex_api_key, # Support for Vertex AI API Key logging_obj=logging, acompletion=acompletion, timeout=timeout, @@ -4701,6 +4707,51 @@ def embedding( # noqa: PLR0915 litellm_params=litellm_params_dict, headers=headers, ) + elif custom_llm_provider == "openrouter": + api_base = ( + api_base + or litellm.api_base + or get_secret_str("OPENROUTER_API_BASE") + or "https://openrouter.ai/api/v1" + ) + + api_key = ( + api_key + or litellm.api_key + or litellm.openrouter_key + or get_secret("OPENROUTER_API_KEY") + or get_secret("OR_API_KEY") + ) + + openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai" + openrouter_app_name = get_secret("OR_APP_NAME") or "liteLLM" + + openrouter_headers = { + "HTTP-Referer": openrouter_site_url, + "X-Title": openrouter_app_name, + } + + _headers = headers or litellm.headers + if _headers: + openrouter_headers.update(_headers) + + headers = openrouter_headers + + response = base_llm_http_handler.embedding( + model=model, + input=input, + custom_llm_provider=custom_llm_provider, + api_base=api_base, + api_key=api_key, + logging_obj=logging, + timeout=timeout, + model_response=EmbeddingResponse(), + optional_params=optional_params, + client=client, + aembedding=aembedding, + litellm_params=litellm_params_dict, + headers=headers, + ) elif custom_llm_provider == "huggingface": api_key = ( api_key diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index fb00f636409..73579db75cd 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -28345,6 +28345,19 @@ "supports_tool_choice": true, "supports_web_search": true }, + "vertex_ai/zai-org/glm-4.7-maas": { + "input_cost_per_token": 3e-07, + "litellm_provider": "vertex_ai-zai_models", + "max_input_tokens": 200000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "vertex_ai/mistral-medium-3": { "input_cost_per_token": 4e-07, "litellm_provider": "vertex_ai-mistral_models", diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 3df3037ed58..df4737cec85 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -216,6 +216,11 @@ def llm_passthrough_route( ) litellm_params_dict = get_litellm_params(**kwargs) + + # Add model_id to litellm_params if present in kwargs (for Bedrock Application Inference Profiles) + if "model_id" in kwargs: + litellm_params_dict["model_id"] = kwargs["model_id"] + litellm_logging_obj.update_environment_variables( model=model, litellm_params=litellm_params_dict, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3a548e203c5..1029f2241a1 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -551,6 +551,7 @@ class MCPServerManager: allowed_tools=getattr(mcp_server, "allowed_tools", None), disallowed_tools=getattr(mcp_server, "disallowed_tools", None), allow_all_keys=mcp_server.allow_all_keys, + updated_at=getattr(mcp_server, "updated_at", None), ) return new_server @@ -697,9 +698,7 @@ class MCPServerManager: results = await asyncio.gather(*tasks) # Flatten results into single list - list_tools_result: List[MCPTool] = [ - tool for tools in results for tool in tools - ] + list_tools_result: List[MCPTool] = [tool for tools in results for tool in tools] verbose_logger.info( f"Successfully fetched {len(list_tools_result)} tools total from all servers" @@ -2059,7 +2058,8 @@ class MCPServerManager: return None - async def _add_mcp_servers_from_db_to_in_memory_registry(self): + async def reload_servers_from_database(self): + """Re-synchronize the in-memory MCP server registry with the database.""" from litellm.proxy._experimental.mcp_server.db import get_all_mcp_servers from litellm.proxy.management_endpoints.mcp_management_endpoints import ( get_prisma_client_or_throw, @@ -2074,15 +2074,34 @@ class MCPServerManager: db_mcp_servers = await get_all_mcp_servers(prisma_client) verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database") - # ensure the global_mcp_server_manager is up to date with the db + previous_registry = self.registry + new_registry: Dict[str, MCPServer] = {} + for server in db_mcp_servers: + existing_server = previous_registry.get(server.server_id) + + if ( + existing_server is not None + and existing_server.updated_at is not None + and server.updated_at is not None + and existing_server.updated_at == server.updated_at + ): + # Re-use existing server instance to avoid re-running build_mcp_server_from_table() + # which can perform network discovery for OAuth2 servers. + new_registry[server.server_id] = existing_server + continue + verbose_logger.debug( - f"Adding server to registry: {server.server_id} ({server.server_name})" + f"Building server from DB: {server.server_id} ({server.server_name})" ) - await self.add_server(server) + new_registry[server.server_id] = await self.build_mcp_server_from_table( + server + ) + + self.registry = new_registry verbose_logger.debug( - f"Registry now contains {len(self.get_registry())} servers" + "MCP registry refreshed (%s servers in registry)", len(new_registry) ) def get_mcp_servers_from_ids(self, server_ids: List[str]) -> List[MCPServer]: @@ -2369,13 +2388,6 @@ class MCPServerManager: servers.append(self._build_mcp_server_table(server)) return servers - async def reload_servers_from_database(self): - """ - Public method to reload all MCP servers from database into registry. - This can be called from management endpoints to ensure registry is up to date. - """ - await self._add_mcp_servers_from_db_to_in_memory_registry() - async def get_all_mcp_servers_with_health_unfiltered( self, server_ids: Optional[List[str]] = None ) -> List[LiteLLM_MCPServerTable]: diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 4c947b99ba3..642cb0cec2d 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -218,10 +218,10 @@ if MCP_AVAILABLE: from fastapi import HTTPException from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException - from litellm.proxy.proxy_server import add_litellm_data_to_request, proxy_config from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) + from litellm.proxy.proxy_server import add_litellm_data_to_request, proxy_config try: data = await request.json() @@ -252,7 +252,12 @@ if MCP_AVAILABLE: if mcp_server_auth_headers: data["mcp_server_auth_headers"] = mcp_server_auth_headers data["raw_headers"] = raw_headers_from_request - + + # Extract user_api_key_auth from metadata and add to top level + # call_mcp_tool expects user_api_key_auth as a top-level parameter + if "metadata" in data and "user_api_key_auth" in data["metadata"]: + data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"] + result = await call_mcp_tool(**data) return result except BlockedPiiEntityError as e: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 4140273ea25..09c8bb562f7 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1942,7 +1942,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): description="connect to a postgres db - needed for generating temporary keys + tracking spend / key", ) database_connection_pool_limit: Optional[int] = Field( - 100, + 10, description="default connection pool for prisma client connecting to postgres db", ) database_connection_timeout: Optional[float] = Field( @@ -2104,6 +2104,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase): rotation_interval: Optional[str] = None # How often to rotate (e.g., "30d", "90d") last_rotation_at: Optional[datetime] = None # When this key was last rotated key_rotation_at: Optional[datetime] = None # When this key should next be rotated + router_settings: Optional[Dict] = None # Router settings for this key (Key > Team > Global precedence) model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 334362a0271..7de6b7fccfc 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -11,7 +11,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.integrations.custom_guardrail import ModifyResponseException from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, - create_streaming_response, + create_response, ) from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.types.utils import TokenCountResponse @@ -106,7 +106,7 @@ async def anthropic_response( # noqa: PLR0915 ) ) - return await create_streaming_response( + return await create_response( generator=selected_data_generator, media_type="text/event-stream", headers={}, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index de4973ecc69..26778ece60e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -147,6 +147,7 @@ async def common_checks( # 3.1. If organization is in budget await _organization_max_budget_check( valid_token=valid_token, + team_object=team_object, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -2310,61 +2311,86 @@ async def _team_max_budget_check( async def _organization_max_budget_check( valid_token: Optional[UserAPIKeyAuth], + team_object: Optional[LiteLLM_TeamTable], prisma_client: Optional[PrismaClient], user_api_key_cache: DualCache, proxy_logging_obj: ProxyLogging, ): """ Check if the organization is over its max budget. + + This function checks the organization budget using: + 1. First, tries to use valid_token.org_id (if key has organization_id set) + 2. Falls back to team_object.organization_id (if key doesn't have org_id but team does) + + This ensures organization budget checks work even when keys don't have organization_id + set directly, as long as their team belongs to an organization. Raises: BudgetExceededError if the organization is over its max budget. Triggers a budget alert if the organization is over its max budget. """ - # Only check if token has organization info and organization_max_budget is set - if ( - valid_token is None - or valid_token.org_id is None - or valid_token.organization_max_budget is None - or valid_token.organization_max_budget <= 0 - ): + if valid_token is None or prisma_client is None: return - # Get organization object to check current spend - if prisma_client is not None: - org_table = await get_org_object( - org_id=valid_token.org_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, + # Determine organization_id: first try from token, then fallback to team + org_id: Optional[str] = None + if valid_token.org_id is not None: + org_id = valid_token.org_id + elif team_object is not None and team_object.organization_id is not None: + org_id = team_object.organization_id + + # If no organization_id found, skip the check + if org_id is None: + return + + # Get organization object with budget table to check current spend and max budget + try: + org_table = await prisma_client.db.litellm_organizationtable.find_unique( + where={"organization_id": org_id}, + include={"litellm_budget_table": True}, + ) + except Exception: + # If organization lookup fails, skip the check + return + + if org_table is None: + return + + # Get max_budget from organization's budget table + org_max_budget: Optional[float] = None + if org_table.litellm_budget_table is not None: + org_max_budget = org_table.litellm_budget_table.max_budget + + # Only check if organization has a valid max_budget set + if org_max_budget is None or org_max_budget <= 0: + return + + # Check if organization spend exceeds max budget + if org_table.spend >= org_max_budget: + # Trigger budget alert + call_info = CallInfo( + token=valid_token.token, + spend=org_table.spend, + max_budget=org_max_budget, + user_id=valid_token.user_id, + team_id=valid_token.team_id, + team_alias=valid_token.team_alias, + organization_id=org_id, + event_group=Litellm_EntityType.ORGANIZATION, + ) + asyncio.create_task( + proxy_logging_obj.budget_alerts( + type="organization_budget", + user_info=call_info, + ) ) - if ( - org_table is not None - and org_table.spend >= valid_token.organization_max_budget - ): - # Trigger budget alert - call_info = CallInfo( - token=valid_token.token, - spend=org_table.spend, - max_budget=valid_token.organization_max_budget, - user_id=valid_token.user_id, - team_id=valid_token.team_id, - team_alias=valid_token.team_alias, - organization_id=valid_token.org_id, - event_group=Litellm_EntityType.ORGANIZATION, - ) - asyncio.create_task( - proxy_logging_obj.budget_alerts( - type="organization_budget", - user_info=call_info, - ) - ) - - raise litellm.BudgetExceededError( - current_cost=org_table.spend, - max_budget=valid_token.organization_max_budget, - message=f"Budget has been exceeded! Organization={valid_token.org_id} Current cost: {org_table.spend}, Max budget: {valid_token.organization_max_budget}", - ) + raise litellm.BudgetExceededError( + current_cost=org_table.spend, + max_budget=org_max_budget, + message=f"Budget has been exceeded! Organization={org_id} Current cost: {org_table.spend}, Max budget: {org_max_budget}", + ) async def _tag_max_budget_check( @@ -2601,4 +2627,4 @@ def _can_object_call_vector_stores( code=status.HTTP_401_UNAUTHORIZED, ) - return True + return True \ No newline at end of file diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 7a71af1da5c..797540deaa4 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -426,38 +426,65 @@ def get_key_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, ) -> Optional[Dict[str, int]]: """ - Get the model rpm limit for a given api key - - check key metadata - - check key model max budget - - check team metadata + Get the model rpm limit for a given api key. + + Priority order (returns first found): + 1. Key metadata (model_rpm_limit) + 2. Key model_max_budget (rpm_limit per model) + 3. Team metadata (model_rpm_limit) """ + # 1. Check key metadata first (takes priority) if user_api_key_dict.metadata: - if "model_rpm_limit" in user_api_key_dict.metadata: - return user_api_key_dict.metadata["model_rpm_limit"] - elif user_api_key_dict.model_max_budget: + result = user_api_key_dict.metadata.get("model_rpm_limit") + if result: + return result + + # 2. Check model_max_budget + if user_api_key_dict.model_max_budget: model_rpm_limit: Dict[str, Any] = {} for model, budget in user_api_key_dict.model_max_budget.items(): - if "rpm_limit" in budget and budget["rpm_limit"] is not None: + if isinstance(budget, dict) and budget.get("rpm_limit") is not None: model_rpm_limit[model] = budget["rpm_limit"] - return model_rpm_limit - elif user_api_key_dict.team_metadata: - if "model_rpm_limit" in user_api_key_dict.team_metadata: - return user_api_key_dict.team_metadata["model_rpm_limit"] + if model_rpm_limit: + return model_rpm_limit + + # 3. Fallback to team metadata + if user_api_key_dict.team_metadata: + return user_api_key_dict.team_metadata.get("model_rpm_limit") + return None def get_key_model_tpm_limit( user_api_key_dict: UserAPIKeyAuth, ) -> Optional[Dict[str, int]]: + """ + Get the model tpm limit for a given api key. + + Priority order (returns first found): + 1. Key metadata (model_tpm_limit) + 2. Key model_max_budget (tpm_limit per model) + 3. Team metadata (model_tpm_limit) + """ + # 1. Check key metadata first (takes priority) if user_api_key_dict.metadata: - if "model_tpm_limit" in user_api_key_dict.metadata: - return user_api_key_dict.metadata["model_tpm_limit"] - elif user_api_key_dict.model_max_budget: - if "tpm_limit" in user_api_key_dict.model_max_budget: - return user_api_key_dict.model_max_budget["tpm_limit"] - elif user_api_key_dict.team_metadata: - if "model_tpm_limit" in user_api_key_dict.team_metadata: - return user_api_key_dict.team_metadata["model_tpm_limit"] + result = user_api_key_dict.metadata.get("model_tpm_limit") + if result: + return result + + # 2. Check model_max_budget (iterate per-model like RPM does) + if user_api_key_dict.model_max_budget: + model_tpm_limit: Dict[str, Any] = {} + for model, budget in user_api_key_dict.model_max_budget.items(): + if isinstance(budget, dict) and budget.get("tpm_limit") is not None: + model_tpm_limit[model] = budget["tpm_limit"] + if model_tpm_limit: + return model_tpm_limit + + # 3. Fallback to team metadata + if user_api_key_dict.team_metadata: + return user_api_key_dict.team_metadata.get("model_tpm_limit") + return None @@ -469,7 +496,8 @@ def get_model_rate_limit_from_metadata( if getattr(user_api_key_dict, metadata_accessor_key): return getattr(user_api_key_dict, metadata_accessor_key).get(rate_limit_key) return None - + + def get_team_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, ) -> Optional[Dict[str, int]]: diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 8cc33ce6cdd..5be44f479b8 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -188,7 +188,7 @@ async def authenticate_user( # noqa: PLR0915 _user_row = cast( Optional[LiteLLM_UserTable], await prisma_client.db.litellm_usertable.find_first( - where={"user_email": {"equals": username}} + where={"user_email": {"equals": username, "mode": "insensitive"}} ), ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 9b53d9a3a80..401aa7fd443 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -138,7 +138,7 @@ def _apply_budget_limits_to_end_user_params( ) -> None: """ Helper function to apply budget limits to end user parameters. - + Args: end_user_params: Dictionary to update with budget parameters budget_info: Budget table object containing limits @@ -146,16 +146,14 @@ def _apply_budget_limits_to_end_user_params( """ if budget_info.tpm_limit is not None: end_user_params["end_user_tpm_limit"] = budget_info.tpm_limit - + if budget_info.rpm_limit is not None: end_user_params["end_user_rpm_limit"] = budget_info.rpm_limit - + if budget_info.max_budget is not None: end_user_params["end_user_max_budget"] = budget_info.max_budget - - verbose_proxy_logger.debug( - f"Applied budget limits to end user {end_user_id}" - ) + + verbose_proxy_logger.debug(f"Applied budget limits to end user {end_user_id}") async def user_api_key_auth_websocket(websocket: WebSocket): @@ -170,12 +168,10 @@ async def user_api_key_auth_websocket(websocket: WebSocket): model = query_params.get("model") - async def return_body(): return _realtime_request_body(model) - - request.body = return_body # type: ignore + request.body = return_body # type: ignore authorization = websocket.headers.get("authorization") # If no Authorization header, try the api-key header @@ -586,7 +582,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if team_membership is not None else None ), - team_metadata=team_object.metadata if team_object is not None else None, + team_metadata=team_object.metadata + if team_object is not None + else None, ) # run through common checks _ = await common_checks( @@ -669,9 +667,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 route=route, ) if _end_user_object is not None: - end_user_params["allowed_model_region"] = ( - _end_user_object.allowed_model_region - ) + end_user_params[ + "allowed_model_region" + ] = _end_user_object.allowed_model_region if _end_user_object.litellm_budget_table is not None: _apply_budget_limits_to_end_user_params( end_user_params=end_user_params, @@ -753,7 +751,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}", type=ProxyErrorTypes.expired_key, code=400, - param=api_key, + param=abbreviate_api_key(api_key=api_key), ) valid_token = update_valid_token_with_end_user_params( valid_token=valid_token, end_user_params=end_user_params @@ -994,7 +992,6 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 # Check 3. Check if user is in their team budget if valid_token.team_member_spend is not None: - if prisma_client is not None: _cache_key = f"{valid_token.team_id}_{valid_token.user_id}" @@ -1055,7 +1052,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}", type=ProxyErrorTypes.expired_key, code=400, - param=api_key, + param=abbreviate_api_key(api_key=api_key), ) # Check 4. Token Spend is under budget @@ -1216,8 +1213,6 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 ) - - @tracer.wrap() async def user_api_key_auth( request: Request, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 537b48f06ed..61bffde3aca 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -17,7 +17,7 @@ from typing import ( import httpx import orjson from fastapi import HTTPException, Request, status -from fastapi.responses import Response, StreamingResponse +from fastapi.responses import JSONResponse, Response, StreamingResponse import litellm from litellm._logging import verbose_proxy_logger @@ -96,16 +96,55 @@ async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional return None -async def create_streaming_response( +def _extract_error_from_sse_chunk(event_line: Union[str, bytes]) -> dict: + """ + Extract error dictionary from SSE format chunk. + + Args: + event_line: SSE format event line, e.g. "data: {"error": {...}}\n\n" + + Returns: + Error dictionary in OpenAI API format + """ + event_line = ( + event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line + ) + + # Default error format + default_error = { + "message": "Unknown error", + "type": "internal_server_error", + "param": None, + "code": "500", + } + + if event_line.startswith("data: "): + json_str = event_line[len("data: ") :].strip() + if not json_str or json_str == "[DONE]": + return default_error + + try: + data = orjson.loads(json_str) + if isinstance(data, dict) and "error" in data: + error_obj = data["error"] + if isinstance(error_obj, dict): + return error_obj + except (orjson.JSONDecodeError, json.JSONDecodeError): + pass + + return default_error + + +async def create_response( generator: AsyncGenerator[str, None], media_type: str, headers: dict, default_status_code: int = status.HTTP_200_OK, -) -> StreamingResponse: +) -> Union[StreamingResponse, JSONResponse]: """ - Creates a StreamingResponse by inspecting the first chunk for an error code. - The entire original generator content is streamed, but the HTTP status code - of the response is set based on the first chunk if it's a recognized error. + Create streaming response, checking if the first chunk is an error. + If the first chunk is an error, return a standard JSON error response. + Otherwise, return StreamingResponse and stream all content. """ first_chunk_value: Optional[str] = None final_status_code = default_status_code @@ -124,9 +163,27 @@ async def create_streaming_response( first_chunk_value ) if error_code_from_chunk is not None: + # First chunk is an error, stream hasn't really started yet + # Should return standard JSON error response instead of SSE format final_status_code = error_code_from_chunk verbose_proxy_logger.debug( - f"Error detected in first stream chunk. Status code set to: {final_status_code}" + f"Error detected in first stream chunk. Returning JSON error response with status code: {final_status_code}" + ) + + # Parse error content + error_dict = _extract_error_from_sse_chunk(first_chunk_value) + + # Consume and close generator (avoid resource leak) + try: + await generator.aclose() + except Exception: + pass + + # Return JSON format error response + return JSONResponse( + status_code=final_status_code, + content={"error": error_dict}, + headers=headers, ) except Exception as e: verbose_proxy_logger.debug(f"Error parsing first chunk value: {e}") @@ -237,7 +294,11 @@ class ProxyBaseLLMRequestProcessing: if response_cost is not None: try: # Convert response_cost to float if it's a string - cost_value = float(response_cost) if isinstance(response_cost, str) else response_cost + cost_value = ( + float(response_cost) + if isinstance(response_cost, str) + else response_cost + ) if cost_value > 0: updated_spend = current_spend + cost_value except (ValueError, TypeError): @@ -376,6 +437,16 @@ class ProxyBaseLLMRequestProcessing: ) -> Tuple[dict, LiteLLMLoggingObj]: start_time = datetime.now() # start before calling guardrail hooks + # Calculate request queue time if arrival_time is available + # Use start_time.timestamp() to avoid extra time.time() call for better performance + proxy_server_request = self.data.get("proxy_server_request", {}) + arrival_time = proxy_server_request.get("arrival_time") + queue_time_seconds = None + if arrival_time is not None: + # Convert start_time (datetime) to timestamp for calculation + processing_start_time = start_time.timestamp() + queue_time_seconds = processing_start_time - arrival_time + self.data = await add_litellm_data_to_request( data=self.data, request=request, @@ -385,6 +456,19 @@ class ProxyBaseLLMRequestProcessing: proxy_config=proxy_config, ) + # Store queue time in metadata after add_litellm_data_to_request to ensure it's preserved + if queue_time_seconds is not None: + from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name + + _metadata_variable_name = _get_metadata_variable_name(request) + if _metadata_variable_name not in self.data: + self.data[_metadata_variable_name] = {} + if not isinstance(self.data[_metadata_variable_name], dict): + self.data[_metadata_variable_name] = {} + self.data[_metadata_variable_name][ + "queue_time_seconds" + ] = queue_time_seconds + self.data["model"] = ( general_settings.get("completion_model", None) # server default or user_model # model name passed via cli args @@ -440,6 +524,29 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type # type: ignore ) + # Apply hierarchical router_settings (Key > Team > Global) + if llm_router is not None and proxy_config is not None: + from litellm.proxy.proxy_server import prisma_client + + router_settings = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) + + # If router_settings found (from key, team, or global), apply them + # This ensures key/team settings override global settings + if router_settings is not None and router_settings: + # Get model_list from current router + model_list = llm_router.get_model_list() + if model_list is not None: + # Create user_config with model_list and router_settings + # This creates a per-request router with the hierarchical settings + user_config = { + "model_list": model_list, + **router_settings + } + self.data["user_config"] = user_config + if "messages" in self.data and self.data["messages"]: logging_obj.update_messages(self.data["messages"]) @@ -647,7 +754,7 @@ class ProxyBaseLLMRequestProcessing: proxy_logging_obj=proxy_logging_obj, ) ) - return await create_streaming_response( + return await create_response( generator=selected_data_generator, media_type="text/event-stream", headers=custom_headers, @@ -658,7 +765,7 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict=user_api_key_dict, request_data=self.data, ) - return await create_streaming_response( + return await create_response( generator=selected_data_generator, media_type="text/event-stream", headers=custom_headers, @@ -900,11 +1007,11 @@ class ProxyBaseLLMRequestProcessing: @staticmethod def _get_pre_call_type( route_type: Literal["acompletion", "aembedding", "aresponses", "allm_passthrough_route"], - ) -> Literal["completion", "embeddings", "responses", "allm_passthrough_route"]: + ) -> Literal["completion", "embedding", "responses", "allm_passthrough_route"]: if route_type == "acompletion": return "completion" elif route_type == "aembedding": - return "embeddings" + return "embedding" elif route_type == "aresponses": return "responses" elif route_type == "allm_passthrough_route": @@ -1155,9 +1262,9 @@ class ProxyBaseLLMRequestProcessing: # Add cache-related fields to **params (handled by Usage.__init__) if cache_creation_input_tokens is not None: - usage_kwargs["cache_creation_input_tokens"] = ( - cache_creation_input_tokens - ) + usage_kwargs[ + "cache_creation_input_tokens" + ] = cache_creation_input_tokens if cache_read_input_tokens is not None: usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index 406ddceabf5..c9c0cfe8f68 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -17,50 +17,141 @@ from litellm.secret_managers.main import str_to_bool class PrismaWrapper: + """ + Wrapper around Prisma client that handles RDS IAM token authentication. + + When iam_token_db_auth is enabled, this wrapper: + 1. Proactively refreshes IAM tokens before they expire (background task) + 2. Falls back to synchronous refresh if a token is found expired + 3. Uses proper locking to prevent race conditions during reconnection + + RDS IAM tokens are valid for 15 minutes. This wrapper refreshes them + 3 minutes before expiration to ensure uninterrupted database connectivity. + """ + + # Buffer time in seconds before token expiration to trigger refresh + # Refresh 3 minutes (180 seconds) before the token expires + TOKEN_REFRESH_BUFFER_SECONDS = 180 + + # Fallback refresh interval if token parsing fails (10 minutes) + FALLBACK_REFRESH_INTERVAL_SECONDS = 600 + def __init__(self, original_prisma: Any, iam_token_db_auth: bool): self._original_prisma = original_prisma self.iam_token_db_auth = iam_token_db_auth + # Background token refresh task management + self._token_refresh_task: Optional[asyncio.Task] = None + self._reconnection_lock = asyncio.Lock() + self._last_refresh_time: Optional[datetime] = None + + def _extract_token_from_db_url(self, db_url: Optional[str]) -> Optional[str]: + """ + Extract the token (password) from the DATABASE_URL. + + The token contains the AWS signature with X-Amz-Date and X-Amz-Expires parameters. + + Important: We must parse the URL while it's still encoded to preserve structure, + then decode the password portion. Otherwise the '?' in the token breaks URL parsing. + """ + if db_url is None: + return None + try: + # Parse URL while still encoded to preserve structure + parsed = urllib.parse.urlparse(db_url) + if parsed.password: + # Now decode just the password/token + return urllib.parse.unquote(parsed.password) + return None + except Exception: + return None + + def _parse_token_expiration(self, token: Optional[str]) -> Optional[datetime]: + """ + Parse the token to extract its expiration time. + + Returns the datetime when the token expires, or None if parsing fails. + """ + if token is None: + return None + + try: + # Token format: ...?X-Amz-Date=YYYYMMDDTHHMMSSZ&X-Amz-Expires=900&... + if "?" not in token: + return None + + query_string = token.split("?", 1)[1] + params = urllib.parse.parse_qs(query_string) + + expires_str = params.get("X-Amz-Expires", [None])[0] + date_str = params.get("X-Amz-Date", [None])[0] + + if not expires_str or not date_str: + return None + + token_created = datetime.strptime(date_str, "%Y%m%dT%H%M%SZ") + expires_in = int(expires_str) + + return token_created + timedelta(seconds=expires_in) + except Exception as e: + verbose_proxy_logger.debug(f"Failed to parse token expiration: {e}") + return None + + def _calculate_seconds_until_refresh(self) -> float: + """ + Calculate exactly how many seconds until we need to refresh the token. + + Uses precise timing: sleeps until (token_expiration - buffer_seconds). + For a 15-minute (900s) token with 180s buffer, this returns ~720s (12 min). + + Returns: + Number of seconds to sleep before the next refresh. + Returns 0 if token should be refreshed immediately. + Returns FALLBACK_REFRESH_INTERVAL_SECONDS if parsing fails. + """ + db_url = os.getenv("DATABASE_URL") + token = self._extract_token_from_db_url(db_url) + expiration_time = self._parse_token_expiration(token) + + if expiration_time is None: + # If we can't parse the token, use fallback interval + verbose_proxy_logger.debug( + f"Could not parse token expiration, using fallback interval of " + f"{self.FALLBACK_REFRESH_INTERVAL_SECONDS}s" + ) + return self.FALLBACK_REFRESH_INTERVAL_SECONDS + + # Calculate when we should refresh (expiration - buffer) + refresh_at = expiration_time - timedelta( + seconds=self.TOKEN_REFRESH_BUFFER_SECONDS + ) + + # How long until refresh time? + now = datetime.utcnow() + seconds_until_refresh = (refresh_at - now).total_seconds() + + # If already past refresh time, return 0 (refresh immediately) + return max(0, seconds_until_refresh) + def is_token_expired(self, token_url: Optional[str]) -> bool: + """Check if the token in the given URL is expired.""" if token_url is None: return True - # Decode the token URL to handle URL-encoded characters - decoded_url = urllib.parse.unquote(token_url) - # Parse the token URL - parsed_url = urllib.parse.urlparse(decoded_url) + token = self._extract_token_from_db_url(token_url) + expiration_time = self._parse_token_expiration(token) - # Parse the query parameters from the path component (if they exist there) - query_params = urllib.parse.parse_qs(parsed_url.query) + if expiration_time is None: + # If we can't parse the token, assume it's expired to trigger refresh + verbose_proxy_logger.debug( + "Could not parse token expiration, treating as expired" + ) + return True - # Get expiration time from the query parameters - expires = query_params.get("X-Amz-Expires", [None])[0] - if expires is None: - raise ValueError("X-Amz-Expires parameter is missing or invalid.") - - expires_int = int(expires) - - # Get the token's creation time from the X-Amz-Date parameter - token_time_str = query_params.get("X-Amz-Date", [""])[0] - if not token_time_str: - raise ValueError("X-Amz-Date parameter is missing or invalid.") - - # Ensure the token time string is parsed correctly - try: - token_time = datetime.strptime(token_time_str, "%Y%m%dT%H%M%SZ") - except ValueError as e: - raise ValueError(f"Invalid X-Amz-Date format: {e}") - - # Calculate the expiration time - expiration_time = token_time + timedelta(seconds=expires_int) - - # Current time in UTC - current_time = datetime.utcnow() - - # Check if the token is expired - return current_time > expiration_time + return datetime.utcnow() > expiration_time def get_rds_iam_token(self) -> Optional[str]: + """Generate a new RDS IAM token and update DATABASE_URL.""" if self.iam_token_db_auth: from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token @@ -74,7 +165,6 @@ class PrismaWrapper: db_host=db_host, db_port=db_port, db_user=db_user ) - # print(f"token: {token}") _db_url = f"postgresql://{db_user}:{token}@{db_host}:{db_port}/{db_name}" if db_schema: _db_url += f"?schema={db_schema}" @@ -86,6 +176,7 @@ class PrismaWrapper: async def recreate_prisma_client( self, new_db_url: str, http_client: Optional[Any] = None ): + """Disconnect and reconnect the Prisma client with a new database URL.""" from prisma import Prisma # type: ignore try: @@ -100,21 +191,159 @@ class PrismaWrapper: await self._original_prisma.connect() + async def start_token_refresh_task(self) -> None: + """ + Start the background token refresh task. + + This task proactively refreshes RDS IAM tokens before they expire, + preventing connection failures. Should be called after the initial + Prisma client connection is established. + """ + if not self.iam_token_db_auth: + verbose_proxy_logger.debug( + "IAM token auth not enabled, skipping token refresh task" + ) + return + + if self._token_refresh_task is not None: + verbose_proxy_logger.debug("Token refresh task already running") + return + + self._token_refresh_task = asyncio.create_task(self._token_refresh_loop()) + verbose_proxy_logger.info( + "Started RDS IAM token proactive refresh background task" + ) + + async def stop_token_refresh_task(self) -> None: + """ + Stop the background token refresh task gracefully. + + Should be called during application shutdown to clean up resources. + """ + if self._token_refresh_task is None: + return + + self._token_refresh_task.cancel() + try: + await self._token_refresh_task + except asyncio.CancelledError: + pass + self._token_refresh_task = None + verbose_proxy_logger.info("Stopped RDS IAM token refresh background task") + + async def _token_refresh_loop(self) -> None: + """ + Background loop that proactively refreshes RDS IAM tokens before expiration. + + Uses precise timing: calculates the exact sleep duration until the token + needs to be refreshed (expiration - 3 minute buffer), then refreshes. + This is more efficient than polling, requiring only 1 wake-up per token cycle. + """ + verbose_proxy_logger.info( + f"RDS IAM token refresh loop started. " + f"Tokens will be refreshed {self.TOKEN_REFRESH_BUFFER_SECONDS}s before expiration." + ) + + while True: + try: + # Calculate exactly how long to sleep until next refresh + sleep_seconds = self._calculate_seconds_until_refresh() + + if sleep_seconds > 0: + verbose_proxy_logger.info( + f"RDS IAM token refresh scheduled in {sleep_seconds:.0f} seconds " + f"({sleep_seconds / 60:.1f} minutes)" + ) + await asyncio.sleep(sleep_seconds) + + # Refresh the token + verbose_proxy_logger.info("Proactively refreshing RDS IAM token...") + await self._safe_refresh_token() + + except asyncio.CancelledError: + verbose_proxy_logger.info("RDS IAM token refresh loop cancelled") + break + except Exception as e: + verbose_proxy_logger.error( + f"Error in RDS IAM token refresh loop: {e}. " + f"Retrying in {self.FALLBACK_REFRESH_INTERVAL_SECONDS}s..." + ) + # On error, wait before retrying to avoid tight error loops + try: + await asyncio.sleep(self.FALLBACK_REFRESH_INTERVAL_SECONDS) + except asyncio.CancelledError: + break + + async def _safe_refresh_token(self) -> None: + """ + Refresh the RDS IAM token with proper locking to prevent race conditions. + + Uses an asyncio lock to ensure only one refresh operation happens at a time, + preventing multiple concurrent reconnection attempts. + """ + async with self._reconnection_lock: + new_db_url = self.get_rds_iam_token() + if new_db_url: + await self.recreate_prisma_client(new_db_url) + self._last_refresh_time = datetime.utcnow() + verbose_proxy_logger.info( + "RDS IAM token refreshed successfully. New token valid for ~15 minutes." + ) + else: + verbose_proxy_logger.error( + "Failed to generate new RDS IAM token during proactive refresh" + ) + def __getattr__(self, name: str): + """ + Proxy attribute access to the underlying Prisma client. + + If IAM token auth is enabled and the token is expired, this method + provides a synchronous fallback to refresh the token. However, this + should rarely be needed since the background task proactively refreshes + tokens before they expire. + + FIXED: Now properly waits for reconnection to complete before returning, + instead of the previous fire-and-forget pattern that caused the bug. + """ original_attr = getattr(self._original_prisma, name) + if self.iam_token_db_auth: db_url = os.getenv("DATABASE_URL") - if self.is_token_expired(db_url): - db_url = self.get_rds_iam_token() - loop = asyncio.get_event_loop() - if db_url: + # Check if token is expired (should be rare if background task is running) + if self.is_token_expired(db_url): + verbose_proxy_logger.warning( + "RDS IAM token expired in __getattr__ - proactive refresh may have failed. " + "Triggering synchronous fallback refresh..." + ) + + new_db_url = self.get_rds_iam_token() + if new_db_url: + loop = asyncio.get_event_loop() + if loop.is_running(): - asyncio.run_coroutine_threadsafe( - self.recreate_prisma_client(db_url), loop + # FIXED: Actually wait for the reconnection to complete! + # The previous code used fire-and-forget which caused the bug. + future = asyncio.run_coroutine_threadsafe( + self.recreate_prisma_client(new_db_url), loop ) + try: + # Wait up to 30 seconds for reconnection + future.result(timeout=30) + verbose_proxy_logger.info( + "Synchronous token refresh completed successfully" + ) + except Exception as e: + verbose_proxy_logger.error( + f"Failed to refresh token synchronously: {e}" + ) + raise else: - asyncio.run(self.recreate_prisma_client(db_url)) + asyncio.run(self.recreate_prisma_client(new_db_url)) + + # Get the NEW attribute after reconnection + original_attr = getattr(self._original_prisma, name) else: raise ValueError("Failed to get RDS IAM token") diff --git a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py index a6971b49f3b..87da11efad0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py @@ -5,6 +5,7 @@ # +-------------------------------------------------------------+ # Qualifire - Evaluate LLM outputs for quality, safety, and reliability +import json import os from typing import Any, Dict, List, Literal, Optional, Type @@ -15,12 +16,17 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.litellm_logging import ( Logging as LiteLLMLoggingObj, ) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel from litellm.types.utils import GenericGuardrailAPIInputs GUARDRAIL_NAME = "qualifire" +DEFAULT_QUALIFIRE_API_BASE = "https://proxy.qualifire.ai" class QualifireGuardrail(CustomGuardrail): @@ -44,7 +50,7 @@ class QualifireGuardrail(CustomGuardrail): Args: api_key: API key for Qualifire (or use QUALIFIRE_API_KEY env var) - api_base: Optional custom API base URL + api_base: Optional custom API base URL (defaults to https://api.qualifire.ai) evaluation_id: Pre-configured evaluation ID from Qualifire dashboard prompt_injections: Enable prompt injection detection (default if no other checks) hallucinations_check: Enable hallucination detection @@ -64,6 +70,7 @@ class QualifireGuardrail(CustomGuardrail): api_base or get_secret_str("QUALIFIRE_BASE_URL") or os.environ.get("QUALIFIRE_BASE_URL") + or DEFAULT_QUALIFIRE_API_BASE ) self.evaluation_id = evaluation_id self.prompt_injections = prompt_injections @@ -79,7 +86,11 @@ class QualifireGuardrail(CustomGuardrail): if not self._has_any_check_enabled() and not self.evaluation_id: self.prompt_injections = True - self._client = None + # Initialize async HTTP client for direct API calls + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + super().__init__(**kwargs) def _has_any_check_enabled(self) -> bool: @@ -96,43 +107,22 @@ class QualifireGuardrail(CustomGuardrail): ] ) - def _get_client(self): - """Lazy initialization of Qualifire client.""" - if self._client is None: - try: - from qualifire.client import Client - except ImportError: - raise ImportError( - "qualifire package is required for QualifireGuardrail. " - "Install it with: pip install qualifire" - ) - - client_kwargs: Dict[str, Any] = {} - if self.qualifire_api_key: - client_kwargs["api_key"] = self.qualifire_api_key - if self.qualifire_api_base: - client_kwargs["base_url"] = self.qualifire_api_base - - self._client = Client(**client_kwargs) - - return self._client - - def _convert_messages_to_qualifire_format( + def _convert_messages_to_api_format( self, messages: List[AllMessageValues] - ) -> List[Any]: + ) -> List[Dict[str, Any]]: """ - Convert LiteLLM messages to Qualifire's LLMMessage format. + Convert LiteLLM messages to Qualifire API format. Supports tool calls for tool_selection_quality_check. - """ - try: - from qualifire.types import LLMMessage, LLMToolCall - except ImportError: - raise ImportError( - "qualifire package is required for QualifireGuardrail. " - "Install it with: pip install qualifire" - ) - qualifire_messages = [] + Returns a list of dicts matching the API's ModelInvocationCanonicalMessage schema: + { + "role": "user" | "assistant" | "system" | "tool", + "content": "...", + "tool_call_id": "...", # optional + "tool_calls": [{"id": "...", "name": "...", "arguments": {...}}] # optional + } + """ + api_messages = [] for msg in messages: role = msg.get("role", "user") content = msg.get("content", "") @@ -147,42 +137,86 @@ class QualifireGuardrail(CustomGuardrail): text_parts.append(part) content = "\n".join(text_parts) - llm_message_kwargs: Dict[str, Any] = { + api_message: Dict[str, Any] = { "role": role, "content": content if isinstance(content, str) else str(content), } + # Handle tool_call_id for tool response messages + tool_call_id = msg.get("tool_call_id") + if tool_call_id: + api_message["tool_call_id"] = tool_call_id + # Handle tool calls if present tool_calls = msg.get("tool_calls") if tool_calls and isinstance(tool_calls, list): - qualifire_tool_calls = [] + api_tool_calls = [] for tc in tool_calls: if isinstance(tc, dict): function_info = tc.get("function", {}) # Arguments can be a string (JSON) or dict args = function_info.get("arguments", {}) if isinstance(args, str): - import json - try: args = json.loads(args) except json.JSONDecodeError: args = {} - qualifire_tool_calls.append( - LLMToolCall( - id=tc.get("id") or "", - name=function_info.get("name") or "", - arguments=args if isinstance(args, dict) else {}, - ) + api_tool_calls.append( + { + "id": tc.get("id") or "", + "name": function_info.get("name") or "", + "arguments": args if isinstance(args, dict) else {}, + } ) - if qualifire_tool_calls: - llm_message_kwargs["tool_calls"] = qualifire_tool_calls + if api_tool_calls: + api_message["tool_calls"] = api_tool_calls - qualifire_messages.append(LLMMessage(**llm_message_kwargs)) + api_messages.append(api_message) - return qualifire_messages + return api_messages - def _check_if_flagged(self, result: Any) -> bool: + def _convert_tools_to_api_format( + self, tools: Optional[List[Any]] + ) -> Optional[List[Dict[str, Any]]]: + """ + Convert OpenAI-format tools to Qualifire API format. + + Returns a list of dicts matching the API's ModelInvocationToolDefinition schema: + { + "name": "...", + "description": "...", + "parameters": {...} + } + """ + if not tools: + return None + + api_tools = [] + for tool in tools: + if isinstance(tool, dict): + # Handle OpenAI function tool format + if tool.get("type") == "function": + function_def = tool.get("function", {}) + api_tools.append( + { + "name": function_def.get("name", ""), + "description": function_def.get("description", ""), + "parameters": function_def.get("parameters", {}), + } + ) + # Handle direct tool format + elif "name" in tool: + api_tools.append( + { + "name": tool.get("name", ""), + "description": tool.get("description", ""), + "parameters": tool.get("parameters", {}), + } + ) + + return api_tools if api_tools else None + + def _check_if_flagged(self, result: Dict[str, Any]) -> bool: """ Check if the Qualifire evaluation result indicates flagged content. @@ -190,65 +224,53 @@ class QualifireGuardrail(CustomGuardrail): A high score (close to 100) indicates GOOD content, low score indicates problems. """ # Check evaluation results for any flagged items - evaluation_results = getattr(result, "evaluationResults", None) or [] - if isinstance(result, dict): - evaluation_results = result.get("evaluationResults", []) or [] + evaluation_results = result.get("evaluationResults", []) or [] for eval_result in evaluation_results: - results: List[Any] = [] - if isinstance(eval_result, dict): - results = eval_result.get("results", []) or [] - else: - results = getattr(eval_result, "results", []) or [] - + results = eval_result.get("results", []) or [] for r in results: - flagged = ( - r.get("flagged") - if isinstance(r, dict) - else getattr(r, "flagged", False) - ) - if flagged: + if r.get("flagged"): return True return False - def _build_evaluate_kwargs( + def _build_evaluate_payload( self, - qualifire_messages: List[Any], + api_messages: List[Dict[str, Any]], output: Optional[str], assertions: Optional[List[str]], - available_tools: Optional[List[Any]], + available_tools: Optional[List[Dict[str, Any]]], ) -> Dict[str, Any]: - """Build kwargs dictionary for the evaluate call.""" - kwargs: Dict[str, Any] = {"messages": qualifire_messages} + """Build payload dictionary for the /api/evaluation/evaluate endpoint.""" + payload: Dict[str, Any] = {"messages": api_messages} if output is not None: - kwargs["output"] = output + payload["output"] = output # Add enabled checks if self.prompt_injections: - kwargs["prompt_injections"] = True + payload["prompt_injections"] = True if self.hallucinations_check: - kwargs["hallucinations_check"] = True + payload["hallucinations_check"] = True if self.grounding_check: - kwargs["grounding_check"] = True + payload["grounding_check"] = True if self.pii_check: - kwargs["pii_check"] = True + payload["pii_check"] = True if self.content_moderation_check: - kwargs["content_moderation_check"] = True + payload["content_moderation_check"] = True if self.tool_selection_quality_check: # Only enable tool_selection_quality_check if available_tools is provided if available_tools: - kwargs["tool_selection_quality_check"] = True - kwargs["available_tools"] = available_tools + payload["tool_selection_quality_check"] = True + payload["available_tools"] = available_tools else: verbose_proxy_logger.debug( "Qualifire Guardrail: tool_selection_quality_check enabled but no available_tools provided, skipping this check" ) if assertions: - kwargs["assertions"] = assertions + payload["assertions"] = assertions - return kwargs + return payload async def _run_qualifire_check( self, @@ -274,11 +296,17 @@ class QualifireGuardrail(CustomGuardrail): assertions = dynamic_params.get("assertions") or self.assertions on_flagged = dynamic_params.get("on_flagged") or self.on_flagged - try: - client = self._get_client() - qualifire_messages = self._convert_messages_to_qualifire_format(messages) + # Prepare headers + headers = { + "X-Qualifire-API-Key": self.qualifire_api_key or "", + "Content-Type": "application/json", + } - # Use invoke_evaluation if evaluation_id is provided + try: + # Convert messages to API format + api_messages = self._convert_messages_to_api_format(messages) + + # Use invoke endpoint if evaluation_id is provided if evaluation_id: # For invoke_evaluation, we need to extract input/output input_text = "" @@ -291,25 +319,47 @@ class QualifireGuardrail(CustomGuardrail): input_text = content break - result = client.invoke_evaluation( - evaluation_id=evaluation_id, - input=input_text, - output=output or "", - ) + payload = { + "evaluation_id": evaluation_id, + "input": input_text, + "output": output or "", + "messages": api_messages, + } + + # Convert tools if provided + api_tools = self._convert_tools_to_api_format(available_tools) + if api_tools: + payload["available_tools"] = api_tools + + url = f"{self.qualifire_api_base}/api/evaluation/invoke" else: - # Use evaluate with individual checks - kwargs = self._build_evaluate_kwargs( - qualifire_messages=qualifire_messages, + # Use evaluate endpoint with individual checks + api_tools = self._convert_tools_to_api_format(available_tools) + payload = self._build_evaluate_payload( + api_messages=api_messages, output=output, assertions=assertions, - available_tools=available_tools, + available_tools=api_tools, ) - result = client.evaluate(**kwargs) + url = f"{self.qualifire_api_base}/api/evaluation/evaluate" - # Convert result to dict for logging + verbose_proxy_logger.debug( + f"Qualifire Guardrail: Making request to {url}" + ) + + # Make the API request + response = await self.async_handler.post( + url=url, + headers=headers, + json=payload, + ) + response.raise_for_status() + result = response.json() + + # Extract response info for logging qualifire_response = { - "score": getattr(result, "score", None), - "status": getattr(result, "status", None), + "score": result.get("score"), + "status": result.get("status"), } verbose_proxy_logger.debug( diff --git a/litellm/proxy/hooks/litellm_skills/main.py b/litellm/proxy/hooks/litellm_skills/main.py index 26d4cbe1de7..c2ad1e29447 100644 --- a/litellm/proxy/hooks/litellm_skills/main.py +++ b/litellm/proxy/hooks/litellm_skills/main.py @@ -336,8 +336,8 @@ class SkillsInjectionHook(CustomLogger): ) # Check if code execution is enabled for this request - litellm_metadata = request_data.get("litellm_metadata", {}) - metadata = request_data.get("metadata", {}) + litellm_metadata = request_data.get("litellm_metadata") or {} + metadata = request_data.get("metadata") or {} code_exec_enabled = ( litellm_metadata.get("_litellm_code_execution_enabled") or diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index c416527990e..4d17cca22ad 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -167,7 +167,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.token_increment_script = None self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) - + # Batch rate limiter (lazy loaded) self._batch_rate_limiter: Optional[Any] = None @@ -1013,7 +1013,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) # Fail safe: enforce limits if we can't check return True - + def get_rate_limiter_for_call_type(self, call_type: str) -> Optional[Any]: """Get the rate limiter for the call type.""" if call_type == "acreate_batch": @@ -1095,9 +1095,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): now = self._get_current_time().timestamp() reset_time = now + self.window_size - reset_time_formatted = datetime.fromtimestamp( - reset_time - ).strftime("%Y-%m-%d %H:%M:%S UTC") + reset_time_formatted = datetime.fromtimestamp(reset_time).strftime( + "%Y-%m-%d %H:%M:%S UTC" + ) remaining_display = max(0, status["limit_remaining"]) rate_limit_type = status["rate_limit_type"] @@ -1137,7 +1137,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Check if the call type has a specific rate limiter # eg. for Batch APIs we need to use the batch rate limiter to read the input file and count the tokens and requests ######################################################### - call_type_specific_rate_limiter = self.get_rate_limiter_for_call_type(call_type=call_type) + call_type_specific_rate_limiter = self.get_rate_limiter_for_call_type( + call_type=call_type + ) if call_type_specific_rate_limiter: return await call_type_specific_rate_limiter.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -1233,26 +1235,58 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return pipeline_operations - def _get_total_tokens_from_usage(self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"]) -> int: - # Get total tokens from response + def _get_total_tokens_from_usage( + self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"] + ) -> int: + """ + Get total tokens from response usage for rate limiting. + + For 'input' and 'total' rate limit types, cached tokens are excluded + because providers like AWS Bedrock don't count cached tokens toward + rate limits. This aligns LiteLLM's TPM calculation with provider behavior. + """ total_tokens = 0 - # spot fix for /responses api + cached_tokens = 0 + if usage: if isinstance(usage, Usage): if rate_limit_type == "output": - total_tokens = usage.completion_tokens + total_tokens = usage.completion_tokens or 0 elif rate_limit_type == "input": - total_tokens = usage.prompt_tokens + total_tokens = usage.prompt_tokens or 0 elif rate_limit_type == "total": - total_tokens = usage.total_tokens + total_tokens = usage.total_tokens or 0 + + # Get cached tokens to exclude from input/total + if rate_limit_type in ("input", "total"): + if ( + hasattr(usage, "prompt_tokens_details") + and usage.prompt_tokens_details is not None + ): + cached_tokens = ( + getattr(usage.prompt_tokens_details, "cached_tokens", 0) + or 0 + ) + elif isinstance(usage, dict): - # Responses API usage comes as a dict in ResponsesAPIResponse + # Responses API usage comes as a dict if rate_limit_type == "output": - total_tokens = usage.get("completion_tokens", 0) + total_tokens = usage.get("completion_tokens", 0) or 0 elif rate_limit_type == "input": - total_tokens = usage.get("prompt_tokens", 0) + total_tokens = usage.get("prompt_tokens", 0) or 0 elif rate_limit_type == "total": - total_tokens = usage.get("total_tokens", 0) + total_tokens = usage.get("total_tokens", 0) or 0 + + # Get cached tokens from dict + if rate_limit_type in ("input", "total"): + prompt_details = usage.get("prompt_tokens_details") or {} + if isinstance(prompt_details, dict): + cached_tokens = prompt_details.get("cached_tokens", 0) or 0 + + # Subtract cached tokens for input/total (providers don't count them) + if cached_tokens > 0: + total_tokens = max(0, total_tokens - cached_tokens) + return total_tokens async def _execute_token_increment_script( @@ -1336,6 +1370,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def get_rate_limit_type(self) -> Literal["output", "input", "total"]: from litellm.proxy.proxy_server import general_settings + specified_rate_limit_type = general_settings.get( "token_rate_limit_type", "total" ) @@ -1381,9 +1416,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): user_api_key_organization_id = standard_logging_metadata.get( "user_api_key_org_id" ) - user_api_key_end_user_id = kwargs.get("user") or standard_logging_metadata.get( - "user_api_key_end_user_id" - ) + user_api_key_end_user_id = kwargs.get( + "user" + ) or standard_logging_metadata.get("user_api_key_end_user_id") model_group = get_model_group_from_litellm_kwargs(kwargs) # Get total tokens from response @@ -1393,7 +1428,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): response_obj, BaseLiteLLMOpenAIResponseObject ): _usage = getattr(response_obj, "usage", None) - total_tokens = self._get_total_tokens_from_usage(usage=_usage, rate_limit_type=rate_limit_type) + total_tokens = self._get_total_tokens_from_usage( + usage=_usage, rate_limit_type=rate_limit_type + ) # Create pipeline operations for TPM increments pipeline_operations: List[RedisPipelineIncrementOperation] = [] @@ -1518,9 +1555,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): from litellm.types.caching import RedisPipelineIncrementOperation try: - litellm_parent_otel_span: Union[ - Span, None - ] = _get_parent_otel_span_from_kwargs(kwargs) + litellm_parent_otel_span: Union[Span, None] = ( + _get_parent_otel_span_from_kwargs(kwargs) + ) # Get metadata from standard_logging_object - this correctly handles both # 'metadata' and 'litellm_metadata' fields from litellm_params standard_logging_object = kwargs.get("standard_logging_object") or {} @@ -1555,7 +1592,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): f"Error in rate limit failure event: {str(e)}" ) - async def async_post_call_success_hook( self, data: dict, user_api_key_dict: UserAPIKeyAuth, response ): diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 5b5723efc3d..3f844f21eb0 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -161,7 +161,6 @@ class KeyAndTeamLoggingSettings: @staticmethod def get_team_dynamic_logging_settings(user_api_key_dict: UserAPIKeyAuth): - if ( user_api_key_dict.team_metadata is not None and "logging" in user_api_key_dict.team_metadata @@ -174,12 +173,12 @@ def _get_dynamic_logging_metadata( user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig ) -> Optional[TeamCallbackMetadata]: callback_settings_obj: Optional[TeamCallbackMetadata] = None - key_dynamic_logging_settings: Optional[dict] = ( - KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) - ) - team_dynamic_logging_settings: Optional[dict] = ( - KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) - ) + key_dynamic_logging_settings: Optional[ + dict + ] = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) + team_dynamic_logging_settings: Optional[ + dict + ] = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) ######################################################################################### # Key-based callbacks ######################################################################################### @@ -462,7 +461,6 @@ class LiteLLMProxyRequestSetup: team_id=user_api_key_dict.team_id, ) # handles aliases, wildcards, etc. ): - _headers = LiteLLMProxyRequestSetup.add_headers_to_llm_call( headers, user_api_key_dict ) @@ -663,11 +661,11 @@ class LiteLLMProxyRequestSetup: ## KEY-LEVEL SPEND LOGS / TAGS if "tags" in key_metadata and key_metadata["tags"] is not None: - data[_metadata_variable_name]["tags"] = ( - LiteLLMProxyRequestSetup._merge_tags( - request_tags=data[_metadata_variable_name].get("tags"), - tags_to_add=key_metadata["tags"], - ) + data[_metadata_variable_name][ + "tags" + ] = LiteLLMProxyRequestSetup._merge_tags( + request_tags=data[_metadata_variable_name].get("tags"), + tags_to_add=key_metadata["tags"], ) if "disable_global_guardrails" in key_metadata and isinstance( key_metadata["disable_global_guardrails"], bool @@ -815,11 +813,14 @@ async def add_litellm_data_to_request( # noqa: PLR0915 # Init - Proxy Server Request # we do this as soon as entering so we track the original request ########################################################## + # Track arrival time for queue time metric + arrival_time = time.time() data["proxy_server_request"] = { "url": str(request.url), "method": request.method, "headers": _headers, "body": copy.copy(data), # use copy instead of deepcopy + "arrival_time": arrival_time, # Track when request arrived at proxy } safe_add_api_version_from_query_params(data, request) @@ -930,9 +931,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data[_metadata_variable_name]["litellm_api_version"] = version if general_settings is not None: - data[_metadata_variable_name]["global_max_parallel_requests"] = ( - general_settings.get("global_max_parallel_requests", None) - ) + data[_metadata_variable_name][ + "global_max_parallel_requests" + ] = general_settings.get("global_max_parallel_requests", None) ### KEY-LEVEL Controls key_metadata = user_api_key_dict.metadata diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 47793c8fc8e..d8816df010a 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -36,12 +36,19 @@ from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import LITELLM_PROXY_ADMIN_NAME from litellm.proxy._experimental.mcp_server.utils import ( + get_server_prefix, validate_and_normalize_mcp_server_payload, ) router = APIRouter(prefix="/v1/mcp", tags=["mcp"]) + MCP_AVAILABLE: bool = True + TEMPORARY_MCP_SERVER_TTL_SECONDS = 300 +DEFAULT_MCP_REGISTRY_VERSION = "1.0.0" +LITELLM_MCP_SERVER_NAME = "litellm-mcp-server" +LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM" + try: importlib.import_module("mcp") except ImportError as e: @@ -57,6 +64,7 @@ if MCP_AVAILABLE: update_mcp_server, ) from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, authorize_with_server, exchange_token_with_server, register_client_with_server, @@ -89,6 +97,66 @@ if MCP_AVAILABLE: server: MCPServer expires_at: datetime + def _is_public_registry_enabled() -> bool: + from litellm.proxy.proxy_server import ( + general_settings as proxy_general_settings, + ) + + return bool(proxy_general_settings.get("enable_mcp_registry")) + + def _build_registry_remote_url(base_url: str, path: str) -> str: + normalized_base = base_url.rstrip("/") + normalized_path = path if path.startswith("/") else f"/{path}" + return f"{normalized_base}{normalized_path}" + + def _build_mcp_registry_server_name(server: MCPServer) -> str: + if server.alias: + return server.alias + if server.server_name: + return server.server_name + return server.server_id + + def _build_mcp_registry_entry_for_server( + server: MCPServer, base_url: str + ) -> Dict[str, Any]: + server_name = _build_mcp_registry_server_name(server) + title = server_name + description = server_name + version = DEFAULT_MCP_REGISTRY_VERSION + + server_prefix = get_server_prefix(server) + if not server_prefix: + raise ValueError("MCP server prefix is missing") + remote_url = _build_registry_remote_url(base_url, f"/{server_prefix}/mcp") + + return { + "name": server_name, + "title": title, + "description": description, + "version": version, + "remotes": [ + { + "type": "streamable-http", + "url": remote_url, + } + ], + } + + def _build_builtin_registry_entry(base_url: str) -> Dict[str, Any]: + remote_url = _build_registry_remote_url(base_url, "/mcp") + return { + "name": LITELLM_MCP_SERVER_NAME, + "title": LITELLM_MCP_SERVER_NAME, + "description": LITELLM_MCP_SERVER_DESCRIPTION, + "version": DEFAULT_MCP_REGISTRY_VERSION, + "remotes": [ + { + "type": "streamable-http", + "url": remote_url, + } + ], + } + _temporary_mcp_servers: Dict[str, _TemporaryMCPServerEntry] = {} def _prune_expired_temporary_mcp_servers() -> None: @@ -302,15 +370,42 @@ if MCP_AVAILABLE: access_groups_list = sorted(list(access_groups)) return {"access_groups": access_groups_list} + @router.get( + "/registry.json", + tags=["mcp"], + description="MCP registry endpoint. Spec: https://github.com/modelcontextprotocol/registry", + ) + async def get_mcp_registry(request: Request): + if not _is_public_registry_enabled(): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="MCP registry is not enabled", + ) + + base_url = get_request_base_url(request) + registry_servers: List[Dict[str, Any]] = [] + registry_servers.append({"server": _build_builtin_registry_entry(base_url)}) + + registered_servers = list(global_mcp_server_manager.get_registry().values()) + registered_servers.sort(key=_build_mcp_registry_server_name) + + for server in registered_servers: + try: + entry = _build_mcp_registry_entry_for_server(server, base_url) + except Exception as e: + verbose_proxy_logger.debug( + f"Skipping MCP server {getattr(server, 'server_id', 'unknown')} in registry: {e}" + ) + continue + registry_servers.append({"server": entry}) + + return {"servers": registry_servers} + ## FastAPI Routes def _get_user_mcp_management_mode() -> UserMCPManagementMode: - proxy_general_settings: dict = {} - try: - from litellm.proxy.proxy_server import ( - general_settings as proxy_general_settings, - ) - except Exception: - pass + from litellm.proxy.proxy_server import ( + general_settings as proxy_general_settings, + ) mode = proxy_general_settings.get("user_mcp_management_mode") if mode == "view_all": diff --git a/litellm/proxy/management_endpoints/router_settings_endpoints.py b/litellm/proxy/management_endpoints/router_settings_endpoints.py index 167160c72d1..4d4c41a3dc0 100644 --- a/litellm/proxy/management_endpoints/router_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/router_settings_endpoints.py @@ -4,6 +4,7 @@ ROUTER SETTINGS MANAGEMENT Endpoints for accessing router configuration and metadata GET /router/settings - Get router configuration including available routing strategies +GET /router/fields - Get router settings field definitions without values (for UI rendering) """ import inspect @@ -37,6 +38,15 @@ class RouterSettingsResponse(BaseModel): ) +class RouterFieldsResponse(BaseModel): + fields: List[RouterSettingsField] = Field( + description="List of all configurable router settings with metadata (without field values)" + ) + routing_strategy_descriptions: Dict[str, str] = Field( + description="Descriptions for each routing strategy option" + ) + + def _get_routing_strategies_from_router_class() -> List[str]: """ Dynamically extract routing strategies from the Router class __init__ method. @@ -120,3 +130,53 @@ async def get_router_settings( ) raise + +@router.get( + "/router/fields", + tags=["Router Settings"], + dependencies=[Depends(user_api_key_auth)], + response_model=RouterFieldsResponse, +) +async def get_router_fields( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get router settings field definitions without values. + + Returns only the field metadata (type, description, default, options) without + populating field_value. This is useful for UI components that need to know + what fields to render, but will get the actual values from a different endpoint. + + Returns: + - fields: List of all configurable router settings with their metadata (type, description, default, options) + The routing_strategy field includes available options extracted from the Router class + Note: field_value will be None for all fields + - routing_strategy_descriptions: Descriptions for each routing strategy option + """ + try: + # Get available routing strategies dynamically from Router class + available_routing_strategies = _get_routing_strategies_from_router_class() + + # Get router settings fields from types file + router_fields = [field.model_copy(deep=True) for field in ROUTER_SETTINGS_FIELDS] + + # Populate routing_strategy field with available options + for field in router_fields: + if field.field_name == "routing_strategy": + field.options = available_routing_strategies + break + + # Ensure field_value is None for all fields (don't populate values) + for field in router_fields: + field.field_value = None + + return RouterFieldsResponse( + fields=router_fields, + routing_strategy_descriptions=ROUTING_STRATEGY_DESCRIPTIONS, + ) + except Exception as e: + verbose_proxy_logger.error( + f"Error fetching router fields: {str(e)}" + ) + raise + diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 84550092d2e..d9798dae690 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -776,6 +776,7 @@ async def handle_bedrock_count_tokens( - /v1/messages/count_tokens - /v1/messages/count-tokens """ + from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler from litellm.proxy.proxy_server import llm_router @@ -822,6 +823,12 @@ async def handle_bedrock_count_tokens( return result + except BedrockError as e: + # Convert BedrockError to HTTPException for FastAPI + verbose_proxy_logger.error(f"BedrockError in handle_bedrock_count_tokens: {str(e)}") + raise HTTPException( + status_code=e.status_code, detail={"error": e.message} + ) except HTTPException: # Re-raise HTTP exceptions as-is raise diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 06525e39133..83d4ab5657e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -229,7 +229,7 @@ from litellm.proxy.batches_endpoints.endpoints import router as batches_router from litellm.proxy.caching_routes import router as caching_router from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, - create_streaming_response, + create_response, ) from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_proxy from litellm.proxy.common_utils.debug_utils import init_verbose_loggers @@ -533,9 +533,9 @@ except ImportError: server_root_path = os.getenv("SERVER_ROOT_PATH", "") _license_check = LicenseCheck() premium_user: bool = _license_check.is_premium() -premium_user_data: Optional["EnterpriseLicenseData"] = ( - _license_check.airgapped_license_data -) +premium_user_data: Optional[ + "EnterpriseLicenseData" +] = _license_check.airgapped_license_data global_max_parallel_request_retries_env: Optional[str] = os.getenv( "LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES" ) @@ -658,7 +658,7 @@ async def _initialize_shared_aiohttp_session(): @asynccontextmanager -async def proxy_startup_event(app: FastAPI): +async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 global prisma_client, master_key, use_background_health_checks, llm_router, llm_model_list, general_settings, proxy_budget_rescheduler_min_time, proxy_budget_rescheduler_max_time, litellm_proxy_admin_name, db_writer_client, store_model_in_db, premium_user, _license_check, proxy_batch_polling_interval, shared_aiohttp_session import json @@ -788,6 +788,17 @@ async def proxy_startup_event(app: FastAPI): except Exception as e: verbose_proxy_logger.error(f"Error closing shared aiohttp session: {e}") + # Shutdown event - stop RDS IAM token refresh background task + if ( + prisma_client is not None + and hasattr(prisma_client, "db") + and hasattr(prisma_client.db, "stop_token_refresh_task") + ): + try: + await prisma_client.db.stop_token_refresh_task() + except Exception as e: + verbose_proxy_logger.error(f"Error stopping token refresh task: {e}") + await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues] @@ -1083,9 +1094,7 @@ try: # In non-root Docker, we restructure in /var/lib/litellm/ui. try: _restructure_ui_html_files(ui_path) - verbose_proxy_logger.info( - f"Restructured UI directory: {ui_path}" - ) + verbose_proxy_logger.info(f"Restructured UI directory: {ui_path}") except PermissionError as e: verbose_proxy_logger.exception( f"Permission error while restructuring UI directory {ui_path}: {e}" @@ -1171,9 +1180,9 @@ master_key: Optional[str] = None config_agents: Optional[List[AgentConfig]] = None otel_logging = False prisma_client: Optional[PrismaClient] = None -shared_aiohttp_session: Optional["ClientSession"] = ( - None # Global shared session for connection reuse -) +shared_aiohttp_session: Optional[ + "ClientSession" +] = None # Global shared session for connection reuse user_api_key_cache = DualCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value ) @@ -1181,9 +1190,9 @@ model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter( dual_cache=user_api_key_cache ) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) -redis_usage_cache: Optional[RedisCache] = ( - None # redis cache used for tracking spend, tpm/rpm limits -) +redis_usage_cache: Optional[ + RedisCache +] = None # redis cache used for tracking spend, tpm/rpm limits polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache user_custom_auth = None @@ -1522,9 +1531,9 @@ async def update_cache( # noqa: PLR0915 _id = "team_id:{}".format(team_id) try: # Fetch the existing cost for the given user - existing_spend_obj: Optional[LiteLLM_TeamTable] = ( - await user_api_key_cache.async_get_cache(key=_id) - ) + existing_spend_obj: Optional[ + LiteLLM_TeamTable + ] = await user_api_key_cache.async_get_cache(key=_id) if existing_spend_obj is None: # do nothing if team not in api key cache return @@ -1876,7 +1885,6 @@ class ProxyConfig: "environment_variables" in config_to_save and config_to_save["environment_variables"] ): - # decrypt the environment_variables - in case a caller function has already encrypted the environment_variables decrypted_env_vars = self._decrypt_and_set_db_env_variables( environment_variables=config_to_save["environment_variables"], @@ -2794,21 +2802,21 @@ class ProxyConfig: verbose_proxy_logger.debug(f"_alerting_callbacks: {general_settings}") if _alerting_callbacks is None: return - + # Ensure proxy_logging_obj.alerting is set for all alerting types _alerting_value = general_settings.get("alerting", None) - verbose_proxy_logger.debug(f"_load_alerting_settings: Calling update_values with alerting={_alerting_value}") + verbose_proxy_logger.debug( + f"_load_alerting_settings: Calling update_values with alerting={_alerting_value}" + ) proxy_logging_obj.update_values( alerting=_alerting_value, alerting_threshold=general_settings.get("alerting_threshold", 600), alert_types=general_settings.get("alert_types", None), - alert_to_webhook_url=general_settings.get( - "alert_to_webhook_url", None - ), + alert_to_webhook_url=general_settings.get("alert_to_webhook_url", None), alerting_args=general_settings.get("alerting_args", None), redis_cache=redis_usage_cache, ) - + for _alert in _alerting_callbacks: if _alert == "slack": # [OLD] v0 implementation - already handled by update_values above @@ -3222,6 +3230,84 @@ class ProxyConfig: decrypted_variables[k] = decrypted_value return decrypted_variables + async def _get_hierarchical_router_settings( + self, + user_api_key_dict: Optional["UserAPIKeyAuth"], + prisma_client: Optional[PrismaClient], + ) -> Optional[dict]: + """ + Get router_settings in priority order: Key > Team > Global + + Returns: + dict: Combined router_settings, or None if no settings found + """ + if prisma_client is None: + return None + + import json + import yaml + + # 1. Try key-level router_settings + if user_api_key_dict is not None: + # Check if router_settings is available on the key object + key_router_settings_value = getattr(user_api_key_dict, "router_settings", None) + if key_router_settings_value is not None: + key_router_settings = None + if isinstance(key_router_settings_value, str): + try: + key_router_settings = yaml.safe_load(key_router_settings_value) + except (yaml.YAMLError, json.JSONDecodeError): + try: + key_router_settings = json.loads(key_router_settings_value) + except json.JSONDecodeError: + pass + elif isinstance(key_router_settings_value, dict): + key_router_settings = key_router_settings_value + + # If key has router_settings (non-empty dict), use it + if key_router_settings is not None and isinstance(key_router_settings, dict) and key_router_settings: + return key_router_settings + + # 2. Try team-level router_settings + if user_api_key_dict is not None and user_api_key_dict.team_id is not None: + try: + team_obj = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": user_api_key_dict.team_id} + ) + if team_obj is not None: + team_router_settings_value = getattr(team_obj, "router_settings", None) + if team_router_settings_value is not None: + team_router_settings = None + if isinstance(team_router_settings_value, str): + try: + team_router_settings = yaml.safe_load(team_router_settings_value) + except (yaml.YAMLError, json.JSONDecodeError): + try: + team_router_settings = json.loads(team_router_settings_value) + except json.JSONDecodeError: + pass + elif isinstance(team_router_settings_value, dict): + team_router_settings = team_router_settings_value + + # If team has router_settings (non-empty dict), use it + if team_router_settings is not None and isinstance(team_router_settings, dict) and team_router_settings: + return team_router_settings + except Exception: + # If team lookup fails, continue to global settings + pass + + # 3. Try global router_settings + try: + db_router_settings = await prisma_client.db.litellm_config.find_first( + where={"param_name": "router_settings"} + ) + if db_router_settings is not None and isinstance(db_router_settings.param_value, dict) and db_router_settings.param_value: + return db_router_settings.param_value + except Exception: + pass + + return None + async def _add_router_settings_from_db_config( self, config_data: dict, @@ -3279,7 +3365,7 @@ class ProxyConfig: proxy_logging_obj: ProxyLogging """ _general_settings = config_data.get("general_settings", {}) - + if _general_settings is not None and "alerting" in _general_settings: if ( general_settings is not None @@ -3294,7 +3380,8 @@ class ProxyConfig: _merged_alerting = list(_yaml_alerting.union(_db_alerting)) # Preserve order: YAML values first, then DB values _merged_alerting = list(general_settings["alerting"]) + [ - item for item in _general_settings["alerting"] + item + for item in _general_settings["alerting"] if item not in general_settings["alerting"] ] verbose_proxy_logger.debug( @@ -3402,8 +3489,8 @@ class ProxyConfig: def _deep_merge_dicts(dst: dict, src: dict) -> None: """ - Deep-merge src into dst, skipping None values from src. - On conflicts, src (DB) wins. + Deep-merge src into dst, skipping None values and empty lists from src. + On conflicts, src (DB) wins, but empty lists are treated as "no value" and don't overwrite. """ stack = [(dst, src)] while stack: @@ -3412,6 +3499,9 @@ class ProxyConfig: if v is None: # Preserve existing config when DB value is None (matches prior behavior) continue + # Skip empty lists - treat them as "no value" to preserve file config + if isinstance(v, list) and len(v) == 0: + continue if isinstance(v, dict) and isinstance(d.get(k), dict): stack.append((d[k], v)) else: @@ -3605,7 +3695,6 @@ class ProxyConfig: await self._init_vector_stores_in_db(prisma_client=prisma_client) if self._should_load_db_object(object_type="vector_store_indexes"): - await self._init_vector_store_indexes_in_db(prisma_client=prisma_client) if self._should_load_db_object(object_type="mcp"): @@ -3804,10 +3893,10 @@ class ProxyConfig: ) try: - guardrails_in_db: List[Guardrail] = ( - await GuardrailRegistry.get_all_guardrails_from_db( - prisma_client=prisma_client - ) + guardrails_in_db: List[ + Guardrail + ] = await GuardrailRegistry.get_all_guardrails_from_db( + prisma_client=prisma_client ) verbose_proxy_logger.debug( "guardrails from the DB %s", str(guardrails_in_db) @@ -3894,7 +3983,7 @@ class ProxyConfig: ) try: - await global_mcp_server_manager._add_mcp_servers_from_db_to_in_memory_registry() + await global_mcp_server_manager.reload_servers_from_database() except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db - {}".format( @@ -4031,6 +4120,23 @@ class ProxyConfig: return [] +async def _reload_mcp_servers_job(): + """Background job entrypoint for MCP registry refreshes.""" + if proxy_config._should_load_db_object(object_type="mcp") is False: + return + + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + await global_mcp_server_manager.reload_servers_from_database() + except Exception as e: + verbose_proxy_logger.exception( + "Failed to reload MCP servers from database: %s", str(e) + ) + + proxy_config = ProxyConfig() @@ -4134,9 +4240,9 @@ async def initialize( # noqa: PLR0915 user_api_base = api_base dynamic_config[user_model]["api_base"] = api_base if api_version: - os.environ["AZURE_API_VERSION"] = ( - api_version # set this for azure - litellm can read this from the env - ) + os.environ[ + "AZURE_API_VERSION" + ] = api_version # set this for azure - litellm can read this from the env if max_tokens: # model-specific param dynamic_config[user_model]["max_tokens"] = max_tokens if temperature: # model-specific param @@ -4568,6 +4674,18 @@ class ProxyStartupEvent: misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, ) await proxy_config.get_credentials(prisma_client=prisma_client) + + from litellm.proxy._experimental.mcp_server.utils import is_mcp_available + + if is_mcp_available(): + scheduler.add_job( + _reload_mcp_servers_job, + "interval", + seconds=30, + id="reload_mcp_servers_job", + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) await cls._initialize_slack_alerting_jobs( scheduler=scheduler, general_settings=general_settings, @@ -4654,10 +4772,14 @@ class ProxyStartupEvent: replace_existing=True, misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, ) - verbose_proxy_logger.info("Responses cost check job scheduled successfully") + verbose_proxy_logger.info( + "Responses cost check job scheduled successfully" + ) except Exception as e: - verbose_proxy_logger.debug(f"Failed to setup responses cost checking: {e}") + verbose_proxy_logger.debug( + f"Failed to setup responses cost checking: {e}" + ) verbose_proxy_logger.debug( "Checking responses cost for LiteLLM Managed Files is an Enterprise Feature. Skipping..." ) @@ -4681,8 +4803,9 @@ class ProxyStartupEvent: """ Initialize the spend tracking and other background jobs 1. CloudZero Background Job - 2. Prometheus Background Job - 3. Key Rotation Background Job + 2. Focus Background Job + 3. Prometheus Background Job + 4. Key Rotation Background Job Args: scheduler: The scheduler to add the background jobs to @@ -4691,11 +4814,17 @@ class ProxyStartupEvent: # CloudZero Background Job ######################################################## from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger + from litellm.integrations.focus.focus_logger import FocusLogger from litellm.proxy.spend_tracking.cloudzero_endpoints import is_cloudzero_setup if await is_cloudzero_setup(): await CloudZeroLogger.init_cloudzero_background_job(scheduler=scheduler) + ######################################################## + # Focus Background Job + ######################################################## + await FocusLogger.init_focus_export_background_job(scheduler=scheduler) + ######################################################## # Prometheus Background Job ######################################################## @@ -4826,6 +4955,14 @@ class ProxyStartupEvent: await prisma_client.connect() + ## Start RDS IAM token refresh background task if enabled ## + # This proactively refreshes IAM tokens before they expire, + # preventing the 15-minute connection failure bug (#16220) + if hasattr(prisma_client, "db") and hasattr( + prisma_client.db, "start_token_refresh_task" + ): + await prisma_client.db.start_token_refresh_task() + ## Add necessary views to proxy ## asyncio.create_task( prisma_client.check_view_exists() @@ -5937,7 +6074,6 @@ async def realtime_websocket_endpoint( ), user_api_key_dict=Depends(user_api_key_auth_websocket), ): - await websocket.accept() # Only use explicit parameters, not all query params @@ -6736,7 +6872,7 @@ async def run_thread( if ( "stream" in data and data["stream"] is True ): # use generate_responses to stream responses - return await create_streaming_response( + return await create_response( generator=async_assistants_data_generator( user_api_key_dict=user_api_key_dict, response=response, @@ -9525,9 +9661,9 @@ async def get_config_list( hasattr(sub_field_info, "description") and sub_field_info.description is not None ): - nested_fields[idx].field_description = ( - sub_field_info.description - ) + nested_fields[ + idx + ].field_description = sub_field_info.description idx += 1 _stored_in_db = None @@ -9762,6 +9898,18 @@ async def get_config(): # noqa: PLR0915 _failure_callbacks = _litellm_settings.get("failure_callback", []) _success_and_failure_callbacks = _litellm_settings.get("callbacks", []) + # Normalize string callbacks to lists + def normalize_callback(callback): + if isinstance(callback, str): + return [callback] + elif callback is None: + return [] + return callback + + _success_callbacks = normalize_callback(_success_callbacks) + _failure_callbacks = normalize_callback(_failure_callbacks) + _success_and_failure_callbacks = normalize_callback(_success_and_failure_callbacks) + _data_to_return = [] """ [ diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d1a78534dae..171898b1631 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -131,6 +131,7 @@ else: unified_guardrail = UnifiedLLMGuardrails() +_anthropic_async_clients = {} def print_verbose(print_statement): """ @@ -961,6 +962,7 @@ class ProxyLogging: Updated data dictionary if guardrail passes, None if guardrail should be skipped """ from litellm.types.guardrails import GuardrailEventHooks + from litellm.integrations.prometheus import PrometheusLogger # Determine the event type based on call type event_type = GuardrailEventHooks.pre_call @@ -973,30 +975,62 @@ class ProxyLogging: guardrail_name = callback.guardrail_name - # Check if load balancing should be used - if guardrail_name and self._should_use_guardrail_load_balancing(guardrail_name): - response = await self._execute_guardrail_with_load_balancing( - guardrail_name=guardrail_name, - hook_type="pre_call", - data=data, - user_api_key_dict=user_api_key_dict, - call_type=call_type, - ) - else: - # Single guardrail - execute directly - response = await self._execute_guardrail_hook( - callback=callback, - hook_type="pre_call", - data=data, - user_api_key_dict=user_api_key_dict, - call_type=call_type, - ) + # Track timing and errors for prometheus metrics + # Use time.perf_counter() for more accurate duration measurements + guardrail_start_time = time.perf_counter() + status = "success" + error_type = None - # Process the response if one was returned - if response is not None: - data = await self.process_pre_call_hook_response( - response=response, data=data, call_type=call_type - ) + try: + # Check if load balancing should be used + if guardrail_name and self._should_use_guardrail_load_balancing(guardrail_name): + response = await self._execute_guardrail_with_load_balancing( + guardrail_name=guardrail_name, + hook_type="pre_call", + data=data, + user_api_key_dict=user_api_key_dict, + call_type=call_type, + ) + else: + # Single guardrail - execute directly + response = await self._execute_guardrail_hook( + callback=callback, + hook_type="pre_call", + data=data, + user_api_key_dict=user_api_key_dict, + call_type=call_type, + ) + + # Process the response if one was returned + if response is not None: + data = await self.process_pre_call_hook_response( + response=response, data=data, call_type=call_type + ) + + except Exception as e: + status = "error" + error_type = type(e).__name__ + # Re-raise the exception to maintain existing behavior + raise + finally: + # Record prometheus metrics + guardrail_end_time = time.perf_counter() + latency_seconds = guardrail_end_time - guardrail_start_time + + # Get guardrail name for metrics (fallback if not set) + metrics_guardrail_name = guardrail_name or getattr(callback, "guardrail_name", callback.__class__.__name__) or "unknown" + + # Find PrometheusLogger in callbacks and record metrics + for prom_callback in litellm.callbacks: + if isinstance(prom_callback, PrometheusLogger): + prom_callback._record_guardrail_metrics( + guardrail_name=metrics_guardrail_name, + latency_seconds=latency_seconds, + status=status, + error_type=error_type, + hook_type="pre_call", + ) + break return data @@ -1195,7 +1229,7 @@ class ProxyLogging: and _callback.__class__.async_pre_call_hook != CustomLogger.async_pre_call_hook ): - if call_type == "mcp_call" and user_api_key_dict is None: + if call_type == "call_mcp_tool" and user_api_key_dict is None: continue response = await _callback.async_pre_call_hook( @@ -4254,11 +4288,16 @@ async def count_tokens_with_anthropic_api( if anthropic_api_key and messages: # Call Anthropic API directly for more accurate token counting - client = anthropic.Anthropic(api_key=anthropic_api_key) + + # Use cached client if available to avoid socket exhaustion + if anthropic_api_key not in _anthropic_async_clients: + _anthropic_async_clients[anthropic_api_key] = anthropic.AsyncAnthropic(api_key=anthropic_api_key) + + client = _anthropic_async_clients[anthropic_api_key] # Call with explicit parameters to satisfy type checking # Type ignore for now since messages come from generic dict input - response = client.beta.messages.count_tokens( + response = await client.beta.messages.count_tokens( model=model_to_use, messages=messages, # type: ignore betas=["token-counting-2024-11-01"], diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 8177b177fe6..07fc3cb02cc 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -577,7 +577,13 @@ def responses( api_base=litellm_params.api_base, api_key=litellm_params.api_key, ) - + + # Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True) + if dynamic_api_key is not None: + litellm_params.api_key = dynamic_api_key + if dynamic_api_base is not None: + litellm_params.api_base = dynamic_api_base + ######################################################### # Update input with provider-specific file IDs if managed files are used ######################################################### @@ -1483,6 +1489,12 @@ def compact_responses( api_key=litellm_params.api_key, ) + # Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True) + if dynamic_api_key is not None: + litellm_params.api_key = dynamic_api_key + if dynamic_api_base is not None: + litellm_params.api_base = dynamic_api_base + if custom_llm_provider is None: raise ValueError("custom_llm_provider is required but passed as None") diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index a92b5d25a37..7667d1bad84 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -443,11 +443,18 @@ class ResponseAPILoggingUtils: completion_tokens=0, total_tokens=0, ) - response_api_usage: ResponseAPIUsage = ( - ResponseAPIUsage(**usage_input) - if isinstance(usage_input, dict) - else usage_input - ) + response_api_usage: ResponseAPIUsage + if isinstance(usage_input, dict): + total_tokens = usage_input.get("total_tokens") + if total_tokens is None: + input_tokens = usage_input.get("input_tokens") + output_tokens = usage_input.get("output_tokens") + if input_tokens is not None and output_tokens is not None: + total_tokens = input_tokens + output_tokens + usage_input["total_tokens"] = total_tokens + response_api_usage = ResponseAPIUsage(**usage_input) + else: + response_api_usage = usage_input prompt_tokens: int = response_api_usage.input_tokens or 0 completion_tokens: int = response_api_usage.output_tokens or 0 prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None diff --git a/litellm/router.py b/litellm/router.py index 98ccf41c96d..364e6719300 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -255,6 +255,7 @@ class Router: ] = {}, enable_pre_call_checks: bool = False, enable_tag_filtering: bool = False, + tag_filtering_match_any: bool = True, retry_after: int = 0, # min time to wait before retrying a failed request retry_policy: Optional[ Union[RetryPolicy, dict] @@ -363,6 +364,7 @@ class Router: self.debug_level = debug_level self.enable_pre_call_checks = enable_pre_call_checks self.enable_tag_filtering = enable_tag_filtering + self.tag_filtering_match_any = tag_filtering_match_any from litellm._service_logger import ServiceLogging self.service_logger_obj: ServiceLogging = ServiceLogging() @@ -384,9 +386,9 @@ class Router: ) # names of models under litellm_params. ex. azure/chatgpt-v-2 self.deployment_latency_map = {} ### CACHING ### - cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = ( - "local" # default to an in-memory cache - ) + cache_type: Literal[ + "local", "redis", "redis-semantic", "s3", "disk" + ] = "local" # default to an in-memory cache redis_cache = None cache_config: Dict[str, Any] = {} @@ -428,9 +430,9 @@ class Router: self.default_max_parallel_requests = default_max_parallel_requests self.provider_default_deployment_ids: List[str] = [] self.pattern_router = PatternMatchRouter() - self.team_pattern_routers: Dict[str, PatternMatchRouter] = ( - {} - ) # {"TEAM_ID": PatternMatchRouter} + self.team_pattern_routers: Dict[ + str, PatternMatchRouter + ] = {} # {"TEAM_ID": PatternMatchRouter} self.auto_routers: Dict[str, "AutoRouter"] = {} # Initialize model_group_alias early since it's used in set_model_list @@ -611,9 +613,9 @@ class Router: ) ) - self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = ( - model_group_retry_policy - ) + self.model_group_retry_policy: Optional[ + Dict[str, RetryPolicy] + ] = model_group_retry_policy self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None if allowed_fails_policy is not None: @@ -720,7 +722,10 @@ class Router: valid_strategy_strings = ["simple-shuffle"] + [s.value for s in RoutingStrategy] if routing_strategy is not None: - is_valid_string = isinstance(routing_strategy, str) and routing_strategy in valid_strategy_strings + is_valid_string = ( + isinstance(routing_strategy, str) + and routing_strategy in valid_strategy_strings + ) is_valid_enum = isinstance(routing_strategy, RoutingStrategy) if not is_valid_string and not is_valid_enum: raise ValueError( @@ -1069,7 +1074,7 @@ class Router: self.delete_container = self.factory_function( delete_container, call_type="delete_container" ) - + # Auto-register JSON-generated container file endpoints for name, func in container_file_endpoints.items(): setattr(self, name, self.factory_function(func, call_type=name)) # type: ignore[arg-type] @@ -1498,10 +1503,7 @@ class Router: async def _acompletion( self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ - ModelResponse, - CustomStreamWrapper, - ]: + ) -> Union[ModelResponse, CustomStreamWrapper,]: """ - Get an available deployment - call it with a semaphore over the call @@ -3019,7 +3021,9 @@ class Router: kwargs["original_generic_function"] = original_function kwargs["original_function"] = self._aguardrail_helper self._update_kwargs_before_fallbacks( - model=guardrail_name, kwargs=kwargs, metadata_variable_name="litellm_metadata" + model=guardrail_name, + kwargs=kwargs, + metadata_variable_name="litellm_metadata", ) verbose_router_logger.debug( f"Inside aguardrail() - guardrail_name: {guardrail_name}; kwargs: {kwargs}" @@ -3312,8 +3316,7 @@ class Router: kwargs["model"] = model kwargs["input"] = input kwargs["original_function"] = self._embedding - kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries) - kwargs.setdefault("metadata", {}).update({"model_group": model}) + self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response = self.function_with_fallbacks(**kwargs) return response except Exception as e: @@ -3615,9 +3618,9 @@ class Router: healthy_deployments=healthy_deployments, responses=responses ) returned_response = cast(OpenAIFileObject, responses[0]) - returned_response._hidden_params["model_file_id_mapping"] = ( - model_file_id_mapping - ) + returned_response._hidden_params[ + "model_file_id_mapping" + ] = model_file_id_mapping return returned_response except Exception as e: verbose_router_logger.exception( @@ -4364,11 +4367,11 @@ class Router: if isinstance(e, litellm.ContextWindowExceededError): if context_window_fallbacks is not None: - context_window_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=context_window_fallbacks, - model_group=model_group, - ) + context_window_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=context_window_fallbacks, + model_group=model_group, ) if context_window_fallback_model_group is None: raise original_exception @@ -4400,11 +4403,11 @@ class Router: e.message += "\n{}".format(error_message) elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: - content_policy_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=content_policy_fallbacks, - model_group=model_group, - ) + content_policy_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=content_policy_fallbacks, + model_group=model_group, ) if content_policy_fallback_model_group is None: raise original_exception @@ -4483,9 +4486,21 @@ class Router: if hasattr(original_exception, "message"): # add the available fallbacks to the exception - original_exception.message += ". Received Model Group={}\nAvailable Model Group Fallbacks={}".format( # type: ignore - model_group, - fallback_model_group, + deployment_info = "" + if kwargs is not None: + metadata = kwargs.get('metadata', {}) + if metadata and 'deployment' in metadata: + deployment_info = f"\nUsed Deployment: {metadata['deployment']}" + if 'model_info' in metadata: + model_info = metadata['model_info'] + if isinstance(model_info, dict): + deployment_info += f"\nDeployment ID: {model_info.get('id', 'unknown')}" + + original_exception.message += ( # type: ignore + f". Received Model Group={model_group}" + f"\nAvailable Model Group Fallbacks={fallback_model_group}" + f"{deployment_info}" + f"\n\n💡 Tip: If using wildcard patterns (e.g., 'openai/*'), ensure all matching deployments have credentials with access to this model." ) if len(fallback_failure_exception_str) > 0: original_exception.message += ( # type: ignore @@ -5667,26 +5682,26 @@ class Router: """ from litellm.router_strategy.auto_router.auto_router import AutoRouter - auto_router_config_path: Optional[str] = ( - deployment.litellm_params.auto_router_config_path - ) + auto_router_config_path: Optional[ + str + ] = deployment.litellm_params.auto_router_config_path auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config if auto_router_config_path is None and auto_router_config is None: raise ValueError( "auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params" ) - default_model: Optional[str] = ( - deployment.litellm_params.auto_router_default_model - ) + default_model: Optional[ + str + ] = deployment.litellm_params.auto_router_default_model if default_model is None: raise ValueError( "auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params" ) - embedding_model: Optional[str] = ( - deployment.litellm_params.auto_router_embedding_model - ) + embedding_model: Optional[ + str + ] = deployment.litellm_params.auto_router_embedding_model if embedding_model is None: raise ValueError( "auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params" @@ -6233,9 +6248,9 @@ class Router: # Add custom_llm_provider if deployment.litellm_params.custom_llm_provider: - credentials["custom_llm_provider"] = ( - deployment.litellm_params.custom_llm_provider - ) + credentials[ + "custom_llm_provider" + ] = deployment.litellm_params.custom_llm_provider elif "/" in deployment.litellm_params.model: # Extract provider from "provider/model" format credentials["custom_llm_provider"] = deployment.litellm_params.model.split( @@ -6929,42 +6944,44 @@ class Router: """ return candidate_id in self.model_id_to_deployment_index_map - def resolve_model_name_from_model_id(self, model_id: Optional[str]) -> Optional[str]: + def resolve_model_name_from_model_id( + self, model_id: Optional[str] + ) -> Optional[str]: """ Resolve model_name from model_id. - + This method attempts to find the correct model_name to use with the router so that litellm_params can be automatically injected from the model config. - + Strategy: 1. First, check if model_id directly matches a model_name or deployment ID 2. If not, search through router's model_list to find a match by litellm_params.model 3. Return the model_name if found, None otherwise - + Args: model_id: The model_id extracted from decoded video_id (could be model_name or litellm_params.model value) - + Returns: model_name if found, None otherwise. If None, the request will fall through to normal flow using environment variables. """ if not model_id: return None - + # Strategy 1: Check if model_id directly matches a model_name or deployment ID if model_id in self.model_names or self.has_model_id(model_id): return model_id - + # Strategy 2: Search through router's model_list to find by litellm_params.model all_models = self.get_model_list(model_name=None) if not all_models: return None - + for deployment in all_models: litellm_params = deployment.get("litellm_params", {}) actual_model = litellm_params.get("model") - + # Match by exact match or by checking if actual_model ends with /model_id or :model_id # e.g., model_id="veo-2.0-generate-001" matches actual_model="vertex_ai/veo-2.0-generate-001" matches = ( @@ -6972,12 +6989,12 @@ class Router: or (actual_model and actual_model.endswith(f"/{model_id}")) or (actual_model and actual_model.endswith(f":{model_id}")) ) - + if matches: model_name = deployment.get("model_name") if model_name: return model_name - + # No match found return None @@ -7662,6 +7679,10 @@ class Router: ) if pattern_deployments: + verbose_router_logger.debug( + f"Pattern match for model='{model}': Found {len(pattern_deployments)} deployments. " + f"Deployment IDs: {[d.get('model_info', {}).get('id', 'unknown') for d in pattern_deployments]}" + ) return model, pattern_deployments if ( @@ -7767,14 +7788,18 @@ class Router: request_kwargs=request_kwargs, ) - verbose_router_logger.debug(f"healthy_deployments after team filter: {healthy_deployments}") + verbose_router_logger.debug( + f"healthy_deployments after team filter: {healthy_deployments}" + ) healthy_deployments = filter_web_search_deployments( healthy_deployments=healthy_deployments, request_kwargs=request_kwargs, ) - verbose_router_logger.debug(f"healthy_deployments after web search filter: {healthy_deployments}") + verbose_router_logger.debug( + f"healthy_deployments after web search filter: {healthy_deployments}" + ) if isinstance(healthy_deployments, dict): return healthy_deployments diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index b25c20eb281..e960e00a68f 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -20,17 +20,28 @@ else: def is_valid_deployment_tag( - deployment_tags: List[str], request_tags: List[str] + deployment_tags: List[str], request_tags: List[str], match_any: bool = True ) -> bool: """ - Check if a tag is valid + Check if a tag is valid, the matching can be either any or all based on `match_any` flag """ + if not request_tags: + return False - if any(tag in deployment_tags for tag in request_tags): + dep_set = set(deployment_tags) + req_set = set(request_tags) + + if match_any: + is_valid_deployment = bool(dep_set & req_set) + else: + is_valid_deployment = req_set.issubset(dep_set) + + if is_valid_deployment: verbose_logger.debug( - "adding deployment with tags: %s, request tags: %s", + "adding deployment with tags: %s, request tags: %s for match_any=%s", deployment_tags, request_tags, + match_any, ) return True return False @@ -68,6 +79,7 @@ async def get_deployments_for_tag( if metadata_variable_name in request_kwargs: metadata = request_kwargs[metadata_variable_name] request_tags = metadata.get("tags") + match_any = llm_router_instance.tag_filtering_match_any new_healthy_deployments = [] default_deployments = [] @@ -76,7 +88,6 @@ async def get_deployments_for_tag( "get_deployments_for_tag routing: router_keys: %s", request_tags ) # example this can be router_keys=["free", "custom"] - # get all deployments that have a superset of these router keys for deployment in healthy_deployments: deployment_litellm_params = deployment.get("litellm_params") deployment_tags = deployment_litellm_params.get("tags") @@ -90,7 +101,7 @@ async def get_deployments_for_tag( if deployment_tags is None: continue - if is_valid_deployment_tag(deployment_tags, request_tags): + if is_valid_deployment_tag(deployment_tags, request_tags, match_any): new_healthy_deployments.append(deployment) if "default" in deployment_tags: diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 6a254fc8252..88dee19ae5a 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -1,7 +1,7 @@ import re from dataclasses import dataclass from enum import Enum -from typing import Dict, List, Literal, Optional, Tuple, Union +from typing import Dict, List, Literal, Optional, Tuple from pydantic import BaseModel, Field from typing_extensions import Annotated @@ -185,6 +185,14 @@ DEFINED_PROMETHEUS_METRICS = Literal[ "litellm_redis_daily_spend_update_queue_size", "litellm_in_memory_spend_update_queue_size", "litellm_redis_spend_update_queue_size", + "litellm_request_queue_time_seconds", + "litellm_guardrail_latency_seconds", + "litellm_guardrail_errors_total", + "litellm_guardrail_requests_total", + # Cache metrics + "litellm_cache_hits_metric", + "litellm_cache_misses_metric", + "litellm_cached_tokens_metric", ] @@ -219,6 +227,23 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, ] + litellm_request_queue_time_seconds = [ + UserAPIKeyLabelNames.END_USER.value, + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.REQUESTED_MODEL.value, + UserAPIKeyLabelNames.TEAM.value, + UserAPIKeyLabelNames.TEAM_ALIAS.value, + UserAPIKeyLabelNames.USER.value, + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + ] + + # Guardrail metrics - these use custom labels (guardrail_name, status, error_type, hook_type) + # which are not part of UserAPIKeyLabelNames + litellm_guardrail_latency_seconds: List[str] = [] + litellm_guardrail_errors_total: List[str] = [] + litellm_guardrail_requests_total: List[str] = [] + litellm_proxy_total_requests_metric = [ UserAPIKeyLabelNames.END_USER.value, UserAPIKeyLabelNames.API_KEY_HASH.value, @@ -436,6 +461,21 @@ class PrometheusMetricLabels: litellm_redis_spend_update_queue_size: List[str] = [] + # Cache metrics - track cache hits, misses, and tokens served from cache + _cache_metric_labels = [ + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.TEAM.value, + UserAPIKeyLabelNames.TEAM_ALIAS.value, + UserAPIKeyLabelNames.END_USER.value, + UserAPIKeyLabelNames.USER.value, + ] + + litellm_cache_hits_metric = _cache_metric_labels + litellm_cache_misses_metric = _cache_metric_labels + litellm_cached_tokens_metric = _cache_metric_labels + @staticmethod def get_labels(label_name: DEFINED_PROMETHEUS_METRICS) -> List[str]: default_labels = getattr(PrometheusMetricLabels, label_name) @@ -460,11 +500,6 @@ class PrometheusMetricLabels: return default_labels + custom_labels -from typing import List, Optional - -from pydantic import BaseModel, Field - - class UserAPIKeyLabelValues(BaseModel): end_user: Annotated[ Optional[str], Field(..., alias=UserAPIKeyLabelNames.END_USER.value) diff --git a/litellm/types/management_endpoints/router_settings_endpoints.py b/litellm/types/management_endpoints/router_settings_endpoints.py index 9e3002ecf45..8b05c1483e8 100644 --- a/litellm/types/management_endpoints/router_settings_endpoints.py +++ b/litellm/types/management_endpoints/router_settings_endpoints.py @@ -184,6 +184,14 @@ ROUTER_SETTINGS_FIELDS: List[RouterSettingsField] = [ field_default=False, ui_field_name="Enable Tag Filtering", link="https://docs.litellm.ai/docs/proxy/tag_routing", + ), + RouterSettingsField( + field_name="tag_filtering_match_any", + field_type="Boolean", + field_value=None, + field_description="Match any tag instead of all tags for tag-based routing", + field_default=True, + ui_field_name="Tag Filtering Match Any", ), RouterSettingsField( field_name="disable_cooldowns", diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 96fd79f466b..94f33ffb297 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -1,3 +1,4 @@ +from datetime import datetime from typing import Any, Dict, List, Optional from pydantic import BaseModel, ConfigDict @@ -50,4 +51,5 @@ class MCPServer(BaseModel): env: Optional[Dict[str, str]] = None access_groups: Optional[List[str]] = None allow_all_keys: bool = False + updated_at: Optional[datetime] = None model_config = ConfigDict(arbitrary_types_allowed=True) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 3817f46c3e2..891826787ef 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3016,6 +3016,7 @@ class LlmProviders(str, Enum): AUTO_ROUTER = "auto_router" VERCEL_AI_GATEWAY = "vercel_ai_gateway" DOTPROMPT = "dotprompt" + MANUS = "manus" WANDB = "wandb" OVHCLOUD = "ovhcloud" LEMONADE = "lemonade" @@ -3028,6 +3029,7 @@ class LlmProviders(str, Enum): NANOGPT = "nano-gpt" POE = "poe" CHUTES = "chutes" + XIAOMI_MIMO = "xiaomi_mimo" diff --git a/litellm/utils.py b/litellm/utils.py index fbbaa94f7a1..42d4a1ac372 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -47,7 +47,6 @@ from tiktoken import Encoding from tokenizers import Tokenizer import litellm - import litellm.litellm_core_utils # audio_utils.utils is lazy-loaded - only imported when needed for transcription calls import litellm.litellm_core_utils.json_validation_rule @@ -2895,6 +2894,7 @@ def get_optional_params_image_gen( litellm.drop_params is True or drop_params is True ) and k not in supported_params: # drop the unsupported non-default values non_default_params.pop(k, None) + passed_params.pop(k, None) elif k not in supported_params: raise UnsupportedParamsError( status_code=500, @@ -7410,15 +7410,184 @@ def validate_chat_completion_tool_choice( class ProviderConfigManager: + # Dictionary mapping for O(1) provider lookup + # Stores tuples of (factory_function, needs_model_parameter) + # This is initialized lazily on first access to avoid circular imports + _PROVIDER_CONFIG_MAP: Optional[dict[LlmProviders, tuple[Callable, bool]]] = None + + @staticmethod + def _build_provider_config_map() -> dict[LlmProviders, tuple[Callable, bool]]: + """Build the provider-to-config mapping dictionary. + + Returns a dict mapping provider to (factory_function, needs_model_parameter). + This avoids expensive inspect.signature() calls at runtime. + """ + return { + # Most common providers first for readability + # Format: (factory_function, needs_model_parameter: bool) + LlmProviders.OPENAI: (lambda: litellm.OpenAIGPTConfig(), False), + LlmProviders.ANTHROPIC: (lambda: litellm.AnthropicConfig(), False), + LlmProviders.AZURE: (lambda model: ProviderConfigManager._get_azure_config(model), True), + LlmProviders.AZURE_AI: (lambda model: ProviderConfigManager._get_azure_ai_config(model), True), + LlmProviders.VERTEX_AI: (lambda model: ProviderConfigManager._get_vertex_ai_config(model), True), + LlmProviders.BEDROCK: (lambda model: ProviderConfigManager._get_bedrock_config(model), True), + LlmProviders.COHERE: (lambda model: ProviderConfigManager._get_cohere_config(model), True), + LlmProviders.COHERE_CHAT: (lambda model: ProviderConfigManager._get_cohere_config(model), True), + # Simple provider mappings (no model parameter needed) + LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False), + LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False), + LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False), + LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False), + LlmProviders.XAI: (lambda: litellm.XAIChatConfig(), False), + LlmProviders.ZAI: (lambda: litellm.ZAIChatConfig(), False), + LlmProviders.LAMBDA_AI: (lambda: litellm.LambdaAIChatConfig(), False), + LlmProviders.LLAMA: (lambda: litellm.LlamaAPIConfig(), False), + LlmProviders.TEXT_COMPLETION_OPENAI: (lambda: litellm.OpenAITextCompletionConfig(), False), + LlmProviders.SNOWFLAKE: (lambda: litellm.SnowflakeConfig(), False), + LlmProviders.CLARIFAI: (lambda: litellm.ClarifaiConfig(), False), + LlmProviders.ANTHROPIC_TEXT: (lambda: litellm.AnthropicTextConfig(), False), + LlmProviders.VERTEX_AI_BETA: (lambda: litellm.VertexGeminiConfig(), False), + LlmProviders.CLOUDFLARE: (lambda: litellm.CloudflareChatConfig(), False), + LlmProviders.SAGEMAKER_CHAT: (lambda: litellm.SagemakerChatConfig(), False), + LlmProviders.SAGEMAKER: (lambda: litellm.SagemakerConfig(), False), + LlmProviders.FIREWORKS_AI: (lambda: litellm.FireworksAIConfig(), False), + LlmProviders.FRIENDLIAI: (lambda: litellm.FriendliaiChatConfig(), False), + LlmProviders.WATSONX: (lambda: litellm.IBMWatsonXChatConfig(), False), + LlmProviders.WATSONX_TEXT: (lambda: litellm.IBMWatsonXAIConfig(), False), + LlmProviders.EMPOWER: (lambda: litellm.EmpowerChatConfig(), False), + LlmProviders.MINIMAX: (lambda: litellm.MinimaxChatConfig(), False), + LlmProviders.GITHUB: (lambda: litellm.GithubChatConfig(), False), + LlmProviders.COMPACTIFAI: (lambda: litellm.CompactifAIChatConfig(), False), + LlmProviders.GITHUB_COPILOT: (lambda: litellm.GithubCopilotConfig(), False), + LlmProviders.GIGACHAT: (lambda: litellm.GigaChatConfig(), False), + LlmProviders.RAGFLOW: (lambda: litellm.RAGFlowConfig(), False), + LlmProviders.CUSTOM: (lambda: litellm.OpenAILikeChatConfig(), False), + LlmProviders.CUSTOM_OPENAI: (lambda: litellm.OpenAILikeChatConfig(), False), + LlmProviders.OPENAI_LIKE: (lambda: litellm.OpenAILikeChatConfig(), False), + LlmProviders.AIOHTTP_OPENAI: (lambda: litellm.AiohttpOpenAIChatConfig(), False), + LlmProviders.HOSTED_VLLM: (lambda: litellm.HostedVLLMChatConfig(), False), + LlmProviders.LLAMAFILE: (lambda: litellm.LlamafileChatConfig(), False), + LlmProviders.LM_STUDIO: (lambda: litellm.LMStudioChatConfig(), False), + LlmProviders.GALADRIEL: (lambda: litellm.GaladrielChatConfig(), False), + LlmProviders.REPLICATE: (lambda: litellm.ReplicateConfig(), False), + LlmProviders.HUGGINGFACE: (lambda: litellm.HuggingFaceChatConfig(), False), + LlmProviders.TOGETHER_AI: (lambda: litellm.TogetherAIConfig(), False), + LlmProviders.OPENROUTER: (lambda: litellm.OpenrouterConfig(), False), + LlmProviders.VERCEL_AI_GATEWAY: (lambda: litellm.VercelAIGatewayConfig(), False), + LlmProviders.COMETAPI: (lambda: litellm.CometAPIConfig(), False), + LlmProviders.DATAROBOT: (lambda: litellm.DataRobotConfig(), False), + LlmProviders.GEMINI: (lambda: litellm.GoogleAIStudioGeminiConfig(), False), + LlmProviders.AI21: (lambda: litellm.AI21ChatConfig(), False), + LlmProviders.AI21_CHAT: (lambda: litellm.AI21ChatConfig(), False), + LlmProviders.AZURE_TEXT: (lambda: litellm.AzureOpenAITextConfig(), False), + LlmProviders.NLP_CLOUD: (lambda: litellm.NLPCloudConfig(), False), + LlmProviders.OOBABOOGA: (lambda: litellm.OobaboogaConfig(), False), + LlmProviders.OLLAMA_CHAT: (lambda: litellm.OllamaChatConfig(), False), + LlmProviders.DEEPINFRA: (lambda: litellm.DeepInfraConfig(), False), + LlmProviders.PERPLEXITY: (lambda: litellm.PerplexityChatConfig(), False), + LlmProviders.MISTRAL: (lambda: litellm.MistralConfig(), False), + LlmProviders.CODESTRAL: (lambda: litellm.MistralConfig(), False), + LlmProviders.NVIDIA_NIM: (lambda: litellm.NvidiaNimConfig(), False), + LlmProviders.CEREBRAS: (lambda: litellm.CerebrasConfig(), False), + LlmProviders.BASETEN: (lambda: litellm.BasetenConfig(), False), + LlmProviders.VOLCENGINE: (lambda: litellm.VolcEngineConfig(), False), + LlmProviders.TEXT_COMPLETION_CODESTRAL: (lambda: litellm.CodestralTextCompletionConfig(), False), + LlmProviders.SAMBANOVA: (lambda: litellm.SambanovaConfig(), False), + LlmProviders.MARITALK: (lambda: litellm.MaritalkConfig(), False), + LlmProviders.VLLM: (lambda: litellm.VLLMConfig(), False), + LlmProviders.OLLAMA: (lambda: litellm.OllamaConfig(), False), + LlmProviders.PREDIBASE: (lambda: litellm.PredibaseConfig(), False), + LlmProviders.TRITON: (lambda: litellm.TritonConfig(), False), + LlmProviders.PETALS: (lambda: litellm.PetalsConfig(), False), + LlmProviders.SAP_GENERATIVE_AI_HUB: (lambda: litellm.GenAIHubOrchestrationConfig(), False), + LlmProviders.FEATHERLESS_AI: (lambda: litellm.FeatherlessAIConfig(), False), + LlmProviders.NOVITA: (lambda: litellm.NovitaConfig(), False), + LlmProviders.NEBIUS: (lambda: litellm.NebiusConfig(), False), + LlmProviders.WANDB: (lambda: litellm.WandbConfig(), False), + LlmProviders.DASHSCOPE: (lambda: litellm.DashScopeChatConfig(), False), + LlmProviders.MOONSHOT: (lambda: litellm.MoonshotChatConfig(), False), + LlmProviders.DOCKER_MODEL_RUNNER: (lambda: litellm.DockerModelRunnerChatConfig(), False), + LlmProviders.V0: (lambda: litellm.V0ChatConfig(), False), + LlmProviders.MORPH: (lambda: litellm.MorphChatConfig(), False), + LlmProviders.LITELLM_PROXY: (lambda: litellm.LiteLLMProxyChatConfig(), False), + LlmProviders.GRADIENT_AI: (lambda: litellm.GradientAIConfig(), False), + LlmProviders.NSCALE: (lambda: litellm.NscaleConfig(), False), + LlmProviders.HEROKU: (lambda: litellm.HerokuChatConfig(), False), + LlmProviders.OCI: (lambda: litellm.OCIChatConfig(), False), + LlmProviders.HYPERBOLIC: (lambda: litellm.HyperbolicChatConfig(), False), + LlmProviders.OVHCLOUD: (lambda: litellm.OVHCloudChatConfig(), False), + LlmProviders.AMAZON_NOVA: (lambda: litellm.AmazonNovaChatConfig(), False), + LlmProviders.LANGGRAPH: (lambda: ProviderConfigManager._get_langgraph_config(), False), + } + + @staticmethod + def _get_azure_config(model: str) -> BaseConfig: + """Get Azure config based on model type.""" + if litellm.AzureOpenAIO1Config().is_o_series_model(model=model): + return litellm.AzureOpenAIO1Config() + if litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model): + return litellm.AzureOpenAIGPT5Config() + return litellm.AzureOpenAIConfig() + + @staticmethod + def _get_azure_ai_config(model: str) -> BaseConfig: + """Get Azure AI config based on model type.""" + if "claude" in model.lower(): + return litellm.AzureAnthropicConfig() + return litellm.AzureAIStudioConfig() + + @staticmethod + def _get_vertex_ai_config(model: str) -> BaseConfig: + """Get Vertex AI config based on model type.""" + if "gemini" in model: + return litellm.VertexGeminiConfig() + elif "claude" in model: + return litellm.VertexAIAnthropicConfig() + elif "gpt-oss" in model: + from litellm.llms.vertex_ai.vertex_ai_partner_models.gpt_oss.transformation import ( + VertexAIGPTOSSTransformation, + ) + return VertexAIGPTOSSTransformation() + elif model in litellm.vertex_mistral_models: + if "codestral" in model: + return litellm.CodestralTextCompletionConfig() + return litellm.MistralConfig() + elif model in litellm.vertex_ai_ai21_models: + return litellm.VertexAIAi21Config() + else: + return litellm.VertexAILlama3Config() + + @staticmethod + def _get_bedrock_config(model: str) -> BaseConfig: + """Get Bedrock config based on model.""" + from litellm.llms.bedrock.common_utils import get_bedrock_chat_config + return get_bedrock_chat_config(model=model) + + @staticmethod + def _get_cohere_config(model: str) -> BaseConfig: + """Get Cohere config based on route.""" + CohereModelInfo = getattr(sys.modules[__name__], 'CohereModelInfo') + route = CohereModelInfo.get_cohere_route(model) + if route == "v2": + return litellm.CohereV2ChatConfig() + return litellm.CohereChatConfig() + + @staticmethod + def _get_langgraph_config() -> BaseConfig: + """Get LangGraph config.""" + from litellm.llms.langgraph.chat.transformation import LangGraphConfig + return LangGraphConfig() + @staticmethod def get_provider_chat_config( # noqa: PLR0915 model: str, provider: LlmProviders ) -> Optional[BaseConfig]: """ Returns the provider config for a given provider. + + Uses O(1) dictionary lookup for fast provider resolution. """ - - # Check JSON providers FIRST + # Check JSON providers FIRST (these override standard mappings) from litellm.llms.openai_like.dynamic_config import create_config_class from litellm.llms.openai_like.json_loader import JSONProviderRegistry @@ -7428,244 +7597,29 @@ class ProviderConfigManager: raise ValueError(f"Provider {provider.value} not found") return create_config_class(provider_config)() - if ( - provider == LlmProviders.OPENAI - and litellm.openaiOSeriesConfig.is_model_o_series_model(model=model) - ): - return litellm.openaiOSeriesConfig - elif ( - provider == LlmProviders.OPENAI - and litellm.OpenAIGPT5Config.is_model_gpt_5_model(model=model) - ): - return litellm.OpenAIGPT5Config() - elif litellm.LlmProviders.DEEPSEEK == provider: - return litellm.DeepSeekChatConfig() - elif litellm.LlmProviders.GROQ == provider: - return litellm.GroqChatConfig() - elif litellm.LlmProviders.BYTEZ == provider: - return litellm.BytezChatConfig() - elif litellm.LlmProviders.DATABRICKS == provider: - return litellm.DatabricksConfig() - elif litellm.LlmProviders.XAI == provider: - return litellm.XAIChatConfig() - elif litellm.LlmProviders.ZAI == provider: - return litellm.ZAIChatConfig() - elif litellm.LlmProviders.LAMBDA_AI == provider: - return litellm.LambdaAIChatConfig() - elif litellm.LlmProviders.LLAMA == provider: - return litellm.LlamaAPIConfig() - elif litellm.LlmProviders.TEXT_COMPLETION_OPENAI == provider: - return litellm.OpenAITextCompletionConfig() - elif ( - litellm.LlmProviders.COHERE_CHAT == provider - or litellm.LlmProviders.COHERE == provider - ): - CohereModelInfo = getattr(sys.modules[__name__], 'CohereModelInfo') - route = CohereModelInfo.get_cohere_route(model) - if route == "v2": - return litellm.CohereV2ChatConfig() - else: + # Handle OpenAI special cases (O-series and GPT-5 models) + if provider == LlmProviders.OPENAI: + if litellm.openaiOSeriesConfig.is_model_o_series_model(model=model): + return litellm.openaiOSeriesConfig + if litellm.OpenAIGPT5Config.is_model_gpt_5_model(model=model): + return litellm.OpenAIGPT5Config() - return litellm.CohereChatConfig() - elif litellm.LlmProviders.SNOWFLAKE == provider: - return litellm.SnowflakeConfig() - elif litellm.LlmProviders.CLARIFAI == provider: - return litellm.ClarifaiConfig() - elif litellm.LlmProviders.ANTHROPIC == provider: - return litellm.AnthropicConfig() - elif litellm.LlmProviders.ANTHROPIC_TEXT == provider: - return litellm.AnthropicTextConfig() - elif litellm.LlmProviders.VERTEX_AI_BETA == provider: - return litellm.VertexGeminiConfig() - elif litellm.LlmProviders.VERTEX_AI == provider: - if "gemini" in model: - return litellm.VertexGeminiConfig() - elif "claude" in model: - return litellm.VertexAIAnthropicConfig() - elif "gpt-oss" in model: - from litellm.llms.vertex_ai.vertex_ai_partner_models.gpt_oss.transformation import ( - VertexAIGPTOSSTransformation, - ) + # Initialize provider config map lazily (avoids circular imports) + if ProviderConfigManager._PROVIDER_CONFIG_MAP is None: + ProviderConfigManager._PROVIDER_CONFIG_MAP = ProviderConfigManager._build_provider_config_map() - return VertexAIGPTOSSTransformation() - elif model in litellm.vertex_mistral_models: - if "codestral" in model: - return litellm.CodestralTextCompletionConfig() - else: - return litellm.MistralConfig() - elif model in litellm.vertex_ai_ai21_models: - return litellm.VertexAIAi21Config() - else: # use generic openai-like param mapping - return litellm.VertexAILlama3Config() - elif litellm.LlmProviders.CLOUDFLARE == provider: - return litellm.CloudflareChatConfig() - elif litellm.LlmProviders.SAGEMAKER_CHAT == provider: - return litellm.SagemakerChatConfig() - elif litellm.LlmProviders.SAGEMAKER == provider: - return litellm.SagemakerConfig() - elif litellm.LlmProviders.FIREWORKS_AI == provider: - return litellm.FireworksAIConfig() - elif litellm.LlmProviders.FRIENDLIAI == provider: - return litellm.FriendliaiChatConfig() - elif litellm.LlmProviders.WATSONX == provider: - return litellm.IBMWatsonXChatConfig() - elif litellm.LlmProviders.WATSONX_TEXT == provider: - return litellm.IBMWatsonXAIConfig() - elif litellm.LlmProviders.EMPOWER == provider: - return litellm.EmpowerChatConfig() - elif litellm.LlmProviders.MINIMAX == provider: - return litellm.MinimaxChatConfig() - elif litellm.LlmProviders.GITHUB == provider: - return litellm.GithubChatConfig() - elif litellm.LlmProviders.COMPACTIFAI == provider: - return litellm.CompactifAIChatConfig() - elif litellm.LlmProviders.GITHUB_COPILOT == provider: - return litellm.GithubCopilotConfig() - elif litellm.LlmProviders.GIGACHAT == provider: - return litellm.GigaChatConfig() - elif litellm.LlmProviders.RAGFLOW == provider: - return litellm.RAGFlowConfig() - elif ( - litellm.LlmProviders.CUSTOM == provider - or litellm.LlmProviders.CUSTOM_OPENAI == provider - or litellm.LlmProviders.OPENAI_LIKE == provider - ): - return litellm.OpenAILikeChatConfig() - elif litellm.LlmProviders.AIOHTTP_OPENAI == provider: - return litellm.AiohttpOpenAIChatConfig() - elif litellm.LlmProviders.HOSTED_VLLM == provider: - return litellm.HostedVLLMChatConfig() - elif litellm.LlmProviders.LLAMAFILE == provider: - return litellm.LlamafileChatConfig() - elif litellm.LlmProviders.LM_STUDIO == provider: - return litellm.LMStudioChatConfig() - elif litellm.LlmProviders.GALADRIEL == provider: - return litellm.GaladrielChatConfig() - elif litellm.LlmProviders.REPLICATE == provider: - return litellm.ReplicateConfig() - elif litellm.LlmProviders.HUGGINGFACE == provider: - return litellm.HuggingFaceChatConfig() - elif litellm.LlmProviders.TOGETHER_AI == provider: - return litellm.TogetherAIConfig() - elif litellm.LlmProviders.OPENROUTER == provider: - return litellm.OpenrouterConfig() - elif litellm.LlmProviders.VERCEL_AI_GATEWAY == provider: - return litellm.VercelAIGatewayConfig() - elif litellm.LlmProviders.COMETAPI == provider: - return litellm.CometAPIConfig() - elif litellm.LlmProviders.DATAROBOT == provider: - return litellm.DataRobotConfig() - elif litellm.LlmProviders.GEMINI == provider: - return litellm.GoogleAIStudioGeminiConfig() - elif ( - litellm.LlmProviders.AI21 == provider - or litellm.LlmProviders.AI21_CHAT == provider - ): - return litellm.AI21ChatConfig() - elif litellm.LlmProviders.AZURE == provider: - if litellm.AzureOpenAIO1Config().is_o_series_model(model=model): - return litellm.AzureOpenAIO1Config() - if litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model): - return litellm.AzureOpenAIGPT5Config() - return litellm.AzureOpenAIConfig() - elif litellm.LlmProviders.AZURE_AI == provider: - if "claude" in model.lower(): - return litellm.AzureAnthropicConfig() - return litellm.AzureAIStudioConfig() - elif litellm.LlmProviders.AZURE_TEXT == provider: - return litellm.AzureOpenAITextConfig() - elif litellm.LlmProviders.HOSTED_VLLM == provider: - return litellm.HostedVLLMChatConfig() - elif litellm.LlmProviders.NLP_CLOUD == provider: - return litellm.NLPCloudConfig() - elif litellm.LlmProviders.OOBABOOGA == provider: - return litellm.OobaboogaConfig() - elif litellm.LlmProviders.OLLAMA_CHAT == provider: - return litellm.OllamaChatConfig() - elif litellm.LlmProviders.DEEPINFRA == provider: - return litellm.DeepInfraConfig() - elif litellm.LlmProviders.PERPLEXITY == provider: - return litellm.PerplexityChatConfig() - elif ( - litellm.LlmProviders.MISTRAL == provider - or litellm.LlmProviders.CODESTRAL == provider - ): - return litellm.MistralConfig() - elif litellm.LlmProviders.NVIDIA_NIM == provider: - return litellm.NvidiaNimConfig() - elif litellm.LlmProviders.CEREBRAS == provider: - return litellm.CerebrasConfig() - elif litellm.LlmProviders.BASETEN == provider: - return litellm.BasetenConfig() - elif litellm.LlmProviders.VOLCENGINE == provider: - return litellm.VolcEngineConfig() - elif litellm.LlmProviders.TEXT_COMPLETION_CODESTRAL == provider: - return litellm.CodestralTextCompletionConfig() - elif litellm.LlmProviders.SAMBANOVA == provider: - return litellm.SambanovaConfig() - elif litellm.LlmProviders.MARITALK == provider: - return litellm.MaritalkConfig() - elif litellm.LlmProviders.CLOUDFLARE == provider: - return litellm.CloudflareChatConfig() - elif litellm.LlmProviders.ANTHROPIC_TEXT == provider: - return litellm.AnthropicTextConfig() - elif litellm.LlmProviders.VLLM == provider: - return litellm.VLLMConfig() - elif litellm.LlmProviders.OLLAMA == provider: - return litellm.OllamaConfig() - elif litellm.LlmProviders.PREDIBASE == provider: - return litellm.PredibaseConfig() - elif litellm.LlmProviders.TRITON == provider: - return litellm.TritonConfig() - elif litellm.LlmProviders.PETALS == provider: - return litellm.PetalsConfig() - elif litellm.LlmProviders.SAP_GENERATIVE_AI_HUB == provider: - return litellm.GenAIHubOrchestrationConfig() - elif litellm.LlmProviders.FEATHERLESS_AI == provider: - return litellm.FeatherlessAIConfig() - elif litellm.LlmProviders.NOVITA == provider: - 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: - return litellm.MoonshotChatConfig() - elif litellm.LlmProviders.DOCKER_MODEL_RUNNER == provider: - return litellm.DockerModelRunnerChatConfig() - elif litellm.LlmProviders.V0 == provider: - return litellm.V0ChatConfig() - elif litellm.LlmProviders.MORPH == provider: - return litellm.MorphChatConfig() - elif litellm.LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.common_utils import get_bedrock_chat_config + # O(1) dictionary lookup + config_entry = ProviderConfigManager._PROVIDER_CONFIG_MAP.get(provider) + if config_entry is None: + return None - return get_bedrock_chat_config(model=model) - elif litellm.LlmProviders.LITELLM_PROXY == provider: - return litellm.LiteLLMProxyChatConfig() - elif litellm.LlmProviders.OPENAI == provider: - return litellm.OpenAIGPTConfig() - elif litellm.LlmProviders.GRADIENT_AI == provider: - return litellm.GradientAIConfig() - elif litellm.LlmProviders.NSCALE == provider: - return litellm.NscaleConfig() - elif litellm.LlmProviders.HEROKU == provider: - return litellm.HerokuChatConfig() - elif litellm.LlmProviders.OCI == provider: - return litellm.OCIChatConfig() - elif litellm.LlmProviders.HYPERBOLIC == provider: - return litellm.HyperbolicChatConfig() - elif litellm.LlmProviders.OVHCLOUD == provider: - return litellm.OVHCloudChatConfig() - elif litellm.LlmProviders.AMAZON_NOVA == provider: - return litellm.AmazonNovaChatConfig() - elif litellm.LlmProviders.LANGGRAPH == provider: - from litellm.llms.langgraph.chat.transformation import LangGraphConfig - - return LangGraphConfig() - return None + # Unpack factory function and whether it needs model parameter + # This avoids expensive inspect.signature() calls at runtime + config_factory, needs_model = config_entry + if needs_model: + return config_factory(model) # type: ignore + else: + return config_factory() # type: ignore @staticmethod def get_provider_embedding_config( @@ -7718,6 +7672,11 @@ class ProviderConfigManager: return litellm.CometAPIEmbeddingConfig() elif litellm.LlmProviders.GITHUB_COPILOT == provider: return litellm.GithubCopilotEmbeddingConfig() + elif litellm.LlmProviders.OPENROUTER == provider: + from litellm.llms.openrouter.embedding.transformation import ( + OpenrouterEmbeddingConfig, + ) + return OpenrouterEmbeddingConfig() elif litellm.LlmProviders.GIGACHAT == provider: return litellm.GigaChatEmbeddingConfig() elif litellm.LlmProviders.SAGEMAKER == provider: @@ -7871,6 +7830,8 @@ class ProviderConfigManager: return litellm.GithubCopilotResponsesAPIConfig() elif litellm.LlmProviders.LITELLM_PROXY == provider: return litellm.LiteLLMProxyResponsesAPIConfig() + elif litellm.LlmProviders.MANUS == provider: + return litellm.ManusResponsesAPIConfig() return None @staticmethod @@ -7939,6 +7900,8 @@ class ProviderConfigManager: return litellm.LemonadeChatConfig() elif LlmProviders.CLARIFAI == provider: return litellm.ClarifaiConfig() + elif LlmProviders.BEDROCK == provider: + return litellm.llms.bedrock.common_utils.BedrockModelInfo() return None @staticmethod diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index fb00f636409..2f34d2a7d3d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -410,8 +410,8 @@ "max_input_tokens": 8172, "max_tokens": 8172, "mode": "embedding", - "input_cost_per_token": 1.35e-7, - "input_cost_per_image": 6e-5, + "input_cost_per_token": 1.35e-07, + "input_cost_per_image": 6e-05, "input_cost_per_video_per_second": 0.0007, "input_cost_per_audio_per_second": 0.00014, "output_cost_per_token": 0.0, @@ -1398,8 +1398,8 @@ "mode": "chat" }, "azure_ai/gpt-oss-120b": { - "input_cost_per_token": 1.5e-7, - "output_cost_per_token": 6e-7, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, "litellm_provider": "azure_ai", "max_input_tokens": 131072, "max_output_tokens": 131072, @@ -2077,7 +2077,7 @@ "litellm_provider": "azure", "max_input_tokens": 4097, "max_output_tokens": 4096, - "max_tokens": 4097, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, "supports_function_calling": true, @@ -2090,7 +2090,7 @@ "litellm_provider": "azure", "max_input_tokens": 4097, "max_output_tokens": 4096, - "max_tokens": 4097, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, "supports_function_calling": true, @@ -2869,7 +2869,7 @@ "/v1/audio/transcriptions" ] }, - "azure/gpt-5.1-2025-11-13": { + "azure/gpt-5.1-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, @@ -2905,7 +2905,7 @@ "supports_service_tier": true, "supports_vision": true }, - "azure/gpt-5.1-chat-2025-11-13": { + "azure/gpt-5.1-chat-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, @@ -2940,7 +2940,7 @@ "supports_tool_choice": false, "supports_vision": true }, - "azure/gpt-5.1-codex-2025-11-13": { + "azure/gpt-5.1-codex-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, @@ -3298,7 +3298,7 @@ "litellm_provider": "azure", "max_input_tokens": 272000, "max_output_tokens": 128000, - "max_tokens": 400000, + "max_tokens": 128000, "mode": "responses", "output_cost_per_token": 0.00012, "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/foundry-models/concepts/models-sold-directly-by-azure?pivots=azure-openai&tabs=global-standard-aoai%2Cstandard-chat-completions%2Cglobal-standard#gpt-5", @@ -3623,7 +3623,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", - "output_cost_per_token": 1.68e-04, + "output_cost_per_token": 0.000168, "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -3654,7 +3654,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", - "output_cost_per_token": 1.68e-04, + "output_cost_per_token": 0.000168, "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -4297,13 +4297,13 @@ "output_cost_per_token": 0.0 }, "azure/speech/azure-tts": { - "input_cost_per_character": 15e-06, + "input_cost_per_character": 1.5e-05, "litellm_provider": "azure", "mode": "audio_speech", "source": "https://azure.microsoft.com/en-us/pricing/calculator/" }, "azure/speech/azure-tts-hd": { - "input_cost_per_character": 30e-06, + "input_cost_per_character": 3e-05, "litellm_provider": "azure", "mode": "audio_speech", "source": "https://azure.microsoft.com/en-us/pricing/calculator/" @@ -5197,7 +5197,7 @@ }, "azure_ai/mistral-document-ai-2505": { "litellm_provider": "azure_ai", - "ocr_cost_per_page": 3e-3, + "ocr_cost_per_page": 0.003, "mode": "ocr", "supported_endpoints": [ "/v1/ocr" @@ -5206,7 +5206,7 @@ }, "azure_ai/doc-intelligence/prebuilt-read": { "litellm_provider": "azure_ai", - "ocr_cost_per_page": 1.5e-3, + "ocr_cost_per_page": 0.0015, "mode": "ocr", "supported_endpoints": [ "/v1/ocr" @@ -5215,7 +5215,7 @@ }, "azure_ai/doc-intelligence/prebuilt-layout": { "litellm_provider": "azure_ai", - "ocr_cost_per_page": 1e-2, + "ocr_cost_per_page": 0.01, "mode": "ocr", "supported_endpoints": [ "/v1/ocr" @@ -5224,7 +5224,7 @@ }, "azure_ai/doc-intelligence/prebuilt-document": { "litellm_provider": "azure_ai", - "ocr_cost_per_page": 1e-2, + "ocr_cost_per_page": 0.01, "mode": "ocr", "supported_endpoints": [ "/v1/ocr" @@ -5298,12 +5298,12 @@ "mode": "rerank", "output_cost_per_token": 0.0 }, - "azure_ai/deepseek-v3.2": { + "azure_ai/deepseek-v3.2": { "input_cost_per_token": 5.8e-07, "litellm_provider": "azure_ai", "max_input_tokens": 163840, "max_output_tokens": 163840, - "max_tokens": 8192, + "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.68e-06, "supports_assistant_prefill": true, @@ -5317,7 +5317,7 @@ "litellm_provider": "azure_ai", "max_input_tokens": 163840, "max_output_tokens": 163840, - "max_tokens": 8192, + "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.68e-06, "supports_assistant_prefill": true, @@ -5452,7 +5452,7 @@ "supports_web_search": true }, "azure_ai/grok-4-fast-non-reasoning": { - "input_cost_per_token": 0.43e-06, + "input_cost_per_token": 4.3e-07, "output_cost_per_token": 1.73e-06, "litellm_provider": "azure_ai", "max_input_tokens": 131072, @@ -5465,7 +5465,7 @@ "supports_web_search": true }, "azure_ai/grok-4-fast-reasoning": { - "input_cost_per_token": 0.43e-06, + "input_cost_per_token": 4.3e-07, "output_cost_per_token": 1.73e-06, "litellm_provider": "azure_ai", "max_input_tokens": 131072, @@ -5623,7 +5623,7 @@ "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, - "max_tokens": 16384, + "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 4e-07 }, @@ -7038,7 +7038,7 @@ "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, - "max_tokens": 1000000, + "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, "output_cost_per_token_above_200k_tokens": 2.25e-05, @@ -7800,7 +7800,7 @@ "litellm_provider": "deepseek", "max_input_tokens": 131072, "max_output_tokens": 8192, - "max_tokens": 131072, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.7e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", @@ -7821,7 +7821,7 @@ "litellm_provider": "deepseek", "max_input_tokens": 131072, "max_output_tokens": 65536, - "max_tokens": 131072, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 1.7e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", @@ -7842,7 +7842,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 1000000, "max_output_tokens": 16384, - "max_tokens": 1000000, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, "source": "https://www.alibabacloud.com/help/en/model-studio/models", @@ -7854,7 +7854,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 32768, - "max_tokens": 1000000, + "max_tokens": 32768, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -7883,7 +7883,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 32768, - "max_tokens": 1000000, + "max_tokens": 32768, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -7913,7 +7913,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 30720, "max_output_tokens": 8192, - "max_tokens": 32768, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 6.4e-06, "source": "https://www.alibabacloud.com/help/en/model-studio/models", @@ -7926,7 +7926,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 129024, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://www.alibabacloud.com/help/en/model-studio/models", @@ -7939,7 +7939,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 129024, "max_output_tokens": 8192, - "max_tokens": 131072, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://www.alibabacloud.com/help/en/model-studio/models", @@ -7952,7 +7952,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 129024, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", "output_cost_per_reasoning_token": 4e-06, "output_cost_per_token": 1.2e-06, @@ -7966,7 +7966,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 129024, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", "output_cost_per_reasoning_token": 4e-06, "output_cost_per_token": 1.2e-06, @@ -7979,7 +7979,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 32768, - "max_tokens": 1000000, + "max_tokens": 32768, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8010,7 +8010,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 32768, - "max_tokens": 1000000, + "max_tokens": 32768, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8041,7 +8041,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 32768, - "max_tokens": 1000000, + "max_tokens": 32768, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8073,7 +8073,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 129024, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", "output_cost_per_reasoning_token": 5e-07, "output_cost_per_token": 2e-07, @@ -8087,7 +8087,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 1000000, "max_output_tokens": 8192, - "max_tokens": 1000000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2e-07, "source": "https://www.alibabacloud.com/help/en/model-studio/models", @@ -8100,7 +8100,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 1000000, "max_output_tokens": 16384, - "max_tokens": 1000000, + "max_tokens": 16384, "mode": "chat", "output_cost_per_reasoning_token": 5e-07, "output_cost_per_token": 2e-07, @@ -8114,7 +8114,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 1000000, "max_output_tokens": 16384, - "max_tokens": 1000000, + "max_tokens": 16384, "mode": "chat", "output_cost_per_reasoning_token": 5e-07, "output_cost_per_token": 2e-07, @@ -8127,7 +8127,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 129024, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8138,7 +8138,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 65536, - "max_tokens": 1000000, + "max_tokens": 65536, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8187,7 +8187,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 65536, - "max_tokens": 1000000, + "max_tokens": 65536, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8232,7 +8232,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 65536, - "max_tokens": 1000000, + "max_tokens": 65536, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8281,7 +8281,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 997952, "max_output_tokens": 65536, - "max_tokens": 1000000, + "max_tokens": 65536, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8326,7 +8326,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 258048, "max_output_tokens": 65536, - "max_tokens": 262144, + "max_tokens": 65536, "mode": "chat", "source": "https://www.alibabacloud.com/help/en/model-studio/models", "supports_function_calling": true, @@ -8364,7 +8364,7 @@ "litellm_provider": "dashscope", "max_input_tokens": 98304, "max_output_tokens": 8192, - "max_tokens": 131072, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.4e-06, "source": "https://www.alibabacloud.com/help/en/model-studio/models", @@ -8393,7 +8393,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 128000, - "max_tokens": 200000, + "max_tokens": 128000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8412,7 +8412,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 64000, - "max_tokens": 200000, + "max_tokens": 64000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8431,7 +8431,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 32000, - "max_tokens": 200000, + "max_tokens": 32000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8450,7 +8450,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 32000, - "max_tokens": 200000, + "max_tokens": 32000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8469,7 +8469,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 64000, - "max_tokens": 200000, + "max_tokens": 64000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8488,7 +8488,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 64000, - "max_tokens": 200000, + "max_tokens": 64000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8507,7 +8507,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 64000, - "max_tokens": 200000, + "max_tokens": 64000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8526,7 +8526,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 64000, - "max_tokens": 200000, + "max_tokens": 64000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8545,7 +8545,7 @@ "litellm_provider": "databricks", "max_input_tokens": 1048576, "max_output_tokens": 65535, - "max_tokens": 1048576, + "max_tokens": 65535, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8562,7 +8562,7 @@ "litellm_provider": "databricks", "max_input_tokens": 1048576, "max_output_tokens": 65536, - "max_tokens": 1048576, + "max_tokens": 65536, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8579,7 +8579,7 @@ "litellm_provider": "databricks", "max_input_tokens": 128000, "max_output_tokens": 32000, - "max_tokens": 128000, + "max_tokens": 32000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8594,7 +8594,7 @@ "litellm_provider": "databricks", "max_input_tokens": 400000, "max_output_tokens": 128000, - "max_tokens": 400000, + "max_tokens": 128000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8609,7 +8609,7 @@ "litellm_provider": "databricks", "max_input_tokens": 400000, "max_output_tokens": 128000, - "max_tokens": 400000, + "max_tokens": 128000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8624,7 +8624,7 @@ "litellm_provider": "databricks", "max_input_tokens": 400000, "max_output_tokens": 128000, - "max_tokens": 400000, + "max_tokens": 128000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8639,7 +8639,7 @@ "litellm_provider": "databricks", "max_input_tokens": 400000, "max_output_tokens": 128000, - "max_tokens": 400000, + "max_tokens": 128000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8747,7 +8747,7 @@ "litellm_provider": "databricks", "max_input_tokens": 200000, "max_output_tokens": 128000, - "max_tokens": 200000, + "max_tokens": 128000, "metadata": { "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." }, @@ -8846,7 +8846,7 @@ "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, - "max_tokens": 16384, + "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 2e-06 }, @@ -10107,7 +10107,7 @@ "litellm_provider": "deepseek", "max_input_tokens": 163840, "max_output_tokens": 163840, - "max_tokens": 8192, + "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 4e-07, "supports_assistant_prefill": true, @@ -10121,7 +10121,7 @@ "litellm_provider": "bedrock_converse", "max_input_tokens": 163840, "max_output_tokens": 81920, - "max_tokens": 163840, + "max_tokens": 81920, "mode": "chat", "output_cost_per_token": 1.68e-06, "supports_function_calling": true, @@ -10202,14 +10202,14 @@ "mode": "search", "tiered_pricing": [ { - "input_cost_per_query": 5e-03, + "input_cost_per_query": 0.005, "max_results_range": [ 0, 25 ] }, { - "input_cost_per_query": 25e-03, + "input_cost_per_query": 0.025, "max_results_range": [ 26, 100 @@ -10222,70 +10222,70 @@ "mode": "search", "tiered_pricing": [ { - "input_cost_per_query": 1.66e-03, + "input_cost_per_query": 0.00166, "max_results_range": [ 1, 10 ] }, { - "input_cost_per_query": 3.32e-03, + "input_cost_per_query": 0.00332, "max_results_range": [ 11, 20 ] }, { - "input_cost_per_query": 4.98e-03, + "input_cost_per_query": 0.00498, "max_results_range": [ 21, 30 ] }, { - "input_cost_per_query": 6.64e-03, + "input_cost_per_query": 0.00664, "max_results_range": [ 31, 40 ] }, { - "input_cost_per_query": 8.3e-03, + "input_cost_per_query": 0.0083, "max_results_range": [ 41, 50 ] }, { - "input_cost_per_query": 9.96e-03, + "input_cost_per_query": 0.00996, "max_results_range": [ 51, 60 ] }, { - "input_cost_per_query": 11.62e-03, + "input_cost_per_query": 0.01162, "max_results_range": [ 61, 70 ] }, { - "input_cost_per_query": 13.28e-03, + "input_cost_per_query": 0.01328, "max_results_range": [ 71, 80 ] }, { - "input_cost_per_query": 14.94e-03, + "input_cost_per_query": 0.01494, "max_results_range": [ 81, 90 ] }, { - "input_cost_per_query": 16.6e-03, + "input_cost_per_query": 0.0166, "max_results_range": [ 91, 100 @@ -10297,7 +10297,7 @@ } }, "perplexity/search": { - "input_cost_per_query": 5e-03, + "input_cost_per_query": 0.005, "litellm_provider": "perplexity", "mode": "search" }, @@ -10395,7 +10395,7 @@ "supports_embedding_image_input": true }, "embed-multilingual-light-v3.0": { - "input_cost_per_token": 1e-04, + "input_cost_per_token": 0.0001, "litellm_provider": "cohere", "max_input_tokens": 1024, "max_tokens": 1024, @@ -10689,7 +10689,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.3e-07, "supports_function_calling": true, @@ -10700,7 +10700,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.9e-07, "supports_function_calling": true, @@ -10711,7 +10711,7 @@ "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-06, "supports_function_calling": true, @@ -10817,14 +10817,14 @@ "litellm_provider": "featherless_ai", "max_input_tokens": 32768, "max_output_tokens": 4096, - "max_tokens": 32768, + "max_tokens": 4096, "mode": "chat" }, "featherless_ai/featherless-ai/Qwerky-QwQ-32B": { "litellm_provider": "featherless_ai", "max_input_tokens": 32768, "max_output_tokens": 4096, - "max_tokens": 32768, + "max_tokens": 4096, "mode": "chat" }, "fireworks-ai-4.1b-to-16b": { @@ -11031,7 +11031,7 @@ "supports_tool_choice": true }, "fireworks_ai/accounts/fireworks/models/glm-4p6": { - "input_cost_per_token": 0.55e-06, + "input_cost_per_token": 5.5e-07, "output_cost_per_token": 2.19e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 202800, @@ -11077,7 +11077,7 @@ "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 2.5e-06, "source": "https://fireworks.ai/models/fireworks/kimi-k2-instruct", @@ -11090,7 +11090,7 @@ "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, - "max_tokens": 262144, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 2.5e-06, "source": "https://app.fireworks.ai/models/fireworks/kimi-k2-instruct-0905", @@ -11337,7 +11337,7 @@ "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, - "max_tokens": 16384, + "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 1.6e-06, "output_cost_per_token_batches": 2e-07 @@ -11348,7 +11348,7 @@ "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, - "max_tokens": 16384, + "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 1e-06 @@ -11638,7 +11638,7 @@ "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 8192, "max_output_tokens": 2048, - "max_tokens": 8192, + "max_tokens": 2048, "mode": "chat", "output_cost_per_character": 3.75e-07, "output_cost_per_token": 1.5e-06, @@ -11655,7 +11655,7 @@ "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 8192, "max_output_tokens": 2048, - "max_tokens": 8192, + "max_tokens": 2048, "mode": "chat", "output_cost_per_character": 3.75e-07, "output_cost_per_token": 1.5e-06, @@ -12585,10 +12585,10 @@ "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, - "max_tokens": 65536, + "max_tokens": 32768, "mode": "image_generation", "output_cost_per_image": 0.134, - "output_cost_per_image_token": 1.2e-04, + "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, "source": "https://ai.google.dev/gemini-api/docs/pricing", @@ -14112,7 +14112,7 @@ "max_input_tokens": 1048576, "max_output_tokens": 65536, "max_pdf_size_mb": 30, - "max_tokens": 8192, + "max_tokens": 65536, "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", @@ -14162,7 +14162,7 @@ "max_input_tokens": 1048576, "max_output_tokens": 65536, "max_pdf_size_mb": 30, - "max_tokens": 8192, + "max_tokens": 65536, "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", @@ -14388,10 +14388,10 @@ "litellm_provider": "gemini", "max_input_tokens": 65536, "max_output_tokens": 32768, - "max_tokens": 65536, + "max_tokens": 32768, "mode": "image_generation", "output_cost_per_image": 0.134, - "output_cost_per_image_token": 1.2e-04, + "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "rpm": 1000, "tpm": 4000000, @@ -15524,7 +15524,7 @@ "max_input_tokens": 1024, "max_tokens": 1024, "mode": "video_generation", - "output_cost_per_second": 0.40, + "output_cost_per_second": 0.4, "source": "https://ai.google.dev/gemini-api/docs/video", "supported_modalities": [ "text" @@ -15552,7 +15552,7 @@ "max_input_tokens": 1024, "max_tokens": 1024, "mode": "video_generation", - "output_cost_per_second": 0.40, + "output_cost_per_second": 0.4, "source": "https://ai.google.dev/gemini-api/docs/video", "supported_modalities": [ "text" @@ -16056,11 +16056,11 @@ "supports_vision": true }, "gpt-3.5-turbo": { - "input_cost_per_token": 0.5e-06, + "input_cost_per_token": 5e-07, "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, - "max_tokens": 4097, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-06, "supports_function_calling": true, @@ -16073,7 +16073,7 @@ "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, - "max_tokens": 16385, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-06, "supports_function_calling": true, @@ -16087,7 +16087,7 @@ "litellm_provider": "openai", "max_input_tokens": 4097, "max_output_tokens": 4096, - "max_tokens": 4097, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, "supports_prompt_caching": true, @@ -16099,7 +16099,7 @@ "litellm_provider": "openai", "max_input_tokens": 4097, "max_output_tokens": 4096, - "max_tokens": 4097, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, "supports_function_calling": true, @@ -16113,7 +16113,7 @@ "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, - "max_tokens": 16385, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, "supports_function_calling": true, @@ -16127,7 +16127,7 @@ "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, - "max_tokens": 16385, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 4e-06, "supports_prompt_caching": true, @@ -16139,7 +16139,7 @@ "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, - "max_tokens": 16385, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 4e-06, "supports_prompt_caching": true, @@ -17191,7 +17191,7 @@ "supports_pdf_input": true }, "high/1024-x-1536/gpt-image-1.5": { - "input_cost_per_image": 0.20, + "input_cost_per_image": 0.2, "litellm_provider": "openai", "mode": "image_generation", "supported_endpoints": [ @@ -17202,7 +17202,7 @@ "supports_pdf_input": true }, "high/1536-x-1024/gpt-image-1.5": { - "input_cost_per_image": 0.20, + "input_cost_per_image": 0.2, "litellm_provider": "openai", "mode": "image_generation", "supported_endpoints": [ @@ -17356,7 +17356,7 @@ "supports_pdf_input": true }, "high/1024-x-1536/gpt-image-1.5-2025-12-16": { - "input_cost_per_image": 0.20, + "input_cost_per_image": 0.2, "litellm_provider": "openai", "mode": "image_generation", "supported_endpoints": [ @@ -17367,7 +17367,7 @@ "supports_pdf_input": true }, "high/1536-x-1024/gpt-image-1.5-2025-12-16": { - "input_cost_per_image": 0.20, + "input_cost_per_image": 0.2, "litellm_provider": "openai", "mode": "image_generation", "supported_endpoints": [ @@ -17704,7 +17704,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", - "output_cost_per_token": 1.68e-04, + "output_cost_per_token": 0.000168, "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -17735,7 +17735,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", - "output_cost_per_token": 1.68e-04, + "output_cost_per_token": 0.000168, "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -17767,7 +17767,7 @@ "max_output_tokens": 272000, "max_tokens": 272000, "mode": "responses", - "output_cost_per_token": 1.2e-04, + "output_cost_per_token": 0.00012, "output_cost_per_token_batches": 6e-05, "supported_endpoints": [ "/v1/batch", @@ -17800,7 +17800,7 @@ "max_output_tokens": 272000, "max_tokens": 272000, "mode": "responses", - "output_cost_per_token": 1.2e-04, + "output_cost_per_token": 0.00012, "output_cost_per_token_batches": 6e-05, "supported_endpoints": [ "/v1/batch", @@ -18503,7 +18503,7 @@ "lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": { "input_cost_per_token": 0, "litellm_provider": "lemonade", - "max_tokens": 262144, + "max_tokens": 32768, "max_input_tokens": 262144, "max_output_tokens": 32768, "mode": "chat", @@ -18515,7 +18515,7 @@ "lemonade/gpt-oss-20b-mxfp4-GGUF": { "input_cost_per_token": 0, "litellm_provider": "lemonade", - "max_tokens": 131072, + "max_tokens": 32768, "max_input_tokens": 131072, "max_output_tokens": 32768, "mode": "chat", @@ -18527,7 +18527,7 @@ "lemonade/gpt-oss-120b-mxfp-GGUF": { "input_cost_per_token": 0, "litellm_provider": "lemonade", - "max_tokens": 131072, + "max_tokens": 32768, "max_input_tokens": 131072, "max_output_tokens": 32768, "mode": "chat", @@ -18539,7 +18539,7 @@ "lemonade/Gemma-3-4b-it-GGUF": { "input_cost_per_token": 0, "litellm_provider": "lemonade", - "max_tokens": 128000, + "max_tokens": 8192, "max_input_tokens": 128000, "max_output_tokens": 8192, "mode": "chat", @@ -18551,7 +18551,7 @@ "lemonade/Qwen3-4B-Instruct-2507-GGUF": { "input_cost_per_token": 0, "litellm_provider": "lemonade", - "max_tokens": 262144, + "max_tokens": 32768, "max_input_tokens": 262144, "max_output_tokens": 32768, "mode": "chat", @@ -18688,11 +18688,11 @@ "groq/moonshotai/kimi-k2-instruct-0905": { "input_cost_per_token": 1e-06, "output_cost_per_token": 3e-06, - "cache_read_input_token_cost": 0.5e-06, + "cache_read_input_token_cost": 5e-07, "litellm_provider": "groq", "max_input_tokens": 262144, "max_output_tokens": 16384, - "max_tokens": 278528, + "max_tokens": 16384, "mode": "chat", "supports_function_calling": true, "supports_response_schema": true, @@ -19353,7 +19353,7 @@ "litellm_provider": "lambda_ai", "max_input_tokens": 131072, "max_output_tokens": 8192, - "max_tokens": 131072, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1e-07, "supports_function_calling": true, @@ -19366,7 +19366,7 @@ "litellm_provider": "lambda_ai", "max_input_tokens": 16384, "max_output_tokens": 8192, - "max_tokens": 16384, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1e-07, "supports_function_calling": true, @@ -19702,7 +19702,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.6e-05, "supports_function_calling": true, @@ -19713,7 +19713,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 9.9e-07, "supports_function_calling": true, @@ -19724,7 +19724,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 2.2e-07, "supports_function_calling": true, @@ -19735,7 +19735,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 3.5e-07, "supports_function_calling": true, @@ -19747,7 +19747,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1e-07, "supports_function_calling": true, @@ -19758,7 +19758,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-07, "supports_function_calling": true, @@ -19769,7 +19769,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, "supports_function_calling": true, @@ -19851,7 +19851,7 @@ "litellm_provider": "meta_llama", "max_input_tokens": 128000, "max_output_tokens": 4028, - "max_tokens": 128000, + "max_tokens": 4028, "mode": "chat", "source": "https://llama.developer.meta.com/docs/models", "supported_modalities": [ @@ -19867,7 +19867,7 @@ "litellm_provider": "meta_llama", "max_input_tokens": 128000, "max_output_tokens": 4028, - "max_tokens": 128000, + "max_tokens": 4028, "mode": "chat", "source": "https://llama.developer.meta.com/docs/models", "supported_modalities": [ @@ -19883,7 +19883,7 @@ "litellm_provider": "meta_llama", "max_input_tokens": 1000000, "max_output_tokens": 4028, - "max_tokens": 128000, + "max_tokens": 4028, "mode": "chat", "source": "https://llama.developer.meta.com/docs/models", "supported_modalities": [ @@ -19900,7 +19900,7 @@ "litellm_provider": "meta_llama", "max_input_tokens": 10000000, "max_output_tokens": 4028, - "max_tokens": 128000, + "max_tokens": 4028, "mode": "chat", "source": "https://llama.developer.meta.com/docs/models", "supported_modalities": [ @@ -19932,7 +19932,7 @@ ] }, "minimax/speech-02-turbo": { - "input_cost_per_character": 0.00006, + "input_cost_per_character": 6e-05, "litellm_provider": "minimax", "mode": "audio_speech", "supported_endpoints": [ @@ -19948,7 +19948,7 @@ ] }, "minimax/speech-2.6-turbo": { - "input_cost_per_character": 0.00006, + "input_cost_per_character": 6e-05, "litellm_provider": "minimax", "mode": "audio_speech", "supported_endpoints": [ @@ -20278,8 +20278,8 @@ }, "mistral/mistral-ocr-latest": { "litellm_provider": "mistral", - "ocr_cost_per_page": 1e-3, - "annotation_cost_per_page": 3e-3, + "ocr_cost_per_page": 0.001, + "annotation_cost_per_page": 0.003, "mode": "ocr", "supported_endpoints": [ "/v1/ocr" @@ -20288,8 +20288,8 @@ }, "mistral/mistral-ocr-2505-completion": { "litellm_provider": "mistral", - "ocr_cost_per_page": 1e-3, - "annotation_cost_per_page": 3e-3, + "ocr_cost_per_page": 0.001, + "annotation_cost_per_page": 0.003, "mode": "ocr", "supported_endpoints": [ "/v1/ocr" @@ -20349,14 +20349,14 @@ "mode": "embedding" }, "mistral/codestral-embed": { - "input_cost_per_token": 0.15e-06, + "input_cost_per_token": 1.5e-07, "litellm_provider": "mistral", "max_input_tokens": 8192, "max_tokens": 8192, "mode": "embedding" }, "mistral/codestral-embed-2505": { - "input_cost_per_token": 0.15e-06, + "input_cost_per_token": 1.5e-07, "litellm_provider": "mistral", "max_input_tokens": 8192, "max_tokens": 8192, @@ -20757,28 +20757,28 @@ "supports_vision": true }, "moonshot/kimi-k2-thinking": { - "cache_read_input_token_cost": 1.5e-7, - "input_cost_per_token": 6e-7, + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 6e-07, "litellm_provider": "moonshot", "max_input_tokens": 262144, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", - "output_cost_per_token": 2.5e-6, + "output_cost_per_token": 2.5e-06, "source": "https://platform.moonshot.ai/docs/pricing/chat#generation-model-kimi-k2", "supports_function_calling": true, "supports_tool_choice": true, "supports_web_search": true }, "moonshot/kimi-k2-thinking-turbo": { - "cache_read_input_token_cost": 1.5e-7, - "input_cost_per_token": 1.15e-6, + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 1.15e-06, "litellm_provider": "moonshot", "max_input_tokens": 262144, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", - "output_cost_per_token": 8e-6, + "output_cost_per_token": 8e-06, "source": "https://platform.moonshot.ai/docs/pricing/chat#generation-model-kimi-k2", "supports_function_calling": true, "supports_tool_choice": true, @@ -21650,7 +21650,7 @@ "litellm_provider": "oci", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 1.068e-05, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", @@ -21662,7 +21662,7 @@ "litellm_provider": "oci", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 2e-06, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", @@ -21674,7 +21674,7 @@ "litellm_provider": "oci", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 7.2e-07, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", @@ -21686,7 +21686,7 @@ "litellm_provider": "oci", "max_input_tokens": 512000, "max_output_tokens": 4000, - "max_tokens": 512000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 7.2e-07, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", @@ -21698,7 +21698,7 @@ "litellm_provider": "oci", "max_input_tokens": 192000, "max_output_tokens": 4000, - "max_tokens": 192000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 7.2e-07, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", @@ -21770,7 +21770,7 @@ "litellm_provider": "oci", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 1.56e-06, "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", @@ -21782,7 +21782,7 @@ "litellm_provider": "oci", "max_input_tokens": 256000, "max_output_tokens": 4000, - "max_tokens": 256000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 1.56e-06, "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", @@ -21794,7 +21794,7 @@ "litellm_provider": "oci", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 1.56e-06, "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", @@ -21806,7 +21806,7 @@ "litellm_provider": "ollama", "max_input_tokens": 32768, "max_output_tokens": 8192, - "max_tokens": 32768, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.0, "supports_function_calling": false @@ -21844,7 +21844,7 @@ "litellm_provider": "ollama", "max_input_tokens": 32768, "max_output_tokens": 8192, - "max_tokens": 32768, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.0, "supports_function_calling": true @@ -21864,12 +21864,12 @@ "litellm_provider": "ollama", "max_input_tokens": 32768, "max_output_tokens": 8192, - "max_tokens": 32768, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.0, "supports_function_calling": true }, - "ollama/deepseek-v3.1:671b-cloud" : { + "ollama/deepseek-v3.1:671b-cloud": { "input_cost_per_token": 0.0, "litellm_provider": "ollama", "max_input_tokens": 163840, @@ -21879,7 +21879,7 @@ "output_cost_per_token": 0.0, "supports_function_calling": true }, - "ollama/gpt-oss:120b-cloud" : { + "ollama/gpt-oss:120b-cloud": { "input_cost_per_token": 0.0, "litellm_provider": "ollama", "max_input_tokens": 131072, @@ -21889,7 +21889,7 @@ "output_cost_per_token": 0.0, "supports_function_calling": true }, - "ollama/gpt-oss:20b-cloud" : { + "ollama/gpt-oss:20b-cloud": { "input_cost_per_token": 0.0, "litellm_provider": "ollama", "max_input_tokens": 131072, @@ -21904,7 +21904,7 @@ "litellm_provider": "ollama", "max_input_tokens": 32768, "max_output_tokens": 8192, - "max_tokens": 32768, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.0, "supports_function_calling": true @@ -21968,7 +21968,7 @@ "litellm_provider": "ollama", "max_input_tokens": 8192, "max_output_tokens": 8192, - "max_tokens": 32768, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.0, "supports_function_calling": true @@ -22026,7 +22026,7 @@ "litellm_provider": "ollama", "max_input_tokens": 65536, "max_output_tokens": 8192, - "max_tokens": 65536, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.0, "supports_function_calling": true @@ -22084,7 +22084,7 @@ "litellm_provider": "openai", "max_input_tokens": 32768, "max_output_tokens": 0, - "max_tokens": 32768, + "max_tokens": 0, "mode": "moderation", "output_cost_per_token": 0.0 }, @@ -22093,7 +22093,7 @@ "litellm_provider": "openai", "max_input_tokens": 32768, "max_output_tokens": 0, - "max_tokens": 32768, + "max_tokens": 0, "mode": "moderation", "output_cost_per_token": 0.0 }, @@ -22102,7 +22102,7 @@ "litellm_provider": "openai", "max_input_tokens": 32768, "max_output_tokens": 0, - "max_tokens": 32768, + "max_tokens": 0, "mode": "moderation", "output_cost_per_token": 0.0 }, @@ -22156,7 +22156,7 @@ "input_cost_per_token": 1.102e-05, "litellm_provider": "openrouter", "max_output_tokens": 8191, - "max_tokens": 100000, + "max_tokens": 8191, "mode": "chat", "output_cost_per_token": 3.268e-05, "supports_tool_choice": true @@ -22296,7 +22296,7 @@ "input_cost_per_token": 1.63e-06, "litellm_provider": "openrouter", "max_output_tokens": 8191, - "max_tokens": 100000, + "max_tokens": 8191, "mode": "chat", "output_cost_per_token": 5.51e-06, "supports_tool_choice": true @@ -22491,7 +22491,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 163840, "max_output_tokens": 163840, - "max_tokens": 8192, + "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 8e-07, "supports_assistant_prefill": true, @@ -22506,7 +22506,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 163840, "max_output_tokens": 163840, - "max_tokens": 8192, + "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 4e-07, "supports_assistant_prefill": true, @@ -22521,7 +22521,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 163840, "max_output_tokens": 163840, - "max_tokens": 8192, + "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 4e-07, "supports_assistant_prefill": true, @@ -22535,7 +22535,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 66000, "max_output_tokens": 4096, - "max_tokens": 8192, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2.8e-07, "supports_prompt_caching": true, @@ -22693,51 +22693,51 @@ "supports_web_search": true }, "openrouter/google/gemini-3-flash-preview": { - "cache_read_input_token_cost": 5e-08, - "input_cost_per_audio_token": 1e-06, - "input_cost_per_token": 5e-07, - "litellm_provider": "openrouter", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_reasoning_token": 3e-06, - "output_cost_per_token": 3e-06, - "rpm": 2000, - "source": "https://ai.google.dev/pricing/gemini-3", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true, - "tpm": 800000 + "cache_read_input_token_cost": 5e-08, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 5e-07, + "litellm_provider": "openrouter", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 3e-06, + "output_cost_per_token": 3e-06, + "rpm": 2000, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_vision": true, + "supports_web_search": true, + "tpm": 800000 }, "openrouter/google/gemini-pro-1.5": { "input_cost_per_image": 0.00265, @@ -22868,13 +22868,13 @@ "supports_tool_choice": true }, "openrouter/minimax/minimax-m2": { - "input_cost_per_token": 2.55e-7, + "input_cost_per_token": 2.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 204800, - "max_tokens": 32768, + "max_tokens": 204800, "mode": "chat", - "output_cost_per_token": 1.02e-6, + "output_cost_per_token": 1.02e-06, "supports_function_calling": true, "supports_prompt_caching": false, "supports_reasoning": true, @@ -22900,7 +22900,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, - "max_tokens": 262144, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 6e-07, "supports_function_calling": true, @@ -23285,7 +23285,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 400000, "max_output_tokens": 128000, - "max_tokens": 400000, + "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.4e-05, "supports_function_calling": true, @@ -23301,7 +23301,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 128000, "max_output_tokens": 16384, - "max_tokens": 128000, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.4e-05, "supports_function_calling": true, @@ -23315,9 +23315,9 @@ "litellm_provider": "openrouter", "max_input_tokens": 400000, "max_output_tokens": 128000, - "max_tokens": 400000, + "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 1.68e-04, + "output_cost_per_token": 0.000168, "supports_function_calling": true, "supports_prompt_caching": false, "supports_reasoning": true, @@ -23474,20 +23474,20 @@ "litellm_provider": "openrouter", "max_input_tokens": 8192, "max_output_tokens": 2048, - "max_tokens": 8192, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 6.3e-07, "supports_tool_choice": true, "supports_vision": true }, "openrouter/qwen/qwen3-coder": { - "input_cost_per_token": 2.2e-7, + "input_cost_per_token": 2.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262100, "max_output_tokens": 262100, "max_tokens": 262100, "mode": "chat", - "output_cost_per_token": 9.5e-7, + "output_cost_per_token": 9.5e-07, "source": "https://openrouter.ai/qwen/qwen3-coder", "supports_tool_choice": true, "supports_function_calling": true @@ -23530,7 +23530,7 @@ "litellm_provider": "openrouter", "max_input_tokens": 2000000, "max_output_tokens": 30000, - "max_tokens": 2000000, + "max_tokens": 30000, "mode": "chat", "output_cost_per_token": 0, "source": "https://openrouter.ai/x-ai/grok-4-fast:free", @@ -23540,26 +23540,26 @@ "supports_web_search": false }, "openrouter/z-ai/glm-4.6": { - "input_cost_per_token": 4.0e-7, + "input_cost_per_token": 4e-07, "litellm_provider": "openrouter", "max_input_tokens": 202800, "max_output_tokens": 131000, - "max_tokens": 202800, + "max_tokens": 131000, "mode": "chat", - "output_cost_per_token": 1.75e-6, + "output_cost_per_token": 1.75e-06, "source": "https://openrouter.ai/z-ai/glm-4.6", "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true }, "openrouter/z-ai/glm-4.6:exacto": { - "input_cost_per_token": 4.5e-7, + "input_cost_per_token": 4.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 202800, "max_output_tokens": 131000, - "max_tokens": 202800, + "max_tokens": 131000, "mode": "chat", - "output_cost_per_token": 1.9e-6, + "output_cost_per_token": 1.9e-06, "source": "https://openrouter.ai/z-ai/glm-4.6:exacto", "supports_function_calling": true, "supports_reasoning": true, @@ -24107,7 +24107,7 @@ "litellm_provider": "publicai", "max_input_tokens": 8192, "max_output_tokens": 4096, - "max_tokens": 8192, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24119,7 +24119,7 @@ "litellm_provider": "publicai", "max_input_tokens": 8192, "max_output_tokens": 4096, - "max_tokens": 8192, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24131,7 +24131,7 @@ "litellm_provider": "publicai", "max_input_tokens": 8192, "max_output_tokens": 4096, - "max_tokens": 8192, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24143,7 +24143,7 @@ "litellm_provider": "publicai", "max_input_tokens": 16384, "max_output_tokens": 4096, - "max_tokens": 16384, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24155,7 +24155,7 @@ "litellm_provider": "publicai", "max_input_tokens": 8192, "max_output_tokens": 4096, - "max_tokens": 8192, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24167,7 +24167,7 @@ "litellm_provider": "publicai", "max_input_tokens": 32768, "max_output_tokens": 4096, - "max_tokens": 32768, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24179,7 +24179,7 @@ "litellm_provider": "publicai", "max_input_tokens": 32768, "max_output_tokens": 4096, - "max_tokens": 32768, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24191,7 +24191,7 @@ "litellm_provider": "publicai", "max_input_tokens": 32768, "max_output_tokens": 4096, - "max_tokens": 32768, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24204,7 +24204,7 @@ "litellm_provider": "publicai", "max_input_tokens": 32768, "max_output_tokens": 4096, - "max_tokens": 32768, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://platform.publicai.co/docs", @@ -24217,7 +24217,7 @@ "litellm_provider": "bedrock_converse", "max_input_tokens": 262000, "max_output_tokens": 65536, - "max_tokens": 262144, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 1.8e-06, "supports_function_calling": true, @@ -24229,7 +24229,7 @@ "litellm_provider": "bedrock_converse", "max_input_tokens": 262144, "max_output_tokens": 131072, - "max_tokens": 262144, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 8.8e-07, "supports_function_calling": true, @@ -24241,9 +24241,9 @@ "litellm_provider": "bedrock_converse", "max_input_tokens": 262144, "max_output_tokens": 131072, - "max_tokens": 262144, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 6.0e-07, + "output_cost_per_token": 6e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true @@ -24253,9 +24253,9 @@ "litellm_provider": "bedrock_converse", "max_input_tokens": 131072, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 6.0e-07, + "output_cost_per_token": 6e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true @@ -24756,12 +24756,11 @@ "supports_reasoning": true, "source": "https://cloud.sambanova.ai/plans/pricing" }, - "snowflake/claude-3-5-sonnet": { "litellm_provider": "snowflake", "max_input_tokens": 18000, "max_output_tokens": 8192, - "max_tokens": 18000, + "max_tokens": 8192, "mode": "chat", "supports_computer_use": true }, @@ -24769,7 +24768,7 @@ "litellm_provider": "snowflake", "max_input_tokens": 32768, "max_output_tokens": 8192, - "max_tokens": 32768, + "max_tokens": 8192, "mode": "chat", "supports_reasoning": true }, @@ -24777,293 +24776,339 @@ "litellm_provider": "snowflake", "max_input_tokens": 8000, "max_output_tokens": 8192, - "max_tokens": 8000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/jamba-1.5-large": { "litellm_provider": "snowflake", "max_input_tokens": 256000, "max_output_tokens": 8192, - "max_tokens": 256000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/jamba-1.5-mini": { "litellm_provider": "snowflake", "max_input_tokens": 256000, "max_output_tokens": 8192, - "max_tokens": 256000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/jamba-instruct": { "litellm_provider": "snowflake", "max_input_tokens": 256000, "max_output_tokens": 8192, - "max_tokens": 256000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama2-70b-chat": { "litellm_provider": "snowflake", "max_input_tokens": 4096, "max_output_tokens": 8192, - "max_tokens": 4096, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3-70b": { "litellm_provider": "snowflake", "max_input_tokens": 8000, "max_output_tokens": 8192, - "max_tokens": 8000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3-8b": { "litellm_provider": "snowflake", "max_input_tokens": 8000, "max_output_tokens": 8192, - "max_tokens": 8000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3.1-405b": { "litellm_provider": "snowflake", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3.1-70b": { "litellm_provider": "snowflake", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3.1-8b": { "litellm_provider": "snowflake", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3.2-1b": { "litellm_provider": "snowflake", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3.2-3b": { "litellm_provider": "snowflake", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/llama3.3-70b": { "litellm_provider": "snowflake", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/mistral-7b": { "litellm_provider": "snowflake", "max_input_tokens": 32000, "max_output_tokens": 8192, - "max_tokens": 32000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/mistral-large": { "litellm_provider": "snowflake", "max_input_tokens": 32000, "max_output_tokens": 8192, - "max_tokens": 32000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/mistral-large2": { "litellm_provider": "snowflake", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/mixtral-8x7b": { "litellm_provider": "snowflake", "max_input_tokens": 32000, "max_output_tokens": 8192, - "max_tokens": 32000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/reka-core": { "litellm_provider": "snowflake", "max_input_tokens": 32000, "max_output_tokens": 8192, - "max_tokens": 32000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/reka-flash": { "litellm_provider": "snowflake", "max_input_tokens": 100000, "max_output_tokens": 8192, - "max_tokens": 100000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/snowflake-arctic": { "litellm_provider": "snowflake", "max_input_tokens": 4096, "max_output_tokens": 8192, - "max_tokens": 4096, + "max_tokens": 8192, "mode": "chat" }, "snowflake/snowflake-llama-3.1-405b": { "litellm_provider": "snowflake", "max_input_tokens": 8000, "max_output_tokens": 8192, - "max_tokens": 8000, + "max_tokens": 8192, "mode": "chat" }, "snowflake/snowflake-llama-3.3-70b": { "litellm_provider": "snowflake", "max_input_tokens": 8000, "max_output_tokens": 8192, - "max_tokens": 8000, + "max_tokens": 8192, "mode": "chat" }, "stability/sd3": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.065, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/sd3-large": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.065, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/sd3-large-turbo": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.04, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/sd3-medium": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.035, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/sd3.5-large": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.065, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/sd3.5-large-turbo": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.04, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/sd3.5-medium": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.035, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/stable-image-ultra": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.08, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability/inpaint": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/outpaint": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.004, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/erase": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/search-and-replace": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/search-and-recolor": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/remove-background": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/replace-background-and-relight": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.008, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/sketch": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/structure": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/style": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.005, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/style-transfer": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.008, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/fast": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.002, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/conservative": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.04, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/creative": { "litellm_provider": "stability", "mode": "image_edit", "output_cost_per_image": 0.06, - "supported_endpoints": ["/v1/images/edits"] + "supported_endpoints": [ + "/v1/images/edits" + ] }, "stability/stable-image-core": { "litellm_provider": "stability", "mode": "image_generation", "output_cost_per_image": 0.03, - "supported_endpoints": ["/v1/images/generations"] + "supported_endpoints": [ + "/v1/images/generations" + ] }, "stability.sd3-5-large-v1:0": { "litellm_provider": "bedrock", @@ -25090,13 +25135,13 @@ "litellm_provider": "bedrock", "max_input_tokens": 77, "mode": "image_edit", - "output_cost_per_image": 0.40 + "output_cost_per_image": 0.4 }, "stability.stable-creative-upscale-v1:0": { "litellm_provider": "bedrock", "max_input_tokens": 77, "mode": "image_edit", - "output_cost_per_image": 0.60 + "output_cost_per_image": 0.6 }, "stability.stable-fast-upscale-v1:0": { "litellm_provider": "bedrock", @@ -25204,12 +25249,12 @@ "output_cost_per_pixel": 0.0 }, "linkup/search": { - "input_cost_per_query": 5.87e-03, + "input_cost_per_query": 0.00587, "litellm_provider": "linkup", "mode": "search" }, "linkup/search-deep": { - "input_cost_per_query": 58.67e-03, + "input_cost_per_query": 0.05867, "litellm_provider": "linkup", "mode": "search" }, @@ -25388,7 +25433,7 @@ "litellm_provider": "openai", "max_input_tokens": 32768, "max_output_tokens": 0, - "max_tokens": 32768, + "max_tokens": 0, "mode": "moderation", "output_cost_per_token": 0.0 }, @@ -25397,7 +25442,7 @@ "litellm_provider": "openai", "max_input_tokens": 32768, "max_output_tokens": 0, - "max_tokens": 32768, + "max_tokens": 0, "mode": "moderation", "output_cost_per_token": 0.0 }, @@ -25406,7 +25451,7 @@ "litellm_provider": "openai", "max_input_tokens": 32768, "max_output_tokens": 0, - "max_tokens": 32768, + "max_tokens": 0, "mode": "moderation", "output_cost_per_token": 0.0 }, @@ -25842,7 +25887,7 @@ "supports_tool_choice": true }, "together_ai/zai-org/GLM-4.6": { - "input_cost_per_token": 0.6e-06, + "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 200000, "max_output_tokens": 200000, @@ -25925,7 +25970,7 @@ "source": "https://aws.amazon.com/polly/pricing/" }, "aws_polly/long-form": { - "input_cost_per_character": 1e-04, + "input_cost_per_character": 0.0001, "litellm_provider": "aws_polly", "mode": "audio_speech", "supported_endpoints": [ @@ -26357,7 +26402,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.6e-05, "supports_function_calling": true, @@ -26368,7 +26413,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 9.9e-07, "supports_function_calling": true, @@ -26379,7 +26424,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 2.2e-07, "supports_function_calling": true, @@ -26390,7 +26435,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 3.5e-07, "supports_function_calling": true, @@ -26402,7 +26447,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1e-07, "supports_function_calling": true, @@ -26413,7 +26458,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-07, "supports_function_calling": true, @@ -26424,7 +26469,7 @@ "litellm_provider": "bedrock", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, "supports_function_calling": true, @@ -26489,7 +26534,7 @@ "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-06, "supports_function_calling": true, @@ -26542,7 +26587,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 40960, "max_output_tokens": 16384, - "max_tokens": 40960, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 2.4e-07 }, @@ -26551,7 +26596,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 40960, "max_output_tokens": 16384, - "max_tokens": 40960, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 6e-07 }, @@ -26560,7 +26605,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 40960, "max_output_tokens": 16384, - "max_tokens": 40960, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3e-07 }, @@ -26569,7 +26614,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 40960, "max_output_tokens": 16384, - "max_tokens": 40960, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3e-07 }, @@ -26578,7 +26623,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 262144, "max_output_tokens": 66536, - "max_tokens": 262144, + "max_tokens": 66536, "mode": "chat", "output_cost_per_token": 1.6e-06 }, @@ -26587,7 +26632,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 300000, "max_output_tokens": 8192, - "max_tokens": 300000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.4e-07 }, @@ -26596,7 +26641,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.4e-07 }, @@ -26605,7 +26650,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 300000, "max_output_tokens": 8192, - "max_tokens": 300000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 3.2e-06 }, @@ -26625,7 +26670,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 4096, - "max_tokens": 200000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.25e-06 }, @@ -26636,7 +26681,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 4096, - "max_tokens": 200000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 7.5e-05 }, @@ -26647,7 +26692,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 8192, - "max_tokens": 200000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 4e-06 }, @@ -26658,7 +26703,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 8192, - "max_tokens": 200000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.5e-05 }, @@ -26669,7 +26714,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 64000, - "max_tokens": 200000, + "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05 }, @@ -26680,7 +26725,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 32000, - "max_tokens": 200000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05 }, @@ -26691,7 +26736,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 64000, - "max_tokens": 200000, + "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05 }, @@ -26700,7 +26745,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 256000, "max_output_tokens": 8000, - "max_tokens": 256000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1e-05 }, @@ -26709,7 +26754,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-07 }, @@ -26718,7 +26763,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1e-05 }, @@ -26736,7 +26781,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.19e-06 }, @@ -26754,7 +26799,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 9e-07 }, @@ -26763,7 +26808,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 1048576, "max_output_tokens": 8192, - "max_tokens": 1048576, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 6e-07 }, @@ -26772,7 +26817,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 1048576, "max_output_tokens": 8192, - "max_tokens": 1048576, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 3e-07 }, @@ -26781,7 +26826,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 1000000, "max_output_tokens": 65536, - "max_tokens": 1000000, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 2.5e-06 }, @@ -26790,7 +26835,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 1048576, "max_output_tokens": 65536, - "max_tokens": 1048576, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 1e-05 }, @@ -26835,7 +26880,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 32000, "max_output_tokens": 16384, - "max_tokens": 32000, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1e-06 }, @@ -26862,7 +26907,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07 }, @@ -26871,7 +26916,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 131000, "max_output_tokens": 131072, - "max_tokens": 131000, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 8e-08 }, @@ -26880,7 +26925,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.6e-07 }, @@ -26889,7 +26934,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1e-07 }, @@ -26898,7 +26943,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.5e-07 }, @@ -26907,7 +26952,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07 }, @@ -26916,7 +26961,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 8192, - "max_tokens": 128000, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07 }, @@ -26925,7 +26970,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 131072, "max_output_tokens": 8192, - "max_tokens": 131072, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 6e-07 }, @@ -26934,7 +26979,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 131072, "max_output_tokens": 8192, - "max_tokens": 131072, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 3e-07 }, @@ -26943,7 +26988,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 256000, "max_output_tokens": 4000, - "max_tokens": 256000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 9e-07 }, @@ -26970,7 +27015,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 64000, - "max_tokens": 128000, + "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06 }, @@ -26979,7 +27024,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 64000, - "max_tokens": 128000, + "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-06 }, @@ -26988,7 +27033,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 4e-08 }, @@ -26997,7 +27042,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 1e-07 }, @@ -27015,7 +27060,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 32000, "max_output_tokens": 4000, - "max_tokens": 32000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 6e-06 }, @@ -27033,7 +27078,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 32000, "max_output_tokens": 4000, - "max_tokens": 32000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 3e-07 }, @@ -27042,7 +27087,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 65536, "max_output_tokens": 2048, - "max_tokens": 65536, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 1.2e-06 }, @@ -27051,7 +27096,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 1.5e-07 }, @@ -27060,7 +27105,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 4000, - "max_tokens": 128000, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 6e-06 }, @@ -27069,7 +27114,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 131072, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 2.2e-06 }, @@ -27078,7 +27123,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 32768, "max_output_tokens": 16384, - "max_tokens": 32768, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06 }, @@ -27087,7 +27132,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 32768, "max_output_tokens": 16384, - "max_tokens": 32768, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.9e-06 }, @@ -27096,7 +27141,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 16385, "max_output_tokens": 4096, - "max_tokens": 16385, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-06 }, @@ -27105,7 +27150,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 8192, "max_output_tokens": 4096, - "max_tokens": 8192, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06 }, @@ -27114,7 +27159,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 4096, - "max_tokens": 128000, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 3e-05 }, @@ -27125,7 +27170,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 1047576, "max_output_tokens": 32768, - "max_tokens": 1047576, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 8e-06 }, @@ -27136,7 +27181,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 1047576, "max_output_tokens": 32768, - "max_tokens": 1047576, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 1.6e-06 }, @@ -27147,7 +27192,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 1047576, "max_output_tokens": 32768, - "max_tokens": 1047576, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-07 }, @@ -27158,7 +27203,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 16384, - "max_tokens": 128000, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1e-05 }, @@ -27169,7 +27214,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 16384, - "max_tokens": 128000, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 6e-07 }, @@ -27180,7 +27225,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 100000, - "max_tokens": 200000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 6e-05 }, @@ -27191,7 +27236,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 100000, - "max_tokens": 200000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 8e-06 }, @@ -27202,7 +27247,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 100000, - "max_tokens": 200000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06 }, @@ -27213,7 +27258,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 100000, - "max_tokens": 200000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06 }, @@ -27249,7 +27294,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 127000, "max_output_tokens": 8000, - "max_tokens": 127000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1e-06 }, @@ -27258,7 +27303,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 200000, "max_output_tokens": 8000, - "max_tokens": 200000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.5e-05 }, @@ -27267,7 +27312,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 127000, "max_output_tokens": 8000, - "max_tokens": 127000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 5e-06 }, @@ -27276,7 +27321,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 127000, "max_output_tokens": 8000, - "max_tokens": 127000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 8e-06 }, @@ -27285,7 +27330,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 32000, - "max_tokens": 128000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 1.5e-05 }, @@ -27294,7 +27339,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 32768, - "max_tokens": 128000, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 1.5e-05 }, @@ -27303,7 +27348,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 131072, "max_output_tokens": 4000, - "max_tokens": 131072, + "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 1e-05 }, @@ -27375,7 +27420,7 @@ "litellm_provider": "vercel_ai_gateway", "max_input_tokens": 128000, "max_output_tokens": 96000, - "max_tokens": 128000, + "max_tokens": 96000, "mode": "chat", "output_cost_per_token": 1.1e-06 }, @@ -27394,7 +27439,7 @@ "supports_tool_choice": true }, "vertex_ai/chirp": { - "input_cost_per_character": 30e-06, + "input_cost_per_character": 3e-05, "litellm_provider": "vertex_ai", "mode": "audio_speech", "source": "https://cloud.google.com/text-to-speech/pricing", @@ -27938,7 +27983,7 @@ "litellm_provider": "vertex_ai-deepseek_models", "max_input_tokens": 163840, "max_output_tokens": 32768, - "max_tokens": 163840, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 5.4e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", @@ -27957,7 +28002,7 @@ "litellm_provider": "vertex_ai-deepseek_models", "max_input_tokens": 163840, "max_output_tokens": 32768, - "max_tokens": 163840, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 1.68e-06, "output_cost_per_token_batches": 8.4e-07, @@ -28042,10 +28087,10 @@ "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, - "max_tokens": 65536, + "max_tokens": 32768, "mode": "image_generation", "output_cost_per_image": 0.134, - "output_cost_per_image_token": 1.2e-04, + "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" @@ -28154,7 +28199,7 @@ "litellm_provider": "vertex_ai-llama_models", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 1.6e-05, "source": "https://console.cloud.google.com/vertex-ai/publishers/meta/model-garden/llama-3.2-90b-vision-instruct-maas", @@ -28167,7 +28212,7 @@ "litellm_provider": "vertex_ai-llama_models", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 0.0, "source": "https://console.cloud.google.com/vertex-ai/publishers/meta/model-garden/llama-3.2-90b-vision-instruct-maas", @@ -28180,7 +28225,7 @@ "litellm_provider": "vertex_ai-llama_models", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "metadata": { "notes": "VertexAI states that The Llama 3.1 API service for llama-3.1-70b-instruct-maas and llama-3.1-8b-instruct-maas are in public preview and at no cost." }, @@ -28196,7 +28241,7 @@ "litellm_provider": "vertex_ai-llama_models", "max_input_tokens": 128000, "max_output_tokens": 2048, - "max_tokens": 128000, + "max_tokens": 2048, "metadata": { "notes": "VertexAI states that The Llama 3.2 API service is at no cost during public preview, and will be priced as per dollar-per-1M-tokens at GA." }, @@ -28345,6 +28390,19 @@ "supports_tool_choice": true, "supports_web_search": true }, + "vertex_ai/zai-org/glm-4.7-maas": { + "input_cost_per_token": 3e-07, + "litellm_provider": "vertex_ai-zai_models", + "max_input_tokens": 200000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "vertex_ai/mistral-medium-3": { "input_cost_per_token": 4e-07, "litellm_provider": "vertex_ai-mistral_models", @@ -28481,7 +28539,7 @@ "vertex_ai/mistral-ocr-2505": { "litellm_provider": "vertex_ai", "mode": "ocr", - "ocr_cost_per_page": 5e-4, + "ocr_cost_per_page": 0.0005, "supported_endpoints": [ "/v1/ocr" ], @@ -28492,7 +28550,7 @@ "mode": "ocr", "input_cost_per_token": 3e-07, "output_cost_per_token": 1.2e-06, - "ocr_cost_per_page": 3e-04, + "ocr_cost_per_page": 0.0003, "source": "https://cloud.google.com/vertex-ai/pricing" }, "vertex_ai/openai/gpt-oss-120b-maas": { @@ -28980,13 +29038,13 @@ "mode": "chat" }, "watsonx/ibm/granite-3-8b-instruct": { - "input_cost_per_token": 0.2e-06, + "input_cost_per_token": 2e-07, "litellm_provider": "watsonx", "max_input_tokens": 8192, "max_output_tokens": 1024, - "max_tokens": 8192, + "max_tokens": 1024, "mode": "chat", - "output_cost_per_token": 0.2e-06, + "output_cost_per_token": 2e-07, "supports_audio_input": false, "supports_audio_output": false, "supports_function_calling": true, @@ -29002,9 +29060,9 @@ "litellm_provider": "watsonx", "max_input_tokens": 131072, "max_output_tokens": 16384, - "max_tokens": 131072, + "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 10e-06, + "output_cost_per_token": 1e-05, "supports_audio_input": false, "supports_audio_output": false, "supports_function_calling": true, @@ -29043,8 +29101,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.6e-06, - "output_cost_per_token": 0.6e-06, + "input_cost_per_token": 6e-07, + "output_cost_per_token": 6e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29055,8 +29113,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.6e-06, - "output_cost_per_token": 0.6e-06, + "input_cost_per_token": 6e-07, + "output_cost_per_token": 6e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29067,8 +29125,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.6e-06, - "output_cost_per_token": 0.6e-06, + "input_cost_per_token": 6e-07, + "output_cost_per_token": 6e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29079,8 +29137,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.2e-06, - "output_cost_per_token": 0.2e-06, + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29091,8 +29149,8 @@ "max_tokens": 20480, "max_input_tokens": 20480, "max_output_tokens": 20480, - "input_cost_per_token": 0.06e-06, - "output_cost_per_token": 0.25e-06, + "input_cost_per_token": 6e-08, + "output_cost_per_token": 2.5e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29103,8 +29161,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.1e-06, - "output_cost_per_token": 0.1e-06, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 1e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29115,8 +29173,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.2e-06, - "output_cost_per_token": 0.2e-06, + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29127,8 +29185,8 @@ "max_tokens": 512, "max_input_tokens": 512, "max_output_tokens": 512, - "input_cost_per_token": 0.38e-06, - "output_cost_per_token": 0.38e-06, + "input_cost_per_token": 3.8e-07, + "output_cost_per_token": 3.8e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29139,8 +29197,8 @@ "max_tokens": 512, "max_input_tokens": 512, "max_output_tokens": 512, - "input_cost_per_token": 0.38e-06, - "output_cost_per_token": 0.38e-06, + "input_cost_per_token": 3.8e-07, + "output_cost_per_token": 3.8e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29151,8 +29209,8 @@ "max_tokens": 512, "max_input_tokens": 512, "max_output_tokens": 512, - "input_cost_per_token": 0.38e-06, - "output_cost_per_token": 0.38e-06, + "input_cost_per_token": 3.8e-07, + "output_cost_per_token": 3.8e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29163,8 +29221,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.1e-06, - "output_cost_per_token": 0.1e-06, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 1e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29175,8 +29233,8 @@ "max_tokens": 128000, "max_input_tokens": 128000, "max_output_tokens": 128000, - "input_cost_per_token": 0.35e-06, - "output_cost_per_token": 0.35e-06, + "input_cost_per_token": 3.5e-07, + "output_cost_per_token": 3.5e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29187,8 +29245,8 @@ "max_tokens": 128000, "max_input_tokens": 128000, "max_output_tokens": 128000, - "input_cost_per_token": 0.1e-06, - "output_cost_per_token": 0.1e-06, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 1e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29199,8 +29257,8 @@ "max_tokens": 128000, "max_input_tokens": 128000, "max_output_tokens": 128000, - "input_cost_per_token": 0.15e-06, - "output_cost_per_token": 0.15e-06, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 1.5e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29223,8 +29281,8 @@ "max_tokens": 128000, "max_input_tokens": 128000, "max_output_tokens": 128000, - "input_cost_per_token": 0.71e-06, - "output_cost_per_token": 0.71e-06, + "input_cost_per_token": 7.1e-07, + "output_cost_per_token": 7.1e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29235,7 +29293,7 @@ "max_tokens": 128000, "max_input_tokens": 128000, "max_output_tokens": 128000, - "input_cost_per_token": 0.35e-06, + "input_cost_per_token": 3.5e-07, "output_cost_per_token": 1.4e-06, "litellm_provider": "watsonx", "mode": "chat", @@ -29247,8 +29305,8 @@ "max_tokens": 128000, "max_input_tokens": 128000, "max_output_tokens": 128000, - "input_cost_per_token": 0.35e-06, - "output_cost_per_token": 0.35e-06, + "input_cost_per_token": 3.5e-07, + "output_cost_per_token": 3.5e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29260,7 +29318,7 @@ "max_input_tokens": 128000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, - "output_cost_per_token": 10e-06, + "output_cost_per_token": 1e-05, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29271,8 +29329,8 @@ "max_tokens": 32000, "max_input_tokens": 32000, "max_output_tokens": 32000, - "input_cost_per_token": 0.1e-06, - "output_cost_per_token": 0.3e-06, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 3e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29283,8 +29341,8 @@ "max_tokens": 32000, "max_input_tokens": 32000, "max_output_tokens": 32000, - "input_cost_per_token": 0.1e-06, - "output_cost_per_token": 0.3e-06, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 3e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": true, @@ -29295,8 +29353,8 @@ "max_tokens": 128000, "max_input_tokens": 128000, "max_output_tokens": 128000, - "input_cost_per_token": 0.35e-06, - "output_cost_per_token": 0.35e-06, + "input_cost_per_token": 3.5e-07, + "output_cost_per_token": 3.5e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29307,8 +29365,8 @@ "max_tokens": 8192, "max_input_tokens": 8192, "max_output_tokens": 8192, - "input_cost_per_token": 0.15e-06, - "output_cost_per_token": 0.6e-06, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, "litellm_provider": "watsonx", "mode": "chat", "supports_function_calling": false, @@ -29598,15 +29656,15 @@ }, "xai/grok-4-fast-reasoning": { "litellm_provider": "xai", - "max_input_tokens": 2e6, - "max_output_tokens": 2e6, - "max_tokens": 2e6, + "max_input_tokens": 2000000.0, + "max_output_tokens": 2000000.0, + "max_tokens": 2000000.0, "mode": "chat", - "input_cost_per_token": 0.2e-06, - "input_cost_per_token_above_128k_tokens": 0.4e-06, - "output_cost_per_token": 0.5e-06, + "input_cost_per_token": 2e-07, + "input_cost_per_token_above_128k_tokens": 4e-07, + "output_cost_per_token": 5e-07, "output_cost_per_token_above_128k_tokens": 1e-06, - "cache_read_input_token_cost": 0.05e-06, + "cache_read_input_token_cost": 5e-08, "source": "https://docs.x.ai/docs/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -29614,14 +29672,14 @@ }, "xai/grok-4-fast-non-reasoning": { "litellm_provider": "xai", - "max_input_tokens": 2e6, - "max_output_tokens": 2e6, - "cache_read_input_token_cost": 0.05e-06, - "max_tokens": 2e6, + "max_input_tokens": 2000000.0, + "max_output_tokens": 2000000.0, + "cache_read_input_token_cost": 5e-08, + "max_tokens": 2000000.0, "mode": "chat", - "input_cost_per_token": 0.2e-06, - "input_cost_per_token_above_128k_tokens": 0.4e-06, - "output_cost_per_token": 0.5e-06, + "input_cost_per_token": 2e-07, + "input_cost_per_token_above_128k_tokens": 4e-07, + "output_cost_per_token": 5e-07, "output_cost_per_token_above_128k_tokens": 1e-06, "source": "https://docs.x.ai/docs/models", "supports_function_calling": true, @@ -29637,7 +29695,7 @@ "max_tokens": 256000, "mode": "chat", "output_cost_per_token": 1.5e-05, - "output_cost_per_token_above_128k_tokens": 30e-06, + "output_cost_per_token_above_128k_tokens": 3e-05, "source": "https://docs.x.ai/docs/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -29652,22 +29710,22 @@ "max_tokens": 256000, "mode": "chat", "output_cost_per_token": 1.5e-05, - "output_cost_per_token_above_128k_tokens": 30e-06, + "output_cost_per_token_above_128k_tokens": 3e-05, "source": "https://docs.x.ai/docs/models", "supports_function_calling": true, "supports_tool_choice": true, "supports_web_search": true }, "xai/grok-4-1-fast": { - "cache_read_input_token_cost": 0.05e-06, - "input_cost_per_token": 0.2e-06, - "input_cost_per_token_above_128k_tokens": 0.4e-06, + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "input_cost_per_token_above_128k_tokens": 4e-07, "litellm_provider": "xai", - "max_input_tokens": 2e6, - "max_output_tokens": 2e6, - "max_tokens": 2e6, + "max_input_tokens": 2000000.0, + "max_output_tokens": 2000000.0, + "max_tokens": 2000000.0, "mode": "chat", - "output_cost_per_token": 0.5e-06, + "output_cost_per_token": 5e-07, "output_cost_per_token_above_128k_tokens": 1e-06, "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", "supports_audio_input": true, @@ -29679,15 +29737,15 @@ "supports_web_search": true }, "xai/grok-4-1-fast-reasoning": { - "cache_read_input_token_cost": 0.05e-06, - "input_cost_per_token": 0.2e-06, - "input_cost_per_token_above_128k_tokens": 0.4e-06, + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "input_cost_per_token_above_128k_tokens": 4e-07, "litellm_provider": "xai", - "max_input_tokens": 2e6, - "max_output_tokens": 2e6, - "max_tokens": 2e6, + "max_input_tokens": 2000000.0, + "max_output_tokens": 2000000.0, + "max_tokens": 2000000.0, "mode": "chat", - "output_cost_per_token": 0.5e-06, + "output_cost_per_token": 5e-07, "output_cost_per_token_above_128k_tokens": 1e-06, "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", "supports_audio_input": true, @@ -29699,15 +29757,15 @@ "supports_web_search": true }, "xai/grok-4-1-fast-reasoning-latest": { - "cache_read_input_token_cost": 0.05e-06, - "input_cost_per_token": 0.2e-06, - "input_cost_per_token_above_128k_tokens": 0.4e-06, + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "input_cost_per_token_above_128k_tokens": 4e-07, "litellm_provider": "xai", - "max_input_tokens": 2e6, - "max_output_tokens": 2e6, - "max_tokens": 2e6, + "max_input_tokens": 2000000.0, + "max_output_tokens": 2000000.0, + "max_tokens": 2000000.0, "mode": "chat", - "output_cost_per_token": 0.5e-06, + "output_cost_per_token": 5e-07, "output_cost_per_token_above_128k_tokens": 1e-06, "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", "supports_audio_input": true, @@ -29719,15 +29777,15 @@ "supports_web_search": true }, "xai/grok-4-1-fast-non-reasoning": { - "cache_read_input_token_cost": 0.05e-06, - "input_cost_per_token": 0.2e-06, - "input_cost_per_token_above_128k_tokens": 0.4e-06, + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "input_cost_per_token_above_128k_tokens": 4e-07, "litellm_provider": "xai", - "max_input_tokens": 2e6, - "max_output_tokens": 2e6, - "max_tokens": 2e6, + "max_input_tokens": 2000000.0, + "max_output_tokens": 2000000.0, + "max_tokens": 2000000.0, "mode": "chat", - "output_cost_per_token": 0.5e-06, + "output_cost_per_token": 5e-07, "output_cost_per_token_above_128k_tokens": 1e-06, "source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning", "supports_audio_input": true, @@ -29738,15 +29796,15 @@ "supports_web_search": true }, "xai/grok-4-1-fast-non-reasoning-latest": { - "cache_read_input_token_cost": 0.05e-06, - "input_cost_per_token": 0.2e-06, - "input_cost_per_token_above_128k_tokens": 0.4e-06, + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 2e-07, + "input_cost_per_token_above_128k_tokens": 4e-07, "litellm_provider": "xai", - "max_input_tokens": 2e6, - "max_output_tokens": 2e6, - "max_tokens": 2e6, + "max_input_tokens": 2000000.0, + "max_output_tokens": 2000000.0, + "max_tokens": 2000000.0, "mode": "chat", - "output_cost_per_token": 0.5e-06, + "output_cost_per_token": 5e-07, "output_cost_per_token_above_128k_tokens": 1e-06, "source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning", "supports_audio_input": true, @@ -29929,7 +29987,7 @@ "source": "https://docs.z.ai/guides/overview/pricing" }, "vertex_ai/search_api": { - "input_cost_per_query": 1.5e-03, + "input_cost_per_query": 0.0015, "litellm_provider": "vertex_ai", "mode": "vector_store" }, @@ -29941,7 +29999,7 @@ "openai/sora-2": { "litellm_provider": "openai", "mode": "video_generation", - "output_cost_per_video_per_second": 0.10, + "output_cost_per_video_per_second": 0.1, "source": "https://platform.openai.com/docs/api-reference/videos", "supported_modalities": [ "text", @@ -29958,7 +30016,7 @@ "openai/sora-2-pro": { "litellm_provider": "openai", "mode": "video_generation", - "output_cost_per_video_per_second": 0.30, + "output_cost_per_video_per_second": 0.3, "source": "https://platform.openai.com/docs/api-reference/videos", "supported_modalities": [ "text", @@ -29975,7 +30033,7 @@ "azure/sora-2": { "litellm_provider": "azure", "mode": "video_generation", - "output_cost_per_video_per_second": 0.10, + "output_cost_per_video_per_second": 0.1, "source": "https://azure.microsoft.com/en-us/products/ai-services/video-generation", "supported_modalities": [ "text" @@ -29991,7 +30049,7 @@ "azure/sora-2-pro": { "litellm_provider": "azure", "mode": "video_generation", - "output_cost_per_video_per_second": 0.30, + "output_cost_per_video_per_second": 0.3, "source": "https://azure.microsoft.com/en-us/products/ai-services/video-generation", "supported_modalities": [ "text" @@ -30007,7 +30065,7 @@ "azure/sora-2-pro-high-res": { "litellm_provider": "azure", "mode": "video_generation", - "output_cost_per_video_per_second": 0.50, + "output_cost_per_video_per_second": 0.5, "source": "https://azure.microsoft.com/en-us/products/ai-services/video-generation", "supported_modalities": [ "text" @@ -32178,6 +32236,1100 @@ "litellm_provider": "fireworks_ai", "mode": "chat" }, + "novita/deepseek/deepseek-v3.2": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.69e-03, + "output_cost_per_token": 4e-03, + "max_input_tokens": 163840, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 1.345e-03, + "input_cost_per_token_cache_hit": 1.345e-03, + "supports_reasoning": true + }, + "novita/minimax/minimax-m2.1": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3e-03, + "output_cost_per_token": 1.2e-02, + "max_input_tokens": 204800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 3e-04, + "input_cost_per_token_cache_hit": 3e-04, + "supports_reasoning": true + }, + "novita/zai-org/glm-4.7": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 6e-03, + "output_cost_per_token": 2.2e-02, + "max_input_tokens": 204800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 1.1e-03, + "input_cost_per_token_cache_hit": 1.1e-03, + "supports_reasoning": true + }, + "novita/xiaomimimo/mimo-v2-flash": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1e-03, + "output_cost_per_token": 3e-03, + "max_input_tokens": 262144, + "max_output_tokens": 32000, + "max_tokens": 32000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 2e-04, + "input_cost_per_token_cache_hit": 2e-04, + "supports_reasoning": true + }, + "novita/zai-org/autoglm-phone-9b-multilingual": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.5e-04, + "output_cost_per_token": 1.38e-03, + "max_input_tokens": 65536, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_vision": true, + "supports_system_messages": true + }, + "novita/moonshotai/kimi-k2-thinking": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.8e-03, + "output_cost_per_token": 2e-02, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/minimax/minimax-m2": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-03, + "output_cost_per_token": 9.6e-03, + "max_input_tokens": 204800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "cache_read_input_token_cost": 2.4e-04, + "input_cost_per_token_cache_hit": 2.4e-04, + "supports_reasoning": true + }, + "novita/paddlepaddle/paddleocr-vl": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.6e-04, + "output_cost_per_token": 1.6e-04, + "max_input_tokens": 16384, + "max_output_tokens": 16384, + "max_tokens": 16384, + "supports_vision": true, + "supports_system_messages": true + }, + "novita/deepseek/deepseek-v3.2-exp": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.16e-03, + "output_cost_per_token": 3.28e-03, + "max_input_tokens": 163840, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/qwen/qwen3-vl-235b-a22b-thinking": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 7.84e-03, + "output_cost_per_token": 3.16e-02, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_vision": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/zai-org/glm-4.6v": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3e-03, + "output_cost_per_token": 9e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 5.5e-04, + "input_cost_per_token_cache_hit": 5.5e-04, + "supports_reasoning": true + }, + "novita/zai-org/glm-4.6": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.4e-03, + "output_cost_per_token": 1.76e-02, + "max_input_tokens": 204800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 8.8e-04, + "input_cost_per_token_cache_hit": 8.8e-04, + "supports_reasoning": true + }, + "novita/qwen/qwen3-next-80b-a3b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.2e-03, + "output_cost_per_token": 1.2e-02, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen3-next-80b-a3b-thinking": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.2e-03, + "output_cost_per_token": 1.2e-02, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/deepseek/deepseek-ocr": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-04, + "output_cost_per_token": 2.4e-04, + "max_input_tokens": 8192, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/deepseek/deepseek-v3.1-terminus": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.16e-03, + "output_cost_per_token": 8e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 1.08e-03, + "input_cost_per_token_cache_hit": 1.08e-03, + "supports_reasoning": true + }, + "novita/qwen/qwen3-vl-235b-a22b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-03, + "output_cost_per_token": 1.2e-02, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen3-max": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.688e-02, + "output_cost_per_token": 6.76e-02, + "max_input_tokens": 262144, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/skywork/r1v4-lite": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2e-03, + "output_cost_per_token": 6e-03, + "max_input_tokens": 262144, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/deepseek/deepseek-v3.1": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.16e-03, + "output_cost_per_token": 8e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 1.08e-03, + "input_cost_per_token_cache_hit": 1.08e-03, + "supports_reasoning": true + }, + "novita/moonshotai/kimi-k2-0905": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.8e-03, + "output_cost_per_token": 2e-02, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen3-coder-480b-a35b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-03, + "output_cost_per_token": 1.04e-02, + "max_input_tokens": 262144, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen3-coder-30b-a3b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 7e-04, + "output_cost_per_token": 2.7e-03, + "max_input_tokens": 160000, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/openai/gpt-oss-120b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4e-04, + "output_cost_per_token": 2e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/moonshotai/kimi-k2-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.56e-03, + "output_cost_per_token": 1.84e-02, + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/deepseek/deepseek-v3-0324": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.16e-03, + "output_cost_per_token": 8.96e-03, + "max_input_tokens": 163840, + "max_output_tokens": 163840, + "max_tokens": 163840, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 1.08e-03, + "input_cost_per_token_cache_hit": 1.08e-03 + }, + "novita/zai-org/glm-4.5": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.8e-03, + "output_cost_per_token": 1.76e-02, + "max_input_tokens": 131072, + "max_output_tokens": 98304, + "max_tokens": 98304, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "cache_read_input_token_cost": 8.8e-04, + "input_cost_per_token_cache_hit": 8.8e-04, + "supports_reasoning": true + }, + "novita/qwen/qwen3-235b-a22b-thinking-2507": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-03, + "output_cost_per_token": 2.4e-02, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/meta-llama/llama-3.1-8b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2e-04, + "output_cost_per_token": 5e-04, + "max_input_tokens": 16384, + "max_output_tokens": 16384, + "max_tokens": 16384, + "supports_system_messages": true + }, + "novita/google/gemma-3-12b-it": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4e-04, + "output_cost_per_token": 8e-04, + "max_input_tokens": 131072, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/zai-org/glm-4.5v": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.8e-03, + "output_cost_per_token": 1.44e-02, + "max_input_tokens": 65536, + "max_output_tokens": 16384, + "max_tokens": 16384, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 8.8e-04, + "input_cost_per_token_cache_hit": 8.8e-04, + "supports_reasoning": true + }, + "novita/openai/gpt-oss-20b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.2e-04, + "output_cost_per_token": 1.2e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/qwen/qwen3-235b-a22b-instruct-2507": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 7.2e-04, + "output_cost_per_token": 4.64e-03, + "max_input_tokens": 131072, + "max_output_tokens": 16384, + "max_tokens": 16384, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/deepseek/deepseek-r1-distill-qwen-14b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.2e-03, + "output_cost_per_token": 1.2e-03, + "max_input_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/meta-llama/llama-3.3-70b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.08e-03, + "output_cost_per_token": 3.2e-03, + "max_input_tokens": 131072, + "max_output_tokens": 120000, + "max_tokens": 120000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true + }, + "novita/qwen/qwen-2.5-72b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.04e-03, + "output_cost_per_token": 3.2e-03, + "max_input_tokens": 32000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/mistralai/mistral-nemo": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.2e-04, + "output_cost_per_token": 1.36e-03, + "max_input_tokens": 60288, + "max_output_tokens": 16000, + "max_tokens": 16000, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/minimaxai/minimax-m1-80k": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.4e-03, + "output_cost_per_token": 1.76e-02, + "max_input_tokens": 1000000, + "max_output_tokens": 40000, + "max_tokens": 40000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/deepseek/deepseek-r1-0528": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.6e-03, + "output_cost_per_token": 2e-02, + "max_input_tokens": 163840, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "cache_read_input_token_cost": 2.8e-03, + "input_cost_per_token_cache_hit": 2.8e-03, + "supports_reasoning": true + }, + "novita/deepseek/deepseek-r1-distill-qwen-32b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-03, + "output_cost_per_token": 2.4e-03, + "max_input_tokens": 64000, + "max_output_tokens": 32000, + "max_tokens": 32000, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/meta-llama/llama-3-8b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.2e-04, + "output_cost_per_token": 3.2e-04, + "max_input_tokens": 8192, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_system_messages": true + }, + "novita/microsoft/wizardlm-2-8x22b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.96e-03, + "output_cost_per_token": 4.96e-03, + "max_input_tokens": 65535, + "max_output_tokens": 8000, + "max_tokens": 8000, + "supports_system_messages": true + }, + "novita/deepseek/deepseek-r1-0528-qwen3-8b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 4.8e-04, + "output_cost_per_token": 7.2e-04, + "max_input_tokens": 128000, + "max_output_tokens": 32000, + "max_tokens": 32000, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/deepseek/deepseek-r1-distill-llama-70b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 6.4e-03, + "output_cost_per_token": 6.4e-03, + "max_input_tokens": 8192, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/meta-llama/llama-3-70b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.1e-03, + "output_cost_per_token": 7.4e-03, + "max_input_tokens": 8192, + "max_output_tokens": 8000, + "max_tokens": 8000, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen3-235b-a22b-fp8": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.6e-03, + "output_cost_per_token": 6.4e-03, + "max_input_tokens": 40960, + "max_output_tokens": 20000, + "max_tokens": 20000, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/meta-llama/llama-4-maverick-17b-128e-instruct-fp8": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.6e-03, + "output_cost_per_token": 7.2e-03, + "max_input_tokens": 1048576, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_vision": true, + "supports_system_messages": true + }, + "novita/meta-llama/llama-4-scout-17b-16e-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 8e-04, + "output_cost_per_token": 4e-03, + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "supports_vision": true, + "supports_system_messages": true + }, + "novita/nousresearch/hermes-2-pro-llama-3-8b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.4e-03, + "output_cost_per_token": 1.4e-03, + "max_input_tokens": 8192, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen2.5-vl-72b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 6.4e-03, + "output_cost_per_token": 6.4e-03, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_vision": true, + "supports_system_messages": true + }, + "novita/sao10k/l3-70b-euryale-v2.1": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.48e-02, + "output_cost_per_token": 1.48e-02, + "max_input_tokens": 8192, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true + }, + "novita/baidu/ernie-4.5-21B-a3b-thinking": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.6e-04, + "output_cost_per_token": 2.24e-03, + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/sao10k/l3-8b-lunaris": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5e-04, + "output_cost_per_token": 5e-04, + "max_input_tokens": 8192, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/baichuan/baichuan-m2-32b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.6e-04, + "output_cost_per_token": 5.6e-04, + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/thudm/glm-4.1v-9b-thinking": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.8e-04, + "output_cost_per_token": 1.104e-03, + "max_input_tokens": 65536, + "max_output_tokens": 8000, + "max_tokens": 8000, + "supports_vision": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/baidu/ernie-4.5-vl-424b-a47b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.36e-03, + "output_cost_per_token": 1e-02, + "max_input_tokens": 123000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "supports_vision": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/baidu/ernie-4.5-300b-a47b-paddle": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.24e-03, + "output_cost_per_token": 8.8e-03, + "max_input_tokens": 123000, + "max_output_tokens": 12000, + "max_tokens": 12000, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/deepseek/deepseek-prover-v2-671b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.6e-03, + "output_cost_per_token": 2e-02, + "max_input_tokens": 160000, + "max_output_tokens": 160000, + "max_tokens": 160000, + "supports_system_messages": true + }, + "novita/qwen/qwen3-32b-fp8": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 8e-04, + "output_cost_per_token": 3.6e-03, + "max_input_tokens": 40960, + "max_output_tokens": 20000, + "max_tokens": 20000, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/qwen/qwen3-30b-a3b-fp8": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 7.2e-04, + "output_cost_per_token": 3.6e-03, + "max_input_tokens": 40960, + "max_output_tokens": 20000, + "max_tokens": 20000, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/google/gemma-3-27b-it": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 9.52e-04, + "output_cost_per_token": 1.6e-03, + "max_input_tokens": 98304, + "max_output_tokens": 16384, + "max_tokens": 16384, + "supports_vision": true, + "supports_system_messages": true + }, + "novita/deepseek/deepseek-v3-turbo": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.2e-03, + "output_cost_per_token": 1.04e-02, + "max_input_tokens": 64000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true + }, + "novita/deepseek/deepseek-r1-turbo": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.6e-03, + "output_cost_per_token": 2e-02, + "max_input_tokens": 64000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/Sao10K/L3-8B-Stheno-v3.2": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5e-04, + "output_cost_per_token": 5e-04, + "max_input_tokens": 8192, + "max_output_tokens": 32000, + "max_tokens": 32000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true + }, + "novita/gryphe/mythomax-l2-13b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 7.2e-04, + "output_cost_per_token": 7.2e-04, + "max_input_tokens": 4096, + "max_output_tokens": 3200, + "max_tokens": 3200, + "supports_system_messages": true + }, + "novita/baidu/ernie-4.5-vl-28b-a3b-thinking": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 3.9e-03, + "output_cost_per_token": 3.9e-03, + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true, + "supports_reasoning": true + }, + "novita/qwen/qwen3-vl-8b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 6.4e-04, + "output_cost_per_token": 4e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/zai-org/glm-4.5-air": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.04e-03, + "output_cost_per_token": 6.8e-03, + "max_input_tokens": 131072, + "max_output_tokens": 98304, + "max_tokens": 98304, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/qwen/qwen3-vl-30b-a3b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.6e-03, + "output_cost_per_token": 5.6e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen3-vl-30b-a3b-thinking": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.6e-03, + "output_cost_per_token": 8e-03, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/qwen/qwen-mt-plus": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2e-03, + "output_cost_per_token": 6e-03, + "max_input_tokens": 16384, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_system_messages": true + }, + "novita/baidu/ernie-4.5-vl-28b-a3b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.12e-03, + "output_cost_per_token": 4.48e-03, + "max_input_tokens": 30000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/baidu/ernie-4.5-21B-a3b": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.6e-04, + "output_cost_per_token": 2.24e-03, + "max_input_tokens": 120000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true + }, + "novita/qwen/qwen3-8b-fp8": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.8e-04, + "output_cost_per_token": 1.104e-03, + "max_input_tokens": 128000, + "max_output_tokens": 20000, + "max_tokens": 20000, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/qwen/qwen3-4b-fp8": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-04, + "output_cost_per_token": 2.4e-04, + "max_input_tokens": 128000, + "max_output_tokens": 20000, + "max_tokens": 20000, + "supports_system_messages": true, + "supports_reasoning": true + }, + "novita/qwen/qwen2.5-7b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 5.6e-04, + "output_cost_per_token": 5.6e-04, + "max_input_tokens": 32000, + "max_output_tokens": 32000, + "max_tokens": 32000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true, + "supports_response_schema": true + }, + "novita/meta-llama/llama-3.2-3b-instruct": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 2.4e-04, + "output_cost_per_token": 4e-04, + "max_input_tokens": 32768, + "max_output_tokens": 32000, + "max_tokens": 32000, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true + }, + "novita/sao10k/l31-70b-euryale-v2.2": { + "litellm_provider": "novita", + "mode": "chat", + "input_cost_per_token": 1.48e-02, + "output_cost_per_token": 1.48e-02, + "max_input_tokens": 8192, + "max_output_tokens": 8192, + "max_tokens": 8192, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_tool_choice": true, + "supports_system_messages": true + }, + "novita/qwen/qwen3-embedding-0.6b": { + "litellm_provider": "novita", + "mode": "embedding", + "input_cost_per_token": 5.6e-04, + "output_cost_per_token": 0, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768 + }, + "novita/qwen/qwen3-embedding-8b": { + "litellm_provider": "novita", + "mode": "embedding", + "input_cost_per_token": 5.6e-04, + "output_cost_per_token": 0, + "max_input_tokens": 32768, + "max_output_tokens": 4096, + "max_tokens": 4096 + }, + "novita/baai/bge-m3": { + "litellm_provider": "novita", + "mode": "embedding", + "input_cost_per_token": 1e-04, + "output_cost_per_token": 1e-04, + "max_input_tokens": 8192, + "max_output_tokens": 96000, + "max_tokens": 96000 + }, + "novita/qwen/qwen3-reranker-8b": { + "litellm_provider": "novita", + "mode": "rerank", + "input_cost_per_token": 4e-04, + "output_cost_per_token": 4e-04, + "max_input_tokens": 32768, + "max_output_tokens": 4096, + "max_tokens": 4096 + }, + "novita/baai/bge-reranker-v2-m3": { + "litellm_provider": "novita", + "mode": "rerank", + "input_cost_per_token": 1e-04, + "output_cost_per_token": 1e-04, + "max_input_tokens": 8000, + "max_output_tokens": 8000, + "max_tokens": 8000 + }, "llamagate/llama-3.1-8b": { "max_tokens": 8192, "max_input_tokens": 131072, @@ -32354,4 +33506,3 @@ "mode": "embedding" } } - diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index f671409175a..8432ce4e874 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -32,6 +32,23 @@ } }, "providers": { + "abliteration": { + "display_name": "Abliteration (`abliteration`)", + "url": "https://docs.litellm.ai/docs/providers/abliteration", + "endpoints": { + "chat_completions": true, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false + } + }, "aiml": { "display_name": "AI/ML API (`aiml`)", "url": "https://docs.litellm.ai/docs/providers/aiml", @@ -1559,7 +1576,7 @@ "chat_completions": true, "messages": true, "responses": true, - "embeddings": false, + "embeddings": true, "image_generations": false, "audio_transcriptions": false, "audio_speech": false, @@ -2304,6 +2321,24 @@ "messages": true, "responses": true } + }, + "manus": { + "display_name": "Manus (`manus`)", + "url": "https://docs.litellm.ai/docs/providers/manus", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": true, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": true, + "interactions": true + } } }, "endpoints": { diff --git a/pyproject.toml b/pyproject.toml index 51ef8650d0c..03a73f38e2d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "1.80.12" +version = "1.80.13" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT" @@ -167,7 +167,7 @@ requires = ["poetry-core", "wheel"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "1.80.12" +version = "1.80.13" version_files = [ "pyproject.toml:^version" ] @@ -179,4 +179,14 @@ plugins = "pydantic.mypy" asyncio_mode = "auto" markers = [ "asyncio: mark test as an asyncio test", + "limit_leaks: mark test with memory limit for leak detection (e.g., '40 MB')", + "no_parallel: mark test to run sequentially (not in parallel) - typically for memory measurement tests", +] +filterwarnings = [ + # Suppress Pydantic serializer warnings from mock server responses (non-critical for memory tests) + # These occur because the mock server returns a simplified response format + "ignore:Pydantic serializer warnings:UserWarning", + "ignore::UserWarning:pydantic.main", + # Suppress pytest-asyncio event loop deprecation warning (handled automatically by pytest-asyncio) + "ignore::DeprecationWarning:pytest_asyncio.plugin", ] diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index 810bd80a5b0..393b4cb67a1 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -154,8 +154,8 @@ class TestAzureAIFlux2ImageEdit(BaseLLMImageEditTest): return { "model": "azure_ai/flux.2-pro", "image": SINGLE_TEST_IMAGE, - "api_base": os.getenv("AZURE_AI_API_BASE", "https://litellm-ci-cd-prod.services.ai.azure.com"), - "api_key": os.getenv("AZURE_AI_API_KEY"), + "api_base": "https://litellm-ci-cd-prod.services.ai.azure.com", + "api_key": os.getenv("AZURE_API_KEY"), "api_version": "preview", } diff --git a/tests/litellm/llms/openai_like/test_abliteration_provider.py b/tests/litellm/llms/openai_like/test_abliteration_provider.py new file mode 100644 index 00000000000..8b8d443fc44 --- /dev/null +++ b/tests/litellm/llms/openai_like/test_abliteration_provider.py @@ -0,0 +1,50 @@ +""" +Unit tests for the Abliteration OpenAI-like provider. +""" + +import os +import sys + +sys.path.insert( + 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) +) + +from litellm.llms.openai_like.dynamic_config import create_config_class +from litellm.llms.openai_like.json_loader import JSONProviderRegistry + +ABLITERATION_BASE_URL = "https://api.abliteration.ai/v1" + + +def _get_config(): + provider = JSONProviderRegistry.get("abliteration") + assert provider is not None + config_class = create_config_class(provider) + return config_class() + + +def test_abliteration_provider_registered(): + provider = JSONProviderRegistry.get("abliteration") + assert provider is not None + assert provider.base_url == ABLITERATION_BASE_URL + assert provider.api_key_env == "ABLITERATION_API_KEY" + + +def test_abliteration_resolves_env_api_key(monkeypatch): + config = _get_config() + monkeypatch.setenv("ABLITERATION_API_KEY", "test-key") + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == ABLITERATION_BASE_URL + assert api_key == "test-key" + + +def test_abliteration_complete_url_appends_endpoint(): + config = _get_config() + url = config.get_complete_url( + api_base=ABLITERATION_BASE_URL, + api_key="test-key", + model="abliteration/abliterated-model", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url == f"{ABLITERATION_BASE_URL}/chat/completions" diff --git a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py index 6d005af28ac..20f48b6f393 100644 --- a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py @@ -7,6 +7,7 @@ sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path from litellm.llms.vertex_ai.gemini import transformation +from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig from litellm.types.llms import openai from litellm.types import completion from litellm.types.llms.vertex_ai import RequestBody @@ -225,4 +226,65 @@ async def test__transform_request_body_image_config_with_image_size(): assert "generationConfig" in rb assert "imageConfig" in rb["generationConfig"] assert rb["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9" - assert rb["generationConfig"]["imageConfig"]["imageSize"] == "4K" \ No newline at end of file + assert rb["generationConfig"]["imageConfig"]["imageSize"] == "4K" + + +def test_map_function_google_search_snake_case(): + """ + Test that google_search tool (snake_case) is properly mapped to googleSearch. + Fixes issue where tools=[{"google_search": {}}] was being stripped. + """ + config = VertexGeminiConfig() + optional_params = {} + + # Test snake_case google_search + tools = [{"google_search": {}}] + result = config._map_function(tools, optional_params) + + assert len(result) == 1 + assert "googleSearch" in result[0] + assert result[0]["googleSearch"] == {} + + +def test_map_function_google_search_camel_case(): + """ + Test that googleSearch tool (camelCase) still works. + """ + config = VertexGeminiConfig() + optional_params = {} + + # Test camelCase googleSearch + tools = [{"googleSearch": {}}] + result = config._map_function(tools, optional_params) + + assert len(result) == 1 + assert "googleSearch" in result[0] + assert result[0]["googleSearch"] == {} + + +def test_map_function_google_search_retrieval_snake_case(): + """ + Test that google_search_retrieval tool (snake_case) is properly mapped. + """ + config = VertexGeminiConfig() + optional_params = {} + + tools = [{"google_search_retrieval": {"dynamic_retrieval_config": {"mode": "MODE_DYNAMIC"}}}] + result = config._map_function(tools, optional_params) + + assert len(result) == 1 + assert "googleSearchRetrieval" in result[0] + + +def test_map_function_enterprise_web_search_snake_case(): + """ + Test that enterprise_web_search tool (snake_case) is properly mapped. + """ + config = VertexGeminiConfig() + optional_params = {} + + tools = [{"enterprise_web_search": {}}] + result = config._map_function(tools, optional_params) + + assert len(result) == 1 + assert "enterpriseWebSearch" in result[0] \ No newline at end of file diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py index f68b373feae..37ed1a9b08c 100644 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ b/tests/llm_responses_api_testing/base_responses_api.py @@ -54,9 +54,10 @@ def validate_responses_api_response(response, final_chunk: bool = False): assert "created_at" in response and isinstance( response["created_at"], int ), "Response should have an integer 'created_at' field" - assert "output" in response and isinstance( - response["output"], list - ), "Response should have a list 'output' field" + if response.get("status") == "completed": + assert "output" in response and isinstance( + response["output"], list + ), "Response should have a list 'output' field" # Optional fields with their expected types optional_fields = { @@ -91,7 +92,7 @@ def validate_responses_api_response(response, final_chunk: bool = False): ), f"Field '{field}' should be of type {expected_type}, but got {type(response[field])}" # Check if output has at least one item - if final_chunk is True: + if final_chunk is True and response.get("status") == "completed": assert ( len(response["output"]) > 0 ), "Response 'output' field should have at least one item" @@ -170,48 +171,57 @@ class BaseResponsesAPITest(ABC): elif event.type == "response.completed": response_completed_event = event - # assert the delta chunks content had len(collected_content_string) > 0 - # this content is typically rendered on chat ui's - assert len(collected_content_string) > 0 - # assert the response completed event is not None assert response_completed_event is not None # assert the response completed event has a response assert response_completed_event.response is not None - # assert the response completed event includes the usage - assert response_completed_event.response.usage is not None + # For async agent APIs (like Manus), the response may be in 'running' state + # without content yet - this is valid behavior + response_status = response_completed_event.response.status + if response_status in ["running", "pending"]: + # Running/pending state is acceptable - task started successfully + print(f"Response is in '{response_status}' state - async agent API behavior") + assert response_completed_event.response.id is not None + else: + # For completed responses, validate content and usage + # assert the delta chunks content had len(collected_content_string) > 0 + # this content is typically rendered on chat ui's + assert len(collected_content_string) > 0 - # basic test assert the usage seems reasonable - print( - "response_completed_event.response.usage=", - response_completed_event.response.usage, - ) - assert ( - response_completed_event.response.usage.input_tokens > 0 - and response_completed_event.response.usage.input_tokens < 100 - ) - assert ( - response_completed_event.response.usage.output_tokens > 0 - and response_completed_event.response.usage.output_tokens < 2000 - ) - assert ( - response_completed_event.response.usage.total_tokens > 0 - and response_completed_event.response.usage.total_tokens < 2000 - ) + # assert the response completed event includes the usage + assert response_completed_event.response.usage is not None - # total tokens should be the sum of input and output tokens - assert ( - response_completed_event.response.usage.total_tokens - == response_completed_event.response.usage.input_tokens - + response_completed_event.response.usage.output_tokens - ) + # basic test assert the usage seems reasonable + print( + "response_completed_event.response.usage=", + response_completed_event.response.usage, + ) + assert ( + response_completed_event.response.usage.input_tokens > 0 + and response_completed_event.response.usage.input_tokens < 100 + ) + assert ( + response_completed_event.response.usage.output_tokens > 0 + and response_completed_event.response.usage.output_tokens < 2000 + ) + assert ( + response_completed_event.response.usage.total_tokens > 0 + and response_completed_event.response.usage.total_tokens < 2000 + ) - # assert the response completed event includes cost when include_cost_in_streaming_usage is True - assert hasattr(response_completed_event.response.usage, "cost"), "Cost should be included in streaming responses API usage object" - assert response_completed_event.response.usage.cost > 0, "Cost should be greater than 0" - print(f"Cost found in streaming response: {response_completed_event.response.usage.cost}") + # total tokens should be the sum of input and output tokens + assert ( + response_completed_event.response.usage.total_tokens + == response_completed_event.response.usage.input_tokens + + response_completed_event.response.usage.output_tokens + ) + + # assert the response completed event includes cost when include_cost_in_streaming_usage is True + assert hasattr(response_completed_event.response.usage, "cost"), "Cost should be included in streaming responses API usage object" + assert response_completed_event.response.usage.cost > 0, "Cost should be greater than 0" + print(f"Cost found in streaming response: {response_completed_event.response.usage.cost}") # Reset the setting litellm.include_cost_in_streaming_usage = False @@ -450,7 +460,13 @@ class BaseResponsesAPITest(ABC): # Additional assertions specific to tool calls assert response is not None assert "output" in response - assert len(response["output"]) > 0 + # For async agent APIs (like Manus), the response may be in 'running' state + # without output yet - this is valid behavior + if response.get("status") in ["running", "pending"]: + print(f"Response is in '{response.get('status')}' state - async agent API behavior") + assert response.get("id") is not None + else: + assert len(response["output"]) > 0 @pytest.mark.asyncio async def test_responses_api_multi_turn_with_reasoning_and_structured_output(self): diff --git a/tests/llm_responses_api_testing/test_manus_responses_api.py b/tests/llm_responses_api_testing/test_manus_responses_api.py new file mode 100644 index 00000000000..338a956bfec --- /dev/null +++ b/tests/llm_responses_api_testing/test_manus_responses_api.py @@ -0,0 +1,115 @@ +import os +import sys +import pytest +import asyncio +from typing import Optional +from unittest.mock import patch, AsyncMock + +sys.path.insert(0, os.path.abspath("../..")) +import litellm +from litellm.integrations.custom_logger import CustomLogger +import json +from litellm.types.utils import StandardLoggingPayload +from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponseAPIUsage, + IncompleteDetails, +) +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from base_responses_api import BaseResponsesAPITest + + +# class TestManusResponsesAPITest(BaseResponsesAPITest): +# def get_base_completion_call_args(self): +# return { +# "model": "manus/manus-1.6", +# "api_key": os.getenv("MANUS_API_KEY"), +# } + +# @pytest.mark.parametrize("sync_mode", [True, False]) +# @pytest.mark.asyncio +# async def test_basic_openai_responses_delete_endpoint(self, sync_mode): +# pytest.skip("DELETE responses is not supported for Manus") + +# @pytest.mark.parametrize("sync_mode", [True, False]) +# @pytest.mark.asyncio +# async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode): +# pytest.skip("DELETE responses is not supported for Manus") + +# # GET responses is now supported for Manus +# @pytest.mark.parametrize("sync_mode", [True, False]) +# @pytest.mark.asyncio +# async def test_basic_openai_responses_get_endpoint(self, sync_mode): +# pytest.skip("GET responses is not supported for Manus") + +# @pytest.mark.parametrize("sync_mode", [True, False]) +# @pytest.mark.asyncio +# async def test_basic_openai_responses_cancel_endpoint(self, sync_mode): +# pytest.skip("CANCEL responses is not supported for Manus") + +# @pytest.mark.parametrize("sync_mode", [True, False]) +# @pytest.mark.asyncio +# async def test_cancel_responses_invalid_response_id(self, sync_mode): +# pytest.skip("CANCEL responses is not supported for Manus") + +# @pytest.mark.asyncio +# async def test_multiturn_responses_api(self): +# pytest.skip("Multiturn responses is not supported for Manus") + + +# @pytest.mark.asyncio +# async def test_manus_responses_api_with_agent_profile(): +# """ +# Test that Manus API correctly extracts agent profile from model name +# and includes task_mode and agent_profile in the request. +# """ +# litellm._turn_on_debug() + +# response = await litellm.aresponses( +# model="manus/manus-1.6", +# input="What's the color of the sky?", +# api_key=os.getenv("MANUS_API_KEY"), +# max_output_tokens=50, +# ) + +# print("Manus response=", json.dumps(response, indent=4, default=str)) + +# # Validate response structure +# assert isinstance(response, ResponsesAPIResponse), "Response should be ResponsesAPIResponse" +# assert response.id is not None, "Response should have an ID" +# assert response.status in ["running", "completed", "pending"], f"Status should be valid, got {response.status}" + +# # Check that metadata includes Manus-specific fields +# if response.metadata: +# assert "task_id" in response.metadata or "task_url" in response.metadata, ( +# "Manus response should include task_id or task_url in metadata" +# ) + + +# @pytest.mark.asyncio +# async def test_manus_responses_api_different_agent_profiles(): +# """ +# Test that different agent profiles work correctly. +# """ +# litellm._turn_on_debug() + +# # Test with different agent profile variants +# agent_profiles = ["manus-1.6", "manus-1.6-lite", "manus-1.6-max"] + +# for profile in agent_profiles: +# try: +# response = await litellm.aresponses( +# model=f"manus/{profile}", +# input="Hello", +# api_key=os.getenv("MANUS_API_KEY"), +# max_output_tokens=20, +# ) + +# assert response.id is not None, f"Response for {profile} should have an ID" +# print(f"✓ {profile} works: {response.id}") +# except Exception as e: +# # Some profiles might not be available, that's okay +# print(f"⚠ {profile} not available: {e}") +# pass + diff --git a/tests/llm_translation/test_bedrock_common_utils.py b/tests/llm_translation/test_bedrock_common_utils.py new file mode 100644 index 00000000000..7b6a05b6988 --- /dev/null +++ b/tests/llm_translation/test_bedrock_common_utils.py @@ -0,0 +1,182 @@ +""" +Unit tests for litellm/llms/bedrock/common_utils.py + +Tests the standalone model name utility functions and BedrockTokenCounter. +""" + +import pytest + +from litellm.llms.bedrock.common_utils import ( + BedrockModelInfo, + extract_model_name_from_bedrock_arn, + get_bedrock_base_model, + get_bedrock_cross_region_inference_regions, + strip_bedrock_routing_prefix, +) +from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter + + +class TestStripBedrockRoutingPrefix: + """Tests for strip_bedrock_routing_prefix function.""" + + def test_strips_bedrock_prefix(self): + assert strip_bedrock_routing_prefix("bedrock/claude-3-sonnet") == "claude-3-sonnet" + + def test_strips_converse_prefix(self): + assert strip_bedrock_routing_prefix("converse/claude-3-sonnet") == "claude-3-sonnet" + + def test_strips_invoke_prefix(self): + assert strip_bedrock_routing_prefix("invoke/claude-3-sonnet") == "claude-3-sonnet" + + def test_strips_openai_prefix(self): + assert strip_bedrock_routing_prefix("openai/gpt-4") == "gpt-4" + + def test_strips_all_known_prefixes(self): + # Function strips all known prefixes iteratively + # bedrock/converse/model -> converse/model -> model + assert strip_bedrock_routing_prefix("bedrock/converse/claude-3") == "claude-3" + + def test_no_prefix_unchanged(self): + assert strip_bedrock_routing_prefix("claude-3-sonnet") == "claude-3-sonnet" + + def test_model_with_dots_unchanged(self): + assert ( + strip_bedrock_routing_prefix("anthropic.claude-3-sonnet-20240229-v1:0") + == "anthropic.claude-3-sonnet-20240229-v1:0" + ) + + +class TestExtractModelNameFromBedrockArn: + """Tests for extract_model_name_from_bedrock_arn function.""" + + def test_extracts_from_provisioned_model_arn(self): + arn = "arn:aws:bedrock:us-east-1:123456789012:provisioned-model/my-model-id" + assert extract_model_name_from_bedrock_arn(arn) == "my-model-id" + + def test_extracts_from_foundation_model_arn(self): + arn = "arn:aws:bedrock:us-west-2:123456789012:foundation-model/anthropic.claude-v2" + assert extract_model_name_from_bedrock_arn(arn) == "anthropic.claude-v2" + + def test_non_arn_unchanged(self): + model = "anthropic.claude-3-sonnet-20240229-v1:0" + assert extract_model_name_from_bedrock_arn(model) == model + + def test_case_insensitive_arn_detection(self): + arn = "ARN:aws:bedrock:us-east-1:123456789012:model/my-model" + assert extract_model_name_from_bedrock_arn(arn) == "my-model" + + +class TestGetBedrockCrossRegionInferenceRegions: + """Tests for get_bedrock_cross_region_inference_regions function.""" + + def test_returns_expected_regions(self): + regions = get_bedrock_cross_region_inference_regions() + assert "us" in regions + assert "eu" in regions + assert "global" in regions + assert "apac" in regions + + def test_returns_list(self): + regions = get_bedrock_cross_region_inference_regions() + assert isinstance(regions, list) + + +class TestGetBedrockBaseModel: + """Tests for get_bedrock_base_model function.""" + + def test_strips_bedrock_prefix(self): + assert get_bedrock_base_model("bedrock/claude-3-sonnet") == "claude-3-sonnet" + + def test_strips_converse_prefix(self): + assert get_bedrock_base_model("bedrock/converse/claude-3-sonnet") == "claude-3-sonnet" + + def test_strips_us_region_prefix(self): + # us.anthropic.model -> anthropic.model + assert ( + get_bedrock_base_model("us.anthropic.claude-3-sonnet-20240229-v1:0") + == "anthropic.claude-3-sonnet-20240229-v1:0" + ) + + def test_strips_eu_region_prefix(self): + assert ( + get_bedrock_base_model("eu.anthropic.claude-3-sonnet-20240229-v1:0") + == "anthropic.claude-3-sonnet-20240229-v1:0" + ) + + def test_extracts_from_arn(self): + arn = "arn:aws:bedrock:us-east-1:123456789012:provisioned-model/my-model" + assert get_bedrock_base_model(arn) == "my-model" + + def test_model_without_prefix_unchanged(self): + model = "anthropic.claude-3-sonnet-20240229-v1:0" + assert get_bedrock_base_model(model) == model + + def test_combined_bedrock_and_region_prefix(self): + # bedrock/us.anthropic.model -> anthropic.model + assert ( + get_bedrock_base_model("bedrock/us.anthropic.claude-3-sonnet-20240229-v1:0") + == "anthropic.claude-3-sonnet-20240229-v1:0" + ) + + +class TestBedrockModelInfoWrappers: + """Tests that BedrockModelInfo methods correctly wrap standalone functions.""" + + def test_get_base_model_matches_standalone(self): + test_cases = [ + "bedrock/claude-3-sonnet", + "us.anthropic.claude-3-sonnet-20240229-v1:0", + "arn:aws:bedrock:us-east-1:123:model/my-model", + ] + for model in test_cases: + assert BedrockModelInfo.get_base_model(model) == get_bedrock_base_model(model) + + def test_extract_model_name_from_arn_matches_standalone(self): + arn = "arn:aws:bedrock:us-east-1:123456789012:provisioned-model/my-model" + assert ( + BedrockModelInfo.extract_model_name_from_arn(arn) + == extract_model_name_from_bedrock_arn(arn) + ) + + def test_get_non_litellm_routing_model_name_matches_standalone(self): + model = "bedrock/converse/claude-3" + assert ( + BedrockModelInfo.get_non_litellm_routing_model_name(model) + == strip_bedrock_routing_prefix(model) + ) + + +class TestBedrockTokenCounter: + """Tests for BedrockTokenCounter class.""" + + def test_should_use_token_counting_api_for_bedrock(self): + counter = BedrockTokenCounter() + assert counter.should_use_token_counting_api("bedrock") is True + + def test_should_not_use_token_counting_api_for_other_providers(self): + counter = BedrockTokenCounter() + assert counter.should_use_token_counting_api("openai") is False + assert counter.should_use_token_counting_api("anthropic") is False + assert counter.should_use_token_counting_api(None) is False + + def test_get_token_counter_returns_bedrock_token_counter(self): + model_info = BedrockModelInfo() + token_counter = model_info.get_token_counter() + assert isinstance(token_counter, BedrockTokenCounter) + + @pytest.mark.asyncio + async def test_count_tokens_returns_none_for_empty_messages(self): + counter = BedrockTokenCounter() + result = await counter.count_tokens( + model_to_use="anthropic.claude-3-sonnet", + messages=None, + contents=None, + ) + assert result is None + + result = await counter.count_tokens( + model_to_use="anthropic.claude-3-sonnet", + messages=[], + contents=None, + ) + assert result is None diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index ec510b8f953..29c2d26981c 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -2835,6 +2835,34 @@ def test_bedrock_invoke_provider(): ) == "nova" ) + assert ( + litellm.AmazonInvokeConfig().get_bedrock_invoke_provider("amazon.nova-pro-v1:0") + == "nova" + ) + assert ( + litellm.AmazonInvokeConfig().get_bedrock_invoke_provider( + "amazon.nova-lite-v1:0" + ) + == "nova" + ) + assert ( + litellm.AmazonInvokeConfig().get_bedrock_invoke_provider( + "amazon.nova-micro-v1:0" + ) + == "nova" + ) + assert ( + litellm.AmazonInvokeConfig().get_bedrock_invoke_provider( + "amazon.nova-premier-v1:0" + ) + == "nova" + ) + assert ( + litellm.AmazonInvokeConfig().get_bedrock_invoke_provider( + "amazon.nova-2-lite-v1:0" + ) + == "nova" + ) def test_bedrock_description_param(): @@ -3488,7 +3516,9 @@ def test_bedrock_openai_imported_model(): url = mock_post.call_args.kwargs["url"] print(f"URL: {url}") assert "bedrock-runtime.us-east-1.amazonaws.com" in url - assert "arn:aws:bedrock:us-east-1:117159858402:imported-model/m4gc1mrfuddy" in url + assert ( + "arn:aws:bedrock:us-east-1:117159858402:imported-model/m4gc1mrfuddy" in url + ) assert "/invoke" in url # Validate request body follows OpenAI format @@ -3517,7 +3547,9 @@ def test_bedrock_openai_imported_model(): # Check image_url content assert user_msg["content"][1]["type"] == "image_url" assert "image_url" in user_msg["content"][1] - assert user_msg["content"][1]["image_url"]["url"].startswith("data:image/jpeg;base64,") + assert user_msg["content"][1]["image_url"]["url"].startswith( + "data:image/jpeg;base64," + ) assert user_msg["content"][2]["type"] == "image_url" assert "image_url" in user_msg["content"][2] @@ -3526,21 +3558,67 @@ def test_bedrock_openai_imported_model(): assert request_body["max_tokens"] == 300 assert request_body["temperature"] == 0.5 + +def test_bedrock_nova_provider_detection(): + """ + Test that Nova models are correctly detected even when prefixed with "amazon." + Regression test for issue #17910 where models like "amazon.nova-pro-v1:0" + were incorrectly identified as "amazon" (Titan) instead of "nova". + """ + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + # Test various Nova model formats + nova_test_cases = [ + ("us.amazon.nova-pro-v1:0", "nova"), + ("us.amazon.nova-lite-v1:0", "nova"), + ("us.amazon.nova-micro-v1:0", "nova"), + ("amazon.nova-pro-v1:0", "nova"), + ("amazon.nova-lite-v1:0", "nova"), + ("amazon.nova-micro-v1:0", "nova"), + ("amazon.nova-premier-v1:0", "nova"), + ("amazon.nova-2-lite-v1:0", "nova"), + ("bedrock/amazon.nova-pro-v1:0", "nova"), + ("bedrock/invoke/amazon.nova-pro-v1:0", "nova"), + ("amazon.Nova-pro-v1:0", "nova"), + ("amazon.NOVA-pro-v1:0", "nova"), + ] + + for model, expected in nova_test_cases: + provider = BaseAWSLLM.get_bedrock_invoke_provider(model) + assert ( + provider == expected + ), f"Failed for model: {model}, expected: {expected}, got: {provider}" + + # Verify that Amazon Titan models still return "amazon" + titan_test_cases = [ + ("amazon.titan-text-express-v1", "amazon"), + ("us.amazon.titan-text-lite-v1", "amazon"), + ] + + for model, expected in titan_test_cases: + provider = BaseAWSLLM.get_bedrock_invoke_provider(model) + assert ( + provider == expected + ), f"Failed for model: {model}, expected: {expected}, got: {provider}" + + def test_bedrock_openai_provider_detection(): """ Test that the OpenAI provider is correctly detected from model strings. """ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM - + # Test various OpenAI model formats test_cases = [ "openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/abc123", "bedrock/openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/xyz789", ] - + for model in test_cases: provider = BaseAWSLLM.get_bedrock_invoke_provider(model) - assert provider == "openai", f"Failed for model: {model}, got provider: {provider}" + assert ( + provider == "openai" + ), f"Failed for model: {model}, got provider: {provider}" print(f"✓ Provider detection works for: {model}") @@ -3549,16 +3627,16 @@ def test_bedrock_openai_model_id_extraction(): Test that the model ID (ARN) is correctly extracted and encoded for OpenAI models. """ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM - - model = "openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/test-model-123" - provider = BaseAWSLLM.get_bedrock_invoke_provider(model) - - model_id = BaseAWSLLM.get_bedrock_model_id( - model=model, - provider=provider, - optional_params={} + + model = ( + "openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/test-model-123" ) - + provider = BaseAWSLLM.get_bedrock_invoke_provider(model) + + model_id = BaseAWSLLM.get_bedrock_model_id( + model=model, provider=provider, optional_params={} + ) + # The ARN should be double URL encoded assert "arn" in model_id assert "imported-model" in model_id @@ -3570,20 +3648,17 @@ def test_bedrock_openai_convert_messages_to_prompt(): Test that convert_messages_to_prompt returns empty string for OpenAI models. """ from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM - + bedrock_llm = BedrockLLM() messages = [ {"role": "system", "content": "You are helpful"}, - {"role": "user", "content": "Hello"} + {"role": "user", "content": "Hello"}, ] - + prompt, chat_history = bedrock_llm.convert_messages_to_prompt( - model="test-model", - messages=messages, - provider="openai", - custom_prompt_dict={} + model="test-model", messages=messages, provider="openai", custom_prompt_dict={} ) - + # OpenAI models use messages directly, no prompt conversion assert prompt == "" assert chat_history is None @@ -3598,37 +3673,33 @@ def test_bedrock_openai_response_parsing(): from litellm import ModelResponse from unittest.mock import Mock import json - + bedrock_llm = BedrockLLM() - + # Mock OpenAI-style response openai_response = { "choices": [ { "message": { "content": "The capital of France is Paris.", - "role": "assistant" + "role": "assistant", }, "finish_reason": "stop", - "index": 0 + "index": 0, } ], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 8, - "total_tokens": 18 - } + "usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18}, } - + mock_response = Mock() mock_response.json.return_value = openai_response mock_response.text = json.dumps(openai_response) mock_response.status_code = 200 mock_response.headers = {} - + model_response = ModelResponse() mock_logging = Mock() - + result = bedrock_llm.process_response( model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", response=mock_response, @@ -3640,18 +3711,18 @@ def test_bedrock_openai_response_parsing(): data={}, messages=[{"role": "user", "content": "What is the capital of France?"}], print_verbose=lambda x: None, - encoding=None + encoding=None, ) - + # Verify response content assert result.choices[0].message.content == "The capital of France is Paris." assert result.choices[0].finish_reason == "stop" - + # Verify usage assert result.usage.prompt_tokens == 10 assert result.usage.completion_tokens == 8 assert result.usage.total_tokens == 18 - + print("✓ OpenAI response parsing works correctly") @@ -3659,45 +3730,47 @@ def test_bedrock_openai_request_transformation(): """ Test that the request is correctly transformed for OpenAI models. """ - from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import AmazonInvokeConfig - + from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( + AmazonInvokeConfig, + ) + config = AmazonInvokeConfig() - + model = "openai/arn:aws:bedrock:us-east-1:123:imported-model/test" messages = [ {"role": "system", "content": "You are helpful"}, - {"role": "user", "content": "Hello"} + {"role": "user", "content": "Hello"}, ] - + optional_params = { "max_tokens": 100, "temperature": 0.7, "top_p": 0.9, - "stream": False + "stream": False, } - + litellm_params = {} headers = {} - - with patch.object(config, 'get_bedrock_invoke_provider', return_value="openai"): + + with patch.object(config, "get_bedrock_invoke_provider", return_value="openai"): result = config.transform_request( model=model, messages=messages, optional_params=optional_params.copy(), litellm_params=litellm_params, - headers=headers + headers=headers, ) - + # Verify the request uses messages format (not prompt) assert "messages" in result assert len(result["messages"]) == 2 assert result["messages"][0]["role"] == "system" assert result["messages"][1]["role"] == "user" - + # Verify parameters are included assert "max_tokens" in result assert "temperature" in result - + print("✓ Request transformation works correctly") @@ -3705,20 +3778,22 @@ def test_bedrock_openai_parameter_filtering(): """ Test that only supported OpenAI parameters are included in the request. """ - from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import AmazonBedrockOpenAIConfig - + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( + AmazonBedrockOpenAIConfig, + ) + config = AmazonBedrockOpenAIConfig() model = "test-model" - + supported_params = config.get_supported_openai_params(model=model) - + # Verify common OpenAI parameters are supported assert "max_tokens" in supported_params assert "temperature" in supported_params assert "top_p" in supported_params assert "stream" in supported_params assert "stop" in supported_params - + print(f"✓ Parameter filtering supports: {len(supported_params)} parameters") print(f" Supported params: {supported_params}") @@ -3728,12 +3803,12 @@ def test_bedrock_openai_route_detection(): Test that the OpenAI route is correctly detected. """ from litellm.llms.bedrock.common_utils import BedrockModelInfo - + test_cases = [ ("openai/arn:aws:bedrock:us-east-1:123:imported-model/test", "openai"), ("bedrock/openai/arn:aws:bedrock:us-east-1:123:imported-model/test", "openai"), ] - + for model, expected_route in test_cases: route = BedrockModelInfo.get_bedrock_route(model) assert route == expected_route, f"Failed for model: {model}, got route: {route}" @@ -3745,15 +3820,30 @@ def test_bedrock_openai_explicit_route_check(): Test the explicit OpenAI route checker helper method. """ from litellm.llms.bedrock.common_utils import BedrockModelInfo - + # Test with openai/ prefix - assert BedrockModelInfo._explicit_openai_route("openai/arn:aws:bedrock:us-east-1:123:imported-model/test") is True - assert BedrockModelInfo._explicit_openai_route("bedrock/openai/arn:aws:bedrock:us-east-1:123:imported-model/test") is True - + assert ( + BedrockModelInfo._explicit_openai_route( + "openai/arn:aws:bedrock:us-east-1:123:imported-model/test" + ) + is True + ) + assert ( + BedrockModelInfo._explicit_openai_route( + "bedrock/openai/arn:aws:bedrock:us-east-1:123:imported-model/test" + ) + is True + ) + # Test without openai/ prefix assert BedrockModelInfo._explicit_openai_route("anthropic.claude-3-sonnet") is False - assert BedrockModelInfo._explicit_openai_route("arn:aws:bedrock:us-east-1:123:imported-model/test") is False - + assert ( + BedrockModelInfo._explicit_openai_route( + "arn:aws:bedrock:us-east-1:123:imported-model/test" + ) + is False + ) + print("✓ Explicit route check works correctly") @@ -3761,16 +3851,18 @@ def test_bedrock_openai_config_initialization(): """ Test that AmazonBedrockOpenAIConfig can be properly initialized. """ - from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import AmazonBedrockOpenAIConfig - + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( + AmazonBedrockOpenAIConfig, + ) + config = AmazonBedrockOpenAIConfig() - + # Verify it has the necessary methods - assert hasattr(config, 'get_supported_openai_params') - assert hasattr(config, 'transform_request') - assert hasattr(config, 'transform_response') - assert hasattr(config, 'map_openai_params') - + assert hasattr(config, "get_supported_openai_params") + assert hasattr(config, "transform_request") + assert hasattr(config, "transform_response") + assert hasattr(config, "map_openai_params") + print("✓ AmazonBedrockOpenAIConfig initializes correctly") @@ -3779,9 +3871,9 @@ def test_bedrock_openai_multiple_message_types(): Test that various message content types are handled correctly. """ from litellm.llms.custom_httpx.http_handler import HTTPHandler - + client = HTTPHandler() - + # Test with mixed content types messages = [ {"role": "system", "content": "You are helpful"}, @@ -3790,11 +3882,14 @@ def test_bedrock_openai_multiple_message_types(): "role": "user", "content": [ {"type": "text", "text": "Complex message with text"}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,iVBORw0KGg"}} - ] - } + { + "type": "image_url", + "image_url": {"url": "data:image/jpeg;base64,iVBORw0KGg"}, + }, + ], + }, ] - + with patch.object(client, "post") as mock_post: try: response = completion( @@ -3805,18 +3900,18 @@ def test_bedrock_openai_multiple_message_types(): ) except Exception as e: pass - + # Verify the request was made if mock_post.called: request_body = json.loads(mock_post.call_args.kwargs["data"]) - + # Verify messages are preserved assert "messages" in request_body assert len(request_body["messages"]) == 3 - + # Verify mixed content is handled assert isinstance(request_body["messages"][2]["content"], list) - + print("✓ Multiple message types handled correctly") @@ -3829,18 +3924,18 @@ def test_bedrock_openai_error_handling(): from litellm.llms.bedrock.common_utils import BedrockError from unittest.mock import Mock import json - + bedrock_llm = BedrockLLM() - + # Mock error response mock_response = Mock() mock_response.json.side_effect = Exception("Invalid JSON") mock_response.text = "Invalid response" mock_response.status_code = 422 - + model_response = ModelResponse() mock_logging = Mock() - + with pytest.raises(BedrockError) as exc_info: bedrock_llm.process_response( model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", @@ -3853,8 +3948,8 @@ def test_bedrock_openai_error_handling(): data={}, messages=[], print_verbose=lambda x: None, - encoding=None + encoding=None, ) - + assert exc_info.value.status_code == 422 print("✓ Error handling works correctly") diff --git a/tests/llm_translation/test_bedrock_moonshot.py b/tests/llm_translation/test_bedrock_moonshot.py new file mode 100644 index 00000000000..c6066c7db42 --- /dev/null +++ b/tests/llm_translation/test_bedrock_moonshot.py @@ -0,0 +1,290 @@ +""" +Tests for Bedrock Moonshot (Kimi K2) integration. + +This test suite verifies: +1. Basic completion functionality +2. Streaming responses +3. System message support +4. Temperature parameter handling +5. Reasoning content extraction from tags +6. Tool calling support (including tool response handling) +7. Parameter validation (e.g., stop sequences not supported) +""" + +from base_llm_unit_tests import BaseLLMChatTest +import pytest +import sys +import os +import json + +sys.path.insert(0, os.path.abspath("../..")) +import litellm +from litellm.llms.bedrock.common_utils import get_bedrock_chat_config + + +class TestBedrockMoonshotInvoke(BaseLLMChatTest): + """ + Test suite for Bedrock Moonshot via invoke route. + Inherits all standard LLM tests from BaseLLMChatTest. + """ + + def get_base_completion_call_args(self) -> dict: + litellm._turn_on_debug() + return { + "model": "bedrock/invoke/moonshot.kimi-k2-thinking", + } + + def test_tool_call_no_arguments(self, tool_call_no_arguments): + """Test that tool calls with no arguments is translated correctly.""" + pass + + +class TestBedrockMoonshotBasic: + """Unit tests for Bedrock Moonshot configuration and transformations.""" + + def test_provider_detection_invoke(self): + """Test that Bedrock Moonshot invoke models are correctly detected.""" + config = get_bedrock_chat_config("bedrock/invoke/moonshot.kimi-k2-thinking") + assert config is not None + assert config.__class__.__name__ == "AmazonMoonshotConfig" + + def test_provider_detection_converse(self): + """Test that Bedrock Moonshot converse models are correctly detected.""" + config = get_bedrock_chat_config("bedrock/moonshot.kimi-k2-thinking") + assert config is not None + + def test_config_initialization(self): + """Test that AmazonMoonshotConfig initializes correctly.""" + config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking") + assert config is not None + assert config.custom_llm_provider == "bedrock" + + def test_supported_params(self): + """Test that supported OpenAI params are correctly defined.""" + config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking") + supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking") + + # Should support these params + assert "temperature" in supported_params + assert "max_tokens" in supported_params + assert "top_p" in supported_params + assert "stream" in supported_params + assert "tools" in supported_params + assert "tool_choice" in supported_params + + # Should NOT support stop sequences on Bedrock + assert "stop" not in supported_params + + # Should NOT support functions (use tools instead) + assert "functions" not in supported_params + + def test_transform_request_strips_model_prefix(self): + """Test that model ID prefixes are correctly stripped in transform_request.""" + from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import ( + AmazonMoonshotConfig, + ) + + config = AmazonMoonshotConfig() + + messages = [{"role": "user", "content": "Hello"}] + + # Test that bedrock/invoke/ prefix is stripped + transformed = config.transform_request( + model="bedrock/invoke/moonshot.kimi-k2-thinking", + messages=messages, + optional_params={}, + litellm_params={}, + headers={} + ) + + # The model ID in the request body should be stripped + assert transformed["model"] == "moonshot.kimi-k2-thinking" + + +class TestBedrockMoonshotReasoningContent: + """Tests for reasoning content extraction.""" + + def test_reasoning_content_extraction(self): + """Test that reasoning content is extracted from tags.""" + from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import ( + AmazonMoonshotConfig, + ) + + config = AmazonMoonshotConfig() + + # Test with reasoning tags + content_with_reasoning = "This is my thought processThis is the answer" + reasoning, content = config._extract_reasoning_from_content(content_with_reasoning) + + assert reasoning == "This is my thought process" + assert content == "This is the answer" + assert "" not in content + + # Test without reasoning tags + content_without_reasoning = "This is just a regular answer" + reasoning, content = config._extract_reasoning_from_content(content_without_reasoning) + + assert reasoning is None + assert content == "This is just a regular answer" + + +class TestBedrockMoonshotToolCalling: + """Unit tests for tool calling functionality.""" + + def test_tool_calling_supported(self): + """Test that tool calling is supported for Kimi K2 Thinking model.""" + config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking") + supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking") + + # Kimi K2 Thinking DOES support tool calls (unlike kimi-thinking-preview) + assert "tools" in supported_params + assert "tool_choice" in supported_params + + def test_tool_call_request_format(self): + """Test that tool call requests are formatted correctly.""" + from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import ( + AmazonMoonshotConfig, + ) + + config = AmazonMoonshotConfig() + + messages = [ + {"role": "user", "content": "What's the weather in San Francisco?"} + ] + + optional_params = { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the current weather", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string"} + }, + "required": ["location"] + } + } + } + ] + } + + transformed = config.transform_request( + model="bedrock/invoke/moonshot.kimi-k2-thinking", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={} + ) + + # Verify model ID is stripped + assert transformed["model"] == "moonshot.kimi-k2-thinking" + + # Verify tools are included + assert "tools" in transformed + assert len(transformed["tools"]) == 1 + assert transformed["tools"][0]["function"]["name"] == "get_weather" + + def test_tool_response_message_format(self): + """Test that tool response messages are formatted correctly.""" + # This tests the proper format for sending tool responses back + tool_response_message = { + "role": "tool", + "tool_call_id": "call_123", + "content": json.dumps({"temperature": 72, "condition": "sunny"}) + } + + # Verify the message structure + assert tool_response_message["role"] == "tool" + assert "tool_call_id" in tool_response_message + assert "content" in tool_response_message + + +class TestBedrockMoonshotParameterValidation: + """Tests for parameter validation and edge cases.""" + + def test_stop_sequences_not_supported(self): + """Test that stop sequences are correctly excluded from supported params.""" + config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking") + supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking") + + # Bedrock Moonshot doesn't support stopSequences field + assert "stop" not in supported_params + + def test_temperature_range(self): + """Test that temperature parameter is handled correctly.""" + # Moonshot models support temperature 0-1 + # This is handled by the parent MoonshotChatConfig class + config = get_bedrock_chat_config("invoke/moonshot.kimi-k2-thinking") + + # Verify config exists and can handle temperature + assert config is not None + supported_params = config.get_supported_openai_params("moonshot.kimi-k2-thinking") + assert "temperature" in supported_params + + +class TestBedrockMoonshotTransformations: + """Tests for request/response transformations.""" + + def test_transform_request_basic(self): + """Test basic request transformation.""" + from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import ( + AmazonMoonshotConfig, + ) + + config = AmazonMoonshotConfig() + + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello!"} + ] + + optional_params = { + "temperature": 0.7, + "max_tokens": 100 + } + + transformed = config.transform_request( + model="bedrock/invoke/moonshot.kimi-k2-thinking", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={} + ) + + # Verify model ID is stripped + assert transformed["model"] == "moonshot.kimi-k2-thinking" + + # Verify messages are included + assert "messages" in transformed + assert len(transformed["messages"]) >= 1 + + # Verify optional params are included + assert transformed["temperature"] == 0.7 + assert transformed["max_tokens"] == 100 + + def test_transform_request_with_system_message(self): + """Test request transformation with system message.""" + from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import ( + AmazonMoonshotConfig, + ) + + config = AmazonMoonshotConfig() + + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello!"} + ] + + transformed = config.transform_request( + model="moonshot.kimi-k2-thinking", + messages=messages, + optional_params={}, + litellm_params={}, + headers={} + ) + + # System messages should be supported + assert "messages" in transformed diff --git a/tests/llm_translation/test_openrouter.py b/tests/llm_translation/test_openrouter.py index 839d08e12bf..105b05d3449 100644 --- a/tests/llm_translation/test_openrouter.py +++ b/tests/llm_translation/test_openrouter.py @@ -32,3 +32,18 @@ def test_completion_openrouter_image_generation(): .message.images[0]["image_url"]["url"] .startswith("data:image/png;base64,") ) + + +def test_openrouter_embedding(): + """Test OpenRouter embeddings support.""" + litellm._turn_on_debug() + resp = litellm.embedding( + model="openrouter/openai/text-embedding-3-small", + input=["Hello world", "How are you?"], + ) + print(resp) + assert resp is not None + assert len(resp.data) == 2 + assert resp.data[0]["embedding"] is not None + assert isinstance(resp.data[0]["embedding"], list) + assert len(resp.data[0]["embedding"]) > 0 diff --git a/tests/load_tests/memory_leak_utils.py b/tests/load_tests/memory_leak_utils.py new file mode 100644 index 00000000000..160a67fa184 --- /dev/null +++ b/tests/load_tests/memory_leak_utils.py @@ -0,0 +1,314 @@ +""" +Memory Leak Testing Utilities + +This module provides reusable utilities, fixtures, and helpers for memory leak +and OOM (Out of Memory) detection tests. It includes: +- Mock server setup for local testing +- Memory tracking fixtures +- Router fixtures configured for testing +- Helper functions for running memory baseline tests + +Usage: + from tests.load_tests.memory_leak_utils import ( + mock_server, + limit_memory, + test_router, + run_memory_baseline_test, + ) +""" + +import gc +import os +import socket +import sys +import time +from threading import Thread + +# Add parent directory to path to import litellm (same pattern as other tests) +filepath = os.path.dirname(os.path.abspath(__file__)) +sys.path.insert(0, os.path.abspath(os.path.join(filepath, "../.."))) + +import httpx +import pytest +import psutil +from fastapi import FastAPI, Request +from fastapi.responses import JSONResponse +from litellm.router import Router + +# Test Configuration Constants +TEST_API_KEY = "sk-1234" +TEST_MODEL_NAME = "gpt-3.5-turbo" + +# Timing Constants (seconds) +GC_STABILIZATION_DELAY = 0.05 + + +# Mock OpenAI-compatible server +def create_mock_server(): + """Create a simple FastAPI mock server that mimics OpenAI API responses.""" + app = FastAPI() + + @app.post("/v1/chat/completions") + @app.post("/chat/completions") + async def chat_completions(request: Request): + """Mock OpenAI chat completions endpoint.""" + request_data = await request.json() + # Return a simple mock response + return JSONResponse({ + "id": "chatcmpl-mock", + "object": "chat.completion", + "created": int(time.time()), + "model": request_data.get("model", TEST_MODEL_NAME), + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Mock response" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15 + } + }) + + # Catch-all route to see what URLs are being requested + @app.api_route("/{path:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"]) + async def catch_all(request: Request, path: str): + """Catch-all route to debug what URLs are being requested.""" + print(f"[Mock Server] Received request: {request.method} {request.url.path}") + # For non-chat-completions, return 404 + return JSONResponse({"detail": "Not Found"}, status_code=404) + + return app + + +def run_server(app, port): + """Run uvicorn server in a thread.""" + import uvicorn + # Use uvicorn.run which blocks - this is fine in a daemon thread + uvicorn.run(app, host="127.0.0.1", port=port, log_level="error", access_log=False) + + +@pytest.fixture(scope="session") +def mock_server(): + """Start a mock server in a separate thread for the test session. + + Yields the server URL (with trailing slash) for use in router configuration. + """ + app = create_mock_server() + port = 18888 + + # Check if port is already in use + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + try: + sock.bind(("127.0.0.1", port)) + sock.close() + except OSError: + # Port already in use, try next port + port = 18889 + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + try: + sock.bind(("127.0.0.1", port)) + sock.close() + except OSError: + pytest.fail(f"Could not find available port for mock server (tried 18888, 18889)") + + # Start server in background thread + thread = Thread(target=lambda: run_server(app, port), daemon=True) + thread.start() + + # Wait for server to start and verify it's accessible + # Ensure api_base has trailing slash (LiteLLM appends /v1/chat/completions) + server_url = f"http://127.0.0.1:{port}/" + max_attempts = 20 # More attempts to ensure server is ready + server_ready = False + for attempt in range(max_attempts): + time.sleep(0.3) # Longer wait between attempts + try: + # Test the actual endpoint we'll use (LiteLLM appends /v1/chat/completions to api_base) + response = httpx.post( + f"{server_url}v1/chat/completions", + json={"model": TEST_MODEL_NAME, "messages": [{"role": "user", "content": "test"}]}, + timeout=2.0 + ) + if response.status_code == 200: + server_ready = True + print(f"[Mock Server] Server ready at {server_url}") + break + except (httpx.ConnectError, httpx.TimeoutException, httpx.NetworkError) as e: + # Server not ready yet, continue waiting + if attempt == max_attempts - 1: + pytest.fail( + f"Mock server failed to start on {server_url} after {max_attempts} attempts. " + f"Could not connect to /v1/chat/completions endpoint. Error: {e}" + ) + continue + except Exception as e: + # Other errors might indicate server is up but endpoint has issues + # If we get a response (even error), server is running + print(f"[Mock Server] Server responded with error (but is running): {e}") + server_ready = True + break + + if not server_ready: + pytest.fail(f"Mock server not accessible at {server_url} after {max_attempts} attempts") + + yield server_url + + # Server will be cleaned up when thread dies (daemon=True) + + +@pytest.fixture +def limit_memory(request): + """Fixture to track memory usage and enforce limits via @pytest.mark.limit_leaks marker. + + Usage: + @pytest.mark.limit_leaks("40 MB") + def test_something(limit_memory): + # Test code here + # Memory will be measured and test will fail if increase exceeds limit + """ + marker = request.node.get_closest_marker("limit_leaks") + if marker: + # Parse limit from marker (e.g., "30 MB" -> 30) + limit_str = marker.args[0] if marker.args else "100 MB" + limit_mb = float(limit_str.split()[0]) + limit_bytes = limit_mb * 1024 * 1024 + + # Measure baseline memory (router will be fresh from fixture) + process = psutil.Process(os.getpid()) + baseline_memory = process.memory_info().rss + + yield + + # Force GC before measuring final memory + gc.collect() + # Small delay for memory to stabilize + time.sleep(GC_STABILIZATION_DELAY) + + # Measure final memory after test + final_memory = process.memory_info().rss + memory_increase = final_memory - baseline_memory + memory_increase_mb = memory_increase / 1024 / 1024 + + # Print memory stats + print(f"\n[Memory Limit Test] Memory usage:") + print(f" Baseline: {baseline_memory / 1024 / 1024:.2f} MB") + print(f" Final: {final_memory / 1024 / 1024:.2f} MB") + print(f" Increase: {memory_increase_mb:+.2f} MB") + print(f" Limit: {limit_mb:.2f} MB") + + # Fail if memory increase exceeds limit + if memory_increase > limit_bytes: + pytest.fail( + f"Memory limit exceeded: {memory_increase_mb:.2f} MB increase > {limit_mb:.2f} MB limit. " + f"Baseline: {baseline_memory / 1024 / 1024:.2f} MB, Final: {final_memory / 1024 / 1024:.2f} MB" + ) + else: + yield + + +@pytest.fixture +def test_router(mock_server): + """Fixture to create a fresh router instance for each test. + + Uses the mock server fixture to avoid external API calls. + Disables cooldowns to prevent deployments from being marked unavailable. + + Usage: + def test_something(test_router, limit_memory): + # Use test_router for making requests + response = await test_router.acompletion(...) + """ + router = Router( + model_list=[ + { + "model_name": TEST_MODEL_NAME, + "litellm_params": { + "model": f"openai/{TEST_MODEL_NAME}", + "api_base": mock_server, + "api_key": TEST_API_KEY, + }, + }, + ], + disable_cooldowns=True, # Disable cooldowns for testing + allowed_fails=1000, # Allow many failures before cooldown (effectively disabled) + ) + yield router + # Cleanup after test + try: + router.discard() + except Exception: + pass # Ignore cleanup errors + + +async def run_memory_baseline_test(num_requests: int, router: Router, limit_memory): + """Helper function to run memory baseline test with specified number of requests. + + Makes requests concurrently in batches for speed, with proper error handling + that doesn't fail the test on individual request failures. + + Args: + num_requests: Number of requests to make. + router: Router instance to use for requests. + limit_memory: Pytest fixture for memory tracking (reference to suppress linter warning). + + Example: + @pytest.mark.asyncio + @pytest.mark.limit_leaks("40 MB") + async def test_memory(test_router, limit_memory): + await run_memory_baseline_test(1000, test_router, limit_memory) + """ + # Fixture is used automatically by pytest - reference it to suppress linter warning + _ = limit_memory + + # Make requests concurrently in batches for speed + # Batch size of 20 provides good balance between speed and memory pressure + BATCH_SIZE = 20 + + for batch_start in range(0, num_requests, BATCH_SIZE): + batch_end = min(batch_start + BATCH_SIZE, num_requests) + # Create concurrent tasks for this batch + tasks = [ + router.acompletion( + model=TEST_MODEL_NAME, + messages=[{"role": "user", "content": f"Test request {i}"}], + ) + for i in range(batch_start, batch_end) + ] + # Execute batch concurrently + # Note: return_exceptions=True allows test to continue even if some requests fail + import asyncio + responses = await asyncio.gather(*tasks, return_exceptions=True) + # Filter out failed requests but continue with test + valid_responses = [] + failed_count = 0 + for i, response in enumerate(responses): + if isinstance(response, Exception): + failed_count += 1 + # Log exception but continue + print(f" Warning: Request {batch_start + i} failed: {type(response).__name__}: {response}") + elif response is None: + failed_count += 1 + print(f" Warning: Request {batch_start + i} returned None") + else: + valid_responses.append(response) + + # Continue with valid responses - don't fail the test + # If all failed, that's logged but test continues (might indicate bigger issue) + if failed_count > 0: + print(f" Note: {failed_count}/{len(responses)} requests failed in batch {batch_start}-{batch_end}, continuing with {len(valid_responses)} valid responses") + + # Use valid_responses for cleanup + responses = valid_responses + # Clean up batch + del responses + del tasks + del valid_responses + # GC after each batch to prevent accumulation + gc.collect() + + print(f"[Simple Memory Test] Completed {num_requests} requests") diff --git a/tests/load_tests/test_linear_memory_growth.py b/tests/load_tests/test_linear_memory_growth.py new file mode 100644 index 00000000000..3b7b8041b90 --- /dev/null +++ b/tests/load_tests/test_linear_memory_growth.py @@ -0,0 +1,121 @@ +""" +Memory Leak Detection Tests - Linear Memory Growth + +Tests that check for linear/progressive memory growth by running different numbers +of requests (1k, 2k, 4k, 10k, 30k) with the same memory limit. If lower request +count tests pass but higher ones fail, it indicates linear memory growth per request. + +These tests will fail if memory leaks are detected, helping catch OOM issues before production. + +IMPORTANT: These tests should be run INDIVIDUALLY, not all together. Running them +together causes memory baseline drift between tests, making it difficult to detect +linear growth accurately. Each test should be run in isolation: + + pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_1k -v + pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_2k -v + # etc. + +NOTE: Not recommended for accurate results: +pytest tests/load_tests/test_linear_memory_growth.py -v +""" + +import pytest + +from tests.load_tests.memory_leak_utils import ( + limit_memory, # noqa: F401 # pytest fixture used via dependency injection + mock_server, # noqa: F401 # pytest fixture used via dependency injection + run_memory_baseline_test, + test_router, # noqa: F401 # pytest fixture used via dependency injection +) + +# Memory limit for all linear memory growth tests +MEMORY_LIMIT = "40 MB" + + +@pytest.mark.asyncio +@pytest.mark.limit_leaks(MEMORY_LIMIT) +@pytest.mark.no_parallel # Must run sequentially - measures process memory +async def test_memory_baseline_1k(test_router, limit_memory): + """ + Memory baseline test with 1,000 requests. + Uses @pytest.mark.limit_leaks("40 MB") to enforce memory limit. + If this passes but higher request count tests fail, indicates progressive memory leak. + + NOTE: This test should be run INDIVIDUALLY, not with other tests in this file. + Running multiple tests together causes memory baseline drift, making it difficult + to accurately detect linear memory growth. Run with: + pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_1k -v + """ + await run_memory_baseline_test(1000, test_router, limit_memory) + + +@pytest.mark.asyncio +@pytest.mark.limit_leaks(MEMORY_LIMIT) +@pytest.mark.no_parallel # Must run sequentially - measures process memory +async def test_memory_baseline_2k(test_router, limit_memory): + """ + Memory baseline test with 2,000 requests. + Uses @pytest.mark.limit_leaks("40 MB") to enforce memory limit. + If this passes but test_memory_baseline_4k fails, indicates progressive memory leak. + + NOTE: This test should be run INDIVIDUALLY, not with other tests in this file. + Running multiple tests together causes memory baseline drift, making it difficult + to accurately detect linear memory growth. Run with: + pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_2k -v + """ + await run_memory_baseline_test(2000, test_router, limit_memory) + + +@pytest.mark.asyncio +@pytest.mark.limit_leaks(MEMORY_LIMIT) +@pytest.mark.no_parallel # Must run sequentially - measures process memory +async def test_memory_baseline_4k(test_router, limit_memory): + """ + Memory baseline test with 4,000 requests. + Uses @pytest.mark.limit_leaks("40 MB") to enforce memory limit. + If test_memory_baseline_1k and test_memory_baseline_2k pass but this fails, + it's a clear sign of sequential/progressive memory growth. + + NOTE: This test should be run INDIVIDUALLY, not with other tests in this file. + Running multiple tests together causes memory baseline drift, making it difficult + to accurately detect linear memory growth. Run with: + pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_4k -v + """ + await run_memory_baseline_test(4000, test_router, limit_memory) + + + +@pytest.mark.asyncio +@pytest.mark.limit_leaks(MEMORY_LIMIT) +@pytest.mark.no_parallel # Must run sequentially - measures process memory +async def test_memory_baseline_10k(test_router, limit_memory): + """ + Memory baseline test with 10,000 requests. + Uses @pytest.mark.limit_leaks("40 MB") to enforce memory limit. + If test_memory_baseline_1k and test_memory_baseline_2k pass but this fails, + it's a clear sign of sequential/progressive memory growth. + + NOTE: This test should be run INDIVIDUALLY, not with other tests in this file. + Running multiple tests together causes memory baseline drift, making it difficult + to accurately detect linear memory growth. Run with: + pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_10k -v + """ + await run_memory_baseline_test(10000, test_router, limit_memory) + + +@pytest.mark.asyncio +@pytest.mark.limit_leaks(MEMORY_LIMIT) +@pytest.mark.no_parallel # Must run sequentially - measures process memory +async def test_memory_baseline_30k(test_router, limit_memory): + """ + Memory baseline test with 30,000 requests. + Uses @pytest.mark.limit_leaks("40 MB") to enforce memory limit. + If test_memory_baseline_1k and test_memory_baseline_2k pass but this fails, + it's a clear sign of sequential/progressive memory growth. + + NOTE: This test should be run INDIVIDUALLY, not with other tests in this file. + Running multiple tests together causes memory baseline drift, making it difficult + to accurately detect linear memory growth. Run with: + pytest tests/load_tests/test_linear_memory_growth.py::test_memory_baseline_30k -v + """ + await run_memory_baseline_test(30000, test_router, limit_memory) diff --git a/tests/router_unit_tests/test_router_embedding_headers.py b/tests/router_unit_tests/test_router_embedding_headers.py new file mode 100644 index 00000000000..6d480792b7c --- /dev/null +++ b/tests/router_unit_tests/test_router_embedding_headers.py @@ -0,0 +1,372 @@ +""" +Test suite for router embedding method header propagation. + +This tests the fix for the issue where the embedding method was not +propagating proxy model configuration headers to the LLM API calls. + +The fix ensures that router.embedding() calls _update_kwargs_before_fallbacks() +just like router.completion() does, which properly sets up metadata and allows +default_litellm_params (including headers) to be propagated. +""" +import os +import sys +from unittest.mock import MagicMock, patch, AsyncMock + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm import Router + + +class TestRouterEmbeddingHeaders: + """Test that embedding methods properly propagate headers from router configuration.""" + + def test_embedding_calls_update_kwargs_before_fallbacks(self): + """ + Test that router.embedding() calls _update_kwargs_before_fallbacks. + + This ensures that metadata is properly set up before the fallback mechanism, + which is necessary for header propagation to work correctly. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + # Mock the _update_kwargs_before_fallbacks method to verify it's called + with patch.object( + router, + "_update_kwargs_before_fallbacks", + wraps=router._update_kwargs_before_fallbacks, + ) as mock_update: + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + # Verify _update_kwargs_before_fallbacks was called + mock_update.assert_called_once() + call_kwargs = mock_update.call_args[1] + assert call_kwargs["model"] == "text-embedding-ada-002" + assert "kwargs" in call_kwargs + + @pytest.mark.asyncio + async def test_aembedding_calls_update_kwargs_before_fallbacks(self): + """ + Test that router.aembedding() calls _update_kwargs_before_fallbacks. + + This ensures consistency between sync and async embedding methods. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + # Mock the _update_kwargs_before_fallbacks method to verify it's called + with patch.object( + router, + "_update_kwargs_before_fallbacks", + wraps=router._update_kwargs_before_fallbacks, + ) as mock_update: + with patch( + "litellm.aembedding", new_callable=AsyncMock + ) as mock_litellm_aembedding: + mock_litellm_aembedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + await router.aembedding( + model="text-embedding-ada-002", input=["test input"] + ) + + # Verify _update_kwargs_before_fallbacks was called + mock_update.assert_called_once() + call_kwargs = mock_update.call_args[1] + assert call_kwargs["model"] == "text-embedding-ada-002" + assert "kwargs" in call_kwargs + + def test_embedding_propagates_default_litellm_params(self): + """ + Test that embedding calls properly propagate default_litellm_params including headers. + + This is the main fix - ensuring that headers set in default_litellm_params + are included in the embedding request. + """ + custom_headers = {"X-Custom-Header": "test-value", "X-API-Version": "v2"} + + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + # Create router with default_litellm_params containing headers + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": custom_headers, + "metadata": {"test_key": "test_value"}, + }, + ) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + # Verify that litellm.embedding was called with the headers + mock_litellm_embedding.assert_called_once() + call_kwargs = mock_litellm_embedding.call_args[1] + + # Check that headers were included + assert "headers" in call_kwargs + assert call_kwargs["headers"] == custom_headers + + # Check that metadata was properly set up + assert "metadata" in call_kwargs + assert "model_group" in call_kwargs["metadata"] + assert call_kwargs["metadata"]["model_group"] == "text-embedding-ada-002" + + @pytest.mark.asyncio + async def test_aembedding_propagates_default_litellm_params(self): + """ + Test that async embedding calls properly propagate default_litellm_params including headers. + """ + custom_headers = {"X-Custom-Header": "test-value", "X-API-Version": "v2"} + + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + # Create router with default_litellm_params containing headers + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": custom_headers, + "metadata": {"test_key": "test_value"}, + }, + ) + + with patch( + "litellm.aembedding", new_callable=AsyncMock + ) as mock_litellm_aembedding: + mock_litellm_aembedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + await router.aembedding( + model="text-embedding-ada-002", input=["test input"] + ) + + # Verify that litellm.aembedding was called with the headers + mock_litellm_aembedding.assert_called_once() + call_kwargs = mock_litellm_aembedding.call_args[1] + + # Check that headers were included + assert "headers" in call_kwargs + assert call_kwargs["headers"] == custom_headers + + # Check that metadata was properly set up + assert "metadata" in call_kwargs + assert "model_group" in call_kwargs["metadata"] + assert call_kwargs["metadata"]["model_group"] == "text-embedding-ada-002" + + def test_embedding_metadata_includes_model_group(self): + """ + Test that embedding calls include model_group in metadata. + + The _update_kwargs_before_fallbacks method should set this up. + """ + model_list = [ + { + "model_name": "test-embedding-model", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="test-embedding-model", input=["test input"]) + + call_kwargs = mock_litellm_embedding.call_args[1] + + # Verify metadata contains model_group + assert "metadata" in call_kwargs + assert "model_group" in call_kwargs["metadata"] + assert call_kwargs["metadata"]["model_group"] == "test-embedding-model" + + def test_embedding_sets_num_retries_from_router(self): + """ + Test that embedding calls inherit num_retries from router configuration. + + This is set by _update_kwargs_before_fallbacks. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + # Create router with num_retries set + router = Router(model_list=model_list, num_retries=3) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + # Verify num_retries was not set in the call (it's handled by function_with_fallbacks) + # The important thing is that it was set in kwargs before being passed to function_with_fallbacks + # We verify this indirectly by checking that _update_kwargs_before_fallbacks was called + mock_litellm_embedding.assert_called_once() + + def test_embedding_sets_litellm_trace_id(self): + """ + Test that embedding calls include a litellm_trace_id. + + This is generated and set by _update_kwargs_before_fallbacks. + """ + model_list = [ + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + } + ] + + router = Router(model_list=model_list) + + with patch("litellm.embedding") as mock_litellm_embedding: + mock_litellm_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + call_kwargs = mock_litellm_embedding.call_args[1] + + # Verify litellm_trace_id was set + assert "litellm_trace_id" in call_kwargs + assert isinstance(call_kwargs["litellm_trace_id"], str) + assert len(call_kwargs["litellm_trace_id"]) > 0 + + def test_embedding_consistency_with_completion(self): + """ + Test that embedding and completion methods handle kwargs similarly. + + Both should call _update_kwargs_before_fallbacks to ensure consistent behavior. + """ + custom_headers = {"X-Test": "value"} + + model_list = [ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": { + "model": "gpt-3.5-turbo", + "api_key": "fake-key", + }, + }, + { + "model_name": "text-embedding-ada-002", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fake-key", + }, + }, + ] + + router = Router( + model_list=model_list, default_litellm_params={"headers": custom_headers} + ) + + # Test completion + with patch("litellm.completion") as mock_completion: + mock_completion.return_value = MagicMock() + + router.completion( + model="gpt-3.5-turbo", messages=[{"role": "user", "content": "test"}] + ) + + completion_kwargs = mock_completion.call_args[1] + + # Test embedding + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2, 0.3]}] + ) + + router.embedding(model="text-embedding-ada-002", input=["test input"]) + + embedding_kwargs = mock_embedding.call_args[1] + + # Both should have headers from default_litellm_params + assert "headers" in completion_kwargs + assert "headers" in embedding_kwargs + assert completion_kwargs["headers"] == custom_headers + assert embedding_kwargs["headers"] == custom_headers + + # Both should have metadata with model_group + assert "metadata" in completion_kwargs + assert "metadata" in embedding_kwargs + assert "model_group" in completion_kwargs["metadata"] + assert "model_group" in embedding_kwargs["metadata"] + + # Both should have litellm_trace_id + assert "litellm_trace_id" in completion_kwargs + assert "litellm_trace_id" in embedding_kwargs + + +if __name__ == "__main__": + # Run a simple test + test = TestRouterEmbeddingHeaders() + test.test_embedding_calls_update_kwargs_before_fallbacks() + test.test_embedding_propagates_default_litellm_params() + test.test_embedding_metadata_includes_model_group() + test.test_embedding_sets_litellm_trace_id() + test.test_embedding_consistency_with_completion() + print("All tests passed!") # noqa: T201 diff --git a/tests/router_unit_tests/test_router_embedding_integration.py b/tests/router_unit_tests/test_router_embedding_integration.py new file mode 100644 index 00000000000..ab2071714a9 --- /dev/null +++ b/tests/router_unit_tests/test_router_embedding_integration.py @@ -0,0 +1,355 @@ +""" +Integration tests for router embedding method with various configurations. + +These tests simulate real-world scenarios where headers and configuration +need to be properly propagated through the router to the LLM API. +""" +import os +import sys +from unittest.mock import MagicMock, patch, AsyncMock + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm import Router + + +class TestRouterEmbeddingIntegration: + """Integration tests for embedding with router configuration.""" + + def test_embedding_with_deployment_specific_headers(self): + """ + Test that deployment-specific headers are propagated. + + This simulates a scenario where different deployments have + different header requirements (e.g., different API versions). + """ + model_list = [ + { + "model_name": "embedding-deployment-1", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-1", + "headers": {"X-Deployment": "deployment-1"}, + }, + }, + { + "model_name": "embedding-deployment-2", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-2", + "headers": {"X-Deployment": "deployment-2"}, + }, + }, + ] + + router = Router(model_list=model_list) + + # Test first deployment + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="embedding-deployment-1", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + assert call_kwargs["api_key"] == "key-1" + + # Test second deployment + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="embedding-deployment-2", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + assert call_kwargs["api_key"] == "key-2" + + def test_embedding_with_router_and_deployment_headers_merge(self): + """ + Test that router-level headers are propagated. + + When no request headers are provided, router default headers should be used. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": { + "X-Router-Header": "router-value", + "X-Common-Header": "router-common", + } + }, + ) + + # Test: No request headers - router headers should be used + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding( + model="test-embedding", + input=["test"], + ) + + call_kwargs = mock_embedding.call_args[1] + + # Router headers should be present + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Router-Header"] == "router-value" + assert call_kwargs["headers"]["X-Common-Header"] == "router-common" + + def test_embedding_metadata_propagation(self): + """ + Test that metadata is properly set up and propagated. + + This is important for logging, tracking, and debugging. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router( + model_list=model_list, + default_litellm_params={ + "metadata": {"environment": "test", "service": "embedding-service"} + }, + ) + + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding( + model="test-embedding", + input=["test"], + metadata={"request_id": "req-123"}, # Additional metadata from request + ) + + call_kwargs = mock_embedding.call_args[1] + + # Check metadata contains all expected fields + assert "metadata" in call_kwargs + metadata = call_kwargs["metadata"] + + # From _update_kwargs_before_fallbacks + assert "model_group" in metadata + assert metadata["model_group"] == "test-embedding" + + # From default_litellm_params + assert "environment" in metadata + assert metadata["environment"] == "test" + assert "service" in metadata + assert metadata["service"] == "embedding-service" + + # From request + assert "request_id" in metadata + assert metadata["request_id"] == "req-123" + + @pytest.mark.asyncio + async def test_async_embedding_with_multiple_retries(self): + """ + Test that async embedding properly uses num_retries from router config. + + This ensures the fix works with the retry mechanism. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router(model_list=model_list, num_retries=2) + + with patch("litellm.aembedding", new_callable=AsyncMock) as mock_aembedding: + mock_aembedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + await router.aembedding(model="test-embedding", input=["test"]) + + # The call should succeed + mock_aembedding.assert_called_once() + + def test_embedding_with_timeout_from_router(self): + """ + Test that timeout settings from router config are propagated. + """ + model_list = [ + { + "model_name": "test-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "test-key", + }, + } + ] + + router = Router(model_list=model_list, timeout=30.0) + + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="test-embedding", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + + # Timeout should be set from router config + assert "timeout" in call_kwargs + assert call_kwargs["timeout"] == 30.0 + + def test_embedding_with_multiple_deployments_load_balancing(self): + """ + Test that headers are correctly propagated when router load balances + between multiple deployments. + """ + model_list = [ + { + "model_name": "shared-embedding-model", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-1", + }, + }, + { + "model_name": "shared-embedding-model", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "key-2", + }, + }, + ] + + router = Router( + model_list=model_list, + default_litellm_params={"headers": {"X-Shared-Header": "shared-value"}}, + ) + + # Make multiple calls and verify headers are always present + for i in range(5): + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock( + data=[{"embedding": [0.1, 0.2]}] + ) + + router.embedding(model="shared-embedding-model", input=[f"test {i}"]) + + call_kwargs = mock_embedding.call_args[1] + + # Headers should always be present regardless of which deployment is chosen + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Shared-Header"] == "shared-value" + + @pytest.mark.asyncio + async def test_embedding_with_fallback_configuration(self): + """ + Test that headers are propagated correctly when using fallback models. + """ + model_list = [ + { + "model_name": "primary-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "primary-key", + }, + }, + { + "model_name": "fallback-embedding", + "litellm_params": { + "model": "text-embedding-ada-002", + "api_key": "fallback-key", + }, + }, + ] + + router = Router( + model_list=model_list, + fallbacks=[{"primary-embedding": ["fallback-embedding"]}], + default_litellm_params={"headers": {"X-Fallback-Test": "test-value"}}, + ) + + # Simulate primary failing, fallback succeeding + with patch("litellm.aembedding", new_callable=AsyncMock) as mock_aembedding: + call_count = 0 + + async def side_effect(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + # First call (primary) fails + raise Exception("Primary failed") + else: + # Second call (fallback) succeeds + return MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + mock_aembedding.side_effect = side_effect + + await router.aembedding(model="primary-embedding", input=["test"]) + + # Both calls should have headers + assert mock_aembedding.call_count == 2 + + # Check that both calls had headers + for call_obj in mock_aembedding.call_args_list: + call_kwargs = call_obj[1] + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Fallback-Test"] == "test-value" + + def test_embedding_with_custom_provider_headers(self): + """ + Test that provider-specific headers are correctly propagated. + + Some providers require specific headers for API versioning, features, etc. + """ + model_list = [ + { + "model_name": "azure-embedding", + "litellm_params": { + "model": "azure/text-embedding-ada-002", + "api_key": "azure-key", + "api_base": "https://example.openai.azure.com", + "api_version": "2024-02-01", + }, + } + ] + + router = Router( + model_list=model_list, + default_litellm_params={ + "headers": {"X-Custom-Azure-Header": "azure-value"} + }, + ) + + with patch("litellm.embedding") as mock_embedding: + mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}]) + + router.embedding(model="azure-embedding", input=["test"]) + + call_kwargs = mock_embedding.call_args[1] + + # Verify Azure-specific params are present + assert call_kwargs["api_base"] == "https://example.openai.azure.com" + assert call_kwargs["api_version"] == "2024-02-01" + + # Verify custom headers are present + assert "headers" in call_kwargs + assert call_kwargs["headers"]["X-Custom-Azure-Header"] == "azure-value" + + +if __name__ == "__main__": + # Run tests + pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 596398e639f..f8a082ee30c 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -1095,3 +1095,186 @@ def test_map_reasoning_effort_adds_summary_detailed(): os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = original_env elif "LITELLM_REASONING_AUTO_SUMMARY" in os.environ: del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] + + +def test_transform_response_preserves_annotations(): + """ + Test that annotations from Responses API are preserved when transforming to Chat Completions format. + + This is a regression test for the bug where annotations (like url_citation) were being + dropped during the transformation from ResponsesAPIResponse to ModelResponse. + + The fix ensures annotations are extracted from ResponseOutputText content items and + passed through to the Message object in the Chat Completions response. + """ + from unittest.mock import Mock + + from openai.types.responses import ResponseOutputMessage, ResponseOutputText + + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + from litellm.types.llms.openai import ( + InputTokensDetails, + OutputTokensDetails, + ResponseAPIUsage, + ResponsesAPIResponse, + ) + from litellm.types.utils import ModelResponse, Usage + + handler = LiteLLMResponsesTransformationHandler() + + # Create annotations similar to what OpenAI Responses API returns + annotations = [ + { + "type": "url_citation", + "start_index": 0, + "end_index": 100, + "title": "Example Article", + "url": "https://example.com/article", + }, + { + "type": "url_citation", + "start_index": 101, + "end_index": 200, + "title": "Another Source", + "url": "https://example.com/source", + }, + ] + + # Create output text with annotations + output_text = ResponseOutputText( + annotations=annotations, + text="Here is some information with citations.", + type="output_text", + logprobs=[], + ) + + # Create output message + output_message = ResponseOutputMessage( + id="msg_test123", + content=[output_text], + role="assistant", + status="completed", + type="message", + ) + + # Create usage information + usage = ResponseAPIUsage( + input_tokens=10, + input_tokens_details=InputTokensDetails( + audio_tokens=None, cached_tokens=0, text_tokens=None + ), + output_tokens=20, + output_tokens_details=OutputTokensDetails( + reasoning_tokens=0, text_tokens=None + ), + total_tokens=30, + cost=None, + ) + + # Create the full ResponsesAPIResponse + raw_response = ResponsesAPIResponse( + id="resp_test123", + created_at=1234567890, + error=None, + incomplete_details=None, + instructions=None, + metadata={}, + model="gpt-5.1", + object="response", + output=[output_message], + parallel_tool_calls=True, + temperature=1.0, + tool_choice="auto", + tools=[], + top_p=1.0, + max_output_tokens=None, + previous_response_id=None, + reasoning=None, + status="completed", + text={"format": {"type": "text"}, "verbosity": "medium"}, + truncation="disabled", + usage=usage, + user=None, + store=True, + background=False, + billing={"payer": "openai"}, + max_tool_calls=None, + prompt_cache_key=None, + safety_identifier=None, + service_tier="default", + top_logprobs=0, + ) + + # Create empty model_response + model_response = ModelResponse( + id="chatcmpl-test123", + created=1234567890, + model=None, + object="chat.completion", + system_fingerprint=None, + choices=[], + usage=Usage(completion_tokens=0, prompt_tokens=0, total_tokens=0), + ) + + # Create mock objects for required parameters + logging_obj = Mock() + messages = [{"role": "user", "content": "Tell me about AI"}] + request_data = {"model": "gpt-5.1"} + optional_params = {} + litellm_params = {"acompletion": False, "api_key": None} + encoding = Mock() + + # Call transform_response + result = handler.transform_response( + model="gpt-5.1", + raw_response=raw_response, + model_response=model_response, + logging_obj=logging_obj, + request_data=request_data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + api_key=None, + json_mode=None, + ) + + # Assertions + assert result.model == "gpt-5.1" + assert len(result.choices) == 1 + + # Check the choice + choice = result.choices[0] + assert choice.finish_reason == "stop" + assert choice.index == 0 + assert choice.message.role == "assistant" + assert choice.message.content == "Here is some information with citations." + + # Check that annotations are preserved + assert hasattr(choice.message, "annotations"), "Message should have annotations attribute" + assert choice.message.annotations is not None, "Annotations should not be None" + assert len(choice.message.annotations) == 2, f"Expected 2 annotations, got {len(choice.message.annotations)}" + + # Verify annotation content + annotation1 = choice.message.annotations[0] + assert annotation1["type"] == "url_citation" + assert annotation1["title"] == "Example Article" + assert annotation1["url"] == "https://example.com/article" + assert annotation1["start_index"] == 0 + assert annotation1["end_index"] == 100 + + annotation2 = choice.message.annotations[1] + assert annotation2["type"] == "url_citation" + assert annotation2["title"] == "Another Source" + assert annotation2["url"] == "https://example.com/source" + assert annotation2["start_index"] == 101 + assert annotation2["end_index"] == 200 + + # Check usage + assert result.usage.prompt_tokens == 10 + assert result.usage.completion_tokens == 20 + assert result.usage.total_tokens == 30 + + print("✓ Annotations from Responses API are correctly preserved in Chat Completions format") diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py index 31a2f6cbf51..b0aac17e7d9 100644 --- a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py +++ b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py @@ -66,16 +66,77 @@ class TestCloudZeroHourlyExport: fake_client = MagicMock() fake_db = MagicMock() - async def query_raw_mock(query: str): - sql_context = pl.SQLContext( - LiteLLM_DailyUserSpend=spend_mock_data, - LiteLLM_VerificationToken=verification_mock_data, - LiteLLM_TeamTable=team_mock_data, - LiteLLM_UserTable=user_mock_data, - ) - result = sql_context.execute(query).collect() + async def query_raw_mock(query: str, *params): + start_time_utc = params[0] if len(params) > 0 else None + end_time_utc = params[1] if len(params) > 1 else None + limit = params[2] if len(params) > 2 else None - return result + spend_df = spend_mock_data.collect() + verification_df = verification_mock_data.collect().rename( + {"key_alias": "api_key_alias"} + ) + team_df = team_mock_data.collect() + user_df = user_mock_data.collect() + + joined = ( + spend_df.join( + verification_df, left_on="api_key", right_on="token", how="left" + ) + .join( + team_df, + left_on="team_id", + right_on="team_id", + how="left", + suffix="_team", + ) + .join( + user_df, + left_on="user_id", + right_on="user_id", + how="left", + suffix="_user", + ) + ) + + for duplicate_column in ("team_id_team", "user_id_user"): + if duplicate_column in joined.columns: + joined = joined.drop(duplicate_column) + + if start_time_utc is not None: + joined = joined.filter(pl.col("updated_at") >= start_time_utc) + if end_time_utc is not None: + joined = joined.filter(pl.col("updated_at") <= end_time_utc) + + joined = joined.select( + [ + "id", + "date", + "user_id", + "api_key", + "model", + "model_group", + "custom_llm_provider", + "prompt_tokens", + "completion_tokens", + "spend", + "api_requests", + "successful_requests", + "failed_requests", + "cache_creation_input_tokens", + "cache_read_input_tokens", + "created_at", + "updated_at", + "team_id", + "api_key_alias", + "team_alias", + "user_email", + ] + ).sort(["date", "created_at"], descending=[True, True]) + + if limit is not None: + joined = joined.head(int(limit)) + + return joined fake_db.query_raw = AsyncMock(side_effect=query_raw_mock) fake_client.db = fake_db diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py b/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py new file mode 100644 index 00000000000..89a5028011c --- /dev/null +++ b/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py @@ -0,0 +1,57 @@ +"""Tests for LiteLLM CloudZero database helper.""" + +from datetime import datetime, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from litellm.integrations.cloudzero.database import LiteLLMDatabase + + +def _setup_db(monkeypatch: pytest.MonkeyPatch, query_return): + """Return a database instance with prisma client mocked out.""" + query_mock = AsyncMock(return_value=query_return) + mock_client = SimpleNamespace(db=SimpleNamespace(query_raw=query_mock)) + db = LiteLLMDatabase() + monkeypatch.setattr(db, "_ensure_prisma_client", lambda: mock_client) + return db, query_mock + + +@pytest.mark.asyncio +async def test_get_usage_data_parameterized(monkeypatch: pytest.MonkeyPatch): + """Start/end filters and limit should be parameterized via placeholders.""" + start = datetime(2024, 5, 1, tzinfo=timezone.utc) + end = datetime(2024, 5, 2, tzinfo=timezone.utc) + db, query_mock = _setup_db(monkeypatch, []) + + await db.get_usage_data(limit=10, start_time_utc=start, end_time_utc=end) + + query_text, *params = query_mock.await_args.args + assert "dus.updated_at >= $1::timestamptz" in query_text + assert "dus.updated_at <= $2::timestamptz" in query_text + assert "LIMIT $3" in query_text + assert params == [start, end, 10] + + +@pytest.mark.asyncio +async def test_get_usage_data_handles_missing_filters(monkeypatch: pytest.MonkeyPatch): + """When no filters provided the params should be None placeholders.""" + db, query_mock = _setup_db(monkeypatch, []) + + await db.get_usage_data() + + query_text, *params = query_mock.await_args.args + assert "LIMIT $3" not in query_text + assert params == [None, None] + + +@pytest.mark.asyncio +async def test_get_usage_data_rejects_invalid_limit(monkeypatch: pytest.MonkeyPatch): + """limit must coerce to int or raise ValueError before hitting the DB.""" + db, query_mock = _setup_db(monkeypatch, []) + + with pytest.raises(ValueError): + await db.get_usage_data(limit="invalid") + + assert query_mock.await_count == 0 diff --git a/tests/test_litellm/integrations/focus/test_focus_database.py b/tests/test_litellm/integrations/focus/test_focus_database.py new file mode 100644 index 00000000000..5ee98cc9dd0 --- /dev/null +++ b/tests/test_litellm/integrations/focus/test_focus_database.py @@ -0,0 +1,74 @@ +"""Tests for FocusLiteLLMDatabase query construction.""" + +from datetime import datetime, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from litellm.integrations.focus.database import FocusLiteLLMDatabase + + +def _setup_db(monkeypatch: pytest.MonkeyPatch, query_return): + """Create a database instance with a stubbed prisma client.""" + query_mock = AsyncMock(return_value=query_return) + mock_client = SimpleNamespace(db=SimpleNamespace(query_raw=query_mock)) + db = FocusLiteLLMDatabase() + monkeypatch.setattr(db, "_ensure_prisma_client", lambda: mock_client) + return db, query_mock + + +@pytest.mark.asyncio +async def test_should_parameterize_filters_and_limit(monkeypatch: pytest.MonkeyPatch): + start = datetime(2024, 1, 1, tzinfo=timezone.utc) + end = datetime(2024, 1, 2, tzinfo=timezone.utc) + db, query_mock = _setup_db(monkeypatch, []) + + await db.get_usage_data(limit=25, start_time_utc=start, end_time_utc=end) + + query_text, *params = query_mock.await_args.args + assert "dus.updated_at >= $1::timestamptz" in query_text + assert "dus.updated_at <= $2::timestamptz" in query_text + assert "LIMIT $3" in query_text + assert params == [start, end, 25] + + +@pytest.mark.asyncio +async def test_should_execute_without_filters(monkeypatch: pytest.MonkeyPatch): + row = { + "id": 1, + "user_id": "user", + "date": datetime(2024, 1, 1, tzinfo=timezone.utc), + } + db, query_mock = _setup_db(monkeypatch, [row]) + + result = await db.get_usage_data() + + query_text, *params = query_mock.await_args.args + assert "WHERE" not in query_text + assert "LIMIT $" not in query_text + assert params == [] + assert result.height == 1 + assert result["id"][0] == 1 + + +@pytest.mark.asyncio +async def test_should_accept_string_timestamps(monkeypatch: pytest.MonkeyPatch): + db, query_mock = _setup_db(monkeypatch, []) + + start = "2024-02-01T00:00:00+00:00" + end = "2024-02-02T00:00:00+00:00" + await db.get_usage_data(start_time_utc=start, end_time_utc=end) + + _, *params = query_mock.await_args.args + assert params == [start, end] + + +@pytest.mark.asyncio +async def test_should_reject_invalid_limit(monkeypatch: pytest.MonkeyPatch): + db, query_mock = _setup_db(monkeypatch, []) + + with pytest.raises(ValueError): + await db.get_usage_data(limit="invalid") + + assert query_mock.await_count == 0 diff --git a/tests/test_litellm/integrations/focus/test_s3_destination.py b/tests/test_litellm/integrations/focus/test_s3_destination.py new file mode 100644 index 00000000000..f915b2c56a3 --- /dev/null +++ b/tests/test_litellm/integrations/focus/test_s3_destination.py @@ -0,0 +1,100 @@ +"""Tests for FocusS3Destination behavior.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from types import SimpleNamespace +from typing import Any, Dict + +import pytest + +import litellm.integrations.focus.destinations.s3_destination as s3_module +from litellm.integrations.focus.destinations.base import FocusTimeWindow +from litellm.integrations.focus.destinations.s3_destination import FocusS3Destination + + +def _window(freq: str = "hourly", hour: int = 5) -> FocusTimeWindow: + start = datetime(2024, 1, 2, hour, tzinfo=timezone.utc) + end = start.replace(hour=hour + 1) + return FocusTimeWindow(start_time=start, end_time=end, frequency=freq) + + +def test_should_require_bucket_name(): + with pytest.raises(ValueError): + FocusS3Destination(prefix="focus", config={}) + + +def test_should_build_hourly_object_key(): + dest = FocusS3Destination(prefix="exports/", config={"bucket_name": "bucket"}) + key = dest._build_object_key( + time_window=_window(freq="hourly", hour=3), filename="data.snappy" + ) + assert key == "exports/date=2024-01-02/hour=03/data.snappy" + + +def test_should_build_daily_key_without_hour_segment(): + dest = FocusS3Destination(prefix="", config={"bucket_name": "bucket"}) + key = dest._build_object_key( + time_window=_window(freq="daily", hour=0), filename="daily.parquet" + ) + assert key == "date=2024-01-02/daily.parquet" + + +@pytest.mark.asyncio +async def test_should_dispatch_upload_via_thread(monkeypatch: pytest.MonkeyPatch): + dest = FocusS3Destination(prefix="focus", config={"bucket_name": "bucket"}) + captured: Dict[str, Any] = {} + + async def fake_to_thread(func, *args, **kwargs): # type: ignore[override] + captured["func"] = func + captured["args"] = args + captured["kwargs"] = kwargs + + monkeypatch.setattr(s3_module.asyncio, "to_thread", fake_to_thread) + + window = _window(freq="hourly", hour=1) + await dest.deliver(content=b"payload", time_window=window, filename="file.bin") + + assert captured["func"] == dest._upload + assert captured["args"][0] == b"payload" + assert captured["args"][1].endswith("/file.bin") + + +def test_should_upload_with_configured_client(monkeypatch: pytest.MonkeyPatch): + config = { + "bucket_name": "bucket", + "region_name": "us-east-2", + "endpoint_url": "http://localhost:4566", + "aws_access_key_id": "key", + "aws_secret_access_key": "secret", + "aws_session_token": "token", + } + dest = FocusS3Destination(prefix="focus", config=config) + captured: Dict[str, Any] = {} + + def fake_client(service: str, **kwargs): + assert service == "s3" + captured["client_kwargs"] = kwargs + + def put_object(**put_kwargs): + captured["put_kwargs"] = put_kwargs + + return SimpleNamespace(put_object=put_object) + + monkeypatch.setattr(s3_module.boto3, "client", fake_client) + + dest._upload(content=b"payload", object_key="path/file.bin") + + assert captured["client_kwargs"] == { + "region_name": "us-east-2", + "endpoint_url": "http://localhost:4566", + "aws_access_key_id": "key", + "aws_secret_access_key": "secret", + "aws_session_token": "token", + } + assert captured["put_kwargs"] == { + "Bucket": "bucket", + "Key": "path/file.bin", + "Body": b"payload", + "ContentType": "application/octet-stream", + } diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index a322dfe9a2b..d7d7720ff43 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -530,7 +530,7 @@ class TestPassthroughCallTypeHandling: ) assert ( ProxyBaseLLMRequestProcessing._get_pre_call_type(route_type="aembedding") - == "embeddings" + == "embedding" ) assert ( ProxyBaseLLMRequestProcessing._get_pre_call_type(route_type="aresponses") diff --git a/tests/test_litellm/integrations/test_prometheus_cache_metrics.py b/tests/test_litellm/integrations/test_prometheus_cache_metrics.py new file mode 100644 index 00000000000..660757673f6 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_cache_metrics.py @@ -0,0 +1,211 @@ +""" +Unit tests for cache Prometheus metrics. + +Run with: poetry run pytest tests/test_litellm/integrations/test_prometheus_cache_metrics.py -v +""" +import pytest +from unittest.mock import MagicMock, patch +from litellm.types.integrations.prometheus import UserAPIKeyLabelValues + + +class TestPrometheusCacheMetrics: + """Tests for cache-related Prometheus metrics""" + + @pytest.fixture + def sample_enum_values(self): + """Create sample enum values for labels""" + return UserAPIKeyLabelValues( + end_user="test-end-user", + hashed_api_key="test-key-hash", + api_key_alias="test-key-alias", + team="test-team", + team_alias="test-team-alias", + user="test-user", + model="gpt-3.5-turbo", + ) + + def test_cache_metrics_defined_in_types(self): + """Test that cache metrics are defined in DEFINED_PROMETHEUS_METRICS""" + from litellm.types.integrations.prometheus import DEFINED_PROMETHEUS_METRICS + from typing import get_args + + defined_metrics = get_args(DEFINED_PROMETHEUS_METRICS) + + assert "litellm_cache_hits_metric" in defined_metrics + assert "litellm_cache_misses_metric" in defined_metrics + assert "litellm_cached_tokens_metric" in defined_metrics + + def test_cache_metric_labels_defined(self): + """Test that cache metric labels are properly defined""" + from litellm.types.integrations.prometheus import PrometheusMetricLabels + + # Verify labels are defined for each cache metric + assert hasattr(PrometheusMetricLabels, "litellm_cache_hits_metric") + assert hasattr(PrometheusMetricLabels, "litellm_cache_misses_metric") + assert hasattr(PrometheusMetricLabels, "litellm_cached_tokens_metric") + + # Verify labels include expected keys + expected_labels = [ + "model", + "hashed_api_key", + "api_key_alias", + "team", + "team_alias", + "end_user", + "user", + ] + for label in expected_labels: + assert label in PrometheusMetricLabels.litellm_cache_hits_metric + assert label in PrometheusMetricLabels.litellm_cache_misses_metric + assert label in PrometheusMetricLabels.litellm_cached_tokens_metric + + def test_increment_cache_metrics_on_cache_hit(self, sample_enum_values): + """Test that cache hit increments the correct metrics""" + # Create mock for PrometheusLogger instance + mock_logger = MagicMock() + + # Import the method directly and bind it to our mock + from litellm.integrations.prometheus import PrometheusLogger + + # Create a mock standard logging payload with cache_hit=True + standard_logging_payload = { + "cache_hit": True, + "total_tokens": 100, + "prompt_tokens": 50, + "completion_tokens": 50, + "model_group": "openai", + "request_tags": [], + } + + # Create mock metrics + mock_logger.litellm_cache_hits_metric = MagicMock() + mock_logger.litellm_cache_misses_metric = MagicMock() + mock_logger.litellm_cached_tokens_metric = MagicMock() + mock_logger.get_labels_for_metric = MagicMock( + return_value=[ + "model", + "hashed_api_key", + "api_key_alias", + "team", + "team_alias", + "end_user", + "user", + ] + ) + + # Call the method using unbound method approach + PrometheusLogger._increment_cache_metrics( + mock_logger, + standard_logging_payload=standard_logging_payload, + enum_values=sample_enum_values, + ) + + # Verify cache hits metric was incremented + mock_logger.litellm_cache_hits_metric.labels.assert_called() + mock_logger.litellm_cache_hits_metric.labels().inc.assert_called_once() + + # Verify cached tokens metric was incremented with total_tokens + mock_logger.litellm_cached_tokens_metric.labels.assert_called() + mock_logger.litellm_cached_tokens_metric.labels().inc.assert_called_once_with( + 100 + ) + + # Verify cache misses metric was NOT called + mock_logger.litellm_cache_misses_metric.labels.assert_not_called() + + def test_increment_cache_metrics_on_cache_miss(self, sample_enum_values): + """Test that cache miss increments the correct metrics""" + # Create mock for PrometheusLogger instance + mock_logger = MagicMock() + + from litellm.integrations.prometheus import PrometheusLogger + + # Create a mock standard logging payload with cache_hit=False + standard_logging_payload = { + "cache_hit": False, + "total_tokens": 100, + "prompt_tokens": 50, + "completion_tokens": 50, + "model_group": "openai", + "request_tags": [], + } + + # Create mock metrics + mock_logger.litellm_cache_hits_metric = MagicMock() + mock_logger.litellm_cache_misses_metric = MagicMock() + mock_logger.litellm_cached_tokens_metric = MagicMock() + mock_logger.get_labels_for_metric = MagicMock( + return_value=[ + "model", + "hashed_api_key", + "api_key_alias", + "team", + "team_alias", + "end_user", + "user", + ] + ) + + # Call the method + PrometheusLogger._increment_cache_metrics( + mock_logger, + standard_logging_payload=standard_logging_payload, + enum_values=sample_enum_values, + ) + + # Verify cache misses metric was incremented + mock_logger.litellm_cache_misses_metric.labels.assert_called() + mock_logger.litellm_cache_misses_metric.labels().inc.assert_called_once() + + # Verify cache hits and cached tokens metrics were NOT called + mock_logger.litellm_cache_hits_metric.labels.assert_not_called() + mock_logger.litellm_cached_tokens_metric.labels.assert_not_called() + + def test_increment_cache_metrics_when_cache_hit_is_none(self, sample_enum_values): + """Test that no metrics are incremented when cache_hit is None""" + # Create mock for PrometheusLogger instance + mock_logger = MagicMock() + + from litellm.integrations.prometheus import PrometheusLogger + + # Create a mock standard logging payload with cache_hit=None + standard_logging_payload = { + "cache_hit": None, + "total_tokens": 100, + "prompt_tokens": 50, + "completion_tokens": 50, + "model_group": "openai", + "request_tags": [], + } + + # Create mock metrics + mock_logger.litellm_cache_hits_metric = MagicMock() + mock_logger.litellm_cache_misses_metric = MagicMock() + mock_logger.litellm_cached_tokens_metric = MagicMock() + mock_logger.get_labels_for_metric = MagicMock( + return_value=[ + "model", + "hashed_api_key", + "api_key_alias", + "team", + "team_alias", + "end_user", + "user", + ] + ) + + # Call the method + PrometheusLogger._increment_cache_metrics( + mock_logger, + standard_logging_payload=standard_logging_payload, + enum_values=sample_enum_values, + ) + + # Verify NO metrics were called + mock_logger.litellm_cache_hits_metric.labels.assert_not_called() + mock_logger.litellm_cache_misses_metric.labels.assert_not_called() + mock_logger.litellm_cached_tokens_metric.labels.assert_not_called() + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/integrations/test_prometheus_invalid_key_filtering.py b/tests/test_litellm/integrations/test_prometheus_invalid_key_filtering.py new file mode 100644 index 00000000000..ff433480d5e --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_invalid_key_filtering.py @@ -0,0 +1,161 @@ +""" +Unit tests for Prometheus invalid API key request filtering. + +Tests functionality that prevents invalid API key requests (401 status codes) +from being recorded in Prometheus metrics. +""" + +import os +import sys +from unittest.mock import Mock, patch + +import pytest +from prometheus_client import REGISTRY + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.integrations.prometheus import PrometheusLogger +from litellm.proxy._types import UserAPIKeyAuth + + +@pytest.fixture(scope="function") +def prometheus_logger(): + """Create a PrometheusLogger instance for testing.""" + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + REGISTRY.unregister(collector) + return PrometheusLogger() + + +class ExceptionWithCode: + """Exception-like object with 'code' attribute (ProxyException pattern).""" + def __init__(self, code): + self.code = code + + +class ExceptionWithStatusCode: + """Exception-like object with 'status_code' attribute.""" + def __init__(self, status_code): + self.status_code = status_code + + +class TestExtractStatusCode: + """Test status code extraction from various sources.""" + + @pytest.mark.parametrize("exception_class,code_value,expected", [ + (ExceptionWithCode, "401", 401), + (ExceptionWithStatusCode, 401, 401), + ]) + def test_extract_from_exception(self, prometheus_logger, exception_class, code_value, expected): + exception = exception_class(code_value) + assert prometheus_logger._extract_status_code(exception=exception) == expected + + def test_extract_from_kwargs(self, prometheus_logger): + exception = ExceptionWithCode("401") + assert prometheus_logger._extract_status_code(kwargs={"exception": exception}) == 401 + + def test_extract_from_enum_values(self, prometheus_logger): + enum_values = Mock(status_code="401") + assert prometheus_logger._extract_status_code(enum_values=enum_values) == 401 + + +class TestInvalidAPIKeyDetection: + """Test invalid API key request detection logic.""" + + @pytest.mark.parametrize("status_code,expected", [ + (401, True), + (200, False), + (500, False), + (None, False), + ]) + def test_status_code_detection(self, prometheus_logger, status_code, expected): + assert prometheus_logger._is_invalid_api_key_request(status_code=status_code) == expected + + def test_auth_error_message_detection(self, prometheus_logger): + exception = AssertionError("LiteLLM Virtual Key expected. Received=invalid-key-12345, expected to start with 'sk-'.") + assert prometheus_logger._is_invalid_api_key_request(status_code=None, exception=exception) is True + + def test_non_auth_exception_not_detected(self, prometheus_logger): + exception = ValueError("Some other error") + assert prometheus_logger._is_invalid_api_key_request(status_code=None, exception=exception) is False + + +class TestSkipMetricsValidation: + """Test high-level validation method that orchestrates detection and extraction.""" + + def test_skip_for_401_exception(self, prometheus_logger): + """Test full flow: extraction -> detection -> skip decision.""" + exception = ExceptionWithCode("401") + assert prometheus_logger._should_skip_metrics_for_invalid_key(exception=exception) is True + + def test_skip_for_auth_error_message(self, prometheus_logger): + """Test full flow: exception message -> detection -> skip decision.""" + exception = AssertionError("expected to start with 'sk-'") + assert prometheus_logger._should_skip_metrics_for_invalid_key(exception=exception) is True + + def test_no_skip_for_valid_request(self, prometheus_logger): + assert prometheus_logger._should_skip_metrics_for_invalid_key() is False + + +class TestAsyncHooks: + """Test async hook methods skip metrics for invalid API keys.""" + + @pytest.fixture + def mock_user_api_key(self): + """Create a mock UserAPIKeyAuth object.""" + user_key = Mock(spec=UserAPIKeyAuth) + user_key.api_key = "test-key" + user_key.end_user_id = None + user_key.user_id = None + user_key.user_email = None + user_key.key_alias = None + user_key.team_id = None + user_key.team_alias = None + user_key.request_route = "/test" + return user_key + + @pytest.mark.asyncio + async def test_post_call_failure_hook_skips_401(self, prometheus_logger, mock_user_api_key): + exception = ExceptionWithCode("401") + exception.__class__.__name__ = "ProxyException" + + with patch.object(prometheus_logger, 'litellm_proxy_failed_requests_metric') as mock_failed, \ + patch.object(prometheus_logger, 'litellm_proxy_total_requests_metric') as mock_total: + + await prometheus_logger.async_post_call_failure_hook( + request_data={"model": "test-model"}, + original_exception=exception, + user_api_key_dict=mock_user_api_key + ) + + mock_failed.labels.assert_not_called() + mock_total.labels.assert_not_called() + + @pytest.mark.asyncio + async def test_log_failure_event_skips_401(self, prometheus_logger): + exception = ExceptionWithCode("401") + kwargs = { + "model": "test-model", + "standard_logging_object": { + "metadata": { + "user_api_key_hash": "test-key", + "user_api_key_user_id": "test-user", + }, + "model_group": "test-model", + }, + "exception": exception, + "litellm_params": {}, + } + + with patch.object(prometheus_logger, 'litellm_llm_api_failed_requests_metric') as mock_failed, \ + patch.object(prometheus_logger, 'set_llm_deployment_failure_metrics') as mock_deployment: + + await prometheus_logger.async_log_failure_event( + kwargs=kwargs, + response_obj=None, + start_time=None, + end_time=None + ) + + mock_failed.labels.assert_not_called() + mock_deployment.assert_not_called() diff --git a/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py b/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py new file mode 100644 index 00000000000..0743a9c7ba2 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py @@ -0,0 +1,424 @@ +""" +Unit tests for prometheus queue time and guardrail metrics +""" +from datetime import datetime +from unittest.mock import MagicMock + +import pytest +from prometheus_client import REGISTRY + +from litellm.integrations.prometheus import PrometheusLogger +from litellm.types.integrations.prometheus import UserAPIKeyLabelValues + + +@pytest.fixture(autouse=True) +def cleanup_prometheus_registry(): + """Clean up prometheus registry between tests""" + # Clear the registry before each test + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + REGISTRY.unregister(collector) + yield + # Clean up after test + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + REGISTRY.unregister(collector) + + +class TestPrometheusQueueTimeMetric: + """Test request queue time metric recording""" + + def test_queue_time_metric_recorded_in_set_latency_metrics(self): + """Test that queue time metric is recorded when queue_time_seconds is present in metadata""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock the metric + mock_metric = MagicMock() + mock_labeled_metric = MagicMock() + mock_metric.labels.return_value = mock_labeled_metric + prometheus_logger.litellm_request_queue_time_metric = mock_metric + + # Create mock kwargs with queue_time_seconds in metadata + queue_time_seconds = 0.5 + + kwargs = { + "litellm_params": {"metadata": {"queue_time_seconds": queue_time_seconds}}, + "model": "gpt-3.5-turbo", + "start_time": datetime.now(), + "end_time": datetime.now(), + } + + enum_values = UserAPIKeyLabelValues( + end_user=None, + hashed_api_key="test-key", + api_key_alias="test-alias", + requested_model="gpt-3.5-turbo", + model_group="gpt-3.5-turbo", + team=None, + team_alias=None, + user=None, + user_email=None, + status_code="200", + model="gpt-3.5-turbo", + litellm_model_name="gpt-3.5-turbo", + tags=[], + model_id="gpt-3.5-turbo", + api_base="https://api.openai.com", + api_provider="openai", + exception_status=None, + exception_class=None, + custom_metadata_labels={}, + route=None, + ) + + # Act + prometheus_logger._set_latency_metrics( + kwargs=kwargs, + model="gpt-3.5-turbo", + user_api_key="test-key", + user_api_key_alias="test-alias", + user_api_team=None, + user_api_team_alias=None, + enum_values=enum_values, + ) + + # Assert - queue time metric should be called + mock_metric.labels.assert_called() + # Check that observe was called on the queue time metric + assert mock_labeled_metric.observe.called + # Verify the observed value + observed_value = None + for call in mock_labeled_metric.observe.call_args_list: + if len(call[0]) > 0: + observed_value = call[0][0] + if observed_value == queue_time_seconds: + break + assert observed_value == queue_time_seconds + assert observed_value >= 0 + + def test_queue_time_metric_not_recorded_when_missing(self): + """Test that queue time metric is not recorded when queue_time_seconds is missing""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock the metric + mock_metric = MagicMock() + mock_labeled_metric = MagicMock() + mock_metric.labels.return_value = mock_labeled_metric + prometheus_logger.litellm_request_queue_time_metric = mock_metric + + # Create mock kwargs without queue_time_seconds + kwargs = { + "litellm_params": {"metadata": {}}, + "model": "gpt-3.5-turbo", + "start_time": datetime.now(), + "end_time": datetime.now(), + } + + enum_values = UserAPIKeyLabelValues( + end_user=None, + hashed_api_key="test-key", + api_key_alias="test-alias", + requested_model="gpt-3.5-turbo", + model_group="gpt-3.5-turbo", + team=None, + team_alias=None, + user=None, + user_email=None, + status_code="200", + model="gpt-3.5-turbo", + litellm_model_name="gpt-3.5-turbo", + tags=[], + model_id="gpt-3.5-turbo", + api_base="https://api.openai.com", + api_provider="openai", + exception_status=None, + exception_class=None, + custom_metadata_labels={}, + route=None, + ) + + # Act + prometheus_logger._set_latency_metrics( + kwargs=kwargs, + model="gpt-3.5-turbo", + user_api_key="test-key", + user_api_key_alias="test-alias", + user_api_team=None, + user_api_team_alias=None, + enum_values=enum_values, + ) + + # Assert - queue time metric should not be called (queue_time_seconds is None) + # We check that observe was not called with queue_time_seconds + queue_time_called = False + for call in mock_labeled_metric.observe.call_args_list: + if len(call[0]) > 0 and call[0][0] == 0.5: # Our test queue time value + queue_time_called = True + break + assert ( + not queue_time_called + ), "Queue time metric should not be recorded when queue_time_seconds is missing" + + def test_queue_time_metric_not_recorded_when_negative(self): + """Test that queue time metric is not recorded when queue_time_seconds is negative""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock the metric + mock_metric = MagicMock() + mock_labeled_metric = MagicMock() + mock_metric.labels.return_value = mock_labeled_metric + prometheus_logger.litellm_request_queue_time_metric = mock_metric + + # Create mock kwargs with negative queue_time_seconds + kwargs = { + "litellm_params": { + "metadata": {"queue_time_seconds": -0.1} # Negative value + }, + "model": "gpt-3.5-turbo", + "start_time": datetime.now(), + "end_time": datetime.now(), + } + + enum_values = UserAPIKeyLabelValues( + end_user=None, + hashed_api_key="test-key", + api_key_alias="test-alias", + requested_model="gpt-3.5-turbo", + model_group="gpt-3.5-turbo", + team=None, + team_alias=None, + user=None, + user_email=None, + status_code="200", + model="gpt-3.5-turbo", + litellm_model_name="gpt-3.5-turbo", + tags=[], + model_id="gpt-3.5-turbo", + api_base="https://api.openai.com", + api_provider="openai", + exception_status=None, + exception_class=None, + custom_metadata_labels={}, + route=None, + ) + + # Act + prometheus_logger._set_latency_metrics( + kwargs=kwargs, + model="gpt-3.5-turbo", + user_api_key="test-key", + user_api_key_alias="test-alias", + user_api_team=None, + user_api_team_alias=None, + enum_values=enum_values, + ) + + # Assert - queue time metric should not be called for negative values + # We check that observe was not called with the negative value + negative_value_called = False + for call in mock_labeled_metric.observe.call_args_list: + if len(call[0]) > 0 and call[0][0] == -0.1: + negative_value_called = True + break + assert ( + not negative_value_called + ), "Queue time metric should not be recorded for negative values" + + +class TestPrometheusGuardrailMetrics: + """Test guardrail metrics recording""" + + def test_record_guardrail_metrics_success(self): + """Test recording guardrail metrics for successful execution""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock metrics + mock_latency_metric = MagicMock() + mock_requests_metric = MagicMock() + mock_errors_metric = MagicMock() + + prometheus_logger.litellm_guardrail_latency_metric = mock_latency_metric + prometheus_logger.litellm_guardrail_requests_total = mock_requests_metric + prometheus_logger.litellm_guardrail_errors_total = mock_errors_metric + + guardrail_name = "test_guardrail" + latency_seconds = 0.15 + status = "success" + error_type = None + hook_type = "pre_call" + + # Act + prometheus_logger._record_guardrail_metrics( + guardrail_name=guardrail_name, + latency_seconds=latency_seconds, + status=status, + error_type=error_type, + hook_type=hook_type, + ) + + # Assert - latency metric should be recorded + mock_latency_metric.labels.assert_called_once_with( + guardrail_name=guardrail_name, + status=status, + error_type="none", + hook_type=hook_type, + ) + mock_latency_metric.labels.return_value.observe.assert_called_once_with( + latency_seconds + ) + + # Assert - requests metric should be incremented + mock_requests_metric.labels.assert_called_once_with( + guardrail_name=guardrail_name, + status=status, + hook_type=hook_type, + ) + mock_requests_metric.labels.return_value.inc.assert_called_once() + + # Assert - errors metric should NOT be called for success + mock_errors_metric.labels.assert_not_called() + + def test_record_guardrail_metrics_error(self): + """Test recording guardrail metrics for failed execution""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock metrics + mock_latency_metric = MagicMock() + mock_requests_metric = MagicMock() + mock_errors_metric = MagicMock() + + prometheus_logger.litellm_guardrail_latency_metric = mock_latency_metric + prometheus_logger.litellm_guardrail_requests_total = mock_requests_metric + prometheus_logger.litellm_guardrail_errors_total = mock_errors_metric + + guardrail_name = "test_guardrail" + latency_seconds = 0.2 + status = "error" + error_type = "ValueError" + hook_type = "pre_call" + + # Act + prometheus_logger._record_guardrail_metrics( + guardrail_name=guardrail_name, + latency_seconds=latency_seconds, + status=status, + error_type=error_type, + hook_type=hook_type, + ) + + # Assert - latency metric should be recorded + mock_latency_metric.labels.assert_called_once_with( + guardrail_name=guardrail_name, + status=status, + error_type=error_type, + hook_type=hook_type, + ) + mock_latency_metric.labels.return_value.observe.assert_called_once_with( + latency_seconds + ) + + # Assert - requests metric should be incremented + mock_requests_metric.labels.assert_called_once_with( + guardrail_name=guardrail_name, + status=status, + hook_type=hook_type, + ) + mock_requests_metric.labels.return_value.inc.assert_called_once() + + # Assert - errors metric should be incremented + mock_errors_metric.labels.assert_called_once_with( + guardrail_name=guardrail_name, + error_type=error_type, + hook_type=hook_type, + ) + mock_errors_metric.labels.return_value.inc.assert_called_once() + + def test_record_guardrail_metrics_during_call_hook(self): + """Test recording guardrail metrics for during_call hook""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock metrics + mock_latency_metric = MagicMock() + mock_requests_metric = MagicMock() + + prometheus_logger.litellm_guardrail_latency_metric = mock_latency_metric + prometheus_logger.litellm_guardrail_requests_total = mock_requests_metric + + guardrail_name = "moderation_guardrail" + latency_seconds = 0.1 + status = "success" + hook_type = "during_call" + + # Act + prometheus_logger._record_guardrail_metrics( + guardrail_name=guardrail_name, + latency_seconds=latency_seconds, + status=status, + error_type=None, + hook_type=hook_type, + ) + + # Assert - hook_type should be "during_call" + mock_latency_metric.labels.assert_called_once() + call_kwargs = mock_latency_metric.labels.call_args[1] + assert call_kwargs["hook_type"] == "during_call" + + def test_record_guardrail_metrics_handles_exception(self): + """Test that _record_guardrail_metrics handles exceptions gracefully""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock metric to raise exception + mock_metric = MagicMock() + mock_metric.labels.side_effect = Exception("Test error") + prometheus_logger.litellm_guardrail_latency_metric = mock_metric + prometheus_logger.litellm_guardrail_requests_total = MagicMock() + + # Act & Assert - should not raise exception + try: + prometheus_logger._record_guardrail_metrics( + guardrail_name="test", + latency_seconds=0.1, + status="success", + error_type=None, + hook_type="pre_call", + ) + except Exception: + pytest.fail("_record_guardrail_metrics should handle exceptions gracefully") + + def test_record_guardrail_metrics_with_guardrail_name_attribute(self): + """Test that guardrail name is extracted from guardrail_name attribute if available""" + # Arrange + prometheus_logger = PrometheusLogger() + + # Mock metrics + mock_latency_metric = MagicMock() + mock_requests_metric = MagicMock() + + prometheus_logger.litellm_guardrail_latency_metric = mock_latency_metric + prometheus_logger.litellm_guardrail_requests_total = mock_requests_metric + + guardrail_name = "custom_guardrail_name" + latency_seconds = 0.1 + status = "success" + hook_type = "pre_call" + + # Act + prometheus_logger._record_guardrail_metrics( + guardrail_name=guardrail_name, + latency_seconds=latency_seconds, + status=status, + error_type=None, + hook_type=hook_type, + ) + + # Assert - guardrail_name should be used + mock_latency_metric.labels.assert_called_once() + call_kwargs = mock_latency_metric.labels.call_args[1] + assert call_kwargs["guardrail_name"] == guardrail_name diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 4914ec0bfb7..42a2b5d0971 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -620,26 +620,26 @@ def test_bedrock_tools_unpack_defs(): def test_bedrock_image_processor_content_type_fallback_url_extension(): """ - Test that _post_call_image_processing falls back to URL extension + Test that _post_call_image_processing falls back to URL extension when content-type is binary/octet-stream or application/octet-stream """ import base64 - + # Create mock response with binary/octet-stream content-type mock_response = MagicMock() mock_response.headers.get.return_value = "binary/octet-stream" - + # Create a simple PNG header (magic bytes) png_header = b"\x89\x50\x4e\x47\x0d\x0a\x1a\x0a" png_content = png_header + b"\x00" * 100 # Add some padding mock_response.content = png_content - + # Test with .png URL image_url = "https://example.com/test-image.png" base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, image_url ) - + assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -650,22 +650,22 @@ def test_bedrock_image_processor_content_type_fallback_binary_detection(): when content-type is missing and URL extension is not recognized """ import base64 - + # Create mock response with no content-type mock_response = MagicMock() mock_response.headers.get.return_value = None - + # Create a JPEG header (magic bytes) jpeg_header = b"\xff\xd8\xff" jpeg_content = jpeg_header + b"\x00" * 100 # Add some padding mock_response.content = jpeg_content - + # Test with URL without extension image_url = "https://example.com/test-image-without-extension" base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, image_url ) - + assert content_type == "image/jpeg" assert base64_bytes == base64.b64encode(jpeg_content).decode("utf-8") @@ -675,22 +675,22 @@ def test_bedrock_image_processor_content_type_fallback_application_octet_stream( Test that _post_call_image_processing handles application/octet-stream correctly """ import base64 - + # Create mock response with application/octet-stream content-type mock_response = MagicMock() mock_response.headers.get.return_value = "application/octet-stream" - + # Create a GIF header (magic bytes) gif_header = b"GIF8" + b"\x00" + b"a" gif_content = gif_header + b"\x00" * 100 # Add some padding mock_response.content = gif_content - + # Test with .gif URL image_url = "https://s3.amazonaws.com/bucket/image.gif" base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, image_url ) - + assert content_type == "image/gif" assert base64_bytes == base64.b64encode(gif_content).decode("utf-8") @@ -700,22 +700,22 @@ def test_bedrock_image_processor_content_type_with_query_params(): Test that _post_call_image_processing correctly extracts extension from URL with query parameters """ import base64 - + # Create mock response with binary/octet-stream content-type mock_response = MagicMock() mock_response.headers.get.return_value = "binary/octet-stream" - + # Create a WebP header (magic bytes) webp_header = b"RIFF" + b"\x00\x00\x00\x00" + b"WEBP" webp_content = webp_header + b"\x00" * 100 # Add some padding mock_response.content = webp_content - + # Test with URL containing query parameters (common in S3 signed URLs) image_url = "https://s3.amazonaws.com/bucket/image.webp?AWSAccessKeyId=123&Expires=456&Signature=789" base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, image_url ) - + assert content_type == "image/webp" assert base64_bytes == base64.b64encode(webp_content).decode("utf-8") @@ -725,21 +725,21 @@ def test_bedrock_image_processor_content_type_normal_header(): Test that _post_call_image_processing works normally when content-type is correctly set """ import base64 - + # Create mock response with correct content-type mock_response = MagicMock() mock_response.headers.get.return_value = "image/png" - + # Create a PNG header png_header = b"\x89\x50\x4e\x47\x0d\x0a\x1a\x0a" png_content = png_header + b"\x00" * 100 mock_response.content = png_content - + image_url = "https://example.com/test-image.png" base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, image_url ) - + assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -751,16 +751,16 @@ def test_bedrock_image_processor_content_type_fallback_failure(): # Create mock response with binary/octet-stream content-type mock_response = MagicMock() mock_response.headers.get.return_value = "binary/octet-stream" - + # Create content with unrecognizable image format mock_response.content = b"\x00" * 100 - + # Test with URL without recognizable extension image_url = "https://example.com/unknown-file" - + with pytest.raises(ValueError) as excinfo: BedrockImageProcessor._post_call_image_processing(mock_response, image_url) - + assert "Unable to determine content type" in str(excinfo.value) @@ -771,18 +771,18 @@ def test_bedrock_image_processor_content_type_jpeg_variants(): # Create mock response with binary/octet-stream mock_response = MagicMock() mock_response.headers.get.return_value = "binary/octet-stream" - + jpeg_header = b"\xff\xd8\xff" jpeg_content = jpeg_header + b"\x00" * 100 mock_response.content = jpeg_content - + # Test with .jpg extension image_url_jpg = "https://example.com/photo.jpg" _, content_type_jpg = BedrockImageProcessor._post_call_image_processing( mock_response, image_url_jpg ) assert content_type_jpg == "image/jpeg" - + # Test with .jpeg extension image_url_jpeg = "https://example.com/photo.jpeg" _, content_type_jpeg = BedrockImageProcessor._post_call_image_processing( @@ -797,22 +797,22 @@ def test_bedrock_image_processor_content_type_pdf_document(): when content-type is binary/octet-stream """ import base64 - + # Create mock response with binary/octet-stream content-type mock_response = MagicMock() mock_response.headers.get.return_value = "binary/octet-stream" - + # Create a PDF header (magic bytes: %PDF) pdf_header = b"%PDF-1.4" pdf_content = pdf_header + b"\x00" * 100 mock_response.content = pdf_content - + # Test with .pdf URL pdf_url = "https://s3.amazonaws.com/bucket/document.pdf" base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, pdf_url ) - + assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -822,12 +822,12 @@ def test_bedrock_image_processor_content_type_document_formats(): Test that _post_call_image_processing handles various document formats """ import base64 - + # Create mock response mock_response = MagicMock() mock_response.headers.get.return_value = "application/octet-stream" mock_response.content = b"\x00" * 100 - + # Test various document formats test_cases = [ ("https://example.com/doc.pdf", "application/pdf"), @@ -837,7 +837,7 @@ def test_bedrock_image_processor_content_type_document_formats(): ("https://example.com/page.html", "text/html"), ("https://example.com/readme.txt", "text/plain"), ] - + for url, expected_mime in test_cases: _, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, url @@ -850,21 +850,21 @@ def test_bedrock_image_processor_content_type_s3_pdf_with_query(): Test that _post_call_image_processing handles S3 PDF with query parameters """ import base64 - + # Create mock response mock_response = MagicMock() mock_response.headers.get.return_value = "binary/octet-stream" - + pdf_content = b"%PDF-1.4" + b"\x00" * 100 mock_response.content = pdf_content - + # S3 signed URL with query parameters s3_url = "https://my-bucket.s3.us-east-1.amazonaws.com/documents/report.pdf?AWSAccessKeyId=AKIAIOSFODNN7EXAMPLE&Expires=1234567890&Signature=abcdef123456" - + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( mock_response, s3_url ) - + assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -1139,6 +1139,170 @@ def test_bedrock_create_bedrock_block_different_document_formats(): assert block["document"]["format"] == format_type +def test_convert_to_anthropic_tool_result_image_with_cache_control(): + """ + Test that cache_control is properly applied to image content in tool results. + This tests the functionality added in the uncommitted changes where + add_cache_control_to_content is called for image_url content types. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_anthropic_tool_result, + ) + + # Test with base64 image data URI + message = { + "role": "tool", + "tool_call_id": "call_test_123", + "content": [ + { + "type": "text", + "text": "Here is the image you requested:", + }, + { + "type": "image_url", + "image_url": "data:image/jpeg;base64,/9j/4AAQSkZJRgABAQAAAQABAAD/2wBDAAgGBgcGBQ", + "cache_control": {"type": "ephemeral"}, + }, + ], + } + + result = convert_to_anthropic_tool_result(message) + + # Verify the result structure + assert result["type"] == "tool_result" + assert result["tool_use_id"] == "call_test_123" + assert isinstance(result["content"], list) + assert len(result["content"]) == 2 + + # Verify text content + assert result["content"][0]["type"] == "text" + assert result["content"][0]["text"] == "Here is the image you requested:" + + # Verify image content with cache_control + assert result["content"][1]["type"] == "image" + assert result["content"][1]["source"]["type"] == "base64" + assert result["content"][1]["source"]["media_type"] == "image/jpeg" + assert "cache_control" in result["content"][1] + assert result["content"][1]["cache_control"]["type"] == "ephemeral" + + +def test_convert_to_anthropic_tool_result_image_without_cache_control(): + """ + Test that images without cache_control in tool results work correctly. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_anthropic_tool_result, + ) + + message = { + "role": "tool", + "tool_call_id": "call_test_456", + "content": [ + { + "type": "image_url", + "image_url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAUA", + }, + ], + } + + result = convert_to_anthropic_tool_result(message) + + # Verify the result structure + assert result["type"] == "tool_result" + assert result["tool_use_id"] == "call_test_456" + assert isinstance(result["content"], list) + assert len(result["content"]) == 1 + + # Verify image content without cache_control (cache_control will be None if not set) + assert result["content"][0]["type"] == "image" + assert result["content"][0]["source"]["type"] == "base64" + assert result["content"][0]["source"]["media_type"] == "image/png" + assert result["content"][0].get("cache_control") is None + + +def test_convert_to_anthropic_tool_result_mixed_content_with_cache_control(): + """ + Test tool results with mixed content types (text and image) where only some have cache_control. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_anthropic_tool_result, + ) + + message = { + "role": "tool", + "tool_call_id": "call_test_789", + "content": [ + { + "type": "text", + "text": "First image:", + "cache_control": {"type": "ephemeral"}, + }, + { + "type": "image_url", + "image_url": "data:image/jpeg;base64,/9j/4AAQSkZJRg", + "cache_control": {"type": "ephemeral"}, + }, + { + "type": "text", + "text": "Second image (no cache):", + }, + { + "type": "image_url", + "image_url": "data:image/png;base64,iVBORw0KGgo", + }, + ], + } + + result = convert_to_anthropic_tool_result(message) + + assert result["type"] == "tool_result" + assert isinstance(result["content"], list) + assert len(result["content"]) == 4 + + # First text with cache_control + assert result["content"][0]["type"] == "text" + assert result["content"][0]["cache_control"]["type"] == "ephemeral" + + # First image with cache_control + assert result["content"][1]["type"] == "image" + assert result["content"][1]["cache_control"]["type"] == "ephemeral" + + # Second text without cache_control (cache_control will be None if not set) + assert result["content"][2]["type"] == "text" + assert result["content"][2].get("cache_control") is None + + # Second image without cache_control (cache_control will be None if not set) + assert result["content"][3]["type"] == "image" + assert result["content"][3].get("cache_control") is None + + +def test_convert_to_anthropic_tool_result_image_url_as_http(): + """ + Test that HTTP/HTTPS URLs with cache_control are handled correctly. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_anthropic_tool_result, + ) + + message = { + "role": "tool", + "tool_call_id": "call_http_001", + "content": [ + { + "type": "image_url", + "image_url": "https://example.com/image.jpg", + "cache_control": {"type": "ephemeral"}, + }, + ], + } + + result = convert_to_anthropic_tool_result(message) + + # Verify image is passed as URL reference with cache_control + assert result["content"][0]["type"] == "image" + assert result["content"][0]["source"]["type"] == "url" + assert result["content"][0]["source"]["url"] == "https://example.com/image.jpg" + assert result["content"][0]["cache_control"]["type"] == "ephemeral" def test_anthropic_messages_pt_server_tool_use_passthrough(): """ Test that anthropic_messages_pt passes through server_tool_use and diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 9b150fd89f4..bae0e5bbb4f 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -393,6 +393,63 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): assert "User-Agent: litellm/1.0.0" in tags +def test_get_request_tags_does_not_mutate_original_tags(): + """ + Test that _get_request_tags does not mutate the original tags list in metadata. + + This is a regression test for a bug where calling _get_request_tags multiple times + would cause User-Agent tags to be duplicated because the function was mutating + the original tags list instead of creating a copy. + """ + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + # Create metadata with original tags + original_tags = ["custom-tag-1", "custom-tag-2"] + metadata = {"tags": original_tags} + litellm_params = {"metadata": metadata} + proxy_server_request = { + "headers": { + "user-agent": "AsyncOpenAI/Python 1.99.9", + } + } + + # Call _get_request_tags multiple times (simulating multiple callbacks) + tags1 = StandardLoggingPayloadSetup._get_request_tags( + litellm_params=litellm_params, + proxy_server_request=proxy_server_request, + ) + tags2 = StandardLoggingPayloadSetup._get_request_tags( + litellm_params=litellm_params, + proxy_server_request=proxy_server_request, + ) + tags3 = StandardLoggingPayloadSetup._get_request_tags( + litellm_params=litellm_params, + proxy_server_request=proxy_server_request, + ) + + # Verify the original tags list was NOT mutated + assert original_tags == ["custom-tag-1", "custom-tag-2"], ( + f"Original tags list was mutated: {original_tags}" + ) + assert metadata["tags"] == ["custom-tag-1", "custom-tag-2"], ( + f"metadata['tags'] was mutated: {metadata['tags']}" + ) + + # Verify each returned list has exactly 2 User-Agent tags (not duplicated) + user_agent_count_1 = len([t for t in tags1 if t.startswith("User-Agent:")]) + user_agent_count_2 = len([t for t in tags2 if t.startswith("User-Agent:")]) + user_agent_count_3 = len([t for t in tags3 if t.startswith("User-Agent:")]) + + assert user_agent_count_1 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_1}" + assert user_agent_count_2 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_2}" + assert user_agent_count_3 == 2, f"Expected 2 User-Agent tags, got {user_agent_count_3}" + + # Verify all returned lists are independent (different objects) + assert tags1 is not tags2 + assert tags2 is not tags3 + assert tags1 is not original_tags + + def test_get_extra_header_tags(): """Test the _get_extra_header_tags method with various scenarios.""" import litellm diff --git a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py index 2836398228a..9f0a1ae8ffe 100644 --- a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py +++ b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py @@ -75,3 +75,47 @@ def test_excluded_keys_exact_match(): assert masked["api_key"] == "sk-1234567890abcdef" # Should NOT be masked assert masked["access_token"] != "token-12345" # Should still be masked assert "*" in masked["access_token"] + + +def test_extra_headers_are_masked_recursively(): + """ + Ensure nested dictionaries (like extra_headers) are masked. + """ + masker = SensitiveDataMasker() + + data = { + "litellm_params": { + "model": "openai/gpt-4", + "extra_headers": { + "rits_api_key": "sk-secret-12345-very-sensitive", + "Authorization": "Bearer token123", + }, + } + } + + masked = masker.mask_dict(data) + extra_headers = masked["litellm_params"]["extra_headers"] + + assert extra_headers["rits_api_key"] != "sk-secret-12345-very-sensitive" + assert "*" in extra_headers["rits_api_key"] + assert extra_headers["Authorization"] != "Bearer token123" + assert "*" in extra_headers["Authorization"] + + +def test_lists_with_sensitive_keys_are_masked(): + """ + Lists belonging to sensitive keys should have their values masked. + """ + masker = SensitiveDataMasker() + data = { + "api_key": ["sk-123", "sk-456"], + "tags": ["prod", "test"], + } + + masked = masker.mask_dict(data) + # sensitive key list entries should be masked + assert masked["api_key"][0] != "sk-123" + assert "*" in masked["api_key"][0] + + # non-sensitive list should remain unchanged + assert masked["tags"] == ["prod", "test"] diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index 9b6d1c6e178..5e58601d589 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -1754,3 +1754,61 @@ def test_transform_request_respects_user_max_tokens(): ) assert result["max_tokens"] == 1000 + + +def test_calculate_usage_completion_tokens_details_always_populated(): + """ + Test that completion_tokens_details is always populated in Usage object, + not just when there's reasoning_content. + + Fixes: https://github.com/BerriAI/litellm/issues/18772 + Bug: completion_tokens_details was None for regular Claude responses without reasoning + """ + config = AnthropicConfig() + + # Test without reasoning_content - completion_tokens_details should still be populated + usage_object = { + "input_tokens": 37, + "output_tokens": 248, + } + usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None) + + # completion_tokens_details should NOT be None + assert usage.completion_tokens_details is not None + assert usage.completion_tokens_details.reasoning_tokens is None + assert usage.completion_tokens_details.text_tokens == 248 + assert usage.completion_tokens == 248 + assert usage.prompt_tokens == 37 + assert usage.total_tokens == 285 + + +def test_calculate_usage_completion_tokens_details_with_reasoning(): + """ + Test that completion_tokens_details correctly splits text_tokens and reasoning_tokens + when reasoning_content is present. + + Fixes: https://github.com/BerriAI/litellm/issues/18772 + """ + config = AnthropicConfig() + + # Test with reasoning_content - should split tokens correctly + usage_object = { + "input_tokens": 100, + "output_tokens": 500, + } + # Simulating reasoning content that would count as ~50 tokens + reasoning_content = "Let me think about this step by step. " * 10 # Roughly 50 tokens + + usage = config.calculate_usage( + usage_object=usage_object, + reasoning_content=reasoning_content + ) + + # completion_tokens_details should be populated with both reasoning and text tokens + assert usage.completion_tokens_details is not None + assert usage.completion_tokens_details.reasoning_tokens is not None + assert usage.completion_tokens_details.reasoning_tokens > 0 + # text_tokens should be total minus reasoning + expected_text_tokens = 500 - usage.completion_tokens_details.reasoning_tokens + assert usage.completion_tokens_details.text_tokens == expected_text_tokens + assert usage.completion_tokens == 500 diff --git a/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py b/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py index 91d664c3216..199a16d8590 100644 --- a/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py +++ b/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py @@ -204,3 +204,62 @@ def test_azure_gpt5_reasoning_effort_none_dropped(config: AzureOpenAIGPT5Config) ) assert "reasoning_effort" not in params or params.get("reasoning_effort") != "none" + +# Logprobs support tests for Azure GPT-5.2 +def test_azure_gpt5_2_supports_logprobs(config: AzureOpenAIGPT5Config): + """Test that Azure GPT-5.2 models support logprobs parameters. + + Only Azure OpenAI GPT-5.2 supports logprobs, unlike OpenAI's GPT-5 or Azure's gpt-5/gpt-5.1. + Tested with gpt-5.2 on api-version 2025-01-01-preview. + """ + supported_params = config.get_supported_openai_params(model="gpt-5.2") + assert "logprobs" in supported_params + assert "top_logprobs" in supported_params + + +def test_azure_gpt5_2_with_prefix_supports_logprobs(config: AzureOpenAIGPT5Config): + """Test that Azure GPT-5.2 with azure/ prefix supports logprobs parameters.""" + supported_params = config.get_supported_openai_params(model="azure/gpt-5.2") + assert "logprobs" in supported_params + assert "top_logprobs" in supported_params + + +def test_azure_gpt5_2_series_supports_logprobs(config: AzureOpenAIGPT5Config): + """Test that Azure GPT-5.2 with gpt5_series prefix supports logprobs.""" + supported_params = config.get_supported_openai_params(model="gpt5_series/gpt-5.2") + assert "logprobs" in supported_params + assert "top_logprobs" in supported_params + + +def test_azure_gpt5_2_logprobs_params_passed_through(config: AzureOpenAIGPT5Config): + """Test that logprobs parameters are correctly passed through to the API for gpt-5.2.""" + params = config.map_openai_params( + non_default_params={"logprobs": True, "top_logprobs": 5}, + optional_params={}, + model="azure/gpt-5.2", + drop_params=False, + api_version="2025-01-01-preview", + ) + assert params["logprobs"] is True + assert params["top_logprobs"] == 5 + + +def test_azure_gpt5_base_does_not_support_logprobs(config: AzureOpenAIGPT5Config): + """Test that Azure GPT-5 (non-5.2) does not support logprobs parameters. + + Only gpt-5.2 has been verified to support logprobs on Azure. + """ + supported_params = config.get_supported_openai_params(model="gpt-5") + assert "logprobs" not in supported_params + assert "top_logprobs" not in supported_params + + +def test_azure_gpt5_1_does_not_support_logprobs(config: AzureOpenAIGPT5Config): + """Test that Azure GPT-5.1 does not support logprobs parameters. + + Only gpt-5.2 has been verified to support logprobs on Azure. + """ + supported_params = config.get_supported_openai_params(model="gpt-5.1") + assert "logprobs" not in supported_params + assert "top_logprobs" not in supported_params + diff --git a/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py b/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py index 998510efcd9..987eb5bf998 100644 --- a/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py +++ b/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py @@ -164,3 +164,90 @@ def test_azure_image_generation_headers_without_api_key(): # Verify api-key is added when api_key is valid assert "api-key" in default_headers_with_key assert default_headers_with_key["api-key"] == "valid-key-123" + + +def test_azure_image_generation_drop_params_response_format(): + """ + Test that unsupported params like response_format are dropped when drop_params=True. + + Azure gpt-image-1.5 doesn't support response_format parameter. When drop_params=True, + this parameter should be completely removed and not appear in the final request body, + including not being added to extra_body. + + This test verifies the fix where: + 1. Unsupported params are removed from non_default_params in _check_valid_arg + 2. Unsupported params are also removed from passed_params to prevent them from + being re-added via extra_body in add_provider_specific_params_to_optional_params + + Without the fix, response_format would be added to extra_body and cause Azure to + return a 400 Bad Request error due to strict schema validation. + """ + from litellm.llms.openai.image_generation.gpt_transformation import ( + GPTImageGenerationConfig, + ) + + # Test with gpt-image-1.5 which doesn't support response_format + config = GPTImageGenerationConfig() + supported_params = config.get_supported_openai_params(model="gpt-image-1.5") + + # Verify response_format is NOT in supported params for gpt-image-1.5 + assert "response_format" not in supported_params + assert "n" in supported_params + assert "size" in supported_params + + # Test get_optional_params_image_gen with drop_params=True + optional_params = get_optional_params_image_gen( + model="gpt-image-1.5", + n=1, + size="1024x1024", + response_format="b64_json", # This should be dropped + custom_llm_provider="azure", + provider_config=config, + drop_params=True, + ) + + # Verify response_format is NOT in optional_params + assert "response_format" not in optional_params, ( + "response_format should be dropped from optional_params" + ) + + # Verify response_format is NOT in extra_body either + if "extra_body" in optional_params: + assert "response_format" not in optional_params["extra_body"], ( + "response_format should not be in extra_body" + ) + + # Verify supported params ARE in optional_params + assert "n" in optional_params + assert optional_params["n"] == 1 + assert "size" in optional_params + assert optional_params["size"] == "1024x1024" + + +def test_azure_image_generation_drop_params_false_raises_error(): + """ + Test that unsupported params raise an error when drop_params=False. + + This verifies that the error handling still works correctly when drop_params + is not enabled. + """ + from litellm.exceptions import UnsupportedParamsError + from litellm.llms.openai.image_generation.gpt_transformation import ( + GPTImageGenerationConfig, + ) + + config = GPTImageGenerationConfig() + + # Test that passing unsupported param with drop_params=False raises error + with pytest.raises(UnsupportedParamsError) as exc_info: + optional_params = get_optional_params_image_gen( + model="gpt-image-1.5", + n=1, + response_format="b64_json", # Unsupported param + custom_llm_provider="azure", + provider_config=config, + drop_params=False, + ) + + # Verify the error message mentions the unsupported parameter + assert "response_format" in str(exc_info.value) diff --git a/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py b/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py index 7cb1ee2b54a..07253a8e09e 100644 --- a/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py +++ b/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py @@ -175,3 +175,132 @@ def test_format_url_handles_trailing_slash_normalization(): assert str(result_with_slash) == "http://proxy.com/bedrockproxy/model/test/invoke" +def test_bedrock_passthrough_with_application_inference_profile(): + """ + Test get_complete_url with Application Inference Profile ARN as model_id. + + This test verifies the fix for GitHub issue #18761 where Bedrock passthrough + was not working with Application Inference Profiles. The model_id (ARN) should + replace the translated model name in the endpoint URL. + """ + config = BedrockPassthroughConfig() + + model = "anthropic.claude-sonnet-4-20250514-v1:0" + model_id = "arn:aws:bedrock:eu-west-1:123456789:application-inference-profile/abcdefgh1234" + endpoint = f"model/{model}/invoke" + + with patch.object(config, '_get_aws_region_name', return_value="eu-west-1"), \ + patch.object(config, 'get_runtime_endpoint', return_value=( + "https://bedrock-runtime.eu-west-1.amazonaws.com", + "https://bedrock-runtime.eu-west-1.amazonaws.com" + )): + + url, api_base = config.get_complete_url( + api_base=None, + api_key=None, + model=model, + endpoint=endpoint, + request_query_params=None, + litellm_params={"model_id": model_id, "aws_region_name": "eu-west-1"} + ) + + # Verify that the URL contains the model_id (ARN) instead of the model name + url_str = str(url) + assert model_id in url_str, f"Expected model_id ARN in URL, but got: {url_str}" + assert model not in url_str, f"Model name should be replaced by model_id, but got: {url_str}" + assert "/invoke" in url_str, "Expected /invoke action in URL" + + # Verify the complete URL structure + expected_url = f"https://bedrock-runtime.eu-west-1.amazonaws.com/model/{model_id}/invoke" + assert url_str == expected_url, f"Expected {expected_url}, but got: {url_str}" + + +def test_bedrock_passthrough_with_inference_profile_converse_endpoint(): + """Test Application Inference Profile with converse endpoint""" + config = BedrockPassthroughConfig() + + model = "anthropic.claude-sonnet-4-20250514-v1:0" + model_id = "arn:aws:bedrock:us-east-1:123456789:application-inference-profile/xyz123" + endpoint = f"model/{model}/converse" + + with patch.object(config, '_get_aws_region_name', return_value="us-east-1"), \ + patch.object(config, 'get_runtime_endpoint', return_value=( + "https://bedrock-runtime.us-east-1.amazonaws.com", + "https://bedrock-runtime.us-east-1.amazonaws.com" + )): + + url, api_base = config.get_complete_url( + api_base=None, + api_key=None, + model=model, + endpoint=endpoint, + request_query_params=None, + litellm_params={"model_id": model_id} + ) + + url_str = str(url) + assert model_id in url_str + assert "/converse" in url_str + assert model not in url_str + + +def test_bedrock_passthrough_without_model_id_backward_compatibility(): + """ + Test that passthrough still works without model_id (backward compatibility). + + When model_id is not provided, the system should use the model name as before. + """ + config = BedrockPassthroughConfig() + + model = "anthropic.claude-3-sonnet" + endpoint = f"model/{model}/invoke" + + with patch.object(config, '_get_aws_region_name', return_value="us-east-1"), \ + patch.object(config, 'get_runtime_endpoint', return_value=( + "https://bedrock-runtime.us-east-1.amazonaws.com", + "https://bedrock-runtime.us-east-1.amazonaws.com" + )): + + url, api_base = config.get_complete_url( + api_base=None, + api_key=None, + model=model, + endpoint=endpoint, + request_query_params=None, + litellm_params={} # No model_id provided + ) + + # Verify that the URL contains the model name (not replaced) + url_str = str(url) + assert model in url_str, f"Expected model name in URL when model_id not provided, but got: {url_str}" + expected_url = f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{model}/invoke" + assert url_str == expected_url + + +def test_bedrock_passthrough_region_extraction_from_inference_profile_arn(): + """Test that AWS region is correctly extracted from Application Inference Profile ARN""" + config = BedrockPassthroughConfig() + + model = "anthropic.claude-sonnet-4-20250514-v1:0" + # ARN contains us-west-2 region + model_id = "arn:aws:bedrock:us-west-2:123456789:application-inference-profile/test123" + endpoint = f"model/{model}/invoke" + + # Don't provide aws_region_name in litellm_params to test ARN extraction + with patch.object(config, 'get_runtime_endpoint', return_value=( + "https://bedrock-runtime.us-west-2.amazonaws.com", + "https://bedrock-runtime.us-west-2.amazonaws.com" + )): + + url, api_base = config.get_complete_url( + api_base=None, + api_key=None, + model=model, + endpoint=endpoint, + request_query_params=None, + litellm_params={"model_id": model_id} # Region should be extracted from ARN + ) + + # Verify that the region from ARN is used in the base URL + assert "us-west-2" in api_base, f"Expected region 'us-west-2' from ARN in base URL, but got: {api_base}" + diff --git a/tests/test_litellm/llms/bedrock/test_bedrock_ssl_verify.py b/tests/test_litellm/llms/bedrock/test_bedrock_ssl_verify.py new file mode 100644 index 00000000000..9142de295ea --- /dev/null +++ b/tests/test_litellm/llms/bedrock/test_bedrock_ssl_verify.py @@ -0,0 +1,349 @@ +""" +Test SSL verification for AWS Bedrock boto3 clients. + +This test ensures that custom CA certificates are properly passed to all boto3 clients +(STS and Bedrock services) to support internal certificate authorities. + +Issue: https://github.com/BerriAI/litellm/issues/XXXX +User reported that SSL_CERT_FILE environment variable and ssl_verify config were not +being applied to boto3 clients, causing "certificate verify failed" errors. +""" + +import os +import sys +import tempfile +from unittest.mock import MagicMock, Mock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock.common_utils import init_bedrock_client + + +class TestBedrockSSLVerify: + """Test suite for SSL verification in Bedrock boto3 clients.""" + + def test_base_aws_llm_get_ssl_verify_default(self): + """Test that _get_ssl_verify returns default value when no custom config is set.""" + base_aws = BaseAWSLLM() + + # Clear any environment variables + os.environ.pop("SSL_VERIFY", None) + os.environ.pop("SSL_CERT_FILE", None) + + # Reset litellm.ssl_verify to default + litellm.ssl_verify = True + + ssl_verify = base_aws._get_ssl_verify() + assert ssl_verify is True + + def test_base_aws_llm_get_ssl_verify_false(self): + """Test that _get_ssl_verify returns False when SSL verification is disabled.""" + base_aws = BaseAWSLLM() + + # Set SSL_VERIFY to False via environment + os.environ["SSL_VERIFY"] = "False" + + ssl_verify = base_aws._get_ssl_verify() + assert ssl_verify is False + + # Clean up + os.environ.pop("SSL_VERIFY", None) + + def test_base_aws_llm_get_ssl_verify_custom_ca_bundle(self): + """Test that _get_ssl_verify returns custom CA bundle path when SSL_CERT_FILE is set.""" + base_aws = BaseAWSLLM() + + # Create a temporary CA bundle file + with tempfile.NamedTemporaryFile(mode="w", suffix=".pem", delete=False) as f: + f.write("-----BEGIN CERTIFICATE-----\n") + f.write("FAKE CERTIFICATE FOR TESTING\n") + f.write("-----END CERTIFICATE-----\n") + ca_bundle_path = f.name + + try: + # Set SSL_CERT_FILE environment variable + os.environ["SSL_CERT_FILE"] = ca_bundle_path + os.environ.pop("SSL_VERIFY", None) + litellm.ssl_verify = True + + ssl_verify = base_aws._get_ssl_verify() + assert ssl_verify == ca_bundle_path + finally: + # Clean up + os.environ.pop("SSL_CERT_FILE", None) + os.unlink(ca_bundle_path) + + def test_base_aws_llm_get_ssl_verify_litellm_config(self): + """Test that _get_ssl_verify uses litellm.ssl_verify when set.""" + base_aws = BaseAWSLLM() + + # Clear environment variables + os.environ.pop("SSL_VERIFY", None) + os.environ.pop("SSL_CERT_FILE", None) + + # Create a temporary CA bundle file + with tempfile.NamedTemporaryFile(mode="w", suffix=".pem", delete=False) as f: + f.write("-----BEGIN CERTIFICATE-----\n") + f.write("FAKE CERTIFICATE FOR TESTING\n") + f.write("-----END CERTIFICATE-----\n") + ca_bundle_path = f.name + + try: + # Set litellm.ssl_verify to custom CA bundle + litellm.ssl_verify = ca_bundle_path + + ssl_verify = base_aws._get_ssl_verify() + # When ssl_verify is a path, it should be returned directly + assert ssl_verify == ca_bundle_path + finally: + # Clean up + litellm.ssl_verify = True + os.unlink(ca_bundle_path) + + @patch("boto3.client") + def test_init_bedrock_client_passes_ssl_verify_to_sts(self, mock_boto3_client): + """Test that init_bedrock_client passes ssl_verify to STS client.""" + # Create a temporary CA bundle file + with tempfile.NamedTemporaryFile(mode="w", suffix=".pem", delete=False) as f: + f.write("-----BEGIN CERTIFICATE-----\n") + f.write("FAKE CERTIFICATE FOR TESTING\n") + f.write("-----END CERTIFICATE-----\n") + ca_bundle_path = f.name + + try: + # Set SSL_CERT_FILE environment variable + os.environ["SSL_CERT_FILE"] = ca_bundle_path + litellm.ssl_verify = True + + # Mock the STS client and Bedrock client + mock_sts_client = MagicMock() + mock_sts_response = { + "Credentials": { + "AccessKeyId": "test_access_key", + "SecretAccessKey": "test_secret_key", + "SessionToken": "test_session_token", + } + } + mock_sts_client.assume_role.return_value = mock_sts_response + + mock_bedrock_client = MagicMock() + + # Configure mock to return different clients based on service name + def side_effect(service_name=None, **kwargs): + if service_name == "sts": + return mock_sts_client + elif service_name == "bedrock-runtime": + return mock_bedrock_client + return MagicMock() + + mock_boto3_client.side_effect = side_effect + + # Call init_bedrock_client with role assumption + client = init_bedrock_client( + aws_region_name="us-west-2", + aws_access_key_id="test_key", + aws_secret_access_key="test_secret", + aws_role_name="arn:aws:iam::123456789012:role/test-role", + aws_session_name="test-session", + ) + + # Verify that boto3.client was called with verify parameter for STS + sts_calls = [ + call for call in mock_boto3_client.call_args_list + if (len(call[0]) > 0 and call[0][0] == "sts") or + ("service_name" not in call[1]) # STS calls don't use service_name kwarg + ] + + assert len(sts_calls) > 0, "STS client should have been created" + + # Check that verify parameter was passed to STS client + sts_call = sts_calls[0] + assert "verify" in sts_call[1], "verify parameter should be passed to STS client" + assert sts_call[1]["verify"] == ca_bundle_path, f"verify should be set to CA bundle path, got {sts_call[1]['verify']}" + + # Verify that boto3.client was called with verify parameter for Bedrock + bedrock_calls = [ + call for call in mock_boto3_client.call_args_list + if "service_name" in call[1] and call[1]["service_name"] == "bedrock-runtime" + ] + + assert len(bedrock_calls) > 0, "Bedrock client should have been created" + + bedrock_call = bedrock_calls[0] + assert "verify" in bedrock_call[1], "verify parameter should be passed to Bedrock client" + assert bedrock_call[1]["verify"] == ca_bundle_path, f"verify should be set to CA bundle path, got {bedrock_call[1]['verify']}" + + finally: + # Clean up + os.environ.pop("SSL_CERT_FILE", None) + os.unlink(ca_bundle_path) + + @patch("boto3.client") + def test_base_aws_llm_auth_with_role_passes_ssl_verify(self, mock_boto3_client): + """Test that _auth_with_aws_role passes ssl_verify to STS client.""" + base_aws = BaseAWSLLM() + + # Create a temporary CA bundle file + with tempfile.NamedTemporaryFile(mode="w", suffix=".pem", delete=False) as f: + f.write("-----BEGIN CERTIFICATE-----\n") + f.write("FAKE CERTIFICATE FOR TESTING\n") + f.write("-----END CERTIFICATE-----\n") + ca_bundle_path = f.name + + try: + # Set SSL_CERT_FILE environment variable + os.environ["SSL_CERT_FILE"] = ca_bundle_path + litellm.ssl_verify = True + + # Mock the STS client + mock_sts_client = MagicMock() + mock_sts_response = { + "Credentials": { + "AccessKeyId": "test_access_key", + "SecretAccessKey": "test_secret_key", + "SessionToken": "test_session_token", + "Expiration": "2025-01-10T00:00:00Z", + } + } + + # Convert Expiration to datetime + from datetime import datetime, timezone + mock_sts_response["Credentials"]["Expiration"] = datetime.now(timezone.utc) + + mock_sts_client.assume_role.return_value = mock_sts_response + mock_boto3_client.return_value = mock_sts_client + + # Call _auth_with_aws_role + credentials, ttl = base_aws._auth_with_aws_role( + aws_access_key_id="test_key", + aws_secret_access_key="test_secret", + aws_session_token=None, + aws_role_name="arn:aws:iam::123456789012:role/test-role", + aws_session_name="test-session", + ) + + # Verify that boto3.client was called with verify parameter + assert mock_boto3_client.called, "boto3.client should have been called" + + call_kwargs = mock_boto3_client.call_args[1] + assert "verify" in call_kwargs, "verify parameter should be passed to STS client" + assert call_kwargs["verify"] == ca_bundle_path, f"verify should be set to CA bundle path, got {call_kwargs['verify']}" + + finally: + # Clean up + os.environ.pop("SSL_CERT_FILE", None) + os.unlink(ca_bundle_path) + + @patch("litellm.llms.bedrock.base_aws_llm.get_secret") + @patch("boto3.client") + def test_base_aws_llm_auth_with_web_identity_passes_ssl_verify(self, mock_boto3_client, mock_get_secret): + """Test that _auth_with_web_identity_token passes ssl_verify to STS client.""" + base_aws = BaseAWSLLM() + + # Create a temporary CA bundle file + with tempfile.NamedTemporaryFile(mode="w", suffix=".pem", delete=False) as f: + f.write("-----BEGIN CERTIFICATE-----\n") + f.write("FAKE CERTIFICATE FOR TESTING\n") + f.write("-----END CERTIFICATE-----\n") + ca_bundle_path = f.name + + try: + # Set SSL_CERT_FILE environment variable + os.environ["SSL_CERT_FILE"] = ca_bundle_path + litellm.ssl_verify = True + + # Mock get_secret to return the token + mock_get_secret.return_value = "mocked_oidc_token" + + # Mock the STS client + mock_sts_client = MagicMock() + mock_sts_response = { + "Credentials": { + "AccessKeyId": "test_access_key", + "SecretAccessKey": "test_secret_key", + "SessionToken": "test_session_token", + }, + "PackedPolicySize": 100, + } + + mock_sts_client.assume_role_with_web_identity.return_value = mock_sts_response + + # Mock boto3.Session + mock_session = MagicMock() + mock_credentials = MagicMock() + mock_session.get_credentials.return_value = mock_credentials + + mock_boto3_client.return_value = mock_sts_client + + with patch("boto3.Session", return_value=mock_session): + # Call _auth_with_web_identity_token + credentials, ttl = base_aws._auth_with_web_identity_token( + aws_web_identity_token="test_token", + aws_role_name="arn:aws:iam::123456789012:role/test-role", + aws_session_name="test-session", + aws_region_name="us-west-2", + aws_sts_endpoint=None, + ) + + # Verify that boto3.client was called with verify parameter + assert mock_boto3_client.called, "boto3.client should have been called" + + call_kwargs = mock_boto3_client.call_args[1] + assert "verify" in call_kwargs, "verify parameter should be passed to STS client" + assert call_kwargs["verify"] == ca_bundle_path, f"verify should be set to CA bundle path, got {call_kwargs['verify']}" + + finally: + # Clean up + os.environ.pop("SSL_CERT_FILE", None) + os.unlink(ca_bundle_path) + + def test_ssl_verify_priority_env_over_litellm_config(self): + """Test that SSL_VERIFY environment variable takes priority over litellm.ssl_verify.""" + base_aws = BaseAWSLLM() + + # Set litellm.ssl_verify to True + litellm.ssl_verify = True + + # Set SSL_VERIFY environment variable to False + os.environ["SSL_VERIFY"] = "False" + + try: + ssl_verify = base_aws._get_ssl_verify() + assert ssl_verify is False, "Environment variable should take priority" + finally: + # Clean up + os.environ.pop("SSL_VERIFY", None) + litellm.ssl_verify = True + + def test_ssl_cert_file_priority_over_default(self): + """Test that SSL_CERT_FILE takes priority when ssl_verify is True.""" + base_aws = BaseAWSLLM() + + # Create a temporary CA bundle file + with tempfile.NamedTemporaryFile(mode="w", suffix=".pem", delete=False) as f: + f.write("-----BEGIN CERTIFICATE-----\n") + f.write("FAKE CERTIFICATE FOR TESTING\n") + f.write("-----END CERTIFICATE-----\n") + ca_bundle_path = f.name + + try: + # Set SSL_CERT_FILE environment variable + os.environ["SSL_CERT_FILE"] = ca_bundle_path + os.environ.pop("SSL_VERIFY", None) + litellm.ssl_verify = True + + ssl_verify = base_aws._get_ssl_verify() + assert ssl_verify == ca_bundle_path, "SSL_CERT_FILE should be used when ssl_verify is True" + finally: + # Clean up + os.environ.pop("SSL_CERT_FILE", None) + os.unlink(ca_bundle_path) + + +if __name__ == "__main__": + # Run tests + pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_litellm/llms/deepinfra/test_deepinfra_chat_transformation.py b/tests/test_litellm/llms/deepinfra/test_deepinfra_chat_transformation.py index fc8cf6dc60f..49d55f920b5 100644 --- a/tests/test_litellm/llms/deepinfra/test_deepinfra_chat_transformation.py +++ b/tests/test_litellm/llms/deepinfra/test_deepinfra_chat_transformation.py @@ -24,3 +24,194 @@ def test_deepseek_supported_openai_params(): supported_openai_params = DeepInfraConfig().get_supported_openai_params(model="deepinfra/deepseek-ai/DeepSeek-V3.1") print(supported_openai_params) assert "reasoning_effort" in supported_openai_params + + +def test_deepinfra_tool_message_content_transformation(): + """ + Test that DeepInfra transforms tool message content from array to string. + + This fixes the issue where LibreChat sends tool messages with content as an array: + {"role": "tool", "content": [{"type": "text", "text": "20"}]} + + DeepInfra requires content to be a string, so we transform it to: + {"role": "tool", "content": "20"} + + Related to issue #13982 + """ + from litellm.llms.deepinfra.chat.transformation import DeepInfraConfig + + config = DeepInfraConfig() + + # Test case 1: Simple single text item in array (common case from LibreChat) + messages_with_array_content = [ + { + "role": "user", + "content": "Calculate 10 + 10" + }, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": { + "name": "calculator", + "arguments": '{"input": "10 + 10"}' + } + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_123", + "name": "calculator", + "content": [{"type": "text", "text": "20"}] # Array format from LibreChat + } + ] + + transformed_messages = config._transform_messages( + messages=messages_with_array_content, + model="deepinfra/Qwen/Qwen3-235B-A22B" + ) + + # Verify the tool message content was converted to string + tool_message = transformed_messages[2] + assert tool_message["role"] == "tool" + assert isinstance(tool_message["content"], str) + assert tool_message["content"] == "20" + print(f"✓ Test case 1 passed: {tool_message['content']}") + + # Test case 2: Complex content array (multiple items) + messages_with_complex_content = [ + { + "role": "user", + "content": "Test" + }, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_456", + "type": "function", + "function": {"name": "test", "arguments": "{}"} + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_456", + "content": [ + {"type": "text", "text": "Result 1"}, + {"type": "text", "text": "Result 2"} + ] + } + ] + + transformed_messages_complex = config._transform_messages( + messages=messages_with_complex_content, + model="deepinfra/Qwen/Qwen3-235B-A22B" + ) + + tool_message_complex = transformed_messages_complex[2] + assert tool_message_complex["role"] == "tool" + assert isinstance(tool_message_complex["content"], str) + # For complex content, it should be JSON stringified + parsed_content = json.loads(tool_message_complex["content"]) + assert len(parsed_content) == 2 + assert parsed_content[0]["text"] == "Result 1" + print(f"✓ Test case 2 passed: {tool_message_complex['content']}") + + # Test case 3: Tool message with string content (should remain unchanged) + messages_with_string_content = [ + { + "role": "user", + "content": "Test" + }, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_789", + "type": "function", + "function": {"name": "test", "arguments": "{}"} + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_789", + "content": "Simple string result" # Already a string + } + ] + + transformed_messages_string = config._transform_messages( + messages=messages_with_string_content, + model="deepinfra/Qwen/Qwen3-235B-A22B" + ) + + tool_message_string = transformed_messages_string[2] + assert tool_message_string["role"] == "tool" + assert isinstance(tool_message_string["content"], str) + assert tool_message_string["content"] == "Simple string result" + print(f"✓ Test case 3 passed: {tool_message_string['content']}") + + print("\n✅ All DeepInfra tool message transformation tests passed!") + + +@pytest.mark.asyncio +async def test_deepinfra_tool_message_content_transformation_async(): + """ + Test that DeepInfra transforms tool message content from array to string in async mode. + + This ensures the async path works correctly when is_async=True. + + Related to issue #13982 + """ + from litellm.llms.deepinfra.chat.transformation import DeepInfraConfig + + config = DeepInfraConfig() + + # Test async transformation with tool message containing array content + messages_with_array_content = [ + { + "role": "user", + "content": "Calculate 10 + 10" + }, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": { + "name": "calculator", + "arguments": '{"input": "10 + 10"}' + } + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_123", + "name": "calculator", + "content": [{"type": "text", "text": "20"}] # Array format from LibreChat + } + ] + + # Call with is_async=True + transformed_messages = await config._transform_messages( + messages=messages_with_array_content, + model="deepinfra/Qwen/Qwen3-235B-A22B", + is_async=True + ) + + # Verify the tool message content was converted to string + tool_message = transformed_messages[2] + assert tool_message["role"] == "tool" + assert isinstance(tool_message["content"], str) + assert tool_message["content"] == "20" + print(f"✓ Async test passed: {tool_message['content']}") + + print("\n✅ DeepInfra async tool message transformation test passed!") diff --git a/tests/test_litellm/llms/manus/__init__.py b/tests/test_litellm/llms/manus/__init__.py new file mode 100644 index 00000000000..d4037b65199 --- /dev/null +++ b/tests/test_litellm/llms/manus/__init__.py @@ -0,0 +1,2 @@ +# Manus provider tests + diff --git a/tests/test_litellm/llms/manus/responses/__init__.py b/tests/test_litellm/llms/manus/responses/__init__.py new file mode 100644 index 00000000000..a7131749c5c --- /dev/null +++ b/tests/test_litellm/llms/manus/responses/__init__.py @@ -0,0 +1,2 @@ +# Manus Responses API tests + diff --git a/tests/test_litellm/llms/manus/responses/test_manus_responses_transformation.py b/tests/test_litellm/llms/manus/responses/test_manus_responses_transformation.py new file mode 100644 index 00000000000..b47ed77156d --- /dev/null +++ b/tests/test_litellm/llms/manus/responses/test_manus_responses_transformation.py @@ -0,0 +1,60 @@ +""" +Tests for Manus Responses API transformation + +Tests the ManusResponsesAPIConfig class that handles Manus-specific +transformations for the Responses API. + +Source: litellm/llms/manus/responses/transformation.py +""" +import os +import sys + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.manus.responses.transformation import ManusResponsesAPIConfig +from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams +from litellm.types.router import GenericLiteLLMParams + + +def test_extract_agent_profile(): + """Test that agent profile is correctly extracted from model name""" + config = ManusResponsesAPIConfig() + + assert config._extract_agent_profile("manus/manus-1.6") == "manus-1.6" + assert config._extract_agent_profile("manus/manus-1.6-lite") == "manus-1.6-lite" + assert config._extract_agent_profile("manus/manus-1.6-max") == "manus-1.6-max" + + +def test_transform_responses_api_request_adds_manus_params(): + """Test that transform_responses_api_request adds task_mode and agent_profile""" + config = ManusResponsesAPIConfig() + + input_param = [ + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": "What's the color of the sky?", + } + ], + } + ] + + optional_params = ResponsesAPIOptionalRequestParams() + litellm_params = GenericLiteLLMParams() + headers = {} + + result = config.transform_responses_api_request( + model="manus/manus-1.6", + input=input_param, + response_api_optional_request_params=dict(optional_params), + litellm_params=litellm_params, + headers=headers, + ) + + assert result["task_mode"] == "agent" + assert result["agent_profile"] == "manus-1.6" + assert "input" in result + assert "model" in result + diff --git a/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py b/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py new file mode 100644 index 00000000000..d025c716a4a --- /dev/null +++ b/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py @@ -0,0 +1,150 @@ +""" +Tests for Xiaomi MiMo provider configuration and integration. +Related to issue #18794 +""" + +import os +import sys +from unittest.mock import MagicMock, patch + +try: + import pytest +except ImportError: + pytest = None + +# Add workspace to path +workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) +sys.path.insert(0, workspace_path) + +import litellm + + +class TestXiaomiMiMoProviderConfig: + """Test Xiaomi MiMo provider configuration""" + + def test_xiaomi_mimo_in_provider_list(self): + """Test that xiaomi_mimo is in the provider list (fixes #18794)""" + from litellm import LlmProviders + + # Verify xiaomi_mimo is in the enum + assert hasattr(LlmProviders, 'XIAOMI_MIMO') + assert LlmProviders.XIAOMI_MIMO.value == 'xiaomi_mimo' + + # Verify it's in the provider list + assert 'xiaomi_mimo' in litellm.provider_list + + def test_xiaomi_mimo_json_config_exists(self): + """Test that xiaomi_mimo is configured in providers.json""" + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + # Verify xiaomi_mimo is loaded + assert JSONProviderRegistry.exists("xiaomi_mimo") + + # Get xiaomi_mimo config + xiaomi_mimo = JSONProviderRegistry.get("xiaomi_mimo") + assert xiaomi_mimo is not None + assert xiaomi_mimo.base_url == "https://api.xiaomimimo.com/v1" + assert xiaomi_mimo.api_key_env == "XIAOMI_MIMO_API_KEY" + assert xiaomi_mimo.param_mappings.get("max_completion_tokens") == "max_tokens" + + def test_xiaomi_mimo_provider_resolution(self): + """Test that provider resolution finds xiaomi_mimo""" + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, api_key, api_base = get_llm_provider( + model="xiaomi_mimo/mimo-v2-flash", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert model == "mimo-v2-flash" + assert provider == "xiaomi_mimo" + assert api_base == "https://api.xiaomimimo.com/v1" + + def test_xiaomi_mimo_router_config(self): + """Test that xiaomi_mimo can be used in Router configuration (fixes #18794)""" + from litellm import Router + + # This should not raise "Unsupported provider - xiaomi_mimo" + router = Router( + model_list=[ + { + "model_name": "mimo-v2-flash", + "litellm_params": { + "model": "xiaomi_mimo/mimo-v2-flash", + "api_key": "test-key", + }, + } + ] + ) + + # Verify the deployment was created successfully + assert len(router.model_list) == 1 + assert router.model_list[0]["model_name"] == "mimo-v2-flash" + + +class TestXiaomiMiMoIntegration: + """Integration tests for Xiaomi MiMo provider""" + + def test_xiaomi_mimo_completion_basic(self): + """Test basic completion call to Xiaomi MiMo""" + # Skip test if API key not set in environment + if not os.environ.get("XIAOMI_MIMO_API_KEY"): + if pytest: + pytest.skip("XIAOMI_MIMO_API_KEY not set") + return + + try: + response = litellm.completion( + model="xiaomi_mimo/mimo-v2-flash", + messages=[{"role": "user", "content": "Say 'test successful' and nothing else"}], + max_tokens=10, + ) + + # Verify response structure + assert response is not None + assert hasattr(response, "choices") + assert len(response.choices) > 0 + assert hasattr(response.choices[0], "message") + assert hasattr(response.choices[0].message, "content") + assert response.choices[0].message.content is not None + + # Check that we got a response + content = response.choices[0].message.content.lower() + assert len(content) > 0 + + print(f"✓ Xiaomi MiMo completion successful: {response.choices[0].message.content}") + + except Exception as e: + if pytest: + pytest.fail(f"Xiaomi MiMo completion failed: {str(e)}") + else: + raise + + +if __name__ == "__main__": + # Run basic tests + print("Testing Xiaomi MiMo Provider...") + + test_config = TestXiaomiMiMoProviderConfig() + + print("\n1. Testing provider in list...") + test_config.test_xiaomi_mimo_in_provider_list() + print(" ✓ xiaomi_mimo in provider list") + + print("\n2. Testing JSON config...") + test_config.test_xiaomi_mimo_json_config_exists() + print(" ✓ xiaomi_mimo JSON config loaded") + + print("\n3. Testing provider resolution...") + test_config.test_xiaomi_mimo_provider_resolution() + print(" ✓ Provider resolution works") + + print("\n4. Testing router configuration...") + test_config.test_xiaomi_mimo_router_config() + print(" ✓ Router configuration works (issue #18794 fixed)") + + print("\n" + "="*50) + print("✓ All configuration tests passed!") + print("="*50) diff --git a/tests/test_litellm/llms/openrouter/test_openrouter_embedding_transformation.py b/tests/test_litellm/llms/openrouter/test_openrouter_embedding_transformation.py new file mode 100644 index 00000000000..714adc346db --- /dev/null +++ b/tests/test_litellm/llms/openrouter/test_openrouter_embedding_transformation.py @@ -0,0 +1,132 @@ +""" +Unit tests for OpenRouter embedding transformation logic. +""" +from litellm.llms.openrouter.embedding.transformation import ( + OpenrouterEmbeddingConfig, +) + + +def test_openrouter_embedding_supported_params(): + """Test that supported OpenAI params are correctly defined.""" + config = OpenrouterEmbeddingConfig() + supported = config.get_supported_openai_params("test-model") + + assert "timeout" in supported + assert "dimensions" in supported + assert "encoding_format" in supported + assert "user" in supported + + +def test_openrouter_embedding_transform_request(): + """Test request transformation logic.""" + config = OpenrouterEmbeddingConfig() + + # Test with string input + result = config.transform_embedding_request( + model="openrouter/google/text-embedding-004", + input="Hello world", + optional_params={}, + headers={}, + ) + + assert result["model"] == "google/text-embedding-004" + assert result["input"] == ["Hello world"] + + # Test with list input + result = config.transform_embedding_request( + model="google/text-embedding-004", + input=["Hello", "World"], + optional_params={"dimensions": 512}, + headers={}, + ) + + assert result["model"] == "google/text-embedding-004" + assert result["input"] == ["Hello", "World"] + assert result["dimensions"] == 512 + + +def test_openrouter_embedding_validate_environment(): + """Test environment validation and header setup.""" + config = OpenrouterEmbeddingConfig() + + # Test with API key + headers = config.validate_environment( + headers={"Custom-Header": "value"}, + model="test-model", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test-api-key", + ) + + # Should include OpenRouter-specific headers + assert "HTTP-Referer" in headers + assert "X-Title" in headers + # Should include Content-Type header + assert "Content-Type" in headers + assert headers["Content-Type"] == "application/json" + # Should include Authorization header + assert "Authorization" in headers + assert headers["Authorization"] == "Bearer test-api-key" + # Should preserve custom headers + assert headers["Custom-Header"] == "value" + + # Test without API key + headers_no_key = config.validate_environment( + headers={}, + model="test-model", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + # Should still include OpenRouter headers but not Authorization + assert "HTTP-Referer" in headers_no_key + assert "X-Title" in headers_no_key + assert "Content-Type" in headers_no_key + assert "Authorization" not in headers_no_key + + +def test_openrouter_embedding_get_complete_url(): + """Test URL construction.""" + config = OpenrouterEmbeddingConfig() + + url = config.get_complete_url( + api_base="https://openrouter.ai/api/v1", + api_key="test-key", + model="test-model", + optional_params={}, + litellm_params={}, + ) + + assert url == "https://openrouter.ai/api/v1/embeddings" + + # Test with trailing slash + url = config.get_complete_url( + api_base="https://openrouter.ai/api/v1/", + api_key="test-key", + model="test-model", + optional_params={}, + litellm_params={}, + ) + + assert url == "https://openrouter.ai/api/v1/embeddings" + + +def test_openrouter_embedding_map_params(): + """Test parameter mapping.""" + config = OpenrouterEmbeddingConfig() + + result = config.map_openai_params( + non_default_params={"dimensions": 512, "timeout": 30, "unsupported": "value"}, + optional_params={}, + model="test-model", + drop_params=False, + ) + + # Supported params should be included + assert result["dimensions"] == 512 + assert result["timeout"] == 30 + # Unsupported params should not be included + assert "unsupported" not in result diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 3c3d68e8be8..70e6e9452e5 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -721,3 +721,189 @@ def test_convert_tool_response_text_only(): # Check inline_data does NOT exist (no image provided) assert "inline_data" not in result + + +def test_file_data_field_order(): + """ + Test that file_data fields are in the correct order (mime_type before file_uri). + + The Gemini API is sensitive to field order in the file_data object. + This test verifies that mime_type comes before file_uri in both: + 1. Dictionary key order + 2. JSON serialization + + Related issue: Gemini API returns 400 INVALID_ARGUMENT when fields are in wrong order. + """ + import json + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_image + + # Test with HTTPS URL and explicit format (audio file) + file_url = "https://generativelanguage.googleapis.com/v1beta/files/test123" + format = "audio/mpeg" + + result = _process_gemini_image(image_url=file_url, format=format) + + # Verify the result has file_data + assert "file_data" in result + file_data = result["file_data"] + + # Verify both fields are present + assert "mime_type" in file_data + assert "file_uri" in file_data + assert file_data["mime_type"] == "audio/mpeg" + assert file_data["file_uri"] == file_url + + # Verify field order by checking dictionary keys + # In Python 3.7+, dict maintains insertion order + file_data_keys = list(file_data.keys()) + assert file_data_keys.index("mime_type") < file_data_keys.index("file_uri"), \ + "mime_type must come before file_uri in the file_data dict" + + # Also verify by serializing to JSON string + json_str = json.dumps(file_data) + mime_type_pos = json_str.find('"mime_type"') + file_uri_pos = json_str.find('"file_uri"') + assert mime_type_pos < file_uri_pos, \ + "mime_type must appear before file_uri in JSON serialization" + + +def test_file_data_field_order_gcs_urls(): + """Test that GCS URLs also maintain correct field order.""" + import json + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_image + + # Test with GCS URL + gcs_url = "gs://bucket/audio.mp3" + + result = _process_gemini_image(image_url=gcs_url) + + # Verify the result has file_data + assert "file_data" in result + file_data = result["file_data"] + + # Verify both fields are present + assert "mime_type" in file_data + assert "file_uri" in file_data + + # Verify field order + file_data_keys = list(file_data.keys()) + assert file_data_keys.index("mime_type") < file_data_keys.index("file_uri"), \ + "mime_type must come before file_uri in the file_data dict" + + +def test_extract_file_data_with_path_object(): + """ + Test that filename is correctly extracted from Path objects for MIME type detection. + + When uploading files using Path objects (e.g., Path("speech.mp3")), the filename + must be extracted to enable proper MIME type detection. Without this, files get + uploaded with 'application/octet-stream' instead of the correct MIME type. + + Related issue: Files uploaded with wrong MIME type cause Gemini API to reject + requests where the specified format doesn't match the uploaded file's MIME type. + """ + from pathlib import Path + from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data + import tempfile + import os + + # Create a temporary MP3 file + with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as tmp: + tmp.write(b"fake mp3 content") + tmp_path = tmp.name + + try: + # Test with Path object + path_obj = Path(tmp_path) + extracted = extract_file_data(path_obj) + + # Verify filename was extracted + assert extracted["filename"] is not None + assert extracted["filename"].endswith(".mp3") + + # Verify MIME type was correctly detected + assert extracted["content_type"] == "audio/mpeg", \ + f"Expected 'audio/mpeg' but got '{extracted['content_type']}'" + + # Verify content was read + assert extracted["content"] == b"fake mp3 content" + + finally: + # Clean up temporary file + os.unlink(tmp_path) + + +def test_extract_file_data_with_string_path(): + """Test that filename is correctly extracted from string paths.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data + import tempfile + import os + + # Create a temporary WAV file + with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: + tmp.write(b"fake wav content") + tmp_path = tmp.name + + try: + # Test with string path + extracted = extract_file_data(tmp_path) + + # Verify filename was extracted + assert extracted["filename"] is not None + assert extracted["filename"].endswith(".wav") + + # Verify MIME type was correctly detected (can be audio/wav or audio/x-wav depending on system) + assert extracted["content_type"] in ["audio/wav", "audio/x-wav"], \ + f"Expected 'audio/wav' or 'audio/x-wav' but got '{extracted['content_type']}'" + + # Verify content was read + assert extracted["content"] == b"fake wav content" + + finally: + # Clean up temporary file + os.unlink(tmp_path) + + +def test_extract_file_data_with_tuple_format(): + """Test that tuple format (with explicit content_type) still works correctly.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data + + # Test with tuple format: (filename, content, content_type) + filename = "test_audio.mp3" + content = b"test audio content" + content_type = "audio/mpeg" + + extracted = extract_file_data((filename, content, content_type)) + + # Verify all fields are correct + assert extracted["filename"] == filename + assert extracted["content"] == content + assert extracted["content_type"] == content_type + + +def test_extract_file_data_fallback_to_octet_stream(): + """Test that unknown file types fall back to application/octet-stream.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data + import tempfile + import os + + # Create a temporary file with unknown extension + with tempfile.NamedTemporaryFile(suffix=".xyz123", delete=False) as tmp: + tmp.write(b"unknown content") + tmp_path = tmp.name + + try: + # Test with unknown file type + extracted = extract_file_data(tmp_path) + + # Verify filename was extracted + assert extracted["filename"] is not None + assert extracted["filename"].endswith(".xyz123") + + # Verify MIME type falls back to octet-stream + assert extracted["content_type"] == "application/octet-stream", \ + f"Expected 'application/octet-stream' for unknown type, got '{extracted['content_type']}'" + + finally: + # Clean up temporary file + os.unlink(tmp_path) diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py new file mode 100644 index 00000000000..0a1ac7e2a54 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py @@ -0,0 +1,38 @@ +import pytest +from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig +from litellm import ModelResponse + +def test_process_candidates_unbound_local_error_fix(): + # Setup + candidates = [ + { + "content": { + "role": "model" + # "parts" is missing intentionally to trigger the issue + }, + "finishReason": "STOP" + } + ] + model_response = ModelResponse() + + # Execution + try: + VertexGeminiConfig._process_candidates( + _candidates=candidates, + model_response=model_response, + standard_optional_params={}, + cumulative_tool_call_index=0 + ) + except UnboundLocalError as e: + pytest.fail(f"UnboundLocalError raised: {e}") + except Exception as e: + # Other exceptions might be okay if they are not UnboundLocalError, + # but ideally it should pass without error or raise a specific error if parts are required. + # However, the goal is to verify thought_signatures doesn't crash. + pass + + # Verify that we didn't crash with UnboundLocalError + +if __name__ == "__main__": + test_process_candidates_unbound_local_error_fix() + print("Test passed!") diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index 591b33911dc..b5637db3e52 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -1175,6 +1175,30 @@ def test_vertex_ai_moonshot_uses_openai_handler(): ) +def test_vertex_ai_zai_uses_openai_handler(): + """ + Ensure ZAI partner models re-use the OpenAI-format handler. + """ + from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( + VertexAIPartnerModels, + ) + + assert VertexAIPartnerModels.should_use_openai_handler( + "zai-org/glm-4.7-maas" + ) + + +def test_vertex_ai_zai_is_partner_model(): + """ + Ensure ZAI models are detected as Vertex AI partner models. + """ + from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( + VertexAIPartnerModels, + ) + + assert VertexAIPartnerModels.is_vertex_partner_model("zai-org/glm-4.7-maas") + + def test_build_vertex_schema_empty_properties(): """ Test _build_vertex_schema handles empty properties objects correctly. diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py index 389c8446135..80d65991acb 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py @@ -13,6 +13,7 @@ sys.path.insert( import litellm from litellm.llms.vertex_ai.vertex_llm_base import VertexBase +from litellm.llms.vertex_ai.common_utils import _get_gemini_url def run_sync(coro): @@ -1048,3 +1049,139 @@ class TestVertexBase: MockCredentials.from_info.assert_called_once_with(json_obj) mock_creds.with_scopes.assert_called_once_with(scopes) assert result == "scoped_creds" + + def test_get_token_and_url_with_api_key(self): + """Test that API key authentication routes to Google AI Studio endpoint""" + vertex_base = VertexBase() + + # Test with API key and no credentials - should use Google AI Studio endpoint + auth_header, url = vertex_base._get_token_and_url( + model="gemini-2.0-flash-exp", + auth_header=None, + gemini_api_key="test-api-key-123", + vertex_project="test-project", + vertex_location="us-central1", + vertex_credentials=None, # No service account credentials + stream=False, + custom_llm_provider="vertex_ai", + api_base=None, + should_use_v1beta1_features=False, + mode="chat", + ) + + # Should route to Google AI Studio endpoint + assert "generativelanguage.googleapis.com" in url + assert "gemini-2.0-flash-exp" in url + assert "key=test-api-key-123" in url + assert auth_header is None # API key is in URL, not header + + def test_get_token_and_url_with_credentials(self): + """Test that service account credentials route to Vertex AI endpoint""" + vertex_base = VertexBase() + + mock_creds = MagicMock() + mock_creds.token = "mock-bearer-token" + mock_creds.expired = False + + with patch.object( + vertex_base, "_ensure_access_token", return_value=("mock-bearer-token", "test-project") + ): + # Test with credentials - should use Vertex AI endpoint + auth_header, url = vertex_base._get_token_and_url( + model="gemini-2.0-flash-exp", + auth_header="mock-bearer-token", + gemini_api_key=None, + vertex_project="test-project", + vertex_location="us-central1", + vertex_credentials={"type": "service_account"}, + stream=False, + custom_llm_provider="vertex_ai", + api_base=None, + should_use_v1beta1_features=False, + mode="chat", + ) + + # Should route to Vertex AI endpoint + assert "aiplatform.googleapis.com" in url + assert "projects/test-project" in url + assert "locations/us-central1" in url + assert auth_header == "mock-bearer-token" + + def test_get_token_and_url_api_key_with_streaming(self): + """Test API key authentication with streaming enabled""" + vertex_base = VertexBase() + + auth_header, url = vertex_base._get_token_and_url( + model="gemini-2.0-flash-exp", + auth_header=None, + gemini_api_key="test-api-key-456", + vertex_project="test-project", + vertex_location="us-central1", + vertex_credentials=None, + stream=True, # Streaming enabled + custom_llm_provider="vertex_ai", + api_base=None, + should_use_v1beta1_features=False, + mode="chat", + ) + + # Should route to Google AI Studio endpoint with streaming + assert "generativelanguage.googleapis.com" in url + assert "streamGenerateContent" in url + assert "key=test-api-key-456" in url + assert "alt=sse" in url + assert auth_header is None + + def test_get_token_and_url_api_key_priority(self): + """Test that credentials take priority over API key when both are provided""" + vertex_base = VertexBase() + + # When both API key and credentials are provided, credentials take priority + mock_creds = MagicMock() + mock_creds.token = "mock-bearer-token" + mock_creds.expired = False + + with patch.object( + vertex_base, "_ensure_access_token", return_value=("mock-bearer-token", "test-project") + ): + auth_header, url = vertex_base._get_token_and_url( + model="gemini-2.0-flash-exp", + auth_header="mock-bearer-token", + gemini_api_key="test-api-key-789", + vertex_project="test-project", + vertex_location="us-central1", + vertex_credentials={"type": "service_account"}, # Credentials provided + stream=False, + custom_llm_provider="vertex_ai", + api_base=None, + should_use_v1beta1_features=False, + mode="chat", + ) + + # Should use Vertex AI endpoint with Bearer token (credentials take priority) + assert "aiplatform.googleapis.com" in url + assert auth_header == "mock-bearer-token" + + def test_get_token_and_url_with_embedding_mode(self): + """Test API key authentication with embedding mode""" + vertex_base = VertexBase() + + auth_header, url = vertex_base._get_token_and_url( + model="text-embedding-004", + auth_header=None, + gemini_api_key="test-embedding-key", + vertex_project="test-project", + vertex_location="us-central1", + vertex_credentials=None, + stream=False, + custom_llm_provider="vertex_ai", + api_base=None, + should_use_v1beta1_features=False, + mode="embedding", + ) + + # Should route to Google AI Studio endpoint for embeddings + assert "generativelanguage.googleapis.com" in url + assert "embedContent" in url + assert "key=test-embedding-key" in url + assert auth_header is None \ No newline at end of file diff --git a/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py b/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py index e36a494998b..fd5f8f3eff8 100644 --- a/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py +++ b/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py @@ -14,6 +14,10 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) import litellm +from litellm.llms.watsonx.audio_transcription.transformation import ( + IBMWatsonXAudioTranscriptionConfig, +) +from litellm.types.utils import TranscriptionResponse class TestWatsonXAudioTranscription: @@ -189,3 +193,72 @@ class TestWatsonXAudioTranscription: # Verify file is sent separately files = captured_request.get("files", {}) assert "file" in files + + def test_transform_audio_transcription_response_removes_model_field(self): + """ + Test that transform_audio_transcription_response removes the 'model' field + from WatsonX response before creating TranscriptionResponse. + + This test ensures that when WatsonX returns a response with a 'model' field, + it is removed before creating the TranscriptionResponse object, since + TranscriptionResponse doesn't accept a 'model' parameter. + """ + handler = IBMWatsonXAudioTranscriptionConfig() + + # Mock response with 'model' field (as WatsonX may return) + mock_response = MagicMock() + mock_response.json.return_value = { + "text": "Hello, this is a test transcription.", + "model": "whisper-large-v3-turbo", # This field should be removed + "duration": 5.5, + } + mock_response.text = '{"text": "Hello, this is a test transcription.", "model": "whisper-large-v3-turbo", "duration": 5.5}' + + # This should not raise a TypeError - model field should be removed + result = handler.transform_audio_transcription_response(mock_response) + + # Verify the result is a TranscriptionResponse + assert isinstance(result, TranscriptionResponse) + + # Verify the text is correct + assert result.text == "Hello, this is a test transcription." + + # Verify duration is set via dictionary assignment + assert result["duration"] == 5.5 + + # Verify the model field is NOT in the serialized result + # Check via model_dump() or dict() to ensure it's not in the output + try: + result_dict = result.model_dump() + except AttributeError: + # Fallback for pydantic v1 + result_dict = result.dict() + + # The 'model' field should not be in the result + assert "model" not in result_dict, "Model field should be removed from response" + + def test_transform_audio_transcription_response_without_model_field(self): + """ + Test that transform_audio_transcription_response works correctly + when WatsonX response doesn't include a 'model' field. + """ + handler = IBMWatsonXAudioTranscriptionConfig() + + # Mock response without 'model' field + mock_response = MagicMock() + mock_response.json.return_value = { + "text": "Hello, this is a test transcription.", + "duration": 5.5, + } + mock_response.text = '{"text": "Hello, this is a test transcription.", "duration": 5.5}' + + result = handler.transform_audio_transcription_response(mock_response) + + # Verify the result is a TranscriptionResponse + assert isinstance(result, TranscriptionResponse) + + # Verify the text is correct + assert result.text == "Hello, this is a test transcription." + + # Verify duration is set via dictionary assignment + assert result["duration"] == 5.5 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 8062243dfdd..f1558ac5791 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 @@ -1,4 +1,5 @@ import asyncio +from datetime import datetime, timedelta from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -7,7 +8,12 @@ from fastapi import HTTPException from mcp import ReadResourceResult, Resource from mcp.types import Prompt, ResourceTemplate, TextResourceContents -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import ( + LiteLLM_MCPServerTable, + MCPTransport, + UserAPIKeyAuth, +) +from litellm.types.mcp_server.mcp_server_manager import MCPServer @pytest.mark.asyncio @@ -1688,3 +1694,99 @@ def test_filter_tools_by_allowed_tools(): assert len(filtered_tools) == 2 assert filtered_tools[0].name == "my_api_mcp-getpetbyid" assert filtered_tools[1].name == "my_api_mcp-findpetsbystatus" + + +def _make_db_mcp_server(server_id: str, updated_at: datetime) -> LiteLLM_MCPServerTable: + return LiteLLM_MCPServerTable( + server_id=server_id, + server_name="server", + alias="server", + url="https://example.com", + transport=MCPTransport.http, + created_at=updated_at, + updated_at=updated_at, + mcp_info={}, + ) + + +class TestMCPServerManagerReload: + @pytest.mark.asyncio + async def test_reuses_existing_server_when_updated_at_matches(self): + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + timestamp = datetime.utcnow() + existing_server = MCPServer( + server_id="server-1", + name="server", + transport=MCPTransport.http, + updated_at=timestamp, + ) + manager.registry = {existing_server.server_id: existing_server} + + db_row = _make_db_mcp_server("server-1", timestamp) + + with patch( + "litellm.proxy._experimental.mcp_server.db.get_all_mcp_servers", + new=AsyncMock(return_value=[db_row]), + ) as mock_get_all, patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=object(), + ), patch.object( + manager, "build_mcp_server_from_table", AsyncMock() + ) as mock_build: + await manager.reload_servers_from_database() + + mock_get_all.assert_awaited_once() + mock_build.assert_not_awaited() + assert manager.registry["server-1"] is existing_server + + @pytest.mark.asyncio + async def test_rebuilds_server_when_updated_at_changes(self): + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + timestamp = datetime.utcnow() + existing_server = MCPServer( + server_id="server-1", + name="server", + transport=MCPTransport.http, + updated_at=timestamp, + ) + manager.registry = {existing_server.server_id: existing_server} + + new_timestamp = timestamp + timedelta(minutes=5) + db_row = _make_db_mcp_server("server-1", new_timestamp) + rebuilt_server = MCPServer( + server_id="server-1", + name="server", + transport=MCPTransport.http, + updated_at=new_timestamp, + ) + + with patch( + "litellm.proxy._experimental.mcp_server.db.get_all_mcp_servers", + new=AsyncMock(return_value=[db_row]), + ) as mock_get_all, patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=object(), + ), patch.object( + manager, + "build_mcp_server_from_table", + AsyncMock(return_value=rebuilt_server), + ) as mock_build: + await manager.reload_servers_from_database() + + mock_get_all.assert_awaited_once() + mock_build.assert_awaited_once_with(db_row) + assert manager.registry["server-1"] is rebuilt_server diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py new file mode 100644 index 00000000000..b1bef63933e --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -0,0 +1,131 @@ +""" +Unit tests for auth_utils functions related to rate limiting. +""" + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.auth_utils import ( + get_key_model_rpm_limit, + get_key_model_tpm_limit, +) + + +class TestGetKeyModelRpmLimit: + """Tests for get_key_model_rpm_limit function.""" + + def test_returns_key_metadata_when_present(self): + """Key metadata takes priority over team metadata.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + metadata={"model_rpm_limit": {"gpt-4": 100}}, + team_metadata={"model_rpm_limit": {"gpt-4": 50}}, + ) + result = get_key_model_rpm_limit(user_api_key_dict) + assert result == {"gpt-4": 100} + + def test_falls_back_to_team_metadata_when_key_has_other_metadata(self): + """Should fall back to team metadata when key metadata exists but has no model_rpm_limit.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + metadata={ + "some_other_key": "value" + }, # Has metadata, but not model_rpm_limit + team_metadata={"model_rpm_limit": {"gpt-4": 50}}, + ) + result = get_key_model_rpm_limit(user_api_key_dict) + assert result == {"gpt-4": 50} + + def test_extracts_from_model_max_budget(self): + """Should extract rpm_limit from model_max_budget when metadata is empty.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + model_max_budget={ + "gpt-4": {"rpm_limit": 100, "tpm_limit": 1000}, + "gpt-3.5-turbo": {"rpm_limit": 200}, + }, + ) + result = get_key_model_rpm_limit(user_api_key_dict) + assert result == {"gpt-4": 100, "gpt-3.5-turbo": 200} + + def test_skips_models_without_rpm_limit(self): + """Should skip models that don't have rpm_limit in model_max_budget.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + model_max_budget={ + "gpt-4": {"rpm_limit": 100}, + "gpt-3.5-turbo": {"tpm_limit": 1000}, # No rpm_limit + }, + ) + result = get_key_model_rpm_limit(user_api_key_dict) + assert result == {"gpt-4": 100} + + def test_returns_none_when_no_limits_configured(self): + """Should return None when no rate limits are configured.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + result = get_key_model_rpm_limit(user_api_key_dict) + assert result is None + + +class TestGetKeyModelTpmLimit: + """Tests for get_key_model_tpm_limit function.""" + + def test_returns_key_metadata_when_present(self): + """Key metadata takes priority over team metadata.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + metadata={"model_tpm_limit": {"gpt-4": 10000}}, + team_metadata={"model_tpm_limit": {"gpt-4": 5000}}, + ) + result = get_key_model_tpm_limit(user_api_key_dict) + assert result == {"gpt-4": 10000} + + def test_falls_back_to_team_metadata_when_key_has_other_metadata(self): + """Should fall back to team metadata when key metadata exists but has no model_tpm_limit.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + metadata={ + "some_other_key": "value" + }, # Has metadata, but not model_tpm_limit + team_metadata={"model_tpm_limit": {"gpt-4": 5000}}, + ) + result = get_key_model_tpm_limit(user_api_key_dict) + assert result == {"gpt-4": 5000} + + def test_extracts_from_model_max_budget(self): + """Should extract tpm_limit from model_max_budget when metadata is empty.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + model_max_budget={ + "gpt-4": {"tpm_limit": 10000, "rpm_limit": 100}, + "gpt-3.5-turbo": {"tpm_limit": 20000}, + }, + ) + result = get_key_model_tpm_limit(user_api_key_dict) + assert result == {"gpt-4": 10000, "gpt-3.5-turbo": 20000} + + def test_skips_models_without_tpm_limit(self): + """Should skip models that don't have tpm_limit in model_max_budget.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + model_max_budget={ + "gpt-4": {"tpm_limit": 10000}, + "gpt-3.5-turbo": {"rpm_limit": 100}, # No tpm_limit + }, + ) + result = get_key_model_tpm_limit(user_api_key_dict) + assert result == {"gpt-4": 10000} + + def test_returns_none_when_no_limits_configured(self): + """Should return None when no rate limits are configured.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + result = get_key_model_tpm_limit(user_api_key_dict) + assert result is None + + def test_model_max_budget_priority_over_team(self): + """model_max_budget should take priority over team_metadata.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + model_max_budget={"gpt-4": {"tpm_limit": 10000}}, + team_metadata={"model_tpm_limit": {"gpt-4": 5000}}, + ) + result = get_key_model_tpm_limit(user_api_key_dict) + assert result == {"gpt-4": 10000} diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index c04b8114939..e7b27908c14 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -248,6 +248,83 @@ async def test_authenticate_user_wrong_password(): assert "Invalid credentials" in exc_info.value.message +@pytest.mark.asyncio +async def test_authenticate_user_email_case_insensitive_login(): + """Test that email lookup is case-insensitive during login""" + master_key = "sk-1234" + stored_email = "testemail@test.com" + login_email_mixed_case = "testEmail@test.com" + correct_password = "correct-password" + hashed_password = hash_token(token=correct_password) + + # `LiteLLM_UserTable` does not define a `password` field, but `authenticate_user()` + # expects `user_row.password` to exist (invite-link login). Use a simple object. + mock_user = MagicMock() + mock_user.user_id = "test-user-123" + mock_user.user_email = stored_email + mock_user.password = hashed_password + mock_user.user_role = LitellmUserRoles.INTERNAL_USER + + def mock_find_first(**kwargs): + where = kwargs.get("where", {}) + user_email = where.get("user_email", {}) + if user_email.get("mode") != "insensitive": + return None + if str(user_email.get("equals", "")).lower() == stored_email.lower(): + return mock_user + return None + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( + side_effect=mock_find_first + ) + + with patch.dict( + os.environ, + { + "DATABASE_URL": "postgresql://test:test@localhost/test", + "UI_USERNAME": "admin", + "UI_PASSWORD": "admin-password", + }, + ): + with patch( + "litellm.proxy.auth.login_utils.expire_previous_ui_session_tokens", + new_callable=AsyncMock, + return_value=None, + ): + with patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", + new_callable=AsyncMock, + ) as mock_generate_key: + mock_generate_key.side_effect = [ + {"token": "token-1"}, + {"token": "token-2"}, + ] + + result_mixed = await authenticate_user( + username=login_email_mixed_case, + password=correct_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + result_lower = await authenticate_user( + username=stored_email, + password=correct_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + + assert result_mixed.user_id == result_lower.user_id == "test-user-123" + assert result_mixed.user_email == result_lower.user_email == stored_email + + calls = mock_prisma_client.db.litellm_usertable.find_first.await_args_list + assert len(calls) == 2 + for call, expected_username in zip(calls, [login_email_mixed_case, stored_email]): + where = call.kwargs["where"] + assert where["user_email"]["equals"] == expected_username + assert where["user_email"]["mode"] == "insensitive" + + @pytest.mark.asyncio async def test_authenticate_user_database_required_for_admin(): """Test that database is required for admin login""" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index fcc8c1f0f2e..5f49db66089 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -338,6 +338,17 @@ async def test_proxy_admin_expired_key_from_cache(): f"Exception message should mention 'Expired Key', got: {exc_info.value.message}" ) + # Verify that the param field does NOT leak the full API key (Issue #18731) + # The param should be abbreviated like "sk-...XXXX" not the full plaintext key + assert exc_info.value.param is not None, "Exception should have 'param' attribute" + assert exc_info.value.param != api_key, ( + f"SECURITY: Full API key should NOT be in param field! " + f"Got: {exc_info.value.param}, Expected abbreviated format like 'sk-...XXXX'" + ) + assert exc_info.value.param.startswith("sk-..."), ( + f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}" + ) + # Verify that cache deletion was called mock_delete_cache.assert_called_once() call_args = mock_delete_cache.call_args @@ -347,3 +358,4 @@ async def test_proxy_admin_expired_key_from_cache(): finally: # Clean up - restore original values if needed pass + diff --git a/tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py b/tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py new file mode 100644 index 00000000000..1492acb0794 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py @@ -0,0 +1,275 @@ +""" +Tests for the RDS IAM token proactive refresh implementation. + +Tests for GitHub Issue #16220: RDS IAM authentication connection failures after 15 minutes. + +The fix implements: +1. Proactive background token refresh (refreshes 3 min before expiration) +2. Precise sleep timing (1 wake-up per token cycle instead of polling) +3. Proper locking during reconnection +4. Fixed __getattr__ fallback that now waits for reconnection + +Run these tests: + poetry run pytest tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py -v -s +""" + +import asyncio +import os +import urllib.parse +from datetime import datetime, timedelta +from unittest.mock import MagicMock, patch + +import pytest + + +class TestPrismaWrapperTokenRefresh: + """Tests for the PrismaWrapper RDS IAM token refresh implementation.""" + + @pytest.fixture + def setup_env(self): + """Setup environment variables for testing.""" + os.environ["DATABASE_HOST"] = "test-host.rds.amazonaws.com" + os.environ["DATABASE_PORT"] = "5432" + os.environ["DATABASE_USER"] = "test_user" + os.environ["DATABASE_NAME"] = "test_db" + os.environ["IAM_TOKEN_DB_AUTH"] = "True" + yield + # Cleanup + for key in [ + "DATABASE_HOST", + "DATABASE_PORT", + "DATABASE_USER", + "DATABASE_NAME", + "DATABASE_URL", + "IAM_TOKEN_DB_AUTH", + "DATABASE_SCHEMA", + ]: + os.environ.pop(key, None) + + def _generate_mock_token(self, expires_in_seconds: int = 900) -> str: + """Generate a mock IAM token with expiration info.""" + now = datetime.utcnow() + date_str = now.strftime("%Y%m%dT%H%M%SZ") + # Build the token like AWS does + token = f"mock-token?X-Amz-Date={date_str}&X-Amz-Expires={expires_in_seconds}&X-Amz-Signature=abc123" + return urllib.parse.quote(token, safe="") + + def _set_database_url_with_token(self, expires_in_seconds: int = 900): + """Set DATABASE_URL with a mock token.""" + token = self._generate_mock_token(expires_in_seconds) + os.environ[ + "DATABASE_URL" + ] = f"postgresql://test_user:{token}@test-host:5432/test_db" + + @pytest.mark.asyncio + async def test_is_token_expired_fresh(self, setup_env): + """Test that fresh token is not detected as expired.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + self._set_database_url_with_token(expires_in_seconds=900) + db_url = os.getenv("DATABASE_URL") + + assert wrapper.is_token_expired(db_url) is False + + @pytest.mark.asyncio + async def test_is_token_expired_old(self, setup_env): + """Test that old token is detected as expired.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + # Create an expired token + old_date = datetime.utcnow() - timedelta(seconds=901) + date_str = old_date.strftime("%Y%m%dT%H%M%SZ") + token = ( + f"mock-token?X-Amz-Date={date_str}&X-Amz-Expires=900&X-Amz-Signature=abc" + ) + encoded_token = urllib.parse.quote(token, safe="") + db_url = f"postgresql://test_user:{encoded_token}@test-host:5432/test_db" + + assert wrapper.is_token_expired(db_url) is True + + @pytest.mark.asyncio + async def test_start_stop_token_refresh_task(self, setup_env): + """Test that token refresh task starts and stops correctly.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + # Set a valid token + self._set_database_url_with_token(expires_in_seconds=900) + + # Start the task + await wrapper.start_token_refresh_task() + assert wrapper._token_refresh_task is not None + assert not wrapper._token_refresh_task.done() + + # Stop the task + await wrapper.stop_token_refresh_task() + assert wrapper._token_refresh_task is None + + @pytest.mark.asyncio + async def test_start_task_not_enabled(self, setup_env): + """Test that task doesn't start when IAM auth is not enabled.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + # IAM auth disabled + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False) + + await wrapper.start_token_refresh_task() + assert wrapper._token_refresh_task is None + + @pytest.mark.asyncio + async def test_is_token_expired_null(self, setup_env): + """Test that None token is treated as expired.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + assert wrapper.is_token_expired(None) is True + + +class TestTokenExpirationParsing: + """Tests for token expiration parsing utilities.""" + + def test_parse_token_expiration_valid(self): + """Test parsing expiration from a valid token.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + # Create a token with known expiration + token = "mock-token?X-Amz-Date=20240101T120000Z&X-Amz-Expires=900&X-Amz-Signature=abc" + + expiration = wrapper._parse_token_expiration(token) + + assert expiration is not None + expected = datetime(2024, 1, 1, 12, 0, 0) + timedelta(seconds=900) + assert expiration == expected + + def test_parse_token_expiration_invalid(self): + """Test that invalid token returns None.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + # Invalid tokens + assert wrapper._parse_token_expiration(None) is None + assert wrapper._parse_token_expiration("no-query-params") is None + assert wrapper._parse_token_expiration("?missing=params") is None + + +class TestBackgroundRefreshLoop: + """Tests for the background refresh loop timing.""" + + @pytest.fixture + def setup_env(self): + """Setup environment variables for testing.""" + os.environ["DATABASE_HOST"] = "test-host.rds.amazonaws.com" + os.environ["DATABASE_PORT"] = "5432" + os.environ["DATABASE_USER"] = "test_user" + os.environ["DATABASE_NAME"] = "test_db" + yield + # Cleanup + for key in [ + "DATABASE_HOST", + "DATABASE_PORT", + "DATABASE_USER", + "DATABASE_NAME", + "DATABASE_URL", + ]: + os.environ.pop(key, None) + + @pytest.mark.asyncio + async def test_calculate_seconds_fallback_when_no_url(self, setup_env): + """Test that fallback is used when DATABASE_URL is not set.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + + mock_prisma = MagicMock() + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + # Don't set DATABASE_URL + seconds = wrapper._calculate_seconds_until_refresh() + + # Should return fallback interval + assert seconds == wrapper.FALLBACK_REFRESH_INTERVAL_SECONDS + + +# ============================================================================ +# DEMONSTRATION SCRIPT +# ============================================================================ + + +async def demonstrate_fix(): + """ + Demonstrates the fix for the RDS IAM token expiration bug. + + Shows how the proactive refresh prevents the 15-minute connection failure. + """ + # Import the actual implementation + try: + from litellm.proxy.db.prisma_client import PrismaWrapper + except ImportError: + return + + # Setup mock environment + os.environ["DATABASE_HOST"] = "mock-rds.region.rds.amazonaws.com" + os.environ["DATABASE_PORT"] = "5432" + os.environ["DATABASE_USER"] = "iam_user" + os.environ["DATABASE_NAME"] = "litellm" + + # Create initial token (expires in 10 seconds for demo) + now = datetime.utcnow() + date_str = now.strftime("%Y%m%dT%H%M%SZ") + token = f"mock-token?X-Amz-Date={date_str}&X-Amz-Expires=10&X-Amz-Signature=abc123" + encoded_token = urllib.parse.quote(token, safe="") + os.environ[ + "DATABASE_URL" + ] = f"postgresql://iam_user:{encoded_token}@mock-rds:5432/litellm" + + # Create mock prisma client + mock_prisma = MagicMock() + + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=True) + + # Override buffer for faster demo + wrapper.TOKEN_REFRESH_BUFFER_SECONDS = 3 + wrapper.FALLBACK_REFRESH_INTERVAL_SECONDS = 5 + _ = wrapper._calculate_seconds_until_refresh() # Verify calculation works + db_url = os.getenv("DATABASE_URL") + is_expired = wrapper.is_token_expired(db_url) + assert is_expired is False, "Fresh token should not be expired!" + + # Mock the _token_refresh_loop to prevent it from actually running + async def mock_loop(): + try: + await asyncio.sleep(1000) + except asyncio.CancelledError: + pass + + with patch.object(wrapper, "_token_refresh_loop", side_effect=mock_loop): + await wrapper.start_token_refresh_task() + await wrapper.stop_token_refresh_task() + + # Cleanup + for key in [ + "DATABASE_HOST", + "DATABASE_PORT", + "DATABASE_USER", + "DATABASE_NAME", + "DATABASE_URL", + ]: + os.environ.pop(key, None) + + +if __name__ == "__main__": + asyncio.run(demonstrate_fix()) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py index 35ed49a84ed..fd72185d1e7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py @@ -2,7 +2,6 @@ Unit tests for Qualifire guardrail integration. """ -import sys from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -75,9 +74,37 @@ class TestQualifireGuardrailInit: assert guardrail.on_flagged == "monitor" + def test_init_with_default_api_base(self): + """Test that default API base is set when not provided.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + DEFAULT_QUALIFIRE_API_BASE, + QualifireGuardrail, + ) + + guardrail = QualifireGuardrail( + api_key="test_key", + guardrail_name="test_guardrail", + ) + + assert guardrail.qualifire_api_base == DEFAULT_QUALIFIRE_API_BASE + + def test_init_with_custom_api_base(self): + """Test initialization with custom API base URL.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) + + guardrail = QualifireGuardrail( + api_key="test_key", + api_base="https://custom.qualifire.ai", + guardrail_name="test_guardrail", + ) + + assert guardrail.qualifire_api_base == "https://custom.qualifire.ai" + class TestQualifireGuardrailMessageConversion: - """Tests for message conversion to Qualifire format.""" + """Tests for message conversion to API format.""" def test_convert_simple_messages(self): """Test conversion of simple text messages.""" @@ -95,15 +122,13 @@ class TestQualifireGuardrailMessageConversion: {"role": "assistant", "content": "Hi there!"}, ] - # Create mock LLMMessage class - mock_llm_message = MagicMock() + result = guardrail._convert_messages_to_api_format(messages) - with patch( - "litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire.QualifireGuardrail._convert_messages_to_qualifire_format" - ) as mock_convert: - mock_convert.return_value = [mock_llm_message, mock_llm_message] - result = guardrail._convert_messages_to_qualifire_format(messages) - assert len(result) == 2 + assert len(result) == 2 + assert result[0]["role"] == "user" + assert result[0]["content"] == "Hello, world!" + assert result[1]["role"] == "assistant" + assert result[1]["content"] == "Hi there!" def test_convert_multimodal_messages(self): """Test conversion of multimodal messages with text parts.""" @@ -126,112 +151,258 @@ class TestQualifireGuardrailMessageConversion: }, ] - with patch( - "litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire.QualifireGuardrail._convert_messages_to_qualifire_format" - ) as mock_convert: - mock_convert.return_value = [MagicMock()] - result = guardrail._convert_messages_to_qualifire_format(messages) - assert len(result) == 1 + result = guardrail._convert_messages_to_api_format(messages) + + assert len(result) == 1 + assert result[0]["role"] == "user" + assert result[0]["content"] == "First part\nSecond part" + + def test_convert_messages_with_tool_calls(self): + """Test conversion of messages with tool calls.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) + + guardrail = QualifireGuardrail( + api_key="test_key", + guardrail_name="test_guardrail", + ) + + messages = [ + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "NYC"}', + }, + } + ], + }, + ] + + result = guardrail._convert_messages_to_api_format(messages) + + assert len(result) == 1 + assert result[0]["role"] == "assistant" + assert "tool_calls" in result[0] + assert len(result[0]["tool_calls"]) == 1 + assert result[0]["tool_calls"][0]["id"] == "call_123" + assert result[0]["tool_calls"][0]["name"] == "get_weather" + assert result[0]["tool_calls"][0]["arguments"] == {"location": "NYC"} -class TestQualifireGuardrailEvaluateKwargs: - """Tests for evaluate kwargs passed to Qualifire client.""" +class TestQualifireGuardrailToolConversion: + """Tests for tool definition conversion.""" + + def test_convert_openai_function_tools(self): + """Test conversion of OpenAI function tool format.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) + + guardrail = QualifireGuardrail( + api_key="test_key", + guardrail_name="test_guardrail", + ) + + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a location", + "parameters": {"type": "object", "properties": {}}, + }, + } + ] + + result = guardrail._convert_tools_to_api_format(tools) + + assert result is not None + assert len(result) == 1 + assert result[0]["name"] == "get_weather" + assert result[0]["description"] == "Get weather for a location" + + def test_convert_empty_tools(self): + """Test that empty tools returns None.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) + + guardrail = QualifireGuardrail( + api_key="test_key", + guardrail_name="test_guardrail", + ) + + result = guardrail._convert_tools_to_api_format(None) + assert result is None + + result = guardrail._convert_tools_to_api_format([]) + assert result is None + + +class TestQualifireGuardrailAPICall: + """Tests for API call with httpx client.""" @pytest.mark.asyncio async def test_evaluate_called_with_prompt_injections(self): - """Test that evaluate is called with prompt_injections enabled.""" - # Mock the qualifire module and its types - mock_qualifire_types = MagicMock() - mock_llm_message = MagicMock() - mock_llm_tool_call = MagicMock() - mock_message_instance = MagicMock() - mock_llm_message.return_value = mock_message_instance - - mock_qualifire_types.LLMMessage = mock_llm_message - mock_qualifire_types.LLMToolCall = mock_llm_tool_call - - with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}): - from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( - QualifireGuardrail, - ) + """Test that evaluate endpoint is called with prompt_injections enabled.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) - guardrail = QualifireGuardrail( - api_key="test_key", - prompt_injections=True, - guardrail_name="test_guardrail", - ) + guardrail = QualifireGuardrail( + api_key="test_key", + prompt_injections=True, + guardrail_name="test_guardrail", + ) - # Mock the client - mock_client = MagicMock() - mock_result = MagicMock() - mock_result.score = 100 - mock_result.status = "completed" - mock_result.evaluationResults = [] - mock_client.evaluate.return_value = mock_result - guardrail._client = mock_client + # Mock the async HTTP handler + mock_response = MagicMock() + mock_response.json.return_value = { + "score": 100, + "status": "completed", + "evaluationResults": [], + } + mock_response.raise_for_status = MagicMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) - messages = [{"role": "user", "content": "Hello, world!"}] + messages = [{"role": "user", "content": "Hello, world!"}] - await guardrail._run_qualifire_check( - messages=messages, output=None, dynamic_params={} - ) + await guardrail._run_qualifire_check( + messages=messages, output=None, dynamic_params={} + ) - # Verify evaluate was called with correct kwargs - mock_client.evaluate.assert_called_once() - call_kwargs = mock_client.evaluate.call_args[1] - assert call_kwargs["prompt_injections"] is True - assert "messages" in call_kwargs + # Verify the API was called + guardrail.async_handler.post.assert_called_once() + call_kwargs = guardrail.async_handler.post.call_args[1] + + assert "json" in call_kwargs + payload = call_kwargs["json"] + assert payload["prompt_injections"] is True + assert "messages" in payload + assert call_kwargs["url"].endswith("/api/evaluation/evaluate") @pytest.mark.asyncio async def test_evaluate_called_with_multiple_checks(self): """Test that evaluate is called with multiple checks enabled.""" - # Mock the qualifire module and its types - mock_qualifire_types = MagicMock() - mock_llm_message = MagicMock() - mock_llm_tool_call = MagicMock() - mock_message_instance = MagicMock() - mock_llm_message.return_value = mock_message_instance - - mock_qualifire_types.LLMMessage = mock_llm_message - mock_qualifire_types.LLMToolCall = mock_llm_tool_call - - with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}): - from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( - QualifireGuardrail, - ) + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) - guardrail = QualifireGuardrail( - api_key="test_key", - prompt_injections=True, - pii_check=True, - hallucinations_check=True, - assertions=["Output must be valid JSON"], - guardrail_name="test_guardrail", - ) + guardrail = QualifireGuardrail( + api_key="test_key", + prompt_injections=True, + pii_check=True, + hallucinations_check=True, + assertions=["Output must be valid JSON"], + guardrail_name="test_guardrail", + ) - # Mock the client - mock_client = MagicMock() - mock_result = MagicMock() - mock_result.score = 100 - mock_result.status = "completed" - mock_result.evaluationResults = [] - mock_client.evaluate.return_value = mock_result - guardrail._client = mock_client + # Mock the async HTTP handler + mock_response = MagicMock() + mock_response.json.return_value = { + "score": 100, + "status": "completed", + "evaluationResults": [], + } + mock_response.raise_for_status = MagicMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) - messages = [{"role": "user", "content": "Hello, world!"}] + messages = [{"role": "user", "content": "Hello, world!"}] - await guardrail._run_qualifire_check( - messages=messages, output="Test output", dynamic_params={} - ) + await guardrail._run_qualifire_check( + messages=messages, output="Test output", dynamic_params={} + ) - # Verify evaluate was called with correct kwargs - mock_client.evaluate.assert_called_once() - call_kwargs = mock_client.evaluate.call_args[1] - assert call_kwargs["prompt_injections"] is True - assert call_kwargs["pii_check"] is True - assert call_kwargs["hallucinations_check"] is True - assert call_kwargs["assertions"] == ["Output must be valid JSON"] - assert call_kwargs["output"] == "Test output" + # Verify the API was called with correct payload + guardrail.async_handler.post.assert_called_once() + call_kwargs = guardrail.async_handler.post.call_args[1] + + payload = call_kwargs["json"] + assert payload["prompt_injections"] is True + assert payload["pii_check"] is True + assert payload["hallucinations_check"] is True + assert payload["assertions"] == ["Output must be valid JSON"] + assert payload["output"] == "Test output" + + @pytest.mark.asyncio + async def test_invoke_endpoint_used_with_evaluation_id(self): + """Test that invoke endpoint is used when evaluation_id is provided.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) + + guardrail = QualifireGuardrail( + api_key="test_key", + evaluation_id="eval_123", + guardrail_name="test_guardrail", + ) + + # Mock the async HTTP handler + mock_response = MagicMock() + mock_response.json.return_value = { + "score": 100, + "status": "completed", + "evaluationResults": [], + } + mock_response.raise_for_status = MagicMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + messages = [{"role": "user", "content": "Hello, world!"}] + + await guardrail._run_qualifire_check( + messages=messages, output="Test output", dynamic_params={} + ) + + # Verify the invoke endpoint was called + guardrail.async_handler.post.assert_called_once() + call_kwargs = guardrail.async_handler.post.call_args[1] + + assert call_kwargs["url"].endswith("/api/evaluation/invoke") + payload = call_kwargs["json"] + assert payload["evaluation_id"] == "eval_123" + assert payload["input"] == "Hello, world!" + assert payload["output"] == "Test output" + + @pytest.mark.asyncio + async def test_correct_headers_sent(self): + """Test that correct headers are sent with the API request.""" + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) + + guardrail = QualifireGuardrail( + api_key="my_api_key", + guardrail_name="test_guardrail", + ) + + # Mock the async HTTP handler + mock_response = MagicMock() + mock_response.json.return_value = { + "score": 100, + "status": "completed", + "evaluationResults": [], + } + mock_response.raise_for_status = MagicMock() + guardrail.async_handler.post = AsyncMock(return_value=mock_response) + + messages = [{"role": "user", "content": "Hello!"}] + + await guardrail._run_qualifire_check( + messages=messages, output=None, dynamic_params={} + ) + + call_kwargs = guardrail.async_handler.post.call_args[1] + headers = call_kwargs["headers"] + + assert headers["X-Qualifire-API-Key"] == "my_api_key" + assert headers["Content-Type"] == "application/json" class TestQualifireGuardrailCheckIfFlagged: @@ -248,12 +419,14 @@ class TestQualifireGuardrailCheckIfFlagged: guardrail_name="test_guardrail", ) - # Mock result with completed status and no flagged items - mock_result = MagicMock() - mock_result.status = "completed" - mock_result.evaluationResults = [] + # Result with completed status and no flagged items (dict format) + result = { + "status": "completed", + "score": 100, + "evaluationResults": [], + } - assert guardrail._check_if_flagged(mock_result) is False + assert guardrail._check_if_flagged(result) is False def test_check_if_flagged_returns_true_for_flagged_content(self): """Test that _check_if_flagged returns True when content is flagged.""" @@ -266,18 +439,25 @@ class TestQualifireGuardrailCheckIfFlagged: guardrail_name="test_guardrail", ) - # Mock result with flagged item - mock_inner_result = MagicMock() - mock_inner_result.flagged = True + # Result with flagged item (dict format matching API response) + result = { + "status": "completed", + "score": 15, + "evaluationResults": [ + { + "type": "prompt_injection", + "results": [ + { + "flagged": True, + "score": 0.15, + "reason": "Prompt injection detected", + } + ], + } + ], + } - mock_eval_result = MagicMock() - mock_eval_result.results = [mock_inner_result] - - mock_result = MagicMock() - mock_result.status = "completed" - mock_result.evaluationResults = [mock_eval_result] - - assert guardrail._check_if_flagged(mock_result) is True + assert guardrail._check_if_flagged(result) is True def test_check_if_flagged_returns_false_when_no_flagged_items(self): """Test that _check_if_flagged returns False when no items are flagged.""" @@ -291,17 +471,24 @@ class TestQualifireGuardrailCheckIfFlagged: ) # Result with evaluation results but nothing flagged - mock_inner_result = MagicMock() - mock_inner_result.flagged = False + result = { + "status": "completed", + "score": 95, + "evaluationResults": [ + { + "type": "prompt_injection", + "results": [ + { + "flagged": False, + "score": 0.95, + "reason": "No issues detected", + } + ], + } + ], + } - mock_eval_result = MagicMock() - mock_eval_result.results = [mock_inner_result] - - mock_result = MagicMock() - mock_result.status = "success" - mock_result.evaluationResults = [mock_eval_result] - - assert guardrail._check_if_flagged(mock_result) is False + assert guardrail._check_if_flagged(result) is False class TestQualifireGuardrailShouldRun: diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index b76957dbf39..134fc84965f 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -247,7 +247,9 @@ async def test_rate_limiter_script_return_values_v3(monkeypatch, time_controller ) @pytest.mark.flaky(reruns=3) @pytest.mark.asyncio -async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object, time_controller): +async def test_normal_router_call_tpm_v3( + monkeypatch, rate_limit_object, time_controller +): """ Test normal router call with parallel request limiter v3 for TPM rate limiting """ @@ -394,8 +396,10 @@ async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object, time_co # Manually increment the token counter to simulate token usage from previous call # This simulates what would happen after a successful call - await local_cache.async_increment_cache(key=counter_key, value=15, ttl=2) # Use up most of our 10 token limit - + await local_cache.async_increment_cache( + key=counter_key, value=15, ttl=2 + ) # Use up most of our 10 token limit + # Make another request to test rate limiting - this should fail as we've consumed tokens with pytest.raises(HTTPException) as exc_info: await parallel_request_handler.async_pre_call_hook( @@ -535,7 +539,9 @@ async def test_async_log_failure_event_v3(): ) # Mock kwargs with user_api_key via standard_logging_object - mock_kwargs = {"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}} + mock_kwargs = { + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}} + } # Capture pipeline operations captured_ops = [] @@ -785,7 +791,7 @@ async def test_tpm_api_key_rate_limits_v3(): tpm_limit_per_model=tpms, models=[], ) - + user_api_key_dict.metadata["model_tpm_limit"] = tpms user_api_key_dict.metadata["model_rpm_limit"] = rpms @@ -804,32 +810,45 @@ async def test_tpm_api_key_rate_limits_v3(): # Return Error response to ensure HTTPException return { "overall_code": "OVER_LIMIT", - "statuses": [{'code': 'OK', 'current_limit': 2, 'limit_remaining': 1, 'rate_limit_type': 'requests', 'descriptor_key': 'model_per_key'}, - {'code': 'OVER_LIMIT', 'current_limit': 2, 'limit_remaining': -18, 'rate_limit_type': 'tokens', 'descriptor_key': 'model_per_key'}] + "statuses": [ + { + "code": "OK", + "current_limit": 2, + "limit_remaining": 1, + "rate_limit_type": "requests", + "descriptor_key": "model_per_key", + }, + { + "code": "OVER_LIMIT", + "current_limit": 2, + "limit_remaining": -18, + "rate_limit_type": "tokens", + "descriptor_key": "model_per_key", + }, + ], } - + parallel_request_handler.should_rate_limit = mock_should_rate_limit - + # Test the pre-call hook error = None try: - await parallel_request_handler.async_pre_call_hook( + await parallel_request_handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=local_cache, data={"model": model}, call_type="", ) except HTTPException as e: - error=e + error = e assert e.status_code == 429 assert "rate_limit_type" in e.headers assert e.headers.get("rate_limit_type") == "tokens" assert "retry-after" in e.headers - - + assert error is not None, "An Exception must be thrown" assert captured_descriptors is not None, "Rate limit descriptors should be captured" - + model_per_key_descriptor = None for descriptor in captured_descriptors: if descriptor["key"] == "model_per_key": @@ -837,9 +856,15 @@ async def test_tpm_api_key_rate_limits_v3(): break assert model_per_key_descriptor is not None, "Api-Key descriptor should be present" - assert model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}", "Api-Key value should combine api_key and model" - assert model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit, "Api-Key RPM limit should be set" - assert model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit, "Api-Key TPM limit should be set" + assert ( + model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}" + ), "Api-Key value should combine api_key and model" + assert ( + model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit + ), "Api-Key RPM limit should be set" + assert ( + model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit + ), "Api-Key TPM limit should be set" @pytest.mark.asyncio @@ -861,7 +886,7 @@ async def test_rpm_api_key_rate_limits_v3(): tpm_limit_per_model=tpms, models=[], ) - + user_api_key_dict.metadata["model_tpm_limit"] = tpms user_api_key_dict.metadata["model_rpm_limit"] = rpms @@ -880,31 +905,45 @@ async def test_rpm_api_key_rate_limits_v3(): # Return Error response to ensure HTTPException return { "overall_code": "OVER_LIMIT", - "statuses": [{'code': 'OVER_LIMIT', 'current_limit': 2, 'limit_remaining': -2, 'rate_limit_type': 'requests', 'descriptor_key': 'model_per_key'}, - {'code': 'OK', 'current_limit': 2, 'limit_remaining': 2, 'rate_limit_type': 'tokens', 'descriptor_key': 'model_per_key'}] + "statuses": [ + { + "code": "OVER_LIMIT", + "current_limit": 2, + "limit_remaining": -2, + "rate_limit_type": "requests", + "descriptor_key": "model_per_key", + }, + { + "code": "OK", + "current_limit": 2, + "limit_remaining": 2, + "rate_limit_type": "tokens", + "descriptor_key": "model_per_key", + }, + ], } - + parallel_request_handler.should_rate_limit = mock_should_rate_limit - + # Test the pre-call hook error = None try: - await parallel_request_handler.async_pre_call_hook( + await parallel_request_handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=local_cache, data={"model": model}, call_type="", ) except HTTPException as e: - error=e + error = e assert e.status_code == 429 assert "rate_limit_type" in e.headers assert e.headers.get("rate_limit_type") == "requests" assert "retry-after" in e.headers - + assert error is not None, "An Exception must be thrown" assert captured_descriptors is not None, "Rate limit descriptors should be captured" - + model_per_key_descriptor = None for descriptor in captured_descriptors: if descriptor["key"] == "model_per_key": @@ -912,9 +951,16 @@ async def test_rpm_api_key_rate_limits_v3(): break assert model_per_key_descriptor is not None, "Api-Key descriptor should be present" - assert model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}", "Api-Key value should combine api_key and model" - assert model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit, "Api-Key RPM limit should be set" - assert model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit, "Api-Key TPM limit should be set" + assert ( + model_per_key_descriptor["value"] == f"{_api_key_hash}:{model}" + ), "Api-Key value should combine api_key and model" + assert ( + model_per_key_descriptor["rate_limit"]["requests_per_unit"] == rpm_limit + ), "Api-Key RPM limit should be set" + assert ( + model_per_key_descriptor["rate_limit"]["tokens_per_unit"] == tpm_limit + ), "Api-Key TPM limit should be set" + @pytest.mark.asyncio async def test_team_member_rate_limits_v3(): @@ -925,7 +971,7 @@ async def test_team_member_rate_limits_v3(): _api_key = hash_token(_api_key) _team_id = "team_123" _user_id = "user_456" - + user_api_key_dict = UserAPIKeyAuth( api_key=_api_key, team_id=_team_id, @@ -933,7 +979,7 @@ async def test_team_member_rate_limits_v3(): team_member_rpm_limit=10, team_member_tpm_limit=1000, ) - + local_cache = DualCache() parallel_request_handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) @@ -947,15 +993,12 @@ async def test_team_member_rate_limits_v3(): nonlocal captured_descriptors captured_descriptors = descriptors # Return OK response to avoid HTTPException - return { - "overall_code": "OK", - "statuses": [] - } + return {"overall_code": "OK", "statuses": []} parallel_request_handler.should_rate_limit = mock_should_rate_limit # Test the pre-call hook - + await parallel_request_handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=local_cache, @@ -965,24 +1008,32 @@ async def test_team_member_rate_limits_v3(): # Verify team member descriptor was created assert captured_descriptors is not None, "Rate limit descriptors should be captured" - + team_member_descriptor = None for descriptor in captured_descriptors: if descriptor["key"] == "team_member": team_member_descriptor = descriptor break - - assert team_member_descriptor is not None, "Team member descriptor should be present" - assert team_member_descriptor["value"] == f"{_team_id}:{_user_id}", "Team member value should combine team_id and user_id" - assert team_member_descriptor["rate_limit"]["requests_per_unit"] == 10, "Team member RPM limit should be set" - assert team_member_descriptor["rate_limit"]["tokens_per_unit"] == 1000, "Team member TPM limit should be set" + + assert ( + team_member_descriptor is not None + ), "Team member descriptor should be present" + assert ( + team_member_descriptor["value"] == f"{_team_id}:{_user_id}" + ), "Team member value should combine team_id and user_id" + assert ( + team_member_descriptor["rate_limit"]["requests_per_unit"] == 10 + ), "Team member RPM limit should be set" + assert ( + team_member_descriptor["rate_limit"]["tokens_per_unit"] == 1000 + ), "Team member TPM limit should be set" @pytest.mark.asyncio async def test_dynamic_rate_limiting_v3(): """ Test that dynamic rate limiting only enforces limits when model has failures. - + When rpm_limit_type is set to "dynamic": - If model has no failures, rate limits should NOT be enforced (allow exceeding) - If model has failures above threshold, rate limits SHOULD be enforced @@ -990,75 +1041,75 @@ async def test_dynamic_rate_limiting_v3(): _api_key = "sk-12345" _api_key_hash = hash_token(_api_key) model = "gpt-3.5-turbo" - + # Set a low RPM limit to make testing easier user_api_key_dict = UserAPIKeyAuth( api_key=_api_key_hash, rpm_limit=2, metadata={"rpm_limit_type": "dynamic"}, ) - + local_cache = DualCache() parallel_request_handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Mock should_rate_limit to track if limits are enforced captured_descriptors = [] - + async def mock_should_rate_limit(descriptors, **kwargs): captured_descriptors.clear() captured_descriptors.extend(descriptors) return {"overall_code": "OK", "statuses": []} - + parallel_request_handler.should_rate_limit = mock_should_rate_limit - + # Test 1: No failures - rate limits should NOT be enforced (rpm_limit should be None) async def mock_check_no_failures(*args, **kwargs): return False - + parallel_request_handler._check_model_has_recent_failures = mock_check_no_failures - + await parallel_request_handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=local_cache, data={"model": model}, call_type="", ) - + # Find the API key descriptor api_key_descriptor = None for descriptor in captured_descriptors: if descriptor["key"] == "api_key": api_key_descriptor = descriptor break - + assert api_key_descriptor is not None, "API key descriptor should be present" assert ( api_key_descriptor["rate_limit"]["requests_per_unit"] is None ), "RPM limit should be None when dynamic mode and no failures" - + # Test 2: With failures - rate limits SHOULD be enforced (rpm_limit should be set) async def mock_check_with_failures(*args, **kwargs): return True - + parallel_request_handler._check_model_has_recent_failures = mock_check_with_failures captured_descriptors.clear() - + await parallel_request_handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=local_cache, data={"model": model}, call_type="", ) - + # Find the API key descriptor again api_key_descriptor = None for descriptor in captured_descriptors: if descriptor["key"] == "api_key": api_key_descriptor = descriptor break - + assert api_key_descriptor is not None, "API key descriptor should be present" assert ( api_key_descriptor["rate_limit"]["requests_per_unit"] == 2 @@ -1069,17 +1120,17 @@ async def test_dynamic_rate_limiting_v3(): async def test_async_increment_tokens_with_ttl_preservation(): """ Test TTL preservation functionality for token increment operations. - + This test verifies that: 1. Keys are created with proper TTL on first increment 2. TTL is preserved on subsequent increments (not reset) 3. Both TTL and non-TTL operations work correctly in the same call - + Environment variables required: - REDIS_HOST: Redis server hostname - REDIS_PORT: Redis server port - REDIS_PASSWORD: Redis password (optional) - + Test scenario: 1. First call: Create keys with TTL=60s and TTL=None 2. Wait 2 seconds @@ -1094,38 +1145,40 @@ async def test_async_increment_tokens_with_ttl_preservation(): # Skip test if Redis environment variables are not set redis_host = os.getenv("REDIS_HOST") - redis_port = os.getenv("REDIS_PORT") + redis_port = os.getenv("REDIS_PORT") redis_password = os.getenv("REDIS_PASSWORD") - + if not redis_host or not redis_port: pytest.skip("Redis environment variables (REDIS_HOST, REDIS_PORT) not set") - + # Setup Redis cache redis_cache = RedisCache( host=redis_host, port=int(redis_port), password=redis_password, ) - + local_cache = DualCache(redis_cache=redis_cache) parallel_request_handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Verify Redis connection is working try: await redis_cache.ping() except Exception as e: pytest.skip(f"Redis connection failed: {str(e)}") - + # Verify the TTL preservation script is registered if parallel_request_handler.token_increment_script is None: - pytest.skip("Token increment script not available - Redis Lua scripting may not be supported") - + pytest.skip( + "Token increment script not available - Redis Lua scripting may not be supported" + ) + # Test keys - use hash tags to ensure they map to same Redis cluster slot test_key_with_ttl = "{test_ttl}:with_ttl" test_key_without_ttl = "{test_ttl}:without_ttl" - + try: # Clean up any existing test keys try: @@ -1134,88 +1187,108 @@ async def test_async_increment_tokens_with_ttl_preservation(): except Exception: # Keys might not exist, ignore cleanup errors pass - + # First increment: Create operations with mixed TTL scenarios pipeline_operations_first = [ RedisPipelineIncrementOperation( - key=test_key_with_ttl, - increment_value=10.0, - ttl=60 + key=test_key_with_ttl, increment_value=10.0, ttl=60 ), RedisPipelineIncrementOperation( - key=test_key_without_ttl, - increment_value=5.0, - ttl=None # No TTL - ) + key=test_key_without_ttl, increment_value=5.0, ttl=None # No TTL + ), ] - + # Execute first increment await parallel_request_handler.async_increment_tokens_with_ttl_preservation( pipeline_operations=pipeline_operations_first ) - + # Small delay to ensure Redis has processed the commands await asyncio.sleep(0.1) - + # Verify keys exist and check initial TTL ttl_after_first = await redis_cache.async_get_ttl(test_key_with_ttl) - value_after_first_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl) - value_after_first_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl) - - assert value_after_first_with_ttl == 10.0, f"First increment should set value to 10.0, got {value_after_first_with_ttl}" - assert value_after_first_without_ttl == 5.0, "First increment should set value to 5.0" - assert ttl_after_first is not None and ttl_after_first > 0, "Key with TTL should have positive TTL after first increment" + value_after_first_with_ttl = await redis_cache.async_get_cache( + test_key_with_ttl + ) + value_after_first_without_ttl = await redis_cache.async_get_cache( + test_key_without_ttl + ) + + assert ( + value_after_first_with_ttl == 10.0 + ), f"First increment should set value to 10.0, got {value_after_first_with_ttl}" + assert ( + value_after_first_without_ttl == 5.0 + ), "First increment should set value to 5.0" + assert ( + ttl_after_first is not None and ttl_after_first > 0 + ), "Key with TTL should have positive TTL after first increment" assert ttl_after_first <= 60, "TTL should not exceed the set value" - + # Check TTL for key without TTL (should be None, meaning no expiry) ttl_no_ttl_key = await redis_cache.async_get_ttl(test_key_without_ttl) - assert ttl_no_ttl_key is None, "Key without TTL should have no expiry (None from async_get_ttl)" - + assert ( + ttl_no_ttl_key is None + ), "Key without TTL should have no expiry (None from async_get_ttl)" + # Wait a moment to ensure TTL decreases await asyncio.sleep(2) - + # Second increment: Same operations to test TTL preservation pipeline_operations_second = [ RedisPipelineIncrementOperation( - key=test_key_with_ttl, - increment_value=15.0, - ttl=60 # Same TTL value + key=test_key_with_ttl, increment_value=15.0, ttl=60 # Same TTL value ), RedisPipelineIncrementOperation( - key=test_key_without_ttl, - increment_value=7.0, - ttl=None # No TTL - ) + key=test_key_without_ttl, increment_value=7.0, ttl=None # No TTL + ), ] - + # Execute second increment await parallel_request_handler.async_increment_tokens_with_ttl_preservation( pipeline_operations=pipeline_operations_second ) - + # Small delay to ensure Redis has processed the commands await asyncio.sleep(0.1) - + # Verify TTL preservation and value updates ttl_after_second = await redis_cache.async_get_ttl(test_key_with_ttl) - value_after_second_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl) - value_after_second_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl) - - assert value_after_second_with_ttl == 25.0, "Second increment should update value to 25.0" - assert value_after_second_without_ttl == 12.0, "Second increment should update value to 12.0" - + value_after_second_with_ttl = await redis_cache.async_get_cache( + test_key_with_ttl + ) + value_after_second_without_ttl = await redis_cache.async_get_cache( + test_key_without_ttl + ) + + assert ( + value_after_second_with_ttl == 25.0 + ), "Second increment should update value to 25.0" + assert ( + value_after_second_without_ttl == 12.0 + ), "Second increment should update value to 12.0" + # Critical test: TTL should be preserved (not reset to 60) assert ttl_after_second is not None, "TTL should still exist" - assert ttl_after_second < ttl_after_first, "TTL should have decreased (not been reset)" + assert ( + ttl_after_second < ttl_after_first + ), "TTL should have decreased (not been reset)" assert ttl_after_second > 0, "TTL should still be positive" - + # TTL should not be close to the original 60 seconds (proving it wasn't reset) - assert ttl_after_second < 59, "TTL should be significantly less than original, proving preservation" - + assert ( + ttl_after_second < 59 + ), "TTL should be significantly less than original, proving preservation" + # Key without TTL should still have no expiry - ttl_no_ttl_key_after_second = await redis_cache.async_get_ttl(test_key_without_ttl) - assert ttl_no_ttl_key_after_second is None, "Key without TTL should still have no expiry" - + ttl_no_ttl_key_after_second = await redis_cache.async_get_ttl( + test_key_without_ttl + ) + assert ( + ttl_no_ttl_key_after_second is None + ), "Key without TTL should still have no expiry" + finally: # Clean up test keys try: @@ -1224,7 +1297,7 @@ async def test_async_increment_tokens_with_ttl_preservation(): except Exception: # Ignore cleanup errors pass - + # Properly close Redis connections to prevent warnings try: await redis_cache.disconnect() @@ -1239,115 +1312,125 @@ async def test_async_increment_tokens_fallback_behavior(): Test fallback behavior when Lua script is not available. """ from litellm.types.caching import RedisPipelineIncrementOperation - + local_cache = DualCache() parallel_request_handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Mock the token_increment_script to None to simulate unavailable script parallel_request_handler.token_increment_script = None - + # Mock the fallback method fallback_called = False - original_method = parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline - + original_method = ( + parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline + ) + async def mock_fallback(*args, **kwargs): nonlocal fallback_called fallback_called = True return await original_method(*args, **kwargs) - - parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_fallback - + + parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_fallback + ) + # Test operations pipeline_operations = [ RedisPipelineIncrementOperation( - key="test_fallback_key", - increment_value=10.0, - ttl=60 + key="test_fallback_key", increment_value=10.0, ttl=60 ) ] - + # Execute increment await parallel_request_handler.async_increment_tokens_with_ttl_preservation( pipeline_operations=pipeline_operations ) - + # Verify fallback was called - assert fallback_called, "Fallback method should be called when Lua script is not available" + assert ( + fallback_called + ), "Fallback method should be called when Lua script is not available" # Redis Cluster Compatibility Tests def test_group_keys_by_hash_tag_regular_redis(): """ Test that keys are correctly grouped for regular Redis (non-cluster). - + For regular Redis, all keys should be grouped together under a single group. """ local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Test keys with different hash tags test_keys = [ "{api_key:sk-123}:window", - "{api_key:sk-123}:requests", + "{api_key:sk-123}:requests", "{api_key:sk-123}:tokens", "{user:user-456}:window", "{user:user-456}:requests", "{team:team-789}:window", "{team:team-789}:tokens", - "no_hash_tag_key" + "no_hash_tag_key", ] - + # Group the keys (should be single group for regular Redis) groups = handler._group_keys_by_hash_tag(test_keys) - + # Verify all keys are in single group for regular Redis assert len(groups) == 1, f"Expected 1 group for regular Redis, got {len(groups)}" assert "all_keys" in groups, "Expected 'all_keys' group for regular Redis" - assert set(groups["all_keys"]) == set(test_keys), "All keys should be in single group" + assert set(groups["all_keys"]) == set( + test_keys + ), "All keys should be in single group" def test_group_keys_by_hash_tag_redis_cluster(): """ Test that keys are correctly grouped by Redis cluster slots when using Redis cluster. - + This ensures that keys are grouped by their slot number for cluster compatibility. """ from unittest.mock import patch - + local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Mock _is_redis_cluster to return True - with patch.object(handler, '_is_redis_cluster', return_value=True): + with patch.object(handler, "_is_redis_cluster", return_value=True): # Test keys with different hash tags test_keys = [ "{api_key:sk-123}:window", - "{api_key:sk-123}:requests", + "{api_key:sk-123}:requests", "{user:user-456}:window", "{user:user-456}:requests", ] - + # Group the keys (should be grouped by slot for Redis cluster) groups = handler._group_keys_by_hash_tag(test_keys) - + # Verify keys are grouped by slot assert len(groups) >= 1, "Should have at least 1 slot group" - + # All group keys should start with "slot_" for group_key in groups.keys(): - assert group_key.startswith("slot_"), f"Group key {group_key} should start with 'slot_'" - + assert group_key.startswith( + "slot_" + ), f"Group key {group_key} should start with 'slot_'" + # Verify all original keys are present across groups all_grouped_keys = [] for group_keys in groups.values(): all_grouped_keys.extend(group_keys) - assert set(all_grouped_keys) == set(test_keys), "All keys should be present in groups" + assert set(all_grouped_keys) == set( + test_keys + ), "All keys should be present in groups" def test_keyslot_for_redis_cluster(): @@ -1358,16 +1441,16 @@ def test_keyslot_for_redis_cluster(): handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Test basic key slot1 = handler.keyslot_for_redis_cluster("user:1000") assert 0 <= slot1 < 16384, "Slot should be in valid range" - + # Test key with hash tag slot2 = handler.keyslot_for_redis_cluster("foo{bar}baz") slot3 = handler.keyslot_for_redis_cluster("{bar}") assert slot2 == slot3, "Keys with same hash tag should have same slot" - + # Test keys with same hash tag should have same slot slot4 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:requests") slot5 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:window") @@ -1379,67 +1462,70 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): """ Test that the Redis batch rate limiter script execution handles cluster compatibility by grouping keys and falling back gracefully on errors. - + This simulates the Redis cluster error scenario and verifies fallback behavior. """ from unittest.mock import AsyncMock, patch - + local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Mock _is_redis_cluster to return True for this test - with patch.object(handler, '_is_redis_cluster', return_value=True): + with patch.object(handler, "_is_redis_cluster", return_value=True): # Mock script that simulates Redis cluster slot conflict mock_script = AsyncMock() mock_script.side_effect = [ - Exception("EVALSHA - all keys must map to the same key slot"), # First group fails - [1234, 1, 1234, 2] # Second group succeeds + Exception( + "EVALSHA - all keys must map to the same key slot" + ), # First group fails + [1234, 1, 1234, 2], # Second group succeeds ] handler.batch_rate_limiter_script = mock_script - + # Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter) handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1]) - + # Test keys from different hash tags (would fail in cluster without grouping) test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", - "{user:user-456}:window", - "{user:user-456}:requests" + "{user:user-456}:window", + "{user:user-456}:requests", ] - + # Execute the method results = await handler._execute_redis_batch_rate_limiter_script( - keys_to_fetch=test_keys, - now_int=1234 + keys_to_fetch=test_keys, now_int=1234 ) - + # Verify results: 2 from fallback + 4 from successful script = 6 total assert len(results) == 6, f"Expected 6 results, got {len(results)}" - + # Verify script was called twice (once per slot group) assert mock_script.call_count == 2 - + # Verify fallback was called for the failed group handler.in_memory_cache_sliding_window.assert_called_once() - + # Verify the calls were made with grouped keys call_args_list = mock_script.call_args_list - + # Both calls should have keys, but we can't predict exact grouping without knowing slots # Just verify that keys were grouped and calls were made assert len(call_args_list) == 2, "Should have made 2 script calls" - + # Verify all keys were processed all_processed_keys = [] for call_args in call_args_list: - all_processed_keys.extend(call_args[1]['keys']) - + all_processed_keys.extend(call_args[1]["keys"]) + # Should have processed all keys (some might be duplicated due to fallback) unique_processed_keys = set(all_processed_keys) - assert len(unique_processed_keys) >= 2, "Should have processed at least some keys" + assert ( + len(unique_processed_keys) >= 2 + ), "Should have processed at least some keys" @pytest.mark.asyncio @@ -1485,23 +1571,23 @@ async def test_multiple_rate_limits_per_descriptor(): "current_limit": 2, "limit_remaining": 1, "rate_limit_type": "requests", - "descriptor_key": "api_key" + "descriptor_key": "api_key", }, { "code": "OK", "current_limit": 10, "limit_remaining": 8, "rate_limit_type": "tokens", - "descriptor_key": "api_key" + "descriptor_key": "api_key", }, { "code": "OVER_LIMIT", "current_limit": 1, "limit_remaining": -1, "rate_limit_type": "max_parallel_requests", - "descriptor_key": "api_key" - } - ] + "descriptor_key": "api_key", + }, + ], } parallel_request_handler.should_rate_limit = mock_should_rate_limit @@ -1560,9 +1646,9 @@ async def test_missing_descriptor_fallback(): "current_limit": 2, "limit_remaining": -1, "rate_limit_type": "requests", - "descriptor_key": "nonexistent_key" # This won't match any descriptor + "descriptor_key": "nonexistent_key", # This won't match any descriptor } - ] + ], } parallel_request_handler.should_rate_limit = mock_should_rate_limit @@ -1597,14 +1683,17 @@ async def test_get_rate_limit_type_default_is_total(monkeypatch): # Mock general_settings to return empty dict (no token_rate_limit_type set) import litellm.proxy.proxy_server as proxy_server - original_settings = getattr(proxy_server, 'general_settings', {}) - monkeypatch.setattr(proxy_server, 'general_settings', {}) + + original_settings = getattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(proxy_server, "general_settings", {}) try: result = parallel_request_handler.get_rate_limit_type() - assert result == "total", f"Default rate limit type should be 'total', got '{result}'" + assert ( + result == "total" + ), f"Default rate limit type should be 'total', got '{result}'" finally: - monkeypatch.setattr(proxy_server, 'general_settings', original_settings) + monkeypatch.setattr(proxy_server, "general_settings", original_settings) @pytest.mark.asyncio @@ -1619,14 +1708,19 @@ async def test_get_rate_limit_type_invalid_falls_back_to_total(monkeypatch): # Mock general_settings to return an invalid token_rate_limit_type import litellm.proxy.proxy_server as proxy_server - original_settings = getattr(proxy_server, 'general_settings', {}) - monkeypatch.setattr(proxy_server, 'general_settings', {'token_rate_limit_type': 'invalid_type'}) + + original_settings = getattr(proxy_server, "general_settings", {}) + monkeypatch.setattr( + proxy_server, "general_settings", {"token_rate_limit_type": "invalid_type"} + ) try: result = parallel_request_handler.get_rate_limit_type() - assert result == "total", f"Invalid rate limit type should fall back to 'total', got '{result}'" + assert ( + result == "total" + ), f"Invalid rate limit type should fall back to 'total', got '{result}'" finally: - monkeypatch.setattr(proxy_server, 'general_settings', original_settings) + monkeypatch.setattr(proxy_server, "general_settings", original_settings) @pytest.mark.parametrize( @@ -1638,7 +1732,9 @@ async def test_get_rate_limit_type_invalid_falls_back_to_total(monkeypatch): ], ) @pytest.mark.asyncio -async def test_async_log_success_event_with_dict_usage(monkeypatch, token_rate_limit_type, expected_field): +async def test_async_log_success_event_with_dict_usage( + monkeypatch, token_rate_limit_type, expected_field +): """ Test that async_log_success_event correctly handles usage as a dict (Responses API format). @@ -1664,13 +1760,13 @@ async def test_async_log_success_event_with_dict_usage(monkeypatch, token_rate_l # Create a mock response object with usage as a dict (Responses API format) from litellm.types.utils import BaseLiteLLMOpenAIResponseObject - + # Use spec to make isinstance checks work correctly with MagicMock mock_response = MagicMock(spec=BaseLiteLLMOpenAIResponseObject) mock_response.usage = { "prompt_tokens": 25, "completion_tokens": 35, - "total_tokens": 60 + "total_tokens": 60, } # Create mock kwargs for the success event @@ -1760,7 +1856,10 @@ async def test_async_log_success_event_with_dict_usage_missing_fields(monkeypatc # total_tokens is missing } from litellm.types.utils import BaseLiteLLMOpenAIResponseObject - mock_response.__class__ = type('MockResponse', (BaseLiteLLMOpenAIResponseObject,), {}) + + mock_response.__class__ = type( + "MockResponse", (BaseLiteLLMOpenAIResponseObject,), {} + ) # Create mock kwargs for the success event mock_kwargs = { @@ -1805,7 +1904,9 @@ async def test_async_log_success_event_with_dict_usage_missing_fields(monkeypatc assert tpm_operation is not None, "Should have a TPM increment operation" # Should default to 0 when field is missing - assert tpm_operation["increment_value"] == 0, "Should default to 0 when completion_tokens is missing" + assert ( + tpm_operation["increment_value"] == 0 + ), "Should default to 0 when completion_tokens is missing" @pytest.mark.asyncio @@ -1813,68 +1914,154 @@ async def test_execute_token_increment_script_cluster_compatibility(): """ Test that token increment script execution handles Redis cluster compatibility by grouping operations by slot. - + This ensures token increments work correctly in cluster environments. """ from typing import List from unittest.mock import AsyncMock, patch from litellm.types.caching import RedisPipelineIncrementOperation - + local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - + # Mock _is_redis_cluster to return True for this test - with patch.object(handler, '_is_redis_cluster', return_value=True): + with patch.object(handler, "_is_redis_cluster", return_value=True): # Mock script mock_script = AsyncMock() handler.token_increment_script = mock_script - + # Create pipeline operations with different hash tags pipeline_operations: List[RedisPipelineIncrementOperation] = [ + {"key": "{api_key:sk-123}:tokens", "increment_value": 100, "ttl": 60}, { - "key": "{api_key:sk-123}:tokens", - "increment_value": 100, - "ttl": 60 - }, - { - "key": "{api_key:sk-123}:max_parallel_requests", + "key": "{api_key:sk-123}:max_parallel_requests", "increment_value": -1, - "ttl": 60 + "ttl": 60, }, - { - "key": "{user:user-456}:tokens", - "increment_value": 50, - "ttl": 60 - } + {"key": "{user:user-456}:tokens", "increment_value": 50, "ttl": 60}, ] - + # Execute the method await handler._execute_token_increment_script(pipeline_operations) - + # Verify script was called (at least once, possibly more depending on slot grouping) assert mock_script.call_count >= 1, "Script should be called at least once" - + call_args_list = mock_script.call_args_list - + # Verify all operations were processed all_processed_keys = [] for call_args in call_args_list: - all_processed_keys.extend(call_args[1]['keys']) - + all_processed_keys.extend(call_args[1]["keys"]) + # Should have processed all 3 keys expected_keys = { "{api_key:sk-123}:tokens", "{api_key:sk-123}:max_parallel_requests", - "{user:user-456}:tokens" + "{user:user-456}:tokens", } - assert set(all_processed_keys) == expected_keys, "All operation keys should be processed" - + assert ( + set(all_processed_keys) == expected_keys + ), "All operation keys should be processed" + # Verify args structure is correct for each call for call_args in call_args_list: - keys = call_args[1]['keys'] - args = call_args[1]['args'] + keys = call_args[1]["keys"] + args = call_args[1]["args"] # Each key should have 2 args (increment_value, ttl) - assert len(args) == len(keys) * 2, f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys" + assert ( + len(args) == len(keys) * 2 + ), f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys" + + +class TestGetTotalTokensFromUsageCacheExclusion: + """ + Tests for _get_total_tokens_from_usage cache token exclusion. + + Issue: AWS Bedrock and similar providers exclude cache tokens from TPM calculation, + but LiteLLM was including them, causing up to 10x difference in rate limiting. + """ + + @pytest.fixture + def handler(self): + """Create a handler instance for testing.""" + local_cache = DualCache() + return _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache), + ) + + def test_excludes_cached_tokens_from_total(self, handler): + """Cached tokens should be excluded from total token count.""" + from litellm.types.utils import PromptTokensDetailsWrapper + + usage = Usage( + prompt_tokens=1000, + completion_tokens=500, + total_tokens=1500, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=800), + ) + + # Total should be 1500 - 800 = 700 + result = handler._get_total_tokens_from_usage(usage, "total") + assert result == 700, f"Expected 700 (1500 - 800 cached), got {result}" + + def test_excludes_cached_tokens_from_input(self, handler): + """Cached tokens should be excluded from input token count.""" + from litellm.types.utils import PromptTokensDetailsWrapper + + usage = Usage( + prompt_tokens=1000, + completion_tokens=500, + total_tokens=1500, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=800), + ) + + # Input should be 1000 - 800 = 200 + result = handler._get_total_tokens_from_usage(usage, "input") + assert result == 200, f"Expected 200 (1000 - 800 cached), got {result}" + + def test_does_not_exclude_cached_tokens_from_output(self, handler): + """Cached tokens should NOT affect output token count.""" + from litellm.types.utils import PromptTokensDetailsWrapper + + usage = Usage( + prompt_tokens=1000, + completion_tokens=500, + total_tokens=1500, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=800), + ) + + # Output tokens should be unchanged + result = handler._get_total_tokens_from_usage(usage, "output") + assert result == 500, f"Expected 500 (no change for output), got {result}" + + def test_handles_no_cached_tokens(self, handler): + """Should work correctly when no cached tokens present.""" + usage = Usage( + prompt_tokens=1000, + completion_tokens=500, + total_tokens=1500, + ) + + result = handler._get_total_tokens_from_usage(usage, "total") + assert result == 1500, f"Expected 1500 (no cache), got {result}" + + def test_handles_dict_usage_with_cached_tokens(self, handler): + """Should handle dict usage format (Responses API) with cached tokens.""" + usage = { + "prompt_tokens": 1000, + "completion_tokens": 500, + "total_tokens": 1500, + "prompt_tokens_details": {"cached_tokens": 600}, + } + + result = handler._get_total_tokens_from_usage(usage, "total") + assert result == 900, f"Expected 900 (1500 - 600 cached), got {result}" + + def test_handles_none_usage(self, handler): + """Should handle None usage gracefully.""" + result = handler._get_total_tokens_from_usage(None, "total") + assert result == 0, f"Expected 0 for None usage, got {result}" diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index f2bae2cb14a..bc223d15d5f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1,20 +1,20 @@ import json import os import sys -from litellm._uuid import uuid +import types from datetime import datetime, timedelta -from typing import List +from typing import List, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import FastAPI from fastapi.testclient import TestClient +from litellm._uuid import uuid sys.path.insert( 0, os.path.abspath("../../../..") ) # Adds the parent directory to the system path -from typing import Optional - from litellm.proxy._types import ( LiteLLM_MCPServerTable, LitellmUserRoles, @@ -118,6 +118,22 @@ def setup_mock_prisma_client( return mock_prisma_client +def create_mcp_router_test_client() -> TestClient: + from litellm.proxy.management_endpoints.mcp_management_endpoints import router + + app = FastAPI() + app.include_router(router) + return TestClient(app) + + +def patch_proxy_general_settings(settings: dict): + fake_proxy_server_module = types.SimpleNamespace(general_settings=settings) + return patch.dict( + sys.modules, + {"litellm.proxy.proxy_server": fake_proxy_server_module}, + ) + + class TestListMCPServers: """Test suite for list MCP servers functionality""" @@ -1082,6 +1098,55 @@ class TestHealthCheckServers: assert result[1]["server_id"] == "server-2" assert result[1]["status"] == "unhealthy" + +class TestMCPRegistryEndpoint: + def test_registry_returns_404_when_flag_missing(self): + client = create_mcp_router_test_client() + + with patch_proxy_general_settings({}): + response = client.get("/v1/mcp/registry.json") + + assert response.status_code == 404 + + def test_registry_returns_404_when_flag_false(self): + client = create_mcp_router_test_client() + + with patch_proxy_general_settings({"enable_mcp_registry": False}): + response = client.get("/v1/mcp/registry.json") + + assert response.status_code == 404 + + def test_registry_returns_entries_when_enabled(self): + client = create_mcp_router_test_client() + + mock_server = generate_mock_mcp_server_config_record( + server_id="server-123", + name="zapier", + url="https://zapier.example.com/mcp", + transport="http", + ) + + mock_manager = MagicMock() + mock_manager.get_registry.return_value = {mock_server.server_id: mock_server} + + with patch_proxy_general_settings({"enable_mcp_registry": True}), patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ): + response = client.get("/v1/mcp/registry.json") + + assert response.status_code == 200 + data = response.json() + assert len(data["servers"]) == 2 # built-in + custom server + + builtin_entry = data["servers"][0]["server"] + assert builtin_entry["name"] == "litellm-mcp-server" + assert builtin_entry["remotes"][0]["url"].endswith("/mcp") + + custom_entry = data["servers"][1]["server"] + assert custom_entry["name"] == "zapier" + assert custom_entry["remotes"][0]["url"].endswith("/zapier/mcp") + @pytest.mark.asyncio async def test_health_check_specific_servers(self): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_router_settings_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_router_settings_endpoints.py new file mode 100644 index 00000000000..1f5473e75d4 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_router_settings_endpoints.py @@ -0,0 +1,71 @@ +""" +Tests for router settings management endpoints. + +Tests the GET endpoints for router settings and router fields. +""" +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../../..") +) + +from litellm.proxy.proxy_server import app + +client = TestClient(app) + + +class TestRouterSettingsEndpoints: + """Test suite for router settings endpoints""" + + @pytest.mark.asyncio + async def test_get_router_fields_success(self): + """ + Test GET /router/fields endpoint successfully returns field definitions without values. + """ + # Make request to router fields endpoint + response = client.get( + "/router/fields", + headers={"Authorization": "Bearer sk-1234"} + ) + + # Verify response + assert response.status_code == 200 + + response_data = response.json() + + # Verify response structure + assert "fields" in response_data + assert "routing_strategy_descriptions" in response_data + + # Verify fields is a list + assert isinstance(response_data["fields"], list) + assert len(response_data["fields"]) > 0 + + # Verify each field has required properties and field_value is None + for field in response_data["fields"]: + assert "field_name" in field + assert "field_type" in field + assert "field_description" in field + assert "field_default" in field + assert "ui_field_name" in field + assert "field_value" in field + assert field["field_value"] is None # Ensure field_value is None + + # Verify routing_strategy_descriptions is a dict + assert isinstance(response_data["routing_strategy_descriptions"], dict) + assert len(response_data["routing_strategy_descriptions"]) > 0 + + # Verify routing_strategy field has options populated + routing_strategy_field = next( + (f for f in response_data["fields"] if f["field_name"] == "routing_strategy"), + None + ) + assert routing_strategy_field is not None + assert "options" in routing_strategy_field + assert isinstance(routing_strategy_field["options"], list) + assert len(routing_strategy_field["options"]) > 0 diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 56bba39e6c3..57ec019e420 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -1242,7 +1242,7 @@ class TestSpendLogsPayload: "model": "claude-3-7-sonnet-20250219", "user": "", "team_id": "", - "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', + "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', "cache_key": "Cache OFF", "spend": 0.01383, "total_tokens": 2598, @@ -1334,7 +1334,7 @@ class TestSpendLogsPayload: "model": "claude-3-7-sonnet-20250219", "user": "", "team_id": "", - "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', + "metadata": '{"applied_guardrails": [], "batch_models": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "guardrail_information": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-3-7-sonnet-20250219", "model_map_value": {"key": "claude-3-7-sonnet-20250219", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}', "cache_key": "Cache OFF", "spend": 0.01383, "total_tokens": 2598, diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index b5d44385698..3d1e9aece41 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -3,7 +3,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import Request, status -from fastapi.responses import StreamingResponse +from fastapi.responses import JSONResponse, StreamingResponse import litellm from litellm._uuid import uuid @@ -11,9 +11,10 @@ from litellm.integrations.opentelemetry import UserAPIKeyAuth from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ProxyConfig, + _extract_error_from_sse_chunk, _get_cost_breakdown_from_logging_obj, _parse_event_data_for_error, - create_streaming_response, + create_response, ) from litellm.proxy.utils import ProxyLogging @@ -75,6 +76,84 @@ class TestProxyBaseLLMRequestProcessing: pytest.fail("litellm_call_id is not a valid UUID") assert data_passed["litellm_call_id"] == returned_data["litellm_call_id"] + @pytest.mark.asyncio + async def test_should_apply_hierarchical_router_settings_to_user_config( + self, monkeypatch + ): + processing_obj = ProxyBaseLLMRequestProcessing(data={}) + mock_request = MagicMock(spec=Request) + mock_request.headers = {} + + async def mock_add_litellm_data_to_request(*args, **kwargs): + return {} + + async def mock_common_processing_pre_call_logic( + user_api_key_dict, data, call_type + ): + data_copy = copy.deepcopy(data) + return data_copy + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.pre_call_hook = AsyncMock( + side_effect=mock_common_processing_pre_call_logic + ) + monkeypatch.setattr( + litellm.proxy.common_request_processing, + "add_litellm_data_to_request", + mock_add_litellm_data_to_request, + ) + + mock_general_settings = {} + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_proxy_config = MagicMock(spec=ProxyConfig) + + mock_router_settings = { + "routing_strategy": "least-busy", + "timeout": 30.0, + "num_retries": 3, + } + mock_proxy_config._get_hierarchical_router_settings = AsyncMock( + return_value=mock_router_settings + ) + + mock_model_list = [ + {"model_name": "gpt-3.5-turbo", "litellm_params": {"model": "gpt-3.5-turbo"}}, + {"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}}, + ] + mock_llm_router = MagicMock() + mock_llm_router.get_model_list = MagicMock(return_value=mock_model_list) + + mock_prisma_client = MagicMock() + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", + mock_prisma_client, + ) + + route_type = "acompletion" + + returned_data, logging_obj = await processing_obj.common_processing_pre_call_logic( + request=mock_request, + general_settings=mock_general_settings, + user_api_key_dict=mock_user_api_key_dict, + proxy_logging_obj=mock_proxy_logging_obj, + proxy_config=mock_proxy_config, + route_type=route_type, + llm_router=mock_llm_router, + ) + + mock_proxy_config._get_hierarchical_router_settings.assert_called_once_with( + user_api_key_dict=mock_user_api_key_dict, + prisma_client=mock_prisma_client, + ) + mock_llm_router.get_model_list.assert_called_once() + + assert "user_config" in returned_data + user_config = returned_data["user_config"] + assert user_config["model_list"] == mock_model_list + assert user_config["routing_strategy"] == "least-busy" + assert user_config["timeout"] == 30.0 + assert user_config["num_retries"] == 3 + @pytest.mark.asyncio async def test_stream_timeout_header_processing(self): """ @@ -602,21 +681,27 @@ class TestCommonRequestProcessingHelpers: assert await _parse_event_data_for_error(event_line) == expected_code async def test_create_streaming_response_first_chunk_is_error(self): + """ + Test that when the first chunk is an error, a JSON error response is returned + instead of an SSE streaming response + """ async def mock_generator(): yield 'data: {"error": {"code": 403, "message": "forbidden"}}\n\n' yield 'data: {"content": "more data"}\n\n' yield "data: [DONE]\n\n" - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) + # Should return JSONResponse instead of StreamingResponse + assert isinstance(response, JSONResponse) assert response.status_code == status.HTTP_403_FORBIDDEN - content = await self.consume_stream(response) - assert content == [ - 'data: {"error": {"code": 403, "message": "forbidden"}}\n\n', - 'data: {"content": "more data"}\n\n', - "data: [DONE]\n\n", - ] + # Verify the response is in standard JSON error format + import json + body = json.loads(response.body.decode()) + assert "error" in body + assert body["error"]["code"] == 403 + assert body["error"]["message"] == "forbidden" async def test_create_streaming_response_first_chunk_not_error(self): async def mock_generator(): @@ -624,7 +709,7 @@ class TestCommonRequestProcessingHelpers: yield 'data: {"content": "second part"}\n\n' yield "data: [DONE]\n\n" - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) assert response.status_code == status.HTTP_200_OK @@ -641,7 +726,7 @@ class TestCommonRequestProcessingHelpers: yield # Implicitly raises StopAsyncIteration - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) assert response.status_code == status.HTTP_200_OK @@ -654,7 +739,7 @@ class TestCommonRequestProcessingHelpers: mock_gen = AsyncMock() mock_gen.__anext__.side_effect = StopAsyncIteration - response = await create_streaming_response(mock_gen, "text/event-stream", {}) + response = await create_response(mock_gen, "text/event-stream", {}) assert response.status_code == status.HTTP_200_OK content = await self.consume_stream(response) assert content == [] @@ -665,7 +750,7 @@ class TestCommonRequestProcessingHelpers: mock_gen = AsyncMock() mock_gen.__anext__.side_effect = ValueError("Test error from generator") - response = await create_streaming_response(mock_gen, "text/event-stream", {}) + response = await create_response(mock_gen, "text/event-stream", {}) assert response.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR content = await self.consume_stream(response) expected_error_data = { @@ -682,19 +767,24 @@ class TestCommonRequestProcessingHelpers: assert content[1] == "data: [DONE]\n\n" async def test_create_streaming_response_first_chunk_error_string_code(self): + """ + Test that when the first chunk contains a string error code, a JSON error response is returned + """ async def mock_generator(): yield 'data: {"error": {"code": "429", "message": "too many requests"}}\n\n' yield "data: [DONE]\n\n" - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) + assert isinstance(response, JSONResponse) assert response.status_code == status.HTTP_429_TOO_MANY_REQUESTS - content = await self.consume_stream(response) - assert content == [ - 'data: {"error": {"code": "429", "message": "too many requests"}}\n\n', - "data: [DONE]\n\n", - ] + # Verify the response is in standard JSON error format + import json + body = json.loads(response.body.decode()) + assert "error" in body + assert body["error"]["code"] == "429" + assert body["error"]["message"] == "too many requests" async def test_create_streaming_response_custom_headers(self): async def mock_generator(): @@ -702,7 +792,7 @@ class TestCommonRequestProcessingHelpers: yield "data: [DONE]\n\n" custom_headers = {"X-Custom-Header": "TestValue"} - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", custom_headers ) assert response.headers["x-custom-header"] == "TestValue" @@ -712,7 +802,7 @@ class TestCommonRequestProcessingHelpers: yield 'data: {"content": "data"}\n\n' yield "data: [DONE]\n\n" - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {}, @@ -729,7 +819,7 @@ class TestCommonRequestProcessingHelpers: async def mock_generator(): yield "data: [DONE]\n\n" - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) assert response.status_code == status.HTTP_200_OK # Default status @@ -742,7 +832,7 @@ class TestCommonRequestProcessingHelpers: yield 'data: {"content": "actual data"}\n\n' yield "data: [DONE]\n\n" - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) assert response.status_code == status.HTTP_200_OK # Default status @@ -773,7 +863,7 @@ class TestCommonRequestProcessingHelpers: # Patch the tracer in the common_request_processing module with patch("litellm.proxy.common_request_processing.tracer", mock_tracer): - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) @@ -810,7 +900,10 @@ class TestCommonRequestProcessingHelpers: ), f"Call {i} should have operation name 'streaming.chunk.yield', got {args[0]}" async def test_create_streaming_response_dd_trace_with_error_chunk(self): - """Test that dd trace is applied even when the first chunk contains an error""" + """ + Test that when the first chunk contains an error, JSONResponse is returned + and tracing is not triggered (since it's not a streaming response) + """ from unittest.mock import patch # Create a mock tracer @@ -827,28 +920,107 @@ class TestCommonRequestProcessingHelpers: # Patch the tracer in the common_request_processing module with patch("litellm.proxy.common_request_processing.tracer", mock_tracer): - response = await create_streaming_response( + response = await create_response( mock_generator(), "text/event-stream", {} ) - # Even with error, status should be set to error code but tracing should still work + # Should return JSONResponse instead of StreamingResponse + assert isinstance(response, JSONResponse) assert response.status_code == 400 - # Consume the stream to trigger the tracer calls - content = await self.consume_stream(response) + # Verify the response is in standard JSON error format + import json + body = json.loads(response.body.decode()) + assert "error" in body + assert body["error"]["code"] == 400 + assert body["error"]["message"] == "bad request" - # Verify all chunks are present - assert len(content) == 3 + # Since JSONResponse is returned instead of StreamingResponse, streaming tracing should not be triggered + # tracer.trace should not be called + assert mock_tracer.trace.call_count == 0 - # Verify that tracer.trace was called for each chunk - assert mock_tracer.trace.call_count == 3 - # Verify that each call was made with the correct operation name - actual_calls = mock_tracer.trace.call_args_list - assert len(actual_calls) == 3 +class TestExtractErrorFromSSEChunk: + """Tests for _extract_error_from_sse_chunk function""" + + def test_extract_error_from_sse_chunk_with_valid_error(self): + """Test extracting error information from a standard SSE chunk""" + chunk = 'data: {"error": {"code": 403, "message": "forbidden", "type": "auth_error", "param": "api_key"}}\n\n' + error = _extract_error_from_sse_chunk(chunk) + + assert error["code"] == 403 + assert error["message"] == "forbidden" + assert error["type"] == "auth_error" + assert error["param"] == "api_key" + + def test_extract_error_from_sse_chunk_with_string_code(self): + """Test error code as string type""" + chunk = 'data: {"error": {"code": "429", "message": "too many requests"}}\n\n' + error = _extract_error_from_sse_chunk(chunk) + + assert error["code"] == "429" + assert error["message"] == "too many requests" + + def test_extract_error_from_sse_chunk_with_bytes(self): + """Test input as bytes type""" + chunk = b'data: {"error": {"code": 500, "message": "internal error"}}\n\n' + error = _extract_error_from_sse_chunk(chunk) + + assert error["code"] == 500 + assert error["message"] == "internal error" + + def test_extract_error_from_sse_chunk_with_done(self): + """Test [DONE] marker should return default error""" + chunk = "data: [DONE]\n\n" + error = _extract_error_from_sse_chunk(chunk) + + assert error["message"] == "Unknown error" + assert error["type"] == "internal_server_error" + assert error["code"] == "500" + assert error["param"] is None + + def test_extract_error_from_sse_chunk_without_error_field(self): + """Test missing error field should return default error""" + chunk = 'data: {"content": "some content"}\n\n' + error = _extract_error_from_sse_chunk(chunk) + + assert error["message"] == "Unknown error" + assert error["type"] == "internal_server_error" + assert error["code"] == "500" + + def test_extract_error_from_sse_chunk_with_invalid_json(self): + """Test invalid JSON should return default error""" + chunk = 'data: {invalid json}\n\n' + error = _extract_error_from_sse_chunk(chunk) + + assert error["message"] == "Unknown error" + assert error["type"] == "internal_server_error" + assert error["code"] == "500" + + def test_extract_error_from_sse_chunk_without_data_prefix(self): + """Test missing 'data:' prefix should return default error""" + chunk = '{"error": {"code": 400, "message": "bad request"}}\n\n' + error = _extract_error_from_sse_chunk(chunk) + + assert error["message"] == "Unknown error" + assert error["type"] == "internal_server_error" + assert error["code"] == "500" + + def test_extract_error_from_sse_chunk_with_empty_string(self): + """Test empty string should return default error""" + chunk = "" + error = _extract_error_from_sse_chunk(chunk) + + assert error["message"] == "Unknown error" + assert error["type"] == "internal_server_error" + assert error["code"] == "500" + + def test_extract_error_from_sse_chunk_with_minimal_error(self): + """Test minimal error object""" + chunk = 'data: {"error": {"message": "error occurred"}}\n\n' + error = _extract_error_from_sse_chunk(chunk) + + assert error["message"] == "error occurred" + # Other fields should be obtained from the original error object (if exists) + - for i, call in enumerate(actual_calls): - args, kwargs = call - assert ( - args[0] == "streaming.chunk.yield" - ), f"Call {i} should have operation name 'streaming.chunk.yield', got {args[0]}" diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 5c7ece04513..751a9033871 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3036,3 +3036,183 @@ def test_get_image_root_case_uses_current_dir(monkeypatch): # Verify FileResponse was called assert mock_file_response.called, "FileResponse should be called" + + +def test_get_config_normalizes_string_callbacks(monkeypatch): + """ + Test that /get/config/callbacks normalizes string callbacks to lists. + """ + from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth + + config_data = { + "litellm_settings": { + "success_callback": "langfuse", + "failure_callback": None, + "callbacks": ["prometheus", "datadog"], + }, + "general_settings": {}, + "environment_variables": {}, + } + + mock_router = MagicMock() + mock_router.get_settings.return_value = {} + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) + monkeypatch.setattr( + proxy_config, "get_config", AsyncMock(return_value=config_data) + ) + + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = lambda: MagicMock() + + client = TestClient(app) + try: + response = client.get("/get/config/callbacks") + finally: + app.dependency_overrides = original_overrides + + assert response.status_code == 200 + callbacks = response.json()["callbacks"] + + success_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "success"] + failure_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "failure"] + success_and_failure_callbacks = [ + cb["name"] for cb in callbacks if cb.get("type") == "success_and_failure" + ] + + assert "langfuse" in success_callbacks + assert len(failure_callbacks) == 0 + assert "prometheus" in success_and_failure_callbacks + assert "datadog" in success_and_failure_callbacks + + +def test_deep_merge_dicts_skips_none_and_empty_lists(monkeypatch): + """ + Test that _update_config_fields deep merge skips None values and empty lists. + """ + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + + current_config = { + "general_settings": { + "max_parallel_requests": 10, + "allowed_models": ["gpt-3.5-turbo", "gpt-4"], + "nested": { + "key1": "value1", + "key2": "value2", + }, + } + } + + db_param_value = { + "max_parallel_requests": None, + "allowed_models": [], + "new_key": "new_value", + "nested": { + "key1": "updated_value1", + "key3": "value3", + }, + } + + result = proxy_config._update_config_fields( + current_config, "general_settings", db_param_value + ) + + assert result["general_settings"]["max_parallel_requests"] == 10 + assert result["general_settings"]["allowed_models"] == ["gpt-3.5-turbo", "gpt-4"] + assert result["general_settings"]["new_key"] == "new_value" + assert result["general_settings"]["nested"]["key1"] == "updated_value1" + assert result["general_settings"]["nested"]["key2"] == "value2" + assert result["general_settings"]["nested"]["key3"] == "value3" + + +@pytest.mark.asyncio +async def test_get_hierarchical_router_settings(): + """ + Test _get_hierarchical_router_settings method's priority order: Key > Team > Global + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + + # Test Case 1: Returns None when prisma_client is None + result = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=None, + prisma_client=None, + ) + assert result is None + + # Test Case 2: Returns key-level router_settings when available (as dict) + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.router_settings = {"routing_strategy": "key-level", "timeout": 10} + mock_user_api_key_dict.team_id = None + + mock_prisma_client = MagicMock() + + result = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=mock_user_api_key_dict, + prisma_client=mock_prisma_client, + ) + assert result == {"routing_strategy": "key-level", "timeout": 10} + + # Test Case 3: Returns key-level router_settings when available (as YAML string) + mock_user_api_key_dict.router_settings = "routing_strategy: key-yaml\ntimeout: 20" + result = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=mock_user_api_key_dict, + prisma_client=mock_prisma_client, + ) + assert result == {"routing_strategy": "key-yaml", "timeout": 20} + + # Test Case 4: Falls back to team-level router_settings when key-level is not available + mock_user_api_key_dict.router_settings = None + mock_user_api_key_dict.team_id = "team-123" + + mock_team_obj = MagicMock() + mock_team_obj.router_settings = {"routing_strategy": "team-level", "timeout": 30} + + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_team_obj + ) + + result = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=mock_user_api_key_dict, + prisma_client=mock_prisma_client, + ) + assert result == {"routing_strategy": "team-level", "timeout": 30} + mock_prisma_client.db.litellm_teamtable.find_unique.assert_called_once_with( + where={"team_id": "team-123"} + ) + + # Test Case 5: Falls back to global router_settings when neither key nor team settings are available + mock_user_api_key_dict.router_settings = None + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + + mock_db_config = MagicMock() + mock_db_config.param_value = {"routing_strategy": "global-level", "timeout": 40} + + mock_prisma_client.db.litellm_config.find_first = AsyncMock( + return_value=mock_db_config + ) + + result = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=mock_user_api_key_dict, + prisma_client=mock_prisma_client, + ) + assert result == {"routing_strategy": "global-level", "timeout": 40} + mock_prisma_client.db.litellm_config.find_first.assert_called_once_with( + where={"param_name": "router_settings"} + ) + + # Test Case 6: Returns None when no settings are found + mock_user_api_key_dict.router_settings = None + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) + + result = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=mock_user_api_key_dict, + prisma_client=mock_prisma_client, + ) + assert result is None diff --git a/tests/test_litellm/responses/test_responses_utils.py b/tests/test_litellm/responses/test_responses_utils.py index 96ac2e2c345..09628cd4a76 100644 --- a/tests/test_litellm/responses/test_responses_utils.py +++ b/tests/test_litellm/responses/test_responses_utils.py @@ -203,3 +203,22 @@ class TestResponseAPILoggingUtils: assert result.prompt_tokens == 0 assert result.completion_tokens == 20 assert result.total_tokens == 20 + + def test_transform_response_api_usage_calculates_total_from_input_and_output_tokens_if_available(self): + """Test transformation calculates total_tokens when it's None and input / output tokens are present""" + # Setup + usage = { + "input_tokens": 15, + "output_tokens": 25, + "total_tokens": None, + } + + # Execute + result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + usage + ) + + # Assert + assert result.prompt_tokens == 15 + assert result.completion_tokens == 25 + assert result.total_tokens == 40 # 15 + 25 diff --git a/tests/test_litellm/responses/test_text_format_conversion.py b/tests/test_litellm/responses/test_text_format_conversion.py index 645f0f2e148..c7a79d9c461 100644 --- a/tests/test_litellm/responses/test_text_format_conversion.py +++ b/tests/test_litellm/responses/test_text_format_conversion.py @@ -34,7 +34,7 @@ class TestTextFormatConversion: Test that when text_format parameter is passed to litellm.aresponses, it gets converted to text parameter in the raw API call to OpenAI. """ - from unittest.mock import AsyncMock, patch + from unittest.mock import AsyncMock, MagicMock, patch class TestResponse(BaseModel): """Test Pydantic model for structured output""" @@ -42,20 +42,8 @@ class TestTextFormatConversion: answer: str confidence: float - class MockResponse: - """Mock response class for testing""" - - def __init__(self, json_data, status_code): - self._json_data = json_data - self.status_code = status_code - self.text = json.dumps(json_data) - self.headers = {} - - def json(self): - return self._json_data - # Mock response from OpenAI - mock_response = { + mock_response_data = { "id": "resp_123", "object": "response", "created_at": 1741476542, @@ -101,13 +89,74 @@ class TestTextFormatConversion: base_completion_call_args = self.get_base_completion_call_args() - with patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - new_callable=AsyncMock, - ) as mock_post: - # Configure the mock to return our response - mock_post.return_value = MockResponse(mock_response, 200) + # Mock the response_api_handler function to capture the request + captured_request = {} + def mock_handler( + model, + input, + responses_api_provider_config, + response_api_optional_request_params, + custom_llm_provider, + litellm_params, + logging_obj, + extra_headers=None, + extra_body=None, + timeout=None, + client=None, + fake_stream=False, + litellm_metadata=None, + shared_session=None, + _is_async=False, + ): + # Capture the request parameters + captured_request["model"] = model + captured_request["input"] = input + captured_request["params"] = response_api_optional_request_params + + # Return a mock ResponsesAPIResponse wrapped in a coroutine if async + async def async_response(): + return ResponsesAPIResponse( + id="resp_123", + object="response", + created_at=1741476542, + status="completed", + model="gpt-4o", + output=mock_response_data["output"], + usage=ResponseAPIUsage( + input_tokens=10, + output_tokens=20, + total_tokens=30, + ), + text=mock_response_data.get("text"), + error=None, + incomplete_details=None, + ) + + if _is_async: + return async_response() + else: + return ResponsesAPIResponse( + id="resp_123", + object="response", + created_at=1741476542, + status="completed", + model="gpt-4o", + output=mock_response_data["output"], + usage=ResponseAPIUsage( + input_tokens=10, + output_tokens=20, + total_tokens=30, + ), + text=mock_response_data.get("text"), + error=None, + incomplete_details=None, + ) + + with patch( + "litellm.responses.main.base_llm_http_handler.response_api_handler", + new=mock_handler, + ): litellm._turn_on_debug() litellm.set_verbose = True @@ -118,21 +167,19 @@ class TestTextFormatConversion: **base_completion_call_args, ) - # Verify the request was made correctly - mock_post.assert_called_once() - request_body = mock_post.call_args.kwargs["json"] - print("Request body:", json.dumps(request_body, indent=4)) + # Verify the captured request + print("Captured request:", json.dumps(captured_request, indent=4, default=str)) # Validate that text_format was converted to text parameter assert ( - "text" in request_body - ), "text parameter should be present in request body" + "text" in captured_request["params"] + ), "text parameter should be present in request params" assert ( - "text_format" not in request_body - ), "text_format should not be in request body" + "text_format" not in captured_request["params"] + ), "text_format should not be in request params" # Validate the text parameter structure - text_param = request_body["text"] + text_param = captured_request["params"]["text"] assert "format" in text_param, "text parameter should have format field" assert ( text_param["format"]["type"] == "json_schema" @@ -156,7 +203,7 @@ class TestTextFormatConversion: ), "schema should have confidence property" # Validate other request parameters - assert request_body["input"] == "What is the capital of France?" + assert captured_request["input"] == "What is the capital of France?" # Validate the response print("Response:", json.dumps(response, indent=4, default=str)) diff --git a/tests/test_litellm/router_strategy/test_router_tag_routing.py b/tests/test_litellm/router_strategy/test_router_tag_routing.py index a3e722eeb85..1fdd3dad4da 100644 --- a/tests/test_litellm/router_strategy/test_router_tag_routing.py +++ b/tests/test_litellm/router_strategy/test_router_tag_routing.py @@ -313,17 +313,31 @@ async def test_error_from_tag_routing(): def test_tag_routing_with_list_of_tags(): """ - Test that the router can handle a list of tags + Test that the router can handle a list of tags with match_any behavior """ from litellm.router_strategy.tag_based_routing import is_valid_deployment_tag assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA"]) assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamB"]) assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamC"]) + assert is_valid_deployment_tag(["teamA"], ["teamA", "teamB"]) assert not is_valid_deployment_tag(["teamA", "teamB"], ["teamC"]) assert not is_valid_deployment_tag(["teamA", "teamB"], []) assert not is_valid_deployment_tag(["default"], ["teamA"]) +def test_tag_routing_with_list_of_tags_match_all(): + """ + Test that the router can handle a list of tags with match_all behavior + """ + from litellm.router_strategy.tag_based_routing import is_valid_deployment_tag + + assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA"], match_any=False) + assert is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamB"], match_any=False) + assert not is_valid_deployment_tag(["teamA", "teamB", "teamC"], ["teamA", "teamD"], match_any=False) + assert not is_valid_deployment_tag(["teamA"], ["teamA", "teamB"], match_any=False) + assert not is_valid_deployment_tag(["teamA", "teamB"], ["teamA", "teamC"], match_any=False) + assert not is_valid_deployment_tag(["teamA", "teamB"], [], match_any=False) + assert not is_valid_deployment_tag(["default"], ["teamA"], match_any=False) @pytest.mark.asyncio() async def test_router_free_paid_tier_with_responses_api(): diff --git a/tests/test_litellm/test_lazy_imports.py b/tests/test_litellm/test_lazy_imports.py index 660933efac5..48d78c0b01b 100644 --- a/tests/test_litellm/test_lazy_imports.py +++ b/tests/test_litellm/test_lazy_imports.py @@ -42,34 +42,45 @@ from litellm._lazy_imports import ( def _clear_names_from_globals(names: tuple): """Clear all names from litellm globals.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ for name in names: - if name in litellm.__dict__: - del litellm.__dict__[name] + if name in litellm_globals: + del litellm_globals[name] def _clear_names_from_utils_globals(names: tuple): """Clear all names from litellm.utils globals.""" + # Get the actual globals dict, not a copy + utils_globals = sys.modules["litellm.utils"].__dict__ for name in names: - if name in litellm.utils.__dict__: - del litellm.utils.__dict__[name] + if name in utils_globals: + del utils_globals[name] def _verify_only_requested_name_imported(name: str, all_names: tuple): """Verify that only the requested name is in globals, not the others.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ for other_name in all_names: if other_name != name: - assert other_name not in litellm.__dict__, f"{other_name} should not be imported when importing {name}" + assert other_name not in litellm_globals, f"{other_name} should not be imported when importing {name}" def _verify_only_requested_name_imported_in_utils(name: str, all_names: tuple): """Verify that only the requested name is in utils globals, not the others.""" + # Get the actual globals dict, not a copy + utils_globals = sys.modules["litellm.utils"].__dict__ for other_name in all_names: if other_name != name: - assert other_name not in litellm.utils.__dict__, f"{other_name} should not be imported when importing {name}" + assert other_name not in utils_globals, f"{other_name} should not be imported when importing {name}" def test_cost_calculator_lazy_imports(): """Test that all cost calculator functions can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + # Test each name individually - only that name should be imported for name in COST_CALCULATOR_NAMES: # Clear all names before importing just one @@ -78,7 +89,7 @@ def test_cost_calculator_lazy_imports(): func = _lazy_import_cost_calculator(name) assert func is not None assert callable(func) - assert name in litellm.__dict__ + assert name in litellm_globals # Verify only the requested name is in globals, not the others _verify_only_requested_name_imported(name, COST_CALCULATOR_NAMES) @@ -86,6 +97,9 @@ def test_cost_calculator_lazy_imports(): def test_litellm_logging_lazy_imports(): """Test that all litellm_logging items can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + # Test each name individually - only that name should be imported for name in LITELLM_LOGGING_NAMES: # Clear all names before importing just one @@ -93,7 +107,7 @@ def test_litellm_logging_lazy_imports(): item = _lazy_import_litellm_logging(name) assert item is not None - assert name in litellm.__dict__ + assert name in litellm_globals # Verify only the requested name is in globals, not the others _verify_only_requested_name_imported(name, LITELLM_LOGGING_NAMES) @@ -101,6 +115,9 @@ def test_litellm_logging_lazy_imports(): def test_utils_lazy_imports(): """Test that all utils functions can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + # Test each name individually - only that name should be imported for name in UTILS_NAMES: # Clear all names before importing just one @@ -108,7 +125,7 @@ def test_utils_lazy_imports(): attr = _lazy_import_utils(name) assert attr is not None - assert name in litellm.__dict__ + assert name in litellm_globals # Verify only the requested name is in globals, not the others _verify_only_requested_name_imported(name, UTILS_NAMES) @@ -116,6 +133,9 @@ def test_utils_lazy_imports(): def test_caching_lazy_imports(): """Test that all caching classes can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + # Test each name individually - only that name should be imported for name in CACHING_NAMES: # Clear all names before importing just one @@ -123,7 +143,7 @@ def test_caching_lazy_imports(): cls = _lazy_import_caching(name) assert cls is not None - assert name in litellm.__dict__ + assert name in litellm_globals # Verify only the requested name is in globals, not the others _verify_only_requested_name_imported(name, CACHING_NAMES) @@ -131,71 +151,89 @@ def test_caching_lazy_imports(): def test_token_counter_lazy_imports(): """Test that token counter utilities can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in TOKEN_COUNTER_NAMES: _clear_names_from_globals(TOKEN_COUNTER_NAMES) func = _lazy_import_token_counter(name) assert func is not None - assert name in litellm.__dict__ + assert name in litellm_globals _verify_only_requested_name_imported(name, TOKEN_COUNTER_NAMES) def test_bedrock_types_lazy_imports(): """Test that Bedrock type aliases can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in BEDROCK_TYPES_NAMES: _clear_names_from_globals(BEDROCK_TYPES_NAMES) alias = _lazy_import_bedrock_types(name) assert alias is not None - assert name in litellm.__dict__ + assert name in litellm_globals _verify_only_requested_name_imported(name, BEDROCK_TYPES_NAMES) def test_types_utils_lazy_imports(): """Test that common types.utils symbols can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in TYPES_UTILS_NAMES: _clear_names_from_globals(TYPES_UTILS_NAMES) obj = _lazy_import_types_utils(name) assert obj is not None - assert name in litellm.__dict__ + assert name in litellm_globals _verify_only_requested_name_imported(name, TYPES_UTILS_NAMES) def test_llm_client_cache_lazy_imports(): """Test that LLM client cache class and singleton can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in LLM_CLIENT_CACHE_NAMES: _clear_names_from_globals(LLM_CLIENT_CACHE_NAMES) obj = _lazy_import_llm_client_cache(name) assert obj is not None - assert name in litellm.__dict__ + assert name in litellm_globals _verify_only_requested_name_imported(name, LLM_CLIENT_CACHE_NAMES) def test_http_handler_lazy_imports(): """Test that HTTP handler singletons can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in HTTP_HANDLER_NAMES: _clear_names_from_globals(HTTP_HANDLER_NAMES) handler = _lazy_import_http_handlers(name) assert handler is not None - assert name in litellm.__dict__ + assert name in litellm_globals _verify_only_requested_name_imported(name, HTTP_HANDLER_NAMES) def test_dotprompt_lazy_imports(): """Test that dotprompt globals can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in DOTPROMPT_NAMES: _clear_names_from_globals(DOTPROMPT_NAMES) obj = _lazy_import_dotprompt(name) - assert name in litellm.__dict__ + assert name in litellm_globals # Only the setter must be callable; others may be None by default if name == "set_global_prompt_directory": @@ -245,12 +283,15 @@ def test_unknown_attribute_raises_error(): def test_llm_config_lazy_imports(): """Test that LLM config classes can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in LLM_CONFIG_NAMES: _clear_names_from_globals(LLM_CONFIG_NAMES) obj = _lazy_import_llm_configs(name) assert obj is not None - assert name in litellm.__dict__ + assert name in litellm_globals # Config classes should be classes/types assert isinstance(obj, type), f"{name} should be a class" @@ -259,12 +300,15 @@ def test_llm_config_lazy_imports(): def test_types_lazy_imports(): """Test that type classes can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in TYPES_NAMES: _clear_names_from_globals(TYPES_NAMES) obj = _lazy_import_types(name) assert obj is not None - assert name in litellm.__dict__ + assert name in litellm_globals # Type classes should be classes/types assert isinstance(obj, type), f"{name} should be a class" @@ -273,25 +317,31 @@ def test_types_lazy_imports(): def test_llm_provider_logic_lazy_imports(): """Test that LLM provider logic functions can be lazy imported.""" + # Get the actual globals dict, not a copy + litellm_globals = sys.modules["litellm"].__dict__ + for name in LLM_PROVIDER_LOGIC_NAMES: _clear_names_from_globals(LLM_PROVIDER_LOGIC_NAMES) func = _lazy_import_llm_provider_logic(name) assert func is not None assert callable(func) - assert name in litellm.__dict__ + assert name in litellm_globals _verify_only_requested_name_imported(name, LLM_PROVIDER_LOGIC_NAMES) def test_utils_module_lazy_imports(): """Test that utils module attributes can be lazy imported.""" + # Get the actual globals dict, not a copy + utils_globals = sys.modules["litellm.utils"].__dict__ + for name in UTILS_MODULE_NAMES: _clear_names_from_utils_globals(UTILS_MODULE_NAMES) obj = _lazy_import_utils_module(name) assert obj is not None - assert name in litellm.utils.__dict__ + assert name in utils_globals _verify_only_requested_name_imported_in_utils(name, UTILS_MODULE_NAMES) diff --git a/tests/test_litellm/test_responses_id_security.py b/tests/test_litellm/test_responses_id_security.py index 6b04479326e..2addf504f7c 100644 --- a/tests/test_litellm/test_responses_id_security.py +++ b/tests/test_litellm/test_responses_id_security.py @@ -42,8 +42,11 @@ class TestIsEncryptedResponseId: def test_is_encrypted_response_id_valid(self, responses_id_security): """Test that a properly encrypted response ID is identified correctly""" - with patch( - "litellm.proxy.hooks.responses_id_security.decrypt_value_helper" + # Patch at the module level where it's imported + import litellm.proxy.hooks.responses_id_security as responses_module + + with patch.object( + responses_module, "decrypt_value_helper" ) as mock_decrypt: mock_decrypt.return_value = f"{SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value}response_id:resp_123;user_id:user-456" @@ -56,8 +59,11 @@ class TestIsEncryptedResponseId: def test_is_encrypted_response_id_invalid(self, responses_id_security): """Test that an unencrypted response ID returns False""" - with patch( - "litellm.proxy.hooks.responses_id_security.decrypt_value_helper" + # Patch at the module level where it's imported + import litellm.proxy.hooks.responses_id_security as responses_module + + with patch.object( + responses_module, "decrypt_value_helper" ) as mock_decrypt: mock_decrypt.return_value = None @@ -71,8 +77,11 @@ class TestDecryptResponseId: def test_decrypt_response_id_valid(self, responses_id_security): """Test decrypting a valid encrypted response ID""" - with patch( - "litellm.proxy.hooks.responses_id_security.decrypt_value_helper" + # Patch at the module level where it's imported + import litellm.proxy.hooks.responses_id_security as responses_module + + with patch.object( + responses_module, "decrypt_value_helper" ) as mock_decrypt: mock_decrypt.return_value = f"{SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value}response_id:resp_original_123;user_id:user-456;team_id:team-789" @@ -86,8 +95,11 @@ class TestDecryptResponseId: def test_decrypt_response_id_no_encryption(self, responses_id_security): """Test decrypting a non-encrypted response ID""" - with patch( - "litellm.proxy.hooks.responses_id_security.decrypt_value_helper" + # Patch at the module level where it's imported + import litellm.proxy.hooks.responses_id_security as responses_module + + with patch.object( + responses_module, "decrypt_value_helper" ) as mock_decrypt: mock_decrypt.return_value = None diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index cd76c438ded..bfa162e019b 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -749,6 +749,57 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): raise AssertionError(error_message) +def test_max_tokens_consistency(): + """ + Test that max_tokens == max_output_tokens for all models. + + According to the spec in model_prices_and_context_window.json: + - max_tokens is a LEGACY parameter + - It should be set to max_output_tokens if the provider specifies it + + This test ensures consistency across all model definitions. + """ + import json + from pathlib import Path + + # Load the model configuration + config_path = Path(__file__).parent.parent.parent / "model_prices_and_context_window.json" + with open(config_path, 'r') as f: + models = json.load(f) + + inconsistencies = [] + + for model_name, config in models.items(): + # Skip the sample_spec + if model_name == "sample_spec": + continue + + # Check if both max_tokens and max_output_tokens exist + if isinstance(config, dict): + max_tokens = config.get('max_tokens') + max_output_tokens = config.get('max_output_tokens') + + # Only validate if both exist + if max_tokens is not None and max_output_tokens is not None: + if max_tokens != max_output_tokens: + inconsistencies.append({ + 'model': model_name, + 'max_tokens': max_tokens, + 'max_output_tokens': max_output_tokens + }) + + if inconsistencies: + error_msg = f"\n\n❌ Found {len(inconsistencies)} models with max_tokens != max_output_tokens:\n\n" + for item in inconsistencies[:10]: # Show first 10 + error_msg += f" {item['model']}: max_tokens={item['max_tokens']}, max_output_tokens={item['max_output_tokens']}\n" + + if len(inconsistencies) > 10: + error_msg += f"\n ... and {len(inconsistencies) - 10} more\n" + + error_msg += "\nTo fix these inconsistencies, run: poetry run python fix_max_tokens_inconsistencies.py" + raise AssertionError(error_msg) + + def test_get_model_info_gemini(): """ Tests if ALL gemini models have 'tpm' and 'rpm' in the model info diff --git a/tests/test_litellm/test_utils_custom.py b/tests/test_litellm/test_utils_custom.py new file mode 100644 index 00000000000..3e924e9c719 --- /dev/null +++ b/tests/test_litellm/test_utils_custom.py @@ -0,0 +1,45 @@ +import pytest +import sys +from unittest.mock import MagicMock, patch, AsyncMock +from litellm.proxy.utils import count_tokens_with_anthropic_api, _anthropic_async_clients + +@pytest.mark.asyncio +async def test_count_tokens_caching(): + """ + Test that count_tokens_with_anthropic_api caches the client. + """ + # Clear cache + _anthropic_async_clients.clear() + + api_key = "sk-ant-test-key" + messages = [{"role": "user", "content": "hello"}] + model = "claude-3-opus-20240229" + + # Create a mock anthropic module + mock_anthropic = MagicMock() + mock_client = MagicMock() + mock_anthropic.AsyncAnthropic.return_value = mock_client + + # Mock response + mock_response = MagicMock() + mock_response.input_tokens = 10 + + # Setup async return for count_tokens + mock_client.beta.messages.count_tokens = AsyncMock(return_value=mock_response) + + # Patch sys.modules to ensure our mock is used when anthropic is imported + with patch.dict(sys.modules, {"anthropic": mock_anthropic}): + # First call + with patch.dict("os.environ", {"ANTHROPIC_API_KEY": api_key}): + await count_tokens_with_anthropic_api(model, messages) + + assert api_key in _anthropic_async_clients + assert _anthropic_async_clients[api_key] == mock_client + mock_anthropic.AsyncAnthropic.assert_called_once() # Should be called once + + # Second call + with patch.dict("os.environ", {"ANTHROPIC_API_KEY": api_key}): + await count_tokens_with_anthropic_api(model, messages) + + # Should still be called once (cached) + mock_anthropic.AsyncAnthropic.assert_called_once() diff --git a/tests/test_litellm/test_video_generation.py b/tests/test_litellm/test_video_generation.py index 87012f05155..73bfa71d20b 100644 --- a/tests/test_litellm/test_video_generation.py +++ b/tests/test_litellm/test_video_generation.py @@ -2,7 +2,7 @@ import asyncio import json import os import sys -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -18,6 +18,7 @@ from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.llms.gemini.videos.transformation import GeminiVideoConfig from litellm.llms.openai.videos.transformation import OpenAIVideoConfig from litellm.types.videos.main import VideoObject, VideoResponse +from litellm.videos import main as videos_main from litellm.videos.main import ( avideo_generation, avideo_status, @@ -31,32 +32,29 @@ class TestVideoGeneration: def test_video_generation_basic(self): """Test basic video generation functionality.""" - # Mock the video generation response - mock_response = VideoObject( - id="video_123", - object="video", - status="queued", - created_at=1712697600, + # Use mock_response parameter for reliable testing + response = video_generation( + prompt="Show them running around the room", model="sora-2", + seconds="8", size="720x1280", - seconds="8" + mock_response={ + "id": "video_123", + "object": "video", + "status": "queued", + "created_at": 1712697600, + "model": "sora-2", + "size": "720x1280", + "seconds": "8" + } ) - with patch('litellm.videos.main.base_llm_http_handler') as mock_handler: - mock_handler.video_generation_handler.return_value = mock_response - - response = video_generation( - prompt="Show them running around the room", - model="sora-2", - seconds="8", - size="720x1280" - ) - - assert isinstance(response, VideoObject) - assert response.id == "video_123" - assert response.model == "sora-2" - assert response.size == "720x1280" - assert response.seconds == "8" + assert isinstance(response, VideoObject) + assert response.id == "video_123" + assert response.status == "queued" + assert response.model == "sora-2" + assert response.size == "720x1280" + assert response.seconds == "8" def test_video_generation_with_mock_response(self): """Test video generation with mock response.""" @@ -97,26 +95,27 @@ class TestVideoGeneration: progress=50 ) - with patch('litellm.videos.main.base_llm_http_handler') as mock_handler: - mock_handler.video_generation_handler.return_value = mock_response - - import asyncio - - async def test_async(): - response = await avideo_generation( - prompt="A cat playing with a ball", - model="sora-2", - seconds="5", - size="720x1280" - ) - return response - - response = asyncio.run(test_async()) - - assert isinstance(response, VideoObject) - assert response.id == "video_async_123" - assert response.status == "processing" - assert response.progress == 50 + # Mock the async_video_generation_handler to return the mock_response + async_mock = AsyncMock(return_value=mock_response) + with patch.object(videos_main.base_llm_http_handler, 'async_video_generation_handler', async_mock): + with patch.object(videos_main.base_llm_http_handler, 'video_generation_handler', side_effect=lambda **kwargs: async_mock(**kwargs)): + import asyncio + + async def test_async(): + response = await avideo_generation( + prompt="A cat playing with a ball", + model="sora-2", + seconds="5", + size="720x1280" + ) + return response + + response = asyncio.run(test_async()) + + assert isinstance(response, VideoObject) + assert response.id == "video_async_123" + assert response.status == "processing" + assert response.progress == 50 def test_video_generation_parameter_validation(self): """Test video generation parameter validation.""" @@ -132,9 +131,7 @@ class TestVideoGeneration: def test_video_generation_error_handling(self): """Test video generation error handling.""" - with patch('litellm.videos.main.base_llm_http_handler') as mock_handler: - mock_handler.video_generation_handler.side_effect = Exception("API Error") - + with patch.object(videos_main.base_llm_http_handler, 'video_generation_handler', side_effect=Exception("API Error")): with pytest.raises(Exception): video_generation( prompt="Test video", @@ -443,32 +440,28 @@ class TestVideoGeneration: def test_video_status_basic(self): """Test basic video status functionality.""" - # Mock the video status response - mock_response = VideoObject( - id="video_123", - object="video", - status="completed", - created_at=1712697600, - completed_at=1712697660, + # Use mock_response parameter for reliable testing + response = video_status( + video_id="video_123", model="sora-2", - progress=100, - size="720x1280", - seconds="8" + mock_response={ + "id": "video_123", + "object": "video", + "status": "completed", + "created_at": 1712697600, + "completed_at": 1712697660, + "model": "sora-2", + "progress": 100, + "size": "720x1280", + "seconds": "8" + } ) - with patch('litellm.videos.main.base_llm_http_handler') as mock_handler: - mock_handler.video_status_handler.return_value = mock_response - - response = video_status( - video_id="video_123", - model="sora-2" - ) - - assert isinstance(response, VideoObject) - assert response.id == "video_123" - assert response.status == "completed" - assert response.progress == 100 - assert response.model == "sora-2" + assert isinstance(response, VideoObject) + assert response.id == "video_123" + assert response.status == "completed" + assert response.progress == 100 + assert response.model == "sora-2" def test_video_status_with_mock_response(self): """Test video status with mock response.""" @@ -506,24 +499,25 @@ class TestVideoGeneration: progress=0 ) - with patch('litellm.videos.main.base_llm_http_handler') as mock_handler: - mock_handler.video_status_handler.return_value = mock_response - - import asyncio - - async def test_async(): - response = await avideo_status( - video_id="video_async_123", - model="sora-2" - ) - return response - - response = asyncio.run(test_async()) - - assert isinstance(response, VideoObject) - assert response.id == "video_async_123" - assert response.status == "queued" - assert response.progress == 0 + # Mock the async_video_status_handler to return the mock_response + async_mock = AsyncMock(return_value=mock_response) + with patch.object(videos_main.base_llm_http_handler, 'async_video_status_handler', async_mock): + with patch.object(videos_main.base_llm_http_handler, 'video_status_handler', side_effect=lambda **kwargs: async_mock(**kwargs)): + import asyncio + + async def test_async(): + response = await avideo_status( + video_id="video_async_123", + model="sora-2" + ) + return response + + response = asyncio.run(test_async()) + + assert isinstance(response, VideoObject) + assert response.id == "video_async_123" + assert response.status == "queued" + assert response.progress == 0 def test_video_status_parameter_validation(self): """Test video status parameter validation.""" @@ -539,9 +533,7 @@ class TestVideoGeneration: def test_video_status_error_handling(self): """Test video status error handling.""" - with patch('litellm.videos.main.base_llm_http_handler') as mock_handler: - mock_handler.video_status_handler.side_effect = Exception("API Error") - + with patch.object(videos_main.base_llm_http_handler, 'video_status_handler', side_effect=Exception("API Error")): with pytest.raises(Exception): video_status( video_id="test_video_id", @@ -672,33 +664,30 @@ class TestVideoGeneration: def test_video_status_async_inside_async_function(self): """Test that sync video_status works inside async functions (no asyncio.run issues).""" - mock_response = VideoObject( - id="video_sync_in_async", - object="video", - status="completed", - created_at=1712697600, - model="sora-2", - progress=100 - ) + import asyncio - with patch('litellm.videos.main.base_llm_http_handler') as mock_handler: - mock_handler.video_status_handler.return_value = mock_response - - import asyncio - - async def test_sync_in_async(): - # This should work without asyncio.run() issues - response = video_status( - video_id="video_sync_in_async", - model="sora-2" - ) - return response - - response = asyncio.run(test_sync_in_async()) - - assert isinstance(response, VideoObject) - assert response.id == "video_sync_in_async" - assert response.status == "completed" + async def test_sync_in_async(): + # This should work without asyncio.run() issues + # Use mock_response parameter for reliable testing + response = video_status( + video_id="video_sync_in_async", + model="sora-2", + mock_response={ + "id": "video_sync_in_async", + "object": "video", + "status": "completed", + "created_at": 1712697600, + "model": "sora-2", + "progress": 100 + } + ) + return response + + response = asyncio.run(test_sync_in_async()) + + assert isinstance(response, VideoObject) + assert response.id == "video_sync_in_async" + assert response.status == "completed" def test_video_status_url_construction(self): """Test video status URL construction.""" diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts b/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts new file mode 100644 index 00000000000..4a4bb64c8ed --- /dev/null +++ b/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts @@ -0,0 +1,38 @@ +import { Page } from "./pages"; + +/** + * Maps sidebar menu item labels to their corresponding page enum values. + * This mapping is for the admin role. + */ +export const menuLabelToPage: Record = { + "Virtual Keys": Page.ApiKeys, + Playground: Page.LlmPlayground, + Models: Page.Models, + "Models + Endpoints": Page.Models, + Usage: Page.NewUsage, + Teams: Page.Teams, + "Internal Users": Page.Users, + "Internal User": Page.Users, // Legacy label support + Organizations: Page.Organizations, + "API Reference": Page.ApiRef, + "AI Hub": Page.ModelHubTable, + "Model Hub": Page.ModelHubTable, + Logs: Page.Logs, + Guardrails: Page.Guardrails, + // Settings submenu items + "Router Settings": Page.RouterSettings, + "Logging & Alerts": Page.LoggingAndAlerts, + "Admin Settings": Page.AdminPanel, + "Cost Tracking": Page.CostTracking, + "UI Theme": Page.UiTheme, + // Experimental submenu items + Caching: Page.Caching, + Prompts: Page.Prompts, + Budgets: Page.Budgets, + "API Playground": Page.TransformRequest, + "Tag Management": Page.TagManagement, + "Old Usage": Page.Usage, + // Tools submenu items + "MCP Servers": Page.McpServers, + "Vector Stores": Page.VectorStores, +}; diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/pages.ts b/ui/litellm-dashboard/e2e_tests/fixtures/pages.ts new file mode 100644 index 00000000000..3ea37718ab5 --- /dev/null +++ b/ui/litellm-dashboard/e2e_tests/fixtures/pages.ts @@ -0,0 +1,33 @@ +/** + * Enum for all page query parameters supported in the app. + * These values correspond to the `page` query parameter used in the URL. + */ +export enum Page { + ApiKeys = "api-keys", + Models = "models", + LlmPlayground = "llm-playground", + Users = "users", + Teams = "teams", + Organizations = "organizations", + AdminPanel = "admin-panel", + ApiRef = "api_ref", + LoggingAndAlerts = "logging-and-alerts", + Budgets = "budgets", + Guardrails = "guardrails", + Agents = "agents", + Prompts = "prompts", + TransformRequest = "transform-request", + RouterSettings = "router-settings", + UiTheme = "ui-theme", + CostTracking = "cost-tracking", + ModelHubTable = "model-hub-table", + Caching = "caching", + PassThroughSettings = "pass-through-settings", + Logs = "logs", + McpServers = "mcp-servers", + SearchTools = "search-tools", + TagManagement = "tag-management", + VectorStores = "vector-stores", + NewUsage = "new_usage", + Usage = "usage", +} diff --git a/ui/litellm-dashboard/e2e_tests/helpers/navigation.ts b/ui/litellm-dashboard/e2e_tests/helpers/navigation.ts new file mode 100644 index 00000000000..919e516b35b --- /dev/null +++ b/ui/litellm-dashboard/e2e_tests/helpers/navigation.ts @@ -0,0 +1,12 @@ +import { Page } from "../fixtures/pages"; +import { Page as PlaywrightPage } from "@playwright/test"; + +/** + * Navigates to a specific page using the page query parameter. + * Uses relative path which will be resolved against the baseURL configured in playwright.config.ts + * @param page - The Playwright page object + * @param pageEnum - The page enum value to navigate to + */ +export async function navigateToPage(page: PlaywrightPage, pageEnum: Page): Promise { + await page.goto(`/ui?page=${pageEnum}`); +} diff --git a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts index c0619cfa845..2ab782d5678 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts @@ -14,7 +14,7 @@ test.describe("Add Model", () => { await providerInputDropdown.fill("Anthropic"); await page.waitForTimeout(1000); await providerInputDropdown.press("Enter"); - await page.waitForTimeout(1000); + await page.waitForTimeout(2000); const providerModelsDropdown = page.locator(".ant-select-selection-overflow").first(); await providerModelsDropdown.click(); diff --git a/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts index c90be698ae1..ce07cc2b83d 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts @@ -1,6 +1,9 @@ import test, { expect } from "@playwright/test"; import { Role } from "../../fixtures/roles"; import { ADMIN_STORAGE_PATH } from "../../constants"; +import { Page } from "../../fixtures/pages"; +import { menuLabelToPage } from "../../fixtures/menuMappings"; +import { navigateToPage } from "../../helpers/navigation"; const sidebarButtons = { [Role.ProxyAdmin]: [ @@ -9,9 +12,7 @@ const sidebarButtons = { "Models", "Usage", "Teams", - "Internal User", - "Settings", - "Experimental", + "Internal Users", "API Reference", "AI Hub", ], @@ -23,13 +24,36 @@ for (const { role, storage } of roles) { test.describe(`${role} sidebar`, () => { test.use({ storageState: storage }); - test("can see and navigate all sidebar buttons", async ({ page }) => { + test("should navigate to correct URL when clicking sidebar menu items from homepage", async ({ page }) => { await page.goto("/ui"); - for (const button of sidebarButtons[role as keyof typeof sidebarButtons]) { - const tab = page.getByRole("menuitem", { name: button }); + + for (const buttonLabel of sidebarButtons[role as keyof typeof sidebarButtons]) { + const expectedPage = menuLabelToPage[buttonLabel]; + + if (!expectedPage) { + throw new Error(`No page mapping found for menu label: ${buttonLabel}`); + } + + const tab = page.getByRole("menuitem", { name: buttonLabel }); await expect(tab).toBeVisible(); + await tab.click(); + + // Verify URL contains the correct page query parameter + await expect(page).toHaveURL(new RegExp(`[?&]page=${expectedPage}(&|$)`)); } }); + + test("should navigate directly to page using navigation helper", async ({ page }) => { + // Test direct navigation to verify the helper function works + await navigateToPage(page, Page.ApiKeys); + await expect(page).toHaveURL(new RegExp(`[?&]page=${Page.ApiKeys}(&|$)`)); + + await navigateToPage(page, Page.Models); + await expect(page).toHaveURL(new RegExp(`[?&]page=${Page.Models}(&|$)`)); + + await navigateToPage(page, Page.LlmPlayground); + await expect(page).toHaveURL(new RegExp(`[?&]page=${Page.LlmPlayground}(&|$)`)); + }); }); } diff --git a/ui/litellm-dashboard/e2e_tests/tests/users/searchUsers.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/users/searchUsers.spec.ts index 01c1e68f1ee..a9b0e329a2b 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/users/searchUsers.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/users/searchUsers.spec.ts @@ -1,6 +1,6 @@ import { test, expect, Page } from "@playwright/test"; import { ADMIN_STORAGE_PATH } from "../../constants"; -test.describe("Internal Users Search", () => { +test.skip("Internal Users Search", () => { test.use({ storageState: ADMIN_STORAGE_PATH }); async function goToInternalUsers(page: Page) { diff --git a/ui/litellm-dashboard/e2e_tests/tests/users/viewInternalUsers.spec.ts b/ui/litellm-dashboard/e2e_tests/tests/users/viewInternalUsers.spec.ts index 4dfd79c9dff..ea61c238c02 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/users/viewInternalUsers.spec.ts +++ b/ui/litellm-dashboard/e2e_tests/tests/users/viewInternalUsers.spec.ts @@ -1,7 +1,7 @@ import { test, expect, Page } from "@playwright/test"; import { ADMIN_STORAGE_PATH } from "../../constants"; -test.describe("Internal Users Page", () => { +test.skip("Internal Users Page", () => { test.use({ storageState: ADMIN_STORAGE_PATH }); async function goToInternalUsers(page: Page) { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/router/useRouterFields.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/router/useRouterFields.test.ts new file mode 100644 index 00000000000..fe4680cedeb --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/router/useRouterFields.test.ts @@ -0,0 +1,387 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { renderHook, waitFor } from "@testing-library/react"; +import React, { ReactNode } from "react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { RouterFieldsResponse, useRouterFields } from "./useRouterFields"; + +// Mock the networking module +vi.mock("@/components/networking", () => ({ + proxyBaseUrl: null, +})); + +// Mock useAuthorized hook +const mockUseAuthorized = vi.fn(); +vi.mock("../useAuthorized", () => ({ + default: () => mockUseAuthorized(), +})); + +// Mock global fetch +const mockFetch = vi.fn(); +global.fetch = mockFetch; + +// Mock console methods to avoid noise in tests +vi.spyOn(console, "log").mockImplementation(() => {}); +vi.spyOn(console, "error").mockImplementation(() => {}); + +// Mock data +const mockRouterFieldsResponse: RouterFieldsResponse = { + fields: [ + { + field_name: "routing_strategy", + field_type: "String", + field_description: "Routing strategy to use for load balancing across deployments", + field_default: "simple-shuffle", + options: ["simple-shuffle", "least-busy", "latency-based-routing"], + ui_field_name: "Routing Strategy", + link: null, + }, + { + field_name: "num_retries", + field_type: "Integer", + field_description: "Number of retries for failed requests", + field_default: 0, + options: null, + ui_field_name: "Number of Retries", + link: null, + }, + ], + routing_strategy_descriptions: { + "simple-shuffle": "Randomly picks a deployment from the list. Simple and fast.", + "least-busy": "Routes to the deployment with the lowest number of ongoing requests.", + "latency-based-routing": "Routes to the deployment with the lowest latency over a sliding window.", + }, +}; + +describe("useRouterFields", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + }, + }, + }); + + // Reset all mocks + vi.clearAllMocks(); + + // Set default mock for useAuthorized (enabled state) + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userRole: "Admin", + userId: "test-user-id", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("should render", () => { + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => mockRouterFieldsResponse, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + expect(result.current).toBeDefined(); + }); + + it("should return router fields data when query is successful", async () => { + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => mockRouterFieldsResponse, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + // Initially loading + expect(result.current.isLoading).toBe(true); + expect(result.current.data).toBeUndefined(); + + // Wait for success + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data).toEqual(mockRouterFieldsResponse); + expect(result.current.error).toBeNull(); + expect(mockFetch).toHaveBeenCalledTimes(1); + expect(mockFetch).toHaveBeenCalledWith("/router/fields", { + method: "GET", + headers: { + Authorization: "Bearer test-access-token", + "Content-Type": "application/json", + }, + }); + }); + + it("should handle error when fetch fails", async () => { + const errorMessage = "Failed to fetch router fields"; + const errorResponse = { error: errorMessage }; + + mockFetch.mockResolvedValueOnce({ + ok: false, + json: async () => errorResponse, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + // Initially loading + expect(result.current.isLoading).toBe(true); + + // Wait for error + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + expect(result.current.isError).toBe(true); + }); + + expect(result.current.error).toBeDefined(); + expect(result.current.data).toBeUndefined(); + expect(mockFetch).toHaveBeenCalledTimes(1); + }); + + it("should not execute query when accessToken is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: null, + userRole: "Admin", + userId: "test-user-id", + token: null, + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + // Query should not execute + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + + // API should not be called + expect(mockFetch).not.toHaveBeenCalled(); + }); + + it("should not execute query when userId is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userRole: "Admin", + userId: null, + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + // Query should not execute + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + + // API should not be called + expect(mockFetch).not.toHaveBeenCalled(); + }); + + it("should not execute query when userRole is missing", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userRole: null, + userId: "test-user-id", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + // Query should not execute + expect(result.current.isLoading).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(result.current.isFetched).toBe(false); + + // API should not be called + expect(mockFetch).not.toHaveBeenCalled(); + }); + + it("should handle network error", async () => { + const networkError = new Error("Network error"); + mockFetch.mockRejectedValueOnce(networkError); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + // Wait for error + await waitFor(() => { + expect(result.current.isError).toBe(true); + }); + + expect(result.current.error).toBeDefined(); + expect(result.current.data).toBeUndefined(); + }); + + it("should use relative URL when proxyBaseUrl is null", async () => { + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => mockRouterFieldsResponse, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + await waitFor(() => { + expect(result.current.isSuccess).toBe(true); + }); + + // When proxyBaseUrl is null, should use relative URL + expect(mockFetch).toHaveBeenCalledWith("/router/fields", { + method: "GET", + headers: { + Authorization: "Bearer test-access-token", + "Content-Type": "application/json", + }, + }); + }); + + it("should handle error response with different error formats", async () => { + const errorFormats = [ + { error: { message: "Error message" } }, + { message: "Error message" }, + { detail: "Error detail" }, + { error: "Error string" }, + { unknown: "format" }, + ]; + + for (const errorFormat of errorFormats) { + vi.clearAllMocks(); + mockFetch.mockResolvedValueOnce({ + ok: false, + json: async () => errorFormat, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + await waitFor(() => { + expect(result.current.isError).toBe(true); + }); + + expect(result.current.error).toBeDefined(); + } + }); + + it("should return empty fields array when API returns empty fields", async () => { + const emptyResponse: RouterFieldsResponse = { + fields: [], + routing_strategy_descriptions: {}, + }; + + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => emptyResponse, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + await waitFor(() => { + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data?.fields).toEqual([]); + expect(result.current.data?.routing_strategy_descriptions).toEqual({}); + }); + + it("should have correct query configuration", async () => { + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => mockRouterFieldsResponse, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + await waitFor(() => { + expect(result.current.isSuccess).toBe(true); + }); + + // Verify the query was called + expect(mockFetch).toHaveBeenCalledTimes(1); + + // The hook should have the expected properties from useQuery + expect(result.current).toHaveProperty("data"); + expect(result.current).toHaveProperty("isLoading"); + expect(result.current).toHaveProperty("isError"); + expect(result.current).toHaveProperty("isSuccess"); + expect(result.current).toHaveProperty("error"); + }); + + it("should handle fields with null options", async () => { + const responseWithNullOptions: RouterFieldsResponse = { + fields: [ + { + field_name: "timeout", + field_type: "Float", + field_description: "Timeout for requests in seconds", + field_default: null, + options: null, + ui_field_name: "Timeout", + link: null, + }, + ], + routing_strategy_descriptions: {}, + }; + + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => responseWithNullOptions, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + await waitFor(() => { + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data?.fields[0].options).toBeNull(); + }); + + it("should handle fields with link property", async () => { + const responseWithLink: RouterFieldsResponse = { + fields: [ + { + field_name: "enable_tag_filtering", + field_type: "Boolean", + field_description: "Enable tag-based routing", + field_default: false, + options: null, + ui_field_name: "Enable Tag Filtering", + link: "https://docs.litellm.ai/docs/proxy/tag_routing", + }, + ], + routing_strategy_descriptions: {}, + }; + + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => responseWithLink, + }); + + const { result } = renderHook(() => useRouterFields(), { wrapper }); + + await waitFor(() => { + expect(result.current.isSuccess).toBe(true); + }); + + expect(result.current.data?.fields[0].link).toBe("https://docs.litellm.ai/docs/proxy/tag_routing"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/router/useRouterFields.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/router/useRouterFields.ts new file mode 100644 index 00000000000..589508c5ddb --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/router/useRouterFields.ts @@ -0,0 +1,69 @@ +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useQuery, UseQueryResult } from "@tanstack/react-query"; +import { createQueryKeys } from "../common/queryKeysFactory"; +import { proxyBaseUrl } from "@/components/networking"; + +export interface RouterSettingsField { + field_name: string; + field_type: string; + field_description: string; + field_default: any; + options: string[] | null; + ui_field_name: string; + link: string | null; +} + +export interface RouterFieldsResponse { + fields: RouterSettingsField[]; + routing_strategy_descriptions: Record; +} + +const routerFieldsKeys = createQueryKeys("routerFields"); + +const deriveErrorMessage = (errorData: any): string => { + return ( + (errorData?.error && (errorData.error.message || errorData.error)) || + errorData?.message || + errorData?.detail || + errorData?.error || + JSON.stringify(errorData) + ); +}; + +const getRouterFields = async (accessToken: string): Promise => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/router/fields` : `/router/fields`; + + console.log("Fetching router fields from:", url); + + const response = await fetch(url, { + method: "GET", + headers: { + Authorization: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + throw new Error(errorMessage); + } + + const data: RouterFieldsResponse = await response.json(); + console.log("Fetched router fields:", data); + return data; + } catch (error) { + console.error("Failed to fetch router fields:", error); + throw error; + } +}; + +export const useRouterFields = (): UseQueryResult => { + const { accessToken, userId, userRole } = useAuthorized(); + return useQuery({ + queryKey: routerFieldsKeys.detail("fields"), + queryFn: async () => await getRouterFields(accessToken!), + enabled: Boolean(accessToken && userId && userRole), + }); +}; diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatImageUtils.test.tsx b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatImageUtils.test.tsx new file mode 100644 index 00000000000..ecc4914b0ab --- /dev/null +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatImageUtils.test.tsx @@ -0,0 +1,187 @@ +import { describe, expect, it, vi, beforeEach } from "vitest"; +import { + convertImageToBase64, + createChatMultimodalMessage, + createChatDisplayMessage, + shouldShowChatAttachedImage, +} from "./ChatImageUtils"; +import { MessageType } from "./types"; + +describe("ChatImageUtils", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + describe("convertImageToBase64", () => { + it("should convert file to base64 data URI", async () => { + const file = new File(["test content"], "test.png", { type: "image/png" }); + const result = await convertImageToBase64(file); + expect(result).toMatch(/^data:image\/png;base64,/); + }); + + it("should handle different file types", async () => { + const jpegFile = new File(["jpeg content"], "test.jpg", { type: "image/jpeg" }); + const result = await convertImageToBase64(jpegFile); + expect(result).toMatch(/^data:image\/jpeg;base64,/); + }); + + it("should reject on file read error", async () => { + const file = new File(["test"], "test.png", { type: "image/png" }); + const originalReadAsDataURL = FileReader.prototype.readAsDataURL; + + FileReader.prototype.readAsDataURL = vi.fn(function (this: FileReader) { + setTimeout(() => { + if (this.onerror) { + this.onerror(new Error("Read error") as any); + } + }, 0); + }); + + await expect(convertImageToBase64(file)).rejects.toThrow(); + + FileReader.prototype.readAsDataURL = originalReadAsDataURL; + }); + }); + + describe("createChatMultimodalMessage", () => { + it("should create multimodal message with text and image", async () => { + const file = new File(["test content"], "test.png", { type: "image/png" }); + const inputMessage = "What is in this image?"; + + const result = await createChatMultimodalMessage(inputMessage, file); + + expect(result.role).toBe("user"); + expect(result.content).toHaveLength(2); + expect(result.content[0]).toEqual({ type: "text", text: inputMessage }); + expect(result.content[1]).toMatchObject({ + type: "image_url", + image_url: { + url: expect.stringMatching(/^data:image\/png;base64,/), + }, + }); + }); + + it("should include base64 data URI in image_url", async () => { + const file = new File(["test content"], "test.png", { type: "image/png" }); + const result = await createChatMultimodalMessage("test", file); + + const imageContent = result.content[1]; + expect(imageContent.type).toBe("image_url"); + if ("image_url" in imageContent && imageContent.image_url) { + expect(imageContent.image_url.url).toMatch(/^data:/); + } + }); + }); + + describe("createChatDisplayMessage", () => { + it("should create display message without file", () => { + const result = createChatDisplayMessage("Hello world", false); + + expect(result.role).toBe("user"); + expect(result.content).toBe("Hello world"); + expect(result.imagePreviewUrl).toBeUndefined(); + }); + + it("should create display message with PDF file", () => { + const filePreviewUrl = "blob:test-url"; + const result = createChatDisplayMessage("Read this", true, filePreviewUrl, "document.pdf"); + + expect(result.content).toBe("Read this [PDF attached]"); + expect(result.imagePreviewUrl).toBe(filePreviewUrl); + }); + + it("should create display message with image file", () => { + const filePreviewUrl = "blob:test-url"; + const result = createChatDisplayMessage("Look at this", true, filePreviewUrl, "photo.jpg"); + + expect(result.content).toBe("Look at this [Image attached]"); + expect(result.imagePreviewUrl).toBe(filePreviewUrl); + }); + + it("should create display message with file but no fileName", () => { + const filePreviewUrl = "blob:test-url"; + const result = createChatDisplayMessage("Check this", true, filePreviewUrl); + + expect(result.content).toBe("Check this "); + expect(result.imagePreviewUrl).toBe(filePreviewUrl); + }); + + it("should create display message with file but no preview URL", () => { + const result = createChatDisplayMessage("See this", true, undefined, "image.png"); + + expect(result.content).toBe("See this [Image attached]"); + expect(result.imagePreviewUrl).toBeUndefined(); + }); + }); + + describe("shouldShowChatAttachedImage", () => { + it("should return true for user message with image attachment", () => { + const message: MessageType = { + role: "user", + content: "Check this [Image attached]", + imagePreviewUrl: "blob:test-url", + }; + + expect(shouldShowChatAttachedImage(message)).toBe(true); + }); + + it("should return true for user message with PDF attachment", () => { + const message: MessageType = { + role: "user", + content: "Read this [PDF attached]", + imagePreviewUrl: "blob:test-url", + }; + + expect(shouldShowChatAttachedImage(message)).toBe(true); + }); + + it("should return false for assistant message", () => { + const message: MessageType = { + role: "assistant", + content: "Here is the image [Image attached]", + imagePreviewUrl: "blob:test-url", + }; + + expect(shouldShowChatAttachedImage(message)).toBe(false); + }); + + it("should return false when content is not a string", () => { + const message: MessageType = { + role: "user", + content: [{ type: "input_text", text: "test" }], + imagePreviewUrl: "blob:test-url", + }; + + expect(shouldShowChatAttachedImage(message)).toBe(false); + }); + + it("should return false when content does not include attachment marker", () => { + const message: MessageType = { + role: "user", + content: "Just regular text", + imagePreviewUrl: "blob:test-url", + }; + + expect(shouldShowChatAttachedImage(message)).toBe(false); + }); + + it("should return false when imagePreviewUrl is missing", () => { + const message: MessageType = { + role: "user", + content: "Check this [Image attached]", + }; + + expect(shouldShowChatAttachedImage(message)).toBe(false); + }); + + it("should return false when imagePreviewUrl is empty string", () => { + const message: MessageType = { + role: "user", + content: "Check this [Image attached]", + imagePreviewUrl: "", + }; + + expect(shouldShowChatAttachedImage(message)).toBe(false); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx index cee9db74276..f747b5b2599 100644 --- a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx @@ -84,10 +84,7 @@ interface ChatUIProps { }; } -const MCP_SUPPORTED_ENDPOINTS = new Set([ - EndpointType.CHAT, - EndpointType.RESPONSES, -]); +const MCP_SUPPORTED_ENDPOINTS = new Set([EndpointType.CHAT, EndpointType.RESPONSES]); const ChatUI: React.FC = ({ accessToken, @@ -130,6 +127,9 @@ const ChatUI: React.FC = ({ return disabledPersonalKeyCreation ? "custom" : "session"; }); const [apiKey, setApiKey] = useState(() => sessionStorage.getItem("apiKey") || ""); + const [customProxyBaseUrl, setCustomProxyBaseUrl] = useState( + () => sessionStorage.getItem("customProxyBaseUrl") || "", + ); const [inputMessage, setInputMessage] = useState(""); const [chatHistory, setChatHistory] = useState(() => { try { @@ -212,7 +212,7 @@ const ChatUI: React.FC = ({ const [temperature, setTemperature] = useState(1.0); const [maxTokens, setMaxTokens] = useState(2048); const [useAdvancedParams, setUseAdvancedParams] = useState(false); - + // Code Interpreter state (using custom hook) const codeInterpreter = useCodeInterpreter(); @@ -241,16 +241,15 @@ const ChatUI: React.FC = ({ try { const response = await listMCPTools(userApiKey, serverId); - setServerToolsMap(prev => ({ + setServerToolsMap((prev) => ({ ...prev, - [serverId]: response.tools || [] + [serverId]: response.tools || [], })); } catch (error) { console.error(`Error fetching tools for server ${serverId}:`, error); } }; - useEffect(() => { if (isGetCodeModalVisible) { const code = generateCodeSnippet({ @@ -392,7 +391,7 @@ const ChatUI: React.FC = ({ const loadAgents = async () => { try { - const agents = await fetchAvailableAgents(userApiKey); + const agents = await fetchAvailableAgents(userApiKey, customProxyBaseUrl || undefined); setAgentInfo(agents); // Clear selection if current agent not in list if (selectedAgent && !agents.some((a) => a.agent_name === selectedAgent)) { @@ -404,7 +403,7 @@ const ChatUI: React.FC = ({ }; loadAgents(); - }, [accessToken, apiKeySource, apiKey, endpointType]); + }, [accessToken, apiKeySource, apiKey, endpointType, customProxyBaseUrl, selectedAgent]); useEffect(() => { // Scroll to the bottom of the chat whenever chatHistory updates @@ -900,6 +899,7 @@ const ChatUI: React.FC = ({ useAdvancedParams ? temperature : undefined, useAdvancedParams ? maxTokens : undefined, updateTotalLatency, + customProxyBaseUrl || undefined, mcpServers, mcpServerToolRestrictions, ); @@ -912,6 +912,7 @@ const ChatUI: React.FC = ({ effectiveApiKey, selectedTags, signal, + customProxyBaseUrl || undefined, ); } else if (endpointType === EndpointType.SPEECH) { // For audio speech @@ -923,6 +924,9 @@ const ChatUI: React.FC = ({ effectiveApiKey, selectedTags, signal, + undefined, // responseFormat + undefined, // speed + customProxyBaseUrl || undefined, ); } else if (endpointType === EndpointType.IMAGE_EDITS) { // For image edits @@ -935,6 +939,7 @@ const ChatUI: React.FC = ({ effectiveApiKey, selectedTags, signal, + customProxyBaseUrl || undefined, ); } } else if (endpointType === EndpointType.RESPONSES) { @@ -973,6 +978,7 @@ const ChatUI: React.FC = ({ handleMCPEvent, // Pass MCP event handler codeInterpreter.enabled, // Enable Code Interpreter tool codeInterpreter.setResult, // Handle code interpreter output + customProxyBaseUrl || undefined, mcpServers, mcpServerToolRestrictions, ); @@ -997,6 +1003,8 @@ const ChatUI: React.FC = ({ traceId, selectedVectorStores.length > 0 ? selectedVectorStores : undefined, selectedGuardrails.length > 0 ? selectedGuardrails : undefined, + selectedMCPServers, // Pass the selected tools array + customProxyBaseUrl || undefined, ); } else if (endpointType === EndpointType.EMBEDDINGS) { await makeOpenAIEmbeddingsRequest( @@ -1005,6 +1013,7 @@ const ChatUI: React.FC = ({ selectedModel, effectiveApiKey, selectedTags, + customProxyBaseUrl || undefined, ); } else if (endpointType === EndpointType.TRANSCRIPTION) { // For audio transcriptions @@ -1016,6 +1025,11 @@ const ChatUI: React.FC = ({ effectiveApiKey, selectedTags, signal, + undefined, // language + undefined, // prompt + undefined, // responseFormat + undefined, // temperature + customProxyBaseUrl || undefined, ); } } @@ -1032,6 +1046,7 @@ const ChatUI: React.FC = ({ updateTimingData, updateTotalLatency, updateA2AMetadata, + customProxyBaseUrl || undefined, ); } } catch (error) { @@ -1156,6 +1171,40 @@ const ChatUI: React.FC = ({ )} +
+
+ + Custom Proxy Base URL + + {customProxyBaseUrl && ( + + )} +
+ { + setCustomProxyBaseUrl(value); + sessionStorage.setItem("customProxyBaseUrl", value); + }} + value={customProxyBaseUrl} + icon={ApiOutlined} + /> + {customProxyBaseUrl && ( + API calls will be sent to: {customProxyBaseUrl} + )} +
+
Endpoint Type @@ -1327,7 +1376,11 @@ const ChatUI: React.FC = ({ optionLabelProp="label" > {agentInfo.map((agent) => ( - +
{agent.agent_name || agent.agent_id} {agent.agent_card_params?.description && ( @@ -1361,10 +1414,7 @@ const ChatUI: React.FC = ({
MCP Servers - + @@ -1380,15 +1430,15 @@ const ChatUI: React.FC = ({ } else { setSelectedMCPServers(value); // Clean up tool restrictions for removed servers - setMCPServerToolRestrictions(prev => { + setMCPServerToolRestrictions((prev) => { const updated = { ...prev }; - Object.keys(updated).forEach(serverId => { + Object.keys(updated).forEach((serverId) => { if (!value.includes(serverId)) delete updated[serverId]; }); return updated; }); // Load tools for newly selected servers - value.forEach(serverId => { + value.forEach((serverId) => { if (!serverToolsMap[serverId]) { loadServerTools(serverId); } @@ -1419,12 +1469,8 @@ const ChatUI: React.FC = ({ disabled={selectedMCPServers.includes("__all__")} >
- - {server.alias || server.server_name || server.server_id} - - {server.description && ( - {server.description} - )} + {server.alias || server.server_name || server.server_id} + {server.description && {server.description}}
))} @@ -1434,40 +1480,40 @@ const ChatUI: React.FC = ({ {selectedMCPServers.length > 0 && !selectedMCPServers.includes("__all__") && MCP_SUPPORTED_ENDPOINTS.has(endpointType as EndpointType) && ( -
- {selectedMCPServers.map(serverId => { - const server = mcpServers.find(s => s.server_id === serverId); - const tools = serverToolsMap[serverId] || []; - if (tools.length === 0) return null; +
+ {selectedMCPServers.map((serverId) => { + const server = mcpServers.find((s) => s.server_id === serverId); + const tools = serverToolsMap[serverId] || []; + if (tools.length === 0) return null; - return ( -
- - Limit tools for {server?.alias || server?.server_name || serverId}: - - { + setMCPServerToolRestrictions((prev) => ({ + ...prev, + [serverId]: selectedTools, + })); + }} + options={tools.map((tool) => ({ + value: tool.name, + label: tool.name, + }))} + maxTagCount={2} + /> +
+ ); + })} +
+ )}
@@ -2018,7 +2064,13 @@ const ChatUI: React.FC = ({ )} {/* Quick Code Interpreter toggle for Responses */} {endpointType === EndpointType.RESPONSES && ( - +