Merge branch 'main' of https://github.com/BerriAI/litellm into litellm_mem_issues_0000001

This commit is contained in:
Alexsander Hamir 2026-01-09 17:13:28 -08:00
commit 8177ce107b
100 changed files with 10093 additions and 2032 deletions

View file

@ -13,6 +13,7 @@ on:
jobs:
publish-migrations:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
services:
postgres:

View file

@ -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.
<img src="https://contrib.rocks/image?repo=BerriAI/litellm" />
</a>

View file

@ -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.
:::

View file

@ -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/)

View file

@ -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

View file

@ -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) |
<br />
## 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())
```

View file

@ -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

View file

@ -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)

View file

@ -67,7 +67,7 @@ Set `litellm.turn_off_message_logging=True` This will prevent the messages and r
<TabItem value="global" label="Global">
**1. Setup config.yaml **
**1. Setup config.yaml**
```yaml
model_list:
- model_name: gpt-3.5-turbo

View file

@ -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",

View file

@ -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,

View file

@ -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",

View file

@ -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)}")

View file

@ -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"
}
}
"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"]
}
}

View file

@ -45,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
@ -238,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",
@ -330,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
@ -801,7 +851,7 @@ class PrometheusLogger(CustomLogger):
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"
)
@ -821,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", {})
@ -832,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 {}),
@ -951,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.
@ -1020,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],
@ -1188,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
@ -1208,7 +1328,7 @@ class PrometheusLogger(CustomLogger):
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"
)
@ -1562,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"
@ -1595,12 +1714,11 @@ 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
@ -1743,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 = ""
@ -2380,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))
@ -2455,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",

View file

@ -4839,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(

View file

@ -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

View file

@ -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,

View file

@ -314,6 +314,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 +329,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,7 +366,7 @@ 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"
@ -416,7 +418,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(".")
@ -962,7 +964,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

View file

@ -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

View file

@ -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
@ -359,6 +359,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 +458,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(

View file

@ -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

View file

@ -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)}",
)

View file

@ -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/"):

View file

@ -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(

View file

@ -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",

View file

@ -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,
)

View file

@ -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:

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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]:

View file

@ -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(

View file

@ -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={},

View file

@ -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

View file

@ -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]]:

View file

@ -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"}}
),
)

View file

@ -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,

View file

@ -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
@ -670,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,
@ -681,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,
@ -923,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":
@ -1178,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

View file

@ -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")

View file

@ -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(

View file

@ -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
):

View file

@ -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

View file

@ -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":

View file

@ -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

View file

@ -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

View file

@ -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
@ -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]
@ -3972,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(
@ -4109,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()
@ -4646,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,
@ -4915,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()
@ -6824,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,

View file

@ -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
@ -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"],

View file

@ -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

View file

@ -386,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] = {}
@ -430,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
@ -613,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:
@ -722,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(
@ -1071,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]
@ -1500,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
@ -3021,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}"
@ -3314,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:
@ -3617,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(
@ -4366,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
@ -4402,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
@ -4485,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
@ -5669,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"
@ -6235,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(
@ -6931,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 = (
@ -6974,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
@ -7664,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 (
@ -7769,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

View file

@ -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)

View file

@ -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)

View file

@ -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:
@ -7941,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

File diff suppressed because it is too large Load diff

View file

@ -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,

View file

@ -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"

View file

@ -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"
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]

View file

@ -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

View file

@ -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")

View file

@ -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

View file

@ -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

View file

@ -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"])

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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"])

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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!")

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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}

View file

@ -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"""

View file

@ -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

View file

@ -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())

View file

@ -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:

View file

@ -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):
"""

View file

@ -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

View file

@ -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
@ -680,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():
@ -702,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
@ -719,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
@ -732,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 == []
@ -743,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 = {
@ -760,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():
@ -780,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"
@ -790,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",
{},
@ -807,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
@ -820,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
@ -851,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", {}
)
@ -888,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
@ -905,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]}"

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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<string, Page> = {
"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,
};

View file

@ -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",
}

View file

@ -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<void> {
await page.goto(`/ui?page=${pageEnum}`);
}

View file

@ -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();

View file

@ -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}(&|$)`));
});
});
}

View file

@ -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");
});
});

View file

@ -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<string, string>;
}
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<RouterFieldsResponse> => {
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<RouterFieldsResponse> => {
const { accessToken, userId, userRole } = useAuthorized();
return useQuery<RouterFieldsResponse>({
queryKey: routerFieldsKeys.detail("fields"),
queryFn: async () => await getRouterFields(accessToken!),
enabled: Boolean(accessToken && userId && userRole),
});
};

View file

@ -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);
});
});
});

View file

@ -0,0 +1,326 @@
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest";
import CodeInterpreterOutput from "./CodeInterpreterOutput";
vi.mock("@/components/networking", () => ({
getProxyBaseUrl: vi.fn(() => "https://example.com"),
}));
global.fetch = vi.fn();
describe("CodeInterpreterOutput", () => {
beforeEach(() => {
vi.clearAllMocks();
URL.createObjectURL = vi.fn((blob) => `blob:${blob}`);
URL.revokeObjectURL = vi.fn();
});
afterEach(() => {
vi.restoreAllMocks();
});
it("should render", () => {
render(<CodeInterpreterOutput code="print('hello')" accessToken="test-token" />);
expect(screen.getByText("Python Code Executed")).toBeInTheDocument();
});
it("should display code in syntax highlighter", async () => {
const user = userEvent.setup();
const code = "print('hello world')";
const { container } = render(<CodeInterpreterOutput code={code} accessToken="test-token" />);
expect(screen.getByText("Python Code Executed")).toBeInTheDocument();
const collapseHeader = screen.getByRole("button");
await user.click(collapseHeader);
await waitFor(() => {
const codeElement = container.querySelector("code.language-python");
expect(codeElement).toBeInTheDocument();
expect(codeElement?.textContent).toContain(code);
});
});
it("should fetch and display images from annotations", async () => {
const mockBlob = new Blob(["image data"], { type: "image/png" });
const mockResponse = {
ok: true,
blob: vi.fn().mockResolvedValue(mockBlob),
};
(global.fetch as any).mockResolvedValue(mockResponse);
const annotations = [
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-1",
filename: "chart.png",
start_index: 0,
end_index: 10,
},
];
render(
<CodeInterpreterOutput
code="import matplotlib.pyplot as plt"
annotations={annotations}
accessToken="test-token"
/>,
);
await waitFor(() => {
expect(global.fetch).toHaveBeenCalledWith(
"https://example.com/v1/containers/container-1/files/file-1/content",
expect.objectContaining({
headers: {
Authorization: "Bearer test-token",
},
}),
);
});
await waitFor(() => {
expect(screen.getByText("chart.png")).toBeInTheDocument();
});
});
it("should show loading state while fetching images", async () => {
const mockBlob = new Blob(["image data"], { type: "image/png" });
let resolveBlob: (value: Blob) => void;
const blobPromise = new Promise<Blob>((resolve) => {
resolveBlob = resolve;
});
const mockResponse = {
ok: true,
blob: vi.fn().mockReturnValue(blobPromise),
};
(global.fetch as any).mockResolvedValue(mockResponse);
const annotations = [
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-1",
filename: "chart.png",
start_index: 0,
end_index: 10,
},
];
render(
<CodeInterpreterOutput
code="import matplotlib.pyplot as plt"
annotations={annotations}
accessToken="test-token"
/>,
);
await waitFor(() => {
expect(screen.getByText("Loading image...")).toBeInTheDocument();
});
resolveBlob!(mockBlob);
await waitFor(() => {
expect(screen.queryByText("Loading image...")).not.toBeInTheDocument();
});
});
it("should handle download for image files", async () => {
const user = userEvent.setup();
const mockBlob = new Blob(["image data"], { type: "image/png" });
const mockResponse = {
ok: true,
blob: vi.fn().mockResolvedValue(mockBlob),
};
(global.fetch as any).mockResolvedValue(mockResponse);
const annotations = [
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-1",
filename: "chart.png",
start_index: 0,
end_index: 10,
},
];
const createElementSpy = vi.spyOn(document, "createElement");
const appendChildSpy = vi.spyOn(document.body, "appendChild");
const removeChildSpy = vi.spyOn(document.body, "removeChild");
render(
<CodeInterpreterOutput
code="import matplotlib.pyplot as plt"
annotations={annotations}
accessToken="test-token"
/>,
);
await waitFor(() => {
expect(screen.getByText("chart.png")).toBeInTheDocument();
});
const downloadButton = screen.getByText("Download");
await user.click(downloadButton);
await waitFor(() => {
expect(global.fetch).toHaveBeenCalledWith(
"https://example.com/v1/containers/container-1/files/file-1/content",
expect.objectContaining({
headers: {
Authorization: "Bearer test-token",
},
}),
);
});
createElementSpy.mockRestore();
appendChildSpy.mockRestore();
removeChildSpy.mockRestore();
});
it("should handle download for non-image files", async () => {
const user = userEvent.setup();
const mockBlob = new Blob(["file data"], { type: "text/plain" });
const mockResponse = {
ok: true,
blob: vi.fn().mockResolvedValue(mockBlob),
};
(global.fetch as any).mockResolvedValue(mockResponse);
const annotations = [
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-1",
filename: "data.csv",
start_index: 0,
end_index: 10,
},
];
render(<CodeInterpreterOutput code="import pandas as pd" annotations={annotations} accessToken="test-token" />);
await waitFor(() => {
expect(screen.getByText("data.csv")).toBeInTheDocument();
});
const downloadButton = screen.getByText("data.csv").closest("button");
expect(downloadButton).toBeInTheDocument();
if (downloadButton) {
await user.click(downloadButton);
}
await waitFor(() => {
expect(global.fetch).toHaveBeenCalledWith(
"https://example.com/v1/containers/container-1/files/file-1/content",
expect.objectContaining({
headers: {
Authorization: "Bearer test-token",
},
}),
);
});
});
it("should return null when no code and no annotations", () => {
const { container } = render(<CodeInterpreterOutput accessToken="test-token" />);
expect(container.firstChild).toBeNull();
});
it("should handle multiple image formats", async () => {
const mockBlob = new Blob(["image data"], { type: "image/png" });
const mockResponse = {
ok: true,
blob: vi.fn().mockResolvedValue(mockBlob),
};
(global.fetch as any).mockResolvedValue(mockResponse);
const annotations = [
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-1",
filename: "image.png",
start_index: 0,
end_index: 10,
},
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-2",
filename: "image.jpg",
start_index: 0,
end_index: 10,
},
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-3",
filename: "image.jpeg",
start_index: 0,
end_index: 10,
},
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-4",
filename: "image.gif",
start_index: 0,
end_index: 10,
},
];
render(
<CodeInterpreterOutput
code="import matplotlib.pyplot as plt"
annotations={annotations}
accessToken="test-token"
/>,
);
await waitFor(() => {
expect(global.fetch).toHaveBeenCalledTimes(4);
});
});
it("should handle fetch errors gracefully", async () => {
const consoleErrorSpy = vi.spyOn(console, "error").mockImplementation(() => {});
(global.fetch as any).mockRejectedValue(new Error("Network error"));
const annotations = [
{
type: "container_file_citation" as const,
container_id: "container-1",
file_id: "file-1",
filename: "chart.png",
start_index: 0,
end_index: 10,
},
];
render(
<CodeInterpreterOutput
code="import matplotlib.pyplot as plt"
annotations={annotations}
accessToken="test-token"
/>,
);
await waitFor(() => {
expect(consoleErrorSpy).toHaveBeenCalled();
});
consoleErrorSpy.mockRestore();
});
});

View file

@ -0,0 +1,161 @@
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi } from "vitest";
import { UnifiedSelector } from "./UnifiedSelector";
import { EndpointId, ENDPOINT_CONFIGS } from "../endpoint_config";
describe("UnifiedSelector", () => {
it("should render", () => {
const onChange = vi.fn();
const options = [
{ value: "option1", label: "Option 1" },
{ value: "option2", label: "Option 2" },
];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
render(<UnifiedSelector value="" options={options} loading={false} config={config} onChange={onChange} />);
const select = screen.getByRole("combobox");
expect(select).toBeInTheDocument();
});
it("should display placeholder when not loading", () => {
const onChange = vi.fn();
const options = [{ value: "option1", label: "Option 1" }];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
const { container } = render(
<UnifiedSelector value="" options={options} loading={false} config={config} onChange={onChange} />,
);
const placeholder = container.querySelector(".ant-select-selection-placeholder");
expect(placeholder).toHaveTextContent(config.selectorPlaceholder);
});
it("should display loading placeholder when loading", () => {
const onChange = vi.fn();
const options = [{ value: "option1", label: "Option 1" }];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
const { container } = render(
<UnifiedSelector value="" options={options} loading={true} config={config} onChange={onChange} />,
);
const placeholder = container.querySelector(".ant-select-selection-placeholder");
expect(placeholder).toHaveTextContent(`Loading ${config.selectorLabel.toLowerCase()}s...`);
});
it("should call onChange when option is selected", async () => {
const user = userEvent.setup();
const onChange = vi.fn();
const options = [
{ value: "option1", label: "Option 1" },
{ value: "option2", label: "Option 2" },
];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
render(<UnifiedSelector value="" options={options} loading={false} config={config} onChange={onChange} />);
const select = screen.getByRole("combobox");
await user.click(select);
await waitFor(() => {
const option = screen.getByText("Option 1");
expect(option).toBeInTheDocument();
});
const option = screen.getByText("Option 1");
await user.click(option);
await waitFor(() => {
expect(onChange).toHaveBeenCalled();
});
const callArgs = onChange.mock.calls[0];
expect(callArgs[0]).toBe("option1");
});
it("should display selected value", () => {
const onChange = vi.fn();
const options = [
{ value: "option1", label: "Option 1" },
{ value: "option2", label: "Option 2" },
];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
const { container } = render(
<UnifiedSelector value="option1" options={options} loading={false} config={config} onChange={onChange} />,
);
const selectedValue = container.querySelector(".ant-select-selection-item");
expect(selectedValue).toHaveTextContent("Option 1");
});
it("should filter options by search input", async () => {
const user = userEvent.setup();
const onChange = vi.fn();
const options = [
{ value: "option1", label: "Option One" },
{ value: "option2", label: "Option Two" },
{ value: "option3", label: "Different" },
];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
render(<UnifiedSelector value="" options={options} loading={false} config={config} onChange={onChange} />);
const select = screen.getByRole("combobox");
await user.click(select);
await user.type(select, "One");
await waitFor(() => {
expect(screen.getByText("Option One")).toBeInTheDocument();
expect(screen.queryByText("Option Two")).not.toBeInTheDocument();
expect(screen.queryByText("Different")).not.toBeInTheDocument();
});
});
it("should show loading spinner in notFoundContent when loading", async () => {
const user = userEvent.setup();
const onChange = vi.fn();
const options: { value: string; label: string }[] = [];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
render(<UnifiedSelector value="" options={options} loading={true} config={config} onChange={onChange} />);
const select = screen.getByRole("combobox");
await user.click(select);
await waitFor(() => {
const spin = document.querySelector(".ant-spin");
expect(spin).toBeInTheDocument();
});
});
it("should show no options message when not loading and no options", async () => {
const user = userEvent.setup();
const onChange = vi.fn();
const options: { value: string; label: string }[] = [];
const config = ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS];
render(<UnifiedSelector value="" options={options} loading={false} config={config} onChange={onChange} />);
const select = screen.getByRole("combobox");
await user.click(select);
await waitFor(() => {
expect(screen.getByText(`No ${config.selectorLabel.toLowerCase()}s available`)).toBeInTheDocument();
});
});
it("should work with agent endpoint config", () => {
const onChange = vi.fn();
const options = [{ value: "agent1", label: "Agent One" }];
const config = ENDPOINT_CONFIGS[EndpointId.A2A_AGENTS];
const { container } = render(
<UnifiedSelector value="" options={options} loading={false} config={config} onChange={onChange} />,
);
const placeholder = container.querySelector(".ant-select-selection-placeholder");
expect(placeholder).toHaveTextContent(config.selectorPlaceholder);
});
});

View file

@ -0,0 +1,104 @@
import { describe, expect, it } from "vitest";
import {
EndpointId,
ENDPOINT_CONFIGS,
getAvailableEndpoints,
getEndpointConfig,
isAgentEndpoint,
isModelEndpoint,
modelOptionsToSelectorOptions,
agentOptionsToSelectorOptions,
getSelectionFieldName,
getComparisonSelection,
hasValidSelection,
} from "./endpoint_config";
import { Agent } from "../llm_calls/fetch_agents";
describe("endpoint_config", () => {
it("should export EndpointId constants", () => {
expect(EndpointId.CHAT_COMPLETIONS).toBe("/v1/chat/completions");
expect(EndpointId.A2A_AGENTS).toBe("/a2a");
});
it("should have endpoint configs for all endpoint IDs", () => {
expect(ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS]).toBeDefined();
expect(ENDPOINT_CONFIGS[EndpointId.A2A_AGENTS]).toBeDefined();
expect(ENDPOINT_CONFIGS[EndpointId.CHAT_COMPLETIONS].selectorType).toBe("model");
expect(ENDPOINT_CONFIGS[EndpointId.A2A_AGENTS].selectorType).toBe("agent");
});
it("should get available endpoints", () => {
const endpoints = getAvailableEndpoints();
expect(endpoints).toHaveLength(2);
expect(endpoints).toContainEqual({
value: EndpointId.CHAT_COMPLETIONS,
label: "/v1/chat/completions",
});
expect(endpoints).toContainEqual({
value: EndpointId.A2A_AGENTS,
label: "/a2a (Agents)",
});
});
it("should get endpoint config by ID", () => {
const config = getEndpointConfig(EndpointId.CHAT_COMPLETIONS);
expect(config.id).toBe(EndpointId.CHAT_COMPLETIONS);
expect(config.selectorType).toBe("model");
expect(config.selectorLabel).toBe("Model");
});
it("should check if endpoint is agent endpoint", () => {
expect(isAgentEndpoint(EndpointId.A2A_AGENTS)).toBe(true);
expect(isAgentEndpoint(EndpointId.CHAT_COMPLETIONS)).toBe(false);
});
it("should check if endpoint is model endpoint", () => {
expect(isModelEndpoint(EndpointId.CHAT_COMPLETIONS)).toBe(true);
expect(isModelEndpoint(EndpointId.A2A_AGENTS)).toBe(false);
});
it("should convert model options to selector options", () => {
const models = ["gpt-4", "gpt-3.5-turbo", "claude-3"];
const options = modelOptionsToSelectorOptions(models);
expect(options).toHaveLength(3);
expect(options[0]).toEqual({ value: "gpt-4", label: "gpt-4" });
expect(options[1]).toEqual({ value: "gpt-3.5-turbo", label: "gpt-3.5-turbo" });
expect(options[2]).toEqual({ value: "claude-3", label: "claude-3" });
});
it("should convert agent options to selector options", () => {
const agents: Agent[] = [
{ agent_id: "agent-1", agent_name: "Agent One" },
{ agent_id: "agent-2", agent_name: "Agent Two" },
{ agent_id: "agent-3", agent_name: undefined as any },
];
const options = agentOptionsToSelectorOptions(agents);
expect(options).toHaveLength(3);
expect(options[0]).toEqual({ value: "Agent One", label: "Agent One" });
expect(options[1]).toEqual({ value: "Agent Two", label: "Agent Two" });
expect(options[2]).toEqual({ value: undefined, label: "agent-3" });
});
it("should get selection field name based on endpoint", () => {
expect(getSelectionFieldName(EndpointId.CHAT_COMPLETIONS)).toBe("model");
expect(getSelectionFieldName(EndpointId.A2A_AGENTS)).toBe("agent");
});
it("should get comparison selection based on endpoint", () => {
const comparison = { model: "gpt-4", agent: "agent-1" };
expect(getComparisonSelection(comparison, EndpointId.CHAT_COMPLETIONS)).toBe("gpt-4");
expect(getComparisonSelection(comparison, EndpointId.A2A_AGENTS)).toBe("agent-1");
});
it("should check if comparison has valid selection", () => {
const comparisonWithModel = { model: "gpt-4", agent: "" };
const comparisonWithAgent = { model: "", agent: "agent-1" };
const comparisonEmpty = { model: "", agent: "" };
const comparisonWhitespace = { model: " ", agent: "" };
expect(hasValidSelection(comparisonWithModel, EndpointId.CHAT_COMPLETIONS)).toBe(true);
expect(hasValidSelection(comparisonWithAgent, EndpointId.A2A_AGENTS)).toBe(true);
expect(hasValidSelection(comparisonEmpty, EndpointId.CHAT_COMPLETIONS)).toBe(false);
expect(hasValidSelection(comparisonWhitespace, EndpointId.CHAT_COMPLETIONS)).toBe(false);
});
});