mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge branch 'BerriAI:main' into kowyo/fix-ollama-think
This commit is contained in:
commit
765252d7db
121 changed files with 4883 additions and 1693 deletions
|
|
@ -616,6 +616,24 @@ jobs:
|
|||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Set DATABASE_URL environment variable
|
||||
command: |
|
||||
echo 'export DATABASE_URL="postgresql://postgres:postgres@localhost:5432/circle_test"' >> $BASH_ENV
|
||||
source $BASH_ENV
|
||||
- run:
|
||||
name: Run Security Scans
|
||||
command: |
|
||||
|
|
@ -2588,6 +2606,8 @@ jobs:
|
|||
-e GEMINI_API_KEY=$GEMINI_API_KEY \
|
||||
-e ANTHROPIC_API_KEY=$ANTHROPIC_API_KEY \
|
||||
-e ASSEMBLYAI_API_KEY=$ASSEMBLYAI_API_KEY \
|
||||
-e AZURE_API_KEY_PASSHROUGH=$AZURE_API_KEY_PASSHROUGH \
|
||||
-e AZURE_API_BASE_PASSHROUGH=$AZURE_API_BASE_PASSHROUGH \
|
||||
-e USE_DDTRACE=True \
|
||||
-e DD_API_KEY=$DD_API_KEY \
|
||||
-e DD_SITE=$DD_SITE \
|
||||
|
|
|
|||
|
|
@ -8,6 +8,10 @@ Track spend for keys, users, and teams across 100+ LLMs.
|
|||
|
||||
LiteLLM automatically tracks spend for all known models. See our [model cost map](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json)
|
||||
|
||||
:::tip Keep Pricing Data Updated
|
||||
[Sync model pricing data from GitHub](../sync_models_github.md) to ensure accurate cost tracking.
|
||||
:::
|
||||
|
||||
### How to Track Spend with LiteLLM
|
||||
|
||||
**Step 1**
|
||||
|
|
|
|||
|
|
@ -19,6 +19,10 @@ model_list:
|
|||
|
||||
Retrieve detailed information about each model listed in the `/model/info` endpoint, including descriptions from the `config.yaml` file, and additional model info (e.g. max tokens, cost per input token, etc.) pulled from the model_info you set and the [litellm model cost map](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json). Sensitive details like API keys are excluded for security purposes.
|
||||
|
||||
:::tip Sync Model Data
|
||||
Keep your model pricing data up to date by [syncing models from GitHub](../sync_models_github.md).
|
||||
:::
|
||||
|
||||
<Tabs
|
||||
defaultValue="curl"
|
||||
values={[
|
||||
|
|
|
|||
|
|
@ -221,6 +221,8 @@ litellm_settings:
|
|||
|
||||
2. Make a request with the custom metadata labels
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="Curl" label="Curl Request">
|
||||
```bash
|
||||
curl -L -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
|
||||
-H 'Content-Type: application/json' \
|
||||
|
|
@ -244,6 +246,34 @@ curl -L -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
|
|||
}
|
||||
}'
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="key" label="on Key">
|
||||
|
||||
```bash
|
||||
curl -L -X POST 'http://0.0.0.0:4000/key/generate' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"metadata": {
|
||||
"foo": "hello world"
|
||||
}
|
||||
}'
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="team" label="on Team">
|
||||
|
||||
```bash
|
||||
curl -L -X POST 'http://0.0.0.0:4000/team/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"metadata": {
|
||||
"foo": "hello world"
|
||||
}
|
||||
}'
|
||||
```
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
3. Check your `/metrics` endpoint for the custom metrics
|
||||
|
||||
|
|
|
|||
61
docs/my-website/docs/proxy/sync_models_github.md
Normal file
61
docs/my-website/docs/proxy/sync_models_github.md
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
# Syncing Models to GitHub model_context_window
|
||||
|
||||
Sync model pricing data from GitHub's `model_prices_and_context_window.json` file outside of the LiteLLM UI.
|
||||
|
||||
> **📹 Video Tutorial**: [Watch how to sync models via the Admin UI](https://www.loom.com/share/ba41acc1882d41b284bbddbb0e9c27ce?sid=bdae351e-2026-4e39-932b-fcb185ff612c)
|
||||
|
||||
## Quick Start
|
||||
|
||||
**Manual sync:**
|
||||
```bash
|
||||
curl -X POST "https://your-proxy-url/reload/model_cost_map" \
|
||||
-H "Authorization: Bearer YOUR_ADMIN_TOKEN" \
|
||||
-H "Content-Type: application/json"
|
||||
```
|
||||
|
||||
**Automatic sync every 6 hours:**
|
||||
```bash
|
||||
curl -X POST "https://your-proxy-url/schedule/model_cost_map_reload?hours=6" \
|
||||
-H "Authorization: Bearer YOUR_ADMIN_TOKEN" \
|
||||
-H "Content-Type: application/json"
|
||||
```
|
||||
|
||||
## API Endpoints
|
||||
|
||||
| Endpoint | Method | Description |
|
||||
|----------|--------|-------------|
|
||||
| `/reload/model_cost_map` | POST | Manual sync |
|
||||
| `/schedule/model_cost_map_reload?hours={hours}` | POST | Schedule periodic sync |
|
||||
| `/schedule/model_cost_map_reload` | DELETE | Cancel scheduled sync |
|
||||
| `/schedule/model_cost_map_reload/status` | GET | Check sync status |
|
||||
|
||||
**Authentication:** Requires admin role or master key
|
||||
|
||||
## Python Example
|
||||
|
||||
```python
|
||||
import requests
|
||||
|
||||
def sync_models(proxy_url, admin_token):
|
||||
response = requests.post(
|
||||
f"{proxy_url}/reload/model_cost_map",
|
||||
headers={"Authorization": f"Bearer {admin_token}"}
|
||||
)
|
||||
return response.json()
|
||||
|
||||
# Usage
|
||||
result = sync_models("https://your-proxy-url", "your-admin-token")
|
||||
print(result['message'])
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
**Custom model cost map URL:**
|
||||
```bash
|
||||
export LITELLM_MODEL_COST_MAP_URL="https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json"
|
||||
```
|
||||
|
||||
**Use local model cost map:**
|
||||
```bash
|
||||
export LITELLM_LOCAL_MODEL_COST_MAP=True
|
||||
```
|
||||
|
|
@ -54,6 +54,20 @@ Allow others to create/delete their own keys.
|
|||
|
||||
[**Go Here**](./self_serve.md)
|
||||
|
||||
## Model Management
|
||||
|
||||
The Admin UI provides comprehensive model management capabilities:
|
||||
|
||||
- **Add Models**: Add new models through the UI without restarting the proxy
|
||||
- **Model Hub**: Make models public for developers to discover available models
|
||||
- **Price Data Sync**: Keep model pricing data up to date by syncing from GitHub
|
||||
|
||||
For detailed information on model management, see [Model Management](./model_management.md).
|
||||
|
||||
:::tip Sync Model Pricing Data
|
||||
[Sync model pricing data from GitHub](./sync_models_github.md) to keep your model cost information current.
|
||||
:::
|
||||
|
||||
## Disable Admin UI
|
||||
|
||||
Set `DISABLE_ADMIN_UI="True"` in your environment to disable the Admin UI.
|
||||
|
|
|
|||
BIN
docs/my-website/img/release_notes/perf_77_5.png
Normal file
BIN
docs/my-website/img/release_notes/perf_77_5.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 1.3 MiB |
BIN
docs/my-website/img/release_notes/perf_77_7.png
Normal file
BIN
docs/my-website/img/release_notes/perf_77_7.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 253 KiB |
BIN
docs/my-website/img/release_notes/schedule_key_rotations.png
Normal file
BIN
docs/my-website/img/release_notes/schedule_key_rotations.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 603 KiB |
|
|
@ -50,6 +50,7 @@ pip install litellm==1.75.5.post2
|
|||
- **Oracle Cloud Infrastructure** - New LLM provider for calling models on Oracle Cloud Infrastructure.
|
||||
- **Digital Ocean's Gradient AI** - New LLM provider for calling models on Digital Ocean's Gradient AI platform.
|
||||
|
||||
---
|
||||
|
||||
### Risk of Upgrade
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
---
|
||||
title: "[Preview] v1.77.5-stable - MCP OAuth 2.0 Support"
|
||||
title: "v1.77.5-stable - MCP OAuth 2.0 Support"
|
||||
slug: "v1-77-5"
|
||||
date: 2025-09-29T10:00:00
|
||||
authors:
|
||||
|
|
@ -11,6 +11,10 @@ authors:
|
|||
title: CTO, LiteLLM
|
||||
url: https://www.linkedin.com/in/reffajnaahsi/
|
||||
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
|
||||
- name: Alexsander Hamir
|
||||
title: Backend Performance Engineer
|
||||
url: https://www.linkedin.com/in/alexsander-baptista/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D5603AQGXnziu4kqNCQ/profile-displayphoto-crop_800_800/B56ZkxEcuOKEAI-/0/1757464874550?e=1762387200&v=beta&t=9SNXLsWhx8OnYPAMQ9fqAr02oevDYEAL2vMYg2f9ieg
|
||||
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
|
@ -28,7 +32,7 @@ import TabItem from '@theme/TabItem';
|
|||
docker run \
|
||||
-e STORE_MODEL_IN_DB=True \
|
||||
-p 4000:4000 \
|
||||
ghcr.io/berriai/litellm:v1.77.5.rc.1
|
||||
ghcr.io/berriai/litellm:v1.77.5-stable
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
|
@ -49,7 +53,54 @@ pip install litellm==1.77.5
|
|||
- **MCP OAuth 2.0 Support** - Enhanced authentication for Model Context Protocol integrations
|
||||
- **Scheduled Key Rotations** - Automated key rotation capabilities for enhanced security
|
||||
- **New Gemini 2.5 Flash & Flash-lite Models** - Latest September 2025 preview models with improved pricing and features
|
||||
- **Performance Improvements** - Critical InMemoryCache unbounded growth resolution
|
||||
- **Performance Improvements** - 54% RPS improvement
|
||||
|
||||
---
|
||||
|
||||
### Scheduled Key Rotations
|
||||
|
||||
<Image img={require('../../img/release_notes/schedule_key_rotations.png')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release brings support for scheduling virtual key rotations on LiteLLM AI Gateway.
|
||||
|
||||
This is great for Proxy Admins looking to enforce Enterprise Grade security for use cases going through LiteLLM AI Gateway.
|
||||
|
||||
From this release you can enforce Virtual Keys to rotate on a schedule of your choice e.g every 15 days/30 days/60 days etc.
|
||||
|
||||
---
|
||||
### Performance Improvements - 54% RPS Improvement
|
||||
|
||||
<Image img={require('../../img/release_notes/perf_77_5.png')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release brings a 54% RPS improvement (1,040 → 1,602 RPS, aggregated) per instance.
|
||||
|
||||
The improvement comes from fixing O(n²) inefficiencies in the LiteLLM Router, primarily caused by repeated use of `in` statements inside loops over large arrays.
|
||||
|
||||
Tests were run with a database-only setup (no cache hits).
|
||||
|
||||
#### Test Setup
|
||||
|
||||
All benchmarks were executed using Locust with 1,000 concurrent users and a ramp-up of 500. The environment was configured to stress the routing layer and eliminate caching as a variable.
|
||||
|
||||
**System Specs**
|
||||
|
||||
- **CPU:** 8 vCPUs
|
||||
- **Memory:** 32 GB RAM
|
||||
|
||||
**Configuration (config.yaml)**
|
||||
|
||||
View the complete configuration: [gist.github.com/AlexsanderHamir/config.yaml](https://gist.github.com/AlexsanderHamir/53f7d554a5d2afcf2c4edb5b6be68ff4)
|
||||
|
||||
**Load Script (no_cache_hits.py)**
|
||||
|
||||
View the complete load testing script: [gist.github.com/AlexsanderHamir/no_cache_hits.py](https://gist.github.com/AlexsanderHamir/42c33d7a4dc7a57f56a78b560dee3a42)
|
||||
|
||||
---
|
||||
|
||||
|
||||
## New Models / Updated Models
|
||||
|
||||
|
|
|
|||
364
docs/my-website/release_notes/v1.77.7-stable/index.md
Normal file
364
docs/my-website/release_notes/v1.77.7-stable/index.md
Normal file
|
|
@ -0,0 +1,364 @@
|
|||
---
|
||||
title: "[Preview] v1.77.7-stable - Claude Sonnet 4.5"
|
||||
slug: "v1-77-7"
|
||||
date: 2025-10-04T10:00:00
|
||||
authors:
|
||||
- name: Krrish Dholakia
|
||||
title: CEO, LiteLLM
|
||||
url: https://www.linkedin.com/in/krish-d/
|
||||
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
|
||||
- name: Ishaan Jaff
|
||||
title: CTO, LiteLLM
|
||||
url: https://www.linkedin.com/in/reffajnaahsi/
|
||||
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
|
||||
- name: Alexsander Hamir
|
||||
title: Backend Performance Engineer
|
||||
url: https://www.linkedin.com/in/alexsander-baptista/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D5603AQGXnziu4kqNCQ/profile-displayphoto-crop_800_800/B56ZkxEcuOKEAI-/0/1757464874550?e=1762387200&v=beta&t=9SNXLsWhx8OnYPAMQ9fqAr02oevDYEAL2vMYg2f9ieg
|
||||
- name: Achintya Srivastava
|
||||
title: Fullstack Engineer
|
||||
url: https://www.linkedin.com/in/achintya-rajan/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D5603AQGdkEeyJTdljw/profile-displayphoto-shrink_800_800/profile-displayphoto-shrink_800_800/0/1716271140869?e=1762387200&v=beta&t=9gOoLPeqR2E5z3KSX61EUj3HVZXmgo87vhVuSHeffjc
|
||||
- name: Sameer Kankute
|
||||
title: Backend Engineer (LLM Translation)
|
||||
url: https://www.linkedin.com/in/sameer-kankute/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D4D03AQHB_loQYd5gjg/profile-displayphoto-shrink_800_800/profile-displayphoto-shrink_800_800/0/1719137160975?e=1762387200&v=beta&t=0jbuX-f4eSnDxBY3olI6meuYr-LMbObhFmFbRcKF5mY
|
||||
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
import Image from '@theme/IdealImage';
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
## Deploy this version
|
||||
|
||||
<Tabs>
|
||||
<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.7.rc.1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="pip" label="Pip">
|
||||
|
||||
``` showLineNumbers title="pip install litellm"
|
||||
pip install litellm==1.77.7.rc.1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
---
|
||||
|
||||
## Key Highlights
|
||||
|
||||
- **Dynamic Rate Limiter v3** - Automatically maximizes throughput when capacity is available (< 80% saturation) by allowing lower-priority requests to use unused capacity, then switches to fair priority-based allocation under high load (≥ 80%) to prevent blocking
|
||||
- **Major Performance Improvements** - 2.9x lower median latency at 1,000 concurrent users.
|
||||
- **Claude Sonnet 4.5** - Support for Anthropic's new Claude Sonnet 4.5 model family with 200K+ context and tiered pricing
|
||||
- **MCP Gateway Enhancements** - Fine-grained tool control, server permissions, and forwardable headers
|
||||
- **AMD Lemonade & Nvidia NIM** - New provider support for AMD Lemonade and Nvidia NIM Rerank
|
||||
- **GitLab Prompt Management** - GitLab-based prompt management integration
|
||||
|
||||
### Performance - 2.9x Lower Median Latency
|
||||
|
||||
<Image img={require('../../img/release_notes/perf_77_7.png')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
||||
<br/>
|
||||
|
||||
This update removes LiteLLM router inefficiencies, reducing complexity from O(M×N) to O(1). Previously, it built a new array and ran repeated checks like data["model"] in llm_router.get_model_ids(). Now, a direct ID-to-deployment map eliminates redundant allocations and scans.
|
||||
|
||||
As a result, performance improved across all latency percentiles:
|
||||
|
||||
- **Median latency:** 320 ms → **110 ms** (−65.6%)
|
||||
- **p95 latency:** 850 ms → **440 ms** (−48.2%)
|
||||
- **p99 latency:** 1,400 ms → **810 ms** (−42.1%)
|
||||
- **Average latency:** 864 ms → **310 ms** (−64%)
|
||||
|
||||
|
||||
#### Test Setup
|
||||
|
||||
**Locust**
|
||||
|
||||
- **Concurrent users:** 1,000
|
||||
- **Ramp-up:** 500
|
||||
|
||||
**System Specs**
|
||||
|
||||
- **CPU:** 4 vCPUs
|
||||
- **Memory:** 8 GB RAM
|
||||
- **LiteLLM Workers:** 4
|
||||
- **Instances**: 4
|
||||
|
||||
**Configuration (config.yaml)**
|
||||
|
||||
View the complete configuration: [gist.github.com/AlexsanderHamir/config.yaml](https://gist.github.com/AlexsanderHamir/53f7d554a5d2afcf2c4edb5b6be68ff4)
|
||||
|
||||
**Load Script (no_cache_hits.py)**
|
||||
|
||||
View the complete load testing script: [gist.github.com/AlexsanderHamir/no_cache_hits.py](https://gist.github.com/AlexsanderHamir/42c33d7a4dc7a57f56a78b560dee3a42)
|
||||
|
||||
## New Models / Updated Models
|
||||
|
||||
#### New Model Support
|
||||
|
||||
| Provider | Model | Context Window | Input ($/1M tokens) | Output ($/1M tokens) | Features |
|
||||
| -------- | ----- | -------------- | ------------------- | -------------------- | -------- |
|
||||
| Anthropic | `claude-sonnet-4-5` | 200K | $3.00 | $15.00 | Chat, reasoning, vision, function calling, prompt caching |
|
||||
| Anthropic | `claude-sonnet-4-5-20250929` | 200K | $3.00 | $15.00 | Chat, reasoning, vision, function calling, prompt caching |
|
||||
| Bedrock | `eu.anthropic.claude-sonnet-4-5-20250929-v1:0` | 200K | $3.00 | $15.00 | Chat, reasoning, vision, function calling, prompt caching |
|
||||
| Azure AI | `azure_ai/grok-4` | 131K | $5.50 | $27.50 | Chat, reasoning, function calling, web search |
|
||||
| Azure AI | `azure_ai/grok-4-fast-reasoning` | 131K | $0.43 | $1.73 | Chat, reasoning, function calling, web search |
|
||||
| Azure AI | `azure_ai/grok-4-fast-non-reasoning` | 131K | $0.43 | $1.73 | Chat, function calling, web search |
|
||||
| Azure AI | `azure_ai/grok-code-fast-1` | 131K | $3.50 | $17.50 | Chat, function calling, web search |
|
||||
| Groq | `groq/moonshotai/kimi-k2-instruct-0905` | Context varies | Pricing varies | Pricing varies | Chat, function calling |
|
||||
| Ollama | Ollama Cloud models | Varies | Free | Free | Self-hosted models via Ollama Cloud |
|
||||
|
||||
#### Features
|
||||
|
||||
- **[Anthropic](../../docs/providers/anthropic)**
|
||||
- Add new claude-sonnet-4-5 model family with tiered pricing above 200K tokens - [PR #15041](https://github.com/BerriAI/litellm/pull/15041)
|
||||
- Add anthropic/claude-sonnet-4-5 to model price json with prompt caching support - [PR #15049](https://github.com/BerriAI/litellm/pull/15049)
|
||||
- Add 200K prices for Sonnet 4.5 - [PR #15140](https://github.com/BerriAI/litellm/pull/15140)
|
||||
- Add cost tracking for /v1/messages in streaming response - [PR #15102](https://github.com/BerriAI/litellm/pull/15102)
|
||||
- Add /v1/messages/count_tokens to Anthropic routes for non-admin user access - [PR #15034](https://github.com/BerriAI/litellm/pull/15034)
|
||||
- **[Gemini](../../docs/providers/gemini)**
|
||||
- Ignore type param for gemini tools - [PR #15022](https://github.com/BerriAI/litellm/pull/15022)
|
||||
- **[Vertex AI](../../docs/providers/vertex)**
|
||||
- Add LiteLLM Overhead metric for VertexAI - [PR #15040](https://github.com/BerriAI/litellm/pull/15040)
|
||||
- Support googlemap grounding in vertex ai - [PR #15179](https://github.com/BerriAI/litellm/pull/15179)
|
||||
- **[Azure](../../docs/providers/azure)**
|
||||
- Add azure_ai grok-4 model family - [PR #15137](https://github.com/BerriAI/litellm/pull/15137)
|
||||
- Use the `extra_query` parameter for GET requests in Azure Batch - [PR #14997](https://github.com/BerriAI/litellm/pull/14997)
|
||||
- Use extra_query for download results (Batch API) - [PR #15025](https://github.com/BerriAI/litellm/pull/15025)
|
||||
- Add support for Azure AD token-based authorization - [PR #14813](https://github.com/BerriAI/litellm/pull/14813)
|
||||
- **[Ollama](../../docs/providers/ollama)**
|
||||
- Add ollama cloud models - [PR #15008](https://github.com/BerriAI/litellm/pull/15008)
|
||||
- **[Groq](../../docs/providers/groq)**
|
||||
- Add groq/moonshotai/kimi-k2-instruct-0905 - [PR #15079](https://github.com/BerriAI/litellm/pull/15079)
|
||||
- **[OpenAI](../../docs/providers/openai)**
|
||||
- Add support for GPT 5 codex models - [PR #14841](https://github.com/BerriAI/litellm/pull/14841)
|
||||
- **[DeepInfra](../../docs/providers/deepinfra)**
|
||||
- Update DeepInfra model data refresh with latest pricing - [PR #14939](https://github.com/BerriAI/litellm/pull/14939)
|
||||
- **[Bedrock](../../docs/providers/bedrock)**
|
||||
- Add JP Cross-Region Inference - [PR #15188](https://github.com/BerriAI/litellm/pull/15188)
|
||||
- Add "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" - [PR #15181](https://github.com/BerriAI/litellm/pull/15181)
|
||||
- Add twelvelabs bedrock Async Invoke Support - [PR #14871](https://github.com/BerriAI/litellm/pull/14871)
|
||||
- **[Nvidia NIM](../../docs/providers/nvidia_nim)**
|
||||
- Add Nvidia NIM Rerank Support - [PR #15152](https://github.com/BerriAI/litellm/pull/15152)
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- **[VLLM](../../docs/providers/vllm)**
|
||||
- Fix response_format bug in hosted vllm audio_transcription - [PR #15010](https://github.com/BerriAI/litellm/pull/15010)
|
||||
- Fix passthrough of atranscription into kwargs going to upstream provider - [PR #15005](https://github.com/BerriAI/litellm/pull/15005)
|
||||
- **[OCI](../../docs/providers/oci)**
|
||||
- Fix OCI Generative AI Integration when using Proxy - [PR #15072](https://github.com/BerriAI/litellm/pull/15072)
|
||||
- **General**
|
||||
- Fix: Authorization header to use correct "Bearer" capitalization - [PR #14764](https://github.com/BerriAI/litellm/pull/14764)
|
||||
- Bug fix: gpt-5-chat-latest has incorrect max_input_tokens value - [PR #15116](https://github.com/BerriAI/litellm/pull/15116)
|
||||
- Update request handling for original exceptions - [PR #15013](https://github.com/BerriAI/litellm/pull/15013)
|
||||
|
||||
#### New Provider Support
|
||||
|
||||
- **[AMD Lemonade](../../docs/providers/lemonade)**
|
||||
- Add AMD Lemonade provider support - [PR #14840](https://github.com/BerriAI/litellm/pull/14840)
|
||||
|
||||
---
|
||||
|
||||
## LLM API Endpoints
|
||||
|
||||
#### Features
|
||||
|
||||
- **[Responses API](../../docs/response_api)**
|
||||
- Return Cost for Responses API Streaming requests - [PR #15053](https://github.com/BerriAI/litellm/pull/15053)
|
||||
|
||||
- **[/generateContent](../../docs/providers/gemini)**
|
||||
- Add full support for native Gemini API translation - [PR #15029](https://github.com/BerriAI/litellm/pull/15029)
|
||||
|
||||
- **Passthrough Gemini Routes**
|
||||
- Add Gemini generateContent passthrough cost tracking - [PR #15014](https://github.com/BerriAI/litellm/pull/15014)
|
||||
- Add streamGenerateContent cost tracking in passthrough - [PR #15199](https://github.com/BerriAI/litellm/pull/15199)
|
||||
|
||||
- **Passthrough Vertex AI Routes**
|
||||
- Add cost tracking for Vertex AI Passthrough `/predict` endpoint - [PR #15019](https://github.com/BerriAI/litellm/pull/15019)
|
||||
- Add cost tracking for Vertex AI Live API WebSocket Passthrough - [PR #14956](https://github.com/BerriAI/litellm/pull/14956)
|
||||
|
||||
- **General**
|
||||
- Preserve Whitespace Characters in Model Response Streams - [PR #15160](https://github.com/BerriAI/litellm/pull/15160)
|
||||
- Add provider name to payload specification - [PR #15130](https://github.com/BerriAI/litellm/pull/15130)
|
||||
- Ensure query params are forwarded from origin url to downstream request - [PR #15087](https://github.com/BerriAI/litellm/pull/15087)
|
||||
|
||||
---
|
||||
|
||||
## Management Endpoints / UI
|
||||
|
||||
#### Features
|
||||
|
||||
- **Virtual Keys**
|
||||
- Ensure LLM_API_KEYs can access pass through routes - [PR #15115](https://github.com/BerriAI/litellm/pull/15115)
|
||||
- Support 'guaranteed_throughput' when setting limits on keys belonging to a team - [PR #15120](https://github.com/BerriAI/litellm/pull/15120)
|
||||
|
||||
- **Models + Endpoints**
|
||||
- Ensure OCI secret fields not shared on /models and /v1/models endpoints - [PR #15085](https://github.com/BerriAI/litellm/pull/15085)
|
||||
- Add snowflake on UI - [PR #15083](https://github.com/BerriAI/litellm/pull/15083)
|
||||
- Make UI theme settings publicly accessible for custom branding - [PR #15074](https://github.com/BerriAI/litellm/pull/15074)
|
||||
|
||||
- **Admin Settings**
|
||||
- Ensure OTEL settings are saved in DB after set on UI - [PR #15118](https://github.com/BerriAI/litellm/pull/15118)
|
||||
- Top api key tags - [PR #15151](https://github.com/BerriAI/litellm/pull/15151), [PR #15156](https://github.com/BerriAI/litellm/pull/15156)
|
||||
|
||||
- **MCP**
|
||||
- show health status of MCP servers - [PR #15185](https://github.com/BerriAI/litellm/pull/15185)
|
||||
- allow setting extra headers on the UI - [PR #15185](https://github.com/BerriAI/litellm/pull/15185)
|
||||
- allow editing allowed tools on the UI - [PR #15185](https://github.com/BerriAI/litellm/pull/15185)
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- **Virtual Keys**
|
||||
- (security) prevent user key from updating other user keys - [PR #15201](https://github.com/BerriAI/litellm/pull/15201)
|
||||
- (security) don't return all keys with blank key alias on /v2/key/info - [PR #15201](https://github.com/BerriAI/litellm/pull/15201)
|
||||
- Fix Session Token Cookie Infinite Logout Loop - [PR #15146](https://github.com/BerriAI/litellm/pull/15146)
|
||||
|
||||
- **Models + Endpoints**
|
||||
- Make UI theme settings publicly accessible for custom branding - [PR #15074](https://github.com/BerriAI/litellm/pull/15074)
|
||||
|
||||
- **Teams**
|
||||
- fix failed copy to clipboard for http ui - [PR #15195](https://github.com/BerriAI/litellm/pull/15195)
|
||||
|
||||
- **Logs**
|
||||
- fix logs page render logs on filter lookup - [PR #15195](https://github.com/BerriAI/litellm/pull/15195)
|
||||
- fix lookup list of end users (migrate to more efficient /customers/list lookup) - [PR #15195](https://github.com/BerriAI/litellm/pull/15195)
|
||||
|
||||
- **Test key**
|
||||
- update selected model on key change - [PR #15197](https://github.com/BerriAI/litellm/pull/15197)
|
||||
|
||||
- **Dashboard**
|
||||
- Fix LiteLLM model name fallback in dashboard overview - [PR #14998](https://github.com/BerriAI/litellm/pull/14998)
|
||||
|
||||
|
||||
---
|
||||
|
||||
## Logging / Guardrail / Prompt Management Integrations
|
||||
|
||||
#### Features
|
||||
|
||||
- **[OpenTelemetry](../../docs/observability/otel)**
|
||||
- Use generation_name for span naming in logging method - [PR #14799](https://github.com/BerriAI/litellm/pull/14799)
|
||||
- **[Langfuse](../../docs/proxy/logging#langfuse)**
|
||||
- Handle non-serializable objects in Langfuse logging - [PR #15148](https://github.com/BerriAI/litellm/pull/15148)
|
||||
- Set usage_details.total in langfuse integration - [PR #15015](https://github.com/BerriAI/litellm/pull/15015)
|
||||
- **[Prometheus](../../docs/proxy/prometheus)**
|
||||
- support custom metadata labels on key/team - [PR #15094](https://github.com/BerriAI/litellm/pull/15094)
|
||||
|
||||
|
||||
#### Guardrails
|
||||
|
||||
- **[Javelin](../../docs/proxy/guardrails)**
|
||||
- Add Javelin standalone guardrails integration for LiteLLM Proxy - [PR #14983](https://github.com/BerriAI/litellm/pull/14983)
|
||||
- Add logging for important status fields in guardrails - [PR #15090](https://github.com/BerriAI/litellm/pull/15090)
|
||||
- Don't run post_call guardrail if no text returned from Bedrock - [PR #15106](https://github.com/BerriAI/litellm/pull/15106)
|
||||
|
||||
#### Prompt Management
|
||||
|
||||
- **[GitLab](../../docs/proxy/prompt_management)**
|
||||
- GitLab based Prompt manager - [PR #14988](https://github.com/BerriAI/litellm/pull/14988)
|
||||
|
||||
---
|
||||
|
||||
## Spend Tracking, Budgets and Rate Limiting
|
||||
|
||||
- **Cost Tracking**
|
||||
- Proxy: end user cost tracking in the responses API - [PR #15124](https://github.com/BerriAI/litellm/pull/15124)
|
||||
- **Parallel Request Limiter v3**
|
||||
- Use well known redis cluster hashing algorithm - [PR #15052](https://github.com/BerriAI/litellm/pull/15052)
|
||||
- Fixes to dynamic rate limiter v3 - add saturation detection - [PR #15119](https://github.com/BerriAI/litellm/pull/15119)
|
||||
- Dynamic Rate Limiter v3 - fixes for detecting saturation + fixes for post saturation behavior - [PR #15192](https://github.com/BerriAI/litellm/pull/15192)
|
||||
- **Teams**
|
||||
- Add model specific tpm/rpm limits to teams on LiteLLM - [PR #15044](https://github.com/BerriAI/litellm/pull/15044)
|
||||
|
||||
---
|
||||
|
||||
## MCP Gateway
|
||||
|
||||
- **Server Configuration**
|
||||
- Specify forwardable headers, specify allowed/disallowed tools for MCP servers - [PR #15002](https://github.com/BerriAI/litellm/pull/15002)
|
||||
- Enforce server permissions on call tools - [PR #15044](https://github.com/BerriAI/litellm/pull/15044)
|
||||
- MCP Gateway Fine-grained Tools Addition - [PR #15153](https://github.com/BerriAI/litellm/pull/15153)
|
||||
- **Bug Fixes**
|
||||
- Remove servername prefix mcp tools tests - [PR #14986](https://github.com/BerriAI/litellm/pull/14986)
|
||||
- Resolve regression with duplicate Mcp-Protocol-Version header - [PR #15050](https://github.com/BerriAI/litellm/pull/15050)
|
||||
- Fix test_mcp_server.py - [PR #15183](https://github.com/BerriAI/litellm/pull/15183)
|
||||
|
||||
---
|
||||
|
||||
## Performance / Loadbalancing / Reliability improvements
|
||||
|
||||
- **Router Optimizations**
|
||||
- **+62.5% P99 Latency Improvement** - Remove router inefficiencies (from O(M*N) to O(1)) - [PR #15046](https://github.com/BerriAI/litellm/pull/15046)
|
||||
- Remove hasattr checks in Router - [PR #15082](https://github.com/BerriAI/litellm/pull/15082)
|
||||
- Remove Double Lookups - [PR #15084](https://github.com/BerriAI/litellm/pull/15084)
|
||||
- Optimize _filter_cooldown_deployments from O(n×m + k×n) to O(n) - [PR #15091](https://github.com/BerriAI/litellm/pull/15091)
|
||||
- Optimize unhealthy deployment filtering in retry path (O(n*m) → O(n+m)) - [PR #15110](https://github.com/BerriAI/litellm/pull/15110)
|
||||
- **Cache Optimizations**
|
||||
- Reduce complexity of InMemoryCache.evict_cache from O(n*log(n)) to O(log(n)) - [PR #15000](https://github.com/BerriAI/litellm/pull/15000)
|
||||
- Avoiding expensive operations when cache isn't available - [PR #15182](https://github.com/BerriAI/litellm/pull/15182)
|
||||
- **Worker Management**
|
||||
- Add proxy CLI option to recycle workers after N requests - [PR #15007](https://github.com/BerriAI/litellm/pull/15007)
|
||||
- **Metrics & Monitoring**
|
||||
- LiteLLM Overhead metric tracking - Add support for tracking litellm overhead on cache hits - [PR #15045](https://github.com/BerriAI/litellm/pull/15045)
|
||||
|
||||
---
|
||||
|
||||
## Documentation Updates
|
||||
|
||||
- **Provider Documentation**
|
||||
- Update litellm docs from latest release - [PR #15004](https://github.com/BerriAI/litellm/pull/15004)
|
||||
- Add missing api_key parameter - [PR #15058](https://github.com/BerriAI/litellm/pull/15058)
|
||||
- **General Documentation**
|
||||
- Use docker compose instead of docker-compose - [PR #15024](https://github.com/BerriAI/litellm/pull/15024)
|
||||
- Add railtracks to projects that are using litellm - [PR #15144](https://github.com/BerriAI/litellm/pull/15144)
|
||||
- Perf: Last week improvement - [PR #15193](https://github.com/BerriAI/litellm/pull/15193)
|
||||
- Sync models GitHub documentation with Loom video and cross-reference - [PR #15191](https://github.com/BerriAI/litellm/pull/15191)
|
||||
|
||||
---
|
||||
|
||||
## Security Fixes
|
||||
|
||||
- **JWT Token Security** - Don't log JWT SSO token on .info() log - [PR #15145](https://github.com/BerriAI/litellm/pull/15145)
|
||||
|
||||
---
|
||||
|
||||
## New Contributors
|
||||
|
||||
* @herve-ves made their first contribution in [PR #14998](https://github.com/BerriAI/litellm/pull/14998)
|
||||
* @wenxi-onyx made their first contribution in [PR #15008](https://github.com/BerriAI/litellm/pull/15008)
|
||||
* @jpetrucciani made their first contribution in [PR #15005](https://github.com/BerriAI/litellm/pull/15005)
|
||||
* @abhijitjavelin made their first contribution in [PR #14983](https://github.com/BerriAI/litellm/pull/14983)
|
||||
* @ZeroClover made their first contribution in [PR #15039](https://github.com/BerriAI/litellm/pull/15039)
|
||||
* @cedarm made their first contribution in [PR #15043](https://github.com/BerriAI/litellm/pull/15043)
|
||||
* @Isydmr made their first contribution in [PR #15025](https://github.com/BerriAI/litellm/pull/15025)
|
||||
* @serializer made their first contribution in [PR #15013](https://github.com/BerriAI/litellm/pull/15013)
|
||||
* @eddierichter-amd made their first contribution in [PR #14840](https://github.com/BerriAI/litellm/pull/14840)
|
||||
* @malags made their first contribution in [PR #15000](https://github.com/BerriAI/litellm/pull/15000)
|
||||
* @henryhwang made their first contribution in [PR #15029](https://github.com/BerriAI/litellm/pull/15029)
|
||||
* @plafleur made their first contribution in [PR #15111](https://github.com/BerriAI/litellm/pull/15111)
|
||||
* @tyler-liner made their first contribution in [PR #14799](https://github.com/BerriAI/litellm/pull/14799)
|
||||
* @Amir-R25 made their first contribution in [PR #15144](https://github.com/BerriAI/litellm/pull/15144)
|
||||
* @georg-wolflein made their first contribution in [PR #15124](https://github.com/BerriAI/litellm/pull/15124)
|
||||
* @niharm made their first contribution in [PR #15140](https://github.com/BerriAI/litellm/pull/15140)
|
||||
* @anthony-liner made their first contribution in [PR #15015](https://github.com/BerriAI/litellm/pull/15015)
|
||||
* @rishiganesh2002 made their first contribution in [PR #15153](https://github.com/BerriAI/litellm/pull/15153)
|
||||
* @danielaskdd made their first contribution in [PR #15160](https://github.com/BerriAI/litellm/pull/15160)
|
||||
* @JVenberg made their first contribution in [PR #15146](https://github.com/BerriAI/litellm/pull/15146)
|
||||
* @speglich made their first contribution in [PR #15072](https://github.com/BerriAI/litellm/pull/15072)
|
||||
* @daily-kim made their first contribution in [PR #14764](https://github.com/BerriAI/litellm/pull/14764)
|
||||
|
||||
---
|
||||
|
||||
## **[Full Changelog](https://github.com/BerriAI/litellm/compare/v1.77.5.rc.4...v1.77.7.rc.1)**
|
||||
|
|
@ -119,6 +119,7 @@ class PagerDutyAlerting(SlackAlerting):
|
|||
user_api_key_end_user_id=_meta.get("user_api_key_end_user_id"),
|
||||
user_api_key_user_email=_meta.get("user_api_key_user_email"),
|
||||
user_api_key_request_route=_meta.get("user_api_key_request_route"),
|
||||
user_api_key_auth_metadata=_meta.get("user_api_key_auth_metadata"),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -196,7 +197,11 @@ class PagerDutyAlerting(SlackAlerting):
|
|||
user_api_key_alias=user_api_key_dict.key_alias,
|
||||
user_api_key_spend=user_api_key_dict.spend,
|
||||
user_api_key_max_budget=user_api_key_dict.max_budget,
|
||||
user_api_key_budget_reset_at=user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None,
|
||||
user_api_key_budget_reset_at=(
|
||||
user_api_key_dict.budget_reset_at.isoformat()
|
||||
if user_api_key_dict.budget_reset_at
|
||||
else None
|
||||
),
|
||||
user_api_key_org_id=user_api_key_dict.org_id,
|
||||
user_api_key_team_id=user_api_key_dict.team_id,
|
||||
user_api_key_user_id=user_api_key_dict.user_id,
|
||||
|
|
@ -204,6 +209,7 @@ class PagerDutyAlerting(SlackAlerting):
|
|||
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
||||
user_api_key_user_email=user_api_key_dict.user_email,
|
||||
user_api_key_request_route=user_api_key_dict.request_route,
|
||||
user_api_key_auth_metadata=user_api_key_dict.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm._logging import print_verbose, verbose_logger
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth
|
||||
from litellm.types.integrations.prometheus import *
|
||||
from litellm.types.integrations.prometheus import _sanitize_prometheus_label_name
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
from litellm.utils import get_end_user_id_for_cost_tracking
|
||||
|
||||
|
|
@ -794,9 +795,16 @@ class PrometheusLogger(CustomLogger):
|
|||
output_tokens = standard_logging_payload["completion_tokens"]
|
||||
tokens_used = standard_logging_payload["total_tokens"]
|
||||
response_cost = standard_logging_payload["response_cost"]
|
||||
_requester_metadata = standard_logging_payload["metadata"].get(
|
||||
_requester_metadata: Optional[dict] = standard_logging_payload["metadata"].get(
|
||||
"requester_metadata"
|
||||
)
|
||||
user_api_key_auth_metadata: Optional[dict] = standard_logging_payload[
|
||||
"metadata"
|
||||
].get("user_api_key_auth_metadata")
|
||||
combined_metadata: Dict[str, Any] = {
|
||||
**(_requester_metadata if _requester_metadata else {}),
|
||||
**(user_api_key_auth_metadata if user_api_key_auth_metadata else {}),
|
||||
}
|
||||
if standard_logging_payload is not None and isinstance(
|
||||
standard_logging_payload, dict
|
||||
):
|
||||
|
|
@ -828,8 +836,7 @@ class PrometheusLogger(CustomLogger):
|
|||
exception_status=None,
|
||||
exception_class=None,
|
||||
custom_metadata_labels=get_custom_labels_from_metadata(
|
||||
metadata=standard_logging_payload["metadata"].get("requester_metadata")
|
||||
or {}
|
||||
metadata=combined_metadata
|
||||
),
|
||||
route=standard_logging_payload["metadata"].get(
|
||||
"user_api_key_request_route"
|
||||
|
|
@ -1649,9 +1656,22 @@ class PrometheusLogger(CustomLogger):
|
|||
api_base: Optional[str],
|
||||
api_provider: str,
|
||||
):
|
||||
self.litellm_deployment_state.labels(
|
||||
litellm_model_name, model_id, api_base, api_provider
|
||||
).set(state)
|
||||
"""
|
||||
Set the deployment state.
|
||||
"""
|
||||
### get labels
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_deployment_state"
|
||||
),
|
||||
enum_values=UserAPIKeyLabelValues(
|
||||
litellm_model_name=litellm_model_name,
|
||||
model_id=model_id,
|
||||
api_base=api_base,
|
||||
api_provider=api_provider,
|
||||
),
|
||||
)
|
||||
self.litellm_deployment_state.labels(**_labels).set(state)
|
||||
|
||||
def set_deployment_healthy(
|
||||
self,
|
||||
|
|
@ -2228,8 +2248,10 @@ def prometheus_label_factory(
|
|||
|
||||
if enum_values.custom_metadata_labels is not None:
|
||||
for key, value in enum_values.custom_metadata_labels.items():
|
||||
if key in supported_enum_labels:
|
||||
filtered_labels[key] = value
|
||||
# check sanitized key
|
||||
sanitized_key = _sanitize_prometheus_label_name(key)
|
||||
if sanitized_key in supported_enum_labels:
|
||||
filtered_labels[sanitized_key] = value
|
||||
|
||||
# Add custom tags if configured
|
||||
if enum_values.tags is not None:
|
||||
|
|
|
|||
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.23-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.23-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.23.tar.gz
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.23.tar.gz
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.25-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.25-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.25.tar.gz
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.25.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -0,0 +1,3 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "allowed_tools" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "extra_headers" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
||||
|
|
@ -178,6 +178,8 @@ model LiteLLM_MCPServerTable {
|
|||
updated_by String?
|
||||
mcp_info Json? @default("{}")
|
||||
mcp_access_groups String[]
|
||||
allowed_tools String[] @default([])
|
||||
extra_headers String[] @default([])
|
||||
// Health check status
|
||||
status String? @default("unknown")
|
||||
last_health_check DateTime?
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.2.22"
|
||||
version = "0.2.25"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
|
|
@ -22,7 +22,7 @@ requires = ["poetry-core"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.2.22"
|
||||
version = "0.2.25"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -290,7 +290,7 @@ banned_keywords_list: Optional[Union[str, List]] = None
|
|||
llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all"
|
||||
guardrail_name_config_map: Dict[str, GuardrailItem] = {}
|
||||
include_cost_in_streaming_usage: bool = False
|
||||
### PROMPTS ###
|
||||
### PROMPTS ####
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
|
||||
prompt_name_config_map: Dict[str, PromptSpec] = {}
|
||||
|
|
@ -367,7 +367,7 @@ disable_add_prefix_to_prompt: bool = (
|
|||
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 ######
|
||||
#### REQUEST PRIORITIZATION #######
|
||||
priority_reservation: Optional[Dict[str, float]] = None
|
||||
priority_reservation_settings: "PriorityReservationSettings" = (
|
||||
PriorityReservationSettings()
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ It utilizes the (RedisCache, s3Cache, RedisSemanticCache, QdrantSemanticCache, I
|
|||
In each method it will call the appropriate method from caching.py
|
||||
"""
|
||||
|
||||
import time
|
||||
import asyncio
|
||||
import datetime
|
||||
import inspect
|
||||
|
|
@ -57,10 +58,16 @@ from litellm.types.utils import (
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
CustomStreamWrapper = Any
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
)
|
||||
|
||||
|
||||
class CachingHandlerResponse(BaseModel):
|
||||
|
|
@ -112,7 +119,7 @@ class LLMCachingHandler:
|
|||
call_type: str,
|
||||
kwargs: Dict[str, Any],
|
||||
args: Optional[Tuple[Any, ...]] = None,
|
||||
) -> CachingHandlerResponse:
|
||||
) -> Optional[CachingHandlerResponse]:
|
||||
"""
|
||||
Internal method to get from the cache.
|
||||
Handles different call types (embeddings, chat/completions, text_completion, transcription)
|
||||
|
|
@ -133,32 +140,27 @@ class LLMCachingHandler:
|
|||
Raises:
|
||||
None
|
||||
"""
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
)
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
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
|
||||
final_embedding_cached_response: Optional[EmbeddingResponse] = None
|
||||
embedding_all_elements_cache_hit: bool = False
|
||||
cached_result: Optional[Any] = None
|
||||
# Check if caching should be performed BEFORE doing expensive operations
|
||||
if (
|
||||
(kwargs.get("caching", None) is None and litellm.cache is not None)
|
||||
or kwargs.get("caching", False) is True
|
||||
) and (
|
||||
kwargs.get("cache", {}).get("no-cache", False) is not True
|
||||
): # allow users to control returning cached responses from the completion function
|
||||
args = args or ()
|
||||
final_embedding_cached_response: Optional[EmbeddingResponse] = None
|
||||
embedding_all_elements_cache_hit: bool = False
|
||||
cached_result: Optional[Any] = None
|
||||
kwargs = kwargs.copy()
|
||||
#########################################################
|
||||
# Init cache timing metrics
|
||||
#########################################################
|
||||
cache_check_start_time = time.perf_counter()
|
||||
cache_check_end_time: Optional[float] = None
|
||||
#########################################################
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
kwargs["parent_otel_span"] = parent_otel_span
|
||||
|
||||
if litellm.cache is not None and self._is_call_type_supported_by_cache(
|
||||
original_function=original_function
|
||||
):
|
||||
|
|
@ -168,7 +170,7 @@ class LLMCachingHandler:
|
|||
kwargs=kwargs,
|
||||
args=args,
|
||||
)
|
||||
cache_check_end_time = datetime.datetime.now()
|
||||
cache_check_end_time = time.perf_counter()
|
||||
|
||||
if cached_result is not None and not isinstance(cached_result, list):
|
||||
verbose_logger.debug("Cache Hit!")
|
||||
|
|
@ -180,7 +182,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
|
||||
cache_duration_ms = (cache_check_end_time - cache_check_start_time) * 1000
|
||||
self._update_litellm_logging_obj_environment(
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
|
|
@ -245,11 +247,14 @@ class LLMCachingHandler:
|
|||
final_embedding_cached_response=final_embedding_cached_response,
|
||||
embedding_all_elements_cache_hit=embedding_all_elements_cache_hit,
|
||||
)
|
||||
verbose_logger.debug(f"CACHE RESULT: {cached_result}")
|
||||
return CachingHandlerResponse(
|
||||
cached_result=cached_result,
|
||||
final_embedding_cached_response=final_embedding_cached_response,
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"CACHE RESULT: {cached_result}")
|
||||
return CachingHandlerResponse(
|
||||
cached_result=cached_result,
|
||||
final_embedding_cached_response=final_embedding_cached_response,
|
||||
)
|
||||
# Caching disabled - return None to indicate no caching attempted
|
||||
return None
|
||||
|
||||
def _sync_get_cache(
|
||||
self,
|
||||
|
|
@ -263,18 +268,22 @@ class LLMCachingHandler:
|
|||
) -> CachingHandlerResponse:
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
args = args or ()
|
||||
new_kwargs = kwargs.copy()
|
||||
new_kwargs.update(
|
||||
convert_args_to_kwargs(
|
||||
self.original_function,
|
||||
args,
|
||||
)
|
||||
)
|
||||
|
||||
cached_result: Optional[Any] = None
|
||||
|
||||
# Check if caching should be performed BEFORE doing expensive kwargs copy
|
||||
if litellm.cache is not None and self._is_call_type_supported_by_cache(
|
||||
original_function=original_function
|
||||
):
|
||||
args = args or ()
|
||||
# Now that we confirmed caching will happen, prepare kwargs
|
||||
new_kwargs = kwargs.copy()
|
||||
new_kwargs.update(
|
||||
convert_args_to_kwargs(
|
||||
self.original_function,
|
||||
args,
|
||||
)
|
||||
)
|
||||
print_verbose("Checking Sync Cache")
|
||||
cached_result = litellm.cache.get_cache(**new_kwargs)
|
||||
if cached_result is not None:
|
||||
|
|
|
|||
|
|
@ -374,7 +374,7 @@ OPENAI_TRANSCRIPTION_PARAMS = [
|
|||
"timestamp_granularities",
|
||||
]
|
||||
|
||||
OPENAI_EMBEDDING_PARAMS = ["dimensions", "encoding_format", "user", "input_type"]
|
||||
OPENAI_EMBEDDING_PARAMS = ["dimensions", "encoding_format", "user"]
|
||||
|
||||
DEFAULT_EMBEDDING_PARAM_VALUES = {
|
||||
**{k: None for k in OPENAI_EMBEDDING_PARAMS},
|
||||
|
|
|
|||
|
|
@ -405,6 +405,7 @@ async def agenerate_content_stream(
|
|||
config=setup_result.generate_content_config_dict,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
tools=tools,
|
||||
stream=True,
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
|
|
@ -485,6 +486,7 @@ def generate_content_stream(
|
|||
config=setup_result.generate_content_config_dict,
|
||||
_is_async=_is_async,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
stream=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -576,7 +576,7 @@ class OpenTelemetry(CustomLogger):
|
|||
return
|
||||
|
||||
litellm_params = kwargs.get("litellm_params", {})
|
||||
metadata = litellm_params.get("metadata", {})
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
generation_name = metadata.get("generation_name")
|
||||
|
||||
raw_span_name = generation_name if generation_name else RAW_REQUEST_SPAN_NAME
|
||||
|
|
@ -1178,7 +1178,7 @@ class OpenTelemetry(CustomLogger):
|
|||
|
||||
def _get_span_name(self, kwargs):
|
||||
litellm_params = kwargs.get("litellm_params", {})
|
||||
metadata = litellm_params.get("metadata", {})
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
generation_name = metadata.get("generation_name")
|
||||
|
||||
if generation_name:
|
||||
|
|
|
|||
|
|
@ -4040,6 +4040,7 @@ class StandardLoggingPayloadSetup:
|
|||
usage_object=usage_object,
|
||||
requester_custom_headers=None,
|
||||
cold_storage_object_key=None,
|
||||
user_api_key_auth_metadata=None,
|
||||
)
|
||||
if isinstance(metadata, dict):
|
||||
# Filter the metadata dictionary to include only the specified keys
|
||||
|
|
@ -4755,6 +4756,7 @@ def get_standard_logging_metadata(
|
|||
requester_custom_headers=None,
|
||||
user_api_key_request_route=None,
|
||||
cold_storage_object_key=None,
|
||||
user_api_key_auth_metadata=None,
|
||||
)
|
||||
if isinstance(metadata, dict):
|
||||
# Update the clean_metadata with values from input metadata that match StandardLoggingMetadata fields
|
||||
|
|
|
|||
|
|
@ -1117,6 +1117,14 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
status_code=422, message="max retries must be an int"
|
||||
)
|
||||
|
||||
if api_key is None and azure_ad_token_provider is not None:
|
||||
azure_ad_token = azure_ad_token_provider()
|
||||
if azure_ad_token:
|
||||
headers.pop(
|
||||
"api-key", None
|
||||
)
|
||||
headers["Authorization"] = f"Bearer {azure_ad_token}"
|
||||
|
||||
# init AzureOpenAI Client
|
||||
azure_client_params: Dict[str, Any] = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
|
|
|
|||
|
|
@ -440,7 +440,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
"""
|
||||
Abbreviations of regions AWS Bedrock supports for cross region inference
|
||||
"""
|
||||
return ["us", "eu", "apac"]
|
||||
return ["us", "eu", "apac", "jp"]
|
||||
|
||||
@staticmethod
|
||||
def get_bedrock_route(
|
||||
|
|
|
|||
|
|
@ -6,9 +6,10 @@ Why separate file? Make it easy to see how transformation works
|
|||
Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo.html
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Union
|
||||
from typing import List, Optional, Union, cast
|
||||
|
||||
from litellm.types.llms.bedrock import (
|
||||
TWELVELABS_EMBEDDING_INPUT_TYPES,
|
||||
TwelveLabsAsyncInvokeRequest,
|
||||
TwelveLabsMarengoEmbeddingRequest,
|
||||
TwelveLabsOutputDataConfig,
|
||||
|
|
@ -89,10 +90,11 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
- Audio inputs (async-invoke only)
|
||||
- S3 URLs for all media types (async-invoke only)
|
||||
"""
|
||||
if inference_params.get("inputType"):
|
||||
input_type = inference_params["inputType"]
|
||||
else:
|
||||
raise ValueError("input_type is required")
|
||||
# Get input_type or default to "text"
|
||||
input_type = cast(
|
||||
TWELVELABS_EMBEDDING_INPUT_TYPES,
|
||||
inference_params.get("inputType") or inference_params.get("input_type") or "text"
|
||||
)
|
||||
|
||||
# Validate that async-invoke is used for video/audio
|
||||
if input_type in ["video", "audio"] and not async_invoke_route:
|
||||
|
|
@ -136,6 +138,7 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
for k, v in inference_params.items():
|
||||
if k not in [
|
||||
"inputType",
|
||||
"input_type", # Exclude both camelCase and snake_case
|
||||
"inputText",
|
||||
"mediaSource",
|
||||
"bucketOwner", # Don't include bucketOwner in the request
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ def validate_environment(
|
|||
"Request-Source": "unspecified:litellm",
|
||||
"accept": "application/json",
|
||||
"content-type": "application/json",
|
||||
"Authorization": "bearer $CO_API_KEY"
|
||||
"Authorization": "Bearer $CO_API_KEY"
|
||||
}
|
||||
"""
|
||||
headers.update(
|
||||
|
|
@ -42,7 +42,7 @@ def validate_environment(
|
|||
}
|
||||
)
|
||||
if api_key:
|
||||
headers["Authorization"] = f"bearer {api_key}"
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return headers
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -86,7 +86,7 @@ class CohereRerankConfig(BaseRerankConfig):
|
|||
)
|
||||
|
||||
default_headers = {
|
||||
"Authorization": f"bearer {api_key}",
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"accept": "application/json",
|
||||
"content-type": "application/json",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -49,7 +49,7 @@ class InfinityRerankConfig(CohereRerankConfig):
|
|||
)
|
||||
|
||||
default_headers = {
|
||||
"Authorization": f"bearer {api_key}",
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"accept": "application/json",
|
||||
"content-type": "application/json",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -98,9 +98,26 @@ class JinaAIRerankConfig(BaseRerankConfig):
|
|||
if _results is None:
|
||||
raise ValueError(f"No results found in the response={_json_response}")
|
||||
|
||||
# Transform Jina AI's response format to match LiteLLM's expected format
|
||||
# Jina AI returns: {"index": 0, "relevance_score": 0.72, "document": "hello"}
|
||||
# LiteLLM expects: {"index": 0, "relevance_score": 0.72, "document": {"text": "hello"}}
|
||||
transformed_results = []
|
||||
for result in _results:
|
||||
transformed_result = {
|
||||
"index": result["index"],
|
||||
"relevance_score": result["relevance_score"],
|
||||
}
|
||||
# Convert document from string to dict format if it exists
|
||||
if "document" in result and isinstance(result["document"], str):
|
||||
transformed_result["document"] = {"text": result["document"]}
|
||||
elif "document" in result:
|
||||
# If it's already a dict, keep it as is
|
||||
transformed_result["document"] = result["document"]
|
||||
transformed_results.append(transformed_result)
|
||||
|
||||
return RerankResponse(
|
||||
id=_json_response.get("id") or str(uuid.uuid4()),
|
||||
results=_results, # type: ignore
|
||||
results=transformed_results, # type: ignore
|
||||
meta=rerank_meta,
|
||||
) # Return response
|
||||
|
||||
|
|
|
|||
|
|
@ -415,8 +415,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
googleSearchRetrieval = self.get_tool_value(tool, VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value)
|
||||
elif tool_name and tool_name == VertexToolName.ENTERPRISE_WEB_SEARCH.value:
|
||||
enterpriseWebSearch = self.get_tool_value(tool, VertexToolName.ENTERPRISE_WEB_SEARCH.value)
|
||||
elif tool_name and tool_name == VertexToolName.URL_CONTEXT.value:
|
||||
urlContext = self.get_tool_value(tool, VertexToolName.URL_CONTEXT.value)
|
||||
elif tool_name and (tool_name == VertexToolName.URL_CONTEXT.value or tool_name == "urlContext"):
|
||||
urlContext = self.get_tool_value(tool, tool_name)
|
||||
elif tool_name and (
|
||||
tool_name == VertexToolName.GOOGLE_MAPS.value or tool_name == "google_maps"
|
||||
):
|
||||
|
|
@ -448,9 +448,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
"Invalid tool={}. Use `litellm.set_verbose` or `litellm --detailed_debug` to see raw request."
|
||||
)
|
||||
|
||||
_tools = Tools(
|
||||
function_declarations=gtool_func_declarations,
|
||||
)
|
||||
# Only include function_declarations if there are actual functions
|
||||
_tools = Tools()
|
||||
if gtool_func_declarations:
|
||||
_tools["function_declarations"] = gtool_func_declarations
|
||||
if googleSearch is not None:
|
||||
_tools[VertexToolName.GOOGLE_SEARCH.value] = googleSearch
|
||||
if googleSearchRetrieval is not None:
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,7 +1,7 @@
|
|||
from litellm._uuid import uuid
|
||||
from typing import Any, Dict, Iterable, List, Optional, Set, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
|
|
@ -30,7 +30,7 @@ def _prepare_mcp_server_data(
|
|||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
# Convert model to dict
|
||||
data_dict = data.model_dump()
|
||||
data_dict = data.model_dump(exclude_none=True)
|
||||
# Ensure alias is always present in the dict (even if None)
|
||||
if "alias" not in data_dict:
|
||||
data_dict["alias"] = getattr(data, "alias", None)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import asyncio
|
|||
import datetime
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
from typing import Any, Dict, List, Optional, Set, Union, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
|
|
@ -240,50 +240,64 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
def add_update_server(self, mcp_server: LiteLLM_MCPServerTable):
|
||||
if mcp_server.server_id not in self.get_registry():
|
||||
_mcp_info: MCPInfo = mcp_server.mcp_info or {}
|
||||
# Use helper to deserialize environment dictionary
|
||||
# Safely access env field which may not exist on Prisma model objects
|
||||
env_data = getattr(mcp_server, "env", None)
|
||||
env_dict = _deserialize_env_dict(env_data)
|
||||
# Use alias for name if present, else server_name
|
||||
name_for_prefix = (
|
||||
mcp_server.alias or mcp_server.server_name or mcp_server.server_id
|
||||
)
|
||||
# Preserve all custom fields from database while setting defaults for core fields
|
||||
mcp_info: MCPInfo = _mcp_info.copy()
|
||||
# Set default values for core fields if not present
|
||||
if "server_name" not in mcp_info:
|
||||
mcp_info["server_name"] = mcp_server.server_name or mcp_server.server_id
|
||||
if "description" not in mcp_info and mcp_server.description:
|
||||
mcp_info["description"] = mcp_server.description
|
||||
try:
|
||||
if mcp_server.server_id not in self.get_registry():
|
||||
_mcp_info: MCPInfo = mcp_server.mcp_info or {}
|
||||
# Use helper to deserialize environment dictionary
|
||||
# Safely access env field which may not exist on Prisma model objects
|
||||
env_data = getattr(mcp_server, "env", None)
|
||||
env_dict = _deserialize_env_dict(env_data)
|
||||
# Use alias for name if present, else server_name
|
||||
name_for_prefix = (
|
||||
mcp_server.alias or mcp_server.server_name or mcp_server.server_id
|
||||
)
|
||||
# Preserve all custom fields from database while setting defaults for core fields
|
||||
mcp_info: MCPInfo = _mcp_info.copy()
|
||||
# Set default values for core fields if not present
|
||||
if "server_name" not in mcp_info:
|
||||
mcp_info["server_name"] = (
|
||||
mcp_server.server_name or mcp_server.server_id
|
||||
)
|
||||
if "description" not in mcp_info and mcp_server.description:
|
||||
mcp_info["description"] = mcp_server.description
|
||||
|
||||
new_server = MCPServer(
|
||||
server_id=mcp_server.server_id,
|
||||
name=name_for_prefix,
|
||||
alias=getattr(mcp_server, "alias", None),
|
||||
server_name=getattr(mcp_server, "server_name", None),
|
||||
url=mcp_server.url,
|
||||
transport=cast(MCPTransportType, mcp_server.transport),
|
||||
auth_type=cast(MCPAuthType, mcp_server.auth_type),
|
||||
mcp_info=mcp_info,
|
||||
extra_headers=getattr(mcp_server, "extra_headers", None),
|
||||
# oauth specific fields
|
||||
client_id=getattr(mcp_server, "client_id", None),
|
||||
client_secret=getattr(mcp_server, "client_secret", None),
|
||||
scopes=getattr(mcp_server, "scopes", None),
|
||||
authorization_url=getattr(mcp_server, "authorization_url", None),
|
||||
token_url=getattr(mcp_server, "token_url", None),
|
||||
# Stdio-specific fields
|
||||
command=getattr(mcp_server, "command", None),
|
||||
args=getattr(mcp_server, "args", None) or [],
|
||||
env=env_dict,
|
||||
access_groups=getattr(mcp_server, "mcp_access_groups", None),
|
||||
allowed_tools=getattr(mcp_server, "allowed_tools", None),
|
||||
disallowed_tools=getattr(mcp_server, "disallowed_tools", None),
|
||||
)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
verbose_logger.debug(f"Added MCP Server: {name_for_prefix}")
|
||||
new_server = MCPServer(
|
||||
server_id=mcp_server.server_id,
|
||||
name=name_for_prefix,
|
||||
alias=getattr(mcp_server, "alias", None),
|
||||
server_name=getattr(mcp_server, "server_name", None),
|
||||
url=mcp_server.url,
|
||||
transport=cast(MCPTransportType, mcp_server.transport),
|
||||
auth_type=cast(MCPAuthType, mcp_server.auth_type),
|
||||
mcp_info=mcp_info,
|
||||
extra_headers=getattr(mcp_server, "extra_headers", None),
|
||||
# oauth specific fields
|
||||
client_id=getattr(mcp_server, "client_id", None),
|
||||
client_secret=getattr(mcp_server, "client_secret", None),
|
||||
scopes=getattr(mcp_server, "scopes", None),
|
||||
authorization_url=getattr(mcp_server, "authorization_url", None),
|
||||
token_url=getattr(mcp_server, "token_url", None),
|
||||
# Stdio-specific fields
|
||||
command=getattr(mcp_server, "command", None),
|
||||
args=getattr(mcp_server, "args", None) or [],
|
||||
env=env_dict,
|
||||
access_groups=getattr(mcp_server, "mcp_access_groups", None),
|
||||
allowed_tools=getattr(mcp_server, "allowed_tools", None),
|
||||
disallowed_tools=getattr(mcp_server, "disallowed_tools", None),
|
||||
)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
verbose_logger.debug(f"Added MCP Server: {name_for_prefix}")
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Failed to add MCP server: {str(e)}")
|
||||
raise e
|
||||
|
||||
def get_all_mcp_server_ids(self) -> Set[str]:
|
||||
"""
|
||||
Get all MCP server IDs
|
||||
"""
|
||||
all_servers = list(self.get_registry().values())
|
||||
return {server.server_id for server in all_servers}
|
||||
|
||||
async def get_allowed_mcp_servers(
|
||||
self, user_api_key_auth: Optional[UserAPIKeyAuth] = None
|
||||
|
|
@ -1118,25 +1132,23 @@ class MCPServerManager:
|
|||
if _server_id in allowed_server_ids:
|
||||
list_mcp_servers.append(
|
||||
LiteLLM_MCPServerTable(
|
||||
server_id=_server_id,
|
||||
server_name=_server_config.name,
|
||||
alias=_server_config.alias,
|
||||
url=_server_config.url,
|
||||
transport=_server_config.transport,
|
||||
auth_type=_server_config.auth_type,
|
||||
created_at=datetime.datetime.now(),
|
||||
updated_at=datetime.datetime.now(),
|
||||
description=(
|
||||
_server_config.mcp_info.get("description")
|
||||
if _server_config.mcp_info
|
||||
else None
|
||||
),
|
||||
mcp_info=_server_config.mcp_info,
|
||||
mcp_access_groups=_server_config.access_groups or [],
|
||||
# Stdio-specific fields
|
||||
command=getattr(_server_config, "command", None),
|
||||
args=getattr(_server_config, "args", None) or [],
|
||||
env=getattr(_server_config, "env", None) or {},
|
||||
**{
|
||||
**_server_config.model_dump(),
|
||||
"created_at": datetime.datetime.now(),
|
||||
"updated_at": datetime.datetime.now(),
|
||||
"description": (
|
||||
_server_config.mcp_info.get("description")
|
||||
if _server_config.mcp_info
|
||||
else None
|
||||
),
|
||||
"allowed_tools": _server_config.allowed_tools or [],
|
||||
"mcp_info": _server_config.mcp_info,
|
||||
"mcp_access_groups": _server_config.access_groups or [],
|
||||
"extra_headers": _server_config.extra_headers or [],
|
||||
"command": getattr(_server_config, "command", None),
|
||||
"args": getattr(_server_config, "args", None) or [],
|
||||
"env": getattr(_server_config, "env", None) or {},
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -1176,44 +1188,19 @@ class MCPServerManager:
|
|||
}
|
||||
)
|
||||
|
||||
# Map servers to their teams and return with health data
|
||||
from typing import cast
|
||||
## mark invalid servers w/ reason for being invalid
|
||||
valid_server_ids = self.get_all_mcp_server_ids()
|
||||
for server in list_mcp_servers:
|
||||
if server.server_id not in valid_server_ids:
|
||||
server.status = "unhealthy"
|
||||
## try adding server to registry to get error
|
||||
try:
|
||||
self.add_update_server(server)
|
||||
except Exception as e:
|
||||
server.health_check_error = str(e)
|
||||
server.health_check_error = "Server is not in in memory registry yet. This could be a temporary sync issue."
|
||||
|
||||
return [
|
||||
LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.server_name,
|
||||
alias=server.alias,
|
||||
description=server.description,
|
||||
url=server.url,
|
||||
transport=server.transport,
|
||||
auth_type=server.auth_type,
|
||||
created_at=server.created_at,
|
||||
created_by=server.created_by,
|
||||
updated_at=server.updated_at,
|
||||
updated_by=server.updated_by,
|
||||
mcp_access_groups=(
|
||||
server.mcp_access_groups
|
||||
if server.mcp_access_groups is not None
|
||||
else []
|
||||
),
|
||||
allowed_tools=(
|
||||
server.allowed_tools
|
||||
if server.allowed_tools is not None
|
||||
else []
|
||||
),
|
||||
mcp_info=server.mcp_info,
|
||||
teams=cast(
|
||||
List[Dict[str, str | None]],
|
||||
server_to_teams_map.get(server.server_id, []),
|
||||
),
|
||||
# Stdio-specific fields
|
||||
command=getattr(server, "command", None),
|
||||
args=getattr(server, "args", None) or [],
|
||||
env=getattr(server, "env", None) or {},
|
||||
)
|
||||
for server in list_mcp_servers
|
||||
]
|
||||
return list_mcp_servers
|
||||
|
||||
async def reload_servers_from_database(self):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
model_list:
|
||||
- model_name: byok-fixed-gpt-4o-mini
|
||||
- model_name: openai/gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o-mini
|
||||
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5"
|
||||
|
|
@ -16,15 +16,18 @@ model_list:
|
|||
api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5"
|
||||
api_key: dummy
|
||||
|
||||
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"]
|
||||
allowed_tools: ["list_tools"]
|
||||
# disallowed_tools: ["repo_delete"]
|
||||
# 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"]
|
||||
# allowed_tools: ["list_tools"]
|
||||
# # disallowed_tools: ["repo_delete"]
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["prometheus"]
|
||||
custom_prometheus_metadata_labels: ["metadata.initiative", "metadata.business-unit"]
|
||||
|
|
@ -731,6 +731,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
|
|||
metadata: Optional[dict] = {}
|
||||
tpm_limit: Optional[int] = None
|
||||
rpm_limit: Optional[int] = None
|
||||
|
||||
budget_duration: Optional[str] = None
|
||||
allowed_cache_controls: Optional[list] = []
|
||||
config: Optional[dict] = {}
|
||||
|
|
@ -755,6 +756,12 @@ class KeyRequestBase(GenerateRequestBase):
|
|||
tags: Optional[List[str]] = None
|
||||
enforced_params: Optional[List[str]] = None
|
||||
allowed_routes: Optional[list] = []
|
||||
rpm_limit_type: Optional[
|
||||
Literal["guaranteed_throughput", "best_effort_throughput"]
|
||||
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating rpm
|
||||
tpm_limit_type: Optional[
|
||||
Literal["guaranteed_throughput", "best_effort_throughput"]
|
||||
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
|
||||
|
||||
|
||||
class LiteLLMKeyType(str, enum.Enum):
|
||||
|
|
@ -918,6 +925,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
mcp_info: Optional[MCPInfo] = None
|
||||
mcp_access_groups: List[str] = Field(default_factory=list)
|
||||
allowed_tools: Optional[List[str]] = None
|
||||
extra_headers: Optional[List[str]] = None
|
||||
# Stdio-specific fields
|
||||
command: Optional[str] = None
|
||||
args: List[str] = Field(default_factory=list)
|
||||
|
|
@ -987,9 +995,10 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
teams: List[Dict[str, Optional[str]]] = Field(default_factory=list)
|
||||
mcp_access_groups: List[str] = Field(default_factory=list)
|
||||
allowed_tools: List[str] = Field(default_factory=list)
|
||||
extra_headers: List[str] = Field(default_factory=list)
|
||||
mcp_info: Optional[MCPInfo] = None
|
||||
# Health check status
|
||||
status: Optional[str] = Field(
|
||||
status: Optional[Literal["healthy", "unhealthy", "unknown"]] = Field(
|
||||
default="unknown",
|
||||
description="Health status: 'healthy', 'unhealthy', 'unknown'",
|
||||
)
|
||||
|
|
@ -3056,6 +3065,8 @@ class PassThroughEndpointLoggingTypedDict(TypedDict):
|
|||
LiteLLM_ManagementEndpoint_MetadataFields = [
|
||||
"model_rpm_limit",
|
||||
"model_tpm_limit",
|
||||
"rpm_limit_type",
|
||||
"tpm_limit_type",
|
||||
"guardrails",
|
||||
"tags",
|
||||
"enforced_params",
|
||||
|
|
@ -3068,6 +3079,7 @@ LiteLLM_ManagementEndpoint_MetadataFields_Premium = [
|
|||
"tags",
|
||||
"team_member_key_duration",
|
||||
"prompts",
|
||||
"logging",
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ from litellm.proxy.common_utils.callback_utils import (
|
|||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.router import Router
|
||||
from litellm.types.utils import ServerToolUse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
|
||||
|
|
@ -897,24 +898,30 @@ class ProxyBaseLLMRequestProcessing:
|
|||
completion_tokens_details = _usage.get("completion_tokens_details")
|
||||
prompt_tokens_details = _usage.get("prompt_tokens_details")
|
||||
|
||||
# Build usage kwargs with only non-None values
|
||||
usage_kwargs = {
|
||||
|
||||
usage_kwargs: dict[str, Any] = {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": total_tokens,
|
||||
}
|
||||
|
||||
# Add optional fields if they exist
|
||||
if cache_creation_input_tokens is not None:
|
||||
usage_kwargs["cache_creation_input_tokens"] = cache_creation_input_tokens
|
||||
if cache_read_input_tokens is not None:
|
||||
usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens
|
||||
if web_search_requests is not None:
|
||||
usage_kwargs["web_search_requests"] = web_search_requests
|
||||
# Add optional named parameters
|
||||
if completion_tokens_details is not None:
|
||||
usage_kwargs["completion_tokens_details"] = completion_tokens_details
|
||||
if prompt_tokens_details is not None:
|
||||
usage_kwargs["prompt_tokens_details"] = prompt_tokens_details
|
||||
|
||||
# Handle web_search_requests by wrapping in ServerToolUse
|
||||
if web_search_requests is not None:
|
||||
usage_kwargs["server_tool_use"] = ServerToolUse(
|
||||
web_search_requests=web_search_requests
|
||||
)
|
||||
|
||||
# Add cache-related fields to **params (handled by Usage.__init__)
|
||||
if cache_creation_input_tokens is not None:
|
||||
usage_kwargs["cache_creation_input_tokens"] = cache_creation_input_tokens
|
||||
if cache_read_input_tokens is not None:
|
||||
usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens
|
||||
|
||||
_mr = ModelResponse(
|
||||
usage=Usage(**usage_kwargs)
|
||||
|
|
|
|||
|
|
@ -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_litellm_params(_litellm_params)) or {}
|
||||
_model_group = _metadata.get("model_group", None) or kwargs.get("model", None)
|
||||
_metadata = _litellm_params.get(get_metadata_variable_name_from_kwargs(kwargs)) or {}
|
||||
_model_group = _metadata.get("model_group", 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_litellm_params(
|
||||
litellm_params: dict
|
||||
def get_metadata_variable_name_from_kwargs(
|
||||
kwargs: 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_litellm_params(
|
|||
- 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 litellm_params else "metadata"
|
||||
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ def decrypt_value_helper(
|
|||
value: str,
|
||||
key: str, # this is just for debug purposes, showing the k,v pair that's invalid. not a signing key.
|
||||
exception_type: Literal["debug", "error"] = "error",
|
||||
return_original_value: bool = False,
|
||||
):
|
||||
signing_key = _get_salt_key()
|
||||
|
||||
|
|
@ -55,14 +56,14 @@ def decrypt_value_helper(
|
|||
error_message = f"Error decrypting value for key: {key}, Did your master_key/salt key change recently? \nError: {str(e)}\nSet permanent salt key - https://docs.litellm.ai/docs/proxy/prod#5-set-litellm-salt-key"
|
||||
if exception_type == "debug":
|
||||
verbose_proxy_logger.debug(error_message)
|
||||
return None
|
||||
return value if return_original_value else None
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Unable to decrypt value={value} for key: {key}, returning None"
|
||||
)
|
||||
verbose_proxy_logger.exception(error_message)
|
||||
# [Non-Blocking Exception. - this should not block decrypting other values]
|
||||
return None
|
||||
return value if return_original_value else None
|
||||
|
||||
|
||||
def encrypt_value(value: str, signing_key: str):
|
||||
|
|
|
|||
|
|
@ -24,11 +24,19 @@ model_list:
|
|||
litellm_params:
|
||||
model: anthropic/*
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
- model_name: openai/*
|
||||
litellm_params:
|
||||
model: openai/*
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
custom_auth: custom_auth_basic.user_api_key_auth
|
||||
pass_through_endpoints:
|
||||
- path: "/azure-config-passthrough"
|
||||
target: os.environ/AZURE_API_BASE
|
||||
target: os.environ/AZURE_API_BASE_PASSHROUGH
|
||||
include_subpath: true
|
||||
headers:
|
||||
Authorization: os.environ/AZURE_API_KEY
|
||||
Authorization: os.environ/AZURE_API_KEY_PASSHROUGH
|
||||
|
||||
litellm_settings:
|
||||
include_cost_in_streaming_usage: true
|
||||
|
|
@ -727,14 +727,21 @@ 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
|
||||
# Check if the ModelResponse has text content in its choices
|
||||
# to avoid sending empty content to Bedrock (e.g., during tool calls)
|
||||
if isinstance(response, litellm.ModelResponse):
|
||||
has_text_content = False
|
||||
for choice in response.choices:
|
||||
if isinstance(choice, litellm.Choices):
|
||||
if choice.message.content and isinstance(choice.message.content, str):
|
||||
has_text_content = True
|
||||
break
|
||||
|
||||
if not has_text_content:
|
||||
verbose_proxy_logger.warning(
|
||||
"Bedrock AI: not running guardrail. No output text in response"
|
||||
)
|
||||
return
|
||||
|
||||
#########################################################
|
||||
########## 1. Make parallel Bedrock API requests ##########
|
||||
|
|
|
|||
|
|
@ -27,15 +27,19 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
Saturation-aware priority-based rate limiter using v3 infrastructure.
|
||||
|
||||
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)
|
||||
1. Model capacity ALWAYS enforced at 100% (prevents over-allocation)
|
||||
2. Priority usage tracked from first request (accurate accounting)
|
||||
3. Priority limits only enforced when saturated >= threshold
|
||||
4. Three-phase checking prevents partial counter increments
|
||||
5. Reuses v3 limiter's Redis-based tracking (multi-instance safe)
|
||||
|
||||
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
|
||||
- Phase 1: Read-only check of ALL limits (no increments)
|
||||
- Phase 2: Decide enforcement based on saturation
|
||||
- Phase 3: Increment counters only if request allowed
|
||||
- When under-saturated: priorities can borrow unused capacity (generous)
|
||||
- When saturated: strict priority-based limits enforced (fair)
|
||||
- Uses v3 limiter's atomic Lua scripts for race-free increments
|
||||
"""
|
||||
def __init__(self, internal_usage_cache: DualCache):
|
||||
self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache)
|
||||
|
|
@ -84,6 +88,46 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
|
||||
return weights
|
||||
|
||||
def _get_priority_allocation(
|
||||
self,
|
||||
model: str,
|
||||
priority: Optional[str],
|
||||
normalized_weights: Dict[str, float],
|
||||
) -> tuple[float, str]:
|
||||
"""
|
||||
Get priority weight and pool key for a given priority.
|
||||
|
||||
For explicit priorities: returns specific allocation and unique pool key
|
||||
For default priority: returns default allocation and shared pool key
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
priority: Priority level (None for default)
|
||||
normalized_weights: Pre-computed normalized weights
|
||||
|
||||
Returns:
|
||||
tuple: (priority_weight, priority_key)
|
||||
"""
|
||||
# Check if this key has an explicit priority in litellm.priority_reservation
|
||||
has_explicit_priority = (
|
||||
priority is not None
|
||||
and litellm.priority_reservation is not None
|
||||
and priority in litellm.priority_reservation
|
||||
)
|
||||
|
||||
if has_explicit_priority and priority is not None:
|
||||
# Explicit priority: get its specific allocation
|
||||
priority_weight = normalized_weights.get(priority, self._get_priority_weight(priority))
|
||||
# Use unique key per priority level
|
||||
priority_key = f"{model}:{priority}"
|
||||
else:
|
||||
# No explicit priority: share the default_priority pool with ALL other default keys
|
||||
priority_weight = litellm.priority_reservation_settings.default_priority
|
||||
# Use shared key for all default-priority requests
|
||||
priority_key = f"{model}:default_pool"
|
||||
|
||||
return priority_weight, priority_key
|
||||
|
||||
async def _check_model_saturation(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -174,7 +218,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
Create rate limit descriptors with normalized priority weights.
|
||||
|
||||
Uses normalized weights to handle over-allocation scenarios.
|
||||
Only called when system is saturated.
|
||||
|
||||
For explicit priorities: each priority gets its own pool (e.g., prod gets 75%)
|
||||
For default priority: ALL keys without explicit priority share ONE pool (e.g., all share 25%)
|
||||
"""
|
||||
descriptors: List[RateLimitDescriptor] = []
|
||||
|
||||
|
|
@ -185,31 +231,24 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
if model_group_info is None:
|
||||
return descriptors
|
||||
|
||||
# Get normalized priority weight (handles over-allocation)
|
||||
# Get normalized priority weight and pool key
|
||||
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
|
||||
priority_key = f"{model}:{priority or 'default'}"
|
||||
priority_weight, priority_key = self._get_priority_allocation(
|
||||
model=model,
|
||||
priority=priority,
|
||||
normalized_weights=normalized_weights,
|
||||
)
|
||||
|
||||
rate_limit_config: RateLimitDescriptorRateLimitObject = {}
|
||||
|
||||
# Apply normalized priority weight to model limits
|
||||
# Apply priority weight to model limits
|
||||
if model_group_info.tpm is not None:
|
||||
# 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 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
|
||||
|
|
@ -257,58 +296,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
},
|
||||
)
|
||||
|
||||
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(
|
||||
async def _check_rate_limits(
|
||||
self,
|
||||
model: str,
|
||||
model_group_info: ModelGroupInfo,
|
||||
|
|
@ -318,9 +307,23 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
data: dict,
|
||||
) -> None:
|
||||
"""
|
||||
Handle rate limiting in strict mode (above saturation threshold).
|
||||
Check rate limits using THREE-PHASE approach to prevent partial increments.
|
||||
|
||||
In this mode, we enforce priority-specific limits using normalized weights.
|
||||
Phase 1: Read-only check of ALL limits (no increments)
|
||||
Phase 2: Decide which limits to enforce based on saturation
|
||||
Phase 3: Increment ALL counters atomically (model + priority)
|
||||
|
||||
This prevents the bug where:
|
||||
- Model counter increments in stage 1
|
||||
- Priority check fails in stage 2
|
||||
- Request blocked but model counter already incremented
|
||||
|
||||
Key behaviors:
|
||||
- All checks performed first (read-only)
|
||||
- Only increment counters if request will be allowed
|
||||
- Model capacity: Always enforced at 100%
|
||||
- Priority limits: Only enforced when saturated >= threshold
|
||||
- Both counters tracked from first request (accurate accounting)
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
|
|
@ -331,63 +334,115 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
data: Request data dictionary
|
||||
|
||||
Raises:
|
||||
HTTPException: If priority-specific limit is exceeded
|
||||
HTTPException: If any limit is exceeded
|
||||
"""
|
||||
# Create priority-based descriptors
|
||||
descriptors = self._create_priority_based_descriptors(
|
||||
import json
|
||||
saturation_threshold = litellm.priority_reservation_settings.saturation_threshold
|
||||
should_enforce_priority = saturation >= saturation_threshold
|
||||
|
||||
# Build ALL descriptors upfront
|
||||
descriptors_to_check: List[RateLimitDescriptor] = []
|
||||
|
||||
# Model-wide descriptor (always enforce)
|
||||
model_wide_descriptor = self._create_model_tracking_descriptor(
|
||||
model=model,
|
||||
model_group_info=model_group_info,
|
||||
high_limit_multiplier=1,
|
||||
)
|
||||
descriptors_to_check.append(model_wide_descriptor)
|
||||
|
||||
# Priority descriptors (always track, conditionally enforce)
|
||||
priority_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,
|
||||
if priority_descriptors:
|
||||
descriptors_to_check.extend(priority_descriptors)
|
||||
|
||||
# PHASE 1: Read-only check of ALL limits (no increments)
|
||||
check_response = await self.v3_limiter.should_rate_limit(
|
||||
descriptors=descriptors_to_check,
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
read_only=True, # CRITICAL: Don't increment counters yet
|
||||
)
|
||||
|
||||
await self.v3_limiter.should_rate_limit(
|
||||
descriptors=[tracking_descriptor],
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
)
|
||||
verbose_proxy_logger.debug(f"Read-only check: {json.dumps(check_response, indent=2)}")
|
||||
|
||||
# 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"]:
|
||||
# PHASE 2: Decide which limits to enforce
|
||||
if check_response["overall_code"] == "OVER_LIMIT":
|
||||
for status in check_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%}",
|
||||
},
|
||||
)
|
||||
descriptor_key = status["descriptor_key"]
|
||||
|
||||
# Model-wide limit exceeded (ALWAYS enforce)
|
||||
if descriptor_key == "model_saturation_check":
|
||||
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",
|
||||
},
|
||||
)
|
||||
|
||||
# Priority limit exceeded (ONLY enforce when saturated)
|
||||
elif descriptor_key == "priority_model" and should_enforce_priority:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Enforcing priority limits for {model}, saturation: {saturation:.1%}, "
|
||||
f"priority: {key_priority}"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": f"Priority-based rate limit exceeded. "
|
||||
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%}",
|
||||
},
|
||||
)
|
||||
|
||||
# PHASE 3: Increment counters separately to avoid early-exit issues
|
||||
# Model counter must ALWAYS increment, but priority counter might be over limit
|
||||
# If we increment them together, v3_limiter's in-memory check will exit early
|
||||
# and skip incrementing the model counter
|
||||
|
||||
# Step 3a: Increment model-wide counter (always)
|
||||
model_increment_response = await self.v3_limiter.should_rate_limit(
|
||||
descriptors=[model_wide_descriptor],
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
read_only=False,
|
||||
)
|
||||
|
||||
# Step 3b: Increment priority counter (may be over limit, but we still track it)
|
||||
if priority_descriptors:
|
||||
priority_increment_response = await self.v3_limiter.should_rate_limit(
|
||||
descriptors=priority_descriptors,
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
read_only=False,
|
||||
)
|
||||
|
||||
# Combine responses for post-call hook
|
||||
combined_response = {
|
||||
"overall_code": model_increment_response["overall_code"],
|
||||
"statuses": model_increment_response["statuses"] + priority_increment_response["statuses"]
|
||||
}
|
||||
data["litellm_proxy_rate_limit_response"] = combined_response
|
||||
else:
|
||||
# Store response for post-call hook
|
||||
data["litellm_proxy_rate_limit_response"] = response
|
||||
data["litellm_proxy_rate_limit_response"] = model_increment_response
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
|
|
@ -409,9 +464,27 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
"""
|
||||
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
|
||||
Flow:
|
||||
1. Check current saturation level
|
||||
2. THREE-PHASE rate limit check:
|
||||
- PHASE 1: Read-only check of ALL limits (no increments)
|
||||
- PHASE 2: Decide which limits to enforce based on saturation
|
||||
- PHASE 3: Increment ALL counters atomically if request allowed
|
||||
|
||||
This three-phase approach ensures:
|
||||
- Model capacity is NEVER exceeded (always enforced at 100%)
|
||||
- Priority usage tracked from first request (accurate metrics)
|
||||
- Counters only increment when request will be allowed (prevents phantom usage)
|
||||
- When under-saturated: priorities can borrow unused capacity (generous)
|
||||
- When saturated: fair allocation based on normalized priority weights (strict)
|
||||
|
||||
Example with 100 RPM model, 60% priority allocation, 80% threshold:
|
||||
- Saturation < 80%: Priority can use up to 100 RPM (model limit enforced only)
|
||||
- Saturation >= 80%: Priority limited to 60 RPM (both limits enforced)
|
||||
|
||||
Prevents bugs where:
|
||||
- Model counter increments but priority check fails → model over-capacity
|
||||
- Priority counter increments but not enforced → inaccurate metrics
|
||||
|
||||
Args:
|
||||
user_api_key_dict: User authentication and metadata
|
||||
|
|
@ -436,8 +509,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
verbose_proxy_logger.debug(f"No model group info for {model}, allowing request")
|
||||
return None
|
||||
|
||||
# Check current saturation level
|
||||
try:
|
||||
# STEP 1: Check current saturation level
|
||||
saturation = await self._check_model_saturation(model, model_group_info)
|
||||
|
||||
saturation_threshold = litellm.priority_reservation_settings.saturation_threshold
|
||||
|
|
@ -449,23 +522,19 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
|
||||
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:
|
||||
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,
|
||||
)
|
||||
# STEP 2: Check rate limits in THREE phases
|
||||
# Phase 1: Read-only check of ALL limits (no increments)
|
||||
# Phase 2: Decide which limits to enforce (based on saturation)
|
||||
# Phase 3: Increment ALL counters only if request will be allowed
|
||||
# This prevents partial increments and ensures accurate tracking
|
||||
await self._check_rate_limits(
|
||||
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
|
||||
|
|
|
|||
|
|
@ -27,7 +27,6 @@ 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
|
||||
|
|
@ -414,6 +413,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
Check if any of the rate limit descriptors should be rate limited.
|
||||
Returns a RateLimitResponse with the overall code and status for each descriptor.
|
||||
Uses batch operations for Redis to improve performance.
|
||||
|
||||
Args:
|
||||
descriptors: List of rate limit descriptors to check
|
||||
parent_otel_span: Optional OpenTelemetry span for tracing
|
||||
read_only: If True, only check limits without incrementing counters
|
||||
"""
|
||||
|
||||
now = datetime.now().timestamp()
|
||||
|
|
@ -486,8 +490,22 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
if rate_limit_response["overall_code"] == "OVER_LIMIT":
|
||||
return rate_limit_response
|
||||
|
||||
## IF under limit, check Redis
|
||||
if self.batch_rate_limiter_script is not None:
|
||||
## IF under limit in-memory, check Redis
|
||||
if read_only:
|
||||
# READ-ONLY MODE: Just read current values without incrementing
|
||||
cache_values = await self.internal_usage_cache.async_batch_get_cache(
|
||||
keys=keys_to_fetch,
|
||||
parent_otel_span=parent_otel_span,
|
||||
local_only=False, # Check Redis too
|
||||
)
|
||||
|
||||
# For keys that don't exist yet, set them to 0
|
||||
if cache_values is None:
|
||||
cache_values = []
|
||||
for _ in keys_to_fetch:
|
||||
cache_values.append(str(now_int) if _.endswith(":window") else 0)
|
||||
elif self.batch_rate_limiter_script is not None:
|
||||
# NORMAL MODE: Increment counters in Redis
|
||||
# Group keys by hash tag for Redis cluster compatibility
|
||||
cache_values = await self._execute_redis_batch_rate_limiter_script(
|
||||
keys_to_fetch=keys_to_fetch,
|
||||
|
|
@ -515,6 +533,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
local_only=True,
|
||||
)
|
||||
else:
|
||||
# NORMAL MODE: In-memory sliding window (no Redis)
|
||||
cache_values = await self.in_memory_cache_sliding_window(
|
||||
keys=keys_to_fetch,
|
||||
now_int=now_int,
|
||||
|
|
@ -846,7 +865,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_litellm_params,
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_model_group_from_litellm_kwargs,
|
||||
)
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
|
@ -864,7 +883,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
# Get metadata from kwargs
|
||||
litellm_metadata = kwargs["litellm_params"].get(
|
||||
get_metadata_variable_name_from_litellm_params(kwargs["litellm_params"]), {}
|
||||
get_metadata_variable_name_from_kwargs(kwargs), {}
|
||||
)
|
||||
if litellm_metadata is None:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -51,7 +51,11 @@ class _ProxyDBLogger(CustomLogger):
|
|||
user_api_key_alias=user_api_key_dict.key_alias,
|
||||
user_api_key_spend=user_api_key_dict.spend,
|
||||
user_api_key_max_budget=user_api_key_dict.max_budget,
|
||||
user_api_key_budget_reset_at=user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None,
|
||||
user_api_key_budget_reset_at=(
|
||||
user_api_key_dict.budget_reset_at.isoformat()
|
||||
if user_api_key_dict.budget_reset_at
|
||||
else None
|
||||
),
|
||||
user_api_key_user_email=user_api_key_dict.user_email,
|
||||
user_api_key_user_id=user_api_key_dict.user_id,
|
||||
user_api_key_team_id=user_api_key_dict.team_id,
|
||||
|
|
@ -59,15 +63,16 @@ class _ProxyDBLogger(CustomLogger):
|
|||
user_api_key_team_alias=user_api_key_dict.team_alias,
|
||||
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
||||
user_api_key_request_route=user_api_key_dict.request_route,
|
||||
user_api_key_auth_metadata=user_api_key_dict.metadata,
|
||||
)
|
||||
)
|
||||
_metadata["user_api_key"] = user_api_key_dict.api_key
|
||||
_metadata["status"] = "failure"
|
||||
_metadata[
|
||||
"error_information"
|
||||
] = StandardLoggingPayloadSetup.get_error_information(
|
||||
original_exception=original_exception,
|
||||
traceback_str=traceback_str,
|
||||
_metadata["error_information"] = (
|
||||
StandardLoggingPayloadSetup.get_error_information(
|
||||
original_exception=original_exception,
|
||||
traceback_str=traceback_str,
|
||||
)
|
||||
)
|
||||
|
||||
existing_metadata: dict = request_data.get("metadata", None) or {}
|
||||
|
|
|
|||
|
|
@ -579,7 +579,12 @@ class LiteLLMProxyRequestSetup:
|
|||
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
||||
user_api_key_user_email=user_api_key_dict.user_email,
|
||||
user_api_key_request_route=user_api_key_dict.request_route,
|
||||
user_api_key_budget_reset_at=user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None,
|
||||
user_api_key_budget_reset_at=(
|
||||
user_api_key_dict.budget_reset_at.isoformat()
|
||||
if user_api_key_dict.budget_reset_at
|
||||
else None
|
||||
),
|
||||
user_api_key_auth_metadata=None,
|
||||
)
|
||||
return user_api_key_logged_metadata
|
||||
|
||||
|
|
@ -607,6 +612,39 @@ class LiteLLMProxyRequestSetup:
|
|||
)
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def add_management_endpoint_metadata_to_request_metadata(
|
||||
data: dict,
|
||||
management_endpoint_metadata: dict,
|
||||
_metadata_variable_name: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Adds the `UserAPIKeyAuth` metadata to the request metadata.
|
||||
|
||||
ignore any sensitive fields like logging, api_key, etc.
|
||||
"""
|
||||
if _metadata_variable_name not in data:
|
||||
return data
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
)
|
||||
|
||||
# ignore any special fields
|
||||
added_metadata = {}
|
||||
for k, v in management_endpoint_metadata.items():
|
||||
if k not in (
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium
|
||||
+ LiteLLM_ManagementEndpoint_MetadataFields
|
||||
):
|
||||
added_metadata[k] = v
|
||||
if data[_metadata_variable_name].get("user_api_key_auth_metadata") is None:
|
||||
data[_metadata_variable_name]["user_api_key_auth_metadata"] = {}
|
||||
data[_metadata_variable_name]["user_api_key_auth_metadata"].update(
|
||||
added_metadata
|
||||
)
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def add_key_level_controls(
|
||||
key_metadata: Optional[dict], data: dict, _metadata_variable_name: str
|
||||
|
|
@ -651,6 +689,13 @@ class LiteLLMProxyRequestSetup:
|
|||
key_metadata["disable_fallbacks"], bool
|
||||
):
|
||||
data["disable_fallbacks"] = key_metadata["disable_fallbacks"]
|
||||
|
||||
## KEY-LEVEL METADATA
|
||||
data = LiteLLMProxyRequestSetup.add_management_endpoint_metadata_to_request_metadata(
|
||||
data=data,
|
||||
management_endpoint_metadata=key_metadata,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
)
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -889,6 +934,15 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
"spend_logs_metadata"
|
||||
]
|
||||
|
||||
## TEAM-LEVEL METADATA
|
||||
data = (
|
||||
LiteLLMProxyRequestSetup.add_management_endpoint_metadata_to_request_metadata(
|
||||
data=data,
|
||||
management_endpoint_metadata=team_metadata,
|
||||
_metadata_variable_name=_metadata_variable_name,
|
||||
)
|
||||
)
|
||||
|
||||
# Team spend, budget - used by prometheus.py
|
||||
data[_metadata_variable_name][
|
||||
"user_api_key_team_max_budget"
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ def _set_object_metadata_field(
|
|||
value: Value to set for the field
|
||||
"""
|
||||
if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
|
||||
_premium_user_check()
|
||||
_premium_user_check(field_name)
|
||||
object_data.metadata = object_data.metadata or {}
|
||||
object_data.metadata[field_name] = value
|
||||
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm.caching import DualCache
|
|||
from litellm.constants import LENGTH_OF_LITELLM_GENERATED_KEY, UI_SESSION_TOKEN_TEAM_ID
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_cache_key_object,
|
||||
_delete_cache_key_object,
|
||||
|
|
@ -90,10 +91,10 @@ def _get_user_in_team(
|
|||
def _calculate_key_rotation_time(rotation_interval: str) -> datetime:
|
||||
"""
|
||||
Helper function to calculate the next rotation time for a key based on the rotation interval.
|
||||
|
||||
|
||||
Args:
|
||||
rotation_interval: String representing the rotation interval (e.g., '30d', '90d', '1h')
|
||||
|
||||
|
||||
Returns:
|
||||
datetime: The calculated next rotation time in UTC
|
||||
"""
|
||||
|
|
@ -102,28 +103,34 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime:
|
|||
return now + timedelta(seconds=interval_seconds)
|
||||
|
||||
|
||||
def _set_key_rotation_fields(data: dict, auto_rotate: bool, rotation_interval: Optional[str]) -> None:
|
||||
def _set_key_rotation_fields(
|
||||
data: dict, auto_rotate: bool, rotation_interval: Optional[str]
|
||||
) -> None:
|
||||
"""
|
||||
Helper function to set rotation fields in key data if auto_rotate is enabled.
|
||||
|
||||
|
||||
Args:
|
||||
data: Dictionary to update with rotation fields
|
||||
auto_rotate: Whether auto rotation is enabled
|
||||
rotation_interval: The rotation interval string (required if auto_rotate is True)
|
||||
"""
|
||||
if auto_rotate and rotation_interval:
|
||||
data.update({
|
||||
"auto_rotate": auto_rotate,
|
||||
"rotation_interval": rotation_interval,
|
||||
"key_rotation_at": _calculate_key_rotation_time(rotation_interval)
|
||||
})
|
||||
data.update(
|
||||
{
|
||||
"auto_rotate": auto_rotate,
|
||||
"rotation_interval": rotation_interval,
|
||||
"key_rotation_at": _calculate_key_rotation_time(rotation_interval),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _is_allowed_to_make_key_request(
|
||||
user_api_key_dict: UserAPIKeyAuth, user_id: Optional[str], team_id: Optional[str]
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
user_id: Optional[str],
|
||||
team_id: Optional[str],
|
||||
) -> bool:
|
||||
"""
|
||||
Assert user only creates keys for themselves
|
||||
Assert user only creates/updates keys for themselves
|
||||
|
||||
Relevant issue: https://github.com/BerriAI/litellm/issues/7336
|
||||
"""
|
||||
|
|
@ -332,6 +339,7 @@ def common_key_access_checks(
|
|||
data: Union[GenerateKeyRequest, UpdateKeyRequest],
|
||||
llm_router: Optional[Router],
|
||||
premium_user: bool,
|
||||
user_id: Optional[str] = None,
|
||||
) -> Literal[True]:
|
||||
"""
|
||||
Check if user is allowed to make a key request, for this key
|
||||
|
|
@ -339,7 +347,7 @@ def common_key_access_checks(
|
|||
try:
|
||||
_is_allowed_to_make_key_request(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
user_id=data.user_id,
|
||||
user_id=user_id or data.user_id,
|
||||
team_id=data.team_id,
|
||||
)
|
||||
except AssertionError as e:
|
||||
|
|
@ -542,6 +550,15 @@ async def _common_key_generation_helper( # noqa: PLR0915
|
|||
value=getattr(data, field),
|
||||
)
|
||||
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields:
|
||||
if getattr(data, field, None) is not None:
|
||||
_set_object_metadata_field(
|
||||
object_data=data,
|
||||
field_name=field,
|
||||
value=getattr(data, field),
|
||||
)
|
||||
delattr(data, field)
|
||||
|
||||
data_json = data.model_dump(exclude_unset=True, exclude_none=True) # type: ignore
|
||||
|
||||
data_json = handle_key_type(data, data_json)
|
||||
|
|
@ -620,6 +637,153 @@ async def _common_key_generation_helper( # noqa: PLR0915
|
|||
return response
|
||||
|
||||
|
||||
def check_team_key_model_specific_limits(
|
||||
keys: List[LiteLLM_VerificationToken],
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
data: Union[GenerateKeyRequest, UpdateKeyRequest],
|
||||
) -> None:
|
||||
"""
|
||||
Check if the team key is allocating model specific limits. If so, raise an error if we're overallocating.
|
||||
"""
|
||||
if data.model_rpm_limit is None and data.model_tpm_limit is None:
|
||||
return
|
||||
# get total model specific tpm/rpm limit
|
||||
model_specific_rpm_limit: Dict[str, int] = {}
|
||||
model_specific_tpm_limit: Dict[str, int] = {}
|
||||
|
||||
for key in keys:
|
||||
if key.metadata.get("model_rpm_limit", None) is not None:
|
||||
for model, rpm_limit in key.metadata.get("model_rpm_limit", {}).items():
|
||||
model_specific_rpm_limit[model] = (
|
||||
model_specific_rpm_limit.get(model, 0) + rpm_limit
|
||||
)
|
||||
if key.metadata.get("model_tpm_limit", None) is not None:
|
||||
for model, tpm_limit in key.metadata.get("model_tpm_limit", {}).items():
|
||||
model_specific_tpm_limit[model] = (
|
||||
model_specific_tpm_limit.get(model, 0) + tpm_limit
|
||||
)
|
||||
if data.model_rpm_limit is not None:
|
||||
for model, rpm_limit in data.model_rpm_limit.items():
|
||||
if (
|
||||
model_specific_rpm_limit.get(model, 0) + rpm_limit
|
||||
> team_table.rpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated RPM limit={model_specific_rpm_limit.get(model, 0)} + Key RPM limit={rpm_limit} is greater than team RPM limit={team_table.rpm_limit}",
|
||||
)
|
||||
elif team_table.metadata and team_table.metadata.get("model_rpm_limit"):
|
||||
team_model_specific_rpm_limit_dict = team_table.metadata.get(
|
||||
"model_rpm_limit", {}
|
||||
)
|
||||
team_model_specific_rpm_limit = team_model_specific_rpm_limit_dict.get(
|
||||
model
|
||||
)
|
||||
if (
|
||||
model_specific_rpm_limit.get(model, 0) + rpm_limit
|
||||
> team_model_specific_rpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated RPM limit={model_specific_rpm_limit.get(model, 0)} + Key RPM limit={rpm_limit} is greater than team RPM limit={team_model_specific_rpm_limit.get(model, 0)}",
|
||||
)
|
||||
if data.model_tpm_limit is not None:
|
||||
for model, tpm_limit in data.model_tpm_limit.items():
|
||||
if (
|
||||
team_table.tpm_limit is not None
|
||||
and model_specific_tpm_limit.get(model, 0) + tpm_limit
|
||||
> team_table.tpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated TPM limit={model_specific_tpm_limit.get(model, 0)} + Key TPM limit={tpm_limit} is greater than team TPM limit={team_table.tpm_limit}",
|
||||
)
|
||||
elif team_table.metadata and team_table.metadata.get("model_tpm_limit"):
|
||||
team_model_specific_tpm_limit_dict = team_table.metadata.get(
|
||||
"model_tpm_limit", {}
|
||||
)
|
||||
team_model_specific_tpm_limit = team_model_specific_tpm_limit_dict.get(
|
||||
model
|
||||
)
|
||||
if (
|
||||
team_model_specific_tpm_limit
|
||||
and model_specific_tpm_limit.get(model, 0) + tpm_limit
|
||||
> team_model_specific_tpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated TPM limit={model_specific_tpm_limit.get(model, 0)} + Key TPM limit={tpm_limit} is greater than team TPM limit={team_model_specific_tpm_limit}",
|
||||
)
|
||||
|
||||
|
||||
def check_team_key_rpm_tpm_limits(
|
||||
keys: List[LiteLLM_VerificationToken],
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
data: Union[GenerateKeyRequest, UpdateKeyRequest],
|
||||
) -> None:
|
||||
"""
|
||||
Check if the team key is allocating rpm/tpm limits. If so, raise an error if we're overallocating.
|
||||
"""
|
||||
if keys is not None and len(keys) > 0:
|
||||
allocated_tpm = sum(key.tpm_limit for key in keys if key.tpm_limit is not None)
|
||||
allocated_rpm = sum(key.rpm_limit for key in keys if key.rpm_limit is not None)
|
||||
else:
|
||||
allocated_tpm = 0
|
||||
allocated_rpm = 0
|
||||
if (
|
||||
data.tpm_limit is not None
|
||||
and team_table.tpm_limit is not None
|
||||
and data.tpm_limit + allocated_tpm > team_table.tpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated TPM limit={allocated_tpm} + Key TPM limit={data.tpm_limit} is greater than team TPM limit={team_table.tpm_limit}",
|
||||
)
|
||||
if (
|
||||
data.rpm_limit is not None
|
||||
and team_table.rpm_limit is not None
|
||||
and data.rpm_limit + allocated_rpm > team_table.rpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Allocated RPM limit={allocated_rpm} + Key RPM limit={data.rpm_limit} is greater than team RPM limit={team_table.rpm_limit}",
|
||||
)
|
||||
|
||||
|
||||
async def _check_team_key_limits(
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
data: Union[GenerateKeyRequest, UpdateKeyRequest],
|
||||
prisma_client: PrismaClient,
|
||||
) -> None:
|
||||
"""
|
||||
Check if the team key is allocating guaranteed throughput limits. If so, raise an error if we're overallocating.
|
||||
|
||||
Only runs check if tpm_limit_type or rpm_limit_type is "guaranteed_throughput"
|
||||
"""
|
||||
if (
|
||||
data.tpm_limit_type != "guaranteed_throughput"
|
||||
and data.rpm_limit_type != "guaranteed_throughput"
|
||||
):
|
||||
return
|
||||
# get all team keys
|
||||
# calculate allocated tpm/rpm limit
|
||||
# check if specified tpm/rpm limit is greater than allocated tpm/rpm limit
|
||||
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"team_id": team_table.team_id},
|
||||
)
|
||||
check_team_key_model_specific_limits(
|
||||
keys=keys,
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
)
|
||||
check_team_key_rpm_tpm_limits(
|
||||
keys=keys,
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/key/generate",
|
||||
tags=["key management"],
|
||||
|
|
@ -661,6 +825,8 @@ async def generate_key_fn(
|
|||
- model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget.
|
||||
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
|
||||
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
|
||||
- tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm). Defaults to "best_effort_throughput".
|
||||
- rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm). Defaults to "best_effort_throughput".
|
||||
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request
|
||||
- blocked: Optional[bool] - Whether the key is blocked.
|
||||
- rpm_limit: Optional[int] - Specify rpm limit for a given key (Requests per minute)
|
||||
|
|
@ -696,12 +862,19 @@ async def generate_key_fn(
|
|||
- user_id: (str) Unique user id - used for tracking spend across multiple keys for same user id.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
user_custom_key_generate,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("entered /key/generate")
|
||||
|
||||
if user_custom_key_generate is not None:
|
||||
|
|
@ -729,7 +902,6 @@ async def generate_key_fn(
|
|||
verbose_proxy_logger.debug(
|
||||
f"Error getting team object in `/key/generate`: {e}"
|
||||
)
|
||||
team_table = None
|
||||
|
||||
key_generation_check(
|
||||
team_table=team_table,
|
||||
|
|
@ -738,12 +910,20 @@ async def generate_key_fn(
|
|||
route=KeyManagementRoutes.KEY_GENERATE,
|
||||
)
|
||||
|
||||
if team_table is not None:
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
return await _common_key_generation_helper(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
team_table=team_table,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.generate_key_fn(): Exception occured - {}".format(
|
||||
|
|
@ -797,6 +977,8 @@ async def generate_service_account_key_fn(
|
|||
- model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget.
|
||||
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
|
||||
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
|
||||
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput" or "guaranteed_throughput"
|
||||
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput" or "guaranteed_throughput"
|
||||
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request
|
||||
- blocked: Optional[bool] - Whether the key is blocked.
|
||||
- rpm_limit: Optional[int] - Specify rpm limit for a given key (Requests per minute)
|
||||
|
|
@ -825,12 +1007,19 @@ async def generate_service_account_key_fn(
|
|||
- user_id: (str) Unique user id - used for tracking spend across multiple keys for same user id.
|
||||
|
||||
"""
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
user_custom_key_generate,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
await validate_team_id_used_in_service_account_request(
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -863,6 +1052,13 @@ async def generate_service_account_key_fn(
|
|||
)
|
||||
team_table = None
|
||||
|
||||
if team_table is not None:
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
key_generation_check(
|
||||
team_table=team_table,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -903,7 +1099,7 @@ def prepare_metadata_fields(
|
|||
if k in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
|
||||
from litellm.proxy.utils import _premium_user_check
|
||||
|
||||
_premium_user_check()
|
||||
_premium_user_check(k)
|
||||
casted_metadata[k] = v
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -1089,6 +1285,8 @@ async def update_key_fn(
|
|||
- rpm_limit: Optional[int] - Requests per minute limit
|
||||
- model_rpm_limit: Optional[dict] - Model-specific RPM limits {"gpt-4": 100, "claude-v1": 200}
|
||||
- model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000}
|
||||
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput" or "guaranteed_throughput"
|
||||
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput" or "guaranteed_throughput"
|
||||
- allowed_cache_controls: Optional[list] - List of allowed cache control values
|
||||
- duration: Optional[str] - Key validity duration ("30d", "1h", etc.)
|
||||
- permissions: Optional[dict] - Key-specific permissions
|
||||
|
|
@ -1136,13 +1334,6 @@ async def update_key_fn(
|
|||
if prisma_client is None:
|
||||
raise Exception("Not connected to DB!")
|
||||
|
||||
common_key_access_checks(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=data,
|
||||
llm_router=llm_router,
|
||||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
existing_key_row = await prisma_client.get_data(
|
||||
token=data.key, table_name="key", query_type="find_unique"
|
||||
)
|
||||
|
|
@ -1153,6 +1344,25 @@ async def update_key_fn(
|
|||
detail={"error": f"Team not found, passed team_id={data.team_id}"},
|
||||
)
|
||||
|
||||
## sanity check - prevent non-proxy admin user from updating key to belong to a different user
|
||||
if (
|
||||
data.user_id is not None
|
||||
and data.user_id != existing_key_row.user_id
|
||||
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"User={data.user_id} is not allowed to update key={key} to belong to user={existing_key_row.user_id}",
|
||||
)
|
||||
|
||||
common_key_access_checks(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=data,
|
||||
user_id=existing_key_row.user_id,
|
||||
llm_router=llm_router,
|
||||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
# check if user has permission to update key
|
||||
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -1162,14 +1372,25 @@ async def update_key_fn(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
# if team change - check if this is possible
|
||||
if is_different_team(data=data, existing_key_row=existing_key_row):
|
||||
# Only check team limits if key has a team_id
|
||||
team_obj: Optional[LiteLLM_TeamTableCachedObj] = None
|
||||
if data.team_id is not None:
|
||||
team_obj = await get_team_object(
|
||||
team_id=cast(str, data.team_id),
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
check_db_only=True,
|
||||
)
|
||||
|
||||
if team_obj is not None:
|
||||
await _check_team_key_limits(
|
||||
team_table=team_obj,
|
||||
data=data,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
# if team change - check if this is possible
|
||||
if is_different_team(data=data, existing_key_row=existing_key_row):
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -1177,6 +1398,14 @@ async def update_key_fn(
|
|||
"error": "LLM router not found. Please set it up by passing in a valid config.yaml or adding models via the UI."
|
||||
},
|
||||
)
|
||||
# team_obj should be set since is_different_team() returns True only when data.team_id is not None
|
||||
if team_obj is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "Team object not found for team change validation"
|
||||
},
|
||||
)
|
||||
validate_key_team_change(
|
||||
key=existing_key_row,
|
||||
team=team_obj,
|
||||
|
|
@ -1198,9 +1427,9 @@ async def update_key_fn(
|
|||
|
||||
# Handle rotation fields if auto_rotate is being enabled
|
||||
_set_key_rotation_fields(
|
||||
non_default_values,
|
||||
non_default_values.get("auto_rotate", False),
|
||||
non_default_values.get("rotation_interval")
|
||||
non_default_values,
|
||||
non_default_values.get("auto_rotate", False),
|
||||
non_default_values.get("rotation_interval"),
|
||||
)
|
||||
|
||||
_data = {**non_default_values, "token": key}
|
||||
|
|
@ -1602,8 +1831,6 @@ def _check_model_access_group(
|
|||
return True
|
||||
|
||||
|
||||
|
||||
|
||||
async def generate_key_helper_fn( # noqa: PLR0915
|
||||
request_type: Literal[
|
||||
"user", "key"
|
||||
|
|
@ -1766,12 +1993,12 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
"allowed_routes": allowed_routes or [],
|
||||
"object_permission_id": object_permission_id,
|
||||
}
|
||||
|
||||
|
||||
# Add rotation fields if auto_rotate is enabled
|
||||
_set_key_rotation_fields(
|
||||
data=key_data,
|
||||
auto_rotate=auto_rotate or False,
|
||||
rotation_interval=rotation_interval
|
||||
rotation_interval=rotation_interval,
|
||||
)
|
||||
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ All /team management endpoints
|
|||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
from litellm._uuid import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
|
|
@ -22,6 +21,7 @@ from pydantic import BaseModel
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import (
|
||||
BlockTeamRequest,
|
||||
CommonProxyErrors,
|
||||
|
|
@ -105,7 +105,7 @@ router = APIRouter()
|
|||
|
||||
class TeamMemberBudgetHandler:
|
||||
"""Helper class to handle team member budget, RPM, and TPM limit operations"""
|
||||
|
||||
|
||||
@staticmethod
|
||||
def should_create_budget(
|
||||
team_member_budget: Optional[float] = None,
|
||||
|
|
@ -113,12 +113,14 @@ class TeamMemberBudgetHandler:
|
|||
team_member_tpm_limit: Optional[int] = None,
|
||||
) -> bool:
|
||||
"""Check if any team member limits are provided"""
|
||||
return any([
|
||||
team_member_budget is not None,
|
||||
team_member_rpm_limit is not None,
|
||||
team_member_tpm_limit is not None,
|
||||
])
|
||||
|
||||
return any(
|
||||
[
|
||||
team_member_budget is not None,
|
||||
team_member_rpm_limit is not None,
|
||||
team_member_tpm_limit is not None,
|
||||
]
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def create_team_member_budget_table(
|
||||
data: Union[NewTeamRequest, LiteLLM_TeamTable],
|
||||
|
|
@ -146,7 +148,7 @@ class TeamMemberBudgetHandler:
|
|||
budget_id=budget_id,
|
||||
budget_duration=data.budget_duration,
|
||||
)
|
||||
|
||||
|
||||
if team_member_budget is not None:
|
||||
budget_request.max_budget = team_member_budget
|
||||
if team_member_rpm_limit is not None:
|
||||
|
|
@ -165,12 +167,12 @@ class TeamMemberBudgetHandler:
|
|||
new_team_data_json["metadata"][
|
||||
"team_member_budget_id"
|
||||
] = team_member_budget_table.budget_id
|
||||
|
||||
|
||||
# Remove team member fields from new_team_data_json
|
||||
TeamMemberBudgetHandler._clean_team_member_fields(new_team_data_json)
|
||||
|
||||
return new_team_data_json
|
||||
|
||||
|
||||
@staticmethod
|
||||
async def upsert_team_member_budget_table(
|
||||
team_table: LiteLLM_TeamTable,
|
||||
|
|
@ -193,14 +195,14 @@ class TeamMemberBudgetHandler:
|
|||
if team_member_budget_id is not None and isinstance(team_member_budget_id, str):
|
||||
# Budget exists - create update request with only provided values
|
||||
budget_request = BudgetNewRequest(budget_id=team_member_budget_id)
|
||||
|
||||
|
||||
if team_member_budget is not None:
|
||||
budget_request.max_budget = team_member_budget
|
||||
if team_member_rpm_limit is not None:
|
||||
budget_request.rpm_limit = team_member_rpm_limit
|
||||
if team_member_tpm_limit is not None:
|
||||
budget_request.tpm_limit = team_member_tpm_limit
|
||||
|
||||
|
||||
budget_row = await update_budget(
|
||||
budget_obj=budget_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -221,11 +223,11 @@ class TeamMemberBudgetHandler:
|
|||
team_member_rpm_limit=team_member_rpm_limit,
|
||||
team_member_tpm_limit=team_member_tpm_limit,
|
||||
)
|
||||
|
||||
|
||||
# Remove team member fields from updated_kv
|
||||
TeamMemberBudgetHandler._clean_team_member_fields(updated_kv)
|
||||
return updated_kv
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _clean_team_member_fields(data_dict: dict) -> None:
|
||||
"""Remove team member fields from data dictionary"""
|
||||
|
|
@ -267,7 +269,6 @@ async def get_all_team_memberships(
|
|||
return returned_tm
|
||||
|
||||
|
||||
|
||||
#### TEAM MANAGEMENT ####
|
||||
@router.post(
|
||||
"/team/new",
|
||||
|
|
@ -383,7 +384,7 @@ async def new_team( # noqa: PLR0915
|
|||
"error": f"Team id = {data.team_id} already exists. Please use a different team id."
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# If max_budget is not explicitly provided in the request,
|
||||
# check for a default value in the proxy configuration.
|
||||
if data.max_budget is None:
|
||||
|
|
@ -503,7 +504,7 @@ async def new_team( # noqa: PLR0915
|
|||
|
||||
# Set Management Endpoint Metadata Fields
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
|
||||
if getattr(data, field) is not None:
|
||||
if getattr(data, field, None) is not None:
|
||||
_set_object_metadata_field(
|
||||
object_data=complete_team_data,
|
||||
field_name=field,
|
||||
|
|
|
|||
|
|
@ -84,7 +84,7 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona
|
|||
|
||||
example header can be
|
||||
|
||||
{"Authorization": "bearer os.environ/COHERE_API_KEY"}
|
||||
{"Authorization": "Bearer os.environ/COHERE_API_KEY"}
|
||||
"""
|
||||
if custom_headers is None:
|
||||
return None
|
||||
|
|
@ -96,9 +96,13 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona
|
|||
# langfuse requires b64 encoded headers - we construct that here
|
||||
_langfuse_public_key = custom_headers["LANGFUSE_PUBLIC_KEY"]
|
||||
_langfuse_secret_key = custom_headers["LANGFUSE_SECRET_KEY"]
|
||||
if isinstance(_langfuse_public_key, str) and _langfuse_public_key.startswith("os.environ/"):
|
||||
if isinstance(
|
||||
_langfuse_public_key, str
|
||||
) and _langfuse_public_key.startswith("os.environ/"):
|
||||
_langfuse_public_key = get_secret_str(_langfuse_public_key)
|
||||
if isinstance(_langfuse_secret_key, str) and _langfuse_secret_key.startswith("os.environ/"):
|
||||
if isinstance(
|
||||
_langfuse_secret_key, str
|
||||
) and _langfuse_secret_key.startswith("os.environ/"):
|
||||
_langfuse_secret_key = get_secret_str(_langfuse_secret_key)
|
||||
headers["Authorization"] = "Basic " + b64encode(
|
||||
f"{_langfuse_public_key}:{_langfuse_secret_key}".encode("utf-8")
|
||||
|
|
@ -107,7 +111,9 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona
|
|||
# for all other headers
|
||||
headers[key] = value
|
||||
if isinstance(value, str) and "os.environ/" in value:
|
||||
verbose_proxy_logger.debug("pass through endpoint - looking up 'os.environ/' variable")
|
||||
verbose_proxy_logger.debug(
|
||||
"pass through endpoint - looking up 'os.environ/' variable"
|
||||
)
|
||||
# get string section that is os.environ/
|
||||
start_index = value.find("os.environ/")
|
||||
_variable_name = value[start_index:]
|
||||
|
|
@ -200,7 +206,9 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
|
|||
# skip router if user passed their key
|
||||
if "api_key" in data:
|
||||
llm_response = asyncio.create_task(litellm.aadapter_completion(**data))
|
||||
elif llm_router is not None and data["model"] in router_model_names: # model in router model list
|
||||
elif (
|
||||
llm_router is not None and data["model"] in router_model_names
|
||||
): # model in router model list
|
||||
llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
|
||||
elif (
|
||||
llm_router is not None
|
||||
|
|
@ -214,8 +222,8 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
|
|||
llm_response = asyncio.create_task(
|
||||
llm_router.aadapter_completion(**data, specific_deployment=True)
|
||||
)
|
||||
elif (
|
||||
llm_router is not None and llm_router.has_model_id(data["model"])
|
||||
elif llm_router is not None and llm_router.has_model_id(
|
||||
data["model"]
|
||||
): # model in router model list
|
||||
llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
|
||||
elif (
|
||||
|
|
@ -229,7 +237,10 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
|
|||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail={"error": "completion: Invalid model name passed in model=" + data.get("model", "")},
|
||||
detail={
|
||||
"error": "completion: Invalid model name passed in model="
|
||||
+ data.get("model", "")
|
||||
},
|
||||
)
|
||||
|
||||
# Await the llm_response task
|
||||
|
|
@ -243,7 +254,9 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
|
|||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success")
|
||||
proxy_logging_obj.update_request_status(
|
||||
litellm_call_id=data.get("litellm_call_id", ""), status="success"
|
||||
)
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("final response: %s", response)
|
||||
|
|
@ -265,7 +278,11 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
verbose_proxy_logger.exception("litellm.proxy.proxy_server.completion(): Exception occured - {}".format(str(e)))
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.completion(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
error_msg = f"{str(e)}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
|
|
@ -284,7 +301,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
) -> dict:
|
||||
excluded_headers = {"transfer-encoding", "content-encoding"}
|
||||
|
||||
return_headers = {key: value for key, value in headers.items() if key.lower() not in excluded_headers}
|
||||
return_headers = {
|
||||
key: value
|
||||
for key, value in headers.items()
|
||||
if key.lower() not in excluded_headers
|
||||
}
|
||||
if litellm_call_id:
|
||||
return_headers["x-litellm-call-id"] = litellm_call_id
|
||||
if custom_headers:
|
||||
|
|
@ -411,8 +432,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
|
||||
for field_name, field_value in form_data.items():
|
||||
if isinstance(field_value, (StarletteUploadFile, UploadFile)):
|
||||
files[field_name] = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
|
||||
upload_file=field_value
|
||||
files[field_name] = (
|
||||
await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
|
||||
upload_file=field_value
|
||||
)
|
||||
)
|
||||
else:
|
||||
form_data_dict[field_name] = field_value
|
||||
|
|
@ -462,8 +485,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
user_api_key_spend=user_api_key_dict.spend,
|
||||
user_api_key_max_budget=user_api_key_dict.max_budget,
|
||||
user_api_key_budget_reset_at=(
|
||||
user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None
|
||||
user_api_key_dict.budget_reset_at.isoformat()
|
||||
if user_api_key_dict.budget_reset_at
|
||||
else None
|
||||
),
|
||||
user_api_key_auth_metadata=user_api_key_dict.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -496,12 +522,16 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
"passthrough_logging_payload": passthrough_logging_payload,
|
||||
}
|
||||
|
||||
logging_obj.model_call_details["passthrough_logging_payload"] = passthrough_logging_payload
|
||||
logging_obj.model_call_details["passthrough_logging_payload"] = (
|
||||
passthrough_logging_payload
|
||||
)
|
||||
|
||||
return kwargs
|
||||
|
||||
@staticmethod
|
||||
def construct_target_url_with_subpath(base_target: str, subpath: str, include_subpath: Optional[bool]) -> str:
|
||||
def construct_target_url_with_subpath(
|
||||
base_target: str, subpath: str, include_subpath: Optional[bool]
|
||||
) -> str:
|
||||
"""
|
||||
Helper function to construct the full target URL with subpath handling.
|
||||
|
||||
|
|
@ -604,7 +634,9 @@ async def pass_through_request( # noqa: PLR0915
|
|||
).encode("ascii")
|
||||
)
|
||||
|
||||
endpoint_type: EndpointType = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url))
|
||||
endpoint_type: EndpointType = HttpPassThroughEndpointHelpers.get_endpoint_type(
|
||||
str(url)
|
||||
)
|
||||
|
||||
if custom_body:
|
||||
_parsed_body = custom_body
|
||||
|
|
@ -665,13 +697,15 @@ async def pass_through_request( # noqa: PLR0915
|
|||
logging_obj.model_call_details["litellm_call_id"] = litellm_call_id
|
||||
|
||||
# combine url with query params for logging
|
||||
requested_query_params: Optional[dict] = (
|
||||
query_params or dict(request.query_params)
|
||||
requested_query_params: Optional[dict] = query_params or dict(
|
||||
request.query_params
|
||||
)
|
||||
|
||||
requested_query_params_str = None
|
||||
if requested_query_params:
|
||||
requested_query_params_str = "&".join(f"{k}={v}" for k, v in requested_query_params.items())
|
||||
requested_query_params_str = "&".join(
|
||||
f"{k}={v}" for k, v in requested_query_params.items()
|
||||
)
|
||||
|
||||
logging_url = str(url)
|
||||
if requested_query_params_str:
|
||||
|
|
@ -689,9 +723,11 @@ async def pass_through_request( # noqa: PLR0915
|
|||
"headers": headers,
|
||||
},
|
||||
)
|
||||
stream = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
|
||||
parsed_body=_parsed_body,
|
||||
stream=stream,
|
||||
stream = (
|
||||
HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
|
||||
parsed_body=_parsed_body,
|
||||
stream=stream,
|
||||
)
|
||||
)
|
||||
|
||||
if stream:
|
||||
|
|
@ -708,7 +744,9 @@ async def pass_through_request( # noqa: PLR0915
|
|||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise HTTPException(status_code=e.response.status_code, detail=await e.response.aread())
|
||||
raise HTTPException(
|
||||
status_code=e.response.status_code, detail=await e.response.aread()
|
||||
)
|
||||
|
||||
return StreamingResponse(
|
||||
PassThroughStreamingHandler.chunk_processor(
|
||||
|
|
@ -730,16 +768,20 @@ async def pass_through_request( # noqa: PLR0915
|
|||
verbose_proxy_logger.debug("request method: {}".format(request.method))
|
||||
verbose_proxy_logger.debug("request url: {}".format(url))
|
||||
verbose_proxy_logger.debug("request headers: {}".format(headers))
|
||||
verbose_proxy_logger.debug("requested_query_params={}".format(requested_query_params))
|
||||
verbose_proxy_logger.debug(
|
||||
"requested_query_params={}".format(requested_query_params)
|
||||
)
|
||||
verbose_proxy_logger.debug("request body: {}".format(_parsed_body))
|
||||
|
||||
response = await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
|
||||
request=request,
|
||||
async_client=async_client,
|
||||
url=url,
|
||||
headers=headers,
|
||||
requested_query_params=requested_query_params,
|
||||
_parsed_body=_parsed_body,
|
||||
response = (
|
||||
await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
|
||||
request=request,
|
||||
async_client=async_client,
|
||||
url=url,
|
||||
headers=headers,
|
||||
requested_query_params=requested_query_params,
|
||||
_parsed_body=_parsed_body,
|
||||
)
|
||||
)
|
||||
verbose_proxy_logger.debug("response.headers= %s", response.headers)
|
||||
|
||||
|
|
@ -747,7 +789,9 @@ async def pass_through_request( # noqa: PLR0915
|
|||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise HTTPException(status_code=e.response.status_code, detail=await e.response.aread())
|
||||
raise HTTPException(
|
||||
status_code=e.response.status_code, detail=await e.response.aread()
|
||||
)
|
||||
|
||||
return StreamingResponse(
|
||||
PassThroughStreamingHandler.chunk_processor(
|
||||
|
|
@ -769,7 +813,9 @@ async def pass_through_request( # noqa: PLR0915
|
|||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise HTTPException(status_code=e.response.status_code, detail=e.response.text)
|
||||
raise HTTPException(
|
||||
status_code=e.response.status_code, detail=e.response.text
|
||||
)
|
||||
|
||||
if response.status_code >= 300:
|
||||
raise HTTPException(status_code=response.status_code, detail=response.text)
|
||||
|
|
@ -822,7 +868,9 @@ async def pass_through_request( # noqa: PLR0915
|
|||
api_base=str(url._uri_reference) if url else None,
|
||||
)
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.pass_through_endpoint(): Exception occured - {}".format(str(e))
|
||||
"litellm.proxy.proxy_server.pass_through_endpoint(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
|
||||
#########################################################
|
||||
|
|
@ -921,12 +969,16 @@ def create_pass_through_route(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
query_params: Optional[dict] = None,
|
||||
custom_body: Optional[dict] = None,
|
||||
stream: Optional[bool] = None, # if pass-through endpoint is a streaming request
|
||||
stream: Optional[
|
||||
bool
|
||||
] = None, # if pass-through endpoint is a streaming request
|
||||
subpath: str = "", # captures sub-paths when include_subpath=True
|
||||
):
|
||||
# Construct the full target URL with subpath if needed
|
||||
full_target = HttpPassThroughEndpointHelpers.construct_target_url_with_subpath(
|
||||
base_target=target, subpath=subpath, include_subpath=include_subpath
|
||||
full_target = (
|
||||
HttpPassThroughEndpointHelpers.construct_target_url_with_subpath(
|
||||
base_target=target, subpath=subpath, include_subpath=include_subpath
|
||||
)
|
||||
)
|
||||
|
||||
return await pass_through_request( # type: ignore
|
||||
|
|
@ -1078,7 +1130,9 @@ async def websocket_passthrough_request( # noqa: PLR0915
|
|||
# Create a dummy request object for WebSocket connections to maintain compatibility
|
||||
# with the existing _init_kwargs_for_pass_through_endpoint function
|
||||
class DummyRequest:
|
||||
def __init__(self, url: str, method: str = "WEBSOCKET", headers: Optional[dict] = None):
|
||||
def __init__(
|
||||
self, url: str, method: str = "WEBSOCKET", headers: Optional[dict] = None
|
||||
):
|
||||
self.url = url
|
||||
self.method = method
|
||||
self.headers = headers or {}
|
||||
|
|
@ -1183,9 +1237,9 @@ async def websocket_passthrough_request( # noqa: PLR0915
|
|||
)
|
||||
if extracted_model:
|
||||
kwargs["model"] = extracted_model
|
||||
kwargs[
|
||||
"custom_llm_provider"
|
||||
] = "vertex_ai-language-models"
|
||||
kwargs["custom_llm_provider"] = (
|
||||
"vertex_ai-language-models"
|
||||
)
|
||||
# Update logging object with correct model
|
||||
logging_obj.model = extracted_model
|
||||
logging_obj.model_call_details[
|
||||
|
|
@ -1251,9 +1305,9 @@ async def websocket_passthrough_request( # noqa: PLR0915
|
|||
# Update logging object with correct model
|
||||
logging_obj.model = extracted_model
|
||||
logging_obj.model_call_details["model"] = extracted_model
|
||||
logging_obj.model_call_details[
|
||||
"custom_llm_provider"
|
||||
] = "vertex_ai_language_models"
|
||||
logging_obj.model_call_details["custom_llm_provider"] = (
|
||||
"vertex_ai_language_models"
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"WebSocket passthrough ({endpoint}): Successfully extracted model '{extracted_model}' and set provider to 'vertex_ai' from server setup response"
|
||||
)
|
||||
|
|
@ -1597,11 +1651,15 @@ class InitPassThroughEndpointHelpers:
|
|||
def remove_endpoint_routes(endpoint_id: str):
|
||||
"""Remove all routes for a specific endpoint ID from the registry"""
|
||||
keys_to_remove = [
|
||||
key for key, value in _registered_pass_through_routes.items() if value["endpoint_id"] == endpoint_id
|
||||
key
|
||||
for key, value in _registered_pass_through_routes.items()
|
||||
if value["endpoint_id"] == endpoint_id
|
||||
]
|
||||
for key in keys_to_remove:
|
||||
del _registered_pass_through_routes[key]
|
||||
verbose_proxy_logger.debug("Removed pass-through route from registry: %s", key)
|
||||
verbose_proxy_logger.debug(
|
||||
"Removed pass-through route from registry: %s", key
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def is_registered_pass_through_route(route: str) -> bool:
|
||||
|
|
@ -1625,11 +1683,13 @@ class InitPassThroughEndpointHelpers:
|
|||
if len(parts) == 3:
|
||||
route_type = parts[1]
|
||||
registered_path = parts[2]
|
||||
|
||||
|
||||
if route_type == "exact" and route == registered_path:
|
||||
return True
|
||||
elif route_type == "subpath":
|
||||
if route == registered_path or route.startswith(registered_path + "/"):
|
||||
if route == registered_path or route.startswith(
|
||||
registered_path + "/"
|
||||
):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
@ -1669,7 +1729,9 @@ async def initialize_pass_through_endpoints(
|
|||
if _path is None:
|
||||
raise ValueError("Path is required for pass-through endpoint")
|
||||
_custom_headers = endpoint.get("headers", None)
|
||||
_custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
|
||||
_custom_headers = await set_env_variables_in_header(
|
||||
custom_headers=_custom_headers
|
||||
)
|
||||
_forward_headers = endpoint.get("forward_headers", None)
|
||||
_merge_query_params = endpoint.get("merge_query_params", None)
|
||||
_auth = endpoint.get("auth", None)
|
||||
|
|
@ -1688,7 +1750,9 @@ async def initialize_pass_through_endpoints(
|
|||
continue
|
||||
|
||||
# Add exact path route
|
||||
verbose_proxy_logger.debug("Initializing pass through endpoint: %s (ID: %s)", _path, endpoint_id)
|
||||
verbose_proxy_logger.debug(
|
||||
"Initializing pass through endpoint: %s (ID: %s)", _path, endpoint_id
|
||||
)
|
||||
InitPassThroughEndpointHelpers.add_exact_path_route(
|
||||
app=app,
|
||||
path=_path,
|
||||
|
|
@ -1715,7 +1779,9 @@ async def initialize_pass_through_endpoints(
|
|||
endpoint_id=endpoint_id,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("Added new pass through endpoint: %s (ID: %s)", _path, endpoint_id)
|
||||
verbose_proxy_logger.debug(
|
||||
"Added new pass through endpoint: %s (ID: %s)", _path, endpoint_id
|
||||
)
|
||||
|
||||
|
||||
async def _get_pass_through_endpoints_from_db(
|
||||
|
|
@ -1819,7 +1885,11 @@ async def update_pass_through_endpoints(
|
|||
# Find the index for updating the list
|
||||
endpoint_index = None
|
||||
for idx, endpoint in enumerate(pass_through_endpoint_data):
|
||||
_endpoint = PassThroughGenericEndpoint(**endpoint) if isinstance(endpoint, dict) else endpoint
|
||||
_endpoint = (
|
||||
PassThroughGenericEndpoint(**endpoint)
|
||||
if isinstance(endpoint, dict)
|
||||
else endpoint
|
||||
)
|
||||
if _endpoint.id == endpoint_id:
|
||||
endpoint_index = idx
|
||||
break
|
||||
|
|
@ -1827,7 +1897,9 @@ async def update_pass_through_endpoints(
|
|||
if endpoint_index is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Could not find index for endpoint with ID '{endpoint_id}'"},
|
||||
detail={
|
||||
"error": f"Could not find index for endpoint with ID '{endpoint_id}'"
|
||||
},
|
||||
)
|
||||
|
||||
# Get the update data as dict, excluding None values for partial updates
|
||||
|
|
@ -1858,9 +1930,13 @@ async def update_pass_through_endpoints(
|
|||
field_value=pass_through_endpoint_data,
|
||||
config_type="general_settings",
|
||||
)
|
||||
await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict)
|
||||
await update_config_general_settings(
|
||||
data=updated_data, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
return PassThroughEndpointResponse(endpoints=[updated_endpoint] if updated_endpoint else [])
|
||||
return PassThroughEndpointResponse(
|
||||
endpoints=[updated_endpoint] if updated_endpoint else []
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
@ -1887,7 +1963,9 @@ async def create_pass_through_endpoints(
|
|||
field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
except Exception:
|
||||
response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None)
|
||||
response = ConfigFieldInfo(
|
||||
field_name="pass_through_endpoints", field_value=None
|
||||
)
|
||||
|
||||
## Auto-generate ID if not provided
|
||||
data_dict = data.model_dump()
|
||||
|
|
@ -1905,7 +1983,9 @@ async def create_pass_through_endpoints(
|
|||
field_value=response.field_value,
|
||||
config_type="general_settings",
|
||||
)
|
||||
await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict)
|
||||
await update_config_general_settings(
|
||||
data=updated_data, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
# Return the created endpoint with the generated ID
|
||||
created_endpoint = PassThroughGenericEndpoint(**data_dict)
|
||||
|
|
@ -1938,7 +2018,9 @@ async def delete_pass_through_endpoints(
|
|||
field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
except Exception:
|
||||
response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None)
|
||||
response = ConfigFieldInfo(
|
||||
field_name="pass_through_endpoints", field_value=None
|
||||
)
|
||||
|
||||
## Update field by removing endpoint
|
||||
pass_through_endpoint_data: Optional[List] = response.field_value
|
||||
|
|
@ -1954,13 +2036,21 @@ async def delete_pass_through_endpoints(
|
|||
if found_endpoint is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "Endpoint with ID '{}' was not found in pass-through endpoint list.".format(endpoint_id)},
|
||||
detail={
|
||||
"error": "Endpoint with ID '{}' was not found in pass-through endpoint list.".format(
|
||||
endpoint_id
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
# Find the index for deleting from the list
|
||||
endpoint_index = None
|
||||
for idx, endpoint in enumerate(pass_through_endpoint_data):
|
||||
_endpoint = PassThroughGenericEndpoint(**endpoint) if isinstance(endpoint, dict) else endpoint
|
||||
_endpoint = (
|
||||
PassThroughGenericEndpoint(**endpoint)
|
||||
if isinstance(endpoint, dict)
|
||||
else endpoint
|
||||
)
|
||||
if _endpoint.id == endpoint_id:
|
||||
endpoint_index = idx
|
||||
break
|
||||
|
|
@ -1968,7 +2058,9 @@ async def delete_pass_through_endpoints(
|
|||
if endpoint_index is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"Could not find index for endpoint with ID '{endpoint_id}'"},
|
||||
detail={
|
||||
"error": f"Could not find index for endpoint with ID '{endpoint_id}'"
|
||||
},
|
||||
)
|
||||
|
||||
# Remove the endpoint
|
||||
|
|
@ -1984,7 +2076,9 @@ async def delete_pass_through_endpoints(
|
|||
field_value=pass_through_endpoint_data,
|
||||
config_type="general_settings",
|
||||
)
|
||||
await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict)
|
||||
await update_config_general_settings(
|
||||
data=updated_data, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
return PassThroughEndpointResponse(endpoints=[response_obj])
|
||||
|
||||
|
|
@ -2022,4 +2116,6 @@ async def initialize_pass_through_endpoints_in_db():
|
|||
Gets all pass-through endpoints from db and initializes them in the proxy server.
|
||||
"""
|
||||
pass_through_endpoints = await _get_pass_through_endpoints_from_db()
|
||||
await initialize_pass_through_endpoints(pass_through_endpoints=pass_through_endpoints)
|
||||
await initialize_pass_through_endpoints(
|
||||
pass_through_endpoints=pass_through_endpoints
|
||||
)
|
||||
|
|
|
|||
|
|
@ -253,9 +253,7 @@ 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,
|
||||
|
|
@ -302,9 +300,7 @@ 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,
|
||||
)
|
||||
|
|
@ -467,9 +463,9 @@ except ImportError:
|
|||
server_root_path = os.getenv("SERVER_ROOT_PATH", "")
|
||||
_license_check = LicenseCheck()
|
||||
premium_user: bool = _license_check.is_premium()
|
||||
premium_user_data: Optional[
|
||||
"EnterpriseLicenseData"
|
||||
] = _license_check.airgapped_license_data
|
||||
premium_user_data: Optional["EnterpriseLicenseData"] = (
|
||||
_license_check.airgapped_license_data
|
||||
)
|
||||
global_max_parallel_request_retries_env: Optional[str] = os.getenv(
|
||||
"LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES"
|
||||
)
|
||||
|
|
@ -966,9 +962,9 @@ model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(
|
|||
dual_cache=user_api_key_cache
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter)
|
||||
redis_usage_cache: Optional[
|
||||
RedisCache
|
||||
] = None # redis cache used for tracking spend, tpm/rpm limits
|
||||
redis_usage_cache: Optional[RedisCache] = (
|
||||
None # redis cache used for tracking spend, tpm/rpm limits
|
||||
)
|
||||
user_custom_auth = None
|
||||
user_custom_key_generate = None
|
||||
user_custom_sso = None
|
||||
|
|
@ -1299,9 +1295,9 @@ async def update_cache( # noqa: PLR0915
|
|||
_id = "team_id:{}".format(team_id)
|
||||
try:
|
||||
# Fetch the existing cost for the given user
|
||||
existing_spend_obj: Optional[
|
||||
LiteLLM_TeamTable
|
||||
] = await user_api_key_cache.async_get_cache(key=_id)
|
||||
existing_spend_obj: Optional[LiteLLM_TeamTable] = (
|
||||
await user_api_key_cache.async_get_cache(key=_id)
|
||||
)
|
||||
if existing_spend_obj is None:
|
||||
# do nothing if team not in api key cache
|
||||
return
|
||||
|
|
@ -1878,9 +1874,7 @@ class ProxyConfig:
|
|||
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,
|
||||
)
|
||||
from litellm.integrations.gitlab import set_global_gitlab_config
|
||||
|
||||
set_global_gitlab_config(value)
|
||||
verbose_proxy_logger.info(
|
||||
|
|
@ -2541,10 +2535,14 @@ class ProxyConfig:
|
|||
_model_list: list = []
|
||||
for m in new_models:
|
||||
_litellm_params = m.litellm_params
|
||||
if isinstance(_litellm_params, BaseModel):
|
||||
_litellm_params = _litellm_params.model_dump()
|
||||
if isinstance(_litellm_params, dict):
|
||||
# decrypt values
|
||||
for k, v in _litellm_params.items():
|
||||
decrypted_value = decrypt_value_helper(value=v, key=k)
|
||||
decrypted_value = decrypt_value_helper(
|
||||
value=v, key=k, return_original_value=True
|
||||
)
|
||||
_litellm_params[k] = decrypted_value
|
||||
_litellm_params = LiteLLM_Params(**_litellm_params)
|
||||
else:
|
||||
|
|
@ -2628,7 +2626,7 @@ class ProxyConfig:
|
|||
) -> 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"])
|
||||
|
|
@ -3153,10 +3151,10 @@ class ProxyConfig:
|
|||
)
|
||||
|
||||
try:
|
||||
guardrails_in_db: List[
|
||||
Guardrail
|
||||
] = await GuardrailRegistry.get_all_guardrails_from_db(
|
||||
prisma_client=prisma_client
|
||||
guardrails_in_db: List[Guardrail] = (
|
||||
await GuardrailRegistry.get_all_guardrails_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"guardrails from the DB %s", str(guardrails_in_db)
|
||||
|
|
@ -3386,9 +3384,9 @@ async def initialize( # noqa: PLR0915
|
|||
user_api_base = api_base
|
||||
dynamic_config[user_model]["api_base"] = api_base
|
||||
if api_version:
|
||||
os.environ[
|
||||
"AZURE_API_VERSION"
|
||||
] = api_version # set this for azure - litellm can read this from the env
|
||||
os.environ["AZURE_API_VERSION"] = (
|
||||
api_version # set this for azure - litellm can read this from the env
|
||||
)
|
||||
if max_tokens: # model-specific param
|
||||
dynamic_config[user_model]["max_tokens"] = max_tokens
|
||||
if temperature: # model-specific param
|
||||
|
|
@ -3888,10 +3886,10 @@ class ProxyStartupEvent:
|
|||
LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS,
|
||||
LITELLM_KEY_ROTATION_ENABLED,
|
||||
)
|
||||
|
||||
|
||||
key_rotation_enabled: Optional[bool] = str_to_bool(LITELLM_KEY_ROTATION_ENABLED)
|
||||
verbose_proxy_logger.debug(f"key_rotation_enabled: {key_rotation_enabled}")
|
||||
|
||||
|
||||
if key_rotation_enabled is True:
|
||||
try:
|
||||
from litellm.proxy.common_utils.key_rotation_manager import (
|
||||
|
|
@ -3902,19 +3900,25 @@ class ProxyStartupEvent:
|
|||
global prisma_client
|
||||
if prisma_client is not None:
|
||||
key_rotation_manager = KeyRotationManager(prisma_client)
|
||||
verbose_proxy_logger.debug(f"Key rotation background job scheduled every {LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS} seconds (LITELLM_KEY_ROTATION_ENABLED=true)")
|
||||
verbose_proxy_logger.debug(
|
||||
f"Key rotation background job scheduled every {LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS} seconds (LITELLM_KEY_ROTATION_ENABLED=true)"
|
||||
)
|
||||
scheduler.add_job(
|
||||
key_rotation_manager.process_rotations,
|
||||
"interval",
|
||||
seconds=LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS,
|
||||
id="key_rotation_job"
|
||||
id="key_rotation_job",
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.warning("Key rotation enabled but prisma_client not available")
|
||||
verbose_proxy_logger.warning(
|
||||
"Key rotation enabled but prisma_client not available"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(f"Failed to setup key rotation job: {e}")
|
||||
else:
|
||||
verbose_proxy_logger.debug("Key rotation disabled (set LITELLM_KEY_ROTATION_ENABLED=true to enable)")
|
||||
verbose_proxy_logger.debug(
|
||||
"Key rotation disabled (set LITELLM_KEY_ROTATION_ENABLED=true to enable)"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def _setup_prisma_client(
|
||||
|
|
@ -8745,9 +8749,9 @@ async def get_config_list(
|
|||
hasattr(sub_field_info, "description")
|
||||
and sub_field_info.description is not None
|
||||
):
|
||||
nested_fields[
|
||||
idx
|
||||
].field_description = sub_field_info.description
|
||||
nested_fields[idx].field_description = (
|
||||
sub_field_info.description
|
||||
)
|
||||
idx += 1
|
||||
|
||||
_stored_in_db = None
|
||||
|
|
|
|||
|
|
@ -179,6 +179,7 @@ model LiteLLM_MCPServerTable {
|
|||
mcp_info Json? @default("{}")
|
||||
mcp_access_groups String[]
|
||||
allowed_tools String[] @default([])
|
||||
extra_headers String[] @default([])
|
||||
// Health check status
|
||||
status String? @default("unknown")
|
||||
last_health_check DateTime?
|
||||
|
|
|
|||
|
|
@ -1395,9 +1395,12 @@ class ProxyLogging:
|
|||
3. /image/generation
|
||||
4. /files
|
||||
"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
try:
|
||||
guardrail_callbacks: List[CustomGuardrail] = []
|
||||
other_callbacks: List[CustomLogger] = []
|
||||
try:
|
||||
for callback in litellm.callbacks:
|
||||
_callback: Optional[CustomLogger] = None
|
||||
if isinstance(callback, str):
|
||||
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
|
||||
|
|
@ -1407,36 +1410,37 @@ class ProxyLogging:
|
|||
_callback = callback # type: ignore
|
||||
|
||||
if _callback is not None:
|
||||
if isinstance(_callback, CustomGuardrail):
|
||||
guardrail_callbacks.append(_callback)
|
||||
else:
|
||||
other_callbacks.append(_callback)
|
||||
############## Handle Guardrails ########################################
|
||||
#############################################################################
|
||||
if isinstance(callback, CustomGuardrail):
|
||||
# Main - V2 Guardrails implementation
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
if (
|
||||
callback.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.post_call
|
||||
)
|
||||
is not True
|
||||
):
|
||||
continue
|
||||
for callback in guardrail_callbacks:
|
||||
# Main - V2 Guardrails implementation
|
||||
if (
|
||||
callback.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.post_call
|
||||
)
|
||||
is not True
|
||||
):
|
||||
continue
|
||||
|
||||
await callback.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=data,
|
||||
response=response,
|
||||
)
|
||||
await callback.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=data,
|
||||
response=response,
|
||||
)
|
||||
|
||||
############ Handle CustomLogger ###############################
|
||||
#################################################################
|
||||
elif isinstance(_callback, CustomLogger):
|
||||
await _callback.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=data,
|
||||
response=response,
|
||||
)
|
||||
except Exception as e:
|
||||
raise e
|
||||
############ Handle CustomLogger ###############################
|
||||
#################################################################
|
||||
for callback in other_callbacks:
|
||||
await callback.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict, data=data, response=response
|
||||
)
|
||||
except Exception as e:
|
||||
raise e
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_hook(
|
||||
|
|
@ -3571,18 +3575,21 @@ def handle_exception_on_proxy(e: Exception) -> ProxyException:
|
|||
)
|
||||
|
||||
|
||||
def _premium_user_check():
|
||||
def _premium_user_check(feature: Optional[str] = None):
|
||||
"""
|
||||
Raises an HTTPException if the user is not a premium user
|
||||
"""
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
||||
if feature:
|
||||
detail_msg = f"This feature is only available for LiteLLM Enterprise users: {feature}. {CommonProxyErrors.not_premium_user.value}"
|
||||
else:
|
||||
detail_msg = f"This feature is only available for LiteLLM Enterprise users. {CommonProxyErrors.not_premium_user.value}"
|
||||
|
||||
if not premium_user:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"This feature is only available for LiteLLM Enterprise users. {CommonProxyErrors.not_premium_user.value}"
|
||||
},
|
||||
detail={"error": detail_msg},
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -426,13 +426,13 @@ class PrometheusMetricLabels:
|
|||
|
||||
# Buffer monitoring metrics - these typically don't need additional labels
|
||||
litellm_pod_lock_manager_size: List[str] = []
|
||||
|
||||
|
||||
litellm_in_memory_daily_spend_update_queue_size: List[str] = []
|
||||
|
||||
|
||||
litellm_redis_daily_spend_update_queue_size: List[str] = []
|
||||
|
||||
|
||||
litellm_in_memory_spend_update_queue_size: List[str] = []
|
||||
|
||||
|
||||
litellm_redis_spend_update_queue_size: List[str] = []
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -1867,6 +1867,7 @@ class StandardLoggingUserAPIKeyMetadata(TypedDict):
|
|||
user_api_key_team_alias: Optional[str]
|
||||
user_api_key_end_user_id: Optional[str]
|
||||
user_api_key_request_route: Optional[str]
|
||||
user_api_key_auth_metadata: Optional[Dict[str, str]]
|
||||
|
||||
|
||||
class StandardLoggingMCPToolCall(TypedDict, total=False):
|
||||
|
|
@ -2077,10 +2078,12 @@ 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
|
||||
|
|
@ -2090,12 +2093,16 @@ class CachingDetails(TypedDict):
|
|||
Duration for reading from cache
|
||||
"""
|
||||
|
||||
|
||||
class CostBreakdown(TypedDict):
|
||||
"""
|
||||
Detailed cost breakdown for a request
|
||||
"""
|
||||
|
||||
input_cost: float # Cost of input/prompt tokens
|
||||
output_cost: float # Cost of output/completion tokens (includes reasoning if applicable)
|
||||
output_cost: (
|
||||
float # Cost of output/completion tokens (includes reasoning if applicable)
|
||||
)
|
||||
total_cost: float # Total cost (input + output + tool usage)
|
||||
tool_usage_cost: float # Cost of usage of built-in tools
|
||||
|
||||
|
|
@ -2702,12 +2709,12 @@ class PriorityReservationSettings(BaseModel):
|
|||
"""
|
||||
|
||||
default_priority: float = Field(
|
||||
default=0.5,
|
||||
default=0.25,
|
||||
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,
|
||||
default=0.50,
|
||||
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."
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1402,7 +1402,7 @@ def client(original_function): # noqa: PLR0915
|
|||
print_verbose(
|
||||
f"ASYNC kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}; kwargs.get('cache'): {kwargs.get('cache', None)}"
|
||||
)
|
||||
_caching_handler_response: CachingHandlerResponse = (
|
||||
_caching_handler_response: Optional[CachingHandlerResponse] = (
|
||||
await _llm_caching_handler._async_get_cache(
|
||||
model=model or "",
|
||||
original_function=original_function,
|
||||
|
|
@ -1414,14 +1414,15 @@ def client(original_function): # noqa: PLR0915
|
|||
)
|
||||
)
|
||||
|
||||
if (
|
||||
_caching_handler_response.cached_result is not None
|
||||
and _caching_handler_response.final_embedding_cached_response is None
|
||||
):
|
||||
return _caching_handler_response.cached_result
|
||||
if _caching_handler_response is not None:
|
||||
if (
|
||||
_caching_handler_response.cached_result is not None
|
||||
and _caching_handler_response.final_embedding_cached_response is None
|
||||
):
|
||||
return _caching_handler_response.cached_result
|
||||
|
||||
elif _caching_handler_response.embedding_all_elements_cache_hit is True:
|
||||
return _caching_handler_response.final_embedding_cached_response
|
||||
elif _caching_handler_response.embedding_all_elements_cache_hit is True:
|
||||
return _caching_handler_response.final_embedding_cached_response
|
||||
|
||||
# CHECK MAX TOKENS
|
||||
if (
|
||||
|
|
@ -1524,6 +1525,7 @@ def client(original_function): # noqa: PLR0915
|
|||
# REBUILD EMBEDDING CACHING
|
||||
if (
|
||||
isinstance(result, EmbeddingResponse)
|
||||
and _caching_handler_response is not None
|
||||
and _caching_handler_response.final_embedding_cached_response
|
||||
is not None
|
||||
):
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
8
poetry.lock
generated
8
poetry.lock
generated
|
|
@ -3804,15 +3804,15 @@ files = [
|
|||
|
||||
[[package]]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.2.22"
|
||||
version = "0.2.25"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
optional = true
|
||||
python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8"
|
||||
groups = ["main"]
|
||||
markers = "extra == \"proxy\""
|
||||
files = [
|
||||
{file = "litellm_proxy_extras-0.2.22-py3-none-any.whl", hash = "sha256:e64b19b48e8d84cad56bb136c7f31d9ae601a10628327c922634d7081803c205"},
|
||||
{file = "litellm_proxy_extras-0.2.22.tar.gz", hash = "sha256:59c395bff3353de57d67b7637e8ce0a8a4e096ce55e2ee2df4d9d4bda94f6ef0"},
|
||||
{file = "litellm_proxy_extras-0.2.25-py3-none-any.whl", hash = "sha256:334ac3c04511258e2cbbd7a1ddb6e30619a6e693b267db92033732f4d981baab"},
|
||||
{file = "litellm_proxy_extras-0.2.25.tar.gz", hash = "sha256:9cf363570a5dc3349bea6ad1fba00ce9aeb90232fc69adc32881e53bec2cbf8f"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -9598,4 +9598,4 @@ utils = ["numpydoc"]
|
|||
[metadata]
|
||||
lock-version = "2.1"
|
||||
python-versions = ">=3.8.1,<4.0, !=3.9.7"
|
||||
content-hash = "dd6b1b42d43c2049fd8fcc95a6627581c5d9c60b3afd5eab60659d8f5d6ae641"
|
||||
content-hash = "ef5f8d965a4d77f6ae7d424306e2c88082708bc7e374896e6b022a13ce7c1962"
|
||||
|
|
|
|||
|
|
@ -59,7 +59,7 @@ websockets = {version = "^13.1.0", optional = true}
|
|||
boto3 = {version = "1.36.0", optional = true}
|
||||
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
|
||||
mcp = {version = "^1.10.0", optional = true, python = ">=3.10"}
|
||||
litellm-proxy-extras = {version = "0.2.22", optional = true}
|
||||
litellm-proxy-extras = {version = "0.2.25", optional = true}
|
||||
rich = {version = "13.7.1", optional = true}
|
||||
litellm-enterprise = {version = "0.1.20", optional = true}
|
||||
diskcache = {version = "^5.6.1", optional = true}
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ sentry_sdk==2.21.0 # for sentry error handling
|
|||
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
|
||||
cryptography==44.0.1
|
||||
tzdata==2025.1 # IANA time zone database
|
||||
litellm-proxy-extras==0.2.22 # for proxy extras - e.g. prisma migrations
|
||||
litellm-proxy-extras==0.2.25 # for proxy extras - e.g. prisma migrations
|
||||
### LITELLM PACKAGE DEPENDENCIES
|
||||
python-dotenv==1.0.1 # for env
|
||||
tiktoken==0.8.0 # for calculating usage
|
||||
|
|
|
|||
|
|
@ -179,6 +179,7 @@ model LiteLLM_MCPServerTable {
|
|||
mcp_info Json? @default("{}")
|
||||
mcp_access_groups String[]
|
||||
allowed_tools String[] @default([])
|
||||
extra_headers String[] @default([])
|
||||
// Health check status
|
||||
status String? @default("unknown")
|
||||
last_health_check DateTime?
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ sys.path.insert(0, os.path.abspath("../.."))
|
|||
|
||||
import asyncio
|
||||
import logging
|
||||
from litellm._uuid import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
|
||||
|
|
@ -16,6 +15,7 @@ from prometheus_client import REGISTRY, CollectorRegistry
|
|||
import litellm
|
||||
from litellm import completion
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.types.utils import (
|
||||
StandardLoggingHiddenParams,
|
||||
|
|
@ -1033,10 +1033,10 @@ def test_deployment_state_management(prometheus_logger):
|
|||
# Test set_deployment_healthy (state=0)
|
||||
prometheus_logger.set_deployment_healthy(**test_params)
|
||||
prometheus_logger.litellm_deployment_state.labels.assert_called_with(
|
||||
test_params["litellm_model_name"],
|
||||
test_params["model_id"],
|
||||
test_params["api_base"],
|
||||
test_params["api_provider"],
|
||||
litellm_model_name=test_params["litellm_model_name"],
|
||||
model_id=test_params["model_id"],
|
||||
api_base=test_params["api_base"],
|
||||
api_provider=test_params["api_provider"],
|
||||
)
|
||||
prometheus_logger.litellm_deployment_state.labels().set.assert_called_with(0)
|
||||
|
||||
|
|
@ -1153,22 +1153,28 @@ def test_get_custom_labels_from_tags_wildcard_patterns(monkeypatch):
|
|||
|
||||
# Configure tags with wildcard patterns
|
||||
monkeypatch.setattr(
|
||||
"litellm.custom_prometheus_tags",
|
||||
["User-Agent: curl/*", "User-Agent: python-requests/*", "Environment: prod*", "Service: api-gateway*", "exact-match"]
|
||||
"litellm.custom_prometheus_tags",
|
||||
[
|
||||
"User-Agent: curl/*",
|
||||
"User-Agent: python-requests/*",
|
||||
"Environment: prod*",
|
||||
"Service: api-gateway*",
|
||||
"exact-match",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# Test tags that should match the wildcard patterns
|
||||
tags = [
|
||||
"User-Agent: curl/7.68.0",
|
||||
"User-Agent: python-requests/2.28.1",
|
||||
"User-Agent: curl/7.68.0",
|
||||
"User-Agent: python-requests/2.28.1",
|
||||
"Environment: production",
|
||||
"Service: api-gateway-v2",
|
||||
"exact-match",
|
||||
"other-tag"
|
||||
"other-tag",
|
||||
]
|
||||
|
||||
|
||||
result = get_custom_labels_from_tags(tags)
|
||||
|
||||
|
||||
expected = {
|
||||
"tag_User_Agent__curl__": "true", # matches "User-Agent: curl/*"
|
||||
"tag_User_Agent__python_requests__": "true", # matches "User-Agent: python-requests/*"
|
||||
|
|
@ -1176,7 +1182,7 @@ def test_get_custom_labels_from_tags_wildcard_patterns(monkeypatch):
|
|||
"tag_Service__api_gateway_": "true", # matches "Service: api-gateway*"
|
||||
"tag_exact_match": "true", # exact match
|
||||
}
|
||||
|
||||
|
||||
assert result == expected
|
||||
|
||||
|
||||
|
|
@ -1186,26 +1192,26 @@ def test_get_custom_labels_from_tags_wildcard_no_matches(monkeypatch):
|
|||
|
||||
# Configure tags with wildcard patterns
|
||||
monkeypatch.setattr(
|
||||
"litellm.custom_prometheus_tags",
|
||||
["User-Agent: firefox/*", "Environment: dev*", "Service: web-app*"]
|
||||
"litellm.custom_prometheus_tags",
|
||||
["User-Agent: firefox/*", "Environment: dev*", "Service: web-app*"],
|
||||
)
|
||||
|
||||
|
||||
# Test tags that should NOT match the wildcard patterns
|
||||
tags = [
|
||||
"User-Agent: curl/7.68.0", # doesn't match "User-Agent: firefox/*"
|
||||
"Environment: production", # doesn't match "Environment: dev*"
|
||||
"Environment: production", # doesn't match "Environment: dev*"
|
||||
"Service: api-gateway-v2", # doesn't match "Service: web-app*"
|
||||
"other-tag"
|
||||
"other-tag",
|
||||
]
|
||||
|
||||
|
||||
result = get_custom_labels_from_tags(tags)
|
||||
|
||||
|
||||
expected = {
|
||||
"tag_User_Agent__firefox__": "false", # no match for "User-Agent: firefox/*"
|
||||
"tag_Environment__dev_": "false", # no match for "Environment: dev*"
|
||||
"tag_Service__web_app_": "false", # no match for "Service: web-app*"
|
||||
}
|
||||
|
||||
|
||||
assert result == expected
|
||||
|
||||
|
||||
|
|
@ -1216,48 +1222,69 @@ def test_tag_matches_wildcard_configured_pattern():
|
|||
)
|
||||
|
||||
# Test cases that should match
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=["User-Agent: curl/7.68.0", "prod", "other"],
|
||||
configured_tag="User-Agent: curl/*"
|
||||
) is True
|
||||
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=["User-Agent: python-requests/2.28.1", "test"],
|
||||
configured_tag="User-Agent: python-requests/*"
|
||||
) is True
|
||||
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=["Environment: production", "debug"],
|
||||
configured_tag="Environment: prod*"
|
||||
) is True
|
||||
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=["User-Agent: curl/7.68.0", "prod", "other"],
|
||||
configured_tag="User-Agent: curl/*",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=["User-Agent: python-requests/2.28.1", "test"],
|
||||
configured_tag="User-Agent: python-requests/*",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=["Environment: production", "debug"],
|
||||
configured_tag="Environment: prod*",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
# Test exact match (no wildcard)
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=["prod", "test"],
|
||||
configured_tag="prod"
|
||||
) is True
|
||||
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=["prod", "test"], configured_tag="prod"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
# Test cases that should NOT match
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=["User-Agent: firefox/98.0", "prod"],
|
||||
configured_tag="User-Agent: curl/*"
|
||||
) is False
|
||||
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=["Environment: development", "test"],
|
||||
configured_tag="Environment: prod*"
|
||||
) is False
|
||||
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=["staging", "test"],
|
||||
configured_tag="prod"
|
||||
) is False
|
||||
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=["User-Agent: firefox/98.0", "prod"],
|
||||
configured_tag="User-Agent: curl/*",
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=["Environment: development", "test"],
|
||||
configured_tag="Environment: prod*",
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=["staging", "test"], configured_tag="prod"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
# Test with empty tags
|
||||
assert _tag_matches_wildcard_configured_pattern(
|
||||
tags=[],
|
||||
configured_tag="User-Agent: curl/*"
|
||||
) is False
|
||||
assert (
|
||||
_tag_matches_wildcard_configured_pattern(
|
||||
tags=[], configured_tag="User-Agent: curl/*"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio(scope="session")
|
||||
|
|
@ -1920,12 +1947,12 @@ def test_set_llm_deployment_success_metrics_with_label_filtering():
|
|||
async def test_prometheus_token_metrics_with_prometheus_config():
|
||||
"""
|
||||
Test that validates the renamed token metrics are incremented correctly with a prometheus config.
|
||||
|
||||
|
||||
This test ensures that after the metric renaming (git diff):
|
||||
- litellm_total_tokens -> litellm_total_tokens_metric
|
||||
- litellm_input_tokens -> litellm_input_tokens_metric
|
||||
- litellm_input_tokens -> litellm_input_tokens_metric
|
||||
- litellm_output_tokens -> litellm_output_tokens_metric
|
||||
|
||||
|
||||
All three metrics should be properly incremented when making a successful completion request.
|
||||
"""
|
||||
from prometheus_client import CollectorRegistry, Counter
|
||||
|
|
@ -1937,39 +1964,39 @@ async def test_prometheus_token_metrics_with_prometheus_config():
|
|||
collectors = list(REGISTRY._collector_to_names.keys())
|
||||
for collector in collectors:
|
||||
REGISTRY.unregister(collector)
|
||||
|
||||
|
||||
# Set up prometheus configuration that includes the token metrics
|
||||
config = [
|
||||
PrometheusMetricsConfig(
|
||||
group="token_metrics_test",
|
||||
metrics=[
|
||||
"litellm_total_tokens_metric",
|
||||
"litellm_input_tokens_metric",
|
||||
"litellm_input_tokens_metric",
|
||||
"litellm_output_tokens_metric",
|
||||
"litellm_requests_metric"
|
||||
"litellm_requests_metric",
|
||||
],
|
||||
include_labels=[
|
||||
"model",
|
||||
"hashed_api_key",
|
||||
"hashed_api_key",
|
||||
"api_key_alias",
|
||||
"team",
|
||||
"team_alias"
|
||||
"team_alias",
|
||||
],
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
# Mock litellm.prometheus_metrics_config
|
||||
with patch("litellm.prometheus_metrics_config", config):
|
||||
# Create PrometheusLogger with the configuration
|
||||
prometheus_logger = PrometheusLogger()
|
||||
|
||||
|
||||
# Test data with specific token counts
|
||||
standard_logging_payload = create_standard_logging_payload()
|
||||
standard_logging_payload["total_tokens"] = 1500
|
||||
standard_logging_payload["prompt_tokens"] = 900
|
||||
standard_logging_payload["completion_tokens"] = 600
|
||||
standard_logging_payload["response_cost"] = 0.075
|
||||
|
||||
|
||||
kwargs = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"stream": False,
|
||||
|
|
@ -1983,7 +2010,7 @@ async def test_prometheus_token_metrics_with_prometheus_config():
|
|||
}
|
||||
},
|
||||
"start_time": datetime.now() - timedelta(seconds=2),
|
||||
"completion_start_time": datetime.now() - timedelta(seconds=1),
|
||||
"completion_start_time": datetime.now() - timedelta(seconds=1),
|
||||
"api_call_start_time": datetime.now() - timedelta(seconds=1.5),
|
||||
"end_time": datetime.now(),
|
||||
"standard_logging_object": standard_logging_payload,
|
||||
|
|
@ -1999,69 +2026,75 @@ async def test_prometheus_token_metrics_with_prometheus_config():
|
|||
|
||||
print("final registry values", REGISTRY._collector_to_names)
|
||||
|
||||
# Get metric collectors directly from registry
|
||||
# Get metric collectors directly from registry
|
||||
metric_collectors = {}
|
||||
for collector, names in REGISTRY._collector_to_names.items():
|
||||
metric_name = names[0] # First name is the base metric name
|
||||
metric_collectors[metric_name] = collector
|
||||
|
||||
print("=== Final Metric Values (Direct Access) ===")
|
||||
|
||||
# Expected values
|
||||
|
||||
# Expected values
|
||||
expected_values = {
|
||||
"litellm_total_tokens_metric": 1500.0,
|
||||
"litellm_input_tokens_metric": 900.0,
|
||||
"litellm_output_tokens_metric": 600.0,
|
||||
"litellm_requests_metric": 1.0
|
||||
"litellm_requests_metric": 1.0,
|
||||
}
|
||||
|
||||
|
||||
expected_label_values = {
|
||||
'api_key_alias': 'test_alias',
|
||||
'hashed_api_key': 'test_hash',
|
||||
'model': 'gpt-3.5-turbo',
|
||||
'team': 'test_team',
|
||||
'team_alias': 'test_team_alias'
|
||||
"api_key_alias": "test_alias",
|
||||
"hashed_api_key": "test_hash",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"team": "test_team",
|
||||
"team_alias": "test_team_alias",
|
||||
}
|
||||
|
||||
# Validate each metric directly
|
||||
for metric_name, expected_value in expected_values.items():
|
||||
if metric_name in metric_collectors:
|
||||
collector = metric_collectors[metric_name]
|
||||
|
||||
|
||||
# Get all samples for this metric
|
||||
samples = list(collector.collect())[0].samples
|
||||
|
||||
|
||||
# Find the _total sample (the actual counter value)
|
||||
total_sample = None
|
||||
for sample in samples:
|
||||
if sample.name.endswith('_total'):
|
||||
if sample.name.endswith("_total"):
|
||||
total_sample = sample
|
||||
break
|
||||
|
||||
|
||||
if total_sample:
|
||||
actual_value = total_sample.value
|
||||
actual_labels = total_sample.labels
|
||||
|
||||
print(f"✓ {metric_name}: expected={expected_value}, actual={actual_value}")
|
||||
|
||||
print(
|
||||
f"✓ {metric_name}: expected={expected_value}, actual={actual_value}"
|
||||
)
|
||||
print(f" Labels: {actual_labels}")
|
||||
|
||||
|
||||
# Validate the value
|
||||
assert actual_value == expected_value, f"Expected {expected_value}, got {actual_value} for {metric_name}"
|
||||
|
||||
assert (
|
||||
actual_value == expected_value
|
||||
), f"Expected {expected_value}, got {actual_value} for {metric_name}"
|
||||
|
||||
# Validate the labels
|
||||
for label_key, expected_label_value in expected_label_values.items():
|
||||
for (
|
||||
label_key,
|
||||
expected_label_value,
|
||||
) in expected_label_values.items():
|
||||
actual_label_value = actual_labels.get(label_key)
|
||||
assert actual_label_value == expected_label_value, f"Expected label {label_key}={expected_label_value}, got {actual_label_value}"
|
||||
|
||||
assert (
|
||||
actual_label_value == expected_label_value
|
||||
), f"Expected label {label_key}={expected_label_value}, got {actual_label_value}"
|
||||
|
||||
print(f" ✓ {metric_name} VALIDATED")
|
||||
else:
|
||||
raise AssertionError(f"No _total sample found for {metric_name}")
|
||||
else:
|
||||
raise AssertionError(f"Metric {metric_name} not found in registry")
|
||||
|
||||
|
||||
print("✓ All token metrics validated successfully!")
|
||||
|
||||
# check final value of metrics in registry
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1384,28 +1384,34 @@ async def test_bedrock_guardrail_post_call_success_hook_no_output_text():
|
|||
guardrailVersion="DRAFT"
|
||||
)
|
||||
|
||||
# Mock Bedrock API with no output text
|
||||
mock_bedrock_response = MagicMock()
|
||||
mock_bedrock_response.status_code = 200
|
||||
mock_bedrock_response.json.return_value = {
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"toolUse": {
|
||||
"toolUseId": "tooluse_kZJMlvQmRJ6eAyJE5GIl7Q",
|
||||
"name": "top_song",
|
||||
"input": {
|
||||
"sign": "WZPZ"
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"stopReason": "tool_use"
|
||||
}
|
||||
# Create a ModelResponse with tool calls (no text content)
|
||||
# This simulates a response where the LLM is making a tool call
|
||||
mock_response = litellm.ModelResponse(
|
||||
id="test-id",
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
index=0,
|
||||
message=litellm.Message(
|
||||
role="assistant",
|
||||
content=None, # No text content
|
||||
tool_calls=[
|
||||
litellm.utils.ChatCompletionMessageToolCall(
|
||||
id="tooluse_kZJMlvQmRJ6eAyJE5GIl7Q",
|
||||
function=litellm.utils.Function(
|
||||
name="top_song",
|
||||
arguments='{"sign": "WZPZ"}'
|
||||
),
|
||||
type="function"
|
||||
)
|
||||
]
|
||||
),
|
||||
finish_reason="tool_calls"
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-4o",
|
||||
object="chat.completion"
|
||||
)
|
||||
|
||||
data = {
|
||||
"model": "gpt-4o",
|
||||
|
|
@ -1415,10 +1421,11 @@ async def test_bedrock_guardrail_post_call_success_hook_no_output_text():
|
|||
}
|
||||
mock_user_api_key_dict = UserAPIKeyAuth()
|
||||
|
||||
return await guardrail.async_post_call_success_hook(
|
||||
result = await guardrail.async_post_call_success_hook(
|
||||
data=data,
|
||||
response=mock_bedrock_response,
|
||||
response=mock_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
# If no error is raised, then the test passes
|
||||
# If no error is raised and result is None, then the test passes
|
||||
assert result is None
|
||||
print("✅ No output text in response test passed")
|
||||
|
|
@ -140,7 +140,8 @@ def test_e2e_bedrock_embedding_image_twelvelabs_marengo():
|
|||
response = litellm.embedding(
|
||||
model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0",
|
||||
input=[duck_img_base64],
|
||||
aws_region_name="us-east-1"
|
||||
aws_region_name="us-east-1",
|
||||
input_type="image"
|
||||
)
|
||||
|
||||
# Validate response structure
|
||||
|
|
@ -252,7 +253,7 @@ async def test_e2e_bedrock_async_invoke_embedding_async_twelvelabs_marengo():
|
|||
|
||||
# Validate hidden params contain invocation ARN
|
||||
assert hasattr(response._hidden_params, '_invocation_arn'), "Hidden params should have _invocation_arn"
|
||||
assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123", "Invocation ARN should be preserved"
|
||||
assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-async-job-456", "Invocation ARN should be preserved"
|
||||
|
||||
print(f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}")
|
||||
|
||||
|
|
|
|||
|
|
@ -465,11 +465,11 @@ def test_gemini_url_context():
|
|||
from litellm import completion
|
||||
|
||||
litellm._turn_on_debug()
|
||||
URL1 = "https://www.foodnetwork.com/recipes/ina-garten/perfect-roast-chicken-recipe-1940592"
|
||||
|
||||
url = "https://ai.google.dev/gemini-api/docs/models"
|
||||
prompt = f"""
|
||||
Summarize this document:
|
||||
{url}
|
||||
Get the recipes listed on the following website
|
||||
{URL1}
|
||||
"""
|
||||
response = completion(
|
||||
model="gemini/gemini-2.5-flash",
|
||||
|
|
@ -482,7 +482,7 @@ def test_gemini_url_context():
|
|||
url_context_metadata = response.model_extra["vertex_ai_url_context_metadata"]
|
||||
assert url_context_metadata is not None
|
||||
urlMetadata = url_context_metadata[0]["urlMetadata"][0]
|
||||
assert urlMetadata["retrievedUrl"] == url
|
||||
assert urlMetadata["retrievedUrl"] == URL1
|
||||
assert urlMetadata["urlRetrievalStatus"] == "URL_RETRIEVAL_STATUS_SUCCESS"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -117,7 +117,7 @@ async def test_pass_through_endpoint_rerank(client):
|
|||
{
|
||||
"path": "/v1/rerank",
|
||||
"target": "https://api.cohere.com/v1/rerank",
|
||||
"headers": {"Authorization": f"bearer {_cohere_api_key}"},
|
||||
"headers": {"Authorization": f"Bearer {_cohere_api_key}"},
|
||||
}
|
||||
]
|
||||
|
||||
|
|
@ -193,7 +193,7 @@ async def test_pass_through_endpoint_rpm_limit(
|
|||
"path": "/v1/rerank",
|
||||
"target": "https://api.cohere.com/v1/rerank",
|
||||
"auth": auth,
|
||||
"headers": {"Authorization": f"bearer {_cohere_api_key}"},
|
||||
"headers": {"Authorization": f"Bearer {_cohere_api_key}"},
|
||||
}
|
||||
]
|
||||
|
||||
|
|
@ -293,7 +293,7 @@ async def test_pass_through_endpoint_sequential_rpm_limit(
|
|||
"path": "/v1/rerank",
|
||||
"target": "https://api.cohere.com/v1/rerank",
|
||||
"auth": auth,
|
||||
"headers": {"Authorization": f"bearer {_cohere_api_key}"},
|
||||
"headers": {"Authorization": f"Bearer {_cohere_api_key}"},
|
||||
}
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -37,7 +37,11 @@
|
|||
"cache_key": null,
|
||||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 3.5e-05,
|
||||
"additional_headers": {}
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 3.5e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -62,7 +66,9 @@
|
|||
"endTime": "2025-01-16T11:28:55.124353-08:00",
|
||||
"completionStartTime": "2025-01-16T11:28:55.124353-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -70,11 +76,12 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
},
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
},
|
||||
"traceId": "litellm-test-6a51ae70-a4e7-499e-afcd-dce2a3b31850"
|
||||
},
|
||||
"timestamp": "2025-01-16T19:28:55.125258Z"
|
||||
|
|
|
|||
|
|
@ -65,11 +65,12 @@
|
|||
"totalCost": 0.00018
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 10
|
||||
}
|
||||
"input": 10,
|
||||
"output": 10,
|
||||
"total": 20,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-05-26T21:13:16.797156Z"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -78,7 +78,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -103,7 +106,9 @@
|
|||
"endTime": "2025-01-22T09:27:51.702048-08:00",
|
||||
"completionStartTime": "2025-01-22T09:27:51.702048-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -111,10 +116,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:27:51.703046Z"
|
||||
|
|
|
|||
|
|
@ -54,7 +54,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 3.5e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 3.5e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -81,7 +84,9 @@
|
|||
"endTime": "2025-01-22T09:19:11.234200-08:00",
|
||||
"completionStartTime": "2025-01-22T09:19:11.234200-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -91,6 +96,7 @@
|
|||
"usageDetails": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
|
|
|
|||
|
|
@ -33,7 +33,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 3.5e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 3.5e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -52,7 +55,9 @@
|
|||
"endTime": "2025-02-06T16:23:27.644253-08:00",
|
||||
"completionStartTime": "2025-02-06T16:23:27.644253-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 10,
|
||||
|
|
@ -60,10 +65,11 @@
|
|||
"totalCost": 1.9999999999999998e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 10
|
||||
"output": 10,
|
||||
"total": 20,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-02-07T00:23:27.670175Z"
|
||||
|
|
|
|||
|
|
@ -75,10 +75,11 @@
|
|||
"totalCost": 1.9999999999999998e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 10
|
||||
"output": 10,
|
||||
"total": 20,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-05-24T17:01:19.408586Z"
|
||||
|
|
@ -91,4 +92,4 @@
|
|||
"sdk_version": "2.44.1",
|
||||
"public_key": "pk-lf-3bfc4db9-217f-48e9-92e0-142566e3c204"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -46,7 +46,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -71,7 +74,9 @@
|
|||
"endTime": "2025-01-22T07:31:28.962389-08:00",
|
||||
"completionStartTime": "2025-01-22T07:31:28.962389-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -79,10 +84,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T15:31:28.964179Z"
|
||||
|
|
|
|||
|
|
@ -46,7 +46,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -71,7 +74,9 @@
|
|||
"endTime": "2025-01-22T08:38:26.015666-08:00",
|
||||
"completionStartTime": "2025-01-22T08:38:26.015666-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -79,10 +84,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T16:38:26.017252Z"
|
||||
|
|
|
|||
|
|
@ -63,10 +63,11 @@
|
|||
"totalCost": 7.5e-06
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 10
|
||||
"output": 10,
|
||||
"total": 20,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-05-26T21:15:40.610953Z"
|
||||
|
|
|
|||
|
|
@ -53,7 +53,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -78,7 +81,9 @@
|
|||
"endTime": "2025-01-22T09:59:39.365756-08:00",
|
||||
"completionStartTime": "2025-01-22T09:59:39.365756-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -86,10 +91,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:59:39.368310Z"
|
||||
|
|
|
|||
|
|
@ -45,7 +45,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -70,7 +73,9 @@
|
|||
"endTime": "2025-01-22T10:06:50.958374-08:00",
|
||||
"completionStartTime": "2025-01-22T10:06:50.958374-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -78,10 +83,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T18:06:50.959850Z"
|
||||
|
|
|
|||
|
|
@ -39,7 +39,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -64,7 +67,9 @@
|
|||
"endTime": "2025-01-22T09:59:32.880691-08:00",
|
||||
"completionStartTime": "2025-01-22T09:59:32.880691-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -72,10 +77,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:59:32.889548Z"
|
||||
|
|
|
|||
|
|
@ -39,7 +39,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -64,7 +67,9 @@
|
|||
"endTime": "2025-01-22T09:59:36.161959-08:00",
|
||||
"completionStartTime": "2025-01-22T09:59:36.161959-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -72,10 +77,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:59:36.162997Z"
|
||||
|
|
|
|||
|
|
@ -39,7 +39,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -64,7 +67,9 @@
|
|||
"endTime": "2025-01-22T09:59:32.880691-08:00",
|
||||
"completionStartTime": "2025-01-22T09:59:32.880691-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -72,10 +77,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:59:32.889548Z"
|
||||
|
|
|
|||
|
|
@ -45,7 +45,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -70,7 +73,9 @@
|
|||
"endTime": "2025-01-22T09:55:28.853979-08:00",
|
||||
"completionStartTime": "2025-01-22T09:55:28.853979-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -78,10 +83,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:55:28.855732Z"
|
||||
|
|
|
|||
|
|
@ -45,7 +45,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -70,7 +73,9 @@
|
|||
"endTime": "2025-01-22T09:53:53.753431-08:00",
|
||||
"completionStartTime": "2025-01-22T09:53:53.753431-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -78,10 +83,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:53:53.754511Z"
|
||||
|
|
|
|||
|
|
@ -49,7 +49,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -74,7 +77,9 @@
|
|||
"endTime": "2025-01-22T09:56:35.476236-08:00",
|
||||
"completionStartTime": "2025-01-22T09:56:35.476236-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -82,10 +87,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:56:35.478171Z"
|
||||
|
|
|
|||
|
|
@ -53,7 +53,10 @@
|
|||
"api_base": "https://api.openai.com",
|
||||
"response_cost": 5.4999999999999995e-05,
|
||||
"additional_headers": {},
|
||||
"litellm_overhead_time_ms": null
|
||||
"litellm_overhead_time_ms": null,
|
||||
"batch_models": null,
|
||||
"litellm_model_name": "gpt-3.5-turbo",
|
||||
"usage_object": null
|
||||
},
|
||||
"litellm_response_cost": 5.4999999999999995e-05,
|
||||
"cache_hit": false,
|
||||
|
|
@ -78,7 +81,9 @@
|
|||
"endTime": "2025-01-22T09:56:38.785762-08:00",
|
||||
"completionStartTime": "2025-01-22T09:56:38.785762-08:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"modelParameters": {"extra_body": "{}"},
|
||||
"modelParameters": {
|
||||
"extra_body": "{}"
|
||||
},
|
||||
"usage": {
|
||||
"input": 10,
|
||||
"output": 20,
|
||||
|
|
@ -86,10 +91,11 @@
|
|||
"totalCost": 3.5e-05
|
||||
},
|
||||
"usageDetails": {
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"input": 10,
|
||||
"output": 20
|
||||
"output": 20,
|
||||
"total": 30,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
},
|
||||
"timestamp": "2025-01-22T17:56:38.787196Z"
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ os.environ["LANGFUSE_DEBUG"] = "True"
|
|||
import time
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
|
||||
def assert_langfuse_request_matches_expected(
|
||||
|
|
@ -116,7 +117,7 @@ def assert_langfuse_request_matches_expected(
|
|||
|
||||
|
||||
class TestLangfuseLogging:
|
||||
@pytest.fixture
|
||||
@pytest_asyncio.fixture
|
||||
async def mock_setup(self):
|
||||
"""Common setup for Langfuse logging tests"""
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -168,7 +169,7 @@ class TestLangfuseLogging:
|
|||
@pytest.mark.asyncio
|
||||
async def test_langfuse_logging_completion(self, mock_setup):
|
||||
"""Test Langfuse logging for chat completion"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
setup = mock_setup
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -183,7 +184,7 @@ class TestLangfuseLogging:
|
|||
@pytest.mark.asyncio
|
||||
async def test_langfuse_logging_completion_with_tags(self, mock_setup):
|
||||
"""Test Langfuse logging for chat completion with tags"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
setup = mock_setup
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -201,7 +202,7 @@ class TestLangfuseLogging:
|
|||
@pytest.mark.asyncio
|
||||
async def test_langfuse_logging_completion_with_tags_stream(self, mock_setup):
|
||||
"""Test Langfuse logging for chat completion with tags"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
setup = mock_setup
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -221,7 +222,7 @@ class TestLangfuseLogging:
|
|||
@pytest.mark.asyncio
|
||||
async def test_langfuse_logging_completion_with_langfuse_metadata(self, mock_setup):
|
||||
"""Test Langfuse logging for chat completion with metadata for langfuse"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
setup = mock_setup
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
await litellm.acompletion(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -259,7 +260,7 @@ class TestLangfuseLogging:
|
|||
last_login: datetime.datetime
|
||||
settings: dict
|
||||
|
||||
setup = await mock_setup
|
||||
setup = mock_setup
|
||||
|
||||
test_metadata = {
|
||||
"user_prefs": UserPreferences(
|
||||
|
|
@ -334,7 +335,7 @@ class TestLangfuseLogging:
|
|||
"""Test Langfuse logging with various metadata types including non-serializable objects"""
|
||||
import threading
|
||||
|
||||
setup = await mock_setup
|
||||
setup = mock_setup
|
||||
|
||||
if test_metadata is not None:
|
||||
test_metadata["trace_id"] = setup["trace_id"]
|
||||
|
|
@ -358,7 +359,7 @@ class TestLangfuseLogging:
|
|||
self, mock_setup
|
||||
):
|
||||
"""Test Langfuse logging for chat completion with malformed LLM response"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
setup = mock_setup
|
||||
litellm._turn_on_debug()
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
mock_response = litellm.ModelResponse(
|
||||
|
|
@ -387,7 +388,7 @@ class TestLangfuseLogging:
|
|||
self, mock_setup
|
||||
):
|
||||
"""Test Langfuse logging for chat completion with malformed LLM response"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
setup = mock_setup
|
||||
litellm._turn_on_debug()
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
mock_response = litellm.ModelResponse(
|
||||
|
|
@ -418,7 +419,7 @@ class TestLangfuseLogging:
|
|||
self, mock_setup
|
||||
):
|
||||
"""Test Langfuse logging for chat completion with malformed LLM response"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
setup = mock_setup
|
||||
litellm._turn_on_debug()
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
mock_response = litellm.ModelResponse(
|
||||
|
|
@ -447,7 +448,6 @@ class TestLangfuseLogging:
|
|||
@pytest.mark.asyncio
|
||||
async def test_langfuse_logging_with_router(self, mock_setup):
|
||||
"""Test Langfuse logging with router"""
|
||||
setup = await mock_setup # Await the fixture
|
||||
litellm._turn_on_debug()
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
|
|
@ -461,7 +461,7 @@ class TestLangfuseLogging:
|
|||
}
|
||||
]
|
||||
)
|
||||
with patch("httpx.Client.post", setup["mock_post"]):
|
||||
with patch("httpx.Client.post", mock_setup["mock_post"]):
|
||||
mock_response = litellm.ModelResponse(
|
||||
choices=[],
|
||||
usage=litellm.Usage(
|
||||
|
|
@ -477,8 +477,8 @@ class TestLangfuseLogging:
|
|||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "Hello!"}],
|
||||
mock_response=mock_response,
|
||||
metadata={"trace_id": setup["trace_id"]},
|
||||
metadata={"trace_id": mock_setup["trace_id"]},
|
||||
)
|
||||
await self._verify_langfuse_call(
|
||||
setup["mock_post"], "completion_with_router.json", setup["trace_id"]
|
||||
mock_setup["mock_post"], "completion_with_router.json", mock_setup["trace_id"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -102,7 +102,7 @@ def test_openai_assistants_e2e_operations_stream():
|
|||
def test_azure_openai_assistants_e2e_operations_stream():
|
||||
from openai import AzureOpenAI
|
||||
client = AzureOpenAI(
|
||||
base_url="http://0.0.0.0:4000/azure-config-passthrough",
|
||||
base_url="http://0.0.0.0:4000/azure-config-passthrough/openai",
|
||||
api_key="sk-1234",
|
||||
api_version="2025-01-01-preview"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,13 +4,13 @@ from fastapi.testclient import TestClient
|
|||
from litellm.proxy.proxy_server import app, ProxyLogging
|
||||
from litellm.caching import DualCache
|
||||
|
||||
TEST_DB_ENV_VAR_NAME = "MASTER_KEY_CHECK_DB_URL"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def override_env_settings(monkeypatch):
|
||||
# Set environment variables only for tests using-monkeypatch (function scope by default).
|
||||
monkeypatch.setenv("DATABASE_URL", os.environ[TEST_DB_ENV_VAR_NAME])
|
||||
# Use DATABASE_URL from environment (set by CircleCI to local postgres)
|
||||
if "DATABASE_URL" not in os.environ:
|
||||
pytest.fail("DATABASE_URL not set - this test requires a local postgres database to be running")
|
||||
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-1234")
|
||||
monkeypatch.setenv("LITELLM_LOG", "DEBUG")
|
||||
|
||||
|
|
@ -38,7 +38,7 @@ async def test_master_key_not_inserted(test_client):
|
|||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
prisma_client = PrismaClient(
|
||||
database_url=os.environ[TEST_DB_ENV_VAR_NAME],
|
||||
database_url=os.environ["DATABASE_URL"],
|
||||
proxy_logging_obj=ProxyLogging(
|
||||
user_api_key_cache=DualCache(), premium_user=True
|
||||
),
|
||||
|
|
|
|||
|
|
@ -1511,7 +1511,10 @@ def test_key_generate_with_custom_auth(prisma_client):
|
|||
asyncio.run(test())
|
||||
except Exception as e:
|
||||
print("Got Exception", e)
|
||||
print(e.message)
|
||||
if hasattr(e, "message"):
|
||||
print(e.message)
|
||||
else:
|
||||
print(e)
|
||||
pytest.fail(f"An exception occurred - {str(e)}")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -338,12 +338,12 @@ def test_twelvelabs_input_type_parameter_mapping_async_invoke():
|
|||
|
||||
|
||||
def test_twelvelabs_missing_input_type_error():
|
||||
"""Test that missing input_type parameter throws an error for TwelveLabs models but not others"""
|
||||
"""Test that missing input_type parameter defaults to 'text' for TwelveLabs models"""
|
||||
litellm.set_verbose = True
|
||||
client = HTTPHandler()
|
||||
test_api_key = "test-bearer-token-12345"
|
||||
|
||||
# Test TwelveLabs model - should throw error
|
||||
# Test TwelveLabs model - should default to 'text' when input_type is missing
|
||||
twelvelabs_model = "bedrock/twelvelabs.marengo-embed-2-7-v1:0"
|
||||
twelvelabs_response = {
|
||||
"data": [{
|
||||
|
|
@ -359,20 +359,24 @@ def test_twelvelabs_missing_input_type_error():
|
|||
mock_response.json = lambda: json.loads(mock_response.text)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Test that missing input_type throws an error for TwelveLabs
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
litellm.embedding(
|
||||
model=twelvelabs_model,
|
||||
input=test_input,
|
||||
client=client,
|
||||
aws_region_name="us-east-1",
|
||||
aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com",
|
||||
api_key=test_api_key
|
||||
# No input_type parameter - should throw an error
|
||||
)
|
||||
# Test that missing input_type defaults to "text" for TwelveLabs
|
||||
response = litellm.embedding(
|
||||
model=twelvelabs_model,
|
||||
input=test_input,
|
||||
client=client,
|
||||
aws_region_name="us-east-1",
|
||||
aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com",
|
||||
api_key=test_api_key
|
||||
# No input_type parameter - should default to "text"
|
||||
)
|
||||
|
||||
# Verify the error message contains the expected text
|
||||
assert "input_type is required" in str(exc_info.value)
|
||||
# Verify the response is successful
|
||||
assert isinstance(response, litellm.EmbeddingResponse)
|
||||
|
||||
# Verify that the request contains inputType: "text" by default
|
||||
request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}"))
|
||||
assert "inputType" in request_body
|
||||
assert request_body["inputType"] == "text"
|
||||
|
||||
# Test Amazon Titan model - should NOT throw error (input_type not required)
|
||||
titan_model = "bedrock/amazon.titan-embed-text-v1"
|
||||
|
|
|
|||
|
|
@ -88,10 +88,14 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
|
|||
working_server = MagicMock()
|
||||
working_server.name = "working_server"
|
||||
working_server.alias = "working"
|
||||
working_server.allowed_tools = None
|
||||
working_server.disallowed_tools = None
|
||||
|
||||
failing_server = MagicMock()
|
||||
failing_server.name = "failing_server"
|
||||
failing_server.alias = "failing"
|
||||
failing_server.allowed_tools = None
|
||||
failing_server.disallowed_tools = None
|
||||
|
||||
# Mock global_mcp_server_manager
|
||||
mock_manager = MagicMock()
|
||||
|
|
@ -586,6 +590,8 @@ async def test_list_tools_single_server_unprefixed_names():
|
|||
server.server_id = "server1"
|
||||
server.name = "Zapier MCP"
|
||||
server.alias = "zapier"
|
||||
server.allowed_tools = None
|
||||
server.disallowed_tools = None
|
||||
|
||||
# Mock manager: allow just one server and return a tool based on add_prefix flag
|
||||
mock_manager = MagicMock()
|
||||
|
|
@ -641,11 +647,15 @@ async def test_list_tools_multiple_servers_prefixed_names():
|
|||
server1.server_id = "server1"
|
||||
server1.name = "Zapier MCP"
|
||||
server1.alias = "zapier"
|
||||
server1.allowed_tools = None
|
||||
server1.disallowed_tools = None
|
||||
|
||||
server2 = MagicMock()
|
||||
server2.server_id = "server2"
|
||||
server2.name = "Jira MCP"
|
||||
server2.alias = "jira"
|
||||
server2.allowed_tools = None
|
||||
server2.disallowed_tools = None
|
||||
|
||||
# Mock manager
|
||||
mock_manager = MagicMock()
|
||||
|
|
|
|||
|
|
@ -654,6 +654,7 @@ class TestMCPServerManager:
|
|||
"Tool tool3 is not allowed for server test-server"
|
||||
in exc_info.value.detail["error"]
|
||||
)
|
||||
|
||||
async def test_get_tools_from_server_add_prefix(self):
|
||||
"""Verify _get_tools_from_server respects add_prefix True/False."""
|
||||
manager = MCPServerManager()
|
||||
|
|
@ -909,6 +910,39 @@ class TestMCPServerManager:
|
|||
assert "tool_1" in tool_names
|
||||
assert "tool_2" in tool_names
|
||||
|
||||
def test_add_db_mcp_server_to_registry(self):
|
||||
"""Test that add_db_mcp_server_to_registry adds a MCP server to the registry"""
|
||||
manager = MCPServerManager()
|
||||
server = LiteLLM_MCPServerTable(
|
||||
**{
|
||||
"server_id": "4c679a81-acd9-4954-9f84-30b739362498",
|
||||
"server_name": "edc_mcp_server",
|
||||
"alias": "edc_mcp_server",
|
||||
"description": None,
|
||||
"url": "fake_mcp_url",
|
||||
"transport": "http",
|
||||
"auth_type": "none",
|
||||
"created_at": "2025-09-30T08:28:31.353000Z",
|
||||
"created_by": "a1248959",
|
||||
"updated_at": "2025-09-30T08:28:31.353000Z",
|
||||
"updated_by": "a1248959",
|
||||
"teams": [],
|
||||
"mcp_access_groups": [],
|
||||
"mcp_info": {
|
||||
"server_name": "edc_mcp_server",
|
||||
"mcp_server_cost_info": None,
|
||||
},
|
||||
"status": "unknown",
|
||||
"last_health_check": None,
|
||||
"health_check_error": None,
|
||||
"command": None,
|
||||
"args": [],
|
||||
"env": {},
|
||||
},
|
||||
)
|
||||
manager.add_update_server(server)
|
||||
assert server.server_id in manager.get_registry()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -409,7 +409,7 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
return 1800 # 1800/2000 = 90% saturation
|
||||
return None
|
||||
|
||||
async def mock_should_rate_limit(descriptors, parent_otel_span=None):
|
||||
async def mock_should_rate_limit(descriptors, parent_otel_span=None, read_only=False):
|
||||
"""Mock rate limiter that handles saturation-aware descriptors."""
|
||||
descriptor = descriptors[0]
|
||||
descriptor_key = descriptor["key"]
|
||||
|
|
@ -431,48 +431,48 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
}
|
||||
|
||||
# Handle priority-specific enforcement in strict mode
|
||||
if descriptor_key == "priority_model":
|
||||
elif descriptor_key == "priority_model":
|
||||
# Extract priority from value like "pre-call-stress-model:premium"
|
||||
priority = descriptor_value.split(":")[-1]
|
||||
|
||||
if priority == "premium":
|
||||
# Allow all premium requests
|
||||
return {
|
||||
"overall_code": "OK",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OK",
|
||||
"descriptor_key": descriptor_value,
|
||||
"rate_limit_type": "tokens_per_unit",
|
||||
"limit_remaining": 1000,
|
||||
}
|
||||
],
|
||||
}
|
||||
else:
|
||||
# Rate limit some standard requests (simulate load)
|
||||
import random
|
||||
|
||||
if random.random() < 0.3: # 30% of standard requests get rate limited
|
||||
return {
|
||||
"overall_code": "OVER_LIMIT",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OVER_LIMIT",
|
||||
"descriptor_key": descriptor_value,
|
||||
"rate_limit_type": "tokens_per_unit",
|
||||
"limit_remaining": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
else:
|
||||
if priority == "premium":
|
||||
# Allow all premium requests
|
||||
return {
|
||||
"overall_code": "OK",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OK",
|
||||
"descriptor_key": descriptor_value,
|
||||
"descriptor_key": descriptor_value,
|
||||
"rate_limit_type": "tokens_per_unit",
|
||||
"limit_remaining": 100,
|
||||
"limit_remaining": 1000,
|
||||
}
|
||||
],
|
||||
}
|
||||
else:
|
||||
# Rate limit some standard requests (simulate load)
|
||||
import random
|
||||
|
||||
if random.random() < 0.3: # 30% of standard requests get rate limited
|
||||
return {
|
||||
"overall_code": "OVER_LIMIT",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OVER_LIMIT",
|
||||
"descriptor_key": descriptor_value,
|
||||
"rate_limit_type": "tokens_per_unit",
|
||||
"limit_remaining": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
else:
|
||||
return {
|
||||
"overall_code": "OK",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OK",
|
||||
"descriptor_key": descriptor_value,
|
||||
"rate_limit_type": "tokens_per_unit",
|
||||
"limit_remaining": 100,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
|
@ -486,9 +486,9 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
"descriptor_key": descriptor_value,
|
||||
"rate_limit_type": "tokens_per_unit",
|
||||
"limit_remaining": 1000,
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
# Create 50 users: 30 premium, 20 standard
|
||||
users = []
|
||||
|
|
@ -509,44 +509,44 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
"""Make a pre-call hook request."""
|
||||
user, priority = user_data
|
||||
|
||||
with patch.object(
|
||||
handler.v3_limiter, "should_rate_limit", side_effect=mock_should_rate_limit
|
||||
), patch.object(
|
||||
handler.internal_usage_cache, "async_get_cache", side_effect=mock_get_cache
|
||||
):
|
||||
try:
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user,
|
||||
cache=DualCache(),
|
||||
data={"model": model},
|
||||
call_type="completion",
|
||||
)
|
||||
try:
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user,
|
||||
cache=DualCache(),
|
||||
data={"model": model},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
# If no exception, request was allowed
|
||||
successful_requests.append(
|
||||
{"user_id": user.user_id, "priority": priority, "result": "allowed"}
|
||||
)
|
||||
return {
|
||||
"status": "success",
|
||||
"user_id": user.user_id,
|
||||
"priority": priority,
|
||||
}
|
||||
# If no exception, request was allowed
|
||||
successful_requests.append(
|
||||
{"user_id": user.user_id, "priority": priority, "result": "allowed"}
|
||||
)
|
||||
return {
|
||||
"status": "success",
|
||||
"user_id": user.user_id,
|
||||
"priority": priority,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
# Request was rate limited
|
||||
rate_limited_requests.append(
|
||||
{"user_id": user.user_id, "priority": priority, "error": str(e)}
|
||||
)
|
||||
return {
|
||||
"status": "rate_limited",
|
||||
"user_id": user.user_id,
|
||||
"priority": priority,
|
||||
}
|
||||
except Exception as e:
|
||||
# Request was rate limited
|
||||
rate_limited_requests.append(
|
||||
{"user_id": user.user_id, "priority": priority, "error": str(e)}
|
||||
)
|
||||
return {
|
||||
"status": "rate_limited",
|
||||
"user_id": user.user_id,
|
||||
"priority": priority,
|
||||
}
|
||||
|
||||
# Run all 50 requests concurrently
|
||||
# Run all 50 requests concurrently with patches applied to the entire batch
|
||||
start_time = time.time()
|
||||
tasks = [make_request(user_data) for user_data in users]
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
with patch.object(
|
||||
handler.v3_limiter, "should_rate_limit", side_effect=mock_should_rate_limit
|
||||
), patch.object(
|
||||
handler.internal_usage_cache, "async_get_cache", side_effect=mock_get_cache
|
||||
):
|
||||
tasks = [make_request(user_data) for user_data in users]
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
end_time = time.time()
|
||||
|
||||
# Analyze results
|
||||
|
|
@ -582,9 +582,13 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
assert (
|
||||
standard_success_rate >= 0.5
|
||||
), f"Standard success rate should be >= 50% (with 30% random limiting, allows for variance), got {standard_success_rate:.2%}"
|
||||
assert (
|
||||
premium_success_rate > standard_success_rate
|
||||
), "Premium should have higher success rate than standard"
|
||||
|
||||
# Allow for the case where both are 100% due to timing/mocking issues
|
||||
# The test is inherently flaky due to random behavior
|
||||
if premium_success_rate < 1.0 or standard_success_rate < 1.0:
|
||||
assert (
|
||||
premium_success_rate >= standard_success_rate
|
||||
), "Premium should have >= success rate than standard"
|
||||
|
||||
total_duration = end_time - start_time
|
||||
|
||||
|
|
@ -604,17 +608,19 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
@pytest.mark.asyncio
|
||||
async def test_fake_calls_case_1_no_rate_limiting_at_capacity():
|
||||
"""
|
||||
Test Case 1: No Rate Limiting When At Capacity
|
||||
Test Case 1: Saturation-Aware Rate Limiting at 50% Threshold
|
||||
|
||||
System: 100 RPM capacity
|
||||
System: 100 RPM capacity, saturation_threshold=50%
|
||||
Key A: priority_reservation=0.75 (75 RPM reserved)
|
||||
Key B: priority_reservation=0.25 (25 RPM reserved)
|
||||
Traffic A: 50 RPM
|
||||
Traffic B: 50 RPM
|
||||
Expected A: 50 RPM (no limiting, under reserved capacity)
|
||||
Expected B: 50 RPM (no limiting, under reserved capacity)
|
||||
Traffic A: 1 request
|
||||
Traffic B: 100 requests
|
||||
|
||||
When traffic is under individual reservations, no rate limiting should occur.
|
||||
Expected behavior:
|
||||
- Key A: 1 request succeeds (low traffic)
|
||||
- Key B: ~25-26 requests succeed (capped at reservation when saturation >= 50%)
|
||||
|
||||
Once saturation hits 50%, strict mode enforces priority-based limits.
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
|
||||
|
|
@ -676,13 +682,13 @@ async def test_fake_calls_case_1_no_rate_limiting_at_capacity():
|
|||
rate_limited_requests[priority_name] += 1
|
||||
return {"status": "rate_limited", "priority": priority_name, "error": str(e)}
|
||||
|
||||
# Send 50 requests from each priority (within capacity)
|
||||
# Send 1 request from key_a, 100 from key_b
|
||||
tasks = []
|
||||
|
||||
for i in range(50):
|
||||
for i in range(1):
|
||||
tasks.append(make_request(key_a_user, "key_a", f"key_a_{i}"))
|
||||
|
||||
for i in range(50):
|
||||
for i in range(100):
|
||||
tasks.append(make_request(key_b_user, "key_b", f"key_b_{i}"))
|
||||
|
||||
start_time = time.time()
|
||||
|
|
@ -693,16 +699,23 @@ async def test_fake_calls_case_1_no_rate_limiting_at_capacity():
|
|||
total_successful = successful_requests["key_a"] + successful_requests["key_b"]
|
||||
total_rate_limited = rate_limited_requests["key_a"] + rate_limited_requests["key_b"]
|
||||
|
||||
print(f"Test Case 1 - No Rate Limiting When At Capacity:")
|
||||
print(f"Test Case 1 - Saturation-Aware Rate Limiting:")
|
||||
print(f" - Duration: {end_time - start_time:.2f}s")
|
||||
print(f" - Key A: {successful_requests['key_a']}/50 successful (reserved 75 RPM)")
|
||||
print(f" - Key B: {successful_requests['key_b']}/50 successful (reserved 25 RPM)")
|
||||
print(f" - Total successful: {total_successful}/100")
|
||||
print(f" - Total rate limited: {total_rate_limited}/100")
|
||||
print(f" - Key A: {successful_requests['key_a']}/1 successful (reserved 75 RPM)")
|
||||
print(f" - Key B: {successful_requests['key_b']}/100 successful (reserved 25 RPM)")
|
||||
print(f" - Total successful: {total_successful}/101")
|
||||
print(f" - Total rate limited: {total_rate_limited}/101")
|
||||
|
||||
# Both keys should get all their requests since they're under capacity
|
||||
assert successful_requests["key_a"] >= 45, f"Key A should get ≥45 requests, got {successful_requests['key_a']}"
|
||||
assert successful_requests["key_b"] >= 45, f"Key B should get ≥45 requests, got {successful_requests['key_b']}"
|
||||
# Key A should get its 1 request
|
||||
assert successful_requests["key_a"] == 1, f"Key A should get 1 request, got {successful_requests['key_a']}"
|
||||
|
||||
# Key B can send until saturation hits 50% (which is ~50 total requests)
|
||||
# After that, strict mode enforces its 25 RPM reservation
|
||||
# Due to race conditions in concurrent execution, allow 45-52 successful requests
|
||||
assert 45 <= successful_requests["key_b"] <= 52, f"Key B should get ~49 requests (45-52), got {successful_requests['key_b']}"
|
||||
|
||||
# Verify approximately half of key_b requests were rate limited
|
||||
assert rate_limited_requests["key_b"] >= 45, f"Key B should have ≥45 rate limited requests, got {rate_limited_requests['key_b']}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1202,3 +1215,89 @@ async def test_fake_calls_case_5_default_value_priority_reservation():
|
|||
if total_successful > 0:
|
||||
key_a_share = successful_requests["key_a"] / total_successful
|
||||
print(f" - Key A got {key_a_share:.1%} of successful requests (expected ~55-62%)")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_priority_shared_pool():
|
||||
"""
|
||||
Test that keys without explicit priority share ONE default pool, not get individual allocations.
|
||||
|
||||
With default_priority=0.25:
|
||||
- Key A, B, C (no priority) should share ONE 25 RPM pool
|
||||
- NOT get 25 RPM each (which would be 75 RPM total)
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
|
||||
litellm.priority_reservation = {"prod": 0.75}
|
||||
litellm.priority_reservation_settings.default_priority = 0.25
|
||||
|
||||
dual_cache = DualCache()
|
||||
handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache)
|
||||
|
||||
model = "test-default-pool"
|
||||
total_rpm = 100
|
||||
|
||||
llm_router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": model,
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "test-key",
|
||||
"api_base": "test-base",
|
||||
"rpm": total_rpm,
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
handler.update_variables(llm_router=llm_router)
|
||||
|
||||
# Create 3 users without explicit priority
|
||||
user_a = UserAPIKeyAuth()
|
||||
user_a.metadata = {}
|
||||
user_a.user_id = "user_a"
|
||||
|
||||
user_b = UserAPIKeyAuth()
|
||||
user_b.metadata = {}
|
||||
user_b.user_id = "user_b"
|
||||
|
||||
user_c = UserAPIKeyAuth()
|
||||
user_c.metadata = {}
|
||||
user_c.user_id = "user_c"
|
||||
|
||||
# Get descriptors for each
|
||||
desc_a = handler._create_priority_based_descriptors(
|
||||
model=model, user_api_key_dict=user_a, priority=None
|
||||
)
|
||||
desc_b = handler._create_priority_based_descriptors(
|
||||
model=model, user_api_key_dict=user_b, priority=None
|
||||
)
|
||||
desc_c = handler._create_priority_based_descriptors(
|
||||
model=model, user_api_key_dict=user_c, priority=None
|
||||
)
|
||||
|
||||
# All should use the SAME shared pool key
|
||||
assert desc_a[0]["value"] == f"{model}:default_pool"
|
||||
assert desc_b[0]["value"] == f"{model}:default_pool"
|
||||
assert desc_c[0]["value"] == f"{model}:default_pool"
|
||||
|
||||
# All should have same limit (25 RPM SHARED, not 25 RPM each)
|
||||
assert desc_a[0]["rate_limit"]["requests_per_unit"] == 25
|
||||
assert desc_b[0]["rate_limit"]["requests_per_unit"] == 25
|
||||
assert desc_c[0]["rate_limit"]["requests_per_unit"] == 25
|
||||
|
||||
# Verify explicit priority uses different pool
|
||||
user_prod = UserAPIKeyAuth()
|
||||
user_prod.metadata = {"priority": "prod"}
|
||||
desc_prod = handler._create_priority_based_descriptors(
|
||||
model=model, user_api_key_dict=user_prod, priority="prod"
|
||||
)
|
||||
|
||||
assert desc_prod[0]["value"] == f"{model}:prod"
|
||||
assert desc_prod[0]["rate_limit"]["requests_per_unit"] == 75
|
||||
assert desc_prod[0]["value"] != desc_a[0]["value"] # Different pools
|
||||
|
||||
print("✅ Default priority test passed:")
|
||||
print(f" - 3 keys without priority share ONE pool: {desc_a[0]['value']}")
|
||||
print(f" - Shared pool limit: {desc_a[0]['rate_limit']['requests_per_unit']} RPM")
|
||||
print(f" - Explicit priority 'prod' uses separate pool: {desc_prod[0]['value']}")
|
||||
|
|
|
|||
|
|
@ -15,14 +15,18 @@ from fastapi import HTTPException
|
|||
|
||||
from litellm.proxy._types import (
|
||||
GenerateKeyRequest,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_VerificationToken,
|
||||
LitellmUserRoles,
|
||||
ProxyException,
|
||||
UpdateKeyRequest,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_check_team_key_limits,
|
||||
_common_key_generation_helper,
|
||||
_list_key_helper,
|
||||
check_team_key_model_specific_limits,
|
||||
generate_key_helper_fn,
|
||||
prepare_key_update_data,
|
||||
validate_key_team_change,
|
||||
|
|
@ -847,17 +851,24 @@ async def test_generate_service_account_key_endpoint_validation():
|
|||
)
|
||||
|
||||
# Test case 1: Missing team_id
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await generate_service_account_key_fn(
|
||||
data=GenerateKeyRequest(team_id=None),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1"
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||||
# Mock prisma_client to be not None so we can reach team_id validation
|
||||
mock_prisma_instance = AsyncMock()
|
||||
mock_prisma.return_value = mock_prisma_instance
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "team_id is required for service account keys" in str(exc_info.value.detail)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await generate_service_account_key_fn(
|
||||
data=GenerateKeyRequest(team_id=None),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1"
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "team_id is required for service account keys" in str(
|
||||
exc_info.value.detail
|
||||
)
|
||||
|
||||
# Test case 2: Team doesn't exist in database
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||||
|
|
@ -1040,7 +1051,7 @@ async def test_unblock_key_invalid_key_format(monkeypatch):
|
|||
def test_validate_key_team_change_with_member_permissions():
|
||||
"""
|
||||
Test validate_key_team_change function with team member permissions.
|
||||
|
||||
|
||||
This test covers the new logic that allows team members with specific
|
||||
permissions to update keys, not just team admins.
|
||||
"""
|
||||
|
|
@ -1054,111 +1065,107 @@ def test_validate_key_team_change_with_member_permissions():
|
|||
mock_key.models = ["gpt-4"]
|
||||
mock_key.tpm_limit = None
|
||||
mock_key.rpm_limit = None
|
||||
|
||||
|
||||
mock_team = MagicMock()
|
||||
mock_team.team_id = "test-team-456"
|
||||
mock_team.team_id = "test-team-456"
|
||||
mock_team.members_with_roles = []
|
||||
mock_team.tpm_limit = None
|
||||
mock_team.rpm_limit = None
|
||||
|
||||
|
||||
mock_change_initiator = MagicMock()
|
||||
mock_change_initiator.user_id = "test-user-123"
|
||||
|
||||
|
||||
mock_router = MagicMock()
|
||||
|
||||
|
||||
# Mock the member object returned by _get_user_in_team
|
||||
mock_member_object = MagicMock()
|
||||
|
||||
with patch('litellm.proxy.management_endpoints.key_management_endpoints.can_team_access_model'):
|
||||
with patch('litellm.proxy.management_endpoints.key_management_endpoints._get_user_in_team') as mock_get_user:
|
||||
with patch('litellm.proxy.management_endpoints.key_management_endpoints._is_user_team_admin') as mock_is_admin:
|
||||
with patch('litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint') as mock_has_perms:
|
||||
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.can_team_access_model"
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._get_user_in_team"
|
||||
) as mock_get_user:
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._is_user_team_admin"
|
||||
) as mock_is_admin:
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint"
|
||||
) as mock_has_perms:
|
||||
|
||||
mock_get_user.return_value = mock_member_object
|
||||
mock_is_admin.return_value = False
|
||||
mock_has_perms.return_value = True
|
||||
|
||||
|
||||
# This should not raise an exception due to member permissions
|
||||
validate_key_team_change(
|
||||
key=mock_key,
|
||||
team=mock_team,
|
||||
change_initiated_by=mock_change_initiator,
|
||||
llm_router=mock_router
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
|
||||
# Verify the permission check was called with correct parameters
|
||||
mock_has_perms.assert_called_once_with(
|
||||
team_member_object=mock_member_object,
|
||||
team_table=mock_team,
|
||||
route=KeyManagementRoutes.KEY_UPDATE.value
|
||||
route=KeyManagementRoutes.KEY_UPDATE.value,
|
||||
)
|
||||
|
||||
|
||||
def test_key_rotation_fields_helper():
|
||||
"""
|
||||
Test the key data update logic for rotation fields.
|
||||
|
||||
|
||||
This test focuses on the core logic that adds rotation fields to key_data
|
||||
when auto_rotate is enabled, without the complexity of full key generation.
|
||||
"""
|
||||
# Test Case 1: With rotation enabled
|
||||
key_data = {
|
||||
"models": ["gpt-3.5-turbo"],
|
||||
"user_id": "test-user"
|
||||
}
|
||||
|
||||
key_data = {"models": ["gpt-3.5-turbo"], "user_id": "test-user"}
|
||||
|
||||
auto_rotate = True
|
||||
rotation_interval = "30d"
|
||||
|
||||
|
||||
# Simulate the rotation logic from generate_key_helper_fn
|
||||
if auto_rotate and rotation_interval:
|
||||
key_data.update({
|
||||
"auto_rotate": auto_rotate,
|
||||
"rotation_interval": rotation_interval
|
||||
})
|
||||
|
||||
key_data.update(
|
||||
{"auto_rotate": auto_rotate, "rotation_interval": rotation_interval}
|
||||
)
|
||||
|
||||
# Verify rotation fields are added
|
||||
assert key_data["auto_rotate"] == True
|
||||
assert key_data["rotation_interval"] == "30d"
|
||||
assert key_data["models"] == ["gpt-3.5-turbo"] # Original fields preserved
|
||||
|
||||
|
||||
# Test Case 2: Without rotation enabled
|
||||
key_data2 = {
|
||||
"models": ["gpt-4"],
|
||||
"user_id": "test-user"
|
||||
}
|
||||
|
||||
key_data2 = {"models": ["gpt-4"], "user_id": "test-user"}
|
||||
|
||||
auto_rotate2 = False
|
||||
rotation_interval2 = None
|
||||
|
||||
|
||||
# Simulate the rotation logic
|
||||
if auto_rotate2 and rotation_interval2:
|
||||
key_data2.update({
|
||||
"auto_rotate": auto_rotate2,
|
||||
"rotation_interval": rotation_interval2
|
||||
})
|
||||
|
||||
key_data2.update(
|
||||
{"auto_rotate": auto_rotate2, "rotation_interval": rotation_interval2}
|
||||
)
|
||||
|
||||
# Verify rotation fields are NOT added
|
||||
assert "auto_rotate" not in key_data2
|
||||
assert "rotation_interval" not in key_data2
|
||||
assert key_data2["models"] == ["gpt-4"] # Original fields preserved
|
||||
|
||||
|
||||
# Test Case 3: auto_rotate=True but no interval
|
||||
key_data3 = {
|
||||
"models": ["claude-3"],
|
||||
"user_id": "test-user"
|
||||
}
|
||||
|
||||
key_data3 = {"models": ["claude-3"], "user_id": "test-user"}
|
||||
|
||||
auto_rotate3 = True
|
||||
rotation_interval3 = None
|
||||
|
||||
|
||||
# Simulate the rotation logic
|
||||
if auto_rotate3 and rotation_interval3:
|
||||
key_data3.update({
|
||||
"auto_rotate": auto_rotate3,
|
||||
"rotation_interval": rotation_interval3
|
||||
})
|
||||
|
||||
key_data3.update(
|
||||
{"auto_rotate": auto_rotate3, "rotation_interval": rotation_interval3}
|
||||
)
|
||||
|
||||
# Verify rotation fields are NOT added (missing interval)
|
||||
assert "auto_rotate" not in key_data3
|
||||
assert "rotation_interval" not in key_data3
|
||||
|
|
@ -1181,27 +1188,24 @@ async def test_update_key_fn_auto_rotate_enable():
|
|||
team_id=None,
|
||||
auto_rotate=False,
|
||||
rotation_interval=None,
|
||||
metadata={}
|
||||
metadata={},
|
||||
)
|
||||
|
||||
|
||||
# Test enabling auto rotation
|
||||
update_request = UpdateKeyRequest(
|
||||
key="test-token",
|
||||
auto_rotate=True,
|
||||
rotation_interval="30d"
|
||||
key="test-token", auto_rotate=True, rotation_interval="30d"
|
||||
)
|
||||
|
||||
|
||||
result = await prepare_key_update_data(
|
||||
data=update_request,
|
||||
existing_key_row=existing_key
|
||||
data=update_request, existing_key_row=existing_key
|
||||
)
|
||||
|
||||
|
||||
# Verify rotation fields are included
|
||||
assert result["auto_rotate"] is True
|
||||
assert result["rotation_interval"] == "30d"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_fn_auto_rotate_disable():
|
||||
"""Test that update_key_fn properly handles disabling auto rotation."""
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
|
||||
|
|
@ -1218,19 +1222,520 @@ async def test_update_key_fn_auto_rotate_disable():
|
|||
team_id=None,
|
||||
auto_rotate=True,
|
||||
rotation_interval="30d",
|
||||
metadata={}
|
||||
metadata={},
|
||||
)
|
||||
|
||||
|
||||
# Test disabling auto rotation
|
||||
update_request = UpdateKeyRequest(
|
||||
key="test-token",
|
||||
auto_rotate=False
|
||||
)
|
||||
|
||||
update_request = UpdateKeyRequest(key="test-token", auto_rotate=False)
|
||||
|
||||
result = await prepare_key_update_data(
|
||||
data=update_request,
|
||||
existing_key_row=existing_key
|
||||
data=update_request, existing_key_row=existing_key
|
||||
)
|
||||
|
||||
|
||||
# Verify auto_rotate is set to False
|
||||
assert result["auto_rotate"] is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_no_existing_keys():
|
||||
"""
|
||||
Test _check_team_key_limits when team has no existing keys.
|
||||
Should allow any TPM/RPM limits within team bounds.
|
||||
"""
|
||||
# Mock prisma client
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
|
||||
# Create team table with limits
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-123",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000,
|
||||
rpm_limit=1000,
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request with limits within team bounds
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=5000,
|
||||
rpm_limit=500,
|
||||
tpm_limit_type="guaranteed_throughput",
|
||||
rpm_limit_type="guaranteed_throughput",
|
||||
)
|
||||
|
||||
# Should not raise any exception
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
# Verify database was queried
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many.assert_called_once_with(
|
||||
where={"team_id": "test-team-123"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_with_existing_keys_within_bounds():
|
||||
"""
|
||||
Test _check_team_key_limits when team has existing keys but total allocation
|
||||
is still within team limits.
|
||||
"""
|
||||
# Create mock existing keys
|
||||
existing_key1 = MagicMock()
|
||||
existing_key1.tpm_limit = 3000
|
||||
existing_key1.rpm_limit = 200
|
||||
|
||||
existing_key2 = MagicMock()
|
||||
existing_key2.tpm_limit = 2000
|
||||
existing_key2.rpm_limit = 300
|
||||
|
||||
existing_key3 = MagicMock()
|
||||
existing_key3.tpm_limit = None # Should be ignored in calculation
|
||||
existing_key3.rpm_limit = None # Should be ignored in calculation
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[existing_key1, existing_key2, existing_key3]
|
||||
)
|
||||
|
||||
# Create team table with limits
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-456",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000, # Total: 3000 + 2000 + 4000 (new) = 9000 < 10000 ✓
|
||||
rpm_limit=1000, # Total: 200 + 300 + 400 (new) = 900 < 1000 ✓
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request that would still be within bounds
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=4000,
|
||||
rpm_limit=400,
|
||||
)
|
||||
|
||||
# Should not raise any exception
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_tpm_overallocation():
|
||||
"""
|
||||
Test _check_team_key_limits when new key would cause TPM overallocation.
|
||||
Should raise HTTPException with appropriate error message.
|
||||
"""
|
||||
# Create mock existing keys with high TPM usage
|
||||
existing_key1 = MagicMock()
|
||||
existing_key1.tpm_limit = 6000
|
||||
existing_key1.rpm_limit = 100
|
||||
existing_key1.metadata = {}
|
||||
|
||||
existing_key2 = MagicMock()
|
||||
existing_key2.tpm_limit = 3000
|
||||
existing_key2.rpm_limit = 200
|
||||
existing_key2.metadata = {}
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[existing_key1, existing_key2]
|
||||
)
|
||||
|
||||
# Create team table with limits
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-789",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000, # Allocated: 6000 + 3000 = 9000, New: 2000, Total: 11000 > 10000 ✗
|
||||
rpm_limit=1000,
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request that would exceed TPM limits
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=2000,
|
||||
rpm_limit=100,
|
||||
tpm_limit_type="guaranteed_throughput",
|
||||
)
|
||||
|
||||
# Should raise HTTPException for TPM overallocation
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert (
|
||||
"Allocated TPM limit=9000 + Key TPM limit=2000 is greater than team TPM limit=10000"
|
||||
in str(exc_info.value.detail)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_rpm_overallocation():
|
||||
"""
|
||||
Test _check_team_key_limits when new key would cause RPM overallocation.
|
||||
Should raise HTTPException with appropriate error message.
|
||||
"""
|
||||
# Create mock existing keys with high RPM usage
|
||||
existing_key1 = MagicMock()
|
||||
existing_key1.tpm_limit = 1000
|
||||
existing_key1.rpm_limit = 600
|
||||
existing_key1.metadata = {}
|
||||
|
||||
existing_key2 = MagicMock()
|
||||
existing_key2.tpm_limit = 2000
|
||||
existing_key2.rpm_limit = 300
|
||||
existing_key2.metadata = {}
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[existing_key1, existing_key2]
|
||||
)
|
||||
|
||||
# Create team table with limits
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-101",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000,
|
||||
rpm_limit=1000, # Allocated: 600 + 300 = 900, New: 200, Total: 1100 > 1000 ✗
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request that would exceed RPM limits
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=1000,
|
||||
rpm_limit=200,
|
||||
rpm_limit_type="guaranteed_throughput",
|
||||
)
|
||||
|
||||
# Should raise HTTPException for RPM overallocation
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert (
|
||||
"Allocated RPM limit=900 + Key RPM limit=200 is greater than team RPM limit=1000"
|
||||
in str(exc_info.value.detail)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_no_team_limits():
|
||||
"""
|
||||
Test _check_team_key_limits when team has no TPM/RPM limits set.
|
||||
Should allow any key limits since there are no team constraints.
|
||||
"""
|
||||
# Create mock existing keys
|
||||
existing_key = MagicMock()
|
||||
existing_key.tpm_limit = 5000
|
||||
existing_key.rpm_limit = 500
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[existing_key]
|
||||
)
|
||||
|
||||
# Create team table with no limits
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-202",
|
||||
team_alias="test-team",
|
||||
tpm_limit=None, # No team limit
|
||||
rpm_limit=None, # No team limit
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request with any limits
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=10000, # High limit should be allowed
|
||||
rpm_limit=2000, # High limit should be allowed
|
||||
)
|
||||
|
||||
# Should not raise any exception
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_no_key_limits():
|
||||
"""
|
||||
Test _check_team_key_limits when new key has no TPM/RPM limits.
|
||||
Should not raise any exceptions since no limits are being allocated.
|
||||
"""
|
||||
# Create mock existing keys
|
||||
existing_key = MagicMock()
|
||||
existing_key.tpm_limit = 8000
|
||||
existing_key.rpm_limit = 800
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[existing_key]
|
||||
)
|
||||
|
||||
# Create team table with limits
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-303",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000,
|
||||
rpm_limit=1000,
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request with no limits
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=None, # No limit being set
|
||||
rpm_limit=None, # No limit being set
|
||||
)
|
||||
|
||||
# Should not raise any exception
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_mixed_scenarios():
|
||||
"""
|
||||
Test _check_team_key_limits with mixed scenarios:
|
||||
- Some existing keys have limits, others don't
|
||||
- New key has only one type of limit
|
||||
- Team has only one type of limit
|
||||
"""
|
||||
# Create mock existing keys with mixed limits
|
||||
existing_key1 = MagicMock()
|
||||
existing_key1.tpm_limit = 3000
|
||||
existing_key1.rpm_limit = None # No RPM limit
|
||||
|
||||
existing_key2 = MagicMock()
|
||||
existing_key2.tpm_limit = None # No TPM limit
|
||||
existing_key2.rpm_limit = 400
|
||||
|
||||
existing_key3 = MagicMock()
|
||||
existing_key3.tpm_limit = 2000
|
||||
existing_key3.rpm_limit = 300
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[existing_key1, existing_key2, existing_key3]
|
||||
)
|
||||
|
||||
# Create team table with only TPM limit
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-404",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000, # Allocated: 3000 + 0 + 2000 = 5000, New: 4000, Total: 9000 < 10000 ✓
|
||||
rpm_limit=None, # No team RPM limit
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request with only TPM limit
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=4000,
|
||||
rpm_limit=None, # No RPM limit being set
|
||||
)
|
||||
|
||||
# Should not raise any exception
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_key_limits_exact_boundary():
|
||||
"""
|
||||
Test _check_team_key_limits when allocation exactly matches team limits.
|
||||
Should allow the allocation (boundary case).
|
||||
"""
|
||||
# Create mock existing keys
|
||||
existing_key = MagicMock()
|
||||
existing_key.tpm_limit = 7000
|
||||
existing_key.rpm_limit = 700
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[existing_key]
|
||||
)
|
||||
|
||||
# Create team table with limits
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-505",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000, # Allocated: 7000, New: 3000, Total: 10000 = 10000 ✓
|
||||
rpm_limit=1000, # Allocated: 700, New: 300, Total: 1000 = 1000 ✓
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
)
|
||||
|
||||
# Create request that exactly matches remaining capacity
|
||||
data = GenerateKeyRequest(
|
||||
tpm_limit=3000,
|
||||
rpm_limit=300,
|
||||
)
|
||||
|
||||
# Should not raise any exception (exact boundary should be allowed)
|
||||
await _check_team_key_limits(
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
|
||||
def test_check_team_key_model_specific_limits_no_limits():
|
||||
"""
|
||||
Test check_team_key_model_specific_limits when no model-specific limits are set.
|
||||
Should return without raising any exceptions.
|
||||
"""
|
||||
# Create existing key with no model-specific limits
|
||||
existing_key = LiteLLM_VerificationToken(
|
||||
token="test-token-1",
|
||||
user_id="test-user",
|
||||
team_id="test-team-123",
|
||||
metadata={},
|
||||
)
|
||||
|
||||
keys = [existing_key]
|
||||
|
||||
# Create team table
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-123",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000,
|
||||
rpm_limit=1000,
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
metadata={},
|
||||
)
|
||||
|
||||
# Create request with no model-specific limits
|
||||
data = GenerateKeyRequest(
|
||||
model_rpm_limit=None,
|
||||
model_tpm_limit=None,
|
||||
)
|
||||
|
||||
# Should not raise any exception
|
||||
check_team_key_model_specific_limits(
|
||||
keys=keys,
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
)
|
||||
|
||||
|
||||
def test_check_team_key_model_specific_limits_rpm_overallocation():
|
||||
"""
|
||||
Test check_team_key_model_specific_limits when model-specific RPM would cause overallocation.
|
||||
Should raise HTTPException with appropriate error message.
|
||||
"""
|
||||
# Create existing keys with model-specific RPM limits
|
||||
existing_key1 = LiteLLM_VerificationToken(
|
||||
token="test-token-1",
|
||||
user_id="test-user-1",
|
||||
team_id="test-team-456",
|
||||
metadata={
|
||||
"model_rpm_limit": {
|
||||
"gpt-4": 500,
|
||||
"gpt-3.5-turbo": 300,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
existing_key2 = LiteLLM_VerificationToken(
|
||||
token="test-token-2",
|
||||
user_id="test-user-2",
|
||||
team_id="test-team-456",
|
||||
metadata={
|
||||
"model_rpm_limit": {
|
||||
"gpt-4": 300,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
keys = [existing_key1, existing_key2]
|
||||
|
||||
# Create team table with RPM limit
|
||||
team_table = LiteLLM_TeamTableCachedObj(
|
||||
team_id="test-team-456",
|
||||
team_alias="test-team",
|
||||
tpm_limit=10000,
|
||||
rpm_limit=1000, # Total team RPM limit
|
||||
max_budget=100.0,
|
||||
spend=0.0,
|
||||
models=[],
|
||||
blocked=False,
|
||||
members_with_roles=[],
|
||||
metadata={},
|
||||
)
|
||||
|
||||
# Create request that would exceed model-specific RPM limits
|
||||
# Existing gpt-4: 500 + 300 = 800, New: 300, Total: 1100 > 1000 (team limit)
|
||||
data = GenerateKeyRequest(
|
||||
model_rpm_limit={
|
||||
"gpt-4": 300, # This would cause overallocation
|
||||
},
|
||||
model_tpm_limit=None,
|
||||
)
|
||||
|
||||
# Should raise HTTPException for model-specific RPM overallocation
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
check_team_key_model_specific_limits(
|
||||
keys=keys,
|
||||
team_table=team_table,
|
||||
data=data,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert (
|
||||
"Allocated RPM limit=800 + Key RPM limit=300 is greater than team RPM limit=1000"
|
||||
in str(exc_info.value.detail)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ dotenv.load_dotenv()
|
|||
async def cohere_rerank(session):
|
||||
url = "http://localhost:4000/v1/rerank"
|
||||
headers = {
|
||||
"Authorization": f"bearer {os.getenv('COHERE_API_KEY')}",
|
||||
"Authorization": f"Bearer {os.getenv('COHERE_API_KEY')}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -17,7 +17,8 @@ from litellm.google_genai import (
|
|||
generate_content_stream,
|
||||
agenerate_content_stream,
|
||||
)
|
||||
from google.genai.types import ContentDict, PartDict, GenerateContentResponse
|
||||
from google.genai.types import ContentDict, PartDict
|
||||
from litellm.types.google_genai.main import GenerateContentResponse
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
|
|
@ -107,11 +108,11 @@ class BaseGoogleGenAITest:
|
|||
|
||||
def _validate_non_streaming_response(self, response: Any):
|
||||
"""Validate non-streaming response structure"""
|
||||
# Handle type checking - response should be a dict for non-streaming
|
||||
# Handle type checking - response should be a GenerateContentResponse for non-streaming
|
||||
if isinstance(response, AsyncIterator):
|
||||
pytest.fail("Expected non-streaming response but got AsyncIterator")
|
||||
|
||||
assert isinstance(response, GenerateContentResponse), f"Expected dict response, got {type(response)}"
|
||||
assert isinstance(response, GenerateContentResponse), f"Expected GenerateContentResponse, got {type(response)}"
|
||||
print(f"Response: {response.model_dump_json(indent=4)}")
|
||||
|
||||
# Basic validation - adjust based on actual Google GenAI response structure
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ import SpendLogsTable from "@/components/view_logs"
|
|||
import ModelHubTable from "@/components/model_hub_table"
|
||||
import NewUsagePage from "@/components/new_usage"
|
||||
import APIRef from "@/components/api_ref"
|
||||
import ChatUI from "@/components/chat_ui"
|
||||
import ChatUI from "@/components/chat_ui/ChatUI"
|
||||
import Sidebar from "@/components/leftnav"
|
||||
import Usage from "@/components/usage"
|
||||
import CacheDashboard from "@/components/cache_dashboard"
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue