mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'origin' into litellm_key_last_active_tracking
This commit is contained in:
commit
c911cfbabf
87 changed files with 15413 additions and 9087 deletions
150
docs/my-website/blog/gemin_3.1/index.md
Normal file
150
docs/my-website/blog/gemin_3.1/index.md
Normal file
|
|
@ -0,0 +1,150 @@
|
|||
---
|
||||
slug: gemini_3_1_pro
|
||||
title: "DAY 0 Support: Gemini 3.1 Pro on LiteLLM"
|
||||
date: 2026-02-19T10:00:00
|
||||
authors:
|
||||
- name: Sameer Kankute
|
||||
title: SWE @ LiteLLM (LLM Translation)
|
||||
url: https://www.linkedin.com/in/sameer-kankute/
|
||||
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
|
||||
- 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
|
||||
description: "Guide to using Gemini 3.1 Pro on LiteLLM Proxy and SDK with day 0 support."
|
||||
tags: [gemini, day 0 support, llms]
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Gemini 3.1 Pro Day 0 Support
|
||||
|
||||
LiteLLM now supports `gemini-3.1-pro-preview` and all the new API changes along with it.
|
||||
|
||||
## 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:main-v1.80.8-stable.1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="pip" label="Pip">
|
||||
|
||||
``` showLineNumbers title="pip install litellm"
|
||||
pip install litellm==1.80.8.post1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## What's New
|
||||
|
||||
### 1. New Thinking Levels: `thinkingLevel` with MINIMAL & MEDIUM
|
||||
|
||||
Gemini 3.1 Pro introduces support for **medium** thinking level
|
||||
|
||||
LiteLLM automatically maps the OpenAI `reasoning_effort` parameter to Gemini's `thinkingLevel`, so you can use familiar `reasoning_effort` values (`minimal`, `low`, `medium`, `high`) without changing your code!
|
||||
|
||||
---
|
||||
## Supported Endpoints
|
||||
|
||||
LiteLLM provides **full end-to-end support** for Gemini 3.1 Pro on:
|
||||
|
||||
- ✅ `/v1/chat/completions` - OpenAI-compatible chat completions endpoint
|
||||
- ✅ `/v1/responses` - OpenAI Responses API endpoint (streaming and non-streaming)
|
||||
- ✅ [`/v1/messages`](../../docs/anthropic_unified) - Anthropic-compatible messages endpoint
|
||||
- ✅ `/v1/generateContent` – [Google Gemini API](../../docs/generateContent.md) compatible endpoint
|
||||
|
||||
All endpoints support:
|
||||
- Streaming and non-streaming responses
|
||||
- Function calling with thought signatures
|
||||
- Multi-turn conversations
|
||||
- All Gemini 3-specific features
|
||||
- Conversion of provider specific thinking related param to thinkingLevel
|
||||
|
||||
## Quick Start
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
**Basic Usage with MEDIUM thinking (NEW)**
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
# No need to make any changes to your code as we map openai reasoning param to thinkingLevel
|
||||
response = completion(
|
||||
model="gemini/gemini-3.1-pro-preview",
|
||||
messages=[{"role": "user", "content": "Solve this complex math problem: 25 * 4 + 10"}],
|
||||
reasoning_effort="medium", # NEW: MEDIUM thinking level
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
**1. Setup config.yaml**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemini-3.1-pro-preview
|
||||
litellm_params:
|
||||
model: gemini/gemini-3.1-pro-preview
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
- model_name: vertex-gemini-3.1-pro-preview
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-3.1-pro-preview
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
**3. Call with MEDIUM thinking**
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer <YOUR-LITELLM-KEY>" \
|
||||
-d '{
|
||||
"model": "gemini-3.1-pro-preview",
|
||||
"messages": [{"role": "user", "content": "Complex reasoning task"}],
|
||||
"reasoning_effort": "medium"
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
---
|
||||
|
||||
## `reasoning_effort` Mapping for Gemini 3+
|
||||
|
||||
| reasoning_effort | thinking_level |
|
||||
|------------------|----------------|
|
||||
| `minimal` | `minimal` |
|
||||
| `low` | `low` |
|
||||
| `medium` | `medium` |
|
||||
| `high` | `high` |
|
||||
| `disable` | `minimal` |
|
||||
| `none` | `minimal` |
|
||||
|
||||
|
|
@ -50,3 +50,51 @@ for chunk in completion:
|
|||
print(chunk.choices[0].delta)
|
||||
|
||||
```
|
||||
|
||||
### Proxy: Always Include Streaming Usage
|
||||
|
||||
When using the LiteLLM Proxy, you can configure it to automatically include usage information in all streaming responses, even if the client doesn't send `stream_options={"include_usage": True}`.
|
||||
|
||||
#### Configuration
|
||||
|
||||
Add the following to your config.yaml:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
always_include_stream_usage: true
|
||||
```
|
||||
|
||||
Alternatively, configure it through the UI:
|
||||
|
||||
1. Navigate to the LiteLLM Proxy UI
|
||||
2. Go to `Settings` > `Router Settings` > `General`
|
||||
3. Find the `always_include_stream_usage` setting
|
||||
4. Toggle it to `true`
|
||||
5. Click `Update` to save
|
||||
|
||||
#### How it works
|
||||
|
||||
When `always_include_stream_usage` is enabled:
|
||||
- All streaming requests will automatically have `stream_options={"include_usage": True}` added
|
||||
- Clients will receive usage information in the final chunk, even if they didn't explicitly request it
|
||||
- If a client already provides `stream_options`, `include_usage: True` will be added without overwriting other options
|
||||
- Non-streaming requests are not affected
|
||||
|
||||
#### Example
|
||||
|
||||
With this setting enabled, a simple streaming request like:
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/v1/chat/completions \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "Hello!"}],
|
||||
"stream": true
|
||||
}'
|
||||
```
|
||||
|
||||
Will automatically receive usage information in the response, without needing to explicitly include `stream_options`.
|
||||
|
||||
```
|
||||
|
|
|
|||
|
|
@ -58,6 +58,17 @@ Configure the required authentication and pricing:
|
|||
- The Bria API requires an `api_token` header
|
||||
- Enter your Bria API key as the value for the `api_token` header
|
||||
|
||||
**Default Query Parameters (Optional):**
|
||||
- Add query parameters that will be automatically sent with every request
|
||||
- Perfect for API versioning, format specifications, or default configurations
|
||||
- Clients can override these parameters by providing their own values
|
||||
- Example: `version=v1`, `format=json`, `timeout=30`
|
||||
|
||||
<Image
|
||||
img={require('../../img/passthrough_query_default.png')}
|
||||
style={{width: '60%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
**Pricing Configuration:**
|
||||
- Set a cost per request (e.g., $12.00 in this example)
|
||||
- This enables cost tracking and billing for your users
|
||||
|
|
@ -112,6 +123,9 @@ general_settings:
|
|||
content-type: application/json
|
||||
accept: application/json
|
||||
forward_headers: true # Forward all incoming headers
|
||||
default_query_params: # Optional: Default query parameters
|
||||
version: "v1" # Always send version=v1
|
||||
format: "json" # Default format (can be overridden)
|
||||
```
|
||||
|
||||
### Start and Test
|
||||
|
|
@ -166,6 +180,9 @@ general_settings:
|
|||
auth: boolean # Enable LiteLLM authentication (Enterprise)
|
||||
forward_headers: boolean # Forward all incoming headers
|
||||
include_subpath: boolean # If true, forwards requests to sub-paths (default: false)
|
||||
methods: list[string] # Optional: HTTP methods (e.g., ["GET", "POST"]). If not specified, all methods are supported.
|
||||
default_query_params: # Optional: Default query parameters sent with every request
|
||||
<param-name>: string # Key-value pairs (e.g., version: "v1", format: "json")
|
||||
headers: # Custom headers to add
|
||||
Authorization: string # Auth header for target API
|
||||
content-type: string # Request content type
|
||||
|
|
@ -177,11 +194,17 @@ general_settings:
|
|||
|
||||
### Header Options
|
||||
- **Authorization**: Authentication for the target API
|
||||
- **content-type**: Request body format specification
|
||||
- **content-type**: Request body format specification
|
||||
- **accept**: Expected response format
|
||||
- **LANGFUSE_PUBLIC_KEY/SECRET_KEY**: For Langfuse integration
|
||||
- **Custom headers**: Any additional key-value pairs
|
||||
|
||||
### Default Query Parameters
|
||||
- **Parameter precedence**: Client params > URL params > default params
|
||||
- **Use cases**: API versioning, authentication tokens, format control, feature flags
|
||||
- **Override capability**: Clients can override any default parameter
|
||||
- **Examples**: `version: "v1"`, `format: "json"`, `timeout: "30"`
|
||||
|
||||
### Sub-path Routing
|
||||
|
||||
By default, pass-through endpoints only match the **exact path** specified. To forward requests to sub-paths, set `include_subpath: true`:
|
||||
|
|
@ -201,6 +224,92 @@ general_settings:
|
|||
|
||||
---
|
||||
|
||||
### Default Query Parameters
|
||||
|
||||
Pass-through endpoints support default query parameters that are automatically added to every request. This is useful for API versioning, format specifications, authentication tokens, or any default configuration.
|
||||
|
||||
#### How It Works
|
||||
|
||||
**Parameter Precedence (highest to lowest priority):**
|
||||
1. **Client-provided parameters** (in the request URL)
|
||||
2. **URL parameters** (from the target URL)
|
||||
3. **Default parameters** (from configuration)
|
||||
|
||||
#### Example Configuration
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
pass_through_endpoints:
|
||||
- path: "/api/v1"
|
||||
target: "https://external-api.com/service?timeout=60" # URL has timeout=60
|
||||
default_query_params:
|
||||
version: "v1" # Always add version=v1
|
||||
format: "json" # Default format=json (can be overridden)
|
||||
auth_level: "basic" # Always add auth_level=basic
|
||||
```
|
||||
|
||||
#### Request Examples
|
||||
|
||||
**Client Request:** `GET /api/v1/users`
|
||||
**Actual Backend Call:** `https://external-api.com/service?version=v1&format=json&auth_level=basic&timeout=60`
|
||||
|
||||
**Client Request:** `GET /api/v1/users?format=xml&custom=value`
|
||||
**Actual Backend Call:** `https://external-api.com/service?version=v1&auth_level=basic&timeout=60&format=xml&custom=value`
|
||||
- Client `format=xml` overrides default `format=json`
|
||||
- Default `version=v1` and `auth_level=basic` are preserved
|
||||
- URL `timeout=60` is preserved
|
||||
- Client `custom=value` is added
|
||||
|
||||
#### Use Cases
|
||||
|
||||
- **API Versioning**: Always send `version=v2` to maintain compatibility
|
||||
- **Authentication**: Add authentication tokens like `api_key=default_key`
|
||||
- **Format Control**: Default to `format=json` but allow client override
|
||||
- **Rate Limiting**: Set `rate_limit=standard` as default
|
||||
- **Feature Flags**: Enable `experimental=false` by default
|
||||
|
||||
---
|
||||
|
||||
You can configure different target URLs for the same path using different HTTP methods. This is useful when different backends handle different operations:
|
||||
|
||||
<Image
|
||||
img={require('../../img/passthrough_method_setup.png')}
|
||||
style={{width: '60%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
pass_through_endpoints:
|
||||
# GET requests to /azure/kb go to read API
|
||||
- path: "/azure/kb"
|
||||
target: "https://read-api.example.com/knowledge-base"
|
||||
methods: ["GET"]
|
||||
headers:
|
||||
Authorization: "bearer os.environ/READ_API_KEY"
|
||||
|
||||
# POST requests to /azure/kb go to write API
|
||||
- path: "/azure/kb"
|
||||
target: "https://write-api.example.com/knowledge-base"
|
||||
methods: ["POST"]
|
||||
headers:
|
||||
Authorization: "bearer os.environ/WRITE_API_KEY"
|
||||
|
||||
# PUT requests to /azure/kb go to update API
|
||||
- path: "/azure/kb"
|
||||
target: "https://update-api.example.com/knowledge-base"
|
||||
methods: ["PUT"]
|
||||
headers:
|
||||
Authorization: "bearer os.environ/UPDATE_API_KEY"
|
||||
```
|
||||
|
||||
**Key Points:**
|
||||
- If `methods` is not specified, the endpoint supports all HTTP methods (GET, POST, PUT, DELETE, PATCH)
|
||||
- Multiple endpoints can share the same path as long as they have different methods
|
||||
- You can specify multiple methods for a single endpoint: `methods: ["GET", "POST"]`
|
||||
- This allows you to route to different backends based on the operation type
|
||||
|
||||
---
|
||||
|
||||
## Advanced: Custom Adapters
|
||||
|
||||
For complex integrations (like Anthropic/Bedrock clients), you can create custom adapters that translate between different API schemas.
|
||||
|
|
|
|||
318
docs/my-website/docs/proxy/project_management.md
Normal file
318
docs/my-website/docs/proxy/project_management.md
Normal file
|
|
@ -0,0 +1,318 @@
|
|||
# [Beta] Project Management
|
||||
|
||||
Projects in LiteLLM sit between teams and keys in the organizational hierarchy, enabling fine-grained access control and budget management for specific use cases or applications.
|
||||
|
||||
```mermaid
|
||||
graph TD
|
||||
A[Organization] --> B[Team 1]
|
||||
A --> C[Team 2]
|
||||
B --> D[Project A]
|
||||
B --> E[Project B]
|
||||
C --> F[Project C]
|
||||
D --> G[API Key 1]
|
||||
D --> H[API Key 2]
|
||||
E --> I[API Key 3]
|
||||
F --> J[API Key 4]
|
||||
|
||||
style A fill:#e1f5ff
|
||||
style B fill:#fff4e6
|
||||
style C fill:#fff4e6
|
||||
style D fill:#f3e5f5
|
||||
style E fill:#f3e5f5
|
||||
style F fill:#f3e5f5
|
||||
style G fill:#e8f5e9
|
||||
style H fill:#e8f5e9
|
||||
style I fill:#e8f5e9
|
||||
style J fill:#e8f5e9
|
||||
```
|
||||
|
||||
**Hierarchy**: `Organizations > Teams > Projects > Keys`
|
||||
|
||||
## Quick Start
|
||||
|
||||
This walkthrough shows how to create a project, generate an API key, make requests, and view project-level spend tracking in the UI.
|
||||
|
||||
### Step 1: Create a Project
|
||||
|
||||
```bash showLineNumbers
|
||||
curl --location 'http://0.0.0.0:4000/project/new' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"project_alias": "flight-search-assistant",
|
||||
"team_id": "ad898803-c8a3-4f4a-976a-a3c372cffa45",
|
||||
"models": ["gpt-4", "gpt-3.5-turbo"],
|
||||
"max_budget": 100,
|
||||
"metadata": {
|
||||
"use_case_id": "SNOW-12345",
|
||||
"responsible_ai_id": "RAI-67890"
|
||||
}
|
||||
}' | jq
|
||||
```
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"project_id": "e402a141-725a-4437-bff5-d47459189716",
|
||||
"project_alias": "flight-search-assistant",
|
||||
"team_id": "ad898803-c8a3-4f4a-976a-a3c372cffa45",
|
||||
"models": ["gpt-4", "gpt-3.5-turbo"],
|
||||
"max_budget": 100,
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
### Step 2: Generate API Key for Project
|
||||
|
||||
```bash showLineNumbers
|
||||
curl 'http://0.0.0.0:4000/key/generate' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data-raw '{
|
||||
"models": ["gpt-3.5-turbo", "gpt-4"],
|
||||
"metadata": {"user": "ishaan@berri.ai"},
|
||||
"project_id": "e402a141-725a-4437-bff5-d47459189716"
|
||||
}' | jq
|
||||
```
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"key": "sk-W8VbscpfuyvHm5TkxRYiXA",
|
||||
"key_name": "sk-...YiXA",
|
||||
"project_id": "e402a141-725a-4437-bff5-d47459189716",
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
### Step 3: Use API Key in Chat Completions
|
||||
|
||||
```bash showLineNumbers
|
||||
curl http://localhost:4000/v1/chat/completions \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-W8VbscpfuyvHm5TkxRYiXA' \
|
||||
--data '{
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "What is litellm?"}]
|
||||
}' | jq
|
||||
```
|
||||
|
||||
### Step 4: View Project Spend in UI
|
||||
|
||||
Navigate to the **Logs** page in the LiteLLM Admin UI. You'll see the `user_api_key_project_id` tracked in the request metadata:
|
||||
|
||||

|
||||
|
||||
As shown above, the spend logs metadata includes:
|
||||
- `"user_api_key_project_id": "e402a141-725a-4437-bff5-d47459189716"` - Links the request to your project
|
||||
- All costs and token usage are automatically attributed to the project
|
||||
- You can query and filter logs by project ID for detailed reporting
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### POST /project/new
|
||||
|
||||
Create a new project.
|
||||
|
||||
**Who can call**: Admins or Team Admins
|
||||
|
||||
**Parameters**:
|
||||
- `project_alias` (string, optional): Human-readable name for the project
|
||||
- `team_id` (string, required): The team this project belongs to
|
||||
- `models` (array, optional): List of models the project can access
|
||||
- `max_budget` (float, optional): Maximum spend budget for the project
|
||||
- `tpm_limit` (int, optional): Tokens per minute limit
|
||||
- `rpm_limit` (int, optional): Requests per minute limit
|
||||
- `budget_duration` (string, optional): Budget reset period (e.g., "30d", "1mo")
|
||||
- `metadata` (object, optional): Custom metadata for the project
|
||||
- `blocked` (boolean, optional): Block all API calls for this project
|
||||
|
||||
**Example**:
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/new' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"project_alias": "hotel-recommendations",
|
||||
"team_id": "team-123",
|
||||
"models": ["claude-3-sonnet"],
|
||||
"max_budget": 200,
|
||||
"tpm_limit": 100000,
|
||||
"metadata": {
|
||||
"use_case_id": "SNOW-12346",
|
||||
"cost_center": "travel-products"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**Response**:
|
||||
|
||||
```json
|
||||
{
|
||||
"project_id": "project-def",
|
||||
"project_alias": "hotel-recommendations",
|
||||
"team_id": "team-123",
|
||||
"models": ["claude-3-sonnet"],
|
||||
"spend": 0.0,
|
||||
"budget_id": "budget-xyz",
|
||||
"metadata": {
|
||||
"use_case_id": "SNOW-12346",
|
||||
"cost_center": "travel-products"
|
||||
},
|
||||
"created_at": "2025-01-15T10:00:00Z",
|
||||
"updated_at": "2025-01-15T10:00:00Z"
|
||||
}
|
||||
```
|
||||
|
||||
### POST /project/update
|
||||
|
||||
Update an existing project.
|
||||
|
||||
**Who can call**: Admins or Team Admins
|
||||
|
||||
**Parameters**:
|
||||
- `project_id` (string, required): The project to update
|
||||
- `project_alias` (string, optional): Updated project name
|
||||
- `team_id` (string, optional): Move project to different team
|
||||
- `models` (array, optional): Updated list of allowed models
|
||||
- `max_budget` (float, optional): Updated budget
|
||||
- `tpm_limit` (int, optional): Updated TPM limit
|
||||
- `rpm_limit` (int, optional): Updated RPM limit
|
||||
- `metadata` (object, optional): Updated metadata
|
||||
- `blocked` (boolean, optional): Updated blocked status
|
||||
|
||||
**Example**:
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/update' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"project_id": "project-abc",
|
||||
"max_budget": 200,
|
||||
"tpm_limit": 200000,
|
||||
"metadata": {
|
||||
"status": "production"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### GET /project/info
|
||||
|
||||
Get information about a specific project.
|
||||
|
||||
**Parameters**:
|
||||
- `project_id` (string, required): Query parameter
|
||||
|
||||
**Example**:
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/info?project_id=project-abc' \
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
|
||||
**Response**:
|
||||
|
||||
```json
|
||||
{
|
||||
"project_id": "project-abc",
|
||||
"project_alias": "flight-search-assistant",
|
||||
"team_id": "team-123",
|
||||
"models": ["gpt-4", "gpt-3.5-turbo"],
|
||||
"spend": 45.67,
|
||||
"model_spend": {
|
||||
"gpt-4": 42.30,
|
||||
"gpt-3.5-turbo": 3.37
|
||||
},
|
||||
"litellm_budget_table": {
|
||||
"budget_id": "budget-xyz",
|
||||
"max_budget": 100.0,
|
||||
"tpm_limit": 100000,
|
||||
"rpm_limit": 100
|
||||
},
|
||||
"metadata": {
|
||||
"use_case_id": "SNOW-12345"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### GET /project/list
|
||||
|
||||
List all projects the user has access to.
|
||||
|
||||
**Example**:
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/list' \
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
|
||||
**Response**:
|
||||
|
||||
```json
|
||||
[
|
||||
{
|
||||
"project_id": "project-abc",
|
||||
"project_alias": "flight-search-assistant",
|
||||
"team_id": "team-123",
|
||||
"spend": 45.67
|
||||
},
|
||||
{
|
||||
"project_id": "project-def",
|
||||
"project_alias": "hotel-recommendations",
|
||||
"team_id": "team-123",
|
||||
"spend": 23.45
|
||||
}
|
||||
]
|
||||
```
|
||||
|
||||
### DELETE /project/delete
|
||||
|
||||
Delete one or more projects.
|
||||
|
||||
**Who can call**: Admins only
|
||||
|
||||
**Parameters**:
|
||||
- `project_ids` (array, required): List of project IDs to delete
|
||||
|
||||
**Example**:
|
||||
|
||||
```bash
|
||||
curl --location --request DELETE 'http://0.0.0.0:4000/project/delete' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"project_ids": ["project-abc", "project-def"]
|
||||
}'
|
||||
```
|
||||
|
||||
**Note**: Projects with associated API keys cannot be deleted. Delete or reassign the keys first.
|
||||
|
||||
## Model-Specific Quotas
|
||||
|
||||
You can set different quotas for different models within a project:
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/new' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"project_alias": "multi-model-project",
|
||||
"team_id": "team-123",
|
||||
"models": ["gpt-4", "gpt-3.5-turbo", "claude-3-sonnet"],
|
||||
"max_budget": 500,
|
||||
"metadata": {
|
||||
"model_tpm_limit": {
|
||||
"gpt-4": 50000,
|
||||
"gpt-3.5-turbo": 200000,
|
||||
"claude-3-sonnet": 100000
|
||||
},
|
||||
"model_rpm_limit": {
|
||||
"gpt-4": 50,
|
||||
"gpt-3.5-turbo": 500,
|
||||
"claude-3-sonnet": 100
|
||||
}
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
|
@ -1047,6 +1047,8 @@ For long-running conversations, you can enable **server-side compaction** so tha
|
|||
|
||||
Supported on the OpenAI Responses API when using the `openai` or `azure` provider. Pass `context_management` with a compaction entry and `compact_threshold` (token count; minimum 1000). When the context crosses the threshold, the server compacts in-stream and continues. Chain turns with `previous_response_id` or by appending output items to your next input array. See [OpenAI Compaction guide](https://developers.openai.com/api/docs/guides/compaction) for details.
|
||||
|
||||
> **Note:** You can use openai `context_management` format with Anthropic models via LiteLLM via responses API. LiteLLM will automatically translate this format for Anthropic and handle context management for you.
|
||||
|
||||
For explicit control over when compaction runs, use the standalone compact endpoint (`POST /v1/responses/compact`) instead.
|
||||
|
||||
### Python SDK
|
||||
|
|
|
|||
BIN
docs/my-website/img/passthrough_method_setup.png
Normal file
BIN
docs/my-website/img/passthrough_method_setup.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 77 KiB |
BIN
docs/my-website/img/passthrough_query_default.png
Normal file
BIN
docs/my-website/img/passthrough_query_default.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 84 KiB |
BIN
docs/my-website/img/project_spend.png
Normal file
BIN
docs/my-website/img/project_spend.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 850 KiB |
|
|
@ -410,6 +410,7 @@ const sidebars = {
|
|||
items: [
|
||||
"proxy/users",
|
||||
"proxy/team_budgets",
|
||||
"proxy/project_management",
|
||||
"proxy/ui_team_soft_budget_alerts",
|
||||
"proxy/tag_budgets",
|
||||
"proxy/customers",
|
||||
|
|
@ -781,13 +782,13 @@ const sidebars = {
|
|||
"providers/bedrock_batches",
|
||||
"providers/bedrock_realtime_with_audio",
|
||||
"providers/aws_polly",
|
||||
"providers/bedrock_vector_store",
|
||||
]
|
||||
},
|
||||
"providers/litellm_proxy",
|
||||
"providers/abliteration",
|
||||
"providers/ai21",
|
||||
"providers/aiml",
|
||||
"providers/bedrock_vector_store",
|
||||
]
|
||||
},
|
||||
"providers/litellm_proxy",
|
||||
"providers/abliteration",
|
||||
"providers/ai21",
|
||||
"providers/aiml",
|
||||
"providers/aleph_alpha",
|
||||
"providers/amazon_nova",
|
||||
"providers/anyscale",
|
||||
|
|
|
|||
BIN
docs/my-website/static/img/project_spend.png
Normal file
BIN
docs/my-website/static/img/project_spend.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 850 KiB |
|
|
@ -1,309 +1,311 @@
|
|||
"""
|
||||
PagerDuty Alerting Integration
|
||||
|
||||
Handles two types of alerts:
|
||||
- High LLM API Failure Rate. Configure X fails in Y seconds to trigger an alert.
|
||||
- High Number of Hanging LLM Requests. Configure X hangs in Y seconds to trigger an alert.
|
||||
|
||||
Note: This is a Free feature on the regular litellm docker image.
|
||||
|
||||
However, this is under the enterprise license
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Literal, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching import DualCache
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.integrations.pagerduty import (
|
||||
AlertingConfig,
|
||||
PagerDutyInternalEvent,
|
||||
PagerDutyPayload,
|
||||
PagerDutyRequestBody,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
CallTypesLiteral,
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingPayloadErrorInformation,
|
||||
)
|
||||
|
||||
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD = 60
|
||||
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS = 60
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS = 60
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS = 600
|
||||
|
||||
|
||||
class PagerDutyAlerting(SlackAlerting):
|
||||
"""
|
||||
Tracks failed requests and hanging requests separately.
|
||||
If threshold is crossed for either type, triggers a PagerDuty alert.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, alerting_args: Optional[Union[AlertingConfig, dict]] = None, **kwargs
|
||||
):
|
||||
super().__init__()
|
||||
_api_key = os.getenv("PAGERDUTY_API_KEY")
|
||||
if not _api_key:
|
||||
raise ValueError("PAGERDUTY_API_KEY is not set")
|
||||
|
||||
self.api_key: str = _api_key
|
||||
alerting_args = alerting_args or {}
|
||||
self.pagerduty_alerting_args: AlertingConfig = AlertingConfig(
|
||||
failure_threshold=alerting_args.get(
|
||||
"failure_threshold", PAGERDUTY_DEFAULT_FAILURE_THRESHOLD
|
||||
),
|
||||
failure_threshold_window_seconds=alerting_args.get(
|
||||
"failure_threshold_window_seconds",
|
||||
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS,
|
||||
),
|
||||
hanging_threshold_seconds=alerting_args.get(
|
||||
"hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
|
||||
),
|
||||
hanging_threshold_window_seconds=alerting_args.get(
|
||||
"hanging_threshold_window_seconds",
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS,
|
||||
),
|
||||
)
|
||||
|
||||
# Separate storage for failures vs. hangs
|
||||
self._failure_events: List[PagerDutyInternalEvent] = []
|
||||
self._hanging_events: List[PagerDutyInternalEvent] = []
|
||||
|
||||
# ------------------ MAIN LOGIC ------------------ #
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Record a failure event. Only send an alert to PagerDuty if the
|
||||
configured *failure* threshold is exceeded in the specified window.
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
if not standard_logging_payload:
|
||||
raise ValueError(
|
||||
"standard_logging_object is required for PagerDutyAlerting"
|
||||
)
|
||||
|
||||
# Extract error details
|
||||
error_info: Optional[StandardLoggingPayloadErrorInformation] = (
|
||||
standard_logging_payload.get("error_information") or {}
|
||||
)
|
||||
_meta = standard_logging_payload.get("metadata") or {}
|
||||
|
||||
self._failure_events.append(
|
||||
PagerDutyInternalEvent(
|
||||
failure_event_type="failed_response",
|
||||
timestamp=now,
|
||||
error_class=error_info.get("error_class"),
|
||||
error_code=error_info.get("error_code"),
|
||||
error_llm_provider=error_info.get("llm_provider"),
|
||||
user_api_key_hash=_meta.get("user_api_key_hash"),
|
||||
user_api_key_alias=_meta.get("user_api_key_alias"),
|
||||
user_api_key_spend=_meta.get("user_api_key_spend"),
|
||||
user_api_key_max_budget=_meta.get("user_api_key_max_budget"),
|
||||
user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"),
|
||||
user_api_key_org_id=_meta.get("user_api_key_org_id"),
|
||||
user_api_key_team_id=_meta.get("user_api_key_team_id"),
|
||||
user_api_key_user_id=_meta.get("user_api_key_user_id"),
|
||||
user_api_key_team_alias=_meta.get("user_api_key_team_alias"),
|
||||
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"),
|
||||
)
|
||||
)
|
||||
|
||||
# Prune + Possibly alert
|
||||
window_seconds = self.pagerduty_alerting_args.get(
|
||||
"failure_threshold_window_seconds", 60
|
||||
)
|
||||
threshold = self.pagerduty_alerting_args.get("failure_threshold", 1)
|
||||
|
||||
# If threshold is crossed, send PD alert for failures
|
||||
await self._send_alert_if_thresholds_crossed(
|
||||
events=self._failure_events,
|
||||
window_seconds=window_seconds,
|
||||
threshold=threshold,
|
||||
alert_prefix="High LLM API Failure Rate",
|
||||
)
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict,
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
"""
|
||||
Example of detecting hanging requests by waiting a given threshold.
|
||||
If the request didn't finish by then, we treat it as 'hanging'.
|
||||
"""
|
||||
verbose_logger.info("Inside Proxy Logging Pre-call hook!")
|
||||
asyncio.create_task(
|
||||
self.hanging_response_handler(
|
||||
request_data=data, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
)
|
||||
return None
|
||||
|
||||
async def hanging_response_handler(
|
||||
self, request_data: Optional[dict], user_api_key_dict: UserAPIKeyAuth
|
||||
):
|
||||
"""
|
||||
Checks if request completed by the time 'hanging_threshold_seconds' elapses.
|
||||
If not, we classify it as a hanging request.
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"Inside Hanging Response Handler!..sleeping for {self.pagerduty_alerting_args.get('hanging_threshold_seconds', PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS)} seconds"
|
||||
)
|
||||
await asyncio.sleep(
|
||||
self.pagerduty_alerting_args.get(
|
||||
"hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
|
||||
)
|
||||
)
|
||||
|
||||
if await self._request_is_completed(request_data=request_data):
|
||||
return # It's not hanging if completed
|
||||
|
||||
# Otherwise, record it as hanging
|
||||
self._hanging_events.append(
|
||||
PagerDutyInternalEvent(
|
||||
failure_event_type="hanging_response",
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
error_class="HangingRequest",
|
||||
error_code="HangingRequest",
|
||||
error_llm_provider="HangingRequest",
|
||||
user_api_key_hash=user_api_key_dict.api_key,
|
||||
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_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,
|
||||
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_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,
|
||||
)
|
||||
)
|
||||
|
||||
# Prune + Possibly alert
|
||||
window_seconds = self.pagerduty_alerting_args.get(
|
||||
"hanging_threshold_window_seconds",
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS,
|
||||
)
|
||||
threshold: int = self.pagerduty_alerting_args.get(
|
||||
"hanging_threshold_fails", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
|
||||
)
|
||||
|
||||
# If threshold is crossed, send PD alert for hangs
|
||||
await self._send_alert_if_thresholds_crossed(
|
||||
events=self._hanging_events,
|
||||
window_seconds=window_seconds,
|
||||
threshold=threshold,
|
||||
alert_prefix="High Number of Hanging LLM Requests",
|
||||
)
|
||||
|
||||
# ------------------ HELPERS ------------------ #
|
||||
|
||||
async def _send_alert_if_thresholds_crossed(
|
||||
self,
|
||||
events: List[PagerDutyInternalEvent],
|
||||
window_seconds: int,
|
||||
threshold: int,
|
||||
alert_prefix: str,
|
||||
):
|
||||
"""
|
||||
1. Prune old events
|
||||
2. If threshold is reached, build alert, send to PagerDuty
|
||||
3. Clear those events
|
||||
"""
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(seconds=window_seconds)
|
||||
pruned = [e for e in events if e.get("timestamp", datetime.min) > cutoff]
|
||||
|
||||
# Update the reference list
|
||||
events.clear()
|
||||
events.extend(pruned)
|
||||
|
||||
# Check threshold
|
||||
verbose_logger.debug(
|
||||
f"Have {len(events)} events in the last {window_seconds} seconds. Threshold is {threshold}"
|
||||
)
|
||||
if len(events) >= threshold:
|
||||
# Build short summary of last N events
|
||||
error_summaries = self._build_error_summaries(events, max_errors=5)
|
||||
alert_message = (
|
||||
f"{alert_prefix}: {len(events)} in the last {window_seconds} seconds."
|
||||
)
|
||||
custom_details = {"recent_errors": error_summaries}
|
||||
|
||||
await self.send_alert_to_pagerduty(
|
||||
alert_message=alert_message,
|
||||
custom_details=custom_details,
|
||||
)
|
||||
|
||||
# Clear them after sending an alert, so we don't spam
|
||||
events.clear()
|
||||
|
||||
def _build_error_summaries(
|
||||
self, events: List[PagerDutyInternalEvent], max_errors: int = 5
|
||||
) -> List[PagerDutyInternalEvent]:
|
||||
"""
|
||||
Build short text summaries for the last `max_errors`.
|
||||
Example: "ValueError (code: 500, provider: openai)"
|
||||
"""
|
||||
recent = events[-max_errors:]
|
||||
summaries = []
|
||||
for fe in recent:
|
||||
# If any of these is None, show "N/A" to avoid messing up the summary string
|
||||
fe.pop("timestamp")
|
||||
summaries.append(fe)
|
||||
return summaries
|
||||
|
||||
async def send_alert_to_pagerduty(self, alert_message: str, custom_details: dict):
|
||||
"""
|
||||
Send [critical] Alert to PagerDuty
|
||||
|
||||
https://developer.pagerduty.com/api-reference/YXBpOjI3NDgyNjU-pager-duty-v2-events-api
|
||||
"""
|
||||
try:
|
||||
verbose_logger.debug(f"Sending alert to PagerDuty: {alert_message}")
|
||||
async_client: AsyncHTTPHandler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
payload: PagerDutyRequestBody = PagerDutyRequestBody(
|
||||
payload=PagerDutyPayload(
|
||||
summary=alert_message,
|
||||
severity="critical",
|
||||
source="LiteLLM Alert",
|
||||
component="LiteLLM",
|
||||
custom_details=custom_details,
|
||||
),
|
||||
routing_key=self.api_key,
|
||||
event_action="trigger",
|
||||
)
|
||||
|
||||
return await async_client.post(
|
||||
url="https://events.pagerduty.com/v2/enqueue",
|
||||
json=dict(payload),
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error sending alert to PagerDuty: {e}")
|
||||
"""
|
||||
PagerDuty Alerting Integration
|
||||
|
||||
Handles two types of alerts:
|
||||
- High LLM API Failure Rate. Configure X fails in Y seconds to trigger an alert.
|
||||
- High Number of Hanging LLM Requests. Configure X hangs in Y seconds to trigger an alert.
|
||||
|
||||
Note: This is a Free feature on the regular litellm docker image.
|
||||
|
||||
However, this is under the enterprise license
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching import DualCache
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.integrations.pagerduty import (
|
||||
AlertingConfig,
|
||||
PagerDutyInternalEvent,
|
||||
PagerDutyPayload,
|
||||
PagerDutyRequestBody,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
CallTypesLiteral,
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingPayloadErrorInformation,
|
||||
)
|
||||
|
||||
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD = 60
|
||||
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS = 60
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS = 60
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS = 600
|
||||
|
||||
|
||||
class PagerDutyAlerting(SlackAlerting):
|
||||
"""
|
||||
Tracks failed requests and hanging requests separately.
|
||||
If threshold is crossed for either type, triggers a PagerDuty alert.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, alerting_args: Optional[Union[AlertingConfig, dict]] = None, **kwargs
|
||||
):
|
||||
super().__init__()
|
||||
_api_key = os.getenv("PAGERDUTY_API_KEY")
|
||||
if not _api_key:
|
||||
raise ValueError("PAGERDUTY_API_KEY is not set")
|
||||
|
||||
self.api_key: str = _api_key
|
||||
alerting_args = alerting_args or {}
|
||||
self.pagerduty_alerting_args: AlertingConfig = AlertingConfig(
|
||||
failure_threshold=alerting_args.get(
|
||||
"failure_threshold", PAGERDUTY_DEFAULT_FAILURE_THRESHOLD
|
||||
),
|
||||
failure_threshold_window_seconds=alerting_args.get(
|
||||
"failure_threshold_window_seconds",
|
||||
PAGERDUTY_DEFAULT_FAILURE_THRESHOLD_WINDOW_SECONDS,
|
||||
),
|
||||
hanging_threshold_seconds=alerting_args.get(
|
||||
"hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
|
||||
),
|
||||
hanging_threshold_window_seconds=alerting_args.get(
|
||||
"hanging_threshold_window_seconds",
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS,
|
||||
),
|
||||
)
|
||||
|
||||
# Separate storage for failures vs. hangs
|
||||
self._failure_events: List[PagerDutyInternalEvent] = []
|
||||
self._hanging_events: List[PagerDutyInternalEvent] = []
|
||||
|
||||
# ------------------ MAIN LOGIC ------------------ #
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Record a failure event. Only send an alert to PagerDuty if the
|
||||
configured *failure* threshold is exceeded in the specified window.
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
if not standard_logging_payload:
|
||||
raise ValueError(
|
||||
"standard_logging_object is required for PagerDutyAlerting"
|
||||
)
|
||||
|
||||
# Extract error details
|
||||
error_info: Optional[StandardLoggingPayloadErrorInformation] = (
|
||||
standard_logging_payload.get("error_information") or {}
|
||||
)
|
||||
_meta = standard_logging_payload.get("metadata") or {}
|
||||
|
||||
self._failure_events.append(
|
||||
PagerDutyInternalEvent(
|
||||
failure_event_type="failed_response",
|
||||
timestamp=now,
|
||||
error_class=error_info.get("error_class"),
|
||||
error_code=error_info.get("error_code"),
|
||||
error_llm_provider=error_info.get("llm_provider"),
|
||||
user_api_key_hash=_meta.get("user_api_key_hash"),
|
||||
user_api_key_alias=_meta.get("user_api_key_alias"),
|
||||
user_api_key_spend=_meta.get("user_api_key_spend"),
|
||||
user_api_key_max_budget=_meta.get("user_api_key_max_budget"),
|
||||
user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"),
|
||||
user_api_key_org_id=_meta.get("user_api_key_org_id"),
|
||||
user_api_key_team_id=_meta.get("user_api_key_team_id"),
|
||||
user_api_key_project_id=_meta.get("user_api_key_project_id"),
|
||||
user_api_key_user_id=_meta.get("user_api_key_user_id"),
|
||||
user_api_key_team_alias=_meta.get("user_api_key_team_alias"),
|
||||
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"),
|
||||
)
|
||||
)
|
||||
|
||||
# Prune + Possibly alert
|
||||
window_seconds = self.pagerduty_alerting_args.get(
|
||||
"failure_threshold_window_seconds", 60
|
||||
)
|
||||
threshold = self.pagerduty_alerting_args.get("failure_threshold", 1)
|
||||
|
||||
# If threshold is crossed, send PD alert for failures
|
||||
await self._send_alert_if_thresholds_crossed(
|
||||
events=self._failure_events,
|
||||
window_seconds=window_seconds,
|
||||
threshold=threshold,
|
||||
alert_prefix="High LLM API Failure Rate",
|
||||
)
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict,
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
"""
|
||||
Example of detecting hanging requests by waiting a given threshold.
|
||||
If the request didn't finish by then, we treat it as 'hanging'.
|
||||
"""
|
||||
verbose_logger.info("Inside Proxy Logging Pre-call hook!")
|
||||
asyncio.create_task(
|
||||
self.hanging_response_handler(
|
||||
request_data=data, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
)
|
||||
return None
|
||||
|
||||
async def hanging_response_handler(
|
||||
self, request_data: Optional[dict], user_api_key_dict: UserAPIKeyAuth
|
||||
):
|
||||
"""
|
||||
Checks if request completed by the time 'hanging_threshold_seconds' elapses.
|
||||
If not, we classify it as a hanging request.
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"Inside Hanging Response Handler!..sleeping for {self.pagerduty_alerting_args.get('hanging_threshold_seconds', PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS)} seconds"
|
||||
)
|
||||
await asyncio.sleep(
|
||||
self.pagerduty_alerting_args.get(
|
||||
"hanging_threshold_seconds", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
|
||||
)
|
||||
)
|
||||
|
||||
if await self._request_is_completed(request_data=request_data):
|
||||
return # It's not hanging if completed
|
||||
|
||||
# Otherwise, record it as hanging
|
||||
self._hanging_events.append(
|
||||
PagerDutyInternalEvent(
|
||||
failure_event_type="hanging_response",
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
error_class="HangingRequest",
|
||||
error_code="HangingRequest",
|
||||
error_llm_provider="HangingRequest",
|
||||
user_api_key_hash=user_api_key_dict.api_key,
|
||||
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_org_id=user_api_key_dict.org_id,
|
||||
user_api_key_team_id=user_api_key_dict.team_id,
|
||||
user_api_key_project_id=user_api_key_dict.project_id,
|
||||
user_api_key_user_id=user_api_key_dict.user_id,
|
||||
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_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,
|
||||
)
|
||||
)
|
||||
|
||||
# Prune + Possibly alert
|
||||
window_seconds = self.pagerduty_alerting_args.get(
|
||||
"hanging_threshold_window_seconds",
|
||||
PAGERDUTY_DEFAULT_HANGING_THRESHOLD_WINDOW_SECONDS,
|
||||
)
|
||||
threshold: int = self.pagerduty_alerting_args.get(
|
||||
"hanging_threshold_fails", PAGERDUTY_DEFAULT_HANGING_THRESHOLD_SECONDS
|
||||
)
|
||||
|
||||
# If threshold is crossed, send PD alert for hangs
|
||||
await self._send_alert_if_thresholds_crossed(
|
||||
events=self._hanging_events,
|
||||
window_seconds=window_seconds,
|
||||
threshold=threshold,
|
||||
alert_prefix="High Number of Hanging LLM Requests",
|
||||
)
|
||||
|
||||
# ------------------ HELPERS ------------------ #
|
||||
|
||||
async def _send_alert_if_thresholds_crossed(
|
||||
self,
|
||||
events: List[PagerDutyInternalEvent],
|
||||
window_seconds: int,
|
||||
threshold: int,
|
||||
alert_prefix: str,
|
||||
):
|
||||
"""
|
||||
1. Prune old events
|
||||
2. If threshold is reached, build alert, send to PagerDuty
|
||||
3. Clear those events
|
||||
"""
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(seconds=window_seconds)
|
||||
pruned = [e for e in events if e.get("timestamp", datetime.min) > cutoff]
|
||||
|
||||
# Update the reference list
|
||||
events.clear()
|
||||
events.extend(pruned)
|
||||
|
||||
# Check threshold
|
||||
verbose_logger.debug(
|
||||
f"Have {len(events)} events in the last {window_seconds} seconds. Threshold is {threshold}"
|
||||
)
|
||||
if len(events) >= threshold:
|
||||
# Build short summary of last N events
|
||||
error_summaries = self._build_error_summaries(events, max_errors=5)
|
||||
alert_message = (
|
||||
f"{alert_prefix}: {len(events)} in the last {window_seconds} seconds."
|
||||
)
|
||||
custom_details = {"recent_errors": error_summaries}
|
||||
|
||||
await self.send_alert_to_pagerduty(
|
||||
alert_message=alert_message,
|
||||
custom_details=custom_details,
|
||||
)
|
||||
|
||||
# Clear them after sending an alert, so we don't spam
|
||||
events.clear()
|
||||
|
||||
def _build_error_summaries(
|
||||
self, events: List[PagerDutyInternalEvent], max_errors: int = 5
|
||||
) -> List[PagerDutyInternalEvent]:
|
||||
"""
|
||||
Build short text summaries for the last `max_errors`.
|
||||
Example: "ValueError (code: 500, provider: openai)"
|
||||
"""
|
||||
recent = events[-max_errors:]
|
||||
summaries = []
|
||||
for fe in recent:
|
||||
# If any of these is None, show "N/A" to avoid messing up the summary string
|
||||
fe.pop("timestamp")
|
||||
summaries.append(fe)
|
||||
return summaries
|
||||
|
||||
async def send_alert_to_pagerduty(self, alert_message: str, custom_details: dict):
|
||||
"""
|
||||
Send [critical] Alert to PagerDuty
|
||||
|
||||
https://developer.pagerduty.com/api-reference/YXBpOjI3NDgyNjU-pager-duty-v2-events-api
|
||||
"""
|
||||
try:
|
||||
verbose_logger.debug(f"Sending alert to PagerDuty: {alert_message}")
|
||||
async_client: AsyncHTTPHandler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
payload: PagerDutyRequestBody = PagerDutyRequestBody(
|
||||
payload=PagerDutyPayload(
|
||||
summary=alert_message,
|
||||
severity="critical",
|
||||
source="LiteLLM Alert",
|
||||
component="LiteLLM",
|
||||
custom_details=custom_details,
|
||||
),
|
||||
routing_key=self.api_key,
|
||||
event_action="trigger",
|
||||
)
|
||||
|
||||
return await async_client.post(
|
||||
url="https://events.pagerduty.com/v2/enqueue",
|
||||
json=dict(payload),
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error sending alert to PagerDuty: {e}")
|
||||
|
|
|
|||
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.42-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.42-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.42.tar.gz
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.42.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -0,0 +1,35 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_ProjectTable" (
|
||||
"project_id" TEXT NOT NULL,
|
||||
"project_alias" TEXT,
|
||||
"team_id" TEXT,
|
||||
"budget_id" TEXT,
|
||||
"metadata" JSONB NOT NULL DEFAULT '{}',
|
||||
"models" TEXT[],
|
||||
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
|
||||
"model_spend" JSONB NOT NULL DEFAULT '{}',
|
||||
"blocked" BOOLEAN NOT NULL DEFAULT false,
|
||||
"object_permission_id" TEXT,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"created_by" TEXT NOT NULL,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_by" TEXT NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_ProjectTable_pkey" PRIMARY KEY ("project_id")
|
||||
);
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_ProjectTable" ADD CONSTRAINT "LiteLLM_ProjectTable_team_id_fkey" FOREIGN KEY ("team_id") REFERENCES "LiteLLM_TeamTable"("team_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_ProjectTable" ADD CONSTRAINT "LiteLLM_ProjectTable_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_ProjectTable" ADD CONSTRAINT "LiteLLM_ProjectTable_object_permission_id_fkey" FOREIGN KEY ("object_permission_id") REFERENCES "LiteLLM_ObjectPermissionTable"("object_permission_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
|
||||
-- AlterTable: Add project_id to LiteLLM_VerificationToken
|
||||
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "project_id" TEXT;
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_VerificationToken" ADD CONSTRAINT "LiteLLM_VerificationToken_project_id_fkey" FOREIGN KEY ("project_id") REFERENCES "LiteLLM_ProjectTable"("project_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
-- AlterTable: Add new fields to LiteLLM_ProjectTable
|
||||
ALTER TABLE "LiteLLM_ProjectTable" ADD COLUMN "description" TEXT;
|
||||
ALTER TABLE "LiteLLM_ProjectTable" ADD COLUMN "model_rpm_limit" JSONB NOT NULL DEFAULT '{}';
|
||||
ALTER TABLE "LiteLLM_ProjectTable" ADD COLUMN "model_tpm_limit" JSONB NOT NULL DEFAULT '{}';
|
||||
|
||||
|
|
@ -24,6 +24,7 @@ model LiteLLM_BudgetTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
updated_by String
|
||||
organization LiteLLM_OrganizationTable[] // multiple orgs can have the same budget
|
||||
projects LiteLLM_ProjectTable[] // multiple projects can have the same budget
|
||||
keys LiteLLM_VerificationToken[] // multiple keys can have the same budget
|
||||
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
|
||||
tags LiteLLM_TagTable[] // multiple tags can have the same budget
|
||||
|
|
@ -135,6 +136,81 @@ model LiteLLM_TeamTable {
|
|||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
projects LiteLLM_ProjectTable[]
|
||||
}
|
||||
|
||||
// Projects sit between teams and keys for use-case management
|
||||
model LiteLLM_ProjectTable {
|
||||
project_id String @id @default(uuid())
|
||||
project_alias String?
|
||||
description String?
|
||||
team_id String?
|
||||
budget_id String?
|
||||
metadata Json @default("{}")
|
||||
models String[]
|
||||
spend Float @default(0.0)
|
||||
model_spend Json @default("{}")
|
||||
model_rpm_limit Json @default("{}")
|
||||
model_tpm_limit Json @default("{}")
|
||||
blocked Boolean @default(false)
|
||||
object_permission_id String?
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
created_by String
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
updated_by String
|
||||
|
||||
// Relations
|
||||
litellm_team_table LiteLLM_TeamTable? @relation(fields: [team_id], references: [team_id])
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
keys LiteLLM_VerificationToken[]
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
}
|
||||
|
||||
// Audit table for deleted teams - preserves spend and team information for historical tracking
|
||||
model LiteLLM_DeletedTeamTable {
|
||||
id String @id @default(uuid())
|
||||
team_id String // Original team_id
|
||||
team_alias String?
|
||||
organization_id String?
|
||||
object_permission_id String?
|
||||
admins String[]
|
||||
members String[]
|
||||
members_with_roles Json @default("{}")
|
||||
metadata Json @default("{}")
|
||||
max_budget Float?
|
||||
soft_budget Float?
|
||||
spend Float @default(0.0)
|
||||
models String[]
|
||||
max_parallel_requests Int?
|
||||
tpm_limit BigInt?
|
||||
rpm_limit BigInt?
|
||||
budget_duration String?
|
||||
budget_reset_at DateTime?
|
||||
blocked Boolean @default(false)
|
||||
model_spend Json @default("{}")
|
||||
model_max_budget Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
team_member_permissions String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
policies String[] @default([])
|
||||
model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases
|
||||
allow_team_guardrail_config Boolean @default(false)
|
||||
|
||||
// Original timestamps from team creation/updates
|
||||
created_at DateTime? @map("created_at")
|
||||
updated_at DateTime? @map("updated_at")
|
||||
|
||||
// Deletion metadata
|
||||
deleted_at DateTime @default(now()) @map("deleted_at")
|
||||
deleted_by String? @map("deleted_by") // User who deleted the team
|
||||
deleted_by_api_key String? @map("deleted_by_api_key") // API key hash that performed the deletion
|
||||
litellm_changed_by String? @map("litellm_changed_by") // From litellm-changed-by header if provided
|
||||
|
||||
@@index([team_id])
|
||||
@@index([deleted_at])
|
||||
@@index([organization_id])
|
||||
@@index([team_alias])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
// Audit table for deleted teams - preserves spend and team information for historical tracking
|
||||
|
|
@ -230,6 +306,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
organizations LiteLLM_OrganizationTable[]
|
||||
users LiteLLM_UserTable[]
|
||||
|
|
@ -284,6 +361,7 @@ model LiteLLM_VerificationToken {
|
|||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
project_id String?
|
||||
permissions Json @default("{}")
|
||||
max_parallel_requests Int?
|
||||
metadata Json @default("{}")
|
||||
|
|
@ -314,6 +392,7 @@ model LiteLLM_VerificationToken {
|
|||
key_rotation_at DateTime? // When this key should next be rotated
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
|
||||
// SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub"
|
||||
|
|
@ -450,7 +529,7 @@ model LiteLLM_SpendLogs {
|
|||
custom_llm_provider String? @default("") // litellm used custom_llm_provider
|
||||
api_base String? @default("")
|
||||
user String? @default("")
|
||||
metadata Json? @default("{}")
|
||||
metadata Json? @default("{}") // project_id stored here
|
||||
cache_hit String? @default("")
|
||||
cache_key String? @default("")
|
||||
request_tags Json? @default("[]")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.41"
|
||||
version = "0.4.42"
|
||||
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.4.41"
|
||||
version = "0.4.42"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -578,7 +578,6 @@ OPENAI_CHAT_COMPLETION_PARAMS = [
|
|||
"thinking",
|
||||
"web_search_options",
|
||||
"service_tier",
|
||||
"store",
|
||||
"prompt_cache_key",
|
||||
"prompt_cache_retention",
|
||||
"safety_identifier",
|
||||
|
|
@ -644,6 +643,7 @@ DEFAULT_CHAT_COMPLETION_PARAM_VALUES = {
|
|||
"prompt_cache_retention": None,
|
||||
"store": None,
|
||||
"metadata": None,
|
||||
"context_management": None,
|
||||
}
|
||||
|
||||
openai_compatible_endpoints: List = [
|
||||
|
|
@ -1474,3 +1474,8 @@ MICROSOFT_USER_FIRST_NAME_ATTRIBUTE = str(
|
|||
MICROSOFT_USER_LAST_NAME_ATTRIBUTE = str(
|
||||
os.getenv("MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "surname")
|
||||
)
|
||||
|
||||
# Policy template enrichment
|
||||
MAX_COMPETITOR_NAMES = int(os.getenv("MAX_COMPETITOR_NAMES", 100))
|
||||
COMPETITOR_LLM_TEMPERATURE = float(os.getenv("COMPETITOR_LLM_TEMPERATURE", 0.3))
|
||||
DEFAULT_COMPETITOR_DISCOVERY_MODEL = "gpt-4o-mini"
|
||||
|
|
|
|||
|
|
@ -74,6 +74,14 @@ class ProjectedLimitExceededAlert(BaseBudgetAlertType):
|
|||
return user_info.token or "default_id"
|
||||
|
||||
|
||||
class ProjectBudgetAlert(BaseBudgetAlertType):
|
||||
def get_event_message(self) -> str:
|
||||
return "Project Budget: "
|
||||
|
||||
def get_id(self, user_info: CallInfo) -> str:
|
||||
return user_info.token or "default_id"
|
||||
|
||||
|
||||
def get_budget_alert_type(
|
||||
type: Literal[
|
||||
"token_budget",
|
||||
|
|
@ -84,6 +92,7 @@ def get_budget_alert_type(
|
|||
"organization_budget",
|
||||
"proxy_budget",
|
||||
"projected_limit_exceeded",
|
||||
"project_budget",
|
||||
],
|
||||
) -> BaseBudgetAlertType:
|
||||
"""Factory function to get the appropriate budget alert type class"""
|
||||
|
|
@ -97,6 +106,7 @@ def get_budget_alert_type(
|
|||
"organization_budget": OrganizationBudgetAlert(),
|
||||
"token_budget": TokenBudgetAlert(),
|
||||
"projected_limit_exceeded": ProjectedLimitExceededAlert(),
|
||||
"project_budget": ProjectBudgetAlert(),
|
||||
}
|
||||
|
||||
if type in alert_types:
|
||||
|
|
|
|||
|
|
@ -538,6 +538,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
"organization_budget",
|
||||
"proxy_budget",
|
||||
"projected_limit_exceeded",
|
||||
"project_budget",
|
||||
],
|
||||
user_info: CallInfo,
|
||||
):
|
||||
|
|
@ -1378,9 +1379,13 @@ Model Info:
|
|||
"""
|
||||
if self.alerting is None:
|
||||
return
|
||||
|
||||
|
||||
# Start periodic flush if not already started
|
||||
if not self.periodic_started and self.alerting is not None and len(self.alerting) > 0:
|
||||
if (
|
||||
not self.periodic_started
|
||||
and self.alerting is not None
|
||||
and len(self.alerting) > 0
|
||||
):
|
||||
asyncio.create_task(self.periodic_flush())
|
||||
self.periodic_started = True
|
||||
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.websearch_interception.tools import (
|
||||
get_litellm_web_search_tool,
|
||||
get_litellm_web_search_tool_openai,
|
||||
is_web_search_tool,
|
||||
is_web_search_tool_chat_completion,
|
||||
)
|
||||
|
|
@ -77,7 +78,13 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
that we can intercept and execute ourselves.
|
||||
"""
|
||||
# Check if this is for an enabled provider
|
||||
custom_llm_provider = kwargs.get("litellm_params", {}).get("custom_llm_provider", "")
|
||||
# Try top-level kwargs first, then nested litellm_params, then derive from model name
|
||||
custom_llm_provider = kwargs.get("custom_llm_provider", "") or kwargs.get("litellm_params", {}).get("custom_llm_provider", "")
|
||||
if not custom_llm_provider:
|
||||
try:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=kwargs.get("model", ""))
|
||||
except Exception:
|
||||
custom_llm_provider = ""
|
||||
if custom_llm_provider not in self.enabled_providers:
|
||||
return None
|
||||
|
||||
|
|
@ -101,7 +108,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
for tool in tools:
|
||||
if is_web_search_tool(tool):
|
||||
# Convert to LiteLLM standard web search tool
|
||||
converted_tool = get_litellm_web_search_tool()
|
||||
converted_tool = get_litellm_web_search_tool_openai()
|
||||
converted_tools.append(converted_tool)
|
||||
verbose_logger.debug(
|
||||
f"WebSearchInterception: Converted {tool.get('name', 'unknown')} "
|
||||
|
|
@ -111,8 +118,9 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
# Keep other tools as-is
|
||||
converted_tools.append(tool)
|
||||
|
||||
# Return modified kwargs with converted tools
|
||||
return {"tools": converted_tools}
|
||||
# Update tools in-place and return full kwargs
|
||||
kwargs["tools"] = converted_tools
|
||||
return kwargs
|
||||
|
||||
@classmethod
|
||||
def from_config_yaml(
|
||||
|
|
|
|||
|
|
@ -49,6 +49,39 @@ def get_litellm_web_search_tool() -> Dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
def get_litellm_web_search_tool_openai() -> Dict[str, Any]:
|
||||
"""
|
||||
Get the standard LiteLLM web search tool definition in OpenAI format.
|
||||
|
||||
Used by async_pre_call_deployment_hook which runs in the chat completions
|
||||
path where tools must be in OpenAI format (type: "function" with
|
||||
function.parameters).
|
||||
|
||||
Returns:
|
||||
Dict containing the OpenAI-style tool definition.
|
||||
"""
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": LITELLM_WEB_SEARCH_TOOL_NAME,
|
||||
"description": (
|
||||
"Search the web for information. Use this when you need current "
|
||||
"information or answers to questions that require up-to-date data."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "The search query to execute"
|
||||
}
|
||||
},
|
||||
"required": ["query"]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def is_web_search_tool_chat_completion(tool: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
Check if a tool is a web search tool for Chat Completions API (strict check).
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -191,6 +191,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"user",
|
||||
"web_search_options",
|
||||
"speed",
|
||||
"context_management",
|
||||
]
|
||||
|
||||
if "claude-3-7-sonnet" in model or supports_reasoning(
|
||||
|
|
@ -825,6 +826,62 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
return hosted_web_search_tool
|
||||
|
||||
@staticmethod
|
||||
def map_openai_context_management_to_anthropic(
|
||||
context_management: Union[List[Dict[str, Any]], Dict[str, Any]]
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
OpenAI format: [{"type": "compaction", "compact_threshold": 200000}]
|
||||
Anthropic format: {
|
||||
"edits": [
|
||||
{
|
||||
"type": "compact_20260112",
|
||||
"trigger": {"type": "input_tokens", "value": 150000}
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
Args:
|
||||
context_management: OpenAI or Anthropic context_management parameter
|
||||
|
||||
Returns:
|
||||
Anthropic-formatted context_management dict, or None if invalid
|
||||
"""
|
||||
# If already in Anthropic format (dict with 'edits'), pass through
|
||||
if isinstance(context_management, dict) and "edits" in context_management:
|
||||
return context_management
|
||||
|
||||
# If in OpenAI format (list), transform to Anthropic format
|
||||
if isinstance(context_management, list):
|
||||
anthropic_edits = []
|
||||
for entry in context_management:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
|
||||
entry_type = entry.get("type")
|
||||
if entry_type == "compaction":
|
||||
anthropic_edit: Dict[str, Any] = {
|
||||
"type": "compact_20260112"
|
||||
}
|
||||
compact_threshold = entry.get("compact_threshold")
|
||||
# Rewrite to 'trigger' with correct nesting if threshold exists
|
||||
if compact_threshold is not None and isinstance(compact_threshold, (int, float)):
|
||||
anthropic_edit["trigger"] = {
|
||||
"type": "input_tokens",
|
||||
"value": int(compact_threshold)
|
||||
}
|
||||
# Map any other keys by passthrough except handled ones
|
||||
for k in entry:
|
||||
if k not in {"type", "compact_threshold"}: # only passthrough other keys
|
||||
anthropic_edit[k] = entry[k]
|
||||
|
||||
anthropic_edits.append(anthropic_edit)
|
||||
|
||||
if anthropic_edits:
|
||||
return {"edits": anthropic_edits}
|
||||
|
||||
return None
|
||||
|
||||
def map_openai_params( # noqa: PLR0915
|
||||
self,
|
||||
non_default_params: dict,
|
||||
|
|
@ -931,9 +988,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
)
|
||||
elif param == "extra_headers":
|
||||
optional_params["extra_headers"] = value
|
||||
elif param == "context_management" and isinstance(value, dict):
|
||||
# Pass through Anthropic-specific context_management parameter
|
||||
optional_params["context_management"] = value
|
||||
elif param == "context_management":
|
||||
# Supports both OpenAI list format and Anthropic dict format
|
||||
if isinstance(value, (list, dict)):
|
||||
anthropic_context_management = self.map_openai_context_management_to_anthropic(value)
|
||||
if anthropic_context_management is not None:
|
||||
optional_params["context_management"] = anthropic_context_management
|
||||
elif param == "speed" and isinstance(value, str):
|
||||
# Pass through Anthropic-specific speed parameter for fast mode
|
||||
optional_params["speed"] = value
|
||||
|
|
@ -1094,32 +1154,39 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
headers["anthropic-beta"] = f"{existing_beta}, {beta_value}"
|
||||
|
||||
def _ensure_context_management_beta_header(
|
||||
self, headers: dict, context_management: dict
|
||||
self, headers: dict, context_management: object
|
||||
) -> None:
|
||||
"""
|
||||
Add appropriate beta headers based on context_management edits.
|
||||
- If any edit has type "compact_20260112", add compact-2026-01-12 header
|
||||
- For all other edits, add context-management-2025-06-27 header
|
||||
"""
|
||||
edits = context_management.get("edits", [])
|
||||
|
||||
edits = []
|
||||
# If anthropic format (dict with "edits" key)
|
||||
if isinstance(context_management, dict) and "edits" in context_management:
|
||||
edits = context_management.get("edits", [])
|
||||
# If OpenAI format: list of context management entries
|
||||
elif isinstance(context_management, list):
|
||||
edits = context_management
|
||||
# Defensive: ignore/fallback if context_management not valid
|
||||
else:
|
||||
return
|
||||
|
||||
has_compact = False
|
||||
has_other = False
|
||||
|
||||
|
||||
for edit in edits:
|
||||
edit_type = edit.get("type", "")
|
||||
if edit_type == "compact_20260112":
|
||||
if edit_type == "compact_20260112" or edit_type == "compaction":
|
||||
has_compact = True
|
||||
else:
|
||||
has_other = True
|
||||
|
||||
# Add compact header if any compact edits exist
|
||||
|
||||
# Add compact header if any compact edits/entries exist
|
||||
if has_compact:
|
||||
self._ensure_beta_header(
|
||||
headers, ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value
|
||||
)
|
||||
|
||||
# Add context management header if any other edits exist
|
||||
|
||||
# Add context management header if any other edits/entries exist
|
||||
if has_other:
|
||||
self._ensure_beta_header(
|
||||
headers, ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
|
||||
|
|
|
|||
|
|
@ -164,6 +164,17 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
# Remove system parameter if all content was filtered out
|
||||
anthropic_messages_optional_request_params.pop("system", None)
|
||||
|
||||
# Transform context_management from OpenAI format to Anthropic format if needed
|
||||
context_management_param = anthropic_messages_optional_request_params.get("context_management")
|
||||
if context_management_param is not None:
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
transformed_context_management = AnthropicConfig.map_openai_context_management_to_anthropic(
|
||||
context_management_param
|
||||
)
|
||||
if transformed_context_management is not None:
|
||||
anthropic_messages_optional_request_params["context_management"] = transformed_context_management
|
||||
|
||||
####### get required params for all anthropic messages requests ######
|
||||
verbose_logger.debug(f"TRANSFORMATION DEBUG - Messages: {messages}")
|
||||
anthropic_messages_request: AnthropicMessagesRequest = AnthropicMessagesRequest(
|
||||
|
|
|
|||
|
|
@ -106,6 +106,7 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
"audio",
|
||||
"web_search_options",
|
||||
"prompt_cache_key",
|
||||
"store",
|
||||
]
|
||||
|
||||
def _is_response_format_supported_model(self, model: str) -> bool:
|
||||
|
|
@ -158,7 +159,6 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
api_version: str = "",
|
||||
) -> dict:
|
||||
supported_openai_params = self.get_supported_openai_params(model)
|
||||
|
||||
api_version_times = api_version.split("-")
|
||||
|
||||
if len(api_version_times) >= 3:
|
||||
|
|
@ -245,7 +245,6 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
optional_params["tools"].extend(value)
|
||||
elif param in supported_openai_params:
|
||||
optional_params[param] = value
|
||||
|
||||
return optional_params
|
||||
|
||||
def transform_request(
|
||||
|
|
|
|||
|
|
@ -4,9 +4,6 @@ Translates from OpenAI's `/v1/chat/completions` to DashScope's `/v1/chat/complet
|
|||
|
||||
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
|
@ -32,10 +29,6 @@ class DashScopeChatConfig(OpenAIGPTConfig):
|
|||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
|
||||
"""
|
||||
DashScope does not support content in list format.
|
||||
"""
|
||||
messages = handle_messages_with_content_list_to_str_conversion(messages)
|
||||
if is_async:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=True
|
||||
|
|
|
|||
|
|
@ -137,10 +137,29 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
|
|||
self, messages: List[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
|
||||
"""
|
||||
Support translating video files from file_id or file_data to video_url
|
||||
Support translating:
|
||||
- video files from file_id or file_data to video_url
|
||||
- thinking_blocks on assistant messages to content blocks
|
||||
"""
|
||||
for message in messages:
|
||||
if message["role"] == "user":
|
||||
if message["role"] == "assistant":
|
||||
thinking_blocks = message.pop("thinking_blocks", None) # type: ignore
|
||||
if thinking_blocks:
|
||||
new_content: list = [
|
||||
{"type": block["type"], "thinking": block.get("thinking", "")}
|
||||
if block.get("type") == "thinking"
|
||||
else {"type": block["type"], "data": block.get("data", "")}
|
||||
for block in thinking_blocks
|
||||
]
|
||||
existing_content = message.get("content")
|
||||
if isinstance(existing_content, str):
|
||||
new_content.append(
|
||||
{"type": "text", "text": existing_content}
|
||||
)
|
||||
elif isinstance(existing_content, list):
|
||||
new_content.extend(existing_content)
|
||||
message["content"] = new_content # type: ignore
|
||||
elif message["role"] == "user":
|
||||
message_content = message.get("content")
|
||||
if message_content and isinstance(message_content, list):
|
||||
replaced_content_items: List[
|
||||
|
|
|
|||
|
|
@ -162,6 +162,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
"service_tier",
|
||||
"safety_identifier",
|
||||
"prompt_cache_key",
|
||||
"store",
|
||||
] # works across all models
|
||||
|
||||
model_specific_params = []
|
||||
|
|
|
|||
|
|
@ -767,14 +767,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
elif reasoning_effort == "low":
|
||||
return {"thinkingLevel": "low", "includeThoughts": True}
|
||||
elif reasoning_effort == "medium":
|
||||
# For gemini-3-flash-preview, medium maps to "medium", otherwise "high"
|
||||
if is_gemini3flash:
|
||||
return {"thinkingLevel": "medium", "includeThoughts": True}
|
||||
else:
|
||||
return {
|
||||
"thinkingLevel": "high",
|
||||
"includeThoughts": True,
|
||||
} # medium is not out yet for other models
|
||||
elif reasoning_effort == "high":
|
||||
return {"thinkingLevel": "high", "includeThoughts": True}
|
||||
elif reasoning_effort == "disable":
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import List, Optional, Tuple
|
||||
from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -11,9 +11,18 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import Choices, ModelResponse, Usage, PromptTokensDetailsWrapper
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from ...openai.chat.gpt_transformation import (
|
||||
OpenAIChatCompletionStreamingHandler,
|
||||
OpenAIGPTConfig,
|
||||
)
|
||||
|
||||
|
||||
class XAIChatConfig(OpenAIGPTConfig):
|
||||
|
|
@ -119,6 +128,18 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
optional_params[param] = value
|
||||
return optional_params
|
||||
|
||||
def get_model_response_iterator(
|
||||
self,
|
||||
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
) -> Any:
|
||||
return XAIChatCompletionStreamingHandler(
|
||||
streaming_response=streaming_response,
|
||||
sync_stream=sync_stream,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -225,3 +246,25 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
usage.prompt_tokens_details.web_search_requests = int(num_sources_used)
|
||||
setattr(usage, "num_sources_used", int(num_sources_used))
|
||||
verbose_logger.debug(f"X.AI web search sources used: {num_sources_used}")
|
||||
|
||||
|
||||
class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
|
||||
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
|
||||
"""
|
||||
Handle xAI-specific streaming behavior.
|
||||
|
||||
xAI Grok sends a final chunk with empty choices array but with usage data
|
||||
when stream_options={"include_usage": True} is set.
|
||||
|
||||
Example from xAI API:
|
||||
{"id":"...","object":"chat.completion.chunk","created":...,"model":"grok-4-1-fast-non-reasoning",
|
||||
"choices":[],"usage":{"prompt_tokens":171,"completion_tokens":2,"total_tokens":173,...}}
|
||||
"""
|
||||
# Handle chunks with empty choices but with usage data
|
||||
choices = chunk.get("choices", [])
|
||||
if len(choices) == 0 and "usage" in chunk:
|
||||
# xAI sends usage in a chunk with empty choices array
|
||||
# Add a dummy choice with empty delta to ensure proper processing
|
||||
chunk["choices"] = [{"index": 0, "delta": {}, "finish_reason": None}]
|
||||
|
||||
return super().chunk_parser(chunk)
|
||||
|
|
|
|||
|
|
@ -14696,6 +14696,108 @@
|
|||
"supports_web_search": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"gemini-3.1-pro-preview-customtools": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"vertex_ai/gemini-3-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
|
|
@ -14789,6 +14891,108 @@
|
|||
"supports_web_search": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"vertex_ai/gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"vertex_ai/gemini-3.1-pro-preview-customtools": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"gemini-2.5-pro-exp-03-25": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
|
|
@ -16751,6 +16955,108 @@
|
|||
"supports_native_streaming": true,
|
||||
"tpm": 800000
|
||||
},
|
||||
"gemini/gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"rpm": 2000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-3.1-pro-preview",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 800000
|
||||
},
|
||||
"gemini/gemini-3.1-pro-preview-customtools": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"rpm": 2000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-3.1-pro-preview",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 800000
|
||||
},
|
||||
"gemini-3-flash-preview": {
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Dict, List, Optional, Union
|
||||
from typing import Dict, List, Mapping, Optional, Union
|
||||
from urllib.parse import parse_qs
|
||||
|
||||
import httpx
|
||||
|
|
@ -9,7 +9,9 @@ from litellm.constants import PASS_THROUGH_HEADER_PREFIX
|
|||
class BasePassthroughUtils:
|
||||
@staticmethod
|
||||
def get_merged_query_parameters(
|
||||
existing_url: httpx.URL, request_query_params: Dict[str, Union[str, list]]
|
||||
existing_url: httpx.URL,
|
||||
request_query_params: Mapping[str, Union[str, list]],
|
||||
default_query_params: Optional[Dict[str, Union[str, list]]] = None
|
||||
) -> Dict[str, Union[str, List[str]]]:
|
||||
# Get the existing query params from the target URL
|
||||
existing_query_string = existing_url.query.decode("utf-8")
|
||||
|
|
@ -19,8 +21,19 @@ class BasePassthroughUtils:
|
|||
updated_existing_query_params = {
|
||||
k: v[0] if len(v) == 1 else v for k, v in existing_query_params.items()
|
||||
}
|
||||
# Merge the query params, giving priority to the existing ones
|
||||
return {**request_query_params, **updated_existing_query_params}
|
||||
|
||||
# Start with default query params (lowest priority)
|
||||
merged_params = {}
|
||||
if default_query_params:
|
||||
merged_params.update(default_query_params)
|
||||
|
||||
# Override with existing URL query params (medium priority)
|
||||
merged_params.update(updated_existing_query_params)
|
||||
|
||||
# Override with request query params (highest priority - client can override anything)
|
||||
merged_params.update(request_query_params)
|
||||
|
||||
return merged_params
|
||||
|
||||
@staticmethod
|
||||
def forward_headers_from_request(
|
||||
|
|
|
|||
|
|
@ -1522,7 +1522,13 @@
|
|||
"guardrails": [
|
||||
"aviation-ops-data-protection",
|
||||
"aviation-safety-topic-filter",
|
||||
"airline-brand-protection-filter"
|
||||
"airline-brand-protection-filter",
|
||||
"competitor-name-input-blocker",
|
||||
"competitor-name-output-blocker",
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"complexity": "High",
|
||||
"parameters": [
|
||||
|
|
@ -1531,9 +1537,14 @@
|
|||
"label": "Your Airline / Brand Name",
|
||||
"type": "text",
|
||||
"required": true,
|
||||
"placeholder": "e.g. Emirates"
|
||||
"placeholder": "e.g. Acme Airlines"
|
||||
}
|
||||
],
|
||||
"llm_enrichment": {
|
||||
"parameter": "brand_name",
|
||||
"prompt": "List the top 30 direct competitors of {{brand_name}} in the airline industry. Include major international carriers, regional competitors, and low-cost carriers that operate on overlapping routes. Return ONLY airline/brand names, one per line, no numbering, no explanations.",
|
||||
"result_key": "competitors"
|
||||
},
|
||||
"guardrailDefinitions": [
|
||||
{
|
||||
"guardrail_name": "aviation-ops-data-protection",
|
||||
|
|
@ -1675,6 +1686,72 @@
|
|||
"guardrail_info": {
|
||||
"description": "Blocks AI-generated fake incident reports, unauthorized statements, and reputation-damaging content about your brand (runs on output)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-name-input-blocker",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs that mention competitor names (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-name-output-blocker",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI outputs that mention competitor names (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user requests asking to recommend competitors (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI from recommending or suggesting competitor services (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs requesting unfavorable brand comparisons (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI outputs with unfavorable brand comparisons (post_call)"
|
||||
}
|
||||
}
|
||||
],
|
||||
"templateData": {
|
||||
|
|
@ -1683,7 +1760,13 @@
|
|||
"guardrails_add": [
|
||||
"aviation-ops-data-protection",
|
||||
"aviation-safety-topic-filter",
|
||||
"airline-brand-protection-filter"
|
||||
"airline-brand-protection-filter",
|
||||
"competitor-name-input-blocker",
|
||||
"competitor-name-output-blocker",
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"guardrails_remove": []
|
||||
},
|
||||
|
|
@ -1812,9 +1895,12 @@
|
|||
"iconColor": "text-orange-500",
|
||||
"iconBg": "bg-orange-50",
|
||||
"guardrails": [
|
||||
"competitor-input-blocker",
|
||||
"competitor-output-blocker",
|
||||
"competitor-recommendation-filter",
|
||||
"competitor-comparison-filter"
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"complexity": "Medium",
|
||||
"parameters": [
|
||||
|
|
@ -1823,15 +1909,26 @@
|
|||
"label": "Your Brand Name",
|
||||
"type": "text",
|
||||
"required": true,
|
||||
"placeholder": "e.g. Emirates"
|
||||
"placeholder": "e.g. Acme Airlines"
|
||||
}
|
||||
],
|
||||
"llm_enrichment": {
|
||||
"parameter": "brand_name",
|
||||
"prompt": "List the top 10 direct competitors of {{brand_name}} in the same industry. Return ONLY company/brand names, one per line, no numbering, no explanations.",
|
||||
"prompt": "List the top 30 direct competitors of {{brand_name}} in the same industry. Return ONLY company/brand names, one per line, no numbering, no explanations.",
|
||||
"result_key": "competitors"
|
||||
},
|
||||
"guardrailDefinitions": [
|
||||
{
|
||||
"guardrail_name": "competitor-input-blocker",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs that mention competitor brands (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-output-blocker",
|
||||
"litellm_params": {
|
||||
|
|
@ -1840,39 +1937,64 @@
|
|||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI outputs that mention or promote competitor brands (auto-discovered via LLM)"
|
||||
"description": "Blocks AI outputs that mention competitor brands (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-filter",
|
||||
"guardrail_name": "competitor-recommendation-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user requests asking to recommend competitors (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI from recommending, suggesting, or directing users to competitor services"
|
||||
"description": "Blocks AI from recommending or suggesting competitor services (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-filter",
|
||||
"guardrail_name": "competitor-comparison-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs requesting unfavorable brand comparisons (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks unfavorable comparisons between your brand and competitors in AI outputs"
|
||||
"description": "Blocks AI outputs with unfavorable brand comparisons (post_call)"
|
||||
}
|
||||
}
|
||||
],
|
||||
"templateData": {
|
||||
"policy_name": "competitor-mention-detection",
|
||||
"description": "Detects and blocks competitor mentions in AI outputs. Uses LLM-powered competitor discovery based on your brand name.",
|
||||
"description": "Detects and blocks competitor mentions in both inputs and outputs. Uses LLM-powered competitor discovery based on your brand name.",
|
||||
"guardrails_add": [
|
||||
"competitor-input-blocker",
|
||||
"competitor-output-blocker",
|
||||
"competitor-recommendation-filter",
|
||||
"competitor-comparison-filter"
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"guardrails_remove": []
|
||||
},
|
||||
|
|
|
|||
|
|
@ -198,6 +198,7 @@ class Litellm_EntityType(enum.Enum):
|
|||
TEAM = "team"
|
||||
TEAM_MEMBER = "team_member"
|
||||
ORGANIZATION = "organization"
|
||||
PROJECT = "project"
|
||||
TAG = "tag"
|
||||
|
||||
# global proxy level entity
|
||||
|
|
@ -237,6 +238,9 @@ class KeyManagementRoutes(str, enum.Enum):
|
|||
# list routes
|
||||
KEY_LIST = "/key/list"
|
||||
|
||||
# team usage routes
|
||||
TEAM_DAILY_ACTIVITY = "/team/daily/activity"
|
||||
|
||||
|
||||
class LiteLLMRoutes(enum.Enum):
|
||||
openai_route_names = [
|
||||
|
|
@ -505,6 +509,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
KeyManagementRoutes.KEY_BLOCK.value,
|
||||
KeyManagementRoutes.KEY_UNBLOCK.value,
|
||||
KeyManagementRoutes.KEY_BULK_UPDATE.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
|
||||
]
|
||||
|
||||
management_routes = [
|
||||
|
|
@ -925,6 +930,7 @@ class GenerateKeyRequest(KeyRequestBase):
|
|||
description="How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True",
|
||||
)
|
||||
organization_id: Optional[str] = None
|
||||
project_id: Optional[str] = None
|
||||
|
||||
|
||||
class GenerateKeyResponse(KeyRequestBase):
|
||||
|
|
@ -934,6 +940,7 @@ class GenerateKeyResponse(KeyRequestBase):
|
|||
user_id: Optional[str] = None
|
||||
token_id: Optional[str] = None
|
||||
organization_id: Optional[str] = None
|
||||
project_id: Optional[str] = None
|
||||
litellm_budget_table: Optional[Any] = None
|
||||
token: Optional[str] = None
|
||||
created_by: Optional[str] = None
|
||||
|
|
@ -1909,6 +1916,10 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase):
|
|||
default={},
|
||||
description="Key-value pairs of headers to be forwarded with the request. You can set any key value pair here and it will be forwarded to your target endpoint",
|
||||
)
|
||||
default_query_params: dict = Field(
|
||||
default={},
|
||||
description="Key-value pairs of default query parameters to be sent with every request to this endpoint. These can be overridden by client-provided query parameters. For example: {'key': 'default_value', 'api_version': '2023-01'}",
|
||||
)
|
||||
include_subpath: bool = Field(
|
||||
default=False,
|
||||
description="If True, requests to subpaths of the path will be forwarded to the target endpoint. For example, if the path is /bria and include_subpath is True, requests to /bria/v1/text-to-image/base/2.3 will be forwarded to the target endpoint.",
|
||||
|
|
@ -1929,6 +1940,10 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase):
|
|||
default=False,
|
||||
description="True if this endpoint is defined in the config file, False if from DB. Config-defined endpoints cannot be edited via the UI.",
|
||||
)
|
||||
methods: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="List of HTTP methods this endpoint handles (e.g., ['GET', 'POST']). If None or empty, all methods (GET, POST, PUT, DELETE, PATCH) are supported for backward compatibility. This allows the same path to have different targets for different HTTP methods.",
|
||||
)
|
||||
|
||||
|
||||
class PassThroughEndpointResponse(LiteLLMPydanticObjectBase):
|
||||
|
|
@ -2171,6 +2186,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase):
|
|||
config: Dict = {}
|
||||
user_id: Optional[str] = None
|
||||
team_id: Optional[str] = None
|
||||
project_id: Optional[str] = None
|
||||
max_parallel_requests: Optional[int] = None
|
||||
metadata: Dict = {}
|
||||
tpm_limit: Optional[int] = None
|
||||
|
|
@ -2522,6 +2538,116 @@ class NewOrganizationResponse(LiteLLM_OrganizationTable):
|
|||
updated_at: datetime
|
||||
|
||||
|
||||
### PROJECT MANAGEMENT TYPES ###
|
||||
|
||||
|
||||
class ProjectBase(LiteLLMPydanticObjectBase):
|
||||
"""Base fields shared by project create/update requests"""
|
||||
|
||||
project_id: Optional[str] = None
|
||||
project_alias: Optional[str] = None
|
||||
team_id: Optional[str] = None
|
||||
metadata: Optional[dict] = None
|
||||
models: Optional[List[str]] = None
|
||||
blocked: bool = False
|
||||
|
||||
|
||||
class NewProjectRequest(LiteLLM_BudgetTable):
|
||||
"""Request model for POST /project/new"""
|
||||
|
||||
project_id: Optional[str] = None
|
||||
project_alias: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
team_id: str
|
||||
budget_id: Optional[str] = None
|
||||
metadata: Optional[dict] = None
|
||||
models: List[str] = []
|
||||
model_rpm_limit: Optional[dict] = None
|
||||
model_tpm_limit: Optional[dict] = None
|
||||
blocked: bool = False
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def set_model_info(cls, values):
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields:
|
||||
if values.get(field) is not None:
|
||||
if values.get("metadata") is None:
|
||||
values.update({"metadata": {}})
|
||||
values["metadata"][field] = values.get(field)
|
||||
values.pop(field)
|
||||
return values
|
||||
|
||||
|
||||
class UpdateProjectRequest(LiteLLM_BudgetTable):
|
||||
"""Request model for POST /project/update"""
|
||||
|
||||
project_id: str
|
||||
project_alias: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
team_id: Optional[str] = None
|
||||
metadata: Optional[dict] = None
|
||||
models: Optional[List[str]] = None
|
||||
model_rpm_limit: Optional[dict] = None
|
||||
model_tpm_limit: Optional[dict] = None
|
||||
blocked: Optional[bool] = None
|
||||
budget_id: Optional[str] = None
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def set_model_info(cls, values):
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields:
|
||||
if values.get(field) is not None:
|
||||
if values.get("metadata") is None:
|
||||
values.update({"metadata": {}})
|
||||
values["metadata"][field] = values.get(field)
|
||||
values.pop(field)
|
||||
return values
|
||||
|
||||
|
||||
class DeleteProjectRequest(LiteLLMPydanticObjectBase):
|
||||
"""Request model for DELETE /project/delete"""
|
||||
|
||||
project_ids: List[str]
|
||||
|
||||
|
||||
class LiteLLM_ProjectTable(LiteLLMPydanticObjectBase):
|
||||
"""Database model representation for project"""
|
||||
|
||||
project_id: str
|
||||
project_alias: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
team_id: Optional[str] = None
|
||||
budget_id: Optional[str] = None
|
||||
metadata: Optional[dict] = None
|
||||
models: List[str] = []
|
||||
spend: float = 0.0
|
||||
model_spend: Optional[dict] = None
|
||||
model_rpm_limit: Optional[dict] = None
|
||||
model_tpm_limit: Optional[dict] = None
|
||||
blocked: bool = False
|
||||
object_permission_id: Optional[str] = None
|
||||
created_by: str
|
||||
updated_by: str
|
||||
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
|
||||
|
||||
|
||||
class NewProjectResponse(LiteLLM_ProjectTable):
|
||||
"""Response model for POST /project/new"""
|
||||
|
||||
project_id: str
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class LiteLLM_ProjectTableCachedObj(LiteLLM_ProjectTable):
|
||||
"""Cached version for auth checks. Mirrors LiteLLM_TeamTableCachedObj pattern."""
|
||||
|
||||
last_refreshed_at: Optional[float] = None
|
||||
|
||||
|
||||
class LiteLLM_UserTableFiltered(BaseModel): # done to avoid exposing sensitive data
|
||||
user_id: str
|
||||
user_email: Optional[str] = None
|
||||
|
|
@ -2893,6 +3019,7 @@ class SpendLogsMetadata(TypedDict):
|
|||
user_api_key: Optional[str]
|
||||
user_api_key_alias: Optional[str]
|
||||
user_api_key_team_id: Optional[str]
|
||||
user_api_key_project_id: Optional[str]
|
||||
user_api_key_org_id: Optional[str]
|
||||
user_api_key_user_id: Optional[str]
|
||||
user_api_key_team_alias: Optional[str]
|
||||
|
|
@ -3130,6 +3257,11 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
Organization does not have access to the model
|
||||
"""
|
||||
|
||||
project_model_access_denied = "project_model_access_denied"
|
||||
"""
|
||||
Project does not have access to the model
|
||||
"""
|
||||
|
||||
expired_key = "expired_key"
|
||||
"""
|
||||
Key has expired
|
||||
|
|
@ -3192,7 +3324,7 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
|
||||
@classmethod
|
||||
def get_model_access_error_type_for_object(
|
||||
cls, object_type: Literal["key", "user", "team", "org"]
|
||||
cls, object_type: Literal["key", "user", "team", "org", "project"]
|
||||
) -> "ProxyErrorTypes":
|
||||
"""
|
||||
Get the model access error type for object_type
|
||||
|
|
@ -3205,6 +3337,8 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
return cls.user_model_access_denied
|
||||
elif object_type == "org":
|
||||
return cls.org_model_access_denied
|
||||
elif object_type == "project":
|
||||
return cls.project_model_access_denied
|
||||
|
||||
@classmethod
|
||||
def get_vector_store_access_error_type_for_object(
|
||||
|
|
|
|||
|
|
@ -11,8 +11,7 @@ Run checks for:
|
|||
import asyncio
|
||||
import re
|
||||
import time
|
||||
from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union,
|
||||
cast)
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -21,27 +20,42 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.dual_cache import LimitedSizeOrderedDict
|
||||
from litellm.constants import (CLI_JWT_EXPIRATION_HOURS, CLI_JWT_TOKEN_NAME,
|
||||
DEFAULT_ACCESS_GROUP_CACHE_TTL,
|
||||
DEFAULT_IN_MEMORY_TTL,
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE)
|
||||
from litellm.constants import (
|
||||
CLI_JWT_EXPIRATION_HOURS,
|
||||
CLI_JWT_TOKEN_NAME,
|
||||
DEFAULT_ACCESS_GROUP_CACHE_TTL,
|
||||
DEFAULT_IN_MEMORY_TTL,
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE,
|
||||
)
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.proxy._types import (RBAC_ROLES, CallInfo,
|
||||
LiteLLM_AccessGroupTable,
|
||||
LiteLLM_BudgetTable, LiteLLM_EndUserTable,
|
||||
Litellm_EntityType, LiteLLM_JWTAuth,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_OrganizationMembershipTable,
|
||||
LiteLLM_OrganizationTable, LiteLLM_TagTable,
|
||||
LiteLLM_TeamMembership, LiteLLM_TeamTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable, LiteLLMRoutes,
|
||||
LitellmUserRoles, NewTeamRequest,
|
||||
ProxyErrorTypes, ProxyException,
|
||||
RoleBasedPermissions, SpecialModelNames,
|
||||
UserAPIKeyAuth)
|
||||
from litellm.proxy._types import (
|
||||
RBAC_ROLES,
|
||||
CallInfo,
|
||||
LiteLLM_AccessGroupTable,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_EndUserTable,
|
||||
Litellm_EntityType,
|
||||
LiteLLM_JWTAuth,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_OrganizationMembershipTable,
|
||||
LiteLLM_OrganizationTable,
|
||||
LiteLLM_TagTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_ProjectTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
LiteLLMRoutes,
|
||||
LitellmUserRoles,
|
||||
NewTeamRequest,
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
RoleBasedPermissions,
|
||||
SpecialModelNames,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
|
||||
|
|
@ -64,6 +78,7 @@ db_cache_expiry = DEFAULT_IN_MEMORY_TTL # refresh every 5s
|
|||
|
||||
all_routes = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
|
||||
|
||||
|
||||
def _log_budget_lookup_failure(entity: str, error: Exception) -> None:
|
||||
"""
|
||||
Log a warning when budget lookup fails; cache will not be populated.
|
||||
|
|
@ -81,38 +96,41 @@ def _log_budget_lookup_failure(entity: str, error: Exception) -> None:
|
|||
x in err_str
|
||||
for x in ("column", "schema", "does not exist", "prisma", "migrate")
|
||||
):
|
||||
hint = " Run `prisma db push` or `prisma migrate deploy` to fix schema mismatches."
|
||||
hint = (
|
||||
" Run `prisma db push` or `prisma migrate deploy` to fix schema mismatches."
|
||||
)
|
||||
verbose_proxy_logger.error(
|
||||
f"Budget lookup failed for {entity}; cache will not be populated. "
|
||||
f"Each request will hit the database. Error: {error}.{hint}"
|
||||
)
|
||||
|
||||
|
||||
def _is_model_cost_zero(
|
||||
model: Optional[Union[str, List[str]]], llm_router: Optional[Router]
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a model has zero cost (no configured pricing).
|
||||
|
||||
|
||||
Uses the router's get_model_group_info method to get pricing information.
|
||||
|
||||
|
||||
Args:
|
||||
model: The model name or list of model names
|
||||
llm_router: The LiteLLM router instance
|
||||
|
||||
|
||||
Returns:
|
||||
bool: True if all costs for the model are zero, False otherwise
|
||||
"""
|
||||
if model is None or llm_router is None:
|
||||
return False
|
||||
|
||||
|
||||
# Handle list of models
|
||||
model_list = [model] if isinstance(model, str) else model
|
||||
|
||||
|
||||
for model_name in model_list:
|
||||
try:
|
||||
# Use router's get_model_group_info method directly for better reliability
|
||||
model_group_info = llm_router.get_model_group_info(model_group=model_name)
|
||||
|
||||
|
||||
if model_group_info is None:
|
||||
# Model not found or no pricing info available
|
||||
# Conservative approach: assume it has cost
|
||||
|
|
@ -120,42 +138,87 @@ def _is_model_cost_zero(
|
|||
f"No model group info found for {model_name}, assuming it has cost"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
# Check costs for this model
|
||||
# Only allow bypass if BOTH costs are explicitly set to 0 (not None)
|
||||
input_cost = model_group_info.input_cost_per_token
|
||||
output_cost = model_group_info.output_cost_per_token
|
||||
|
||||
|
||||
# If costs are not explicitly configured (None), assume it has cost
|
||||
if input_cost is None or output_cost is None:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Model {model_name} has undefined cost (input: {input_cost}, output: {output_cost}), assuming it has cost"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
# If either cost is non-zero, return False
|
||||
if input_cost > 0 or output_cost > 0:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Model {model_name} has non-zero cost (input: {input_cost}, output: {output_cost})"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
# This model has zero cost explicitly configured
|
||||
verbose_proxy_logger.debug(
|
||||
f"Model {model_name} has zero cost explicitly configured (input: {input_cost}, output: {output_cost})"
|
||||
)
|
||||
|
||||
|
||||
except Exception as e:
|
||||
# If we can't determine the cost, assume it has cost (conservative approach)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Error checking cost for model {model_name}: {str(e)}, assuming it has cost"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
# All models checked have zero cost
|
||||
return True
|
||||
|
||||
|
||||
async def _run_project_checks(
|
||||
project_object: Optional[LiteLLM_ProjectTableCachedObj],
|
||||
_model: Optional[Union[str, List[str]]],
|
||||
llm_router: Optional[Router],
|
||||
skip_budget_checks: bool,
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> None:
|
||||
"""
|
||||
Run all project-level checks: blocked, model access, budget, soft budget.
|
||||
Extracted from common_checks() to keep statement count manageable.
|
||||
"""
|
||||
if project_object is None:
|
||||
return
|
||||
|
||||
# 1.1. If project is blocked
|
||||
if project_object.blocked is True:
|
||||
raise Exception(
|
||||
f"Project={project_object.project_id} is blocked. Update via `/project/update` if you're an admin."
|
||||
)
|
||||
|
||||
# 2.2 If project can call model
|
||||
if _model and len(project_object.models) > 0:
|
||||
can_project_access_model(
|
||||
model=_model,
|
||||
project_object=project_object,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
if not skip_budget_checks:
|
||||
# 3.0.2. If project is in budget
|
||||
await _project_max_budget_check(
|
||||
project_object=project_object,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# 3.0.3. If project is over soft budget (alert only, doesn't block)
|
||||
await _project_soft_budget_check(
|
||||
project_object=project_object,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
async def common_checks(
|
||||
request_body: dict,
|
||||
team_object: Optional[LiteLLM_TeamTable],
|
||||
|
|
@ -169,13 +232,18 @@ async def common_checks(
|
|||
valid_token: Optional[UserAPIKeyAuth],
|
||||
request: Request,
|
||||
skip_budget_checks: bool = False,
|
||||
project_object: Optional[LiteLLM_ProjectTableCachedObj] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Common checks across jwt + key-based auth.
|
||||
|
||||
1. If team is blocked
|
||||
1.1. If project is blocked
|
||||
2. If team can call model
|
||||
2.2 If project can call model
|
||||
3. If team is in budget
|
||||
3.0.2. If project is in budget
|
||||
3.0.3. If project is over soft budget (alert only)
|
||||
4. If user passed in (JWT or key.user_id) - is in budget
|
||||
5. If end_user (either via JWT or 'user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
6. [OPTIONAL] If 'enforce_end_user' enabled - did developer pass in 'user' param for openai endpoints
|
||||
|
|
@ -220,6 +288,16 @@ async def common_checks(
|
|||
user_object=user_object,
|
||||
)
|
||||
|
||||
# 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget)
|
||||
await _run_project_checks(
|
||||
project_object=project_object,
|
||||
_model=_model,
|
||||
llm_router=llm_router,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# If this is a free model, skip all budget checks
|
||||
if not skip_budget_checks:
|
||||
# 3. If team is in budget
|
||||
|
|
@ -279,7 +357,10 @@ async def common_checks(
|
|||
)
|
||||
|
||||
# 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
if end_user_object is not None and end_user_object.litellm_budget_table is not None:
|
||||
if (
|
||||
end_user_object is not None
|
||||
and end_user_object.litellm_budget_table is not None
|
||||
):
|
||||
end_user_budget = end_user_object.litellm_budget_table.max_budget
|
||||
if end_user_budget is not None and end_user_object.spend > end_user_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -353,8 +434,7 @@ async def common_checks(
|
|||
_request_metadata: dict = request_body.get("metadata", {}) or {}
|
||||
if _request_metadata.get("guardrails"):
|
||||
# check if team allowed to modify guardrails
|
||||
from litellm.proxy.guardrails.guardrail_helpers import \
|
||||
can_modify_guardrails
|
||||
from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails
|
||||
|
||||
can_modify: bool = can_modify_guardrails(team_object)
|
||||
if can_modify is False:
|
||||
|
|
@ -529,11 +609,7 @@ def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool:
|
|||
|
||||
|
||||
def allowed_routes_check(
|
||||
user_role: Literal[
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.TEAM,
|
||||
LitellmUserRoles.INTERNAL_USER,
|
||||
],
|
||||
user_role: LitellmUserRoles,
|
||||
user_route: str,
|
||||
litellm_proxy_roles: LiteLLM_JWTAuth,
|
||||
) -> bool:
|
||||
|
|
@ -1358,7 +1434,7 @@ async def _get_team_object_from_user_api_key_cache(
|
|||
raise Exception
|
||||
|
||||
_response = LiteLLM_TeamTableCachedObj(**response.dict())
|
||||
|
||||
|
||||
# Load object_permission if object_permission_id exists but object_permission is not loaded
|
||||
if _response.object_permission_id and not _response.object_permission:
|
||||
try:
|
||||
|
|
@ -1373,7 +1449,7 @@ async def _get_team_object_from_user_api_key_cache(
|
|||
verbose_proxy_logger.debug(
|
||||
f"Failed to load object_permission for team {team_id} with object_permission_id={_response.object_permission_id}: {e}"
|
||||
)
|
||||
|
||||
|
||||
# save the team object to cache
|
||||
await _cache_team_object(
|
||||
team_id=team_id,
|
||||
|
|
@ -1800,8 +1876,9 @@ class ExperimentalUIJWTToken:
|
|||
def get_experimental_ui_login_jwt_auth_token(user_info: LiteLLM_UserTable) -> str:
|
||||
from datetime import timedelta
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import \
|
||||
encrypt_value_helper
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
if user_info.user_role is None:
|
||||
raise Exception("User role is required for experimental UI login")
|
||||
|
|
@ -1847,8 +1924,9 @@ class ExperimentalUIJWTToken:
|
|||
"""
|
||||
from datetime import timedelta
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import \
|
||||
encrypt_value_helper
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
if user_info.user_role is None:
|
||||
raise Exception("User role is required for CLI JWT login")
|
||||
|
|
@ -1887,8 +1965,9 @@ class ExperimentalUIJWTToken:
|
|||
import json
|
||||
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import \
|
||||
decrypt_value_helper
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
decrypt_value_helper,
|
||||
)
|
||||
|
||||
decrypted_token = decrypt_value_helper(
|
||||
hashed_token, key="ui_hash_key", exception_type="debug"
|
||||
|
|
@ -2136,10 +2215,8 @@ async def _get_resources_from_access_groups(
|
|||
# Lazy import to avoid circular imports
|
||||
if prisma_client is None or user_api_key_cache is None:
|
||||
from litellm.proxy.proxy_server import prisma_client as _prisma_client
|
||||
from litellm.proxy.proxy_server import \
|
||||
proxy_logging_obj as _proxy_logging_obj
|
||||
from litellm.proxy.proxy_server import \
|
||||
user_api_key_cache as _user_api_key_cache
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging_obj
|
||||
from litellm.proxy.proxy_server import user_api_key_cache as _user_api_key_cache
|
||||
|
||||
prisma_client = prisma_client or _prisma_client
|
||||
user_api_key_cache = user_api_key_cache or _user_api_key_cache
|
||||
|
|
@ -2280,7 +2357,7 @@ def _can_object_call_model(
|
|||
models: List[str],
|
||||
team_model_aliases: Optional[Dict[str, str]] = None,
|
||||
team_id: Optional[str] = None,
|
||||
object_type: Literal["user", "team", "key", "org"] = "user",
|
||||
object_type: Literal["user", "team", "key", "org", "project"] = "user",
|
||||
fallback_depth: int = 0,
|
||||
) -> Literal[True]:
|
||||
"""
|
||||
|
|
@ -2474,6 +2551,24 @@ async def can_team_access_model(
|
|||
raise
|
||||
|
||||
|
||||
def can_project_access_model(
|
||||
model: Union[str, List[str]],
|
||||
project_object: LiteLLM_ProjectTableCachedObj,
|
||||
llm_router: Optional[Router],
|
||||
) -> Literal[True]:
|
||||
"""
|
||||
Returns True if the project can access a specific model.
|
||||
|
||||
Raises ProxyException if access is denied.
|
||||
"""
|
||||
return _can_object_call_model(
|
||||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=project_object.models if project_object else [],
|
||||
object_type="project",
|
||||
)
|
||||
|
||||
|
||||
async def can_user_call_model(
|
||||
model: Union[str, List[str]],
|
||||
llm_router: Optional[Router],
|
||||
|
|
@ -2774,14 +2869,26 @@ async def _team_soft_budget_check(
|
|||
if valid_token:
|
||||
# Extract alert emails from team metadata
|
||||
alert_emails: Optional[List[str]] = None
|
||||
if team_object.metadata is not None and isinstance(team_object.metadata, dict):
|
||||
soft_budget_alert_emails = team_object.metadata.get("soft_budget_alerting_emails")
|
||||
if team_object.metadata is not None and isinstance(
|
||||
team_object.metadata, dict
|
||||
):
|
||||
soft_budget_alert_emails = team_object.metadata.get(
|
||||
"soft_budget_alerting_emails"
|
||||
)
|
||||
if soft_budget_alert_emails is not None:
|
||||
if isinstance(soft_budget_alert_emails, list):
|
||||
alert_emails = [email for email in soft_budget_alert_emails if isinstance(email, str) and email.strip()]
|
||||
alert_emails = [
|
||||
email
|
||||
for email in soft_budget_alert_emails
|
||||
if isinstance(email, str) and email.strip()
|
||||
]
|
||||
elif isinstance(soft_budget_alert_emails, str):
|
||||
# Handle comma-separated string
|
||||
alert_emails = [email.strip() for email in soft_budget_alert_emails.split(",") if email.strip()]
|
||||
alert_emails = [
|
||||
email.strip()
|
||||
for email in soft_budget_alert_emails.split(",")
|
||||
if email.strip()
|
||||
]
|
||||
# Filter out empty strings
|
||||
if alert_emails:
|
||||
alert_emails = [email for email in alert_emails if email]
|
||||
|
|
@ -2820,6 +2927,150 @@ async def _team_soft_budget_check(
|
|||
)
|
||||
|
||||
|
||||
async def _project_max_budget_check(
|
||||
project_object: Optional[LiteLLM_ProjectTableCachedObj],
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
):
|
||||
"""
|
||||
Check if the project is over its max budget.
|
||||
|
||||
Raises:
|
||||
BudgetExceededError if the project is over its max budget.
|
||||
Triggers a budget alert if the project is over its max budget.
|
||||
"""
|
||||
if project_object is None:
|
||||
return
|
||||
|
||||
max_budget = None
|
||||
if project_object.litellm_budget_table is not None:
|
||||
max_budget = project_object.litellm_budget_table.max_budget
|
||||
|
||||
if (
|
||||
max_budget is not None
|
||||
and project_object.spend is not None
|
||||
and project_object.spend > max_budget
|
||||
):
|
||||
if valid_token:
|
||||
call_info = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=project_object.spend,
|
||||
max_budget=max_budget,
|
||||
user_id=valid_token.user_id,
|
||||
team_id=valid_token.team_id,
|
||||
team_alias=valid_token.team_alias,
|
||||
organization_id=valid_token.org_id,
|
||||
event_group=Litellm_EntityType.PROJECT,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="project_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=project_object.spend,
|
||||
max_budget=max_budget,
|
||||
message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_object.spend}, Max budget: {max_budget}",
|
||||
)
|
||||
|
||||
|
||||
async def _project_soft_budget_check(
|
||||
project_object: Optional[LiteLLM_ProjectTableCachedObj],
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
):
|
||||
"""
|
||||
Triggers a budget alert if the project is over its soft budget.
|
||||
|
||||
Mirrors _team_soft_budget_check() pattern.
|
||||
"""
|
||||
if project_object is None:
|
||||
return
|
||||
|
||||
soft_budget = None
|
||||
if project_object.litellm_budget_table is not None:
|
||||
soft_budget = project_object.litellm_budget_table.soft_budget
|
||||
|
||||
if (
|
||||
soft_budget is not None
|
||||
and project_object.spend is not None
|
||||
and project_object.spend >= soft_budget
|
||||
):
|
||||
verbose_proxy_logger.debug(
|
||||
"Crossed Soft Budget for project %s, spend %s, soft_budget %s",
|
||||
project_object.project_id,
|
||||
project_object.spend,
|
||||
soft_budget,
|
||||
)
|
||||
if valid_token:
|
||||
call_info = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=project_object.spend,
|
||||
max_budget=None,
|
||||
soft_budget=soft_budget,
|
||||
user_id=valid_token.user_id,
|
||||
team_id=valid_token.team_id,
|
||||
team_alias=valid_token.team_alias,
|
||||
organization_id=valid_token.org_id,
|
||||
event_group=Litellm_EntityType.PROJECT,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="soft_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def get_project_object(
|
||||
project_id: str,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
user_api_key_cache: DualCache,
|
||||
proxy_logging_obj: Optional[ProxyLogging] = None,
|
||||
) -> Optional[LiteLLM_ProjectTableCachedObj]:
|
||||
"""
|
||||
Fetch project object from cache or DB.
|
||||
|
||||
Follows get_team_object() caching pattern with TTL and last_refreshed_at.
|
||||
|
||||
Returns LiteLLM_ProjectTableCachedObj or None if not found.
|
||||
"""
|
||||
if prisma_client is None:
|
||||
return None
|
||||
|
||||
# Check cache first
|
||||
cache_key = "project_id:{}".format(project_id)
|
||||
cached_obj = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached_obj is not None:
|
||||
if isinstance(cached_obj, dict):
|
||||
return LiteLLM_ProjectTableCachedObj(**cached_obj)
|
||||
elif isinstance(cached_obj, LiteLLM_ProjectTableCachedObj):
|
||||
return cached_obj
|
||||
|
||||
# Fetch from DB
|
||||
project_row = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": project_id},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
if project_row is None:
|
||||
return None
|
||||
|
||||
project_obj = LiteLLM_ProjectTableCachedObj(**project_row.model_dump())
|
||||
|
||||
# Cache with TTL following _cache_management_object pattern
|
||||
project_obj.last_refreshed_at = time.time()
|
||||
await _cache_management_object(
|
||||
key=cache_key,
|
||||
value=project_obj,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return project_obj
|
||||
|
||||
|
||||
async def _organization_max_budget_check(
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
team_object: Optional[LiteLLM_TeamTable],
|
||||
|
|
@ -2921,8 +3172,7 @@ async def _tag_max_budget_check(
|
|||
BudgetExceededError if any tag is over its max budget.
|
||||
Triggers a budget alert if any tag is over its max budget.
|
||||
"""
|
||||
from litellm.proxy.common_utils.http_parsing_utils import \
|
||||
get_tags_from_request_body
|
||||
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ class LicenseCheck:
|
|||
self.license_str = os.getenv("LITELLM_LICENSE", None)
|
||||
verbose_proxy_logger.debug("License Str value - {}".format(self.license_str))
|
||||
self.http_handler = HTTPHandler(timeout=NON_LLM_CONNECTION_TIMEOUT)
|
||||
self._premium_check_logged = False
|
||||
self.public_key = None
|
||||
self.read_public_key()
|
||||
self.airgapped_license_data: Optional["EnterpriseLicenseData"] = None
|
||||
|
|
@ -99,20 +100,23 @@ class LicenseCheck:
|
|||
2. _verify: checks if license is valid calling litellm API. This is the old way we were generating/validating license
|
||||
"""
|
||||
try:
|
||||
verbose_proxy_logger.debug(
|
||||
"litellm.proxy.auth.litellm_license.py::is_premium() - ENTERING 'IS_PREMIUM' - LiteLLM License={}".format(
|
||||
self.license_str
|
||||
if not self._premium_check_logged:
|
||||
verbose_proxy_logger.debug(
|
||||
"litellm.proxy.auth.litellm_license.py::is_premium() - ENTERING 'IS_PREMIUM' - LiteLLM License={}".format(
|
||||
self.license_str
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
if self.license_str is None:
|
||||
self.license_str = os.getenv("LITELLM_LICENSE", None)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"litellm.proxy.auth.litellm_license.py::is_premium() - Updated 'self.license_str' - {}".format(
|
||||
self.license_str
|
||||
if not self._premium_check_logged:
|
||||
verbose_proxy_logger.debug(
|
||||
"litellm.proxy.auth.litellm_license.py::is_premium() - Updated 'self.license_str' - {}".format(
|
||||
self.license_str
|
||||
)
|
||||
)
|
||||
)
|
||||
self._premium_check_logged = True
|
||||
|
||||
if self.license_str is None:
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
common_checks,
|
||||
get_end_user_object,
|
||||
get_key_object,
|
||||
get_project_object,
|
||||
get_team_object,
|
||||
get_user_object,
|
||||
is_valid_fallback_model,
|
||||
|
|
@ -120,12 +121,12 @@ def _get_bearer_token_or_received_api_key(api_key: str) -> str:
|
|||
# Handle AWS Signature V4 format from LangChain
|
||||
# Format: AWS4-HMAC-SHA256 Credential=Bearer sk-12345/date/region/service/aws4_request, SignedHeaders=..., Signature=...
|
||||
# Extract the Bearer token from the Credential field
|
||||
match = re.search(r'Credential=Bearer\s+([^/\s,]+)', api_key)
|
||||
match = re.search(r"Credential=Bearer\s+([^/\s,]+)", api_key)
|
||||
if match:
|
||||
api_key = match.group(1)
|
||||
else:
|
||||
# If no Bearer token found in Credential, try to extract just the credential value
|
||||
match = re.search(r'Credential=([^/\s,]+)', api_key)
|
||||
match = re.search(r"Credential=([^/\s,]+)", api_key)
|
||||
if match:
|
||||
api_key = match.group(1)
|
||||
|
||||
|
|
@ -145,12 +146,12 @@ def _get_bearer_token(
|
|||
# Handle AWS Signature V4 format from LangChain
|
||||
# Format: AWS4-HMAC-SHA256 Credential=Bearer sk-12345/date/region/service/aws4_request, SignedHeaders=..., Signature=...
|
||||
# Extract the Bearer token from the Credential field
|
||||
match = re.search(r'Credential=Bearer\s+([^/\s,]+)', api_key)
|
||||
match = re.search(r"Credential=Bearer\s+([^/\s,]+)", api_key)
|
||||
if match:
|
||||
api_key = match.group(1)
|
||||
else:
|
||||
# If no Bearer token found in Credential, try to extract just the credential value
|
||||
match = re.search(r'Credential=([^/\s,]+)', api_key)
|
||||
match = re.search(r"Credential=([^/\s,]+)", api_key)
|
||||
if match:
|
||||
api_key = match.group(1)
|
||||
else:
|
||||
|
|
@ -274,7 +275,9 @@ async def get_global_proxy_spend(
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
) -> Optional[float]:
|
||||
global_proxy_spend = None
|
||||
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
|
||||
if (
|
||||
litellm.max_budget > 0 and prisma_client is not None
|
||||
): # user set proxy max budget
|
||||
# Use event-driven coordination to prevent cache stampede
|
||||
cache_key = "{}:spend".format(litellm_proxy_admin_name)
|
||||
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
|
||||
|
|
@ -650,7 +653,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
skip_budget_checks = False
|
||||
if model is not None and llm_router is not None:
|
||||
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
|
||||
|
||||
|
||||
skip_budget_checks = _is_model_cost_zero(
|
||||
model=model, llm_router=llm_router
|
||||
)
|
||||
|
|
@ -658,7 +661,17 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
verbose_proxy_logger.info(
|
||||
f"Skipping all budget checks for zero-cost model: {model}"
|
||||
)
|
||||
|
||||
|
||||
# Fetch project object for JWT path if project_id is set
|
||||
_jwt_project_obj = None
|
||||
if valid_token.project_id is not None:
|
||||
_jwt_project_obj = await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# run through common checks
|
||||
_ = await common_checks(
|
||||
request=request,
|
||||
|
|
@ -673,6 +686,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
project_object=_jwt_project_obj,
|
||||
)
|
||||
|
||||
# return UserAPIKeyAuth object
|
||||
|
|
@ -1072,7 +1086,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
skip_budget_checks = False
|
||||
if model is not None and llm_router is not None:
|
||||
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
|
||||
|
||||
|
||||
skip_budget_checks = _is_model_cost_zero(
|
||||
model=model, llm_router=llm_router
|
||||
)
|
||||
|
|
@ -1217,6 +1231,16 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
key=valid_token.team_id, value=_team_obj
|
||||
) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py
|
||||
|
||||
# Fetch project object if key belongs to a project
|
||||
_project_obj = None
|
||||
if valid_token.project_id is not None:
|
||||
_project_obj = await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
global_proxy_spend = None
|
||||
if (
|
||||
litellm.max_budget > 0 and prisma_client is not None
|
||||
|
|
@ -1256,6 +1280,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
project_object=_project_obj,
|
||||
)
|
||||
# Token passed all checks
|
||||
if valid_token is None:
|
||||
|
|
|
|||
|
|
@ -619,6 +619,23 @@ class ProxyBaseLLMRequestProcessing:
|
|||
self.data["litellm_call_id"] = request.headers.get(
|
||||
"x-litellm-call-id", str(uuid.uuid4())
|
||||
)
|
||||
|
||||
### AUTO STREAM USAGE TRACKING ###
|
||||
# If always_include_stream_usage is enabled and this is a streaming request
|
||||
# automatically add stream_options={'include_usage': True} if not already set
|
||||
if (
|
||||
general_settings.get("always_include_stream_usage", False) is True
|
||||
and self.data.get("stream", False) is True
|
||||
):
|
||||
# Only set if stream_options is not already provided by the client
|
||||
if "stream_options" not in self.data:
|
||||
self.data["stream_options"] = {"include_usage": True}
|
||||
elif (
|
||||
isinstance(self.data["stream_options"], dict)
|
||||
and "include_usage" not in self.data["stream_options"]
|
||||
):
|
||||
self.data["stream_options"]["include_usage"] = True
|
||||
|
||||
### CALL HOOKS ### - modify/reject incoming data before calling the model
|
||||
|
||||
## LOGGING OBJECT ## - initialize logging object for logging success/failure events for call
|
||||
|
|
|
|||
|
|
@ -11,7 +11,11 @@ identifier_words:
|
|||
- "airline"
|
||||
- "carrier"
|
||||
- "airways"
|
||||
- "flight"
|
||||
- "flights"
|
||||
- "Emirates"
|
||||
- "Etihad"
|
||||
- "Qatar Airways"
|
||||
|
||||
additional_block_words:
|
||||
# Fake incident reports
|
||||
|
|
@ -57,6 +61,25 @@ additional_block_words:
|
|||
- "unsafe"
|
||||
- "dangerous"
|
||||
- "negligence"
|
||||
# Competitor recommendations
|
||||
- "consider etihad"
|
||||
- "try qatar"
|
||||
- "switch to"
|
||||
- "better value"
|
||||
- "better product"
|
||||
- "best product"
|
||||
- "caught up"
|
||||
- "competitors have"
|
||||
- "recommend qatar"
|
||||
- "recommend etihad"
|
||||
- "consider switching"
|
||||
# Ops data leakage
|
||||
- "crew roster"
|
||||
- "crew schedule"
|
||||
- "crew list"
|
||||
- "pilot names"
|
||||
- "gate assignment"
|
||||
- "cabin crew"
|
||||
|
||||
always_block_keywords:
|
||||
- keyword: "plane crash"
|
||||
|
|
@ -73,13 +96,20 @@ always_block_keywords:
|
|||
severity: "high"
|
||||
- keyword: "fleet grounded"
|
||||
severity: "high"
|
||||
- keyword: "crew roster"
|
||||
severity: "high"
|
||||
- keyword: "crew schedule"
|
||||
severity: "high"
|
||||
- keyword: "pilot names"
|
||||
severity: "high"
|
||||
- keyword: "gate assignment"
|
||||
severity: "high"
|
||||
|
||||
exceptions:
|
||||
- "how safe is"
|
||||
- "safety record"
|
||||
- "safety rating"
|
||||
- "what is"
|
||||
- "explain"
|
||||
- "what is the baggage"
|
||||
- "historical"
|
||||
- "aviation history"
|
||||
- "customer review"
|
||||
|
|
|
|||
|
|
@ -51,6 +51,11 @@ additional_block_words:
|
|||
- "break into cockpit"
|
||||
- "cockpit intrusion"
|
||||
- "cockpit breach"
|
||||
- "locking mechanism"
|
||||
- "open the cockpit"
|
||||
- "cockpit door open"
|
||||
- "enter the cockpit"
|
||||
- "access the cockpit"
|
||||
# Aircraft vulnerabilities
|
||||
- "vulnerability"
|
||||
- "exploit"
|
||||
|
|
@ -115,6 +120,3 @@ exceptions:
|
|||
- "how to report"
|
||||
- "incident report"
|
||||
- "safety culture"
|
||||
- "what is"
|
||||
- "explain"
|
||||
- "research on"
|
||||
|
|
|
|||
|
|
@ -1,287 +1,295 @@
|
|||
import asyncio
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from typing import Any, List, Optional, Union, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
get_litellm_metadata_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import log_db_metrics
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.utils import ProxyUpdateSpend
|
||||
from litellm.types.utils import (
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingUserAPIKeyMetadata,
|
||||
)
|
||||
from litellm.utils import get_end_user_id_for_cost_tracking
|
||||
|
||||
|
||||
class _ProxyDBLogger(CustomLogger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._PROXY_track_cost_callback(
|
||||
kwargs, response_obj, start_time, end_time
|
||||
)
|
||||
|
||||
async def async_post_call_failure_hook(
|
||||
self,
|
||||
request_data: dict,
|
||||
original_exception: Exception,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
traceback_str: Optional[str] = None,
|
||||
):
|
||||
request_route = user_api_key_dict.request_route
|
||||
if _ProxyDBLogger._should_track_errors_in_db() is False:
|
||||
return
|
||||
elif request_route is not None and not RouteChecks.is_llm_api_route(
|
||||
route=request_route
|
||||
):
|
||||
return
|
||||
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
_metadata = dict(
|
||||
StandardLoggingUserAPIKeyMetadata(
|
||||
user_api_key_hash=user_api_key_dict.api_key,
|
||||
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_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,
|
||||
user_api_key_org_id=user_api_key_dict.org_id,
|
||||
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,
|
||||
)
|
||||
)
|
||||
|
||||
existing_metadata: dict = request_data.get("metadata", None) or {}
|
||||
existing_metadata.update(_metadata)
|
||||
|
||||
if "litellm_params" not in request_data:
|
||||
request_data["litellm_params"] = {}
|
||||
|
||||
existing_litellm_params = request_data.get("litellm_params", {})
|
||||
existing_litellm_metadata = existing_litellm_params.get("metadata", {}) or {}
|
||||
|
||||
# Preserve tags from existing metadata
|
||||
if existing_litellm_metadata.get("tags"):
|
||||
existing_metadata["tags"] = existing_litellm_metadata.get("tags")
|
||||
|
||||
request_data["litellm_params"]["proxy_server_request"] = (
|
||||
request_data.get("proxy_server_request") or existing_litellm_params.get("proxy_server_request") or {}
|
||||
)
|
||||
request_data["litellm_params"]["metadata"] = existing_metadata
|
||||
|
||||
# Preserve model name and custom_llm_provider
|
||||
if "model" not in request_data:
|
||||
request_data["model"] = existing_litellm_params.get("model") or request_data.get("model", "")
|
||||
if "custom_llm_provider" not in request_data:
|
||||
request_data["custom_llm_provider"] = existing_litellm_params.get("custom_llm_provider") or request_data.get("custom_llm_provider", "")
|
||||
|
||||
await proxy_logging_obj.db_spend_update_writer.update_database(
|
||||
token=user_api_key_dict.api_key,
|
||||
response_cost=0.0,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
end_user_id=user_api_key_dict.end_user_id,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
kwargs=request_data,
|
||||
completion_response=original_exception,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
org_id=user_api_key_dict.org_id,
|
||||
)
|
||||
|
||||
@log_db_metrics
|
||||
async def _PROXY_track_cost_callback(
|
||||
self,
|
||||
kwargs, # kwargs to completion
|
||||
completion_response: Optional[
|
||||
Union[litellm.ModelResponse, Any]
|
||||
], # response from completion
|
||||
start_time=None,
|
||||
end_time=None, # start/end time for completion
|
||||
):
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj, update_cache
|
||||
|
||||
verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback")
|
||||
try:
|
||||
verbose_proxy_logger.debug(
|
||||
f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}"
|
||||
)
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
|
||||
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
|
||||
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
|
||||
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
|
||||
org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
|
||||
key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None))
|
||||
end_user_max_budget = metadata.get("user_api_end_user_max_budget", None)
|
||||
sl_object: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object", None
|
||||
)
|
||||
response_cost = (
|
||||
sl_object.get("response_cost", None)
|
||||
if sl_object is not None
|
||||
else kwargs.get("response_cost", None)
|
||||
)
|
||||
tags: Optional[List[str]] = (
|
||||
sl_object.get("request_tags", None) if sl_object is not None else None
|
||||
)
|
||||
|
||||
if response_cost is not None:
|
||||
user_api_key = metadata.get("user_api_key", None)
|
||||
if kwargs.get("cache_hit", False) is True:
|
||||
response_cost = 0.0
|
||||
verbose_proxy_logger.debug(
|
||||
f"Cache Hit: response_cost {response_cost}, for user_id {user_id}"
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}"
|
||||
)
|
||||
if _should_track_cost_callback(
|
||||
user_api_key=user_api_key,
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
end_user_id=end_user_id,
|
||||
):
|
||||
## UPDATE DATABASE
|
||||
await proxy_logging_obj.db_spend_update_writer.update_database(
|
||||
token=user_api_key,
|
||||
response_cost=response_cost,
|
||||
user_id=user_id,
|
||||
end_user_id=end_user_id,
|
||||
team_id=team_id,
|
||||
kwargs=kwargs,
|
||||
completion_response=completion_response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
org_id=org_id,
|
||||
)
|
||||
|
||||
# update cache
|
||||
asyncio.create_task(
|
||||
update_cache(
|
||||
token=user_api_key,
|
||||
user_id=user_id,
|
||||
end_user_id=end_user_id,
|
||||
response_cost=response_cost,
|
||||
team_id=team_id,
|
||||
parent_otel_span=parent_otel_span,
|
||||
tags=tags,
|
||||
)
|
||||
)
|
||||
|
||||
await proxy_logging_obj.slack_alerting_instance.customer_spend_alert(
|
||||
token=user_api_key,
|
||||
key_alias=key_alias,
|
||||
end_user_id=end_user_id,
|
||||
response_cost=response_cost,
|
||||
max_budget=end_user_max_budget,
|
||||
)
|
||||
else:
|
||||
# Non-model call types (health checks, afile_delete) have no model or standard_logging_object.
|
||||
# Use .get() for "stream" to avoid KeyError on health checks.
|
||||
if sl_object is None and not kwargs.get("model"):
|
||||
verbose_proxy_logger.warning(
|
||||
"Cost tracking - skipping, no standard_logging_object and no model for call_type=%s",
|
||||
kwargs.get("call_type", "unknown"),
|
||||
)
|
||||
return
|
||||
if kwargs.get("stream") is not True or (
|
||||
kwargs.get("stream") is True and "complete_streaming_response" in kwargs
|
||||
):
|
||||
if sl_object is not None:
|
||||
cost_tracking_failure_debug_info: Union[dict, str] = (
|
||||
sl_object["response_cost_failure_debug_info"] # type: ignore
|
||||
or "response_cost_failure_debug_info is None in standard_logging_object"
|
||||
)
|
||||
else:
|
||||
cost_tracking_failure_debug_info = (
|
||||
"standard_logging_object not found"
|
||||
)
|
||||
model = kwargs.get("model")
|
||||
raise Exception(
|
||||
f"Cost tracking failed for model={model}.\nDebug info - {cost_tracking_failure_debug_info}\nAdd custom pricing - https://docs.litellm.ai/docs/proxy/custom_pricing"
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = f"Error in tracking cost callback - {str(e)}\n Traceback:{traceback.format_exc()}"
|
||||
model = kwargs.get("model", "")
|
||||
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
|
||||
litellm_metadata = kwargs.get("litellm_params", {}).get(
|
||||
"litellm_metadata", {}
|
||||
)
|
||||
old_metadata = kwargs.get("litellm_params", {}).get("metadata", {})
|
||||
call_type = kwargs.get("call_type", "")
|
||||
error_msg += f"\n Args to _PROXY_track_cost_callback\n model: {model}\n chosen_metadata: {metadata}\n litellm_metadata: {litellm_metadata}\n old_metadata: {old_metadata}\n call_type: {call_type}\n"
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.failed_tracking_alert(
|
||||
error_message=error_msg,
|
||||
failing_model=model,
|
||||
)
|
||||
)
|
||||
|
||||
verbose_proxy_logger.exception(
|
||||
"Error in tracking cost callback - %s", str(e)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _should_track_errors_in_db():
|
||||
"""
|
||||
Returns True if errors should be tracked in the database
|
||||
|
||||
By default, errors are tracked in the database
|
||||
|
||||
If users want to disable error tracking, they can set the disable_error_logs flag in the general_settings
|
||||
"""
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
if general_settings.get("disable_error_logs") is True:
|
||||
return False
|
||||
return
|
||||
|
||||
|
||||
def _should_track_cost_callback(
|
||||
user_api_key: Optional[str],
|
||||
user_id: Optional[str],
|
||||
team_id: Optional[str],
|
||||
end_user_id: Optional[str],
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if the cost callback should be tracked based on the kwargs
|
||||
"""
|
||||
|
||||
# don't run track cost callback if user opted into disabling spend
|
||||
if ProxyUpdateSpend.disable_spend_updates() is True:
|
||||
return False
|
||||
|
||||
if (
|
||||
user_api_key is not None
|
||||
or user_id is not None
|
||||
or team_id is not None
|
||||
or end_user_id is not None
|
||||
):
|
||||
return True
|
||||
return False
|
||||
import asyncio
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from typing import Any, List, Optional, Union, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
get_litellm_metadata_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import log_db_metrics
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.utils import ProxyUpdateSpend
|
||||
from litellm.types.utils import (
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingUserAPIKeyMetadata,
|
||||
)
|
||||
from litellm.utils import get_end_user_id_for_cost_tracking
|
||||
|
||||
|
||||
class _ProxyDBLogger(CustomLogger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._PROXY_track_cost_callback(
|
||||
kwargs, response_obj, start_time, end_time
|
||||
)
|
||||
|
||||
async def async_post_call_failure_hook(
|
||||
self,
|
||||
request_data: dict,
|
||||
original_exception: Exception,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
traceback_str: Optional[str] = None,
|
||||
):
|
||||
request_route = user_api_key_dict.request_route
|
||||
if _ProxyDBLogger._should_track_errors_in_db() is False:
|
||||
return
|
||||
elif request_route is not None and not RouteChecks.is_llm_api_route(
|
||||
route=request_route
|
||||
):
|
||||
return
|
||||
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
_metadata = dict(
|
||||
StandardLoggingUserAPIKeyMetadata(
|
||||
user_api_key_hash=user_api_key_dict.api_key,
|
||||
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_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,
|
||||
user_api_key_org_id=user_api_key_dict.org_id,
|
||||
user_api_key_project_id=user_api_key_dict.project_id,
|
||||
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,
|
||||
)
|
||||
|
||||
existing_metadata: dict = request_data.get("metadata", None) or {}
|
||||
existing_metadata.update(_metadata)
|
||||
|
||||
if "litellm_params" not in request_data:
|
||||
request_data["litellm_params"] = {}
|
||||
|
||||
existing_litellm_params = request_data.get("litellm_params", {})
|
||||
existing_litellm_metadata = existing_litellm_params.get("metadata", {}) or {}
|
||||
|
||||
# Preserve tags from existing metadata
|
||||
if existing_litellm_metadata.get("tags"):
|
||||
existing_metadata["tags"] = existing_litellm_metadata.get("tags")
|
||||
|
||||
request_data["litellm_params"]["proxy_server_request"] = (
|
||||
request_data.get("proxy_server_request")
|
||||
or existing_litellm_params.get("proxy_server_request")
|
||||
or {}
|
||||
)
|
||||
request_data["litellm_params"]["metadata"] = existing_metadata
|
||||
|
||||
# Preserve model name and custom_llm_provider
|
||||
if "model" not in request_data:
|
||||
request_data["model"] = existing_litellm_params.get(
|
||||
"model"
|
||||
) or request_data.get("model", "")
|
||||
if "custom_llm_provider" not in request_data:
|
||||
request_data["custom_llm_provider"] = existing_litellm_params.get(
|
||||
"custom_llm_provider"
|
||||
) or request_data.get("custom_llm_provider", "")
|
||||
|
||||
await proxy_logging_obj.db_spend_update_writer.update_database(
|
||||
token=user_api_key_dict.api_key,
|
||||
response_cost=0.0,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
end_user_id=user_api_key_dict.end_user_id,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
kwargs=request_data,
|
||||
completion_response=original_exception,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
org_id=user_api_key_dict.org_id,
|
||||
)
|
||||
|
||||
@log_db_metrics
|
||||
async def _PROXY_track_cost_callback(
|
||||
self,
|
||||
kwargs, # kwargs to completion
|
||||
completion_response: Optional[
|
||||
Union[litellm.ModelResponse, Any]
|
||||
], # response from completion
|
||||
start_time=None,
|
||||
end_time=None, # start/end time for completion
|
||||
):
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj, update_cache
|
||||
|
||||
verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback")
|
||||
try:
|
||||
verbose_proxy_logger.debug(
|
||||
f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}"
|
||||
)
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
|
||||
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
|
||||
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
|
||||
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
|
||||
org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
|
||||
key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None))
|
||||
end_user_max_budget = metadata.get("user_api_end_user_max_budget", None)
|
||||
sl_object: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object", None
|
||||
)
|
||||
response_cost = (
|
||||
sl_object.get("response_cost", None)
|
||||
if sl_object is not None
|
||||
else kwargs.get("response_cost", None)
|
||||
)
|
||||
tags: Optional[List[str]] = (
|
||||
sl_object.get("request_tags", None) if sl_object is not None else None
|
||||
)
|
||||
|
||||
if response_cost is not None:
|
||||
user_api_key = metadata.get("user_api_key", None)
|
||||
if kwargs.get("cache_hit", False) is True:
|
||||
response_cost = 0.0
|
||||
verbose_proxy_logger.debug(
|
||||
f"Cache Hit: response_cost {response_cost}, for user_id {user_id}"
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}"
|
||||
)
|
||||
if _should_track_cost_callback(
|
||||
user_api_key=user_api_key,
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
end_user_id=end_user_id,
|
||||
):
|
||||
## UPDATE DATABASE
|
||||
await proxy_logging_obj.db_spend_update_writer.update_database(
|
||||
token=user_api_key,
|
||||
response_cost=response_cost,
|
||||
user_id=user_id,
|
||||
end_user_id=end_user_id,
|
||||
team_id=team_id,
|
||||
kwargs=kwargs,
|
||||
completion_response=completion_response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
org_id=org_id,
|
||||
)
|
||||
|
||||
# update cache
|
||||
asyncio.create_task(
|
||||
update_cache(
|
||||
token=user_api_key,
|
||||
user_id=user_id,
|
||||
end_user_id=end_user_id,
|
||||
response_cost=response_cost,
|
||||
team_id=team_id,
|
||||
parent_otel_span=parent_otel_span,
|
||||
tags=tags,
|
||||
)
|
||||
)
|
||||
|
||||
await proxy_logging_obj.slack_alerting_instance.customer_spend_alert(
|
||||
token=user_api_key,
|
||||
key_alias=key_alias,
|
||||
end_user_id=end_user_id,
|
||||
response_cost=response_cost,
|
||||
max_budget=end_user_max_budget,
|
||||
)
|
||||
else:
|
||||
# Non-model call types (health checks, afile_delete) have no model or standard_logging_object.
|
||||
# Use .get() for "stream" to avoid KeyError on health checks.
|
||||
if sl_object is None and not kwargs.get("model"):
|
||||
verbose_proxy_logger.warning(
|
||||
"Cost tracking - skipping, no standard_logging_object and no model for call_type=%s",
|
||||
kwargs.get("call_type", "unknown"),
|
||||
)
|
||||
return
|
||||
if kwargs.get("stream") is not True or (
|
||||
kwargs.get("stream") is True
|
||||
and "complete_streaming_response" in kwargs
|
||||
):
|
||||
if sl_object is not None:
|
||||
cost_tracking_failure_debug_info: Union[dict, str] = (
|
||||
sl_object["response_cost_failure_debug_info"] # type: ignore
|
||||
or "response_cost_failure_debug_info is None in standard_logging_object"
|
||||
)
|
||||
else:
|
||||
cost_tracking_failure_debug_info = (
|
||||
"standard_logging_object not found"
|
||||
)
|
||||
model = kwargs.get("model")
|
||||
raise Exception(
|
||||
f"Cost tracking failed for model={model}.\nDebug info - {cost_tracking_failure_debug_info}\nAdd custom pricing - https://docs.litellm.ai/docs/proxy/custom_pricing"
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = f"Error in tracking cost callback - {str(e)}\n Traceback:{traceback.format_exc()}"
|
||||
model = kwargs.get("model", "")
|
||||
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
|
||||
litellm_metadata = kwargs.get("litellm_params", {}).get(
|
||||
"litellm_metadata", {}
|
||||
)
|
||||
old_metadata = kwargs.get("litellm_params", {}).get("metadata", {})
|
||||
call_type = kwargs.get("call_type", "")
|
||||
error_msg += f"\n Args to _PROXY_track_cost_callback\n model: {model}\n chosen_metadata: {metadata}\n litellm_metadata: {litellm_metadata}\n old_metadata: {old_metadata}\n call_type: {call_type}\n"
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.failed_tracking_alert(
|
||||
error_message=error_msg,
|
||||
failing_model=model,
|
||||
)
|
||||
)
|
||||
|
||||
verbose_proxy_logger.exception(
|
||||
"Error in tracking cost callback - %s", str(e)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _should_track_errors_in_db():
|
||||
"""
|
||||
Returns True if errors should be tracked in the database
|
||||
|
||||
By default, errors are tracked in the database
|
||||
|
||||
If users want to disable error tracking, they can set the disable_error_logs flag in the general_settings
|
||||
"""
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
if general_settings.get("disable_error_logs") is True:
|
||||
return False
|
||||
return
|
||||
|
||||
|
||||
def _should_track_cost_callback(
|
||||
user_api_key: Optional[str],
|
||||
user_id: Optional[str],
|
||||
team_id: Optional[str],
|
||||
end_user_id: Optional[str],
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if the cost callback should be tracked based on the kwargs
|
||||
"""
|
||||
|
||||
# don't run track cost callback if user opted into disabling spend
|
||||
if ProxyUpdateSpend.disable_spend_updates() is True:
|
||||
return False
|
||||
|
||||
if (
|
||||
user_api_key is not None
|
||||
or user_id is not None
|
||||
or team_id is not None
|
||||
or end_user_id is not None
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -591,6 +591,7 @@ class LiteLLMProxyRequestSetup:
|
|||
user_api_key_spend=user_api_key_dict.spend,
|
||||
user_api_key_max_budget=user_api_key_dict.max_budget,
|
||||
user_api_key_team_id=user_api_key_dict.team_id,
|
||||
user_api_key_project_id=user_api_key_dict.project_id,
|
||||
user_api_key_user_id=user_api_key_dict.user_id,
|
||||
user_api_key_org_id=user_api_key_dict.org_id,
|
||||
user_api_key_team_alias=user_api_key_dict.team_alias,
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from litellm.proxy._types import (
|
|||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
LiteLLM_OrganizationTable,
|
||||
LiteLLM_ProjectTable,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_UserTable,
|
||||
LitellmUserRoles,
|
||||
|
|
@ -37,6 +38,25 @@ def _is_user_team_admin(
|
|||
return False
|
||||
|
||||
|
||||
def _team_member_has_permission(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
permission: str,
|
||||
) -> bool:
|
||||
"""Check if a non-admin team member has a specific permission on a team."""
|
||||
if not team_obj.team_member_permissions:
|
||||
return False
|
||||
if permission not in team_obj.team_member_permissions:
|
||||
return False
|
||||
for member in team_obj.members_with_roles:
|
||||
if (
|
||||
member.user_id is not None
|
||||
and member.user_id == user_api_key_dict.user_id
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
async def _user_has_admin_privileges(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: Optional["PrismaClient"] = None,
|
||||
|
|
@ -262,6 +282,7 @@ def _set_object_metadata_field(
|
|||
LiteLLM_TeamTable,
|
||||
KeyRequestBase,
|
||||
LiteLLM_OrganizationTable,
|
||||
LiteLLM_ProjectTable,
|
||||
],
|
||||
field_name: str,
|
||||
value: Any,
|
||||
|
|
@ -270,7 +291,7 @@ def _set_object_metadata_field(
|
|||
Helper function to set metadata fields that require premium user checks
|
||||
|
||||
Args:
|
||||
object_data: The team data object to modify
|
||||
object_data: The team/key/organization/project data object to modify
|
||||
field_name: Name of the metadata field to set
|
||||
value: Value to set for the field
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -573,7 +573,7 @@ def get_user_id_from_request(request: Request) -> Optional[str]:
|
|||
"/user/info",
|
||||
tags=["Internal User management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
# response_model=UserInfoResponse,
|
||||
response_model=UserInfoResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def user_info(
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
can_team_access_model,
|
||||
get_key_object,
|
||||
get_org_object,
|
||||
get_project_object,
|
||||
get_team_object,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import abbreviate_api_key
|
||||
|
|
@ -890,6 +891,61 @@ async def _check_team_key_limits(
|
|||
)
|
||||
|
||||
|
||||
async def _check_project_key_limits(
|
||||
project_id: str,
|
||||
data: Union[GenerateKeyRequest, UpdateKeyRequest],
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: DualCache,
|
||||
) -> None:
|
||||
"""
|
||||
Validate that key's models and budget respect its project's limits.
|
||||
|
||||
- Key models must be a subset of project models
|
||||
- Key max_budget must be <= project max_budget
|
||||
"""
|
||||
project_obj = await get_project_object(
|
||||
project_id=project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
if project_obj is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Project not found, project_id={project_id}"},
|
||||
)
|
||||
|
||||
# Validate key models are a subset of project models
|
||||
if data.models and len(project_obj.models) > 0:
|
||||
for m in data.models:
|
||||
if m not in project_obj.models:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Model '{m}' not in project's allowed models. Project allowed models={project_obj.models}. Project: {project_id}"
|
||||
},
|
||||
)
|
||||
|
||||
# Validate key max_budget <= project max_budget
|
||||
project_max_budget = None
|
||||
if project_obj.litellm_budget_table is not None:
|
||||
project_max_budget = getattr(
|
||||
project_obj.litellm_budget_table, "max_budget", None
|
||||
)
|
||||
|
||||
if (
|
||||
data.max_budget is not None
|
||||
and project_max_budget is not None
|
||||
and data.max_budget > project_max_budget
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Key max_budget ({data.max_budget}) exceeds project's max_budget ({project_max_budget}). Project: {project_id}"
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def check_org_key_model_specific_limits(
|
||||
keys: List[LiteLLM_VerificationToken],
|
||||
org_table: LiteLLM_OrganizationTable,
|
||||
|
|
@ -1145,6 +1201,15 @@ async def generate_key_fn(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
# Validate key against project limits if project_id is set
|
||||
if data.project_id is not None:
|
||||
await _check_project_key_limits(
|
||||
project_id=data.project_id,
|
||||
data=data,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
return await _common_key_generation_helper(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -1820,6 +1885,20 @@ async def update_key_fn(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
# Validate key against project limits if project_id is being set
|
||||
_project_id_to_check = getattr(data, "project_id", None) or getattr(
|
||||
existing_key_row, "project_id", None
|
||||
)
|
||||
if _project_id_to_check is not None and (
|
||||
data.models is not None or data.max_budget is not None
|
||||
):
|
||||
await _check_project_key_limits(
|
||||
project_id=_project_id_to_check,
|
||||
data=data,
|
||||
prisma_client=prisma_client,
|
||||
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):
|
||||
if llm_router is None:
|
||||
|
|
@ -2475,6 +2554,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
prompts: Optional[list] = None,
|
||||
teams: Optional[list] = None,
|
||||
organization_id: Optional[str] = None,
|
||||
project_id: Optional[str] = None,
|
||||
table_name: Optional[Literal["key", "user"]] = None,
|
||||
send_invite_email: Optional[bool] = None,
|
||||
created_by: Optional[str] = None,
|
||||
|
|
@ -2588,6 +2668,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
"max_budget": key_max_budget,
|
||||
"user_id": user_id,
|
||||
"team_id": team_id,
|
||||
"project_id": project_id,
|
||||
"max_parallel_requests": max_parallel_requests,
|
||||
"metadata": metadata_json,
|
||||
"tpm_limit": tpm_limit,
|
||||
|
|
|
|||
|
|
@ -9,14 +9,29 @@ All /policy management endpoints
|
|||
/policy/templates - Get policy templates (GitHub with local fallback)
|
||||
"""
|
||||
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, List, Literal, Optional, TypedDict, cast
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
AsyncIterator,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
TypedDict,
|
||||
cast,
|
||||
)
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
COMPETITOR_LLM_TEMPERATURE,
|
||||
DEFAULT_COMPETITOR_DISCOVERY_MODEL,
|
||||
MAX_COMPETITOR_NAMES,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
|
@ -511,6 +526,48 @@ async def get_policy_templates(
|
|||
class EnrichTemplateRequest(BaseModel):
|
||||
template_id: str
|
||||
parameters: dict
|
||||
model: Optional[str] = None
|
||||
competitors: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
max_length=MAX_COMPETITOR_NAMES,
|
||||
description="Optional list of competitor names",
|
||||
)
|
||||
instruction: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Refinement instruction for modifying the competitor list (e.g. 'add 10 more from Asia')",
|
||||
)
|
||||
|
||||
|
||||
def _validate_enrichment_request(data: EnrichTemplateRequest) -> tuple[dict, dict, str]:
|
||||
"""
|
||||
Validate enrichment request and return (template, llm_enrichment, brand_name).
|
||||
|
||||
Raises HTTPException on validation failure.
|
||||
"""
|
||||
templates = _load_policy_templates_from_local_backup()
|
||||
template = next((t for t in templates if t.get("id") == data.template_id), None)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
|
||||
|
||||
llm_enrichment = template.get("llm_enrichment")
|
||||
if llm_enrichment is None:
|
||||
raise HTTPException(status_code=400, detail="Template does not support LLM enrichment")
|
||||
|
||||
# Validate competitors list size if provided
|
||||
if data.competitors and len(data.competitors) > MAX_COMPETITOR_NAMES:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"competitors list exceeds maximum of {MAX_COMPETITOR_NAMES}",
|
||||
)
|
||||
|
||||
brand_name = data.parameters.get(llm_enrichment["parameter"], "")
|
||||
if not brand_name:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Parameter '{llm_enrichment['parameter']}' is required",
|
||||
)
|
||||
|
||||
return template, llm_enrichment, brand_name
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
@ -530,108 +587,311 @@ async def enrich_policy_template(
|
|||
Calls an onboarded LLM to discover competitors for the given brand name,
|
||||
then returns enriched guardrailDefinitions with the discovered data populated.
|
||||
"""
|
||||
templates = _load_policy_templates_from_local_backup()
|
||||
template = next((t for t in templates if t.get("id") == data.template_id), None)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
|
||||
template, llm_enrichment, brand_name = _validate_enrichment_request(data)
|
||||
model = data.model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
|
||||
|
||||
llm_enrichment = template.get("llm_enrichment")
|
||||
if llm_enrichment is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Template does not support LLM enrichment",
|
||||
if data.competitors:
|
||||
competitors = data.competitors
|
||||
else:
|
||||
prompt = llm_enrichment["prompt"].replace(
|
||||
"{{" + llm_enrichment["parameter"] + "}}", brand_name
|
||||
)
|
||||
competitors = await _discover_competitors_via_llm(prompt, model=model)
|
||||
|
||||
brand_name = data.parameters.get(llm_enrichment["parameter"], "")
|
||||
if not brand_name:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Parameter '{llm_enrichment['parameter']}' is required",
|
||||
)
|
||||
|
||||
prompt = llm_enrichment["prompt"].replace(
|
||||
"{{" + llm_enrichment["parameter"] + "}}", brand_name
|
||||
)
|
||||
|
||||
competitors = await _discover_competitors_via_llm(prompt)
|
||||
variations_map = await _generate_competitor_variations(competitors, model=model)
|
||||
|
||||
enriched_definitions = _build_competitor_guardrail_definitions(
|
||||
template.get("guardrailDefinitions", []),
|
||||
competitors,
|
||||
brand_name,
|
||||
variations_map,
|
||||
)
|
||||
|
||||
return {"guardrailDefinitions": enriched_definitions, "competitors": competitors}
|
||||
return {
|
||||
"guardrailDefinitions": enriched_definitions,
|
||||
"competitors": competitors,
|
||||
"competitor_variations": variations_map,
|
||||
}
|
||||
|
||||
|
||||
async def _discover_competitors_via_llm(prompt: str) -> list:
|
||||
"""Call an onboarded LLM to discover competitor names."""
|
||||
import litellm
|
||||
def _build_refinement_prompt(
|
||||
instruction: str,
|
||||
existing_competitors: list[str],
|
||||
brand_name: str,
|
||||
) -> str:
|
||||
"""Build a prompt for refining the competitor list based on user instruction."""
|
||||
existing_list = ", ".join(existing_competitors)
|
||||
return (
|
||||
f"I have a brand called '{brand_name}' and the following competitor list:\n"
|
||||
f"{existing_list}\n\n"
|
||||
f"User instruction: {instruction}\n\n"
|
||||
"Return ONLY the NEW names to add (not the existing ones), one per line, "
|
||||
"no numbering, no explanations. If the instruction asks to remove names, "
|
||||
"return nothing."
|
||||
)
|
||||
|
||||
|
||||
async def _stream_llm_competitor_names(
|
||||
prompt: str,
|
||||
model: str,
|
||||
existing: list[str],
|
||||
) -> AsyncIterator[tuple[Optional[str], bool]]:
|
||||
"""
|
||||
Stream competitor names from LLM. Yields (name, is_error) tuples.
|
||||
|
||||
Deduplicates against existing names (case-insensitive).
|
||||
"""
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is None:
|
||||
raise ValueError("LLM router not initialized")
|
||||
|
||||
existing_lower = {n.lower() for n in existing}
|
||||
response = await llm_router.acompletion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
temperature=COMPETITOR_LLM_TEMPERATURE,
|
||||
stream=True,
|
||||
)
|
||||
buffer = ""
|
||||
count = len(existing)
|
||||
async for chunk in response: # type: ignore[union-attr]
|
||||
delta = chunk.choices[0].delta.content or ""
|
||||
buffer += delta
|
||||
while "\n" in buffer:
|
||||
line, buffer = buffer.split("\n", 1)
|
||||
name = _clean_competitor_line(line)
|
||||
if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES:
|
||||
existing_lower.add(name.lower())
|
||||
count += 1
|
||||
yield name, False
|
||||
# Handle remaining buffer
|
||||
name = _clean_competitor_line(buffer)
|
||||
if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES:
|
||||
yield name, False
|
||||
|
||||
|
||||
async def _stream_competitor_events(
|
||||
data: EnrichTemplateRequest,
|
||||
template: dict,
|
||||
llm_enrichment: dict,
|
||||
brand_name: str,
|
||||
model: str,
|
||||
) -> AsyncIterator[str]:
|
||||
"""Stream competitor names as SSE events, then emit a final 'done' event."""
|
||||
competitors: list[str] = list(data.competitors or [])
|
||||
|
||||
if data.instruction and competitors:
|
||||
# Refinement mode: keep existing, stream only new names
|
||||
for comp in competitors:
|
||||
yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n"
|
||||
|
||||
refinement_prompt = _build_refinement_prompt(
|
||||
data.instruction, competitors, brand_name
|
||||
)
|
||||
try:
|
||||
async for name, _ in _stream_llm_competitor_names(
|
||||
refinement_prompt, model, competitors
|
||||
):
|
||||
if name:
|
||||
competitors.append(name)
|
||||
yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("LLM competitor refinement failed: %s", e)
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
|
||||
return
|
||||
elif data.competitors and not data.instruction:
|
||||
# Free-form mode (no instruction): just emit existing
|
||||
for comp in competitors:
|
||||
yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n"
|
||||
else:
|
||||
# Initial discovery mode
|
||||
prompt = llm_enrichment["prompt"].replace(
|
||||
"{{" + llm_enrichment["parameter"] + "}}", brand_name
|
||||
)
|
||||
try:
|
||||
async for name, _ in _stream_llm_competitor_names(
|
||||
prompt, model, []
|
||||
):
|
||||
if name:
|
||||
competitors.append(name)
|
||||
yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("LLM competitor streaming failed: %s", e)
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
|
||||
return
|
||||
|
||||
yield f"data: {json.dumps({'type': 'status', 'message': f'Generating alternate spellings for {len(competitors)} competitors...'})}\n\n"
|
||||
variations_map = await _generate_competitor_variations(competitors, model=model)
|
||||
|
||||
total_variations = sum(len(v) for v in variations_map.values())
|
||||
yield f"data: {json.dumps({'type': 'status', 'message': f'Building guardrail definitions with {total_variations} variations...'})}\n\n"
|
||||
enriched_definitions = _build_competitor_guardrail_definitions(
|
||||
template.get("guardrailDefinitions", []),
|
||||
competitors,
|
||||
brand_name,
|
||||
variations_map,
|
||||
)
|
||||
|
||||
yield f"data: {json.dumps({'type': 'done', 'competitors': competitors, 'competitor_variations': variations_map, 'guardrailDefinitions': enriched_definitions})}\n\n"
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policy/templates/enrich/stream",
|
||||
tags=["policy management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def enrich_policy_template_stream(
|
||||
data: EnrichTemplateRequest,
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Stream competitor names as SSE events as the LLM generates them.
|
||||
|
||||
Events:
|
||||
- data: {"type": "competitor", "name": "..."} — each competitor as discovered
|
||||
- data: {"type": "done", "competitors": [...], "competitor_variations": {...}, "guardrailDefinitions": [...]}
|
||||
"""
|
||||
template, llm_enrichment, brand_name = _validate_enrichment_request(data)
|
||||
model = data.model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
|
||||
|
||||
return StreamingResponse(
|
||||
_stream_competitor_events(data, template, llm_enrichment, brand_name, model),
|
||||
media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||||
)
|
||||
|
||||
|
||||
def _clean_competitor_line(line: str) -> Optional[str]:
|
||||
"""Strip numbering, bullets, and whitespace from a competitor name line."""
|
||||
name = line.strip().strip(".-) ").strip()
|
||||
return name if name and len(name) > 1 else None
|
||||
|
||||
|
||||
async def _generate_competitor_variations(
|
||||
competitors: list, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL
|
||||
) -> dict:
|
||||
"""Generate common misspellings, abbreviations, and alternate names for each competitor."""
|
||||
if not competitors:
|
||||
return {}
|
||||
|
||||
# Cap the list to prevent oversized prompts
|
||||
capped = competitors[:MAX_COMPETITOR_NAMES]
|
||||
names_list = "\n".join(capped)
|
||||
prompt = (
|
||||
"For each company/brand name below, list 3-5 common misspellings, abbreviations, "
|
||||
"and alternate names that people might type. Include typos, missing spaces, "
|
||||
"wrong suffixes (e.g. 'Airlines' vs 'Airways' vs 'Airline'), and common shortcuts.\n\n"
|
||||
f"Names:\n{names_list}\n\n"
|
||||
"Return the result as one line per variation in the format:\n"
|
||||
"OriginalName: variation1, variation2, variation3\n"
|
||||
"Use the EXACT original name before the colon. No numbering, no extra text."
|
||||
)
|
||||
|
||||
try:
|
||||
response = await litellm.acompletion(
|
||||
model="gpt-4o-mini",
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is None:
|
||||
raise ValueError("LLM router not initialized")
|
||||
response = await llm_router.acompletion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
temperature=0.3,
|
||||
temperature=COMPETITOR_LLM_TEMPERATURE,
|
||||
)
|
||||
raw = response.choices[0].message.content or "" # type: ignore
|
||||
return _parse_variations_response(raw, capped)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("LLM competitor variation generation failed: %s", e)
|
||||
return {}
|
||||
|
||||
|
||||
def _parse_variations_response(raw: str, competitors: list) -> dict[str, list[str]]:
|
||||
"""Parse the LLM response for competitor variations into a name -> variations map."""
|
||||
# Build a lowercase lookup for case-insensitive matching
|
||||
lower_to_canonical = {comp.lower(): comp for comp in competitors}
|
||||
variations_map: dict[str, list[str]] = {}
|
||||
|
||||
for line in raw.strip().split("\n"):
|
||||
if ":" not in line:
|
||||
continue
|
||||
name, _, variations_str = line.partition(":")
|
||||
canonical = lower_to_canonical.get(name.strip().lower())
|
||||
if canonical is None:
|
||||
continue
|
||||
variations = [
|
||||
v.strip()
|
||||
for v in variations_str.split(",")
|
||||
if v.strip() and v.strip().lower() != canonical.lower()
|
||||
]
|
||||
variations_map[canonical] = variations
|
||||
|
||||
return variations_map
|
||||
|
||||
|
||||
async def _discover_competitors_via_llm(
|
||||
prompt: str, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL
|
||||
) -> list:
|
||||
"""Call an onboarded LLM to discover competitor names."""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is None:
|
||||
raise ValueError("LLM router not initialized")
|
||||
response = await llm_router.acompletion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
temperature=COMPETITOR_LLM_TEMPERATURE,
|
||||
)
|
||||
raw = response.choices[0].message.content or "" # type: ignore
|
||||
competitors = [
|
||||
line.strip().strip(".-) ").strip()
|
||||
name
|
||||
for line in raw.strip().split("\n")
|
||||
if line.strip() and len(line.strip()) > 1
|
||||
if (name := _clean_competitor_line(line)) is not None
|
||||
]
|
||||
return competitors[:15]
|
||||
return competitors[:MAX_COMPETITOR_NAMES]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("LLM competitor discovery failed: %s", e)
|
||||
return []
|
||||
|
||||
|
||||
def _build_all_names_per_competitor(
|
||||
competitors: list[str], variations_map: dict[str, list[str]]
|
||||
) -> dict[str, list[str]]:
|
||||
"""Build canonical + variation name lists for each competitor."""
|
||||
return {
|
||||
comp: [comp] + variations_map.get(comp, [])
|
||||
for comp in competitors
|
||||
}
|
||||
|
||||
|
||||
def _build_competitor_guardrail_definitions(
|
||||
definitions: list,
|
||||
competitors: list,
|
||||
brand_name: str,
|
||||
variations_map: Optional[dict] = None,
|
||||
) -> list:
|
||||
"""Build enriched guardrailDefinitions with competitor names populated."""
|
||||
import copy
|
||||
|
||||
"""Build enriched guardrailDefinitions with competitor names and variations populated."""
|
||||
variations_map = variations_map or {}
|
||||
enriched = copy.deepcopy(definitions)
|
||||
all_names = _build_all_names_per_competitor(competitors, variations_map)
|
||||
|
||||
output_blocked = [
|
||||
{"keyword": comp, "action": "BLOCK", "description": f"Competitor: {comp}"}
|
||||
for comp in competitors
|
||||
]
|
||||
|
||||
recommendation_blocked = []
|
||||
for comp in competitors:
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"try {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"use {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"switch to {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
recommendation_blocked.append(
|
||||
{"keyword": f"consider {comp}", "action": "BLOCK", "description": "Recommendation to competitor"}
|
||||
)
|
||||
|
||||
comparison_blocked = []
|
||||
for comp in competitors:
|
||||
comparison_blocked.append(
|
||||
{"keyword": f"{comp} is better", "action": "BLOCK", "description": "Unfavorable comparison"}
|
||||
)
|
||||
comparison_blocked.append(
|
||||
{"keyword": f"better than {brand_name}", "action": "BLOCK", "description": "Unfavorable comparison"}
|
||||
)
|
||||
comparison_blocked.append(
|
||||
{"keyword": f"{brand_name} is worse", "action": "BLOCK", "description": "Unfavorable comparison"}
|
||||
)
|
||||
output_blocked = _build_name_blocked_words(competitors, all_names)
|
||||
recommendation_blocked = _build_recommendation_blocked_words(competitors, all_names)
|
||||
comparison_blocked = _build_comparison_blocked_words(competitors, all_names, brand_name)
|
||||
|
||||
blocked_words_map = {
|
||||
"competitor-output-blocker": output_blocked,
|
||||
"competitor-input-blocker": output_blocked,
|
||||
"competitor-name-blocker": output_blocked,
|
||||
"competitor-name-input-blocker": output_blocked,
|
||||
"competitor-name-output-blocker": output_blocked,
|
||||
"competitor-recommendation-filter": recommendation_blocked,
|
||||
"competitor-recommendation-input-filter": recommendation_blocked,
|
||||
"competitor-recommendation-output-filter": recommendation_blocked,
|
||||
"competitor-comparison-filter": comparison_blocked,
|
||||
"competitor-comparison-input-filter": comparison_blocked,
|
||||
"competitor-comparison-output-filter": comparison_blocked,
|
||||
}
|
||||
|
||||
for defn in enriched:
|
||||
|
|
@ -640,3 +900,59 @@ def _build_competitor_guardrail_definitions(
|
|||
defn["litellm_params"]["blocked_words"] = blocked_words_map[guardrail_name]
|
||||
|
||||
return enriched
|
||||
|
||||
|
||||
def _build_name_blocked_words(
|
||||
competitors: list[str], all_names: dict[str, list[str]]
|
||||
) -> list[dict]:
|
||||
"""Build blocked word entries for direct competitor name mentions."""
|
||||
result = []
|
||||
for comp in competitors:
|
||||
for name in all_names[comp]:
|
||||
desc = f"Competitor: {comp}" if name == comp else f"Competitor variation ({comp}): {name}"
|
||||
result.append({"keyword": name, "action": "BLOCK", "description": desc})
|
||||
return result
|
||||
|
||||
|
||||
def _build_recommendation_blocked_words(
|
||||
competitors: list[str], all_names: dict[str, list[str]]
|
||||
) -> list[dict]:
|
||||
"""Build blocked word entries for competitor recommendations."""
|
||||
result = []
|
||||
for comp in competitors:
|
||||
for name in all_names[comp]:
|
||||
for prefix in ["try", "use", "switch to", "consider"]:
|
||||
result.append({
|
||||
"keyword": f"{prefix} {name}",
|
||||
"action": "BLOCK",
|
||||
"description": f"Recommendation to competitor ({comp})",
|
||||
})
|
||||
return result
|
||||
|
||||
|
||||
def _build_comparison_blocked_words(
|
||||
competitors: list[str], all_names: dict[str, list[str]], brand_name: str
|
||||
) -> list[dict]:
|
||||
"""Build blocked word entries for unfavorable competitor comparisons."""
|
||||
result = []
|
||||
for comp in competitors:
|
||||
for name in all_names[comp]:
|
||||
result.append({
|
||||
"keyword": f"{name} is better",
|
||||
"action": "BLOCK",
|
||||
"description": f"Unfavorable comparison ({comp})",
|
||||
})
|
||||
|
||||
# Brand-level comparisons (only need one entry each, not per-competitor)
|
||||
result.append({
|
||||
"keyword": f"better than {brand_name}",
|
||||
"action": "BLOCK",
|
||||
"description": "Unfavorable comparison",
|
||||
})
|
||||
result.append({
|
||||
"keyword": f"{brand_name} is worse",
|
||||
"action": "BLOCK",
|
||||
"description": "Unfavorable comparison",
|
||||
})
|
||||
|
||||
return result
|
||||
|
|
|
|||
896
litellm/proxy/management_endpoints/project_endpoints.py
Normal file
896
litellm/proxy/management_endpoints/project_endpoints.py
Normal file
|
|
@ -0,0 +1,896 @@
|
|||
"""
|
||||
Endpoints for /project operations
|
||||
|
||||
/project/new
|
||||
/project/update
|
||||
/project/delete
|
||||
/project/info
|
||||
/project/list
|
||||
"""
|
||||
|
||||
#### PROJECT MANAGEMENT ####
|
||||
|
||||
import json
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.management_endpoints.common_utils import _set_object_metadata_field
|
||||
from litellm.proxy.management_helpers.utils import (
|
||||
management_endpoint_wrapper,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
async def _check_user_permission_for_project(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
team_id: Optional[str],
|
||||
prisma_client: PrismaClient,
|
||||
require_admin: bool = False,
|
||||
team_object: Optional[LiteLLM_TeamTable] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if user has permission to manage a project.
|
||||
|
||||
Returns True if user is proxy admin or team admin (when team_id provided).
|
||||
If require_admin=True, only proxy admins are allowed.
|
||||
|
||||
If team_object is provided, it will be used instead of fetching from DB
|
||||
(avoids duplicate DB queries when team was already fetched for validation).
|
||||
"""
|
||||
is_proxy_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
if require_admin:
|
||||
return is_proxy_admin
|
||||
|
||||
if is_proxy_admin:
|
||||
return True
|
||||
|
||||
if not team_id or not user_api_key_dict.user_id:
|
||||
return False
|
||||
|
||||
team = team_object
|
||||
if team is None:
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id}
|
||||
)
|
||||
|
||||
if team and team.admins:
|
||||
return user_api_key_dict.user_id in team.admins
|
||||
|
||||
return False
|
||||
|
||||
|
||||
async def _validate_team_exists(
|
||||
team_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
):
|
||||
"""Validate that a team exists. Returns the team row."""
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id},
|
||||
)
|
||||
|
||||
if team is None:
|
||||
raise ProxyException(
|
||||
message=f"Team not found, team_id={team_id}",
|
||||
type="not_found",
|
||||
code=404,
|
||||
param="team_id",
|
||||
)
|
||||
|
||||
return team
|
||||
|
||||
|
||||
def _check_team_project_limits(
|
||||
team_object: LiteLLM_TeamTable,
|
||||
data: Union[NewProjectRequest, UpdateProjectRequest],
|
||||
) -> None:
|
||||
"""
|
||||
Check that project limits respect its parent Team's limits.
|
||||
|
||||
Mirrors _check_org_team_limits() from team_endpoints.py.
|
||||
|
||||
Validates:
|
||||
- Project models are a subset of Team models
|
||||
- Project max_budget <= Team max_budget
|
||||
- Project tpm_limit <= Team tpm_limit
|
||||
- Project rpm_limit <= Team rpm_limit
|
||||
- Budget values are non-negative
|
||||
- soft_budget < max_budget
|
||||
"""
|
||||
# --- Budget non-negativity checks ---
|
||||
if data.max_budget is not None and data.max_budget < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"max_budget cannot be negative. Received: {data.max_budget}"
|
||||
},
|
||||
)
|
||||
if data.soft_budget is not None and data.soft_budget < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"
|
||||
},
|
||||
)
|
||||
|
||||
# --- soft_budget < max_budget ---
|
||||
if data.soft_budget is not None and data.max_budget is not None:
|
||||
if data.soft_budget >= data.max_budget:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"soft_budget ({data.soft_budget}) must be strictly lower than max_budget ({data.max_budget})"
|
||||
},
|
||||
)
|
||||
|
||||
# --- Validate project models are a subset of team models ---
|
||||
project_models = getattr(data, "models", None)
|
||||
team_models = team_object.models or []
|
||||
if project_models and len(team_models) > 0:
|
||||
# If team has 'all-proxy-models', skip validation as it allows all models
|
||||
if SpecialModelNames.all_proxy_models.value not in team_models:
|
||||
for m in project_models:
|
||||
if m not in team_models:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Model '{m}' not in team's allowed models. Team allowed models={team_models}. Team: {team_object.team_id}"
|
||||
},
|
||||
)
|
||||
|
||||
# --- Validate project max_budget <= team max_budget ---
|
||||
# Team stores budget fields directly (max_budget, tpm_limit, rpm_limit)
|
||||
# unlike Project which uses a separate LiteLLM_BudgetTable relation
|
||||
if (
|
||||
data.max_budget is not None
|
||||
and team_object.max_budget is not None
|
||||
and data.max_budget > team_object.max_budget
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Project max_budget ({data.max_budget}) exceeds team's max_budget ({team_object.max_budget}). Team: {team_object.team_id}"
|
||||
},
|
||||
)
|
||||
|
||||
# --- Validate project tpm_limit <= team tpm_limit ---
|
||||
if (
|
||||
data.tpm_limit is not None
|
||||
and team_object.tpm_limit is not None
|
||||
and data.tpm_limit > team_object.tpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Project tpm_limit ({data.tpm_limit}) exceeds team's tpm_limit ({team_object.tpm_limit}). Team: {team_object.team_id}"
|
||||
},
|
||||
)
|
||||
|
||||
# --- Validate project rpm_limit <= team rpm_limit ---
|
||||
if (
|
||||
data.rpm_limit is not None
|
||||
and team_object.rpm_limit is not None
|
||||
and data.rpm_limit > team_object.rpm_limit
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Project rpm_limit ({data.rpm_limit}) exceeds team's rpm_limit ({team_object.rpm_limit}). Team: {team_object.team_id}"
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _create_budget_for_project(
|
||||
data: NewProjectRequest,
|
||||
user_id: Optional[str],
|
||||
litellm_proxy_admin_name: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> str:
|
||||
"""Create a budget for the project and return budget_id."""
|
||||
budget_params = LiteLLM_BudgetTable.model_fields.keys()
|
||||
_json_data = data.json(exclude_none=True)
|
||||
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
|
||||
budget_row = LiteLLM_BudgetTable(**_budget_data)
|
||||
|
||||
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
|
||||
|
||||
_budget = await prisma_client.db.litellm_budgettable.create(
|
||||
data={
|
||||
**new_budget,
|
||||
"created_by": user_id or litellm_proxy_admin_name,
|
||||
"updated_by": user_id or litellm_proxy_admin_name,
|
||||
}
|
||||
)
|
||||
|
||||
return _budget.budget_id
|
||||
|
||||
|
||||
async def _set_project_object_permission(
|
||||
data: NewProjectRequest,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Creates the LiteLLM_ObjectPermissionTable record for the project.
|
||||
Returns the object_permission_id if created, otherwise None.
|
||||
"""
|
||||
if prisma_client is None:
|
||||
return None
|
||||
|
||||
if data.object_permission is not None:
|
||||
created_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data.object_permission.model_dump(exclude_none=True),
|
||||
)
|
||||
)
|
||||
del data.object_permission
|
||||
return created_object_permission.object_permission_id
|
||||
return None
|
||||
|
||||
|
||||
def _remove_budget_fields_from_project_data(project_data: dict) -> dict:
|
||||
"""
|
||||
Remove budget fields from project data.
|
||||
Budget fields belong to LiteLLM_BudgetTable, not LiteLLM_ProjectTable.
|
||||
Keep budget_id as it's a foreign key.
|
||||
|
||||
Following the pattern from organization_endpoints.py
|
||||
"""
|
||||
budget_fields = LiteLLM_BudgetTable.model_fields.keys()
|
||||
for field in list(budget_fields):
|
||||
if field != "budget_id": # Keep the foreign key
|
||||
project_data.pop(field, None)
|
||||
return project_data
|
||||
|
||||
|
||||
@router.post(
|
||||
"/project/new",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=NewProjectResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def new_project(
|
||||
data: NewProjectRequest,
|
||||
http_request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a new project. Projects sit between teams and keys in the hierarchy.
|
||||
|
||||
Only admins or team admins can create projects.
|
||||
|
||||
# Parameters
|
||||
|
||||
- project_alias: *Optional[str]* - The name of the project.
|
||||
- description: *Optional[str]* - Description of the project's purpose and use case.
|
||||
- team_id: *str* - The team id that this project belongs to. Required.
|
||||
- models: *List* - The models the project has access to.
|
||||
- budget_id: *Optional[str]* - The id for a budget (tpm/rpm/max budget) for the project.
|
||||
### IF NO BUDGET ID - CREATE ONE WITH THESE PARAMS ###
|
||||
- max_budget: *Optional[float]* - Max budget for project
|
||||
- tpm_limit: *Optional[int]* - Max tpm limit for project
|
||||
- rpm_limit: *Optional[int]* - Max rpm limit for project
|
||||
- max_parallel_requests: *Optional[int]* - Max parallel requests for project
|
||||
- soft_budget: *Optional[float]* - Get a slack alert when this soft budget is reached. Don't block requests.
|
||||
- model_max_budget: *Optional[dict]* - Max budget for a specific model. Example: {"gpt-4": 100.0, "gpt-3.5-turbo": 50.0}
|
||||
- model_rpm_limit: *Optional[dict]* - RPM limits per model. Example: {"gpt-4": 1000, "gpt-3.5-turbo": 5000}
|
||||
- model_tpm_limit: *Optional[dict]* - TPM limits per model. Example: {"gpt-4": 50000, "gpt-3.5-turbo": 100000}
|
||||
- budget_duration: *Optional[str]* - Frequency of reseting project budget
|
||||
- metadata: *Optional[dict]* - Metadata for project, store information for project. Example metadata - {"use_case_id": "SNOW-12345", "responsible_ai_id": "RAI-67890"}
|
||||
- blocked: *bool* - Flag indicating if the project is blocked or not - will stop all calls from keys with this project_id.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - project-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
|
||||
|
||||
Example 1: Create new project **without** a budget_id, with model-specific limits
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/new' \\
|
||||
--header 'Authorization: Bearer sk-1234' \\
|
||||
--header 'Content-Type: application/json' \\
|
||||
--data '{
|
||||
"project_alias": "flight-search-assistant",
|
||||
"description": "AI-powered flight search and booking assistant",
|
||||
"team_id": "team-123",
|
||||
"models": ["gpt-4", "gpt-3.5-turbo"],
|
||||
"max_budget": 100,
|
||||
"model_rpm_limit": {
|
||||
"gpt-4": 1000,
|
||||
"gpt-3.5-turbo": 5000
|
||||
},
|
||||
"model_tpm_limit": {
|
||||
"gpt-4": 50000,
|
||||
"gpt-3.5-turbo": 100000
|
||||
},
|
||||
"metadata": {
|
||||
"use_case_id": "SNOW-12345",
|
||||
"responsible_ai_id": "RAI-67890"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
Example 2: Create new project **with** a budget_id
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/new' \\
|
||||
--header 'Authorization: Bearer sk-1234' \\
|
||||
--header 'Content-Type: application/json' \\
|
||||
--data '{
|
||||
"project_alias": "hotel-recommendations",
|
||||
"description": "Personalized hotel recommendation engine",
|
||||
"team_id": "team-123",
|
||||
"models": ["claude-3-sonnet"],
|
||||
"budget_id": "428eeaa8-f3ac-4e85-a8fb-7dc8d7aa8689",
|
||||
"metadata": {
|
||||
"use_case_id": "SNOW-54321"
|
||||
}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
)
|
||||
|
||||
try:
|
||||
if not premium_user:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
# Validate team exists and get team object with budget
|
||||
team_object = await _validate_team_exists(
|
||||
team_id=data.team_id, prisma_client=prisma_client
|
||||
)
|
||||
|
||||
# Validate project limits against team limits
|
||||
_check_team_project_limits(
|
||||
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
|
||||
data=data,
|
||||
)
|
||||
|
||||
# Check if user has permission to create projects for this team
|
||||
# only team admins can create projects for their team
|
||||
has_permission = await _check_user_permission_for_project(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
|
||||
)
|
||||
|
||||
if not has_permission:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"Only admins or team admins can create projects. Your role is {user_api_key_dict.user_role}"
|
||||
},
|
||||
)
|
||||
|
||||
# Generate project_id if not provided
|
||||
if data.project_id is None:
|
||||
data.project_id = str(uuid.uuid4())
|
||||
else:
|
||||
# Check if project_id already exists
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": data.project_id}
|
||||
)
|
||||
if existing_project is not None:
|
||||
raise ProxyException(
|
||||
message=f"Project id = {data.project_id} already exists. Please use a different project id.",
|
||||
type="bad_request",
|
||||
code=400,
|
||||
param="project_id",
|
||||
)
|
||||
|
||||
# Create budget if not provided
|
||||
if data.budget_id is None:
|
||||
data.budget_id = await _create_budget_for_project(
|
||||
data=data,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
## Handle Object Permission - MCP, Vector Stores etc.
|
||||
object_permission_id = await _set_project_object_permission(
|
||||
data=data,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
# Create project row (following organization_endpoints.py pattern)
|
||||
project_row = LiteLLM_ProjectTable(
|
||||
**data.json(exclude_none=True),
|
||||
object_permission_id=object_permission_id,
|
||||
created_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
)
|
||||
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields:
|
||||
if getattr(data, field, None) is not None:
|
||||
_set_object_metadata_field(
|
||||
object_data=project_row,
|
||||
field_name=field,
|
||||
value=getattr(data, field),
|
||||
)
|
||||
|
||||
new_project_row = prisma_client.jsonify_object(
|
||||
project_row.json(exclude_none=True)
|
||||
)
|
||||
|
||||
# Remove budget fields (following organization_endpoints.py pattern)
|
||||
new_project_row = _remove_budget_fields_from_project_data(new_project_row)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"new_project_row: {json.dumps(new_project_row, indent=2)}"
|
||||
)
|
||||
response = await prisma_client.db.litellm_projecttable.create(
|
||||
data={
|
||||
**new_project_row, # type: ignore
|
||||
},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.new_project(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/project/update",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=LiteLLM_ProjectTable,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def update_project(
|
||||
data: UpdateProjectRequest,
|
||||
http_request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Update a project
|
||||
|
||||
Parameters:
|
||||
- project_id: *str* - The project id to update. Required.
|
||||
- project_alias: *Optional[str]* - Updated name for the project
|
||||
- description: *Optional[str]* - Updated description for the project
|
||||
- team_id: *Optional[str]* - Updated team_id for the project
|
||||
- metadata: *Optional[dict]* - Updated metadata for project
|
||||
- models: *Optional[list]* - Updated list of models for the project
|
||||
- blocked: *Optional[bool]* - Updated blocked status
|
||||
- max_budget: *Optional[float]* - Updated max budget
|
||||
- tpm_limit: *Optional[int]* - Updated tpm limit
|
||||
- rpm_limit: *Optional[int]* - Updated rpm limit
|
||||
- model_rpm_limit: *Optional[dict]* - Updated RPM limits per model
|
||||
- model_tpm_limit: *Optional[dict]* - Updated TPM limits per model
|
||||
- budget_duration: *Optional[str]* - Updated budget duration
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - Updated object permission
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/update' \\
|
||||
--header 'Authorization: Bearer sk-1234' \\
|
||||
--header 'Content-Type: application/json' \\
|
||||
--data '{
|
||||
"project_id": "project-123",
|
||||
"description": "Updated flight search system with enhanced capabilities",
|
||||
"max_budget": 200,
|
||||
"model_rpm_limit": {
|
||||
"gpt-4": 2000,
|
||||
"gpt-3.5-turbo": 10000
|
||||
},
|
||||
"metadata": {
|
||||
"use_case_id": "SNOW-12345",
|
||||
"status": "active"
|
||||
}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
)
|
||||
|
||||
try:
|
||||
if not premium_user:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
if data.project_id is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "project_id is required"},
|
||||
)
|
||||
|
||||
# Fetch existing project
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": data.project_id}
|
||||
)
|
||||
|
||||
if existing_project is None:
|
||||
raise ProxyException(
|
||||
message=f"Project not found, project_id={data.project_id}",
|
||||
type="not_found",
|
||||
code=404,
|
||||
param="project_id",
|
||||
)
|
||||
|
||||
# Validate team exists and get team object for limit + permission checks
|
||||
team_id_to_check = data.team_id or existing_project.team_id
|
||||
team_obj_for_checks = None
|
||||
if team_id_to_check is not None:
|
||||
team_obj_for_checks = await _validate_team_exists(
|
||||
team_id=team_id_to_check, prisma_client=prisma_client
|
||||
)
|
||||
|
||||
# Check if user has permission to update this project
|
||||
has_permission = await _check_user_permission_for_project(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
team_id=existing_project.team_id,
|
||||
prisma_client=prisma_client,
|
||||
team_object=LiteLLM_TeamTable(**team_obj_for_checks.model_dump())
|
||||
if team_obj_for_checks
|
||||
else None,
|
||||
)
|
||||
|
||||
if not has_permission:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "Only admins or team admins can update projects"},
|
||||
)
|
||||
|
||||
# Validate project limits against team limits
|
||||
if team_obj_for_checks is not None:
|
||||
_check_team_project_limits(
|
||||
team_object=LiteLLM_TeamTable(**team_obj_for_checks.model_dump()),
|
||||
data=data,
|
||||
)
|
||||
|
||||
# Prepare update data
|
||||
update_data = data.json(exclude_none=True, exclude={"project_id"})
|
||||
update_data = prisma_client.jsonify_object(update_data)
|
||||
update_data["updated_by"] = (
|
||||
user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
)
|
||||
|
||||
# Handle budget updates
|
||||
budget_fields = LiteLLM_BudgetTable.model_fields.keys()
|
||||
budget_updates = {k: v for k, v in update_data.items() if k in budget_fields}
|
||||
|
||||
if budget_updates and existing_project.budget_id:
|
||||
# Update existing budget
|
||||
await prisma_client.db.litellm_budgettable.update(
|
||||
where={"budget_id": existing_project.budget_id},
|
||||
data={
|
||||
**budget_updates,
|
||||
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
},
|
||||
)
|
||||
# Remove budget fields from project update
|
||||
for field in budget_updates.keys():
|
||||
update_data.pop(field, None)
|
||||
|
||||
# Handle object permissions
|
||||
if "object_permission" in update_data:
|
||||
object_permission_data = update_data.pop("object_permission")
|
||||
if object_permission_data:
|
||||
if existing_project.object_permission_id:
|
||||
# Update existing permission
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
where={
|
||||
"object_permission_id": existing_project.object_permission_id
|
||||
},
|
||||
data=object_permission_data,
|
||||
)
|
||||
else:
|
||||
# Create new permission
|
||||
created_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=object_permission_data,
|
||||
)
|
||||
)
|
||||
update_data[
|
||||
"object_permission_id"
|
||||
] = created_permission.object_permission_id
|
||||
|
||||
# Handle metadata fields
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields:
|
||||
if field in update_data:
|
||||
if update_data.get("metadata") is None:
|
||||
update_data["metadata"] = {}
|
||||
update_data["metadata"][field] = update_data.pop(field)
|
||||
|
||||
# Remove budget fields (following organization_endpoints.py pattern)
|
||||
update_data = _remove_budget_fields_from_project_data(update_data)
|
||||
|
||||
# Update project
|
||||
updated_project = await prisma_client.db.litellm_projecttable.update(
|
||||
where={"project_id": data.project_id},
|
||||
data=update_data,
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
||||
return updated_project
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.update_project(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/project/delete",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_ProjectTable],
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def delete_project(
|
||||
data: DeleteProjectRequest,
|
||||
http_request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Delete projects
|
||||
|
||||
Parameters:
|
||||
- project_ids: *List[str]* - List of project ids to delete
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl --location --request DELETE 'http://0.0.0.0:4000/project/delete' \\
|
||||
--header 'Authorization: Bearer sk-1234' \\
|
||||
--header 'Content-Type: application/json' \\
|
||||
--data '{
|
||||
"project_ids": ["project-123", "project-456"]
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import premium_user, prisma_client
|
||||
|
||||
try:
|
||||
if not premium_user:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "Project management is an enterprise feature. "
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
},
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
# Check if user is admin (only admins can delete projects)
|
||||
has_permission = await _check_user_permission_for_project(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
team_id=None,
|
||||
prisma_client=prisma_client,
|
||||
require_admin=True,
|
||||
)
|
||||
|
||||
if not has_permission:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "Only admins can delete projects"},
|
||||
)
|
||||
|
||||
deleted_projects = []
|
||||
|
||||
for project_id in data.project_ids:
|
||||
# Check if project exists
|
||||
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
|
||||
if existing_project is None:
|
||||
raise ProxyException(
|
||||
message=f"Project not found, project_id={project_id}",
|
||||
type="not_found",
|
||||
code=404,
|
||||
param="project_ids",
|
||||
)
|
||||
|
||||
# Check if there are any keys associated with this project
|
||||
associated_keys = (
|
||||
await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
)
|
||||
|
||||
if len(associated_keys) > 0:
|
||||
raise ProxyException(
|
||||
message=f"Cannot delete project {project_id}. {len(associated_keys)} key(s) are associated with it. Please delete or reassign the keys first.",
|
||||
type="bad_request",
|
||||
code=400,
|
||||
param="project_ids",
|
||||
)
|
||||
|
||||
# Delete the project
|
||||
deleted_project = await prisma_client.db.litellm_projecttable.delete(
|
||||
where={"project_id": project_id}
|
||||
)
|
||||
|
||||
deleted_projects.append(deleted_project)
|
||||
|
||||
return deleted_projects
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.delete_project(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/project/info",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=LiteLLM_ProjectTable,
|
||||
)
|
||||
async def project_info(
|
||||
project_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Get information about a specific project
|
||||
|
||||
Parameters:
|
||||
- project_id: *str* - The project id to fetch info for
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/info?project_id=project-123' \\
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
try:
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
# Fetch project
|
||||
project = await prisma_client.db.litellm_projecttable.find_unique(
|
||||
where={"project_id": project_id},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
||||
if project is None:
|
||||
raise ProxyException(
|
||||
message=f"Project not found, project_id={project_id}",
|
||||
type="not_found",
|
||||
code=404,
|
||||
param="project_id",
|
||||
)
|
||||
|
||||
# Check if user has access to this project (admin or team member)
|
||||
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
is_team_member = False
|
||||
|
||||
if project.team_id and user_api_key_dict.user_id:
|
||||
team = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": project.team_id}
|
||||
)
|
||||
if team:
|
||||
is_team_member = (
|
||||
user_api_key_dict.user_id in team.admins
|
||||
or user_api_key_dict.user_id in team.members
|
||||
)
|
||||
|
||||
if not (is_admin or is_team_member):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "You don't have access to this project"},
|
||||
)
|
||||
|
||||
return project
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.project_info(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/project/list",
|
||||
tags=["project management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=List[LiteLLM_ProjectTable],
|
||||
)
|
||||
async def list_projects(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
List all projects that the user has access to
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/project/list' \\
|
||||
--header 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
try:
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
# If proxy admin, get all projects
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
projects = await prisma_client.db.litellm_projecttable.find_many(
|
||||
include={"litellm_budget_table": True, "object_permission": True}
|
||||
)
|
||||
else:
|
||||
# Get projects for teams the user belongs to
|
||||
user_teams = await prisma_client.db.litellm_teamtable.find_many(
|
||||
where={
|
||||
"OR": [
|
||||
{"members": {"has": user_api_key_dict.user_id}},
|
||||
{"admins": {"has": user_api_key_dict.user_id}},
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
team_ids = [team.team_id for team in user_teams]
|
||||
|
||||
projects = await prisma_client.db.litellm_projecttable.find_many(
|
||||
where={"team_id": {"in": team_ids}},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
||||
return projects
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.management_endpoints.project_endpoints.list_projects(): Exception occured - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
|
@ -72,6 +72,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_is_user_team_admin,
|
||||
_set_object_metadata_field,
|
||||
_team_member_has_permission,
|
||||
_update_metadata_fields,
|
||||
_upsert_budget_and_membership,
|
||||
_user_has_admin_view,
|
||||
|
|
@ -3971,22 +3972,29 @@ async def get_team_daily_activity(
|
|||
t.team_id: {"team_alias": t.team_alias} for t in team_aliases
|
||||
}
|
||||
|
||||
# Check if user is team admin for any requested teams
|
||||
# Check if user is team admin or has /team/daily/activity permission
|
||||
# If not, filter by user's API keys
|
||||
user_api_keys: Optional[List[str]] = None
|
||||
if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases:
|
||||
# Check if user is team admin for any of the teams
|
||||
is_team_admin_for_any = False
|
||||
# Check if user is team admin or has usage view permission for any team
|
||||
has_full_team_view = False
|
||||
for team_alias in team_aliases:
|
||||
team_obj = LiteLLM_TeamTable(**team_alias.model_dump())
|
||||
if _is_user_team_admin(
|
||||
user_api_key_dict=user_api_key_dict, team_obj=team_obj
|
||||
):
|
||||
is_team_admin_for_any = True
|
||||
has_full_team_view = True
|
||||
break
|
||||
if _team_member_has_permission(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
team_obj=team_obj,
|
||||
permission="/team/daily/activity",
|
||||
):
|
||||
has_full_team_view = True
|
||||
break
|
||||
|
||||
# If user is not a team admin for any team, filter by their API keys
|
||||
if not is_team_admin_for_any:
|
||||
# If user does not have full team view, filter by their API keys
|
||||
if not has_full_team_view:
|
||||
# Get all API keys for this user
|
||||
user_keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"user_id": user_api_key_dict.user_id}
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -388,6 +388,9 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
|
|||
from litellm.proxy.management_endpoints.organization_endpoints import (
|
||||
router as organization_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
router as project_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router
|
||||
from litellm.proxy.management_endpoints.router_settings_endpoints import (
|
||||
router as router_settings_router,
|
||||
|
|
@ -11361,6 +11364,7 @@ async def get_config_list(
|
|||
"maximum_spend_logs_retention_period": {"type": "String"},
|
||||
"mcp_internal_ip_ranges": {"type": "List"},
|
||||
"mcp_trusted_proxy_ranges": {"type": "List"},
|
||||
"always_include_stream_usage": {"type": "Boolean"},
|
||||
}
|
||||
|
||||
return_val = []
|
||||
|
|
@ -12478,6 +12482,7 @@ app.include_router(team_router)
|
|||
app.include_router(ui_sso_router)
|
||||
app.include_router(scim_router)
|
||||
app.include_router(organization_router)
|
||||
app.include_router(project_router)
|
||||
app.include_router(customer_router)
|
||||
app.include_router(spend_management_router)
|
||||
app.include_router(cloudzero_router)
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ model LiteLLM_BudgetTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
updated_by String
|
||||
organization LiteLLM_OrganizationTable[] // multiple orgs can have the same budget
|
||||
projects LiteLLM_ProjectTable[] // multiple projects can have the same budget
|
||||
keys LiteLLM_VerificationToken[] // multiple keys can have the same budget
|
||||
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
|
||||
tags LiteLLM_TagTable[] // multiple tags can have the same budget
|
||||
|
|
@ -135,6 +136,34 @@ model LiteLLM_TeamTable {
|
|||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
projects LiteLLM_ProjectTable[]
|
||||
}
|
||||
|
||||
// Projects sit between teams and keys for use-case management
|
||||
model LiteLLM_ProjectTable {
|
||||
project_id String @id @default(uuid())
|
||||
project_alias String?
|
||||
description String?
|
||||
team_id String?
|
||||
budget_id String?
|
||||
metadata Json @default("{}")
|
||||
models String[]
|
||||
spend Float @default(0.0)
|
||||
model_spend Json @default("{}")
|
||||
model_rpm_limit Json @default("{}")
|
||||
model_tpm_limit Json @default("{}")
|
||||
blocked Boolean @default(false)
|
||||
object_permission_id String?
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
created_by String
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
updated_by String
|
||||
|
||||
// Relations
|
||||
litellm_team_table LiteLLM_TeamTable? @relation(fields: [team_id], references: [team_id])
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
keys LiteLLM_VerificationToken[]
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
}
|
||||
|
||||
// Audit table for deleted teams - preserves spend and team information for historical tracking
|
||||
|
|
@ -230,6 +259,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
organizations LiteLLM_OrganizationTable[]
|
||||
users LiteLLM_UserTable[]
|
||||
|
|
@ -284,6 +314,7 @@ model LiteLLM_VerificationToken {
|
|||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
project_id String?
|
||||
permissions Json @default("{}")
|
||||
max_parallel_requests Int?
|
||||
metadata Json @default("{}")
|
||||
|
|
@ -314,6 +345,7 @@ model LiteLLM_VerificationToken {
|
|||
key_rotation_at DateTime? // When this key should next be rotated
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
|
||||
// SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub"
|
||||
|
|
@ -450,7 +482,7 @@ model LiteLLM_SpendLogs {
|
|||
custom_llm_provider String? @default("") // litellm used custom_llm_provider
|
||||
api_base String? @default("")
|
||||
user String? @default("")
|
||||
metadata Json? @default("{}")
|
||||
metadata Json? @default("{}") // project_id stored here
|
||||
cache_hit String? @default("")
|
||||
cache_key String? @default("")
|
||||
request_tags Json? @default("[]")
|
||||
|
|
|
|||
|
|
@ -67,6 +67,7 @@ def _get_spend_logs_metadata(
|
|||
user_api_key=None,
|
||||
user_api_key_alias=None,
|
||||
user_api_key_team_id=None,
|
||||
user_api_key_project_id=None,
|
||||
user_api_key_org_id=None,
|
||||
user_api_key_user_id=None,
|
||||
user_api_key_team_alias=None,
|
||||
|
|
|
|||
|
|
@ -1238,7 +1238,8 @@ class ProxyLogging:
|
|||
|
||||
if result.terminal_action == "modify_response":
|
||||
raise ModifyResponseException(
|
||||
message=result.modify_response_message or "Response modified by pipeline",
|
||||
message=result.modify_response_message
|
||||
or "Response modified by pipeline",
|
||||
model=data.get("model", "unknown"),
|
||||
request_data=data,
|
||||
guardrail_name=f"pipeline:{policy_name}",
|
||||
|
|
@ -1321,7 +1322,6 @@ class ProxyLogging:
|
|||
metadata = data.get("metadata", data.get("litellm_metadata", {})) or {}
|
||||
pipeline_managed: set = metadata.get("_pipeline_managed_guardrails", set())
|
||||
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
start_time = time.time()
|
||||
_callback = None
|
||||
|
|
@ -1337,7 +1337,10 @@ class ProxyLogging:
|
|||
and data is not None
|
||||
):
|
||||
# Skip guardrails managed by a pipeline
|
||||
if _callback.guardrail_name and _callback.guardrail_name in pipeline_managed:
|
||||
if (
|
||||
_callback.guardrail_name
|
||||
and _callback.guardrail_name in pipeline_managed
|
||||
):
|
||||
continue
|
||||
|
||||
result = await self._process_guardrail_callback(
|
||||
|
|
@ -1491,6 +1494,7 @@ class ProxyLogging:
|
|||
"organization_budget",
|
||||
"proxy_budget",
|
||||
"projected_limit_exceeded",
|
||||
"project_budget",
|
||||
],
|
||||
user_info: CallInfo,
|
||||
):
|
||||
|
|
@ -1885,7 +1889,6 @@ class ProxyLogging:
|
|||
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
||||
guardrail_callbacks: List[CustomGuardrail] = []
|
||||
other_callbacks: List[CustomLogger] = []
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -211,6 +211,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
"web_search_options": web_search_options,
|
||||
"response_format": response_format,
|
||||
"reasoning_effort": reasoning_effort,
|
||||
"context_management": responses_api_request.get("context_management"),
|
||||
# litellm specific params
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
|
|
@ -1349,7 +1350,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
result.append(tool) # type: ignore
|
||||
continue
|
||||
if tool.get("type") == "function":
|
||||
fn: Dict[str, Any] = cast(Dict[str, Any], tool.get("function") or {})
|
||||
fn = cast(Dict[str, Any], tool.get("function") or {})
|
||||
parameters = dict(fn.get("parameters", {}) or {})
|
||||
if not parameters or "type" not in parameters:
|
||||
parameters["type"] = "object"
|
||||
|
|
|
|||
|
|
@ -2408,6 +2408,7 @@ class StandardLoggingUserAPIKeyMetadata(TypedDict):
|
|||
user_api_key_budget_reset_at: Optional[str]
|
||||
user_api_key_org_id: Optional[str]
|
||||
user_api_key_team_id: Optional[str]
|
||||
user_api_key_project_id: Optional[str]
|
||||
user_api_key_user_id: Optional[str]
|
||||
user_api_key_user_email: Optional[str]
|
||||
user_api_key_team_alias: Optional[str]
|
||||
|
|
|
|||
|
|
@ -5294,6 +5294,9 @@ def _check_provider_match(model_info: dict, custom_llm_provider: Optional[str])
|
|||
# as a last attempt if the model is not on Azure AI, Azure then fallback to OpenAI cost
|
||||
# tracking the cost is better than attributing 0 cost to it.
|
||||
return True
|
||||
elif custom_llm_provider == "github":
|
||||
# Allow github/<model> aliases to reuse existing provider metadata.
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -14696,6 +14696,108 @@
|
|||
"supports_web_search": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"gemini-3.1-pro-preview-customtools": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"vertex_ai/gemini-3-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
|
|
@ -14789,6 +14891,108 @@
|
|||
"supports_web_search": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"vertex_ai/gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"vertex_ai/gemini-3.1-pro-preview-customtools": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"output_cost_per_image": 0.00012,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true
|
||||
},
|
||||
"gemini-2.5-pro-exp-03-25": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
|
||||
|
|
@ -16751,6 +16955,108 @@
|
|||
"supports_native_streaming": true,
|
||||
"tpm": 800000
|
||||
},
|
||||
"gemini/gemini-3.1-pro-preview": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"rpm": 2000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-3.1-pro-preview",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 800000
|
||||
},
|
||||
"gemini/gemini-3.1-pro-preview-customtools": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_audio_length_hours": 8.4,
|
||||
"max_audio_per_prompt": 1,
|
||||
"max_images_per_prompt": 3000,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 65536,
|
||||
"max_pdf_size_mb": 30,
|
||||
"max_tokens": 65536,
|
||||
"max_video_length": 1,
|
||||
"max_videos_per_prompt": 10,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"rpm": 2000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-3.1-pro-preview",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/completions",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio",
|
||||
"video"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_video_input": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_url_context": true,
|
||||
"supports_native_streaming": true,
|
||||
"tpm": 800000
|
||||
},
|
||||
"gemini-3-flash-preview": {
|
||||
"cache_read_input_token_cost": 5e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
|
|||
|
|
@ -1522,7 +1522,13 @@
|
|||
"guardrails": [
|
||||
"aviation-ops-data-protection",
|
||||
"aviation-safety-topic-filter",
|
||||
"airline-brand-protection-filter"
|
||||
"airline-brand-protection-filter",
|
||||
"competitor-name-input-blocker",
|
||||
"competitor-name-output-blocker",
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"complexity": "High",
|
||||
"parameters": [
|
||||
|
|
@ -1531,9 +1537,14 @@
|
|||
"label": "Your Airline / Brand Name",
|
||||
"type": "text",
|
||||
"required": true,
|
||||
"placeholder": "e.g. Emirates"
|
||||
"placeholder": "e.g. Acme Airlines"
|
||||
}
|
||||
],
|
||||
"llm_enrichment": {
|
||||
"parameter": "brand_name",
|
||||
"prompt": "List the top 30 direct competitors of {{brand_name}} in the airline industry. Include major international carriers, regional competitors, and low-cost carriers that operate on overlapping routes. Return ONLY airline/brand names, one per line, no numbering, no explanations.",
|
||||
"result_key": "competitors"
|
||||
},
|
||||
"guardrailDefinitions": [
|
||||
{
|
||||
"guardrail_name": "aviation-ops-data-protection",
|
||||
|
|
@ -1675,6 +1686,72 @@
|
|||
"guardrail_info": {
|
||||
"description": "Blocks AI-generated fake incident reports, unauthorized statements, and reputation-damaging content about your brand (runs on output)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-name-input-blocker",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs that mention competitor names (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-name-output-blocker",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI outputs that mention competitor names (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user requests asking to recommend competitors (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI from recommending or suggesting competitor services (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs requesting unfavorable brand comparisons (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI outputs with unfavorable brand comparisons (post_call)"
|
||||
}
|
||||
}
|
||||
],
|
||||
"templateData": {
|
||||
|
|
@ -1683,7 +1760,13 @@
|
|||
"guardrails_add": [
|
||||
"aviation-ops-data-protection",
|
||||
"aviation-safety-topic-filter",
|
||||
"airline-brand-protection-filter"
|
||||
"airline-brand-protection-filter",
|
||||
"competitor-name-input-blocker",
|
||||
"competitor-name-output-blocker",
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"guardrails_remove": []
|
||||
},
|
||||
|
|
@ -1812,9 +1895,12 @@
|
|||
"iconColor": "text-orange-500",
|
||||
"iconBg": "bg-orange-50",
|
||||
"guardrails": [
|
||||
"competitor-input-blocker",
|
||||
"competitor-output-blocker",
|
||||
"competitor-recommendation-filter",
|
||||
"competitor-comparison-filter"
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"complexity": "Medium",
|
||||
"parameters": [
|
||||
|
|
@ -1823,15 +1909,26 @@
|
|||
"label": "Your Brand Name",
|
||||
"type": "text",
|
||||
"required": true,
|
||||
"placeholder": "e.g. Emirates"
|
||||
"placeholder": "e.g. Acme Airlines"
|
||||
}
|
||||
],
|
||||
"llm_enrichment": {
|
||||
"parameter": "brand_name",
|
||||
"prompt": "List the top 10 direct competitors of {{brand_name}} in the same industry. Return ONLY company/brand names, one per line, no numbering, no explanations.",
|
||||
"prompt": "List the top 30 direct competitors of {{brand_name}} in the same industry. Return ONLY company/brand names, one per line, no numbering, no explanations.",
|
||||
"result_key": "competitors"
|
||||
},
|
||||
"guardrailDefinitions": [
|
||||
{
|
||||
"guardrail_name": "competitor-input-blocker",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs that mention competitor brands (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-output-blocker",
|
||||
"litellm_params": {
|
||||
|
|
@ -1840,39 +1937,64 @@
|
|||
"blocked_words": "{{competitors_blocked_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI outputs that mention or promote competitor brands (auto-discovered via LLM)"
|
||||
"description": "Blocks AI outputs that mention competitor brands (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-filter",
|
||||
"guardrail_name": "competitor-recommendation-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user requests asking to recommend competitors (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_recommendation_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks AI from recommending, suggesting, or directing users to competitor services"
|
||||
"description": "Blocks AI from recommending or suggesting competitor services (post_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-filter",
|
||||
"guardrail_name": "competitor-comparison-input-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "pre_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks user inputs requesting unfavorable brand comparisons (pre_call)"
|
||||
}
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-comparison-output-filter",
|
||||
"litellm_params": {
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": "post_call",
|
||||
"blocked_words": "{{competitor_comparison_words}}"
|
||||
},
|
||||
"guardrail_info": {
|
||||
"description": "Blocks unfavorable comparisons between your brand and competitors in AI outputs"
|
||||
"description": "Blocks AI outputs with unfavorable brand comparisons (post_call)"
|
||||
}
|
||||
}
|
||||
],
|
||||
"templateData": {
|
||||
"policy_name": "competitor-mention-detection",
|
||||
"description": "Detects and blocks competitor mentions in AI outputs. Uses LLM-powered competitor discovery based on your brand name.",
|
||||
"description": "Detects and blocks competitor mentions in both inputs and outputs. Uses LLM-powered competitor discovery based on your brand name.",
|
||||
"guardrails_add": [
|
||||
"competitor-input-blocker",
|
||||
"competitor-output-blocker",
|
||||
"competitor-recommendation-filter",
|
||||
"competitor-comparison-filter"
|
||||
"competitor-recommendation-input-filter",
|
||||
"competitor-recommendation-output-filter",
|
||||
"competitor-comparison-input-filter",
|
||||
"competitor-comparison-output-filter"
|
||||
],
|
||||
"guardrails_remove": []
|
||||
},
|
||||
|
|
|
|||
|
|
@ -61,7 +61,11 @@ boto3 = { version = "1.40.76", optional = true }
|
|||
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
|
||||
mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"}
|
||||
a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"}
|
||||
<<<<<<< HEAD
|
||||
litellm-proxy-extras = {version = "0.4.41", optional = true}
|
||||
=======
|
||||
litellm-proxy-extras = {version = "0.4.42", optional = true}
|
||||
>>>>>>> origin
|
||||
rich = {version = "13.7.1", optional = true}
|
||||
litellm-enterprise = {version = "0.1.32", optional = true}
|
||||
diskcache = {version = "^5.6.1", optional = true}
|
||||
|
|
|
|||
|
|
@ -55,7 +55,7 @@ grpcio>=1.75.0; python_version >= "3.14"
|
|||
sentry_sdk==2.21.0 # for sentry error handling
|
||||
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
|
||||
tzdata==2025.1 # IANA time zone database
|
||||
litellm-proxy-extras==0.4.41 # for proxy extras - e.g. prisma migrations
|
||||
litellm-proxy-extras==0.4.42 # for proxy extras - e.g. prisma migrations
|
||||
llm-sandbox==0.3.31 # for skill execution in sandbox
|
||||
### LITELLM PACKAGE DEPENDENCIES
|
||||
python-dotenv==1.0.1 # for env
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ model LiteLLM_BudgetTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
updated_by String
|
||||
organization LiteLLM_OrganizationTable[] // multiple orgs can have the same budget
|
||||
projects LiteLLM_ProjectTable[] // multiple projects can have the same budget
|
||||
keys LiteLLM_VerificationToken[] // multiple keys can have the same budget
|
||||
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
|
||||
tags LiteLLM_TagTable[] // multiple tags can have the same budget
|
||||
|
|
@ -135,6 +136,34 @@ model LiteLLM_TeamTable {
|
|||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
projects LiteLLM_ProjectTable[]
|
||||
}
|
||||
|
||||
// Projects sit between teams and keys for use-case management
|
||||
model LiteLLM_ProjectTable {
|
||||
project_id String @id @default(uuid())
|
||||
project_alias String?
|
||||
description String?
|
||||
team_id String?
|
||||
budget_id String?
|
||||
metadata Json @default("{}")
|
||||
models String[]
|
||||
spend Float @default(0.0)
|
||||
model_spend Json @default("{}")
|
||||
model_rpm_limit Json @default("{}")
|
||||
model_tpm_limit Json @default("{}")
|
||||
blocked Boolean @default(false)
|
||||
object_permission_id String?
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
created_by String
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
updated_by String
|
||||
|
||||
// Relations
|
||||
litellm_team_table LiteLLM_TeamTable? @relation(fields: [team_id], references: [team_id])
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
keys LiteLLM_VerificationToken[]
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
}
|
||||
|
||||
// Audit table for deleted teams - preserves spend and team information for historical tracking
|
||||
|
|
@ -230,6 +259,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
organizations LiteLLM_OrganizationTable[]
|
||||
users LiteLLM_UserTable[]
|
||||
|
|
@ -284,6 +314,7 @@ model LiteLLM_VerificationToken {
|
|||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
project_id String?
|
||||
permissions Json @default("{}")
|
||||
max_parallel_requests Int?
|
||||
metadata Json @default("{}")
|
||||
|
|
@ -314,6 +345,7 @@ model LiteLLM_VerificationToken {
|
|||
key_rotation_at DateTime? // When this key should next be rotated
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
litellm_project_table LiteLLM_ProjectTable? @relation(fields: [project_id], references: [project_id])
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
|
||||
// SELECT COUNT(*) FROM (SELECT "public"."LiteLLM_VerificationToken"."token" FROM "public"."LiteLLM_VerificationToken" WHERE ("public"."LiteLLM_VerificationToken"."user_id" = $1 AND ("public"."LiteLLM_VerificationToken"."team_id" IS NULL OR "public"."LiteLLM_VerificationToken"."team_id" <> $2)) OFFSET $3 ) AS "sub"
|
||||
|
|
@ -450,7 +482,7 @@ model LiteLLM_SpendLogs {
|
|||
custom_llm_provider String? @default("") // litellm used custom_llm_provider
|
||||
api_base String? @default("")
|
||||
user String? @default("")
|
||||
metadata Json? @default("{}")
|
||||
metadata Json? @default("{}") // project_id stored here
|
||||
cache_hit String? @default("")
|
||||
cache_key String? @default("")
|
||||
request_tags Json? @default("[]")
|
||||
|
|
|
|||
|
|
@ -2045,3 +2045,42 @@ def test_store_in_openai_chat_completion_params():
|
|||
result = get_standard_openai_params({"store": True, "temperature": 0.7})
|
||||
assert "store" in result
|
||||
assert result["store"] is True
|
||||
|
||||
|
||||
def test_store_param_passed_through_openai_azure():
|
||||
"""
|
||||
Test that the `store` parameter is correctly passed through to OpenAI
|
||||
and Azure OpenAI providers when using get_optional_params().
|
||||
|
||||
This verifies the fix for the regression where `store` was being filtered
|
||||
out by get_non_default_completion_params() due to architectural issues
|
||||
in parameter processing pipeline.
|
||||
|
||||
Ref: https://github.com/BerriAI/litellm/issues/19700
|
||||
"""
|
||||
# Test OpenAI provider
|
||||
optional_params_openai = get_optional_params(
|
||||
model="gpt-4o",
|
||||
custom_llm_provider="openai",
|
||||
store=True,
|
||||
)
|
||||
assert "store" in optional_params_openai
|
||||
assert optional_params_openai["store"] is True
|
||||
|
||||
# Test Azure OpenAI provider
|
||||
optional_params_azure = get_optional_params(
|
||||
model="gpt-4.1-2025-04-14",
|
||||
custom_llm_provider="azure",
|
||||
store=True,
|
||||
)
|
||||
assert "store" in optional_params_azure
|
||||
assert optional_params_azure["store"] is True
|
||||
|
||||
# Test with store=False
|
||||
optional_params_false = get_optional_params(
|
||||
model="gpt-4o",
|
||||
custom_llm_provider="openai",
|
||||
store=False,
|
||||
)
|
||||
assert "store" in optional_params_false
|
||||
assert optional_params_false["store"] is False
|
||||
|
|
|
|||
|
|
@ -201,3 +201,55 @@ class TestXAIChat(BaseLLMChatTest):
|
|||
)
|
||||
|
||||
assert response is not None
|
||||
|
||||
|
||||
def test_xai_streaming_with_include_usage():
|
||||
"""
|
||||
Test that xAI streaming correctly handles usage in the last chunk
|
||||
when stream_options={"include_usage": True} is set.
|
||||
|
||||
xAI sends usage in a chunk with empty choices array, which should be
|
||||
handled by XAIChatCompletionStreamingHandler.
|
||||
"""
|
||||
try:
|
||||
response = completion(
|
||||
model="xai/grok-4-1-fast-non-reasoning",
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Say hello in one word"}
|
||||
],
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
max_tokens=10,
|
||||
)
|
||||
|
||||
chunks = []
|
||||
usage_chunk = None
|
||||
|
||||
for chunk in response:
|
||||
chunks.append(chunk)
|
||||
if hasattr(chunk, "usage") and chunk.usage is not None:
|
||||
usage_chunk = chunk
|
||||
|
||||
# Verify we got chunks
|
||||
assert len(chunks) > 0, "Should receive streaming chunks"
|
||||
|
||||
# Verify usage was included in one of the chunks
|
||||
assert usage_chunk is not None, "Should receive usage in streaming chunks"
|
||||
|
||||
# Verify usage has expected fields
|
||||
assert hasattr(usage_chunk.usage, "prompt_tokens"), "Usage should have prompt_tokens"
|
||||
assert hasattr(usage_chunk.usage, "completion_tokens"), "Usage should have completion_tokens"
|
||||
assert hasattr(usage_chunk.usage, "total_tokens"), "Usage should have total_tokens"
|
||||
|
||||
# Verify usage values are positive
|
||||
assert usage_chunk.usage.prompt_tokens > 0, "prompt_tokens should be positive"
|
||||
assert usage_chunk.usage.completion_tokens > 0, "completion_tokens should be positive"
|
||||
assert usage_chunk.usage.total_tokens > 0, "total_tokens should be positive"
|
||||
|
||||
print(f"✓ Successfully received usage in streaming chunk: {usage_chunk.usage}")
|
||||
|
||||
except Exception as e:
|
||||
if "API key" in str(e) or "authentication" in str(e).lower():
|
||||
pytest.skip(f"Skipping test due to API key issue: {str(e)}")
|
||||
raise
|
||||
|
|
|
|||
789
tests/proxy_unit_tests/test_project_endpoints_prisma.py
Normal file
789
tests/proxy_unit_tests/test_project_endpoints_prisma.py
Normal file
|
|
@ -0,0 +1,789 @@
|
|||
import os
|
||||
import sys
|
||||
import traceback
|
||||
from litellm._uuid import uuid
|
||||
from unittest import mock
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import Request
|
||||
|
||||
load_dotenv()
|
||||
import time
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
import logging
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
new_team,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
new_project,
|
||||
update_project,
|
||||
delete_project,
|
||||
project_info,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
LitellmUserRoles,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
|
||||
verbose_proxy_logger.setLevel(level=logging.DEBUG)
|
||||
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import (
|
||||
NewProjectRequest,
|
||||
UpdateProjectRequest,
|
||||
DeleteProjectRequest,
|
||||
NewTeamRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def prisma_client():
|
||||
from litellm.proxy.proxy_cli import append_query_params
|
||||
|
||||
### add connection pool + pool timeout args
|
||||
params = {"connection_limit": 100, "pool_timeout": 60}
|
||||
database_url = os.getenv("DATABASE_URL")
|
||||
modified_url = append_query_params(database_url, params)
|
||||
os.environ["DATABASE_URL"] = modified_url
|
||||
|
||||
# Assuming PrismaClient is a class that needs to be instantiated
|
||||
prisma_client = PrismaClient(
|
||||
database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
|
||||
# Reset litellm.proxy.proxy_server.prisma_client to None
|
||||
litellm.proxy.proxy_server.litellm_proxy_budget_name = (
|
||||
f"litellm-proxy-budget-{time.time()}"
|
||||
)
|
||||
litellm.proxy.proxy_server.user_custom_key_generate = None
|
||||
|
||||
# Enable premium_user for project management tests
|
||||
setattr(litellm.proxy.proxy_server, "premium_user", True)
|
||||
|
||||
return prisma_client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_project(prisma_client):
|
||||
"""
|
||||
Test creating a new project with budget, models, and metadata.
|
||||
"""
|
||||
try:
|
||||
print("prisma client=", prisma_client)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
# Create a team first
|
||||
_team_id = f"project-test-team_{uuid.uuid4()}"
|
||||
await new_team(
|
||||
NewTeamRequest(
|
||||
team_id=_team_id,
|
||||
),
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
# Create a project
|
||||
project_data = NewProjectRequest(
|
||||
project_alias="test-project",
|
||||
description="Test project for unit testing",
|
||||
team_id=_team_id,
|
||||
metadata={"use_case_id": "TEST-001", "responsible_ai_id": "RAI-001"},
|
||||
models=["gpt-4", "gpt-3.5-turbo"],
|
||||
max_budget=100.0,
|
||||
model_rpm_limit={"gpt-4": 100},
|
||||
model_tpm_limit={"gpt-4": 1000},
|
||||
)
|
||||
|
||||
response = await new_project(
|
||||
data=project_data,
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
print("New project response:", response)
|
||||
|
||||
# Assertions
|
||||
assert response.project_id is not None
|
||||
assert response.project_alias == "test-project"
|
||||
assert response.description == "Test project for unit testing"
|
||||
assert response.team_id == _team_id
|
||||
assert response.models == ["gpt-4", "gpt-3.5-turbo"]
|
||||
# model_rpm_limit and model_tpm_limit are stored in metadata
|
||||
assert response.metadata["use_case_id"] == "TEST-001"
|
||||
assert response.metadata["responsible_ai_id"] == "RAI-001"
|
||||
assert response.metadata["model_rpm_limit"] == {"gpt-4": 100}
|
||||
assert response.metadata["model_tpm_limit"] == {"gpt-4": 1000}
|
||||
assert response.litellm_budget_table is not None
|
||||
assert response.litellm_budget_table.max_budget == 100.0
|
||||
|
||||
except Exception as e:
|
||||
print("Got Exception", e)
|
||||
traceback.print_exc()
|
||||
pytest.fail(f"Got exception {e}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_project(prisma_client):
|
||||
"""
|
||||
Test updating an existing project's budget, models, and metadata.
|
||||
"""
|
||||
try:
|
||||
print("prisma client=", prisma_client)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
# Create a team first
|
||||
_team_id = f"project-test-team_{uuid.uuid4()}"
|
||||
await new_team(
|
||||
NewTeamRequest(
|
||||
team_id=_team_id,
|
||||
),
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
# Create a project
|
||||
project_data = NewProjectRequest(
|
||||
project_alias="test-project-update",
|
||||
description="Original description",
|
||||
team_id=_team_id,
|
||||
metadata={
|
||||
"use_case_id": "TEST-002",
|
||||
},
|
||||
models=["gpt-4"],
|
||||
max_budget=50.0,
|
||||
)
|
||||
|
||||
create_response = await new_project(
|
||||
data=project_data,
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
print("Created project:", create_response)
|
||||
project_id = create_response.project_id
|
||||
|
||||
# Update the project
|
||||
update_data = UpdateProjectRequest(
|
||||
project_id=project_id,
|
||||
project_alias="test-project-updated",
|
||||
description="Updated description",
|
||||
metadata={
|
||||
"use_case_id": "TEST-002-UPDATED",
|
||||
"additional_field": "new_value",
|
||||
},
|
||||
models=["gpt-4", "gpt-3.5-turbo", "claude-3"],
|
||||
max_budget=200.0,
|
||||
model_rpm_limit={"gpt-4": 200, "claude-3": 50},
|
||||
model_tpm_limit={"gpt-4": 2000, "claude-3": 500},
|
||||
)
|
||||
|
||||
update_response = await update_project(
|
||||
data=update_data,
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
print("Updated project response:", update_response)
|
||||
|
||||
# Assertions
|
||||
assert update_response.project_id == project_id
|
||||
assert update_response.project_alias == "test-project-updated"
|
||||
assert update_response.description == "Updated description"
|
||||
assert update_response.models == ["gpt-4", "gpt-3.5-turbo", "claude-3"]
|
||||
# model_rpm_limit and model_tpm_limit are stored in metadata
|
||||
assert update_response.metadata["use_case_id"] == "TEST-002-UPDATED"
|
||||
assert update_response.metadata["additional_field"] == "new_value"
|
||||
assert update_response.metadata["model_rpm_limit"] == {
|
||||
"gpt-4": 200,
|
||||
"claude-3": 50,
|
||||
}
|
||||
assert update_response.metadata["model_tpm_limit"] == {
|
||||
"gpt-4": 2000,
|
||||
"claude-3": 500,
|
||||
}
|
||||
assert update_response.litellm_budget_table is not None
|
||||
assert update_response.litellm_budget_table.max_budget == 200.0
|
||||
|
||||
except Exception as e:
|
||||
print("Got Exception", e)
|
||||
traceback.print_exc()
|
||||
pytest.fail(f"Got exception {e}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_project(prisma_client):
|
||||
"""
|
||||
Test deleting a project.
|
||||
"""
|
||||
try:
|
||||
print("prisma client=", prisma_client)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
# Create a team first
|
||||
_team_id = f"project-test-team_{uuid.uuid4()}"
|
||||
await new_team(
|
||||
NewTeamRequest(
|
||||
team_id=_team_id,
|
||||
),
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
# Create a project
|
||||
project_data = NewProjectRequest(
|
||||
project_alias="test-project-delete",
|
||||
team_id=_team_id,
|
||||
models=["gpt-4"],
|
||||
max_budget=50.0,
|
||||
)
|
||||
|
||||
create_response = await new_project(
|
||||
data=project_data,
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
print("Created project:", create_response)
|
||||
project_id = create_response.project_id
|
||||
|
||||
# Delete the project
|
||||
delete_data = DeleteProjectRequest(project_ids=[project_id])
|
||||
|
||||
delete_response = await delete_project(
|
||||
data=delete_data,
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
print("Delete project response:", delete_response)
|
||||
|
||||
# Assertions - delete_project returns a list of deleted project objects
|
||||
assert isinstance(delete_response, list)
|
||||
assert len(delete_response) == 1
|
||||
assert delete_response[0].project_id == project_id
|
||||
|
||||
# Try to get info on the deleted project - should fail or return None
|
||||
try:
|
||||
await project_info(
|
||||
project_id=project_id,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
pytest.fail("Expected to fail when fetching deleted project")
|
||||
except Exception as e:
|
||||
print("Expected error when fetching deleted project:", e)
|
||||
# This is expected behavior
|
||||
|
||||
except Exception as e:
|
||||
print("Got Exception", e)
|
||||
traceback.print_exc()
|
||||
pytest.fail(f"Got exception {e}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_project_info(prisma_client):
|
||||
"""
|
||||
Test getting project info.
|
||||
"""
|
||||
try:
|
||||
print("prisma client=", prisma_client)
|
||||
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
|
||||
# Create a team first
|
||||
_team_id = f"project-test-team_{uuid.uuid4()}"
|
||||
await new_team(
|
||||
NewTeamRequest(
|
||||
team_id=_team_id,
|
||||
),
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
# Create a project
|
||||
project_data = NewProjectRequest(
|
||||
project_alias="test-project-info",
|
||||
description="Test project info endpoint",
|
||||
team_id=_team_id,
|
||||
metadata={"use_case_id": "TEST-003", "cost_center": "engineering"},
|
||||
models=["gpt-4", "claude-3"],
|
||||
max_budget=150.0,
|
||||
model_rpm_limit={"gpt-4": 150},
|
||||
model_tpm_limit={"gpt-4": 1500},
|
||||
)
|
||||
|
||||
create_response = await new_project(
|
||||
data=project_data,
|
||||
http_request=Request(scope={"type": "http"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
print("Created project:", create_response)
|
||||
project_id = create_response.project_id
|
||||
|
||||
# Get project info
|
||||
info_response = await project_info(
|
||||
project_id=project_id,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="1234",
|
||||
),
|
||||
)
|
||||
|
||||
print("Project info response:", info_response)
|
||||
|
||||
# Assertions - project_info returns the project object directly
|
||||
assert info_response.project_id == project_id
|
||||
assert info_response.project_alias == "test-project-info"
|
||||
assert info_response.description == "Test project info endpoint"
|
||||
assert info_response.team_id == _team_id
|
||||
assert info_response.models == ["gpt-4", "claude-3"]
|
||||
# model_rpm_limit and model_tpm_limit are stored in metadata
|
||||
assert info_response.metadata["use_case_id"] == "TEST-003"
|
||||
assert info_response.metadata["cost_center"] == "engineering"
|
||||
assert info_response.metadata["model_rpm_limit"] == {"gpt-4": 150}
|
||||
assert info_response.metadata["model_tpm_limit"] == {"gpt-4": 1500}
|
||||
assert info_response.litellm_budget_table is not None
|
||||
assert info_response.litellm_budget_table.max_budget == 150.0
|
||||
|
||||
except Exception as e:
|
||||
print("Got Exception", e)
|
||||
traceback.print_exc()
|
||||
pytest.fail(f"Got exception {e}")
|
||||
|
||||
|
||||
### VALIDATION TESTS ###
|
||||
|
||||
|
||||
def test_check_team_project_limits_models_not_in_team():
|
||||
"""
|
||||
Test that creating a project with models not in the team raises an error.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
_check_team_project_limits,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="test-team",
|
||||
models=["gpt-4", "gpt-3.5-turbo"],
|
||||
)
|
||||
|
||||
data = NewProjectRequest(
|
||||
team_id="test-team",
|
||||
models=["gpt-4", "claude-3"], # claude-3 not in team
|
||||
)
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
_check_team_project_limits(team_object=team, data=data)
|
||||
|
||||
assert "claude-3" in str(exc_info.value.detail)
|
||||
assert "not in team's allowed models" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def test_check_team_project_limits_budget_exceeds_team():
|
||||
"""
|
||||
Test that creating a project with budget > team budget raises an error.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
_check_team_project_limits,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
max_budget=100.0,
|
||||
)
|
||||
|
||||
data = NewProjectRequest(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
max_budget=150.0, # exceeds team's 100.0
|
||||
)
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
_check_team_project_limits(team_object=team, data=data)
|
||||
|
||||
assert "exceeds team's max_budget" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def test_check_team_project_limits_valid_subset():
|
||||
"""
|
||||
Test that a valid project (models subset, budget within limit) passes.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
_check_team_project_limits,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="test-team",
|
||||
models=["gpt-4", "gpt-3.5-turbo", "claude-3"],
|
||||
max_budget=1000.0,
|
||||
)
|
||||
|
||||
data = NewProjectRequest(
|
||||
team_id="test-team",
|
||||
models=["gpt-4", "gpt-3.5-turbo"],
|
||||
max_budget=500.0,
|
||||
)
|
||||
|
||||
# Should not raise
|
||||
_check_team_project_limits(team_object=team, data=data)
|
||||
|
||||
|
||||
def test_check_team_project_limits_all_proxy_models():
|
||||
"""
|
||||
Test that team with 'all-proxy-models' allows any project models.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
_check_team_project_limits,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="test-team",
|
||||
models=["all-proxy-models"],
|
||||
)
|
||||
|
||||
data = NewProjectRequest(
|
||||
team_id="test-team",
|
||||
models=["gpt-4", "claude-3", "anything-goes"],
|
||||
)
|
||||
|
||||
# Should not raise - team allows all models
|
||||
_check_team_project_limits(team_object=team, data=data)
|
||||
|
||||
|
||||
def test_check_team_project_limits_tpm_exceeds_team():
|
||||
"""
|
||||
Test that project tpm_limit exceeding team tpm_limit raises an error.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
_check_team_project_limits,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
tpm_limit=10000,
|
||||
)
|
||||
|
||||
data = NewProjectRequest(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
tpm_limit=20000, # exceeds team's 10000
|
||||
)
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
_check_team_project_limits(team_object=team, data=data)
|
||||
|
||||
assert "exceeds team's tpm_limit" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def test_check_team_project_limits_negative_budget():
|
||||
"""
|
||||
Test that negative budget values raise an error.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
_check_team_project_limits,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
)
|
||||
|
||||
data = NewProjectRequest(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
max_budget=-10.0,
|
||||
)
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
_check_team_project_limits(team_object=team, data=data)
|
||||
|
||||
assert "cannot be negative" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def test_check_team_project_limits_soft_budget_gte_max():
|
||||
"""
|
||||
Test that soft_budget >= max_budget raises an error.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
_check_team_project_limits,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
)
|
||||
|
||||
data = NewProjectRequest(
|
||||
team_id="test-team",
|
||||
models=["gpt-4"],
|
||||
max_budget=100.0,
|
||||
soft_budget=100.0, # equal to max, should fail
|
||||
)
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
_check_team_project_limits(team_object=team, data=data)
|
||||
|
||||
assert "must be strictly lower" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def test_premium_user_gate():
|
||||
"""
|
||||
Test that project endpoints require premium_user=True.
|
||||
"""
|
||||
|
||||
# This test just validates the premium_user check exists
|
||||
# The actual endpoint test would need prisma, but we can verify
|
||||
# the import path works
|
||||
setattr(litellm.proxy.proxy_server, "premium_user", False)
|
||||
|
||||
# Verify that CommonProxyErrors.not_premium_user exists
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
|
||||
assert hasattr(CommonProxyErrors, "not_premium_user")
|
||||
|
||||
# Reset
|
||||
setattr(litellm.proxy.proxy_server, "premium_user", True)
|
||||
|
||||
|
||||
def test_project_model_access_denied_error_type():
|
||||
"""
|
||||
Test that ProxyErrorTypes.project_model_access_denied exists.
|
||||
"""
|
||||
from litellm.proxy._types import ProxyErrorTypes
|
||||
|
||||
assert hasattr(ProxyErrorTypes, "project_model_access_denied")
|
||||
assert (
|
||||
ProxyErrorTypes.project_model_access_denied.value
|
||||
== "project_model_access_denied"
|
||||
)
|
||||
|
||||
# Test the classmethod resolves correctly
|
||||
result = ProxyErrorTypes.get_model_access_error_type_for_object("project")
|
||||
assert result == ProxyErrorTypes.project_model_access_denied
|
||||
|
||||
|
||||
def test_project_cached_obj_has_last_refreshed_at():
|
||||
"""
|
||||
Test that LiteLLM_ProjectTableCachedObj has last_refreshed_at field
|
||||
matching LiteLLM_TeamTableCachedObj pattern.
|
||||
"""
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ProjectTableCachedObj,
|
||||
LiteLLM_ProjectTable,
|
||||
)
|
||||
|
||||
# Verify inheritance
|
||||
assert issubclass(LiteLLM_ProjectTableCachedObj, LiteLLM_ProjectTable)
|
||||
|
||||
# Verify last_refreshed_at field exists and defaults to None
|
||||
obj = LiteLLM_ProjectTableCachedObj(
|
||||
project_id="test",
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
)
|
||||
assert obj.last_refreshed_at is None
|
||||
|
||||
# Verify it can be set
|
||||
obj.last_refreshed_at = 1234567890.0
|
||||
assert obj.last_refreshed_at == 1234567890.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_project_max_budget_check_fires_alert():
|
||||
"""
|
||||
Test that _project_max_budget_check fires a budget alert
|
||||
when project exceeds its max budget (matches _team_max_budget_check pattern).
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _project_max_budget_check
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_ProjectTableCachedObj,
|
||||
)
|
||||
|
||||
project = LiteLLM_ProjectTableCachedObj(
|
||||
project_id="test-project",
|
||||
spend=150.0,
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(max_budget=100.0),
|
||||
)
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
user_id="user-1",
|
||||
team_id="team-1",
|
||||
)
|
||||
|
||||
mock_proxy_logging = mock.AsyncMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.budget_alerts = mock.AsyncMock()
|
||||
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
||||
await _project_max_budget_check(
|
||||
project_object=project,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
assert "Project=test-project" in str(exc_info.value)
|
||||
assert "150.0" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_project_soft_budget_check():
|
||||
"""
|
||||
Test that _project_soft_budget_check triggers alert when soft budget is exceeded.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _project_soft_budget_check
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_ProjectTableCachedObj,
|
||||
)
|
||||
|
||||
project = LiteLLM_ProjectTableCachedObj(
|
||||
project_id="test-project",
|
||||
spend=80.0,
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(soft_budget=75.0),
|
||||
)
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
user_id="user-1",
|
||||
team_id="team-1",
|
||||
)
|
||||
|
||||
mock_proxy_logging = mock.AsyncMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.budget_alerts = mock.AsyncMock()
|
||||
|
||||
# Should not raise (soft budget only alerts, doesn't block)
|
||||
await _project_soft_budget_check(
|
||||
project_object=project,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_project_soft_budget_check_no_alert_under_budget():
|
||||
"""
|
||||
Test that _project_soft_budget_check does NOT trigger alert when under soft budget.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _project_soft_budget_check
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_ProjectTableCachedObj,
|
||||
)
|
||||
|
||||
project = LiteLLM_ProjectTableCachedObj(
|
||||
project_id="test-project",
|
||||
spend=50.0,
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(soft_budget=75.0),
|
||||
)
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
user_id="user-1",
|
||||
team_id="team-1",
|
||||
)
|
||||
|
||||
mock_proxy_logging = mock.AsyncMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.budget_alerts = mock.AsyncMock()
|
||||
|
||||
# Should not raise and should not alert
|
||||
await _project_soft_budget_check(
|
||||
project_object=project,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
|
||||
def test_litellm_entity_type_has_project():
|
||||
"""
|
||||
Test that Litellm_EntityType has PROJECT member for budget alerts.
|
||||
"""
|
||||
from litellm.proxy._types import Litellm_EntityType
|
||||
|
||||
assert hasattr(Litellm_EntityType, "PROJECT")
|
||||
assert Litellm_EntityType.PROJECT.value == "project"
|
||||
|
|
@ -100,3 +100,176 @@ async def test_internal_flags_filtered_from_followup_kwargs():
|
|||
# Verify regular kwargs are preserved
|
||||
assert kwargs_for_followup["temperature"] == 0.7
|
||||
assert kwargs_for_followup["max_tokens"] == 1024
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_deployment_hook_provider_from_top_level_kwargs():
|
||||
"""Test that async_pre_call_deployment_hook finds custom_llm_provider at top-level kwargs.
|
||||
|
||||
Regression test for bug where the hook only checked kwargs["litellm_params"]["custom_llm_provider"]
|
||||
but the router places custom_llm_provider at the top level of kwargs.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
|
||||
# Simulate kwargs as they arrive from the router path:
|
||||
# custom_llm_provider is at the TOP LEVEL (not nested under litellm_params)
|
||||
kwargs = {
|
||||
"model": "anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||
"messages": [{"role": "user", "content": "Search the web for LiteLLM"}],
|
||||
"tools": [
|
||||
{"type": "web_search_20250305", "name": "web_search", "max_uses": 3},
|
||||
{"type": "function", "function": {"name": "other_tool", "parameters": {}}},
|
||||
],
|
||||
"custom_llm_provider": "bedrock",
|
||||
"api_key": "fake-key",
|
||||
}
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(kwargs=kwargs, call_type=None)
|
||||
|
||||
# Should NOT be None — the hook should have triggered
|
||||
assert result is not None
|
||||
# The web_search tool should be converted to litellm_web_search (OpenAI format)
|
||||
assert any(
|
||||
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
|
||||
for t in result["tools"]
|
||||
)
|
||||
# The non-web-search tool should be preserved
|
||||
assert any(
|
||||
t.get("type") == "function" and t.get("function", {}).get("name") == "other_tool"
|
||||
for t in result["tools"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_deployment_hook_returns_full_kwargs():
|
||||
"""Test that async_pre_call_deployment_hook returns the full kwargs dict, not a partial one.
|
||||
|
||||
Regression test for bug where the hook returned {"tools": converted_tools} instead of
|
||||
the full kwargs dict, causing model/messages/api_key/etc. to be lost.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["openai"])
|
||||
|
||||
kwargs = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "Search for something"}],
|
||||
"tools": [
|
||||
{"type": "web_search_20250305", "name": "web_search"},
|
||||
],
|
||||
"custom_llm_provider": "openai",
|
||||
"api_key": "sk-fake",
|
||||
"temperature": 0.7,
|
||||
"metadata": {"user": "test"},
|
||||
}
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(kwargs=kwargs, call_type=None)
|
||||
|
||||
assert result is not None
|
||||
# All original keys must be preserved
|
||||
assert result["model"] == "gpt-4o"
|
||||
assert result["messages"] == [{"role": "user", "content": "Search for something"}]
|
||||
assert result["api_key"] == "sk-fake"
|
||||
assert result["temperature"] == 0.7
|
||||
assert result["metadata"] == {"user": "test"}
|
||||
assert result["custom_llm_provider"] == "openai"
|
||||
# Tools should be converted
|
||||
assert any(
|
||||
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
|
||||
for t in result["tools"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_deployment_hook_skips_disabled_provider():
|
||||
"""Test that the hook returns None for providers not in enabled_providers."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
|
||||
kwargs = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"tools": [{"type": "web_search_20250305", "name": "web_search"}],
|
||||
"custom_llm_provider": "openai", # Not in enabled_providers
|
||||
}
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(kwargs=kwargs, call_type=None)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_deployment_hook_skips_no_websearch_tools():
|
||||
"""Test that the hook returns None when no web search tools are present."""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["openai"])
|
||||
|
||||
kwargs = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": "calculator", "parameters": {}}},
|
||||
],
|
||||
"custom_llm_provider": "openai",
|
||||
}
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(kwargs=kwargs, call_type=None)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_deployment_hook_nested_litellm_params_fallback():
|
||||
"""Test that the hook still works when custom_llm_provider is in nested litellm_params.
|
||||
|
||||
This is the Anthropic experimental pass-through path where litellm_params is
|
||||
explicitly constructed with custom_llm_provider inside it.
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
|
||||
|
||||
kwargs = {
|
||||
"model": "anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"tools": [{"type": "web_search_20250305", "name": "web_search"}],
|
||||
"litellm_params": {
|
||||
"custom_llm_provider": "bedrock",
|
||||
},
|
||||
}
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(kwargs=kwargs, call_type=None)
|
||||
|
||||
assert result is not None
|
||||
assert any(
|
||||
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
|
||||
for t in result["tools"]
|
||||
)
|
||||
# Full kwargs preserved
|
||||
assert result["model"] == "anthropic.claude-3-5-sonnet-20241022-v2:0"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_deployment_hook_provider_derived_from_model_name():
|
||||
"""Test that async_pre_call_deployment_hook derives custom_llm_provider from the model name.
|
||||
|
||||
Regression test for the router _acompletion path where custom_llm_provider is NOT
|
||||
in kwargs at all — neither at top-level nor in litellm_params. The hook must derive
|
||||
the provider from the model name (e.g., "openai/gpt-4o-mini" → "openai").
|
||||
"""
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["openai"])
|
||||
|
||||
# Simulate kwargs as they arrive from router._acompletion:
|
||||
# NO custom_llm_provider key anywhere — only model name contains the provider
|
||||
kwargs = {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "Search the web for LiteLLM"}],
|
||||
"tools": [
|
||||
{"type": "web_search_20250305", "name": "web_search", "max_uses": 3},
|
||||
],
|
||||
"api_key": "fake-key",
|
||||
}
|
||||
|
||||
result = await logger.async_pre_call_deployment_hook(kwargs=kwargs, call_type=None)
|
||||
|
||||
# Should NOT be None — the hook should derive "openai" from "openai/gpt-4o-mini"
|
||||
assert result is not None
|
||||
assert any(
|
||||
t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search"
|
||||
for t in result["tools"]
|
||||
)
|
||||
# Full kwargs preserved
|
||||
assert result["model"] == "openai/gpt-4o-mini"
|
||||
assert result["api_key"] == "fake-key"
|
||||
|
|
|
|||
|
|
@ -2582,6 +2582,94 @@ def test_compaction_block_with_other_content_types():
|
|||
assert tool_calls[0]["function"]["name"] == "get_weather"
|
||||
|
||||
|
||||
def test_map_openai_context_management_to_anthropic():
|
||||
"""
|
||||
Test mapping OpenAI Responses API context_management format to Anthropic format.
|
||||
"""
|
||||
config = AnthropicConfig()
|
||||
|
||||
# Test OpenAI list format with compaction
|
||||
openai_format = [{"type": "compaction", "compact_threshold": 200000}]
|
||||
result = config.map_openai_context_management_to_anthropic(openai_format)
|
||||
|
||||
assert result is not None
|
||||
assert "edits" in result
|
||||
assert len(result["edits"]) == 1
|
||||
assert result["edits"][0]["type"] == "compact_20260112"
|
||||
assert result["edits"][0]["trigger"]["type"] == "input_tokens"
|
||||
assert result["edits"][0]["trigger"]["value"] == 200000
|
||||
|
||||
# Test OpenAI format with instructions
|
||||
openai_format_with_instructions = [{
|
||||
"type": "compaction",
|
||||
"compact_threshold": 150000,
|
||||
"instructions": "Focus on preserving code snippets"
|
||||
}]
|
||||
result = config.map_openai_context_management_to_anthropic(openai_format_with_instructions)
|
||||
|
||||
assert result is not None
|
||||
assert result["edits"][0]["trigger"]["value"] == 150000
|
||||
assert result["edits"][0]["instructions"] == "Focus on preserving code snippets"
|
||||
|
||||
# Test Anthropic format (should pass through)
|
||||
anthropic_format = {
|
||||
"edits": [{
|
||||
"type": "compact_20260112",
|
||||
"trigger": {"type": "input_tokens", "value": 150000}
|
||||
}]
|
||||
}
|
||||
result = config.map_openai_context_management_to_anthropic(anthropic_format)
|
||||
|
||||
assert result == anthropic_format
|
||||
|
||||
|
||||
def test_map_openai_params_with_context_management():
|
||||
"""
|
||||
Test that map_openai_params correctly transforms context_management from OpenAI to Anthropic format.
|
||||
"""
|
||||
config = AnthropicConfig()
|
||||
|
||||
# Test with OpenAI list format
|
||||
non_default_params = {
|
||||
"context_management": [{"type": "compaction", "compact_threshold": 200000}]
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
result = config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model="claude-opus-4-6",
|
||||
drop_params=False
|
||||
)
|
||||
|
||||
assert "context_management" in result
|
||||
assert "edits" in result["context_management"]
|
||||
assert result["context_management"]["edits"][0]["type"] == "compact_20260112"
|
||||
assert result["context_management"]["edits"][0]["trigger"]["value"] == 200000
|
||||
|
||||
# Test with Anthropic dict format (should pass through)
|
||||
non_default_params_anthropic = {
|
||||
"context_management": {
|
||||
"edits": [{
|
||||
"type": "compact_20260112",
|
||||
"trigger": {"type": "input_tokens", "value": 150000},
|
||||
"instructions": "Focus on preserving code"
|
||||
}]
|
||||
}
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
result = config.map_openai_params(
|
||||
non_default_params=non_default_params_anthropic,
|
||||
optional_params=optional_params,
|
||||
model="claude-opus-4-6",
|
||||
drop_params=False
|
||||
)
|
||||
|
||||
assert "context_management" in result
|
||||
assert result["context_management"] == non_default_params_anthropic["context_management"]
|
||||
|
||||
|
||||
def test_compaction_block_empty_list_not_added():
|
||||
"""
|
||||
Test that empty compaction_blocks list is not added to provider_specific_fields.
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ sys.path.insert(
|
|||
0, os.path.abspath("../../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
|
@ -111,3 +112,35 @@ class TestDashScopeConfig:
|
|||
# Check for specific content in the response
|
||||
assert "```python" in response.choices[0].message.content
|
||||
assert "Hey from LiteLLM" in response.choices[0].message.content
|
||||
|
||||
def test_dashscope_no_longer_transforms_content_list(self):
|
||||
"""
|
||||
Test that DashScopeChatConfig does not transform content lists to strings.
|
||||
This ensures that the transformation logic specific to content lists is not applied,
|
||||
as DashScope should handle content in list format natively.
|
||||
"""
|
||||
config = DashScopeChatConfig()
|
||||
|
||||
# Create a message with content in list format
|
||||
messages: list[AllMessageValues] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Hello"},
|
||||
{"type": "text", "text": "World"},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
# Call the _transform_messages method directly
|
||||
transformed_messages = config._transform_messages(
|
||||
messages=messages, model="qwen-turbo", is_async=False
|
||||
)
|
||||
|
||||
# Verify that the content is still in list format and has not been transformed to a string
|
||||
assert isinstance(transformed_messages[0]["content"], list)
|
||||
assert len(transformed_messages[0]["content"]) == 2
|
||||
assert transformed_messages[0]["content"][0]["type"] == "text"
|
||||
assert transformed_messages[0]["content"][0]["text"] == "Hello"
|
||||
assert transformed_messages[0]["content"][1]["type"] == "text"
|
||||
assert transformed_messages[0]["content"][1]["text"] == "World"
|
||||
|
|
|
|||
|
|
@ -165,3 +165,95 @@ def test_hosted_vllm_supports_thinking():
|
|||
drop_params=False,
|
||||
)
|
||||
assert optional_params["reasoning_effort"] == "low"
|
||||
|
||||
|
||||
def test_hosted_vllm_thinking_blocks_prepended_to_assistant_content():
|
||||
"""
|
||||
Test that thinking_blocks on assistant messages are converted to content
|
||||
blocks prepended before the existing content.
|
||||
"""
|
||||
config = HostedVLLMChatConfig()
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Here is my answer.",
|
||||
"thinking_blocks": [
|
||||
{
|
||||
"type": "thinking",
|
||||
"thinking": "Let me reason about this...",
|
||||
"signature": "abc123",
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Follow up question",
|
||||
},
|
||||
]
|
||||
transformed = config.transform_request(
|
||||
model="hosted_vllm/llama-3.1-70b-instruct",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assistant_msg = transformed["messages"][1]
|
||||
assert assistant_msg["role"] == "assistant"
|
||||
assert isinstance(assistant_msg["content"], list)
|
||||
assert assistant_msg["content"][0] == {
|
||||
"type": "thinking",
|
||||
"thinking": "Let me reason about this...",
|
||||
}
|
||||
assert assistant_msg["content"][1] == {
|
||||
"type": "text",
|
||||
"text": "Here is my answer.",
|
||||
}
|
||||
assert "thinking_blocks" not in assistant_msg
|
||||
|
||||
|
||||
def test_hosted_vllm_thinking_blocks_with_list_content():
|
||||
"""
|
||||
Test thinking_blocks prepended when assistant content is already a list.
|
||||
"""
|
||||
config = HostedVLLMChatConfig()
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "Response text"}],
|
||||
"thinking_blocks": [
|
||||
{
|
||||
"type": "thinking",
|
||||
"thinking": "Step 1 reasoning",
|
||||
"signature": "sig1",
|
||||
},
|
||||
{
|
||||
"type": "thinking",
|
||||
"thinking": "Step 2 reasoning",
|
||||
"signature": "sig2",
|
||||
},
|
||||
],
|
||||
},
|
||||
]
|
||||
transformed = config.transform_request(
|
||||
model="hosted_vllm/llama-3.1-70b-instruct",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assistant_msg = transformed["messages"][0]
|
||||
assert len(assistant_msg["content"]) == 3
|
||||
assert assistant_msg["content"][0] == {
|
||||
"type": "thinking",
|
||||
"thinking": "Step 1 reasoning",
|
||||
}
|
||||
assert assistant_msg["content"][1] == {
|
||||
"type": "thinking",
|
||||
"thinking": "Step 2 reasoning",
|
||||
}
|
||||
assert assistant_msg["content"][2] == {"type": "text", "text": "Response text"}
|
||||
assert "thinking_blocks" not in assistant_msg
|
||||
|
|
|
|||
|
|
@ -1972,7 +1972,7 @@ def test_reasoning_effort_maps_to_thinking_level_gemini_3():
|
|||
model=model,
|
||||
drop_params=False,
|
||||
)
|
||||
assert result["thinkingConfig"]["thinkingLevel"] == "high"
|
||||
assert result["thinkingConfig"]["thinkingLevel"] == "medium"
|
||||
assert result["thinkingConfig"]["includeThoughts"] is True
|
||||
|
||||
# Test high -> high + includeThoughts=True
|
||||
|
|
@ -2061,7 +2061,7 @@ def test_reasoning_effort_dict_format_gemini_3():
|
|||
model=model,
|
||||
drop_params=False,
|
||||
)
|
||||
assert result["thinkingConfig"]["thinkingLevel"] == "high"
|
||||
assert result["thinkingConfig"]["thinkingLevel"] == "medium"
|
||||
assert result["thinkingConfig"]["includeThoughts"] is True
|
||||
|
||||
# Test dict format without effort key - should fall back to Gemini 3 default (low)
|
||||
|
|
|
|||
|
|
@ -587,3 +587,256 @@ class TestApplyPoliciesDirectGuardrailNames:
|
|||
# Sorted order: direct_guardrail then from_policy; final output is from_policy
|
||||
assert result["inputs"] == first_output
|
||||
assert result["guardrail_errors"] == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests for competitor enrichment helper functions
|
||||
# ---------------------------------------------------------------------------
|
||||
from litellm.proxy.management_endpoints.policy_endpoints import (
|
||||
_build_all_names_per_competitor,
|
||||
_build_comparison_blocked_words,
|
||||
_build_competitor_guardrail_definitions,
|
||||
_build_name_blocked_words,
|
||||
_build_recommendation_blocked_words,
|
||||
_build_refinement_prompt,
|
||||
_clean_competitor_line,
|
||||
_parse_variations_response,
|
||||
)
|
||||
|
||||
|
||||
class TestCleanCompetitorLine:
|
||||
"""Tests for _clean_competitor_line."""
|
||||
|
||||
def test_strips_bullets_and_dashes(self):
|
||||
assert _clean_competitor_line("- United Airlines") == "United Airlines"
|
||||
assert _clean_competitor_line(" - JetBlue ") == "JetBlue"
|
||||
|
||||
def test_strips_trailing_punctuation(self):
|
||||
assert _clean_competitor_line("Delta Airlines.") == "Delta Airlines"
|
||||
assert _clean_competitor_line("Southwest)") == "Southwest"
|
||||
|
||||
def test_returns_none_for_empty(self):
|
||||
assert _clean_competitor_line("") is None
|
||||
assert _clean_competitor_line(" ") is None
|
||||
|
||||
def test_returns_none_for_single_char(self):
|
||||
assert _clean_competitor_line("A") is None
|
||||
assert _clean_competitor_line(" - ") is None
|
||||
|
||||
def test_plain_name(self):
|
||||
assert _clean_competitor_line("Qatar Airways") == "Qatar Airways"
|
||||
|
||||
|
||||
class TestParseVariationsResponse:
|
||||
"""Tests for _parse_variations_response."""
|
||||
|
||||
def test_parses_standard_format(self):
|
||||
raw = "Delta Airlines: Delta Air Lines, DeltaAirlines, Delta\nUnited Airlines: United, UAL"
|
||||
competitors = ["Delta Airlines", "United Airlines"]
|
||||
result = _parse_variations_response(raw, competitors)
|
||||
assert "Delta Airlines" in result
|
||||
assert "Delta Air Lines" in result["Delta Airlines"]
|
||||
assert "United" in result["United Airlines"]
|
||||
|
||||
def test_case_insensitive_matching(self):
|
||||
raw = "delta airlines: Delta Air Lines, DeltaAirlines"
|
||||
competitors = ["Delta Airlines"]
|
||||
result = _parse_variations_response(raw, competitors)
|
||||
assert "Delta Airlines" in result
|
||||
assert len(result["Delta Airlines"]) == 2
|
||||
|
||||
def test_skips_lines_without_colon(self):
|
||||
raw = "This is a header\nDelta Airlines: Delta Air Lines"
|
||||
competitors = ["Delta Airlines"]
|
||||
result = _parse_variations_response(raw, competitors)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_skips_unknown_competitors(self):
|
||||
raw = "Unknown Corp: Foo, Bar\nDelta Airlines: Delta"
|
||||
competitors = ["Delta Airlines"]
|
||||
result = _parse_variations_response(raw, competitors)
|
||||
assert "Unknown Corp" not in result
|
||||
assert "Delta Airlines" in result
|
||||
|
||||
def test_filters_out_self_reference(self):
|
||||
raw = "Delta Airlines: Delta Airlines, Delta Air Lines"
|
||||
competitors = ["Delta Airlines"]
|
||||
result = _parse_variations_response(raw, competitors)
|
||||
# "Delta Airlines" should be filtered out (same as canonical)
|
||||
assert "Delta Airlines" not in result["Delta Airlines"]
|
||||
assert "Delta Air Lines" in result["Delta Airlines"]
|
||||
|
||||
def test_empty_input(self):
|
||||
assert _parse_variations_response("", []) == {}
|
||||
|
||||
|
||||
class TestBuildRefinementPrompt:
|
||||
"""Tests for _build_refinement_prompt."""
|
||||
|
||||
def test_includes_brand_name(self):
|
||||
prompt = _build_refinement_prompt("add 10 more", ["Delta"], "Emirates")
|
||||
assert "Emirates" in prompt
|
||||
|
||||
def test_includes_existing_competitors(self):
|
||||
prompt = _build_refinement_prompt("add more", ["Delta", "United"], "Emirates")
|
||||
assert "Delta" in prompt
|
||||
assert "United" in prompt
|
||||
|
||||
def test_includes_instruction(self):
|
||||
prompt = _build_refinement_prompt("add 10 from Asia", ["Delta"], "Emirates")
|
||||
assert "add 10 from Asia" in prompt
|
||||
|
||||
def test_asks_for_new_names_only(self):
|
||||
prompt = _build_refinement_prompt("add more", ["Delta"], "Emirates")
|
||||
assert "NEW" in prompt
|
||||
|
||||
|
||||
class TestBuildAllNamesPerCompetitor:
|
||||
"""Tests for _build_all_names_per_competitor."""
|
||||
|
||||
def test_includes_canonical_and_variations(self):
|
||||
result = _build_all_names_per_competitor(
|
||||
["Delta Airlines"], {"Delta Airlines": ["Delta", "DeltaAir"]}
|
||||
)
|
||||
assert result["Delta Airlines"] == ["Delta Airlines", "Delta", "DeltaAir"]
|
||||
|
||||
def test_no_variations(self):
|
||||
result = _build_all_names_per_competitor(["Delta Airlines"], {})
|
||||
assert result["Delta Airlines"] == ["Delta Airlines"]
|
||||
|
||||
def test_multiple_competitors(self):
|
||||
result = _build_all_names_per_competitor(
|
||||
["Delta", "United"],
|
||||
{"Delta": ["DL"], "United": ["UA"]},
|
||||
)
|
||||
assert len(result) == 2
|
||||
assert result["Delta"] == ["Delta", "DL"]
|
||||
assert result["United"] == ["United", "UA"]
|
||||
|
||||
|
||||
class TestBuildNameBlockedWords:
|
||||
"""Tests for _build_name_blocked_words."""
|
||||
|
||||
def test_basic_output(self):
|
||||
all_names = {"Delta": ["Delta", "DL"]}
|
||||
result = _build_name_blocked_words(["Delta"], all_names)
|
||||
keywords = [r["keyword"] for r in result]
|
||||
assert "Delta" in keywords
|
||||
assert "DL" in keywords
|
||||
assert all(r["action"] == "BLOCK" for r in result)
|
||||
|
||||
def test_descriptions_differ_for_variations(self):
|
||||
all_names = {"Delta": ["Delta", "DL"]}
|
||||
result = _build_name_blocked_words(["Delta"], all_names)
|
||||
descs = {r["keyword"]: r["description"] for r in result}
|
||||
assert "Competitor: Delta" == descs["Delta"]
|
||||
assert "variation" in descs["DL"].lower()
|
||||
|
||||
|
||||
class TestBuildRecommendationBlockedWords:
|
||||
"""Tests for _build_recommendation_blocked_words."""
|
||||
|
||||
def test_generates_prefix_combinations(self):
|
||||
all_names = {"Delta": ["Delta"]}
|
||||
result = _build_recommendation_blocked_words(["Delta"], all_names)
|
||||
keywords = [r["keyword"] for r in result]
|
||||
assert "try Delta" in keywords
|
||||
assert "use Delta" in keywords
|
||||
assert "switch to Delta" in keywords
|
||||
assert "consider Delta" in keywords
|
||||
|
||||
def test_includes_variations(self):
|
||||
all_names = {"Delta": ["Delta", "DL"]}
|
||||
result = _build_recommendation_blocked_words(["Delta"], all_names)
|
||||
keywords = [r["keyword"] for r in result]
|
||||
assert "try DL" in keywords
|
||||
|
||||
|
||||
class TestBuildComparisonBlockedWords:
|
||||
"""Tests for _build_comparison_blocked_words."""
|
||||
|
||||
def test_generates_competitor_comparisons(self):
|
||||
all_names = {"Delta": ["Delta"]}
|
||||
result = _build_comparison_blocked_words(["Delta"], all_names, "Emirates")
|
||||
keywords = [r["keyword"] for r in result]
|
||||
assert "Delta is better" in keywords
|
||||
|
||||
def test_generates_brand_comparisons_once(self):
|
||||
all_names = {"Delta": ["Delta"], "United": ["United"]}
|
||||
result = _build_comparison_blocked_words(["Delta", "United"], all_names, "Emirates")
|
||||
keywords = [r["keyword"] for r in result]
|
||||
# Brand-level entries should appear exactly once
|
||||
assert keywords.count("better than Emirates") == 1
|
||||
assert keywords.count("Emirates is worse") == 1
|
||||
|
||||
def test_includes_variation_comparisons(self):
|
||||
all_names = {"Delta": ["Delta", "DL"]}
|
||||
result = _build_comparison_blocked_words(["Delta"], all_names, "Emirates")
|
||||
keywords = [r["keyword"] for r in result]
|
||||
assert "DL is better" in keywords
|
||||
|
||||
|
||||
class TestBuildCompetitorGuardrailDefinitions:
|
||||
"""Tests for _build_competitor_guardrail_definitions."""
|
||||
|
||||
def test_populates_blocked_words_for_known_guardrail_names(self):
|
||||
definitions = [
|
||||
{
|
||||
"guardrail_name": "competitor-name-blocker",
|
||||
"litellm_params": {"blocked_words": []},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-recommendation-filter",
|
||||
"litellm_params": {"blocked_words": []},
|
||||
},
|
||||
]
|
||||
result = _build_competitor_guardrail_definitions(
|
||||
definitions, ["Delta"], "Emirates", {"Delta": ["DL"]}
|
||||
)
|
||||
# Name blocker should have entries
|
||||
name_blocker = next(d for d in result if d["guardrail_name"] == "competitor-name-blocker")
|
||||
assert len(name_blocker["litellm_params"]["blocked_words"]) > 0
|
||||
|
||||
# Recommendation filter should have entries
|
||||
rec_filter = next(d for d in result if d["guardrail_name"] == "competitor-recommendation-filter")
|
||||
assert len(rec_filter["litellm_params"]["blocked_words"]) > 0
|
||||
|
||||
def test_does_not_modify_unknown_guardrail_names(self):
|
||||
definitions = [
|
||||
{
|
||||
"guardrail_name": "some-other-guardrail",
|
||||
"litellm_params": {"blocked_words": ["original"]},
|
||||
},
|
||||
]
|
||||
result = _build_competitor_guardrail_definitions(
|
||||
definitions, ["Delta"], "Emirates"
|
||||
)
|
||||
assert result[0]["litellm_params"]["blocked_words"] == ["original"]
|
||||
|
||||
def test_does_not_mutate_original_definitions(self):
|
||||
definitions = [
|
||||
{
|
||||
"guardrail_name": "competitor-name-blocker",
|
||||
"litellm_params": {"blocked_words": []},
|
||||
},
|
||||
]
|
||||
_build_competitor_guardrail_definitions(definitions, ["Delta"], "Emirates")
|
||||
# Original should be unchanged
|
||||
assert definitions[0]["litellm_params"]["blocked_words"] == []
|
||||
|
||||
def test_handles_input_and_output_blocker_variants(self):
|
||||
definitions = [
|
||||
{
|
||||
"guardrail_name": "competitor-name-input-blocker",
|
||||
"litellm_params": {"blocked_words": []},
|
||||
},
|
||||
{
|
||||
"guardrail_name": "competitor-name-output-blocker",
|
||||
"litellm_params": {"blocked_words": []},
|
||||
},
|
||||
]
|
||||
result = _build_competitor_guardrail_definitions(
|
||||
definitions, ["Delta"], "Emirates"
|
||||
)
|
||||
for defn in result:
|
||||
assert len(defn["litellm_params"]["blocked_words"]) > 0
|
||||
|
|
|
|||
|
|
@ -5496,6 +5496,190 @@ async def test_get_team_daily_activity_team_admin_sees_all_spend(mock_db_client)
|
|||
assert False, "API keys should not be fetched for team admin users"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_daily_activity_member_with_permission_sees_all_spend(
|
||||
mock_db_client,
|
||||
):
|
||||
"""
|
||||
Test that non-admin team members with /team/daily/activity permission
|
||||
can see all team spend (no API key filtering), same as team admins.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
get_team_daily_activity,
|
||||
)
|
||||
|
||||
# Create a non-admin user
|
||||
user_id = "test_user_with_perm_123"
|
||||
team_id = "test_team_789"
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
# Mock user info
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
user_id=user_id,
|
||||
teams=[team_id],
|
||||
max_budget=1000.0,
|
||||
spend=0.0,
|
||||
user_email="member@example.com",
|
||||
user_role="internal_user",
|
||||
)
|
||||
|
||||
# Mock team with user as non-admin member AND /team/daily/activity permission
|
||||
mock_team_member = Member(user_id=user_id, role="user")
|
||||
mock_team = MagicMock(spec=LiteLLM_TeamTable)
|
||||
mock_team.team_id = team_id
|
||||
mock_team.team_alias = "Test Team"
|
||||
mock_team.members_with_roles = [mock_team_member]
|
||||
mock_team.team_member_permissions = ["/team/daily/activity"]
|
||||
mock_team.model_dump.return_value = {
|
||||
"team_id": team_id,
|
||||
"team_alias": "Test Team",
|
||||
"members_with_roles": [{"user_id": user_id, "role": "user"}],
|
||||
"team_member_permissions": ["/team/daily/activity"],
|
||||
}
|
||||
|
||||
# Setup mocks
|
||||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(
|
||||
return_value=[mock_team]
|
||||
)
|
||||
|
||||
# Mock get_user_object
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_user_object:
|
||||
mock_get_user_object.return_value = mock_user_info
|
||||
|
||||
# Mock get_daily_activity to capture the api_key parameter
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_daily_activity:
|
||||
mock_get_daily_activity.return_value = MagicMock()
|
||||
|
||||
# Call the endpoint
|
||||
await get_team_daily_activity(
|
||||
team_ids=team_id,
|
||||
start_date="2024-01-01",
|
||||
end_date="2024-01-02",
|
||||
model=None,
|
||||
api_key=None,
|
||||
page=1,
|
||||
page_size=10,
|
||||
exclude_team_ids=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify get_daily_activity was called WITHOUT API key filtering
|
||||
mock_get_daily_activity.assert_called_once()
|
||||
call_kwargs = mock_get_daily_activity.call_args[1]
|
||||
assert call_kwargs["api_key"] is None
|
||||
assert call_kwargs["entity_id"] == [team_id]
|
||||
|
||||
# Verify user's API keys were NOT fetched
|
||||
if hasattr(
|
||||
mock_db_client.db.litellm_verificationtoken, "find_many"
|
||||
) and mock_db_client.db.litellm_verificationtoken.find_many.called:
|
||||
assert (
|
||||
False
|
||||
), "API keys should not be fetched for members with /team/daily/activity permission"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_daily_activity_member_without_permission_filters_by_keys(
|
||||
mock_db_client,
|
||||
):
|
||||
"""
|
||||
Test that non-admin team members WITHOUT /team/daily/activity permission
|
||||
still have their results filtered by their own API keys.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
get_team_daily_activity,
|
||||
)
|
||||
|
||||
# Create a non-admin user
|
||||
user_id = "test_user_no_perm_123"
|
||||
team_id = "test_team_789"
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
# Mock user info
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
user_id=user_id,
|
||||
teams=[team_id],
|
||||
max_budget=1000.0,
|
||||
spend=0.0,
|
||||
user_email="member@example.com",
|
||||
user_role="internal_user",
|
||||
)
|
||||
|
||||
# Mock team with user as non-admin member and NO usage permission
|
||||
mock_team_member = Member(user_id=user_id, role="user")
|
||||
mock_team = MagicMock(spec=LiteLLM_TeamTable)
|
||||
mock_team.team_id = team_id
|
||||
mock_team.team_alias = "Test Team"
|
||||
mock_team.members_with_roles = [mock_team_member]
|
||||
mock_team.team_member_permissions = ["/key/info"]
|
||||
mock_team.model_dump.return_value = {
|
||||
"team_id": team_id,
|
||||
"team_alias": "Test Team",
|
||||
"members_with_roles": [{"user_id": user_id, "role": "user"}],
|
||||
"team_member_permissions": ["/key/info"],
|
||||
}
|
||||
|
||||
# Mock user's API keys
|
||||
user_api_key_1 = MagicMock()
|
||||
user_api_key_1.token = "user_key_abc"
|
||||
user_api_key_2 = MagicMock()
|
||||
user_api_key_2.token = "user_key_def"
|
||||
|
||||
# Setup mocks
|
||||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(
|
||||
return_value=[mock_team]
|
||||
)
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[user_api_key_1, user_api_key_2]
|
||||
)
|
||||
|
||||
# Mock get_user_object
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_user_object:
|
||||
mock_get_user_object.return_value = mock_user_info
|
||||
|
||||
# Mock get_daily_activity to capture the api_key parameter
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_daily_activity",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_daily_activity:
|
||||
mock_get_daily_activity.return_value = MagicMock()
|
||||
|
||||
# Call the endpoint
|
||||
await get_team_daily_activity(
|
||||
team_ids=team_id,
|
||||
start_date="2024-01-01",
|
||||
end_date="2024-01-02",
|
||||
model=None,
|
||||
api_key=None,
|
||||
page=1,
|
||||
page_size=10,
|
||||
exclude_team_ids=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify get_daily_activity was called WITH API key filtering
|
||||
mock_get_daily_activity.assert_called_once()
|
||||
call_kwargs = mock_get_daily_activity.call_args[1]
|
||||
assert call_kwargs["api_key"] == ["user_key_abc", "user_key_def"]
|
||||
assert call_kwargs["entity_id"] == [team_id]
|
||||
|
||||
# Verify user's API keys were fetched
|
||||
mock_db_client.db.litellm_verificationtoken.find_many.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,155 @@
|
|||
"""
|
||||
Test method-specific routing for pass-through endpoints.
|
||||
|
||||
This test demonstrates the ability to configure different targets
|
||||
for the same path but different HTTP methods.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._types import PassThroughGenericEndpoint
|
||||
|
||||
|
||||
def test_pass_through_endpoint_with_methods():
|
||||
"""Test creating pass-through endpoints with specific methods"""
|
||||
|
||||
# Create endpoint for GET /azure/kb
|
||||
get_endpoint = PassThroughGenericEndpoint(
|
||||
id="get-azure-kb",
|
||||
path="/azure/kb",
|
||||
target="https://api1.example.com/knowledge-base",
|
||||
methods=["GET"],
|
||||
headers={"Authorization": "Bearer token1"},
|
||||
)
|
||||
|
||||
assert get_endpoint.path == "/azure/kb"
|
||||
assert get_endpoint.methods == ["GET"]
|
||||
assert get_endpoint.target == "https://api1.example.com/knowledge-base"
|
||||
|
||||
# Create endpoint for POST /azure/kb
|
||||
post_endpoint = PassThroughGenericEndpoint(
|
||||
id="post-azure-kb",
|
||||
path="/azure/kb",
|
||||
target="https://api2.example.com/knowledge-base",
|
||||
methods=["POST"],
|
||||
headers={"Authorization": "Bearer token2"},
|
||||
)
|
||||
|
||||
assert post_endpoint.path == "/azure/kb"
|
||||
assert post_endpoint.methods == ["POST"]
|
||||
assert post_endpoint.target == "https://api2.example.com/knowledge-base"
|
||||
|
||||
# These should be different endpoints despite same path
|
||||
assert get_endpoint.id != post_endpoint.id
|
||||
assert get_endpoint.target != post_endpoint.target
|
||||
|
||||
|
||||
def test_pass_through_endpoint_multiple_methods():
|
||||
"""Test creating endpoint with multiple methods"""
|
||||
|
||||
endpoint = PassThroughGenericEndpoint(
|
||||
id="multi-method",
|
||||
path="/azure/kb",
|
||||
target="https://api.example.com/kb",
|
||||
methods=["GET", "POST", "PUT"],
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert len(endpoint.methods) == 3
|
||||
assert "GET" in endpoint.methods
|
||||
assert "POST" in endpoint.methods
|
||||
assert "PUT" in endpoint.methods
|
||||
|
||||
|
||||
def test_pass_through_endpoint_no_methods_backward_compatibility():
|
||||
"""Test that endpoints without methods field work (backward compatibility)"""
|
||||
|
||||
# When methods is None, all methods should be supported
|
||||
endpoint = PassThroughGenericEndpoint(
|
||||
id="all-methods",
|
||||
path="/azure/kb",
|
||||
target="https://api.example.com/kb",
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert endpoint.methods is None # Default is None for backward compatibility
|
||||
|
||||
|
||||
def test_pass_through_endpoint_serialization():
|
||||
"""Test that endpoints with methods can be serialized/deserialized"""
|
||||
|
||||
endpoint = PassThroughGenericEndpoint(
|
||||
id="test-endpoint",
|
||||
path="/test",
|
||||
target="https://api.example.com",
|
||||
methods=["GET", "POST"],
|
||||
headers={"key": "value"},
|
||||
cost_per_request=0.5,
|
||||
)
|
||||
|
||||
# Serialize to dict
|
||||
endpoint_dict = endpoint.model_dump()
|
||||
assert endpoint_dict["methods"] == ["GET", "POST"]
|
||||
|
||||
# Deserialize from dict
|
||||
restored_endpoint = PassThroughGenericEndpoint(**endpoint_dict)
|
||||
assert restored_endpoint.methods == ["GET", "POST"]
|
||||
assert restored_endpoint.path == "/test"
|
||||
assert restored_endpoint.target == "https://api.example.com"
|
||||
|
||||
|
||||
def test_route_key_generation_with_methods():
|
||||
"""Test that route keys include methods for uniqueness"""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
)
|
||||
|
||||
# Simulate how route keys are generated
|
||||
endpoint_id_1 = "endpoint-1"
|
||||
path = "/azure/kb"
|
||||
methods_1 = ["GET"]
|
||||
methods_str_1 = ",".join(sorted(methods_1))
|
||||
route_key_1 = f"{endpoint_id_1}:exact:{path}:{methods_str_1}"
|
||||
|
||||
endpoint_id_2 = "endpoint-2"
|
||||
methods_2 = ["POST"]
|
||||
methods_str_2 = ",".join(sorted(methods_2))
|
||||
route_key_2 = f"{endpoint_id_2}:exact:{path}:{methods_str_2}"
|
||||
|
||||
# Keys should be different even though path is the same
|
||||
assert route_key_1 != route_key_2
|
||||
assert route_key_1 == "endpoint-1:exact:/azure/kb:GET"
|
||||
assert route_key_2 == "endpoint-2:exact:/azure/kb:POST"
|
||||
|
||||
|
||||
def test_config_yaml_example():
|
||||
"""
|
||||
Example configuration for config.yaml showing method-specific routing:
|
||||
|
||||
general_settings:
|
||||
pass_through_endpoints:
|
||||
# GET endpoint for retrieving knowledge base
|
||||
- id: "get-azure-kb"
|
||||
path: "/azure/kb"
|
||||
target: "https://read-api.example.com/kb"
|
||||
methods: ["GET"]
|
||||
headers:
|
||||
Authorization: "bearer os.environ/READ_API_KEY"
|
||||
|
||||
# POST endpoint for creating knowledge base entries
|
||||
- id: "post-azure-kb"
|
||||
path: "/azure/kb"
|
||||
target: "https://write-api.example.com/kb"
|
||||
methods: ["POST"]
|
||||
headers:
|
||||
Authorization: "bearer os.environ/WRITE_API_KEY"
|
||||
|
||||
# PUT endpoint for updating knowledge base
|
||||
- id: "put-azure-kb"
|
||||
path: "/azure/kb"
|
||||
target: "https://update-api.example.com/kb"
|
||||
methods: ["PUT"]
|
||||
headers:
|
||||
Authorization: "bearer os.environ/UPDATE_API_KEY"
|
||||
"""
|
||||
pass
|
||||
|
|
@ -56,6 +56,55 @@ def test_check_provider_match_azure_ai_allows_openai_and_azure():
|
|||
) is False
|
||||
|
||||
|
||||
def test_check_provider_match_github_allows_upstream_provider_metadata():
|
||||
"""
|
||||
Test that github provider can match upstream provider metadata.
|
||||
GitHub Models can provide models from multiple providers.
|
||||
"""
|
||||
assert _check_provider_match(
|
||||
model_info={"litellm_provider": "openai"},
|
||||
custom_llm_provider="github",
|
||||
) is True
|
||||
|
||||
assert _check_provider_match(
|
||||
model_info={"litellm_provider": "github"},
|
||||
custom_llm_provider="github",
|
||||
) is True
|
||||
|
||||
assert _check_provider_match(
|
||||
model_info={"litellm_provider": "anthropic"},
|
||||
custom_llm_provider="github",
|
||||
) is True
|
||||
|
||||
|
||||
def test_supports_function_calling_github_openai_alias():
|
||||
assert litellm.utils.supports_function_calling(model="github/gpt-4o-mini") is True
|
||||
assert (
|
||||
litellm.utils.supports_function_calling(
|
||||
model="gpt-4o-mini", custom_llm_provider="github"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_supports_function_calling_github_anthropic_alias():
|
||||
assert (
|
||||
litellm.utils.supports_function_calling(
|
||||
model="github/claude-3-5-sonnet-latest"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_supports_function_calling_unknown_github_alias_returns_false():
|
||||
assert (
|
||||
litellm.utils.supports_function_calling(
|
||||
model="github/non-existent-model-for-capability-check"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_get_optional_params_image_gen():
|
||||
from litellm.llms.azure.image_generation import AzureGPTImageGenerationConfig
|
||||
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ import {
|
|||
ApiOutlined,
|
||||
} from "@ant-design/icons";
|
||||
import KeyValueInput from "./key_value_input";
|
||||
import QueryParamInput from "./query_param_input";
|
||||
import { passThroughItem } from "./pass_through_settings";
|
||||
import RoutePreview from "./route_preview";
|
||||
import NotificationsManager from "./molecules/notifications_manager";
|
||||
|
|
@ -30,6 +31,8 @@ import PassThroughSecuritySection from "./common_components/PassThroughSecurityS
|
|||
import PassThroughGuardrailsSection from "./common_components/PassThroughGuardrailsSection";
|
||||
const { Option } = Select2;
|
||||
|
||||
const HTTP_METHODS = ["GET", "POST", "PUT", "DELETE", "PATCH"];
|
||||
|
||||
interface AddFallbacksProps {
|
||||
// models: string[] | undefined;
|
||||
accessToken: string;
|
||||
|
|
@ -52,12 +55,14 @@ const AddPassThroughEndpoint: React.FC<AddFallbacksProps> = ({
|
|||
const [targetValue, setTargetValue] = useState("");
|
||||
const [includeSubpath, setIncludeSubpath] = useState(true);
|
||||
const [authEnabled, setAuthEnabled] = useState(false);
|
||||
const [selectedMethods, setSelectedMethods] = useState<string[]>([]);
|
||||
const [guardrails, setGuardrails] = useState<Record<string, { request_fields?: string[]; response_fields?: string[] } | null>>({});
|
||||
const handleCancel = () => {
|
||||
form.resetFields();
|
||||
setPathValue("");
|
||||
setTargetValue("");
|
||||
setIncludeSubpath(true);
|
||||
setSelectedMethods([]);
|
||||
setGuardrails({});
|
||||
setIsModalVisible(false);
|
||||
};
|
||||
|
|
@ -86,6 +91,11 @@ const AddPassThroughEndpoint: React.FC<AddFallbacksProps> = ({
|
|||
formValues.guardrails = guardrails;
|
||||
}
|
||||
|
||||
// Add methods to formValues (only if specific methods are selected)
|
||||
if (selectedMethods && selectedMethods.length > 0) {
|
||||
formValues.methods = selectedMethods;
|
||||
}
|
||||
|
||||
console.log(`formValues: ${JSON.stringify(formValues)}`);
|
||||
|
||||
const response = await createPassThroughEndpoint(accessToken, formValues);
|
||||
|
|
@ -101,6 +111,7 @@ const AddPassThroughEndpoint: React.FC<AddFallbacksProps> = ({
|
|||
setPathValue("");
|
||||
setTargetValue("");
|
||||
setIncludeSubpath(true);
|
||||
setSelectedMethods([]);
|
||||
setGuardrails({});
|
||||
setIsModalVisible(false);
|
||||
} catch (error) {
|
||||
|
|
@ -204,6 +215,41 @@ const AddPassThroughEndpoint: React.FC<AddFallbacksProps> = ({
|
|||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
HTTP Methods (Optional)
|
||||
<Tooltip title="Select specific HTTP methods. Leave empty to support all methods (GET, POST, PUT, DELETE, PATCH). Useful when the same path needs different targets for different methods.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="methods"
|
||||
extra={
|
||||
<div className="text-xs text-gray-500 mt-1">
|
||||
{selectedMethods.length === 0
|
||||
? "All HTTP methods supported (default)"
|
||||
: `Only ${selectedMethods.join(", ")} requests will be routed to this endpoint`}
|
||||
</div>
|
||||
}
|
||||
className="mb-4"
|
||||
>
|
||||
<Select2
|
||||
mode="multiple"
|
||||
placeholder="Select methods (leave empty for all)"
|
||||
value={selectedMethods}
|
||||
onChange={setSelectedMethods}
|
||||
allowClear
|
||||
style={{ width: "100%" }}
|
||||
>
|
||||
{HTTP_METHODS.map((method) => (
|
||||
<Option key={method} value={method}>
|
||||
{method}
|
||||
</Option>
|
||||
))}
|
||||
</Select2>
|
||||
</Form.Item>
|
||||
|
||||
<div className="flex items-center justify-between py-3">
|
||||
<div>
|
||||
<div className="text-sm font-medium text-gray-700">Include Subpaths</div>
|
||||
|
|
@ -250,6 +296,34 @@ const AddPassThroughEndpoint: React.FC<AddFallbacksProps> = ({
|
|||
</Form.Item>
|
||||
</Card>
|
||||
|
||||
{/* Default Query Parameters Section */}
|
||||
<Card className="p-6">
|
||||
<Title className="text-lg font-semibold text-gray-900 mb-2">Default Query Parameters</Title>
|
||||
<Subtitle className="text-gray-600 mb-6">
|
||||
Add query parameters that will be automatically sent with every request to the target API
|
||||
</Subtitle>
|
||||
|
||||
<Form.Item
|
||||
label={
|
||||
<span className="text-sm font-medium text-gray-700 flex items-center">
|
||||
Default Query Parameters (Optional)
|
||||
<Tooltip title="Query parameters that will be added to all requests. Clients can override these by providing their own values.">
|
||||
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="default_query_params"
|
||||
extra={
|
||||
<div className="text-xs text-gray-500 mt-2">
|
||||
<div className="font-medium mb-1">Parameters are sent with all GET, POST, PUT, PATCH requests</div>
|
||||
<div>Client parameters override defaults. Examples: version=v1, format=json, key=default</div>
|
||||
</div>
|
||||
}
|
||||
>
|
||||
<QueryParamInput />
|
||||
</Form.Item>
|
||||
</Card>
|
||||
|
||||
{/* Security Section */}
|
||||
<PassThroughSecuritySection
|
||||
premiumUser={premiumUser}
|
||||
|
|
|
|||
|
|
@ -12,6 +12,11 @@ interface PassThroughRoutesSelectorProps {
|
|||
teamId?: string | null;
|
||||
}
|
||||
|
||||
interface PassThroughEndpoint {
|
||||
path: string;
|
||||
methods?: string[];
|
||||
}
|
||||
|
||||
const PassThroughRoutesSelector: React.FC<PassThroughRoutesSelectorProps> = ({
|
||||
onChange,
|
||||
value,
|
||||
|
|
@ -21,7 +26,7 @@ const PassThroughRoutesSelector: React.FC<PassThroughRoutesSelectorProps> = ({
|
|||
disabled = false,
|
||||
teamId,
|
||||
}) => {
|
||||
const [passThroughRoutes, setPassThroughRoutes] = useState<string[]>([]);
|
||||
const [passThroughRoutes, setPassThroughRoutes] = useState<Array<{ label: string; value: string }>>([]);
|
||||
const [loading, setLoading] = useState(false);
|
||||
|
||||
useEffect(() => {
|
||||
|
|
@ -32,7 +37,24 @@ const PassThroughRoutesSelector: React.FC<PassThroughRoutesSelectorProps> = ({
|
|||
try {
|
||||
const response = await getPassThroughEndpointsCall(accessToken, teamId);
|
||||
if (response.endpoints) {
|
||||
const routes = response.endpoints.map((route: { path: string }) => route.path);
|
||||
const routes = response.endpoints.flatMap((endpoint: PassThroughEndpoint) => {
|
||||
const path = endpoint.path;
|
||||
const methods = endpoint.methods;
|
||||
|
||||
// If methods are specified, create one entry per method
|
||||
if (methods && methods.length > 0) {
|
||||
return methods.map((method) => ({
|
||||
label: `${method} ${path}`,
|
||||
value: path, // Keep value as path for backward compatibility
|
||||
}));
|
||||
}
|
||||
|
||||
// If no methods specified, show just the path (all methods supported)
|
||||
return [{
|
||||
label: path,
|
||||
value: path,
|
||||
}];
|
||||
});
|
||||
setPassThroughRoutes(routes);
|
||||
}
|
||||
} catch (error) {
|
||||
|
|
@ -54,10 +76,7 @@ const PassThroughRoutesSelector: React.FC<PassThroughRoutesSelectorProps> = ({
|
|||
loading={loading}
|
||||
className={className}
|
||||
allowClear
|
||||
options={passThroughRoutes.map((route) => ({
|
||||
label: route,
|
||||
value: route,
|
||||
}))}
|
||||
options={passThroughRoutes}
|
||||
optionFilterProp="label"
|
||||
showSearch
|
||||
style={{ width: "100%" }}
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import {
|
|||
Text,
|
||||
Button,
|
||||
Icon,
|
||||
Switch,
|
||||
} from "@tremor/react";
|
||||
import { TabPanel, TabPanels, TabGroup, TabList, Tab } from "@tremor/react";
|
||||
import {
|
||||
|
|
@ -163,7 +164,12 @@ const GeneralSettings: React.FC<GeneralSettingsPageProps> = ({ accessToken, user
|
|||
<InputNumber
|
||||
step={1}
|
||||
value={value.field_value}
|
||||
onChange={(newValue) => handleInputChange(value.field_name, newValue)} // Handle value change
|
||||
onChange={(newValue) => handleInputChange(value.field_name, newValue)}
|
||||
/>
|
||||
) : value.field_type == "Boolean" ? (
|
||||
<Switch
|
||||
checked={value.field_value === true || value.field_value === "true"}
|
||||
onChange={(checked) => handleInputChange(value.field_name, checked)}
|
||||
/>
|
||||
) : null}
|
||||
</TableCell>
|
||||
|
|
|
|||
|
|
@ -5555,19 +5555,24 @@ export const getPolicyTemplates = async (accessToken: string) => {
|
|||
export const enrichPolicyTemplate = async (
|
||||
accessToken: string,
|
||||
templateId: string,
|
||||
parameters: Record<string, string>
|
||||
parameters: Record<string, string>,
|
||||
model?: string,
|
||||
competitors?: string[]
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/policy/templates/enrich`
|
||||
: `/policy/templates/enrich`;
|
||||
const body: any = { template_id: templateId, parameters };
|
||||
if (model) body.model = model;
|
||||
if (competitors) body.competitors = competitors;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify({ template_id: templateId, parameters }),
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
|
|
@ -5585,6 +5590,78 @@ export const enrichPolicyTemplate = async (
|
|||
}
|
||||
};
|
||||
|
||||
export const enrichPolicyTemplateStream = async (
|
||||
accessToken: string,
|
||||
templateId: string,
|
||||
parameters: Record<string, string>,
|
||||
model: string,
|
||||
onCompetitor: (name: string) => void,
|
||||
onDone: (result: {
|
||||
competitors: string[];
|
||||
competitor_variations: Record<string, string[]>;
|
||||
guardrailDefinitions: any[];
|
||||
}) => void,
|
||||
onError?: (error: string) => void,
|
||||
options?: { instruction?: string; existingCompetitors?: string[] },
|
||||
onStatus?: (message: string) => void
|
||||
) => {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/policy/templates/enrich/stream`
|
||||
: `/policy/templates/enrich/stream`;
|
||||
const body: any = { template_id: templateId, parameters, model };
|
||||
if (options?.instruction) body.instruction = options.instruction;
|
||||
if (options?.existingCompetitors) body.competitors = options.existingCompetitors;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
const reader = response.body?.getReader();
|
||||
if (!reader) throw new Error("No response body");
|
||||
|
||||
const decoder = new TextDecoder();
|
||||
let buffer = "";
|
||||
|
||||
while (true) {
|
||||
const { done, value } = await reader.read();
|
||||
if (done) break;
|
||||
|
||||
buffer += decoder.decode(value, { stream: true });
|
||||
const lines = buffer.split("\n");
|
||||
buffer = lines.pop() || "";
|
||||
|
||||
for (const line of lines) {
|
||||
if (!line.startsWith("data: ")) continue;
|
||||
try {
|
||||
const event = JSON.parse(line.slice(6));
|
||||
if (event.type === "competitor") {
|
||||
onCompetitor(event.name);
|
||||
} else if (event.type === "status") {
|
||||
onStatus?.(event.message);
|
||||
} else if (event.type === "done") {
|
||||
onDone(event);
|
||||
} else if (event.type === "error") {
|
||||
onError?.(event.message);
|
||||
}
|
||||
} catch {
|
||||
// skip malformed events
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
export const createPolicyCall = async (accessToken: string, policyData: any) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/policies` : `/policies`;
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import {
|
|||
TabPanels,
|
||||
TextInput,
|
||||
} from "@tremor/react";
|
||||
import { Button, Form, Input, Switch, InputNumber } from "antd";
|
||||
import { Button, Form, Input, Switch, InputNumber, Select } from "antd";
|
||||
import { updatePassThroughEndpoint, deletePassThroughEndpointsCall } from "./networking";
|
||||
import { Eye, EyeOff } from "lucide-react";
|
||||
import RoutePreview from "./route_preview";
|
||||
|
|
@ -21,6 +21,9 @@ import NotificationsManager from "./molecules/notifications_manager";
|
|||
import PassThroughSecuritySection from "./common_components/PassThroughSecuritySection";
|
||||
import PassThroughGuardrailsSection from "./common_components/PassThroughGuardrailsSection";
|
||||
|
||||
const HTTP_METHODS = ["GET", "POST", "PUT", "DELETE", "PATCH"];
|
||||
const { Option } = Select;
|
||||
|
||||
export interface PassThroughInfoProps {
|
||||
endpointData: PassThroughEndpoint;
|
||||
onClose: () => void;
|
||||
|
|
@ -38,6 +41,7 @@ interface PassThroughEndpoint {
|
|||
include_subpath?: boolean;
|
||||
cost_per_request?: number;
|
||||
auth?: boolean;
|
||||
methods?: string[];
|
||||
guardrails?: Record<string, { request_fields?: string[]; response_fields?: string[] } | null>;
|
||||
}
|
||||
|
||||
|
|
@ -70,6 +74,7 @@ const PassThroughInfoView: React.FC<PassThroughInfoProps> = ({
|
|||
const [loading, setLoading] = useState(false);
|
||||
const [isEditing, setIsEditing] = useState(false);
|
||||
const [authEnabled, setAuthEnabled] = useState(initialEndpointData?.auth || false);
|
||||
const [selectedMethods, setSelectedMethods] = useState<string[]>(initialEndpointData?.methods || []);
|
||||
const [guardrails, setGuardrails] = useState<Record<string, { request_fields?: string[]; response_fields?: string[] } | null>>(
|
||||
initialEndpointData?.guardrails || {}
|
||||
);
|
||||
|
|
@ -97,6 +102,7 @@ const PassThroughInfoView: React.FC<PassThroughInfoProps> = ({
|
|||
include_subpath: values.include_subpath,
|
||||
cost_per_request: values.cost_per_request,
|
||||
auth: premiumUser ? values.auth : undefined,
|
||||
methods: selectedMethods && selectedMethods.length > 0 ? selectedMethods : undefined,
|
||||
guardrails: guardrails && Object.keys(guardrails).length > 0 ? guardrails : undefined,
|
||||
};
|
||||
|
||||
|
|
@ -191,6 +197,23 @@ const PassThroughInfoView: React.FC<PassThroughInfoProps> = ({
|
|||
{endpointData.auth ? "Auth Required" : "No Auth"}
|
||||
</Badge>
|
||||
</div>
|
||||
{endpointData.methods && endpointData.methods.length > 0 && (
|
||||
<div>
|
||||
<Text className="text-xs text-gray-500">HTTP Methods:</Text>
|
||||
<div className="flex flex-wrap gap-1 mt-1">
|
||||
{endpointData.methods.map((method) => (
|
||||
<Badge key={method} color="indigo" size="sm">
|
||||
{method}
|
||||
</Badge>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
{(!endpointData.methods || endpointData.methods.length === 0) && (
|
||||
<div>
|
||||
<Text className="text-xs text-gray-500">All HTTP methods supported</Text>
|
||||
</div>
|
||||
)}
|
||||
{endpointData.cost_per_request !== undefined && (
|
||||
<div>
|
||||
<Text>Cost per request: ${endpointData.cost_per_request}</Text>
|
||||
|
|
@ -277,6 +300,7 @@ const PassThroughInfoView: React.FC<PassThroughInfoProps> = ({
|
|||
include_subpath: endpointData.include_subpath || false,
|
||||
cost_per_request: endpointData.cost_per_request,
|
||||
auth: endpointData.auth || false,
|
||||
methods: endpointData.methods || [],
|
||||
}}
|
||||
layout="vertical"
|
||||
>
|
||||
|
|
@ -295,6 +319,31 @@ const PassThroughInfoView: React.FC<PassThroughInfoProps> = ({
|
|||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="HTTP Methods (Optional)"
|
||||
name="methods"
|
||||
extra={
|
||||
selectedMethods.length === 0
|
||||
? "All HTTP methods supported (default)"
|
||||
: `Only ${selectedMethods.join(", ")} requests will be routed to this endpoint`
|
||||
}
|
||||
>
|
||||
<Select
|
||||
mode="multiple"
|
||||
placeholder="Select methods (leave empty for all)"
|
||||
value={selectedMethods}
|
||||
onChange={setSelectedMethods}
|
||||
allowClear
|
||||
style={{ width: "100%" }}
|
||||
>
|
||||
{HTTP_METHODS.map((method) => (
|
||||
<Option key={method} value={method}>
|
||||
{method}
|
||||
</Option>
|
||||
))}
|
||||
</Select>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item label="Include Subpath" name="include_subpath" valuePropName="checked">
|
||||
<Switch />
|
||||
</Form.Item>
|
||||
|
|
|
|||
|
|
@ -39,7 +39,9 @@ export interface passThroughItem {
|
|||
include_subpath?: boolean;
|
||||
cost_per_request?: number;
|
||||
auth?: boolean;
|
||||
methods?: string[];
|
||||
guardrails?: Record<string, { request_fields?: string[]; response_fields?: string[] } | null>;
|
||||
default_query_params?: Record<string, string>;
|
||||
}
|
||||
|
||||
// Password field component for headers
|
||||
|
|
@ -147,6 +149,32 @@ const PassThroughSettings: React.FC<GeneralSettingsPageProps> = ({ accessToken,
|
|||
accessorKey: "target",
|
||||
cell: (info: any) => <Text>{info.getValue()}</Text>,
|
||||
},
|
||||
{
|
||||
header: () => (
|
||||
<div className="flex items-center gap-1">
|
||||
<span>Methods</span>
|
||||
<Tooltip title="HTTP methods supported by this endpoint">
|
||||
<InformationCircleIcon className="w-4 h-4 text-gray-400 cursor-help" />
|
||||
</Tooltip>
|
||||
</div>
|
||||
),
|
||||
accessorKey: "methods",
|
||||
cell: (info: any) => {
|
||||
const methods = info.getValue();
|
||||
if (!methods || methods.length === 0) {
|
||||
return <Badge color="blue">ALL</Badge>;
|
||||
}
|
||||
return (
|
||||
<div className="flex flex-wrap gap-1">
|
||||
{methods.map((method: string) => (
|
||||
<Badge key={method} color="indigo" className="text-xs">
|
||||
{method}
|
||||
</Badge>
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
},
|
||||
},
|
||||
{
|
||||
header: () => (
|
||||
<div className="flex items-center gap-1">
|
||||
|
|
|
|||
|
|
@ -222,6 +222,31 @@ const GuardrailSelectionModal: React.FC<GuardrailSelectionModalProps> = ({
|
|||
</div>
|
||||
)}
|
||||
|
||||
{/* Discovered Competitors */}
|
||||
{template?.discoveredCompetitors?.length > 0 && (
|
||||
<>
|
||||
<Divider />
|
||||
<div className="p-3 bg-purple-50 rounded-lg border border-purple-100">
|
||||
<div className="flex items-center gap-2 mb-2">
|
||||
<span className="text-lg">✨</span>
|
||||
<span className="font-medium text-purple-900 text-sm">
|
||||
AI-Discovered Competitors ({template.discoveredCompetitors.length})
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex flex-wrap gap-1.5">
|
||||
{template.discoveredCompetitors.map((name: string) => (
|
||||
<Tag key={name} color="purple" className="text-xs">
|
||||
{name}
|
||||
</Tag>
|
||||
))}
|
||||
</div>
|
||||
<p className="text-xs text-purple-600 mt-2">
|
||||
These competitor names will be automatically blocked by the competitor-name-blocker guardrail.
|
||||
</p>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
|
||||
<Divider />
|
||||
|
||||
{/* Selected Summary */}
|
||||
|
|
|
|||
|
|
@ -228,7 +228,10 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
return JSON.parse(templateStr);
|
||||
};
|
||||
|
||||
const handleParameterConfirm = async (parameters: Record<string, string>) => {
|
||||
const handleParameterConfirm = async (
|
||||
parameters: Record<string, string>,
|
||||
enrichmentOptions?: { model?: string; competitors?: string[] }
|
||||
) => {
|
||||
if (!accessToken || !pendingTemplate) return;
|
||||
|
||||
setIsEnrichingTemplate(true);
|
||||
|
|
@ -237,14 +240,20 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
let enrichedTemplate = pendingTemplate;
|
||||
|
||||
if (pendingTemplate.llm_enrichment) {
|
||||
// Call backend to enrich template with LLM-discovered data
|
||||
// Call backend to enrich template with LLM-discovered data (or user-provided competitors)
|
||||
const enrichResult = await enrichPolicyTemplate(
|
||||
accessToken,
|
||||
pendingTemplate.id,
|
||||
parameters
|
||||
parameters,
|
||||
enrichmentOptions?.model,
|
||||
enrichmentOptions?.competitors
|
||||
);
|
||||
// The backend returns the enriched guardrailDefinitions
|
||||
enrichedTemplate = { ...pendingTemplate, guardrailDefinitions: enrichResult.guardrailDefinitions };
|
||||
// The backend returns the enriched guardrailDefinitions + discovered competitors
|
||||
enrichedTemplate = {
|
||||
...pendingTemplate,
|
||||
guardrailDefinitions: enrichResult.guardrailDefinitions,
|
||||
discoveredCompetitors: enrichResult.competitors || [],
|
||||
};
|
||||
}
|
||||
|
||||
// Substitute parameters in template
|
||||
|
|
@ -491,6 +500,7 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
onConfirm={handleParameterConfirm}
|
||||
onCancel={handleParameterCancel}
|
||||
isLoading={isEnrichingTemplate}
|
||||
accessToken={accessToken || ""}
|
||||
/>
|
||||
</TabPanel>
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import React, { useState, useEffect } from "react";
|
||||
import { Modal, Spin } from "antd";
|
||||
import { Modal, Spin, Radio, Select } from "antd";
|
||||
import { Button, TextInput } from "@tremor/react";
|
||||
import { modelHubCall, enrichPolicyTemplateStream } from "../networking";
|
||||
|
||||
interface TemplateParameter {
|
||||
name: string;
|
||||
|
|
@ -13,9 +14,13 @@ interface TemplateParameter {
|
|||
interface TemplateParameterModalProps {
|
||||
visible: boolean;
|
||||
template: any;
|
||||
onConfirm: (parameters: Record<string, string>) => void;
|
||||
onConfirm: (
|
||||
parameters: Record<string, string>,
|
||||
enrichmentOptions?: { model?: string; competitors?: string[] }
|
||||
) => void;
|
||||
onCancel: () => void;
|
||||
isLoading?: boolean;
|
||||
accessToken: string;
|
||||
}
|
||||
|
||||
const TemplateParameterModal: React.FC<TemplateParameterModalProps> = ({
|
||||
|
|
@ -24,10 +29,28 @@ const TemplateParameterModal: React.FC<TemplateParameterModalProps> = ({
|
|||
onConfirm,
|
||||
onCancel,
|
||||
isLoading = false,
|
||||
accessToken,
|
||||
}) => {
|
||||
const [parameterValues, setParameterValues] = useState<Record<string, string>>({});
|
||||
const [competitorMode, setCompetitorMode] = useState<"ai" | "manual">("ai");
|
||||
const [selectedModel, setSelectedModel] = useState<string | undefined>(undefined);
|
||||
const [availableModels, setAvailableModels] = useState<string[]>([]);
|
||||
const [isLoadingModels, setIsLoadingModels] = useState(false);
|
||||
const [competitorTags, setCompetitorTags] = useState<string[]>([]);
|
||||
const [variationsMap, setVariationsMap] = useState<Record<string, string[]>>({});
|
||||
const [isGenerating, setIsGenerating] = useState(false);
|
||||
const [refinementInput, setRefinementInput] = useState("");
|
||||
const [isRefining, setIsRefining] = useState(false);
|
||||
const [hasGenerated, setHasGenerated] = useState(false);
|
||||
const [statusMessage, setStatusMessage] = useState("");
|
||||
|
||||
const parameters: TemplateParameter[] = template?.parameters || [];
|
||||
const hasEnrichment = !!template?.llm_enrichment;
|
||||
const enrichmentParam = hasEnrichment ? template.llm_enrichment.parameter : null;
|
||||
|
||||
const nonEnrichmentParams = hasEnrichment
|
||||
? parameters.filter((p) => p.name !== enrichmentParam)
|
||||
: parameters;
|
||||
|
||||
useEffect(() => {
|
||||
if (visible && template) {
|
||||
|
|
@ -36,15 +59,136 @@ const TemplateParameterModal: React.FC<TemplateParameterModalProps> = ({
|
|||
initial[p.name] = "";
|
||||
});
|
||||
setParameterValues(initial);
|
||||
setCompetitorMode("ai");
|
||||
setSelectedModel(undefined);
|
||||
setCompetitorTags([]);
|
||||
setVariationsMap({});
|
||||
setIsGenerating(false);
|
||||
setRefinementInput("");
|
||||
setIsRefining(false);
|
||||
setHasGenerated(false);
|
||||
setStatusMessage("");
|
||||
}
|
||||
}, [visible, template]);
|
||||
|
||||
const allRequiredFilled = parameters
|
||||
useEffect(() => {
|
||||
if (visible && hasEnrichment && competitorMode === "ai" && availableModels.length === 0) {
|
||||
loadModels();
|
||||
}
|
||||
}, [visible, hasEnrichment, competitorMode]);
|
||||
|
||||
const loadModels = async () => {
|
||||
if (!accessToken) return;
|
||||
setIsLoadingModels(true);
|
||||
try {
|
||||
const fetchedModels = await modelHubCall(accessToken);
|
||||
if (fetchedModels?.data?.length > 0) {
|
||||
const models = fetchedModels.data
|
||||
.map((item: any) => item.model_group as string)
|
||||
.sort();
|
||||
setAvailableModels(models);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Error fetching models:", error);
|
||||
} finally {
|
||||
setIsLoadingModels(false);
|
||||
}
|
||||
};
|
||||
|
||||
const handleGenerateNames = async () => {
|
||||
if (!accessToken || !selectedModel || !template) return;
|
||||
const brandName = (parameterValues[enrichmentParam || "brand_name"] || "").trim();
|
||||
if (!brandName) return;
|
||||
|
||||
setIsGenerating(true);
|
||||
setCompetitorTags([]);
|
||||
setVariationsMap({});
|
||||
setStatusMessage("");
|
||||
try {
|
||||
await enrichPolicyTemplateStream(
|
||||
accessToken,
|
||||
template.id,
|
||||
parameterValues,
|
||||
selectedModel,
|
||||
(name) => {
|
||||
setCompetitorTags((prev) => [...prev, name]);
|
||||
},
|
||||
(result) => {
|
||||
setCompetitorTags(result.competitors);
|
||||
setVariationsMap(result.competitor_variations || {});
|
||||
setIsGenerating(false);
|
||||
setHasGenerated(true);
|
||||
setStatusMessage("");
|
||||
},
|
||||
(error) => {
|
||||
console.error("Streaming error:", error);
|
||||
setIsGenerating(false);
|
||||
setStatusMessage("");
|
||||
},
|
||||
undefined,
|
||||
(status) => setStatusMessage(status),
|
||||
);
|
||||
} catch (error) {
|
||||
console.error("Error generating competitor names:", error);
|
||||
setIsGenerating(false);
|
||||
}
|
||||
};
|
||||
|
||||
const handleRefine = async () => {
|
||||
if (!accessToken || !selectedModel || !template || !refinementInput.trim()) return;
|
||||
|
||||
setIsRefining(true);
|
||||
setStatusMessage("");
|
||||
try {
|
||||
await enrichPolicyTemplateStream(
|
||||
accessToken,
|
||||
template.id,
|
||||
parameterValues,
|
||||
selectedModel,
|
||||
(name) => {
|
||||
setCompetitorTags((prev) => {
|
||||
if (prev.some((t) => t.toLowerCase() === name.toLowerCase())) return prev;
|
||||
return [...prev, name];
|
||||
});
|
||||
},
|
||||
(result) => {
|
||||
setCompetitorTags(result.competitors);
|
||||
setVariationsMap(result.competitor_variations || {});
|
||||
setIsRefining(false);
|
||||
setRefinementInput("");
|
||||
setStatusMessage("");
|
||||
},
|
||||
(error) => {
|
||||
console.error("Refinement error:", error);
|
||||
setIsRefining(false);
|
||||
setStatusMessage("");
|
||||
},
|
||||
{
|
||||
instruction: refinementInput.trim(),
|
||||
existingCompetitors: competitorTags,
|
||||
},
|
||||
(status) => setStatusMessage(status),
|
||||
);
|
||||
} catch (error) {
|
||||
console.error("Error refining competitor names:", error);
|
||||
setIsRefining(false);
|
||||
}
|
||||
};
|
||||
|
||||
const allNonEnrichmentFilled = nonEnrichmentParams
|
||||
.filter((p) => p.required)
|
||||
.every((p) => (parameterValues[p.name] || "").trim().length > 0);
|
||||
|
||||
const brandNameFilled = enrichmentParam
|
||||
? (parameterValues[enrichmentParam] || "").trim().length > 0
|
||||
: true;
|
||||
|
||||
const canContinue = hasEnrichment
|
||||
? allNonEnrichmentFilled && brandNameFilled && competitorTags.length > 0
|
||||
: allNonEnrichmentFilled && brandNameFilled;
|
||||
|
||||
const handleConfirm = () => {
|
||||
onConfirm(parameterValues);
|
||||
onConfirm(parameterValues, { competitors: competitorTags });
|
||||
};
|
||||
|
||||
return (
|
||||
|
|
@ -53,15 +197,13 @@ const TemplateParameterModal: React.FC<TemplateParameterModalProps> = ({
|
|||
<div>
|
||||
<h3 className="text-lg font-semibold mb-1">{template?.title}</h3>
|
||||
<p className="text-sm text-gray-500 font-normal">
|
||||
{template?.llm_enrichment
|
||||
? "Enter your brand name to auto-discover competitors and configure guardrails"
|
||||
: "Configure template parameters"}
|
||||
Configure competitor blocking for your brand
|
||||
</p>
|
||||
</div>
|
||||
}
|
||||
open={visible}
|
||||
onCancel={onCancel}
|
||||
width={500}
|
||||
width={700}
|
||||
footer={[
|
||||
<Button key="cancel" variant="secondary" onClick={onCancel} disabled={isLoading}>
|
||||
Cancel
|
||||
|
|
@ -70,18 +212,14 @@ const TemplateParameterModal: React.FC<TemplateParameterModalProps> = ({
|
|||
key="confirm"
|
||||
onClick={handleConfirm}
|
||||
loading={isLoading}
|
||||
disabled={!allRequiredFilled || isLoading}
|
||||
disabled={!canContinue || isLoading}
|
||||
>
|
||||
{isLoading
|
||||
? template?.llm_enrichment
|
||||
? "Discovering competitors..."
|
||||
: "Processing..."
|
||||
: "Continue"}
|
||||
{isLoading ? "Creating guardrails..." : "Continue"}
|
||||
</Button>,
|
||||
]}
|
||||
>
|
||||
<div className="py-4 space-y-4">
|
||||
{parameters.map((param) => (
|
||||
{nonEnrichmentParams.map((param) => (
|
||||
<div key={param.name}>
|
||||
<label className="block text-sm font-medium text-gray-700 mb-1">
|
||||
{param.label}
|
||||
|
|
@ -100,25 +238,168 @@ const TemplateParameterModal: React.FC<TemplateParameterModalProps> = ({
|
|||
</div>
|
||||
))}
|
||||
|
||||
{template?.llm_enrichment && (
|
||||
<div className="mt-4 p-3 bg-blue-50 rounded-lg border border-blue-100">
|
||||
<p className="text-sm text-blue-800">
|
||||
This template uses AI to automatically discover your competitors and configure
|
||||
guardrails. An onboarded LLM will be called to identify competitor names.
|
||||
</p>
|
||||
</div>
|
||||
{hasEnrichment && (
|
||||
<>
|
||||
<div>
|
||||
<label className="block text-sm font-medium text-gray-700 mb-2">
|
||||
Competitor Discovery
|
||||
</label>
|
||||
<Radio.Group
|
||||
value={competitorMode}
|
||||
onChange={(e) => setCompetitorMode(e.target.value)}
|
||||
className="w-full"
|
||||
>
|
||||
<div className="flex gap-3">
|
||||
<Radio.Button value="ai" className="flex-1 text-center">
|
||||
✨ Use AI
|
||||
</Radio.Button>
|
||||
<Radio.Button value="manual" className="flex-1 text-center">
|
||||
Enter Manually
|
||||
</Radio.Button>
|
||||
</div>
|
||||
</Radio.Group>
|
||||
</div>
|
||||
|
||||
{/* Brand Name */}
|
||||
<div>
|
||||
<label className="block text-sm font-medium text-gray-700 mb-1">
|
||||
Your Brand Name
|
||||
<span className="text-red-500 ml-1">*</span>
|
||||
</label>
|
||||
<TextInput
|
||||
placeholder="e.g. Acme Airlines"
|
||||
value={parameterValues[enrichmentParam || "brand_name"] || ""}
|
||||
onChange={(e) =>
|
||||
setParameterValues((prev) => ({
|
||||
...prev,
|
||||
[enrichmentParam || "brand_name"]: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{competitorMode === "ai" && (
|
||||
<>
|
||||
<div>
|
||||
<label className="block text-sm font-medium text-gray-700 mb-1">
|
||||
Select Model
|
||||
<span className="text-red-500 ml-1">*</span>
|
||||
</label>
|
||||
<Select
|
||||
placeholder="Select a model to generate names"
|
||||
value={selectedModel}
|
||||
onChange={(value) => setSelectedModel(value)}
|
||||
loading={isLoadingModels}
|
||||
showSearch
|
||||
className="w-full"
|
||||
options={availableModels.map((m) => ({ label: m, value: m }))}
|
||||
filterOption={(input, option) =>
|
||||
(option?.label ?? "").toLowerCase().includes(input.toLowerCase())
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<Button
|
||||
onClick={handleGenerateNames}
|
||||
loading={isGenerating}
|
||||
disabled={!selectedModel || !brandNameFilled || isGenerating}
|
||||
className="w-full"
|
||||
>
|
||||
{isGenerating ? "✨ Generating names..." : "✨ Generate Competitor Names"}
|
||||
</Button>
|
||||
</>
|
||||
)}
|
||||
|
||||
{/* Competitor Tags */}
|
||||
<div>
|
||||
<label className="block text-sm font-medium text-gray-700 mb-1">
|
||||
Competitor Names
|
||||
{competitorTags.length > 0 && (
|
||||
<span className="text-gray-400 font-normal ml-2">
|
||||
({competitorTags.length})
|
||||
</span>
|
||||
)}
|
||||
</label>
|
||||
<Select
|
||||
mode="tags"
|
||||
style={{ width: "100%" }}
|
||||
placeholder="Type a name and press Enter to add"
|
||||
value={competitorTags}
|
||||
onChange={(values) => setCompetitorTags(values)}
|
||||
tokenSeparators={[","]}
|
||||
open={false}
|
||||
suffixIcon={null}
|
||||
/>
|
||||
<p className="text-xs text-gray-500 mt-1">
|
||||
Type a name and press Enter to add. Click ✕ to remove.
|
||||
</p>
|
||||
{statusMessage && (
|
||||
<div className="flex items-center gap-2 mt-2 p-2 bg-blue-50 rounded border border-blue-100">
|
||||
<Spin size="small" />
|
||||
<span className="text-xs text-blue-700">{statusMessage}</span>
|
||||
</div>
|
||||
)}
|
||||
{Object.keys(variationsMap).length > 0 && !statusMessage && (
|
||||
<p className="text-xs text-green-600 mt-1">
|
||||
✓ {Object.values(variationsMap).flat().length} alternate spellings & variations auto-generated for guardrail matching
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* Refinement input — shown after initial generation in AI mode */}
|
||||
{competitorMode === "ai" && hasGenerated && competitorTags.length > 0 && (
|
||||
<div>
|
||||
<label className="block text-sm font-medium text-gray-700 mb-1">
|
||||
Refine List
|
||||
</label>
|
||||
<div className="flex gap-2">
|
||||
<TextInput
|
||||
placeholder="e.g. add 10 more from Asia, increase to 50 total..."
|
||||
value={refinementInput}
|
||||
onChange={(e) => setRefinementInput(e.target.value)}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && refinementInput.trim() && !isRefining) {
|
||||
handleRefine();
|
||||
}
|
||||
}}
|
||||
disabled={isRefining}
|
||||
/>
|
||||
<Button
|
||||
onClick={handleRefine}
|
||||
loading={isRefining}
|
||||
disabled={!refinementInput.trim() || isRefining}
|
||||
size="xs"
|
||||
>
|
||||
{isRefining ? "..." : "Send"}
|
||||
</Button>
|
||||
</div>
|
||||
<p className="text-xs text-gray-400 mt-1">
|
||||
Give instructions to add, remove, or change competitors. Press Enter to send.
|
||||
</p>
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
|
||||
{isLoading && (
|
||||
<div className="flex items-center gap-3 mt-4 p-3 bg-gray-50 rounded-lg">
|
||||
<Spin size="small" />
|
||||
<span className="text-sm text-gray-600">
|
||||
{template?.llm_enrichment
|
||||
? "Using AI to discover competitors..."
|
||||
: "Processing template..."}
|
||||
</span>
|
||||
</div>
|
||||
)}
|
||||
{!hasEnrichment &&
|
||||
parameters.map((param) => (
|
||||
<div key={param.name}>
|
||||
<label className="block text-sm font-medium text-gray-700 mb-1">
|
||||
{param.label}
|
||||
{param.required && <span className="text-red-500 ml-1">*</span>}
|
||||
</label>
|
||||
<TextInput
|
||||
placeholder={param.placeholder || ""}
|
||||
value={parameterValues[param.name] || ""}
|
||||
onChange={(e) =>
|
||||
setParameterValues((prev) => ({
|
||||
...prev,
|
||||
[param.name]: e.target.value,
|
||||
}))
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</Modal>
|
||||
);
|
||||
|
|
|
|||
57
ui/litellm-dashboard/src/components/query_param_input.tsx
Normal file
57
ui/litellm-dashboard/src/components/query_param_input.tsx
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
import React, { useState } from "react";
|
||||
import { Button, Space } from "antd";
|
||||
import { MinusCircleOutlined, PlusOutlined } from "@ant-design/icons";
|
||||
import { TextInput } from "@tremor/react";
|
||||
|
||||
interface QueryParamInputProps {
|
||||
value?: Record<string, string>;
|
||||
onChange?: (value: Record<string, string>) => void;
|
||||
}
|
||||
|
||||
const QueryParamInput: React.FC<QueryParamInputProps> = ({ value = {}, onChange }) => {
|
||||
const [pairs, setPairs] = useState<[string, string][]>(Object.entries(value));
|
||||
|
||||
const handleAdd = () => {
|
||||
setPairs([...pairs, ["", ""]]);
|
||||
};
|
||||
|
||||
const handleRemove = (index: number) => {
|
||||
const newPairs = pairs.filter((_, i) => i !== index);
|
||||
setPairs(newPairs);
|
||||
onChange?.(Object.fromEntries(newPairs));
|
||||
};
|
||||
|
||||
const handleChange = (index: number, key: string, val: string) => {
|
||||
const newPairs = [...pairs];
|
||||
newPairs[index] = [key, val];
|
||||
setPairs(newPairs);
|
||||
onChange?.(Object.fromEntries(newPairs));
|
||||
};
|
||||
|
||||
return (
|
||||
<div>
|
||||
{pairs.map(([key, val], index) => (
|
||||
<Space key={index} style={{ display: "flex", marginBottom: 8 }} align="center">
|
||||
<TextInput
|
||||
placeholder="Parameter Name (e.g., version)"
|
||||
value={key}
|
||||
onChange={(e) => handleChange(index, e.target.value, val)}
|
||||
/>
|
||||
<TextInput
|
||||
placeholder="Parameter Value (e.g., v1)"
|
||||
value={val}
|
||||
onChange={(e) => handleChange(index, key, e.target.value)}
|
||||
/>
|
||||
<div style={{ display: "flex", alignItems: "center", justifyContent: "center", height: "100%" }}>
|
||||
<MinusCircleOutlined onClick={() => handleRemove(index)} style={{ cursor: "pointer" }} />
|
||||
</div>
|
||||
</Space>
|
||||
))}
|
||||
<Button type="dashed" onClick={handleAdd} icon={<PlusOutlined />}>
|
||||
Add Query Parameter
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default QueryParamInput;
|
||||
|
|
@ -97,6 +97,20 @@ describe("MemberPermissions", () => {
|
|||
}
|
||||
});
|
||||
|
||||
it("should render team daily activity permission with correct method and description", async () => {
|
||||
vi.mocked(networking.getTeamPermissionsCall).mockResolvedValue({
|
||||
all_available_permissions: ["/key/generate", "/team/daily/activity"],
|
||||
team_member_permissions: [],
|
||||
});
|
||||
|
||||
renderWithProviders(<MemberPermissions teamId="team-123" accessToken="token-123" canEditTeam={true} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("/team/daily/activity")).toBeInTheDocument();
|
||||
expect(screen.getByText("Member can view all team usage data (not just their own)")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should not show save button when canEditTeam is false", async () => {
|
||||
vi.mocked(networking.getTeamPermissionsCall).mockResolvedValue({
|
||||
all_available_permissions: ["/key/generate", "/key/list"],
|
||||
|
|
|
|||
|
|
@ -11,6 +11,10 @@ describe("permission_definitions", () => {
|
|||
expect(getMethodForEndpoint("/key/list")).toBe("GET");
|
||||
});
|
||||
|
||||
it("should return GET for activity endpoints", () => {
|
||||
expect(getMethodForEndpoint("/team/daily/activity")).toBe("GET");
|
||||
});
|
||||
|
||||
it("should return POST for other endpoints", () => {
|
||||
expect(getMethodForEndpoint("/key/generate")).toBe("POST");
|
||||
expect(getMethodForEndpoint("/key/update")).toBe("POST");
|
||||
|
|
@ -48,6 +52,14 @@ describe("permission_definitions", () => {
|
|||
expect(result.description).toBe(PERMISSION_DESCRIPTIONS["/key/service-account/generate"]);
|
||||
});
|
||||
|
||||
it("should return correct info for team daily activity permission", () => {
|
||||
const result = getPermissionInfo("/team/daily/activity");
|
||||
expect(result.method).toBe("GET");
|
||||
expect(result.endpoint).toBe("/team/daily/activity");
|
||||
expect(result.description).toBe(PERMISSION_DESCRIPTIONS["/team/daily/activity"]);
|
||||
expect(result.route).toBe("/team/daily/activity");
|
||||
});
|
||||
|
||||
it("should return fallback description for unknown permission", () => {
|
||||
const result = getPermissionInfo("/unknown/endpoint");
|
||||
expect(result.method).toBe("POST");
|
||||
|
|
@ -56,4 +68,11 @@ describe("permission_definitions", () => {
|
|||
expect(result.route).toBe("/unknown/endpoint");
|
||||
});
|
||||
});
|
||||
|
||||
describe("PERMISSION_DESCRIPTIONS", () => {
|
||||
it("should include team daily activity permission", () => {
|
||||
expect(PERMISSION_DESCRIPTIONS["/team/daily/activity"]).toBeDefined();
|
||||
expect(PERMISSION_DESCRIPTIONS["/team/daily/activity"]).toContain("team usage");
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -20,13 +20,15 @@ export const PERMISSION_DESCRIPTIONS: Record<string, string> = {
|
|||
"/key/list": "Member can list virtual keys belonging to this team",
|
||||
"/key/block": "Member can block a virtual key belonging to this team",
|
||||
"/key/unblock": "Member can unblock a virtual key belonging to this team",
|
||||
"/team/daily/activity":
|
||||
"Member can view all team usage data (not just their own)",
|
||||
};
|
||||
|
||||
/**
|
||||
* Determines the HTTP method for a given permission endpoint
|
||||
*/
|
||||
export const getMethodForEndpoint = (endpoint: string): string => {
|
||||
if (endpoint.includes("/info") || endpoint.includes("/list")) {
|
||||
if (endpoint.includes("/info") || endpoint.includes("/list") || endpoint.includes("/activity")) {
|
||||
return "GET";
|
||||
}
|
||||
return "POST";
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue