Merge branch 'BerriAI:main' into main

This commit is contained in:
Mubashir Osmani 2025-10-02 21:40:11 -04:00 • committed by GitHub
commit 42ed6ad907
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
151 changed files with 12130 additions and 1712 deletions

View file

@ -273,7 +273,7 @@ echo 'LITELLM_SALT_KEY="sk-1234"' >> .env
source .env
# Start
docker-compose up
docker compose up
```

View file

@ -28,7 +28,7 @@ Replace `your-secret-key` with a strong, randomly generated secret.
Once you have set the `MASTER_KEY`, you can build and run the containers using the following command:
```bash
docker-compose up -d --build
docker compose up -d --build
```
This command will:
@ -42,13 +42,13 @@ This command will:
You can check the status of the running containers with the following command:
```bash
docker-compose ps
docker compose ps
```
To view the logs of the `litellm` container, run:
```bash
docker-compose logs -f litellm
docker compose logs -f litellm
```
### 4. Stopping the Application
@ -56,7 +56,7 @@ docker-compose logs -f litellm
To stop the running containers, use the following command:
```bash
docker-compose down
docker compose down
```
## Troubleshooting

View file

@ -13,9 +13,6 @@ git clone https://github.com/BerriAI/litellm.git
Tell the proxy where the UI is located
```bash
export PROXY_BASE_URL="http://localhost:3000/"
### ALSO ### - set the basic env variables
DATABASE_URL = "postgresql://<user>:<password>@<host>:<port>/<dbname>"
LITELLM_MASTER_KEY = "sk-1234"
STORE_MODEL_IN_DB = "True"
@ -30,7 +27,7 @@ python3 proxy_cli.py --config /path/to/config.yaml --port 4000
Set the mode as development (this will assume the proxy is running on localhost:4000)
```bash
export NODE_ENV="development"
npm install # install dependencies
```
```bash

View file

@ -1045,6 +1045,25 @@ curl --location 'http://localhost:4000/github_mcp/mcp' \
---
## MCP Oauth
LiteLLM v 1.77.6 added support for OAuth 2.0 Client Credentials for MCP servers.
This configuration is currently available on the config.yaml, with UI support coming soon.
```yaml
mcp_servers:
github_mcp:
url: "https://api.githubcopilot.com/mcp"
auth_type: oauth2
authorization_url: https://github.com/login/oauth/authorize
token_url: https://github.com/login/oauth/access_token
client_id: os.environ/GITHUB_OAUTH_CLIENT_ID
client_secret: os.environ/GITHUB_OAUTH_CLIENT_SECRET
scopes: ["public_repo", "user:email"]
```
## Using your MCP with client side credentials
Use this if you want to pass a client side authentication token to LiteLLM to then pass to your MCP to auth to your MCP.

View file

@ -15,10 +15,11 @@ Pass-through endpoints for Vertex AI - call provider-specific endpoint, in nativ
## Supported Endpoints
LiteLLM supports 2 vertex ai passthrough routes:
LiteLLM supports 3 vertex ai passthrough routes:
1. `/vertex_ai` → routes to `https://{vertex_location}-aiplatform.googleapis.com/`
2. `/vertex_ai/discovery` → routes to [`https://discoveryengine.googleapis.com`](https://discoveryengine.googleapis.com/)
3. `/vertex_ai/live` → upgrades to the Vertex AI Live API WebSocket (`google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent`)
## How to use
@ -170,6 +171,50 @@ generateContent();
</Tabs>
## Vertex AI Live API WebSocket
LiteLLM can now proxy the Vertex AI Live API to help you experiment with streaming audio/text from Gemini Live models without exposing Google credentials to clients.
- Configure default Vertex credentials via `default_vertex_config` or environment variables (see examples above).
- Connect to `wss://<PROXY_URL>/vertex_ai/live`. LiteLLM will exchange your saved credentials for a short-lived access token and forward messages bidirectionally.
- Optional query params `vertex_project`, `vertex_location`, and `model` let you override defaults for multi-project setups or global-only models.
```python title="client.py"
import asyncio
import json
from websockets.asyncio.client import connect
async def main() -> None:
headers = {
"x-litellm-api-key": "Bearer sk-your-litellm-key",
"Content-Type": "application/json",
}
async with connect(
"ws://localhost:4000/vertex_ai/live",
additional_headers=headers,
) as ws:
await ws.send(
json.dumps(
{
"setup": {
"model": "projects/your-project/locations/us-central1/publishers/google/models/gemini-2.0-flash-live-preview-04-09",
"generation_config": {"response_modalities": ["TEXT"]},
}
}
)
)
async for message in ws:
print("server:", message)
if __name__ == "__main__":
asyncio.run(main())
```
## Quick Start
Let's call the Vertex AI [`/generateContent` endpoint](https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/inference)
@ -415,4 +460,4 @@ generateContent();
```
</TabItem>
</Tabs>
</Tabs>

View file

@ -0,0 +1,284 @@
# Vertex AI Live API WebSocket Passthrough
LiteLLM now supports WebSocket passthrough for the Vertex AI Live API, enabling real-time bidirectional communication with Gemini models.
## Overview
The Vertex AI Live API WebSocket passthrough allows you to:
- Connect to Vertex AI Live API through LiteLLM proxy
- Use existing Vertex AI authentication methods
- Pass through all WebSocket messages bidirectionally
- Support text, audio, video, and multimodal interactions
- Track costs automatically for all usage types
## Configuration
### Environment Variables
Set the following environment variables for Vertex AI authentication:
```bash
# Required
DEFAULT_VERTEXAI_PROJECT=your-project-id
DEFAULT_VERTEXAI_LOCATION=us-central1
# Optional - use one of these for authentication
DEFAULT_GOOGLE_APPLICATION_CREDENTIALS=/path/to/service-account.json
# OR run: gcloud auth application-default login
```
### Configuration File
Alternatively, configure in your `config.yaml`:
```yaml
litellm_settings:
default_vertex_config:
vertex_project: "your-project-id"
vertex_location: "us-central1"
vertex_credentials: "os.environ/GOOGLE_APPLICATION_CREDENTIALS"
```
## Usage
### WebSocket Endpoints
- `ws://your-proxy-host/v1/vertex-ai/live`
- `ws://your-proxy-host/vertex-ai/live`
### Query Parameters
- `project_id` (optional): Google Cloud project ID (can be set in config)
- `location` (optional): Vertex AI location (can be set in config, default: us-central1)
### Example Connection
```javascript
// If project_id and location are set in config, you can connect without query params
const ws = new WebSocket('ws://localhost:4000/v1/vertex-ai/live');
// Or specify them explicitly
const ws = new WebSocket('ws://localhost:4000/v1/vertex-ai/live?project_id=your-project-id&location=us-central1');
```
## Cost Tracking
The WebSocket passthrough automatically tracks costs for all usage types based on the [Vertex AI pricing](https://cloud.google.com/vertex-ai/generative-ai/pricing#model-optimizer-pricing):
### Supported Cost Tracking
- **Text**: Character-based or token-based pricing depending on model
- **Audio**: Per-second pricing for audio input/output
- **Video**: Per-second pricing for video input
- **Images**: Per-image pricing for image input
### Cost Calculation
Costs are calculated using the same methods as other Vertex AI models in LiteLLM:
- Uses `cost_per_character` for Gemini models
- Uses `cost_per_token` for partner models (Claude, Llama, etc.)
- Includes audio, video, and image costs when applicable
### Cost Logging
Costs are automatically logged to:
- LiteLLM proxy logs
- Database (if configured)
- Spend tracking system
- Admin dashboard
Example log output:
```
Vertex AI Live WebSocket session cost: $0.001234 (input: $0.000800, output: $0.000434) tokens: 150, characters: 1200, duration: 45.2s
```
## API Reference
### Setup Message
Send this message first to initialize the session:
```json
{
"setup": {
"model": "projects/your-project-id/locations/us-central1/publishers/google/models/gemini-2.0-flash-live-preview-04-09",
"generation_config": {
"response_modalities": ["TEXT"]
}
}
}
```
### Text Input
```json
{
"client_content": {
"turns": [
{
"role": "user",
"parts": [{"text": "Hello! How are you?"}]
}
],
"turn_complete": true
}
}
```
### Audio Input
```json
{
"realtime_input": {
"media_chunks": [
{
"data": "base64-encoded-audio-data",
"mime_type": "audio/pcm"
}
]
}
}
```
## Supported Features
### Response Modalities
- **TEXT**: Text responses
- **AUDIO**: Audio responses with voice synthesis
### Tools
- **Function Calling**: Define and use custom functions
- **Code Execution**: Execute Python code
- **Google Search**: Search the web
- **Voice Activity Detection**: Detect when user is speaking
### Advanced Features
- **Audio Transcription**: Transcribe input and output audio
- **Proactive Audio**: Model responds only when relevant
- **Affective Dialog**: Understand emotional expressions
## Examples
### Python Client
```python
import asyncio
import json
import websockets
async def chat_with_gemini():
uri = "ws://localhost:4000/v1/vertex-ai/live?project_id=your-project-id"
async with websockets.connect(uri) as websocket:
# Setup
setup = {
"setup": {
"model": "projects/your-project-id/locations/us-central1/publishers/google/models/gemini-2.0-flash-live-preview-04-09",
"generation_config": {"response_modalities": ["TEXT"]}
}
}
await websocket.send(json.dumps(setup))
# Wait for setup response
response = await websocket.recv()
print(f"Setup: {response}")
# Send message
message = {
"client_content": {
"turns": [{"role": "user", "parts": [{"text": "Hello!"}]}],
"turn_complete": True
}
}
await websocket.send(json.dumps(message))
# Receive response
async for response in websocket:
print(f"Response: {response}")
# Check if turn is complete
data = json.loads(response)
if data.get("serverContent", {}).get("turnComplete"):
break
asyncio.run(chat_with_gemini())
```
### JavaScript Client
```javascript
const ws = new WebSocket('ws://localhost:4000/v1/vertex-ai/live?project_id=your-project-id');
ws.onopen = function() {
// Send setup
const setup = {
setup: {
model: "projects/your-project-id/locations/us-central1/publishers/google/models/gemini-2.0-flash-live-preview-04-09",
generation_config: { response_modalities: ["TEXT"] }
}
};
ws.send(JSON.stringify(setup));
};
ws.onmessage = function(event) {
const data = JSON.parse(event.data);
console.log('Received:', data);
// Check if setup is complete
if (data.setupComplete) {
// Send a message
const message = {
client_content: {
turns: [{ role: "user", parts: [{ text: "Hello!" }] }],
turn_complete: true
}
};
ws.send(JSON.stringify(message));
}
};
```
## Error Handling
The WebSocket connection may close with these codes:
- `4001`: Vertex AI credentials not configured
- `4002`: Project ID not provided
- `1011`: Internal server error
## Authentication
The WebSocket passthrough uses the same authentication as other LiteLLM endpoints:
1. **API Key**: Pass `Authorization: Bearer your-api-key` header
2. **Vertex AI Credentials**: Set environment variables or config file
## Limitations
- Requires valid Google Cloud project with Vertex AI API enabled
- WebSocket connections are not persistent across server restarts
- Rate limits apply based on your Google Cloud quotas
## Troubleshooting
### Common Issues
1. **Authentication Error**: Ensure Vertex AI credentials are properly configured
2. **Project Not Found**: Verify the project ID exists and has Vertex AI enabled
3. **Connection Refused**: Check that the LiteLLM proxy server is running
### Debug Mode
Enable debug logging to see detailed connection information:
```bash
export LITELLM_LOG=DEBUG
```
## Related Documentation
- [Vertex AI Live API Reference](https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-live)
- [LiteLLM Proxy Configuration](../proxy/)
- [Vertex AI Passthrough Endpoints](./vertex_ai.md)

View file

@ -0,0 +1,7 @@
# Railtracks
`Railtracks` is an open-source agentic framework that helps developers build resilient agentic systems offering local and remote monitoring tools.
- [Github](https://github.com/RailtownAI/railtracks)
- [Docs](https://railtownai.github.io/railtracks/)
- [Railtracks](https://railtracks.org/)

View file

@ -4,6 +4,7 @@ import TabItem from '@theme/TabItem';
# Anthropic
LiteLLM supports all anthropic models.
- `claude-sonnet-4-5-20250929`
- `claude-opus-4-1-20250805`
- `claude-4` (`claude-opus-4-20250514`, `claude-sonnet-4-20250514`)
- `claude-3.7` (`claude-3-7-sonnet-20250219`)
@ -268,6 +269,7 @@ print(response)
| Model Name | Function Call |
|------------------|--------------------------------------------|
| claude-sonnet-4-5 | `completion('claude-sonnet-4-5-20250929', messages)` | `os.environ['ANTHROPIC_API_KEY']` |
| claude-opus-4 | `completion('claude-opus-4-20250514', messages)` | `os.environ['ANTHROPIC_API_KEY']` |
| claude-sonnet-4 | `completion('claude-sonnet-4-20250514', messages)` | `os.environ['ANTHROPIC_API_KEY']` |
| claude-3.7 | `completion('claude-3-7-sonnet-20250219', messages)` | `os.environ['ANTHROPIC_API_KEY']` |

View file

@ -101,6 +101,7 @@ aws_profile_name: Optional[str],
aws_role_name: Optional[str],
aws_web_identity_token: Optional[str],
aws_bedrock_runtime_endpoint: Optional[str],
api_key: Optional[str],
```
### 2. Start the proxy
@ -1857,6 +1858,7 @@ Here's an example of using a bedrock model with LiteLLM. For a complete list, re
| GPT-OSS 20B | `completion(model='bedrock/converse/openai.gpt-oss-20b-1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
| GPT-OSS 120B | `completion(model='bedrock/converse/openai.gpt-oss-120b-1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
| Deepseek R1 | `completion(model='bedrock/us.deepseek.r1-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` |
| Anthropic Claude Sonnet 4.5 | `completion(model='bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` |
| Anthropic Claude-V3.5 Sonnet | `completion(model='bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` |
| Anthropic Claude-V3 sonnet | `completion(model='bedrock/anthropic.claude-3-sonnet-20240229-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` |
| Anthropic Claude-V3 Haiku | `completion(model='bedrock/anthropic.claude-3-haiku-20240307-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` |

View file

@ -0,0 +1,191 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Lemonade
[Lemonade Server](https://lemonade-server.ai/) is an OpenAI-compatible local language model inference provider optimized for AMD GPUs and NPUs. The `lemonade` litellm provider supports standard chat completions with full OpenAI API compatibility.
| Property | Details |
|-------|-------|
| Description | OpenAI-compatible AI provider for local and cloud-based language model inference |
| Provider Route on LiteLLM | `lemonade/` (add this prefix to the model name - e.g. `lemonade/your-model-name`) |
| API Endpoint for Provider | http://localhost:8000/api/v1 (default) |
| Supported Endpoints | `/chat/completions` |
## Supported OpenAI Parameters
Lemonade is fully OpenAI-compatible and supports the following parameters:
```
"repeat_penalty"
"functions"
"logit_bias"
"max_tokens"
"max_completion_tokens"
"presence_penalty"
"stop"
"temperature"
"top_p"
"top_k"
"response_format"
"tools"
```
## API Key Setup
Lemonade can be configured with custom API URLs and doesn't require strict API key validation. Set the `LEMONADE_API_BASE` environment variable to modify the base URL.
## Usage
<Tabs>
<TabItem value="sdk" label="SDK">
```python
from litellm import completion
import os
# Optional: Set custom API base. Useful if your lemonade server is on
# a different port
os.environ['LEMONADE_API_BASE'] = "http://localhost:8000/api/v1"
response = completion(
model="lemonade/your-model-name",
messages=[
{"role": "user", "content": "Hello from LiteLLM!"}
],
)
print(response)
```
## Streaming
```python
from litellm import completion
import os
# Optional: Set custom API base. Useful if your lemonade server is on
# a different port
os.environ['LEMONADE_API_BASE'] = "http://localhost:8000/api/v1"
response = completion(
model="lemonade/your-model-name",
messages=[
{"role": "user", "content": "Write a short story"}
],
stream=True
)
for chunk in response:
print(chunk.choices[0].delta.content, end='', flush=True)
```
## Advanced Usage
### Custom Parameters
Lemonade supports additional parameters beyond the standard OpenAI set:
```python
from litellm import completion
response = completion(
model="lemonade/your-model-name",
messages=[{"role": "user", "content": "Explain quantum computing"}],
temperature=0.7,
max_tokens=500,
top_p=0.9,
top_k=50,
repeat_penalty=1.1,
stop=["Human:", "AI:"]
)
print(response)
```
### Function Calling
Lemonade supports OpenAI-compatible function calling:
```python
from litellm import completion
functions = [
{
"name": "get_weather",
"description": "Get current weather information",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state"
}
},
"required": ["location"]
}
}
]
response = completion(
model="lemonade/your-model-name",
messages=[{"role": "user", "content": "What's the weather in San Francisco?"}],
tools=[{"type": "function", "function": f} for f in functions],
tool_choice="auto"
)
print(response)
```
### Response Format
Lemonade supports structured output with response format:
```python
from litellm import completion
import json
# Define schema in response_format
response = completion(
model="lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF",
messages=[{"role": "user", "content": "Generate JSON data for a person with their name, age, and city."}],
response_format={
"type": "json_schema",
"json_schema": {
"name": "person",
"schema": {
"type": "object",
"properties": {
"name": {"type": "string"},
"age": {"type": "integer"},
"city": {"type": "string"}
},
"required": ["name", "age"]
}
}
}
)
print(f"Model: {response.model}")
print(f"JSON Output:")
json_data = json.loads(response.choices[0].message.content)
print(json.dumps(json_data, indent=2))
```
## Available Models
Lemonade automatically validates available models by querying the `/models` endpoint. You can check available models programmatically:
```python
import httpx
api_base = "http://localhost:8000" # or your custom base
response = httpx.get(f"{api_base}/api/v1/models")
models = response.json()
print("Available models:", [model['id'] for model in models.get('data', [])])
```
## Support
For more information regarding Lemonade please go to to the [Lemonade website](https://lemonade-server.ai/) or [Lemonade repository](https://github.com/lemonade-sdk/lemonade).
</TabItem>
</Tabs>

View file

@ -1299,8 +1299,6 @@ litellm.vertex_location = "us-central1 # Your Location
| gemini-2.5-pro | `completion('gemini-2.5-pro', messages)`, `completion('vertex_ai/gemini-2.5-pro', messages)` |
| gemini-2.5-flash-preview-09-2025 | `completion('gemini-2.5-flash-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-preview-09-2025', messages)` |
| gemini-2.5-flash-lite-preview-09-2025 | `completion('gemini-2.5-flash-lite-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-lite-preview-09-2025', messages)` |
| gemini-flash-latest | `completion('gemini-flash-latest', messages)`, `completion('vertex_ai/gemini-flash-latest', messages)` |
| gemini-flash-lite-latest | `completion('gemini-flash-lite-latest', messages)`, `completion('vertex_ai/gemini-flash-lite-latest', messages)` |
## Fine-tuned Models

View file

@ -27,7 +27,7 @@ echo 'LITELLM_SALT_KEY="sk-1234"' >> .env
source .env
# Start
docker-compose up
docker compose up
```

View file

@ -55,7 +55,7 @@ echo 'LITELLM_SALT_KEY="sk-1234"' >> .env
source .env
# Start
docker-compose up
docker compose up
```
</TabItem>

View file

@ -14,6 +14,7 @@ Found under `kwargs["standard_logging_object"]`. This is a standard payload, log
| `cost_breakdown` | `Optional[CostBreakdown]` | Detailed cost breakdown object |
| `response_cost_failure_debug_info` | `StandardLoggingModelCostFailureDebugInformation` | Debug information if cost tracking fails |
| `status` | `StandardLoggingPayloadStatus` | Status of the payload |
| `status_fields` | `StandardLoggingPayloadStatusFields` | Typed status fields for easy filtering and analytics |
| `total_tokens` | `int` | Total number of tokens |
| `prompt_tokens` | `int` | Number of prompt tokens |
| `completion_tokens` | `int` | Number of completion tokens |
@ -162,17 +163,89 @@ A literal type with two possible values:
## StandardLoggingGuardrailInformation
| Field | Type | Description |
|-----------------------|------|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| `guardrail_name` | `Optional[str]` | Guardrail name |
| `guardrail_provider` | `Optional[str]` | Guardrail provider |
| `guardrail_mode` | `Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]]` | Guardrail mode |
| `guardrail_request` | `Optional[dict]` | Guardrail request |
| `guardrail_response` | `Optional[Union[dict, str, List[dict]]]` | Guardrail response |
| `guardrail_status` | `Literal["success", "failure", "blocked"]` | Guardrail execution status: `success` = no violations detected, `blocked` = content blocked/modified due to policy violations, `failure` = technical error or API failure |
| `start_time` | `Optional[float]` | Start time of the guardrail |
| `end_time` | `Optional[float]` | End time of the guardrail |
| `duration` | `Optional[float]` | Duration of the guardrail in seconds |
| `masked_entity_count` | `Optional[Dict[str, int]]` | Count of masked entities |
## StandardLoggingPayloadStatusFields
Typed status fields for easy filtering and analytics.
| Field | Type | Description |
|-------|------|-------------|
| `guardrail_name` | `Optional[str]` | Guardrail name |
| `guardrail_mode` | `Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]]` | Guardrail mode |
| `guardrail_request` | `Optional[dict]` | Guardrail request |
| `guardrail_response` | `Optional[Union[dict, str, List[dict]]]` | Guardrail response |
| `guardrail_status` | `Literal["success", "failure"]` | Guardrail status |
| `start_time` | `Optional[float]` | Start time of the guardrail |
| `end_time` | `Optional[float]` | End time of the guardrail |
| `duration` | `Optional[float]` | Duration of the guardrail in seconds |
| `masked_entity_count` | `Optional[Dict[str, int]]` | Count of masked entities |
| `llm_api_status` | `StandardLoggingPayloadStatus` | Status of the LLM API call: `"success"` if completed successfully, `"failure"` if errored |
| `guardrail_status` | `GuardrailStatus` | Status of guardrail execution (see below) |
### StandardLoggingPayloadStatus
A literal type with two possible values:
- `"success"` - The LLM API request completed successfully
- `"failure"` - The LLM API request failed
### GuardrailStatus
A literal type with four possible values:
- `"success"` - Guardrail ran and allowed content through (no violations detected)
- `"guardrail_intervened"` - Guardrail blocked or modified content due to policy violations
- `"guardrail_failed_to_respond"` - Guardrail had a technical failure or API error
- `"not_run"` - No guardrail was executed for this request
### Usage Examples
Filter logs for requests where guardrails intervened:
```json
{
"status_fields": {
"guardrail_status": "guardrail_intervened"
}
}
```
Find guardrail technical failures:
```json
{
"status_fields": {
"guardrail_status": "guardrail_failed_to_respond"
}
}
```
Get successful LLM requests:
```json
{
"status_fields": {
"llm_api_status": "success"
}
}
```
Find requests where guardrails ran successfully without intervention:
```json
{
"status_fields": {
"guardrail_status": "success",
"llm_api_status": "success"
}
}
```
Find requests where no guardrail was run:
```json
{
"status_fields": {
"guardrail_status": "not_run"
}
}
```
## StandardLoggingPromptManagementMetadata

View file

@ -9,7 +9,7 @@ Store prompts as `.prompt` files in your repository and use them directly with L
- **File System**: Store `.prompt` files locally
- **BitBucket**: Store `.prompt` files in BitBucket repositories with team-based access control
- **Gitlab**: Store `.prompt` files in Gitlab repositories with team-based access control
## Quick Start
<Tabs>
@ -90,6 +90,51 @@ response = litellm.completion(
```
</TabItem>
<TabItem value="gitlab" label="GITLAB">
**1. Create a .prompt file in a gitlab repo**
Create `prompts/hello.prompt` in your gitlab repository:
```yaml
---
model: gpt-4
temperature: 0.7
---
System: You are a helpful assistant.
User: {{user_message}}
```
**2. Configure Gitlab access**
```python
import litellm
# Configure gitlab access
gitlab_config = {
"workspace": "your-workspace",
"repository": "your-repo",
"access_token": "your-access-token",
"branch": "main"
}
# Set global gitlab configuration
litellm.set_global_gitlab_config(gitlab_config)
```
**3. Use with LiteLLM**
```python
response = litellm.completion(
model="gitlab/gpt-4",
prompt_id="hello",
prompt_variables={"user_message": "What is the capital of France?"}
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
**1. Create a .prompt file**
@ -124,6 +169,12 @@ litellm_settings:
repository: "your-repo"
access_token: "your-access-token"
branch: "main"
# Or use Gitlab for team-based prompt management
global_gitlab_config:
workspace: "your-workspace"
repository: "your-repo"
access_token: "your-access-token"
branch: "main"
```
**3. Start the proxy**
@ -213,6 +264,14 @@ prompt_variables: Optional[dict] # optional - variables for template rendering
bitbucket_config: Optional[dict] # optional - BitBucket configuration (if not set globally)
```
**Gitlab:**
```
model: gitlab/<base_model> # required (e.g., gitlab/gpt-4)
prompt_id: str # required - the .prompt filename without extension
prompt_variables: Optional[dict] # optional - variables for template rendering
gitlab_config: Optional[dict] # optional - Gitlab configuration (if not set globally)
```
**Example API calls:**
```python
@ -235,4 +294,18 @@ response = litellm.completion(
"access_token": "your-token"
}
)
# Gitlab integration
response = litellm.completion(
model="gitlab/gpt-4",
prompt_id="hello",
prompt_variables={"user_message": "Hello world"},
gitlab_config={
"project": "a/b/<repo_name>",
"access_token": "your-access-token",
"base_url": "gitlab url",
"prompts_path": "src/prompts", # folder to point to, defaults to root
"branch":"main" # optional, defaults to main
}
)
```

View file

@ -243,6 +243,18 @@ curl --location 'http://0.0.0.0:4000/v1/messages' \
}'
```
---
## Tutorial - Add Azure OpenAI Assistants API as a Pass Through Endpoint
In this video, we'll add the Azure OpenAI Assistants API as a pass through endpoint to LiteLLM Proxy.
<iframe width="840" height="500" src="https://www.loom.com/embed/12965cb299d24fc0bd7b6b413ab6d0ad" frameborder="0" webkitallowfullscreen mozallowfullscreen allowfullscreen></iframe>
<br/>
<br/>
---
## Troubleshooting

View file

@ -25,6 +25,10 @@ import TabItem from '@theme/TabItem';
<TabItem value="docker" label="Docker">
``` showLineNumbers title="docker run litellm"
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
ghcr.io/berriai/litellm:v1.77.5.rc.1
```
</TabItem>
@ -32,6 +36,7 @@ import TabItem from '@theme/TabItem';
<TabItem value="pip" label="Pip">
``` showLineNumbers title="pip install litellm"
pip install litellm==1.77.5
```
</TabItem>

View file

@ -477,6 +477,7 @@ const sidebars = {
"providers/fireworks_ai",
"providers/clarifai",
"providers/compactifai",
"providers/lemonade",
"providers/vllm",
"providers/llamafile",
"providers/infinity",
@ -698,7 +699,8 @@ const sidebars = {
"projects/llm_cord",
"projects/pgai",
"projects/GPTLocalhost",
"projects/HolmesGPT"
"projects/HolmesGPT",
"projects/Railtracks",
],
},
"extras/code_quality",

View file

@ -152,6 +152,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"vector_store_pre_call_hook",
"dotprompt",
"bitbucket",
"gitlab",
"cloudzero",
"posthog",
]
@ -172,22 +173,22 @@ prometheus_initialize_budget_metrics: Optional[bool] = False
require_auth_for_metrics_endpoint: Optional[bool] = False
argilla_batch_size: Optional[int] = None
datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload.
gcs_pub_sub_use_v1: Optional[bool] = (
False # if you want to use v1 gcs pubsub logged payload
)
generic_api_use_v1: Optional[bool] = (
False # if you want to use v1 generic api logged payload
)
gcs_pub_sub_use_v1: Optional[
bool
] = False # if you want to use v1 gcs pubsub logged payload
generic_api_use_v1: Optional[
bool
] = False # if you want to use v1 generic api logged payload
argilla_transformation_object: Optional[Dict[str, Any]] = None
_async_input_callback: List[Union[str, Callable, CustomLogger]] = (
[]
) # internal variable - async custom callbacks are routed here.
_async_success_callback: List[Union[str, Callable, CustomLogger]] = (
[]
) # internal variable - async custom callbacks are routed here.
_async_failure_callback: List[Union[str, Callable, CustomLogger]] = (
[]
) # internal variable - async custom callbacks are routed here.
_async_input_callback: List[
Union[str, Callable, CustomLogger]
] = [] # internal variable - async custom callbacks are routed here.
_async_success_callback: List[
Union[str, Callable, CustomLogger]
] = [] # internal variable - async custom callbacks are routed here.
_async_failure_callback: List[
Union[str, Callable, CustomLogger]
] = [] # internal variable - async custom callbacks are routed here.
pre_call_rules: List[Callable] = []
post_call_rules: List[Callable] = []
turn_off_message_logging: Optional[bool] = False
@ -195,18 +196,18 @@ log_raw_request_response: bool = False
redact_messages_in_exceptions: Optional[bool] = False
redact_user_api_key_info: Optional[bool] = False
filter_invalid_headers: Optional[bool] = False
add_user_information_to_llm_headers: Optional[bool] = (
None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
)
add_user_information_to_llm_headers: Optional[
bool
] = None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
store_audit_logs = False # Enterprise feature, allow users to see audit logs
### end of callbacks #############
email: Optional[str] = (
None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
token: Optional[str] = (
None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
email: Optional[
str
] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
token: Optional[
str
] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
telemetry = True
max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults
drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False))
@ -250,6 +251,7 @@ wandb_key: Optional[str] = None
heroku_key: Optional[str] = None
cometapi_key: Optional[str] = None
ovhcloud_key: Optional[str] = None
lemonade_key: Optional[str] = None
common_cloud_provider_auth_params: dict = {
"params": ["project", "region_name", "token"],
"providers": ["vertex_ai", "bedrock", "watsonx", "azure", "vertex_ai_beta"],
@ -306,24 +308,20 @@ enable_loadbalancing_on_batch_endpoints: Optional[bool] = None
enable_caching_on_provider_specific_optional_params: bool = (
False # feature-flag for caching on optional params - e.g. 'top_k'
)
caching: bool = (
False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
caching_with_models: bool = (
False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
cache: Optional[Cache] = (
None # cache object <- use this - https://docs.litellm.ai/docs/caching
)
caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
cache: Optional[
Cache
] = None # cache object <- use this - https://docs.litellm.ai/docs/caching
default_in_memory_ttl: Optional[float] = None
default_redis_ttl: Optional[float] = None
default_redis_batch_cache_expiry: Optional[float] = None
model_alias_map: Dict[str, str] = {}
model_group_settings: Optional["ModelGroupSettings"] = None
max_budget: float = 0.0 # set the max budget across all providers
budget_duration: Optional[str] = (
None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
)
budget_duration: Optional[
str
] = None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
default_soft_budget: float = (
DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0
)
@ -332,15 +330,11 @@ forward_traceparent_to_llm_provider: bool = False
_current_cost = 0.0 # private variable, used if max budget is set
error_logs: Dict = {}
add_function_to_prompt: bool = (
False # if function calling not supported by api, append function call details to system prompt
)
add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt
client_session: Optional[httpx.Client] = None
aclient_session: Optional[httpx.AsyncClient] = None
model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks'
model_cost_map_url: str = (
"https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json"
)
model_cost_map_url: str = "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json"
suppress_debug_info = False
dynamodb_table_name: Optional[str] = None
s3_callback_params: Optional[Dict] = None
@ -370,9 +364,7 @@ prometheus_metrics_config: Optional[List] = None
disable_add_prefix_to_prompt: bool = (
False # used by anthropic, to disable adding prefix to prompt
)
disable_copilot_system_to_assistant: bool = (
False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
)
disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
public_model_groups: Optional[List[str]] = None
public_model_groups_links: Dict[str, str] = {}
#### REQUEST PRIORITIZATION ######
@ -383,17 +375,13 @@ priority_reservation_settings: "PriorityReservationSettings" = (
######## Networking Settings ########
use_aiohttp_transport: bool = (
True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead.
)
use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead.
aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings
disable_aiohttp_transport: bool = False # Set this to true to use httpx instead
disable_aiohttp_trust_env: bool = (
False # When False, aiohttp will respect HTTP(S)_PROXY env vars
)
force_ipv4: bool = (
False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6.
)
force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6.
module_level_aclient = AsyncHTTPHandler(
timeout=request_timeout, client_alias="module level aclient"
)
@ -407,13 +395,13 @@ fallbacks: Optional[List] = None
context_window_fallbacks: Optional[List] = None
content_policy_fallbacks: Optional[List] = None
allowed_fails: int = 3
num_retries_per_request: Optional[int] = (
None # for the request overall (incl. fallbacks + model retries)
)
num_retries_per_request: Optional[
int
] = None # for the request overall (incl. fallbacks + model retries)
####### SECRET MANAGERS #####################
secret_manager_client: Optional[Any] = (
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
)
secret_manager_client: Optional[
Any
] = None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
_google_kms_resource_name: Optional[str] = None
_key_management_system: Optional[KeyManagementSystem] = None
_key_management_settings: KeyManagementSettings = KeyManagementSettings()
@ -536,6 +524,7 @@ volcengine_models: Set = set()
wandb_models: Set = set(WANDB_MODELS)
ovhcloud_models: Set = set()
ovhcloud_embedding_models: Set = set()
lemonade_models: Set = set()
def is_bedrock_pricing_only_model(key: str) -> bool:
@ -756,6 +745,8 @@ def add_known_models():
ovhcloud_models.add(key)
elif value.get("litellm_provider") == "ovhcloud-embedding-models":
ovhcloud_embedding_models.add(key)
elif value.get("litellm_provider") == "lemonade":
lemonade_models.add(key)
add_known_models()
@ -852,6 +843,7 @@ model_list = list(
| volcengine_models
| wandb_models
| ovhcloud_models
| lemonade_models
)
model_list_set = set(model_list)
@ -935,6 +927,7 @@ models_by_provider: dict = {
"volcengine": volcengine_models,
"wandb": wandb_models,
"ovhcloud": ovhcloud_models | ovhcloud_embedding_models,
"lemonade": lemonade_models,
}
# mapping for those models which have larger equivalents
@ -1284,6 +1277,7 @@ from .llms.hyperbolic.chat.transformation import HyperbolicChatConfig
from .llms.vercel_ai_gateway.chat.transformation import VercelAIGatewayConfig
from .llms.ovhcloud.chat.transformation import OVHCloudChatConfig
from .llms.ovhcloud.embedding.transformation import OVHCloudEmbeddingConfig
from .llms.lemonade.chat.transformation import LemonadeChatConfig
from .main import * # type: ignore
from .integrations import *
from .llms.custom_httpx.async_client_cleanup import close_litellm_async_clients
@ -1342,12 +1336,12 @@ from .types.llms.custom_llm import CustomLLMItem
from .types.utils import GenericStreamingChunk
custom_provider_map: List[CustomLLMItem] = []
_custom_providers: List[str] = (
[]
) # internal helper util, used to track names of custom providers
disable_hf_tokenizer_download: Optional[bool] = (
None # disable huggingface tokenizer download. Defaults to openai clk100
)
_custom_providers: List[
str
] = [] # internal helper util, used to track names of custom providers
disable_hf_tokenizer_download: Optional[
bool
] = None # disable huggingface tokenizer download. Defaults to openai clk100
global_disable_no_log_param: bool = False
### CLI UTILITIES ###
@ -1355,6 +1349,7 @@ from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_k
### PASSTHROUGH ###
from .passthrough import allm_passthrough_route, llm_passthrough_route
from .google_genai import agenerate_content
### GLOBAL CONFIG ###
global_bitbucket_config: Optional[Dict[str, Any]] = None
@ -1364,3 +1359,11 @@ def set_global_bitbucket_config(config: Dict[str, Any]) -> None:
"""Set global BitBucket configuration for prompt management."""
global global_bitbucket_config
global_bitbucket_config = config
### GLOBAL CONFIG ###
global_gitlab_config: Optional[Dict[str, Any]] = None
def set_global_gitlab_config(config: Dict[str, Any]) -> None:
"""Set global BitBucket configuration for prompt management."""
global global_gitlab_config
global_gitlab_config = config

View file

@ -36,12 +36,16 @@ import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.caching import InMemoryCache
from litellm.caching.caching import S3Cache
from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
update_response_metadata,
)
from litellm.litellm_core_utils.logging_utils import (
_assemble_complete_response_from_streaming_chunks,
)
from litellm.types.caching import CachedEmbedding
from litellm.types.rerank import RerankResponse
from litellm.types.utils import (
CachingDetails,
CallTypes,
Embedding,
EmbeddingResponse,
@ -136,6 +140,13 @@ class LLMCachingHandler:
kwargs = kwargs.copy()
args = args or ()
#########################################################
# Init cache timing metrics
#########################################################
cache_check_start_time = datetime.datetime.now()
cache_check_end_time = None
#########################################################
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
kwargs["parent_otel_span"] = parent_otel_span
@ -157,6 +168,7 @@ class LLMCachingHandler:
kwargs=kwargs,
args=args,
)
cache_check_end_time = datetime.datetime.now()
if cached_result is not None and not isinstance(cached_result, list):
verbose_logger.debug("Cache Hit!")
@ -168,6 +180,7 @@ class LLMCachingHandler:
api_base=kwargs.get("api_base", None),
api_key=kwargs.get("api_key", None),
)
cache_duration_ms = (cache_check_end_time - cache_check_start_time).total_seconds() * 1000
self._update_litellm_logging_obj_environment(
logging_obj=logging_obj,
model=model,
@ -175,10 +188,12 @@ class LLMCachingHandler:
cached_result=cached_result,
is_async=True,
custom_llm_provider=custom_llm_provider,
cache_duration_ms=cache_duration_ms,
)
call_type = original_function.__name__
cached_result = self._convert_cached_result_to_model_response(
cached_result=cached_result,
call_type=call_type,
@ -716,6 +731,18 @@ class LLMCachingHandler:
and isinstance(cached_result._hidden_params, dict)
):
cached_result._hidden_params["cache_hit"] = True
#########################################################
# Add final timing metrics to the cached result
#########################################################
update_response_metadata(
result=cached_result,
logging_obj=logging_obj,
model=model,
kwargs=kwargs,
start_time=self.start_time,
end_time=datetime.datetime.now(),
)
return cached_result
def _convert_cached_stream_response(
@ -944,6 +971,7 @@ class LLMCachingHandler:
is_async: bool,
is_embedding: bool = False,
custom_llm_provider: Optional[str] = None,
cache_duration_ms: Optional[float] = None,
):
"""
Helper function to update the LiteLLMLoggingObj environment variables.
@ -995,6 +1023,11 @@ class LLMCachingHandler:
custom_llm_provider=custom_llm_provider,
)
logging_obj.caching_details = CachingDetails(
cache_hit=True,
cache_duration_ms=cache_duration_ms,
)
def convert_args_to_kwargs(
original_function: Callable,

View file

@ -11,6 +11,7 @@ Has 4 methods:
import json
import sys
import time
import heapq
from typing import TYPE_CHECKING, Any, List, Optional
if TYPE_CHECKING:
@ -46,6 +47,7 @@ class InMemoryCache(BaseCache):
# in-memory cache
self.cache_dict: dict = {}
self.ttl_dict: dict = {}
self.expiration_heap: list[tuple[float, str]] = []
def check_value_size(self, value: Any):
"""
@ -114,19 +116,27 @@ class InMemoryCache(BaseCache):
"""
current_time = time.time()
# Step 1: Remove expired items
expired_keys = [key for key, ttl in self.ttl_dict.items() if current_time > ttl]
for key in expired_keys:
self._remove_key(key)
# Step 2: If cache is still full, evict items with earliest expiration times
if len(self.cache_dict) >= self.max_size_in_memory:
# Sort by expiration time (earliest first) and evict until we're under the limit
items_by_expiration = sorted(self.ttl_dict.items(), key=lambda x: x[1])
keys_to_evict = items_by_expiration[:len(self.cache_dict) - self.max_size_in_memory + 1]
for key, _ in keys_to_evict:
# Step 1: Remove expired or outdated items
while self.expiration_heap:
expiration_time, key = self.expiration_heap[0]
# Case 1: Heap entry is outdated
if expiration_time != self.ttl_dict.get(key):
heapq.heappop(self.expiration_heap)
# Case 2: Entry is valid but expired
elif expiration_time <= current_time:
heapq.heappop(self.expiration_heap)
self._remove_key(key)
else:
# Case 3: Entry is valid and not expired
break
# Step 2: Evict if cache is still full
while len(self.cache_dict) >= self.max_size_in_memory:
expiration_time, key = heapq.heappop(self.expiration_heap)
# Skip if key was removed or updated
if self.ttl_dict.get(key) == expiration_time:
self._remove_key(key)
# de-reference the removed item
@ -150,7 +160,7 @@ class InMemoryCache(BaseCache):
# Handle the edge case where max_size_in_memory is 0
if self.max_size_in_memory == 0:
return # Don't cache anything if max size is 0
if len(self.cache_dict) >= self.max_size_in_memory:
# only evict when cache is full
self.evict_cache()
@ -161,8 +171,10 @@ class InMemoryCache(BaseCache):
if self.allow_ttl_override(key): # if ttl is not set, set it to default ttl
if "ttl" in kwargs and kwargs["ttl"] is not None:
self.ttl_dict[key] = time.time() + float(kwargs["ttl"])
heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key))
else:
self.ttl_dict[key] = time.time() + self.default_ttl
heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key))
async def async_set_cache(self, key, value, **kwargs):
self.set_cache(key=key, value=value, **kwargs)
@ -253,6 +265,7 @@ class InMemoryCache(BaseCache):
def flush_cache(self):
self.cache_dict.clear()
self.ttl_dict.clear()
self.expiration_heap.clear()
async def disconnect(self):
pass

View file

@ -315,6 +315,7 @@ LITELLM_CHAT_PROVIDERS = [
"vercel_ai_gateway",
"wandb",
"ovhcloud",
"lemonade"
]
LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS = [
@ -819,6 +820,7 @@ BEDROCK_CONVERSE_MODELS = [
"deepseek.v3-v1:0",
"openai.gpt-oss-20b-1:0",
"openai.gpt-oss-120b-1:0",
"anthropic.claude-sonnet-4-5-20250929-v1:0",
"anthropic.claude-opus-4-1-20250805-v1:0",
"anthropic.claude-opus-4-20250514-v1:0",
"anthropic.claude-sonnet-4-20250514-v1:0",

View file

@ -58,6 +58,9 @@ from litellm.llms.vertex_ai.cost_calculator import (
)
from litellm.llms.vertex_ai.cost_calculator import cost_router as google_cost_router
from litellm.llms.xai.cost_calculator import cost_per_token as xai_cost_per_token
from litellm.llms.lemonade.cost_calculator import (
cost_per_token as lemonade_cost_per_token,
)
from litellm.responses.utils import ResponseAPILoggingUtils
from litellm.types.llms.openai import (
HttpxBinaryResponseContent,
@ -347,6 +350,8 @@ def cost_per_token( # noqa: PLR0915
return perplexity_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "xai":
return xai_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "lemonade":
return lemonade_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "dashscope":
from litellm.llms.dashscope.cost_calculator import (
cost_per_token as dashscope_cost_per_token,

View file

@ -194,7 +194,7 @@ class MCPClient:
def _get_auth_headers(self) -> dict:
"""Generate authentication headers based on auth type."""
headers = {"MCP-Protocol-Version": "2025-06-18"}
headers = {}
if self._mcp_auth_value:
if isinstance(self._mcp_auth_value, str):

View file

@ -72,15 +72,24 @@ class GenerateContentToCompletionHandler:
completion_response = await litellm.acompletion(**completion_kwargs)
if stream:
# Transform streaming completion response to generate_content format
transformed_stream = (
GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
# Check if completion_response is actually a stream or a ModelResponse
# This can happen in error cases or when stream is not properly supported
if not hasattr(completion_response, "__aiter__"):
# If it's not a stream, treat it as a regular response
generate_content_response = (
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
cast(ModelResponse, completion_response)
)
)
return generate_content_response
else:
# Transform streaming completion response to generate_content format
transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
completion_response
)
)
if transformed_stream is not None:
return transformed_stream
raise ValueError("Failed to transform streaming response")
if transformed_stream is not None:
return transformed_stream
raise ValueError("Failed to transform streaming response")
else:
# Transform completion response back to generate_content format
generate_content_response = (
@ -136,15 +145,24 @@ class GenerateContentToCompletionHandler:
completion_response = litellm.completion(**completion_kwargs)
if stream:
# Transform streaming completion response to generate_content format
transformed_stream = (
GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
# Check if completion_response is actually a stream or a ModelResponse
# This can happen in error cases or when stream is not properly supported
if not hasattr(completion_response, "__iter__"):
# If it's not a stream, treat it as a regular response
generate_content_response = (
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
cast(ModelResponse, completion_response)
)
)
return generate_content_response
else:
# Transform streaming completion response to generate_content format
transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
completion_response
)
)
if transformed_stream is not None:
return transformed_stream
raise ValueError("Failed to transform streaming response")
if transformed_stream is not None:
return transformed_stream
raise ValueError("Failed to transform streaming response")
else:
# Transform completion response back to generate_content format
generate_content_response = (

View file

@ -1,12 +1,15 @@
import json
from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Union, cast
from litellm import verbose_logger
from litellm.litellm_core_utils.json_validation_rule import normalize_tool_schema
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionAssistantMessage,
ChatCompletionAssistantToolCall,
ChatCompletionRequest,
ChatCompletionSystemMessage,
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolChoiceValues,
ChatCompletionToolMessage,
@ -36,43 +39,103 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
def __init__(self, completion_stream: Any):
self.sent_first_chunk = False
self.accumulated_tool_calls = {}
self._returned_response = False
super().__init__(completion_stream)
def __next__(self):
try:
if not hasattr(self.completion_stream, "__iter__"):
if self._returned_response:
raise StopIteration
self._returned_response = True
return GoogleGenAIAdapter().translate_completion_to_generate_content(
self.completion_stream
)
for chunk in self.completion_stream:
if chunk == "None" or chunk is None:
continue
# Transform OpenAI streaming chunk to Google GenAI format
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
chunk, self
)
if transformed_chunk: # Only return non-empty chunks
if transformed_chunk:
return transformed_chunk
raise StopIteration
except StopIteration:
raise StopIteration
raise
except Exception:
raise StopIteration
async def __anext__(self):
try:
if not hasattr(self.completion_stream, "__aiter__"):
if self._returned_response:
raise StopAsyncIteration
self._returned_response = True
return GoogleGenAIAdapter().translate_completion_to_generate_content(
self.completion_stream
)
async for chunk in self.completion_stream:
if chunk == "None" or chunk is None:
continue
# Transform OpenAI streaming chunk to Google GenAI format
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
chunk, self
)
if transformed_chunk: # Only return non-empty chunks
if transformed_chunk:
return transformed_chunk
# After the stream is exhausted, check for any remaining accumulated tool calls
if self.accumulated_tool_calls:
try:
parts = []
for (
tool_call_index,
tool_call_data,
) in self.accumulated_tool_calls.items():
try:
# For tool calls with no arguments, accumulated_args will be "", which is not valid JSON.
# We default to an empty JSON object in this case.
parsed_args = json.loads(
tool_call_data["arguments"] or "{}"
)
function_call_part = {
"functionCall": {
"name": tool_call_data["name"]
or "undefined_tool_name",
"args": parsed_args,
}
}
parts.append(function_call_part)
except json.JSONDecodeError:
# This can happen if the stream is abruptly cut off mid-argument string.
verbose_logger.warning(
f"Could not parse tool call arguments at end of stream for index {tool_call_index}. "
f"Name: {tool_call_data['name']}. "
f"Partial args: {tool_call_data['arguments']}"
)
pass
if parts:
final_chunk = {
"candidates": [
{
"content": {"parts": parts, "role": "model"},
"finishReason": "STOP",
"index": 0,
"safetyRatings": [],
}
]
}
return final_chunk
finally:
# Ensure the accumulator is always cleared to prevent memory leaks
self.accumulated_tool_calls.clear()
raise StopAsyncIteration
except StopAsyncIteration:
raise StopAsyncIteration
raise
except Exception:
raise StopAsyncIteration
@ -107,9 +170,14 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
payload = f"data: {json.dumps(transformed_chunk)}\n\n"
yield payload.encode()
else:
raise ValueError(f"Invalid chunk 1: {chunk}")
# For empty chunks, continue to next iteration
continue
else:
raise ValueError(f"Invalid chunk 2: {chunk}")
# For other chunk types, yield them directly
if hasattr(chunk, "encode"):
yield chunk.encode()
else:
yield str(chunk).encode()
class GoogleGenAIAdapter:
@ -133,12 +201,19 @@ class GoogleGenAIAdapter:
model: The model name
contents: Generate content contents (can be list or single dict)
config: Optional config parameters
**kwargs: Additional parameters
**kwargs: Additional parameters from the original request
Returns:
Dict in OpenAI format
"""
# Extract top-level fields from kwargs
system_instruction = kwargs.get("systemInstruction") or kwargs.get(
"system_instruction"
)
tools = kwargs.get("tools")
tool_config = kwargs.get("toolConfig") or kwargs.get("tool_config")
# Normalize contents to list format
if isinstance(contents, dict):
contents_list = [contents]
@ -146,7 +221,9 @@ class GoogleGenAIAdapter:
contents_list = contents
# Transform contents to OpenAI messages format
messages = self._transform_contents_to_messages(contents_list)
messages = self._transform_contents_to_messages(
contents_list, system_instruction=system_instruction
)
# Create base request as dict (which is compatible with ChatCompletionRequest)
completion_request: ChatCompletionRequest = {
@ -182,20 +259,19 @@ class GoogleGenAIAdapter:
completion_request["stop"] = config["stopSequences"]
# Handle tools transformation
if "tools" in kwargs:
tools = kwargs["tools"]
if tools:
# Check if tools are already in OpenAI format or Google GenAI format
if isinstance(tools, list) and len(tools) > 0:
# Tools are in Google GenAI format, transform them
openai_tools = self._transform_google_genai_tools_to_openai(tools)
if openai_tools:
completion_request["tools"] = openai_tools
# Handle tool_config (tool choice)
if "tool_config" in kwargs:
if tool_config:
tool_choice = self._transform_google_genai_tool_config_to_openai(
kwargs["tool_config"]
tool_config
)
if tool_choice:
completion_request["tool_choice"] = tool_choice
@ -235,7 +311,8 @@ class GoogleGenAIAdapter:
return completion_request_dict
def translate_completion_output_params_streaming(
self, completion_stream: Any
self,
completion_stream: Any,
) -> Union[AsyncIterator[bytes], None]:
"""Transform streaming completion output to Google GenAI format"""
google_genai_wrapper = GoogleGenAIStreamWrapper(
@ -245,7 +322,8 @@ class GoogleGenAIAdapter:
return google_genai_wrapper.async_google_genai_sse_wrapper()
def _transform_google_genai_tools_to_openai(
self, tools: List[Dict[str, Any]]
self,
tools: List[Dict[str, Any]],
) -> List[ChatCompletionToolParam]:
"""Transform Google GenAI tools to OpenAI tools format"""
openai_tools: List[Dict[str, Any]] = []
@ -259,8 +337,8 @@ class GoogleGenAIAdapter:
if "description" in func_decl:
function_chunk["description"] = func_decl["description"]
if "parameters" in func_decl:
function_chunk["parameters"] = func_decl["parameters"]
if "parametersJsonSchema" in func_decl:
function_chunk["parameters"] = func_decl["parametersJsonSchema"]
openai_tool = {"type": "function", "function": function_chunk}
openai_tools.append(openai_tool)
@ -271,7 +349,8 @@ class GoogleGenAIAdapter:
return cast(List[ChatCompletionToolParam], normalized_tools)
def _transform_google_genai_tool_config_to_openai(
self, tool_config: Dict[str, Any]
self,
tool_config: Dict[str, Any],
) -> Optional[ChatCompletionToolChoiceValues]:
"""Transform Google GenAI tool_config to OpenAI tool_choice"""
function_calling_config = tool_config.get("functionCallingConfig", {})
@ -283,11 +362,23 @@ class GoogleGenAIAdapter:
return cast(ChatCompletionToolChoiceValues, tool_choice)
def _transform_contents_to_messages(
self, contents: List[Dict[str, Any]]
self,
contents: List[Dict[str, Any]],
system_instruction: Optional[Dict[str, Any]] = None,
) -> List[AllMessageValues]:
"""Transform Google GenAI contents to OpenAI messages format"""
messages: List[AllMessageValues] = []
# Handle system instruction
if system_instruction:
system_parts = system_instruction.get("parts", [])
if system_parts and "text" in system_parts[0]:
messages.append(
ChatCompletionSystemMessage(
role="system", content=system_parts[0]["text"]
)
)
for content in contents:
role = content.get("role", "user")
parts = content.get("parts", [])
@ -364,7 +455,8 @@ class GoogleGenAIAdapter:
return messages
def translate_completion_to_generate_content(
self, response: ModelResponse
self,
response: ModelResponse,
) -> Dict[str, Any]:
"""
Transform litellm completion response to Google GenAI generate_content format
@ -376,6 +468,7 @@ class GoogleGenAIAdapter:
Dict in Google GenAI generate_content response format
"""
# Extract the main response content
choice = response.choices[0] if response.choices else None
if not choice:
@ -388,12 +481,6 @@ class GoogleGenAIAdapter:
"Invalid completion response: no message found in choice"
)
parts = self._transform_openai_message_to_google_genai_parts(choice.message)
elif isinstance(choice, StreamingChoices):
if not choice.delta:
raise ValueError(
"Invalid completion response: no delta found in streaming choice"
)
parts = self._transform_openai_delta_to_google_genai_parts(choice.delta)
else:
# Fallback for generic choice objects
message_content = getattr(choice, "message", {}).get(
@ -438,7 +525,7 @@ class GoogleGenAIAdapter:
self,
response: Union[ModelResponse, ModelResponseStream],
wrapper: GoogleGenAIStreamWrapper,
) -> Dict[str, Any]:
) -> Optional[Dict[str, Any]]:
"""
Transform streaming litellm completion chunk to Google GenAI generate_content format
@ -454,7 +541,7 @@ class GoogleGenAIAdapter:
choice = response.choices[0] if response.choices else None
if not choice:
# Return empty chunk if no choices
return {}
return None
# Handle streaming choice
if isinstance(choice, StreamingChoices):
@ -473,7 +560,7 @@ class GoogleGenAIAdapter:
# Only create response chunk if we have parts or it's the final chunk
if not parts and not finish_reason:
return {}
return None
# Create Google GenAI streaming format response
streaming_chunk: Dict[str, Any] = {
@ -515,7 +602,8 @@ class GoogleGenAIAdapter:
return streaming_chunk
def _transform_openai_message_to_google_genai_parts(
self, message: Any
self,
message: Any,
) -> List[Dict[str, Any]]:
"""Transform OpenAI message to Google GenAI parts format"""
parts: List[Dict[str, Any]] = []
@ -537,112 +625,94 @@ class GoogleGenAIAdapter:
except json.JSONDecodeError:
args = {}
function_call_part = {
"functionCall": {"name": tool_call.function.name, "args": args}
}
parts.append(function_call_part)
return parts if parts else [{"text": ""}]
def _transform_openai_delta_to_google_genai_parts(
self, delta: Any
) -> List[Dict[str, Any]]:
"""Transform OpenAI delta to Google GenAI parts format for streaming"""
parts: List[Dict[str, Any]] = []
# Add text content if present
if hasattr(delta, "content") and delta.content:
parts.append({"text": delta.content})
# Add tool calls if present (for streaming tool calls)
if hasattr(delta, "tool_calls") and delta.tool_calls:
for tool_call in delta.tool_calls:
if hasattr(tool_call, "function") and tool_call.function:
# For streaming, we might get partial function arguments
args_str = getattr(tool_call.function, "arguments", "") or ""
try:
args = json.loads(args_str) if args_str else {}
except json.JSONDecodeError:
# For partial JSON in streaming, return as text for now
args = {"partial": args_str}
function_call_part = {
"functionCall": {
"name": getattr(tool_call.function, "name", "") or "",
"name": tool_call.function.name or "undefined_tool_name",
"args": args,
}
}
parts.append(function_call_part)
return parts
return parts if parts else [{"text": ""}]
def _transform_openai_delta_to_google_genai_parts_with_accumulation(
self, delta: Any, wrapper: GoogleGenAIStreamWrapper
) -> List[Dict[str, Any]]:
"""Transform OpenAI delta to Google GenAI parts format with tool call accumulation"""
"""Transforms OpenAI delta to Google GenAI parts, accumulating streaming tool calls."""
# 1. Initialize wrapper state if it doesn't exist
if not hasattr(wrapper, "accumulated_tool_calls"):
wrapper.accumulated_tool_calls = {}
parts: List[Dict[str, Any]] = []
# Add text content if present
if hasattr(delta, "content") and delta.content:
parts.append({"text": delta.content})
# Handle tool calls with accumulation for streaming
if hasattr(delta, "tool_calls") and delta.tool_calls:
for tool_call in delta.tool_calls:
if hasattr(tool_call, "function") and tool_call.function:
tool_call_id = getattr(tool_call, "id", "") or "call_unknown"
function_name = getattr(tool_call.function, "name", "") or ""
args_str = getattr(tool_call.function, "arguments", "") or ""
# 2. Ensure tool_calls is iterable
tool_calls = delta.tool_calls or []
# Initialize accumulation for this tool call if not exists
if tool_call_id not in wrapper.accumulated_tool_calls:
wrapper.accumulated_tool_calls[tool_call_id] = {
"name": "",
"arguments": "",
"complete": False,
}
for tool_call in tool_calls:
if not hasattr(tool_call, "function"):
continue
# Accumulate function name if provided
if function_name:
wrapper.accumulated_tool_calls[tool_call_id][
"name"
] = function_name
# 3. Use `index` as the primary key for accumulation
tool_call_index = getattr(tool_call, "index", None)
if tool_call_index is None:
continue # Index is essential for tracking streaming tool calls
# Accumulate arguments if provided
if args_str:
wrapper.accumulated_tool_calls[tool_call_id][
"arguments"
] += args_str
# Initialize accumulator for this index if it's new
if tool_call_index not in wrapper.accumulated_tool_calls:
wrapper.accumulated_tool_calls[tool_call_index] = {
"name": "",
"arguments": "",
}
# Try to parse the accumulated arguments as JSON
accumulated_args = wrapper.accumulated_tool_calls[tool_call_id][
"arguments"
]
try:
if accumulated_args:
parsed_args = json.loads(accumulated_args)
# JSON is valid, mark as complete and create function call part
wrapper.accumulated_tool_calls[tool_call_id][
"complete"
] = True
# Accumulate name and arguments
function_name = getattr(tool_call.function, "name", None)
args_chunk = getattr(tool_call.function, "arguments", None)
function_call_part = {
"functionCall": {
"name": wrapper.accumulated_tool_calls[
tool_call_id
]["name"],
"args": parsed_args,
}
}
parts.append(function_call_part)
# Optimization: Skip chunks that have no new data
if not function_name and not args_chunk:
verbose_logger.debug(
f"Skipping empty tool call chunk for index: {tool_call_index}"
)
continue
# Clean up completed tool call
del wrapper.accumulated_tool_calls[tool_call_id]
if function_name:
wrapper.accumulated_tool_calls[tool_call_index]["name"] = function_name
except json.JSONDecodeError:
# JSON is still incomplete, continue accumulating
# Don't add to parts yet
pass
if args_chunk:
wrapper.accumulated_tool_calls[tool_call_index][
"arguments"
] += args_chunk
# Attempt to parse and emit a complete tool call
accumulated_data = wrapper.accumulated_tool_calls[tool_call_index]
accumulated_name = accumulated_data["name"]
accumulated_args = accumulated_data["arguments"]
# 5. Attempt to parse arguments even if name hasn't arrived.
try:
# Attempt to parse the accumulated arguments string
parsed_args = json.loads(accumulated_args)
# If parsing succeeds, but we don't have a name yet, wait.
# The part will be created by a later chunk that brings the name.
if accumulated_name:
# If successful, create the part and clean up
function_call_part = {
"functionCall": {"name": accumulated_name, "args": parsed_args}
}
parts.append(function_call_part)
# Remove the completed tool call from the accumulator
del wrapper.accumulated_tool_calls[tool_call_index]
except json.JSONDecodeError:
# The JSON for arguments is still incomplete.
# We will continue to accumulate and wait for more chunks.
pass
return parts

View file

@ -85,7 +85,6 @@ class GenerateContentHelper:
contents: GenerateContentContentListUnionDict,
config: Optional[GenerateContentConfigDict] = None,
custom_llm_provider: Optional[str] = None,
stream: bool = False,
tools: Optional[ToolConfigDict] = None,
**kwargs,
) -> GenerateContentSetupResult:
@ -97,8 +96,7 @@ class GenerateContentHelper:
contents: The content to generate from
config: Optional configuration
custom_llm_provider: Optional custom LLM provider
stream: Whether this is a streaming call
local_vars: Local variables from the calling function
tools: Optional tools
**kwargs: Additional keyword arguments
Returns:
@ -114,7 +112,7 @@ class GenerateContentHelper:
## MOCK RESPONSE LOGIC (only for non-streaming)
if (
not stream
not kwargs.get("stream", False)
and litellm_params.mock_response
and isinstance(litellm_params.mock_response, str)
):
@ -289,7 +287,7 @@ def generate_content(
"""
local_vars = locals()
try:
_is_async = kwargs.pop("agenerate_content", False) is True
_is_async = kwargs.pop("agenerate_content", False)
# Handle generationConfig parameter from kwargs for backward compatibility
if "generationConfig" in kwargs and config is None:
@ -309,7 +307,6 @@ def generate_content(
contents=contents,
config=config,
custom_llm_provider=custom_llm_provider,
stream=False,
tools=tools,
**kwargs,
)
@ -321,7 +318,7 @@ def generate_content(
model=model,
contents=contents, # type: ignore
config=setup_result.generate_content_config_dict,
stream=False,
tools=tools,
_is_async=_is_async,
litellm_params=setup_result.litellm_params,
**kwargs,
@ -342,7 +339,6 @@ def generate_content(
timeout=timeout or request_timeout,
_is_async=_is_async,
client=kwargs.get("client"),
stream=False,
litellm_metadata=kwargs.get("litellm_metadata", {}),
)
@ -391,15 +387,12 @@ async def agenerate_content_stream(
# Setup the call
setup_result = GenerateContentHelper.setup_generate_content_call(
**{
"model": model,
"contents": contents,
"config": config,
"custom_llm_provider": custom_llm_provider,
"stream": True,
"tools": tools,
**kwargs,
}
model=model,
contents=contents,
config=config,
custom_llm_provider=custom_llm_provider,
tools=tools,
**kwargs,
)
# Check if we should use the adapter (when provider config is None)
@ -411,7 +404,7 @@ async def agenerate_content_stream(
contents=contents, # type: ignore
config=setup_result.generate_content_config_dict,
litellm_params=setup_result.litellm_params,
stream=True,
tools=tools,
**kwargs,
)
)
@ -479,7 +472,6 @@ def generate_content_stream(
contents=contents,
config=config,
custom_llm_provider=custom_llm_provider,
stream=True,
tools=tools,
**kwargs,
)
@ -491,7 +483,6 @@ def generate_content_stream(
model=model,
contents=contents, # type: ignore
config=setup_result.generate_content_config_dict,
stream=True,
_is_async=_is_async,
litellm_params=setup_result.litellm_params,
**kwargs,

View file

@ -1,5 +1,5 @@
from datetime import datetime
from typing import Any, Dict, List, Literal, Optional, Type, Union, get_args
from typing import Any, Dict, List, Optional, Type, Union, get_args
from litellm._logging import verbose_logger
from litellm.caching import DualCache
@ -14,6 +14,7 @@ from litellm.types.guardrails import (
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
from litellm.types.utils import (
CallTypes,
GuardrailStatus,
LLMResponseTypes,
StandardLoggingGuardrailInformation,
)
@ -352,7 +353,7 @@ class CustomGuardrail(CustomLogger):
self,
guardrail_json_response: Union[Exception, str, dict, List[dict]],
request_data: dict,
guardrail_status: Literal["success", "failure", "blocked"],
guardrail_status: GuardrailStatus,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
@ -460,7 +461,7 @@ class CustomGuardrail(CustomLogger):
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=e,
request_data=request_data,
guardrail_status="failure",
guardrail_status="guardrail_failed_to_respond",
duration=duration,
start_time=start_time,
end_time=end_time,

View file

@ -0,0 +1,317 @@
# LiteLLM gitlab Prompt Management
A powerful prompt management system for LiteLLM that fetches `.prompt` files from gitlab repositories. This enables team-based prompt management with gitlab's built-in access control and version control capabilities.
## Features
- **🏢 Team-based access control**: Leverage gitlab's workspace and repository permissions
- **📁 Repository-based prompt storage**: Store prompts in gitlab repositories
- **🔐 Multiple authentication methods**: Support for access tokens and basic auth
- **🎯 YAML frontmatter**: Define model, parameters, and schemas in file headers
- **🔧 Handlebars templating**: Use `{{variable}}` syntax with Jinja2 backend
- **✅ Input validation**: Automatic validation against defined schemas
- **🔗 LiteLLM integration**: Works seamlessly with `litellm.completion()`
- **💬 Smart message parsing**: Converts prompts to proper chat messages
- **⚙️ Parameter extraction**: Automatically applies model settings from prompts
## Quick Start
### 1. Set up gitlab Repository
Create a repository in your gitlab workspace and add `.prompt` files:
```
your-repo/
├── prompts/
│ ├── chat_assistant.prompt
│ ├── code_reviewer.prompt
│ └── data_analyst.prompt
```
### 2. Create a `.prompt` file
Create a file called `prompts/chat_assistant.prompt`:
```yaml
---
model: gpt-4
temperature: 0.7
max_tokens: 150
input:
schema:
user_message: string
system_context?: string
---
{% if system_context %}System: {{system_context}}
{% endif %}User: {{user_message}}
```
### 3. Configure gitlab Access
#### Option A: Access Token (Recommended)
```python
import litellm
# Configure gitlab access
gitlab_config = {
"project": "a/b/<repo_name>",
"access_token": "your-access-token",
"base_url": "gitlab url",
"prompts_path": "src/prompts", # folder to point to, defaults to root
"branch":"main" # optional, defaults to main
}
# Set global gitlab configuration
litellm.set_global_gitlab_config(gitlab_config)
```
#### Option B: Basic Authentication
```python
import litellm
# Configure gitlab access with basic auth
gitlab_config = {
"project": "a/b/<repo_name>",
"base_url": "base url",
"access_token": "your-app-password", # Use app password for basic auth
"branch": "main",
"prompts_path": "src/prompts", # folder to point to, defaults to root
}
litellm.set_global_gitlab_config(gitlab_config)
```
### 4. Use with LiteLLM
```python
# Use with completion - the model prefix 'gitlab/' tells LiteLLM to use gitlab prompt management
response = litellm.completion(
model="gitlab/gpt-4", # The actual model comes from the .prompt file
prompt_id="prompts/chat_assistant", # Location of the prompt file
prompt_variables={
"user_message": "What is machine learning?",
"system_context": "You are a helpful AI tutor."
},
# Any additional messages will be appended after the prompt
messages=[{"role": "user", "content": "Please explain it simply."}]
)
print(response.choices[0].message.content)
```
## Proxy Server Configuration
### 1. Create a `.prompt` file
Create `prompts/hello.prompt`:
```yaml
---
model: gpt-4
temperature: 0.7
---
System: You are a helpful assistant.
User: {{user_message}}
```
### 2. Setup config.yaml
```yaml
model_list:
- model_name: my-gitlab-model
litellm_params:
model: gitlab/gpt-4
prompt_id: "prompts/hello"
api_key: os.environ/OPENAI_API_KEY
litellm_settings:
global_gitlab_config:
workspace: "your-workspace"
repository: "your-repo"
access_token: "your-access-token"
branch: "main"
```
### 3. Start the proxy
```bash
litellm --config config.yaml --detailed_debug
```
### 4. Test it!
```bash
curl -L -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-d '{
"model": "my-gitlab-model",
"messages": [{"role": "user", "content": "IGNORED"}],
"prompt_variables": {
"user_message": "What is the capital of France?"
}
}'
```
## Prompt File Format
### Basic Structure
```yaml
---
# Model configuration
model: gpt-4
temperature: 0.7
max_tokens: 500
# Input schema (optional)
input:
schema:
user_message: string
system_context?: string
---
System: You are a helpful {{role}} assistant.
User: {{user_message}}
```
### Advanced Features
**Multi-role conversations:**
```yaml
---
model: gpt-4
temperature: 0.3
---
System: You are a helpful coding assistant.
User: {{user_question}}
```
**Dynamic model selection:**
```yaml
---
model: "{{preferred_model}}" # Model can be a variable
temperature: 0.7
---
System: You are a helpful assistant specialized in {{domain}}.
User: {{user_message}}
```
## Team-Based Access Control
gitlab's built-in permission system provides team-based access control:
1. **Workspace-level permissions**: Control access to entire workspaces
2. **Repository-level permissions**: Control access to specific repositories
3. **Branch-level permissions**: Control access to specific branches
4. **User and group management**: Manage team members and their access levels
### Setting up Team Access
1. **Create workspaces for each team**:
```
team-a-prompts/
team-b-prompts/
team-c-prompts/
```
2. **Configure repository permissions**:
- Grant read access to team members
- Grant write access to prompt maintainers
- Use branch protection rules for production prompts
3. **Use different access tokens**:
- Each team can have their own access token
- Tokens can be scoped to specific repositories
- Use app passwords for additional security
## API Reference
### gitlab Configuration
```python
gitlab_config = {
"workspace": str, # Required: gitlab workspace name
"repository": str, # Required: Repository name
"access_token": str, # Required: gitlab access token or app password
"branch": str, # Optional: Branch to fetch from (default: "main")
"base_url": str, # Optional: Custom gitlab API URL
"auth_method": str, # Optional: "token" or "basic" (default: "token")
"username": str, # Optional: Username for basic auth
"base_url" : str # Optional: Incase where the base url is not https://api.gitlab.org/2.0
}
```
### LiteLLM Integration
```python
response = litellm.completion(
model="gitlab/<base_model>", # required (e.g., gitlab/gpt-4)
prompt_id=str, # required - the .prompt filename without extension
prompt_variables=dict, # optional - variables for template rendering
gitlab_config=dict, # optional - gitlab configuration (if not set globally)
messages=list, # optional - additional messages
)
```
## Error Handling
The gitlab integration provides detailed error messages for common issues:
- **Authentication errors**: Invalid access tokens or credentials
- **Permission errors**: Insufficient access to workspace/repository
- **File not found**: Missing .prompt files
- **Network errors**: Connection issues with gitlab API
## Security Considerations
1. **Access Token Security**: Store access tokens securely using environment variables or secret management systems
2. **Repository Permissions**: Use gitlab's permission system to control access
3. **Branch Protection**: Protect main branches from unauthorized changes
4. **Audit Logging**: gitlab provides audit logs for all repository access
## Troubleshooting
### Common Issues
1. **"Access denied" errors**: Check your gitlab permissions for the workspace and repository
2. **"Authentication failed" errors**: Verify your access token or credentials
3. **"File not found" errors**: Ensure the .prompt file exists in the specified branch
4. **Template rendering errors**: Check your Handlebars syntax in the .prompt file
### Debug Mode
Enable debug logging to troubleshoot issues:
```python
import litellm
litellm.set_verbose = True
# Your gitlab prompt calls will now show detailed logs
response = litellm.completion(
model="gitlab/gpt-4",
prompt_id="your_prompt",
prompt_variables={"key": "value"}
)
```
## Migration from File-Based Prompts
If you're currently using file-based prompts with the dotprompt integration, you can easily migrate to gitlab:
1. **Upload your .prompt files** to a gitlab repository
2. **Update your configuration** to use gitlab instead of local files
3. **Set up team access** using gitlab's permission system
4. **Update your code** to use `gitlab/` model prefix instead of `dotprompt/`
This provides better collaboration, version control, and team-based access control for your prompts.

View file

@ -0,0 +1,95 @@
from typing import TYPE_CHECKING, Optional, Dict, Any
if TYPE_CHECKING:
from .gitlab_prompt_manager import GitLabPromptManager
from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.types.prompts.init_prompts import SupportedPromptIntegrations
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.types.prompts.init_prompts import PromptSpec, PromptLiteLLMParams
from .gitlab_prompt_manager import GitLabPromptManager
# Global instances
global_gitlab_config: Optional[dict] = None
def set_global_gitlab_config(config: dict) -> None:
"""
Set the global BitBucket configuration for prompt management.
Args:
config: Dictionary containing BitBucket configuration
- workspace: BitBucket workspace name
- repository: Repository name
- access_token: BitBucket access token
- branch: Branch to fetch prompts from (default: main)
"""
import litellm
litellm.global_gitlab_config = config # type: ignore
def prompt_initializer(
litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec"
) -> "CustomPromptManagement":
"""
Initialize a prompt from a BitBucket repository.
"""
gitlab_config = getattr(litellm_params, "gitlab_config", None)
prompt_id = getattr(litellm_params, "prompt_id", None)
if not gitlab_config:
raise ValueError(
"bitbucket_config is required for BitBucket prompt integration"
)
try:
bitbucket_prompt_manager = GitLabPromptManager(
gitlab_config=gitlab_config,
prompt_id=prompt_id,
)
return bitbucket_prompt_manager
except Exception as e:
raise e
def _gitlab_prompt_initializer(
litellm_params: PromptLiteLLMParams,
prompt: PromptSpec,
) -> CustomPromptManagement:
"""
Build a GitLab-backed prompt manager for this prompt.
Expected fields on litellm_params:
- prompt_integration="gitlab" (handled by the caller)
- gitlab_config: Dict[str, Any] (project/access_token/branch/prompts_path/etc.)
- git_ref (optional): per-prompt tag/branch/SHA override
"""
# You can store arbitrary integration-specific config on PromptLiteLLMParams.
# If your dataclass doesn't have these attributes, add them or put inside
# `litellm_params.extra` and pull them from there.
gitlab_config: Dict[str, Any] = getattr(litellm_params, "gitlab_config", None) or {}
git_ref: Optional[str] = getattr(litellm_params, "git_ref", None)
if not gitlab_config:
raise ValueError("gitlab_config is required for gitlab prompt integration")
# prompt.prompt_id can map to a file path under prompts_path (e.g. "chat/greet/hi")
return GitLabPromptManager(
gitlab_config=gitlab_config,
prompt_id=prompt.prompt_id,
ref=git_ref,
)
prompt_initializer_registry = {
SupportedPromptIntegrations.GITLAB.value: _gitlab_prompt_initializer,
}
# Export public API
__all__ = [
"GitLabPromptManager",
"set_global_gitlab_config",
"global_gitlab_config",
]

View file

@ -0,0 +1,285 @@
"""
GitLab API client for fetching files from GitLab repositories.
Now supports selecting a tag via `config["tag"]`; falls back to branch ("main").
"""
import base64
from typing import Any, Dict, List, Optional
from urllib.parse import quote
from litellm.llms.custom_httpx.http_handler import HTTPHandler
class GitLabClient:
"""
Client for interacting with the GitLab API to fetch files.
Supports:
- Authentication with personal/access tokens or OAuth bearer tokens
- Fetching file contents from repositories (raw endpoint with JSON fallback)
- Namespace/project path or numeric project ID addressing
- Ref selection via tag (preferred) or branch (default "main")
- Directory listing via the repository tree API
"""
def __init__(self, config: Dict[str, Any]):
"""
Initialize the GitLab client.
Args:
config: Dictionary containing:
- project: Project path ("group/subgroup/repo") or numeric project ID (str|int) [required]
- access_token: GitLab personal/access token or OAuth token [required] (str)
- auth_method: 'token' (default; sends Private-Token) or 'oauth' (Authorization: Bearer)
- tag: Tag name to fetch from (takes precedence over branch if provided)
- branch: Branch to fetch from (default: "main")
- base_url: Base GitLab API URL (default: "https://gitlab.com/api/v4")
"""
project = config.get("project")
access_token = config.get("access_token")
if project is None or access_token is None:
raise ValueError("project and access_token are required")
self.project: str | int = project
self.access_token: str = str(access_token)
self.auth_method = config.get("auth_method", "token") # 'token' or 'oauth'
self.branch = config.get("branch", None)
if not self.branch:
self.branch = 'main'
self.tag = config.get("tag")
self.base_url = config.get("base_url", "https://gitlab.com/api/v4")
if not all([self.project, self.access_token]):
raise ValueError("project and access_token are required")
# Effective ref: prefer tag if provided, else branch ("main")
self.ref = str(self.tag or self.branch)
# Build headers
self.headers = {
"Accept": "application/json",
"Content-Type": "application/json",
}
if self.auth_method == "oauth":
self.headers["Authorization"] = f"Bearer {self.access_token}"
else:
# Default GitLab token header
self.headers["Private-Token"] = self.access_token
# Project identifier must be URL-encoded (slashes become %2F)
self._project_enc = quote(str(self.project), safe="")
# HTTP handler
self.http_handler = HTTPHandler()
# ------------------------
# Core helpers
# ------------------------
def _file_raw_url(self, file_path: str, *, ref: Optional[str] = None) -> str:
file_enc = quote(file_path, safe="")
ref_q = quote(ref or self.ref, safe="")
return f"{self.base_url}/projects/{self._project_enc}/repository/files/{file_enc}/raw?ref={ref_q}"
def _file_json_url(self, file_path: str, *, ref: Optional[str] = None) -> str:
file_enc = quote(file_path, safe="")
ref_q = quote(ref or self.ref, safe="")
return f"{self.base_url}/projects/{self._project_enc}/repository/files/{file_enc}?ref={ref_q}"
def _tree_url(self, directory_path: str = "", recursive: bool = False, *, ref: Optional[str] = None) -> str:
path_q = f"&path={quote(directory_path, safe='')}" if directory_path else ""
rec_q = "&recursive=true" if recursive else ""
ref_q = quote(ref or self.ref, safe="")
return f"{self.base_url}/projects/{self._project_enc}/repository/tree?ref={ref_q}{path_q}{rec_q}"
# ------------------------
# Public API
# ------------------------
def set_ref(self, ref: str) -> None:
"""Override the default ref (tag/branch) for subsequent calls."""
if not ref:
raise ValueError("ref must be a non-empty string")
self.ref = ref
def get_file_content(self, file_path: str, *, ref: Optional[str] = None) -> Optional[str]:
"""
Fetch the content of a file from the GitLab repository at the given ref
(tag, branch, or commit SHA). If `ref` is None, uses self.ref.
Strategy:
1) Try the RAW endpoint (returns bytes of the file)
2) Fallback to the JSON endpoint (returns base64-encoded content)
Returns:
File content as UTF-8 string, or None if file not found.
"""
raw_url = self._file_raw_url(file_path, ref=ref)
try:
resp = self.http_handler.get(raw_url, headers=self.headers)
if resp.status_code == 404:
# Fallback to JSON endpoint
return self._get_file_content_via_json(file_path, ref=ref)
resp.raise_for_status()
ctype = (resp.headers.get("content-type") or "").lower()
if ctype.startswith("text/") or "charset=" in ctype or ctype.startswith("application/json"):
return resp.text
try:
return resp.content.decode("utf-8")
except Exception:
return resp.content.decode("utf-8", errors="replace")
except Exception as e:
status = getattr(getattr(e, "response", None), "status_code", None)
if status == 404:
return None
if status == 403:
raise Exception(
f"Access denied to file '{file_path}'. Check your GitLab permissions for project '{self.project}'."
)
if status == 401:
raise Exception("Authentication failed. Check your GitLab token and auth_method.")
raise Exception(f"Failed to fetch file '{file_path}': {e}")
def _get_file_content_via_json(self, file_path: str, *, ref: Optional[str] = None) -> Optional[str]:
"""
Fallback for get_file_content(): use the JSON file API which returns base64 content.
"""
json_url = self._file_json_url(file_path, ref=ref)
try:
resp = self.http_handler.get(json_url, headers=self.headers)
if resp.status_code == 404:
return None
resp.raise_for_status()
data = resp.json()
content = data.get("content")
encoding = data.get("encoding", "")
if content and encoding == "base64":
try:
return base64.b64decode(content).decode("utf-8")
except Exception:
return base64.b64decode(content).decode("utf-8", errors="replace")
return content
except Exception as e:
status = getattr(getattr(e, "response", None), "status_code", None)
if status == 404:
return None
if status == 403:
raise Exception(
f"Access denied to file '{file_path}'. Check your GitLab permissions for project '{self.project}'."
)
if status == 401:
raise Exception("Authentication failed. Check your GitLab token and auth_method.")
raise Exception(f"Failed to fetch file '{file_path}' via JSON endpoint: {e}")
def list_files(
self,
directory_path: str = "",
file_extension: str = ".prompt",
recursive: bool = False,
*,
ref: Optional[str] = None,
) -> List[str]:
"""
List files in a directory with a specific extension using the repository tree API.
Args:
directory_path: Directory path in the repository (empty for repo root)
file_extension: File extension to filter by (default: .prompt)
recursive: If True, traverses subdirectories
ref: Optional override (tag/branch/SHA). Defaults to self.ref.
Returns:
List of file paths (relative to repo root)
"""
url = self._tree_url(directory_path, recursive=recursive, ref=ref)
try:
resp = self.http_handler.get(url, headers=self.headers)
if resp.status_code == 404:
return []
resp.raise_for_status()
data = resp.json() or []
files: List[str] = []
for item in data:
if item.get("type") == "blob":
file_path = item.get("path", "")
if not file_extension or file_path.endswith(file_extension):
files.append(file_path)
return files
except Exception as e:
status = getattr(getattr(e, "response", None), "status_code", None)
if status == 404:
return []
if status == 403:
raise Exception(
f"Access denied to directory '{directory_path}'. Check your GitLab permissions for project '{self.project}'."
)
if status == 401:
raise Exception("Authentication failed. Check your GitLab token and auth_method.")
raise Exception(f"Failed to list files in '{directory_path}': {e}")
def get_repository_info(self) -> Dict[str, Any]:
"""Get information about the project/repository."""
url = f"{self.base_url}/projects/{self._project_enc}"
try:
resp = self.http_handler.get(url, headers=self.headers)
resp.raise_for_status()
return resp.json()
except Exception as e:
raise Exception(f"Failed to get repository info: {e}")
def test_connection(self) -> bool:
"""Test the connection to the GitLab project."""
try:
self.get_repository_info()
return True
except Exception:
return False
def get_branches(self) -> List[Dict[str, Any]]:
"""Get list of branches in the repository."""
url = f"{self.base_url}/projects/{self._project_enc}/repository/branches"
try:
resp = self.http_handler.get(url, headers=self.headers)
resp.raise_for_status()
data = resp.json()
return data if isinstance(data, list) else []
except Exception as e:
raise Exception(f"Failed to get branches: {e}")
def get_file_metadata(self, file_path: str, *, ref: Optional[str] = None) -> Optional[Dict[str, Any]]:
"""
Get minimal metadata about a file via RAW endpoint headers at a given ref.
Args:
file_path: Path to the file in the repository.
ref: Optional override (tag/branch/SHA). Defaults to self.ref.
"""
url = self._file_raw_url(file_path, ref=ref)
try:
headers = dict(self.headers)
headers["Range"] = "bytes=0-0"
resp = self.http_handler.get(url, headers=headers)
if resp.status_code == 404:
return None
resp.raise_for_status()
return {
"content_type": resp.headers.get("content-type"),
"content_length": resp.headers.get("content-length"),
"last_modified": resp.headers.get("last-modified"),
}
except Exception as e:
status = getattr(getattr(e, "response", None), "status_code", None)
if status == 404:
return None
raise Exception(f"Failed to get file metadata for '{file_path}': {e}")
def close(self):
"""Close the HTTP handler to free resources."""
if hasattr(self, "http_handler"):
self.http_handler.close()

View file

@ -0,0 +1,488 @@
"""
GitLab prompt manager with configurable prompts folder.
"""
from typing import Any, Dict, List, Optional, Tuple, Union
from jinja2 import DictLoader, Environment, select_autoescape
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.integrations.prompt_management_base import (
PromptManagementBase,
PromptManagementClient,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import StandardCallbackDynamicParams
from litellm.integrations.gitlab.gitlab_client import GitLabClient
class GitLabPromptTemplate:
def __init__(
self,
template_id: str,
content: str,
metadata: Dict[str, Any],
model: Optional[str] = None,
):
self.template_id = template_id
self.content = content
self.metadata = metadata
self.model = model or metadata.get("model")
self.temperature = metadata.get("temperature")
self.max_tokens = metadata.get("max_tokens")
self.input_schema = metadata.get("input", {}).get("schema", {})
self.optional_params = {
k: v for k, v in metadata.items() if k not in ["model", "input", "content"]
}
def __repr__(self):
return f"GitLabPromptTemplate(id='{self.template_id}', model='{self.model}')"
class GitLabTemplateManager:
"""
Manager for loading and rendering .prompt files from GitLab repositories.
New: supports `prompts_path` (or `folder`) in gitlab_config to scope where prompts live.
"""
def __init__(
self,
gitlab_config: Dict[str, Any],
prompt_id: Optional[str] = None,
ref: Optional[str] = None,
gitlab_client: Optional[GitLabClient] = None
):
self.gitlab_config = dict(gitlab_config)
self.prompt_id = prompt_id
self.prompts: Dict[str, GitLabPromptTemplate] = {}
self.gitlab_client = gitlab_client or GitLabClient(self.gitlab_config)
if ref:
self.gitlab_client.set_ref(ref)
# Folder inside repo to look for prompts (e.g., "prompts" or "prompts/chat")
self.prompts_path: str = (
self.gitlab_config.get("prompts_path")
or self.gitlab_config.get("folder")
or ""
).strip("/")
self.jinja_env = Environment(
loader=DictLoader({}),
autoescape=select_autoescape(["html", "xml"]),
variable_start_string="{{",
variable_end_string="}}",
block_start_string="{%",
block_end_string="%}",
comment_start_string="{#",
comment_end_string="#}",
)
if self.prompt_id:
self._load_prompt_from_gitlab(self.prompt_id)
# ---------- path helpers ----------
def _id_to_repo_path(self, prompt_id: str) -> str:
"""Map a prompt_id to a repo path (respects prompts_path and adds .prompt)."""
if self.prompts_path:
return f"{self.prompts_path}/{prompt_id}.prompt"
return f"{prompt_id}.prompt"
def _repo_path_to_id(self, repo_path: str) -> str:
"""
Map a repo path like 'prompts/chat/greeting.prompt' to an ID relative
to prompts_path without the extension (e.g., 'chat/greeting').
"""
path = repo_path.strip("/")
if self.prompts_path and path.startswith(self.prompts_path.strip("/") + "/"):
path = path[len(self.prompts_path.strip("/")) + 1 :]
if path.endswith(".prompt"):
path = path[: -len(".prompt")]
return path
# ---------- loading ----------
def _load_prompt_from_gitlab(self, prompt_id: str, *, ref: Optional[str] = None) -> None:
"""Load a specific .prompt file from GitLab (scoped under prompts_path if set)."""
try:
file_path = self._id_to_repo_path(prompt_id)
prompt_content = self.gitlab_client.get_file_content(file_path, ref=ref)
if prompt_content:
template = self._parse_prompt_file(prompt_content, prompt_id)
self.prompts[prompt_id] = template
except Exception as e:
raise Exception(f"Failed to load prompt '{prompt_id}' from GitLab: {e}")
def load_all_prompts(self, *, recursive: bool = True) -> List[str]:
"""
Eagerly load all .prompt files from prompts_path. Returns loaded IDs.
"""
files = self.list_templates(recursive=recursive) # reuse logic
loaded: List[str] = []
for pid in files:
if pid not in self.prompts:
self._load_prompt_from_gitlab(pid)
loaded.append(pid)
return loaded
# ---------- parsing & rendering ----------
def _parse_prompt_file(
self, content: str, prompt_id: str
) -> GitLabPromptTemplate:
if content.startswith("---"):
parts = content.split("---", 2)
if len(parts) >= 3:
frontmatter_str = parts[1].strip()
template_content = parts[2].strip()
else:
frontmatter_str = ""
template_content = content
else:
frontmatter_str = ""
template_content = content
metadata: Dict[str, Any] = {}
if frontmatter_str:
try:
import yaml
metadata = yaml.safe_load(frontmatter_str) or {}
except ImportError:
metadata = self._parse_yaml_basic(frontmatter_str)
except Exception:
metadata = {}
return GitLabPromptTemplate(
template_id=prompt_id,
content=template_content,
metadata=metadata,
)
def _parse_yaml_basic(self, yaml_str: str) -> Dict[str, Any]:
result: Dict[str, Any] = {}
for line in yaml_str.split("\n"):
line = line.strip()
if ":" in line and not line.startswith("#"):
key, value = line.split(":", 1)
key = key.strip()
value = value.strip()
if value.lower() in ["true", "false"]:
result[key] = value.lower() == "true"
elif value.isdigit():
result[key] = int(value)
elif value.replace(".", "").isdigit():
try:
result[key] = float(value)
except Exception:
result[key] = value
else:
result[key] = value.strip("\"'")
return result
def render_template(
self, template_id: str, variables: Optional[Dict[str, Any]] = None
) -> str:
if template_id not in self.prompts:
raise ValueError(f"Template '{template_id}' not found")
template = self.prompts[template_id]
jinja_template = self.jinja_env.from_string(template.content)
return jinja_template.render(**(variables or {}))
def get_template(self, template_id: str) -> Optional[GitLabPromptTemplate]:
return self.prompts.get(template_id)
def list_templates(self, *, recursive: bool = True) -> List[str]:
"""
List available prompt IDs discovered under prompts_path (no extension, relative to prompts_path).
"""
"""
List available prompt IDs under prompts_path (no extension).
Compatible with both list_files signatures:
- list_files(directory_path=..., file_extension=..., recursive=...)
- list_files(path=..., ref=None, recursive=...)
"""
# First try the "new" signature (directory_path/file_extension)
try:
files = self.gitlab_client.list_files(
directory_path=self.prompts_path,
file_extension=".prompt",
recursive=recursive,
)
base = self.prompts_path.strip("/")
out: List[str] = []
for p in files or []:
path = str(p).strip("/")
if base and not path.startswith(base + "/"):
# if the client returns extra files outside the folder, skip them
continue
if not path.endswith(".prompt"):
continue
out.append(self._repo_path_to_id(path))
return out
except TypeError:
# Fallback to the "classic" signature
raw = self.gitlab_client.list_files(
directory_path=self.prompts_path or "",
ref=None,
recursive=recursive,
)
# Classic returns GitLab tree entries; filter *.prompt blobs
files = []
for f in (raw or []):
if isinstance(f, dict) and f.get("type") == "blob" and str(f.get("path", "")).endswith(".prompt") and 'path' in f:
files.append(f['path'])
return [self._repo_path_to_id(p) for p in files]
class GitLabPromptManager(CustomPromptManagement):
"""
GitLab prompt manager with folder support.
Example config:
gitlab_config = {
"project": "group/subgroup/repo",
"access_token": "glpat_***",
"tag": "v1.2.3", # optional; takes precedence
"branch": "main", # default fallback
"prompts_path": "prompts/chat" # <--- NEW
}
"""
def __init__(
self,
gitlab_config: Dict[str, Any],
prompt_id: Optional[str] = None,
ref: Optional[str] = None, # tag/branch/SHA override
gitlab_client: Optional[GitLabClient] = None
):
self.gitlab_config = gitlab_config
self.prompt_id = prompt_id
self._prompt_manager: Optional[GitLabTemplateManager] = None
self._ref_override = ref
self._injected_gitlab_client = gitlab_client
if self.prompt_id:
self._prompt_manager = GitLabTemplateManager(
gitlab_config=self.gitlab_config,
prompt_id=self.prompt_id,
ref=self._ref_override,
)
@property
def integration_name(self) -> str:
return "gitlab"
@property
def prompt_manager(self) -> GitLabTemplateManager:
if self._prompt_manager is None:
self._prompt_manager = GitLabTemplateManager(
gitlab_config=self.gitlab_config,
prompt_id=self.prompt_id,
ref=self._ref_override,
gitlab_client=self._injected_gitlab_client
)
return self._prompt_manager
def get_prompt_template(
self,
prompt_id: str,
prompt_variables: Optional[Dict[str, Any]] = None,
*,
ref: Optional[str] = None,
) -> Tuple[str, Dict[str, Any]]:
if prompt_id not in self.prompt_manager.prompts:
self.prompt_manager._load_prompt_from_gitlab(prompt_id, ref=ref)
template = self.prompt_manager.get_template(prompt_id)
if not template:
raise ValueError(f"Prompt template '{prompt_id}' not found")
rendered_prompt = self.prompt_manager.render_template(
prompt_id, prompt_variables or {}
)
metadata = {
"model": template.model,
"temperature": template.temperature,
"max_tokens": template.max_tokens,
**template.optional_params,
}
return rendered_prompt, metadata
def pre_call_hook(
self,
user_id: Optional[str],
messages: List[AllMessageValues],
function_call: Optional[Union[Dict[str, Any], str]] = None,
litellm_params: Optional[Dict[str, Any]] = None,
prompt_id: Optional[str] = None,
prompt_variables: Optional[Dict[str, Any]] = None,
prompt_version: Optional[str] = None,
**kwargs,
) -> Tuple[List[AllMessageValues], Optional[Dict[str, Any]]]:
if not prompt_id:
return messages, litellm_params
try:
# Precedence: explicit prompt_version → per-call git_ref kwarg → manager override → config default
git_ref = prompt_version or kwargs.get("git_ref") or self._ref_override
rendered_prompt, prompt_metadata = self.get_prompt_template(
prompt_id, prompt_variables, ref=git_ref
)
parsed_messages = self._parse_prompt_to_messages(rendered_prompt)
if parsed_messages:
final_messages: List[AllMessageValues] = parsed_messages
else:
final_messages = [{"role": "user", "content": rendered_prompt}] + messages # type: ignore
if litellm_params is None:
litellm_params = {}
if prompt_metadata.get("model"):
litellm_params["model"] = prompt_metadata["model"]
for param in ["temperature", "max_tokens", "top_p", "frequency_penalty", "presence_penalty"]:
if param in prompt_metadata:
litellm_params[param] = prompt_metadata[param]
return final_messages, litellm_params
except Exception as e:
import litellm
litellm._logging.verbose_proxy_logger.error(f"Error in GitLab prompt pre_call_hook: {e}")
return messages, litellm_params
def _parse_prompt_to_messages(self, prompt_content: str) -> List[AllMessageValues]:
messages: List[AllMessageValues] = []
lines = prompt_content.strip().split("\n")
current_role: Optional[str] = None
current_content: List[str] = []
for raw in lines:
line = raw.strip()
if not line:
continue
low = line.lower()
if low.startswith("system:"):
if current_role and current_content:
messages.append({"role": current_role, "content": "\n".join(current_content).strip()}) # type: ignore
current_role = "system"
current_content = [line[7:].strip()]
elif low.startswith("user:"):
if current_role and current_content:
messages.append({"role": current_role, "content": "\n".join(current_content).strip()}) # type: ignore
current_role = "user"
current_content = [line[5:].strip()]
elif low.startswith("assistant:"):
if current_role and current_content:
messages.append({"role": current_role, "content": "\n".join(current_content).strip()}) # type: ignore
current_role = "assistant"
current_content = [line[10:].strip()]
else:
current_content.append(line)
if current_role and current_content:
messages.append({"role": current_role, "content": "\n".join(current_content).strip()}) # type: ignore
if not messages and prompt_content.strip():
messages = [{"role": "user", "content": prompt_content.strip()}] # type: ignore
return messages
def post_call_hook(
self,
user_id: Optional[str],
response: Any,
input_messages: List[AllMessageValues],
function_call: Optional[Union[Dict[str, Any], str]] = None,
litellm_params: Optional[Dict[str, Any]] = None,
prompt_id: Optional[str] = None,
prompt_variables: Optional[Dict[str, Any]] = None,
**kwargs,
) -> Any:
return response
def get_available_prompts(self) -> List[str]:
"""
Return prompt IDs. Prefer already-loaded templates in memory to avoid
unnecessary network calls (and to make tests deterministic).
"""
ids = set(self.prompt_manager.prompts.keys())
try:
ids.update(self.prompt_manager.list_templates())
except Exception:
# If GitLab list fails (auth, network), still return what we've loaded.
pass
return sorted(ids)
def reload_prompts(self) -> None:
if self.prompt_id:
self._prompt_manager = None
_ = self.prompt_manager # trigger re-init/load
def should_run_prompt_management(
self,
prompt_id: str,
dynamic_callback_params: StandardCallbackDynamicParams,
) -> bool:
return True
def _compile_prompt_helper(
self,
prompt_id: str,
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> PromptManagementClient:
try:
if prompt_id not in self.prompt_manager.prompts:
git_ref = getattr(dynamic_callback_params, "extra", {}).get("git_ref") if hasattr(dynamic_callback_params, "extra") else None
self.prompt_manager._load_prompt_from_gitlab(prompt_id, ref=git_ref)
rendered_prompt, prompt_metadata = self.get_prompt_template(
prompt_id, prompt_variables
)
messages = self._parse_prompt_to_messages(rendered_prompt)
template_model = prompt_metadata.get("model")
optional_params: Dict[str, Any] = {}
for param in ["temperature", "max_tokens", "top_p", "frequency_penalty", "presence_penalty"]:
if param in prompt_metadata:
optional_params[param] = prompt_metadata[param]
return PromptManagementClient(
prompt_id=prompt_id,
prompt_template=messages,
prompt_template_model=template_model,
prompt_template_optional_params=optional_params,
completed_messages=None,
)
except Exception as e:
raise ValueError(f"Error compiling prompt '{prompt_id}': {e}")
def get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> Tuple[str, List[AllMessageValues], dict]:
return PromptManagementBase.get_chat_completion_prompt(
self,
model,
messages,
non_default_params,
prompt_id,
prompt_variables,
dynamic_callback_params,
prompt_label,
prompt_version,
)

View file

@ -1,6 +1,5 @@
#### What this does ####
# On success, logs events to Langfuse
import copy
import os
import traceback
from datetime import datetime
@ -11,6 +10,7 @@ from packaging.version import Version
import litellm
from litellm._logging import verbose_logger
from litellm.constants import MAX_LANGFUSE_INITIALIZED_CLIENTS
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
from litellm.secret_managers.main import str_to_bool
@ -222,7 +222,7 @@ class LangFuseLogger:
litellm_params.get("metadata", {}) or {}
) # if litellm_params['metadata'] == None
metadata = self.add_metadata_from_header(litellm_params, metadata)
optional_params = copy.deepcopy(kwargs.get("optional_params", {}))
optional_params = safe_deep_copy(kwargs.get("optional_params", {}))
prompt = {"messages": kwargs.get("messages")}
@ -690,6 +690,7 @@ class LangFuseLogger:
}
usage_details = LangfuseUsageDetails(input=_usage_obj.prompt_tokens,
output=_usage_obj.completion_tokens,
total=_usage_obj.total_tokens,
cache_creation_input_tokens=_usage_obj.get('cache_creation_input_tokens', 0),
cache_read_input_tokens=_usage_obj.get('cache_read_input_tokens', 0))

View file

@ -575,9 +575,16 @@ class OpenTelemetry(CustomLogger):
if litellm.turn_off_message_logging or not self.message_logging:
return
litellm_params = kwargs.get("litellm_params", {})
metadata = litellm_params.get("metadata", {})
generation_name = metadata.get("generation_name")
raw_span_name = generation_name if generation_name else RAW_REQUEST_SPAN_NAME
otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs)
raw_span = otel_tracer.start_span(
name=RAW_REQUEST_SPAN_NAME,
name=raw_span_name,
start_time=self._to_ns(start_time),
context=trace.set_span_in_context(parent_span),
)
@ -645,7 +652,7 @@ class OpenTelemetry(CustomLogger):
if not self.config.enable_events:
return
from opentelemetry._logs import get_logger, LogRecord
from opentelemetry._logs import LogRecord, get_logger
otel_logger = get_logger(LITELLM_LOGGER_NAME)
parent_ctx = span.get_span_context()
@ -1115,56 +1122,68 @@ class OpenTelemetry(CustomLogger):
span.set_attribute(key, primitive_value)
def set_raw_request_attributes(self, span: Span, kwargs, response_obj):
kwargs.get("optional_params", {})
litellm_params = kwargs.get("litellm_params", {}) or {}
custom_llm_provider = litellm_params.get("custom_llm_provider", "Unknown")
try:
kwargs.get("optional_params", {})
litellm_params = kwargs.get("litellm_params", {}) or {}
custom_llm_provider = litellm_params.get("custom_llm_provider", "Unknown")
_raw_response = kwargs.get("original_response")
_additional_args = kwargs.get("additional_args", {}) or {}
complete_input_dict = _additional_args.get("complete_input_dict")
#############################################
########## LLM Request Attributes ###########
#############################################
_raw_response = kwargs.get("original_response")
_additional_args = kwargs.get("additional_args", {}) or {}
complete_input_dict = _additional_args.get("complete_input_dict")
#############################################
########## LLM Request Attributes ###########
#############################################
# OTEL Attributes for the RAW Request to https://docs.anthropic.com/en/api/messages
if complete_input_dict and isinstance(complete_input_dict, dict):
for param, val in complete_input_dict.items():
self.safe_set_attribute(
span=span, key=f"llm.{custom_llm_provider}.{param}", value=val
)
# OTEL Attributes for the RAW Request to https://docs.anthropic.com/en/api/messages
if complete_input_dict and isinstance(complete_input_dict, dict):
for param, val in complete_input_dict.items():
self.safe_set_attribute(
span=span, key=f"llm.{custom_llm_provider}.{param}", value=val
)
#############################################
########## LLM Response Attributes ##########
#############################################
if _raw_response and isinstance(_raw_response, str):
# cast sr -> dict
import json
#############################################
########## LLM Response Attributes ##########
#############################################
if _raw_response and isinstance(_raw_response, str):
# cast sr -> dict
import json
try:
_raw_response = json.loads(_raw_response)
for param, val in _raw_response.items():
self.safe_set_attribute(
span=span,
key=f"llm.{custom_llm_provider}.{param}",
value=val,
)
except json.JSONDecodeError:
verbose_logger.debug(
"litellm.integrations.opentelemetry.py::set_raw_request_attributes() - raw_response not json string - {}".format(
_raw_response
)
)
try:
_raw_response = json.loads(_raw_response)
for param, val in _raw_response.items():
self.safe_set_attribute(
span=span,
key=f"llm.{custom_llm_provider}.{param}",
value=val,
key=f"llm.{custom_llm_provider}.stringified_raw_response",
value=_raw_response,
)
except json.JSONDecodeError:
verbose_logger.debug(
"litellm.integrations.opentelemetry.py::set_raw_request_attributes() - raw_response not json string - {}".format(
_raw_response
)
)
self.safe_set_attribute(
span=span,
key=f"llm.{custom_llm_provider}.stringified_raw_response",
value=_raw_response,
)
except Exception as e:
verbose_logger.exception(
"OpenTelemetry logging error in set_raw_request_attributes %s", str(e)
)
def _to_ns(self, dt):
return int(dt.timestamp() * 1e9)
def _get_span_name(self, kwargs):
litellm_params = kwargs.get("litellm_params", {})
metadata = litellm_params.get("metadata", {})
generation_name = metadata.get("generation_name")
if generation_name:
return generation_name
return LITELLM_REQUEST_SPAN_NAME
def get_traceparent_from_header(self, headers):

View file

@ -16,6 +16,7 @@ from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheCont
from litellm.integrations.argilla import ArgillaLogger
from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger
from litellm.integrations.bitbucket import BitBucketPromptManager
from litellm.integrations.gitlab import GitLabPromptManager
from litellm.integrations.braintrust_logging import BraintrustLogger
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
@ -92,6 +93,7 @@ class CustomLoggerRegistry:
"vector_store_pre_call_hook": VectorStorePreCallHook,
"dotprompt": DotpromptManager,
"bitbucket": BitBucketPromptManager,
"gitlab": GitLabPromptManager,
"cloudzero": CloudZeroLogger,
"posthog": PostHogLogger,
}

View file

@ -1498,7 +1498,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"CohereException - {original_exception.message}",
llm_provider="cohere",
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
raise original_exception
elif custom_llm_provider == "huggingface":
@ -1573,7 +1573,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"HuggingfaceException - {original_exception.message}",
llm_provider="huggingface",
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
elif custom_llm_provider == "ai21":
if hasattr(original_exception, "message"):
@ -1632,7 +1632,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"AI21Exception - {original_exception.message}",
llm_provider="ai21",
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
elif custom_llm_provider == "nlp_cloud":
if "detail" in error_str:
@ -1659,7 +1659,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"NLPCloudException - {error_str}",
model=model,
llm_provider="nlp_cloud",
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
if hasattr(
original_exception, "status_code"
@ -1719,7 +1719,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"NLPCloudException - {original_exception.message}",
llm_provider="nlp_cloud",
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
elif (
original_exception.status_code == 504
@ -1739,7 +1739,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"NLPCloudException - {original_exception.message}",
llm_provider="nlp_cloud",
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
elif custom_llm_provider == "together_ai":
try:
@ -1848,7 +1848,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"TogetherAIException - {original_exception.message}",
llm_provider="together_ai",
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
elif custom_llm_provider == "aleph_alpha":
if (
@ -1953,7 +1953,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"VLLMException - {original_exception.message}",
llm_provider="vllm",
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
elif custom_llm_provider == "azure" or custom_llm_provider == "azure_text":
message = get_error_message(error_obj=original_exception)
@ -2208,7 +2208,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message=f"APIError: {exception_provider} - {error_str}",
llm_provider=custom_llm_provider,
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
litellm_debug_info=extra_information,
)
else:
@ -2243,7 +2243,7 @@ def exception_type( # type: ignore # noqa: PLR0915
message="{} - {}".format(exception_provider, error_str),
llm_provider=custom_llm_provider,
model=model,
request=original_exception.request,
request=getattr(original_exception, "request", None),
)
else:
raise APIConnectionError(

View file

@ -368,6 +368,8 @@ def get_llm_provider( # noqa: PLR0915
# bytez models
elif model.startswith("bytez/"):
custom_llm_provider = "bytez"
elif model.startswith("lemonade/"):
custom_llm_provider = "lemonade"
elif model.startswith("heroku/"):
custom_llm_provider = "heroku"
# cometapi models
@ -379,6 +381,8 @@ def get_llm_provider( # noqa: PLR0915
custom_llm_provider = "compactifai"
elif model.startswith("ovhcloud/"):
custom_llm_provider = "ovhcloud"
elif model.startswith("lemonade/"):
custom_llm_provider = "lemonade"
if not custom_llm_provider:
if litellm.suppress_debug_info is False:
print() # noqa
@ -783,6 +787,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
or "https://api.inference.wandb.ai/v1"
) # type: ignore
dynamic_api_key = api_key or get_secret_str("WANDB_API_KEY")
elif custom_llm_provider == "lemonade":
(
api_base,
dynamic_api_key,
) = litellm.LemonadeChatConfig()._get_openai_compatible_provider_info(
api_base, api_key
)
if api_base is not None and not isinstance(api_base, str):
raise Exception("api base needs to be a string. api_base={}".format(api_base))

View file

@ -83,11 +83,13 @@ from litellm.types.mcp import MCPPostCallResponseObject
from litellm.types.rerank import RerankResponse
from litellm.types.router import CustomPricingLiteLLMParams
from litellm.types.utils import (
CachingDetails,
CallTypes,
CostBreakdown,
CostResponseTypes,
DynamicPromptManagementParamLiteral,
EmbeddingResponse,
GuardrailStatus,
ImageResponse,
LiteLLMBatch,
LiteLLMLoggingBaseClass,
@ -106,6 +108,7 @@ from litellm.types.utils import (
StandardLoggingPayload,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingPayloadStatusFields,
StandardLoggingPromptManagementMetadata,
StandardLoggingVectorStoreRequest,
TextCompletionResponse,
@ -348,6 +351,9 @@ class Logging(LiteLLMLoggingBaseClass):
# Initialize cost breakdown field
self.cost_breakdown: Optional[CostBreakdown] = None
# Init Caching related details
self.caching_details: Optional[CachingDetails] = None
self.model_call_details: Dict[str, Any] = {
"litellm_trace_id": litellm_trace_id,
"litellm_call_id": litellm_call_id,
@ -3663,6 +3669,25 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
bitbucket_logger = BitBucketPromptManager(bitbucket_config=bitbucket_config)
_in_memory_loggers.append(bitbucket_logger)
return bitbucket_logger # type: ignore
elif logging_integration == "gitlab":
from litellm.integrations.gitlab.gitlab_prompt_manager import (
GitLabPromptManager,
)
for callback in _in_memory_loggers:
if isinstance(callback, GitLabPromptManager):
return callback
# Get global BitBucket config
gitlab_config = getattr(litellm, "global_gitlab_config", None)
if gitlab_config is None:
raise ValueError(
"Gitlab configuration not found. Please set litellm.global_gitlab_config first."
)
gitlab_logger = GitLabPromptManager(gitlab_config=gitlab_config)
_in_memory_loggers.append(gitlab_logger)
return gitlab_logger # type: ignore
return None
except Exception as e:
verbose_logger.exception(
@ -4421,6 +4446,51 @@ class StandardLoggingPayloadSetup:
return request_tags
def _get_status_fields(
status: StandardLoggingPayloadStatus,
guardrail_information: Optional[dict],
error_str: Optional[str]
) -> "StandardLoggingPayloadStatusFields":
"""
Determine status fields based on request status and guardrail information.
Args:
status: Overall request status ("success" or "failure")
guardrail_information: Guardrail information from metadata
error_str: Error string if any
Returns:
StandardLoggingPayloadStatusFields with llm_api_status and guardrail_status
"""
# Mapping for legacy guardrail status values to new GuardrailStatus values
GUARDRAIL_STATUS_MAP: Dict[str, GuardrailStatus] = {
"success": "success",
"blocked": "guardrail_intervened", # legacy
"guardrail_intervened": "guardrail_intervened", # direct
"failure": "guardrail_failed_to_respond", # legacy
"guardrail_failed_to_respond": "guardrail_failed_to_respond", # direct
"not_run": "not_run"
}
# Set LLM API status
llm_api_status: StandardLoggingPayloadStatus = status
#########################################################
# Map - guardrail_information.guardrail_status to guardrail_status
#########################################################
guardrail_status: GuardrailStatus = "not_run"
if guardrail_information and isinstance(guardrail_information, dict):
raw_status = guardrail_information.get("guardrail_status", "not_run")
guardrail_status = GUARDRAIL_STATUS_MAP.get(raw_status, "not_run")
return StandardLoggingPayloadStatusFields(
llm_api_status=llm_api_status,
guardrail_status=guardrail_status
)
def get_standard_logging_object_payload(
kwargs: Optional[dict],
init_response_obj: Union[Any, BaseModel, dict],
@ -4530,7 +4600,6 @@ def get_standard_logging_object_payload(
start_time=start_time,
response_id=id,
)
_request_body = proxy_server_request.get("body", {})
end_user_id = clean_metadata["user_api_key_end_user_id"] or _request_body.get(
"user", None
@ -4586,6 +4655,11 @@ def get_standard_logging_object_payload(
cache_hit=cache_hit,
stream=stream,
status=status,
status_fields=_get_status_fields(
status=status,
guardrail_information=metadata.get("standard_logging_guardrail_information", None),
error_str=error_str
),
custom_llm_provider=cast(Optional[str], kwargs.get("custom_llm_provider")),
saved_cache_cost=saved_cache_cost,
startTime=start_time_float,

View file

@ -85,15 +85,37 @@ class ResponseMetadata:
# Set total response time if supported
if self.supports_response_time:
self.result._response_ms = total_response_time_ms
#########################################################
# 1. Add _response_ms total duration
#########################################################
self._update_hidden_params(
{
"_response_ms": total_response_time_ms,
}
)
# Calculate LiteLLM overhead
#########################################################
# 2. Add LiteLLM overhead duration
#########################################################
llm_api_duration_ms = logging_obj.model_call_details.get("llm_api_duration_ms")
if llm_api_duration_ms is not None:
overhead_ms = round(total_response_time_ms - llm_api_duration_ms, 4)
self._update_hidden_params(
{
"litellm_overhead_time_ms": overhead_ms,
"_response_ms": total_response_time_ms,
}
)
#########################################################
# 3. Add duration for reading from cache
# In this case overhead from litellm is the difference between the cache read duration and the total response time
#########################################################
if logging_obj.caching_details is not None and logging_obj.caching_details.get("cache_hit") is True and (cache_duration_ms := logging_obj.caching_details.get("cache_duration_ms")) is not None:
overhead_ms = total_response_time_ms - cache_duration_ms
self._update_hidden_params(
{
"litellm_overhead_time_ms": overhead_ms,
}
)
@ -113,6 +135,10 @@ def update_response_metadata(
) -> None:
"""
Updates response metadata including hidden params and timing metrics
Updates response metadata, adds the following:
- response._hidden_params
- response._hidden_params["litellm_overhead_time_ms"]
- response.response_time_ms
"""
if result is None:
return

View file

@ -21,6 +21,8 @@ class SensitiveDataMasker:
"access",
"private",
"certificate",
"fingerprint",
"tenancy",
}
self.visible_prefix = visible_prefix
@ -42,7 +44,14 @@ class SensitiveDataMasker:
def is_sensitive_key(self, key: str) -> bool:
key_lower = str(key).lower()
result = any(pattern in key_lower for pattern in self.sensitive_patterns)
# Split on underscores and check if any segment matches the pattern
# This avoids false positives like "max_tokens" matching "token"
# but still catches "api_key", "access_token", etc.
key_segments = key_lower.replace('-', '_').split('_')
result = any(
pattern in key_segments
for pattern in self.sensitive_patterns
)
return result
def mask_dict(

View file

@ -0,0 +1,149 @@
"""
Translate from OpenAI's `/v1/chat/completions` to Lemonade's `/v1/chat/completions`
"""
from typing import Any, List, Optional, Tuple, Union
import httpx
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
AllMessageValues,
)
from litellm.types.utils import ModelResponse
from ...openai_like.chat.transformation import OpenAILikeChatConfig
class LemonadeChatConfig(OpenAILikeChatConfig):
repeat_penalty: Optional[float] = None
functions: Optional[list] = None
logit_bias: Optional[dict] = None
max_tokens: Optional[int] = None
max_completion_tokens: Optional[int] = None
n: Optional[int] = None
presence_penalty: Optional[int] = None
stop: Optional[Union[str, list]] = None
temperature: Optional[int] = None
top_p: Optional[int] = None
top_k: Optional[int] = None
response_format: Optional[dict] = None
tools: Optional[list] = None
def __init__(
self,
repeat_penalty: Optional[float] = None,
functions: Optional[list] = None,
logit_bias: Optional[dict] = None,
max_completion_tokens: Optional[int] = None,
max_tokens: Optional[int] = None,
n: Optional[int] = None,
presence_penalty: Optional[int] = None,
stop: Optional[Union[str, list]] = None,
temperature: Optional[int] = None,
top_p: Optional[int] = None,
top_k: Optional[int] = None,
response_format: Optional[dict] = None,
tools: Optional[list] = None,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@property
def custom_llm_provider(self) -> Optional[str]:
return "lemonade"
@classmethod
def get_config(cls):
return super().get_config()
def get_models(self, api_key: Optional[str] = None, api_base: Optional[str] = None):
"""
Get available models from Lemonade API.
This method queries the Lemonade /models endpoint to retrieve the list of available models.
Args:
api_key: Optional API key (Lemonade doesn't require authentication)
api_base: Optional API base URL (defaults to LEMONADE_API_BASE env var or http://localhost:8000)
Returns:
List of model names prefixed with "lemonade/"
"""
api_base, api_key = self._get_openai_compatible_provider_info(
api_base=api_base, api_key=api_key
)
if api_base is None:
raise ValueError(
"LEMONADE_API_BASE is not set. Please set the environment variable to query Lemonade's /models endpoint."
)
# Getting the list of models from lemonade
try:
response = litellm.module_level_client.get(
url=f"{api_base}/models",
)
except Exception as e:
raise ValueError(
f"Failed to fetch models from Lemonade. Set Lemonade API Base via `LEMONADE_API_BASE` environment variable. Error: {e}"
)
if response.status_code != 200:
raise ValueError(
f"Failed to fetch models from Lemonade. Status code: {response.status_code}, Response: {response.text}"
)
model_list = response.json().get("data", [])
return ["lemonade/" + model["id"] for model in model_list]
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
# lemonade is openai compatible, we just need to set this to custom_openai and have the api_base be lemonade's endpoint
api_base = (
api_base
or get_secret_str("LEMONADE_API_BASE")
or "http://localhost:8000/api/v1"
) # type: ignore
# Lemonade doesn't check the key
key = "lemonade"
return api_base, key
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
model_response = super().transform_response(
model=model,
model_response=model_response,
raw_response=raw_response,
messages=messages,
logging_obj=logging_obj,
request_data=request_data,
encoding=encoding,
optional_params=optional_params,
json_mode=json_mode,
litellm_params=litellm_params,
api_key=api_key,
)
# Storing lemonade in the model response for easier cost calculation later
setattr(model_response, "model", "lemonade/" + model)
return model_response

View file

@ -0,0 +1,35 @@
"""
Cost calculation for Lemonade LLM provider.
Since Lemonade is a local/self-hosted service, all costs default to 0.
This prevents cost calculation errors when using models not in model_prices_and_context_window.json
"""
from typing import Tuple
from litellm.types.utils import Usage
def cost_per_token(
model: str,
usage: Usage,
) -> Tuple[float, float]:
"""
Calculate cost per token for Lemonade models.
Since Lemonade is a local/self-hosted deployment, there are no per-token costs.
This function returns (0.0, 0.0) for all models to allow cost tracking to work
without errors for any Lemonade model, regardless of whether it's in the
model_prices_and_context_window.json file.
Args:
model: The model name (with or without "lemonade/" prefix)
usage: Usage object containing token counts
Returns:
Tuple of (prompt_cost, completion_cost) - always (0.0, 0.0) for Lemonade
"""
# Lemonade is self-hosted/local, so cost is always 0
prompt_cost = 0.0
completion_cost = 0.0
return prompt_cost, completion_cost

View file

@ -8,18 +8,24 @@ from .gpt_transformation import OpenAIGPTConfig
class OpenAIGPT5Config(OpenAIGPTConfig):
"""Configuration for gpt-5 models.
"""Configuration for gpt-5 models including GPT-5-Codex variants.
Handles OpenAI API quirks for the gpt-5 series like:
- Mapping ``max_tokens`` -> ``max_completion_tokens``.
- Dropping unsupported ``temperature`` values when requested.
- Support for GPT-5-Codex models optimized for code generation.
"""
@classmethod
def is_model_gpt_5_model(cls, model: str) -> bool:
return "gpt-5" in model
@classmethod
def is_model_gpt_5_codex_model(cls, model: str) -> bool:
"""Check if the model is specifically a GPT-5 Codex variant."""
return "gpt-5-codex" in model
def get_supported_openai_params(self, model: str) -> list:
from litellm.utils import supports_tool_choice
@ -38,7 +44,9 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
]
return [
param for param in base_gpt_series_params if param not in non_supported_params
param
for param in base_gpt_series_params
if param not in non_supported_params
]
def map_openai_params(
@ -67,7 +75,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
else:
raise litellm.utils.UnsupportedParamsError(
message=(
"gpt-5 models don't support temperature={}. Only temperature=1 is supported. To drop unsupported params set `litellm.drop_params = True`"
"gpt-5 models (including gpt-5-codex) don't support temperature={}. Only temperature=1 is supported. To drop unsupported params set `litellm.drop_params = True`"
).format(temperature_value),
status_code=400,
)

View file

@ -3,7 +3,6 @@
## Initial implementation - covers gemini + image gen calls
import json
import time
from litellm._uuid import uuid
from copy import deepcopy
from functools import partial
from typing import (
@ -25,6 +24,7 @@ import litellm
import litellm.litellm_core_utils
import litellm.litellm_core_utils.litellm_logging
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.constants import (
DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
@ -32,8 +32,8 @@ from litellm.constants import (
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH,
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO,
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE,
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO,
)
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
@ -313,9 +313,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return None
for tool in value:
openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = (
None
)
openai_function_object: Optional[
ChatCompletionToolParamFunctionChunk
] = None
if "function" in tool: # tools list
_openai_function_object = ChatCompletionToolParamFunctionChunk( # type: ignore
**tool["function"]
@ -335,6 +335,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
elif "name" in tool: # functions list
openai_function_object = ChatCompletionToolParamFunctionChunk(**tool) # type: ignore
# Handle tools with 'type' field (OpenAI spec compliance) Ignore this field -> https://github.com/BerriAI/litellm/issues/14644#issuecomment-3342061838
if "type" in tool:
del tool["type"] # type: ignore
tool_name = list(tool.keys())[0] if len(tool.keys()) == 1 else None
if tool_name and (
tool_name == "codeExecution" or tool_name == "code_execution"
@ -437,7 +441,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
elif model and "gemini-2.5-pro" in model.lower():
budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO
elif model and "gemini-2.5-flash" in model.lower():
budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH
budget = (
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH
)
else:
budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET
@ -621,16 +627,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
elif param == "seed":
optional_params["seed"] = value
elif param == "reasoning_effort" and isinstance(value, str):
optional_params["thinkingConfig"] = (
VertexGeminiConfig._map_reasoning_effort_to_thinking_budget(
value, model
)
optional_params[
"thinkingConfig"
] = VertexGeminiConfig._map_reasoning_effort_to_thinking_budget(
value, model
)
elif param == "thinking":
optional_params["thinkingConfig"] = (
VertexGeminiConfig._map_thinking_param(
cast(AnthropicThinkingParam, value)
)
optional_params[
"thinkingConfig"
] = VertexGeminiConfig._map_thinking_param(
cast(AnthropicThinkingParam, value)
)
elif param == "modalities" and isinstance(value, list):
response_modalities = self.map_response_modalities(value)
@ -1066,7 +1072,6 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
GenerateContentResponseBody, BidiGenerateContentServerMessage
],
) -> Usage:
if (
completion_response is not None
and "usageMetadata" not in completion_response
@ -1502,28 +1507,28 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
## ADD METADATA TO RESPONSE ##
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata)
model_response._hidden_params["vertex_ai_grounding_metadata"] = (
grounding_metadata
)
model_response._hidden_params[
"vertex_ai_grounding_metadata"
] = grounding_metadata
setattr(
model_response, "vertex_ai_url_context_metadata", url_context_metadata
)
model_response._hidden_params["vertex_ai_url_context_metadata"] = (
url_context_metadata
)
model_response._hidden_params[
"vertex_ai_url_context_metadata"
] = url_context_metadata
setattr(model_response, "vertex_ai_safety_results", safety_ratings)
model_response._hidden_params["vertex_ai_safety_results"] = (
safety_ratings # older approach - maintaining to prevent regressions
)
model_response._hidden_params[
"vertex_ai_safety_results"
] = safety_ratings # older approach - maintaining to prevent regressions
## ADD CITATION METADATA ##
setattr(model_response, "vertex_ai_citation_metadata", citation_metadata)
model_response._hidden_params["vertex_ai_citation_metadata"] = (
citation_metadata # older approach - maintaining to prevent regressions
)
model_response._hidden_params[
"vertex_ai_citation_metadata"
] = citation_metadata # older approach - maintaining to prevent regressions
except Exception as e:
raise VertexAIError(
@ -1596,7 +1601,7 @@ async def make_call(
)
try:
response = await client.post(api_base, headers=headers, data=data, stream=True)
response = await client.post(api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj)
response.raise_for_status()
except httpx.HTTPStatusError as e:
exception_string = str(await e.response.aread())
@ -1643,7 +1648,7 @@ def make_sync_call(
if client is None:
client = HTTPHandler() # Create a new client if none provided
response = client.post(api_base, headers=headers, data=data, stream=True)
response = client.post(api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj)
if response.status_code != 200 and response.status_code != 201:
raise VertexAIError(
@ -1842,7 +1847,7 @@ class VertexLLM(VertexBase):
try:
response = await client.post(
api_base, headers=headers, json=cast(dict, request_body)
api_base, headers=headers, json=cast(dict, request_body), logging_obj=logging_obj
) # type: ignore
response.raise_for_status()
except httpx.HTTPStatusError as err:
@ -2045,7 +2050,7 @@ class VertexLLM(VertexBase):
client = client
try:
response = client.post(url=url, headers=headers, json=data) # type: ignore
response = client.post(url=url, headers=headers, json=data, logging_obj=logging_obj) # type: ignore
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code = err.response.status_code

View file

@ -24,6 +24,7 @@ from functools import partial
from typing import (
TYPE_CHECKING,
Any,
AsyncIterator,
Callable,
Coroutine,
Dict,
@ -149,6 +150,7 @@ from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
from .llms.bedrock.embed.embedding import BedrockEmbedding
from .llms.bedrock.image.image_handler import BedrockImageGeneration
from .llms.bytez.chat.transformation import BytezChatConfig
from .llms.lemonade.chat.transformation import LemonadeChatConfig
from .llms.codestral.completion.handler import CodestralTextCompletion
from .llms.cohere.embed import handler as cohere_embed
from .llms.custom_httpx.aiohttp_handler import BaseLLMAIOHTTPHandler
@ -267,6 +269,7 @@ bytez_transformation = BytezChatConfig()
heroku_transformation = HerokuChatConfig()
oci_transformation = OCIChatConfig()
ovhcloud_transformation = OVHCloudChatConfig()
lemonade_transformation = LemonadeChatConfig()
####### COMPLETION ENDPOINTS ################
@ -3545,6 +3548,35 @@ def completion( # type: ignore # noqa: PLR0915
)
pass
elif custom_llm_provider == "lemonade":
api_key = (
api_key
or litellm.lemonade_key
or get_secret_str("LEMONADE_API_KEY")
or litellm.api_key
)
response = base_llm_http_handler.completion(
model=model,
messages=messages,
headers=headers,
model_response=model_response,
api_key=api_key,
api_base=api_base,
acompletion=acompletion,
logging_obj=logging,
optional_params=optional_params,
litellm_params=litellm_params,
timeout=timeout, # type: ignore
client=client,
custom_llm_provider=custom_llm_provider,
encoding=encoding,
stream=stream,
provider_config=lemonade_transformation,
)
pass
elif custom_llm_provider == "ovhcloud" or model in litellm.ovhcloud_models:
api_key = (
@ -5139,6 +5171,21 @@ async def aadapter_completion(
except Exception as e:
raise e
async def aadapter_generate_content(
**kwargs,
) -> Union[Dict[str, Any], AsyncIterator[bytes]]:
from litellm.google_genai.adapters.handler import (
GenerateContentToCompletionHandler,
)
coro = cast(
Coroutine[Any, Any, Union[Dict[str, Any], AsyncIterator[bytes]]],
GenerateContentToCompletionHandler.generate_content_handler(
**kwargs, _is_async=True
),
)
return await coro
def adapter_completion(
*, adapter_id: str, **kwargs

View file

@ -2004,9 +2004,9 @@
"cache_read_input_token_cost": 1.25e-07,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supported_endpoints": [
@ -3308,6 +3308,64 @@
"supports_tool_choice": true,
"supports_web_search": true
},
"azure_ai/grok-4": {
"input_cost_per_token": 5.5e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
"source": "https://azure.microsoft.com/en-us/blog/grok-4-is-now-available-in-azure-ai-foundry-unlock-frontier-intelligence-and-business-ready-capabilities/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"azure_ai/grok-4-fast-non-reasoning": {
"input_cost_per_token": 5e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.5e-03,
"source": "https://azure.microsoft.com/en-us/blog/grok-4-is-now-available-in-azure-ai-foundry-unlock-frontier-intelligence-and-business-ready-capabilities/",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"azure_ai/grok-4-fast-reasoning": {
"input_cost_per_token": 5.8e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.9e-03,
"source": "https://azure.microsoft.com/en-us/blog/grok-4-is-now-available-in-azure-ai-foundry-unlock-frontier-intelligence-and-business-ready-capabilities/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"azure_ai/grok-code-fast-1": {
"input_cost_per_token": 3.5e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"source": "https://azure.microsoft.com/en-us/blog/grok-4-is-now-available-in-azure-ai-foundry-unlock-frontier-intelligence-and-business-ready-capabilities/",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_web_search": true
},
"azure_ai/jais-30b-chat": {
"input_cost_per_token": 0.0032,
"litellm_provider": "azure_ai",
@ -4739,6 +4797,58 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
"claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "anthropic",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"claude-sonnet-4-5-20250929": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "anthropic",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"claude-opus-4-1": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
@ -9396,96 +9506,6 @@
"supports_vision": true,
"supports_web_search": true
},
"gemini-flash-latest": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
"litellm_provider": "vertex_ai-language-models",
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_images_per_prompt": 3000,
"max_input_tokens": 1048576,
"max_output_tokens": 65535,
"max_pdf_size_mb": 30,
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_reasoning_token": 2.5e-06,
"output_cost_per_token": 2.5e-06,
"source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/completions",
"/v1/batch"
],
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_audio_output": false,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": true
},
"gemini-flash-lite-latest": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_audio_token": 3e-07,
"input_cost_per_token": 1e-07,
"litellm_provider": "vertex_ai-language-models",
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_images_per_prompt": 3000,
"max_input_tokens": 1048576,
"max_output_tokens": 65535,
"max_pdf_size_mb": 30,
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_reasoning_token": 4e-07,
"output_cost_per_token": 4e-07,
"source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/completions",
"/v1/batch"
],
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_audio_output": false,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": true
},
"gemini-2.5-flash-lite-preview-06-17": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_audio_token": 5e-07,
@ -12765,6 +12785,34 @@
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5-codex": {
"cache_read_input_token_cost": 1.25e-07,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-5-2025-08-07": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_flex": 6.25e-08,
@ -12840,9 +12888,9 @@
"cache_read_input_token_cost": 1.25e-07,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supported_endpoints": [
@ -13294,6 +13342,18 @@
],
"supports_tool_choice": false
},
"lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": {
"input_cost_per_token": 0,
"litellm_provider": "lemonade",
"max_tokens": 32768,
"max_input_tokens": 32768,
"max_output_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"groq/deepseek-r1-distill-llama-70b": {
"input_cost_per_token": 7.5e-07,
"litellm_provider": "groq",
@ -13583,6 +13643,19 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"groq/moonshotai/kimi-k2-instruct-0905": {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 3e-06,
"cache_read_input_token_cost": 0.5e-06,
"litellm_provider": "groq",
"max_input_tokens": 262144,
"max_output_tokens": 16384,
"max_tokens": 278528,
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"groq/openai/gpt-oss-120b": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "groq",
@ -19643,6 +19716,32 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
"us.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"us.anthropic.claude-opus-4-20250514-v1:0": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
@ -20983,6 +21082,50 @@
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_batches": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-sonnet-4-5@20250929": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_batches": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-opus-4@20250514": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,

View file

@ -294,6 +294,9 @@ class MCPRequestHandler:
) -> List[str]:
"""
Get list of allowed MCP servers for the given user/key based on permissions
Returns:
List[str]: List of allowed MCP servers by server id
"""
from typing import List
@ -330,11 +333,30 @@ class MCPRequestHandler:
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
return []
@staticmethod
def is_tool_allowed(
allowed_mcp_servers: List[str],
server_name: str,
) -> bool:
"""
Check if the tool is allowed for the given user/key based on permissions
"""
if len(allowed_mcp_servers) == 0:
return True
elif server_name in allowed_mcp_servers:
return True
return False
@staticmethod
async def _get_allowed_mcp_servers_for_key(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.auth.auth_checks import get_object_permission
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if user_api_key_auth is None:
return []
@ -347,12 +369,12 @@ class MCPRequestHandler:
return []
try:
key_object_permission = (
await prisma_client.db.litellm_objectpermissiontable.find_unique(
where={
"object_permission_id": user_api_key_auth.object_permission_id
},
)
key_object_permission = await get_object_permission(
object_permission_id=user_api_key_auth.object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if key_object_permission is None:
return []
@ -386,7 +408,12 @@ class MCPRequestHandler:
first we check if the team has a object_permission_id attached
- if it does then we look up the object_permission for the team
"""
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if user_api_key_auth is None:
return []
@ -399,10 +426,12 @@ class MCPRequestHandler:
return []
try:
team_obj: Optional[LiteLLM_TeamTable] = (
await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": user_api_key_auth.team_id},
)
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
team_id=user_api_key_auth.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if team_obj is None:
verbose_logger.debug("team_obj is None")
@ -534,7 +563,12 @@ class MCPRequestHandler:
async def _get_mcp_access_groups_for_key(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.auth.auth_checks import get_object_permission
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if user_api_key_auth is None:
return []
@ -546,15 +580,21 @@ class MCPRequestHandler:
verbose_logger.debug("prisma_client is None")
return []
key_object_permission = (
await prisma_client.db.litellm_objectpermissiontable.find_unique(
where={"object_permission_id": user_api_key_auth.object_permission_id},
try:
key_object_permission = await get_object_permission(
object_permission_id=user_api_key_auth.object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
)
if key_object_permission is None:
return []
if key_object_permission is None:
return []
return key_object_permission.mcp_access_groups or []
return key_object_permission.mcp_access_groups or []
except Exception as e:
verbose_logger.warning(f"Failed to get MCP access groups for key: {str(e)}")
return []
@staticmethod
async def _get_mcp_access_groups_for_team(
@ -563,7 +603,12 @@ class MCPRequestHandler:
"""
Get MCP access groups for the team
"""
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if user_api_key_auth is None:
return []
@ -575,20 +620,28 @@ class MCPRequestHandler:
verbose_logger.debug("prisma_client is None")
return []
team_obj: Optional[LiteLLM_TeamTable] = (
await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": user_api_key_auth.team_id},
try:
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
team_id=user_api_key_auth.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
)
if team_obj is None:
verbose_logger.debug("team_obj is None")
return []
if team_obj is None:
verbose_logger.debug("team_obj is None")
return []
object_permissions = team_obj.object_permission
if object_permissions is None:
return []
object_permissions = team_obj.object_permission
if object_permissions is None:
return []
return object_permissions.mcp_access_groups or []
return object_permissions.mcp_access_groups or []
except Exception as e:
verbose_logger.warning(
f"Failed to get MCP access groups for team: {str(e)}"
)
return []
@staticmethod
def get_mcp_access_groups_from_headers(headers: Headers) -> Optional[List[str]]:

View file

@ -919,6 +919,14 @@ class MCPServerManager:
return server
return None
def get_mcp_server_names_from_ids(self, server_ids: List[str]) -> List[str]:
server_names = []
registry = self.get_registry()
for server in registry.values():
if server.server_id in server_ids:
server_names.append(server.name)
return server_names
def get_mcp_server_by_name(self, server_name: str) -> Optional[MCPServer]:
"""
Get the MCP Server from the server name

View file

@ -425,7 +425,7 @@ if MCP_AVAILABLE:
continue
# Get server-specific auth header if available
server_auth_header = None
server_auth_header: Optional[Union[Dict[str, str], str]] = None
if mcp_server_auth_headers and server.alias is not None:
server_auth_header = mcp_server_auth_headers.get(server.alias)
elif mcp_server_auth_headers and server.server_name is not None:
@ -560,6 +560,25 @@ if MCP_AVAILABLE:
name
)
## CHECK IF USER IS ALLOWED TO CALL THIS TOOL
allowed_mcp_server_ids = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
)
allowed_mcp_servers = global_mcp_server_manager.get_mcp_server_names_from_ids(
allowed_mcp_server_ids
)
if not MCPRequestHandler.is_tool_allowed(
allowed_mcp_servers=allowed_mcp_servers,
server_name=server_name_from_prefix,
):
raise HTTPException(
status_code=403,
detail=f"User not allowed to call this tool. Allowed MCP servers: {allowed_mcp_servers}",
)
standard_logging_mcp_tool_call: StandardLoggingMCPToolCall = (
_get_standard_logging_mcp_tool_call(
name=original_tool_name, # Use original name for logging
@ -571,16 +590,16 @@ if MCP_AVAILABLE:
"litellm_logging_obj", None
)
if litellm_logging_obj:
litellm_logging_obj.model_call_details[
"mcp_tool_call_metadata"
] = standard_logging_mcp_tool_call
litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = (
standard_logging_mcp_tool_call
)
litellm_logging_obj.model = f"MCP: {name}"
# Try managed server tool first (pass the full prefixed name)
# Primary and recommended way to use MCP servers
#########################################################
mcp_server: Optional[
MCPServer
] = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
mcp_server: Optional[MCPServer] = (
global_mcp_server_manager._get_mcp_server_from_tool_name(name)
)
if mcp_server:
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (
mcp_server.mcp_info or {}

View file

@ -330,6 +330,7 @@ class LiteLLMRoutes(enum.Enum):
anthropic_routes = [
"/v1/messages",
"/v1/messages/count_tokens",
]
mcp_routes = [

View file

@ -120,7 +120,7 @@ async def anthropic_response( # noqa: PLR0915
): # model in router deployments, calling a specific deployment on the router
llm_coro = llm_router.aanthropic_messages(**data, specific_deployment=True)
elif (
llm_router is not None and data["model"] in llm_router.get_model_ids()
llm_router is not None and llm_router.has_model_id(data["model"])
): # model in router model list
llm_coro = llm_router.aanthropic_messages(**data)
elif (

View file

@ -41,12 +41,12 @@ from litellm.proxy._types import (
LiteLLM_UserTable,
LiteLLMRoutes,
LitellmUserRoles,
NewTeamRequest,
ProxyErrorTypes,
ProxyException,
RoleBasedPermissions,
SpecialModelNames,
UserAPIKeyAuth,
NewTeamRequest,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.route_llm_request import route_request
@ -474,7 +474,7 @@ async def get_end_user_object(
return return_obj
# else, check db
try:
try:
response = await prisma_client.db.litellm_endusertable.find_unique(
where={"user_id": end_user_id},
include={"litellm_budget_table": True},
@ -817,7 +817,9 @@ async def _cache_management_object(
):
await user_api_key_cache.async_set_cache(
key=key, value=value, ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
key=key,
value=value,
ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
)
@ -892,7 +894,9 @@ async def _get_team_db_check(
system_admin_user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
created_team_dict = await new_team(
data=new_team_data, http_request=mock_request, user_api_key_dict=system_admin_user
data=new_team_data,
http_request=mock_request,
user_api_key_dict=system_admin_user,
)
response = LiteLLM_TeamTable(**created_team_dict)
return response
@ -1166,6 +1170,54 @@ async def get_key_object(
return _response
@log_db_metrics
async def get_object_permission(
object_permission_id: str,
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
parent_otel_span: Optional[Span] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> Optional[LiteLLM_ObjectPermissionTable]:
"""
- Check if object permission id in proxy ObjectPermissionTable
- if valid, return LiteLLM_ObjectPermissionTable object
- if not, then raise an error
"""
if prisma_client is None:
raise Exception(
"No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys"
)
# check if in cache
key = "object_permission_id:{}".format(object_permission_id)
cached_obj_permission = await user_api_key_cache.async_get_cache(key=key)
if cached_obj_permission is not None:
if isinstance(cached_obj_permission, dict):
return LiteLLM_ObjectPermissionTable(**cached_obj_permission)
elif isinstance(cached_obj_permission, LiteLLM_ObjectPermissionTable):
return cached_obj_permission
# else, check db
try:
response = await prisma_client.db.litellm_objectpermissiontable.find_unique(
where={"object_permission_id": object_permission_id}
)
if response is None:
return None
# save the object permission to cache
await user_api_key_cache.async_set_cache(
key=key,
value=response.model_dump(),
ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
)
return LiteLLM_ObjectPermissionTable(**response.dict())
except Exception:
return None
@log_db_metrics
async def get_org_object(
org_id: str,

View file

@ -417,6 +417,12 @@ def bytes_to_mb(bytes_value: int):
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
"""
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"]
@ -426,7 +432,9 @@ def get_key_model_rpm_limit(
if "rpm_limit" in budget and budget["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"]
return None
@ -439,7 +447,9 @@ def get_key_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"]
return None
@ -473,6 +483,7 @@ def _has_user_setup_sso():
return sso_setup
def get_customer_user_header_from_mapping(user_id_mapping) -> Optional[str]:
"""Return the header_name mapped to CUSTOMER role, if any (dict-based)."""
if not user_id_mapping:
@ -522,7 +533,11 @@ def get_end_user_id_from_request_body(
for header_name, header_value in request_headers.items():
if header_name.lower() == custom_header_name_to_check.lower():
user_id_from_header = header_value
user_id_str = str(user_id_from_header) if user_id_from_header is not None else ""
user_id_str = (
str(user_id_from_header)
if user_id_from_header is not None
else ""
)
if user_id_str.strip():
return user_id_str

View file

@ -62,12 +62,22 @@ class RouteChecks:
for allowed_route in valid_token.allowed_routes
):
for allowed_route in valid_token.allowed_routes:
if allowed_route in LiteLLMRoutes._member_names_:
if allowed_route in LiteLLMRoutes._member_names_:
if RouteChecks.check_route_access(
route=route,
allowed_routes=LiteLLMRoutes._member_map_[allowed_route].value,
):
return True
################################################
# For llm_api_routes, also check registered pass-through endpoints
################################################
if allowed_route == "llm_api_routes":
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
)
if InitPassThroughEndpointHelpers.is_registered_pass_through_route(route=route):
return True
# check if wildcard pattern is allowed
for allowed_route in valid_token.allowed_routes:

View file

@ -379,6 +379,7 @@ class ProxyBaseLLMRequestProcessing:
user_api_base: Optional[str] = None,
version: Optional[str] = None,
is_streaming_request: Optional[bool] = False,
contents: Optional[list] = None, # Add contents parameter
) -> Any:
"""
Common request processing logic for both chat completions and responses API endpoints
@ -417,6 +418,10 @@ class ProxyBaseLLMRequestProcessing:
)
)
# Pass contents if provided
if contents:
self.data["contents"] = contents
### ROUTE THE REQUEST ###
# Do not change this - it should be a constant time fetch - ALWAYS
llm_call = await route_request(

View file

@ -289,8 +289,8 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
def get_model_group_from_litellm_kwargs(kwargs: dict) -> Optional[str]:
_litellm_params = kwargs.get("litellm_params", None) or {}
_metadata = _litellm_params.get(get_metadata_variable_name_from_kwargs(kwargs)) or {}
_model_group = _metadata.get("model_group", None)
_metadata = _litellm_params.get(get_metadata_variable_name_from_litellm_params(_litellm_params)) or {}
_model_group = _metadata.get("model_group", None) or kwargs.get("model", None)
if _model_group is not None:
return _model_group
@ -367,8 +367,8 @@ def add_guardrail_to_applied_guardrails_header(
_metadata["applied_guardrails"] = [guardrail_name]
def get_metadata_variable_name_from_kwargs(
kwargs: dict
def get_metadata_variable_name_from_litellm_params(
litellm_params: dict
) -> Literal["metadata", "litellm_metadata"]:
"""
Helper to return what the "metadata" field should be called in the request data
@ -381,4 +381,4 @@ def get_metadata_variable_name_from_kwargs(
- OpenAI then started using this field for their metadata
- LiteLLM is now moving to using `litellm_metadata` for our metadata
"""
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
return "litellm_metadata" if "litellm_metadata" in litellm_params else "metadata"

View file

@ -6,8 +6,11 @@ from typing import Optional
from fastapi import Request
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
SENSITIVE_DATA_MASKER = SensitiveDataMasker()
def remove_sensitive_info_from_deployment(deployment_dict: dict) -> dict:
"""
@ -25,6 +28,8 @@ def remove_sensitive_info_from_deployment(deployment_dict: dict) -> dict:
deployment_dict["litellm_params"].pop("aws_access_key_id", None)
deployment_dict["litellm_params"].pop("aws_secret_access_key", None)
deployment_dict["litellm_params"] = SENSITIVE_DATA_MASKER.mask_dict(deployment_dict["litellm_params"])
return deployment_dict

View file

@ -26,4 +26,9 @@ model_list:
api_key: os.environ/ANTHROPIC_API_KEY
general_settings:
master_key: sk-1234
custom_auth: custom_auth_basic.user_api_key_auth
custom_auth: custom_auth_basic.user_api_key_auth
pass_through_endpoints:
- path: "/azure-config-passthrough"
target: os.environ/AZURE_API_BASE
headers:
Authorization: os.environ/AZURE_API_KEY

View file

@ -1,8 +1,10 @@
from fastapi import APIRouter, Depends, Request, Response
from fastapi import APIRouter, Depends, Request, Response, HTTPException
from fastapi.responses import StreamingResponse
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.types.llms.vertex_ai import TokenCountDetailsResponse
router = APIRouter(
@ -10,140 +12,63 @@ router = APIRouter(
)
@router.post("/v1beta/models/{model_name}:generateContent", dependencies=[Depends(user_api_key_auth)])
@router.post("/models/{model_name}:generateContent", dependencies=[Depends(user_api_key_auth)])
@router.post(
"/v1beta/models/{model_name}:generateContent",
dependencies=[Depends(user_api_key_auth)],
)
@router.post(
"/models/{model_name}:generateContent", dependencies=[Depends(user_api_key_auth)]
)
async def google_generate_content(
request: Request,
model_name: str,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Not Implemented, this is a placeholder for the google genai generateContent endpoint.
"""
from litellm.proxy.proxy_server import (
_read_request_body,
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
from litellm.proxy.proxy_server import llm_router
data = await _read_request_body(request=request)
if "model" not in data:
data["model"] = model_name
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="agenerate_content",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=None,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content(**data)
return response
class GoogleAIStudioDataGenerator:
"""
Ensures SSE data generator is used for Google AI Studio streaming responses
Thin wrapper around ProxyBaseLLMRequestProcessing.async_sse_data_generator
"""
@staticmethod
def _select_data_generator(response, user_api_key_dict, request_data):
from litellm.proxy.proxy_server import proxy_logging_obj
return ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
proxy_logging_obj=proxy_logging_obj,
)
@router.post("/v1beta/models/{model_name}:streamGenerateContent", dependencies=[Depends(user_api_key_auth)])
@router.post("/models/{model_name}:streamGenerateContent", dependencies=[Depends(user_api_key_auth)])
@router.post(
"/v1beta/models/{model_name}:streamGenerateContent",
dependencies=[Depends(user_api_key_auth)],
)
@router.post(
"/models/{model_name}:streamGenerateContent",
dependencies=[Depends(user_api_key_auth)],
)
async def google_stream_generate_content(
request: Request,
model_name: str,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Not Implemented, this is a placeholder for the google genai streamGenerateContent endpoint.
"""
from litellm.proxy.proxy_server import (
_read_request_body,
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
from litellm.proxy.proxy_server import llm_router
data = await _read_request_body(request=request)
if "model" not in data:
data["model"] = model_name
data["stream"] = True # enforce streaming for this endpoint
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="agenerate_content_stream",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=GoogleAIStudioDataGenerator._select_data_generator,
model=None,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
is_streaming_request=True,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content(**data)
# Check if response is an async iterator (streaming response)
if hasattr(response, "__aiter__"):
return StreamingResponse(response, media_type="text/event-stream")
return response
@router.post(
@ -171,13 +96,13 @@ async def google_count_tokens(request: Request, model_name: str):
}
```
"""
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.proxy_server import token_counter as internal_token_counter
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
data = await _read_request_body(request=request)
contents = data.get("contents", [])
#Create TokenCountRequest for the internal endpoint
# Create TokenCountRequest for the internal endpoint
from litellm.proxy._types import TokenCountRequest
# Translate contents to openai format messages using the adapter

View file

@ -41,6 +41,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
)
from litellm.types.utils import (
Choices,
GuardrailStatus,
ModelResponse,
ModelResponseStream,
StreamingChoices,
@ -361,11 +362,30 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
prepared_request.headers,
)
httpx_response = await self.async_handler.post(
url=prepared_request.url,
data=prepared_request.body, # type: ignore
headers=prepared_request.headers, # type: ignore
)
try:
httpx_response = await self.async_handler.post(
url=prepared_request.url,
data=prepared_request.body, # type: ignore
headers=prepared_request.headers, # type: ignore
)
except Exception as e:
# Endpoint down, timeout, or other HTTP/network errors
verbose_proxy_logger.error(
"Bedrock AI: failed to make guardrail request: %s", str(e)
)
# Add guardrail information with failure status
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response={"error": str(e)},
request_data=request_data or {},
guardrail_status="guardrail_failed_to_respond",
start_time=start_time.timestamp(),
end_time=datetime.now().timestamp(),
duration=(datetime.now() - start_time).total_seconds(),
)
# Re-raise the exception to maintain existing behavior
raise
#########################################################
# Add guardrail information to request trace
#########################################################
@ -437,15 +457,30 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
def _get_bedrock_guardrail_response_status(
self, response: httpx.Response
) -> Literal["success", "failure"]:
) -> GuardrailStatus:
"""
Get the status of the bedrock guardrail response.
Returns:
"success": Content allowed through with no violations
"guardrail_intervened": Content blocked due to policy violations
"guardrail_failed_to_respond": Technical error or API failure
"""
if response.status_code == 200:
if self._check_bedrock_response_for_exception(response):
return "failure"
return "guardrail_failed_to_respond"
# Check if the guardrail would block content
try:
_json_response = response.json()
bedrock_guardrail_response = BedrockGuardrailResponse(**_json_response)
if self._should_raise_guardrail_blocked_exception(bedrock_guardrail_response):
return "guardrail_intervened"
except Exception:
pass
return "success"
return "failure"
return "guardrail_failed_to_respond"
def _get_http_exception_for_blocked_guardrail(
self, response: BedrockGuardrailResponse
@ -692,6 +727,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
return
outputs: List[BedrockGuardrailOutput] = (
response.get("outputs", []) or []
)
if not any(output.get("text") for output in outputs):
verbose_proxy_logger.warning(
"Bedrock AI: not running guardrail. No output text in response"
)
return
#########################################################
########## 1. Make parallel Bedrock API requests ##########
#########################################################

View file

@ -1,5 +1,7 @@
from datetime import datetime
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Union, Type
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Type, Union
from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
@ -12,11 +14,11 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.javelin import (
JavelinGuardInput,
JavelinGuardRequest,
JavelinGuardResponse,
JavelinGuardInput,
)
from fastapi import HTTPException
from litellm.types.utils import GuardrailStatus
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -95,7 +97,7 @@ class JavelinGuardrail(CustomGuardrail):
if self.application:
headers["x-javelin-application"] = self.application
status: Literal["success", "failure", "blocked"] = "failure"
status: GuardrailStatus = "guardrail_failed_to_respond"
javelin_response: Optional[JavelinGuardResponse] = None
exception_str = ""
@ -122,7 +124,7 @@ class JavelinGuardrail(CustomGuardrail):
status = "success"
return javelin_response
except Exception as e:
status = "failure"
status = "guardrail_failed_to_respond"
exception_str = str(e)
return {"assessments": []}
finally:
@ -178,12 +180,12 @@ class JavelinGuardrail(CustomGuardrail):
"""
Pre-call hook for the Javelin guardrail.
"""
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_last_user_message,
)
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
verbose_proxy_logger.debug("Javelin Guardrail: pre_call_hook")
verbose_proxy_logger.debug("Javelin Guardrail: Request data: %s", data)

View file

@ -20,6 +20,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import (
LakeraAIRequest,
LakeraAIResponse,
)
from litellm.types.utils import GuardrailStatus
class LakeraAIGuardrail(CustomGuardrail):
@ -70,7 +71,7 @@ class LakeraAIGuardrail(CustomGuardrail):
"""
Call the Lakera AI v2 guard API.
"""
status: Literal["success", "failure"] = "success"
status: GuardrailStatus = "success"
exception_str: str = ""
start_time: datetime = datetime.now()
lakera_response: Optional[LakeraAIResponse] = None
@ -99,7 +100,7 @@ class LakeraAIGuardrail(CustomGuardrail):
lakera_response = LakeraAIResponse(**response.json())
return lakera_response, masked_entity_count
except Exception as e:
status = "failure"
status = "guardrail_failed_to_respond"
exception_str = str(e)
raise e
finally:

View file

@ -30,6 +30,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
Choices,
GuardrailStatus,
ModelResponse,
ModelResponseStream,
)
@ -329,14 +330,14 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
guardrail_response = metadata.get("_model_armor_response", {})
# Determine status – default to "success" but prefer the explicit value if present.
guardrail_status: Literal["success", "failure", "blocked"] = metadata.get(
guardrail_status: GuardrailStatus = metadata.get(
"_model_armor_status", "success"
) # type: ignore
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=guardrail_response,
request_data=request_data,
guardrail_status=guardrail_status, # type: ignore
guardrail_status=guardrail_status,
duration=duration,
start_time=start_time,
end_time=end_time,

View file

@ -8,7 +8,8 @@
import asyncio
import copy
import os
from typing import Any, Dict, Final, Literal, Optional, Union, Type, TYPE_CHECKING
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, Final, Literal, Optional, Type, Union
from urllib.parse import urljoin
from fastapi import HTTPException
@ -23,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import EmbeddingResponse, ImageResponse
from litellm.types.utils import EmbeddingResponse, GuardrailStatus, ImageResponse
# Constants
USER_ROLE: Final[Literal["user"]] = "user"
@ -204,6 +205,7 @@ class NomaGuardrail(CustomGuardrail):
user_auth: UserAPIKeyAuth,
) -> Optional[str]:
"""Shared logic for processing user message checks"""
start_time = datetime.now()
extra_data = self.get_guardrail_dynamic_request_body_params(request_data)
user_message = await self._extract_user_message(request_data)
@ -218,6 +220,23 @@ class NomaGuardrail(CustomGuardrail):
user_auth=user_auth,
extra_data=extra_data,
)
end_time = datetime.now()
duration = (end_time - start_time).total_seconds()
# Determine guardrail status based on response
guardrail_status = self._determine_guardrail_status(response_json)
# Always log guardrail information for consistency
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider="noma",
guardrail_json_response=response_json,
request_data=request_data,
guardrail_status=guardrail_status,
start_time=start_time.timestamp(),
end_time=end_time.timestamp(),
duration=duration,
)
if self.monitor_mode:
await self._handle_verdict_background(
@ -248,6 +267,8 @@ class NomaGuardrail(CustomGuardrail):
user_auth: UserAPIKeyAuth,
) -> Optional[str]:
"""Shared logic for processing LLM response checks"""
start_time = datetime.now()
extra_data = self.get_guardrail_dynamic_request_body_params(request_data)
if not isinstance(response, litellm.ModelResponse):
@ -271,6 +292,23 @@ class NomaGuardrail(CustomGuardrail):
user_auth=user_auth,
extra_data=extra_data,
)
end_time = datetime.now()
duration = (end_time - start_time).total_seconds()
# Determine guardrail status based on response
guardrail_status = self._determine_guardrail_status(response_json)
# Always log guardrail information for consistency
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider="noma",
guardrail_json_response=response_json,
request_data=request_data,
guardrail_status=guardrail_status,
start_time=start_time.timestamp(),
end_time=end_time.timestamp(),
duration=duration,
)
if self.monitor_mode:
await self._handle_verdict_background(
@ -294,6 +332,41 @@ class NomaGuardrail(CustomGuardrail):
await self._check_verdict(ASSISTANT_ROLE, content, response_json)
return content
def _determine_guardrail_status(self, response_json: dict) -> GuardrailStatus:
"""
Determine the guardrail status based on NOMA API response.
Args:
response_json: Response from NOMA API
Returns:
"success": Content allowed through with no violations
"guardrail_intervened": Content blocked due to policy violations
"guardrail_failed_to_respond": Technical error or API failure
"""
try:
# Check if we got a valid response structure
if not isinstance(response_json, dict):
return "guardrail_failed_to_respond"
# Get the verdict from the response
verdict = response_json.get("verdict", True)
# If verdict is True, content is allowed
if verdict is True:
return "success"
# If verdict is False, content is blocked/flagged
if verdict is False:
return "guardrail_intervened"
# If verdict is missing or invalid, treat as failure
return "guardrail_failed_to_respond"
except Exception as e:
verbose_proxy_logger.error(f"Error determining NOMA guardrail status: {str(e)}")
return "guardrail_failed_to_respond"
def _should_only_sensitive_data_failed(self, classification_obj: dict) -> bool:
"""
Check if only sensitive data detectors (PII, PCI, secrets) have result=true in the classification.
@ -539,8 +612,22 @@ class NomaGuardrail(CustomGuardrail):
try:
return await self._check_user_message(data, user_api_key_dict)
except NomaBlockedMessage:
# Blocked requests were already logged in _process_user_message_check with "blocked" status
raise
except Exception as e:
# Log technical failures
from datetime import datetime
start_time = datetime.now()
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider="noma",
guardrail_json_response=str(e),
request_data=data,
guardrail_status="guardrail_failed_to_respond",
start_time=start_time.timestamp(),
end_time=start_time.timestamp(),
duration=0.0,
)
verbose_proxy_logger.error(f"Noma pre-call hook failed: {str(e)}")
if self.block_failures:
@ -580,8 +667,22 @@ class NomaGuardrail(CustomGuardrail):
try:
return await self._check_user_message(data, user_api_key_dict)
except NomaBlockedMessage:
# Blocked requests were already logged in _process_user_message_check with "blocked" status
raise
except Exception as e:
# Log technical failures
from datetime import datetime
start_time = datetime.now()
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider="noma",
guardrail_json_response=str(e),
request_data=data,
guardrail_status="guardrail_failed_to_respond",
start_time=start_time.timestamp(),
end_time=start_time.timestamp(),
duration=0.0,
)
verbose_proxy_logger.error(f"Noma moderation hook failed: {str(e)}")
if self.block_failures:
@ -615,8 +716,22 @@ class NomaGuardrail(CustomGuardrail):
try:
return await self._check_llm_response(data, response, user_api_key_dict)
except NomaBlockedMessage:
# Blocked requests were already logged in _process_llm_response_check with "blocked" status
raise
except Exception as e:
# Log technical failures
from datetime import datetime
start_time = datetime.now()
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider="noma",
guardrail_json_response=str(e),
request_data=data,
guardrail_status="guardrail_failed_to_respond",
start_time=start_time.timestamp(),
end_time=start_time.timestamp(),
duration=0.0,
)
verbose_proxy_logger.error(f"Noma post-call hook failed: {str(e)}")
if self.block_failures:
raise

View file

@ -10,14 +10,12 @@
import asyncio
import json
from litellm._uuid import uuid
from datetime import datetime
from typing import (
Any,
AsyncGenerator,
Dict,
List,
Literal,
Optional,
Tuple,
Union,
@ -29,6 +27,7 @@ import aiohttp
import litellm # noqa: E401
from litellm import get_secret
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.caching.caching import DualCache
from litellm.exceptions import BlockedPiiEntityError
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -45,6 +44,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.presidio import (
PresidioAnalyzeResponseItem,
)
from litellm.types.utils import CallTypes as LitellmCallTypes
from litellm.types.utils import GuardrailStatus
from litellm.utils import (
EmbeddingResponse,
ImageResponse,
@ -324,7 +324,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
"""
start_time = datetime.now()
analyze_results: Optional[Union[List[PresidioAnalyzeResponseItem], Dict]] = None
status: Literal["success", "failure"] = "success"
status: GuardrailStatus = "success"
masked_entity_count: Dict[str, int] = {}
exception_str: str = ""
try:
@ -356,7 +356,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
)
return redacted_text["text"]
except Exception as e:
status = "failure"
status = "guardrail_failed_to_respond"
exception_str = str(e)
raise e
finally:

View file

@ -0,0 +1,170 @@
# Dynamic Rate Limiter v3 - Saturation-Aware Priority-Based Rate Limiting
## Overview
The v3 dynamic rate limiter implements saturation-aware rate limiting with priority-based allocation. It balances resource efficiency (allowing unused capacity to be borrowed) with fairness guarantees (enforcing priorities during high load).
**Key Behavior:**
- When system is under 80% capacity: Generous mode - allows priority borrowing
- When system is at/above 80% capacity: Strict mode - enforces normalized priority limits
## How It Works
### Flow Diagram
```
┌─────────────────────────────────────────────────────────────┐
│ Incoming Request │
└────────────────────────┬────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────┐
│ 1. Check Model Saturation │
│ - Query v3 limiter's Redis counters │
│ - Calculate: current_usage / capacity │
│ - Returns: 0.0 (empty) to 1.0+ (saturated) │
└────────────────────────┬────────────────────────────────────┘
│
▼
┌────────┴────────┐
│ Saturation? │
└────────┬────────┘
│
┌───────────────┴───────────────┐
│ │
▼ ▼
< 80% (Generous) >= 80% (Strict)
│ │
▼ ▼
┌─────────────────────┐ ┌─────────────────────┐
│ Generous Mode │ │ Strict Mode │
│ │ │ │
│ - Enforce model- │ │ - Normalize │
│ wide capacity │ │ priority weights │
│ - No priority │ │ (if over 1.0) │
│ restrictions │ │ │
│ - Allows borrowing │ │ - Create priority- │
│ │ │ specific │
│ - First-come- │ │ descriptors │
│ first-served │ │ │
│ until capacity │ │ - Enforce strict │
│ │ │ limits per │
│ │ │ priority │
└──────────┬──────────┘ └──────────┬──────────┘
│ │
│ ▼
│ ┌──────────────────────┐
│ │ Track model usage │
│ │ for future │
│ │ saturation checks │
│ └──────────┬───────────┘
│ │
└───────────────┬───────────────┘
│
▼
┌──────────────┐
│ v3 Limiter │
│ Check │
└──────┬───────┘
│
┌───────────────┴───────────────┐
│ │
▼ ▼
OVER_LIMIT OK
│ │
▼ ▼
Return 429 Error Allow Request
```
## Configuration
### Priority Reservation
Set priority weights in your proxy configuration:
```python
litellm.priority_reservation = {
"premium": 0.75, # 75% of capacity
"standard": 0.25 # 25% of capacity
}
```
### Priority Reservation Settings
Configure saturation-aware behavior:
```python
litellm.priority_reservation_settings = PriorityReservationSettings(
default_priority=0.5, # Default weight for users without explicit priority
saturation_threshold=0.80, # 80% - threshold for strict mode enforcement
tracking_multiplier=10 # 10x - multiplier for non-blocking tracking in strict mode
)
```
**Settings:**
- `default_priority` (default: 0.5) - Priority weight for users without explicit priority metadata
- `saturation_threshold` (default: 0.80) - Saturation level (0.0-1.0) at which strict priority enforcement begins
- `tracking_multiplier` (default: 10) - Multiplier for model-wide tracking limits in strict mode
### User Priority Assignment
Set priority in user metadata:
```python
user_api_key_dict.metadata = {"priority": "premium"}
```
## Priority Weight Normalization
If priorities sum to > 1.0, they are automatically normalized:
```
Input: {key_a: 0.60, key_b: 0.80} = 1.40 total
Output: {key_a: 0.43, key_b: 0.57} = 1.00 total
```
This ensures total allocation never exceeds model capacity.
## Implementation Details
### Saturation Detection
- Queries v3 limiter's Redis counters for model-wide usage
- Checks both RPM and TPM, returns higher saturation value
- Non-blocking reads (doesn't increment counters)
### Mode Selection
**Generous Mode (< 80% saturation):**
- Creates single model-wide descriptor
- Enforces total capacity only
- Allows any priority to use available capacity
- Prevents over-subscription via model-wide limit
**Strict Mode (>= 80% saturation):**
- Creates priority-specific descriptors with normalized weights
- Each priority gets its reserved allocation
- Tracks model-wide usage separately (non-blocking, 10x multiplier)
- Ensures fairness under load
Test scenarios covered:
1. No rate limiting when under capacity
2. Priority queue behavior during saturation
3. Spillover capacity for default keys
4. Over-allocated priorities with normalization
5. Default priority value handling
### `_PROXY_DynamicRateLimitHandlerV3`
Main handler class inheriting from `CustomLogger`.
**Key Methods:**
- `async_pre_call_hook()` - Main entry point, routes to generous/strict mode
- `_check_model_saturation()` - Queries Redis for current usage
- `_handle_generous_mode()` - Enforces model-wide capacity only
- `_handle_strict_mode()` - Enforces normalized priority limits
- `_normalize_priority_weights()` - Handles over-allocation
- `_create_priority_based_descriptors()` - Creates rate limit descriptors

View file

@ -1,9 +1,9 @@
"""
Dynamic rate limiter v3
Dynamic rate limiter v3 - Saturation-aware priority-based rate limiting
"""
import os
from typing import List, Literal, Optional, Union
from typing import Dict, List, Literal, Optional, Union
from fastapi import HTTPException
@ -24,12 +24,18 @@ from litellm.types.router import ModelGroupInfo
class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
"""
Simple validation version that uses v3 parallel request limiter for priority-based rate limiting.
Saturation-aware priority-based rate limiter using v3 infrastructure.
Key differences from original:
1. Uses v3 limiter's sliding window approach instead of per-minute cache buckets
2. Leverages Redis Lua scripts for atomic operations under high traffic
3. Creates priority-specific rate limit descriptors
Key features:
1. Reuses v3 limiter's Redis-based tracking (works across multiple instances)
2. Only enforces priority limits when model is saturated (>80% usage)
3. When under capacity, allows all requests (generous behavior)
4. When saturated, enforces strict priority-based limits (fairness)
How it works:
- Uses v3 limiter's counter keys to check model-wide saturation
- Saturation check reads existing counters without incrementing
- Priority enforcement reuses v3 limiter's atomic Lua scripts
"""
def __init__(self, internal_usage_cache: DualCache):
self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache)
@ -57,6 +63,107 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
weight = litellm.priority_reservation[priority]
return weight
def _normalize_priority_weights(self) -> Dict[str, float]:
"""
Normalize priority weights if they sum to > 1.0
Handles over-allocation: {key_a: 0.60, key_b: 0.80} -> {key_a: 0.43, key_b: 0.57}
"""
if litellm.priority_reservation is None:
return {}
weights = dict(litellm.priority_reservation)
total_weight = sum(weights.values())
if total_weight > 1.0:
normalized = {k: v / total_weight for k, v in weights.items()}
verbose_proxy_logger.debug(
f"Normalized over-allocated priorities: {weights} -> {normalized}"
)
return normalized
return weights
async def _check_model_saturation(
self,
model: str,
model_group_info: ModelGroupInfo,
) -> float:
"""
Check current saturation by directly querying v3 limiter's cache keys.
Reuses v3 limiter's Redis-based tracking (works across multiple instances).
Reads counters WITHOUT incrementing them.
Returns:
float: Saturation ratio (0.0 = empty, 1.0 = at capacity, >1.0 = over)
"""
try:
max_saturation = 0.0
# Query RPM saturation
if model_group_info.rpm is not None and model_group_info.rpm > 0:
# Use v3 limiter's key format: {key:value}:rate_limit_type
counter_key = self.v3_limiter.create_rate_limit_keys(
key="model_saturation_check",
value=model,
rate_limit_type="requests",
)
# Query cache for current counter value
counter_value = await self.internal_usage_cache.async_get_cache(
key=counter_key,
litellm_parent_otel_span=None,
local_only=False, # Check Redis too
)
if counter_value is not None:
current_requests = int(counter_value)
rpm_saturation = current_requests / model_group_info.rpm
max_saturation = max(max_saturation, rpm_saturation)
verbose_proxy_logger.debug(
f"Model {model} RPM: {current_requests}/{model_group_info.rpm} "
f"({rpm_saturation:.1%})"
)
# Query TPM saturation
if model_group_info.tpm is not None and model_group_info.tpm > 0:
counter_key = self.v3_limiter.create_rate_limit_keys(
key="model_saturation_check",
value=model,
rate_limit_type="tokens",
)
counter_value = await self.internal_usage_cache.async_get_cache(
key=counter_key,
litellm_parent_otel_span=None,
local_only=False,
)
if counter_value is not None:
current_tokens = float(counter_value)
tpm_saturation = current_tokens / model_group_info.tpm
max_saturation = max(max_saturation, tpm_saturation)
verbose_proxy_logger.debug(
f"Model {model} TPM: {current_tokens}/{model_group_info.tpm} "
f"({tpm_saturation:.1%})"
)
verbose_proxy_logger.debug(
f"Model {model} overall saturation: {max_saturation:.1%}"
)
return max_saturation
except Exception as e:
verbose_proxy_logger.error(
f"Error checking saturation for {model}: {str(e)}"
)
# Fail open: assume not saturated on error
return 0.0
def _create_priority_based_descriptors(
self,
model: str,
@ -64,11 +171,10 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
priority: Optional[str],
) -> List[RateLimitDescriptor]:
"""
Create rate limit descriptors based on priority and model group limits.
Create rate limit descriptors with normalized priority weights.
This is the key change: instead of calculating dynamic quotas based on active projects,
we create descriptors with priority-adjusted limits and let the v3 limiter handle
the actual rate limiting with its sliding window approach.
Uses normalized weights to handle over-allocation scenarios.
Only called when system is saturated.
"""
descriptors: List[RateLimitDescriptor] = []
@ -79,8 +185,13 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
if model_group_info is None:
return descriptors
# Get priority weight
priority_weight = self._get_priority_weight(priority)
# Get normalized priority weight (handles over-allocation)
normalized_weights = self._normalize_priority_weights()
priority_weight = normalized_weights.get(priority, None) if priority else None
if priority_weight is None:
# Fallback to non-normalized weight
priority_weight = self._get_priority_weight(priority)
# Create priority-specific rate limits
# Use model:priority as the key to separate different priority levels
@ -88,16 +199,17 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
rate_limit_config: RateLimitDescriptorRateLimitObject = {}
# Apply priority weight to model limits
# Apply normalized priority weight to model limits
if model_group_info.tpm is not None:
# Reserve portion of TPM based on priority
# Reserve portion of TPM based on normalized priority
reserved_tpm = int(model_group_info.tpm * priority_weight)
rate_limit_config["tokens_per_unit"] = reserved_tpm
if model_group_info.rpm is not None:
# Reserve portion of RPM based on priority
# Reserve portion of RPM based on normalized priority
reserved_rpm = int(model_group_info.rpm * priority_weight)
rate_limit_config["requests_per_unit"] = reserved_rpm
if rate_limit_config:
rate_limit_config["window_size"] = self.v3_limiter.window_size
@ -112,6 +224,171 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
return descriptors
def _create_model_tracking_descriptor(
self,
model: str,
model_group_info: ModelGroupInfo,
high_limit_multiplier: int = 1,
) -> RateLimitDescriptor:
"""
Create a descriptor for tracking model-wide usage.
Args:
model: Model name
model_group_info: Model configuration with RPM/TPM limits
high_limit_multiplier: Multiplier for limits (use >1 for tracking-only)
Returns:
Rate limit descriptor for model-wide tracking
"""
return RateLimitDescriptor(
key="model_saturation_check",
value=model,
rate_limit={
"requests_per_unit": (
model_group_info.rpm * high_limit_multiplier
if model_group_info.rpm else None
),
"tokens_per_unit": (
model_group_info.tpm * high_limit_multiplier
if model_group_info.tpm else None
),
"window_size": self.v3_limiter.window_size,
},
)
async def _handle_generous_mode(
self,
model: str,
model_group_info: ModelGroupInfo,
user_api_key_dict: UserAPIKeyAuth,
key_priority: Optional[str],
) -> None:
"""
Handle rate limiting in generous mode (under saturation threshold).
In this mode, we enforce model-wide capacity but NOT priority-specific limits.
This allows lower-priority users to borrow unused capacity from higher-priority users.
Args:
model: Model name
model_group_info: Model configuration
user_api_key_dict: User authentication info
key_priority: User's priority level
Raises:
HTTPException: If model capacity is reached
"""
descriptor = self._create_model_tracking_descriptor(
model=model,
model_group_info=model_group_info,
high_limit_multiplier=1, # Enforce actual limits in generous mode
)
response = await self.v3_limiter.should_rate_limit(
descriptors=[descriptor],
parent_otel_span=user_api_key_dict.parent_otel_span,
)
if response["overall_code"] == "OVER_LIMIT":
for status in response["statuses"]:
if status["code"] == "OVER_LIMIT":
raise HTTPException(
status_code=429,
detail={
"error": f"Model capacity reached for {model}. "
f"Priority: {key_priority}, "
f"Rate limit type: {status['rate_limit_type']}, "
f"Remaining: {status['limit_remaining']}"
},
headers={
"retry-after": str(self.v3_limiter.window_size),
"rate_limit_type": str(status["rate_limit_type"]),
"x-litellm-priority": key_priority or "default",
},
)
async def _handle_strict_mode(
self,
model: str,
model_group_info: ModelGroupInfo,
user_api_key_dict: UserAPIKeyAuth,
key_priority: Optional[str],
saturation: float,
data: dict,
) -> None:
"""
Handle rate limiting in strict mode (above saturation threshold).
In this mode, we enforce priority-specific limits using normalized weights.
Args:
model: Model name
model_group_info: Model configuration
user_api_key_dict: User authentication info
key_priority: User's priority level
saturation: Current saturation level
data: Request data dictionary
Raises:
HTTPException: If priority-specific limit is exceeded
"""
# Create priority-based descriptors
descriptors = self._create_priority_based_descriptors(
model=model,
user_api_key_dict=user_api_key_dict,
priority=key_priority,
)
if not descriptors:
verbose_proxy_logger.debug("No rate limit descriptors created, allowing request")
return
# Track model-wide usage for future saturation checks
# Why tracking_multiplier: v3_limiter.should_rate_limit() both increments AND checks limits.
# We need the increment (for saturation detection) but NOT the limit check (priority limits handle enforcement).
# Setting limit to 10x capacity ensures tracking never blocks while keeping accurate counters.
tracking_multiplier = litellm.priority_reservation_settings.tracking_multiplier
tracking_descriptor = self._create_model_tracking_descriptor(
model=model,
model_group_info=model_group_info,
high_limit_multiplier=tracking_multiplier,
)
await self.v3_limiter.should_rate_limit(
descriptors=[tracking_descriptor],
parent_otel_span=user_api_key_dict.parent_otel_span,
)
# Enforce priority-specific limits
response = await self.v3_limiter.should_rate_limit(
descriptors=descriptors,
parent_otel_span=user_api_key_dict.parent_otel_span,
)
if response["overall_code"] == "OVER_LIMIT":
for status in response["statuses"]:
if status["code"] == "OVER_LIMIT":
raise HTTPException(
status_code=429,
detail={
"error": f"Priority-based rate limit exceeded for {status['descriptor_key']}. "
f"Priority: {key_priority}, "
f"Rate limit type: {status['rate_limit_type']}, "
f"Remaining: {status['limit_remaining']}, "
f"Model saturation: {saturation:.1%}"
},
headers={
"retry-after": str(self.v3_limiter.window_size),
"rate_limit_type": str(status["rate_limit_type"]),
"x-litellm-priority": key_priority or "default",
"x-litellm-saturation": f"{saturation:.2%}",
},
)
else:
# Store response for post-call hook
data["litellm_proxy_rate_limit_response"] = response
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -130,60 +407,73 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
],
) -> Optional[Union[Exception, str, dict]]:
"""
Pre-call hook using v3 limiter for priority-based rate limiting.
Saturation-aware pre-call hook for priority-based rate limiting.
This hook implements a two-mode rate limiting strategy:
- Generous mode (< 80% saturation): Enforces model capacity, allows priority borrowing
- Strict mode (>= 80% saturation): Enforces normalized priority-based limits
Args:
user_api_key_dict: User authentication and metadata
cache: Dual cache instance
data: Request data containing model name
call_type: Type of API call being made
Returns:
None if request is allowed, otherwise raises HTTPException
"""
if "model" not in data:
return None
model = data["model"]
key_priority: Optional[str] = user_api_key_dict.metadata.get("priority", None)
# Create priority-based descriptors
descriptors = self._create_priority_based_descriptors(
model=data["model"],
user_api_key_dict=user_api_key_dict,
priority=key_priority,
# Get model configuration
model_group_info: Optional[ModelGroupInfo] = self.llm_router.get_model_group_info(
model_group=model
)
if not descriptors:
verbose_proxy_logger.debug("No rate limit descriptors created, allowing request")
if model_group_info is None:
verbose_proxy_logger.debug(f"No model group info for {model}, allowing request")
return None
# Check current saturation level
try:
# Use v3 limiter to check rate limits
response = await self.v3_limiter.should_rate_limit(
descriptors=descriptors,
parent_otel_span=user_api_key_dict.parent_otel_span,
saturation = await self._check_model_saturation(model, model_group_info)
saturation_threshold = litellm.priority_reservation_settings.saturation_threshold
verbose_proxy_logger.debug(
f"[Dynamic Rate Limiter] Model={model}, Saturation={saturation:.1%}, "
f"Threshold={saturation_threshold:.1%}, Priority={key_priority}"
)
if response["overall_code"] == "OVER_LIMIT":
# Find which descriptor hit the limit
for status in response["statuses"]:
if status["code"] == "OVER_LIMIT":
raise HTTPException(
status_code=429,
detail={
"error": f"Priority-based rate limit exceeded for {status['descriptor_key']}. "
f"Priority: {key_priority}, "
f"Rate limit type: {status['rate_limit_type']}, "
f"Remaining: {status['limit_remaining']}"
},
headers={
"retry-after": str(self.v3_limiter.window_size),
"rate_limit_type": str(status["rate_limit_type"]),
"x-litellm-priority": key_priority or "default",
},
)
data["litellm_model_saturation"] = saturation
# Route to appropriate mode based on saturation
if saturation < saturation_threshold:
await self._handle_generous_mode(
model=model,
model_group_info=model_group_info,
user_api_key_dict=user_api_key_dict,
key_priority=key_priority,
)
else:
# Store response for post-call hook
data["litellm_proxy_rate_limit_response"] = response
await self._handle_strict_mode(
model=model,
model_group_info=model_group_info,
user_api_key_dict=user_api_key_dict,
key_priority=key_priority,
saturation=saturation,
data=data,
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception(
f"Error in dynamic rate limiter v3 pre-call hook: {str(e)}"
verbose_proxy_logger.error(
f"Error in dynamic rate limiter: {str(e)}, allowing request"
)
# Allow request to proceed on unexpected errors
# Fail open on unexpected errors
return None
return None

View file

@ -4,6 +4,7 @@ This is a rate limiter implementation based on a similar one by Envoy proxy.
This is currently in development and not yet ready for production.
"""
import binascii
import os
from datetime import datetime
from math import floor
@ -19,13 +20,12 @@ from typing import (
cast,
)
from fastapi import HTTPException
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
from fastapi import HTTPException
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -97,6 +97,9 @@ end
return results
"""
# Redis cluster slot count
REDIS_CLUSTER_SLOTS = 16384
REDIS_NODE_HASHTAG_NAME = "all_keys"
class RateLimitDescriptorRateLimitObject(TypedDict, total=False):
requests_per_unit: Optional[int]
@ -149,6 +152,20 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60))
def _is_redis_cluster(self) -> bool:
"""
Check if the dual cache is using Redis cluster.
Returns:
bool: True if using Redis cluster, False otherwise.
"""
from litellm.caching.redis_cluster_cache import RedisClusterCache
return (
self.internal_usage_cache.dual_cache.redis_cache is not None
and isinstance(self.internal_usage_cache.dual_cache.redis_cache, RedisClusterCache)
)
async def in_memory_cache_sliding_window(
self,
keys: List[str],
@ -291,26 +308,55 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
return RateLimitResponse(overall_code=overall_code, statuses=statuses)
def keyslot_for_redis_cluster(self, key: str) -> int:
"""
Compute the Redis Cluster slot for a given key.
Simple implementation of `HASH_SLOT = CRC16(key) mod 16384`
Read more about hash slots here: https://medium.com/@linz07m/how-hash-slots-power-data-distribution-in-redis-cluster-bc5b7e74ca7d
Args:
key (str): The Redis key.
Returns:
int: The slot number (0-16383).
"""
# Handle hash tags: use substring between { and }
start = key.find('{')
if start != -1:
end = key.find('}', start + 1)
if end != -1 and end != start + 1:
key = key[start + 1:end]
# Compute CRC16 and mod 16384
crc = binascii.crc_hqx(key.encode('utf-8'), 0)
return crc % REDIS_CLUSTER_SLOTS
def _group_keys_by_hash_tag(self, keys: List[str]) -> Dict[str, List[str]]:
"""
Group keys by their Redis hash tag to ensure cluster compatibility.
Keys with the same hash tag will be processed together.
For Redis clusters, uses slot calculation to group keys that belong to the same slot.
For regular Redis, no grouping is needed - all keys can be processed together.
"""
groups: Dict[str, List[str]] = {}
for key in keys:
# Extract hash tag from key like "{api_key:sk-123}:requests"
if "{" in key and "}" in key:
start = key.find("{")
end = key.find("}", start)
hash_tag = key[start : end + 1]
else:
# Fallback for keys without hash tags
hash_tag = "no_hash_tag"
if hash_tag not in groups:
groups[hash_tag] = []
groups[hash_tag].append(key)
# Use slot calculation for Redis clusters only
if self._is_redis_cluster():
for key in keys:
slot = self.keyslot_for_redis_cluster(key)
slot_key = f"slot_{slot}"
if slot_key not in groups:
groups[slot_key] = []
groups[slot_key].append(key)
else:
# For regular Redis, no grouping needed - process all keys together
groups[REDIS_NODE_HASHTAG_NAME] = keys
return groups
@ -798,7 +844,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
_get_parent_otel_span_from_kwargs,
)
from litellm.proxy.common_utils.callback_utils import (
get_metadata_variable_name_from_kwargs,
get_metadata_variable_name_from_litellm_params,
get_model_group_from_litellm_kwargs,
)
from litellm.types.caching import RedisPipelineIncrementOperation
@ -816,7 +862,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# Get metadata from kwargs
litellm_metadata = kwargs["litellm_params"].get(
get_metadata_variable_name_from_kwargs(kwargs), {}
get_metadata_variable_name_from_litellm_params(kwargs["litellm_params"]), {}
)
if litellm_metadata is None:
return

View file

@ -1594,7 +1594,6 @@ class SSOAuthenticationHandler:
master_key or "",
algorithm="HS256",
)
verbose_proxy_logger.info(f"user_id: {user_id}; jwt_token: {jwt_token}")
if user_id is not None and isinstance(user_id, str):
litellm_dashboard_ui += "?login=success"
verbose_proxy_logger.info(f"Redirecting to {litellm_dashboard_ui}")

View file

@ -31,6 +31,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_body,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.utils import ProxyLogging, is_known_model
from litellm.router import Router
@ -237,6 +238,7 @@ async def create_file(
file_content = await file.read()
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_query(request=request)
or await get_custom_llm_provider_from_request_body(request=request)
or "openai"
)
@ -425,6 +427,7 @@ async def get_file_content(
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_query(request=request)
or await get_custom_llm_provider_from_request_body(request=request)
or "openai"
)
@ -591,6 +594,7 @@ async def get_file(
try:
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_query(request=request)
or await get_custom_llm_provider_from_request_body(request=request)
or "openai"
)
@ -733,6 +737,7 @@ async def delete_file(
try:
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_query(request=request)
or await get_custom_llm_provider_from_request_body(request=request)
or "openai"
)
@ -917,6 +922,7 @@ async def list_files(
else:
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_query(request=request)
or await get_custom_llm_provider_from_request_body(request=request)
or "openai"
)

View file

@ -6,12 +6,14 @@ Provider-specific Pass-Through Endpoints
Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
"""
import json
import os
from typing import Optional, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
from fastapi.responses import StreamingResponse
from starlette.websockets import WebSocketState
import litellm
from litellm._logging import verbose_proxy_logger
@ -19,7 +21,9 @@ from litellm.constants import BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.auth.user_api_key_auth import (
user_api_key_auth,
)
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
get_form_data,
@ -28,6 +32,8 @@ from litellm.proxy.common_utils.http_parsing_utils import (
from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
create_pass_through_route,
create_websocket_passthrough_route,
websocket_passthrough_request,
)
from litellm.proxy.utils import is_known_model
from litellm.secret_managers.main import get_secret_str
@ -51,9 +57,7 @@ def create_request_copy(request: Request):
}
def is_passthrough_request_using_router_model(
request_body: dict, llm_router: Optional[litellm.Router]
) -> bool:
def is_passthrough_request_using_router_model(request_body: dict, llm_router: Optional[litellm.Router]) -> bool:
"""
Returns True if the model is in the llm_router model names
"""
@ -89,16 +93,12 @@ async def llm_passthrough_factory_proxy_route(
model=None,
)
if provider_config is None:
raise HTTPException(
status_code=404, detail=f"Provider {custom_llm_provider} not found"
)
raise HTTPException(status_code=404, detail=f"Provider {custom_llm_provider} not found")
base_target_url = provider_config.get_api_base()
if base_target_url is None:
raise HTTPException(
status_code=404, detail=f"Provider {custom_llm_provider} api base not found"
)
raise HTTPException(status_code=404, detail=f"Provider {custom_llm_provider} api base not found")
encoded_endpoint = httpx.URL(endpoint).path
@ -143,7 +143,7 @@ async def llm_passthrough_factory_proxy_route(
_request_body = await request.json()
else:
_request_body = await get_form_data(request)
if _request_body.get("stream"):
is_streaming_request = True
@ -177,17 +177,11 @@ async def gemini_proxy_route(
[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)
"""
## CHECK FOR LITELLM API KEY IN THE QUERY PARAMS - ?..key=LITELLM_API_KEY
google_ai_studio_api_key = request.query_params.get("key") or request.headers.get(
"x-goog-api-key"
)
google_ai_studio_api_key = request.query_params.get("key") or request.headers.get("x-goog-api-key")
user_api_key_dict = await user_api_key_auth(
request=request, api_key=f"Bearer {google_ai_studio_api_key}"
)
user_api_key_dict = await user_api_key_auth(request=request, api_key=f"Bearer {google_ai_studio_api_key}")
base_target_url = (
os.getenv("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com"
)
base_target_url = os.getenv("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com"
encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction
@ -220,6 +214,7 @@ async def gemini_proxy_route(
endpoint_func = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
custom_llm_provider="gemini",
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
@ -304,9 +299,7 @@ async def vllm_proxy_route(
from litellm.proxy.proxy_server import llm_router
request_body = await get_request_body(request)
is_router_model = is_passthrough_request_using_router_model(
request_body, llm_router
)
is_router_model = is_passthrough_request_using_router_model(request_body, llm_router)
is_streaming_request = is_passthrough_request_streaming(request_body)
if is_router_model and llm_router:
result = cast(
@ -321,11 +314,7 @@ async def vllm_proxy_route(
content=None,
data=None,
files=None,
json=(
request_body
if request.headers.get("content-type") == "application/json"
else None
),
json=(request_body if request.headers.get("content-type") == "application/json" else None),
params=None,
headers=None,
cookies=None,
@ -503,9 +492,7 @@ async def handle_bedrock_count_tokens(
# Extract model from request body
model = request_body.get("model")
if not model:
raise HTTPException(
status_code=400, detail={"error": "Model is required in request body"}
)
raise HTTPException(status_code=400, detail={"error": "Model is required in request body"})
# Get model parameters from router
litellm_params = {"user_api_key_dict": user_api_key_dict}
@ -544,9 +531,7 @@ async def handle_bedrock_count_tokens(
raise
except Exception as e:
verbose_proxy_logger.error(f"Error in handle_bedrock_count_tokens: {str(e)}")
raise HTTPException(
status_code=500, detail={"error": f"CountTokens processing error: {str(e)}"}
)
raise HTTPException(status_code=500, detail={"error": f"CountTokens processing error: {str(e)}"})
async def bedrock_llm_proxy_route(
@ -598,8 +583,7 @@ async def bedrock_llm_proxy_route(
raise HTTPException(
status_code=400,
detail={
"error": "Model missing from endpoint. Expected format: /model/<Model>/<endpoint>. Got: "
+ endpoint,
"error": "Model missing from endpoint. Expected format: /model/<Model>/<endpoint>. Got: " + endpoint,
},
)
@ -663,9 +647,7 @@ async def bedrock_proxy_route(
aws_region_name = litellm.utils.get_secret(secret_name="AWS_REGION_NAME")
if _is_bedrock_agent_runtime_route(endpoint=endpoint): # handle bedrock agents
base_target_url = (
f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com"
)
base_target_url = f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com"
else:
return await bedrock_llm_proxy_route(
endpoint=endpoint,
@ -686,7 +668,8 @@ async def bedrock_proxy_route(
# Add or update query parameters
from litellm.llms.bedrock.chat import BedrockConverseLLM
credentials: Credentials = BedrockConverseLLM().get_credentials()
bedrock_llm = BedrockConverseLLM()
credentials: Credentials = bedrock_llm.get_credentials() # type: ignore
sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name)
headers = {"Content-Type": "application/json"}
# Assuming the body contains JSON data, parse it
@ -694,9 +677,7 @@ async def bedrock_proxy_route(
data = await request.json()
except Exception as e:
raise HTTPException(status_code=400, detail={"error": e})
_request = AWSRequest(
method="POST", url=str(updated_url), data=json.dumps(data), headers=headers
)
_request = AWSRequest(method="POST", url=str(updated_url), data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped = _request.prepare()
@ -757,14 +738,8 @@ async def assemblyai_proxy_route(
[Docs](https://api.assemblyai.com)
"""
# Set base URL based on the route
assembly_region = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(
url=str(request.url)
)
base_target_url = (
AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(
region=assembly_region
)
)
assembly_region = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=str(request.url))
base_target_url = AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(region=assembly_region)
encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction
if not encoded_endpoint.startswith("/"):
@ -822,18 +797,14 @@ async def azure_proxy_route(
"""
base_target_url = get_secret_str(secret_name="AZURE_API_BASE")
if base_target_url is None:
raise Exception(
"Required 'AZURE_API_BASE' in environment to make pass-through calls to Azure."
)
raise Exception("Required 'AZURE_API_BASE' in environment to make pass-through calls to Azure.")
# Add or update query parameters
azure_api_key = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.AZURE.value,
region_name=None,
)
if azure_api_key is None:
raise Exception(
"Required 'AZURE_API_KEY' in environment to make pass-through calls to Azure."
)
raise Exception("Required 'AZURE_API_KEY' in environment to make pass-through calls to Azure.")
return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler(
endpoint=endpoint,
@ -857,9 +828,7 @@ class BaseVertexAIPassThroughHandler(ABC):
@staticmethod
@abstractmethod
def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
) -> str:
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str:
pass
@ -869,9 +838,7 @@ class VertexAIDiscoveryPassThroughHandler(BaseVertexAIPassThroughHandler):
return "https://discoveryengine.googleapis.com/"
@staticmethod
def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
) -> str:
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str:
return base_target_url
@ -881,9 +848,7 @@ class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler):
return get_vertex_base_url(vertex_location)
@staticmethod
def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
) -> str:
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: Optional[str]) -> str:
return get_vertex_base_url(vertex_location)
@ -949,18 +914,14 @@ async def _base_vertex_proxy_route(
location=vertex_location,
)
base_target_url = get_vertex_pass_through_handler.get_default_base_target_url(
vertex_location
)
base_target_url = get_vertex_pass_through_handler.get_default_base_target_url(vertex_location)
headers_passed_through = False
# Use headers from the incoming request if no vertex credentials are found
if vertex_credentials is None or vertex_credentials.vertex_project is None:
headers = dict(request.headers) or {}
headers_passed_through = True
verbose_proxy_logger.debug(
"default_vertex_config not set, incoming request headers %s", headers
)
verbose_proxy_logger.debug("default_vertex_config not set, incoming request headers %s", headers)
headers.pop("content-length", None)
headers.pop("host", None)
else:
@ -1126,9 +1087,7 @@ async def openai_proxy_route(
region_name=None,
)
if openai_api_key is None:
raise Exception(
"Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI."
)
raise Exception("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.")
return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler(
endpoint=endpoint,
@ -1174,9 +1133,7 @@ class BaseOpenAIPassThroughHandler:
endpoint_func = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers=BaseOpenAIPassThroughHandler._assemble_headers(
api_key=api_key, request=request
),
custom_headers=BaseOpenAIPassThroughHandler._assemble_headers(api_key=api_key, request=request),
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
@ -1193,10 +1150,7 @@ class BaseOpenAIPassThroughHandler:
"""
Appends the OpenAI-Beta header to the headers if the request is an OpenAI Assistants API request
"""
if (
RouteChecks._is_assistants_api_request(request) is True
and "OpenAI-Beta" not in headers
):
if RouteChecks._is_assistants_api_request(request) is True and "OpenAI-Beta" not in headers:
headers["OpenAI-Beta"] = "assistants=v2"
return headers
@ -1212,9 +1166,7 @@ class BaseOpenAIPassThroughHandler:
)
@staticmethod
def _join_url_paths(
base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders
) -> str:
def _join_url_paths(base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders) -> str:
"""
Properly joins a base URL with a path, preserving any existing path in the base URL.
"""
@ -1230,13 +1182,192 @@ class BaseOpenAIPassThroughHandler:
joined_path_str = str(base_url.copy_with(path=full_path))
# Apply OpenAI-specific path handling for both branches
if (
custom_llm_provider == litellm.LlmProviders.OPENAI
and "/v1/" not in joined_path_str
):
if custom_llm_provider == litellm.LlmProviders.OPENAI and "/v1/" not in joined_path_str:
# Insert v1 after api.openai.com for OpenAI requests
joined_path_str = joined_path_str.replace(
"api.openai.com/", "api.openai.com/v1/"
)
joined_path_str = joined_path_str.replace("api.openai.com/", "api.openai.com/v1/")
return joined_path_str
async def vertex_ai_live_websocket_passthrough(
websocket: WebSocket,
model: Optional[str] = None,
vertex_project: Optional[str] = None,
vertex_location: Optional[str] = None,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
):
"""
Vertex AI Live API WebSocket Pass-through Function
This function provides WebSocket passthrough functionality for Vertex AI Live API,
allowing real-time communication with Google's Live API service.
Note: This function should be registered in proxy_server.py using:
app.websocket("/vertex_ai/live")(vertex_ai_live_websocket_passthrough)
"""
from litellm.proxy.proxy_server import proxy_logging_obj
_ = user_api_key_dict # passthrough route already authenticated; avoid lint warnings
await websocket.accept()
incoming_headers = dict(websocket.headers)
vertex_credentials_config = passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
location=vertex_location,
)
if vertex_credentials_config is None:
# Attempt to load defaults from environment/config if not already initialised
passthrough_endpoint_router.set_default_vertex_config()
vertex_credentials_config = passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
location=vertex_location,
)
resolved_project = vertex_project
resolved_location: Optional[str] = vertex_location
credentials_value: Optional[str] = None
if vertex_credentials_config is not None:
resolved_project = resolved_project or vertex_credentials_config.vertex_project
temp_location = (
resolved_location or vertex_credentials_config.vertex_location
)
# Ensure resolved_location is a string
if isinstance(temp_location, dict):
resolved_location = str(temp_location)
elif temp_location is not None:
resolved_location = str(temp_location)
else:
resolved_location = None
credentials_value = str(vertex_credentials_config.vertex_credentials) if vertex_credentials_config.vertex_credentials is not None else None
try:
resolved_location = resolved_location or (
vertex_llm_base.get_default_vertex_location()
)
if model:
resolved_location = vertex_llm_base.get_vertex_region(
vertex_region=resolved_location,
model=model,
)
(
access_token,
resolved_project,
) = await vertex_llm_base._ensure_access_token_async(
credentials=credentials_value,
project_id=resolved_project,
custom_llm_provider="vertex_ai_beta",
)
except Exception as e:
verbose_proxy_logger.exception(
"Failed to prepare Vertex AI credentials for live passthrough"
)
# Log the authentication failure using proxy_logging_obj
if proxy_logging_obj and user_api_key_dict:
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
request_data={},
)
if websocket.client_state != WebSocketState.DISCONNECTED:
await websocket.close(code=1011, reason="Vertex AI authentication failed")
return
host_location = resolved_location or vertex_llm_base.get_default_vertex_location()
host = (
"aiplatform.googleapis.com"
if host_location == "global"
else f"{host_location}-aiplatform.googleapis.com"
)
service_url = (
f"wss://{host}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
)
upstream_headers = {
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
}
if resolved_project:
upstream_headers["x-goog-user-project"] = resolved_project
# Forward any custom x-goog-* headers provided by the caller if we haven't overridden them
for header_name, header_value in incoming_headers.items():
lower_header = header_name.lower()
if lower_header.startswith("x-goog-") and header_name not in upstream_headers:
upstream_headers[header_name] = header_value
# Use the new WebSocket passthrough pattern
if user_api_key_dict is None:
raise ValueError("user_api_key_dict is required for WebSocket passthrough")
return await websocket_passthrough_request(
websocket=websocket,
target=service_url,
custom_headers=upstream_headers,
user_api_key_dict=user_api_key_dict,
forward_headers=False,
endpoint="/vertex_ai/live",
accept_websocket=False,
)
def create_vertex_ai_live_websocket_endpoint():
"""
Create a Vertex AI Live WebSocket endpoint using the new passthrough pattern.
This demonstrates how to use the create_websocket_passthrough_route function
for a provider-specific WebSocket endpoint.
"""
# This would be used like:
# endpoint_func = create_vertex_ai_live_websocket_endpoint()
# app.websocket("/vertex_ai/live")(endpoint_func)
# For now, we'll keep the existing implementation since it has
# provider-specific logic for Vertex AI credentials and headers
return vertex_ai_live_websocket_passthrough
def create_generic_websocket_passthrough_endpoint(
provider: str,
target_url: str,
custom_headers: Optional[dict] = None,
forward_headers: bool = False,
cost_per_request: Optional[float] = None,
):
"""
Create a generic WebSocket passthrough endpoint for any provider.
This demonstrates the new WebSocket passthrough pattern that's similar to
the HTTP create_pass_through_route function.
Args:
provider: The provider name (e.g., "anthropic", "cohere")
target_url: The target WebSocket URL
custom_headers: Custom headers to include
forward_headers: Whether to forward incoming headers
Returns:
A WebSocket endpoint function that can be registered with app.websocket()
Example usage:
# Create a WebSocket endpoint for Anthropic
anthropic_ws_func = create_generic_websocket_passthrough_endpoint(
provider="anthropic",
target_url="wss://api.anthropic.com/v1/ws",
custom_headers={"x-api-key": "your-api-key"},
forward_headers=True
)
# Register it in proxy_server.py
app.websocket("/anthropic/ws")(anthropic_ws_func)
"""
return create_websocket_passthrough_route(
endpoint=f"/{provider}/ws",
target=target_url,
custom_headers=custom_headers,
_forward_headers=forward_headers,
cost_per_request=cost_per_request,
)

View file

@ -0,0 +1,204 @@
import json
import re
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator as GeminiModelResponseIterator,
)
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.types.utils import (
ModelResponse,
TextCompletionResponse,
)
if TYPE_CHECKING:
from ..success_handler import PassThroughEndpointLogging
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
else:
PassThroughEndpointLogging = Any
EndpointType = Any
class GeminiPassthroughLoggingHandler:
@staticmethod
def gemini_passthrough_handler(
httpx_response: httpx.Response,
response_body: dict,
logging_obj: LiteLLMLoggingObj,
url_route: str,
result: str,
start_time: datetime,
end_time: datetime,
cache_hit: bool,
request_body: dict,
**kwargs,
) -> PassThroughEndpointLoggingTypedDict:
if "generateContent" in url_route:
model = GeminiPassthroughLoggingHandler.extract_model_from_url(url_route)
# Use Gemini config for transformation
instance_of_gemini_llm = litellm.GoogleAIStudioGeminiConfig()
litellm_model_response: ModelResponse = instance_of_gemini_llm.transform_response(
model=model,
messages=[{"role": "user", "content": "no-message-pass-through-endpoint"}],
raw_response=httpx_response,
model_response=litellm.ModelResponse(),
logging_obj=logging_obj,
optional_params={},
litellm_params={},
api_key="",
request_data={},
encoding=litellm.encoding,
)
kwargs = GeminiPassthroughLoggingHandler._create_gemini_response_logging_payload_for_generate_content(
litellm_model_response=litellm_model_response,
model=model,
kwargs=kwargs,
start_time=start_time,
end_time=end_time,
logging_obj=logging_obj,
custom_llm_provider="gemini",
)
return {
"result": litellm_model_response,
"kwargs": kwargs,
}
else:
return {
"result": None,
"kwargs": kwargs,
}
@staticmethod
def _handle_logging_gemini_collected_chunks(
litellm_logging_obj: LiteLLMLoggingObj,
passthrough_success_handler_obj: PassThroughEndpointLogging,
url_route: str,
request_body: dict,
endpoint_type: EndpointType,
start_time: datetime,
all_chunks: List[str],
model: Optional[str],
end_time: datetime,
) -> PassThroughEndpointLoggingTypedDict:
"""
Takes raw chunks from Gemini passthrough endpoint and logs them in litellm callbacks
- Builds complete response from chunks
- Creates standard logging object
- Logs in litellm callbacks
"""
kwargs: Dict[str, Any] = {}
model = model or GeminiPassthroughLoggingHandler.extract_model_from_url(url_route)
complete_streaming_response = GeminiPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
url_route=url_route,
)
if complete_streaming_response is None:
verbose_proxy_logger.error(
"Unable to build complete streaming response for Gemini passthrough endpoint, not logging..."
)
return {
"result": None,
"kwargs": kwargs,
}
kwargs = GeminiPassthroughLoggingHandler._create_gemini_response_logging_payload_for_generate_content(
litellm_model_response=complete_streaming_response,
model=model,
kwargs=kwargs,
start_time=start_time,
end_time=end_time,
logging_obj=litellm_logging_obj,
custom_llm_provider="gemini",
)
return {
"result": complete_streaming_response,
"kwargs": kwargs,
}
@staticmethod
def _build_complete_streaming_response(
all_chunks: List[str],
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
url_route: str,
) -> Optional[Union[ModelResponse, TextCompletionResponse]]:
parsed_chunks = []
if "generateContent" in url_route or "streamGenerateContent" in url_route:
gemini_iterator: Any = GeminiModelResponseIterator(
streaming_response=None,
sync_stream=False,
logging_obj=litellm_logging_obj,
)
chunk_parsing_logic: Any = gemini_iterator._common_chunk_parsing_logic
parsed_chunks = [chunk_parsing_logic(chunk) for chunk in all_chunks]
else:
return None
if len(parsed_chunks) == 0:
return None
all_openai_chunks = []
for parsed_chunk in parsed_chunks:
if parsed_chunk is None:
continue
all_openai_chunks.append(parsed_chunk)
complete_streaming_response = litellm.stream_chunk_builder(chunks=all_openai_chunks)
return complete_streaming_response
@staticmethod
def extract_model_from_url(url: str) -> str:
pattern = r"/models/([^:]+)"
match = re.search(pattern, url)
if match:
return match.group(1)
return "unknown"
@staticmethod
def _create_gemini_response_logging_payload_for_generate_content(
litellm_model_response: Union[ModelResponse, TextCompletionResponse],
model: str,
kwargs: dict,
start_time: datetime,
end_time: datetime,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str,
):
"""
Create the standard logging object for Gemini passthrough generateContent (streaming and non-streaming)
"""
response_cost = litellm.completion_cost(
completion_response=litellm_model_response,
model=model,
custom_llm_provider="gemini",
)
kwargs["response_cost"] = response_cost
kwargs["model"] = model
kwargs["custom_llm_provider"] = custom_llm_provider
# pretty print standard logging object
verbose_proxy_logger.debug("kwargs= %s", json.dumps(kwargs, indent=4))
# set litellm_call_id to logging response object
litellm_model_response.id = logging_obj.litellm_call_id
logging_obj.model = litellm_model_response.model or model
logging_obj.model_call_details["model"] = logging_obj.model
logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
logging_obj.model_call_details["response_cost"] = response_cost
return kwargs

View file

@ -0,0 +1,398 @@
"""
Vertex AI Live API WebSocket Passthrough Logging Handler
Handles cost tracking and logging for Vertex AI Live API WebSocket passthrough endpoints.
Supports different modalities: text, audio, video, and web search.
"""
from datetime import datetime
from typing import Any, Dict, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import (
BasePassthroughLoggingHandler,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
PassThroughEndpointLoggingTypedDict,
)
from litellm.types.utils import LlmProviders, ModelResponse, Usage
from litellm.utils import get_model_info
class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""
Handles cost tracking and logging for Vertex AI Live API WebSocket passthrough.
Supports:
- Text tokens (input/output)
- Audio tokens (input/output)
- Video tokens (input/output)
- Web search requests
- Tool use tokens
"""
def _build_complete_streaming_response(self, *args, **kwargs):
"""Not applicable for WebSocket passthrough."""
return None
def get_provider_config(self, model: str):
"""Return Vertex AI provider configuration."""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
return VertexGeminiConfig()
@property
def llm_provider_name(self) -> LlmProviders:
"""Return the LLM provider name."""
return LlmProviders.VERTEX_AI
@staticmethod
def _extract_usage_metadata_from_websocket_messages(
websocket_messages: List[Dict],
) -> Optional[Dict]:
"""
Extract and aggregate usage metadata from a list of WebSocket messages.
Args:
websocket_messages: List of WebSocket messages from the Live API
Returns:
Dictionary containing aggregated usage metadata, or None if not found
"""
all_usage_metadata = []
# Collect all usage metadata messages
for message in websocket_messages:
if isinstance(message, dict) and "usageMetadata" in message:
all_usage_metadata.append(message["usageMetadata"])
if not all_usage_metadata:
return None
# If only one usage metadata, return it as-is
if len(all_usage_metadata) == 1:
return all_usage_metadata[0]
# Aggregate multiple usage metadata messages
aggregated: Dict[str, Any] = {
"promptTokenCount": 0,
"candidatesTokenCount": 0,
"totalTokenCount": 0,
"promptTokensDetails": [],
"candidatesTokensDetails": [],
}
# Aggregate token counts
for usage in all_usage_metadata:
aggregated["promptTokenCount"] += usage.get("promptTokenCount", 0)
aggregated["candidatesTokenCount"] += usage.get("candidatesTokenCount", 0)
aggregated["totalTokenCount"] += usage.get("totalTokenCount", 0)
# Aggregate token details by modality
modality_totals = {}
for usage in all_usage_metadata:
# Process prompt tokens details
for detail in usage.get("promptTokensDetails", []):
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality not in modality_totals:
modality_totals[modality] = {"prompt": 0, "candidate": 0}
modality_totals[modality]["prompt"] += token_count
# Process candidate tokens details
for detail in usage.get("candidatesTokensDetails", []):
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality not in modality_totals:
modality_totals[modality] = {"prompt": 0, "candidate": 0}
modality_totals[modality]["candidate"] += token_count
# Convert aggregated modality totals back to details format
for modality, totals in modality_totals.items():
if totals["prompt"] > 0:
aggregated["promptTokensDetails"].append(
{"modality": modality, "tokenCount": totals["prompt"]}
)
if totals["candidate"] > 0:
aggregated["candidatesTokensDetails"].append(
{"modality": modality, "tokenCount": totals["candidate"]}
)
# Add any additional fields from the first usage metadata
first_usage = all_usage_metadata[0]
for key, value in first_usage.items():
if key not in aggregated:
aggregated[key] = value
return aggregated
@staticmethod
def _calculate_live_api_cost(
model: str,
usage_metadata: Dict,
custom_llm_provider: str = "vertex_ai",
) -> float:
"""
Calculate cost for Vertex AI Live API based on usage metadata.
Args:
model: The model name (e.g., "gemini-2.0-flash-live-preview-04-09")
usage_metadata: Usage metadata from the Live API response
custom_llm_provider: The LLM provider (default: "vertex_ai")
Returns:
Total cost in USD
"""
try:
# Get model pricing information
model_info = get_model_info(
model=model, custom_llm_provider=custom_llm_provider
)
verbose_proxy_logger.debug(
f"Vertex AI Live API model info for '{model}': {model_info}"
)
# Check if pricing info is available
if not model_info or not model_info.get("input_cost_per_token"):
verbose_proxy_logger.error(
f"No pricing info found for {model} in local model pricing database"
)
return 0.0
total_cost = 0.0
# Extract token counts from usage metadata
prompt_token_count = usage_metadata.get("promptTokenCount", 0)
candidates_token_count = usage_metadata.get("candidatesTokenCount", 0)
# Calculate base text token costs
input_cost_per_token = model_info.get("input_cost_per_token", 0.0)
output_cost_per_token = model_info.get("output_cost_per_token", 0.0)
total_cost += prompt_token_count * input_cost_per_token
total_cost += candidates_token_count * output_cost_per_token
# Handle modality-specific costs if present
prompt_tokens_details = usage_metadata.get("promptTokensDetails", [])
candidates_tokens_details = usage_metadata.get(
"candidatesTokensDetails", []
)
# Process prompt tokens by modality
for detail in prompt_tokens_details:
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality == "AUDIO":
audio_cost_per_token = model_info.get(
"input_cost_per_audio_token", 0.0
)
total_cost += token_count * audio_cost_per_token
elif modality == "VIDEO":
# Video tokens are typically per second, but we'll treat as per token for now
video_cost_per_token = model_info.get(
"input_cost_per_video_per_second", 0.0
)
total_cost += token_count * video_cost_per_token
# TEXT tokens are already handled above
# Process candidate tokens by modality
for detail in candidates_tokens_details:
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality == "AUDIO":
audio_cost_per_token = model_info.get(
"output_cost_per_audio_token", 0.0
)
total_cost += token_count * audio_cost_per_token
elif modality == "VIDEO":
# Video tokens are typically per second, but we'll treat as per token for now
video_cost_per_token = model_info.get(
"output_cost_per_video_per_second", 0.0
)
total_cost += token_count * video_cost_per_token
# TEXT tokens are already handled above
# Handle web search costs if present
tool_use_prompt_token_count = usage_metadata.get(
"toolUsePromptTokenCount", 0
)
if tool_use_prompt_token_count > 0:
# Web search typically has a fixed cost per request
web_search_cost = model_info.get("web_search_cost_per_request", 0.0)
if isinstance(web_search_cost, (int, float)) and web_search_cost > 0:
total_cost += web_search_cost
else:
# Fallback to token-based pricing for tool use
total_cost += tool_use_prompt_token_count * input_cost_per_token
verbose_proxy_logger.debug(
f"Vertex AI Live API cost calculation - Model: {model}, "
f"Prompt tokens: {prompt_token_count}, "
f"Candidate tokens: {candidates_token_count}, "
f"Total cost: ${total_cost:.6f}"
)
return total_cost
except Exception as e:
verbose_proxy_logger.error(
f"Error calculating Vertex AI Live API cost: {e}"
)
return 0.0
@staticmethod
def _create_usage_object_from_metadata(
usage_metadata: Dict,
model: str,
) -> Usage:
"""
Create a LiteLLM Usage object from Live API usage metadata.
Args:
usage_metadata: Usage metadata from the Live API response
model: The model name
Returns:
LiteLLM Usage object
"""
prompt_tokens = usage_metadata.get("promptTokenCount", 0)
completion_tokens = usage_metadata.get("candidatesTokenCount", 0)
total_tokens = usage_metadata.get("totalTokenCount", 0)
# Create modality-specific token details if available
prompt_tokens_details = usage_metadata.get("promptTokensDetails", [])
candidates_tokens_details = usage_metadata.get("candidatesTokensDetails", [])
# Extract text tokens from details
text_prompt_tokens = 0
text_completion_tokens = 0
for detail in prompt_tokens_details:
if detail.get("modality") == "TEXT":
text_prompt_tokens = detail.get("tokenCount", 0)
break
for detail in candidates_tokens_details:
if detail.get("modality") == "TEXT":
text_completion_tokens = detail.get("tokenCount", 0)
break
# If no text tokens found in details, use total counts
if text_prompt_tokens == 0:
text_prompt_tokens = prompt_tokens
if text_completion_tokens == 0:
text_completion_tokens = completion_tokens
return Usage(
prompt_tokens=text_prompt_tokens,
completion_tokens=text_completion_tokens,
total_tokens=total_tokens,
)
def vertex_ai_live_passthrough_handler(
self,
websocket_messages: List[Dict],
logging_obj,
url_route: str,
start_time: datetime,
end_time: datetime,
request_body: dict,
**kwargs,
) -> PassThroughEndpointLoggingTypedDict:
"""
Handle cost tracking and logging for Vertex AI Live API WebSocket passthrough.
Args:
websocket_messages: List of WebSocket messages from the Live API
logging_obj: LiteLLM logging object
url_route: The URL route that was called
start_time: Request start time
end_time: Request end time
request_body: The original request body
**kwargs: Additional keyword arguments
Returns:
Dictionary containing the result and kwargs for logging
"""
try:
# Extract model from request body or kwargs
model = kwargs.get("model", "gemini-2.0-flash-live-preview-04-09")
custom_llm_provider = kwargs.get("custom_llm_provider", "vertex_ai")
verbose_proxy_logger.debug(
f"Vertex AI Live API model: {model}, custom_llm_provider: {custom_llm_provider}"
)
# Extract usage metadata from WebSocket messages
usage_metadata = self._extract_usage_metadata_from_websocket_messages(
websocket_messages
)
if not usage_metadata:
verbose_proxy_logger.warning(
"No usage metadata found in Vertex AI Live API WebSocket messages"
)
return {
"result": None,
"kwargs": kwargs,
}
# Calculate cost using Live API specific pricing
response_cost = self._calculate_live_api_cost(
model=model,
usage_metadata=usage_metadata,
custom_llm_provider=custom_llm_provider,
)
# Create Usage object for standard LiteLLM logging
usage = self._create_usage_object_from_metadata(
usage_metadata=usage_metadata,
model=model,
)
# Create a mock ModelResponse for standard logging
litellm_model_response = ModelResponse(
id=f"vertex-ai-live-{start_time.timestamp()}",
object="chat.completion",
created=int(start_time.timestamp()),
model=model,
usage=usage,
choices=[],
)
# Update kwargs with cost information
kwargs["response_cost"] = response_cost
kwargs["model"] = model
kwargs["custom_llm_provider"] = custom_llm_provider
# Safely log the model name: only allow known safe formats, redact otherwise.
import re
allowed_pattern = re.compile(r"^[A-Za-z0-9._\-:]+$")
safe_model = model if isinstance(model, str) and allowed_pattern.match(model) else "[REDACTED]"
verbose_proxy_logger.debug(
f"Vertex AI Live API passthrough cost tracking - "
f"Model: {safe_model}, Cost: ${response_cost:.6f}, "
f"Prompt tokens: {usage.prompt_tokens}, "
f"Completion tokens: {usage.completion_tokens}"
)
return {
"result": litellm_model_response,
"kwargs": kwargs,
}
except Exception as e:
verbose_proxy_logger.error(
f"Error in Vertex AI Live API passthrough handler: {e}"
)
return {
"result": None,
"kwargs": kwargs,
}

View file

@ -110,7 +110,7 @@ class VertexPassthroughLoggingHandler:
PassthroughCallTypes.passthrough_image_generation.value
)
elif VertexPassthroughLoggingHandler._is_multimodal_embedding_response(
json_response=_json_response,
json_response=_json_response,
):
# Use multimodal embedding transformation
vertex_multimodal_config = VertexAIMultimodalEmbeddingConfig()
@ -137,6 +137,15 @@ class VertexPassthroughLoggingHandler:
logging_obj.model = model
logging_obj.model_call_details["model"] = logging_obj.model
response_cost = litellm.completion_cost(
completion_response=litellm_prediction_response,
model=model,
custom_llm_provider="vertex_ai",
)
kwargs["response_cost"] = response_cost
kwargs["model"] = model
logging_obj.model_call_details["response_cost"] = response_cost
return {
"result": litellm_prediction_response,
@ -221,7 +230,9 @@ class VertexPassthroughLoggingHandler:
- Logs in litellm callbacks
"""
kwargs: Dict[str, Any] = {}
model = model or VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
model = model or VertexPassthroughLoggingHandler.extract_model_from_url(
url_route
)
complete_streaming_response = (
VertexPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=all_chunks,
@ -340,13 +351,13 @@ class VertexPassthroughLoggingHandler:
"""
Detect if the response is from a multimodal embedding request.
Check if the response contains multimodal embedding fields:
- Docs: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-embeddings-api#response-body
Check if the response contains multimodal embedding fields:
- Docs: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-embeddings-api#response-body
Args:
json_response: The JSON response from Vertex AI
Returns:
bool: True if this is a multimodal embedding response
"""
@ -358,10 +369,14 @@ class VertexPassthroughLoggingHandler:
# Check for multimodal embedding response fields
if any(
key in prediction
for key in ["textEmbedding", "imageEmbedding", "videoEmbeddings"]
for key in [
"textEmbedding",
"imageEmbedding",
"videoEmbeddings",
]
):
return True
return False
@staticmethod

View file

@ -25,6 +25,9 @@ from .llm_provider_handlers.cohere_passthrough_logging_handler import (
from .llm_provider_handlers.vertex_passthrough_logging_handler import (
VertexPassthroughLoggingHandler,
)
from .llm_provider_handlers.gemini_passthrough_logging_handler import (
GeminiPassthroughLoggingHandler,
)
cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler()
@ -44,13 +47,17 @@ class PassThroughEndpointLogging:
# Cohere
self.TRACKED_COHERE_ROUTES = ["/v2/chat"]
self.assemblyai_passthrough_logging_handler = (
AssemblyAIPassthroughLoggingHandler()
)
self.assemblyai_passthrough_logging_handler = AssemblyAIPassthroughLoggingHandler()
# Langfuse
self.TRACKED_LANGFUSE_ROUTES = ["/langfuse/"]
# Gemini
self.TRACKED_GEMINI_ROUTES = ["generateContent", "streamGenerateContent"]
# Vertex AI Live API WebSocket
self.TRACKED_VERTEX_AI_LIVE_ROUTES = ["/vertex_ai/live"]
async def _handle_logging(
self,
logging_obj: LiteLLMLoggingObj,
@ -78,11 +85,7 @@ class PassThroughEndpointLogging:
# Handle async logging
await logging_obj.async_success_handler(
result=(
json.dumps(result)
if isinstance(result, dict)
else standard_logging_response_object
),
result=(json.dumps(result) if isinstance(result, dict) else standard_logging_response_object),
start_time=start_time,
end_time=end_time,
cache_hit=False,
@ -100,6 +103,7 @@ class PassThroughEndpointLogging:
start_time: datetime,
end_time: datetime,
cache_hit: bool,
custom_llm_provider: Optional[str] = None,
**kwargs,
):
return_dict = {
@ -107,22 +111,34 @@ class PassThroughEndpointLogging:
"kwargs": kwargs,
}
standard_logging_response_object: Optional[Any] = None
if self.is_vertex_route(url_route):
vertex_passthrough_logging_handler_result = (
VertexPassthroughLoggingHandler.vertex_passthrough_handler(
httpx_response=httpx_response,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
**kwargs,
)
if self.is_gemini_route(url_route, custom_llm_provider):
gemini_passthrough_logging_handler_result = GeminiPassthroughLoggingHandler.gemini_passthrough_handler(
httpx_response=httpx_response,
response_body=response_body or {},
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)
standard_logging_response_object = (
vertex_passthrough_logging_handler_result["result"]
standard_logging_response_object = gemini_passthrough_logging_handler_result["result"]
kwargs = gemini_passthrough_logging_handler_result["kwargs"]
elif self.is_vertex_route(url_route):
vertex_passthrough_logging_handler_result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
httpx_response=httpx_response,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
**kwargs,
)
standard_logging_response_object = vertex_passthrough_logging_handler_result["result"]
kwargs = vertex_passthrough_logging_handler_result["kwargs"]
elif self.is_anthropic_route(url_route):
anthropic_passthrough_logging_handler_result = (
@ -139,55 +155,72 @@ class PassThroughEndpointLogging:
)
)
standard_logging_response_object = (
anthropic_passthrough_logging_handler_result["result"]
)
standard_logging_response_object = anthropic_passthrough_logging_handler_result["result"]
kwargs = anthropic_passthrough_logging_handler_result["kwargs"]
elif self.is_cohere_route(url_route):
cohere_passthrough_logging_handler_result = (
cohere_passthrough_logging_handler.passthrough_chat_handler(
httpx_response=httpx_response,
response_body=response_body or {},
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)
)
standard_logging_response_object = (
cohere_passthrough_logging_handler_result["result"]
cohere_passthrough_logging_handler_result = cohere_passthrough_logging_handler.passthrough_chat_handler(
httpx_response=httpx_response,
response_body=response_body or {},
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)
standard_logging_response_object = cohere_passthrough_logging_handler_result["result"]
kwargs = cohere_passthrough_logging_handler_result["kwargs"]
elif self.is_openai_route(url_route) and self._is_supported_openai_endpoint(url_route):
elif self.is_openai_route(url_route) and self._is_supported_openai_endpoint(
url_route
):
from .llm_provider_handlers.openai_passthrough_logging_handler import (
OpenAIPassthroughLoggingHandler,
)
openai_passthrough_logging_handler_result = (
OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
httpx_response=httpx_response,
response_body=response_body or {},
openai_passthrough_logging_handler_result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
httpx_response=httpx_response,
response_body=response_body or {},
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)
standard_logging_response_object = openai_passthrough_logging_handler_result["result"]
kwargs = openai_passthrough_logging_handler_result["kwargs"]
elif self.is_vertex_ai_live_route(url_route):
from .llm_provider_handlers.vertex_ai_live_passthrough_logging_handler import (
VertexAILivePassthroughLoggingHandler,
)
vertex_ai_live_handler = VertexAILivePassthroughLoggingHandler()
# For WebSocket responses, response_body should be a list of messages
websocket_messages: list[dict[str, Any]] = response_body if isinstance(response_body, list) else []
vertex_ai_live_handler_result = (
vertex_ai_live_handler.vertex_ai_live_passthrough_handler(
websocket_messages=websocket_messages,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)
)
standard_logging_response_object = (
openai_passthrough_logging_handler_result["result"]
)
kwargs = openai_passthrough_logging_handler_result["kwargs"]
standard_logging_response_object = vertex_ai_live_handler_result["result"]
kwargs = vertex_ai_live_handler_result["kwargs"]
return_dict[
"standard_logging_response_object"
] = standard_logging_response_object
return_dict["kwargs"] = kwargs
return return_dict
@ -203,21 +236,13 @@ class PassThroughEndpointLogging:
cache_hit: bool,
request_body: dict,
passthrough_logging_payload: PassthroughStandardLoggingPayload,
custom_llm_provider: Optional[str] = None,
**kwargs,
):
standard_logging_response_object: Optional[
PassThroughEndpointLoggingResultValues
] = None
logging_obj.model_call_details[
"passthrough_logging_payload"
] = passthrough_logging_payload
standard_logging_response_object: Optional[PassThroughEndpointLoggingResultValues] = None
logging_obj.model_call_details["passthrough_logging_payload"] = passthrough_logging_payload
if self.is_assemblyai_route(url_route):
if (
AssemblyAIPassthroughLoggingHandler._should_log_request(
httpx_response.request.method
)
is not True
):
if AssemblyAIPassthroughLoggingHandler._should_log_request(httpx_response.request.method) is not True:
return
self.assemblyai_passthrough_logging_handler.assemblyai_passthrough_logging_handler(
httpx_response=httpx_response,
@ -235,30 +260,25 @@ class PassThroughEndpointLogging:
# Don't log langfuse pass-through requests
return
else:
normalized_llm_passthrough_logging_payload = (
self.normalize_llm_passthrough_logging_payload(
httpx_response=httpx_response,
response_body=response_body,
request_body=request_body,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
**kwargs,
)
)
standard_logging_response_object = (
normalized_llm_passthrough_logging_payload[
"standard_logging_response_object"
]
normalized_llm_passthrough_logging_payload = self.normalize_llm_passthrough_logging_payload(
httpx_response=httpx_response,
response_body=response_body,
request_body=request_body,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
custom_llm_provider=custom_llm_provider,
**kwargs,
)
standard_logging_response_object = normalized_llm_passthrough_logging_payload[
"standard_logging_response_object"
]
kwargs = normalized_llm_passthrough_logging_payload["kwargs"]
if standard_logging_response_object is None:
standard_logging_response_object = StandardPassThroughResponseObject(
response=httpx_response.text
)
standard_logging_response_object = StandardPassThroughResponseObject(response=httpx_response.text)
kwargs = self._set_cost_per_request(
logging_obj=logging_obj,
@ -309,26 +329,43 @@ class PassThroughEndpointLogging:
return True
return False
def is_vertex_ai_live_route(self, url_route: str):
"""Check if the URL route is a Vertex AI Live API WebSocket route."""
if not url_route:
return False
for route in self.TRACKED_VERTEX_AI_LIVE_ROUTES:
if route in url_route:
return True
return False
def is_openai_route(self, url_route: str):
"""Check if the URL route is an OpenAI API route."""
if not url_route:
return False
parsed_url = urlparse(url_route)
return parsed_url.hostname and (
"api.openai.com" in parsed_url.hostname
or "openai.azure.com" in parsed_url.hostname
"api.openai.com" in parsed_url.hostname or "openai.azure.com" in parsed_url.hostname
)
def is_gemini_route(self, url_route: str, custom_llm_provider: Optional[str] = None):
"""Check if the URL route is a Gemini API route."""
for route in self.TRACKED_GEMINI_ROUTES:
if route in url_route and custom_llm_provider == "gemini":
return True
return False
def _is_supported_openai_endpoint(self, url_route: str) -> bool:
"""Check if the OpenAI endpoint is supported by the passthrough logging handler."""
from .llm_provider_handlers.openai_passthrough_logging_handler import (
OpenAIPassthroughLoggingHandler,
)
return (
OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route(url_route) or
OpenAIPassthroughLoggingHandler.is_openai_image_generation_route(url_route) or
OpenAIPassthroughLoggingHandler.is_openai_image_editing_route(url_route)
OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route(url_route)
or OpenAIPassthroughLoggingHandler.is_openai_image_generation_route(
url_route
)
or OpenAIPassthroughLoggingHandler.is_openai_image_editing_route(url_route)
)
def _set_cost_per_request(
@ -347,11 +384,7 @@ class PassThroughEndpointLogging:
# Check if cost per request is set
#########################################################
if passthrough_logging_payload.get("cost_per_request") is not None:
kwargs["response_cost"] = passthrough_logging_payload.get(
"cost_per_request"
)
logging_obj.model_call_details[
"response_cost"
] = passthrough_logging_payload.get("cost_per_request")
kwargs["response_cost"] = passthrough_logging_payload.get("cost_per_request")
logging_obj.model_call_details["response_cost"] = passthrough_logging_payload.get("cost_per_request")
return kwargs

View file

@ -175,4 +175,4 @@ class InMemoryPromptRegistry:
return self.prompt_id_to_custom_prompt.get(prompt_id)
IN_MEMORY_PROMPT_REGISTRY = InMemoryPromptRegistry()
IN_MEMORY_PROMPT_REGISTRY = InMemoryPromptRegistry()

View file

@ -23,22 +23,35 @@ model_list:
litellm_params:
model: gemini/*
api_key: os.environ/GEMINI_API_KEY
- model_name: vertex_ai/*
litellm_params:
model: vertex_ai/*
- model_name: "grok-4"
model_info:
mode: completion
litellm_params:
model: oci/xai.grok-4
oci_key: ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
oci_region: us-phoenix-1
oci_user: ocid1.user.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
oci_fingerprint: aa:bb:cc:dd:ee:ff:11:22:33:44:55:66:77:88:99:00
oci_tenancy: ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
oci_key_file: /path/to/oci_api_key.pem
oci_compartment_id: ocid1.compartment.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
drop_params: True
guardrails:
- guardrail_name: lakera
- guardrail_name: "bedrock-pre-guard"
litellm_params:
guardrail: lakera_v2
mode: pre_call
api_key: os.environ/LAKERA_API_KEY
default_on: false
project_id: project-9770817088
breakdown: true
payload: true
dev_info: true
guardrail: bedrock # supported values: "aporia", "bedrock", "lakera"
mode: "during_call"
guardrailIdentifier: ff6ujrregl1q
guardrailVersion: "DRAFT"
litellm_settings:
callbacks: ["datadog"]
include_cost_in_streaming_usage: true
datadog_params:
turn_off_message_logging: true
datadog_llm_observability_params:

View file

@ -151,6 +151,7 @@ from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
router as mcp_discoverable_endpoints_router,
)
@ -252,7 +253,9 @@ from litellm.proxy.management_endpoints.customer_endpoints import (
from litellm.proxy.management_endpoints.internal_user_endpoints import (
router as internal_user_router,
)
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
from litellm.proxy.management_endpoints.internal_user_endpoints import (
user_update,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
delete_verification_tokens,
duration_in_seconds,
@ -299,13 +302,18 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi
from litellm.proxy.openai_files_endpoints.files_endpoints import (
router as openai_files_router,
)
from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config
from litellm.proxy.openai_files_endpoints.files_endpoints import (
set_files_config,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
passthrough_endpoint_router,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
router as llm_passthrough_router,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
vertex_ai_live_websocket_passthrough,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
initialize_pass_through_endpoints,
)
@ -411,6 +419,8 @@ from fastapi import (
Request,
Response,
UploadFile,
WebSocket,
WebSocketDisconnect,
applications,
status,
)
@ -685,6 +695,8 @@ app = FastAPI(
lifespan=proxy_startup_event,
)
vertex_live_passthrough_vertex_base = VertexBase()
### CUSTOM API DOCS [ENTERPRISE FEATURE] ###
# Custom OpenAPI schema generator to include only selected routes
@ -1865,6 +1877,15 @@ class ProxyConfig:
verbose_proxy_logger.info(
f"{blue_color_code}Set Global BitBucket Config on LiteLLM Proxy{reset_color_code}"
)
elif key == "global_gitlab_config":
from litellm.integrations.gitlab import (
set_global_gitlab_config,
)
set_global_gitlab_config(value)
verbose_proxy_logger.info(
f"{blue_color_code}Set Global Gitlab Config on LiteLLM Proxy{reset_color_code}"
)
elif key == "callbacks":
initialize_callbacks_on_proxy(
value=value,
@ -2599,6 +2620,31 @@ class ProxyConfig:
proxy_logging_obj=proxy_logging_obj,
)
def _add_callback_from_db_to_in_memory_litellm_callbacks(
self,
callback: str,
event_types: List[Literal["success", "failure"]],
existing_callbacks: list,
) -> None:
"""
Helper method to add a single callback to litellm for specified event types.
Args:
callback: The callback name to add
event_types: List of event types (e.g., ["success"], ["failure"], or ["success", "failure"])
existing_callbacks: The existing callback list to check against
"""
if callback in litellm._known_custom_logger_compatible_callbacks:
for event_type in event_types:
_add_custom_logger_callback_to_specific_event(callback, event_type)
elif callback not in existing_callbacks:
if event_types == ["success"]:
litellm.logging_callback_manager.add_litellm_success_callback(callback)
elif event_types == ["failure"]:
litellm.logging_callback_manager.add_litellm_failure_callback(callback)
else: # Both success and failure
litellm.logging_callback_manager.add_litellm_callback(callback)
def _add_callbacks_from_db_config(self, config_data: dict) -> None:
"""
Adds callbacks from DB config to litellm
@ -2606,35 +2652,31 @@ class ProxyConfig:
litellm_settings = config_data.get("litellm_settings", {}) or {}
success_callbacks = litellm_settings.get("success_callback", None)
failure_callbacks = litellm_settings.get("failure_callback", None)
callbacks = litellm_settings.get("callbacks", None)
if success_callbacks is not None and isinstance(success_callbacks, list):
for success_callback in success_callbacks:
if (
success_callback
in litellm._known_custom_logger_compatible_callbacks
):
_add_custom_logger_callback_to_specific_event(
success_callback, "success"
)
elif success_callback not in litellm.success_callback:
litellm.logging_callback_manager.add_litellm_success_callback(
success_callback
)
self._add_callback_from_db_to_in_memory_litellm_callbacks(
callback=success_callback,
event_types=["success"],
existing_callbacks=litellm.success_callback,
)
# Add failure callbacks from DB to litellm
if failure_callbacks is not None and isinstance(failure_callbacks, list):
for failure_callback in failure_callbacks:
if (
failure_callback
in litellm._known_custom_logger_compatible_callbacks
):
_add_custom_logger_callback_to_specific_event(
failure_callback, "failure"
)
elif failure_callback not in litellm.failure_callback:
litellm.logging_callback_manager.add_litellm_failure_callback(
failure_callback
)
self._add_callback_from_db_to_in_memory_litellm_callbacks(
callback=failure_callback,
event_types=["failure"],
existing_callbacks=litellm.failure_callback,
)
if callbacks is not None and isinstance(callbacks, list):
for callback in callbacks:
self._add_callback_from_db_to_in_memory_litellm_callbacks(
callback=callback,
event_types=["success", "failure"],
existing_callbacks=litellm.callbacks,
)
def _encrypt_env_variables(
self, environment_variables: dict, new_encryption_key: Optional[str] = None
@ -4935,13 +4977,49 @@ async def audio_transcriptions(
)
######################################################################
# Vertex AI Live API WebSocket Pass-through
######################################################################
@app.websocket("/vertex_ai/live")
async def vertex_ai_live_passthrough_endpoint(
websocket: WebSocket,
model: Optional[str] = fastapi.Query(
None,
description="Optional model name, used to determine Vertex region for global models.",
),
vertex_project: Optional[str] = fastapi.Query(
None,
description="Override the Vertex AI project id used for the upstream connection.",
),
vertex_location: Optional[str] = fastapi.Query(
None,
description="Override the Vertex AI region (for example, 'us-central1').",
),
user_api_key_dict=Depends(user_api_key_auth_websocket),
):
"""
Vertex AI Live API WebSocket Pass-through Endpoint
This endpoint delegates to the WebSocket function defined in llm_passthrough_endpoints.py
"""
return await vertex_ai_live_websocket_passthrough(
websocket=websocket,
model=model,
vertex_project=vertex_project,
vertex_location=vertex_location,
user_api_key_dict=user_api_key_dict,
)
######################################################################
# /v1/realtime Endpoints
######################################################################
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from litellm import _arealtime

View file

@ -130,7 +130,7 @@ async def route_request(
elif (
data["model"] in router_model_names
or data["model"] in llm_router.get_model_ids()
or llm_router.has_model_id(data["model"])
):
return getattr(llm_router, f"{route_type}")(**data)

View file

@ -539,13 +539,15 @@ async def update_sso_settings(sso_config: SSOConfig):
@router.get(
"/get/ui_theme_settings",
tags=["UI Theme Settings"],
dependencies=[Depends(user_api_key_auth)],
response_model=UIThemeSettingsResponse,
)
async def get_ui_theme_settings():
"""
Get UI theme configuration from the litellm_settings.
Returns current logo settings for UI customization.
Note: This endpoint is public (no authentication required) so all users can see custom branding.
Only the /update/ui_theme_settings endpoint requires authentication for admins to change settings.
"""
from litellm.proxy.proxy_server import proxy_config

View file

@ -49,6 +49,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
self.litellm_metadata: Optional[dict] = litellm_metadata or {}
self.collected_chat_completion_chunks: List[ModelResponseStream] = []
self.finished: bool = False
self.litellm_logging_obj = litellm_custom_stream_wrapper.logging_obj
async def __anext__(
self,
@ -167,8 +168,16 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
def _emit_response_completed_event(self) -> Optional[ResponseCompletedEvent]:
litellm_model_response: Optional[
Union[ModelResponse, TextCompletionResponse]
] = stream_chunk_builder(chunks=self.collected_chat_completion_chunks)
] = stream_chunk_builder(chunks=self.collected_chat_completion_chunks, logging_obj=self.litellm_logging_obj)
if litellm_model_response and isinstance(litellm_model_response, ModelResponse):
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and self.litellm_logging_obj is not None:
usage = getattr(litellm_model_response, "usage", None)
if usage is not None:
setattr(
usage, "cost", self.litellm_logging_obj._response_cost_calculator(result=litellm_model_response)
)
# Transform the response
responses_api_response = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
request_input=self.request_input,

View file

@ -851,8 +851,15 @@ class LiteLLMCompletionResponsesConfig:
output_tokens=0,
total_tokens=0,
)
return ResponseAPIUsage(
response_usage = ResponseAPIUsage(
input_tokens=usage.prompt_tokens,
output_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
)
# Preserve cost field if it exists (for streaming usage with cost calculation)
if hasattr(usage, "cost") and usage.cost is not None:
setattr(response_usage, "cost", usage.cost)
return response_usage

View file

@ -5,6 +5,7 @@ from typing import Any, Dict, Optional
import httpx
import litellm
from litellm.constants import STREAM_SSE_DONE_STRING
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -13,6 +14,7 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
OutputTextDeltaEvent,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
@ -95,6 +97,20 @@ class BaseResponsesAPIStreamingIterator:
== ResponsesAPIStreamEvents.RESPONSE_COMPLETED
):
self.completed_response = openai_responses_api_chunk
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and self.logging_obj is not None:
response_obj: Optional[ResponsesAPIResponse] = getattr(openai_responses_api_chunk, "response", None)
if response_obj:
usage_obj: Optional[ResponseAPIUsage] = getattr(response_obj, "usage", None)
if usage_obj is not None:
try:
cost: Optional[float] = self.logging_obj._response_cost_calculator(result=response_obj)
if cost is not None:
setattr(usage_obj, "cost", cost)
except Exception:
# If cost calculation fails, continue without cost
pass
self._handle_logging_completed_response()
return openai_responses_api_chunk

View file

@ -337,13 +337,13 @@ class Router:
```
"""
from litellm._service_logger import ServiceLogging
self.set_verbose = set_verbose
self.ignore_invalid_deployments = ignore_invalid_deployments
self.debug_level = debug_level
self.enable_pre_call_checks = enable_pre_call_checks
self.enable_tag_filtering = enable_tag_filtering
from litellm._service_logger import ServiceLogging
self.service_logger_obj: ServiceLogging = ServiceLogging()
litellm.suppress_debug_info = True # prevents 'Give Feedback/Get help' message from being emitted on Router - Relevant Issue: https://github.com/BerriAI/litellm/issues/5942
if self.set_verbose is True:
if debug_level == "INFO":
@ -360,9 +360,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] = {}
@ -404,18 +404,22 @@ 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
self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = (
model_group_alias or {}
) # dict to store aliases for router, ex. {"gpt-4": "gpt-3.5-turbo"}, all requests with gpt-4 -> get routed to gpt-3.5-turbo group
# Initialize model ID to deployment index mapping for O(1) lookups
self.model_id_to_deployment_index_map: Dict[str, int] = {}
if model_list is not None:
# Build model index immediately to enable O(1) lookups from the start
self._build_model_id_to_deployment_index_map(model_list)
model_list = copy.deepcopy(model_list)
self.set_model_list(model_list)
self.healthy_deployments: List = self.model_list # type: ignore
for m in model_list:
@ -495,9 +499,6 @@ class Router:
self.previous_models: List = (
[]
) # list to store failed calls (passed in as metadata to next call)
self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = (
model_group_alias or {}
) # dict to store aliases for router, ex. {"gpt-4": "gpt-3.5-turbo"}, all requests with gpt-4 -> get routed to gpt-3.5-turbo group
# make Router.chat.completions.create compatible for openai.chat.completions.create
default_litellm_params = default_litellm_params or {}
@ -562,15 +563,6 @@ class Router:
)
else:
litellm.failure_callback = [self.deployment_callback_on_failure]
verbose_router_logger.debug(
f"Intialized router with Routing strategy: {self.routing_strategy}\n\n"
f"Routing enable_pre_call_checks: {self.enable_pre_call_checks}\n\n"
f"Routing fallbacks: {self.fallbacks}\n\n"
f"Routing content fallbacks: {self.content_policy_fallbacks}\n\n"
f"Routing context window fallbacks: {self.context_window_fallbacks}\n\n"
f"Router Redis Caching={self.cache.redis_cache}\n"
)
self.service_logger_obj = ServiceLogging()
self.routing_strategy_args = routing_strategy_args
self.provider_budget_config = provider_budget_config
self.router_budget_logger: Optional[RouterBudgetLimiting] = None
@ -593,9 +585,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:
@ -700,7 +692,7 @@ class Router:
or routing_strategy == RoutingStrategy.LEAST_BUSY
):
self.leastbusy_logger = LeastBusyLoggingHandler(
router_cache=self.cache, model_list=self.model_list
router_cache=self.cache
)
## add callback
if isinstance(litellm.input_callback, list):
@ -715,7 +707,6 @@ class Router:
):
self.lowesttpm_logger = LowestTPMLoggingHandler(
router_cache=self.cache,
model_list=self.model_list,
routing_args=routing_strategy_args,
)
if isinstance(litellm.callbacks, list):
@ -726,7 +717,6 @@ class Router:
):
self.lowesttpm_logger_v2 = LowestTPMLoggingHandler_v2(
router_cache=self.cache,
model_list=self.model_list,
routing_args=routing_strategy_args,
)
if isinstance(litellm.callbacks, list):
@ -737,7 +727,6 @@ class Router:
):
self.lowestlatency_logger = LowestLatencyLoggingHandler(
router_cache=self.cache,
model_list=self.model_list,
routing_args=routing_strategy_args,
)
if isinstance(litellm.callbacks, list):
@ -748,7 +737,6 @@ class Router:
):
self.lowestcost_logger = LowestCostLoggingHandler(
router_cache=self.cache,
model_list=self.model_list,
routing_args={},
)
if isinstance(litellm.callbacks, list):
@ -774,6 +762,14 @@ class Router:
self.aanthropic_messages = self.factory_function(
litellm.anthropic_messages, call_type="anthropic_messages"
)
self.agenerate_content = self.factory_function(
litellm.agenerate_content, call_type="agenerate_content"
)
self.aadapter_generate_content = self.factory_function(
litellm.aadapter_generate_content, call_type="aadapter_generate_content"
)
self.aresponses = self.factory_function(
litellm.aresponses, call_type="aresponses"
)
@ -972,7 +968,7 @@ class Router:
### DEPLOYMENT-SPECIFIC PRE-CALL CHECKS ### (e.g. update rpm pre-call. Raise error, if deployment over limit)
## only run if model group given, not model id
if model not in self.get_model_ids():
if not self.has_model_id(model):
self.routing_strategy_pre_call_checks(deployment=deployment)
response = litellm.completion(
@ -1217,10 +1213,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
@ -3177,9 +3170,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(
@ -3742,11 +3735,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
@ -3778,11 +3771,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
@ -4500,16 +4493,17 @@ class Router:
try:
exception = kwargs.get("exception", None)
exception_status = getattr(exception, "status_code", "")
_model_info = kwargs.get("litellm_params", {}).get("model_info", {})
# Cache litellm_params to avoid repeated dict lookups
litellm_params = kwargs.get("litellm_params", {})
_model_info = litellm_params.get("model_info", {})
exception_headers = litellm.litellm_core_utils.exception_mapping_utils._get_response_headers(
original_exception=exception
)
# Determine cooldown time with priority: deployment config > response header > router default
deployment_cooldown = kwargs.get("litellm_params", {}).get(
"cooldown_time", None
)
deployment_cooldown = litellm_params.get("cooldown_time", None)
header_cooldown = None
if exception_headers is not None:
@ -4989,7 +4983,9 @@ class Router:
model = deployment.to_json(exclude_none=True)
self._add_model_to_list_and_index_map(model=model, model_id=deployment.model_info.id)
self._add_model_to_list_and_index_map(
model=model, model_id=deployment.model_info.id
)
return deployment
except Exception as e:
if self.ignore_invalid_deployments:
@ -5018,26 +5014,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"
@ -5331,7 +5327,8 @@ class Router:
"""
# check if deployment already exists
if deployment.model_info.id in self.get_model_ids():
_deployment_model_id = deployment.model_info.id
if _deployment_model_id and self.has_model_id(_deployment_model_id):
return None
# add to model list
@ -5340,14 +5337,18 @@ class Router:
self._add_deployment(deployment=deployment)
# add to model names
self._add_model_to_list_and_index_map(model=_deployment, model_id=deployment.model_info.id)
self._add_model_to_list_and_index_map(
model=_deployment, model_id=deployment.model_info.id
)
self.model_names.append(deployment.model_name)
return deployment
def _update_deployment_indices_after_removal(self, model_id: str, removal_idx: int) -> None:
def _update_deployment_indices_after_removal(
self, model_id: str, removal_idx: int
) -> None:
"""
Helper method to update deployment indices after a deployment has been removed from model_list.
Parameters:
- model_id: str - the id of the deployment that was removed
- removal_idx: int - the index where the deployment was removed from model_list
@ -5360,11 +5361,12 @@ class Router:
if model_id in self.model_id_to_deployment_index_map:
del self.model_id_to_deployment_index_map[model_id]
def _add_model_to_list_and_index_map(self, model: dict, model_id: Optional[str] = None) -> None:
def _add_model_to_list_and_index_map(
self, model: dict, model_id: Optional[str] = None
) -> None:
"""
Helper method to add a model to the model_list and update the model_id_to_deployment_index_map.
Parameters:
- model: dict - the model to add to the list
- model_id: Optional[str] - the model ID to use for indexing. If None, will try to get from model["model_info"]["id"]
@ -5374,7 +5376,9 @@ class Router:
if model_id is not None:
self.model_id_to_deployment_index_map[model_id] = len(self.model_list) - 1
elif model.get("model_info", {}).get("id") is not None:
self.model_id_to_deployment_index_map[model["model_info"]["id"]] = len(self.model_list) - 1
self.model_id_to_deployment_index_map[model["model_info"]["id"]] = (
len(self.model_list) - 1
)
def upsert_deployment(self, deployment: Deployment) -> Optional[Deployment]:
"""
@ -5403,13 +5407,15 @@ class Router:
removal_idx: Optional[int] = None
deployment_id = deployment.model_info.id
deployment_fast_mapping = self.model_id_to_deployment_index_map
if deployment_id in deployment_fast_mapping:
removal_idx = deployment_fast_mapping[deployment_id]
if removal_idx is not None:
self.model_list.pop(removal_idx)
self._update_deployment_indices_after_removal(model_id=deployment_id, removal_idx=removal_idx)
self._update_deployment_indices_after_removal(
model_id=deployment_id, removal_idx=removal_idx
)
# if the model_id is not in router
self.add_deployment(deployment=deployment)
@ -5440,7 +5446,9 @@ class Router:
if deployment_idx is not None:
# Pop the item from the list first
item = self.model_list.pop(deployment_idx)
self._update_deployment_indices_after_removal(model_id=id, removal_idx=deployment_idx)
self._update_deployment_indices_after_removal(
model_id=id, removal_idx=deployment_idx
)
return item
else:
return None
@ -5463,7 +5471,7 @@ class Router:
return model
else:
raise Exception("Model invalid format - {}".format(type(model)))
return None
def get_deployment_credentials(self, model_id: str) -> Optional[dict]:
@ -5709,27 +5717,32 @@ class Router:
configurable_clientside_auth_params = (
litellm_params.configurable_clientside_auth_params
)
# Cache nested dict access to avoid repeated temporary dict allocations
model_litellm_params = model.get("litellm_params", {})
model_info_dict = model.get("model_info", {})
# get model tpm
_deployment_tpm: Optional[int] = None
if _deployment_tpm is None:
_deployment_tpm = model.get("tpm", None) # type: ignore
if _deployment_tpm is None:
_deployment_tpm = model.get("litellm_params", {}).get("tpm", None) # type: ignore
_deployment_tpm = model_litellm_params.get("tpm", None) # type: ignore
if _deployment_tpm is None:
_deployment_tpm = model.get("model_info", {}).get("tpm", None) # type: ignore
_deployment_tpm = model_info_dict.get("tpm", None) # type: ignore
# get model rpm
_deployment_rpm: Optional[int] = None
if _deployment_rpm is None:
_deployment_rpm = model.get("rpm", None) # type: ignore
if _deployment_rpm is None:
_deployment_rpm = model.get("litellm_params", {}).get("rpm", None) # type: ignore
_deployment_rpm = model_litellm_params.get("rpm", None) # type: ignore
if _deployment_rpm is None:
_deployment_rpm = model.get("model_info", {}).get("rpm", None) # type: ignore
_deployment_rpm = model_info_dict.get("rpm", None) # type: ignore
# get model info
try:
model_id = model.get("model_info", {}).get("id", None)
model_id = model_info_dict.get("id", None)
if model_id is not None:
model_info = self.get_deployment_model_info(
model_id=model_id, model_name=litellm_params.model
@ -6093,7 +6106,7 @@ class Router:
# Extract model_info from the model dict
model_info = model.get("model_info", {})
model_id = model_info.get("id")
# If no ID exists, generate one using the same logic as set_model_list
if model_id is None:
model_name = model.get("model_name", "")
@ -6103,7 +6116,7 @@ class Router:
if "model_info" not in model:
model["model_info"] = {}
model["model_info"]["id"] = model_id
self._add_model_to_list_and_index_map(model=model, model_id=model_id)
def get_model_ids(
@ -6113,7 +6126,7 @@ class Router:
if 'model_name' is none, returns all.
Returns list of model id's.
"""
"""
ids = []
for model in self.model_list:
if "model_info" in model and "id" in model["model_info"]:
@ -6126,6 +6139,19 @@ class Router:
ids.append(id)
return ids
def has_model_id(self, candidate_id: str) -> bool:
"""
O(1) membership check for a deployment ID without allocating large lists.
Note: Call sites may pass a variable named `model` when it actually
contains a deployment ID. This helper expects the deployment ID string.
Uses the existing `model_id_to_deployment_index_map` which is kept
in sync by `_build_model_id_to_deployment_index_map` and model-list
mutation helpers.
"""
return candidate_id in self.model_id_to_deployment_index_map
def map_team_model(self, team_model_name: str, team_id: str) -> Optional[str]:
"""
Map a team model name to a team-specific model name.
@ -6282,45 +6308,41 @@ class Router:
if team_id specified, returns matching team-specific models
"""
# Note: model_list and model_group_alias are always initialized in __init__
# so hasattr checks are unnecessary
returned_models: List[DeploymentTypedDict] = []
if hasattr(self, "model_list"):
returned_models: List[DeploymentTypedDict] = []
if model_name is not None:
returned_models.extend(
self._get_all_deployments(model_name=model_name, team_id=team_id)
)
if model_name is not None:
returned_models.extend(
self._get_all_deployments(model_name=model_name, team_id=team_id)
returned_models.extend(
self.get_model_list_from_model_alias(model_name=model_name)
)
if len(returned_models) == 0: # check if wildcard route
potential_wildcard_models = self.pattern_router.route(model_name) or []
## check for team-specific wildcard models
if team_id is not None and team_id in self.team_pattern_routers:
potential_team_only_wildcard_models = (
self.team_pattern_routers[team_id].route(model_name) or []
)
potential_wildcard_models.extend(
potential_team_only_wildcard_models
)
if hasattr(self, "model_group_alias"):
returned_models.extend(
self.get_model_list_from_model_alias(model_name=model_name)
)
if model_name is not None and potential_wildcard_models is not None:
for m in potential_wildcard_models:
deployment_typed_dict = DeploymentTypedDict(**m) # type: ignore
deployment_typed_dict["model_name"] = model_name
returned_models.append(deployment_typed_dict)
if len(returned_models) == 0: # check if wildcard route
potential_wildcard_models = self.pattern_router.route(model_name) or []
if model_name is None:
returned_models += self.model_list
## check for team-specific wildcard models
if team_id is not None and team_id in self.team_pattern_routers:
potential_team_only_wildcard_models = (
self.team_pattern_routers[team_id].route(model_name) or []
)
potential_wildcard_models.extend(
potential_team_only_wildcard_models
)
if model_name is not None and potential_wildcard_models is not None:
for m in potential_wildcard_models:
deployment_typed_dict = DeploymentTypedDict(**m) # type: ignore
deployment_typed_dict["model_name"] = model_name
returned_models.append(deployment_typed_dict)
if model_name is None:
returned_models += self.model_list
return returned_models
return returned_models
return None
return returned_models
def get_model_access_groups(
self,
@ -6567,19 +6589,19 @@ class Router:
or {}
) # check the in-memory cache used by lowest_latency and usage-based routing. Only check the local cache.
for idx, deployment in enumerate(_returned_deployments):
# Cache nested dict access to avoid repeated temporary dict allocations
_litellm_params = deployment.get("litellm_params", {})
_model_info = deployment.get("model_info", {})
# see if we have the info for this model
try:
base_model = deployment.get("model_info", {}).get("base_model", None)
base_model = _model_info.get("base_model", None)
if base_model is None:
base_model = deployment.get("litellm_params", {}).get(
"base_model", None
)
base_model = _litellm_params.get("base_model", None)
model_info = self.get_router_model_info(
deployment=deployment, received_model_name=model
)
model = base_model or deployment.get("litellm_params", {}).get(
"model", None
)
model = base_model or _litellm_params.get("model", None)
if (
isinstance(model_info, dict)
@ -6600,8 +6622,7 @@ class Router:
except Exception as e:
verbose_router_logger.exception("An error occurs - {}".format(str(e)))
_litellm_params = deployment.get("litellm_params", {})
model_id = deployment.get("model_info", {}).get("id", "")
model_id = _model_info.get("id", "")
## RPM CHECK ##
### get local router cache ###
current_request_cache_local = (
@ -6762,14 +6783,13 @@ class Router:
# check if aliases set on litellm model alias map
if specific_deployment is True:
return model, self._get_deployment_by_litellm_model(model=model)
elif model in self.get_model_ids():
elif self.has_model_id(model):
deployment = self.get_deployment(model_id=model)
if deployment is not None:
deployment_model = deployment.litellm_params.model
return deployment_model, deployment.model_dump(exclude_none=True)
raise ValueError(
f"LiteLLM Router: Trying to call specific deployment, but Model ID :{model} does not exist in \
Model ID List: {self.get_model_ids}"
f"LiteLLM Router: Trying to call specific deployment, but Model ID :{model} does not exist in Model ID map"
)
_model_from_alias = self._get_model_from_alias(model=model)
@ -7248,19 +7268,13 @@ class Router:
Returns:
List of healthy deployments
"""
# filter out the deployments currently cooling down
deployments_to_remove = []
verbose_router_logger.debug(f"cooldown deployments: {cooldown_deployments}")
# Find deployments in model_list whose model_id is cooling down
for deployment in healthy_deployments:
deployment_id = deployment["model_info"]["id"]
if deployment_id in cooldown_deployments:
deployments_to_remove.append(deployment)
# remove unhealthy deployments from healthy deployments
for deployment in deployments_to_remove:
healthy_deployments.remove(deployment)
return healthy_deployments
# Convert to set for O(1) lookup and use list comprehension for O(n) filtering
cooldown_set = set(cooldown_deployments)
return [
deployment for deployment in healthy_deployments
if deployment["model_info"]["id"] not in cooldown_set
]
def _track_deployment_metrics(
self, deployment, parent_otel_span: Optional[Span], response=None

View file

@ -18,10 +18,9 @@ class LeastBusyLoggingHandler(CustomLogger):
logged_success: int = 0
logged_failure: int = 0
def __init__(self, router_cache: DualCache, model_list: list):
def __init__(self, router_cache: DualCache):
self.router_cache = router_cache
self.mapping_deployment_to_id: dict = {}
self.model_list = model_list
def log_pre_api_call(self, model, messages, kwargs):
"""

View file

@ -16,10 +16,9 @@ class LowestCostLoggingHandler(CustomLogger):
logged_failure: int = 0
def __init__(
self, router_cache: DualCache, model_list: list, routing_args: dict = {}
self, router_cache: DualCache, routing_args: dict = {}
):
self.router_cache = router_cache
self.model_list = model_list
def log_success_event(self, kwargs, response_obj, start_time, end_time):
try:

View file

@ -32,10 +32,9 @@ class LowestLatencyLoggingHandler(CustomLogger):
logged_failure: int = 0
def __init__(
self, router_cache: DualCache, model_list: list, routing_args: dict = {}
self, router_cache: DualCache, routing_args: dict = {}
):
self.router_cache = router_cache
self.model_list = model_list
self.routing_args = RoutingArgs(**routing_args)
def log_success_event( # noqa: PLR0915

View file

@ -23,10 +23,9 @@ class LowestTPMLoggingHandler(CustomLogger):
default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour
def __init__(
self, router_cache: DualCache, model_list: list, routing_args: dict = {}
self, router_cache: DualCache, routing_args: dict = {}
):
self.router_cache = router_cache
self.model_list = model_list
self.routing_args = RoutingArgs(**routing_args)
def log_success_event(self, kwargs, response_obj, start_time, end_time):

View file

@ -48,10 +48,9 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour
def __init__(
self, router_cache: DualCache, model_list: list, routing_args: dict = {}
self, router_cache: DualCache, routing_args: dict = {}
):
self.router_cache = router_cache
self.model_list = model_list
self.routing_args = RoutingArgs(**routing_args)
BaseRoutingStrategy.__init__(
self,

View file

@ -1,28 +1,58 @@
# Import types from the Google GenAI SDK
from typing import TYPE_CHECKING, Any, List, Optional, TypeAlias
from typing import TYPE_CHECKING, Any, Dict, List, Optional, TypeAlias
# During static type-checking we can rely on the real google-genai types.
from google.genai import types as _genai_types # type: ignore
from pydantic import BaseModel
from typing_extensions import TypedDict
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
ContentListUnion = _genai_types.ContentListUnion
ContentListUnionDict = _genai_types.ContentListUnionDict
GenerateContentConfigOrDict = _genai_types.GenerateContentConfigOrDict
GoogleGenAIGenerateContentResponse = _genai_types.GenerateContentResponse
# During static type-checking we can rely on the real google-genai types.
if TYPE_CHECKING:
from google.genai import types as _genai_types # type: ignore
GenerateContentContentListUnionDict = _genai_types.ContentListUnionDict
GenerateContentConfigDict = _genai_types.GenerateContentConfigDict
GenerateContentRequestParametersDict = _genai_types._GenerateContentParametersDict
ToolConfigDict = _genai_types.ToolConfigDict
ContentListUnion = _genai_types.ContentListUnion
ContentListUnionDict = _genai_types.ContentListUnionDict
GenerateContentConfigOrDict = _genai_types.GenerateContentConfigOrDict
GoogleGenAIGenerateContentResponse = _genai_types.GenerateContentResponse
GenerateContentContentListUnionDict = _genai_types.ContentListUnionDict
GenerateContentConfigDict = _genai_types.GenerateContentConfigDict
GenerateContentRequestParametersDict = _genai_types._GenerateContentParametersDict
ToolConfigDict = _genai_types.ToolConfigDict
class GenerateContentRequestDict(GenerateContentRequestParametersDict): # type: ignore[misc]
generationConfig: Optional[Any]
tools: Optional[ToolConfigDict] # type: ignore[assignment]
class GenerateContentRequestDict(GenerateContentRequestParametersDict): # type: ignore[misc]
generationConfig: Optional[Any]
tools: Optional[ToolConfigDict] # type: ignore[assignment]
class GenerateContentResponse(GoogleGenAIGenerateContentResponse, BaseLiteLLMOpenAIResponseObject): # type: ignore[misc]
_hidden_params: dict = {}
pass
else:
# Fallback types when google.genai is not available
ContentListUnion = Any
ContentListUnionDict = Dict[str, Any]
GenerateContentConfigOrDict = Dict[str, Any]
GoogleGenAIGenerateContentResponse = Dict[str, Any]
GenerateContentContentListUnionDict = Dict[str, Any]
class GenerateContentResponse(GoogleGenAIGenerateContentResponse, BaseLiteLLMOpenAIResponseObject): # type: ignore[misc]
_hidden_params: dict = {}
pass
# Create a proper fallback class that can be instantiated
class GenerateContentConfigDict(dict): # type: ignore[misc]
def __init__(self, **kwargs): # type: ignore
super().__init__(**kwargs)
class GenerateContentRequestParametersDict(dict): # type: ignore[misc]
def __init__(self, **kwargs): # type: ignore
super().__init__(**kwargs)
ToolConfigDict = Dict[str, Any]
class GenerateContentRequestDict(GenerateContentRequestParametersDict): # type: ignore[misc]
def __init__(self, **kwargs): # type: ignore
# Extract specific fields
self.generationConfig = kwargs.get('generationConfig')
self.tools = kwargs.get('tools')
super().__init__(**kwargs)
class GenerateContentResponse(BaseLiteLLMOpenAIResponseObject): # type: ignore[misc]
def __init__(self, **kwargs): # type: ignore
super().__init__(**kwargs)
self._hidden_params = kwargs.get('_hidden_params', {})

View file

@ -12,5 +12,6 @@ class LangfuseLoggingConfig(TypedDict):
class LangfuseUsageDetails(TypedDict):
input: Optional[int]
output: Optional[int]
total: Optional[int]
cache_creation_input_tokens: Optional[int]
cache_read_input_tokens: Optional[int]

View file

@ -1033,6 +1033,9 @@ class ResponseAPIUsage(BaseLiteLLMOpenAIResponseObject):
total_tokens: int
"""The total number of tokens used."""
cost: Optional[float] = None
"""The cost of the request."""
model_config = {"extra": "allow"}

View file

@ -10,6 +10,7 @@ class SupportedPromptIntegrations(str, Enum):
LANGFUSE = "langfuse"
CUSTOM = "custom"
BITBUCKET = "bitbucket"
GITLAB = "gitlab"
class PromptInfo(BaseModel):

View file

@ -2031,6 +2031,13 @@ class GuardrailMode(TypedDict, total=False):
default: Optional[str]
GuardrailStatus = Literal[
"success",
"guardrail_intervened",
"guardrail_failed_to_respond",
"not_run"
]
class StandardLoggingGuardrailInformation(TypedDict, total=False):
guardrail_name: Optional[str]
guardrail_provider: Optional[str]
@ -2039,7 +2046,7 @@ class StandardLoggingGuardrailInformation(TypedDict, total=False):
]
guardrail_request: Optional[dict]
guardrail_response: Optional[Union[dict, str, List[dict]]]
guardrail_status: Literal["success", "failure", "blocked"]
guardrail_status: GuardrailStatus
start_time: Optional[float]
end_time: Optional[float]
duration: Optional[float]
@ -2059,6 +2066,18 @@ class StandardLoggingGuardrailInformation(TypedDict, total=False):
StandardLoggingPayloadStatus = Literal["success", "failure"]
class CachingDetails(TypedDict):
"""
Track all caching related metrics, fields for a given request
"""
cache_hit: Optional[bool]
"""
Whether the request hit the cache
"""
cache_duration_ms: Optional[float]
"""
Duration for reading from cache
"""
class CostBreakdown(TypedDict):
"""
@ -2070,6 +2089,20 @@ class CostBreakdown(TypedDict):
tool_usage_cost: float # Cost of usage of built-in tools
class StandardLoggingPayloadStatusFields(TypedDict, total=False):
"""Status fields for easy filtering and analytics"""
llm_api_status: StandardLoggingPayloadStatus
"""Status of the LLM API call - 'success' if completed, 'failure' if errored"""
guardrail_status: GuardrailStatus
"""
Status of guardrail execution:
- 'success': Guardrail ran and allowed content through
- 'guardrail_intervened': Guardrail blocked or modified content
- 'guardrail_failed_to_respond': Guardrail had technical failure
- 'not_run': No guardrail was run
"""
class StandardLoggingPayload(TypedDict):
id: str
trace_id: str # Trace multiple LLM calls belonging to same overall request (e.g. fallbacks/retries)
@ -2081,6 +2114,7 @@ class StandardLoggingPayload(TypedDict):
StandardLoggingModelCostFailureDebugInformation
]
status: StandardLoggingPayloadStatus
status_fields: StandardLoggingPayloadStatusFields
custom_llm_provider: Optional[str]
total_tokens: int
prompt_tokens: int
@ -2416,6 +2450,7 @@ class LlmProviders(str, Enum):
DOTPROMPT = "dotprompt"
WANDB = "wandb"
OVHCLOUD = "ovhcloud"
LEMONADE = "lemonade"
# Create a set of all provider values for quick lookup
@ -2657,5 +2692,15 @@ class PriorityReservationSettings(BaseModel):
default=0.5,
description="Priority level to assign to API keys without explicit priority metadata. Should match a key in litellm.priority_reservation."
)
saturation_threshold: float = Field(
default=0.80,
description="Saturation threshold (0.0-1.0) at which strict priority enforcement begins. Below this threshold, generous mode allows priority borrowing. Above this threshold, strict mode enforces normalized priority limits."
)
tracking_multiplier: int = Field(
default=10,
description="Multiplier for model-wide tracking limits in strict mode. Set to 10x because v3_limiter.should_rate_limit() both increments counters AND enforces limits - we need the counter increment (for saturation checks) but not the enforcement (priority limits handle that). High multiplier ensures tracking never blocks."
)
model_config = ConfigDict(protected_namespaces=())

View file

@ -7,7 +7,6 @@
#
# Thank you users! We ❤️ you! - Krrish & Ishaan
from io import StringIO
import ast
import asyncio
import base64
@ -37,6 +36,7 @@ from dataclasses import dataclass, field
from functools import lru_cache, wraps
from importlib import resources
from inspect import iscoroutine
from io import StringIO
from os.path import abspath, dirname, join
import aiohttp
@ -90,6 +90,7 @@ from litellm.litellm_core_utils.cached_imports import (
get_set_callbacks,
)
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
map_finish_reason,
process_response_headers,
)
@ -232,6 +233,9 @@ from typing import (
from openai import OpenAIError as OriginalError
from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
update_response_metadata,
)
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new
from litellm.llms.base_llm.anthropic_messages.transformation import (
@ -1677,30 +1681,6 @@ def _is_streaming_request(
return False
def update_response_metadata(
result: Any,
logging_obj: LiteLLMLoggingObject,
model: Optional[str],
kwargs: dict,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
"""
Updates response metadata, adds the following:
- response._hidden_params
- response._hidden_params["litellm_overhead_time_ms"]
- response.response_time_ms
"""
if result is None:
return
metadata = ResponseMetadata(result)
metadata.set_hidden_params(logging_obj=logging_obj, model=model, kwargs=kwargs)
metadata.set_timing_metrics(
start_time=start_time, end_time=end_time, logging_obj=logging_obj
)
metadata.apply()
def _select_tokenizer(
model: str, custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None
@ -7338,6 +7318,8 @@ class ProviderConfigManager:
)
return VLLMModelInfo()
elif LlmProviders.LEMONADE == provider:
return litellm.LemonadeChatConfig()
return None
@staticmethod
@ -7601,7 +7583,7 @@ def get_end_user_id_for_cost_tracking(
service_type: "litellm_logging" or "prometheus" - used to allow prometheus only disable cost tracking.
"""
_metadata = cast(dict, litellm_params.get("metadata", {}) or {})
_metadata = cast(dict, get_litellm_metadata_from_kwargs(dict(litellm_params=litellm_params)))
end_user_id = cast(
Optional[str],

View file

@ -2004,9 +2004,9 @@
"cache_read_input_token_cost": 1.25e-07,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supported_endpoints": [
@ -4797,6 +4797,66 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
"claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "anthropic",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"claude-sonnet-4-5-20250929": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "anthropic",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"claude-opus-4-1": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
@ -9454,96 +9514,6 @@
"supports_vision": true,
"supports_web_search": true
},
"gemini-flash-latest": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 3e-07,
"litellm_provider": "vertex_ai-language-models",
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_images_per_prompt": 3000,
"max_input_tokens": 1048576,
"max_output_tokens": 65535,
"max_pdf_size_mb": 30,
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_reasoning_token": 2.5e-06,
"output_cost_per_token": 2.5e-06,
"source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/completions",
"/v1/batch"
],
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_audio_output": false,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": true
},
"gemini-flash-lite-latest": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_audio_token": 3e-07,
"input_cost_per_token": 1e-07,
"litellm_provider": "vertex_ai-language-models",
"max_audio_length_hours": 8.4,
"max_audio_per_prompt": 1,
"max_images_per_prompt": 3000,
"max_input_tokens": 1048576,
"max_output_tokens": 65535,
"max_pdf_size_mb": 30,
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_reasoning_token": 4e-07,
"output_cost_per_token": 4e-07,
"source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/completions",
"/v1/batch"
],
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_audio_output": false,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": true
},
"gemini-2.5-flash-lite-preview-06-17": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_audio_token": 5e-07,
@ -12823,6 +12793,34 @@
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5-codex": {
"cache_read_input_token_cost": 1.25e-07,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 400000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-5-2025-08-07": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_flex": 6.25e-08,
@ -12898,9 +12896,9 @@
"cache_read_input_token_cost": 1.25e-07,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"max_input_tokens": 128000,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supported_endpoints": [
@ -13352,6 +13350,18 @@
],
"supports_tool_choice": false
},
"lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": {
"input_cost_per_token": 0,
"litellm_provider": "lemonade",
"max_tokens": 32768,
"max_input_tokens": 32768,
"max_output_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"groq/deepseek-r1-distill-llama-70b": {
"input_cost_per_token": 7.5e-07,
"litellm_provider": "groq",
@ -13641,6 +13651,19 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"groq/moonshotai/kimi-k2-instruct-0905": {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 3e-06,
"cache_read_input_token_cost": 0.5e-06,
"litellm_provider": "groq",
"max_input_tokens": 262144,
"max_output_tokens": 16384,
"max_tokens": 278528,
"mode": "chat",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"groq/openai/gpt-oss-120b": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "groq",
@ -19701,6 +19724,36 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 159
},
"us.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"us.anthropic.claude-opus-4-20250514-v1:0": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,
@ -21041,6 +21094,58 @@
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token_batches": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-sonnet-4-5@20250929": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token_batches": 1.5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 200000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-opus-4@20250514": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_read_input_token_cost": 1.5e-06,

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.77.6"
version = "1.77.7"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -157,7 +157,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.77.6"
version = "1.77.7"
version_files = [
"pyproject.toml:^version"
]

View file

@ -126,8 +126,13 @@ class SensitiveLogDetector(ast.NodeVisitor):
for value in arg.values:
if isinstance(value, ast.FormattedValue):
value_str = self._get_arg_string(value.value).lower()
if any(pattern in value_str for pattern in
['request', 'response', 'data', 'body', 'content', 'messages']):
# Check for any sensitive data patterns in f-string interpolations
sensitive_f_string_patterns = [
'request', 'response', 'data', 'body', 'content', 'messages',
'token', 'jwt', 'auth', 'api_key', 'apikey', 'credential',
'secret', 'password', 'passwd'
]
if any(pattern in value_str for pattern in sensitive_f_string_patterns):
return True
# Check for .format() calls
@ -137,10 +142,14 @@ class SensitiveLogDetector(ast.NodeVisitor):
base_str = self._get_arg_string(arg.func.value).lower()
if "{}" in base_str or "{" in base_str:
# Check format arguments for sensitive data
sensitive_format_patterns = [
'request', 'response', 'data', 'body', 'content',
'token', 'jwt', 'auth', 'api_key', 'apikey', 'credential',
'secret', 'password', 'passwd'
]
for format_arg in arg.args:
format_str = self._get_arg_string(format_arg).lower()
if any(pattern in format_str for pattern in
['request', 'response', 'data', 'body', 'content']):
if any(pattern in format_str for pattern in sensitive_format_patterns):
return True
return False
@ -171,7 +180,9 @@ class SensitiveLogDetector(ast.NodeVisitor):
"""Get a human-readable reason for the violation"""
arg_str = self._get_arg_string(arg).lower()
if 'request' in arg_str:
if any(pattern in arg_str for pattern in ['jwt', 'token', 'api_key', 'apikey', 'auth', 'credential', 'secret', 'password', 'passwd']):
return "Potentially logging authentication/secret data (JWT, token, API key, etc.)"
elif 'request' in arg_str:
return "Potentially logging request data"
elif 'response' in arg_str:
return "Potentially logging response data"
@ -179,8 +190,6 @@ class SensitiveLogDetector(ast.NodeVisitor):
return "Potentially logging sensitive data/body/content"
elif any(pattern in arg_str for pattern in ['messages', 'input', 'output']):
return "Potentially logging message/input/output data"
elif any(pattern in arg_str for pattern in ['api_key', 'token', 'auth', 'credentials']):
return "Potentially logging authentication data"
else:
return "Potentially logging sensitive data"

View file

@ -0,0 +1,79 @@
# conftest.py
import importlib
import os
import sys
import pytest
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
import asyncio
@pytest.fixture(scope="session")
def event_loop():
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
yield loop
loop.close()
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
"""
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
"""
curr_dir = os.getcwd() # Get the current working directory
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the project directory to the system path
import litellm
from litellm import Router
import asyncio
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
# flush all logs
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
import asyncio
loop = asyncio.get_event_loop_policy().new_event_loop()
asyncio.set_event_loop(loop)
print(litellm)
# from litellm import Router, completion, aembedding, acompletion, embedding
yield
# Teardown code (executes after the yield point)
loop.close() # Close the loop created earlier
asyncio.set_event_loop(None) # Remove the reference to the loop
def pytest_collection_modifyitems(config, items):
# Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests
custom_logger_tests = [
item for item in items if "custom_logger" in item.parent.name
]
other_tests = [item for item in items if "custom_logger" not in item.parent.name]
# Sort tests based on their names
custom_logger_tests.sort(key=lambda x: x.name)
other_tests.sort(key=lambda x: x.name)
# Reorder the items list
items[:] = custom_logger_tests + other_tests

Some files were not shown because too many files have changed in this diff Show more