mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge branch 'BerriAI:main' into main
This commit is contained in:
commit
42ed6ad907
151 changed files with 12130 additions and 1712 deletions
|
|
@ -273,7 +273,7 @@ echo 'LITELLM_SALT_KEY="sk-1234"' >> .env
|
|||
source .env
|
||||
|
||||
# Start
|
||||
docker-compose up
|
||||
docker compose up
|
||||
```
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
284
docs/my-website/docs/pass_through/vertex_ai_live_websocket.md
Normal file
284
docs/my-website/docs/pass_through/vertex_ai_live_websocket.md
Normal 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)
|
||||
7
docs/my-website/docs/projects/Railtracks.md
Normal file
7
docs/my-website/docs/projects/Railtracks.md
Normal 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/)
|
||||
|
|
@ -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']` |
|
||||
|
|
|
|||
|
|
@ -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']` |
|
||||
|
|
|
|||
191
docs/my-website/docs/providers/lemonade.md
Normal file
191
docs/my-website/docs/providers/lemonade.md
Normal 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>
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ echo 'LITELLM_SALT_KEY="sk-1234"' >> .env
|
|||
source .env
|
||||
|
||||
# Start
|
||||
docker-compose up
|
||||
docker compose up
|
||||
```
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -55,7 +55,7 @@ echo 'LITELLM_SALT_KEY="sk-1234"' >> .env
|
|||
source .env
|
||||
|
||||
# Start
|
||||
docker-compose up
|
||||
docker compose up
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
)
|
||||
```
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
317
litellm/integrations/gitlab/README.md
Normal file
317
litellm/integrations/gitlab/README.md
Normal 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.
|
||||
95
litellm/integrations/gitlab/__init__.py
Normal file
95
litellm/integrations/gitlab/__init__.py
Normal 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",
|
||||
]
|
||||
285
litellm/integrations/gitlab/gitlab_client.py
Normal file
285
litellm/integrations/gitlab/gitlab_client.py
Normal 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()
|
||||
488
litellm/integrations/gitlab/gitlab_prompt_manager.py
Normal file
488
litellm/integrations/gitlab/gitlab_prompt_manager.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
149
litellm/llms/lemonade/chat/transformation.py
Normal file
149
litellm/llms/lemonade/chat/transformation.py
Normal 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
|
||||
|
||||
35
litellm/llms/lemonade/cost_calculator.py
Normal file
35
litellm/llms/lemonade/cost_calculator.py
Normal 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
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -330,6 +330,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
|
||||
anthropic_routes = [
|
||||
"/v1/messages",
|
||||
"/v1/messages/count_tokens",
|
||||
]
|
||||
|
||||
mcp_routes = [
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ##########
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
170
litellm/proxy/hooks/README.dynamic_rate_limiter_v3.md
Normal file
170
litellm/proxy/hooks/README.dynamic_rate_limiter_v3.md
Normal 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
|
||||
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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', {})
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ class SupportedPromptIntegrations(str, Enum):
|
|||
LANGFUSE = "langfuse"
|
||||
CUSTOM = "custom"
|
||||
BITBUCKET = "bitbucket"
|
||||
GITLAB = "gitlab"
|
||||
|
||||
|
||||
class PromptInfo(BaseModel):
|
||||
|
|
|
|||
|
|
@ -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=())
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
79
tests/guardrails_tests/conftest.py
Normal file
79
tests/guardrails_tests/conftest.py
Normal 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
Loading…
Add table
Reference in a new issue