mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
commit
7d790b39be
180 changed files with 11865 additions and 1958 deletions
|
|
@ -109,6 +109,8 @@ Key files:
|
|||
- `litellm/proxy/auth/` - Authentication logic
|
||||
- `litellm/proxy/management_endpoints/` - Admin API endpoints
|
||||
|
||||
**Database (proxy)**: Use Prisma model methods (`prisma_client.db.<model>.upsert`, `.find_many`, `.find_unique`, etc.), not raw SQL (`execute_raw`/`query_raw`). See COMMON PITFALLS for details.
|
||||
|
||||
## MCP (MODEL CONTEXT PROTOCOL) SUPPORT
|
||||
|
||||
LiteLLM supports MCP for agent workflows:
|
||||
|
|
@ -176,6 +178,7 @@ When opening issues or pull requests, follow these templates:
|
|||
5. **Dependencies**: Keep dependencies minimal and well-justified
|
||||
6. **UI/Backend Contract Mismatch**: When adding a new entity type to the UI, always check whether the backend endpoint accepts a single value or an array. Match the UI control accordingly (single-select vs. multi-select) to avoid silently dropping user selections
|
||||
7. **Missing Tests for New Entity Types**: When adding a new entity type (e.g., in `EntityUsage`, `UsageViewSelect`), always add corresponding tests in the existing test files and update any icon/component mocks
|
||||
8. **Raw SQL in proxy DB code**: Do not use `execute_raw` or `query_raw` for proxy database access. Use Prisma model methods (e.g. `prisma_client.db.litellm_tooltable.upsert()`, `.find_many()`, `.find_unique()`) so behavior stays consistent with the schema, the client stays mockable in tests, and you avoid the pitfalls of hand-written SQL (parameter ordering, type casting, schema drift)
|
||||
|
||||
8. **Do not hardcode model-specific flags**: Put model-specific capability flags in `model_prices_and_context_window.json` and read them via `get_model_info` (or existing helpers like `supports_reasoning`). This prevents users from needing to upgrade LiteLLM each time a new model supports a feature.
|
||||
|
||||
|
|
|
|||
|
|
@ -107,6 +107,10 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components:
|
|||
- Migration files auto-generated with `prisma migrate dev`
|
||||
- Always test migrations against both PostgreSQL and SQLite
|
||||
|
||||
### Proxy database access
|
||||
- **Do not write raw SQL** for proxy DB operations. Use Prisma model methods instead of `execute_raw` / `query_raw`.
|
||||
- Use the generated client: `prisma_client.db.<model>` (e.g. `litellm_tooltable`, `litellm_usertable`) with `.upsert()`, `.find_many()`, `.find_unique()`, `.update()`, `.update_many()` as appropriate. This avoids schema/client drift, keeps code testable with simple mocks, and matches patterns used in spend logs and other proxy code.
|
||||
|
||||
### Enterprise Features
|
||||
- Enterprise-specific code in `enterprise/` directory
|
||||
- Optional features enabled via environment variables
|
||||
|
|
|
|||
13
dev_config.yaml
Normal file
13
dev_config.yaml
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
model_list:
|
||||
- model_name: fake-openai-endpoint
|
||||
litellm_params:
|
||||
model: openai/fake-model
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
telemetry: False
|
||||
175
docs/my-website/blog/gemini_3_1_flash_lite/index.md
Normal file
175
docs/my-website/blog/gemini_3_1_flash_lite/index.md
Normal file
|
|
@ -0,0 +1,175 @@
|
|||
---
|
||||
slug: gemini_3_1_flash_lite_preview
|
||||
title: "DAY 0 Support: Gemini 3.1 Flash Lite Preview on LiteLLM"
|
||||
date: 2026-03-03T08: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 Flash Lite Preview on LiteLLM Proxy and SDK with day 0 support."
|
||||
tags: [gemini, day 0 support, llms, supernova]
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Gemini 3.1 Flash Lite Preview Day 0 Support
|
||||
|
||||
LiteLLM now supports `gemini-3.1-flash-lite-preview` with full day 0 support!
|
||||
|
||||
:::note
|
||||
If you only want cost tracking, you need no change in your current Litellm version. But if you want the support for new features introduced along with it like thinking levels, you will need to use v1.80.8-stable.1 or above.
|
||||
:::
|
||||
|
||||
## 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==v1.80.8-stable.1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## What's New
|
||||
|
||||
Supports all four thinking levels:
|
||||
- **MINIMAL**: Ultra-fast responses with minimal reasoning
|
||||
- **LOW**: Simple instruction following
|
||||
- **MEDIUM**: Balanced reasoning for complex tasks
|
||||
- **HIGH**: Maximum reasoning depth (dynamic)
|
||||
|
||||
---
|
||||
|
||||
## Quick Start
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
**Basic Usage**
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
response = completion(
|
||||
model="gemini/gemini-3.1-flash-lite-preview",
|
||||
messages=[{"role": "user", "content": "Extract key entities from this text: ..."}],
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
**With Thinking Levels**
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
# Use MEDIUM thinking for complex reasoning tasks
|
||||
response = completion(
|
||||
model="gemini/gemini-3.1-flash-lite-preview",
|
||||
messages=[{"role": "user", "content": "Analyze this dataset and identify patterns"}],
|
||||
reasoning_effort="medium", # low, medium , high
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
**1. Setup config.yaml**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemini-3.1-flash-lite
|
||||
litellm_params:
|
||||
model: gemini/gemini-3.1-flash-lite-preview
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
# Or use Vertex AI
|
||||
- model_name: vertex-gemini-3.1-flash-lite
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-3.1-flash-lite-preview
|
||||
vertex_project: your-project-id
|
||||
vertex_location: us-central1
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
**3. Make requests**
|
||||
|
||||
```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-flash-lite",
|
||||
"messages": [{"role": "user", "content": "Extract structured data from this text"}],
|
||||
"reasoning_effort": "low"
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
---
|
||||
|
||||
## Supported Endpoints
|
||||
|
||||
LiteLLM provides **full end-to-end support** for Gemini 3.1 Flash Lite Preview 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 (thinking levels, thought signatures)
|
||||
- Full multimodal support (text, image, audio, video)
|
||||
|
||||
---
|
||||
|
||||
## `reasoning_effort` Mapping for Gemini 3.1
|
||||
|
||||
LiteLLM automatically maps OpenAI's `reasoning_effort` parameter to Gemini's `thinkingLevel`:
|
||||
|
||||
| reasoning_effort | thinking_level | Use Case |
|
||||
|------------------|----------------|----------|
|
||||
| `minimal` | `minimal` | Ultra-fast responses, simple queries |
|
||||
| `low` | `low` | Basic instruction following |
|
||||
| `medium` | `medium` | Balanced reasoning for moderate complexity |
|
||||
| `high` | `high` | Maximum reasoning depth, complex problems |
|
||||
| `disable` | `minimal` | Disable extended reasoning |
|
||||
| `none` | `minimal` | No extended reasoning |
|
||||
|
|
@ -0,0 +1,321 @@
|
|||
---
|
||||
slug: responses-api-encrypted-content-incident
|
||||
title: "Incident Report: Encrypted Content Failures in Multi-Region Responses API Load Balancing"
|
||||
date: 2026-02-24T10: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
|
||||
tags: [incident-report, proxy, responses-api, load-balancing]
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
**Date:** Feb 24, 2026
|
||||
**Duration:** Ongoing (until fix deployed)
|
||||
**Severity:** High (for users load balancing Responses API across different API keys)
|
||||
**Status:** Resolved
|
||||
|
||||
## Summary
|
||||
|
||||
When load balancing OpenAI's Responses API across deployments with **different API keys** (e.g., different Azure regions or OpenAI organizations), follow-up requests containing encrypted content items (like `rs_...` reasoning items) would fail with:
|
||||
|
||||
```json
|
||||
{
|
||||
"error": {
|
||||
"message": "The encrypted content for item rs_0d09d6e56879e76500699d6feee41c8197bd268aae76141f87 could not be verified. Reason: Encrypted content organization_id did not match the target organization.",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_encrypted_content"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Encrypted content items are cryptographically tied to the API key's organization that created them. When the router load balanced a follow-up request to a deployment with a different API key, decryption failed.
|
||||
|
||||
- **Responses API calls with encrypted content:** Complete failure when routed to wrong deployment
|
||||
- **Initial requests:** Unaffected — only follow-up requests containing encrypted items failed
|
||||
- **Other API endpoints:** No impact — chat completions, embeddings, etc. functioned normally
|
||||
|
||||
{/* truncate */}
|
||||
|
||||
---
|
||||
|
||||
## Background
|
||||
|
||||
OpenAI's Responses API can return encrypted "reasoning items" (with IDs like `rs_...`) that contain intermediate reasoning steps. These items are encrypted with the organization's key and can only be decrypted by the same organization's API key.
|
||||
|
||||
When load balancing across deployments with different API keys, the existing affinity mechanisms were insufficient:
|
||||
|
||||
- **`responses_api_deployment_check`**: Requires `previous_response_id` which some clients (like Codex) don't provide
|
||||
- **`deployment_affinity`**: Too broad — pins *all* requests from a user to one deployment, reducing effective quota by the number of users
|
||||
- **`session_affinity`**: Requires explicit session IDs and still reduces quota
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
A["1. Initial request to Responses API
|
||||
router.aresponses()"] --> B["2. Router load balances to Deployment A
|
||||
(API Key 1, Azure East US)"]
|
||||
B --> C["3. Response contains encrypted item
|
||||
rs_abc123 (encrypted with Org 1 key)"]
|
||||
C --> D["4. Follow-up request includes rs_abc123 in input"]
|
||||
D --> E["5. Router load balances to Deployment B
|
||||
(API Key 2, Azure West Europe)"]
|
||||
E -->|"Different API key"| F["6. ❌ Deployment B cannot decrypt rs_abc123
|
||||
Error: invalid_encrypted_content"]
|
||||
|
||||
D -.->|"With encrypted_content_affinity"| G["5b. Router detects rs_abc123 was created by Deployment A"]
|
||||
G --> H["6b. ✅ Routes to Deployment A (bypasses rate limits)
|
||||
Request succeeds"]
|
||||
|
||||
style F fill:#f8d7da,stroke:#dc3545
|
||||
style H fill:#d4edda,stroke:#28a745
|
||||
style E fill:#fff3cd,stroke:#ffc107
|
||||
style G fill:#d4edda,stroke:#28a745
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Root Cause
|
||||
|
||||
LiteLLM's router had no mechanism to track which deployment created specific encrypted content items and route follow-up requests accordingly. The router treated all deployments as interchangeable, leading to decryption failures when encrypted content crossed organizational boundaries.
|
||||
|
||||
**The Problem Flow:**
|
||||
|
||||
1. User calls `router.aresponses()` with model `gpt-5.1-codex`
|
||||
2. Router load balances to Deployment A (Azure East US, API Key 1)
|
||||
3. Response contains encrypted reasoning item `rs_abc123` (encrypted with Org 1's key)
|
||||
4. User makes follow-up request with `rs_abc123` in the input
|
||||
5. Router load balances to Deployment B (Azure West Europe, API Key 2)
|
||||
6. Deployment B tries to decrypt `rs_abc123` with Org 2's key → **fails**
|
||||
|
||||
**Why Existing Solutions Didn't Work:**
|
||||
|
||||
- **`previous_response_id`**: Not provided by all clients (e.g., Codex)
|
||||
- **`deployment_affinity`**: Pins *all* user requests to one deployment → reduces quota to 1/N where N = number of deployments
|
||||
- **`session_affinity`**: Requires explicit session management and still reduces quota
|
||||
|
||||
**Timeline:**
|
||||
|
||||
1. Users configured multi-region Responses API load balancing with different API keys
|
||||
2. Initial requests succeeded, but follow-up requests with encrypted content failed intermittently
|
||||
3. Error rate correlated with number of deployments (more deployments = higher chance of routing to wrong one)
|
||||
4. Investigation revealed encrypted content was organization-bound
|
||||
5. Existing affinity mechanisms deemed unsuitable (quota reduction, missing `previous_response_id`)
|
||||
6. New solution designed and implemented: `encrypted_content_affinity`
|
||||
|
||||
---
|
||||
|
||||
## The Fix
|
||||
|
||||
Implemented a new `encrypted_content_affinity` pre-call check that intelligently tracks encrypted content and routes follow-up requests **only when necessary**.
|
||||
|
||||
### Implementation
|
||||
|
||||
**1. Encoding `model_id` into output items** ([`responses/utils.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/responses/utils.py))
|
||||
|
||||
The same approach used for `previous_response_id` affinity — no cache needed. When a response contains output items with `encrypted_content`, LiteLLM encodes the originating deployment's `model_id` in **two places** for redundancy:
|
||||
|
||||
1. **Into the item ID** (if present): `rs_abc123` → `encitem_{base64("litellm:model_id:{model_id};item_id:rs_abc123")}`
|
||||
2. **Into the encrypted_content itself**: Wraps the content with `litellm_enc:{base64("model_id:{model_id}")};{original_encrypted_content}`
|
||||
|
||||
```python
|
||||
# Encoding item IDs (when present)
|
||||
def _build_encrypted_item_id(model_id: str, item_id: str) -> str:
|
||||
assembled = f"litellm:model_id:{model_id};item_id:{item_id}"
|
||||
encoded = base64.b64encode(assembled.encode("utf-8")).decode("utf-8")
|
||||
return f"encitem_{encoded}"
|
||||
|
||||
# Wrapping encrypted_content (always, for redundancy)
|
||||
def _wrap_encrypted_content_with_model_id(encrypted_content: str, model_id: str) -> str:
|
||||
metadata = f"model_id:{model_id}"
|
||||
encoded_metadata = base64.b64encode(metadata.encode("utf-8")).decode("utf-8")
|
||||
return f"litellm_enc:{encoded_metadata};{encrypted_content}"
|
||||
```
|
||||
|
||||
**Why wrap encrypted_content directly?** Some clients (like Codex) don't consistently send item IDs in follow-up requests, but they always send the `encrypted_content` itself. By embedding `model_id` into the content, affinity works even when IDs are missing.
|
||||
|
||||
**Streaming responses:** The wrapping logic is applied to both:
|
||||
- Final response objects (non-streaming)
|
||||
- Individual streaming events (`response.output_item.added`, `response.output_item.done`)
|
||||
|
||||
This ensures clients receiving streaming responses get wrapped content they can send back.
|
||||
|
||||
Before forwarding to the upstream provider, LiteLLM restores the original item IDs and unwraps encrypted_content so the provider never sees the encoded form:
|
||||
|
||||
```python
|
||||
# In responses/main.py — before calling the handler
|
||||
input = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(input)
|
||||
```
|
||||
|
||||
**2. `EncryptedContentAffinityCheck` — routing only** ([`encrypted_content_affinity_check.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py))
|
||||
|
||||
No `async_log_success_event` or cache lookups — the `model_id` is decoded directly from the item ID or encrypted_content:
|
||||
|
||||
```python
|
||||
class EncryptedContentAffinityCheck(CustomLogger):
|
||||
async def async_filter_deployments(self, model, healthy_deployments, ...):
|
||||
"""Extract model_id from input items (ID or encrypted_content) and pin to that deployment."""
|
||||
for item in request_kwargs.get("input", []):
|
||||
# Try to extract model_id from two sources:
|
||||
model_id = self._extract_model_id_from_input(item)
|
||||
|
||||
if model_id:
|
||||
deployment = self._find_deployment_by_model_id(
|
||||
healthy_deployments, model_id
|
||||
)
|
||||
if deployment:
|
||||
request_kwargs["_encrypted_content_affinity_pinned"] = True
|
||||
return [deployment]
|
||||
return healthy_deployments
|
||||
|
||||
def _extract_model_id_from_input(self, item: dict) -> Optional[str]:
|
||||
"""Extract model_id from either encoded ID or wrapped encrypted_content."""
|
||||
# 1. Try decoding from item ID (if present)
|
||||
item_id = item.get("id", "")
|
||||
if item_id:
|
||||
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(item_id)
|
||||
if decoded:
|
||||
return decoded["model_id"]
|
||||
|
||||
# 2. Try unwrapping from encrypted_content (fallback for clients that omit IDs)
|
||||
encrypted_content = item.get("encrypted_content", "")
|
||||
if encrypted_content and encrypted_content.startswith("litellm_enc:"):
|
||||
model_id, _ = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(
|
||||
encrypted_content
|
||||
)
|
||||
return model_id
|
||||
|
||||
return None
|
||||
```
|
||||
|
||||
**3. Rate Limit Bypass** ([`router.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/router.py))
|
||||
|
||||
When encrypted content requires a specific deployment, RPM/TPM limits are bypassed (the request would fail on any other deployment anyway):
|
||||
|
||||
```python
|
||||
# In async_get_available_deployment, after filtering healthy deployments:
|
||||
if (
|
||||
request_kwargs.get("_encrypted_content_affinity_pinned")
|
||||
and len(healthy_deployments) == 1
|
||||
):
|
||||
return healthy_deployments[0] # Bypass routing strategy (RPM/TPM checks)
|
||||
```
|
||||
|
||||
**3. Configuration**
|
||||
|
||||
```yaml
|
||||
router_settings:
|
||||
routing_strategy: usage-based-routing-v2
|
||||
enable_pre_call_checks: true
|
||||
optional_pre_call_checks:
|
||||
- encrypted_content_affinity
|
||||
deployment_affinity_ttl_seconds: 86400 # 24 hours
|
||||
```
|
||||
|
||||
### Key Benefits
|
||||
|
||||
✅ **No quota reduction**: Only pins requests containing encrypted items
|
||||
✅ **Bypasses rate limits**: When encrypted content requires a specific deployment, RPM/TPM limits don't block it
|
||||
✅ **No `previous_response_id` required**: Works by encoding `model_id` directly into the item ID
|
||||
✅ **No cache required**: `model_id` is decoded on-the-fly from the item ID — no Redis, no TTL
|
||||
✅ **Globally safe**: Can be enabled for all models; non-Responses-API calls are unaffected
|
||||
✅ **Surgical precision**: Normal requests continue to load balance freely
|
||||
|
||||
---
|
||||
|
||||
## Remediation
|
||||
|
||||
| # | Action | Status | Code |
|
||||
|---|---|---|---|
|
||||
| 1 | Encode `model_id` into encrypted-content item IDs on response | ✅ Done | [`responses/utils.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/responses/utils.py) |
|
||||
| 2 | Restore original item IDs before forwarding to upstream provider | ✅ Done | [`responses/main.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/responses/main.py) |
|
||||
| 3 | `EncryptedContentAffinityCheck`: decode item IDs to route (no cache) | ✅ Done | [`encrypted_content_affinity_check.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py) |
|
||||
| 4 | Add `encrypted_content_affinity` to `OptionalPreCallChecks` type | ✅ Done | [`types/router.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/types/router.py) |
|
||||
| 5 | Implement rate limit bypass for affinity-pinned requests | ✅ Done | [`router.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/router.py) |
|
||||
| 6 | Unit tests: encoding/decoding utilities, routing, RPM bypass | ✅ Done | [`test_encrypted_content_affinity_check.py`](https://github.com/BerriAI/litellm/blob/main/litellm/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py) |
|
||||
| 7 | Documentation: Responses API guide, load balancing guide, config reference | ✅ Done | [Docs](https://docs.litellm.ai/docs/response_api#encrypted-content-affinity-multi-region-load-balancing) |
|
||||
| 8 | **[Mar 3]** Fix streaming events to wrap encrypted_content | ✅ Done | [`responses/streaming_iterator.py`](https://github.com/BerriAI/litellm/blob/main/litellm/litellm/responses/streaming_iterator.py) |
|
||||
|
||||
---
|
||||
|
||||
## Follow-up Fix: Streaming Responses (Mar 3, 2026)
|
||||
|
||||
### The Issue
|
||||
|
||||
After the initial fix was deployed, users reported that the `invalid_encrypted_content` error **still occurred** when using streaming responses with clients like Codex. Investigation revealed:
|
||||
|
||||
- ✅ Non-streaming responses: `encrypted_content` was correctly wrapped with `litellm_enc:` prefix
|
||||
- ❌ Streaming responses: Individual `response.output_item.added` and `response.output_item.done` events contained **raw, unwrapped** `encrypted_content`
|
||||
|
||||
Since Codex and other clients consume responses as streams, they received unwrapped content in these events and sent it back in follow-up requests, causing the affinity check to fail.
|
||||
|
||||
### The Root Cause
|
||||
|
||||
The `_update_encrypted_content_item_ids_in_response` function only modified the **final** response object, which is used for non-streaming responses. For streaming responses, individual chunks are processed by `ResponsesAPIStreamingIterator._process_chunk`, which was **not** applying the wrapping logic to streaming events.
|
||||
|
||||
### The Fix
|
||||
|
||||
Modified `litellm/litellm/responses/streaming_iterator.py` to wrap `encrypted_content` in streaming events:
|
||||
|
||||
```python
|
||||
# In ResponsesAPIStreamingIterator._process_chunk
|
||||
if (
|
||||
self.litellm_metadata
|
||||
and self.litellm_metadata.get("encrypted_content_affinity_enabled")
|
||||
):
|
||||
event_type = getattr(openai_responses_api_chunk, "type", None)
|
||||
if event_type in (
|
||||
ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
):
|
||||
item = getattr(openai_responses_api_chunk, "item", None)
|
||||
if item:
|
||||
encrypted_content = getattr(item, "encrypted_content", None)
|
||||
if encrypted_content and isinstance(encrypted_content, str):
|
||||
model_id = (
|
||||
self.litellm_metadata.get("model_info", {}).get("id")
|
||||
if self.litellm_metadata
|
||||
else None
|
||||
)
|
||||
if model_id:
|
||||
wrapped_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
encrypted_content, model_id
|
||||
)
|
||||
setattr(item, "encrypted_content", wrapped_content)
|
||||
```
|
||||
|
||||
This ensures that **all** `encrypted_content` sent to clients (streaming or non-streaming) is wrapped with `model_id` metadata, enabling consistent affinity routing.
|
||||
|
||||
---
|
||||
|
||||
## Migration Guide
|
||||
|
||||
### Before (Using `deployment_affinity`)
|
||||
|
||||
```yaml
|
||||
router_settings:
|
||||
optional_pre_call_checks:
|
||||
- deployment_affinity # ❌ Reduces quota by number of users
|
||||
```
|
||||
|
||||
**Problem:** All requests from a user pin to one deployment, reducing effective quota to 1/N.
|
||||
|
||||
### After (Using `encrypted_content_affinity`)
|
||||
|
||||
```yaml
|
||||
router_settings:
|
||||
optional_pre_call_checks:
|
||||
- encrypted_content_affinity # ✅ Only pins requests with encrypted content
|
||||
```
|
||||
|
||||
**Benefit:** Normal requests load balance freely, only encrypted content requests pin when necessary.
|
||||
|
||||
---
|
||||
|
|
@ -2041,6 +2041,7 @@ response = litellm.completion(
|
|||
| gemini-2.0-flash-lite-preview-02-05 | `completion(model='gemini/gemini-2.0-flash-lite-preview-02-05', messages)` | `os.environ['GEMINI_API_KEY']` |
|
||||
| gemini-2.5-flash-preview-09-2025 | `completion(model='gemini/gemini-2.5-flash-preview-09-2025', messages)` | `os.environ['GEMINI_API_KEY']` |
|
||||
| gemini-2.5-flash-lite-preview-09-2025 | `completion(model='gemini/gemini-2.5-flash-lite-preview-09-2025', messages)` | `os.environ['GEMINI_API_KEY']` |
|
||||
| gemini-3.1-flash-lite-preview | `completion(model='gemini/gemini-3.1-flash-lite-preview', messages)` | `os.environ['GEMINI_API_KEY']` |
|
||||
| gemini-flash-latest | `completion(model='gemini/gemini-flash-latest', messages)` | `os.environ['GEMINI_API_KEY']` |
|
||||
| gemini-flash-lite-latest | `completion(model='gemini/gemini-flash-lite-latest', messages)` | `os.environ['GEMINI_API_KEY']` |
|
||||
|
||||
|
|
|
|||
|
|
@ -191,6 +191,7 @@ os.environ["OPENAI_BASE_URL"] = "https://your_host/v1" # OPTIONAL
|
|||
| gpt-5.2 | `response = completion(model="gpt-5.2", messages=messages)` |
|
||||
| gpt-5.2-2025-12-11 | `response = completion(model="gpt-5.2-2025-12-11", messages=messages)` |
|
||||
| gpt-5.2-chat-latest | `response = completion(model="gpt-5.2-chat-latest", messages=messages)` |
|
||||
| gpt-5.3-chat-latest | `response = completion(model="gpt-5.3-chat-latest", messages=messages)` |
|
||||
| gpt-5.2-pro | `response = completion(model="gpt-5.2-pro", messages=messages)` |
|
||||
| gpt-5.2-pro-2025-12-11 | `response = completion(model="gpt-5.2-pro-2025-12-11", messages=messages)` |
|
||||
| gpt-5.1 | `response = completion(model="gpt-5.1", messages=messages)` |
|
||||
|
|
|
|||
|
|
@ -1685,6 +1685,7 @@ litellm.vertex_location = "us-central1 # Your Location
|
|||
| gemini-2.5-pro | `completion('gemini-2.5-pro', messages)`, `completion('vertex_ai/gemini-2.5-pro', messages)` |
|
||||
| gemini-2.5-flash-preview-09-2025 | `completion('gemini-2.5-flash-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-preview-09-2025', messages)` |
|
||||
| gemini-2.5-flash-lite-preview-09-2025 | `completion('gemini-2.5-flash-lite-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-lite-preview-09-2025', messages)` |
|
||||
| gemini-3.1-flash-lite-preview | `completion('gemini-3.1-flash-lite-preview', messages)`, `completion('vertex_ai/gemini-3.1-flash-lite-preview', messages)` |
|
||||
|
||||
## Private Service Connect (PSC) Endpoints
|
||||
|
||||
|
|
|
|||
|
|
@ -360,7 +360,7 @@ router_settings:
|
|||
| redis_url | str | URL for Redis server. **Known performance issue with Redis URL.** |
|
||||
| cache_responses | boolean | Flag to enable caching LLM Responses, if cache set under `router_settings`. If true, caches responses. Defaults to False. |
|
||||
| router_general_settings | RouterGeneralSettings | [SDK-Only] Router general settings - contains optimizations like 'async_only_mode'. [Docs](../routing.md#router-general-settings) |
|
||||
| optional_pre_call_checks | List[str] | List of pre-call checks to add to the router. Supported: `router_budget_limiting`, `prompt_caching`, `responses_api_deployment_check`, `deployment_affinity`, `forward_client_headers_by_model_group` |
|
||||
| optional_pre_call_checks | List[str] | List of pre-call checks to add to the router. Supported: `router_budget_limiting`, `prompt_caching`, `responses_api_deployment_check`, `encrypted_content_affinity`, `deployment_affinity`, `session_affinity`, `forward_client_headers_by_model_group` |
|
||||
| deployment_affinity_ttl_seconds | int | TTL (seconds) for user-key → deployment affinity mapping when `deployment_affinity` is enabled (configured at Router init / proxy startup). Defaults to `3600` (1 hour). |
|
||||
| ignore_invalid_deployments | boolean | If true, ignores invalid deployments. Default for proxy is True - to prevent invalid models from blocking other models from being loaded. |
|
||||
| search_tools | List[SearchToolTypedDict] | List of search tool configurations for Search API integration. Each tool specifies a search_tool_name and litellm_params with search_provider, api_key, api_base, etc. [Further Docs](../search.md) |
|
||||
|
|
|
|||
|
|
@ -358,13 +358,13 @@ response = client.chat.completions.create(
|
|||
}
|
||||
],
|
||||
extra_body={
|
||||
"guardrails": [
|
||||
"guardrails": {
|
||||
"aporia-pre-guard": {
|
||||
"extra_body": {
|
||||
"success_threshold": 0.9
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
)
|
||||
|
|
@ -387,13 +387,13 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
|
|||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
"guardrails": [
|
||||
"guardrails": {
|
||||
"aporia-pre-guard": {
|
||||
"extra_body": {
|
||||
"success_threshold": 0.9
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
}'
|
||||
```
|
||||
</TabItem>
|
||||
|
|
@ -451,7 +451,6 @@ curl -X POST 'http://0.0.0.0:4000/key/generate' \
|
|||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"guardrails": ["aporia-pre-guard", "aporia-post-guard"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
|
|
@ -465,7 +464,6 @@ curl --location 'http://0.0.0.0:4000/key/update' \
|
|||
--data '{
|
||||
"key": "sk-jNm1Zar7XfNdZXp49Z1kSQ",
|
||||
"guardrails": ["aporia-pre-guard", "aporia-post-guard"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
|
|
@ -499,6 +497,11 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
|
|||
|
||||
Run guardrails based on the user-agent header. This is useful for running pre-call checks on OpenWebUI but only masking in logs for Claude CLI.
|
||||
|
||||
`default` can be a single mode string or a list of modes.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="single" label="Single Default Mode">
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
|
|
@ -519,6 +522,32 @@ guardrails:
|
|||
default_on: true # run on every request
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="multi" label="Multiple Default Modes">
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: gpt-3.5-turbo
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "guardrails_ai-guard"
|
||||
litellm_params:
|
||||
guardrail: guardrails_ai
|
||||
guard_name: "pii_detect"
|
||||
mode:
|
||||
tags:
|
||||
"User-Agent: claude-cli": "logging_only"
|
||||
default: ["pre_call", "post_call"] # Run on both pre and post call when no tags match
|
||||
api_base: os.environ/GUARDRAILS_AI_API_BASE
|
||||
default_on: true
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
||||
### ✨ Model-level Guardrails
|
||||
|
||||
|
|
@ -640,13 +669,22 @@ guardrails:
|
|||
|
||||
Mode Specification
|
||||
|
||||
`default` accepts either a single string or a list of strings.
|
||||
|
||||
```python
|
||||
from litellm.types.guardrails import Mode
|
||||
|
||||
# Single default mode
|
||||
mode = Mode(
|
||||
tags={"User-Agent: claude-cli": "logging_only"},
|
||||
default="logging_only"
|
||||
)
|
||||
|
||||
# Multiple default modes
|
||||
mode = Mode(
|
||||
tags={"User-Agent: claude-cli": "logging_only"},
|
||||
default=["pre_call", "post_call"]
|
||||
)
|
||||
```
|
||||
|
||||
### `guardrails` Request Parameter
|
||||
|
|
|
|||
|
|
@ -347,3 +347,36 @@ If `order=1` deployment is unavailable (e.g., rate-limited), the router falls ba
|
|||
- **Higher throughput**: More requests handled simultaneously across deployments
|
||||
- **Improved reliability**: If one deployment fails, traffic automatically routes to healthy ones
|
||||
- **Better resource utilization**: Load spread evenly across all available deployments
|
||||
|
||||
## Special Considerations for Responses API
|
||||
|
||||
When load balancing OpenAI's Responses API across deployments with **different API keys** (e.g., different Azure regions or organizations), encrypted content items (like `rs_...` reasoning items) can only be decrypted by the originating API key.
|
||||
|
||||
**Solution:** Use the `encrypted_content_affinity` pre-call check to automatically route follow-up requests containing encrypted items to the correct deployment:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-5.1-codex
|
||||
litellm_params:
|
||||
model: azure/gpt-5.1-codex
|
||||
api_base: https://eastus.openai.azure.com/
|
||||
api_key: os.environ/AZURE_API_KEY_EASTUS
|
||||
model_info:
|
||||
id: "deployment-eastus"
|
||||
|
||||
- model_name: gpt-5.1-codex
|
||||
litellm_params:
|
||||
model: azure/gpt-5.1-codex
|
||||
api_base: https://westeurope.openai.azure.com/
|
||||
api_key: os.environ/AZURE_API_KEY_WESTEUROPE
|
||||
model_info:
|
||||
id: "deployment-westeurope"
|
||||
|
||||
router_settings:
|
||||
optional_pre_call_checks:
|
||||
- encrypted_content_affinity # 👈 Prevents invalid_encrypted_content errors
|
||||
```
|
||||
|
||||
This ensures requests containing encrypted content are routed to the deployment that created them, while other requests continue to load balance normally.
|
||||
|
||||
**[Learn more about Encrypted Content Affinity →](../response_api.md#encrypted-content-affinity-multi-region-load-balancing)**
|
||||
|
|
|
|||
|
|
@ -920,12 +920,17 @@ follow_up = await router.aresponses(
|
|||
To enable session continuity for Responses API in your LiteLLM proxy, set `optional_pre_call_checks` in your proxy config.yaml.
|
||||
|
||||
- `responses_api_deployment_check`: high priority routing when `previous_response_id` is provided
|
||||
- `encrypted_content_affinity`: **[Recommended]** content-aware routing for encrypted items (e.g., `rs_...` reasoning items)
|
||||
- `session_affinity`: sticky sessions based on session id (takes priority over `deployment_affinity`)
|
||||
- `deployment_affinity`: sticky sessions based on user key (applies even without `previous_response_id`)
|
||||
|
||||
:::tip Recommended: Use `encrypted_content_affinity`
|
||||
For Responses API with load balancing across deployments with **different API keys**, use `encrypted_content_affinity` instead of `deployment_affinity`. It only pins requests that contain encrypted content, avoiding quota reduction while preventing `invalid_encrypted_content` errors.
|
||||
:::
|
||||
|
||||
Notes:
|
||||
- User-key affinity is keyed on `metadata.user_api_key_hash` (the API key hash). The OpenAI `user` request parameter is an end-user identifier and is intentionally not used for deployment affinity.
|
||||
- Session-ID affinity is keyed on `metadata.session_id`. For proxy requests, this can be passed via the `x-litellm-session-id` HTTP header. For Python SDK requests, you can pass it via `litellm_metadata={"session_id": "value"}` in request args.
|
||||
- Session-ID affinity is keyed on `metadata.session_id`. For proxy requests, this can be passed via the `x-litellm-session-id` or `x-litellm-trace-id` HTTP header (they are interchangeable for call chaining). For Python SDK requests, you can pass it via `litellm_metadata={"session_id": "value"}` in request args.
|
||||
- `user_api_key_hash` is already SHA-256, and is used as-is (no double hashing).
|
||||
- Affinity is scoped by a stable model identifier (the model-map key, e.g. `model_map_information.model_map_key`) so model aliases map to the same stickiness bucket.
|
||||
- The mapping TTL is controlled by `deployment_affinity_ttl_seconds` (configured on Router init / proxy startup).
|
||||
|
|
@ -983,6 +988,142 @@ follow_up = client.responses.create(
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Encrypted Content Affinity (Multi-Region Load Balancing)
|
||||
|
||||
When load balancing Responses API across deployments with **different API keys** (e.g., different Azure regions or OpenAI organizations), encrypted content items (like `rs_...` reasoning items) can only be decrypted by the API key that created them.
|
||||
|
||||
### The Problem
|
||||
|
||||
```json
|
||||
{
|
||||
"error": {
|
||||
"message": "The encrypted content for item rs_0d09d6e56879e76500699d6feee41c8197bd268aae76141f87 could not be verified. Reason: Encrypted content organization_id did not match the target organization.",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_encrypted_content"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
This error occurs when:
|
||||
1. Initial request goes to Deployment A (API Key 1) → produces encrypted item `rs_xyz`
|
||||
2. Follow-up request with `rs_xyz` in input gets load balanced to Deployment B (API Key 2)
|
||||
3. Deployment B cannot decrypt content created by Deployment A → **request fails**
|
||||
|
||||
### The Solution: `encrypted_content_affinity`
|
||||
|
||||
The `encrypted_content_affinity` pre-call check routes follow-up requests containing encrypted items to the originating deployment **only when necessary**
|
||||
|
||||
**Key Benefits:**
|
||||
- ✅ **No quota reduction**: Unlike `deployment_affinity`, only pins requests that contain encrypted items
|
||||
- ✅ **Bypasses rate limits**: When encrypted content requires a specific deployment, RPM/TPM limits are bypassed (the request would fail on any other deployment anyway)
|
||||
- ✅ **No `previous_response_id` required**: Works by encoding `model_id` directly into item IDs
|
||||
- ✅ **No cache required**: `model_id` is decoded on-the-fly — no Redis dependency, no TTL to manage
|
||||
- ✅ **Globally safe**: Can be enabled for all models; non-Responses-API calls (chat, embeddings) are unaffected
|
||||
|
||||
### How It Works
|
||||
|
||||
1. **Encoding Phase** (on response):
|
||||
- For each output item that contains `encrypted_content`, LiteLLM rewrites the item ID to embed the originating `model_id`: `rs_xyz` → `encitem_{base64("litellm:model_id:{model_id};item_id:rs_xyz")}`
|
||||
- The original item ID is restored before forwarding the request to the upstream provider
|
||||
|
||||
2. **Routing Phase** (before request):
|
||||
- Scans request `input` for `encitem_` prefixed IDs
|
||||
- If found → decodes `model_id`, pins to originating deployment, bypasses rate limits
|
||||
- If no encoded items → normal load balancing
|
||||
|
||||
### Configuration
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="Python SDK">
|
||||
|
||||
```python
|
||||
from litellm import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5.1-codex",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-5.1-codex",
|
||||
"api_key": "org-1-api-key", # Different API key
|
||||
},
|
||||
"model_info": {"id": "deployment-us-east"},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-5.1-codex",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-5.1-codex",
|
||||
"api_key": "org-2-api-key", # Different API key
|
||||
},
|
||||
"model_info": {"id": "deployment-eu-west"},
|
||||
},
|
||||
],
|
||||
optional_pre_call_checks=["encrypted_content_affinity"],
|
||||
)
|
||||
|
||||
# Initial request - routes to any deployment
|
||||
response1 = await router.aresponses(
|
||||
model="gpt-5.1-codex",
|
||||
input="Explain quantum computing",
|
||||
)
|
||||
|
||||
# Follow-up with encrypted items - automatically routes to same deployment
|
||||
response2 = await router.aresponses(
|
||||
model="gpt-5.1-codex",
|
||||
input=response1.output, # Contains encrypted items from response1
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="Proxy Server">
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gpt-5.1-codex
|
||||
litellm_params:
|
||||
model: azure/gpt-5.1-codex
|
||||
api_base: https://eastus.openai.azure.com/
|
||||
api_key: os.environ/AZURE_API_KEY_EASTUS
|
||||
rpm: 600
|
||||
tpm: 100000
|
||||
model_info:
|
||||
id: "gpt-5.1-codex-eastus"
|
||||
|
||||
- model_name: gpt-5.1-codex
|
||||
litellm_params:
|
||||
model: azure/gpt-5.1-codex
|
||||
api_base: https://westeurope.openai.azure.com/
|
||||
api_key: os.environ/AZURE_API_KEY_WESTEUROPE
|
||||
rpm: 600
|
||||
tpm: 100000
|
||||
model_info:
|
||||
id: "gpt-5.1-codex-westeurope"
|
||||
|
||||
router_settings:
|
||||
routing_strategy: usage-based-routing-v2
|
||||
enable_pre_call_checks: true
|
||||
optional_pre_call_checks:
|
||||
- encrypted_content_affinity
|
||||
```
|
||||
|
||||
**Start proxy:**
|
||||
```bash
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### When to Use Each Affinity Type
|
||||
|
||||
| Affinity Type | Use Case | Scope | Quota Impact |
|
||||
|---------------|----------|-------|--------------|
|
||||
| **`encrypted_content_affinity`** | **[Recommended]** Multi-region Responses API with different API keys | Only requests with tracked encrypted items | ✅ None (surgical pinning) |
|
||||
| `responses_api_deployment_check` | When `previous_response_id` is available | Requests with `previous_response_id` | ✅ None |
|
||||
| `session_affinity` | Session-based applications | All requests with same `session_id` | ⚠️ Reduces quota by # of sessions |
|
||||
| `deployment_affinity` | Simple sticky sessions | All requests from same API key | ❌ Reduces quota by # of users |
|
||||
|
||||
|
||||
## Calling non-Responses API endpoints (`/responses` to `/chat/completions` Bridge)
|
||||
|
||||
LiteLLM allows you to call non-Responses API models via a bridge to LiteLLM's `/chat/completions` endpoint. This is useful for calling Anthropic, Gemini and even non-Responses API OpenAI models.
|
||||
|
|
|
|||
|
|
@ -10,10 +10,15 @@ class EnterpriseCustomGuardrailHelper:
|
|||
event_hook: Optional[
|
||||
Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]
|
||||
],
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
) -> Optional[bool]:
|
||||
"""
|
||||
Assumes check for event match is done in `should_run_guardrail`
|
||||
Returns True if the guardrail should be run by tag
|
||||
Returns True if the guardrail should be run for this request and event_type.
|
||||
|
||||
Logic:
|
||||
- If a request tag matches a Mode tag key, only run if event_type matches
|
||||
the tag's value (the mode for that tag).
|
||||
- If no request tag matches, fall back to default mode(s).
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
StandardLoggingPayloadSetup,
|
||||
|
|
@ -36,11 +41,29 @@ class EnterpriseCustomGuardrailHelper:
|
|||
proxy_server_request=proxy_server_request,
|
||||
)
|
||||
|
||||
if request_tags and any(tag in event_hook.tags for tag in request_tags):
|
||||
return True
|
||||
elif event_hook.default and any(
|
||||
tag in event_hook.default for tag in request_tags
|
||||
):
|
||||
# Check if any request tag matches a Mode tag key
|
||||
matched_mode = None
|
||||
if request_tags:
|
||||
for tag in request_tags:
|
||||
if tag in event_hook.tags:
|
||||
matched_mode = event_hook.tags[tag]
|
||||
break
|
||||
|
||||
if matched_mode is not None:
|
||||
# Tag matched: only run if event_type matches the tag's mode value
|
||||
if event_type is not None:
|
||||
return event_type.value == matched_mode
|
||||
return True
|
||||
|
||||
# No tag matched: fall back to default mode(s)
|
||||
if event_hook.default is not None:
|
||||
if event_type is not None:
|
||||
default_list = (
|
||||
event_hook.default
|
||||
if isinstance(event_hook.default, list)
|
||||
else [event_hook.default]
|
||||
)
|
||||
return event_type.value in default_list
|
||||
return False
|
||||
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -1,13 +1,13 @@
|
|||
"""
|
||||
AUDIT LOGGING
|
||||
|
||||
All /audit logging endpoints. Attempting to write these as CRUD endpoints.
|
||||
All /audit logging endpoints. Attempting to write these as CRUD endpoints.
|
||||
|
||||
GET - /audit/{id} - Get audit log by id
|
||||
GET - /audit - Get all audit logs
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
#### AUDIT LOGGING ####
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
|
@ -22,6 +22,27 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
def _build_json_field_or_condition(json_key: str, value: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Build an OR condition that matches a value inside a JSON column at the
|
||||
given key, checking both before_value and updated_values.
|
||||
|
||||
Uses Prisma's JSON path filtering (PostgreSQL only).
|
||||
|
||||
Example result (team_id="t1"):
|
||||
{"OR": [
|
||||
{"before_value": {"path": ["team_id"], "string_contains": "t1"}},
|
||||
{"updated_values": {"path": ["team_id"], "string_contains": "t1"}},
|
||||
]}
|
||||
"""
|
||||
return {
|
||||
"OR": [
|
||||
{"before_value": {"path": [json_key], "string_contains": value}},
|
||||
{"updated_values": {"path": [json_key], "string_contains": value}},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/audit",
|
||||
tags=["Audit Logging"],
|
||||
|
|
@ -49,6 +70,14 @@ async def get_audit_logs(
|
|||
),
|
||||
start_date: Optional[str] = Query(None, description="Filter logs after this date"),
|
||||
end_date: Optional[str] = Query(None, description="Filter logs before this date"),
|
||||
object_team_id: Optional[str] = Query(
|
||||
None,
|
||||
description="Filter by team_id present in before_value or updated_values JSON (PostgreSQL only)",
|
||||
),
|
||||
object_key_hash: Optional[str] = Query(
|
||||
None,
|
||||
description="Filter by token (key hash) present in before_value or updated_values JSON (PostgreSQL only)",
|
||||
),
|
||||
# Sorting parameters
|
||||
sort_by: Optional[str] = Query(
|
||||
None,
|
||||
|
|
@ -60,6 +89,9 @@ async def get_audit_logs(
|
|||
Get all audit logs with filtering and pagination.
|
||||
|
||||
Returns a paginated response of audit logs matching the specified filters.
|
||||
|
||||
Note: object_team_id and object_key_hash use Prisma JSON path filtering,
|
||||
which requires PostgreSQL.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
|
|
@ -82,18 +114,29 @@ async def get_audit_logs(
|
|||
if object_id:
|
||||
where_conditions["object_id"] = object_id
|
||||
if start_date or end_date:
|
||||
date_filter = {}
|
||||
date_filter: Dict[str, Any] = {}
|
||||
if start_date:
|
||||
date_filter["gte"] = start_date
|
||||
if end_date:
|
||||
date_filter["lte"] = end_date
|
||||
where_conditions["updated_at"] = date_filter
|
||||
|
||||
# JSON field filters (PostgreSQL only) — each filter is AND'd with the
|
||||
# others, but checks both before_value and updated_values internally (OR).
|
||||
if object_team_id:
|
||||
where_conditions["AND"] = where_conditions.get("AND", []) + [
|
||||
_build_json_field_or_condition("team_id", object_team_id)
|
||||
]
|
||||
if object_key_hash:
|
||||
where_conditions["AND"] = where_conditions.get("AND", []) + [
|
||||
_build_json_field_or_condition("token", object_key_hash)
|
||||
]
|
||||
|
||||
# Build sort conditions
|
||||
order_by = {}
|
||||
order_by: Dict[str, Any] = {}
|
||||
if sort_by and isinstance(sort_by, str):
|
||||
order_by[sort_by] = sort_order
|
||||
elif sort_order and isinstance(sort_order, str):
|
||||
else:
|
||||
order_by["updated_at"] = sort_order # Default sort by updated_at
|
||||
|
||||
# Get paginated results
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "blocked_tools" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -0,0 +1,11 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_SpendLogToolIndex" (
|
||||
"request_id" TEXT NOT NULL,
|
||||
"tool_name" TEXT NOT NULL,
|
||||
"start_time" TIMESTAMP(3) NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_SpendLogToolIndex_pkey" PRIMARY KEY ("request_id","tool_name")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_SpendLogToolIndex_tool_name_start_time_idx" ON "LiteLLM_SpendLogToolIndex"("tool_name", "start_time");
|
||||
|
|
@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
vector_stores String[] @default([])
|
||||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
@ -928,6 +929,16 @@ model LiteLLM_SpendLogGuardrailIndex {
|
|||
@@index([policy_id, start_time])
|
||||
}
|
||||
|
||||
// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production
|
||||
model LiteLLM_SpendLogToolIndex {
|
||||
request_id String
|
||||
tool_name String // matches LiteLLM_ToolTable.tool_name; join for input_policy/output_policy etc.
|
||||
start_time DateTime
|
||||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
model LiteLLM_PromptTable {
|
||||
id String @id @default(uuid())
|
||||
|
|
@ -1065,26 +1076,31 @@ model LiteLLM_PolicyAttachmentTable {
|
|||
updated_by String?
|
||||
}
|
||||
|
||||
// Global tool registry - auto-discovered from LLM responses; admins set call_policy here
|
||||
// Global tool registry - auto-discovered from LLM responses; admins set input_policy/output_policy here
|
||||
model LiteLLM_ToolTable {
|
||||
tool_id String @id @default(uuid())
|
||||
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
|
||||
origin String? // MCP server name or "user_defined"
|
||||
call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked"
|
||||
call_count Int @default(0) // cumulative number of times this tool was seen
|
||||
assignments Json? @default("{}")
|
||||
key_hash String? // hash of the virtual key that first called this tool
|
||||
team_id String? // team that first called this tool
|
||||
key_alias String? // human-readable alias of the virtual key
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
updated_at DateTime @default(now()) @updatedAt
|
||||
updated_by String?
|
||||
tool_id String @id @default(uuid())
|
||||
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
|
||||
origin String? // MCP server name or "user_defined"
|
||||
input_policy String @default("untrusted") // "trusted" | "untrusted" | "blocked"
|
||||
output_policy String @default("untrusted") // "trusted" | "untrusted"
|
||||
call_count Int @default(0) // cumulative number of times this tool was seen
|
||||
assignments Json? @default("{}")
|
||||
key_hash String? // hash of the virtual key that first called this tool
|
||||
team_id String? // team that first called this tool
|
||||
key_alias String? // human-readable alias of the virtual key
|
||||
user_agent String? // user-agent of the first request that discovered this tool
|
||||
last_used_at DateTime? // timestamp of the most recent call
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
updated_at DateTime @default(now()) @updatedAt
|
||||
updated_by String?
|
||||
|
||||
@@index([call_policy])
|
||||
@@index([input_policy])
|
||||
@@index([output_policy])
|
||||
@@index([team_id])
|
||||
}
|
||||
|
||||
// Per-(tool, team/key) policy overrides. When present, override replaces global tool policy for that scope.
|
||||
//Unified Access Groups table for storing unified access groups
|
||||
model LiteLLM_AccessGroupTable {
|
||||
access_group_id String @id @default(uuid())
|
||||
|
|
|
|||
|
|
@ -24,11 +24,7 @@ from litellm.utils import client
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.client import A2AClient as A2AClientType
|
||||
from a2a.types import (
|
||||
AgentCard,
|
||||
SendMessageRequest,
|
||||
SendStreamingMessageRequest,
|
||||
)
|
||||
from a2a.types import AgentCard, SendMessageRequest, SendStreamingMessageRequest
|
||||
|
||||
# Runtime imports with availability check
|
||||
A2A_SDK_AVAILABLE = False
|
||||
|
|
@ -124,13 +120,48 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str:
|
|||
litellm_logging_obj.model = model
|
||||
litellm_logging_obj.custom_llm_provider = custom_llm_provider
|
||||
litellm_logging_obj.model_call_details["model"] = model
|
||||
litellm_logging_obj.model_call_details[
|
||||
"custom_llm_provider"
|
||||
] = custom_llm_provider
|
||||
litellm_logging_obj.model_call_details["custom_llm_provider"] = (
|
||||
custom_llm_provider
|
||||
)
|
||||
|
||||
return agent_name
|
||||
|
||||
|
||||
async def _send_message_via_completion_bridge(
|
||||
request: "SendMessageRequest",
|
||||
custom_llm_provider: str,
|
||||
api_base: Optional[str],
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
"""
|
||||
Route a send_message through the LiteLLM completion bridge (e.g. LangGraph, Bedrock AgentCore).
|
||||
|
||||
Requires request; api_base is optional for providers that derive endpoint from model.
|
||||
"""
|
||||
verbose_logger.info(
|
||||
f"A2A using completion bridge: provider={custom_llm_provider}, api_base={api_base}"
|
||||
)
|
||||
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
params = (
|
||||
request.params.model_dump(mode="json")
|
||||
if hasattr(request.params, "model_dump")
|
||||
else dict(request.params)
|
||||
)
|
||||
|
||||
response_dict = await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
request_id=str(request.id),
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
return LiteLLMSendMessageResponse.from_dict(response_dict)
|
||||
|
||||
|
||||
@client
|
||||
async def asend_message(
|
||||
a2a_client: Optional["A2AClientType"] = None,
|
||||
|
|
@ -193,39 +224,21 @@ async def asend_message(
|
|||
```
|
||||
"""
|
||||
litellm_params = litellm_params or {}
|
||||
logging_obj = kwargs.get("litellm_logging_obj")
|
||||
trace_id = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
|
||||
# Route through completion bridge if custom_llm_provider is set
|
||||
if custom_llm_provider:
|
||||
if request is None:
|
||||
raise ValueError("request is required for completion bridge")
|
||||
# api_base is optional for providers that derive endpoint from model (e.g., bedrock/agentcore)
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A using completion bridge: provider={custom_llm_provider}, api_base={api_base}"
|
||||
)
|
||||
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
# Extract params from request
|
||||
params = (
|
||||
request.params.model_dump(mode="json")
|
||||
if hasattr(request.params, "model_dump")
|
||||
else dict(request.params)
|
||||
)
|
||||
|
||||
response_dict = await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
request_id=str(request.id),
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
return await _send_message_via_completion_bridge(
|
||||
request=request,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
# Convert to LiteLLMSendMessageResponse
|
||||
return LiteLLMSendMessageResponse.from_dict(response_dict)
|
||||
|
||||
# Standard A2A client flow
|
||||
if request is None:
|
||||
raise ValueError("request is required")
|
||||
|
|
@ -236,11 +249,13 @@ async def asend_message(
|
|||
raise ValueError(
|
||||
"Either a2a_client or api_base is required for standard A2A flow"
|
||||
)
|
||||
trace_id = str(uuid.uuid4())
|
||||
trace_id = trace_id or str(uuid.uuid4())
|
||||
extra_headers = {"X-LiteLLM-Trace-Id": trace_id}
|
||||
if agent_id:
|
||||
extra_headers["X-LiteLLM-Agent-Id"] = agent_id
|
||||
a2a_client = await create_a2a_client(base_url=api_base, extra_headers=extra_headers)
|
||||
a2a_client = await create_a2a_client(
|
||||
base_url=api_base, extra_headers=extra_headers
|
||||
)
|
||||
|
||||
# Type assertion: a2a_client is guaranteed to be non-None here
|
||||
assert a2a_client is not None
|
||||
|
|
@ -255,6 +270,10 @@ async def asend_message(
|
|||
)
|
||||
card_url = getattr(agent_card, "url", None) if agent_card else None
|
||||
|
||||
context_id = trace_id or str(uuid.uuid4())
|
||||
if request.params.message.context_id is None:
|
||||
request.params.message.context_id = context_id
|
||||
|
||||
# Retry loop: if connection fails due to localhost URL in agent card, retry with fixed URL
|
||||
a2a_response = None
|
||||
for _ in range(2): # max 2 attempts: original + 1 retry
|
||||
|
|
@ -606,7 +625,9 @@ async def create_a2a_client(
|
|||
|
||||
if extra_headers:
|
||||
httpx_client.headers.update(extra_headers)
|
||||
verbose_proxy_logger.debug(f"A2A client created with extra_headers={extra_headers}")
|
||||
verbose_proxy_logger.debug(
|
||||
f"A2A client created with extra_headers={extra_headers}"
|
||||
)
|
||||
|
||||
# Resolve agent card
|
||||
resolver = A2ACardResolver(
|
||||
|
|
|
|||
|
|
@ -112,6 +112,7 @@ async def acreate_batch(
|
|||
metadata: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
output_expires_after: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> LiteLLMBatch:
|
||||
"""
|
||||
|
|
@ -133,6 +134,7 @@ async def acreate_batch(
|
|||
metadata,
|
||||
extra_headers,
|
||||
extra_body,
|
||||
output_expires_after,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -152,7 +154,7 @@ async def acreate_batch(
|
|||
|
||||
|
||||
@client
|
||||
def create_batch(
|
||||
def create_batch( # noqa: PLR0915
|
||||
completion_window: Literal["24h"],
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"],
|
||||
input_file_id: str,
|
||||
|
|
@ -160,6 +162,7 @@ def create_batch(
|
|||
metadata: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
output_expires_after: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]:
|
||||
"""
|
||||
|
|
@ -215,6 +218,8 @@ def create_batch(
|
|||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
if output_expires_after is not None:
|
||||
_create_batch_request["output_expires_after"] = output_expires_after
|
||||
if model is not None:
|
||||
provider_config = ProviderConfigManager.get_provider_batches_config(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -278,7 +278,7 @@ async def afile_retrieve(
|
|||
@client
|
||||
def file_retrieve(
|
||||
file_id: str,
|
||||
custom_llm_provider: Literal["openai", "azure", "hosted_vllm", "manus"] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus"] = "openai",
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
**kwargs,
|
||||
|
|
|
|||
|
|
@ -34,6 +34,44 @@ vertex_fine_tuning_apis_instance = VertexFineTuningAPI()
|
|||
#################################################
|
||||
|
||||
|
||||
def _prepare_azure_extra_body(
|
||||
extra_body: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
azure_specific_hyperparams: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Prepare extra_body for Azure fine-tuning API by combining Azure-specific parameters.
|
||||
|
||||
Azure fine-tuning API accepts additional parameters beyond the standard OpenAI spec:
|
||||
- trainingType: Type of training (e.g., 1 for supervised fine-tuning)
|
||||
- prompt_loss_weight: Weight for prompt loss in training
|
||||
|
||||
These parameters must be passed in the extra_body field when calling the Azure OpenAI SDK.
|
||||
|
||||
Args:
|
||||
extra_body: Optional existing extra_body dict
|
||||
kwargs: Request kwargs that may contain Azure-specific parameters
|
||||
azure_specific_hyperparams: Dict of Azure-specific hyperparameters already extracted
|
||||
|
||||
Returns:
|
||||
Dict containing all Azure-specific parameters to be passed in extra_body
|
||||
"""
|
||||
if extra_body is None:
|
||||
extra_body = {}
|
||||
|
||||
# Azure-specific root-level parameters
|
||||
azure_specific_params = ["trainingType"]
|
||||
for param in azure_specific_params:
|
||||
if param in kwargs:
|
||||
extra_body[param] = kwargs[param]
|
||||
|
||||
# Add Azure-specific hyperparameters
|
||||
if azure_specific_hyperparams:
|
||||
extra_body.update(azure_specific_hyperparams)
|
||||
|
||||
return extra_body
|
||||
|
||||
|
||||
@client
|
||||
async def acreate_fine_tuning_job(
|
||||
model: str,
|
||||
|
|
@ -114,6 +152,15 @@ def create_fine_tuning_job(
|
|||
|
||||
# handle hyperparameters
|
||||
hyperparameters = hyperparameters or {} # original hyperparameters
|
||||
|
||||
# For Azure, extract Azure-specific hyperparameters before creating OpenAI-spec hyperparameters
|
||||
azure_specific_hyperparams = {}
|
||||
if custom_llm_provider == "azure":
|
||||
azure_hyperparameter_keys = ["prompt_loss_weight"]
|
||||
for key in azure_hyperparameter_keys:
|
||||
if key in hyperparameters:
|
||||
azure_specific_hyperparams[key] = hyperparameters.pop(key)
|
||||
|
||||
_oai_hyperparameters: Hyperparameters = Hyperparameters(
|
||||
**hyperparameters
|
||||
) # Typed Hyperparameters for OpenAI Spec
|
||||
|
|
@ -207,6 +254,10 @@ def create_fine_tuning_job(
|
|||
extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
get_secret_str("AZURE_AD_TOKEN") # type: ignore
|
||||
|
||||
# Prepare Azure-specific parameters for extra_body
|
||||
extra_body = _prepare_azure_extra_body(extra_body, kwargs, azure_specific_hyperparams)
|
||||
|
||||
create_fine_tuning_job_data = FineTuningJobCreate(
|
||||
model=model,
|
||||
training_file=training_file,
|
||||
|
|
@ -220,6 +271,10 @@ def create_fine_tuning_job(
|
|||
create_fine_tuning_job_data_dict = create_fine_tuning_job_data.model_dump(
|
||||
exclude_none=True
|
||||
)
|
||||
|
||||
# Add extra_body if it has Azure-specific parameters
|
||||
if extra_body:
|
||||
create_fine_tuning_job_data_dict["extra_body"] = extra_body
|
||||
|
||||
response = azure_fine_tuning_apis_instance.create_fine_tuning_job(
|
||||
api_base=api_base,
|
||||
|
|
|
|||
|
|
@ -235,8 +235,13 @@ class CustomGuardrail(CustomLogger):
|
|||
list(event_hook.tags.values()), supported_event_hooks
|
||||
)
|
||||
if event_hook.default:
|
||||
default_list = (
|
||||
event_hook.default
|
||||
if isinstance(event_hook.default, list)
|
||||
else [event_hook.default]
|
||||
)
|
||||
_validate_event_hook_list_is_in_supported_event_hooks(
|
||||
[event_hook.default], supported_event_hooks
|
||||
default_list, supported_event_hooks
|
||||
)
|
||||
elif isinstance(event_hook, GuardrailEventHooks):
|
||||
if event_hook not in supported_event_hooks:
|
||||
|
|
@ -415,7 +420,7 @@ class CustomGuardrail(CustomLogger):
|
|||
"Setting tag-based guardrails is only available in litellm-enterprise. You must be a premium user to use this feature."
|
||||
)
|
||||
result = EnterpriseCustomGuardrailHelper._should_run_if_mode_by_tag(
|
||||
data, self.event_hook
|
||||
data, self.event_hook, event_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
|
|
@ -442,7 +447,7 @@ class CustomGuardrail(CustomLogger):
|
|||
"Setting tag-based guardrails is only available in litellm-enterprise. You must be a premium user to use this feature."
|
||||
)
|
||||
result = EnterpriseCustomGuardrailHelper._should_run_if_mode_by_tag(
|
||||
data, self.event_hook
|
||||
data, self.event_hook, event_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
|
|
@ -461,7 +466,16 @@ class CustomGuardrail(CustomLogger):
|
|||
if isinstance(self.event_hook, list):
|
||||
return event_type.value in self.event_hook
|
||||
if isinstance(self.event_hook, Mode):
|
||||
return event_type.value in self.event_hook.tags.values()
|
||||
if event_type.value in self.event_hook.tags.values():
|
||||
return True
|
||||
if self.event_hook.default:
|
||||
default_list = (
|
||||
self.event_hook.default
|
||||
if isinstance(self.event_hook.default, list)
|
||||
else [self.event_hook.default]
|
||||
)
|
||||
return event_type.value in default_list
|
||||
return False
|
||||
return self.event_hook == event_type.value
|
||||
|
||||
def get_guardrail_dynamic_request_body_params(self, request_data: dict) -> dict:
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
import json
|
||||
import re
|
||||
import traceback
|
||||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
import re
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -443,6 +443,27 @@ def exception_type( # type: ignore # noqa: PLR0915
|
|||
response=getattr(original_exception, "response", None),
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif "invalid_encrypted_content" in error_str or "could not be verified" in error_str:
|
||||
exception_mapping_worked = True
|
||||
helpful_message = (
|
||||
f"{exception_provider} - {message}\n\n"
|
||||
" This error occurs when load balancing Responses API across deployments with different API keys.\n"
|
||||
" Encrypted content is tied to the organization that created it and cannot be decrypted by other organizations.\n\n"
|
||||
" Solution: Enable 'encrypted_content_affinity' to route follow-up requests to the correct deployment:\n\n"
|
||||
" router_settings:\n"
|
||||
" enable_pre_call_checks: true\n"
|
||||
" optional_pre_call_checks:\n"
|
||||
" - encrypted_content_affinity\n\n"
|
||||
" Learn more: https://docs.litellm.ai/docs/response_api#encrypted-content-affinity-multi-region-load-balancing"
|
||||
)
|
||||
raise BadRequestError(
|
||||
message=helpful_message,
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
litellm_debug_info=extra_information,
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
elif (
|
||||
"invalid_request_error" in error_str
|
||||
and "Incorrect API key provided" not in error_str
|
||||
|
|
@ -2126,7 +2147,27 @@ def exception_type( # type: ignore # noqa: PLR0915
|
|||
extra_information=extra_information,
|
||||
original_exception=original_exception,
|
||||
)
|
||||
|
||||
elif azure_error_code == "invalid_encrypted_content" or "could not be verified" in error_str:
|
||||
exception_mapping_worked = True
|
||||
helpful_message = (
|
||||
f"AzureException - {message}\n\n"
|
||||
"This error occurs when load balancing Responses API across deployments with different API keys.\n"
|
||||
" Encrypted content is tied to the organization that created it and cannot be decrypted by other organizations.\n\n"
|
||||
" Solution: Enable 'encrypted_content_affinity' to route follow-up requests to the correct deployment:\n\n"
|
||||
" router_settings:\n"
|
||||
" enable_pre_call_checks: true\n"
|
||||
" optional_pre_call_checks:\n"
|
||||
" - encrypted_content_affinity\n\n"
|
||||
" Learn more: https://docs.litellm.ai/docs/response_api#encrypted-content-affinity-multi-region-load-balancing"
|
||||
)
|
||||
raise BadRequestError(
|
||||
message=helpful_message,
|
||||
llm_provider="azure",
|
||||
model=model,
|
||||
litellm_debug_info=extra_information,
|
||||
response=getattr(original_exception, "response", None),
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
elif "invalid_request_error" in error_str:
|
||||
exception_mapping_worked = True
|
||||
raise BadRequestError(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import Optional
|
||||
|
||||
|
||||
# Pre-define optional kwargs keys as frozenset for O(1) lookups
|
||||
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
|
||||
_OPTIONAL_KWARGS_KEYS = frozenset({
|
||||
|
|
@ -95,6 +94,13 @@ def get_litellm_params(
|
|||
litellm_request_debug: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
# Derive litellm_session_id / litellm_trace_id from metadata when not provided (call chaining)
|
||||
_meta = metadata or {}
|
||||
if litellm_session_id is None:
|
||||
litellm_session_id = _meta.get("session_id") or _meta.get("trace_id")
|
||||
if litellm_trace_id is None:
|
||||
litellm_trace_id = _meta.get("trace_id") or _meta.get("session_id")
|
||||
|
||||
# Build base dict with explicit parameters (always included)
|
||||
litellm_params = {
|
||||
"acompletion": acompletion,
|
||||
|
|
|
|||
|
|
@ -133,8 +133,8 @@ from ..integrations.azure_sentinel.azure_sentinel import AzureSentinelLogger
|
|||
from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger
|
||||
from ..integrations.custom_prompt_management import CustomPromptManagement
|
||||
from ..integrations.datadog.datadog import DataDogLogger
|
||||
from ..integrations.datadog.datadog_metrics import DatadogMetricsLogger
|
||||
from ..integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
|
||||
from ..integrations.datadog.datadog_metrics import DatadogMetricsLogger
|
||||
from ..integrations.dotprompt import DotpromptManager
|
||||
from ..integrations.dynamodb import DyanmoDBLogger
|
||||
from ..integrations.galileo import GalileoObserve
|
||||
|
|
@ -352,9 +352,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
self.function_id = function_id
|
||||
self.streaming_chunks: List[Any] = [] # for generating complete stream response
|
||||
self.sync_streaming_chunks: List[
|
||||
Any
|
||||
] = [] # for generating complete stream response
|
||||
self.sync_streaming_chunks: List[Any] = (
|
||||
[]
|
||||
) # for generating complete stream response
|
||||
self.log_raw_request_response = log_raw_request_response
|
||||
|
||||
# Initialize dynamic callbacks
|
||||
|
|
@ -746,9 +746,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
prompt_spec=prompt_spec,
|
||||
dynamic_callback_params=dynamic_callback_params,
|
||||
):
|
||||
self.model_call_details[
|
||||
"prompt_integration"
|
||||
] = logger.__class__.__name__
|
||||
self.model_call_details["prompt_integration"] = (
|
||||
logger.__class__.__name__
|
||||
)
|
||||
return logger
|
||||
except Exception:
|
||||
# If check fails, continue to next logger
|
||||
|
|
@ -816,9 +816,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if anthropic_cache_control_logger := AnthropicCacheControlHook.get_custom_logger_for_anthropic_cache_control_hook(
|
||||
non_default_params
|
||||
):
|
||||
self.model_call_details[
|
||||
"prompt_integration"
|
||||
] = anthropic_cache_control_logger.__class__.__name__
|
||||
self.model_call_details["prompt_integration"] = (
|
||||
anthropic_cache_control_logger.__class__.__name__
|
||||
)
|
||||
return anthropic_cache_control_logger
|
||||
|
||||
#########################################################
|
||||
|
|
@ -830,9 +830,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
internal_usage_cache=None,
|
||||
llm_router=None,
|
||||
)
|
||||
self.model_call_details[
|
||||
"prompt_integration"
|
||||
] = vector_store_custom_logger.__class__.__name__
|
||||
self.model_call_details["prompt_integration"] = (
|
||||
vector_store_custom_logger.__class__.__name__
|
||||
)
|
||||
# Add to global callbacks so post-call hooks are invoked
|
||||
if (
|
||||
vector_store_custom_logger
|
||||
|
|
@ -892,9 +892,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
model
|
||||
): # if model name was changes pre-call, overwrite the initial model call name with the new one
|
||||
self.model_call_details["model"] = model
|
||||
self.model_call_details["litellm_params"][
|
||||
"api_base"
|
||||
] = self._get_masked_api_base(additional_args.get("api_base", ""))
|
||||
self.model_call_details["litellm_params"]["api_base"] = (
|
||||
self._get_masked_api_base(additional_args.get("api_base", ""))
|
||||
)
|
||||
|
||||
def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915
|
||||
# Log the exact input to the LLM API
|
||||
|
|
@ -923,10 +923,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
try:
|
||||
# [Non-blocking Extra Debug Information in metadata]
|
||||
if turn_off_message_logging is True:
|
||||
_metadata[
|
||||
"raw_request"
|
||||
] = "redacted by litellm. \
|
||||
_metadata["raw_request"] = (
|
||||
"redacted by litellm. \
|
||||
'litellm.turn_off_message_logging=True'"
|
||||
)
|
||||
else:
|
||||
curl_command = self._get_request_curl_command(
|
||||
api_base=additional_args.get("api_base", ""),
|
||||
|
|
@ -937,34 +937,34 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
_metadata["raw_request"] = str(curl_command)
|
||||
# split up, so it's easier to parse in the UI
|
||||
self.model_call_details[
|
||||
"raw_request_typed_dict"
|
||||
] = RawRequestTypedDict(
|
||||
raw_request_api_base=str(
|
||||
additional_args.get("api_base") or ""
|
||||
),
|
||||
raw_request_body=self._get_raw_request_body(
|
||||
additional_args.get("complete_input_dict", {})
|
||||
),
|
||||
# NOTE: setting ignore_sensitive_headers to True will cause
|
||||
# the Authorization header to be leaked when calls to the health
|
||||
# endpoint are made and fail.
|
||||
raw_request_headers=self._get_masked_headers(
|
||||
additional_args.get("headers", {}) or {},
|
||||
),
|
||||
error=None,
|
||||
self.model_call_details["raw_request_typed_dict"] = (
|
||||
RawRequestTypedDict(
|
||||
raw_request_api_base=str(
|
||||
additional_args.get("api_base") or ""
|
||||
),
|
||||
raw_request_body=self._get_raw_request_body(
|
||||
additional_args.get("complete_input_dict", {})
|
||||
),
|
||||
# NOTE: setting ignore_sensitive_headers to True will cause
|
||||
# the Authorization header to be leaked when calls to the health
|
||||
# endpoint are made and fail.
|
||||
raw_request_headers=self._get_masked_headers(
|
||||
additional_args.get("headers", {}) or {},
|
||||
),
|
||||
error=None,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
self.model_call_details[
|
||||
"raw_request_typed_dict"
|
||||
] = RawRequestTypedDict(
|
||||
error=str(e),
|
||||
self.model_call_details["raw_request_typed_dict"] = (
|
||||
RawRequestTypedDict(
|
||||
error=str(e),
|
||||
)
|
||||
)
|
||||
_metadata[
|
||||
"raw_request"
|
||||
] = "Unable to Log \
|
||||
_metadata["raw_request"] = (
|
||||
"Unable to Log \
|
||||
raw request: {}".format(
|
||||
str(e)
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
|
||||
try:
|
||||
|
|
@ -1265,13 +1265,13 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
for callback in callbacks:
|
||||
try:
|
||||
if isinstance(callback, CustomLogger):
|
||||
response: Optional[
|
||||
MCPPostCallResponseObject
|
||||
] = await callback.async_post_mcp_tool_call_hook(
|
||||
kwargs=kwargs,
|
||||
response_obj=post_mcp_tool_call_response_obj,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
response: Optional[MCPPostCallResponseObject] = (
|
||||
await callback.async_post_mcp_tool_call_hook(
|
||||
kwargs=kwargs,
|
||||
response_obj=post_mcp_tool_call_response_obj,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
######################################################################
|
||||
# if any of the callbacks modify the response, use the modified response
|
||||
|
|
@ -1466,9 +1466,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
verbose_logger.debug(
|
||||
f"response_cost_failure_debug_information: {debug_info}"
|
||||
)
|
||||
self.model_call_details[
|
||||
"response_cost_failure_debug_information"
|
||||
] = debug_info
|
||||
self.model_call_details["response_cost_failure_debug_information"] = (
|
||||
debug_info
|
||||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
|
|
@ -1494,9 +1494,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
verbose_logger.debug(
|
||||
f"response_cost_failure_debug_information: {debug_info}"
|
||||
)
|
||||
self.model_call_details[
|
||||
"response_cost_failure_debug_information"
|
||||
] = debug_info
|
||||
self.model_call_details["response_cost_failure_debug_information"] = (
|
||||
debug_info
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
|
|
@ -1652,10 +1652,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
result=logging_result
|
||||
)
|
||||
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = self._build_standard_logging_payload(
|
||||
logging_result, start_time, end_time
|
||||
self.model_call_details["standard_logging_object"] = (
|
||||
self._build_standard_logging_payload(logging_result, start_time, end_time)
|
||||
)
|
||||
|
||||
if (
|
||||
|
|
@ -1734,9 +1732,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
end_time = datetime.datetime.now()
|
||||
if self.completion_start_time is None:
|
||||
self.completion_start_time = end_time
|
||||
self.model_call_details[
|
||||
"completion_start_time"
|
||||
] = self.completion_start_time
|
||||
self.model_call_details["completion_start_time"] = (
|
||||
self.completion_start_time
|
||||
)
|
||||
|
||||
self.model_call_details["log_event_type"] = "successful_api_call"
|
||||
self.model_call_details["end_time"] = end_time
|
||||
|
|
@ -1773,10 +1771,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
end_time=end_time,
|
||||
)
|
||||
elif isinstance(result, dict) or isinstance(result, list):
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = self._build_standard_logging_payload(
|
||||
result, start_time, end_time
|
||||
self.model_call_details["standard_logging_object"] = (
|
||||
self._build_standard_logging_payload(
|
||||
result, start_time, end_time
|
||||
)
|
||||
)
|
||||
if (
|
||||
standard_logging_payload := self.model_call_details.get(
|
||||
|
|
@ -1785,9 +1783,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
) is not None:
|
||||
emit_standard_logging_payload(standard_logging_payload)
|
||||
elif standard_logging_object is not None:
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = standard_logging_object
|
||||
self.model_call_details["standard_logging_object"] = (
|
||||
standard_logging_object
|
||||
)
|
||||
else:
|
||||
self.model_call_details["response_cost"] = None
|
||||
|
||||
|
|
@ -1945,17 +1943,17 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
verbose_logger.debug(
|
||||
"Logging Details LiteLLM-Success Call streaming complete"
|
||||
)
|
||||
self.model_call_details[
|
||||
"complete_streaming_response"
|
||||
] = complete_streaming_response
|
||||
self.model_call_details[
|
||||
"response_cost"
|
||||
] = self._response_cost_calculator(result=complete_streaming_response)
|
||||
self.model_call_details["complete_streaming_response"] = (
|
||||
complete_streaming_response
|
||||
)
|
||||
self.model_call_details["response_cost"] = (
|
||||
self._response_cost_calculator(result=complete_streaming_response)
|
||||
)
|
||||
## STANDARDIZED LOGGING PAYLOAD
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = self._build_standard_logging_payload(
|
||||
complete_streaming_response, start_time, end_time
|
||||
self.model_call_details["standard_logging_object"] = (
|
||||
self._build_standard_logging_payload(
|
||||
complete_streaming_response, start_time, end_time
|
||||
)
|
||||
)
|
||||
if (
|
||||
standard_logging_payload := self.model_call_details.get(
|
||||
|
|
@ -2289,10 +2287,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
else:
|
||||
if self.stream and complete_streaming_response:
|
||||
self.model_call_details[
|
||||
"complete_response"
|
||||
] = self.model_call_details.get(
|
||||
"complete_streaming_response", {}
|
||||
self.model_call_details["complete_response"] = (
|
||||
self.model_call_details.get(
|
||||
"complete_streaming_response", {}
|
||||
)
|
||||
)
|
||||
result = self.model_call_details["complete_response"]
|
||||
openMeterLogger.log_success_event(
|
||||
|
|
@ -2316,10 +2314,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
else:
|
||||
if self.stream and complete_streaming_response:
|
||||
self.model_call_details[
|
||||
"complete_response"
|
||||
] = self.model_call_details.get(
|
||||
"complete_streaming_response", {}
|
||||
self.model_call_details["complete_response"] = (
|
||||
self.model_call_details.get(
|
||||
"complete_streaming_response", {}
|
||||
)
|
||||
)
|
||||
result = self.model_call_details["complete_response"]
|
||||
|
||||
|
|
@ -2458,9 +2456,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if complete_streaming_response is not None:
|
||||
print_verbose("Async success callbacks: Got a complete streaming response")
|
||||
|
||||
self.model_call_details[
|
||||
"async_complete_streaming_response"
|
||||
] = complete_streaming_response
|
||||
self.model_call_details["async_complete_streaming_response"] = (
|
||||
complete_streaming_response
|
||||
)
|
||||
|
||||
try:
|
||||
if self.model_call_details.get("cache_hit", False) is True:
|
||||
|
|
@ -2471,10 +2469,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
model_call_details=self.model_call_details
|
||||
)
|
||||
# base_model defaults to None if not set on model_info
|
||||
self.model_call_details[
|
||||
"response_cost"
|
||||
] = self._response_cost_calculator(
|
||||
result=complete_streaming_response
|
||||
self.model_call_details["response_cost"] = (
|
||||
self._response_cost_calculator(
|
||||
result=complete_streaming_response
|
||||
)
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
|
|
@ -2487,10 +2485,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details["response_cost"] = None
|
||||
|
||||
## STANDARDIZED LOGGING PAYLOAD
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = self._build_standard_logging_payload(
|
||||
complete_streaming_response, start_time, end_time
|
||||
self.model_call_details["standard_logging_object"] = (
|
||||
self._build_standard_logging_payload(
|
||||
complete_streaming_response, start_time, end_time
|
||||
)
|
||||
)
|
||||
|
||||
# print standard logging payload
|
||||
|
|
@ -2517,10 +2515,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
# _success_handler_helper_fn
|
||||
if self.model_call_details.get("standard_logging_object") is None:
|
||||
## STANDARDIZED LOGGING PAYLOAD
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = self._build_standard_logging_payload(
|
||||
result, start_time, end_time
|
||||
self.model_call_details["standard_logging_object"] = (
|
||||
self._build_standard_logging_payload(result, start_time, end_time)
|
||||
)
|
||||
|
||||
# print standard logging payload
|
||||
|
|
@ -2764,18 +2760,18 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
## STANDARDIZED LOGGING PAYLOAD
|
||||
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = get_standard_logging_object_payload(
|
||||
kwargs=self.model_call_details,
|
||||
init_response_obj={},
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
logging_obj=self,
|
||||
status="failure",
|
||||
error_str=str(exception),
|
||||
original_exception=exception,
|
||||
standard_built_in_tools_params=self.standard_built_in_tools_params,
|
||||
self.model_call_details["standard_logging_object"] = (
|
||||
get_standard_logging_object_payload(
|
||||
kwargs=self.model_call_details,
|
||||
init_response_obj={},
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
logging_obj=self,
|
||||
status="failure",
|
||||
error_str=str(exception),
|
||||
original_exception=exception,
|
||||
standard_built_in_tools_params=self.standard_built_in_tools_params,
|
||||
)
|
||||
)
|
||||
return start_time, end_time
|
||||
|
||||
|
|
@ -3739,9 +3735,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
service_name=arize_config.project_name,
|
||||
)
|
||||
|
||||
os.environ[
|
||||
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
|
||||
] = f"space_id={arize_config.space_key or arize_config.space_id},api_key={arize_config.api_key}"
|
||||
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
|
||||
f"space_id={arize_config.space_key or arize_config.space_id},api_key={arize_config.api_key}"
|
||||
)
|
||||
for callback in _in_memory_loggers:
|
||||
if (
|
||||
isinstance(callback, ArizeLogger)
|
||||
|
|
@ -3767,13 +3763,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
existing_attrs = os.environ.get("OTEL_RESOURCE_ATTRIBUTES", "")
|
||||
# Add openinference.project.name attribute
|
||||
if existing_attrs:
|
||||
os.environ[
|
||||
"OTEL_RESOURCE_ATTRIBUTES"
|
||||
] = f"{existing_attrs},openinference.project.name={arize_phoenix_config.project_name}"
|
||||
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
|
||||
f"{existing_attrs},openinference.project.name={arize_phoenix_config.project_name}"
|
||||
)
|
||||
else:
|
||||
os.environ[
|
||||
"OTEL_RESOURCE_ATTRIBUTES"
|
||||
] = f"openinference.project.name={arize_phoenix_config.project_name}"
|
||||
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
|
||||
f"openinference.project.name={arize_phoenix_config.project_name}"
|
||||
)
|
||||
|
||||
# Set Phoenix project name from environment variable
|
||||
phoenix_project_name = os.environ.get("PHOENIX_PROJECT_NAME", None)
|
||||
|
|
@ -3781,19 +3777,19 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
existing_attrs = os.environ.get("OTEL_RESOURCE_ATTRIBUTES", "")
|
||||
# Add openinference.project.name attribute
|
||||
if existing_attrs:
|
||||
os.environ[
|
||||
"OTEL_RESOURCE_ATTRIBUTES"
|
||||
] = f"{existing_attrs},openinference.project.name={phoenix_project_name}"
|
||||
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
|
||||
f"{existing_attrs},openinference.project.name={phoenix_project_name}"
|
||||
)
|
||||
else:
|
||||
os.environ[
|
||||
"OTEL_RESOURCE_ATTRIBUTES"
|
||||
] = f"openinference.project.name={phoenix_project_name}"
|
||||
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
|
||||
f"openinference.project.name={phoenix_project_name}"
|
||||
)
|
||||
|
||||
# auth can be disabled on local deployments of arize phoenix
|
||||
if arize_phoenix_config.otlp_auth_headers is not None:
|
||||
os.environ[
|
||||
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
|
||||
] = arize_phoenix_config.otlp_auth_headers
|
||||
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
|
||||
arize_phoenix_config.otlp_auth_headers
|
||||
)
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if (
|
||||
|
|
@ -3969,9 +3965,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
exporter="otlp_http",
|
||||
endpoint="https://langtrace.ai/api/trace",
|
||||
)
|
||||
os.environ[
|
||||
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
|
||||
] = f"api_key={os.getenv('LANGTRACE_API_KEY')}"
|
||||
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
|
||||
f"api_key={os.getenv('LANGTRACE_API_KEY')}"
|
||||
)
|
||||
for callback in _in_memory_loggers:
|
||||
if (
|
||||
isinstance(callback, OpenTelemetry)
|
||||
|
|
@ -4204,8 +4200,7 @@ def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list) -> None:
|
|||
litellm.logging_callback_manager.add_litellm_callback(phoenix_logger)
|
||||
|
||||
verbose_logger.info(
|
||||
"Auto-initialized Arize Phoenix logger alongside otel "
|
||||
"(endpoint=%s)",
|
||||
"Auto-initialized Arize Phoenix logger alongside otel " "(endpoint=%s)",
|
||||
arize_phoenix_config.endpoint,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
@ -4768,9 +4763,11 @@ class StandardLoggingPayloadSetup:
|
|||
).model_dump()
|
||||
if isinstance(_raw, dict):
|
||||
if ResponseAPILoggingUtils._is_response_api_usage(_raw):
|
||||
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
_raw
|
||||
).model_dump()
|
||||
return (
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
_raw
|
||||
).model_dump()
|
||||
)
|
||||
return _raw
|
||||
if isinstance(_raw, Usage):
|
||||
return _raw.model_dump()
|
||||
|
|
@ -4884,10 +4881,10 @@ class StandardLoggingPayloadSetup:
|
|||
for key in StandardLoggingHiddenParams.__annotations__.keys():
|
||||
if key in hidden_params:
|
||||
if key == "additional_headers":
|
||||
clean_hidden_params[
|
||||
"additional_headers"
|
||||
] = StandardLoggingPayloadSetup.get_additional_headers(
|
||||
hidden_params[key]
|
||||
clean_hidden_params["additional_headers"] = (
|
||||
StandardLoggingPayloadSetup.get_additional_headers(
|
||||
hidden_params[key]
|
||||
)
|
||||
)
|
||||
else:
|
||||
clean_hidden_params[key] = hidden_params[key] # type: ignore
|
||||
|
|
@ -5039,14 +5036,22 @@ class StandardLoggingPayloadSetup:
|
|||
dynamic_litellm_session_id = litellm_params.get("litellm_session_id")
|
||||
dynamic_litellm_trace_id = litellm_params.get("litellm_trace_id")
|
||||
|
||||
|
||||
# Note: we recommend using `litellm_session_id` for session tracking
|
||||
# `litellm_trace_id` is an internal litellm param
|
||||
if dynamic_litellm_session_id:
|
||||
return str(dynamic_litellm_session_id)
|
||||
elif dynamic_litellm_trace_id:
|
||||
return str(dynamic_litellm_trace_id)
|
||||
else:
|
||||
return logging_obj.litellm_trace_id
|
||||
# Fallback: use metadata.session_id or metadata.trace_id for call chaining
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
metadata_session_id = metadata.get("session_id")
|
||||
metadata_trace_id = metadata.get("trace_id")
|
||||
if metadata_session_id:
|
||||
return str(metadata_session_id)
|
||||
if metadata_trace_id:
|
||||
return str(metadata_trace_id)
|
||||
return logging_obj.litellm_trace_id
|
||||
|
||||
@staticmethod
|
||||
def _get_user_agent_tags(proxy_server_request: dict) -> Optional[List[str]]:
|
||||
|
|
@ -5502,9 +5507,9 @@ def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]):
|
|||
):
|
||||
for k, v in metadata["user_api_key_metadata"].items():
|
||||
if k == "logging": # prevent logging user logging keys
|
||||
cleaned_user_api_key_metadata[
|
||||
k
|
||||
] = "scrubbed_by_litellm_for_sensitive_keys"
|
||||
cleaned_user_api_key_metadata[k] = (
|
||||
"scrubbed_by_litellm_for_sensitive_keys"
|
||||
)
|
||||
else:
|
||||
cleaned_user_api_key_metadata[k] = v
|
||||
|
||||
|
|
@ -5616,4 +5621,3 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
|
|||
model_parameters={"stream": True},
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -75,7 +75,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if messages is None:
|
||||
return data
|
||||
|
||||
chat_completion_compatible_request, tool_name_mapping = (
|
||||
chat_completion_compatible_request, _tool_name_mapping = (
|
||||
LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
# Use a shallow copy to avoid mutating request data (pop on litellm_metadata).
|
||||
anthropic_message_request=cast(AnthropicMessagesRequest, data.copy())
|
||||
|
|
@ -141,6 +141,14 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""Extract tool names from Anthropic messages request (tools[].name)."""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if isinstance(tool, dict) and tool.get("name"):
|
||||
names.append(str(tool["name"]))
|
||||
return names
|
||||
|
||||
def _extract_input_text_and_images(
|
||||
self,
|
||||
message: Dict[str, Any],
|
||||
|
|
|
|||
|
|
@ -41,7 +41,6 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
type="text",
|
||||
text="",
|
||||
)
|
||||
pending_new_content_block: bool = False
|
||||
chunk_queue: deque = deque() # Queue for buffering multiple chunks
|
||||
|
||||
def __init__(
|
||||
|
|
@ -80,38 +79,40 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
from .transformation import LiteLLMAnthropicMessagesAdapter
|
||||
|
||||
try:
|
||||
# Always return queued chunks first
|
||||
if self.chunk_queue:
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
# Queue initial chunks if not sent yet
|
||||
if self.sent_first_chunk is False:
|
||||
self.sent_first_chunk = True
|
||||
return {
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_{}".format(uuid.uuid4()),
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"model": self.model,
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": self._create_initial_usage_delta(),
|
||||
},
|
||||
}
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_{}".format(uuid.uuid4()),
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"model": self.model,
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": self._create_initial_usage_delta(),
|
||||
},
|
||||
}
|
||||
)
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
if self.sent_content_block_start is False:
|
||||
self.sent_content_block_start = True
|
||||
return {
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
|
||||
# Handle pending new content block start
|
||||
if self.pending_new_content_block:
|
||||
self.pending_new_content_block = False
|
||||
self.sent_content_block_finish = False # Reset for new block
|
||||
return {
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": self.current_content_block_start,
|
||||
}
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
)
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
|
|
@ -126,45 +127,65 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
current_content_block_index=self.current_content_block_index,
|
||||
)
|
||||
|
||||
# Check if we need to start a new content block
|
||||
# This is where you'd add your logic to detect when a new content block should start
|
||||
# For example, if the chunk indicates a tool call or different content type
|
||||
|
||||
if should_start_new_block and not self.sent_content_block_finish:
|
||||
# End current content block and prepare for new one
|
||||
self.holding_chunk = processed_chunk
|
||||
self.sent_content_block_finish = True
|
||||
self.pending_new_content_block = True
|
||||
return {
|
||||
"type": "content_block_stop",
|
||||
"index": max(self.current_content_block_index - 1, 0),
|
||||
}
|
||||
# Queue the sequence: content_block_stop -> content_block_start
|
||||
# The trigger chunk itself is not emitted as a delta since the
|
||||
# content_block_start already carries the relevant information.
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_stop",
|
||||
"index": max(self.current_content_block_index - 1, 0),
|
||||
}
|
||||
)
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": self.current_content_block_index,
|
||||
"content_block": self.current_content_block_start,
|
||||
}
|
||||
)
|
||||
self.sent_content_block_finish = False
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
if (
|
||||
processed_chunk["type"] == "message_delta"
|
||||
and self.sent_content_block_finish is False
|
||||
):
|
||||
self.holding_chunk = processed_chunk
|
||||
# Queue both the content_block_stop and the message_delta
|
||||
self.chunk_queue.append(
|
||||
{
|
||||
"type": "content_block_stop",
|
||||
"index": self.current_content_block_index,
|
||||
}
|
||||
)
|
||||
self.sent_content_block_finish = True
|
||||
return {
|
||||
"type": "content_block_stop",
|
||||
"index": self.current_content_block_index,
|
||||
}
|
||||
self.chunk_queue.append(processed_chunk)
|
||||
return self.chunk_queue.popleft()
|
||||
elif self.holding_chunk is not None:
|
||||
return_chunk = self.holding_chunk
|
||||
self.holding_chunk = processed_chunk
|
||||
return return_chunk
|
||||
self.chunk_queue.append(self.holding_chunk)
|
||||
self.chunk_queue.append(processed_chunk)
|
||||
self.holding_chunk = None
|
||||
return self.chunk_queue.popleft()
|
||||
else:
|
||||
return processed_chunk
|
||||
self.chunk_queue.append(processed_chunk)
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
# Handle any remaining held chunks after stream ends
|
||||
if self.holding_chunk is not None:
|
||||
return_chunk = self.holding_chunk
|
||||
self.chunk_queue.append(self.holding_chunk)
|
||||
self.holding_chunk = None
|
||||
return return_chunk
|
||||
if self.sent_last_message is False:
|
||||
|
||||
if not self.sent_last_message:
|
||||
self.sent_last_message = True
|
||||
return {"type": "message_stop"}
|
||||
self.chunk_queue.append({"type": "message_stop"})
|
||||
|
||||
if self.chunk_queue:
|
||||
return self.chunk_queue.popleft()
|
||||
|
||||
raise StopIteration
|
||||
except StopIteration:
|
||||
if self.chunk_queue:
|
||||
return self.chunk_queue.popleft()
|
||||
if self.sent_last_message is False:
|
||||
self.sent_last_message = True
|
||||
return {"type": "message_stop"}
|
||||
|
|
@ -265,7 +286,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
|
||||
if not self.queued_usage_chunk:
|
||||
if should_start_new_block and not self.sent_content_block_finish:
|
||||
# Queue the sequence: content_block_stop -> content_block_start -> current_chunk
|
||||
# Queue the sequence: content_block_stop -> content_block_start
|
||||
# The trigger chunk itself is not emitted as a delta since the
|
||||
# content_block_start already carries the relevant information.
|
||||
|
||||
# 1. Stop current content block
|
||||
self.chunk_queue.append(
|
||||
|
|
@ -284,9 +307,6 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
}
|
||||
)
|
||||
|
||||
# 3. Queue the current chunk (don't lose it!)
|
||||
self.chunk_queue.append(processed_chunk)
|
||||
|
||||
# Reset state for new block
|
||||
self.sent_content_block_finish = False
|
||||
|
||||
|
|
|
|||
|
|
@ -43,8 +43,12 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
|
|||
if "tool_choice" not in params:
|
||||
params.append("tool_choice")
|
||||
|
||||
# Only gpt-5.2 has been verified to support logprobs on Azure
|
||||
if self.is_model_gpt_5_2_model(model):
|
||||
# Only gpt-5.2 has been verified to support logprobs on Azure.
|
||||
# The base OpenAI class includes logprobs for gpt-5.1+, but Azure
|
||||
# hasn't verified support for gpt-5.1, so remove them unless gpt-5.2.
|
||||
if self.is_model_gpt_5_1_model(model) and not self.is_model_gpt_5_2_model(model):
|
||||
params = [p for p in params if p not in ["logprobs", "top_logprobs"]]
|
||||
elif self.is_model_gpt_5_2_model(model):
|
||||
azure_supported_params = ["logprobs", "top_logprobs"]
|
||||
params.extend(azure_supported_params)
|
||||
|
||||
|
|
|
|||
|
|
@ -98,3 +98,10 @@ class BaseTranslation(ABC):
|
|||
Optional to override in subclasses.
|
||||
"""
|
||||
return responses_so_far
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""
|
||||
Extract tool names from the request body for allowlist/policy checks.
|
||||
Override in tool-capable handlers; default returns [].
|
||||
"""
|
||||
return []
|
||||
|
|
|
|||
|
|
@ -166,7 +166,8 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
|||
contents: Optional[List[Dict[str, Any]]],
|
||||
deployment: Optional[Dict[str, Any]] = None,
|
||||
request_model: str = "",
|
||||
**kwargs,
|
||||
tools: Optional[List[Dict[str, Any]]] = None,
|
||||
system: Optional[Any] = None,
|
||||
) -> Optional[TokenCountResponse]:
|
||||
import copy
|
||||
|
||||
|
|
|
|||
|
|
@ -135,6 +135,19 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""Extract tool names from OpenAI chat completions request (tools[].function.name, functions[].name)."""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if isinstance(tool, dict) and tool.get("type") == "function":
|
||||
fn = tool.get("function")
|
||||
if isinstance(fn, dict) and fn.get("name"):
|
||||
names.append(str(fn["name"]))
|
||||
for fn in data.get("functions") or []:
|
||||
if isinstance(fn, dict) and fn.get("name"):
|
||||
names.append(str(fn["name"]))
|
||||
return names
|
||||
|
||||
def _extract_inputs(
|
||||
self,
|
||||
message: Dict[str, Any],
|
||||
|
|
|
|||
|
|
@ -30,27 +30,22 @@ Output: response.output is List[GenericResponseOutputItem] where each has:
|
|||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
|
||||
from openai.types.responses.response_function_tool_call import \
|
||||
ResponseFunctionToolCall
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
OpenAiResponsesToChatCompletionStreamIterator,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
OutputFunctionToolCall,
|
||||
OutputText,
|
||||
)
|
||||
OpenAiResponsesToChatCompletionStreamIterator)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import \
|
||||
BaseTranslation
|
||||
from litellm.responses.litellm_completion_transformation.transformation import \
|
||||
LiteLLMCompletionResponsesConfig
|
||||
from litellm.types.llms.openai import (ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam)
|
||||
from litellm.types.responses.main import (GenericResponseOutputItem,
|
||||
OutputFunctionToolCall, OutputText)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -188,6 +183,18 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""Extract tool names from Responses API request (tools[].name for function, tools[].server_label for mcp)."""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
if tool.get("type") == "function" and tool.get("name"):
|
||||
names.append(str(tool["name"]))
|
||||
elif tool.get("type") == "mcp" and tool.get("server_label"):
|
||||
names.append(str(tool["server_label"]))
|
||||
return names
|
||||
|
||||
def _extract_and_transform_tools(
|
||||
self,
|
||||
tools: List[Dict[str, Any]],
|
||||
|
|
|
|||
|
|
@ -115,9 +115,10 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
data=json.dumps(vertex_batch_request),
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_body = e.response.text if hasattr(e, 'response') else "N/A"
|
||||
error_body = e.response.text
|
||||
litellm.verbose_logger.error(
|
||||
f"Vertex AI batch create failed: status={e.response.status_code}, body={error_body[:1000]}"
|
||||
"Vertex AI batch create failed: status=%s, body=%s",
|
||||
e.response.status_code, error_body[:1000],
|
||||
)
|
||||
raise
|
||||
if response.status_code != 200:
|
||||
|
|
|
|||
|
|
@ -1054,7 +1054,8 @@ class VertexAITokenCounter(BaseTokenCounter):
|
|||
contents: Optional[List[Dict[str, Any]]],
|
||||
deployment: Optional[Dict[str, Any]] = None,
|
||||
request_model: str = "",
|
||||
**kwargs,
|
||||
tools: Optional[List[Dict[str, Any]]] = None,
|
||||
system: Optional[Any] = None,
|
||||
) -> Optional[TokenCountResponse]:
|
||||
import copy
|
||||
|
||||
|
|
|
|||
|
|
@ -408,10 +408,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
file_id = "deleted"
|
||||
if hasattr(raw_response, "request") and raw_response.request:
|
||||
url = str(raw_response.request.url)
|
||||
if "/o/" in url:
|
||||
if "/b/" in url and "/o/" in url:
|
||||
import urllib.parse
|
||||
bucket_part = url.split("/b/")[-1].split("/o/")[0]
|
||||
encoded_name = url.split("/o/")[-1].split("?")[0]
|
||||
file_id = f"gs://{urllib.parse.unquote(encoded_name)}"
|
||||
file_id = f"gs://{bucket_part}/{urllib.parse.unquote(encoded_name)}"
|
||||
return FileDeleted(id=file_id, deleted=True, object="file")
|
||||
|
||||
def transform_list_files_request(
|
||||
|
|
|
|||
|
|
@ -1136,23 +1136,6 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
if VertexGeminiConfig._is_gemini_3_or_newer(model):
|
||||
if "temperature" not in optional_params:
|
||||
optional_params["temperature"] = 1.0
|
||||
# Only add thinkingLevel if model supports it (exclude image models)
|
||||
if "image" not in model.lower():
|
||||
thinking_config = optional_params.get("thinkingConfig", {})
|
||||
if (
|
||||
"thinkingLevel" not in thinking_config
|
||||
and "thinkingBudget" not in thinking_config
|
||||
):
|
||||
# For gemini-3-flash-preview, default to "minimal" to match Gemini 2.5 Flash behavior
|
||||
# For other Gemini 3 models, default to "low"
|
||||
is_gemini3flash = (
|
||||
"gemini-3-flash-preview" in model.lower()
|
||||
or "gemini-3-flash" in model.lower()
|
||||
)
|
||||
thinking_config["thinkingLevel"] = (
|
||||
"minimal" if is_gemini3flash else "low"
|
||||
)
|
||||
optional_params["thinkingConfig"] = thinking_config
|
||||
|
||||
return optional_params
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -642,6 +642,7 @@ class MCPServerManager:
|
|||
available_on_public_internet=bool(
|
||||
getattr(mcp_server, "available_on_public_internet", True)
|
||||
),
|
||||
created_at=getattr(mcp_server, "created_at", None),
|
||||
updated_at=getattr(mcp_server, "updated_at", None),
|
||||
)
|
||||
return new_server
|
||||
|
|
@ -2540,8 +2541,8 @@ class MCPServerManager:
|
|||
url=server.url,
|
||||
transport=server.transport,
|
||||
auth_type=server.auth_type,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
created_at=server.created_at,
|
||||
updated_at=server.updated_at,
|
||||
teams=[],
|
||||
mcp_access_groups=server.access_groups or [],
|
||||
allowed_tools=server.allowed_tools or [],
|
||||
|
|
@ -2620,8 +2621,6 @@ class MCPServerManager:
|
|||
return list_mcp_servers
|
||||
|
||||
def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable:
|
||||
from datetime import datetime
|
||||
|
||||
return LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.server_name,
|
||||
|
|
@ -2633,8 +2632,8 @@ class MCPServerManager:
|
|||
spec_path=server.spec_path,
|
||||
transport=server.transport,
|
||||
auth_type=server.auth_type,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
created_at=server.created_at,
|
||||
updated_at=server.updated_at,
|
||||
teams=[],
|
||||
mcp_access_groups=server.access_groups or [],
|
||||
allowed_tools=server.allowed_tools or [],
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ LiteLLM MCP Server Routes
|
|||
|
||||
import asyncio
|
||||
import contextlib
|
||||
|
||||
import traceback
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
|
@ -44,7 +43,10 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
|||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
LiteLLMProxyRequestSetup,
|
||||
get_chain_id_from_headers,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
|
||||
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
|
||||
|
|
@ -331,6 +333,11 @@ if MCP_AVAILABLE:
|
|||
try:
|
||||
# Create a body date for logging
|
||||
body_data = {"name": name, "arguments": arguments}
|
||||
# Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A)
|
||||
chain_id = get_chain_id_from_headers(raw_headers)
|
||||
if chain_id:
|
||||
body_data["litellm_trace_id"] = chain_id
|
||||
body_data["litellm_session_id"] = chain_id
|
||||
|
||||
request = Request(
|
||||
scope={
|
||||
|
|
@ -884,6 +891,10 @@ if MCP_AVAILABLE:
|
|||
# This is intentionally minimal: only async_success_handler / post_call_failure_hook
|
||||
rules_obj = Rules()
|
||||
list_tools_call_id = str(uuid.uuid4())
|
||||
# Derive trace_id from raw_headers when not explicitly passed (same as A2A / MCP call_tool)
|
||||
effective_litellm_trace_id = litellm_trace_id or get_chain_id_from_headers(
|
||||
raw_headers
|
||||
)
|
||||
spend_logs_metadata: Dict[str, Any] = {
|
||||
"mcp_operation": "list_tools",
|
||||
}
|
||||
|
|
@ -896,7 +907,7 @@ if MCP_AVAILABLE:
|
|||
"model": "MCP: list_tools",
|
||||
"call_type": CallTypes.list_mcp_tools.value,
|
||||
"litellm_call_id": list_tools_call_id,
|
||||
"litellm_trace_id": litellm_trace_id,
|
||||
"litellm_trace_id": effective_litellm_trace_id,
|
||||
"metadata": {
|
||||
"spend_logs_metadata": spend_logs_metadata,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -23,33 +23,11 @@ model_list:
|
|||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "airline-competitor-intent"
|
||||
guardrail_id: "airline-competitor-intent"
|
||||
- guardrail_name: "tool_policy"
|
||||
litellm_params:
|
||||
guardrail: litellm_content_filter
|
||||
mode: pre_call
|
||||
default_on: false
|
||||
competitor_intent_config:
|
||||
brand_self:
|
||||
- emirates
|
||||
- ek
|
||||
competitors:
|
||||
- qatar airways
|
||||
- qatar
|
||||
- etihad
|
||||
locations:
|
||||
- qatar
|
||||
- doha
|
||||
- doh
|
||||
competitor_aliases:
|
||||
qatar airways: [qr, doha airline]
|
||||
qatar: [qr]
|
||||
policy:
|
||||
competitor_comparison: refuse
|
||||
possible_competitor_comparison: reframe
|
||||
threshold_high: 0.70
|
||||
threshold_medium: 0.45
|
||||
threshold_low: 0.30
|
||||
guardrail: tool_policy
|
||||
mode: [pre_call, post_call]
|
||||
default_on: true
|
||||
|
||||
mcp_servers:
|
||||
my_http_server:
|
||||
|
|
|
|||
|
|
@ -77,6 +77,7 @@ class SupportedDBObjectType(str, enum.Enum):
|
|||
PASS_THROUGH_ENDPOINTS = "pass_through_endpoints"
|
||||
PROMPTS = "prompts"
|
||||
MODEL_COST_MAP = "model_cost_map"
|
||||
TOOLS = "tools"
|
||||
|
||||
def __str__(self):
|
||||
return str(self.value)
|
||||
|
|
@ -512,6 +513,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
KeyManagementRoutes.KEY_UNBLOCK.value,
|
||||
KeyManagementRoutes.KEY_BULK_UPDATE.value,
|
||||
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
|
||||
KeyManagementRoutes.KEY_RESET_SPEND.value,
|
||||
]
|
||||
|
||||
management_routes = [
|
||||
|
|
@ -1551,6 +1553,8 @@ class NewTeamRequest(TeamBase):
|
|||
] = None # allow user to set TPM limit for all team members
|
||||
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
|
||||
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
|
||||
enforced_batch_output_expires_after: Optional[dict] = None
|
||||
enforced_file_expires_after: Optional[dict] = None
|
||||
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
|
|
@ -1606,6 +1610,8 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
model_rpm_limit: Optional[Dict[str, int]] = None
|
||||
model_tpm_limit: Optional[Dict[str, int]] = None
|
||||
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
|
||||
enforced_batch_output_expires_after: Optional[dict] = None
|
||||
enforced_file_expires_after: Optional[dict] = None
|
||||
router_settings: Optional[dict] = None
|
||||
access_group_ids: Optional[List[str]] = None
|
||||
|
||||
|
|
@ -2128,7 +2134,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
user_header_mappings: Optional[List[UserHeaderMapping]] = None
|
||||
supported_db_objects: Optional[List[SupportedDBObjectType]] = Field(
|
||||
None,
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).",
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map', 'tools'. If not set, all objects are loaded (default behavior).",
|
||||
)
|
||||
user_mcp_management_mode: Optional[UserMCPManagementMode] = Field(
|
||||
None,
|
||||
|
|
@ -3372,6 +3378,11 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
Team member is already in team
|
||||
"""
|
||||
|
||||
tool_access_denied = "tool_access_denied"
|
||||
"""
|
||||
Tool is not in the allowed tools list for this key/team
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def get_model_access_error_type_for_object(
|
||||
cls, object_type: Literal["key", "user", "team", "org", "project"]
|
||||
|
|
@ -3783,6 +3794,8 @@ LiteLLM_ManagementEndpoint_MetadataFields = [
|
|||
"temp_budget_increase",
|
||||
"temp_budget_expiry",
|
||||
"allowed_vector_store_indexes",
|
||||
"enforced_batch_output_expires_after",
|
||||
"enforced_file_expires_after",
|
||||
]
|
||||
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium = [
|
||||
|
|
@ -4154,6 +4167,7 @@ class ToolDiscoveryQueueItem(TypedDict, total=False):
|
|||
key_hash: Optional[str] # hash of virtual key that triggered discovery
|
||||
team_id: Optional[str] # team that triggered discovery
|
||||
key_alias: Optional[str] # human-readable key alias
|
||||
user_agent: Optional[str] # HTTP User-Agent of the caller
|
||||
|
||||
|
||||
class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ async def _handle_stream_message(
|
|||
from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE
|
||||
|
||||
if not A2A_SDK_AVAILABLE:
|
||||
|
||||
async def _error_stream():
|
||||
yield json.dumps(
|
||||
{
|
||||
|
|
@ -106,7 +107,12 @@ async def _handle_stream_message(
|
|||
proxy_server_request=proxy_server_request,
|
||||
)
|
||||
|
||||
if use_proxy_hooks and user_api_key_dict is not None and request_data is not None and proxy_logging_obj is not None:
|
||||
if (
|
||||
use_proxy_hooks
|
||||
and user_api_key_dict is not None
|
||||
and request_data is not None
|
||||
and proxy_logging_obj is not None
|
||||
):
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
)
|
||||
|
|
@ -119,20 +125,27 @@ async def _handle_stream_message(
|
|||
return json.dumps(obj) + "\n"
|
||||
|
||||
def _ndjson_error(proxy_exc: Any) -> str:
|
||||
return json.dumps(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"error": {
|
||||
"code": -32603,
|
||||
"message": getattr(
|
||||
proxy_exc, "message", f"Streaming error: {proxy_exc!s}"
|
||||
),
|
||||
},
|
||||
}
|
||||
) + "\n"
|
||||
return (
|
||||
json.dumps(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"error": {
|
||||
"code": -32603,
|
||||
"message": getattr(
|
||||
proxy_exc,
|
||||
"message",
|
||||
f"Streaming error: {proxy_exc!s}",
|
||||
),
|
||||
},
|
||||
}
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
async for line in ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
|
||||
async for (
|
||||
line
|
||||
) in ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
|
||||
response=a2a_stream,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
|
|
@ -151,7 +164,12 @@ async def _handle_stream_message(
|
|||
yield json.dumps(chunk) + "\n"
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error streaming A2A response: {e}")
|
||||
if use_proxy_hooks and proxy_logging_obj is not None and user_api_key_dict is not None and request_data is not None:
|
||||
if (
|
||||
use_proxy_hooks
|
||||
and proxy_logging_obj is not None
|
||||
and user_api_key_dict is not None
|
||||
and request_data is not None
|
||||
):
|
||||
transformed_exception = await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=e,
|
||||
|
|
@ -382,6 +400,7 @@ async def invoke_agent_a2a(
|
|||
agent_id=agent.agent_id,
|
||||
metadata=data.get("metadata", {}),
|
||||
proxy_server_request=data.get("proxy_server_request"),
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
|
|
|
|||
|
|
@ -58,6 +58,10 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.guardrails.tool_name_extraction import (
|
||||
TOOL_CAPABLE_CALL_TYPES,
|
||||
extract_request_tool_names,
|
||||
)
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
|
||||
from litellm.router import Router
|
||||
|
|
@ -220,7 +224,48 @@ async def _run_project_checks(
|
|||
)
|
||||
|
||||
|
||||
async def common_checks(
|
||||
async def check_tools_allowlist(
|
||||
request_body: dict,
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
team_object: Optional[LiteLLM_TeamTable],
|
||||
route: str,
|
||||
) -> None:
|
||||
"""
|
||||
Enforce key/team tool allowlist (metadata.allowed_tools). No DB in hot path —
|
||||
effective allowlist is read from valid_token.metadata and valid_token.team_metadata.
|
||||
Raises ProxyException with tool_access_denied if a tool is not allowed.
|
||||
"""
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import (
|
||||
get_call_types_for_route,
|
||||
)
|
||||
|
||||
if valid_token is None:
|
||||
return
|
||||
call_types = get_call_types_for_route(route)
|
||||
if not call_types or not any(ct.value in TOOL_CAPABLE_CALL_TYPES for ct in call_types):
|
||||
return
|
||||
tool_names = extract_request_tool_names(route, request_body)
|
||||
if not tool_names:
|
||||
return
|
||||
key_meta = (valid_token.metadata or {}) if isinstance(valid_token.metadata, dict) else {}
|
||||
team_meta = (valid_token.team_metadata or {}) if isinstance(valid_token.team_metadata, dict) else {}
|
||||
key_allowed = key_meta.get("allowed_tools")
|
||||
team_allowed = team_meta.get("allowed_tools")
|
||||
effective = key_allowed if (isinstance(key_allowed, list) and len(key_allowed) > 0) else team_allowed
|
||||
if not isinstance(effective, list) or len(effective) == 0:
|
||||
return
|
||||
allowed_set = {str(t) for t in effective}
|
||||
disallowed = [n for n in tool_names if n not in allowed_set]
|
||||
if disallowed:
|
||||
raise ProxyException(
|
||||
message=f"Tool(s) {disallowed} are not in the allowed tools list for this key/team.",
|
||||
type=ProxyErrorTypes.tool_access_denied,
|
||||
param="tools",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
|
||||
async def common_checks( # noqa: PLR0915
|
||||
request_body: dict,
|
||||
team_object: Optional[LiteLLM_TeamTable],
|
||||
user_object: Optional[LiteLLM_UserTable],
|
||||
|
|
@ -473,6 +518,14 @@ async def common_checks(
|
|||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
# 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path)
|
||||
await check_tools_allowlist(
|
||||
request_body=request_body,
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
route=route,
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1752,20 +1752,21 @@ async def _run_post_custom_auth_checks(
|
|||
if _project_obj is not None:
|
||||
valid_token.project_metadata = _project_obj.metadata
|
||||
|
||||
_ = await common_checks(
|
||||
request=request,
|
||||
request_body=request_data,
|
||||
team_object=_team_obj,
|
||||
user_object=user_object,
|
||||
end_user_object=end_user_object,
|
||||
general_settings=general_settings,
|
||||
global_proxy_spend=None,
|
||||
route=route,
|
||||
llm_router=llm_router,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
skip_budget_checks=False,
|
||||
project_object=_project_obj,
|
||||
)
|
||||
if general_settings.get("custom_auth_run_common_checks", False):
|
||||
_ = await common_checks(
|
||||
request=request,
|
||||
request_body=request_data,
|
||||
team_object=_team_obj,
|
||||
user_object=user_object,
|
||||
end_user_object=end_user_object,
|
||||
general_settings=general_settings,
|
||||
global_proxy_spend=None,
|
||||
route=route,
|
||||
llm_router=llm_router,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
skip_budget_checks=False,
|
||||
project_object=_project_obj,
|
||||
)
|
||||
|
||||
return valid_token
|
||||
|
|
|
|||
|
|
@ -119,6 +119,22 @@ async def create_batch( # noqa: PLR0915
|
|||
or "openai"
|
||||
)
|
||||
_create_batch_data = LiteLLMBatchCreateRequest(**data)
|
||||
|
||||
# Apply team-level batch output expiry enforcement
|
||||
team_metadata = user_api_key_dict.team_metadata or {}
|
||||
enforced_batch_expiry = team_metadata.get(
|
||||
"enforced_batch_output_expires_after"
|
||||
)
|
||||
if enforced_batch_expiry is not None:
|
||||
if "anchor" not in enforced_batch_expiry or "seconds" not in enforced_batch_expiry:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "enforced_batch_output_expires_after must contain 'anchor' and 'seconds' keys",
|
||||
},
|
||||
)
|
||||
_create_batch_data["output_expires_after"] = enforced_batch_expiry
|
||||
|
||||
input_file_id = _create_batch_data.get("input_file_id", None)
|
||||
unified_file_id: Union[str, Literal[False]] = False
|
||||
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ from litellm.constants import (
|
|||
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG,
|
||||
STREAM_SSE_DATA_PREFIX,
|
||||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import set_active_span_tag, tracer
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.llm_response_utils.get_headers import (
|
||||
get_response_headers,
|
||||
|
|
@ -41,6 +41,7 @@ from litellm.proxy.common_utils.callback_utils import (
|
|||
get_logging_caching_headers,
|
||||
get_remaining_tokens_and_requests_from_request_data,
|
||||
)
|
||||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.router import Router
|
||||
|
|
@ -245,26 +246,6 @@ async def create_response(
|
|||
)
|
||||
|
||||
|
||||
def _add_dd_apm_tags_for_litellm_call_id(litellm_call_id: Optional[str]) -> None:
|
||||
"""
|
||||
Attach LiteLLM call id to the active Datadog APM span.
|
||||
|
||||
This enables searching APM traces by LiteLLM call id returned in
|
||||
`x-litellm-call-id`.
|
||||
"""
|
||||
if not litellm_call_id:
|
||||
return
|
||||
|
||||
try:
|
||||
set_active_span_tag("litellm.call_id", str(litellm_call_id))
|
||||
except Exception:
|
||||
# Tagging is best-effort and should never impact request processing.
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to tag active ddtrace span with litellm.call_id",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
def _override_openai_response_model(
|
||||
*,
|
||||
response_obj: Any,
|
||||
|
|
@ -662,7 +643,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
self.data["litellm_call_id"] = request.headers.get(
|
||||
"x-litellm-call-id", str(uuid.uuid4())
|
||||
)
|
||||
_add_dd_apm_tags_for_litellm_call_id(self.data.get("litellm_call_id"))
|
||||
DDSpanTagger.tag_call_id(self.data.get("litellm_call_id"))
|
||||
DDSpanTagger.tag_request(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
requested_model=self.data.get("model"),
|
||||
)
|
||||
|
||||
### AUTO STREAM USAGE TRACKING ###
|
||||
# If always_include_stream_usage is enabled and this is a streaming request
|
||||
|
|
|
|||
|
|
@ -13,49 +13,36 @@ import random
|
|||
import time
|
||||
import traceback
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Union,
|
||||
cast,
|
||||
overload,
|
||||
)
|
||||
from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union,
|
||||
cast, overload)
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import DualCache, RedisCache
|
||||
from litellm.constants import DB_SPEND_UPDATE_JOB_NAME
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy._types import (
|
||||
DB_CONNECTION_ERROR_TYPES,
|
||||
BaseDailySpendTransaction,
|
||||
DailyAgentSpendTransaction,
|
||||
DailyEndUserSpendTransaction,
|
||||
DailyOrganizationSpendTransaction,
|
||||
DailyTagSpendTransaction,
|
||||
DailyTeamSpendTransaction,
|
||||
DailyUserSpendTransaction,
|
||||
DBSpendUpdateTransactions,
|
||||
Litellm_EntityType,
|
||||
LiteLLM_UserTable,
|
||||
SpendLogsMetadata,
|
||||
SpendLogsPayload,
|
||||
SpendUpdateQueueItem,
|
||||
ToolDiscoveryQueueItem,
|
||||
)
|
||||
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
|
||||
DailySpendUpdateQueue,
|
||||
)
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
|
||||
from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue
|
||||
from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import (
|
||||
ToolDiscoveryQueue,
|
||||
)
|
||||
from litellm.proxy._types import (DB_CONNECTION_ERROR_TYPES,
|
||||
BaseDailySpendTransaction,
|
||||
DailyAgentSpendTransaction,
|
||||
DailyEndUserSpendTransaction,
|
||||
DailyOrganizationSpendTransaction,
|
||||
DailyTagSpendTransaction,
|
||||
DailyTeamSpendTransaction,
|
||||
DailyUserSpendTransaction,
|
||||
DBSpendUpdateTransactions,
|
||||
Litellm_EntityType, LiteLLM_UserTable,
|
||||
SpendLogsMetadata, SpendLogsPayload,
|
||||
SpendUpdateQueueItem, ToolDiscoveryQueueItem)
|
||||
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import \
|
||||
DailySpendUpdateQueue
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import \
|
||||
PodLockManager
|
||||
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import \
|
||||
RedisUpdateBuffer
|
||||
from litellm.proxy.db.db_transaction_queue.spend_update_queue import \
|
||||
SpendUpdateQueue
|
||||
from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import \
|
||||
ToolDiscoveryQueue
|
||||
from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -104,12 +91,10 @@ class DBSpendUpdateWriter:
|
|||
end_time: Optional[datetime],
|
||||
response_cost: Optional[float],
|
||||
):
|
||||
from litellm.proxy.proxy_server import (
|
||||
disable_spend_logs,
|
||||
litellm_proxy_budget_name,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (disable_spend_logs,
|
||||
litellm_proxy_budget_name,
|
||||
prisma_client,
|
||||
user_api_key_cache)
|
||||
from litellm.proxy.utils import ProxyUpdateSpend, hash_token
|
||||
|
||||
try:
|
||||
|
|
@ -124,9 +109,8 @@ class DBSpendUpdateWriter:
|
|||
hashed_token = token
|
||||
|
||||
## CREATE SPEND LOG PAYLOAD ##
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import (
|
||||
get_logging_payload,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import \
|
||||
get_logging_payload
|
||||
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -230,6 +214,7 @@ class DBSpendUpdateWriter:
|
|||
_litellm_params = kwargs.get("litellm_params") or {}
|
||||
_metadata = _litellm_params.get("metadata") or {}
|
||||
key_alias = _metadata.get("user_api_key_alias") or None
|
||||
user_agent = _metadata.get("user_agent") or None
|
||||
|
||||
def _enqueue(tool_name: str, origin: str = "user_defined") -> None:
|
||||
self.tool_discovery_queue.add_update(
|
||||
|
|
@ -239,17 +224,20 @@ class DBSpendUpdateWriter:
|
|||
key_hash=hashed_token,
|
||||
team_id=team_id,
|
||||
key_alias=key_alias,
|
||||
user_agent=user_agent,
|
||||
)
|
||||
)
|
||||
|
||||
# --- MCP tool calls ---
|
||||
sl_object = kwargs.get("standard_logging_object")
|
||||
if sl_object is not None:
|
||||
mcp_metadata = (
|
||||
sl_object.get("metadata", {}) or {}
|
||||
).get("mcp_tool_call_metadata")
|
||||
mcp_metadata = (sl_object.get("metadata", {}) or {}).get(
|
||||
"mcp_tool_call_metadata"
|
||||
)
|
||||
if mcp_metadata and isinstance(mcp_metadata, dict):
|
||||
tool_name = mcp_metadata.get("namespaced_tool_name") or mcp_metadata.get("name")
|
||||
tool_name = mcp_metadata.get(
|
||||
"namespaced_tool_name"
|
||||
) or mcp_metadata.get("name")
|
||||
mcp_server_name = mcp_metadata.get("mcp_server_name")
|
||||
if tool_name:
|
||||
_enqueue(tool_name, origin=mcp_server_name or "user_defined")
|
||||
|
|
@ -280,7 +268,9 @@ class DBSpendUpdateWriter:
|
|||
_enqueue(name)
|
||||
|
||||
# --- Response tool_calls (OpenAI format; Anthropic pass-through converts tool_use here) ---
|
||||
if completion_response is not None and hasattr(completion_response, "choices"):
|
||||
if completion_response is not None and hasattr(
|
||||
completion_response, "choices"
|
||||
):
|
||||
for choice in completion_response.choices or []:
|
||||
message = getattr(choice, "message", None)
|
||||
if message is None:
|
||||
|
|
@ -768,19 +758,46 @@ class DBSpendUpdateWriter:
|
|||
daily_end_user_spend_update_transactions,
|
||||
daily_agent_spend_update_transactions,
|
||||
daily_tag_spend_update_transactions,
|
||||
) = await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
|
||||
) = (
|
||||
await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
|
||||
)
|
||||
|
||||
if db_spend_update_transactions is not None:
|
||||
verbose_proxy_logger.info(
|
||||
"Spend tracking - committing spend updates from Redis to DB: "
|
||||
"keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d",
|
||||
len(db_spend_update_transactions.get("key_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("user_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("team_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("org_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("end_user_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("team_member_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("tag_list_transactions") or {}),
|
||||
len(
|
||||
db_spend_update_transactions.get("key_list_transactions")
|
||||
or {}
|
||||
),
|
||||
len(
|
||||
db_spend_update_transactions.get("user_list_transactions")
|
||||
or {}
|
||||
),
|
||||
len(
|
||||
db_spend_update_transactions.get("team_list_transactions")
|
||||
or {}
|
||||
),
|
||||
len(
|
||||
db_spend_update_transactions.get("org_list_transactions")
|
||||
or {}
|
||||
),
|
||||
len(
|
||||
db_spend_update_transactions.get(
|
||||
"end_user_list_transactions"
|
||||
)
|
||||
or {}
|
||||
),
|
||||
len(
|
||||
db_spend_update_transactions.get(
|
||||
"team_member_list_transactions"
|
||||
)
|
||||
or {}
|
||||
),
|
||||
len(
|
||||
db_spend_update_transactions.get("tag_list_transactions")
|
||||
or {}
|
||||
),
|
||||
)
|
||||
await self._commit_spend_updates_to_db(
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -985,10 +1002,8 @@ class DBSpendUpdateWriter:
|
|||
Commits all the spend `UPDATE` transactions to the Database
|
||||
|
||||
"""
|
||||
from litellm.proxy.utils import (
|
||||
ProxyUpdateSpend,
|
||||
_raise_failed_update_spend_exception,
|
||||
)
|
||||
from litellm.proxy.utils import (ProxyUpdateSpend,
|
||||
_raise_failed_update_spend_exception)
|
||||
|
||||
### UPDATE USER TABLE ###
|
||||
user_list_transactions = db_spend_update_transactions["user_list_transactions"]
|
||||
|
|
@ -1523,14 +1538,14 @@ class DBSpendUpdateWriter:
|
|||
|
||||
# Add cache-related fields if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
common_data[
|
||||
"cache_read_input_tokens"
|
||||
] = transaction.get("cache_read_input_tokens", 0)
|
||||
common_data["cache_read_input_tokens"] = (
|
||||
transaction.get("cache_read_input_tokens", 0)
|
||||
)
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
common_data[
|
||||
"cache_creation_input_tokens"
|
||||
] = transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
common_data["cache_creation_input_tokens"] = (
|
||||
transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
)
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
|
|
|
|||
|
|
@ -49,14 +49,18 @@ class SpendLogCleanup:
|
|||
|
||||
try:
|
||||
if isinstance(retention_setting, int):
|
||||
retention_setting = str(retention_setting)
|
||||
verbose_proxy_logger.warning(
|
||||
f"maximum_spend_logs_retention_period is an integer ({retention_setting}); treating as days. "
|
||||
"Use a string like '3d' to be explicit."
|
||||
)
|
||||
retention_setting = f"{retention_setting}d"
|
||||
self.retention_seconds = duration_in_seconds(retention_setting)
|
||||
verbose_proxy_logger.info(
|
||||
f"Retention period set to {self.retention_seconds} seconds"
|
||||
)
|
||||
return True
|
||||
except ValueError as e:
|
||||
verbose_proxy_logger.error(
|
||||
verbose_proxy_logger.warning(
|
||||
f"Invalid maximum_spend_logs_retention_period value: {retention_setting}, error: {str(e)}"
|
||||
)
|
||||
return False
|
||||
|
|
@ -112,13 +116,11 @@ class SpendLogCleanup:
|
|||
If pod_lock_manager is available, ensures only one pod runs cleanup.
|
||||
If no pod_lock_manager, runs cleanup without distributed locking.
|
||||
"""
|
||||
lock_acquired = False
|
||||
try:
|
||||
verbose_proxy_logger.info(f"Cleanup job triggered at {datetime.now()}")
|
||||
|
||||
if not self._should_delete_spend_logs():
|
||||
verbose_proxy_logger.info(
|
||||
"Skipping cleanup — invalid or missing retention setting."
|
||||
)
|
||||
return
|
||||
|
||||
if self.retention_seconds is None:
|
||||
|
|
@ -155,8 +157,8 @@ class SpendLogCleanup:
|
|||
verbose_proxy_logger.error(f"Error during cleanup: {str(e)}")
|
||||
return # Return after error handling
|
||||
finally:
|
||||
# Always release the lock if we have a pod lock manager
|
||||
if self.pod_lock_manager and self.pod_lock_manager.redis_cache:
|
||||
# Only release the lock if it was actually acquired
|
||||
if lock_acquired and self.pod_lock_manager and self.pod_lock_manager.redis_cache:
|
||||
await self.pod_lock_manager.release_lock(
|
||||
cronjob_id=SPEND_LOG_CLEANUP_JOB_NAME
|
||||
)
|
||||
|
|
|
|||
147
litellm/proxy/db/spend_log_tool_index.py
Normal file
147
litellm/proxy/db/spend_log_tool_index.py
Normal file
|
|
@ -0,0 +1,147 @@
|
|||
"""
|
||||
Track tool usage for the dashboard: insert into SpendLogToolIndex when spend logs
|
||||
are written, so "last N requests for tool X" and "how is this tool called in production"
|
||||
queries are fast.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Set
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
def _add_tool_calls_to_set(tool_calls: Any, out: Set[str]) -> None:
|
||||
"""Extract tool names from OpenAI-style tool_calls list into out."""
|
||||
if not isinstance(tool_calls, list):
|
||||
return
|
||||
for tc in tool_calls:
|
||||
if not isinstance(tc, dict):
|
||||
continue
|
||||
fn = tc.get("function")
|
||||
if isinstance(fn, dict):
|
||||
name = fn.get("name")
|
||||
if name and isinstance(name, str) and name.strip():
|
||||
out.add(name.strip())
|
||||
|
||||
|
||||
def _parse_tool_names_from_payload(payload: Dict[str, Any]) -> Set[str]:
|
||||
"""
|
||||
Extract deduplicated tool names from a spend log payload.
|
||||
Sources: mcp_namespaced_tool_name, response (tool_calls), proxy_server_request (tools).
|
||||
"""
|
||||
tool_names: Set[str] = set()
|
||||
|
||||
# Top-level MCP tool name (single tool per request for that flow)
|
||||
mcp_name = payload.get("mcp_namespaced_tool_name")
|
||||
if mcp_name and isinstance(mcp_name, str) and mcp_name.strip():
|
||||
tool_names.add(mcp_name.strip())
|
||||
|
||||
# Response: OpenAI-style tool_calls[].function.name or choices[0].message.tool_calls
|
||||
response_raw = payload.get("response")
|
||||
if response_raw:
|
||||
response_obj = (
|
||||
safe_json_loads(response_raw, default=None)
|
||||
if isinstance(response_raw, str)
|
||||
else response_raw
|
||||
)
|
||||
if isinstance(response_obj, dict):
|
||||
_add_tool_calls_to_set(response_obj.get("tool_calls"), tool_names)
|
||||
choices = response_obj.get("choices")
|
||||
if isinstance(choices, list) and choices:
|
||||
msg = choices[0].get("message") if isinstance(choices[0], dict) else None
|
||||
if isinstance(msg, dict):
|
||||
_add_tool_calls_to_set(msg.get("tool_calls"), tool_names)
|
||||
|
||||
# Request body: tools[].function.name
|
||||
request_raw = payload.get("proxy_server_request")
|
||||
if request_raw:
|
||||
request_obj = (
|
||||
safe_json_loads(request_raw, default=None)
|
||||
if isinstance(request_raw, str)
|
||||
else request_raw
|
||||
)
|
||||
if isinstance(request_obj, dict):
|
||||
body = request_obj.get("body", request_obj)
|
||||
if isinstance(body, dict):
|
||||
request_obj = body
|
||||
if isinstance(request_obj, dict):
|
||||
tools = request_obj.get("tools")
|
||||
if isinstance(tools, list):
|
||||
for t in tools:
|
||||
if isinstance(t, dict):
|
||||
fn = t.get("function")
|
||||
if isinstance(fn, dict):
|
||||
name = fn.get("name")
|
||||
if name and isinstance(name, str) and name.strip():
|
||||
tool_names.add(name.strip())
|
||||
|
||||
return tool_names
|
||||
|
||||
|
||||
async def process_spend_logs_tool_usage(
|
||||
prisma_client: PrismaClient,
|
||||
logs_to_process: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
"""
|
||||
After spend logs are written: insert SpendLogToolIndex rows from each payload.
|
||||
Extracts tool names from mcp_namespaced_tool_name, response tool_calls, and
|
||||
proxy_server_request tools.
|
||||
"""
|
||||
if not logs_to_process:
|
||||
return
|
||||
|
||||
index_rows: List[Dict[str, Any]] = []
|
||||
|
||||
for payload in logs_to_process:
|
||||
request_id = payload.get("request_id")
|
||||
start_time = payload.get("startTime")
|
||||
if not request_id or not start_time:
|
||||
continue
|
||||
if isinstance(start_time, str):
|
||||
try:
|
||||
start_time = datetime.fromisoformat(
|
||||
start_time.replace("Z", "+00:00")
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
continue
|
||||
if start_time.tzinfo is None:
|
||||
start_time = start_time.replace(tzinfo=timezone.utc)
|
||||
|
||||
tool_names = _parse_tool_names_from_payload(payload)
|
||||
for tool_name in tool_names:
|
||||
index_rows.append({
|
||||
"request_id": request_id,
|
||||
"tool_name": tool_name,
|
||||
"start_time": start_time,
|
||||
})
|
||||
|
||||
if not index_rows:
|
||||
return
|
||||
|
||||
try:
|
||||
index_data = []
|
||||
for r in index_rows:
|
||||
st = r["start_time"]
|
||||
if isinstance(st, str):
|
||||
try:
|
||||
st = datetime.fromisoformat(st.replace("Z", "+00:00"))
|
||||
except (ValueError, TypeError):
|
||||
continue
|
||||
if st.tzinfo is None:
|
||||
st = st.replace(tzinfo=timezone.utc)
|
||||
index_data.append({
|
||||
"request_id": r["request_id"],
|
||||
"tool_name": r["tool_name"],
|
||||
"start_time": st,
|
||||
})
|
||||
if index_data:
|
||||
await prisma_client.db.litellm_spendlogtoolindex.create_many(
|
||||
data=index_data,
|
||||
skip_duplicates=True,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Tool usage tracking (SpendLogToolIndex) failed (non-fatal): %s", e
|
||||
)
|
||||
|
|
@ -2,36 +2,64 @@
|
|||
DB helpers for LiteLLM_ToolTable — the global tool registry.
|
||||
|
||||
Tools are auto-discovered from LLM responses and upserted here.
|
||||
Admins use the management endpoints to read and update call_policy.
|
||||
|
||||
NOTE: Uses raw SQL (query_raw / execute_raw) instead of Prisma model methods
|
||||
because the generated Prisma Python client may not have LiteLLM_ToolTable
|
||||
when running against an older generated schema.
|
||||
Admins use the management endpoints to read and update input_policy / output_policy.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ToolDiscoveryQueueItem
|
||||
from litellm.types.tool_management import LiteLLM_ToolTableRow, ToolCallPolicy
|
||||
from litellm.types.tool_management import (
|
||||
LiteLLM_ToolTableRow,
|
||||
ToolPolicyOverrideRow,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
def _row_to_model(row: dict) -> LiteLLM_ToolTableRow:
|
||||
def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow:
|
||||
"""Convert a Prisma model instance or dict to LiteLLM_ToolTableRow."""
|
||||
model_dump = getattr(row, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
row = model_dump()
|
||||
elif not isinstance(row, dict):
|
||||
row = {
|
||||
k: getattr(row, k, None)
|
||||
for k in (
|
||||
"tool_id",
|
||||
"tool_name",
|
||||
"origin",
|
||||
"input_policy",
|
||||
"output_policy",
|
||||
"call_count",
|
||||
"assignments",
|
||||
"key_hash",
|
||||
"team_id",
|
||||
"key_alias",
|
||||
"user_agent",
|
||||
"last_used_at",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"created_by",
|
||||
"updated_by",
|
||||
)
|
||||
}
|
||||
return LiteLLM_ToolTableRow(
|
||||
tool_id=row.get("tool_id", ""),
|
||||
tool_name=row.get("tool_name", ""),
|
||||
origin=row.get("origin"),
|
||||
call_policy=row.get("call_policy", "untrusted"),
|
||||
input_policy=row.get("input_policy") or "untrusted",
|
||||
output_policy=row.get("output_policy") or "untrusted",
|
||||
call_count=int(row.get("call_count") or 0),
|
||||
assignments=row.get("assignments"),
|
||||
key_hash=row.get("key_hash"),
|
||||
team_id=row.get("team_id"),
|
||||
key_alias=row.get("key_alias"),
|
||||
user_agent=row.get("user_agent"),
|
||||
last_used_at=row.get("last_used_at"),
|
||||
created_at=row.get("created_at"),
|
||||
updated_at=row.get("updated_at"),
|
||||
created_by=row.get("created_by"),
|
||||
|
|
@ -44,10 +72,10 @@ async def batch_upsert_tools(
|
|||
items: List[ToolDiscoveryQueueItem],
|
||||
) -> None:
|
||||
"""
|
||||
Batch-upsert tool registry rows via raw SQL.
|
||||
Batch-upsert tool registry rows via Prisma.
|
||||
|
||||
On first insert: sets call_policy = "untrusted" (schema default), call_count = 1.
|
||||
On conflict: increments call_count; preserves existing call_policy.
|
||||
On first insert: sets input_policy/output_policy = "untrusted" (default), call_count = 1.
|
||||
On conflict: increments call_count; preserves existing policies.
|
||||
"""
|
||||
if not items:
|
||||
return
|
||||
|
|
@ -55,6 +83,8 @@ async def batch_upsert_tools(
|
|||
data = [item for item in items if item.get("tool_name")]
|
||||
if not data:
|
||||
return
|
||||
now = datetime.now(timezone.utc)
|
||||
table = prisma_client.db.litellm_tooltable
|
||||
for item in data:
|
||||
tool_name = item.get("tool_name", "")
|
||||
origin = item.get("origin") or "user_defined"
|
||||
|
|
@ -62,49 +92,52 @@ async def batch_upsert_tools(
|
|||
key_hash = item.get("key_hash")
|
||||
team_id = item.get("team_id")
|
||||
key_alias = item.get("key_alias")
|
||||
now = datetime.now(timezone.utc).isoformat()
|
||||
await prisma_client.db.execute_raw(
|
||||
'INSERT INTO "LiteLLM_ToolTable" '
|
||||
"(tool_id, tool_name, origin, call_policy, call_count, created_by, updated_by, key_hash, team_id, key_alias, created_at, updated_at) "
|
||||
"VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8, $8) "
|
||||
"ON CONFLICT (tool_name) DO UPDATE SET "
|
||||
"call_count = \"LiteLLM_ToolTable\".call_count + 1, "
|
||||
"updated_at = $8",
|
||||
tool_name,
|
||||
origin,
|
||||
created_by,
|
||||
key_hash,
|
||||
team_id,
|
||||
key_alias,
|
||||
str(uuid.uuid4()),
|
||||
now,
|
||||
user_agent = item.get("user_agent")
|
||||
await table.upsert(
|
||||
where={"tool_name": tool_name},
|
||||
data={
|
||||
"create": {
|
||||
"tool_id": str(uuid.uuid4()),
|
||||
"tool_name": tool_name,
|
||||
"origin": origin,
|
||||
"input_policy": "untrusted",
|
||||
"output_policy": "untrusted",
|
||||
"call_count": 1,
|
||||
"created_by": created_by,
|
||||
"updated_by": created_by,
|
||||
"key_hash": key_hash,
|
||||
"team_id": team_id,
|
||||
"key_alias": key_alias,
|
||||
"user_agent": user_agent,
|
||||
"last_used_at": now,
|
||||
},
|
||||
"update": {
|
||||
"call_count": {"increment": 1},
|
||||
"updated_at": now,
|
||||
"last_used_at": now,
|
||||
},
|
||||
},
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"tool_registry_writer: upserted %d tool(s)", len(data)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e)
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer batch_upsert_tools error: %s", e
|
||||
)
|
||||
|
||||
|
||||
async def list_tools(
|
||||
prisma_client: "PrismaClient",
|
||||
call_policy: Optional[ToolCallPolicy] = None,
|
||||
input_policy: Optional[str] = None,
|
||||
) -> List[LiteLLM_ToolTableRow]:
|
||||
"""Return all tools, optionally filtered by call_policy."""
|
||||
"""Return all tools, optionally filtered by input_policy."""
|
||||
try:
|
||||
if call_policy is not None:
|
||||
rows = await prisma_client.db.query_raw(
|
||||
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
|
||||
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
|
||||
'FROM "LiteLLM_ToolTable" WHERE call_policy = $1 ORDER BY created_at DESC',
|
||||
call_policy,
|
||||
)
|
||||
else:
|
||||
rows = await prisma_client.db.query_raw(
|
||||
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
|
||||
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
|
||||
'FROM "LiteLLM_ToolTable" ORDER BY created_at DESC',
|
||||
)
|
||||
where = {"input_policy": input_policy} if input_policy is not None else {}
|
||||
rows = await prisma_client.db.litellm_tooltable.find_many(
|
||||
where=where,
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
return [_row_to_model(row) for row in rows]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e)
|
||||
|
|
@ -117,15 +150,12 @@ async def get_tool(
|
|||
) -> Optional[LiteLLM_ToolTableRow]:
|
||||
"""Return a single tool row by tool_name."""
|
||||
try:
|
||||
rows = await prisma_client.db.query_raw(
|
||||
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
|
||||
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
|
||||
'FROM "LiteLLM_ToolTable" WHERE tool_name = $1',
|
||||
tool_name,
|
||||
row = await prisma_client.db.litellm_tooltable.find_unique(
|
||||
where={"tool_name": tool_name},
|
||||
)
|
||||
if not rows:
|
||||
if row is None:
|
||||
return None
|
||||
return _row_to_model(rows[0])
|
||||
return _row_to_model(row)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e)
|
||||
return None
|
||||
|
|
@ -134,46 +164,279 @@ async def get_tool(
|
|||
async def update_tool_policy(
|
||||
prisma_client: "PrismaClient",
|
||||
tool_name: str,
|
||||
call_policy: ToolCallPolicy,
|
||||
updated_by: Optional[str],
|
||||
input_policy: Optional[str] = None,
|
||||
output_policy: Optional[str] = None,
|
||||
) -> Optional[LiteLLM_ToolTableRow]:
|
||||
"""Update the call_policy for a tool. Upserts the row if it does not exist yet."""
|
||||
"""Update input_policy and/or output_policy for a tool. Upserts the row if it does not exist yet."""
|
||||
try:
|
||||
_updated_by = updated_by or "system"
|
||||
now = datetime.now(timezone.utc).isoformat()
|
||||
await prisma_client.db.execute_raw(
|
||||
'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, call_policy, created_by, updated_by, created_at, updated_at) '
|
||||
"VALUES ($4, $1, $2, $3, $3, $5, $5) "
|
||||
"ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5",
|
||||
tool_name,
|
||||
call_policy,
|
||||
_updated_by,
|
||||
str(uuid.uuid4()),
|
||||
now,
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
create_data: dict = {
|
||||
"tool_id": str(uuid.uuid4()),
|
||||
"tool_name": tool_name,
|
||||
"input_policy": input_policy or "untrusted",
|
||||
"output_policy": output_policy or "untrusted",
|
||||
"created_by": _updated_by,
|
||||
"updated_by": _updated_by,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
update_data: dict = {
|
||||
"updated_by": _updated_by,
|
||||
"updated_at": now,
|
||||
}
|
||||
if input_policy is not None:
|
||||
update_data["input_policy"] = input_policy
|
||||
if output_policy is not None:
|
||||
update_data["output_policy"] = output_policy
|
||||
|
||||
await prisma_client.db.litellm_tooltable.upsert(
|
||||
where={"tool_name": tool_name},
|
||||
data={
|
||||
"create": create_data,
|
||||
"update": update_data,
|
||||
},
|
||||
)
|
||||
return await get_tool(prisma_client, tool_name)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("tool_registry_writer update_tool_policy error: %s", e)
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer update_tool_policy error: %s", e
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
async def get_tools_by_names(
|
||||
prisma_client: "PrismaClient",
|
||||
tool_names: List[str],
|
||||
) -> Dict[str, str]:
|
||||
) -> Dict[str, Tuple[str, str]]:
|
||||
"""
|
||||
Return a {tool_name: call_policy} map for the given tool names.
|
||||
Used by the policy enforcement guardrail — single batch query, never N+1.
|
||||
Return a {tool_name: (input_policy, output_policy)} map for the given tool names.
|
||||
"""
|
||||
if not tool_names:
|
||||
return {}
|
||||
try:
|
||||
placeholders = ", ".join(f"${i+1}" for i in range(len(tool_names)))
|
||||
rows = await prisma_client.db.query_raw(
|
||||
f'SELECT tool_name, call_policy FROM "LiteLLM_ToolTable" WHERE tool_name IN ({placeholders})',
|
||||
*tool_names,
|
||||
rows = await prisma_client.db.litellm_tooltable.find_many(
|
||||
where={"tool_name": {"in": tool_names}},
|
||||
)
|
||||
return {row["tool_name"]: row["call_policy"] for row in rows}
|
||||
return {
|
||||
row.tool_name: (
|
||||
getattr(row, "input_policy", "untrusted") or "untrusted",
|
||||
getattr(row, "output_policy", "untrusted") or "untrusted",
|
||||
)
|
||||
for row in rows
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("tool_registry_writer get_tools_by_names error: %s", e)
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer get_tools_by_names error: %s", e
|
||||
)
|
||||
return {}
|
||||
|
||||
|
||||
async def list_overrides_for_tool(
|
||||
prisma_client: "PrismaClient",
|
||||
tool_name: str,
|
||||
) -> List[ToolPolicyOverrideRow]:
|
||||
"""
|
||||
Return override-like rows for a tool by finding object permissions that have
|
||||
this tool in blocked_tools, then resolving each permission to key/team scope for display.
|
||||
"""
|
||||
out: List[ToolPolicyOverrideRow] = []
|
||||
try:
|
||||
perms = await prisma_client.db.litellm_objectpermissiontable.find_many(
|
||||
where={"blocked_tools": {"has": tool_name}},
|
||||
include={
|
||||
"verification_tokens": True,
|
||||
"teams": True,
|
||||
},
|
||||
)
|
||||
for perm in perms:
|
||||
op_id = getattr(perm, "object_permission_id", None) or ""
|
||||
tokens = getattr(perm, "verification_tokens", []) or []
|
||||
teams = getattr(perm, "teams", []) or []
|
||||
for t in tokens:
|
||||
out.append(
|
||||
ToolPolicyOverrideRow(
|
||||
override_id=op_id,
|
||||
tool_name=tool_name,
|
||||
team_id=None,
|
||||
key_hash=getattr(t, "token", None),
|
||||
input_policy="blocked",
|
||||
key_alias=getattr(t, "key_alias", None),
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
)
|
||||
)
|
||||
for team in teams:
|
||||
out.append(
|
||||
ToolPolicyOverrideRow(
|
||||
override_id=op_id,
|
||||
tool_name=tool_name,
|
||||
team_id=getattr(team, "team_id", None),
|
||||
key_hash=None,
|
||||
input_policy="blocked",
|
||||
key_alias=getattr(team, "team_alias", None),
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
)
|
||||
)
|
||||
return out
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer list_overrides_for_tool error: %s", e
|
||||
)
|
||||
return []
|
||||
|
||||
|
||||
class ToolPolicyRegistry:
|
||||
"""
|
||||
In-memory registry of tool policies synced from DB.
|
||||
Hot path uses get_effective_policies only — no DB, no cache.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._tool_input_policies: Dict[str, str] = {}
|
||||
self._tool_output_policies: Dict[str, str] = {}
|
||||
self._blocked_tools_by_op_id: Dict[str, List[str]] = {}
|
||||
self._initialized: bool = False
|
||||
|
||||
def is_initialized(self) -> bool:
|
||||
return self._initialized
|
||||
|
||||
async def sync_tool_policy_from_db(self, prisma_client: "PrismaClient") -> None:
|
||||
"""Load all tool policies and object-permission blocked_tools from DB."""
|
||||
try:
|
||||
tools = await prisma_client.db.litellm_tooltable.find_many()
|
||||
self._tool_input_policies = {
|
||||
row.tool_name: getattr(row, "input_policy", "untrusted") or "untrusted"
|
||||
for row in tools
|
||||
}
|
||||
self._tool_output_policies = {
|
||||
row.tool_name: getattr(row, "output_policy", "untrusted") or "untrusted"
|
||||
for row in tools
|
||||
}
|
||||
|
||||
perms = await prisma_client.db.litellm_objectpermissiontable.find_many()
|
||||
self._blocked_tools_by_op_id = {}
|
||||
for row in perms:
|
||||
op_id = getattr(row, "object_permission_id", None)
|
||||
blocked = getattr(row, "blocked_tools", None) or []
|
||||
if op_id:
|
||||
self._blocked_tools_by_op_id[op_id] = list(blocked)
|
||||
|
||||
self._initialized = True
|
||||
verbose_proxy_logger.info(
|
||||
"ToolPolicyRegistry: synced %d tool policies and %d object permissions from DB",
|
||||
len(self._tool_input_policies),
|
||||
len(self._blocked_tools_by_op_id),
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"ToolPolicyRegistry sync_tool_policy_from_db error: %s", e
|
||||
)
|
||||
raise
|
||||
|
||||
def get_input_policy(self, tool_name: str) -> str:
|
||||
return self._tool_input_policies.get(tool_name, "untrusted")
|
||||
|
||||
def get_output_policy(self, tool_name: str) -> str:
|
||||
return self._tool_output_policies.get(tool_name, "untrusted")
|
||||
|
||||
def get_effective_policies(
|
||||
self,
|
||||
tool_names: List[str],
|
||||
object_permission_id: Optional[str] = None,
|
||||
team_object_permission_id: Optional[str] = None,
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
Return effective input_policy per tool from in-memory state.
|
||||
If tool is in key or team blocked_tools -> "blocked", else global input_policy or "untrusted".
|
||||
"""
|
||||
if not tool_names:
|
||||
return {}
|
||||
blocked: set = set()
|
||||
for op_id in (object_permission_id, team_object_permission_id):
|
||||
if op_id and op_id.strip():
|
||||
blocked.update(
|
||||
self._blocked_tools_by_op_id.get(op_id.strip(), [])
|
||||
)
|
||||
result: Dict[str, str] = {}
|
||||
for name in tool_names:
|
||||
if name in blocked:
|
||||
result[name] = "blocked"
|
||||
else:
|
||||
result[name] = self._tool_input_policies.get(name, "untrusted")
|
||||
return result
|
||||
|
||||
|
||||
_tool_policy_registry: Optional[ToolPolicyRegistry] = None
|
||||
|
||||
|
||||
def get_tool_policy_registry() -> ToolPolicyRegistry:
|
||||
"""Return the global ToolPolicyRegistry singleton."""
|
||||
global _tool_policy_registry
|
||||
if _tool_policy_registry is None:
|
||||
_tool_policy_registry = ToolPolicyRegistry()
|
||||
return _tool_policy_registry
|
||||
|
||||
|
||||
async def add_tool_to_object_permission_blocked(
|
||||
prisma_client: "PrismaClient",
|
||||
object_permission_id: str,
|
||||
tool_name: str,
|
||||
) -> bool:
|
||||
"""Add tool_name to the permission's blocked_tools if not already present."""
|
||||
if not object_permission_id or not tool_name:
|
||||
return False
|
||||
try:
|
||||
row = await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
)
|
||||
if row is None:
|
||||
return False
|
||||
current = list(getattr(row, "blocked_tools", []) or [])
|
||||
if tool_name in current:
|
||||
return True
|
||||
current.append(tool_name)
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
data={"blocked_tools": current},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer add_tool_to_object_permission_blocked error: %s", e
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
async def remove_tool_from_object_permission_blocked(
|
||||
prisma_client: "PrismaClient",
|
||||
object_permission_id: str,
|
||||
tool_name: str,
|
||||
) -> bool:
|
||||
"""Remove tool_name from the permission's blocked_tools. Returns False if tool was not in list."""
|
||||
if not object_permission_id or not tool_name:
|
||||
return False
|
||||
try:
|
||||
row = await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
)
|
||||
if row is None:
|
||||
return False
|
||||
current = list(getattr(row, "blocked_tools", []) or [])
|
||||
if tool_name not in current:
|
||||
return False
|
||||
current = [t for t in current if t != tool_name]
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
where={"object_permission_id": object_permission_id},
|
||||
data={"blocked_tools": current},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer remove_tool_from_object_permission_blocked error: %s",
|
||||
e,
|
||||
)
|
||||
return False
|
||||
|
|
|
|||
60
litellm/proxy/dd_span_tagger.py
Normal file
60
litellm/proxy/dd_span_tagger.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
from typing import Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.dd_tracing import set_active_span_tag
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
class DDSpanTagger:
|
||||
"""Best-effort helpers for tagging the active Datadog APM span with LiteLLM request metadata."""
|
||||
|
||||
@staticmethod
|
||||
def tag_call_id(litellm_call_id: Optional[str]) -> None:
|
||||
"""
|
||||
Attach LiteLLM call id to the active Datadog APM span.
|
||||
|
||||
This enables searching APM traces by LiteLLM call id returned in
|
||||
`x-litellm-call-id`.
|
||||
"""
|
||||
if not litellm_call_id:
|
||||
return
|
||||
try:
|
||||
set_active_span_tag("litellm.call_id", str(litellm_call_id))
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to tag active ddtrace span with litellm.call_id",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def tag_request(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
requested_model: Optional[str],
|
||||
) -> None:
|
||||
"""
|
||||
Attach key and model tags to the active Datadog APM span.
|
||||
|
||||
Tags set (all best-effort, skipped when value is absent):
|
||||
- ``litellm.key_alias`` — human-readable alias for the API key
|
||||
- ``litellm.key_hash`` — hashed API key (safe to log; never the raw secret)
|
||||
- ``litellm.requested_model``— model name as sent by the client
|
||||
|
||||
Use cases:
|
||||
- Trace all requests from a specific user/key: filter by ``litellm.key_alias`` or
|
||||
``litellm.key_hash``.
|
||||
- Trace all requests for a specific model: filter by ``litellm.requested_model``.
|
||||
|
||||
Note: key_alias / key_hash are not available for unauthenticated (e.g. 401) requests.
|
||||
"""
|
||||
try:
|
||||
if user_api_key_dict.key_alias:
|
||||
set_active_span_tag("litellm.key_alias", str(user_api_key_dict.key_alias))
|
||||
if user_api_key_dict.token:
|
||||
set_active_span_tag("litellm.key_hash", str(user_api_key_dict.token))
|
||||
if requested_model:
|
||||
set_active_span_tag("litellm.requested_model", str(requested_model))
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to tag active ddtrace span with key/model tags",
|
||||
exc_info=True,
|
||||
)
|
||||
|
|
@ -1624,11 +1624,11 @@ def _build_field_dict(
|
|||
# Determine the field type from annotation
|
||||
field_type = _get_field_type_from_annotation(field_annotation)
|
||||
|
||||
# Check for custom UI type override (ui_type preferred; "type" leaks into OpenAPI and breaks schema)
|
||||
field_json_schema_extra = getattr(field, "json_schema_extra", {}) or {}
|
||||
# Check for custom UI type override
|
||||
field_json_schema_extra = getattr(field, "json_schema_extra", {})
|
||||
if field_json_schema_extra and "ui_type" in field_json_schema_extra:
|
||||
ut = field_json_schema_extra["ui_type"]
|
||||
field_type = ut if isinstance(ut, str) else getattr(ut, "value", ut)
|
||||
ui_type = field_json_schema_extra["ui_type"]
|
||||
field_type = ui_type.value if hasattr(ui_type, "value") else ui_type
|
||||
elif field_json_schema_extra and "type" in field_json_schema_extra:
|
||||
field_type = field_json_schema_extra["type"]
|
||||
|
||||
|
|
|
|||
|
|
@ -1,13 +1,16 @@
|
|||
"""
|
||||
Tool Policy Guardrail
|
||||
|
||||
Reads call_policy from LiteLLM_ToolTable and enforces it on LLM requests/responses.
|
||||
Reads input_policy / output_policy from LiteLLM_ToolTable and enforces them.
|
||||
|
||||
Policy values:
|
||||
"trusted" - allow through (no action)
|
||||
"untrusted" - allow through (no action; default for newly discovered tools)
|
||||
Input policy values:
|
||||
"untrusted" - allow through (default for newly discovered tools)
|
||||
"trusted" - only allow if conversation contains no untrusted tool output
|
||||
"blocked" - raise HTTPException, preventing the tool call
|
||||
"dual_llm" - (Phase 3) send to second LLM for verification; currently treated as allowed
|
||||
|
||||
Output policy values:
|
||||
"untrusted" - output may be tainted (default)
|
||||
"trusted" - output is verified safe
|
||||
|
||||
Configuration in proxy config YAML:
|
||||
guardrails:
|
||||
|
|
@ -15,25 +18,18 @@ Configuration in proxy config YAML:
|
|||
litellm_params:
|
||||
guardrail: tool_policy
|
||||
mode: post_call
|
||||
|
||||
or both pre and post call:
|
||||
- guardrail_name: "tool_policy"
|
||||
litellm_params:
|
||||
guardrail: tool_policy
|
||||
mode: during_call # runs before LLM and on response
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.proxy.guardrails.tool_name_extraction import extract_request_tool_names
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
|
@ -43,12 +39,71 @@ if TYPE_CHECKING:
|
|||
GUARDRAIL_NAME = "tool_policy"
|
||||
|
||||
|
||||
def _get_request_object_permission_ids(
|
||||
request_data: dict,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""Extract object_permission_id and team_object_permission_id from request_data."""
|
||||
if not request_data:
|
||||
return None, None
|
||||
for key in ("litellm_metadata", "metadata"):
|
||||
meta = request_data.get(key)
|
||||
if not isinstance(meta, dict):
|
||||
continue
|
||||
auth = meta.get("user_api_key_auth")
|
||||
if auth is not None and hasattr(auth, "object_permission_id"):
|
||||
key_op = getattr(auth, "object_permission_id", None)
|
||||
team_op = getattr(auth, "team_object_permission_id", None)
|
||||
if key_op is not None or team_op is not None:
|
||||
return (
|
||||
str(key_op).strip() if key_op else None,
|
||||
str(team_op).strip() if team_op else None,
|
||||
)
|
||||
key_op = meta.get("user_api_key_object_permission_id")
|
||||
team_op = meta.get("user_api_key_team_object_permission_id")
|
||||
if key_op is not None or team_op is not None:
|
||||
return (
|
||||
str(key_op).strip() if key_op else None,
|
||||
str(team_op).strip() if team_op else None,
|
||||
)
|
||||
return None, None
|
||||
|
||||
|
||||
def _get_request_route_from_data(request_data: dict) -> Optional[str]:
|
||||
"""Get request route from request_data (metadata or top-level)."""
|
||||
route = request_data.get("user_api_key_request_route")
|
||||
if route:
|
||||
return route
|
||||
meta = request_data.get("metadata") or request_data.get("litellm_metadata") or {}
|
||||
return meta.get("user_api_key_request_route")
|
||||
|
||||
|
||||
def _resolve_tool_names_from_messages(messages: List[dict]) -> Dict[str, str]:
|
||||
"""
|
||||
Build a map of tool_call_id -> tool_name from assistant messages' tool_calls.
|
||||
Used to resolve which tool produced each tool result in the conversation.
|
||||
"""
|
||||
mapping: Dict[str, str] = {}
|
||||
for msg in messages:
|
||||
if msg.get("role") != "assistant":
|
||||
continue
|
||||
tool_calls = msg.get("tool_calls") or []
|
||||
for tc in tool_calls:
|
||||
if isinstance(tc, dict):
|
||||
tc_id = tc.get("id")
|
||||
fn = (tc.get("function") or {}).get("name")
|
||||
else:
|
||||
tc_id = getattr(tc, "id", None)
|
||||
fn_obj = getattr(tc, "function", None)
|
||||
fn = getattr(fn_obj, "name", None) if fn_obj else None
|
||||
if tc_id and fn:
|
||||
mapping[tc_id] = fn
|
||||
return mapping
|
||||
|
||||
|
||||
class ToolPolicyGuardrail(CustomGuardrail):
|
||||
"""
|
||||
Guardrail that enforces per-tool call policies stored in LiteLLM_ToolTable.
|
||||
|
||||
Tools with call_policy="blocked" are rejected before/after the LLM call.
|
||||
Tools with call_policy="trusted" or "untrusted" pass through unchanged.
|
||||
Guardrail that enforces per-tool input/output policies from the in-memory
|
||||
ToolPolicyRegistry (synced from DB).
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
|
|
@ -59,7 +114,6 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
GuardrailEventHooks.during_call,
|
||||
]
|
||||
super().__init__(**kwargs)
|
||||
self._policy_cache: DualCache = DualCache()
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
|
|
@ -70,12 +124,7 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""
|
||||
Enforce tool policies on both request tools and response tool_calls.
|
||||
|
||||
- input_type="request": check inputs["tools"] (tool definitions in the LLM request)
|
||||
- input_type="response": check inputs["tool_calls"] (tool_calls in the LLM response)
|
||||
|
||||
Raises HTTPException (400) if any tool is "blocked".
|
||||
Enforce input_policy and output_policy trust chain on request tools / response tool_calls.
|
||||
"""
|
||||
if input_type == "request":
|
||||
tools = inputs.get("tools") or []
|
||||
|
|
@ -86,7 +135,11 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
and isinstance(t.get("function"), dict)
|
||||
and t["function"].get("name")
|
||||
]
|
||||
else: # response
|
||||
if not tool_names:
|
||||
route = _get_request_route_from_data(request_data)
|
||||
if route:
|
||||
tool_names = extract_request_tool_names(route, request_data)
|
||||
else:
|
||||
tool_calls = inputs.get("tool_calls") or []
|
||||
tool_names = []
|
||||
for tc in tool_calls:
|
||||
|
|
@ -101,12 +154,25 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
if not tool_names:
|
||||
return inputs
|
||||
|
||||
policy_map = await self._get_policies_cached(tool_names)
|
||||
object_permission_id, team_object_permission_id = (
|
||||
_get_request_object_permission_ids(request_data)
|
||||
)
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool_policy_registry
|
||||
|
||||
registry = get_tool_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
return inputs
|
||||
|
||||
# Stage 1: Check for blocked tools (input_policy=blocked or per-key/team override)
|
||||
policy_map = registry.get_effective_policies(
|
||||
tool_names,
|
||||
object_permission_id=object_permission_id,
|
||||
team_object_permission_id=team_object_permission_id,
|
||||
)
|
||||
blocked = [name for name in tool_names if policy_map.get(name) == "blocked"]
|
||||
if blocked:
|
||||
verbose_proxy_logger.warning(
|
||||
"ToolPolicyGuardrail: blocking tool(s) %s (policy=blocked)", blocked
|
||||
"ToolPolicyGuardrail: blocking tool(s) %s (input_policy=blocked)", blocked
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -117,47 +183,47 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
},
|
||||
)
|
||||
|
||||
# Stage 2: Trust chain enforcement (response path only)
|
||||
# For each tool with input_policy=trusted, check if conversation
|
||||
# contains output from tools with output_policy=untrusted
|
||||
if input_type == "response":
|
||||
trusted_input_tools = [
|
||||
name for name in tool_names if policy_map.get(name) == "trusted"
|
||||
]
|
||||
if trusted_input_tools:
|
||||
messages = request_data.get("messages") or []
|
||||
tc_id_to_name = _resolve_tool_names_from_messages(messages)
|
||||
|
||||
untrusted_sources: List[str] = []
|
||||
for msg in messages:
|
||||
if msg.get("role") != "tool":
|
||||
continue
|
||||
tool_call_id = msg.get("tool_call_id")
|
||||
source_tool = tc_id_to_name.get(tool_call_id, "") if tool_call_id else ""
|
||||
if not source_tool:
|
||||
continue
|
||||
if registry.get_output_policy(source_tool) == "untrusted":
|
||||
if source_tool not in untrusted_sources:
|
||||
untrusted_sources.append(source_tool)
|
||||
|
||||
if untrusted_sources:
|
||||
verbose_proxy_logger.warning(
|
||||
"ToolPolicyGuardrail: trust chain violation — %s require trusted input "
|
||||
"but conversation has untrusted output from %s",
|
||||
trusted_input_tools,
|
||||
untrusted_sources,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Violated tool policy",
|
||||
"blocked_tools": trusted_input_tools,
|
||||
"untrusted_sources": untrusted_sources,
|
||||
"message": (
|
||||
f"{', '.join(trusted_input_tools)} requires trusted input but "
|
||||
f"conversation contains untrusted output from {', '.join(untrusted_sources)}."
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
async def _get_policies_cached(self, tool_names: List[str]) -> Dict[str, str]:
|
||||
"""
|
||||
Batch-fetch call_policy for the given tool names.
|
||||
|
||||
Caches per individual tool name (not per combination) so that adding
|
||||
a new tool to a request doesn't invalidate the cached policies for all
|
||||
the other tools already in the cache.
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import get_tools_by_names
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if not tool_names or prisma_client is None:
|
||||
return {}
|
||||
|
||||
result: Dict[str, str] = {}
|
||||
cache_misses: List[str] = []
|
||||
|
||||
for name in tool_names:
|
||||
cached = await self._policy_cache.async_get_cache(f"tool_policy:{name}")
|
||||
if cached is not None and isinstance(cached, str):
|
||||
result[name] = cached
|
||||
else:
|
||||
cache_misses.append(name)
|
||||
|
||||
if cache_misses:
|
||||
fetched = await get_tools_by_names(
|
||||
prisma_client=prisma_client, tool_names=cache_misses
|
||||
)
|
||||
for name, policy in fetched.items():
|
||||
result[name] = policy
|
||||
await self._policy_cache.async_set_cache(
|
||||
key=f"tool_policy:{name}",
|
||||
value=policy,
|
||||
ttl=TOOL_POLICY_CACHE_TTL_SECONDS,
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"ToolPolicyGuardrail: fetched %d policies from DB (cache hits: %d)",
|
||||
len(cache_misses),
|
||||
len(tool_names) - len(cache_misses),
|
||||
)
|
||||
|
||||
return result
|
||||
|
|
|
|||
85
litellm/proxy/guardrails/tool_name_extraction.py
Normal file
85
litellm/proxy/guardrails/tool_name_extraction.py
Normal file
|
|
@ -0,0 +1,85 @@
|
|||
"""
|
||||
Extract tool names from request body by route/call type.
|
||||
|
||||
Used by auth (check_tools_allowlist) and ToolPolicyGuardrail so tool-format
|
||||
knowledge lives in one place. Uses guardrail translation handlers where available,
|
||||
with standalone extractors for generate_content and MCP.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
# Call types that have no guardrail translation handler; we use standalone extractors
|
||||
STANDALONE_EXTRACTORS: Dict[str, Any] = {}
|
||||
|
||||
|
||||
def _extract_generate_content_tool_names(data: dict) -> List[str]:
|
||||
"""Google generateContent: tools[].functionDeclarations[].name"""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
for decl in tool.get("functionDeclarations") or []:
|
||||
if isinstance(decl, dict) and decl.get("name"):
|
||||
names.append(str(decl["name"]))
|
||||
return names
|
||||
|
||||
|
||||
def _extract_mcp_tool_names(data: dict) -> List[str]:
|
||||
"""MCP call_tool: name or mcp_tool_name in body"""
|
||||
names: List[str] = []
|
||||
name = data.get("name") or data.get("mcp_tool_name")
|
||||
if name:
|
||||
names.append(str(name))
|
||||
return names
|
||||
|
||||
|
||||
def _register_standalone_extractors() -> None:
|
||||
if STANDALONE_EXTRACTORS:
|
||||
return
|
||||
STANDALONE_EXTRACTORS[CallTypes.generate_content.value] = _extract_generate_content_tool_names
|
||||
STANDALONE_EXTRACTORS[CallTypes.agenerate_content.value] = _extract_generate_content_tool_names
|
||||
STANDALONE_EXTRACTORS[CallTypes.call_mcp_tool.value] = _extract_mcp_tool_names
|
||||
|
||||
|
||||
# Tool-capable call types (routes that can send tools in the request)
|
||||
TOOL_CAPABLE_CALL_TYPES = frozenset({
|
||||
CallTypes.completion.value,
|
||||
CallTypes.acompletion.value,
|
||||
CallTypes.responses.value,
|
||||
CallTypes.aresponses.value,
|
||||
CallTypes.anthropic_messages.value,
|
||||
CallTypes.generate_content.value,
|
||||
CallTypes.agenerate_content.value,
|
||||
CallTypes.call_mcp_tool.value,
|
||||
})
|
||||
|
||||
|
||||
def extract_request_tool_names(route: str, data: dict) -> List[str]:
|
||||
"""
|
||||
Extract tool names from the request body for the given route.
|
||||
Uses guardrail translation handlers when available, else standalone extractors
|
||||
for generate_content and MCP. Returns [] for non-tool-capable routes or when
|
||||
no tools are present.
|
||||
"""
|
||||
call_types = get_call_types_for_route(route)
|
||||
if not call_types:
|
||||
return []
|
||||
_register_standalone_extractors()
|
||||
mappings = load_guardrail_translation_mappings()
|
||||
for call_type in call_types:
|
||||
if not isinstance(call_type, CallTypes):
|
||||
continue
|
||||
if call_type.value not in TOOL_CAPABLE_CALL_TYPES:
|
||||
continue
|
||||
if call_type.value in STANDALONE_EXTRACTORS:
|
||||
return STANDALONE_EXTRACTORS[call_type.value](data)
|
||||
handler_cls = mappings.get(call_type)
|
||||
if handler_cls is not None:
|
||||
names = handler_cls().extract_request_tool_names(data)
|
||||
if names:
|
||||
return names
|
||||
return []
|
||||
|
|
@ -89,6 +89,25 @@ def _get_metadata_variable_name(request: Request) -> str:
|
|||
return "metadata"
|
||||
|
||||
|
||||
def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str]:
|
||||
"""
|
||||
Extract chain id for call chaining from request headers.
|
||||
|
||||
x-litellm-trace-id and x-litellm-session-id are interchangeable; when both
|
||||
are present, x-litellm-trace-id takes precedence. Header keys are matched
|
||||
case-insensitively so this works with raw header dicts from any transport.
|
||||
|
||||
Used by MCP (and other paths that have raw_headers but no Request) to set
|
||||
litellm_trace_id/litellm_session_id for spend logs and logging consistency.
|
||||
"""
|
||||
if not headers:
|
||||
return None
|
||||
normalized = {k.lower(): v for k, v in headers.items() if isinstance(k, str)}
|
||||
return normalized.get("x-litellm-trace-id") or normalized.get(
|
||||
"x-litellm-session-id"
|
||||
)
|
||||
|
||||
|
||||
def safe_add_api_version_from_query_params(data: dict, request: Request):
|
||||
try:
|
||||
if hasattr(request, "query_params"):
|
||||
|
|
@ -177,12 +196,12 @@ def _get_dynamic_logging_metadata(
|
|||
user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig
|
||||
) -> Optional[TeamCallbackMetadata]:
|
||||
callback_settings_obj: Optional[TeamCallbackMetadata] = None
|
||||
key_dynamic_logging_settings: Optional[
|
||||
dict
|
||||
] = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict)
|
||||
team_dynamic_logging_settings: Optional[
|
||||
dict
|
||||
] = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict)
|
||||
key_dynamic_logging_settings: Optional[dict] = (
|
||||
KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict)
|
||||
)
|
||||
team_dynamic_logging_settings: Optional[dict] = (
|
||||
KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict)
|
||||
)
|
||||
#########################################################################################
|
||||
# Key-based callbacks
|
||||
#########################################################################################
|
||||
|
|
@ -576,9 +595,13 @@ class LiteLLMProxyRequestSetup:
|
|||
#########################################################################################
|
||||
# Finally update the requests metadata with the `metadata_from_headers`
|
||||
#########################################################################################
|
||||
|
||||
agent_id_from_header = headers.get("x-litellm-agent-id")
|
||||
trace_id_from_header = headers.get("x-litellm-trace-id")
|
||||
session_id_from_header = headers.get("x-litellm-session-id")
|
||||
# x-litellm-trace-id and x-litellm-session-id are interchangeable for call chaining
|
||||
chain_id = headers.get("x-litellm-trace-id") or headers.get(
|
||||
"x-litellm-session-id"
|
||||
)
|
||||
|
||||
|
||||
if agent_id_from_header:
|
||||
metadata_from_headers["agent_id"] = agent_id_from_header
|
||||
|
|
@ -586,16 +609,13 @@ class LiteLLMProxyRequestSetup:
|
|||
f"Extracted agent_id from header: {agent_id_from_header}"
|
||||
)
|
||||
|
||||
if trace_id_from_header:
|
||||
metadata_from_headers["trace_id"] = trace_id_from_header
|
||||
if chain_id:
|
||||
metadata_from_headers["trace_id"] = chain_id
|
||||
metadata_from_headers["session_id"] = chain_id
|
||||
data["litellm_session_id"] = chain_id
|
||||
data["litellm_trace_id"] = chain_id
|
||||
verbose_proxy_logger.debug(
|
||||
f"Extracted trace_id from header: {trace_id_from_header}"
|
||||
)
|
||||
|
||||
if session_id_from_header:
|
||||
metadata_from_headers["session_id"] = session_id_from_header
|
||||
verbose_proxy_logger.debug(
|
||||
f"Extracted session_id from header: {session_id_from_header}"
|
||||
f"Extracted chain_id from header (trace-id/session-id): {chain_id}"
|
||||
)
|
||||
|
||||
if isinstance(data[_metadata_variable_name], dict):
|
||||
|
|
@ -702,11 +722,11 @@ class LiteLLMProxyRequestSetup:
|
|||
|
||||
## KEY-LEVEL SPEND LOGS / TAGS
|
||||
if "tags" in key_metadata and key_metadata["tags"] is not None:
|
||||
data[_metadata_variable_name][
|
||||
"tags"
|
||||
] = LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=data[_metadata_variable_name].get("tags"),
|
||||
tags_to_add=key_metadata["tags"],
|
||||
data[_metadata_variable_name]["tags"] = (
|
||||
LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=data[_metadata_variable_name].get("tags"),
|
||||
tags_to_add=key_metadata["tags"],
|
||||
)
|
||||
)
|
||||
if "disable_global_guardrails" in key_metadata and isinstance(
|
||||
key_metadata["disable_global_guardrails"], bool
|
||||
|
|
@ -839,14 +859,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
"""
|
||||
|
||||
from litellm.proxy.proxy_server import llm_router, premium_user
|
||||
from litellm.types.proxy.litellm_pre_call_utils import (
|
||||
RedactedDict,
|
||||
SecretFields,
|
||||
)
|
||||
from litellm.types.proxy.litellm_pre_call_utils import RedactedDict, SecretFields
|
||||
|
||||
_raw_headers: Dict[str, str] = RedactedDict(
|
||||
_safe_get_request_headers(request)
|
||||
)
|
||||
_raw_headers: Dict[str, str] = RedactedDict(_safe_get_request_headers(request))
|
||||
|
||||
forward_llm_auth = False
|
||||
if general_settings:
|
||||
|
|
@ -986,9 +1001,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
data[_metadata_variable_name]["litellm_api_version"] = version
|
||||
|
||||
if general_settings is not None:
|
||||
data[_metadata_variable_name][
|
||||
"global_max_parallel_requests"
|
||||
] = general_settings.get("global_max_parallel_requests", None)
|
||||
data[_metadata_variable_name]["global_max_parallel_requests"] = (
|
||||
general_settings.get("global_max_parallel_requests", None)
|
||||
)
|
||||
|
||||
### KEY-LEVEL Controls
|
||||
key_metadata = user_api_key_dict.metadata
|
||||
|
|
@ -1076,6 +1091,15 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
] = user_api_key_dict.user_max_budget
|
||||
|
||||
data[_metadata_variable_name]["user_api_key_metadata"] = user_api_key_dict.metadata
|
||||
data[_metadata_variable_name]["user_api_key_team_metadata"] = (
|
||||
user_api_key_dict.team_metadata
|
||||
)
|
||||
data[_metadata_variable_name]["user_api_key_object_permission_id"] = (
|
||||
getattr(user_api_key_dict, "object_permission_id", None)
|
||||
)
|
||||
data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = (
|
||||
getattr(user_api_key_dict, "team_object_permission_id", None)
|
||||
)
|
||||
data[_metadata_variable_name]["headers"] = _headers
|
||||
data[_metadata_variable_name]["endpoint"] = str(request.url)
|
||||
|
||||
|
|
|
|||
|
|
@ -54,16 +54,27 @@ def _resolve_model_for_cost_lookup(model: str) -> Tuple[str, Optional[str]]:
|
|||
deployments = llm_router.get_model_list(model_name=model)
|
||||
|
||||
if deployments and len(deployments) > 0:
|
||||
# Get the first deployment's litellm model
|
||||
first_deployment = deployments[0]
|
||||
litellm_params = first_deployment.get("litellm_params", {})
|
||||
model_info = first_deployment.get("model_info", {})
|
||||
|
||||
# Check base_model first (needed for Azure custom deployment names)
|
||||
base_model = model_info.get("base_model") or litellm_params.get(
|
||||
"base_model"
|
||||
)
|
||||
if base_model:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Resolved model '{model}' to base_model '{base_model}' from router"
|
||||
)
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
return base_model, custom_llm_provider
|
||||
|
||||
resolved_model = litellm_params.get("model")
|
||||
|
||||
if resolved_model:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Resolved model '{model}' to '{resolved_model}' from router"
|
||||
)
|
||||
# Extract custom_llm_provider if present
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
return resolved_model, custom_llm_provider
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -4,27 +4,87 @@ TOOL POLICY MANAGEMENT
|
|||
All /tool management endpoints
|
||||
|
||||
GET /v1/tool/list - List all discovered tools and their policies
|
||||
GET /v1/tool/policy/options - List available input/output policy options with descriptions
|
||||
GET /v1/tool/{tool_name} - Get a single tool's details
|
||||
POST /v1/tool/policy - Update the call_policy for a tool
|
||||
POST /v1/tool/policy - Update the input_policy / output_policy for a tool
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.tool_management import (
|
||||
LiteLLM_ToolTableRow,
|
||||
ToolCallPolicy,
|
||||
ToolDetailResponse,
|
||||
ToolInputPolicy,
|
||||
ToolListResponse,
|
||||
ToolOutputPolicy,
|
||||
ToolPolicyOption,
|
||||
ToolPolicyOptionsResponse,
|
||||
ToolPolicyUpdateRequest,
|
||||
ToolPolicyUpdateResponse,
|
||||
ToolUsageLogEntry,
|
||||
ToolUsageLogsResponse,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
TOOL_POLICY_OPTIONS = ToolPolicyOptionsResponse(
|
||||
input_policies=[
|
||||
ToolPolicyOption(
|
||||
value="untrusted",
|
||||
label="Untrusted",
|
||||
description="Tool accepts any input, including data from untrusted tool outputs. Default for newly discovered tools.",
|
||||
),
|
||||
ToolPolicyOption(
|
||||
value="trusted",
|
||||
label="Trusted",
|
||||
description="Tool requires trusted input. Blocked if the conversation contains output from any tool with output_policy=untrusted.",
|
||||
),
|
||||
ToolPolicyOption(
|
||||
value="blocked",
|
||||
label="Blocked",
|
||||
description="Tool is completely prohibited. Any attempt to call it is rejected.",
|
||||
),
|
||||
],
|
||||
output_policies=[
|
||||
ToolPolicyOption(
|
||||
value="untrusted",
|
||||
label="Untrusted",
|
||||
description="Tool output may contain unsafe content (prompt injection, risky code). Downstream tools with input_policy=trusted will be blocked.",
|
||||
),
|
||||
ToolPolicyOption(
|
||||
value="trusted",
|
||||
label="Trusted",
|
||||
description="Tool output is verified safe. Will not trigger trust-chain blocks on downstream tools.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/policy/options",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ToolPolicyOptionsResponse,
|
||||
)
|
||||
async def get_tool_policy_options(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Return the available input and output policy options with descriptions.
|
||||
Static data — no DB call.
|
||||
"""
|
||||
return TOOL_POLICY_OPTIONS
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/list",
|
||||
|
|
@ -33,14 +93,14 @@ router = APIRouter()
|
|||
response_model=ToolListResponse,
|
||||
)
|
||||
async def list_tools(
|
||||
call_policy: Optional[ToolCallPolicy] = None,
|
||||
input_policy: Optional[ToolInputPolicy] = None,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
List all auto-discovered tools and their call policies.
|
||||
List all auto-discovered tools and their policies.
|
||||
|
||||
Parameters:
|
||||
- call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked"
|
||||
- input_policy: Optional filter — one of "trusted", "untrusted", "blocked"
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
|
@ -51,13 +111,201 @@ async def list_tools(
|
|||
)
|
||||
|
||||
try:
|
||||
tools = await db_list_tools(prisma_client=prisma_client, call_policy=call_policy)
|
||||
tools = await db_list_tools(
|
||||
prisma_client=prisma_client, input_policy=input_policy
|
||||
)
|
||||
return ToolListResponse(tools=tools, total=len(tools))
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error listing tools: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/{tool_name:path}/detail",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ToolDetailResponse,
|
||||
)
|
||||
async def get_tool_detail(
|
||||
tool_name: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Get a single tool with its policy overrides (for UI detail view).
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool
|
||||
from litellm.proxy.db.tool_registry_writer import list_overrides_for_tool
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name)
|
||||
if tool is None:
|
||||
raise HTTPException(status_code=404, detail=f"Tool '{tool_name}' not found")
|
||||
overrides = await list_overrides_for_tool(
|
||||
prisma_client=prisma_client, tool_name=tool_name
|
||||
)
|
||||
return ToolDetailResponse(tool=tool, overrides=overrides)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error getting tool detail: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> Optional[str]:
|
||||
"""Short snippet from messages or proxy_server_request for tool usage log row."""
|
||||
if sl is None:
|
||||
return None
|
||||
messages = getattr(sl, "messages", None)
|
||||
if messages is not None:
|
||||
s = _snippet_str(messages, max_len)
|
||||
if s:
|
||||
return s
|
||||
psr = getattr(sl, "proxy_server_request", None)
|
||||
if not psr:
|
||||
return None
|
||||
if isinstance(psr, str):
|
||||
import json
|
||||
|
||||
try:
|
||||
psr = json.loads(psr)
|
||||
except Exception:
|
||||
return _snippet_str(psr, max_len)
|
||||
if isinstance(psr, dict):
|
||||
msgs = psr.get("messages")
|
||||
if msgs is None and isinstance(psr.get("body"), dict):
|
||||
msgs = psr["body"].get("messages")
|
||||
s = _snippet_str(msgs, max_len)
|
||||
if s:
|
||||
return s
|
||||
return _snippet_str(psr, max_len)
|
||||
|
||||
|
||||
def _snippet_str(text: Any, max_len: int = 200) -> Optional[str]:
|
||||
if text is None:
|
||||
return None
|
||||
if isinstance(text, str):
|
||||
s = text
|
||||
elif isinstance(text, list):
|
||||
parts = []
|
||||
for item in text:
|
||||
if isinstance(item, dict) and "content" in item:
|
||||
c = item["content"]
|
||||
parts.append(c if isinstance(c, str) else str(c))
|
||||
else:
|
||||
parts.append(str(item))
|
||||
s = " ".join(parts)
|
||||
else:
|
||||
s = str(text)
|
||||
if not s or s == "{}":
|
||||
return None
|
||||
return (s[:max_len] + "...") if len(s) > max_len else s
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/{tool_name:path}/logs",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ToolUsageLogsResponse,
|
||||
)
|
||||
async def get_tool_usage_logs(
|
||||
tool_name: str,
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(50, ge=1, le=100),
|
||||
start_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
end_date: Optional[str] = Query(None, description="YYYY-MM-DD"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Return paginated spend logs for requests that used this tool (from SpendLogToolIndex).
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
where: dict = {"tool_name": tool_name}
|
||||
if start_date or end_date:
|
||||
start_time_filter: Optional[datetime] = None
|
||||
end_time_filter: Optional[datetime] = None
|
||||
if start_date:
|
||||
try:
|
||||
start_time_filter = datetime.strptime(
|
||||
start_date + "T00:00:00", "%Y-%m-%dT%H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc)
|
||||
except ValueError:
|
||||
pass
|
||||
if end_date:
|
||||
try:
|
||||
end_time_filter = datetime.strptime(
|
||||
end_date + "T23:59:59", "%Y-%m-%dT%H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc)
|
||||
except ValueError:
|
||||
pass
|
||||
if start_time_filter is not None or end_time_filter is not None:
|
||||
where["start_time"] = {}
|
||||
if start_time_filter is not None:
|
||||
where["start_time"]["gte"] = start_time_filter
|
||||
if end_time_filter is not None:
|
||||
where["start_time"]["lte"] = end_time_filter
|
||||
|
||||
total = await prisma_client.db.litellm_spendlogtoolindex.count(where=where)
|
||||
index_rows = await prisma_client.db.litellm_spendlogtoolindex.find_many(
|
||||
where=where,
|
||||
order={"start_time": "desc"},
|
||||
skip=(page - 1) * page_size,
|
||||
take=page_size,
|
||||
)
|
||||
request_ids = [r.request_id for r in index_rows]
|
||||
if not request_ids:
|
||||
return ToolUsageLogsResponse(
|
||||
logs=[], total=total, page=page, page_size=page_size
|
||||
)
|
||||
|
||||
spend_logs = await prisma_client.db.litellm_spendlogs.find_many(
|
||||
where={"request_id": {"in": request_ids}}
|
||||
)
|
||||
log_by_id = {s.request_id: s for s in spend_logs}
|
||||
|
||||
logs_out: List[ToolUsageLogEntry] = []
|
||||
for r in index_rows:
|
||||
sl = log_by_id.get(r.request_id)
|
||||
if not sl:
|
||||
continue
|
||||
ts = (
|
||||
sl.startTime.isoformat()
|
||||
if hasattr(sl.startTime, "isoformat")
|
||||
else str(sl.startTime)
|
||||
)
|
||||
logs_out.append(
|
||||
ToolUsageLogEntry(
|
||||
id=sl.request_id,
|
||||
timestamp=ts,
|
||||
model=getattr(sl, "model", None) or None,
|
||||
spend=getattr(sl, "spend", None),
|
||||
total_tokens=getattr(sl, "total_tokens", None),
|
||||
input_snippet=_input_snippet_for_tool_log(sl),
|
||||
)
|
||||
)
|
||||
|
||||
return ToolUsageLogsResponse(
|
||||
logs=logs_out, total=total, page=page, page_size=page_size
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error getting tool usage logs: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/{tool_name:path}",
|
||||
tags=["tool management"],
|
||||
|
|
@ -70,9 +318,6 @@ async def get_tool(
|
|||
):
|
||||
"""
|
||||
Get details for a single tool.
|
||||
|
||||
Parameters:
|
||||
- tool_name: The tool name (supports namespaced names with slashes)
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
|
@ -85,9 +330,7 @@ async def get_tool(
|
|||
try:
|
||||
tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name)
|
||||
if tool is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Tool '{tool_name}' not found"
|
||||
)
|
||||
raise HTTPException(status_code=404, detail=f"Tool '{tool_name}' not found")
|
||||
return tool
|
||||
except HTTPException:
|
||||
raise
|
||||
|
|
@ -96,6 +339,80 @@ async def get_tool(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
async def _resolve_key_hash_to_object_permission_id(
|
||||
prisma_client: "PrismaClient",
|
||||
key_hash: str,
|
||||
) -> Optional[str]:
|
||||
"""Resolve key (hash or raw) to object_permission_id; create permission if key has none."""
|
||||
from litellm.proxy.proxy_server import hash_token
|
||||
|
||||
hashed = key_hash if "sk-" not in (key_hash or "") else hash_token(key_hash)
|
||||
if not hashed:
|
||||
return None
|
||||
row = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed}
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
op_id = getattr(row, "object_permission_id", None)
|
||||
if op_id:
|
||||
return op_id
|
||||
new_id = str(uuid.uuid4())
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
)
|
||||
updated_count = await prisma_client.db.litellm_verificationtoken.update_many(
|
||||
where={"token": hashed, "object_permission_id": None},
|
||||
data={"object_permission_id": new_id},
|
||||
)
|
||||
if updated_count == 0:
|
||||
await prisma_client.db.litellm_objectpermissiontable.delete(
|
||||
where={"object_permission_id": new_id}
|
||||
)
|
||||
row = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed}
|
||||
)
|
||||
return getattr(row, "object_permission_id", None) if row else None
|
||||
return new_id
|
||||
|
||||
|
||||
async def _resolve_team_id_to_object_permission_id(
|
||||
prisma_client: "PrismaClient",
|
||||
team_id: str,
|
||||
) -> Optional[str]:
|
||||
"""Resolve team_id to object_permission_id; create permission if team has none."""
|
||||
if not team_id or not team_id.strip():
|
||||
return None
|
||||
team_id_clean = team_id.strip()
|
||||
row = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id_clean},
|
||||
select={"object_permission_id": True},
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
op_id = getattr(row, "object_permission_id", None)
|
||||
if op_id:
|
||||
return op_id
|
||||
new_id = str(uuid.uuid4())
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
)
|
||||
updated_count = await prisma_client.db.litellm_teamtable.update_many(
|
||||
where={"team_id": team_id_clean, "object_permission_id": None},
|
||||
data={"object_permission_id": new_id},
|
||||
)
|
||||
if updated_count == 0:
|
||||
await prisma_client.db.litellm_objectpermissiontable.delete(
|
||||
where={"object_permission_id": new_id}
|
||||
)
|
||||
row = await prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id_clean},
|
||||
select={"object_permission_id": True},
|
||||
)
|
||||
return getattr(row, "object_permission_id", None) if row else None
|
||||
return new_id
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/tool/policy",
|
||||
tags=["tool management"],
|
||||
|
|
@ -107,15 +424,20 @@ async def update_tool_policy(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Set the call policy for a tool.
|
||||
Set the input_policy and/or output_policy for a tool (global), or block for a specific team/key (override).
|
||||
|
||||
Parameters:
|
||||
- tool_name: str - The tool to update
|
||||
- call_policy: "trusted" | "untrusted" | "dual_llm" | "blocked"
|
||||
|
||||
Setting a tool to "blocked" will cause the ToolPolicyGuardrail to remove
|
||||
that tool_call from LLM responses before returning them to the client.
|
||||
- input_policy: optional - "trusted" | "untrusted" | "blocked"
|
||||
- output_policy: optional - "trusted" | "untrusted"
|
||||
- team_id: optional - if set, create/update override for this team only
|
||||
- key_hash: optional - if set, create/update override for this key only
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
add_tool_to_object_permission_blocked,
|
||||
get_tool_policy_registry,
|
||||
remove_tool_from_object_permission_blocked,
|
||||
)
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
update_tool_policy as db_update_tool_policy,
|
||||
)
|
||||
|
|
@ -127,19 +449,80 @@ async def update_tool_policy(
|
|||
)
|
||||
|
||||
try:
|
||||
if data.team_id is not None or data.key_hash is not None:
|
||||
if data.team_id is not None and data.key_hash is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Provide either team_id or key_hash, not both",
|
||||
)
|
||||
if data.key_hash is not None:
|
||||
op_id = await _resolve_key_hash_to_object_permission_id(
|
||||
prisma_client, data.key_hash
|
||||
)
|
||||
else:
|
||||
op_id = await _resolve_team_id_to_object_permission_id(
|
||||
prisma_client, data.team_id or ""
|
||||
)
|
||||
if op_id is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail="Key or team not found for the given identifier",
|
||||
)
|
||||
is_blocking = data.input_policy == "blocked"
|
||||
if is_blocking:
|
||||
ok = await add_tool_to_object_permission_blocked(
|
||||
prisma_client=prisma_client,
|
||||
object_permission_id=op_id,
|
||||
tool_name=data.tool_name,
|
||||
)
|
||||
else:
|
||||
ok = await remove_tool_from_object_permission_blocked(
|
||||
prisma_client=prisma_client,
|
||||
object_permission_id=op_id,
|
||||
tool_name=data.tool_name,
|
||||
)
|
||||
if not ok:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to update policy override for tool '{data.tool_name}'",
|
||||
)
|
||||
registry = get_tool_policy_registry()
|
||||
if registry.is_initialized():
|
||||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return ToolPolicyUpdateResponse(
|
||||
tool_name=data.tool_name,
|
||||
input_policy=data.input_policy,
|
||||
output_policy=data.output_policy,
|
||||
updated=True,
|
||||
team_id=data.team_id,
|
||||
key_hash=data.key_hash,
|
||||
)
|
||||
|
||||
if data.input_policy is None and data.output_policy is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="At least one of input_policy or output_policy must be provided",
|
||||
)
|
||||
|
||||
updated = await db_update_tool_policy(
|
||||
prisma_client=prisma_client,
|
||||
tool_name=data.tool_name,
|
||||
call_policy=data.call_policy,
|
||||
updated_by=user_api_key_dict.user_id,
|
||||
input_policy=data.input_policy,
|
||||
output_policy=data.output_policy,
|
||||
)
|
||||
if updated is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Failed to update policy for tool '{data.tool_name}'"
|
||||
status_code=500,
|
||||
detail=f"Failed to update policy for tool '{data.tool_name}'",
|
||||
)
|
||||
registry = get_tool_policy_registry()
|
||||
if registry.is_initialized():
|
||||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return ToolPolicyUpdateResponse(
|
||||
tool_name=updated.tool_name,
|
||||
call_policy=updated.call_policy,
|
||||
input_policy=updated.input_policy,
|
||||
output_policy=updated.output_policy,
|
||||
updated=True,
|
||||
)
|
||||
except HTTPException:
|
||||
|
|
@ -147,3 +530,77 @@ async def update_tool_policy(
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error updating tool policy: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/v1/tool/{tool_name:path}/overrides",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def delete_tool_policy_override(
|
||||
tool_name: str,
|
||||
team_id: Optional[str] = Query(
|
||||
None, description="Team ID of the override to remove"
|
||||
),
|
||||
key_hash: Optional[str] = Query(
|
||||
None, description="Key hash of the override to remove"
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Remove a policy override for a tool. Specify the override by team_id or key_hash
|
||||
(exactly one required).
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
get_tool_policy_registry,
|
||||
remove_tool_from_object_permission_blocked,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
if team_id is None and key_hash is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="At least one of team_id or key_hash is required to identify the override",
|
||||
)
|
||||
if team_id is not None and key_hash is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Provide either team_id or key_hash, not both",
|
||||
)
|
||||
try:
|
||||
if key_hash is not None:
|
||||
op_id = await _resolve_key_hash_to_object_permission_id(
|
||||
prisma_client, key_hash
|
||||
)
|
||||
else:
|
||||
op_id = await _resolve_team_id_to_object_permission_id(
|
||||
prisma_client, team_id or ""
|
||||
)
|
||||
if op_id is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail="Key or team not found for the given identifier",
|
||||
)
|
||||
deleted = await remove_tool_from_object_permission_blocked(
|
||||
prisma_client=prisma_client,
|
||||
object_permission_id=op_id,
|
||||
tool_name=tool_name,
|
||||
)
|
||||
if not deleted:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"No override found for tool '{tool_name}' with the given scope",
|
||||
)
|
||||
registry = get_tool_policy_registry()
|
||||
if registry.is_initialized():
|
||||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return {"deleted": True, "tool_name": tool_name}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error deleting tool policy override: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
|
|
|||
|
|
@ -454,8 +454,35 @@ async def create_file( # noqa: PLR0915
|
|||
model=router_model, llm_router=llm_router
|
||||
)
|
||||
|
||||
# Apply team-level file expiry enforcement
|
||||
team_metadata = user_api_key_dict.team_metadata or {}
|
||||
enforced_file_expiry = team_metadata.get("enforced_file_expires_after")
|
||||
if enforced_file_expiry is not None:
|
||||
if "anchor" not in enforced_file_expiry or "seconds" not in enforced_file_expiry:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "enforced_file_expires_after must contain 'anchor' and 'seconds' keys",
|
||||
},
|
||||
)
|
||||
if enforced_file_expiry["anchor"] != "created_at":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"enforced_file_expires_after anchor must be 'created_at', got '{enforced_file_expiry['anchor']}'",
|
||||
},
|
||||
)
|
||||
expires_after = FileExpiresAfter(
|
||||
anchor="created_at",
|
||||
seconds=enforced_file_expiry["seconds"],
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"create_file expires_after: %s", expires_after
|
||||
)
|
||||
|
||||
_create_file_request = CreateFileRequest(
|
||||
file=file_data,
|
||||
file=file_data,
|
||||
purpose=cast(CREATE_FILE_REQUESTS_PURPOSE, purpose),
|
||||
expires_after=expires_after,
|
||||
**data
|
||||
|
|
|
|||
|
|
@ -4411,6 +4411,9 @@ class ProxyConfig:
|
|||
if self._should_load_db_object(object_type="search_tools"):
|
||||
await self._init_search_tools_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="tools"):
|
||||
await self._init_tool_policy_in_db(prisma_client=prisma_client)
|
||||
|
||||
if self._should_load_db_object(object_type="model_cost_map"):
|
||||
await self._check_and_reload_model_cost_map(prisma_client=prisma_client)
|
||||
|
||||
|
|
@ -4847,6 +4850,24 @@ class ProxyConfig:
|
|||
)
|
||||
)
|
||||
|
||||
async def _init_tool_policy_in_db(self, prisma_client: PrismaClient):
|
||||
"""
|
||||
Initialize tool policy from database into the in-memory registry.
|
||||
Synced periodically by add_deployment -> _init_non_llm_objects_in_db.
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool_policy_registry
|
||||
|
||||
try:
|
||||
registry = get_tool_policy_registry()
|
||||
await registry.sync_tool_policy_from_db(prisma_client=prisma_client)
|
||||
verbose_proxy_logger.debug("Successfully synced tool policy from DB")
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"litellm.proxy.proxy_server.py::ProxyConfig:_init_tool_policy_in_db - {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
|
||||
async def _init_vector_stores_in_db(self, prisma_client: PrismaClient):
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
|
||||
|
|
@ -10577,6 +10598,12 @@ async def async_queue_request(
|
|||
data["metadata"]["user_api_key_team_id"] = getattr(
|
||||
user_api_key_dict, "team_id", None
|
||||
)
|
||||
data["metadata"]["user_api_key_object_permission_id"] = getattr(
|
||||
user_api_key_dict, "object_permission_id", None
|
||||
)
|
||||
data["metadata"]["user_api_key_team_object_permission_id"] = getattr(
|
||||
user_api_key_dict, "team_object_permission_id", None
|
||||
)
|
||||
data["metadata"]["endpoint"] = str(request.url)
|
||||
|
||||
global user_temperature, user_request_timeout, user_max_tokens, user_api_base
|
||||
|
|
@ -11093,9 +11120,7 @@ async def get_favicon():
|
|||
|
||||
if favicon_url.startswith(("http://", "https://")):
|
||||
try:
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
async_client = get_async_httpx_client(
|
||||
|
|
|
|||
|
|
@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable {
|
|||
vector_stores String[] @default([])
|
||||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
|
||||
teams LiteLLM_TeamTable[]
|
||||
projects LiteLLM_ProjectTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
|
|
@ -928,6 +929,16 @@ model LiteLLM_SpendLogGuardrailIndex {
|
|||
@@index([policy_id, start_time])
|
||||
}
|
||||
|
||||
// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production
|
||||
model LiteLLM_SpendLogToolIndex {
|
||||
request_id String
|
||||
tool_name String // matches LiteLLM_ToolTable.tool_name; join for input_policy/output_policy etc.
|
||||
start_time DateTime
|
||||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
model LiteLLM_PromptTable {
|
||||
id String @id @default(uuid())
|
||||
|
|
@ -1065,23 +1076,27 @@ model LiteLLM_PolicyAttachmentTable {
|
|||
updated_by String?
|
||||
}
|
||||
|
||||
// Global tool registry - auto-discovered from LLM responses; admins set call_policy here
|
||||
// Global tool registry - auto-discovered from LLM responses; admins set input/output policies here
|
||||
model LiteLLM_ToolTable {
|
||||
tool_id String @id @default(uuid())
|
||||
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
|
||||
origin String? // MCP server name or "user_defined"
|
||||
call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked"
|
||||
call_count Int @default(0) // cumulative number of times this tool was seen
|
||||
assignments Json? @default("{}")
|
||||
key_hash String? // hash of the virtual key that first called this tool
|
||||
team_id String? // team that first called this tool
|
||||
key_alias String? // human-readable alias of the virtual key
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
updated_at DateTime @default(now()) @updatedAt
|
||||
updated_by String?
|
||||
tool_id String @id @default(uuid())
|
||||
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
|
||||
origin String? // MCP server name or "user_defined"
|
||||
input_policy String @default("untrusted") // "trusted" | "untrusted" | "blocked"
|
||||
output_policy String @default("untrusted") // "trusted" | "untrusted"
|
||||
call_count Int @default(0) // cumulative number of times this tool was seen
|
||||
assignments Json? @default("{}")
|
||||
key_hash String? // hash of the virtual key that first called this tool
|
||||
team_id String? // team that first called this tool
|
||||
key_alias String? // human-readable alias of the virtual key
|
||||
user_agent String? // user-agent of the first request that discovered this tool
|
||||
last_used_at DateTime? // timestamp of the most recent call
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
updated_at DateTime @default(now()) @updatedAt
|
||||
updated_by String?
|
||||
|
||||
@@index([call_policy])
|
||||
@@index([input_policy])
|
||||
@@index([output_policy])
|
||||
@@index([team_id])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -11,26 +11,21 @@ from pydantic import BaseModel
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
MAX_STRING_LENGTH_PROMPT_IN_DB as DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB,
|
||||
)
|
||||
from litellm.constants import \
|
||||
MAX_STRING_LENGTH_PROMPT_IN_DB as DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_litellm_metadata_from_kwargs,
|
||||
reconstruct_model_name,
|
||||
)
|
||||
get_litellm_metadata_from_kwargs, reconstruct_model_name)
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
|
||||
from litellm.proxy.utils import PrismaClient, hash_token
|
||||
from litellm.types.utils import (
|
||||
CostBreakdown,
|
||||
StandardLoggingGuardrailInformation,
|
||||
StandardLoggingMCPToolCall,
|
||||
StandardLoggingModelInformation,
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingVectorStoreRequest,
|
||||
VectorStoreSearchResponse,
|
||||
)
|
||||
from litellm.types.utils import (CostBreakdown,
|
||||
StandardLoggingGuardrailInformation,
|
||||
StandardLoggingMCPToolCall,
|
||||
StandardLoggingModelInformation,
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingVectorStoreRequest,
|
||||
VectorStoreSearchResponse)
|
||||
from litellm.utils import get_end_user_id_for_cost_tracking
|
||||
|
||||
|
||||
|
|
@ -116,16 +111,15 @@ def _get_spend_logs_metadata(
|
|||
# Filter the metadata dictionary to include only the specified keys
|
||||
clean_metadata = SpendLogsMetadata(
|
||||
**{ # type: ignore
|
||||
key: metadata.get(key)
|
||||
for key in SpendLogsMetadata.__annotations__.keys()
|
||||
key: metadata.get(key) for key in SpendLogsMetadata.__annotations__.keys()
|
||||
}
|
||||
)
|
||||
clean_metadata["applied_guardrails"] = applied_guardrails
|
||||
clean_metadata["batch_models"] = batch_models
|
||||
clean_metadata["mcp_tool_call_metadata"] = mcp_tool_call_metadata
|
||||
clean_metadata[
|
||||
"vector_store_request_metadata"
|
||||
] = _get_vector_store_request_for_spend_logs_payload(vector_store_request_metadata)
|
||||
clean_metadata["vector_store_request_metadata"] = (
|
||||
_get_vector_store_request_for_spend_logs_payload(vector_store_request_metadata)
|
||||
)
|
||||
clean_metadata["guardrail_information"] = guardrail_information
|
||||
clean_metadata["usage_object"] = usage_object
|
||||
clean_metadata["model_map_information"] = model_map_information
|
||||
|
|
@ -372,9 +366,11 @@ def get_logging_payload( # noqa: PLR0915
|
|||
guardrail_information=(
|
||||
standard_logging_payload.get("guardrail_information", None)
|
||||
if standard_logging_payload is not None
|
||||
else metadata.get("standard_logging_guardrail_information", None)
|
||||
if metadata is not None
|
||||
else None
|
||||
else (
|
||||
metadata.get("standard_logging_guardrail_information", None)
|
||||
if metadata is not None
|
||||
else None
|
||||
)
|
||||
),
|
||||
cold_storage_object_key=(
|
||||
standard_logging_payload["metadata"].get("cold_storage_object_key", None)
|
||||
|
|
@ -501,6 +497,7 @@ def _get_session_id_for_spend_log(
|
|||
"""
|
||||
from litellm._uuid import uuid
|
||||
|
||||
|
||||
if (
|
||||
standard_logging_payload is not None
|
||||
and standard_logging_payload.get("trace_id") is not None
|
||||
|
|
@ -515,9 +512,7 @@ def _get_session_id_for_spend_log(
|
|||
return str(uuid.uuid4())
|
||||
|
||||
|
||||
def _get_request_duration_ms(
|
||||
start_time: datetime, end_time: datetime
|
||||
) -> Optional[int]:
|
||||
def _get_request_duration_ms(start_time: datetime, end_time: datetime) -> Optional[int]:
|
||||
"""Compute request duration in milliseconds from start and end times."""
|
||||
try:
|
||||
return int((end_time - start_time).total_seconds() * 1000)
|
||||
|
|
@ -709,20 +704,20 @@ def _convert_to_json_serializable_dict(
|
|||
if max_depth <= 0:
|
||||
# Return a placeholder if max depth is exceeded
|
||||
return "<max_depth_exceeded>"
|
||||
|
||||
|
||||
if visited is None:
|
||||
visited = set()
|
||||
|
||||
|
||||
# Get the object's memory address to track visited objects
|
||||
obj_id = id(obj)
|
||||
if obj_id in visited:
|
||||
# Circular reference detected, return placeholder
|
||||
return "<circular_reference>"
|
||||
|
||||
|
||||
# Only track mutable objects (dict, list, objects with __dict__)
|
||||
if isinstance(obj, (dict, list)) or hasattr(obj, "__dict__"):
|
||||
visited.add(obj_id)
|
||||
|
||||
|
||||
try:
|
||||
if isinstance(obj, BaseModel):
|
||||
# Use Pydantic's model_dump() instead of pickle
|
||||
|
|
@ -741,7 +736,9 @@ def _convert_to_json_serializable_dict(
|
|||
]
|
||||
elif hasattr(obj, "__dict__"):
|
||||
# Handle objects with __dict__ attribute
|
||||
return _convert_to_json_serializable_dict(obj.__dict__, visited, max_depth - 1)
|
||||
return _convert_to_json_serializable_dict(
|
||||
obj.__dict__, visited, max_depth - 1
|
||||
)
|
||||
else:
|
||||
# Primitives (str, int, float, bool, None) pass through
|
||||
return obj
|
||||
|
|
@ -777,9 +774,7 @@ def _get_proxy_server_request_for_spend_logs_payload(
|
|||
# Apply message redaction if turn_off_message_logging is enabled
|
||||
if kwargs is not None:
|
||||
from litellm.litellm_core_utils.redact_messages import (
|
||||
perform_redaction,
|
||||
should_redact_message_logging,
|
||||
)
|
||||
perform_redaction, should_redact_message_logging)
|
||||
|
||||
# Build model_call_details dict to check redaction settings
|
||||
model_call_details = {
|
||||
|
|
@ -788,12 +783,12 @@ def _get_proxy_server_request_for_spend_logs_payload(
|
|||
"standard_callback_dynamic_params"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
# If redaction is enabled, convert to serializable dict before redacting
|
||||
if should_redact_message_logging(model_call_details=model_call_details):
|
||||
_request_body = _convert_to_json_serializable_dict(_request_body)
|
||||
perform_redaction(model_call_details=_request_body, result=None)
|
||||
|
||||
|
||||
_request_body = _sanitize_request_body_for_spend_logs_payload(_request_body)
|
||||
_request_body_json_str = json.dumps(_request_body, default=str)
|
||||
return _request_body_json_str
|
||||
|
|
@ -845,10 +840,8 @@ def _get_response_for_spend_logs_payload(
|
|||
# Apply message redaction if turn_off_message_logging is enabled
|
||||
if kwargs is not None:
|
||||
from litellm.litellm_core_utils.redact_messages import (
|
||||
perform_redaction,
|
||||
should_redact_message_logging,
|
||||
)
|
||||
|
||||
perform_redaction, should_redact_message_logging)
|
||||
|
||||
litellm_params = kwargs.get("litellm_params", {})
|
||||
model_call_details = {
|
||||
"litellm_params": litellm_params,
|
||||
|
|
@ -856,11 +849,13 @@ def _get_response_for_spend_logs_payload(
|
|||
"standard_callback_dynamic_params"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
# If redaction is enabled, convert to serializable dict before redacting
|
||||
if should_redact_message_logging(model_call_details=model_call_details):
|
||||
response_obj = _convert_to_json_serializable_dict(response_obj)
|
||||
response_obj = perform_redaction(model_call_details={}, result=response_obj)
|
||||
response_obj = perform_redaction(
|
||||
model_call_details={}, result=response_obj
|
||||
)
|
||||
|
||||
sanitized_wrapper = _sanitize_request_body_for_spend_logs_payload(
|
||||
{"response": response_obj}
|
||||
|
|
@ -882,7 +877,7 @@ def _should_store_prompts_and_responses_in_spend_logs() -> bool:
|
|||
|
||||
# Check general_settings (from DB or proxy_config.yaml)
|
||||
store_prompts_value = general_settings.get("store_prompts_in_spend_logs")
|
||||
|
||||
|
||||
# Normalize case: handle True/true/TRUE, False/false/FALSE, None/null
|
||||
if store_prompts_value is True:
|
||||
return True
|
||||
|
|
@ -890,7 +885,7 @@ def _should_store_prompts_and_responses_in_spend_logs() -> bool:
|
|||
# Case-insensitive string comparison
|
||||
if store_prompts_value.lower() == "true":
|
||||
return True
|
||||
|
||||
|
||||
# Also check environment variable
|
||||
return get_secret_bool("STORE_PROMPTS_IN_SPEND_LOGS") is True
|
||||
|
||||
|
|
|
|||
|
|
@ -3583,8 +3583,9 @@ class PrismaClient:
|
|||
def _get_engine_pid(self) -> int:
|
||||
try:
|
||||
engine = self.db._original_prisma._engine # type: ignore[attr-defined]
|
||||
if engine is not None and engine.process is not None:
|
||||
return engine.process.pid
|
||||
process = getattr(engine, "process", None) if engine is not None else None
|
||||
if process is not None:
|
||||
return process.pid
|
||||
except (AttributeError, TypeError):
|
||||
pass
|
||||
return 0
|
||||
|
|
@ -4688,6 +4689,19 @@ async def update_spend_logs_job(
|
|||
guardrail_tracking_err,
|
||||
)
|
||||
|
||||
# Tool usage tracking (same batch): SpendLogToolIndex for "last N requests for tool X"
|
||||
try:
|
||||
from litellm.proxy.db.spend_log_tool_index import process_spend_logs_tool_usage
|
||||
await process_spend_logs_tool_usage(
|
||||
prisma_client=prisma_client,
|
||||
logs_to_process=logs_to_process,
|
||||
)
|
||||
except Exception as tool_tracking_err:
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend tracking - tool usage tracking failed (non-fatal): %s",
|
||||
tool_tracking_err,
|
||||
)
|
||||
|
||||
|
||||
async def _monitor_spend_logs_queue(
|
||||
prisma_client: PrismaClient,
|
||||
|
|
|
|||
|
|
@ -745,6 +745,11 @@ def responses(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
# Decode any litellm-encoded encrypted-content item IDs back to their original IDs
|
||||
input = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
|
||||
input
|
||||
)
|
||||
|
||||
# Call the handler with _is_async flag instead of directly calling the async handler
|
||||
response = base_llm_http_handler.response_api_handler(
|
||||
model=model,
|
||||
|
|
@ -1617,6 +1622,12 @@ def compact_responses(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
# Decode any litellm-encoded encrypted-content item IDs back to their original IDs
|
||||
# before forwarding to the upstream provider.
|
||||
input = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
|
||||
input
|
||||
)
|
||||
|
||||
# Call the handler with _is_async flag instead of directly calling the async handler
|
||||
response = base_llm_http_handler.compact_response_api_handler(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -8,7 +8,10 @@ from typing import Any, Dict, Optional
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.constants import LITELLM_MAX_STREAMING_DURATION_SECONDS, STREAM_SSE_DONE_STRING
|
||||
from litellm.constants import (
|
||||
LITELLM_MAX_STREAMING_DURATION_SECONDS,
|
||||
STREAM_SSE_DONE_STRING,
|
||||
)
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -137,6 +140,31 @@ class BaseResponsesAPIStreamingIterator:
|
|||
)
|
||||
setattr(openai_responses_api_chunk, "response", response)
|
||||
|
||||
# Wrap encrypted_content in streaming events (output_item.added, output_item.done)
|
||||
if (
|
||||
self.litellm_metadata
|
||||
and self.litellm_metadata.get("encrypted_content_affinity_enabled")
|
||||
):
|
||||
event_type = getattr(openai_responses_api_chunk, "type", None)
|
||||
if event_type in (
|
||||
ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
):
|
||||
item = getattr(openai_responses_api_chunk, "item", None)
|
||||
if item:
|
||||
encrypted_content = getattr(item, "encrypted_content", None)
|
||||
if encrypted_content and isinstance(encrypted_content, str):
|
||||
model_id = (
|
||||
self.litellm_metadata.get("model_info", {}).get("id")
|
||||
if self.litellm_metadata
|
||||
else None
|
||||
)
|
||||
if model_id:
|
||||
wrapped_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
encrypted_content, model_id
|
||||
)
|
||||
setattr(item, "encrypted_content", wrapped_content)
|
||||
|
||||
# Store the completed response
|
||||
if (
|
||||
openai_responses_api_chunk
|
||||
|
|
|
|||
|
|
@ -217,8 +217,204 @@ class ResponsesAPIRequestUtils:
|
|||
responses_api_response["id"] = updated_id
|
||||
else:
|
||||
responses_api_response.id = updated_id
|
||||
|
||||
if litellm_metadata.get("encrypted_content_affinity_enabled"):
|
||||
responses_api_response = (
|
||||
ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(
|
||||
response=responses_api_response,
|
||||
model_id=model_id,
|
||||
)
|
||||
)
|
||||
|
||||
return responses_api_response
|
||||
|
||||
@staticmethod
|
||||
def _build_encrypted_item_id(model_id: str, item_id: str) -> str:
|
||||
"""Encode model_id into an output item ID for encrypted-content items.
|
||||
|
||||
Format: ``encitem_{base64("litellm:model_id:{model_id};item_id:{original_id}")}``
|
||||
"""
|
||||
assembled = f"litellm:model_id:{model_id};item_id:{item_id}"
|
||||
encoded = base64.b64encode(assembled.encode("utf-8")).decode("utf-8")
|
||||
return f"encitem_{encoded}"
|
||||
|
||||
@staticmethod
|
||||
def _decode_encrypted_item_id(encoded_id: str) -> Optional[Dict[str, str]]:
|
||||
"""Decode a litellm-encoded encrypted-content item ID.
|
||||
|
||||
Returns a dict with ``model_id`` and ``item_id`` keys, or ``None`` if
|
||||
the string is not a litellm-encoded item ID.
|
||||
"""
|
||||
if not encoded_id.startswith("encitem_"):
|
||||
return None
|
||||
try:
|
||||
cleaned = encoded_id[len("encitem_"):]
|
||||
# Restore any padding that may have been stripped in transit
|
||||
missing = len(cleaned) % 4
|
||||
if missing:
|
||||
cleaned += "=" * (4 - missing)
|
||||
decoded = base64.b64decode(cleaned.encode("utf-8")).decode("utf-8")
|
||||
# Split on first ";" only so that semicolons inside item_id are preserved
|
||||
parts = decoded.split(";", 1)
|
||||
if len(parts) < 2:
|
||||
return None
|
||||
model_id = parts[0].replace("litellm:model_id:", "")
|
||||
item_id = parts[1].replace("item_id:", "")
|
||||
return {"model_id": model_id, "item_id": item_id}
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _wrap_encrypted_content_with_model_id(
|
||||
encrypted_content: str, model_id: str
|
||||
) -> str:
|
||||
"""Wrap encrypted_content with model_id metadata for affinity routing.
|
||||
|
||||
When Codex or other clients send items with encrypted_content but no ID,
|
||||
we encode the model_id directly into the encrypted_content itself.
|
||||
|
||||
Format: ``litellm_enc:{base64("model_id:{model_id}")};{original_encrypted_content}``
|
||||
"""
|
||||
metadata = f"model_id:{model_id}"
|
||||
encoded_metadata = base64.b64encode(metadata.encode("utf-8")).decode("utf-8")
|
||||
return f"litellm_enc:{encoded_metadata};{encrypted_content}"
|
||||
|
||||
@staticmethod
|
||||
def _unwrap_encrypted_content_with_model_id(
|
||||
wrapped_content: str,
|
||||
) -> tuple[Optional[str], str]:
|
||||
"""Unwrap encrypted_content to extract model_id and original content.
|
||||
|
||||
Returns:
|
||||
Tuple of (model_id, original_encrypted_content).
|
||||
If not wrapped, returns (None, original_content).
|
||||
"""
|
||||
if not wrapped_content.startswith("litellm_enc:"):
|
||||
return None, wrapped_content
|
||||
|
||||
try:
|
||||
# Split on first ";" to separate metadata from content
|
||||
parts = wrapped_content.split(";", 1)
|
||||
if len(parts) < 2:
|
||||
return None, wrapped_content
|
||||
|
||||
metadata_b64 = parts[0].replace("litellm_enc:", "")
|
||||
original_content = parts[1]
|
||||
|
||||
# Restore padding if needed
|
||||
missing = len(metadata_b64) % 4
|
||||
if missing:
|
||||
metadata_b64 += "=" * (4 - missing)
|
||||
|
||||
decoded_metadata = base64.b64decode(metadata_b64.encode("utf-8")).decode(
|
||||
"utf-8"
|
||||
)
|
||||
model_id = decoded_metadata.replace("model_id:", "")
|
||||
return model_id, original_content
|
||||
except Exception:
|
||||
return None, wrapped_content
|
||||
|
||||
@staticmethod
|
||||
def _update_encrypted_content_item_ids_in_response(
|
||||
response: Union["ResponsesAPIResponse", Dict[str, Any]],
|
||||
model_id: Optional[str],
|
||||
) -> Union["ResponsesAPIResponse", Dict[str, Any]]:
|
||||
"""Rewrite item IDs for output items that contain ``encrypted_content``.
|
||||
|
||||
Encodes ``model_id`` into the item ID so that follow-up requests can be
|
||||
routed back to the originating deployment without any cache lookup.
|
||||
|
||||
For items without an ID (e.g., from Codex), encodes model_id directly
|
||||
into the encrypted_content itself.
|
||||
"""
|
||||
if not model_id:
|
||||
return response
|
||||
|
||||
output: Optional[list] = None
|
||||
if isinstance(response, dict):
|
||||
output = response.get("output")
|
||||
else:
|
||||
output = getattr(response, "output", None)
|
||||
|
||||
if not isinstance(output, list):
|
||||
return response
|
||||
|
||||
for item in output:
|
||||
if isinstance(item, dict):
|
||||
item_id = item.get("id")
|
||||
encrypted_content = item.get("encrypted_content")
|
||||
|
||||
if encrypted_content and isinstance(encrypted_content, str):
|
||||
# Always wrap encrypted_content with model_id for redundancy
|
||||
item["encrypted_content"] = (
|
||||
ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
encrypted_content, model_id
|
||||
)
|
||||
)
|
||||
# Also encode the ID if present
|
||||
if item_id and isinstance(item_id, str):
|
||||
item["id"] = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
model_id, item_id
|
||||
)
|
||||
else:
|
||||
item_id = getattr(item, "id", None)
|
||||
encrypted_content = getattr(item, "encrypted_content", None)
|
||||
|
||||
if encrypted_content and isinstance(encrypted_content, str):
|
||||
# Always wrap encrypted_content with model_id for redundancy
|
||||
try:
|
||||
item.encrypted_content = (
|
||||
ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
encrypted_content, model_id
|
||||
)
|
||||
)
|
||||
except AttributeError:
|
||||
pass
|
||||
# Also encode the ID if present
|
||||
if item_id and isinstance(item_id, str):
|
||||
try:
|
||||
item.id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
model_id, item_id
|
||||
)
|
||||
except AttributeError:
|
||||
pass
|
||||
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _restore_encrypted_content_item_ids_in_input(request_input: Any) -> Any:
|
||||
"""Decode litellm-encoded item IDs in request input back to original IDs.
|
||||
|
||||
Called before forwarding the request to the upstream provider so the
|
||||
provider receives the original item IDs and unwrapped encrypted_content.
|
||||
|
||||
Handles both:
|
||||
1. Items with encoded IDs (encitem_...)
|
||||
2. Items with wrapped encrypted_content (litellm_enc:...)
|
||||
"""
|
||||
if not isinstance(request_input, list):
|
||||
return request_input
|
||||
|
||||
for item in request_input:
|
||||
if isinstance(item, dict):
|
||||
item_id = item.get("id")
|
||||
if item_id and isinstance(item_id, str):
|
||||
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(item_id)
|
||||
if decoded:
|
||||
item["id"] = decoded["item_id"]
|
||||
|
||||
encrypted_content = item.get("encrypted_content")
|
||||
if encrypted_content and isinstance(encrypted_content, str):
|
||||
_, unwrapped = (
|
||||
ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(
|
||||
encrypted_content
|
||||
)
|
||||
)
|
||||
if unwrapped != encrypted_content:
|
||||
item["encrypted_content"] = unwrapped
|
||||
|
||||
return request_input
|
||||
|
||||
@staticmethod
|
||||
def _build_responses_api_response_id(
|
||||
custom_llm_provider: Optional[str],
|
||||
|
|
|
|||
|
|
@ -115,6 +115,9 @@ from litellm.router_utils.handle_error import (
|
|||
from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
|
||||
DeploymentAffinityCheck,
|
||||
)
|
||||
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
|
||||
EncryptedContentAffinityCheck,
|
||||
)
|
||||
from litellm.router_utils.pre_call_checks.model_rate_limit_check import (
|
||||
ModelRateLimitingCheck,
|
||||
)
|
||||
|
|
@ -1248,6 +1251,26 @@ class Router:
|
|||
self.optional_callbacks.append(affinity_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(affinity_callback)
|
||||
|
||||
# ---------------------------------------------------------------------
|
||||
# Encrypted content affinity
|
||||
# ---------------------------------------------------------------------
|
||||
if "encrypted_content_affinity" in optional_pre_call_checks:
|
||||
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
|
||||
EncryptedContentAffinityCheck,
|
||||
)
|
||||
|
||||
if self.optional_callbacks is None:
|
||||
self.optional_callbacks = []
|
||||
|
||||
already_registered = any(
|
||||
isinstance(cb, EncryptedContentAffinityCheck)
|
||||
for cb in self.optional_callbacks
|
||||
)
|
||||
if not already_registered:
|
||||
ec_callback = EncryptedContentAffinityCheck()
|
||||
self.optional_callbacks.append(ec_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(ec_callback)
|
||||
|
||||
# ---------------------------------------------------------------------
|
||||
# Remaining optional pre-call checks
|
||||
# ---------------------------------------------------------------------
|
||||
|
|
@ -1257,6 +1280,7 @@ class Router:
|
|||
"deployment_affinity",
|
||||
"responses_api_deployment_check",
|
||||
"session_affinity",
|
||||
"encrypted_content_affinity",
|
||||
):
|
||||
continue
|
||||
if pre_call_check == "prompt_caching":
|
||||
|
|
@ -8808,6 +8832,13 @@ class Router:
|
|||
if isinstance(healthy_deployments, dict):
|
||||
return healthy_deployments
|
||||
|
||||
# When encrypted content affinity pins to a specific deployment,
|
||||
if (
|
||||
request_kwargs.get("_encrypted_content_affinity_pinned")
|
||||
and len(healthy_deployments) == 1
|
||||
):
|
||||
return healthy_deployments[0]
|
||||
|
||||
start_time = time.time()
|
||||
if (
|
||||
self.routing_strategy == "usage-based-routing-v2"
|
||||
|
|
|
|||
|
|
@ -0,0 +1,172 @@
|
|||
"""
|
||||
Encrypted-content-aware deployment affinity for the Router.
|
||||
|
||||
When Codex or other models use `store: false` with `include: ["reasoning.encrypted_content"]`,
|
||||
the response output items contain encrypted reasoning tokens tied to the originating
|
||||
organization's API key. If a follow-up request containing those items is routed to a
|
||||
different deployment (different org), OpenAI rejects it with an `invalid_encrypted_content`
|
||||
error because the organization_id doesn't match.
|
||||
|
||||
This callback solves the problem by encoding the originating deployment's ``model_id``
|
||||
into the response output items that carry ``encrypted_content``. Two encoding strategies:
|
||||
|
||||
1. **Items with IDs**: Encode model_id into the item ID itself (e.g., ``encitem_...``)
|
||||
2. **Items without IDs** (Codex): Wrap the encrypted_content with model_id metadata
|
||||
(e.g., ``litellm_enc:{base64_metadata};{original_encrypted_content}``)
|
||||
|
||||
The encoded model_id is decoded on the next request so the router can pin to the correct
|
||||
deployment without any cache lookup.
|
||||
|
||||
Response post-processing (encoding) is handled by
|
||||
``ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response`` which is
|
||||
called inside ``_update_responses_api_response_id_with_model_id`` in ``responses/utils.py``.
|
||||
|
||||
Request pre-processing (ID/content restoration before forwarding to upstream) is handled by
|
||||
``ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input`` which is called
|
||||
in ``get_optional_params_responses_api``.
|
||||
|
||||
This pre-call check is responsible only for the routing decision: it reads the encoded
|
||||
``model_id`` from either item IDs or wrapped encrypted_content and pins the request to
|
||||
the matching deployment.
|
||||
|
||||
Safe to enable globally:
|
||||
- Only activates when encoded markers appear in the request ``input``.
|
||||
- No effect on embedding models, chat completions, or first-time requests.
|
||||
- No quota reduction -- first requests are fully load balanced.
|
||||
- No cache required.
|
||||
"""
|
||||
|
||||
from typing import Any, List, Optional, cast
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
class EncryptedContentAffinityCheck(CustomLogger):
|
||||
"""
|
||||
Routes follow-up Responses API requests to the deployment that produced
|
||||
the encrypted output items they reference.
|
||||
|
||||
The ``model_id`` is decoded directly from the litellm-encoded item IDs –
|
||||
no caching or TTL management needed.
|
||||
|
||||
Wired via ``Router(optional_pre_call_checks=["encrypted_content_affinity"])``.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _extract_model_id_from_input(request_input: Any) -> Optional[str]:
|
||||
"""
|
||||
Scan ``input`` items for litellm-encoded encrypted-content markers and
|
||||
return the ``model_id`` embedded in the first one found.
|
||||
|
||||
Checks both:
|
||||
1. Encoded item IDs (encitem_...) - for clients that send IDs
|
||||
2. Wrapped encrypted_content (litellm_enc:...) - for clients like Codex that don't send IDs
|
||||
|
||||
``input`` can be:
|
||||
- a plain string -> no encoded markers
|
||||
- a list of items -> check each item's ``id`` and ``encrypted_content`` fields
|
||||
"""
|
||||
if not isinstance(request_input, list):
|
||||
return None
|
||||
|
||||
for item in request_input:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
|
||||
# First, try to decode from item ID (if present)
|
||||
item_id = item.get("id")
|
||||
if item_id and isinstance(item_id, str):
|
||||
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(item_id)
|
||||
if decoded:
|
||||
return decoded.get("model_id")
|
||||
|
||||
# If no encoded ID, check if encrypted_content itself is wrapped
|
||||
encrypted_content = item.get("encrypted_content")
|
||||
if encrypted_content and isinstance(encrypted_content, str):
|
||||
(
|
||||
model_id,
|
||||
_,
|
||||
) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(
|
||||
encrypted_content
|
||||
)
|
||||
if model_id:
|
||||
return model_id
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _find_deployment_by_model_id(
|
||||
healthy_deployments: List[dict], model_id: str
|
||||
) -> Optional[dict]:
|
||||
for deployment in healthy_deployments:
|
||||
model_info = deployment.get("model_info")
|
||||
if not isinstance(model_info, dict):
|
||||
continue
|
||||
deployment_model_id = model_info.get("id")
|
||||
if deployment_model_id is not None and str(deployment_model_id) == str(
|
||||
model_id
|
||||
):
|
||||
return deployment
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Request routing (pre-call filter)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def async_filter_deployments(
|
||||
self,
|
||||
model: str,
|
||||
healthy_deployments: List,
|
||||
messages: Optional[List[AllMessageValues]],
|
||||
request_kwargs: Optional[dict] = None,
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
) -> List[dict]:
|
||||
"""
|
||||
If the request ``input`` contains litellm-encoded item IDs, decode the
|
||||
embedded ``model_id`` and pin the request to that deployment.
|
||||
"""
|
||||
request_kwargs = request_kwargs or {}
|
||||
typed_healthy_deployments = cast(List[dict], healthy_deployments)
|
||||
|
||||
# Signal to the response post-processor that encrypted item IDs should be
|
||||
# encoded in the output of this request.
|
||||
litellm_metadata = request_kwargs.setdefault("litellm_metadata", {})
|
||||
litellm_metadata["encrypted_content_affinity_enabled"] = True
|
||||
|
||||
request_input = request_kwargs.get("input")
|
||||
model_id = self._extract_model_id_from_input(request_input)
|
||||
if not model_id:
|
||||
return typed_healthy_deployments
|
||||
|
||||
verbose_router_logger.debug(
|
||||
"EncryptedContentAffinityCheck: decoded model_id=%s from input item IDs",
|
||||
model_id,
|
||||
)
|
||||
|
||||
deployment = self._find_deployment_by_model_id(
|
||||
healthy_deployments=typed_healthy_deployments,
|
||||
model_id=model_id,
|
||||
)
|
||||
if deployment is not None:
|
||||
verbose_router_logger.debug(
|
||||
"EncryptedContentAffinityCheck: pinning -> deployment=%s",
|
||||
model_id,
|
||||
)
|
||||
request_kwargs["_encrypted_content_affinity_pinned"] = True
|
||||
return [deployment]
|
||||
|
||||
verbose_router_logger.error(
|
||||
"EncryptedContentAffinityCheck: decoded deployment=%s not found in healthy_deployments",
|
||||
model_id,
|
||||
)
|
||||
return typed_healthy_deployments
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue