Merge branch 'BerriAI:main' into main
|
|
@ -676,18 +676,16 @@ jobs:
|
|||
working_directory: ~/project
|
||||
steps:
|
||||
- checkout
|
||||
- run:
|
||||
name: Install PostgreSQL
|
||||
command: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install postgresql postgresql-contrib
|
||||
echo 'export PATH=/usr/lib/postgresql/*/bin:$PATH' >> $BASH_ENV
|
||||
- setup_google_dns
|
||||
- run:
|
||||
name: Show git commit hash
|
||||
command: |
|
||||
echo "Git commit hash: $CIRCLE_SHA1"
|
||||
|
||||
- run:
|
||||
name: Install PostgreSQL
|
||||
command: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y postgresql-14 postgresql-contrib-14
|
||||
- restore_cache:
|
||||
keys:
|
||||
- v1-dependencies-{{ checksum ".circleci/requirements.txt" }}
|
||||
|
|
@ -2375,6 +2373,25 @@ jobs:
|
|||
pip install "pytest-mock==3.12.0"
|
||||
pip install "pytest-asyncio==0.21.1"
|
||||
pip install "assemblyai==0.37.0"
|
||||
- run:
|
||||
name: Install dockerize
|
||||
command: |
|
||||
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
rm dockerize-linux-amd64-v0.6.1.tar.gz
|
||||
- run:
|
||||
name: Start PostgreSQL Database
|
||||
command: |
|
||||
docker run -d \
|
||||
--name postgres-db \
|
||||
-e POSTGRES_USER=postgres \
|
||||
-e POSTGRES_PASSWORD=postgres \
|
||||
-e POSTGRES_DB=circle_test \
|
||||
-p 5432:5432 \
|
||||
postgres:14
|
||||
- run:
|
||||
name: Wait for PostgreSQL to be ready
|
||||
command: dockerize -wait tcp://localhost:5432 -timeout 1m
|
||||
- run:
|
||||
name: Build Docker image
|
||||
command: docker build -t my-app:latest -f ./docker/Dockerfile.database .
|
||||
|
|
@ -2385,10 +2402,11 @@ jobs:
|
|||
command: |
|
||||
docker run -d \
|
||||
-p 4000:4000 \
|
||||
-e DATABASE_URL=$CLEAN_STORE_MODEL_IN_DB_DATABASE_URL \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e STORE_MODEL_IN_DB="True" \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/store_model_db_config.yaml:/app/config.yaml \
|
||||
my-app:latest \
|
||||
|
|
@ -2418,7 +2436,16 @@ jobs:
|
|||
python -m pytest -vv tests/store_model_in_db_tests -x --junitxml=test-results/junit.xml --durations=5
|
||||
no_output_timeout:
|
||||
120m
|
||||
# Clean up first container
|
||||
- run:
|
||||
name: Stop and remove containers
|
||||
command: |
|
||||
docker stop my-app || true
|
||||
docker rm my-app || true
|
||||
docker stop postgres-db || true
|
||||
docker rm postgres-db || true
|
||||
when: always
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
|
||||
proxy_build_from_pip_tests:
|
||||
# Change from docker to machine executor
|
||||
|
|
|
|||
2
.gitignore
vendored
|
|
@ -97,3 +97,5 @@ litellm_config.yaml
|
|||
.vscode/launch.json
|
||||
litellm/proxy/to_delete_loadtest_work/*
|
||||
update_model_cost_map.py
|
||||
tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
|
||||
litellm/proxy/_experimental/out/guardrails/index.html
|
||||
|
|
|
|||
|
|
@ -68,8 +68,11 @@ run_grype_scans() {
|
|||
|
||||
# Allowlist of CVEs to be ignored in failure threshold/reporting
|
||||
# - CVE-2025-8869: Not applicable on Python >=3.13 (PEP 706 implemented); pip fallback unused; no OS-level fix
|
||||
# - GHSA-4xh5-x5gv-qwph: GitHub Security Advisory alias for CVE-2025-8869
|
||||
ALLOWED_CVES=(
|
||||
"CVE-2025-8869"
|
||||
"GHSA-4xh5-x5gv-qwph"
|
||||
"CVE-2025-8291" # no fix available as of Oct 11, 2025
|
||||
)
|
||||
|
||||
# Build JSON array of allowlisted CVE IDs for jq
|
||||
|
|
@ -77,6 +80,26 @@ run_grype_scans() {
|
|||
|
||||
echo "Checking for vulnerabilities with CVSS score >= 4.0..."
|
||||
echo "Allowlisted CVEs (ignored in threshold): ${ALLOWED_CVES[*]}"
|
||||
echo ""
|
||||
|
||||
# Show all high-severity vulnerabilities for transparency
|
||||
TOTAL_HIGH_SEVERITY=$(grype litellm:latest -o json | jq -r '
|
||||
.matches[]
|
||||
| select(.vulnerability.cvss[]?.metrics.baseScore >= 4.0)
|
||||
| .vulnerability.id' | wc -l)
|
||||
|
||||
if [ "$TOTAL_HIGH_SEVERITY" -gt 0 ]; then
|
||||
echo "Total vulnerabilities found with CVSS >= 4.0: $TOTAL_HIGH_SEVERITY"
|
||||
echo ""
|
||||
echo "All high-severity vulnerabilities (including allowlisted):"
|
||||
grype litellm:latest -o json | jq --argjson allow "$ALLOWED_IDS_JSON" -r '
|
||||
["Package", "Version", "Vulnerability ID", "CVSS Score", "Allowlisted"],
|
||||
(.matches[]
|
||||
| select(.vulnerability.cvss[]?.metrics.baseScore >= 4.0)
|
||||
| [.artifact.name, .artifact.version, .vulnerability.id, .vulnerability.cvss[0].metrics.baseScore, (if (.vulnerability.id as $id | $allow | index($id)) then "YES" else "NO" end)])
|
||||
| @tsv' | column -t -s $'\t'
|
||||
echo ""
|
||||
fi
|
||||
|
||||
HIGH_SEVERITY_COUNT=$(grype litellm:latest -o json | jq --argjson allow "$ALLOWED_IDS_JSON" -r '
|
||||
.matches[]
|
||||
|
|
@ -85,8 +108,17 @@ run_grype_scans() {
|
|||
| .vulnerability.id' | wc -l)
|
||||
|
||||
if [ "$HIGH_SEVERITY_COUNT" -gt 0 ]; then
|
||||
echo "ERROR: Found $HIGH_SEVERITY_COUNT vulnerabilities with CVSS score >= 4.0 in litellm:latest"
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "ERROR: Security Scan Failed"
|
||||
echo "=========================================="
|
||||
echo "Found $HIGH_SEVERITY_COUNT non-allowlisted vulnerabilities with CVSS score >= 4.0 in litellm:latest"
|
||||
echo ""
|
||||
echo "These vulnerabilities are NOT in the allowlist and must be addressed."
|
||||
echo "Current allowlisted CVEs: ${ALLOWED_CVES[*]}"
|
||||
echo ""
|
||||
echo "Detailed vulnerability report:"
|
||||
echo ""
|
||||
grype litellm:latest -o json | jq --argjson allow "$ALLOWED_IDS_JSON" -r '
|
||||
["Package", "Version", "Vulnerability ID", "CVSS Score", "Severity", "Fix Version", "Description"],
|
||||
(.matches[]
|
||||
|
|
@ -94,6 +126,19 @@ run_grype_scans() {
|
|||
| select((.vulnerability.id as $id | $allow | index($id) | not))
|
||||
| [.artifact.name, .artifact.version, .vulnerability.id, .vulnerability.cvss[0].metrics.baseScore, .vulnerability.severity, (.vulnerability.fix.versions[0] // "No fix available"), .vulnerability.description])
|
||||
| @tsv' | column -t -s $'\t'
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo "Action Required:"
|
||||
echo "=========================================="
|
||||
echo "1. If a fix is available, update the package to the fixed version"
|
||||
echo "2. If the vulnerability is not applicable or has no fix:"
|
||||
echo " - Add the CVE/GHSA ID to ALLOWED_CVES array in ci_cd/security_scans.sh"
|
||||
echo " - Add a comment explaining why it's safe to ignore"
|
||||
echo ""
|
||||
echo "Note: Some vulnerabilities may have multiple IDs (CVE-XXXX and GHSA-XXXX)."
|
||||
echo "Add all relevant IDs to the allowlist if they refer to the same issue."
|
||||
echo "=========================================="
|
||||
echo ""
|
||||
exit 1
|
||||
else
|
||||
echo "No high-severity vulnerabilities (CVSS >= 4.0) found in litellm:latest"
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ services:
|
|||
#########################################
|
||||
## Uncomment these lines to start proxy with a config.yaml file ##
|
||||
# volumes:
|
||||
# - ./config.yaml:/app/config.yaml <<- this is missing in the docker-compose file currently
|
||||
# - ./config.yaml:/app/config.yaml
|
||||
# command:
|
||||
# - "--config=/app/config.yaml"
|
||||
##############################################
|
||||
|
|
|
|||
|
|
@ -16,19 +16,17 @@ model_list:
|
|||
api_key: "test"
|
||||
```
|
||||
|
||||
### 1 Instance LiteLLM Proxy
|
||||
### 2 Instance LiteLLM Proxy
|
||||
|
||||
In these tests the baseline latency characteristics are measured against a fake-openai-endpoint.
|
||||
|
||||
#### Performance Metrics
|
||||
|
||||
| Metric | Value |
|
||||
|--------|-------|
|
||||
| **Requests per Second (RPS)** | 475 |
|
||||
| **End-to-End Latency P50 (ms)** | 100 |
|
||||
| **LiteLLM Overhead P50 (ms)** | 3 |
|
||||
| **LiteLLM Overhead P90 (ms)** | 17 |
|
||||
| **LiteLLM Overhead P99 (ms)** | 31 |
|
||||
| **Type** | **Name** | **Median (ms)** | **95%ile (ms)** | **99%ile (ms)** | **Average (ms)** | **Current RPS** |
|
||||
| --- | --- | --- | --- | --- | --- | --- |
|
||||
| POST | /chat/completions | 200 | 630 | 1200 | 262.46 | 1035.7 |
|
||||
| Custom | LiteLLM Overhead Duration (ms) | 12 | 29 | 43 | 14.74 | 1035.7 |
|
||||
| | Aggregated | 100 | 430 | 930 | 138.6 | 2071.4 |
|
||||
|
||||
<!-- <Image img={require('../img/1_instance_proxy.png')} /> -->
|
||||
|
||||
|
|
@ -36,28 +34,32 @@ In these tests the baseline latency characteristics are measured against a fake-
|
|||
|
||||
<Image img={require('../img/instances_vs_rps.png')} /> -->
|
||||
|
||||
|
||||
### 4 Instances
|
||||
|
||||
| **Type** | **Name** | **Median (ms)** | **95%ile (ms)** | **99%ile (ms)** | **Average (ms)** | **Current RPS** |
|
||||
| --- | --- | --- | --- | --- | --- | --- |
|
||||
| POST | /chat/completions | 100 | 150 | 240 | 111.73 | 1170 |
|
||||
| Custom | LiteLLM Overhead Duration (ms) | 2 | 8 | 13 | 3.32 | 1170 |
|
||||
| | Aggregated | 77 | 130 | 180 | 57.53 | 2340 |
|
||||
|
||||
#### Key Findings
|
||||
- Single instance: 475 RPS @ 100ms median latency
|
||||
- LiteLLM adds 3ms P50 overhead, 17ms P90 overhead, 31ms P99 overhead
|
||||
- 2 LiteLLM instances: 950 RPS @ 100ms latency
|
||||
- 4 LiteLLM instances: 1900 RPS @ 100ms latency
|
||||
|
||||
### 2 Instances
|
||||
|
||||
**Adding 1 instance, will double the RPS and maintain the `100ms-110ms` median latency.**
|
||||
|
||||
| Metric | Litellm Proxy (2 Instances) |
|
||||
|--------|------------------------|
|
||||
| Median Latency (ms) | 100 |
|
||||
| RPS | 950 |
|
||||
|
||||
- Doubling from 2 to 4 LiteLLM instances halves median latency: 200 ms → 100 ms.
|
||||
- High-percentile latencies drop significantly: P95 630 ms → 150 ms, P99 1,200 ms → 240 ms.
|
||||
- Setting workers equal to CPU count gives optimal performance.
|
||||
|
||||
## Machine Spec used for testing
|
||||
|
||||
Each machine deploying LiteLLM had the following specs:
|
||||
|
||||
- 2 CPU
|
||||
- 4GB RAM
|
||||
- 4 CPU
|
||||
- 8GB RAM
|
||||
|
||||
|
||||
## Locust Settings
|
||||
|
||||
- 1000 Users
|
||||
- 500 user Ramp Up
|
||||
|
||||
## How to measure LiteLLM Overhead
|
||||
|
||||
|
|
@ -137,10 +139,3 @@ Using LangSmith has **no impact on latency, RPS compared to Basic Litellm Proxy*
|
|||
|--------|------------------------|---------------------|
|
||||
| RPS | 1133.2 | 1135 |
|
||||
| Median Latency (ms) | 140 | 132 |
|
||||
|
||||
|
||||
|
||||
## Locust Settings
|
||||
|
||||
- 2500 Users
|
||||
- 100 user Ramp Up
|
||||
|
|
|
|||
|
|
@ -180,11 +180,11 @@ def completion(
|
|||
|
||||
- `function`: *object* - Required.
|
||||
|
||||
- `tool_choice`: *string or object (optional)* - Controls which (if any) function is called by the model. none means the model will not call a function and instead generates a message. auto means the model can pick between generating a message or calling a function. Specifying a particular function via `{"type: "function", "function": {"name": "my_function"}}` forces the model to call that function.
|
||||
- `tool_choice`: *string or object (optional)* - Controls which (if any) function is called by the model. none means the model will not call a function and instead generates a message. auto means the model can pick between generating a message or calling a function. Specifying a particular function via `{"type": "function", "function": {"name": "my_function"}}` forces the model to call that function.
|
||||
|
||||
- `none` is the default when no functions are present. `auto` is the default if functions are present.
|
||||
|
||||
- `parallel_tool_calls`: *boolean (optional)* - Whether to enable parallel function calling during tool use.. OpenAI default is true.
|
||||
- `parallel_tool_calls`: *boolean (optional)* - Whether to enable parallel function calling during tool use. OpenAI default is true.
|
||||
|
||||
- `frequency_penalty`: *number or null (optional)* - It is used to penalize new tokens based on their frequency in the text so far.
|
||||
|
||||
|
|
|
|||
|
|
@ -246,8 +246,203 @@ litellm_settings:
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## MCP Tool Filtering
|
||||
## Converting OpenAPI Specs to MCP Servers
|
||||
|
||||
LiteLLM can automatically convert OpenAPI specifications into MCP servers, allowing you to expose any REST API as MCP tools. This is useful when you have existing APIs with OpenAPI/Swagger documentation and want to make them available as MCP tools.
|
||||
|
||||
### Benefits
|
||||
|
||||
- **Rapid Integration**: Convert existing APIs to MCP tools without writing custom MCP server code
|
||||
- **Automatic Tool Generation**: LiteLLM automatically generates MCP tools from your OpenAPI spec
|
||||
- **Unified Interface**: Use the same MCP interface for both native MCP servers and OpenAPI-based APIs
|
||||
- **Easy Testing**: Test and iterate on API integrations quickly
|
||||
|
||||
### Configuration
|
||||
|
||||
Add your OpenAPI-based MCP server to your `config.yaml`:
|
||||
|
||||
```yaml title="config.yaml - OpenAPI to MCP" showLineNumbers
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: sk-xxxxxxx
|
||||
|
||||
mcp_servers:
|
||||
# OpenAPI Spec Example - Petstore API
|
||||
petstore_mcp:
|
||||
url: "https://petstore.swagger.io/v2"
|
||||
spec_path: "/path/to/openapi.json"
|
||||
auth_type: "none"
|
||||
|
||||
# OpenAPI Spec with API Key Authentication
|
||||
my_api_mcp:
|
||||
url: "http://0.0.0.0:8090"
|
||||
spec_path: "/path/to/openapi.json"
|
||||
auth_type: "api_key"
|
||||
auth_value: "your-api-key-here"
|
||||
|
||||
# OpenAPI Spec with Bearer Token
|
||||
secured_api_mcp:
|
||||
url: "https://api.example.com"
|
||||
spec_path: "/path/to/openapi.json"
|
||||
auth_type: "bearer_token"
|
||||
auth_value: "your-bearer-token"
|
||||
```
|
||||
|
||||
### Configuration Parameters
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `url` | Yes | The base URL of your API endpoint |
|
||||
| `spec_path` | Yes | Path or URL to your OpenAPI specification file (JSON or YAML) |
|
||||
| `auth_type` | No | Authentication type: `none`, `api_key`, `bearer_token`, `basic`, `authorization` |
|
||||
| `auth_value` | No | Authentication value (required if `auth_type` is set) |
|
||||
| `description` | No | Optional description for the MCP server |
|
||||
| `allowed_tools` | No | List of specific tools to allow (see [MCP Tool Filtering](#mcp-tool-filtering)) |
|
||||
| `disallowed_tools` | No | List of specific tools to block (see [MCP Tool Filtering](#mcp-tool-filtering)) |
|
||||
|
||||
### Usage Example
|
||||
|
||||
Once configured, you can use the OpenAPI-based MCP server just like any other MCP server:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="fastmcp" label="Python FastMCP">
|
||||
|
||||
```python title="Using OpenAPI-based MCP Server" showLineNumbers
|
||||
from fastmcp import Client
|
||||
import asyncio
|
||||
|
||||
# Standard MCP configuration
|
||||
config = {
|
||||
"mcpServers": {
|
||||
"petstore": {
|
||||
"url": "http://localhost:4000/petstore_mcp/mcp",
|
||||
"headers": {
|
||||
"x-litellm-api-key": "Bearer sk-1234"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
# Create a client that connects to the server
|
||||
client = Client(config)
|
||||
|
||||
async def main():
|
||||
async with client:
|
||||
# List available tools generated from OpenAPI spec
|
||||
tools = await client.list_tools()
|
||||
print(f"Available tools: {[tool.name for tool in tools]}")
|
||||
|
||||
# Example: Get a pet by ID (from Petstore API)
|
||||
response = await client.call_tool(
|
||||
name="getpetbyid",
|
||||
arguments={"petId": "1"}
|
||||
)
|
||||
print(f"Response:\n{response}\n")
|
||||
|
||||
# Example: Find pets by status
|
||||
response = await client.call_tool(
|
||||
name="findpetsbystatus",
|
||||
arguments={"status": "available"}
|
||||
)
|
||||
print(f"Response:\n{response}\n")
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="cursor" label="Cursor IDE">
|
||||
|
||||
```json title="Cursor MCP Configuration for OpenAPI Server" showLineNumbers
|
||||
{
|
||||
"mcpServers": {
|
||||
"Petstore": {
|
||||
"url": "http://localhost:4000/petstore_mcp/mcp",
|
||||
"headers": {
|
||||
"x-litellm-api-key": "Bearer $LITELLM_API_KEY"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="openai" label="OpenAI Responses API">
|
||||
|
||||
```bash title="Using OpenAPI MCP Server with OpenAI" showLineNumbers
|
||||
curl --location 'https://api.openai.com/v1/responses' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header "Authorization: Bearer $OPENAI_API_KEY" \
|
||||
--data '{
|
||||
"model": "gpt-4o",
|
||||
"tools": [
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_label": "petstore",
|
||||
"server_url": "http://localhost:4000/petstore_mcp/mcp",
|
||||
"require_approval": "never",
|
||||
"headers": {
|
||||
"x-litellm-api-key": "Bearer YOUR_LITELLM_API_KEY"
|
||||
}
|
||||
}
|
||||
],
|
||||
"input": "Find all available pets in the petstore",
|
||||
"tool_choice": "required"
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### How It Works
|
||||
|
||||
1. **Spec Loading**: LiteLLM loads your OpenAPI specification from the provided `spec_path`
|
||||
2. **Tool Generation**: Each API endpoint in the spec becomes an MCP tool
|
||||
3. **Parameter Mapping**: OpenAPI parameters are automatically mapped to MCP tool parameters
|
||||
4. **Request Handling**: When a tool is called, LiteLLM converts the MCP request to the appropriate HTTP request
|
||||
5. **Response Translation**: API responses are converted back to MCP format
|
||||
|
||||
### OpenAPI Spec Requirements
|
||||
|
||||
Your OpenAPI specification should follow standard OpenAPI/Swagger conventions:
|
||||
- **Supported versions**: OpenAPI 3.0.x, OpenAPI 3.1.x, Swagger 2.0
|
||||
- **Required fields**: `paths`, `info` sections should be properly defined
|
||||
- **Operation IDs**: Each operation should have a unique `operationId` (this becomes the tool name)
|
||||
- **Parameters**: Request parameters should be properly documented with types and descriptions
|
||||
|
||||
### Example OpenAPI Spec Structure
|
||||
|
||||
```yaml title="sample-openapi.yaml" showLineNumbers
|
||||
openapi: 3.0.0
|
||||
info:
|
||||
title: My API
|
||||
version: 1.0.0
|
||||
paths:
|
||||
/pets/{petId}:
|
||||
get:
|
||||
operationId: getPetById
|
||||
summary: Get a pet by ID
|
||||
parameters:
|
||||
- name: petId
|
||||
in: path
|
||||
required: true
|
||||
schema:
|
||||
type: integer
|
||||
responses:
|
||||
'200':
|
||||
description: Successful response
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
type: object
|
||||
```
|
||||
|
||||
## Allow/Disallow MCP Tools
|
||||
|
||||
Control which tools are available from your MCP servers. You can either allow only specific tools or block dangerous ones.
|
||||
|
||||
<Tabs>
|
||||
|
|
@ -306,6 +501,119 @@ mcp_servers:
|
|||
- If you specify both `allowed_tools` and `disallowed_tools`, the allowed list takes priority
|
||||
- Tool names are case-sensitive
|
||||
|
||||
---
|
||||
|
||||
## Allow/Disallow MCP Tool Parameters
|
||||
|
||||
Control which parameters are allowed for specific MCP tools using the `allowed_params` configuration. This provides fine-grained control over tool usage by restricting the parameters that can be passed to each tool.
|
||||
|
||||
### Configuration
|
||||
|
||||
`allowed_params` is a dictionary that maps tool names to lists of allowed parameter names. When configured, only the specified parameters will be accepted for that tool - any other parameters will be rejected with a 403 error.
|
||||
|
||||
```yaml title="config.yaml with allowed_params" showLineNumbers
|
||||
mcp_servers:
|
||||
deepwiki_mcp:
|
||||
url: https://mcp.deepwiki.com/mcp
|
||||
transport: "http"
|
||||
auth_type: "none"
|
||||
allowed_params:
|
||||
# Tool name: list of allowed parameters
|
||||
read_wiki_contents: ["status"]
|
||||
|
||||
my_api_mcp:
|
||||
url: "https://my-api-server.com"
|
||||
auth_type: "api_key"
|
||||
auth_value: "my-key"
|
||||
allowed_params:
|
||||
# Using unprefixed tool name
|
||||
getpetbyid: ["status"]
|
||||
# Using prefixed tool name (both formats work)
|
||||
my_api_mcp-findpetsbystatus: ["status", "limit"]
|
||||
# Another tool with multiple allowed params
|
||||
create_issue: ["title", "body", "labels"]
|
||||
```
|
||||
|
||||
### How It Works
|
||||
|
||||
1. **Tool-specific filtering**: Each tool can have its own list of allowed parameters
|
||||
2. **Flexible naming**: Tool names can be specified with or without the server prefix (e.g., both `"getpetbyid"` and `"my_api_mcp-getpetbyid"` work)
|
||||
3. **Whitelist approach**: Only parameters in the allowed list are permitted
|
||||
4. **Unlisted tools**: If `allowed_params` is not set, all parameters are allowed
|
||||
5. **Error handling**: Requests with disallowed parameters receive a 403 error with details about which parameters are allowed
|
||||
|
||||
### Example Request Behavior
|
||||
|
||||
With the configuration above, here's how requests would be handled:
|
||||
|
||||
**✅ Allowed Request:**
|
||||
```json
|
||||
{
|
||||
"tool": "read_wiki_contents",
|
||||
"arguments": {
|
||||
"status": "active"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**❌ Rejected Request:**
|
||||
```json
|
||||
{
|
||||
"tool": "read_wiki_contents",
|
||||
"arguments": {
|
||||
"status": "active",
|
||||
"limit": 10 // This parameter is not allowed
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Error Response:**
|
||||
```json
|
||||
{
|
||||
"error": "Parameters ['limit'] are not allowed for tool read_wiki_contents. Allowed parameters: ['status']. Contact proxy admin to allow these parameters."
|
||||
}
|
||||
```
|
||||
|
||||
### Use Cases
|
||||
|
||||
- **Security**: Prevent users from accessing sensitive parameters or dangerous operations
|
||||
- **Cost control**: Restrict expensive parameters (e.g., limiting result counts)
|
||||
- **Compliance**: Enforce parameter usage policies for regulatory requirements
|
||||
- **Staged rollouts**: Gradually enable parameters as tools are tested
|
||||
- **Multi-tenant isolation**: Different parameter access for different user groups
|
||||
|
||||
### Combining with Tool Filtering
|
||||
|
||||
`allowed_params` works alongside `allowed_tools` and `disallowed_tools` for complete control:
|
||||
|
||||
```yaml title="Combined filtering example" showLineNumbers
|
||||
mcp_servers:
|
||||
github_mcp:
|
||||
url: "https://api.githubcopilot.com/mcp"
|
||||
auth_type: oauth2
|
||||
authorization_url: https://github.com/login/oauth/authorize
|
||||
token_url: https://github.com/login/oauth/access_token
|
||||
client_id: os.environ/GITHUB_OAUTH_CLIENT_ID
|
||||
client_secret: os.environ/GITHUB_OAUTH_CLIENT_SECRET
|
||||
scopes: ["public_repo", "user:email"]
|
||||
# Only allow specific tools
|
||||
allowed_tools: ["create_issue", "list_issues", "search_issues"]
|
||||
# Block dangerous operations
|
||||
disallowed_tools: ["delete_repo"]
|
||||
# Restrict parameters per tool
|
||||
allowed_params:
|
||||
create_issue: ["title", "body", "labels"]
|
||||
list_issues: ["state", "sort", "perPage"]
|
||||
search_issues: ["query", "sort", "order", "perPage"]
|
||||
```
|
||||
|
||||
This configuration ensures that:
|
||||
1. Only the three listed tools are available
|
||||
2. The `delete_repo` tool is explicitly blocked
|
||||
3. Each tool can only use its specified parameters
|
||||
|
||||
---
|
||||
|
||||
## MCP Server Access Control
|
||||
|
||||
LiteLLM Proxy provides two methods for controlling access to specific MCP servers:
|
||||
|
|
@ -896,6 +1204,8 @@ mcp_servers:
|
|||
scopes: ["public_repo", "user:email"]
|
||||
```
|
||||
|
||||
[**See Claude Code Tutorial**](./tutorials/claude_responses_api#connecting-mcp-servers)
|
||||
|
||||
## Using your MCP with client side credentials
|
||||
|
||||
Use this if you want to pass a client side authentication token to LiteLLM to then pass to your MCP to auth to your MCP.
|
||||
|
|
|
|||
|
|
@ -55,6 +55,26 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
|
|||
}'
|
||||
```
|
||||
|
||||
### Team-Based Logging
|
||||
|
||||
Configure different PostHog credentials per team using the team callback settings:
|
||||
|
||||
```bash
|
||||
curl -X POST 'http://localhost:4000/team/{team_id}/callback' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"callback_name": "posthog",
|
||||
"callback_type": "success",
|
||||
"callback_vars": {
|
||||
"posthog_api_key": "ph_team_specific_key",
|
||||
"posthog_api_url": "https://custom.posthog.com"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
Now all requests from that team will be logged to their specific PostHog project.
|
||||
|
||||
## Usage with LiteLLM Python SDK
|
||||
|
||||
### Quick Start
|
||||
|
|
@ -142,6 +162,31 @@ response = client.chat.completions.create(
|
|||
)
|
||||
```
|
||||
|
||||
#### Per-Request Credentials
|
||||
|
||||
You can override PostHog credentials on a per-request basis:
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
litellm.success_callback = ["posthog"]
|
||||
|
||||
# Use custom PostHog credentials for this specific request
|
||||
response = litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[
|
||||
{"role": "user", "content": "Hello world"}
|
||||
],
|
||||
posthog_api_key="ph_custom_project_key",
|
||||
posthog_api_url="https://custom.posthog.com"
|
||||
)
|
||||
```
|
||||
|
||||
This is useful when you need to:
|
||||
- Log different teams/projects to separate PostHog instances
|
||||
- Use different PostHog projects for staging vs production
|
||||
- Route logs based on customer or tenant
|
||||
|
||||
#### Disable Logging for Specific Calls
|
||||
|
||||
Use the `no-log` flag to prevent logging for specific calls:
|
||||
|
|
|
|||
|
|
@ -6,18 +6,27 @@ LiteLLM supports the following models for OCI on-demand GenAI API.
|
|||
|
||||
Check the [OCI Models List](https://docs.oracle.com/en-us/iaas/Content/generative-ai/pretrained-models.htm) to see if the model is available for your region.
|
||||
|
||||
## Supported Models
|
||||
|
||||
### Meta Llama Models
|
||||
- `meta.llama-4-maverick-17b-128e-instruct-fp8`
|
||||
- `meta.llama-4-scout-17b-16e-instruct`
|
||||
- `meta.llama-3.3-70b-instruct`
|
||||
- `meta.llama-3.2-90b-vision-instruct`
|
||||
- `meta.llama-3.1-405b-instruct`
|
||||
|
||||
### xAI Grok Models
|
||||
- `xai.grok-4`
|
||||
- `xai.grok-3`
|
||||
- `xai.grok-3-fast`
|
||||
- `xai.grok-3-mini`
|
||||
- `xai.grok-3-mini-fast`
|
||||
|
||||
### Cohere Models
|
||||
- `cohere.command-latest`
|
||||
- `cohere.command-a-03-2025`
|
||||
- `cohere.command-plus-latest`
|
||||
|
||||
## Authentication
|
||||
|
||||
LiteLLM uses OCI signing key authentication. Follow the [official Oracle tutorial](https://docs.oracle.com/en-us/iaas/Content/API/Concepts/apisigningkey.htm) to create a signing key and obtain the following parameters:
|
||||
|
|
@ -83,3 +92,24 @@ response = completion(
|
|||
for chunk in response:
|
||||
print(chunk["choices"][0]["delta"]["content"]) # same as openai format
|
||||
```
|
||||
|
||||
## Usage Examples by Model Type
|
||||
|
||||
### Using Cohere Models
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
messages = [{"role": "user", "content": "Explain quantum computing"}]
|
||||
response = completion(
|
||||
model="oci/cohere.command-latest",
|
||||
messages=messages,
|
||||
oci_region="us-chicago-1",
|
||||
oci_user=<your_oci_user>,
|
||||
oci_fingerprint=<your_oci_fingerprint>,
|
||||
oci_tenancy=<your_oci_tenancy>,
|
||||
oci_key=<string_with_content_of_oci_key>,
|
||||
oci_compartment_id=<oci_compartment_id>,
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
|
@ -339,6 +339,72 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
|||
| fine tuned `gpt-3.5-turbo-1106` | `response = completion(model="ft:gpt-3.5-turbo-1106", messages=messages)` |
|
||||
| fine tuned `gpt-3.5-turbo-0613` | `response = completion(model="ft:gpt-3.5-turbo-0613", messages=messages)` |
|
||||
|
||||
## Getting Reasoning Content in `/chat/completions`
|
||||
|
||||
GPT-5 models return reasoning content when called via the Responses API. You can call these models via the `/chat/completions` endpoint by using the `openai/responses/` prefix.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
```python
|
||||
response = litellm.completion(
|
||||
model="openai/responses/gpt-5-mini", # tells litellm to call the model via the Responses API
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
reasoning_effort="low",
|
||||
)
|
||||
```
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
```bash
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-d '{
|
||||
"model": "openai/responses/gpt-5-mini",
|
||||
"messages": [{"role": "user", "content": "What is the capital of France?"}],
|
||||
"reasoning_effort": "low"
|
||||
}'
|
||||
```
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
Expected Response:
|
||||
```json
|
||||
{
|
||||
"id": "chatcmpl-6382a222-43c9-40c4-856b-22e105d88075",
|
||||
"created": 1760146746,
|
||||
"model": "gpt-5-mini",
|
||||
"object": "chat.completion",
|
||||
"system_fingerprint": null,
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"message": {
|
||||
"content": "Paris",
|
||||
"role": "assistant",
|
||||
"tool_calls": null,
|
||||
"function_call": null,
|
||||
"reasoning_content": "**Identifying the capital**\n\nThe user wants me to think of the capital of France and write it down. That's pretty straightforward: it's Paris. There aren't any safety issues to consider here. I think it would be best to keep it concise, so maybe just \"Paris\" would suffice. I feel confident that I should just stick to that without adding anything else. So, let's write it down!",
|
||||
"provider_specific_fields": null
|
||||
}
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"completion_tokens": 7,
|
||||
"prompt_tokens": 18,
|
||||
"total_tokens": 25,
|
||||
"completion_tokens_details": null,
|
||||
"prompt_tokens_details": {
|
||||
"audio_tokens": null,
|
||||
"cached_tokens": 0,
|
||||
"text_tokens": null,
|
||||
"image_tokens": null
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
```
|
||||
|
||||
## OpenAI Chat Completion to Responses API Bridge
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
## **Batch APIs**
|
||||
# Vertex Batch APIs
|
||||
|
||||
Just add the following Vertex env vars to your environment.
|
||||
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ import TabItem from '@theme/TabItem';
|
|||
| AI21 (Jamba) | `vertex_ai/jamba-*` | [Vertex AI - AI21 Models](https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/ai21) |
|
||||
| Qwen | `vertex_ai/qwen/*` | [Vertex AI - Qwen Models](https://cloud.google.com/vertex-ai/generative-ai/docs/maas/qwen) |
|
||||
| OpenAI (GPT-OSS) | `vertex_ai/openai/gpt-oss-*` | [Vertex AI - GPT-OSS Models](https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/) |
|
||||
| Model Garden | `vertex_ai/openai/{MODEL_ID}` or `vertex_ai/{MODEL_ID}` | [Vertex Model Garden](https://cloud.google.com/model-garden?hl=en) |
|
||||
|
||||
## Vertex AI - Anthropic (Claude)
|
||||
|
||||
|
|
@ -793,112 +792,3 @@ curl http://0.0.0.0:4000/v1/chat/completions \
|
|||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Model Garden
|
||||
|
||||
:::tip
|
||||
|
||||
All OpenAI compatible models from Vertex Model Garden are supported.
|
||||
|
||||
:::
|
||||
|
||||
#### Using Model Garden
|
||||
|
||||
**Almost all Vertex Model Garden models are OpenAI compatible.**
|
||||
|
||||
<Tabs>
|
||||
|
||||
<TabItem value="openai" label="OpenAI Compatible Models">
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `vertex_ai/openai/{MODEL_ID}` |
|
||||
| Vertex Documentation | [Model Garden LiteLLM Inference](https://github.com/GoogleCloudPlatform/generative-ai/blob/main/open-models/use-cases/model_garden_litellm_inference.ipynb), [Vertex Model Garden](https://cloud.google.com/model-garden?hl=en) |
|
||||
| Supported Operations | `/chat/completions`, `/embeddings` |
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/openai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="Proxy">
|
||||
|
||||
|
||||
**1. Add to config**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: llama3-1-8b-instruct
|
||||
litellm_params:
|
||||
model: vertex_ai/openai/5464397967697903616
|
||||
vertex_ai_project: "my-test-project"
|
||||
vertex_ai_location: "us-east-1"
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING at http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**3. Test it!**
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "llama3-1-8b-instruct", # 👈 the 'model_name' in config
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
}'
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="non-openai" label="Non-OpenAI Compatible Models">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
|
|
|||
229
docs/my-website/docs/providers/vertex_self_deployed.md
Normal file
|
|
@ -0,0 +1,229 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Vertex AI - Self Deployed Models
|
||||
|
||||
Deploy and use your own models on Vertex AI through Model Garden or custom endpoints.
|
||||
|
||||
## Model Garden
|
||||
|
||||
:::tip
|
||||
|
||||
All OpenAI compatible models from Vertex Model Garden are supported.
|
||||
|
||||
:::
|
||||
|
||||
### Using Model Garden
|
||||
|
||||
**Almost all Vertex Model Garden models are OpenAI compatible.**
|
||||
|
||||
<Tabs>
|
||||
|
||||
<TabItem value="openai" label="OpenAI Compatible Models">
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `vertex_ai/openai/{MODEL_ID}` |
|
||||
| Vertex Documentation | [Model Garden LiteLLM Inference](https://github.com/GoogleCloudPlatform/generative-ai/blob/main/open-models/use-cases/model_garden_litellm_inference.ipynb), [Vertex Model Garden](https://cloud.google.com/model-garden?hl=en) |
|
||||
| Supported Operations | `/chat/completions`, `/embeddings` |
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/openai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="Proxy">
|
||||
|
||||
|
||||
**1. Add to config**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: llama3-1-8b-instruct
|
||||
litellm_params:
|
||||
model: vertex_ai/openai/5464397967697903616
|
||||
vertex_ai_project: "my-test-project"
|
||||
vertex_ai_location: "us-east-1"
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING at http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**3. Test it!**
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "llama3-1-8b-instruct", # 👈 the 'model_name' in config
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
}'
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="non-openai" label="Non-OpenAI Compatible Models">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
## Gemma Models (Custom Endpoints)
|
||||
|
||||
Deploy Gemma models on custom Vertex AI prediction endpoints with OpenAI-compatible format.
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `vertex_ai/gemma/{MODEL_NAME}` |
|
||||
| Vertex Documentation | [Vertex AI Prediction](https://cloud.google.com/vertex-ai/docs/predictions/get-predictions) |
|
||||
| Required Parameter | `api_base` - Full prediction endpoint URL |
|
||||
|
||||
**Proxy Usage:**
|
||||
|
||||
**1. Add to config.yaml**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemma-model
|
||||
litellm_params:
|
||||
model: vertex_ai/gemma/gemma-3-12b-it-1222199011122
|
||||
api_base: https://ENDPOINT.us-central1-PROJECT.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict
|
||||
vertex_project: "my-project-id"
|
||||
vertex_location: "us-central1"
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
**3. Test it**
|
||||
|
||||
```bash
|
||||
curl http://0.0.0.0:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "gemma-model",
|
||||
"messages": [{"role": "user", "content": "What is machine learning?"}],
|
||||
"max_tokens": 100
|
||||
}'
|
||||
```
|
||||
|
||||
**SDK Usage:**
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/gemma/gemma-3-12b-it-1222199011122",
|
||||
messages=[{"role": "user", "content": "What is machine learning?"}],
|
||||
api_base="https://ENDPOINT.us-central1-PROJECT.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict",
|
||||
vertex_project="my-project-id",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
```
|
||||
|
||||
## MedGemma Models (Custom Endpoints)
|
||||
|
||||
Deploy MedGemma models on custom Vertex AI prediction endpoints with OpenAI-compatible format. MedGemma models use the same `vertex_ai/gemma/` route.
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `vertex_ai/gemma/{MODEL_NAME}` |
|
||||
| Vertex Documentation | [Vertex AI Prediction](https://cloud.google.com/vertex-ai/docs/predictions/get-predictions) |
|
||||
| Required Parameter | `api_base` - Full prediction endpoint URL |
|
||||
|
||||
**Proxy Usage:**
|
||||
|
||||
**1. Add to config.yaml**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: medgemma-model
|
||||
litellm_params:
|
||||
model: vertex_ai/gemma/medgemma-2b-v1
|
||||
api_base: https://ENDPOINT.us-central1-PROJECT.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict
|
||||
vertex_project: "my-project-id"
|
||||
vertex_location: "us-central1"
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
**3. Test it**
|
||||
|
||||
```bash
|
||||
curl http://0.0.0.0:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "medgemma-model",
|
||||
"messages": [{"role": "user", "content": "What are the symptoms of hypertension?"}],
|
||||
"max_tokens": 100
|
||||
}'
|
||||
```
|
||||
|
||||
**SDK Usage:**
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/gemma/medgemma-2b-v1",
|
||||
messages=[{"role": "user", "content": "What are the symptoms of hypertension?"}],
|
||||
api_base="https://ENDPOINT.us-central1-PROJECT.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict",
|
||||
vertex_project="my-project-id",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
```
|
||||
196
docs/my-website/docs/providers/wandb_inference.md
Normal file
|
|
@ -0,0 +1,196 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Weights & Biases Inference
|
||||
https://weave-docs.wandb.ai/quickstart-inference
|
||||
|
||||
:::tip
|
||||
|
||||
Litellm provides support to all models from W&B Inference service. To use a model, set `model=wandb/<any-model-on-wandb-inference-dashboard>` as a prefix for litellm requests. The full list of supported models is provided at https://docs.wandb.ai/guides/inference/models/
|
||||
|
||||
:::
|
||||
|
||||
## API Key
|
||||
|
||||
You can get an API key for W&B Inference at - https://wandb.ai/authorize
|
||||
|
||||
```python
|
||||
import os
|
||||
# env variable
|
||||
os.environ['WANDB_API_KEY']
|
||||
```
|
||||
|
||||
## Sample Usage: Text Generation
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
os.environ['WANDB_API_KEY'] = "insert-your-wandb-api-key"
|
||||
response = completion(
|
||||
model="wandb/Qwen/Qwen3-235B-A22B-Instruct-2507",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What character was Wall-e in love with?",
|
||||
}
|
||||
],
|
||||
max_tokens=10,
|
||||
response_format={ "type": "json_object" },
|
||||
seed=123,
|
||||
temperature=0.6, # either set temperature or `top_p`
|
||||
top_p=0.01, # to get as deterministic results as possible
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
## Sample Usage - Streaming
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
os.environ['WANDB_API_KEY'] = ""
|
||||
response = completion(
|
||||
model="wandb/Qwen/Qwen3-235B-A22B-Instruct-2507",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What character was Wall-e in love with?",
|
||||
}
|
||||
],
|
||||
stream=True,
|
||||
max_tokens=10,
|
||||
response_format={ "type": "json_object" },
|
||||
seed=123,
|
||||
temperature=0.6, # either set temperature or `top_p`
|
||||
top_p=0.01, # to get as deterministic results as possible
|
||||
)
|
||||
|
||||
for chunk in response:
|
||||
print(chunk)
|
||||
```
|
||||
|
||||
:::tip
|
||||
|
||||
The above examples may not work if the model has been taken offline. Check the full list of available models at https://docs.wandb.ai/guides/inference/models/.
|
||||
|
||||
:::
|
||||
|
||||
## Usage with LiteLLM Proxy Server
|
||||
|
||||
Here's how to call a W&B Inference model with the LiteLLM Proxy Server
|
||||
|
||||
1. Modify the config.yaml
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: my-model
|
||||
litellm_params:
|
||||
model: wandb/<your-model-name> # add wandb/ prefix to use W&B Inference as provider
|
||||
api_key: api-key # api key to send your model
|
||||
```
|
||||
2. Start the proxy
|
||||
```bash
|
||||
$ litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Send Request to LiteLLM Proxy Server
|
||||
|
||||
<Tabs>
|
||||
|
||||
<TabItem value="openai" label="OpenAI Python v1.0.0+">
|
||||
|
||||
```python
|
||||
import openai
|
||||
client = openai.OpenAI(
|
||||
api_key="litellm-proxy-key", # pass litellm proxy key, if you're using virtual keys
|
||||
base_url="http://0.0.0.0:4000" # litellm-proxy-base url
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="my-model",
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What character was Wall-e in love with?"
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="curl" label="curl">
|
||||
|
||||
```shell
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Authorization: litellm-proxy-key' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "my-model",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What character was Wall-e in love with?"
|
||||
}
|
||||
],
|
||||
}'
|
||||
```
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
## Supported Parameters
|
||||
|
||||
The W&B Inference provider supports the following parameters:
|
||||
|
||||
### Chat Completion Parameters
|
||||
|
||||
| Parameter | Type | Description |
|
||||
| --------- | ---- | ----------- |
|
||||
| frequency_penalty | number | Penalizes new tokens based on their frequency in the text |
|
||||
| function_call | string/object | Controls how the model calls functions |
|
||||
| functions | array | List of functions for which the model may generate JSON inputs |
|
||||
| logit_bias | map | Modifies the likelihood of specified tokens |
|
||||
| max_tokens | integer | Maximum number of tokens to generate |
|
||||
| n | integer | Number of completions to generate |
|
||||
| presence_penalty | number | Penalizes tokens based on if they appear in the text so far |
|
||||
| response_format | object | Format of the response, e.g., `{"type": "json"}` |
|
||||
| seed | integer | Sampling seed for deterministic results |
|
||||
| stop | string/array | Sequences where the API will stop generating tokens |
|
||||
| stream | boolean | Whether to stream the response |
|
||||
| temperature | number | Controls randomness (0-2) |
|
||||
| top_p | number | Controls nucleus sampling |
|
||||
|
||||
|
||||
## Error Handling
|
||||
|
||||
The integration uses the standard LiteLLM error handling. Further, here's a list of commonly encountered errors with the W&B Inference API -
|
||||
|
||||
| Error Code | Message | Cause | Solution |
|
||||
| ---------- | ------- | ----- | -------- |
|
||||
| 401 | Authentication failed | Your authentication credentials are incorrect or your W&B project entity and/or name are incorrect. | Ensure you're using the correct API key and that your W&B project name and entity are correct. |
|
||||
| 403 | Country, region, or territory not supported | Accessing the API from an unsupported location. | Please see [Geographic restrictions](https://docs.wandb.ai/guides/inference/usage-limits/#geographic-restrictions) |
|
||||
| 429 | Concurrency limit reached for requests | Too many concurrent requests. | Reduce the number of concurrent requests or increase your limits. For more information, see [Usage information and limits](https://docs.wandb.ai/guides/inference/usage-limits/). |
|
||||
| 429 | You exceeded your current quota, please check your plan and billing details | Out of credits or reached monthly spending cap. | Get more credits or increase your limits. For more information, see [Usage information and limits](https://docs.wandb.ai/guides/inference/usage-limits/). |
|
||||
| 429 | W&B Inference isn't available for personal accounts. | Switch to a non-personal account. | Follow [the instructions below](#error-429-personal-entities-unsupported) for a work around. |
|
||||
| 500 | The server had an error while processing your request | Internal server error. | Retry after a brief wait and contact support if it persists. |
|
||||
| 503 | The engine is currently overloaded, please try again later | Server is experiencing high traffic. | Retry your request after a short delay. |
|
||||
|
||||
|
||||
### Error 429: Personal entities unsupported
|
||||
|
||||
The user is on a personal account, which doesn't have access to W&B Inference. If one isn't available, create a Team to create a non-personal account.
|
||||
|
||||
Once done, add the `openai-project` header to your request as shown below:
|
||||
|
||||
```python
|
||||
response = completion(
|
||||
model="...",
|
||||
extra_headers={"openai-project": "team_name/project_name"},
|
||||
...
|
||||
```
|
||||
|
||||
For more information, see [Personal entities unsupported](https://docs.wandb.ai/guides/inference/usage-limits/#personal-entities-unsupported).
|
||||
|
||||
You can find more ways of using custom headers with LiteLLM here - https://docs.litellm.ai/docs/proxy/request_headers.
|
||||
|
|
@ -81,6 +81,23 @@ MICROSOFT_TENANT="5a39737
|
|||
http://localhost:4000/sso/callback
|
||||
```
|
||||
|
||||
**Using App Roles for User Permissions**
|
||||
|
||||
You can assign user roles directly from Entra ID using App Roles. LiteLLM will automatically read the app roles from the JWT token and assign the corresponding role to the user.
|
||||
|
||||
Supported roles:
|
||||
- `proxy_admin` - Admin over the platform
|
||||
- `proxy_admin_viewer` - Can login, view all keys, view all spend (read-only)
|
||||
- `internal_user` - Normal user. Can login, view spend and depending on team-member permissions - view/create/delete their own keys.
|
||||
|
||||
|
||||
To set up app roles:
|
||||
1. Navigate to your App Registration on https://portal.azure.com/
|
||||
2. Go to "App roles" and create a new app role
|
||||
3. Use one of the supported role names above (e.g., `proxy_admin`)
|
||||
4. Assign users to these roles in your Enterprise Application
|
||||
5. When users sign in via SSO, LiteLLM will automatically assign them the corresponding role
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="Generic" label="Generic SSO Provider">
|
||||
|
|
|
|||
|
|
@ -278,6 +278,8 @@ Set either `REDIS_URL` or the `REDIS_HOST` in your os environment, to enable cac
|
|||
REDIS_HOST = "" # REDIS_HOST='redis-18841.c274.us-east-1-3.ec2.cloud.redislabs.com'
|
||||
REDIS_PORT = "" # REDIS_PORT='18841'
|
||||
REDIS_PASSWORD = "" # REDIS_PASSWORD='liteLlmIsAmazing'
|
||||
REDIS_USERNAME = "" # REDIS_USERNAME='my-redis-username' [OPTIONAL] if your redis server requires a username
|
||||
REDIS_SSL = "True" # REDIS_SSL='True' to enable SSL by default is False
|
||||
```
|
||||
|
||||
**Additional kwargs**
|
||||
|
|
|
|||
|
|
@ -353,7 +353,10 @@ router_settings:
|
|||
| AGENTOPS_SERVICE_NAME | Service Name for AgentOps logging integration
|
||||
| AISPEND_ACCOUNT_ID | Account ID for AI Spend
|
||||
| AISPEND_API_KEY | API Key for AI Spend
|
||||
| AIOHTTP_CONNECTOR_LIMIT | Connection limit for aiohttp connector. When set to 0, no limit is applied. **Default is 0**
|
||||
| AIOHTTP_KEEPALIVE_TIMEOUT | Keep-alive timeout for aiohttp connections in seconds. **Default is 120**
|
||||
| AIOHTTP_TRUST_ENV | Flag to enable aiohttp trust environment. When this is set to True, aiohttp will respect HTTP(S)_PROXY env vars. **Default is False**
|
||||
| AIOHTTP_TTL_DNS_CACHE | DNS cache time-to-live for aiohttp in seconds. **Default is 300**
|
||||
| ALLOWED_EMAIL_DOMAINS | List of email domains allowed for access
|
||||
| ARIZE_API_KEY | API key for Arize platform integration
|
||||
| ARIZE_SPACE_KEY | Space key for Arize platform
|
||||
|
|
@ -506,6 +509,8 @@ router_settings:
|
|||
| EMAIL_SIGNATURE | Custom HTML footer/signature for all emails. Can include HTML tags for formatting and links.
|
||||
| EMAIL_SUBJECT_INVITATION | Custom subject template for invitation emails.
|
||||
| EMAIL_SUBJECT_KEY_CREATED | Custom subject template for key creation emails.
|
||||
| ENKRYPTAI_API_BASE | Base URL for EnkryptAI Guardrails API. **Default is https://api.enkryptai.com**
|
||||
| ENKRYPTAI_API_KEY | API key for EnkryptAI Guardrails service
|
||||
| EXPERIMENTAL_MULTI_INSTANCE_RATE_LIMITING | Flag to enable new multi-instance rate limiting. **Default is False**
|
||||
| FIREWORKS_AI_4_B | Size parameter for Fireworks AI 4B model. Default is 4
|
||||
| FIREWORKS_AI_16_B | Size parameter for Fireworks AI 16B model. Default is 16
|
||||
|
|
@ -629,6 +634,7 @@ router_settings:
|
|||
| LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development)
|
||||
| LITELLM_RATE_LIMIT_WINDOW_SIZE | Rate limit window size for LiteLLM. Default is 60
|
||||
| LITELLM_SALT_KEY | Salt key for encryption in LiteLLM
|
||||
| LITELLM_SSL_CIPHERS | SSL/TLS cipher configuration for faster handshakes. Controls cipher suite preferences for OpenSSL connections.
|
||||
| LITELLM_SECRET_AWS_KMS_LITELLM_LICENSE | AWS KMS encrypted license for LiteLLM
|
||||
| LITELLM_TOKEN | Access token for LiteLLM integration
|
||||
| LITELLM_PRINT_STANDARD_LOGGING_PAYLOAD | If true, prints the standard logging payload to the console - useful for debugging
|
||||
|
|
@ -774,6 +780,8 @@ router_settings:
|
|||
| USE_AWS_KMS | Flag to enable AWS Key Management Service for encryption
|
||||
| USE_PRISMA_MIGRATE | Flag to use prisma migrate instead of prisma db push. Recommended for production environments.
|
||||
| WEBHOOK_URL | URL for receiving webhooks from external services
|
||||
| SPEND_LOG_RUN_LOOPS | Constant for setting how many runs of 1000 batch deletes should spend_log_cleanup task run |
|
||||
| SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000 |
|
||||
| COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY | Maximum size for CoroutineChecker in-memory cache. Default is 1000 |
|
||||
| SPEND_LOG_RUN_LOOPS | Constant for setting how many runs of 1000 batch deletes should spend_log_cleanup task run
|
||||
| SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000
|
||||
| COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY | Maximum size for CoroutineChecker in-memory cache. Default is 1000
|
||||
| DEFAULT_SHARED_HEALTH_CHECK_TTL | Time-to-live in seconds for cached health check results in shared health check mode. Default is 300 (5 minutes)
|
||||
| DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL | Time-to-live in seconds for health check lock in shared health check mode. Default is 60 (1 minute)
|
||||
|
|
@ -788,6 +788,30 @@ docker run --name litellm-proxy \
|
|||
## Platform-specific Guide
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="AWS ECS" label="AWS ECS - Elastic Container Service>
|
||||
|
||||
### Terraform-based ECS Deployment
|
||||
|
||||
LiteLLM maintains a dedicated Terraform tutorial for deploying the proxy on ECS. Follow the step-by-step guide in the [litellm-ecs-deployment repository](https://github.com/BerriAI/litellm-ecs-deployment) to provision the required ECS services, task definitions, and supporting AWS resources.
|
||||
|
||||
1. Clone the tutorial repository to review the Terraform modules and variables.
|
||||
```bash
|
||||
git clone https://github.com/BerriAI/litellm-ecs-deployment.git
|
||||
cd litellm-ecs-deployment
|
||||
```
|
||||
|
||||
2. Initialize and validate the Terraform project before applying it to your chosen workspace/account.
|
||||
```bash
|
||||
terraform init
|
||||
terraform plan
|
||||
terraform apply
|
||||
```
|
||||
|
||||
3. Once `terraform apply` completes, do `./build.sh` to push the repository on ECR and update the ECS cluster. Use that endpoint (port `4000` by default) for API requests to your LiteLLM proxy.
|
||||
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="AWS EKS" label="AWS EKS - Kubernetes">
|
||||
|
||||
### Kubernetes (AWS EKS)
|
||||
|
|
|
|||
276
docs/my-website/docs/proxy/guardrails/enkryptai.md
Normal file
|
|
@ -0,0 +1,276 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# EnkryptAI Guardrails
|
||||
|
||||
LiteLLM supports EnkryptAI guardrails for content moderation and safety checks on LLM inputs and outputs.
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Define Guardrails on your LiteLLM config.yaml
|
||||
|
||||
Define your guardrails under the `guardrails` section:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: openai/gpt-3.5-turbo
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "enkryptai-guard"
|
||||
litellm_params:
|
||||
guardrail: enkryptai
|
||||
mode: "pre_call"
|
||||
api_key: os.environ/ENKRYPTAI_API_KEY
|
||||
detectors:
|
||||
toxicity:
|
||||
enabled: true
|
||||
nsfw:
|
||||
enabled: true
|
||||
pii:
|
||||
enabled: true
|
||||
entities: ["email", "phone", "secrets"]
|
||||
injection_attack:
|
||||
enabled: true
|
||||
```
|
||||
|
||||
#### Supported values for `mode`
|
||||
|
||||
- `pre_call` - Run **before** LLM call, on **input**
|
||||
- `post_call` - Run **after** LLM call, on **output**
|
||||
- `during_call` - Run **during** LLM call, on **input**. Same as `pre_call` but runs in parallel as LLM call
|
||||
|
||||
#### Available Detectors
|
||||
|
||||
EnkryptAI supports multiple content detection types:
|
||||
|
||||
- **toxicity** - Detect toxic language
|
||||
- **nsfw** - Detect NSFW (Not Safe For Work) content
|
||||
- **pii** - Detect personally identifiable information
|
||||
- Configure entities: `["pii", "email", "phone", "secrets", "ip_address", "url"]`
|
||||
- **injection_attack** - Detect prompt injection attempts
|
||||
- **keyword_detector** - Detect custom keywords/phrases
|
||||
- **policy_violation** - Detect policy violations
|
||||
- **bias** - Detect biased content
|
||||
- **sponge_attack** - Detect sponge attacks
|
||||
|
||||
### 2. Set Environment Variables
|
||||
|
||||
```bash
|
||||
export ENKRYPTAI_API_KEY="your-api-key"
|
||||
```
|
||||
|
||||
### 3. Start LiteLLM Gateway
|
||||
|
||||
```shell
|
||||
litellm --config config.yaml --detailed_debug
|
||||
```
|
||||
|
||||
### 4. Test Request
|
||||
|
||||
**[Langchain, OpenAI SDK Usage Examples](../proxy/user_keys#request-format)**
|
||||
|
||||
<Tabs>
|
||||
<TabItem label="Successful Call" value="allowed">
|
||||
|
||||
```shell
|
||||
curl -i http://localhost:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how can you help me today?"}
|
||||
],
|
||||
"guardrails": ["enkryptai-guard"]
|
||||
}'
|
||||
```
|
||||
|
||||
**Response: HTTP 200 Success**
|
||||
|
||||
Content passes all detector checks and is allowed through.
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem label="Unsuccessful Call" value="not-allowed">
|
||||
|
||||
Expect this to fail if content violates detector policies:
|
||||
|
||||
```shell
|
||||
curl -i http://localhost:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [
|
||||
{"role": "user", "content": "My email is test@example.com and my SSN is 123-45-6789"}
|
||||
],
|
||||
"guardrails": ["enkryptai-guard"]
|
||||
}'
|
||||
```
|
||||
|
||||
**Expected Response on Failure: HTTP 400 Error**
|
||||
|
||||
```json
|
||||
{
|
||||
"error": {
|
||||
"message": {
|
||||
"error": "Content blocked by EnkryptAI guardrail",
|
||||
"detected": true,
|
||||
"violations": ["pii"],
|
||||
"response": {
|
||||
"summary": {
|
||||
"pii": 1
|
||||
},
|
||||
"details": {
|
||||
"pii": {
|
||||
"detected": ["email", "ssn"]
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"type": "None",
|
||||
"param": "None",
|
||||
"code": "400"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Video Walkthrough
|
||||
|
||||
<iframe width="840" height="500" src="https://www.loom.com/embed/ff222211e0864937aee4aeef0f28c3b7" frameborder="0" webkitallowfullscreen mozallowfullscreen allowfullscreen></iframe>
|
||||
|
||||
## Advanced Configuration
|
||||
|
||||
### Using Custom Policies
|
||||
|
||||
You can specify a custom EnkryptAI policy:
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: "enkryptai-custom"
|
||||
litellm_params:
|
||||
guardrail: enkryptai
|
||||
mode: "pre_call"
|
||||
api_key: os.environ/ENKRYPTAI_API_KEY
|
||||
policy_name: "my-custom-policy" # Sent via x-enkrypt-policy header
|
||||
detectors:
|
||||
toxicity:
|
||||
enabled: true
|
||||
```
|
||||
|
||||
### Using Deployments
|
||||
|
||||
Specify an EnkryptAI deployment:
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: "enkryptai-deployment"
|
||||
litellm_params:
|
||||
guardrail: enkryptai
|
||||
mode: "pre_call"
|
||||
api_key: os.environ/ENKRYPTAI_API_KEY
|
||||
deployment_name: "production" # Sent via X-Enkrypt-Deployment header
|
||||
detectors:
|
||||
toxicity:
|
||||
enabled: true
|
||||
```
|
||||
|
||||
### Monitor Mode (Logging Without Blocking)
|
||||
|
||||
Set `block_on_violation: false` to log violations without blocking requests:
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: "enkryptai-monitor"
|
||||
litellm_params:
|
||||
guardrail: enkryptai
|
||||
mode: "pre_call"
|
||||
api_key: os.environ/ENKRYPTAI_API_KEY
|
||||
block_on_violation: false # Log violations but don't block
|
||||
detectors:
|
||||
toxicity:
|
||||
enabled: true
|
||||
nsfw:
|
||||
enabled: true
|
||||
```
|
||||
|
||||
In monitor mode, all violations are logged but requests are never blocked.
|
||||
|
||||
### Input and Output Guardrails
|
||||
|
||||
Configure separate guardrails for input and output:
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
# Input guardrail
|
||||
- guardrail_name: "enkryptai-input"
|
||||
litellm_params:
|
||||
guardrail: enkryptai
|
||||
mode: "pre_call"
|
||||
api_key: os.environ/ENKRYPTAI_API_KEY
|
||||
detectors:
|
||||
pii:
|
||||
enabled: true
|
||||
entities: ["email", "phone", "ssn"]
|
||||
injection_attack:
|
||||
enabled: true
|
||||
|
||||
# Output guardrail
|
||||
- guardrail_name: "enkryptai-output"
|
||||
litellm_params:
|
||||
guardrail: enkryptai
|
||||
mode: "post_call"
|
||||
api_key: os.environ/ENKRYPTAI_API_KEY
|
||||
detectors:
|
||||
toxicity:
|
||||
enabled: true
|
||||
nsfw:
|
||||
enabled: true
|
||||
```
|
||||
|
||||
## Configuration Options
|
||||
|
||||
| Parameter | Type | Description | Default |
|
||||
|-----------|------|-------------|---------|
|
||||
| `api_key` | string | EnkryptAI API key | `ENKRYPTAI_API_KEY` env var |
|
||||
| `api_base` | string | EnkryptAI API base URL | `https://api.enkryptai.com` |
|
||||
| `policy_name` | string | Custom policy name (sent via `x-enkrypt-policy` header) | None |
|
||||
| `deployment_name` | string | Deployment name (sent via `X-Enkrypt-Deployment` header) | None |
|
||||
| `detectors` | object | Detector configuration | `{}` |
|
||||
| `block_on_violation` | boolean | Block requests on violations | `true` |
|
||||
| `mode` | string | When to run: `pre_call`, `post_call`, or `during_call` | Required |
|
||||
|
||||
## Observability
|
||||
|
||||
EnkryptAI guardrail logs include:
|
||||
|
||||
- **guardrail_status**: `success`, `guardrail_intervened`, or `guardrail_failed_to_respond`
|
||||
- **guardrail_provider**: `enkryptai`
|
||||
- **guardrail_json_response**: Full API response with detection details
|
||||
- **duration**: Time taken for guardrail check
|
||||
- **start_time** and **end_time**: Timestamps
|
||||
|
||||
These logs are available through your configured LiteLLM logging callbacks.
|
||||
|
||||
## Error Handling
|
||||
|
||||
The guardrail handles errors gracefully:
|
||||
|
||||
- **API Failures**: Logs error and raises exception
|
||||
- **Rate Limits (429)**: Logs error and raises exception
|
||||
- **Invalid Configuration**: Raises `ValueError` on initialization
|
||||
|
||||
Set `block_on_violation: false` to continue processing even when violations are detected (monitor mode).
|
||||
|
||||
## Support
|
||||
|
||||
For more information about EnkryptAI:
|
||||
- Documentation: [https://docs.enkryptai.com](https://docs.enkryptai.com)
|
||||
- Website: [https://enkryptai.com](https://enkryptai.com)
|
||||
|
||||
|
|
@ -9,13 +9,32 @@ Use this to health check all LLMs defined in your config.yaml
|
|||
| `/health/readiness` | **Load balancer health checks** | Ready to accept traffic - includes DB connection status |
|
||||
| `/health` | **Model health monitoring** | Comprehensive LLM model health - makes actual API calls |
|
||||
| `/health/services` | **Service debugging** | Check specific integrations (datadog, langfuse, etc.) |
|
||||
| `/health/shared-status` | **Multi-pod coordination** | Monitor shared health check state across pods |
|
||||
|
||||
## Summary
|
||||
|
||||
The proxy exposes:
|
||||
* a /health endpoint which returns the health of the LLM APIs
|
||||
* a /health/readiness endpoint for returning if the proxy is ready to accept requests
|
||||
* a /health/liveliness endpoint for returning if the proxy is alive
|
||||
* a /health/liveliness endpoint for returning if the proxy is alive
|
||||
* a /health/shared-status endpoint for monitoring shared health check coordination across pods
|
||||
|
||||
## Shared Health Check State
|
||||
|
||||
When running multiple LiteLLM proxy pods, you can enable shared health check state to coordinate health checks across pods and avoid duplicate API calls. This is especially beneficial for expensive models like Gemini 2.5-pro.
|
||||
|
||||
**Key Benefits:**
|
||||
- Reduces duplicate health checks across pods
|
||||
- Saves costs on expensive model API calls
|
||||
- Reduces monitoring noise and logging
|
||||
- Improves resource efficiency
|
||||
|
||||
**Requirements:**
|
||||
- Redis for shared state coordination
|
||||
- Background health checks enabled
|
||||
- Multiple proxy pods
|
||||
|
||||
For detailed configuration and usage, see [Shared Health Check State](./shared_health_check.md).
|
||||
|
||||
## `/health`
|
||||
#### Request
|
||||
|
|
|
|||
310
docs/my-website/docs/proxy/shared_health_check.md
Normal file
|
|
@ -0,0 +1,310 @@
|
|||
# Shared Health Check State Across Pods
|
||||
|
||||
This feature enables coordination of health checks across multiple LiteLLM proxy pods to avoid duplicate health checks and reduce costs.
|
||||
|
||||
## Overview
|
||||
|
||||
When running multiple LiteLLM proxy pods (e.g., in Kubernetes), each pod typically runs its own independent health checks on every model. This can result in:
|
||||
|
||||
- **Duplicate health checks** across pods
|
||||
- **Increased costs** for expensive models (e.g., Gemini 2.5-pro)
|
||||
- **Redundant monitoring/logging noise**
|
||||
- **Inefficient resource usage**
|
||||
|
||||
The shared health check state feature solves this by:
|
||||
|
||||
- **Coordinating health checks** across pods using Redis
|
||||
- **Caching results** with configurable TTL
|
||||
- **Using distributed locks** to ensure only one pod runs health checks at a time
|
||||
- **Allowing other pods** to read cached results instead of running redundant checks
|
||||
|
||||
## How It Works
|
||||
|
||||
### 1. Lock Acquisition
|
||||
When a pod needs to run health checks:
|
||||
- It attempts to acquire a Redis lock
|
||||
- If successful, it runs the health checks
|
||||
- If failed, it waits briefly and checks for cached results
|
||||
|
||||
### 2. Result Caching
|
||||
After running health checks:
|
||||
- Results are cached in Redis with a configurable TTL
|
||||
- Other pods can read these cached results
|
||||
- Cache includes timestamp and pod ID for tracking
|
||||
|
||||
### 3. Fallback Behavior
|
||||
If Redis is unavailable or cache is expired:
|
||||
- Pods fall back to running health checks locally
|
||||
- System continues to function normally
|
||||
|
||||
## Configuration
|
||||
|
||||
### Enable Shared Health Check
|
||||
|
||||
Add to your `proxy_config.yaml`:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
# Enable background health checks (required)
|
||||
background_health_checks: true
|
||||
|
||||
# Enable shared health check state across pods
|
||||
use_shared_health_check: true
|
||||
|
||||
# Health check interval (seconds)
|
||||
health_check_interval: 300 # 5 minutes
|
||||
|
||||
# Redis configuration (required for shared health check)
|
||||
litellm_settings:
|
||||
cache: true
|
||||
cache_params:
|
||||
type: redis
|
||||
host: your-redis-host
|
||||
port: 6379
|
||||
password: your-redis-password
|
||||
```
|
||||
|
||||
### Environment Variables
|
||||
|
||||
You can also configure using environment variables:
|
||||
|
||||
```bash
|
||||
# Enable shared health check
|
||||
export USE_SHARED_HEALTH_CHECK=true
|
||||
|
||||
# Health check TTL (seconds)
|
||||
export DEFAULT_SHARED_HEALTH_CHECK_TTL=300
|
||||
|
||||
# Lock TTL (seconds)
|
||||
export DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL=60
|
||||
```
|
||||
|
||||
## Requirements
|
||||
|
||||
- **Redis**: Required for shared state coordination
|
||||
- **Background Health Checks**: Must be enabled (`background_health_checks: true`)
|
||||
- **Multiple Pods**: Most beneficial with 2+ proxy instances
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### Check Shared Health Check Status
|
||||
|
||||
```bash
|
||||
GET /health/shared-status
|
||||
```
|
||||
|
||||
Returns information about the shared health check coordination:
|
||||
|
||||
```json
|
||||
{
|
||||
"shared_health_check_enabled": true,
|
||||
"status": {
|
||||
"pod_id": "pod_1703123456789",
|
||||
"redis_available": true,
|
||||
"lock_ttl": 60,
|
||||
"cache_ttl": 300,
|
||||
"lock_owner": "pod_1703123456788",
|
||||
"lock_in_progress": true,
|
||||
"cache_available": true,
|
||||
"cache_age_seconds": 45.2,
|
||||
"last_checked_by": "pod_1703123456788"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Monitoring
|
||||
|
||||
### Health Check Status
|
||||
|
||||
Monitor the shared health check status to ensure proper coordination:
|
||||
|
||||
```bash
|
||||
curl -H "Authorization: Bearer your-api-key" \
|
||||
http://your-proxy-host/health/shared-status
|
||||
```
|
||||
|
||||
### Logs
|
||||
|
||||
Look for these log messages:
|
||||
|
||||
```
|
||||
INFO: Initialized shared health check manager
|
||||
INFO: Pod pod_123 acquired health check lock
|
||||
INFO: Pod pod_123 released health check lock
|
||||
INFO: Cached health check results for 5 healthy and 0 unhealthy endpoints
|
||||
DEBUG: Using cached health check results
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Common Issues
|
||||
|
||||
#### 1. Shared Health Check Not Working
|
||||
|
||||
**Symptoms**: Each pod still runs independent health checks
|
||||
|
||||
**Solutions**:
|
||||
- Verify Redis is configured and accessible
|
||||
- Check that `use_shared_health_check: true` is set
|
||||
- Ensure `background_health_checks: true` is enabled
|
||||
- Check Redis connectivity in logs
|
||||
|
||||
#### 2. Redis Connection Issues
|
||||
|
||||
**Symptoms**: Health checks fall back to local execution
|
||||
|
||||
**Solutions**:
|
||||
- Verify Redis host, port, and credentials
|
||||
- Check network connectivity between pods and Redis
|
||||
- Monitor Redis server logs for errors
|
||||
|
||||
#### 3. Lock Not Released
|
||||
|
||||
**Symptoms**: One pod holds the lock indefinitely
|
||||
|
||||
**Solutions**:
|
||||
- Lock has automatic TTL (default 60 seconds)
|
||||
- Check pod logs for lock release messages
|
||||
- Verify Redis TTL settings
|
||||
|
||||
### Debug Mode
|
||||
|
||||
Enable debug logging to see detailed coordination:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
set_verbose: true
|
||||
```
|
||||
|
||||
## Performance Impact
|
||||
|
||||
### Benefits
|
||||
|
||||
- **Reduced API calls**: Only one pod runs health checks per interval
|
||||
- **Lower costs**: Especially significant for expensive models
|
||||
- **Better resource utilization**: Less redundant work across pods
|
||||
- **Cleaner monitoring**: Reduced noise in logs and metrics
|
||||
|
||||
### Overhead
|
||||
|
||||
- **Redis operations**: Minimal overhead for lock/cache operations
|
||||
- **Network latency**: Small delay for Redis communication
|
||||
- **Memory usage**: Negligible additional memory usage
|
||||
|
||||
## Best Practices
|
||||
|
||||
### 1. Redis Configuration
|
||||
|
||||
- Use Redis with persistence enabled
|
||||
- Configure appropriate memory limits
|
||||
- Set up Redis monitoring and alerts
|
||||
|
||||
### 2. TTL Settings
|
||||
|
||||
- Set `health_check_interval` to your desired check frequency
|
||||
- Use default TTL values unless you have specific requirements
|
||||
- Consider model-specific timeouts for expensive models
|
||||
|
||||
### 3. Monitoring
|
||||
|
||||
- Monitor shared health check status endpoint
|
||||
- Set up alerts for Redis connectivity issues
|
||||
- Track health check costs and frequency
|
||||
|
||||
### 4. Scaling
|
||||
|
||||
- Feature works with any number of pods
|
||||
- More pods = better coordination benefits
|
||||
- Consider Redis cluster for high availability
|
||||
|
||||
## Example Configuration
|
||||
|
||||
### Complete Example
|
||||
|
||||
```yaml
|
||||
# proxy_config.yaml
|
||||
model_list:
|
||||
- model_name: gpt-4
|
||||
litellm_params:
|
||||
model: gpt-4
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
model_info:
|
||||
health_check_timeout: 30 # 30 second timeout for health checks
|
||||
|
||||
general_settings:
|
||||
# Enable background health checks
|
||||
background_health_checks: true
|
||||
|
||||
# Enable shared health check coordination
|
||||
use_shared_health_check: true
|
||||
|
||||
# Health check interval (5 minutes)
|
||||
health_check_interval: 300
|
||||
|
||||
# Health check details
|
||||
health_check_details: true
|
||||
|
||||
litellm_settings:
|
||||
# Redis configuration
|
||||
cache: true
|
||||
cache_params:
|
||||
type: redis
|
||||
host: redis-cluster.example.com
|
||||
port: 6379
|
||||
password: os.environ/REDIS_PASSWORD
|
||||
ssl: true
|
||||
```
|
||||
|
||||
### Kubernetes Example
|
||||
|
||||
```yaml
|
||||
# deployment.yaml
|
||||
apiVersion: apps/v1
|
||||
kind: Deployment
|
||||
metadata:
|
||||
name: litellm-proxy
|
||||
spec:
|
||||
replicas: 3 # Multiple pods for coordination
|
||||
template:
|
||||
spec:
|
||||
containers:
|
||||
- name: litellm-proxy
|
||||
image: ghcr.io/berriai/litellm:latest
|
||||
env:
|
||||
- name: USE_SHARED_HEALTH_CHECK
|
||||
value: "true"
|
||||
- name: REDIS_HOST
|
||||
value: "redis-service"
|
||||
- name: REDIS_PASSWORD
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: redis-secret
|
||||
key: password
|
||||
```
|
||||
|
||||
## Migration
|
||||
|
||||
### From Independent Health Checks
|
||||
|
||||
1. **Enable Redis**: Ensure Redis is configured and accessible
|
||||
2. **Enable Background Health Checks**: Set `background_health_checks: true`
|
||||
3. **Enable Shared Health Check**: Set `use_shared_health_check: true`
|
||||
4. **Deploy**: Update your proxy configuration
|
||||
5. **Monitor**: Check `/health/shared-status` endpoint
|
||||
|
||||
### Rollback
|
||||
|
||||
To disable shared health check:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
use_shared_health_check: false
|
||||
# background_health_checks can remain true for independent checks
|
||||
```
|
||||
|
||||
## Related Features
|
||||
|
||||
- [Background Health Checks](./health.md#background-health-checks)
|
||||
- [Redis Caching](./caching.md)
|
||||
- [High Availability Setup](./db_deadlocks.md)
|
||||
- [Health Check Endpoints](./health.md#health-endpoints)
|
||||
277
docs/my-website/docs/proxy/tag_budgets.md
Normal file
|
|
@ -0,0 +1,277 @@
|
|||
import Image from '@theme/IdealImage';
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Setting Tag Budgets
|
||||
|
||||
Track spend and set budgets for your API requests using tags. Tags allow you to categorize and monitor costs across different cost centers, projects, and departments.
|
||||
|
||||
## Pre-Requisites
|
||||
|
||||
- You must set up a Postgres database (e.g. Supabase, Neon, etc.)
|
||||
|
||||
## What are Tags?
|
||||
|
||||
Tags are labels you can attach to your LLM requests to track and limit spending by category.
|
||||
|
||||
**Common Use Cases:**
|
||||
- **Cost Center Tracking**: Allocate LLM costs to specific departments or business units (e.g., "engineering", "marketing", "customer-support")
|
||||
- **Project-based Budgeting**: Set budgets for different projects or initiatives (e.g., "project-alpha", "chatbot-v2")
|
||||
- **Customer Attribution**: Track spend per customer or client (e.g., "customer-acme", "customer-techcorp")
|
||||
- **Feature Monitoring**: Monitor costs for specific features (e.g., "feature-chat", "feature-summarization")
|
||||
|
||||
Tags are added to each request in the `metadata` field to track and enforce budget limits.
|
||||
|
||||
## Setting Tag Budgets
|
||||
|
||||
### 1. Create a tag with budget
|
||||
|
||||
Create a tag to represent a cost center, project, or any budget category. Set `max_budget` ($ value allowed) and `budget_duration` (how frequently the budget resets).
|
||||
|
||||
**Example:** Create a tag for your Engineering department with a monthly $500 budget
|
||||
|
||||
#### API
|
||||
|
||||
Create a new tag and set `max_budget` and `budget_duration`
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/tag/new' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"name": "engineering",
|
||||
"description": "Engineering department cost center",
|
||||
"max_budget": 500.0,
|
||||
"budget_duration": "30d"
|
||||
}'
|
||||
```
|
||||
|
||||
**Request Body Parameters:**
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
|-----------|------|----------|-------------|
|
||||
| `name` | string | Yes | Unique name for the tag (e.g., cost center name) |
|
||||
| `description` | string | No | Description of what this tag tracks |
|
||||
| `models` | list[string] | No | Restrict tag to specific models |
|
||||
| `max_budget` | float | No | Maximum budget in USD |
|
||||
| `budget_duration` | string | No | How often budget resets (e.g., "30d", "1d") |
|
||||
| `soft_budget` | float | No | Soft budget limit for warnings |
|
||||
|
||||
**Response:**
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "engineering",
|
||||
"description": "Engineering department cost center",
|
||||
"max_budget": 500.0,
|
||||
"budget_duration": "30d",
|
||||
"budget_reset_at": "2025-11-10T00:00:00Z",
|
||||
"created_at": "2025-10-11T00:00:00Z"
|
||||
}
|
||||
```
|
||||
|
||||
#### LiteLLM Admin UI
|
||||
|
||||
Navigate to the **Tag Management** page and click **Create New Tag**. Fill in the tag details and set your budget:
|
||||
|
||||
<Image
|
||||
img={require('../../img/tag_budget1.png')}
|
||||
style={{width: '80%', display: 'block', margin: '0'}}
|
||||
/>
|
||||
|
||||
<br />
|
||||
|
||||
|
||||
**Possible values for `budget_duration`:**
|
||||
|
||||
| `budget_duration` | When Budget will reset |
|
||||
| --- | --- |
|
||||
| `budget_duration="1s"` | every 1 second |
|
||||
| `budget_duration="1m"` | every 1 minute |
|
||||
| `budget_duration="1h"` | every 1 hour |
|
||||
| `budget_duration="1d"` | every 1 day |
|
||||
| `budget_duration="7d"` | every 1 week |
|
||||
| `budget_duration="30d"` | every 1 month |
|
||||
|
||||
### 2. Use the tag in your requests
|
||||
|
||||
Add tags to your API requests in the `metadata` field:
|
||||
|
||||
:::info Tags Budgets on API Keys
|
||||
|
||||
Currently, tag budget enforcement is only supported per request. If you'd like to set tags on API keys so all requests automatically inherit the tags budgets, please [create a feature request on GitHub](https://github.com/BerriAI/litellm/issues/new?assignees=&labels=enhancement&projects=&template=feature_request.yml&title=%5BFeat%5D%3A).
|
||||
|
||||
:::
|
||||
|
||||
<Tabs>
|
||||
|
||||
<TabItem value="openai" label="OpenAI SDK">
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(
|
||||
api_key="sk-1234", # Your LiteLLM proxy key
|
||||
base_url="http://0.0.0.0:4000"
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="gpt-4",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
extra_body={
|
||||
"metadata": {
|
||||
"tags": ["engineering"]
|
||||
}
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="curl" label="cURL">
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"metadata": {
|
||||
"tags": ["engineering"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
### 3. Test It
|
||||
|
||||
Make requests until the budget is exceeded:
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"metadata": {
|
||||
"tags": ["engineering"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**When budget is exceeded, you'll see:**
|
||||
|
||||
```json
|
||||
{
|
||||
"error": {
|
||||
"message": "Budget has been exceeded! Tag=engineering Current cost: 505.50, Max budget: 500.0",
|
||||
"type": "budget_exceeded",
|
||||
"param": null,
|
||||
"code": "400"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Managing Tags
|
||||
|
||||
### View Tag Information
|
||||
|
||||
Get information about specific tags:
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/tag/info' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"names": ["engineering", "marketing"]
|
||||
}'
|
||||
```
|
||||
|
||||
**Response:**
|
||||
|
||||
```json
|
||||
{
|
||||
"engineering": {
|
||||
"name": "engineering",
|
||||
"description": "Engineering department cost center",
|
||||
"spend": 245.50,
|
||||
"max_budget": 500.0,
|
||||
"budget_duration": "30d",
|
||||
"budget_reset_at": "2025-11-10T00:00:00Z",
|
||||
"created_at": "2025-10-11T00:00:00Z",
|
||||
"updated_at": "2025-10-11T12:30:00Z"
|
||||
},
|
||||
"marketing": {
|
||||
"name": "marketing",
|
||||
"description": "Marketing department cost center",
|
||||
"spend": 89.20,
|
||||
"max_budget": 300.0,
|
||||
"budget_duration": "30d",
|
||||
"budget_reset_at": "2025-11-10T00:00:00Z",
|
||||
"created_at": "2025-10-11T00:00:00Z",
|
||||
"updated_at": "2025-10-11T12:30:00Z"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Update Tag Budget
|
||||
|
||||
Update an existing tag's budget:
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/tag/update' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"name": "engineering",
|
||||
"max_budget": 750.0,
|
||||
"budget_duration": "30d"
|
||||
}'
|
||||
```
|
||||
|
||||
### Delete Tag
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/tag/delete' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"name": "engineering"
|
||||
}'
|
||||
```
|
||||
|
||||
## Multiple Tags per Request
|
||||
|
||||
You can apply multiple tags to a single request to track costs across different dimensions simultaneously. For example, track both the cost center and the specific project:
|
||||
|
||||
```python
|
||||
response = client.chat.completions.create(
|
||||
model="gpt-4",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
extra_body={
|
||||
"metadata": {
|
||||
"tags": ["engineering", "project-alpha", "customer-acme"]
|
||||
}
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
```shell
|
||||
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"metadata": {
|
||||
"tags": ["engineering", "project-alpha", "customer-acme"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
**Budget Enforcement:** If any tag exceeds its budget, the request will be rejected.
|
||||
|
|
@ -1,5 +1,6 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
import Image from '@theme/IdealImage';
|
||||
|
||||
# Virtual Keys
|
||||
Track Spend, and control model access via virtual keys for the proxy
|
||||
|
|
@ -66,50 +67,6 @@ curl 'http://0.0.0.0:4000/key/generate' \
|
|||
--data-raw '{"models": ["gpt-3.5-turbo", "gpt-4"], "metadata": {"user": "ishaan@berri.ai"}}'
|
||||
```
|
||||
|
||||
## 🔁 Scheduled Key Rotations (NEW in v1.77.5)
|
||||
|
||||
LiteLLM can now rotate **virtual keys automatically** on a schedule you define.
|
||||
|
||||
### How it works
|
||||
1. When creating a virtual key you set `rotation_schedule` – a [cron expression](https://crontab.guru/).
|
||||
2. LiteLLM stores the schedule in the DB and runs a background job that regenerates the key at the specified time.
|
||||
3. Existing key string is invalidated; a **notification webhook** (if configured) is sent with the new key value.
|
||||
|
||||
### Create a key with rotation
|
||||
|
||||
```bash
|
||||
curl 'http://0.0.0.0:4000/key/generate' \
|
||||
-H 'Authorization: Bearer <your-master-key>' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"models": ["gpt-4o"],
|
||||
"rotation_schedule": "0 0 * * SUN", # rotate every Sunday at 00:00 UTC
|
||||
"webhook_url": "https://example.com/key-rotated"
|
||||
}'
|
||||
```
|
||||
|
||||
### Enable globally via env
|
||||
|
||||
Set these env vars when starting the proxy:
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `LITELLM_KEY_ROTATION_ENABLED` | Enable the rotation worker | `false` |
|
||||
| `LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS` | How often to scan for keys to rotate | `86400` |
|
||||
|
||||
### Webhook payload
|
||||
|
||||
```json
|
||||
{
|
||||
"event": "virtual_key.rotated",
|
||||
"old_key_id": "sk-abc...",
|
||||
"new_key": "sk-def...",
|
||||
"rotation_time": "2025-10-05T00:00:00Z"
|
||||
}
|
||||
```
|
||||
|
||||
If no `webhook_url` is provided the new key value is returned in the response of the `/key/rotate` REST call instead.
|
||||
|
||||
## Spend Tracking
|
||||
|
||||
Get spend per:
|
||||
|
|
@ -604,6 +561,94 @@ curl 'http://localhost:4000/key/sk-1234/regenerate' \
|
|||
[**👉 API REFERENCE DOCS**](https://litellm-api.up.railway.app/#/key%20management/regenerate_key_fn_key__key__regenerate_post)
|
||||
|
||||
|
||||
### Scheduled Key Rotations
|
||||
|
||||
LiteLLM can rotate **virtual keys automatically** based on time intervals you define.
|
||||
|
||||
#### Prerequisites
|
||||
|
||||
1. **Database connection required** - Key rotation requires a connected database to track rotation schedules
|
||||
2. **Enable the rotation worker** - Set environment variable `LITELLM_KEY_ROTATION_ENABLED=true`
|
||||
3. **Configure check interval** - Optionally set `LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS` (default: 86400 seconds / 24 hours)
|
||||
|
||||
#### How it works
|
||||
|
||||
1. When creating a virtual key, set `auto_rotate: true` and `rotation_interval` (duration string)
|
||||
2. LiteLLM calculates the next rotation time as `now + rotation_interval` and stores it in the database
|
||||
3. A background job periodically checks for keys where the rotation time has passed
|
||||
4. When a key is due for rotation, LiteLLM automatically regenerates it and invalidates the old key string
|
||||
5. The new rotation time is calculated and the cycle continues
|
||||
|
||||
#### Create a key with auto rotation
|
||||
|
||||
**API**
|
||||
```bash
|
||||
curl 'http://0.0.0.0:4000/key/generate' \
|
||||
-H 'Authorization: Bearer <your-master-key>' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"models": ["gpt-4o"],
|
||||
"auto_rotate": true,
|
||||
"rotation_interval": "30d"
|
||||
}'
|
||||
```
|
||||
|
||||
**LiteLLM UI**
|
||||
|
||||
On the LiteLLM UI, Navigate to the Keys page and click on `Generate Key` > `Key Lifecycle` > `Enable Auto Rotation`
|
||||
<Image
|
||||
img={require('../../img/key_r.png')}
|
||||
style={{width: '30%', display: 'block', margin: '0'}}
|
||||
/>
|
||||
|
||||
**Valid rotation_interval formats:**
|
||||
- `"30s"` - 30 seconds
|
||||
- `"30m"` - 30 minutes
|
||||
- `"30h"` - 30 hours
|
||||
- `"30d"` - 30 days
|
||||
- `"90d"` - 90 days
|
||||
|
||||
#### Update existing key to enable rotation
|
||||
|
||||
**API**
|
||||
|
||||
```bash
|
||||
curl 'http://0.0.0.0:4000/key/update' \
|
||||
-H 'Authorization: Bearer <your-master-key>' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"key": "sk-existing-key",
|
||||
"auto_rotate": true,
|
||||
"rotation_interval": "90d"
|
||||
}'
|
||||
```
|
||||
|
||||
**LiteLLM UI**
|
||||
|
||||
On the LiteLLM UI, Navigate to the Keys page. Select the key you want to update and click on `Edit Settings` > `Auto-Rotation Settings`
|
||||
|
||||
<Image
|
||||
img={require('../../img/key_u.png')}
|
||||
style={{width: '30%', display: 'block', margin: '0'}}
|
||||
/>
|
||||
|
||||
#### Environment variables
|
||||
|
||||
Set these environment variables when starting the proxy:
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `LITELLM_KEY_ROTATION_ENABLED` | Enable the rotation worker | `false` |
|
||||
| `LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS` | How often to scan for keys to rotate (in seconds) | `86400` (24 hours) |
|
||||
|
||||
**Example:**
|
||||
```bash
|
||||
export LITELLM_KEY_ROTATION_ENABLED=true
|
||||
export LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS=3600 # Check every hour
|
||||
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
### Temporary Budget Increase
|
||||
|
||||
Use the `/key/update` endpoint to increase the budget of an existing key.
|
||||
|
|
|
|||
|
|
@ -105,7 +105,7 @@ LITELLM_MASTER_KEY gives claude access to all proxy models, whereas a virtual ke
|
|||
Alternatively, use the Anthropic pass-through endpoint:
|
||||
|
||||
```bash
|
||||
export ANTHROPIC_BASE_URL="http://0.0.0.0:4000/anthropic"
|
||||
export ANTHROPIC_BASE_URL="http://0.0.0.0:4000"
|
||||
export ANTHROPIC_AUTH_TOKEN="$LITELLM_MASTER_KEY"
|
||||
```
|
||||
|
||||
|
|
@ -209,4 +209,81 @@ claude --model claude-bedrock
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
<Image img={require('../../img/release_notes/claude_code_demo.png')} style={{ width: '500px', height: 'auto' }} />
|
||||
<Image img={require('../../img/release_notes/claude_code_demo.png')} style={{ width: '500px', height: 'auto' }} />
|
||||
|
||||
|
||||
## Connecting MCP Servers
|
||||
|
||||
You can also connect MCP servers to Claude Code via LiteLLM Proxy.
|
||||
|
||||
:::note
|
||||
|
||||
Limitations:
|
||||
|
||||
- Currently, only HTTP MCP servers are supported
|
||||
- Does not work in Cursor IDE yet.
|
||||
|
||||
:::
|
||||
|
||||
1. Add the MCP server to your `config.yaml`
|
||||
|
||||
In this example, we'll add the Github MCP server to our `config.yaml`
|
||||
|
||||
```yaml title="config.yaml" showLineNumbers
|
||||
mcp_servers:
|
||||
github_mcp:
|
||||
url: "https://api.githubcopilot.com/mcp"
|
||||
auth_type: oauth2
|
||||
authorization_url: https://github.com/login/oauth/authorize
|
||||
token_url: https://github.com/login/oauth/access_token
|
||||
client_id: os.environ/GITHUB_OAUTH_CLIENT_ID
|
||||
client_secret: os.environ/GITHUB_OAUTH_CLIENT_SECRET
|
||||
scopes: ["public_repo", "user:email"]
|
||||
```
|
||||
|
||||
2. Start LiteLLM Proxy
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
3. Use the MCP server in Claude Code
|
||||
|
||||
```bash
|
||||
claude mcp add --transport http litellm_proxy http://0.0.0.0:4000 --header "Authorization: Bearer sk-LITELLM_VIRTUAL_KEY"
|
||||
```
|
||||
|
||||
4. Authenticate via Claude Code
|
||||
|
||||
a. Start Claude Code
|
||||
|
||||
```bash
|
||||
claude
|
||||
```
|
||||
|
||||
b. Authenticate via Claude Code
|
||||
|
||||
```bash
|
||||
/mcp
|
||||
```
|
||||
|
||||
c. Select the MCP server
|
||||
|
||||
```bash
|
||||
> litellm_proxy
|
||||
```
|
||||
|
||||
d. Start Oauth flow via Claude Code
|
||||
|
||||
```bash
|
||||
> 1. Authenticate
|
||||
2. Reconnect
|
||||
3. Disable
|
||||
```
|
||||
|
||||
e. Once completed, you should see this success message:
|
||||
|
||||
<Image img={require('../../img/oauth_2_success.png')} style={{ width: '500px', height: 'auto' }} />
|
||||
|
||||
|
|
|
|||
|
|
@ -140,6 +140,54 @@ litellm_settings:
|
|||
<Image img={require('../../img/msft_default_settings.png')} style={{ width: '900px', height: 'auto' }} />
|
||||
|
||||
|
||||
## 4. Using Entra ID App Roles for User Permissions
|
||||
|
||||
You can assign user roles directly from Entra ID using App Roles. LiteLLM will automatically read the app roles from the JWT token during SSO sign-in and assign the corresponding role to the user.
|
||||
|
||||
### 4.1 Supported Roles
|
||||
|
||||
LiteLLM supports the following app roles (case-insensitive):
|
||||
|
||||
- `proxy_admin` - Admin over the entire LiteLLM platform
|
||||
- `proxy_admin_viewer` - Read-only admin access (can view all keys and spend)
|
||||
- `org_admin` - Admin over a specific organization (can create teams and users within their org)
|
||||
- `internal_user` - Standard user (can create/view/delete their own keys and view their own spend)
|
||||
|
||||
### 4.2 Create App Roles in Entra ID
|
||||
|
||||
1. Navigate to your App Registration on https://portal.azure.com/
|
||||
2. Go to **App roles** > **Create app role**
|
||||
|
||||
3. Configure the app role:
|
||||
- **Display name**: Proxy Admin (or your preferred display name)
|
||||
- **Value**: `proxy_admin` (use one of the supported role values above)
|
||||
- **Description**: Administrator access to LiteLLM proxy
|
||||
- **Allowed member types**: Users/Groups
|
||||
|
||||
|
||||
4. Click **Apply** to save the role
|
||||
|
||||
### 4.3 Assign Users to App Roles
|
||||
|
||||
1. Navigate to **Enterprise Applications** on https://portal.azure.com/
|
||||
2. Select your LiteLLM application
|
||||
3. Go to **Users and groups** > **Add user/group**
|
||||
4. Select the user and assign them to one of the app roles you created
|
||||
|
||||
|
||||
### 4.4 Test the Role Assignment
|
||||
|
||||
1. Sign in to LiteLLM UI via SSO as a user with an assigned app role
|
||||
2. LiteLLM will automatically extract the app role from the JWT token
|
||||
3. The user will be assigned the corresponding LiteLLM role in the database
|
||||
4. The user's permissions will reflect their assigned role
|
||||
|
||||
**How it works:**
|
||||
- When a user signs in via Microsoft SSO, LiteLLM extracts the `roles` claim from the JWT `id_token`
|
||||
- If any of the roles match a valid LiteLLM role (case-insensitive), that role is assigned to the user
|
||||
- If multiple roles are present, LiteLLM uses the first valid role it finds
|
||||
- This role assignment persists in the LiteLLM database and determines the user's access level
|
||||
|
||||
## Video Walkthrough
|
||||
|
||||
This walks through setting up sso auto-add for **Microsoft Entra ID**
|
||||
|
|
|
|||
BIN
docs/my-website/img/key_r.png
Normal file
|
After Width: | Height: | Size: 125 KiB |
BIN
docs/my-website/img/key_u.png
Normal file
|
After Width: | Height: | Size: 216 KiB |
BIN
docs/my-website/img/mcp_updates.jpg
Normal file
|
After Width: | Height: | Size: 913 KiB |
BIN
docs/my-website/img/oauth_2_success.png
Normal file
|
After Width: | Height: | Size: 67 KiB |
BIN
docs/my-website/img/release_notes/1_78_0_perf.png
Normal file
|
After Width: | Height: | Size: 134 KiB |
BIN
docs/my-website/img/release_notes/tool_control.png
Normal file
|
After Width: | Height: | Size: 798 KiB |
BIN
docs/my-website/img/tag_budget1.png
Normal file
|
After Width: | Height: | Size: 305 KiB |
BIN
docs/my-website/img/tag_budget2.png
Normal file
|
After Width: | Height: | Size: 416 KiB |
|
|
@ -57,19 +57,6 @@ pip install litellm==1.77.5
|
|||
|
||||
---
|
||||
|
||||
### Scheduled Key Rotations
|
||||
|
||||
<Image img={require('../../img/release_notes/schedule_key_rotations.png')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release brings support for scheduling virtual key rotations on LiteLLM AI Gateway.
|
||||
|
||||
This is great for Proxy Admins looking to enforce Enterprise Grade security for use cases going through LiteLLM AI Gateway.
|
||||
|
||||
From this release you can enforce Virtual Keys to rotate on a schedule of your choice e.g every 15 days/30 days/60 days etc.
|
||||
|
||||
---
|
||||
### Performance Improvements - 54% RPS Improvement
|
||||
|
||||
<Image img={require('../../img/release_notes/perf_77_5.png')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
---
|
||||
title: "[Preview] v1.77.7-stable - Claude Sonnet 4.5"
|
||||
title: "v1.77.7-stable - 2.9x Lower Median Latency"
|
||||
slug: "v1-77-7"
|
||||
date: 2025-10-04T10:00:00
|
||||
authors:
|
||||
|
|
@ -15,7 +15,7 @@ authors:
|
|||
title: Backend Performance Engineer
|
||||
url: https://www.linkedin.com/in/alexsander-baptista/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D5603AQGXnziu4kqNCQ/profile-displayphoto-crop_800_800/B56ZkxEcuOKEAI-/0/1757464874550?e=1762387200&v=beta&t=9SNXLsWhx8OnYPAMQ9fqAr02oevDYEAL2vMYg2f9ieg
|
||||
- name: Achintya Srivastava
|
||||
- name: Achintya Rajan
|
||||
title: Fullstack Engineer
|
||||
url: https://www.linkedin.com/in/achintya-rajan/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D5603AQGdkEeyJTdljw/profile-displayphoto-shrink_800_800/profile-displayphoto-shrink_800_800/0/1716271140869?e=1762387200&v=beta&t=9gOoLPeqR2E5z3KSX61EUj3HVZXmgo87vhVuSHeffjc
|
||||
|
|
@ -103,6 +103,31 @@ View the complete configuration: [gist.github.com/AlexsanderHamir/config.yaml](h
|
|||
|
||||
View the complete load testing script: [gist.github.com/AlexsanderHamir/no_cache_hits.py](https://gist.github.com/AlexsanderHamir/42c33d7a4dc7a57f56a78b560dee3a42)
|
||||
|
||||
### MCP OAuth 2.0 Support
|
||||
|
||||
<Image img={require('../../img/mcp_updates.jpg')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release adds support for OAuth 2.0 Client Credentials for MCP servers. This is great for **Internal Dev Tools** use-cases, as it enables your users to call MCP servers, with their own credentials. E.g. Allowing your developers to call the Github MCP, with their own credentials.
|
||||
|
||||
[Set it up today on Claude Code](../../docs/tutorials/claude_responses_api#connecting-mcp-servers)
|
||||
|
||||
### Scheduled Key Rotations
|
||||
|
||||
<Image img={require('../../img/release_notes/schedule_key_rotations.png')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release brings support for scheduling virtual key rotations on LiteLLM AI Gateway.
|
||||
|
||||
From this release you can enforce Virtual Keys to rotate on a schedule of your choice e.g every 15 days/30 days/60 days etc.
|
||||
|
||||
This is great for Proxy Admins who need to enforce security policies for production workloads.
|
||||
|
||||
[Get Started](../../docs/proxy/virtual_keys#scheduled-key-rotations)
|
||||
|
||||
|
||||
## New Models / Updated Models
|
||||
|
||||
#### New Model Support
|
||||
|
|
|
|||
394
docs/my-website/release_notes/v1.78.0-stable/index.md
Normal file
|
|
@ -0,0 +1,394 @@
|
|||
---
|
||||
title: "[Preview] v1.78.0-stable - MCP Gateway: Control Tool Access by Team, Key"
|
||||
slug: "v1-78-0"
|
||||
date: 2025-10-11T10:00:00
|
||||
authors:
|
||||
- name: Krrish Dholakia
|
||||
title: CEO, LiteLLM
|
||||
url: https://www.linkedin.com/in/krish-d/
|
||||
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
|
||||
- name: Ishaan Jaff
|
||||
title: CTO, LiteLLM
|
||||
url: https://www.linkedin.com/in/reffajnaahsi/
|
||||
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
|
||||
- name: Alexsander Hamir
|
||||
title: Backend Performance Engineer
|
||||
url: https://www.linkedin.com/in/alexsander-baptista/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D5603AQGXnziu4kqNCQ/profile-displayphoto-crop_800_800/B56ZkxEcuOKEAI-/0/1757464874550?e=1762387200&v=beta&t=9SNXLsWhx8OnYPAMQ9fqAr02oevDYEAL2vMYg2f9ieg
|
||||
- name: Achintya Rajan
|
||||
title: Fullstack Engineer
|
||||
url: https://www.linkedin.com/in/achintya-rajan/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D5603AQGdkEeyJTdljw/profile-displayphoto-shrink_800_800/profile-displayphoto-shrink_800_800/0/1716271140869?e=1762387200&v=beta&t=9gOoLPeqR2E5z3KSX61EUj3HVZXmgo87vhVuSHeffjc
|
||||
- name: Sameer Kankute
|
||||
title: Backend Engineer (LLM Translation)
|
||||
url: https://www.linkedin.com/in/sameer-kankute/
|
||||
image_url: https://media.licdn.com/dms/image/v2/D4D03AQHB_loQYd5gjg/profile-displayphoto-shrink_800_800/profile-displayphoto-shrink_800_800/0/1719137160975?e=1762387200&v=beta&t=0jbuX-f4eSnDxBY3olI6meuYr-LMbObhFmFbRcKF5mY
|
||||
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
import Image from '@theme/IdealImage';
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
## Deploy this version
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="docker" label="Docker">
|
||||
|
||||
``` showLineNumbers title="docker run litellm"
|
||||
docker run \
|
||||
-e STORE_MODEL_IN_DB=True \
|
||||
-p 4000:4000 \
|
||||
ghcr.io/berriai/litellm:v1.78.0.rc.1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="pip" label="Pip">
|
||||
|
||||
``` showLineNumbers title="pip install litellm"
|
||||
pip install litellm==1.78.0.rc.1
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
---
|
||||
|
||||
## Key Highlights
|
||||
|
||||
- **MCP Gateway - Control Tool Access by Team, Key** - Control MCP tool access by team/key.
|
||||
- **Performance Improvements** - 70% Lower p99 Latency
|
||||
- **GPT-5 Pro & GPT-Image-1-Mini** - Day 0 support for OpenAI's GPT-5 Pro (400K context) and gpt-image-1-mini image generation
|
||||
- **EnkryptAI Guardrails** - New guardrail integration for content moderation
|
||||
- **Tag-Based Budgets** - Support for setting budgets based on request tags
|
||||
|
||||
---
|
||||
|
||||
### MCP Gateway - Control Tool Access by Team, Key
|
||||
|
||||
<Image
|
||||
img={require('../../img/release_notes/tool_control.png')}
|
||||
style={{width: '100%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
<br/>
|
||||
|
||||
Proxy admins can now control MCP tool access by team or key. This makes it easy to grant different teams selective access to tools from the same MCP server.
|
||||
|
||||
For example, you can now give your Engineering team access to `list_repositories`, `create_issue`, and `search_code` tools, while Sales only gets `search_code` and `close_issue` tools.
|
||||
|
||||
This makes it easier for Proxy Admins to govern MCP Tool Access.
|
||||
|
||||
[Get Started](../../docs/mcp_control#set-allowed-tools-for-a-key-team-or-organization)
|
||||
|
||||
---
|
||||
|
||||
## Performance - 70% Lower p99 Latency
|
||||
|
||||
<Image img={require('../../img/release_notes/1_78_0_perf.png')} style={{ width: '800px', height: 'auto' }} />
|
||||
|
||||
<br/>
|
||||
|
||||
This release cuts p99 latency by 70% on LiteLLM AI Gateway, making it even better for low-latency use cases.
|
||||
|
||||
These gains come from two key enhancements:
|
||||
|
||||
**Reliable Sessions**
|
||||
|
||||
Added support for shared sessions with aiohttp. The shared_session parameter is now consistently used across all calls, enabling connection pooling.
|
||||
|
||||
**Faster Routing**
|
||||
|
||||
A new `model_name_to_deployment_indices` hash map replaces O(n) list scans in `_get_all_deployments()` with O(1) hash lookups, boosting routing performance and scalability.
|
||||
|
||||
As a result, performance improved across all latency percentiles:
|
||||
|
||||
- **Median latency:** 110 ms → **100 ms** (−9.1%)
|
||||
- **p95 latency:** 440 ms → **150 ms** (−65.9%)
|
||||
- **p99 latency:** 810 ms → **240 ms** (−70.4%)
|
||||
- **Average latency:** 310 ms → **111.73 ms** (−64.0%)
|
||||
|
||||
### **Test Setup**
|
||||
|
||||
**Locust**
|
||||
|
||||
- **Concurrent users:** 1,000
|
||||
- **Ramp-up:** 500
|
||||
|
||||
**System Specs**
|
||||
|
||||
- **Database was used**
|
||||
- **CPU:** 4 vCPUs
|
||||
- **Memory:** 8 GB RAM
|
||||
- **LiteLLM Workers:** 4
|
||||
- **Instances**: 4
|
||||
|
||||
**Configuration (config.yaml)**
|
||||
|
||||
View the complete configuration: [gist.github.com/AlexsanderHamir/config.yaml](https://gist.github.com/AlexsanderHamir/53f7d554a5d2afcf2c4edb5b6be68ff4)
|
||||
|
||||
**Load Script (no_cache_hits.py)**
|
||||
|
||||
View the complete load testing script: [gist.github.com/AlexsanderHamir/no_cache_hits.py](https://gist.github.com/AlexsanderHamir/42c33d7a4dc7a57f56a78b560dee3a42)
|
||||
|
||||
---
|
||||
|
||||
## New Models / Updated Models
|
||||
|
||||
#### New Model Support
|
||||
|
||||
| Provider | Model | Context Window | Input ($/1M tokens) | Output ($/1M tokens) | Features |
|
||||
| -------- | ----- | -------------- | ------------------- | -------------------- | -------- |
|
||||
| OpenAI | `gpt-5-pro` | 400K | $15.00 | $120.00 | Responses API, reasoning, vision, function calling, prompt caching, web search |
|
||||
| OpenAI | `gpt-5-pro-2025-10-06` | 400K | $15.00 | $120.00 | Responses API, reasoning, vision, function calling, prompt caching, web search |
|
||||
| OpenAI | `gpt-image-1-mini` | - | $2.00/img | - | Image generation and editing |
|
||||
| OpenAI | `gpt-realtime-mini` | 128K | $0.60 | $2.40 | Realtime audio, function calling |
|
||||
| Azure AI | `azure_ai/Phi-4-mini-reasoning` | 131K | $0.08 | $0.32 | Function calling |
|
||||
| Azure AI | `azure_ai/Phi-4-reasoning` | 32K | $0.125 | $0.50 | Function calling, reasoning |
|
||||
| Azure AI | `azure_ai/MAI-DS-R1` | 128K | $1.35 | $5.40 | Reasoning, function calling |
|
||||
| Bedrock | `au.anthropic.claude-sonnet-4-5-20250929-v1:0` | 200K | $3.30 | $16.50 | Chat, reasoning, vision, function calling, prompt caching |
|
||||
| Bedrock | `global.anthropic.claude-sonnet-4-5-20250929-v1:0` | 200K | $3.00 | $15.00 | Chat, reasoning, vision, function calling, prompt caching |
|
||||
| Bedrock | `global.anthropic.claude-sonnet-4-20250514-v1:0` | 1M | $3.00 | $15.00 | Chat, reasoning, vision, function calling, prompt caching |
|
||||
| Bedrock | `cohere.embed-v4:0` | 128K | $0.12 | - | Embeddings, image input support |
|
||||
| OCI | `oci/cohere.command-latest` | 128K | $1.56 | $1.56 | Function calling |
|
||||
| OCI | `oci/cohere.command-a-03-2025` | 256K | $1.56 | $1.56 | Function calling |
|
||||
| OCI | `oci/cohere.command-plus-latest` | 128K | $1.56 | $1.56 | Function calling |
|
||||
| Together AI | `together_ai/moonshotai/Kimi-K2-Instruct-0905` | 262K | $1.00 | $3.00 | Function calling |
|
||||
| Together AI | `together_ai/Qwen/Qwen3-Next-80B-A3B-Instruct` | 262K | $0.15 | $1.50 | Function calling |
|
||||
| Together AI | `together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking` | 262K | $0.15 | $1.50 | Function calling |
|
||||
| Vertex AI | MedGemma models | Varies | Varies | Varies | Medical-focused Gemma models on custom endpoints |
|
||||
| Watson X | 27 new foundation models | Varies | Varies | Varies | Granite, Llama, Mistral families |
|
||||
|
||||
#### Features
|
||||
|
||||
- **[OpenAI](../../docs/providers/openai)**
|
||||
- Add GPT-5 Pro model configuration and documentation - [PR #15258](https://github.com/BerriAI/litellm/pull/15258)
|
||||
- Add stop parameter to non-supported params for GPT-5 - [PR #15244](https://github.com/BerriAI/litellm/pull/15244)
|
||||
- Day 0 Support, Add gpt-image-1-mini - [PR #15259](https://github.com/BerriAI/litellm/pull/15259)
|
||||
- Add gpt-realtime-mini support - [PR #15283](https://github.com/BerriAI/litellm/pull/15283)
|
||||
- Add gpt-5-pro-2025-10-06 to model costs - [PR #15344](https://github.com/BerriAI/litellm/pull/15344)
|
||||
- Minimal fix: gpt5 models should not go on cooldown when called with temperature!=1 - [PR #15330](https://github.com/BerriAI/litellm/pull/15330)
|
||||
|
||||
- **[Snowflake Cortex](../../docs/providers/snowflake)**
|
||||
- Add function calling support for Snowflake Cortex REST API - [PR #15221](https://github.com/BerriAI/litellm/pull/15221)
|
||||
|
||||
- **[Gemini](../../docs/providers/gemini)**
|
||||
- Fix header forwarding for Gemini/Vertex AI providers in proxy mode - [PR #15231](https://github.com/BerriAI/litellm/pull/15231)
|
||||
|
||||
- **[Azure](../../docs/providers/azure)**
|
||||
- Removed stop param from unsupported azure models - [PR #15229](https://github.com/BerriAI/litellm/pull/15229)
|
||||
- Fix(azure/responses): remove invalid status param from azure call - [PR #15253](https://github.com/BerriAI/litellm/pull/15253)
|
||||
- Add new Azure AI models with pricing details - [PR #15387](https://github.com/BerriAI/litellm/pull/15387)
|
||||
- AzureAD Default credentials - select credential type based on environment - [PR #14470](https://github.com/BerriAI/litellm/pull/14470)
|
||||
|
||||
- **[Bedrock](../../docs/providers/bedrock)**
|
||||
- Add Global Cross-Region Inference - [PR #15210](https://github.com/BerriAI/litellm/pull/15210)
|
||||
- Add Cohere Embed v4 support for AWS Bedrock - [PR #15298](https://github.com/BerriAI/litellm/pull/15298)
|
||||
- Fix(bedrock): include cacheWriteInputTokens in prompt_tokens calculation - [PR #15292](https://github.com/BerriAI/litellm/pull/15292)
|
||||
- Add Bedrock AU Cross-Region Inference for Claude Sonnet 4.5 - [PR #15402](https://github.com/BerriAI/litellm/pull/15402)
|
||||
- Converse → /v1/messages streaming doesn't handle parallel tool calls with Claude models - [PR #15315](https://github.com/BerriAI/litellm/pull/15315)
|
||||
|
||||
- **[Vertex AI](../../docs/providers/vertex)**
|
||||
- Implement Context Caching for Vertex AI provider - [PR #15226](https://github.com/BerriAI/litellm/pull/15226)
|
||||
- Support for Vertex AI Gemma Models on Custom Endpoints - [PR #15397](https://github.com/BerriAI/litellm/pull/15397)
|
||||
- VertexAI - gemma model family support (custom endpoints) - [PR #15419](https://github.com/BerriAI/litellm/pull/15419)
|
||||
- VertexAI Gemma model family streaming support + Added MedGemma - [PR #15427](https://github.com/BerriAI/litellm/pull/15427)
|
||||
|
||||
- **[OCI](../../docs/providers/oci)**
|
||||
- Add OCI Cohere support with tool calling and streaming capabilities - [PR #15365](https://github.com/BerriAI/litellm/pull/15365)
|
||||
|
||||
- **[Watson X](../../docs/providers/watsonx)**
|
||||
- Add Watson X foundation model definitions to model_prices_and_context_window.json - [PR #15219](https://github.com/BerriAI/litellm/pull/15219)
|
||||
- Watsonx - Apply correct prompt templates for openai/gpt-oss model family - [PR #15341](https://github.com/BerriAI/litellm/pull/15341)
|
||||
|
||||
- **[OpenRouter](../../docs/providers/openrouter)**
|
||||
- Fix - (openrouter): move cache_control to content blocks for claude/gemini - [PR #15345](https://github.com/BerriAI/litellm/pull/15345)
|
||||
- Fix - OpenRouter cache_control to only apply to last content block - [PR #15395](https://github.com/BerriAI/litellm/pull/15395)
|
||||
|
||||
- **[Together AI](../../docs/providers/togetherai)**
|
||||
- Add new together models - [PR #15383](https://github.com/BerriAI/litellm/pull/15383)
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- **General**
|
||||
- Bug fix: gpt-5-chat-latest has incorrect max_input_tokens value - [PR #15116](https://github.com/BerriAI/litellm/pull/15116)
|
||||
- Fix reasoning response ID - [PR #15265](https://github.com/BerriAI/litellm/pull/15265)
|
||||
- Fix issue with parsing assistant messages - [PR #15320](https://github.com/BerriAI/litellm/pull/15320)
|
||||
- Fix litellm_param based costing - [PR #15336](https://github.com/BerriAI/litellm/pull/15336)
|
||||
- Fix lint errors - [PR #15406](https://github.com/BerriAI/litellm/pull/15406)
|
||||
|
||||
---
|
||||
|
||||
## LLM API Endpoints
|
||||
|
||||
#### Features
|
||||
|
||||
- **[Responses API](../../docs/response_api)**
|
||||
- Added streaming support for response api streaming image generation - [PR #15269](https://github.com/BerriAI/litellm/pull/15269)
|
||||
- Add native Responses API support for litellm_proxy provider - [PR #15347](https://github.com/BerriAI/litellm/pull/15347)
|
||||
- Temporarily relax ResponsesAPIResponse parsing to support custom backends (e.g., vLLM) - [PR #15362](https://github.com/BerriAI/litellm/pull/15362)
|
||||
|
||||
- **[Files API](../../docs/files_api)**
|
||||
- Feat(files): add @client decorator to file operations - [PR #15339](https://github.com/BerriAI/litellm/pull/15339)
|
||||
|
||||
- **[/generateContent](../../docs/providers/gemini)**
|
||||
- Fix gemini cli by actually streaming the response - [PR #15264](https://github.com/BerriAI/litellm/pull/15264)
|
||||
|
||||
- **[Azure Passthrough](../../docs/pass_through/azure)**
|
||||
- Azure - passthrough support with router models - [PR #15240](https://github.com/BerriAI/litellm/pull/15240)
|
||||
|
||||
#### Bugs
|
||||
|
||||
- **General**
|
||||
- Fix x-litellm-cache-key header not being returned on cache hit - [PR #15348](https://github.com/BerriAI/litellm/pull/15348)
|
||||
|
||||
---
|
||||
|
||||
## Management Endpoints / UI
|
||||
|
||||
#### Features
|
||||
|
||||
- **Proxy CLI Auth**
|
||||
- Proxy CLI - dont store existing key in the URL, store it in the state param - [PR #15290](https://github.com/BerriAI/litellm/pull/15290)
|
||||
|
||||
- **Models + Endpoints**
|
||||
- Make PATCH `/model/{model_id}/update` handle `team_id` consistently with POST `/model/new` - [PR #15297](https://github.com/BerriAI/litellm/pull/15297)
|
||||
- Feature: adds Infinity as a provider in the UI - [PR #15285](https://github.com/BerriAI/litellm/pull/15285)
|
||||
- Fix: model + endpoints page crash when config file contains router_settings.model_group_alias - [PR #15308](https://github.com/BerriAI/litellm/pull/15308)
|
||||
- Models & Endpoints Initial Refactor - [PR #15435](https://github.com/BerriAI/litellm/pull/15435)
|
||||
- Litellm UI API Reference page updates - [PR #15438](https://github.com/BerriAI/litellm/pull/15438)
|
||||
|
||||
- **Teams**
|
||||
- Teams page: new column "Your Role" on the teams table - [PR #15384](https://github.com/BerriAI/litellm/pull/15384)
|
||||
- LiteLLM Dashboard Teams UI refactor - [PR #15418](https://github.com/BerriAI/litellm/pull/15418)
|
||||
|
||||
- **UI Infrastructure**
|
||||
- Added prettier to autoformat frontend - [PR #15215](https://github.com/BerriAI/litellm/pull/15215)
|
||||
- Adds turbopack to the npm run dev command in UI to build faster during development - [PR #15250](https://github.com/BerriAI/litellm/pull/15250)
|
||||
- (perf) fix: Replaces bloated key list calls with lean key aliases endpoint - [PR #15252](https://github.com/BerriAI/litellm/pull/15252)
|
||||
- Potentially fixes a UI spasm issue with an expired cookie - [PR #15309](https://github.com/BerriAI/litellm/pull/15309)
|
||||
- LiteLLM UI Refactor Infrastructure - [PR #15236](https://github.com/BerriAI/litellm/pull/15236)
|
||||
- Enforces removal of unused imports from UI - [PR #15416](https://github.com/BerriAI/litellm/pull/15416)
|
||||
- Fix: usage page >> Model Activity >> spend per day graph: y-axis clipping on large spend values - [PR #15389](https://github.com/BerriAI/litellm/pull/15389)
|
||||
- Updates guardrail provider logos - [PR #15421](https://github.com/BerriAI/litellm/pull/15421)
|
||||
|
||||
- **Admin Settings**
|
||||
- Fix: Router settings do not update despite success message - [PR #15249](https://github.com/BerriAI/litellm/pull/15249)
|
||||
- Fix: Prevents DB from accidentally overriding config file values if they are empty in DB - [PR #15340](https://github.com/BerriAI/litellm/pull/15340)
|
||||
|
||||
- **SSO**
|
||||
- SSO - support EntraID app roles - [PR #15351](https://github.com/BerriAI/litellm/pull/15351)
|
||||
|
||||
---
|
||||
|
||||
## Logging / Guardrail / Prompt Management Integrations
|
||||
|
||||
#### Features
|
||||
|
||||
- **[PostHog](../../docs/observability/posthog)**
|
||||
- Feat: posthog per request api key - [PR #15379](https://github.com/BerriAI/litellm/pull/15379)
|
||||
|
||||
#### Guardrails
|
||||
|
||||
- **[EnkryptAI](../../docs/proxy/guardrails)**
|
||||
- Add EnkryptAI Guardrails on LiteLLM - [PR #15390](https://github.com/BerriAI/litellm/pull/15390)
|
||||
|
||||
---
|
||||
|
||||
## Spend Tracking, Budgets and Rate Limiting
|
||||
|
||||
- **Tag Management**
|
||||
- Tag Management - Add support for setting tag based budgets - [PR #15433](https://github.com/BerriAI/litellm/pull/15433)
|
||||
|
||||
- **Dynamic Rate Limiter v3**
|
||||
- QA/Fixes - Dynamic Rate Limiter v3 - final QA - [PR #15311](https://github.com/BerriAI/litellm/pull/15311)
|
||||
- Fix dynamic Rate limiter v3 - inserting litellm_model_saturation - [PR #15394](https://github.com/BerriAI/litellm/pull/15394)
|
||||
|
||||
- **Shared Health Check**
|
||||
- Implement Shared Health Check State Across Pods - [PR #15380](https://github.com/BerriAI/litellm/pull/15380)
|
||||
|
||||
---
|
||||
|
||||
## MCP Gateway
|
||||
|
||||
- **Tool Control**
|
||||
- MCP Gateway - UI - Select allowed tools for Key, Teams - [PR #15241](https://github.com/BerriAI/litellm/pull/15241)
|
||||
- MCP Gateway - Backend - Allow storing allowed tools by team/key - [PR #15243](https://github.com/BerriAI/litellm/pull/15243)
|
||||
- MCP Gateway - Fine-grained Database Object Storage Control - [PR #15255](https://github.com/BerriAI/litellm/pull/15255)
|
||||
- MCP Gateway - Litellm mcp fixes team control - [PR #15304](https://github.com/BerriAI/litellm/pull/15304)
|
||||
- MCP Gateway - QA/Fixes - Ensure Team/Key level enforcement works for MCPs - [PR #15305](https://github.com/BerriAI/litellm/pull/15305)
|
||||
- Feature: Include server_name in /v1/mcp/server/health endpoint response - [PR #15431](https://github.com/BerriAI/litellm/pull/15431)
|
||||
|
||||
- **OpenAPI Integration**
|
||||
- MCP - support converting OpenAPI specs to MCP servers - [PR #15343](https://github.com/BerriAI/litellm/pull/15343)
|
||||
- MCP - specify allowed params per tool - [PR #15346](https://github.com/BerriAI/litellm/pull/15346)
|
||||
|
||||
- **Configuration**
|
||||
- MCP - support setting CA_BUNDLE_PATH - [PR #15253](https://github.com/BerriAI/litellm/pull/15253)
|
||||
- Fix: Ensure MCP client stays open during tool call - [PR #15391](https://github.com/BerriAI/litellm/pull/15391)
|
||||
- Remove hardcoded "public" schema in migration.sql - [PR #15363](https://github.com/BerriAI/litellm/pull/15363)
|
||||
|
||||
---
|
||||
|
||||
## Performance / Loadbalancing / Reliability improvements
|
||||
|
||||
- **Router Optimizations**
|
||||
- Fix - Router: add model_name index for O(1) deployment lookups - [PR #15113](https://github.com/BerriAI/litellm/pull/15113)
|
||||
- Refactor Utils: extract inner function from client - [PR #15234](https://github.com/BerriAI/litellm/pull/15234)
|
||||
- Fix Networking: remove limitations - [PR #15302](https://github.com/BerriAI/litellm/pull/15302)
|
||||
|
||||
- **Session Management**
|
||||
- Fix - Sessions not being shared - [PR #15388](https://github.com/BerriAI/litellm/pull/15388)
|
||||
- Fix: remove panic from hot path - [PR #15396](https://github.com/BerriAI/litellm/pull/15396)
|
||||
- Fix - shared session parsing and usage issue - [PR #15440](https://github.com/BerriAI/litellm/pull/15440)
|
||||
- Fix: handle closed aiohttp sessions - [PR #15442](https://github.com/BerriAI/litellm/pull/15442)
|
||||
- Fix: prevent session leaks when recreating aiohttp sessions - [PR #15443](https://github.com/BerriAI/litellm/pull/15443)
|
||||
|
||||
- **SSL/TLS Performance**
|
||||
- Perf: optimize SSL/TLS handshake performance with prioritized cipher - [PR #15398](https://github.com/BerriAI/litellm/pull/15398)
|
||||
|
||||
- **Dependencies**
|
||||
- Upgrades tenacity version to 8.5.0 - [PR #15303](https://github.com/BerriAI/litellm/pull/15303)
|
||||
|
||||
- **Data Masking**
|
||||
- Fix - SensitiveDataMasker converts lists to string - [PR #15420](https://github.com/BerriAI/litellm/pull/15420)
|
||||
|
||||
---
|
||||
|
||||
|
||||
## General AI Gateway Improvements
|
||||
|
||||
#### Security
|
||||
|
||||
- **General**
|
||||
- Fix: redact AWS credentials when redact_user_api_key_info enabled - [PR #15321](https://github.com/BerriAI/litellm/pull/15321)
|
||||
|
||||
---
|
||||
|
||||
## Documentation Updates
|
||||
|
||||
- **Provider Documentation**
|
||||
- Update doc: perf update - [PR #15211](https://github.com/BerriAI/litellm/pull/15211)
|
||||
- Add W&B Inference documentation - [PR #15278](https://github.com/BerriAI/litellm/pull/15278)
|
||||
|
||||
- **Deployment**
|
||||
- Deletion of docker-compose buggy comment that cause `config.yaml` based startup fail - [PR #15425](https://github.com/BerriAI/litellm/pull/15425)
|
||||
|
||||
---
|
||||
|
||||
## New Contributors
|
||||
|
||||
* @Gal-bloch made their first contribution in [PR #15219](https://github.com/BerriAI/litellm/pull/15219)
|
||||
* @lcfyi made their first contribution in [PR #15315](https://github.com/BerriAI/litellm/pull/15315)
|
||||
* @ashengstd made their first contribution in [PR #15362](https://github.com/BerriAI/litellm/pull/15362)
|
||||
* @vkolehmainen made their first contribution in [PR #15363](https://github.com/BerriAI/litellm/pull/15363)
|
||||
* @jlan-nl made their first contribution in [PR #15330](https://github.com/BerriAI/litellm/pull/15330)
|
||||
* @BCook98 made their first contribution in [PR #15402](https://github.com/BerriAI/litellm/pull/15402)
|
||||
* @PabloGmz96 made their first contribution in [PR #15425](https://github.com/BerriAI/litellm/pull/15425)
|
||||
|
||||
---
|
||||
|
||||
## **[Full Changelog](https://github.com/BerriAI/litellm/compare/v1.77.7.rc.1...v1.78.0.rc.1)**
|
||||
|
||||
|
|
@ -36,6 +36,7 @@ const sidebars = {
|
|||
"proxy/guardrails/aporia_api",
|
||||
"proxy/guardrails/azure_content_guardrail",
|
||||
"proxy/guardrails/bedrock",
|
||||
"proxy/guardrails/enkryptai",
|
||||
"proxy/guardrails/lasso_security",
|
||||
"proxy/guardrails/guardrails_ai",
|
||||
"proxy/guardrails/lakera_ai",
|
||||
|
|
@ -93,10 +94,10 @@ const sidebars = {
|
|||
|
||||
{
|
||||
type: "category",
|
||||
label: "LiteLLM Proxy Server",
|
||||
label: "LiteLLM AI Gateway",
|
||||
link: {
|
||||
type: "generated-index",
|
||||
title: "LiteLLM Proxy Server (LLM Gateway)",
|
||||
title: "LiteLLM AI Gateway (LLM Proxy)",
|
||||
description: `OpenAI Proxy Server (LLM Gateway) to call 100+ LLMs in a unified interface & track spend, set budgets per virtual key/user`,
|
||||
slug: "/simple_proxy",
|
||||
},
|
||||
|
|
@ -188,12 +189,13 @@ const sidebars = {
|
|||
type: "category",
|
||||
label: "Budgets + Rate Limits",
|
||||
items: [
|
||||
"proxy/users",
|
||||
"proxy/team_budgets",
|
||||
"proxy/tag_budgets",
|
||||
"proxy/customers",
|
||||
"proxy/dynamic_rate_limit",
|
||||
"proxy/rate_limit_tiers",
|
||||
"proxy/team_budgets",
|
||||
"proxy/temporary_budget_increase",
|
||||
"proxy/users"
|
||||
],
|
||||
},
|
||||
"proxy/caching",
|
||||
|
|
@ -422,6 +424,7 @@ const sidebars = {
|
|||
items: [
|
||||
"providers/vertex",
|
||||
"providers/vertex_partner",
|
||||
"providers/vertex_self_deployed",
|
||||
"providers/vertex_image",
|
||||
"providers/vertex_batch",
|
||||
]
|
||||
|
|
@ -532,22 +535,17 @@ const sidebars = {
|
|||
"providers/oci",
|
||||
"providers/datarobot",
|
||||
"providers/ovhcloud",
|
||||
"providers/wandb_inference",
|
||||
],
|
||||
},
|
||||
{
|
||||
type: "category",
|
||||
label: "Guides",
|
||||
items: [
|
||||
{
|
||||
type: "category",
|
||||
label: "Tools",
|
||||
items: [
|
||||
"completion/computer_use",
|
||||
"completion/web_search",
|
||||
"completion/web_fetch",
|
||||
"completion/function_call",
|
||||
]
|
||||
},
|
||||
"completion/computer_use",
|
||||
"completion/web_search",
|
||||
"completion/web_fetch",
|
||||
"completion/function_call",
|
||||
"completion/audio",
|
||||
"completion/document_understanding",
|
||||
"completion/drop_params",
|
||||
|
|
|
|||
|
|
@ -36,6 +36,8 @@ async def apply_guardrail(
|
|||
if active_guardrail is None:
|
||||
raise Exception(f"Guardrail {request.guardrail_name} not found")
|
||||
|
||||
return await active_guardrail.apply_guardrail(
|
||||
response_text = await active_guardrail.apply_guardrail(
|
||||
text=request.text, language=request.language, entities=request.entities
|
||||
)
|
||||
|
||||
return ApplyGuardrailResponse(response_text=response_text)
|
||||
|
|
|
|||
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.26-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.2.26.tar.gz
vendored
Normal file
|
|
@ -5,4 +5,4 @@
|
|||
|
||||
*/
|
||||
-- AlterTable
|
||||
ALTER TABLE "public"."LiteLLM_MCPServerTable" DROP COLUMN "spec_version";
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" DROP COLUMN "spec_version";
|
||||
|
|
|
|||
|
|
@ -0,0 +1,18 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_TagTable" (
|
||||
"tag_name" TEXT NOT NULL,
|
||||
"description" TEXT,
|
||||
"models" TEXT[],
|
||||
"model_info" JSONB,
|
||||
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
|
||||
"budget_id" TEXT,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"created_by" TEXT,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_TagTable_pkey" PRIMARY KEY ("tag_name")
|
||||
);
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_TagTable" ADD CONSTRAINT "LiteLLM_TagTable_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
|
||||
|
|
@ -25,6 +25,7 @@ model LiteLLM_BudgetTable {
|
|||
organization LiteLLM_OrganizationTable[] // multiple orgs can have the same budget
|
||||
keys LiteLLM_VerificationToken[] // multiple keys can have the same budget
|
||||
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
|
||||
tags LiteLLM_TagTable[] // multiple tags can have the same budget
|
||||
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
|
||||
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
|
||||
}
|
||||
|
|
@ -245,6 +246,20 @@ model LiteLLM_EndUserTable {
|
|||
blocked Boolean @default(false)
|
||||
}
|
||||
|
||||
// Track tags with budgets and spend
|
||||
model LiteLLM_TagTable {
|
||||
tag_name String @id
|
||||
description String?
|
||||
models String[]
|
||||
model_info Json? // maps model_id to model_name
|
||||
spend Float @default(0.0)
|
||||
budget_id String?
|
||||
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
created_by String?
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
}
|
||||
|
||||
// store proxy config.yaml
|
||||
model LiteLLM_Config {
|
||||
param_name String @id
|
||||
|
|
|
|||
|
|
@ -131,7 +131,9 @@ class ProxyExtrasDBManager:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_all_migrations(migrations_dir: str, schema_path: str):
|
||||
def _resolve_all_migrations(
|
||||
migrations_dir: str, schema_path: str, mark_all_applied: bool = True
|
||||
):
|
||||
"""
|
||||
1. Compare the current database state to schema.prisma and generate a migration for the diff.
|
||||
2. Run prisma migrate deploy to apply any pending migrations.
|
||||
|
|
@ -210,6 +212,8 @@ class ProxyExtrasDBManager:
|
|||
logger.warning("Migration diff application timed out.")
|
||||
|
||||
# 3. Mark all migrations as applied
|
||||
if not mark_all_applied:
|
||||
return
|
||||
migration_names = ProxyExtrasDBManager._get_migration_names(migrations_dir)
|
||||
logger.info(f"Resolving {len(migration_names)} migrations")
|
||||
for migration_name in migration_names:
|
||||
|
|
@ -263,6 +267,13 @@ class ProxyExtrasDBManager:
|
|||
logger.info(f"prisma migrate deploy stdout: {result.stdout}")
|
||||
|
||||
logger.info("prisma migrate deploy completed")
|
||||
|
||||
# Run sanity check to ensure DB matches schema
|
||||
logger.info("Running post-migration sanity check...")
|
||||
ProxyExtrasDBManager._resolve_all_migrations(
|
||||
migrations_dir, schema_path, mark_all_applied=False
|
||||
)
|
||||
logger.info("✅ Post-migration sanity check completed")
|
||||
return True
|
||||
except subprocess.CalledProcessError as e:
|
||||
logger.info(f"prisma db error: {e.stderr}, e: {e.stdout}")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.2.25"
|
||||
version = "0.2.26"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
|
|
@ -22,7 +22,7 @@ requires = ["poetry-core"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.2.25"
|
||||
version = "0.2.26"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -1190,6 +1190,9 @@ from .llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig
|
|||
from .llms.azure.responses.o_series_transformation import (
|
||||
AzureOpenAIOSeriesResponsesAPIConfig,
|
||||
)
|
||||
from .llms.litellm_proxy.responses.transformation import (
|
||||
LiteLLMProxyResponsesAPIConfig,
|
||||
)
|
||||
from .llms.openai.chat.o_series_transformation import (
|
||||
OpenAIOSeriesConfig as OpenAIO1Config, # maintain backwards compatibility
|
||||
OpenAIOSeriesConfig,
|
||||
|
|
|
|||
|
|
@ -177,14 +177,21 @@ def get_redis_url_from_environment():
|
|||
raise ValueError(
|
||||
"Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis."
|
||||
)
|
||||
|
||||
if "REDIS_PASSWORD" in os.environ:
|
||||
redis_password = f":{os.environ['REDIS_PASSWORD']}@"
|
||||
|
||||
if "REDIS_SSL" in os.environ and os.environ["REDIS_SSL"].lower() == "true":
|
||||
redis_protocol = "rediss"
|
||||
else:
|
||||
redis_password = ""
|
||||
|
||||
redis_protocol = "redis"
|
||||
|
||||
# Build authentication part of URL
|
||||
auth_part = ""
|
||||
if "REDIS_USERNAME" in os.environ and "REDIS_PASSWORD" in os.environ:
|
||||
auth_part = f"{os.environ['REDIS_USERNAME']}:{os.environ['REDIS_PASSWORD']}@"
|
||||
elif "REDIS_PASSWORD" in os.environ:
|
||||
auth_part = f"{os.environ['REDIS_PASSWORD']}@"
|
||||
|
||||
return (
|
||||
f"redis://{redis_password}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}"
|
||||
f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -14,10 +14,10 @@ It utilizes the (RedisCache, s3Cache, RedisSemanticCache, QdrantSemanticCache, I
|
|||
In each method it will call the appropriate method from caching.py
|
||||
"""
|
||||
|
||||
import time
|
||||
import asyncio
|
||||
import datetime
|
||||
import inspect
|
||||
import time
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
|
|
@ -62,12 +62,10 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
|
||||
class CachingHandlerResponse(BaseModel):
|
||||
|
|
@ -214,9 +212,7 @@ class LLMCachingHandler:
|
|||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
cache_key = litellm.cache._get_preset_cache_key_from_kwargs(
|
||||
**kwargs
|
||||
)
|
||||
cache_key = litellm.cache.get_cache_key(**kwargs)
|
||||
if (
|
||||
isinstance(cached_result, BaseModel)
|
||||
or isinstance(cached_result, CustomStreamWrapper)
|
||||
|
|
@ -330,9 +326,7 @@ class LLMCachingHandler:
|
|||
end_time=end_time,
|
||||
cache_hit=cache_hit
|
||||
)
|
||||
cache_key = litellm.cache._get_preset_cache_key_from_kwargs(
|
||||
**kwargs
|
||||
)
|
||||
cache_key = litellm.cache.get_cache_key(**kwargs)
|
||||
if (
|
||||
isinstance(cached_result, BaseModel)
|
||||
or isinstance(cached_result, CustomStreamWrapper)
|
||||
|
|
|
|||
|
|
@ -18,13 +18,15 @@ from typing import (
|
|||
cast,
|
||||
)
|
||||
|
||||
from openai.types.responses.tool_param import FunctionToolParam
|
||||
|
||||
from litellm import ModelResponse
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.bridges.completion_transformation import (
|
||||
CompletionTransformationBridge,
|
||||
)
|
||||
from litellm.types.llms.openai import Reasoning
|
||||
from litellm.types.llms.openai import ChatCompletionToolParamFunctionChunk, Reasoning
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from openai.types.responses import ResponseInputImageParam
|
||||
|
|
@ -50,6 +52,45 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
def __init__(self):
|
||||
pass
|
||||
|
||||
def _handle_raw_dict_response_item(
|
||||
self, item: Dict[str, Any], index: int
|
||||
) -> Tuple[Optional[Any], int]:
|
||||
"""
|
||||
Handle raw dict response items from Responses API (e.g., GPT-5 Codex format).
|
||||
|
||||
Args:
|
||||
item: Raw dict response item with 'type' field
|
||||
index: Current choice index
|
||||
|
||||
Returns:
|
||||
Tuple of (Choice object or None, updated index)
|
||||
"""
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
item_type = item.get("type")
|
||||
|
||||
# Ignore reasoning items for now
|
||||
if item_type == "reasoning":
|
||||
return None, index
|
||||
|
||||
# Handle message items with output_text content
|
||||
if item_type == "message":
|
||||
content_list = item.get("content", [])
|
||||
for content_item in content_list:
|
||||
if isinstance(content_item, dict):
|
||||
content_type = content_item.get("type")
|
||||
if content_type == "output_text":
|
||||
response_text = content_item.get("text", "")
|
||||
msg = Message(
|
||||
role=item.get("role", "assistant"),
|
||||
content=response_text if response_text else "",
|
||||
)
|
||||
choice = Choices(message=msg, finish_reason="stop", index=index)
|
||||
return choice, index + 1
|
||||
|
||||
# Unknown or unsupported type
|
||||
return None, index
|
||||
|
||||
def convert_chat_completion_messages_to_responses_api(
|
||||
self, messages: List["AllMessageValues"]
|
||||
) -> Tuple[List[Any], Optional[str]]:
|
||||
|
|
@ -201,6 +242,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if value is not None:
|
||||
if key == "instructions" and instructions:
|
||||
request_data["instructions"] = instructions
|
||||
elif key == "stream_options" and isinstance(value, dict):
|
||||
request_data["stream_options"] = value.get("include_obfuscation")
|
||||
elif key == "user": # string can't be longer than 64 characters
|
||||
if isinstance(value, str) and len(value) <= 64:
|
||||
request_data["user"] = value
|
||||
else:
|
||||
request_data[key] = value
|
||||
|
||||
|
|
@ -221,7 +267,6 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
json_mode: Optional[bool] = None,
|
||||
) -> "ModelResponse":
|
||||
"""Transform Responses API response to chat completion response"""
|
||||
|
||||
from openai.types.responses import (
|
||||
ResponseFunctionToolCall,
|
||||
ResponseOutputMessage,
|
||||
|
|
@ -240,19 +285,35 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
choices: List[Choices] = []
|
||||
index = 0
|
||||
|
||||
reasoning_content: Optional[str] = None
|
||||
|
||||
for item in raw_response.output:
|
||||
|
||||
if isinstance(item, ResponseReasoningItem):
|
||||
pass # ignore for now.
|
||||
|
||||
for summary_item in item.summary:
|
||||
response_text = getattr(summary_item, "text", "")
|
||||
reasoning_content = response_text if response_text else ""
|
||||
|
||||
elif isinstance(item, ResponseOutputMessage):
|
||||
for content in item.content:
|
||||
response_text = getattr(content, "text", "")
|
||||
msg = Message(
|
||||
role=item.role, content=response_text if response_text else ""
|
||||
role=item.role,
|
||||
content=response_text if response_text else "",
|
||||
reasoning_content=reasoning_content,
|
||||
)
|
||||
|
||||
choices.append(
|
||||
Choices(message=msg, finish_reason="stop", index=index)
|
||||
Choices(
|
||||
message=msg,
|
||||
finish_reason="stop",
|
||||
index=index,
|
||||
)
|
||||
)
|
||||
|
||||
reasoning_content = None # flush reasoning content
|
||||
index += 1
|
||||
elif isinstance(item, ResponseFunctionToolCall):
|
||||
msg = Message(
|
||||
|
|
@ -267,12 +328,21 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"type": "function",
|
||||
}
|
||||
],
|
||||
reasoning_content=reasoning_content,
|
||||
)
|
||||
|
||||
choices.append(
|
||||
Choices(message=msg, finish_reason="tool_calls", index=index)
|
||||
)
|
||||
reasoning_content = None # flush reasoning content
|
||||
index += 1
|
||||
elif isinstance(item, dict):
|
||||
# Handle raw dict responses (e.g., from GPT-5 Codex)
|
||||
choice, index = self._handle_raw_dict_response_item(
|
||||
item=item, index=index
|
||||
)
|
||||
if choice is not None:
|
||||
choices.append(choice)
|
||||
else:
|
||||
pass # don't fail request if item in list is not supported
|
||||
|
||||
|
|
@ -447,9 +517,25 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
self, tools: List[Dict[str, Any]]
|
||||
) -> List["ALL_RESPONSES_API_TOOL_PARAMS"]:
|
||||
"""Convert chat completion tools to responses API tools format"""
|
||||
responses_tools = []
|
||||
responses_tools: List["ALL_RESPONSES_API_TOOL_PARAMS"] = []
|
||||
for tool in tools:
|
||||
responses_tools.append(tool)
|
||||
# convert function tool from chat completion to responses API format
|
||||
if tool.get("type") == "function":
|
||||
function_tool = cast(
|
||||
ChatCompletionToolParamFunctionChunk, tool.get("function")
|
||||
)
|
||||
responses_tools.append(
|
||||
FunctionToolParam(
|
||||
name=function_tool["name"],
|
||||
parameters=function_tool.get("parameters"),
|
||||
strict=function_tool.get("strict"),
|
||||
type="function",
|
||||
description=function_tool.get("description"),
|
||||
)
|
||||
)
|
||||
else:
|
||||
responses_tools.append(tool) # type: ignore
|
||||
|
||||
return cast(List["ALL_RESPONSES_API_TOOL_PARAMS"], responses_tools)
|
||||
|
||||
def _map_reasoning_effort(self, reasoning_effort: str) -> Optional[Reasoning]:
|
||||
|
|
|
|||
|
|
@ -87,6 +87,35 @@ MAX_TOKEN_TRIMMING_ATTEMPTS = int(
|
|||
########## Networking constants ##############################################################
|
||||
_DEFAULT_TTL_FOR_HTTPX_CLIENTS = 3600 # 1 hour, re-use the same httpx client for 1 hour
|
||||
|
||||
# Aiohttp connection pooling constants
|
||||
AIOHTTP_CONNECTOR_LIMIT = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 0))
|
||||
AIOHTTP_KEEPALIVE_TIMEOUT = int(os.getenv("AIOHTTP_KEEPALIVE_TIMEOUT", 120))
|
||||
AIOHTTP_TTL_DNS_CACHE = int(os.getenv("AIOHTTP_TTL_DNS_CACHE", 300))
|
||||
|
||||
# SSL/TLS cipher configuration for faster handshakes
|
||||
# Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones
|
||||
# This balances performance with broad compatibility
|
||||
DEFAULT_SSL_CIPHERS = os.getenv(
|
||||
"LITELLM_SSL_CIPHERS",
|
||||
# Priority 1: TLS 1.3 ciphers (fastest, ~50ms handshake)
|
||||
"TLS_AES_256_GCM_SHA384:" # Fastest observed in testing
|
||||
"TLS_AES_128_GCM_SHA256:" # Slightly faster than 256-bit
|
||||
"TLS_CHACHA20_POLY1305_SHA256:" # Fast on ARM/mobile
|
||||
# Priority 2: TLS 1.2 ECDHE+GCM (fast, ~100ms handshake, widely supported)
|
||||
"ECDHE-RSA-AES256-GCM-SHA384:"
|
||||
"ECDHE-RSA-AES128-GCM-SHA256:"
|
||||
"ECDHE-ECDSA-AES256-GCM-SHA384:"
|
||||
"ECDHE-ECDSA-AES128-GCM-SHA256:"
|
||||
# Priority 3: Additional modern ciphers (good balance)
|
||||
"ECDHE-RSA-CHACHA20-POLY1305:"
|
||||
"ECDHE-ECDSA-CHACHA20-POLY1305:"
|
||||
# Priority 4: Widely compatible fallbacks (slower but universally supported)
|
||||
"ECDHE-RSA-AES256-SHA384:" # Common fallback
|
||||
"ECDHE-RSA-AES128-SHA256:" # Very widely supported
|
||||
"AES256-GCM-SHA384:" # Non-PFS fallback (compatibility)
|
||||
"AES128-GCM-SHA256", # Last resort (maximum compatibility)
|
||||
)
|
||||
|
||||
########### v2 Architecture constants for managing writing updates to the database ###########
|
||||
REDIS_UPDATE_BUFFER_KEY = "litellm_spend_update_buffer"
|
||||
REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_spend_update_buffer"
|
||||
|
|
@ -1028,6 +1057,12 @@ PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds
|
|||
DEFAULT_HEALTH_CHECK_INTERVAL = int(
|
||||
os.getenv("DEFAULT_HEALTH_CHECK_INTERVAL", 300)
|
||||
) # 5 minutes
|
||||
DEFAULT_SHARED_HEALTH_CHECK_TTL = int(
|
||||
os.getenv("DEFAULT_SHARED_HEALTH_CHECK_TTL", 300)
|
||||
) # 5 minutes - TTL for cached health check results
|
||||
DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL = int(
|
||||
os.getenv("DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL", 60)
|
||||
) # 1 minute - TTL for health check lock
|
||||
PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS = int(
|
||||
os.getenv("PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS", 9)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -86,8 +86,15 @@ class MCPClient:
|
|||
async def connect(self):
|
||||
"""Initialize the transport and session."""
|
||||
if self._session:
|
||||
verbose_logger.debug(
|
||||
f"MCP client already connected to {self.server_url or 'stdio'}"
|
||||
)
|
||||
return # Already connected
|
||||
|
||||
verbose_logger.info(
|
||||
f"MCP client connecting to {self.server_url or 'stdio'} via {self.transport_type}"
|
||||
)
|
||||
|
||||
try:
|
||||
if self.transport_type == MCPTransport.stdio:
|
||||
# For stdio transport, use stdio_client with command-line parameters
|
||||
|
|
@ -107,6 +114,9 @@ class MCPClient:
|
|||
)
|
||||
self._session = await self._session_ctx.__aenter__()
|
||||
await self._session.initialize()
|
||||
verbose_logger.info(
|
||||
f"MCP client successfully connected via stdio: {self.stdio_config.get('command', '')}"
|
||||
)
|
||||
elif self.transport_type == MCPTransport.sse:
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
|
|
@ -122,6 +132,9 @@ class MCPClient:
|
|||
)
|
||||
self._session = await self._session_ctx.__aenter__()
|
||||
await self._session.initialize()
|
||||
verbose_logger.info(
|
||||
f"MCP client successfully connected via SSE to {self.server_url}"
|
||||
)
|
||||
else: # http
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
|
|
@ -140,6 +153,9 @@ class MCPClient:
|
|||
)
|
||||
self._session = await self._session_ctx.__aenter__()
|
||||
await self._session.initialize()
|
||||
verbose_logger.info(
|
||||
f"MCP client successfully connected via HTTP to {self.server_url}"
|
||||
)
|
||||
except ValueError as e:
|
||||
# Re-raise ValueError exceptions (like missing stdio_config)
|
||||
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
|
||||
|
|
@ -159,7 +175,12 @@ class MCPClient:
|
|||
|
||||
async def disconnect(self):
|
||||
"""Clean up session and connections."""
|
||||
verbose_logger.info(
|
||||
f"MCP client disconnecting from {self.server_url or 'stdio'}"
|
||||
)
|
||||
|
||||
if self._task and not self._task.done():
|
||||
verbose_logger.debug("MCP client cancelling background task")
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
|
|
@ -168,16 +189,24 @@ class MCPClient:
|
|||
|
||||
if self._session:
|
||||
try:
|
||||
verbose_logger.debug("MCP client closing session")
|
||||
await self._session_ctx.__aexit__(None, None, None) # type: ignore
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Error closing MCP session: {type(e).__name__}: {str(e)}"
|
||||
)
|
||||
pass
|
||||
self._session = None
|
||||
self._session_ctx = None
|
||||
|
||||
if self._transport_ctx:
|
||||
try:
|
||||
verbose_logger.debug("MCP client closing transport")
|
||||
await self._transport_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"Error closing MCP transport: {type(e).__name__}: {str(e)}"
|
||||
)
|
||||
pass
|
||||
self._transport_ctx = None
|
||||
self._transport = None
|
||||
|
|
@ -261,25 +290,55 @@ class MCPClient:
|
|||
|
||||
async def list_tools(self) -> List[MCPTool]:
|
||||
"""List available tools from the server."""
|
||||
verbose_logger.debug(
|
||||
f"MCP client listing tools from {self.server_url or 'stdio'}"
|
||||
)
|
||||
|
||||
if not self._session:
|
||||
verbose_logger.debug("MCP client session not found, attempting to connect")
|
||||
try:
|
||||
await self.connect()
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
|
||||
verbose_logger.error(
|
||||
f"MCP client connection failed during list_tools: {type(e).__name__}: {str(e)}"
|
||||
)
|
||||
return []
|
||||
|
||||
if self._session is None:
|
||||
verbose_logger.warning("MCP client session is not initialized")
|
||||
verbose_logger.error(
|
||||
"MCP client session is not initialized after connection attempt"
|
||||
)
|
||||
return []
|
||||
|
||||
try:
|
||||
result = await self._session.list_tools()
|
||||
tool_count = len(result.tools)
|
||||
tool_names = [tool.name for tool in result.tools]
|
||||
verbose_logger.info(
|
||||
f"MCP client listed {tool_count} tools from {self.server_url or 'stdio'}: {tool_names}"
|
||||
)
|
||||
return result.tools
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client list_tools was cancelled")
|
||||
await self.disconnect()
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client list_tools failed: {str(e)}")
|
||||
error_type = type(e).__name__
|
||||
verbose_logger.error(
|
||||
f"MCP client list_tools failed - "
|
||||
f"Error Type: {error_type}, "
|
||||
f"Error: {str(e)}, "
|
||||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream during list_tools - "
|
||||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
)
|
||||
|
||||
await self.disconnect()
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
|
@ -290,17 +349,28 @@ class MCPClient:
|
|||
"""
|
||||
Call an MCP Tool.
|
||||
"""
|
||||
verbose_logger.info(
|
||||
f"MCP client calling tool '{call_tool_request_params.name}' with arguments: {call_tool_request_params.arguments}"
|
||||
)
|
||||
|
||||
if not self._session:
|
||||
verbose_logger.warning(
|
||||
"MCP client session not found, attempting to connect"
|
||||
)
|
||||
try:
|
||||
await self.connect()
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client connection failed: {str(e)}")
|
||||
verbose_logger.error(
|
||||
f"MCP client connection failed before tool call: {type(e).__name__}: {str(e)}"
|
||||
)
|
||||
return MCPCallToolResult(
|
||||
content=[TextContent(type="text", text=f"{str(e)}")], isError=True
|
||||
)
|
||||
|
||||
if self._session is None:
|
||||
verbose_logger.warning("MCP client session is not initialized")
|
||||
verbose_logger.error(
|
||||
"MCP client session is not initialized after connection attempt"
|
||||
)
|
||||
return MCPCallToolResult(
|
||||
content=[
|
||||
TextContent(
|
||||
|
|
@ -310,22 +380,59 @@ class MCPClient:
|
|||
isError=True,
|
||||
)
|
||||
|
||||
# Check session and transport state before calling tool
|
||||
verbose_logger.debug(
|
||||
f"MCP client state before tool call - "
|
||||
f"session: {'active' if self._session else 'none'}, "
|
||||
f"transport: {'active' if self._transport else 'none'}, "
|
||||
f"session_ctx: {'active' if self._session_ctx else 'none'}, "
|
||||
f"transport_ctx: {'active' if self._transport_ctx else 'none'}"
|
||||
)
|
||||
|
||||
try:
|
||||
verbose_logger.debug("MCP client sending tool call to session")
|
||||
tool_result = await self._session.call_tool(
|
||||
name=call_tool_request_params.name,
|
||||
arguments=call_tool_request_params.arguments,
|
||||
)
|
||||
verbose_logger.info(
|
||||
f"MCP client tool call '{call_tool_request_params.name}' completed successfully"
|
||||
)
|
||||
return tool_result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client tool call was cancelled")
|
||||
await self.disconnect()
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"MCP client call_tool failed: {str(e)}")
|
||||
import traceback
|
||||
|
||||
error_trace = traceback.format_exc()
|
||||
verbose_logger.debug(f"MCP client tool call traceback:\n{error_trace}")
|
||||
|
||||
# Log detailed error information
|
||||
error_type = type(e).__name__
|
||||
verbose_logger.error(
|
||||
f"MCP client call_tool failed - "
|
||||
f"Error Type: {error_type}, "
|
||||
f"Error: {str(e)}, "
|
||||
f"Tool: {call_tool_request_params.name}, "
|
||||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream - "
|
||||
"the MCP server may have crashed, disconnected, or timed out. "
|
||||
"Session and transport will be disconnected."
|
||||
)
|
||||
|
||||
await self.disconnect()
|
||||
# Return a default error result instead of raising
|
||||
return MCPCallToolResult(
|
||||
content=[
|
||||
TextContent(type="text", text=f"{str(e)}")
|
||||
TextContent(type="text", text=f"{error_type}: {str(e)}")
|
||||
], # Empty content for error case
|
||||
isError=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm import get_secret_str
|
|||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure.files.handler import AzureOpenAIFilesAPI
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.openai.openai import FileDeleted, FileObject, OpenAIFilesAPI
|
||||
from litellm.llms.vertex_ai.files.handler import VertexAIFilesHandler
|
||||
|
|
@ -268,6 +269,7 @@ def create_file(
|
|||
raise e
|
||||
|
||||
|
||||
@client
|
||||
async def afile_retrieve(
|
||||
file_id: str,
|
||||
custom_llm_provider: Literal["openai", "azure"] = "openai",
|
||||
|
|
@ -308,6 +310,7 @@ async def afile_retrieve(
|
|||
raise e
|
||||
|
||||
|
||||
@client
|
||||
def file_retrieve(
|
||||
file_id: str,
|
||||
custom_llm_provider: Literal["openai", "azure"] = "openai",
|
||||
|
|
@ -422,6 +425,7 @@ def file_retrieve(
|
|||
|
||||
|
||||
# Delete file
|
||||
@client
|
||||
async def afile_delete(
|
||||
file_id: str,
|
||||
custom_llm_provider: Literal["openai", "azure"] = "openai",
|
||||
|
|
@ -462,6 +466,7 @@ async def afile_delete(
|
|||
raise e
|
||||
|
||||
|
||||
@client
|
||||
def file_delete(
|
||||
file_id: str,
|
||||
custom_llm_provider: Literal["openai", "azure"] = "openai",
|
||||
|
|
@ -577,6 +582,7 @@ def file_delete(
|
|||
|
||||
|
||||
# List files
|
||||
@client
|
||||
async def afile_list(
|
||||
custom_llm_provider: Literal["openai", "azure"] = "openai",
|
||||
purpose: Optional[str] = None,
|
||||
|
|
@ -617,6 +623,7 @@ async def afile_list(
|
|||
raise e
|
||||
|
||||
|
||||
@client
|
||||
def file_list(
|
||||
custom_llm_provider: Literal["openai", "azure"] = "openai",
|
||||
purpose: Optional[str] = None,
|
||||
|
|
@ -729,6 +736,7 @@ def file_list(
|
|||
raise e
|
||||
|
||||
|
||||
@client
|
||||
async def afile_content(
|
||||
file_id: str,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
|
||||
|
|
@ -771,6 +779,7 @@ async def afile_content(
|
|||
raise e
|
||||
|
||||
|
||||
@client
|
||||
def file_content(
|
||||
file_id: str,
|
||||
model: Optional[str] = None,
|
||||
|
|
|
|||
|
|
@ -11,11 +11,10 @@ For batching specific details see CustomBatchLogger class
|
|||
|
||||
import asyncio
|
||||
import os
|
||||
from litellm._uuid import uuid
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
_get_httpx_client,
|
||||
|
|
@ -26,7 +25,7 @@ from litellm.types.integrations.posthog import (
|
|||
POSTHOG_MAX_BATCH_SIZE,
|
||||
PostHogEventPayload,
|
||||
)
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
from litellm.types.utils import StandardCallbackDynamicParams, StandardLoggingPayload
|
||||
|
||||
|
||||
class PostHogLogger(CustomBatchLogger):
|
||||
|
|
@ -72,17 +71,21 @@ class PostHogLogger(CustomBatchLogger):
|
|||
verbose_logger.debug(
|
||||
"PostHog: Sync logging - Enters logging function for model %s", kwargs
|
||||
)
|
||||
|
||||
|
||||
api_key, api_url = self._get_credentials_for_request(kwargs)
|
||||
if api_key is None or api_url is None:
|
||||
raise Exception("PostHog credentials not found in kwargs")
|
||||
event_payload = self.create_posthog_event_payload(kwargs)
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
payload = self._create_posthog_payload([event_payload])
|
||||
payload = self._create_posthog_payload([event_payload], api_key)
|
||||
capture_url = f"{api_url.rstrip('/')}/batch/"
|
||||
|
||||
response = self.sync_client.post(
|
||||
url=self.capture_url,
|
||||
url=capture_url,
|
||||
json=payload,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -92,9 +95,9 @@ class PostHogLogger(CustomBatchLogger):
|
|||
raise Exception(
|
||||
f"Response from PostHog API status_code: {response.status_code}, text: {response.text}"
|
||||
)
|
||||
|
||||
|
||||
verbose_logger.debug("PostHog: Sync event successfully sent")
|
||||
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"PostHog Sync Layer Error - {str(e)}")
|
||||
|
||||
|
|
@ -122,9 +125,15 @@ class PostHogLogger(CustomBatchLogger):
|
|||
|
||||
async def _log_async_event(self, kwargs, response_obj=None, start_time=0.0, end_time=0.0):
|
||||
# Note: response_obj, start_time, end_time not used - all data comes from kwargs
|
||||
api_key, api_url = self._get_credentials_for_request(kwargs)
|
||||
event_payload = self.create_posthog_event_payload(kwargs)
|
||||
|
||||
self.log_queue.append(event_payload)
|
||||
# Store event with its credentials for batch sending
|
||||
self.log_queue.append({
|
||||
"event": event_payload,
|
||||
"api_key": api_key,
|
||||
"api_url": api_url
|
||||
})
|
||||
verbose_logger.debug(
|
||||
f"PostHog, event added to queue. Will flush in {self.flush_interval} seconds..."
|
||||
)
|
||||
|
|
@ -257,16 +266,42 @@ class PostHogLogger(CustomBatchLogger):
|
|||
metadata = self._extract_metadata(kwargs)
|
||||
user_id = self._safe_get(metadata, "user_id")
|
||||
if user_id:
|
||||
return str(user_id)
|
||||
return str(user_id)
|
||||
end_user = self._safe_get(standard_logging_object, "end_user")
|
||||
if end_user:
|
||||
return str(end_user)
|
||||
trace_id = self._safe_get(standard_logging_object, "trace_id")
|
||||
if trace_id:
|
||||
return str(trace_id)
|
||||
|
||||
return str(trace_id)
|
||||
|
||||
return self._safe_uuid()
|
||||
|
||||
def _get_credentials_for_request(self, kwargs: Dict[str, Any]) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Get PostHog credentials for this request.
|
||||
|
||||
Checks for per-request credentials in standard_callback_dynamic_params,
|
||||
falls back to instance defaults from environment variables.
|
||||
|
||||
Args:
|
||||
kwargs: Request kwargs containing standard_callback_dynamic_params
|
||||
|
||||
Returns:
|
||||
tuple[str, str]: (api_key, api_url)
|
||||
"""
|
||||
standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = (
|
||||
kwargs.get("standard_callback_dynamic_params", None)
|
||||
)
|
||||
|
||||
if standard_callback_dynamic_params is not None:
|
||||
api_key = standard_callback_dynamic_params.get("posthog_api_key") or self.POSTHOG_API_KEY
|
||||
api_url = standard_callback_dynamic_params.get("posthog_api_url") or self.posthog_host
|
||||
else:
|
||||
api_key = self.POSTHOG_API_KEY
|
||||
api_url = self.posthog_host
|
||||
|
||||
return api_key, api_url
|
||||
|
||||
async def async_send_batch(self):
|
||||
"""
|
||||
Sends the in memory logs queue to PostHog API
|
||||
|
|
@ -282,23 +317,34 @@ class PostHogLogger(CustomBatchLogger):
|
|||
f"PostHog: Sending batch of {len(self.log_queue)} events"
|
||||
)
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
# Group events by credentials for batch sending
|
||||
batches_by_credentials: Dict[tuple[str, str], list] = {}
|
||||
for item in self.log_queue:
|
||||
key = (item["api_key"], item["api_url"])
|
||||
if key not in batches_by_credentials:
|
||||
batches_by_credentials[key] = []
|
||||
batches_by_credentials[key].append(item["event"])
|
||||
|
||||
payload = self._create_posthog_payload(list(self.log_queue))
|
||||
# Send each batch to its respective PostHog instance
|
||||
for (api_key, api_url), events in batches_by_credentials.items():
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
response = await self.async_client.post(
|
||||
url=self.capture_url,
|
||||
json=payload,
|
||||
headers=headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
payload = self._create_posthog_payload(events, api_key)
|
||||
capture_url = f"{api_url.rstrip('/')}/batch/"
|
||||
|
||||
if response.status_code != 200:
|
||||
raise Exception(
|
||||
f"Response from PostHog API status_code: {response.status_code}, text: {response.text}"
|
||||
response = await self.async_client.post(
|
||||
url=capture_url,
|
||||
json=payload,
|
||||
headers=headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
if response.status_code != 200:
|
||||
raise Exception(
|
||||
f"Response from PostHog API status_code: {response.status_code}, text: {response.text}"
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"PostHog: Batch of {len(self.log_queue)} events successfully sent"
|
||||
|
|
@ -324,8 +370,8 @@ class PostHogLogger(CustomBatchLogger):
|
|||
def _safe_uuid(self) -> str:
|
||||
return str(uuid.uuid4())
|
||||
|
||||
def _create_posthog_payload(self, events: list) -> Dict[str, Any]:
|
||||
return {"api_key": self.POSTHOG_API_KEY, "batch": events}
|
||||
def _create_posthog_payload(self, events: list, api_key: str) -> Dict[str, Any]:
|
||||
return {"api_key": api_key, "batch": events}
|
||||
|
||||
def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
|
||||
if obj is None or not hasattr(obj, 'get'):
|
||||
|
|
|
|||
|
|
@ -81,12 +81,12 @@ from litellm.types.llms.openai import (
|
|||
)
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
from litellm.types.rerank import RerankResponse
|
||||
from litellm.types.router import CustomPricingLiteLLMParams
|
||||
from litellm.types.utils import (
|
||||
CachingDetails,
|
||||
CallTypes,
|
||||
CostBreakdown,
|
||||
CostResponseTypes,
|
||||
CustomPricingLiteLLMParams,
|
||||
DynamicPromptManagementParamLiteral,
|
||||
EmbeddingResponse,
|
||||
GuardrailStatus,
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import copy
|
|||
import json
|
||||
import mimetypes
|
||||
import re
|
||||
from litellm._uuid import uuid
|
||||
import xml.etree.ElementTree as ET
|
||||
from enum import Enum
|
||||
from typing import Any, List, Optional, Tuple, cast, overload
|
||||
|
|
@ -13,6 +12,7 @@ import litellm
|
|||
import litellm.types
|
||||
import litellm.types.llms
|
||||
from litellm import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
|
||||
from litellm.types.files import get_file_extension_from_mime_type
|
||||
from litellm.types.llms.anthropic import *
|
||||
|
|
@ -232,7 +232,6 @@ def ollama_pt(
|
|||
## MERGE CONSECUTIVE ASSISTANT CONTENT ##
|
||||
while msg_i < len(messages) and messages[msg_i]["role"] == "assistant":
|
||||
assistant_content_str += convert_content_list_to_str(messages[msg_i])
|
||||
msg_i += 1
|
||||
|
||||
tool_calls = messages[msg_i].get("tool_calls")
|
||||
ollama_tool_calls = []
|
||||
|
|
@ -258,7 +257,7 @@ def ollama_pt(
|
|||
f"Tool Calls: {json.dumps(ollama_tool_calls, indent=2)}"
|
||||
)
|
||||
|
||||
msg_i += 1
|
||||
msg_i += 1
|
||||
|
||||
if assistant_content_str:
|
||||
prompt += f"### Assistant:\n{assistant_content_str}\n\n"
|
||||
|
|
@ -365,62 +364,20 @@ def phind_codellama_pt(messages):
|
|||
return prompt
|
||||
|
||||
|
||||
def hf_chat_template( # noqa: PLR0915
|
||||
model: str, messages: list, chat_template: Optional[Any] = None
|
||||
):
|
||||
# Define Jinja2 environment
|
||||
env = ImmutableSandboxedEnvironment()
|
||||
|
||||
def raise_exception(message):
|
||||
raise Exception(f"Error message - {message}")
|
||||
|
||||
# Create a template object from the template text
|
||||
env.globals["raise_exception"] = raise_exception
|
||||
|
||||
## get the tokenizer config from huggingface
|
||||
bos_token = ""
|
||||
eos_token = ""
|
||||
if chat_template is None:
|
||||
|
||||
def _get_tokenizer_config(hf_model_name):
|
||||
try:
|
||||
url = f"https://huggingface.co/{hf_model_name}/raw/main/tokenizer_config.json"
|
||||
# Make a GET request to fetch the JSON data
|
||||
client = HTTPHandler(concurrent_limit=1)
|
||||
|
||||
response = client.get(url)
|
||||
except Exception as e:
|
||||
raise e
|
||||
if response.status_code == 200:
|
||||
# Parse the JSON data
|
||||
tokenizer_config = json.loads(response.content)
|
||||
return {"status": "success", "tokenizer": tokenizer_config}
|
||||
else:
|
||||
return {"status": "failure"}
|
||||
|
||||
if model in litellm.known_tokenizer_config:
|
||||
tokenizer_config = litellm.known_tokenizer_config[model]
|
||||
else:
|
||||
tokenizer_config = _get_tokenizer_config(model)
|
||||
litellm.known_tokenizer_config.update({model: tokenizer_config})
|
||||
|
||||
if (
|
||||
tokenizer_config["status"] == "failure"
|
||||
or "chat_template" not in tokenizer_config["tokenizer"]
|
||||
):
|
||||
raise Exception("No chat template found")
|
||||
## read the bos token, eos token and chat template from the json
|
||||
tokenizer_config = tokenizer_config["tokenizer"] # type: ignore
|
||||
|
||||
bos_token = tokenizer_config["bos_token"] # type: ignore
|
||||
if bos_token is not None and not isinstance(bos_token, str):
|
||||
if isinstance(bos_token, dict):
|
||||
bos_token = bos_token.get("content", None)
|
||||
eos_token = tokenizer_config["eos_token"] # type: ignore
|
||||
if eos_token is not None and not isinstance(eos_token, str):
|
||||
if isinstance(eos_token, dict):
|
||||
eos_token = eos_token.get("content", None)
|
||||
chat_template = tokenizer_config["chat_template"] # type: ignore
|
||||
def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: str, messages: list) -> str:
|
||||
"""
|
||||
Shared template rendering logic for both sync and async hf_chat_template
|
||||
|
||||
Args:
|
||||
env: Jinja2 environment
|
||||
chat_template: Chat template string
|
||||
bos_token: Beginning of sequence token
|
||||
eos_token: End of sequence token
|
||||
messages: Messages to render
|
||||
|
||||
Returns:
|
||||
Rendered template string
|
||||
"""
|
||||
try:
|
||||
template = env.from_string(chat_template) # type: ignore
|
||||
except Exception as e:
|
||||
|
|
@ -435,7 +392,6 @@ def hf_chat_template( # noqa: PLR0915
|
|||
bos_token="<bos>",
|
||||
)
|
||||
return True
|
||||
|
||||
# This will be raised if Jinja attempts to render the system message and it can't
|
||||
except Exception:
|
||||
return False
|
||||
|
|
@ -469,7 +425,7 @@ def hf_chat_template( # noqa: PLR0915
|
|||
)
|
||||
except Exception as e:
|
||||
if "Conversation roles must alternate user/assistant" in str(e):
|
||||
# reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, add a blank 'user' or 'assistant' message to ensure compatibility
|
||||
# reformat messages to ensure user/assistant are alternating
|
||||
new_messages = []
|
||||
for i in range(len(reformatted_messages) - 1):
|
||||
new_messages.append(reformatted_messages[i])
|
||||
|
|
@ -495,6 +451,188 @@ def hf_chat_template( # noqa: PLR0915
|
|||
) # don't use verbose_logger.exception, if exception is raised
|
||||
|
||||
|
||||
async def _afetch_and_extract_template(
|
||||
model: str, chat_template: Optional[Any], get_config_fn, get_template_fn
|
||||
) -> Tuple[str, str, str]:
|
||||
"""
|
||||
Async version: Fetch template and tokens from HuggingFace.
|
||||
|
||||
Returns: (chat_template, bos_token, eos_token)
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
|
||||
_extract_token_value,
|
||||
)
|
||||
|
||||
bos_token = ""
|
||||
eos_token = ""
|
||||
|
||||
if chat_template is None:
|
||||
# Fetch or retrieve cached tokenizer config
|
||||
if model in litellm.known_tokenizer_config:
|
||||
tokenizer_config = litellm.known_tokenizer_config[model]
|
||||
else:
|
||||
tokenizer_config = await get_config_fn(hf_model_name=model)
|
||||
litellm.known_tokenizer_config.update({model: tokenizer_config})
|
||||
|
||||
# Try to get chat template from tokenizer_config.json first
|
||||
if (
|
||||
tokenizer_config.get("status") == "success"
|
||||
and "tokenizer" in tokenizer_config
|
||||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
and "chat_template" in tokenizer_config["tokenizer"]
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
bos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("bos_token")
|
||||
)
|
||||
eos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("eos_token")
|
||||
)
|
||||
chat_template = tokenizer_data["chat_template"]
|
||||
else:
|
||||
# Fallback: Try to fetch chat template from separate .jinja file
|
||||
template_result = await get_template_fn(hf_model_name=model)
|
||||
if template_result.get("status") == "success":
|
||||
chat_template = template_result["chat_template"]
|
||||
# Still try to get tokens from tokenizer_config if available
|
||||
if (
|
||||
tokenizer_config.get("status") == "success"
|
||||
and "tokenizer" in tokenizer_config
|
||||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
bos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("bos_token")
|
||||
)
|
||||
eos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("eos_token")
|
||||
)
|
||||
else:
|
||||
raise Exception("No chat template found")
|
||||
|
||||
return chat_template, bos_token, eos_token # type: ignore
|
||||
|
||||
|
||||
def _fetch_and_extract_template(
|
||||
model: str, chat_template: Optional[Any], get_config_fn, get_template_fn
|
||||
) -> Tuple[str, str, str]:
|
||||
"""
|
||||
Sync version: Fetch template and tokens from HuggingFace.
|
||||
|
||||
Returns: (chat_template, bos_token, eos_token)
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
|
||||
_extract_token_value,
|
||||
)
|
||||
|
||||
bos_token = ""
|
||||
eos_token = ""
|
||||
|
||||
if chat_template is None:
|
||||
# Fetch or retrieve cached tokenizer config
|
||||
if model in litellm.known_tokenizer_config:
|
||||
tokenizer_config = litellm.known_tokenizer_config[model]
|
||||
else:
|
||||
tokenizer_config = get_config_fn(hf_model_name=model)
|
||||
litellm.known_tokenizer_config.update({model: tokenizer_config})
|
||||
|
||||
# Try to get chat template from tokenizer_config.json first
|
||||
if (
|
||||
tokenizer_config.get("status") == "success"
|
||||
and "tokenizer" in tokenizer_config
|
||||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
and "chat_template" in tokenizer_config["tokenizer"]
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
bos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("bos_token")
|
||||
)
|
||||
eos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("eos_token")
|
||||
)
|
||||
chat_template = tokenizer_data["chat_template"]
|
||||
else:
|
||||
# Fallback: Try to fetch chat template from separate .jinja file
|
||||
template_result = get_template_fn(hf_model_name=model)
|
||||
if template_result.get("status") == "success":
|
||||
chat_template = template_result["chat_template"]
|
||||
# Still try to get tokens from tokenizer_config if available
|
||||
if (
|
||||
tokenizer_config.get("status") == "success"
|
||||
and "tokenizer" in tokenizer_config
|
||||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
bos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("bos_token")
|
||||
)
|
||||
eos_token = _extract_token_value(
|
||||
token_value=tokenizer_data.get("eos_token")
|
||||
)
|
||||
else:
|
||||
raise Exception("No chat template found")
|
||||
|
||||
return chat_template, bos_token, eos_token # type: ignore
|
||||
|
||||
|
||||
async def ahf_chat_template(
|
||||
model: str, messages: list, chat_template: Optional[Any] = None
|
||||
):
|
||||
"""HuggingFace chat template (async version)"""
|
||||
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
|
||||
_aget_chat_template_file,
|
||||
_aget_tokenizer_config,
|
||||
strftime_now,
|
||||
)
|
||||
|
||||
env = ImmutableSandboxedEnvironment()
|
||||
env.globals["raise_exception"] = lambda msg: Exception(f"Error message - {msg}")
|
||||
env.globals["strftime_now"] = strftime_now
|
||||
|
||||
template, bos_token, eos_token = await _afetch_and_extract_template(
|
||||
model=model,
|
||||
chat_template=chat_template,
|
||||
get_config_fn=_aget_tokenizer_config,
|
||||
get_template_fn=_aget_chat_template_file,
|
||||
)
|
||||
return _render_chat_template(
|
||||
env=env,
|
||||
chat_template=template,
|
||||
bos_token=bos_token,
|
||||
eos_token=eos_token,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
|
||||
def hf_chat_template(
|
||||
model: str, messages: list, chat_template: Optional[Any] = None
|
||||
):
|
||||
"""HuggingFace chat template (sync version)"""
|
||||
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
|
||||
_get_chat_template_file,
|
||||
_get_tokenizer_config,
|
||||
strftime_now,
|
||||
)
|
||||
|
||||
env = ImmutableSandboxedEnvironment()
|
||||
env.globals["raise_exception"] = lambda msg: Exception(f"Error message - {msg}")
|
||||
env.globals["strftime_now"] = strftime_now
|
||||
|
||||
template, bos_token, eos_token = _fetch_and_extract_template(
|
||||
model=model,
|
||||
chat_template=chat_template,
|
||||
get_config_fn=_get_tokenizer_config,
|
||||
get_template_fn=_get_chat_template_file,
|
||||
)
|
||||
return _render_chat_template(
|
||||
env=env,
|
||||
chat_template=template,
|
||||
bos_token=bos_token,
|
||||
eos_token=eos_token,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
|
||||
def deepseek_r1_pt(messages):
|
||||
return hf_chat_template(
|
||||
model="deepseek-r1/deepseek-r1-7b-instruct", messages=messages
|
||||
|
|
@ -4032,33 +4170,9 @@ def prompt_factory(
|
|||
elif custom_llm_provider == "azure_text":
|
||||
return azure_text_pt(messages=messages)
|
||||
elif custom_llm_provider == "watsonx":
|
||||
if "granite" in model and "chat" in model:
|
||||
# granite-13b-chat-v1 and granite-13b-chat-v2 use a specific prompt template
|
||||
return ibm_granite_pt(messages=messages)
|
||||
elif "ibm-mistral" in model and "instruct" in model:
|
||||
# models like ibm-mistral/mixtral-8x7b-instruct-v01-q use the mistral instruct prompt template
|
||||
return mistral_instruct_pt(messages=messages)
|
||||
elif "meta-llama/llama-3" in model and "instruct" in model:
|
||||
# https://llama.meta.com/docs/model-cards-and-prompt-formats/meta-llama-3/
|
||||
return custom_prompt(
|
||||
role_dict={
|
||||
"system": {
|
||||
"pre_message": "<|start_header_id|>system<|end_header_id|>\n",
|
||||
"post_message": "<|eot_id|>",
|
||||
},
|
||||
"user": {
|
||||
"pre_message": "<|start_header_id|>user<|end_header_id|>\n",
|
||||
"post_message": "<|eot_id|>",
|
||||
},
|
||||
"assistant": {
|
||||
"pre_message": "<|start_header_id|>assistant<|end_header_id|>\n",
|
||||
"post_message": "<|eot_id|>",
|
||||
},
|
||||
},
|
||||
messages=messages,
|
||||
initial_prompt_value="<|begin_of_text|>",
|
||||
final_prompt_value="<|start_header_id|>assistant<|end_header_id|>\n",
|
||||
)
|
||||
from litellm.llms.watsonx.chat.transformation import IBMWatsonXChatConfig
|
||||
return IBMWatsonXChatConfig.apply_prompt_template(model=model, messages=messages)
|
||||
|
||||
try:
|
||||
if "meta-llama/llama-2" in model and "chat" in model:
|
||||
return llama_2_chat_pt(messages=messages)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,139 @@
|
|||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Union
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
_get_httpx_client,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
|
||||
def strftime_now(fmt: str) -> str:
|
||||
"""
|
||||
Custom function for templates that need current date/time formatting (e.g., gpt-oss)
|
||||
|
||||
Args:
|
||||
fmt: Format string for datetime.now().strftime()
|
||||
|
||||
Returns:
|
||||
Formatted string
|
||||
"""
|
||||
return datetime.now().strftime(fmt)
|
||||
|
||||
|
||||
def _get_tokenizer_config(hf_model_name: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Fetch tokenizer_config.json from HuggingFace (sync)
|
||||
|
||||
Args:
|
||||
hf_model_name: HuggingFace model name (e.g., 'openai/gpt-oss-120b')
|
||||
|
||||
Returns:
|
||||
Dict with 'status' and optionally 'tokenizer' keys
|
||||
"""
|
||||
try:
|
||||
url = f"https://huggingface.co/{hf_model_name}/raw/main/tokenizer_config.json"
|
||||
client = _get_httpx_client()
|
||||
response = client.get(url=url)
|
||||
except Exception as e:
|
||||
raise e
|
||||
if response.status_code == 200:
|
||||
tokenizer_config = json.loads(response.content)
|
||||
return {"status": "success", "tokenizer": tokenizer_config}
|
||||
else:
|
||||
return {"status": "failure"}
|
||||
|
||||
|
||||
async def _aget_tokenizer_config(hf_model_name: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Fetch tokenizer_config.json from HuggingFace (async)
|
||||
|
||||
Args:
|
||||
hf_model_name: HuggingFace model name (e.g., 'openai/gpt-oss-120b')
|
||||
|
||||
Returns:
|
||||
Dict with 'status' and optionally 'tokenizer' keys
|
||||
"""
|
||||
try:
|
||||
url = f"https://huggingface.co/{hf_model_name}/raw/main/tokenizer_config.json"
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.PromptFactory,
|
||||
)
|
||||
response = await client.get(url=url)
|
||||
except Exception as e:
|
||||
raise e
|
||||
if response.status_code == 200:
|
||||
tokenizer_config = json.loads(response.content)
|
||||
return {"status": "success", "tokenizer": tokenizer_config}
|
||||
else:
|
||||
return {"status": "failure"}
|
||||
|
||||
|
||||
def _get_chat_template_file(hf_model_name: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Fetch chat template from separate .jinja file (sync)
|
||||
|
||||
Args:
|
||||
hf_model_name: HuggingFace model name (e.g., 'openai/gpt-oss-120b')
|
||||
|
||||
Returns:
|
||||
Dict with 'status' and optionally 'chat_template' keys
|
||||
"""
|
||||
template_filenames = ["chat_template.jinja", "chat_template.jinja2"]
|
||||
client = _get_httpx_client()
|
||||
|
||||
for filename in template_filenames:
|
||||
try:
|
||||
url = f"https://huggingface.co/{hf_model_name}/raw/main/{filename}"
|
||||
response = client.get(url=url)
|
||||
if response.status_code == 200:
|
||||
return {"status": "success", "chat_template": response.content.decode("utf-8")}
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
return {"status": "failure"}
|
||||
|
||||
|
||||
async def _aget_chat_template_file(hf_model_name: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Fetch chat template from separate .jinja file (async)
|
||||
|
||||
Args:
|
||||
hf_model_name: HuggingFace model name (e.g., 'openai/gpt-oss-120b')
|
||||
|
||||
Returns:
|
||||
Dict with 'status' and optionally 'chat_template' keys
|
||||
"""
|
||||
template_filenames = ["chat_template.jinja", "chat_template.jinja2"]
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.PromptFactory,
|
||||
)
|
||||
|
||||
for filename in template_filenames:
|
||||
try:
|
||||
url = f"https://huggingface.co/{hf_model_name}/raw/main/{filename}"
|
||||
response = await client.get(url=url)
|
||||
if response.status_code == 200:
|
||||
return {"status": "success", "chat_template": response.content.decode("utf-8")}
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
return {"status": "failure"}
|
||||
|
||||
|
||||
def _extract_token_value(token_value: Union[None, str, Dict[str, Any]]) -> str:
|
||||
"""
|
||||
Extract token string from various formats (string, dict, etc.)
|
||||
|
||||
Args:
|
||||
token_value: Token value in various formats (None, str, or dict with 'content' key)
|
||||
|
||||
Returns:
|
||||
Extracted token string
|
||||
"""
|
||||
if token_value is None or isinstance(token_value, str):
|
||||
return token_value or ""
|
||||
if isinstance(token_value, dict):
|
||||
return token_value.get("content", "")
|
||||
return ""
|
||||
|
|
@ -75,7 +75,7 @@ class SensitiveDataMasker:
|
|||
masked_data[k] = self._mask_value(str_value)
|
||||
else:
|
||||
masked_data[k] = (
|
||||
v if isinstance(v, (int, float, bool, str)) else str(v)
|
||||
v if isinstance(v, (int, float, bool, str, list)) else str(v)
|
||||
)
|
||||
except Exception:
|
||||
masked_data[k] = "<unable to serialize>"
|
||||
|
|
@ -89,12 +89,14 @@ masker = SensitiveDataMasker()
|
|||
data = {
|
||||
"api_key": "sk-1234567890abcdef",
|
||||
"redis_password": "very_secret_pass",
|
||||
"port": 6379
|
||||
"port": 6379,
|
||||
"tags": ["East US 2", "production", "test"]
|
||||
}
|
||||
masked = masker.mask_dict(data)
|
||||
# Result: {
|
||||
# "api_key": "sk-1****cdef",
|
||||
# "redis_password": "very****pass",
|
||||
# "port": 6379
|
||||
# "port": 6379,
|
||||
# "tags": ["East US 2", "production", "test"]
|
||||
# }
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -133,7 +133,6 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
**kwargs,
|
||||
) -> Union[AnthropicMessagesResponse, AsyncIterator]:
|
||||
"""Handle non-Anthropic models asynchronously using the adapter"""
|
||||
|
||||
completion_kwargs = (
|
||||
LiteLLMMessagesToCompletionTransformationHandler._prepare_completion_kwargs(
|
||||
max_tokens=max_tokens,
|
||||
|
|
|
|||
|
|
@ -270,7 +270,6 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
processed_chunk.get("delta", {}).get("stop_reason")
|
||||
is not None
|
||||
):
|
||||
|
||||
self.holding_stop_reason_chunk = processed_chunk
|
||||
else:
|
||||
self.chunk_queue.append(processed_chunk)
|
||||
|
|
@ -380,4 +379,11 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
self.current_content_block_start = content_block_start
|
||||
return True
|
||||
|
||||
# For parallel tool calls, we'll necessarily have a new content block
|
||||
# if we get a function name since it signals a new tool call
|
||||
if block_type == "tool_use" and content_block_start.get("name"):
|
||||
self.current_content_block_type = block_type
|
||||
self.current_content_block_start = content_block_start
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -365,6 +365,11 @@ def get_azure_ad_token(
|
|||
azure_ad_token_provider = get_azure_ad_token_provider(azure_scope=scope)
|
||||
except ValueError:
|
||||
verbose_logger.debug("Azure AD Token Provider could not be used.")
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Error calling Azure AD token provider: {str(e)}. Follow docs - https://docs.litellm.ai/docs/providers/azure/#azure-ad-token-refresh---defaultazurecredential"
|
||||
)
|
||||
raise e
|
||||
|
||||
#########################################################
|
||||
# If litellm.enable_azure_ad_token_refresh is True and no other token provider is available,
|
||||
|
|
@ -561,7 +566,9 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
"Using Azure AD token provider based on Service Principal with Secret workflow for Azure Auth"
|
||||
)
|
||||
try:
|
||||
azure_ad_token_provider = get_azure_ad_token_provider(azure_scope=scope)
|
||||
azure_ad_token_provider = get_azure_ad_token_provider(
|
||||
azure_scope=scope,
|
||||
)
|
||||
except ValueError:
|
||||
verbose_logger.debug("Azure AD Token Provider could not be used.")
|
||||
if api_version is None:
|
||||
|
|
@ -665,6 +672,10 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
) -> dict:
|
||||
litellm_params = litellm_params or GenericLiteLLMParams()
|
||||
|
||||
# Check if api-key is already in headers; if so, use it
|
||||
if "api-key" in headers:
|
||||
return headers
|
||||
|
||||
api_key = (
|
||||
litellm_params.api_key
|
||||
or litellm.api_key
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
from openai.types.responses import ResponseReasoningItem
|
||||
|
|
@ -41,12 +41,10 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
|
||||
def _handle_reasoning_item(self, item: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle reasoning items specifically to filter out status=None using OpenAI's model.
|
||||
Handle reasoning items to filter out the status field.
|
||||
Issue: https://github.com/BerriAI/litellm/issues/13484
|
||||
OpenAI API does not accept ReasoningItem(status=None), so we need to:
|
||||
1. Check if the item is a reasoning type
|
||||
2. Create a ResponseReasoningItem object with the item data
|
||||
3. Convert it back to dict with exclude_none=True to filter None values
|
||||
|
||||
Azure OpenAI API does not accept 'status' field in reasoning input items.
|
||||
"""
|
||||
if item.get("type") == "reasoning":
|
||||
try:
|
||||
|
|
@ -82,6 +80,32 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
}
|
||||
return filtered_item
|
||||
return item
|
||||
|
||||
def _validate_input_param(
|
||||
self, input: Union[str, ResponseInputParam]
|
||||
) -> Union[str, ResponseInputParam]:
|
||||
"""
|
||||
Override parent method to also filter out 'status' field from message items.
|
||||
Azure OpenAI API does not accept 'status' field in input messages.
|
||||
"""
|
||||
from typing import cast
|
||||
|
||||
# First call parent's validation
|
||||
validated_input = super()._validate_input_param(input)
|
||||
|
||||
# Then filter out status from message items
|
||||
if isinstance(validated_input, list):
|
||||
filtered_input: List[Any] = []
|
||||
for item in validated_input:
|
||||
if isinstance(item, dict) and item.get("type") == "message":
|
||||
# Filter out status field from message items
|
||||
filtered_item = {k: v for k, v in item.items() if k != "status"}
|
||||
filtered_input.append(filtered_item)
|
||||
else:
|
||||
filtered_input.append(item)
|
||||
return cast(ResponseInputParam, filtered_input)
|
||||
|
||||
return validated_input
|
||||
|
||||
def transform_responses_api_request(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -210,6 +210,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
client=None,
|
||||
aembedding=None,
|
||||
max_retries: Optional[int] = None,
|
||||
shared_session=None,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
- Separate image url from text
|
||||
|
|
@ -275,6 +276,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
else None
|
||||
),
|
||||
aembedding=aembedding,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
text_embedding_responses = response.data
|
||||
|
|
|
|||
|
|
@ -440,7 +440,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
"""
|
||||
Abbreviations of regions AWS Bedrock supports for cross region inference
|
||||
"""
|
||||
return ["global", "us", "eu", "apac", "jp"]
|
||||
return ["global", "us", "eu", "apac", "jp", "au"]
|
||||
|
||||
@staticmethod
|
||||
def get_bedrock_route(
|
||||
|
|
|
|||
|
|
@ -152,6 +152,16 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
|
||||
# If we don't have a client or it's not a ClientSession, create one
|
||||
if not isinstance(self.client, ClientSession):
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
# Don't return yet - check if the newly created session is valid
|
||||
|
||||
# Check if the session itself is closed
|
||||
if self.client.closed:
|
||||
verbose_logger.debug("Session is closed, creating new session")
|
||||
# Create a new session
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
|
|
@ -169,14 +179,17 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
or session_loop != current_loop
|
||||
or session_loop.is_closed()
|
||||
):
|
||||
# Clean up the old session
|
||||
# Close old session to prevent leaks
|
||||
old_session = self.client
|
||||
try:
|
||||
# Note: not awaiting close() here as it might be from a different loop
|
||||
# The session will be garbage collected
|
||||
pass
|
||||
if not old_session.closed:
|
||||
try:
|
||||
asyncio.create_task(old_session.close())
|
||||
except RuntimeError:
|
||||
# Different event loop - can't schedule task, rely on GC
|
||||
verbose_logger.debug("Old session from different loop, relying on GC")
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error closing old session: {e}")
|
||||
pass
|
||||
|
||||
# Create a new session in the current event loop
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
|
|
@ -193,13 +206,58 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
|
||||
return self.client
|
||||
|
||||
async def _make_aiohttp_request(
|
||||
self,
|
||||
client_session: ClientSession,
|
||||
request: httpx.Request,
|
||||
timeout: dict,
|
||||
proxy: Optional[str],
|
||||
sni_hostname: Optional[str],
|
||||
) -> ClientResponse:
|
||||
"""
|
||||
Helper function to make an aiohttp request with the given parameters.
|
||||
|
||||
Args:
|
||||
client_session: The aiohttp ClientSession to use
|
||||
request: The httpx Request to send
|
||||
timeout: Timeout settings dict with 'connect', 'read', 'pool' keys
|
||||
proxy: Optional proxy URL
|
||||
sni_hostname: Optional SNI hostname for SSL
|
||||
|
||||
Returns:
|
||||
ClientResponse from aiohttp
|
||||
"""
|
||||
from aiohttp import ClientTimeout
|
||||
from yarl import URL as YarlURL
|
||||
|
||||
try:
|
||||
data = request.content
|
||||
except httpx.RequestNotRead:
|
||||
data = request.stream # type: ignore
|
||||
request.headers.pop("transfer-encoding", None) # handled by aiohttp
|
||||
|
||||
response = await client_session.request(
|
||||
method=request.method,
|
||||
url=YarlURL(str(request.url), encoded=True),
|
||||
headers=request.headers,
|
||||
data=data,
|
||||
allow_redirects=False,
|
||||
auto_decompress=False,
|
||||
timeout=ClientTimeout(
|
||||
sock_connect=timeout.get("connect"),
|
||||
sock_read=timeout.get("read"),
|
||||
connect=timeout.get("pool"),
|
||||
),
|
||||
proxy=proxy,
|
||||
server_hostname=sni_hostname,
|
||||
).__aenter__()
|
||||
|
||||
return response
|
||||
|
||||
async def handle_async_request(
|
||||
self,
|
||||
request: httpx.Request,
|
||||
) -> httpx.Response:
|
||||
from aiohttp import ClientTimeout
|
||||
from yarl import URL as YarlURL
|
||||
|
||||
timeout = request.extensions.get("timeout", {})
|
||||
sni_hostname = request.extensions.get("sni_hostname")
|
||||
|
||||
|
|
@ -209,28 +267,38 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
# Resolve proxy settings from environment variables
|
||||
proxy = await self._get_proxy_settings(request)
|
||||
|
||||
with map_aiohttp_exceptions():
|
||||
try:
|
||||
data = request.content
|
||||
except httpx.RequestNotRead:
|
||||
data = request.stream # type: ignore
|
||||
request.headers.pop("transfer-encoding", None) # handled by aiohttp
|
||||
|
||||
response = await client_session.request(
|
||||
method=request.method,
|
||||
url=YarlURL(str(request.url), encoded=True),
|
||||
headers=request.headers,
|
||||
data=data,
|
||||
allow_redirects=False,
|
||||
auto_decompress=False,
|
||||
timeout=ClientTimeout(
|
||||
sock_connect=timeout.get("connect"),
|
||||
sock_read=timeout.get("read"),
|
||||
connect=timeout.get("pool"),
|
||||
),
|
||||
proxy=proxy,
|
||||
server_hostname=sni_hostname,
|
||||
).__aenter__()
|
||||
try:
|
||||
with map_aiohttp_exceptions():
|
||||
response = await self._make_aiohttp_request(
|
||||
client_session=client_session,
|
||||
request=request,
|
||||
timeout=timeout,
|
||||
proxy=proxy,
|
||||
sni_hostname=sni_hostname,
|
||||
)
|
||||
except RuntimeError as e:
|
||||
# Handle the case where session was closed between our check and actual use
|
||||
if "Session is closed" in str(e):
|
||||
verbose_logger.debug(f"Session closed during request, retrying with new session: {e}")
|
||||
# Force creation of a new session
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
client_session = self.client
|
||||
|
||||
# Retry the request with the new session
|
||||
with map_aiohttp_exceptions():
|
||||
response = await self._make_aiohttp_request(
|
||||
client_session=client_session,
|
||||
request=request,
|
||||
timeout=timeout,
|
||||
proxy=proxy,
|
||||
sni_hostname=sni_hostname,
|
||||
)
|
||||
else:
|
||||
# Re-raise if it's a different RuntimeError
|
||||
raise
|
||||
|
||||
return httpx.Response(
|
||||
status_code=response.status,
|
||||
|
|
|
|||
|
|
@ -12,7 +12,13 @@ from httpx._types import RequestFiles
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import _DEFAULT_TTL_FOR_HTTPX_CLIENTS
|
||||
from litellm.constants import (
|
||||
_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
AIOHTTP_CONNECTOR_LIMIT,
|
||||
AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
AIOHTTP_TTL_DNS_CACHE,
|
||||
DEFAULT_SSL_CIPHERS
|
||||
)
|
||||
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
|
||||
from litellm.types.llms.custom_http import *
|
||||
|
||||
|
|
@ -94,10 +100,19 @@ def get_ssl_configuration(
|
|||
|
||||
if ssl_verify is not False:
|
||||
custom_ssl_context = ssl.create_default_context(cafile=cafile)
|
||||
# If security level is set, apply it to the SSL context
|
||||
|
||||
# Optimize SSL handshake performance
|
||||
# Set minimum TLS version to 1.2 for better performance
|
||||
custom_ssl_context.minimum_version = ssl.TLSVersion.TLSv1_2
|
||||
|
||||
# Configure cipher suites for optimal performance
|
||||
if ssl_security_level and isinstance(ssl_security_level, str):
|
||||
# Create a custom SSL context with reduced security level
|
||||
# User provided custom cipher configuration (e.g., via SSL_SECURITY_LEVEL env var)
|
||||
custom_ssl_context.set_ciphers(ssl_security_level)
|
||||
else:
|
||||
# Use optimized cipher list that strongly prefers fast ciphers
|
||||
# but falls back to widely compatible ones
|
||||
custom_ssl_context.set_ciphers(DEFAULT_SSL_CIPHERS)
|
||||
|
||||
# Use our custom SSL context instead of the original ssl_verify value
|
||||
return custom_ssl_context
|
||||
|
|
@ -651,7 +666,13 @@ class AsyncHTTPHandler:
|
|||
)
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=lambda: ClientSession(
|
||||
connector=TCPConnector(limit=0, **connector_kwargs), # 0 = unlimited connections per host
|
||||
connector=TCPConnector(
|
||||
limit=AIOHTTP_CONNECTOR_LIMIT,
|
||||
keepalive_timeout=AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
ttl_dns_cache=AIOHTTP_TTL_DNS_CACHE,
|
||||
enable_cleanup_closed=True,
|
||||
**connector_kwargs
|
||||
),
|
||||
trust_env=trust_env,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -426,6 +426,7 @@ class BaseLLMHTTPHandler:
|
|||
),
|
||||
json_mode=json_mode,
|
||||
signed_json_body=signed_json_body,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
if stream is True:
|
||||
|
|
@ -1532,6 +1533,7 @@ class BaseLLMHTTPHandler:
|
|||
data=data,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
|
|
|
|||
48
litellm/llms/litellm_proxy/responses/transformation.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
"""
|
||||
Responses API transformation for LiteLLM Proxy provider.
|
||||
|
||||
LiteLLM Proxy supports the OpenAI Responses API natively when the underlying model supports it.
|
||||
This config enables pass-through behavior to the proxy's /v1/responses endpoint.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
|
||||
class LiteLLMProxyResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
"""
|
||||
Configuration for LiteLLM Proxy Responses API support.
|
||||
|
||||
Extends OpenAI's config since the proxy follows OpenAI's API spec,
|
||||
but uses LITELLM_PROXY_API_BASE for the base URL.
|
||||
"""
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.LITELLM_PROXY
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
litellm_params: dict,
|
||||
) -> str:
|
||||
"""
|
||||
Get the endpoint for LiteLLM Proxy responses API.
|
||||
|
||||
Uses LITELLM_PROXY_API_BASE environment variable if api_base is not provided.
|
||||
"""
|
||||
api_base = api_base or get_secret_str("LITELLM_PROXY_API_BASE")
|
||||
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"api_base not set for LiteLLM Proxy responses API. "
|
||||
"Set via api_base parameter or LITELLM_PROXY_API_BASE environment variable"
|
||||
)
|
||||
|
||||
# Remove trailing slashes
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
return f"{api_base}/responses"
|
||||
|
|
@ -19,6 +19,13 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.llms.oci.common_utils import OCIError
|
||||
from litellm.types.llms.oci import (
|
||||
CohereChatRequest,
|
||||
CohereMessage,
|
||||
CohereChatResult,
|
||||
CohereParameterDefinition,
|
||||
CohereStreamChunk,
|
||||
CohereTool,
|
||||
CohereToolCall,
|
||||
OCIChatRequestPayload,
|
||||
OCICompletionPayload,
|
||||
OCICompletionResponse,
|
||||
|
|
@ -37,13 +44,13 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.utils import (
|
||||
Delta,
|
||||
LlmProviders,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
from litellm.utils import (
|
||||
ChatCompletionMessageToolCall,
|
||||
CustomStreamWrapper,
|
||||
ModelResponse,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
|
@ -170,13 +177,16 @@ class OCIChatConfig(BaseConfig):
|
|||
"web_search_options": False,
|
||||
}
|
||||
|
||||
# Cohere and Gemini use the same parameter mapping as GENERIC
|
||||
self.openai_to_oci_cohere_param_map = self.openai_to_oci_generic_param_map.copy()
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
supported_params = []
|
||||
vendor = get_vendor_from_model(model)
|
||||
if vendor == OCIVendors.COHERE:
|
||||
raise ValueError(
|
||||
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
|
||||
)
|
||||
open_ai_to_oci_param_map = self.openai_to_oci_cohere_param_map
|
||||
open_ai_to_oci_param_map.pop("tool_choice")
|
||||
open_ai_to_oci_param_map.pop("max_retries")
|
||||
else:
|
||||
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
|
||||
for key, value in open_ai_to_oci_param_map.items():
|
||||
|
|
@ -195,9 +205,7 @@ class OCIChatConfig(BaseConfig):
|
|||
adapted_params = {}
|
||||
vendor = get_vendor_from_model(model)
|
||||
if vendor == OCIVendors.COHERE:
|
||||
raise ValueError(
|
||||
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
|
||||
)
|
||||
open_ai_to_oci_param_map = self.openai_to_oci_cohere_param_map
|
||||
else:
|
||||
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
|
||||
|
||||
|
|
@ -416,21 +424,130 @@ class OCIChatConfig(BaseConfig):
|
|||
def _get_optional_params(self, vendor: OCIVendors, optional_params: dict) -> Dict:
|
||||
selected_params = {}
|
||||
if vendor == OCIVendors.COHERE:
|
||||
raise ValueError(
|
||||
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
|
||||
)
|
||||
open_ai_to_oci_param_map = self.openai_to_oci_cohere_param_map
|
||||
# remove tool_choice from the map
|
||||
open_ai_to_oci_param_map.pop("tool_choice")
|
||||
# Add default values for Cohere API
|
||||
selected_params = {
|
||||
"maxTokens": 600,
|
||||
"temperature": 1,
|
||||
"topK": 0,
|
||||
"topP": 0.75,
|
||||
"frequencyPenalty": 0
|
||||
}
|
||||
else:
|
||||
open_ai_to_oci_param_map = self.openai_to_oci_generic_param_map
|
||||
|
||||
for value in open_ai_to_oci_param_map.values():
|
||||
if value in optional_params:
|
||||
selected_params[value] = optional_params[value]
|
||||
# Map OpenAI params to OCI params
|
||||
for openai_key, oci_key in open_ai_to_oci_param_map.items():
|
||||
if oci_key and openai_key in optional_params:
|
||||
selected_params[oci_key] = optional_params[openai_key] # type: ignore[index]
|
||||
|
||||
# Also check for already-mapped OCI params (for backward compatibility)
|
||||
for oci_value in open_ai_to_oci_param_map.values():
|
||||
if oci_value and oci_value in optional_params and oci_value not in selected_params:
|
||||
selected_params[oci_value] = optional_params[oci_value] # type: ignore[index]
|
||||
|
||||
if "tools" in selected_params:
|
||||
selected_params["tools"] = adapt_tool_definition_to_oci_standard(
|
||||
selected_params["tools"], vendor
|
||||
)
|
||||
if vendor == OCIVendors.COHERE:
|
||||
selected_params["tools"] = self.adapt_tool_definitions_to_cohere_standard( # type: ignore[assignment]
|
||||
selected_params["tools"] # type: ignore[arg-type]
|
||||
)
|
||||
else:
|
||||
selected_params["tools"] = adapt_tool_definition_to_oci_standard( # type: ignore[assignment]
|
||||
selected_params["tools"], vendor # type: ignore[arg-type]
|
||||
)
|
||||
return selected_params
|
||||
|
||||
def adapt_messages_to_cohere_standard(self, messages: List[AllMessageValues]) -> List[CohereMessage]:
|
||||
"""Build chat history for Cohere models."""
|
||||
chat_history = []
|
||||
for msg in messages[:-1]: # All messages except the last one
|
||||
role = msg.get("role")
|
||||
content = msg.get("content")
|
||||
|
||||
if isinstance(content, list):
|
||||
# Extract text from content array
|
||||
text_content = ""
|
||||
for content_item in content:
|
||||
if isinstance(content_item, dict) and content_item.get("type") == "text":
|
||||
text_content += content_item.get("text", "")
|
||||
content = text_content
|
||||
|
||||
# Ensure content is a string
|
||||
if not isinstance(content, str):
|
||||
content = str(content) if content is not None else ""
|
||||
|
||||
# Handle tool calls
|
||||
tool_calls: Optional[List[CohereToolCall]] = None
|
||||
if role == "assistant" and "tool_calls" in msg and msg.get("tool_calls"): # type: ignore[union-attr,typeddict-item]
|
||||
tool_calls = []
|
||||
for tool_call in msg["tool_calls"]: # type: ignore[union-attr,typeddict-item]
|
||||
# Parse arguments if they're a JSON string
|
||||
raw_arguments: Any = tool_call.get("function", {}).get("arguments", {})
|
||||
if isinstance(raw_arguments, str):
|
||||
try:
|
||||
arguments: Dict[str, Any] = json.loads(raw_arguments)
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
else:
|
||||
arguments = raw_arguments
|
||||
|
||||
tool_calls.append(CohereToolCall(
|
||||
name=str(tool_call.get("function", {}).get("name", "")),
|
||||
parameters=arguments
|
||||
))
|
||||
|
||||
if role == "user":
|
||||
chat_history.append(CohereMessage(role="USER", message=content))
|
||||
elif role == "assistant":
|
||||
chat_history.append(CohereMessage(role="CHATBOT", message=content, toolCalls=tool_calls))
|
||||
elif role == "tool":
|
||||
# Tool messages need special handling
|
||||
chat_history.append(CohereMessage(
|
||||
role="TOOL",
|
||||
message=content,
|
||||
toolCalls=None # Tool messages don't have tool calls
|
||||
))
|
||||
|
||||
return chat_history
|
||||
|
||||
def adapt_tool_definitions_to_cohere_standard(self, tools: List[Dict[str, Any]]) -> List[CohereTool]:
|
||||
"""Adapt tool definitions to Cohere format."""
|
||||
cohere_tools = []
|
||||
for tool in tools:
|
||||
function_def = tool.get("function", {})
|
||||
parameters = function_def.get("parameters", {}).get("properties", {})
|
||||
required = function_def.get("parameters", {}).get("required", [])
|
||||
|
||||
parameter_definitions = {}
|
||||
for param_name, param_schema in parameters.items():
|
||||
parameter_definitions[param_name] = CohereParameterDefinition(
|
||||
description=param_schema.get("description", ""),
|
||||
type=param_schema.get("type", "string"),
|
||||
isRequired=param_name in required
|
||||
)
|
||||
|
||||
cohere_tools.append(CohereTool(
|
||||
name=function_def.get("name", ""),
|
||||
description=function_def.get("description", ""),
|
||||
parameterDefinitions=parameter_definitions
|
||||
))
|
||||
|
||||
return cohere_tools
|
||||
|
||||
def _extract_text_content(self, content: Any) -> str:
|
||||
"""Extract text content from message content."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
elif isinstance(content, list):
|
||||
text_content = ""
|
||||
for content_item in content:
|
||||
if isinstance(content_item, dict) and content_item.get("type") == "text":
|
||||
text_content += content_item.get("text", "")
|
||||
return text_content
|
||||
return str(content)
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -445,28 +562,47 @@ class OCIChatConfig(BaseConfig):
|
|||
|
||||
vendor = get_vendor_from_model(model)
|
||||
|
||||
if vendor == OCIVendors.COHERE:
|
||||
oci_serving_mode = optional_params.get("oci_serving_mode", "ON_DEMAND")
|
||||
if oci_serving_mode not in ["ON_DEMAND", "DEDICATED"]:
|
||||
raise Exception(
|
||||
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
|
||||
"kwarg `oci_serving_mode` must be either 'ON_DEMAND' or 'DEDICATED'"
|
||||
)
|
||||
|
||||
if oci_serving_mode == "DEDICATED":
|
||||
servingMode = OCIServingMode(
|
||||
servingType="DEDICATED",
|
||||
endpointId=model,
|
||||
)
|
||||
else:
|
||||
oci_serving_mode = optional_params.get("oci_serving_mode", "ON_DEMAND")
|
||||
if oci_serving_mode not in ["ON_DEMAND", "DEDICATED"]:
|
||||
raise Exception(
|
||||
"kwarg `oci_serving_mode` must be either 'ON_DEMAND' or 'DEDICATED'"
|
||||
)
|
||||
servingMode = OCIServingMode(
|
||||
servingType="ON_DEMAND",
|
||||
modelId=model,
|
||||
)
|
||||
|
||||
if oci_serving_mode == "DEDICATED":
|
||||
servingMode = OCIServingMode(
|
||||
servingType="DEDICATED",
|
||||
endpointId=model,
|
||||
)
|
||||
else:
|
||||
servingMode = OCIServingMode(
|
||||
servingType="ON_DEMAND",
|
||||
modelId=model,
|
||||
)
|
||||
# Build request based on vendor type
|
||||
if vendor == OCIVendors.COHERE:
|
||||
# For Cohere, we need to use the specific Cohere format
|
||||
# Extract the last user message as the main message
|
||||
user_messages = [msg for msg in messages if msg.get("role") == "user"]
|
||||
if not user_messages:
|
||||
raise Exception("No user message found for Cohere model")
|
||||
|
||||
|
||||
# Create Cohere-specific chat request
|
||||
chat_request = CohereChatRequest(
|
||||
apiFormat="COHERE",
|
||||
message=self._extract_text_content(user_messages[-1]["content"]),
|
||||
chatHistory=self.adapt_messages_to_cohere_standard(messages),
|
||||
**self._get_optional_params(OCIVendors.COHERE, optional_params)
|
||||
)
|
||||
|
||||
data = OCICompletionPayload(
|
||||
compartmentId=oci_compartment_id,
|
||||
servingMode=servingMode,
|
||||
chatRequest=chat_request
|
||||
)
|
||||
else:
|
||||
# Use generic format for other vendors
|
||||
data = OCICompletionPayload(
|
||||
compartmentId=oci_compartment_id,
|
||||
servingMode=servingMode,
|
||||
|
|
@ -479,6 +615,111 @@ class OCIChatConfig(BaseConfig):
|
|||
|
||||
return data.model_dump(exclude_none=True)
|
||||
|
||||
def _handle_cohere_response(
|
||||
self,
|
||||
json_response: dict,
|
||||
model: str,
|
||||
model_response: ModelResponse
|
||||
) -> ModelResponse:
|
||||
"""Handle Cohere-specific response format."""
|
||||
cohere_response = CohereChatResult(**json_response)
|
||||
# Cohere response format (uses camelCase)
|
||||
model_id = model
|
||||
|
||||
# Set basic response info
|
||||
model_response.model = model_id
|
||||
model_response.created = int(datetime.datetime.now().timestamp())
|
||||
|
||||
# Extract the response text
|
||||
response_text = cohere_response.chatResponse.text
|
||||
oci_finish_reason = cohere_response.chatResponse.finishReason
|
||||
|
||||
# Map finish reason
|
||||
if oci_finish_reason == "COMPLETE":
|
||||
finish_reason = "stop"
|
||||
elif oci_finish_reason == "MAX_TOKENS":
|
||||
finish_reason = "length"
|
||||
else:
|
||||
finish_reason = "stop"
|
||||
|
||||
# Handle tool calls
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None
|
||||
if cohere_response.chatResponse.toolCalls:
|
||||
tool_calls = []
|
||||
for tool_call in cohere_response.chatResponse.toolCalls:
|
||||
tool_calls.append({
|
||||
"id": f"call_{len(tool_calls)}", # Generate a simple ID
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tool_call.name,
|
||||
"arguments": json.dumps(tool_call.parameters)
|
||||
}
|
||||
})
|
||||
|
||||
# Create choice
|
||||
from litellm.types.utils import Choices
|
||||
choice = Choices(
|
||||
index=0,
|
||||
message={
|
||||
"role": "assistant",
|
||||
"content": response_text,
|
||||
"tool_calls": tool_calls
|
||||
},
|
||||
finish_reason=finish_reason
|
||||
)
|
||||
model_response.choices = [choice]
|
||||
|
||||
# Extract usage info
|
||||
usage_info = cohere_response.chatResponse.usage
|
||||
from litellm.types.utils import Usage
|
||||
model_response.usage = Usage( # type: ignore[attr-defined]
|
||||
prompt_tokens=usage_info.promptTokens, # type: ignore[union-attr]
|
||||
completion_tokens=usage_info.completionTokens, # type: ignore[union-attr]
|
||||
total_tokens=usage_info.totalTokens # type: ignore[union-attr]
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
def _handle_generic_response(
|
||||
self,
|
||||
json: dict,
|
||||
model: str,
|
||||
model_response: ModelResponse,
|
||||
raw_response: httpx.Response
|
||||
) -> ModelResponse:
|
||||
"""Handle generic OCI response format."""
|
||||
try:
|
||||
completion_response = OCICompletionResponse(**json)
|
||||
except TypeError as e:
|
||||
raise OCIError(
|
||||
message=f"Response cannot be casted to OCICompletionResponse: {str(e)}",
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
iso_str = completion_response.chatResponse.timeCreated
|
||||
dt = datetime.datetime.fromisoformat(iso_str.replace("Z", "+00:00"))
|
||||
model_response.created = int(dt.timestamp())
|
||||
|
||||
model_response.model = completion_response.modelId
|
||||
|
||||
message = model_response.choices[0].message # type: ignore
|
||||
response_message = completion_response.chatResponse.choices[0].message
|
||||
if response_message.content and response_message.content[0].type == "TEXT":
|
||||
message.content = response_message.content[0].text
|
||||
if response_message.toolCalls:
|
||||
message.tool_calls = adapt_tools_to_openai_standard(
|
||||
response_message.toolCalls
|
||||
)
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=completion_response.chatResponse.usage.promptTokens,
|
||||
completion_tokens=completion_response.chatResponse.usage.completionTokens,
|
||||
total_tokens=completion_response.chatResponse.usage.totalTokens,
|
||||
)
|
||||
model_response.usage = usage # type: ignore
|
||||
|
||||
return model_response
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -509,46 +750,13 @@ class OCIChatConfig(BaseConfig):
|
|||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
try:
|
||||
completion_response = OCICompletionResponse(**json)
|
||||
except TypeError as e:
|
||||
raise OCIError(
|
||||
message=f"Response cannot be casted to OCICompletionResponse: {str(e)}",
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
vendor = get_vendor_from_model(model)
|
||||
|
||||
# Handle response based on vendor type
|
||||
if vendor == OCIVendors.COHERE:
|
||||
raise ValueError(
|
||||
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
|
||||
)
|
||||
model_response = self._handle_cohere_response(json, model, model_response)
|
||||
else:
|
||||
iso_str = completion_response.chatResponse.timeCreated
|
||||
dt = datetime.datetime.fromisoformat(iso_str.replace("Z", "+00:00"))
|
||||
model_response.created = int(dt.timestamp())
|
||||
|
||||
model_response.model = completion_response.modelId
|
||||
|
||||
message = model_response.choices[0].message # type: ignore
|
||||
if vendor == OCIVendors.COHERE:
|
||||
raise ValueError(
|
||||
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
|
||||
)
|
||||
else:
|
||||
response_message = completion_response.chatResponse.choices[0].message
|
||||
if response_message.content and response_message.content[0].type == "TEXT":
|
||||
message.content = response_message.content[0].text
|
||||
if response_message.toolCalls:
|
||||
message.tool_calls = adapt_tools_to_openai_standard(
|
||||
response_message.toolCalls
|
||||
)
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=completion_response.chatResponse.usage.promptTokens,
|
||||
completion_tokens=completion_response.chatResponse.usage.completionTokens,
|
||||
total_tokens=completion_response.chatResponse.usage.totalTokens,
|
||||
)
|
||||
model_response.usage = usage # type: ignore
|
||||
model_response = self._handle_generic_response(json, model, model_response, raw_response)
|
||||
|
||||
model_response._hidden_params["additional_headers"] = raw_response.headers
|
||||
|
||||
|
|
@ -818,26 +1026,21 @@ def adapt_messages_to_generic_oci_standard(
|
|||
|
||||
def adapt_tool_definition_to_oci_standard(tools: List[Dict], vendor: OCIVendors):
|
||||
new_tools = []
|
||||
if vendor == OCIVendors.COHERE:
|
||||
raise ValueError(
|
||||
"Cohere models are not yet supported in the litellm OCI chat completion endpoint. Use the Cohere API directly."
|
||||
for tool in tools:
|
||||
if tool["type"] != "function":
|
||||
raise Exception("OCI only supports function tools")
|
||||
|
||||
tool_function = tool.get("function")
|
||||
if not isinstance(tool_function, dict):
|
||||
raise Exception("Prop `function` is not a dictionary")
|
||||
|
||||
new_tool = OCIToolDefinition(
|
||||
type="FUNCTION",
|
||||
name=tool_function.get("name"),
|
||||
description=tool_function.get("description", ""),
|
||||
parameters=tool_function.get("parameters", {}),
|
||||
)
|
||||
else:
|
||||
for tool in tools:
|
||||
if tool["type"] != "function":
|
||||
raise Exception("OCI only supports function tools")
|
||||
|
||||
tool_function = tool.get("function")
|
||||
if not isinstance(tool_function, dict):
|
||||
raise Exception("Prop `function` is not a dictionary")
|
||||
|
||||
new_tool = OCIToolDefinition(
|
||||
type="FUNCTION",
|
||||
name=tool_function.get("name"),
|
||||
description=tool_function.get("description", ""),
|
||||
parameters=tool_function.get("parameters", {}),
|
||||
)
|
||||
new_tools.append(new_tool)
|
||||
new_tools.append(new_tool)
|
||||
|
||||
return new_tools
|
||||
|
||||
|
|
@ -877,6 +1080,58 @@ class OCIStreamWrapper(CustomStreamWrapper):
|
|||
if not chunk.startswith("data:"):
|
||||
raise ValueError(f"Chunk does not start with 'data:': {chunk}")
|
||||
dict_chunk = json.loads(chunk[5:]) # Remove 'data: ' prefix and parse JSON
|
||||
|
||||
# Check if this is a Cohere stream chunk
|
||||
if "apiFormat" in dict_chunk and dict_chunk.get("apiFormat") == "COHERE":
|
||||
return self._handle_cohere_stream_chunk(dict_chunk)
|
||||
else:
|
||||
return self._handle_generic_stream_chunk(dict_chunk)
|
||||
|
||||
def _handle_cohere_stream_chunk(self, dict_chunk: dict):
|
||||
"""Handle Cohere-specific streaming chunks."""
|
||||
try:
|
||||
typed_chunk = CohereStreamChunk(**dict_chunk)
|
||||
except TypeError as e:
|
||||
raise ValueError(f"Chunk cannot be casted to CohereStreamChunk: {str(e)}")
|
||||
|
||||
if typed_chunk.index is None:
|
||||
typed_chunk.index = 0
|
||||
|
||||
# Extract text content
|
||||
text = typed_chunk.text or ""
|
||||
|
||||
# Map finish reason to standard format
|
||||
finish_reason = typed_chunk.finishReason
|
||||
if finish_reason == "COMPLETE":
|
||||
finish_reason = "stop"
|
||||
elif finish_reason == "MAX_TOKENS":
|
||||
finish_reason = "length"
|
||||
elif finish_reason is None:
|
||||
finish_reason = None
|
||||
else:
|
||||
finish_reason = "stop"
|
||||
|
||||
# For Cohere, we don't have tool calls in the streaming format
|
||||
tool_calls = None
|
||||
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=typed_chunk.index if typed_chunk.index else 0,
|
||||
delta=Delta(
|
||||
content=text,
|
||||
tool_calls=tool_calls,
|
||||
provider_specific_fields=None,
|
||||
thinking_blocks=None,
|
||||
reasoning_content=None,
|
||||
),
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
def _handle_generic_stream_chunk(self, dict_chunk: dict):
|
||||
"""Handle generic OCI streaming chunks."""
|
||||
try:
|
||||
typed_chunk = OCIStreamChunk(**dict_chunk)
|
||||
except TypeError as e:
|
||||
|
|
|
|||
|
|
@ -184,9 +184,12 @@ class OllamaChatConfig(BaseConfig):
|
|||
):
|
||||
if value.get("json_schema") and value["json_schema"].get("schema"):
|
||||
optional_params["format"] = value["json_schema"]["schema"]
|
||||
### FUNCTION CALLING LOGIC ###
|
||||
if param == "reasoning_effort" and value is not None:
|
||||
optional_params["think"] = True
|
||||
if model.startswith("gpt-oss"):
|
||||
optional_params["think"] = value
|
||||
else:
|
||||
optional_params["think"] = True
|
||||
### FUNCTION CALLING LOGIC ###
|
||||
if param == "tools":
|
||||
## CHECK IF MODEL SUPPORTS TOOL CALLING ##
|
||||
try:
|
||||
|
|
@ -281,6 +284,7 @@ class OllamaChatConfig(BaseConfig):
|
|||
stream = optional_params.pop("stream", False)
|
||||
format = optional_params.pop("format", None)
|
||||
keep_alive = optional_params.pop("keep_alive", None)
|
||||
think = optional_params.pop("think", None)
|
||||
function_name = optional_params.pop("function_name", None)
|
||||
litellm_params["function_name"] = function_name
|
||||
tools = optional_params.pop("tools", None)
|
||||
|
|
@ -344,6 +348,8 @@ class OllamaChatConfig(BaseConfig):
|
|||
data["tools"] = tools
|
||||
if keep_alive is not None:
|
||||
data["keep_alive"] = keep_alive
|
||||
if think is not None:
|
||||
data["think"] = think
|
||||
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -180,7 +180,10 @@ class OllamaConfig(BaseConfig):
|
|||
elif param == "stop":
|
||||
optional_params["stop"] = value
|
||||
elif param == "reasoning_effort" and value is not None:
|
||||
optional_params["think"] = True
|
||||
if model.startswith("gpt-oss"):
|
||||
optional_params["think"] = value
|
||||
else:
|
||||
optional_params["think"] = True
|
||||
elif param == "response_format" and isinstance(value, dict):
|
||||
if value["type"] == "json_object":
|
||||
optional_params["format"] = "json"
|
||||
|
|
@ -412,6 +415,7 @@ class OllamaConfig(BaseConfig):
|
|||
stream = optional_params.pop("stream", False)
|
||||
format = optional_params.pop("format", None)
|
||||
images = optional_params.pop("images", None)
|
||||
think = optional_params.pop("think", None)
|
||||
data = {
|
||||
"model": model,
|
||||
"prompt": ollama_prompt,
|
||||
|
|
@ -425,6 +429,8 @@ class OllamaConfig(BaseConfig):
|
|||
data["images"] = [
|
||||
_convert_image(convert_to_ollama_image(image)) for image in images
|
||||
]
|
||||
if think is not None:
|
||||
data["think"] = think
|
||||
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ def get_ollama_response( # noqa: PLR0915
|
|||
stream = optional_params.pop("stream", False)
|
||||
format = optional_params.pop("format", None)
|
||||
keep_alive = optional_params.pop("keep_alive", None)
|
||||
think = optional_params.pop("think", None)
|
||||
function_name = optional_params.pop("function_name", None)
|
||||
tools = optional_params.pop("tools", None)
|
||||
|
||||
|
|
@ -98,6 +99,8 @@ def get_ollama_response( # noqa: PLR0915
|
|||
data["tools"] = tools
|
||||
if keep_alive is not None:
|
||||
data["keep_alive"] = keep_alive
|
||||
if think is not None:
|
||||
data["think"] = think
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=None,
|
||||
|
|
|
|||
|
|
@ -397,13 +397,13 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
)
|
||||
from litellm.types.llms.openai import ChatCompletionToolParam
|
||||
|
||||
for message in messages:
|
||||
message = cast(
|
||||
for i, message in enumerate(messages):
|
||||
messages[i] = cast(
|
||||
AllMessageValues, filter_value_from_dict(message, "cache_control") # type: ignore
|
||||
)
|
||||
if tools is not None:
|
||||
for tool in tools:
|
||||
tool = cast(
|
||||
for i, tool in enumerate(tools):
|
||||
tools[i] = cast(
|
||||
ChatCompletionToolParam,
|
||||
filter_value_from_dict(tool, "cache_control"), # type: ignore
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1125,6 +1125,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
api_base: Optional[str] = None,
|
||||
client: Optional[AsyncOpenAI] = None,
|
||||
max_retries=None,
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
):
|
||||
try:
|
||||
openai_aclient: AsyncOpenAI = self._get_openai_client( # type: ignore
|
||||
|
|
@ -1134,6 +1135,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
client=client,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
headers, response = await self.make_openai_embedding_request(
|
||||
openai_aclient=openai_aclient,
|
||||
|
|
@ -1197,6 +1199,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
client=None,
|
||||
aembedding=None,
|
||||
max_retries: Optional[int] = None,
|
||||
shared_session: Optional["ClientSession"] = None,
|
||||
) -> EmbeddingResponse:
|
||||
super().embedding()
|
||||
try:
|
||||
|
|
@ -1223,6 +1226,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
timeout=timeout,
|
||||
client=client,
|
||||
max_retries=max_retries,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
openai_client: OpenAI = self._get_openai_client( # type: ignore
|
||||
|
|
|
|||
|
|
@ -161,6 +161,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
) -> ResponsesAPIResponse:
|
||||
"""No transform applied since outputs are in OpenAI spec already"""
|
||||
try:
|
||||
logging_obj.post_call(
|
||||
original_response=raw_response.text,
|
||||
additional_args={"complete_input_dict": {}},
|
||||
)
|
||||
raw_response_json = raw_response.json()
|
||||
raw_response_json["created_at"] = _safe_convert_created_field(
|
||||
raw_response_json["created_at"]
|
||||
|
|
@ -169,7 +173,13 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
raise OpenAIError(
|
||||
message=raw_response.text, status_code=raw_response.status_code
|
||||
)
|
||||
return ResponsesAPIResponse(**raw_response_json)
|
||||
try:
|
||||
return ResponsesAPIResponse(**raw_response_json)
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct"
|
||||
)
|
||||
return ResponsesAPIResponse.model_construct(**raw_response_json)
|
||||
|
||||
def validate_environment(
|
||||
self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]
|
||||
|
|
|
|||
|
|
@ -6,7 +6,8 @@ Calls done in OpenAI/openai.py as OpenRouter is openai-compatible.
|
|||
Docs: https://openrouter.ai/docs/parameters
|
||||
"""
|
||||
|
||||
from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple, Union
|
||||
from enum import Enum
|
||||
from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple, Union, cast
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -20,6 +21,12 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
|||
from ..common_utils import OpenRouterException
|
||||
|
||||
|
||||
class CacheControlSupportedModels(str, Enum):
|
||||
"""Models that support cache_control in content blocks."""
|
||||
CLAUDE = "claude"
|
||||
GEMINI = "gemini"
|
||||
|
||||
|
||||
class OpenrouterConfig(OpenAIGPTConfig):
|
||||
def map_openai_params(
|
||||
self,
|
||||
|
|
@ -48,19 +55,76 @@ class OpenrouterConfig(OpenAIGPTConfig):
|
|||
)
|
||||
return mapped_openai_params
|
||||
|
||||
def _supports_cache_control_in_content(self, model: str) -> bool:
|
||||
"""
|
||||
Check if the model supports cache_control in content blocks.
|
||||
|
||||
Returns:
|
||||
bool: True if model supports cache_control (Claude or Gemini models)
|
||||
"""
|
||||
model_lower = model.lower()
|
||||
return any(
|
||||
supported_model.value in model_lower
|
||||
for supported_model in CacheControlSupportedModels
|
||||
)
|
||||
|
||||
def remove_cache_control_flag_from_messages_and_tools(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
tools: Optional[List["ChatCompletionToolParam"]] = None,
|
||||
) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]:
|
||||
if "claude" in model.lower(): # don't remove 'cache_control' flag
|
||||
if self._supports_cache_control_in_content(model):
|
||||
return messages, tools
|
||||
else:
|
||||
return super().remove_cache_control_flag_from_messages_and_tools(
|
||||
model, messages, tools
|
||||
)
|
||||
|
||||
def _move_cache_control_to_content(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> List[AllMessageValues]:
|
||||
"""
|
||||
Move cache_control from message level to content blocks.
|
||||
OpenRouter requires cache_control to be inside content blocks, not at message level.
|
||||
|
||||
To avoid exceeding Anthropic's limit of 4 cache breakpoints, cache_control is only
|
||||
added to the LAST content block in each message.
|
||||
"""
|
||||
transformed_messages: List[AllMessageValues] = []
|
||||
for message in messages:
|
||||
message_dict = dict(message)
|
||||
cache_control = message_dict.pop("cache_control", None)
|
||||
|
||||
if cache_control is not None:
|
||||
content = message_dict.get("content")
|
||||
|
||||
if isinstance(content, list):
|
||||
# Content is already a list, add cache_control only to the last block
|
||||
if len(content) > 0:
|
||||
content_copy = []
|
||||
for i, block in enumerate(content):
|
||||
block_dict = dict(block)
|
||||
# Only add cache_control to the last content block
|
||||
if i == len(content) - 1:
|
||||
block_dict["cache_control"] = cache_control
|
||||
content_copy.append(block_dict)
|
||||
message_dict["content"] = content_copy
|
||||
else:
|
||||
# Content is a string, convert to structured format
|
||||
message_dict["content"] = [
|
||||
{
|
||||
"type": "text",
|
||||
"text": content,
|
||||
"cache_control": cache_control,
|
||||
}
|
||||
]
|
||||
|
||||
# Cast back to AllMessageValues after modification
|
||||
transformed_messages.append(cast(AllMessageValues, message_dict))
|
||||
|
||||
return transformed_messages
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -75,13 +139,78 @@ class OpenrouterConfig(OpenAIGPTConfig):
|
|||
Returns:
|
||||
dict: The transformed request. Sent as the body of the API call.
|
||||
"""
|
||||
if self._supports_cache_control_in_content(model):
|
||||
messages = self._move_cache_control_to_content(messages)
|
||||
|
||||
extra_body = optional_params.pop("extra_body", {})
|
||||
response = super().transform_request(
|
||||
model, messages, optional_params, litellm_params, headers
|
||||
)
|
||||
response.update(extra_body)
|
||||
|
||||
# ALWAYS add usage parameter to get cost data from OpenRouter
|
||||
# This ensures cost tracking works for all OpenRouter models
|
||||
if "usage" not in response:
|
||||
response["usage"] = {"include": True}
|
||||
|
||||
return response
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: Any,
|
||||
request_data: dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
Transform the response from OpenRouter API.
|
||||
|
||||
Extracts cost information from response headers if available.
|
||||
|
||||
Returns:
|
||||
ModelResponse: The transformed response with cost information.
|
||||
"""
|
||||
# Call parent transform_response to get the standard ModelResponse
|
||||
model_response = super().transform_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
request_data=request_data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
api_key=api_key,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
# Extract cost from OpenRouter response body
|
||||
# OpenRouter returns cost information in the usage object when usage.include=true
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
if "usage" in response_json and response_json["usage"]:
|
||||
response_cost = response_json["usage"].get("cost")
|
||||
if response_cost is not None:
|
||||
# Store cost in hidden params for the cost calculator to use
|
||||
if not hasattr(model_response, "_hidden_params"):
|
||||
model_response._hidden_params = {}
|
||||
if "additional_headers" not in model_response._hidden_params:
|
||||
model_response._hidden_params["additional_headers"] = {}
|
||||
model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] = float(response_cost)
|
||||
except Exception:
|
||||
# If we can't extract cost, continue without it - don't fail the response
|
||||
pass
|
||||
|
||||
return model_response
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import re
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, get_type_hints
|
||||
|
||||
import httpx
|
||||
|
|
@ -24,6 +25,68 @@ class VertexAIError(BaseLLMException):
|
|||
super().__init__(message=message, status_code=status_code, headers=headers)
|
||||
|
||||
|
||||
class VertexAIModelRoute(str, Enum):
|
||||
"""Enum for Vertex AI model routing"""
|
||||
PARTNER_MODELS = "partner_models"
|
||||
GEMINI = "gemini"
|
||||
GEMMA = "gemma"
|
||||
MODEL_GARDEN = "model_garden"
|
||||
NON_GEMINI = "non_gemini"
|
||||
|
||||
|
||||
def get_vertex_ai_model_route(model: str, litellm_params: Optional[dict] = None) -> VertexAIModelRoute:
|
||||
"""
|
||||
Determine which handler to use for a Vertex AI model based on the model name.
|
||||
|
||||
Args:
|
||||
model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "openai/gpt-oss-120b")
|
||||
litellm_params: Optional litellm parameters dict that may contain base_model for routing
|
||||
|
||||
Returns:
|
||||
VertexAIModelRoute: The route enum indicating which handler should be used
|
||||
|
||||
Examples:
|
||||
>>> get_vertex_ai_model_route("llama3-405b")
|
||||
VertexAIModelRoute.PARTNER_MODELS
|
||||
|
||||
>>> get_vertex_ai_model_route("gemini-pro")
|
||||
VertexAIModelRoute.GEMINI
|
||||
|
||||
>>> get_vertex_ai_model_route("gemma/gemma-3-12b-it")
|
||||
VertexAIModelRoute.GEMMA
|
||||
|
||||
>>> get_vertex_ai_model_route("openai/gpt-oss-120b")
|
||||
VertexAIModelRoute.MODEL_GARDEN
|
||||
"""
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
|
||||
VertexAIPartnerModels,
|
||||
)
|
||||
|
||||
# Check base_model in litellm_params for gemini override
|
||||
if litellm_params and litellm_params.get("base_model") is not None:
|
||||
if "gemini" in litellm_params["base_model"]:
|
||||
return VertexAIModelRoute.GEMINI
|
||||
|
||||
# Check for partner models (llama, mistral, claude, etc.)
|
||||
if VertexAIPartnerModels.is_vertex_partner_model(model=model):
|
||||
return VertexAIModelRoute.PARTNER_MODELS
|
||||
|
||||
# Check for gemma models
|
||||
if "gemma/" in model:
|
||||
return VertexAIModelRoute.GEMMA
|
||||
|
||||
# Check for model garden openai models
|
||||
if "openai" in model:
|
||||
return VertexAIModelRoute.MODEL_GARDEN
|
||||
|
||||
# Check for gemini models
|
||||
if "gemini" in model:
|
||||
return VertexAIModelRoute.GEMINI
|
||||
|
||||
# Default to non-gemini (legacy vertex models like chat-bison, text-bison, etc.)
|
||||
return VertexAIModelRoute.NON_GEMINI
|
||||
|
||||
|
||||
def get_supports_system_message(
|
||||
model: str, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"]
|
||||
) -> bool:
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ def cost_router(
|
|||
or "mistral" in model
|
||||
or "jamba" in model
|
||||
or "codestral" in model
|
||||
or "gemma" in model
|
||||
):
|
||||
return "cost_per_token"
|
||||
elif custom_llm_provider == "vertex_ai" and (
|
||||
|
|
|
|||
2
litellm/llms/vertex_ai/vertex_gemma_models/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
"""Vertex AI Gemma-AI Models Handler"""
|
||||
|
||||
145
litellm/llms/vertex_ai/vertex_gemma_models/main.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
"""
|
||||
API Handler for calling Vertex AI Gemma Models
|
||||
|
||||
These models use a custom prediction endpoint format that wraps messages in 'instances'
|
||||
with @requestFormat: "chatCompletions" and returns responses wrapped in 'predictions'.
|
||||
|
||||
Usage:
|
||||
|
||||
response = litellm.completion(
|
||||
model="vertex_ai/gemma/gemma-3-12b-it-1222199011122",
|
||||
messages=[{"role": "user", "content": "What is machine learning?"}],
|
||||
vertex_project="your-project-id",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
|
||||
Sent to this route when `model` is in the format `vertex_ai/gemma/{MODEL_NAME}`
|
||||
|
||||
The API expects a custom endpoint URL format:
|
||||
https://{ENDPOINT_NUMBER}.{location}-{REGION_NUMBER}.prediction.vertexai.goog/v1/projects/{PROJECT_ID}/locations/{location}/endpoints/{ENDPOINT_ID}:predict
|
||||
"""
|
||||
|
||||
from typing import Callable, Optional, Union
|
||||
|
||||
import httpx # type: ignore
|
||||
|
||||
from litellm.utils import ModelResponse
|
||||
|
||||
from ..common_utils import VertexAIError
|
||||
from ..vertex_llm_base import VertexBase
|
||||
|
||||
|
||||
class VertexAIGemmaModels(VertexBase):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
encoding,
|
||||
logging_obj,
|
||||
api_base: Optional[str],
|
||||
optional_params: dict,
|
||||
custom_prompt_dict: dict,
|
||||
headers: Optional[dict],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
litellm_params: dict,
|
||||
vertex_project=None,
|
||||
vertex_location=None,
|
||||
vertex_credentials=None,
|
||||
logger_fn=None,
|
||||
acompletion: bool = False,
|
||||
client=None,
|
||||
):
|
||||
"""
|
||||
Handles calling Vertex AI Gemma Models
|
||||
|
||||
Sent to this route when `model` is in the format `vertex_ai/gemma/{MODEL_NAME}`
|
||||
"""
|
||||
try:
|
||||
import vertexai
|
||||
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexLLM,
|
||||
)
|
||||
from litellm.llms.vertex_ai.vertex_gemma_models.transformation import (
|
||||
VertexGemmaConfig,
|
||||
)
|
||||
except Exception as e:
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message=f"""vertexai import failed please run `pip install -U "google-cloud-aiplatform>=1.38"`. Got error: {e}""",
|
||||
)
|
||||
|
||||
if not (
|
||||
hasattr(vertexai, "preview") or hasattr(vertexai.preview, "language_models")
|
||||
):
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""",
|
||||
)
|
||||
try:
|
||||
model = model.replace("gemma/", "")
|
||||
vertex_httpx_logic = VertexLLM()
|
||||
|
||||
access_token, project_id = vertex_httpx_logic._ensure_access_token(
|
||||
credentials=vertex_credentials,
|
||||
project_id=vertex_project,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
gemma_transformation = VertexGemmaConfig()
|
||||
|
||||
## CONSTRUCT API BASE
|
||||
stream: bool = optional_params.get("stream", False) or False
|
||||
optional_params["stream"] = stream
|
||||
|
||||
# If api_base is not provided, it should be set as an environment variable
|
||||
# or passed explicitly because the endpoint URL is unique per deployment
|
||||
if api_base is None:
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message="api_base is required for Vertex AI Gemma models. Please provide the full endpoint URL.",
|
||||
)
|
||||
|
||||
# Check if we need to append :predict
|
||||
if not api_base.endswith(":predict"):
|
||||
_, api_base = self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
custom_llm_provider="vertex_ai",
|
||||
gemini_api_key=None,
|
||||
endpoint="predict",
|
||||
stream=stream,
|
||||
auth_header=None,
|
||||
url=api_base,
|
||||
)
|
||||
# If api_base already ends with :predict, use it as-is
|
||||
|
||||
# Use the custom transformation handler for gemma models
|
||||
return gemma_transformation.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
api_key=access_token,
|
||||
custom_prompt_dict=custom_prompt_dict,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
acompletion=acompletion,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
encoding=encoding,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
if hasattr(e, "status_code"):
|
||||
raise e
|
||||
raise VertexAIError(status_code=500, message=str(e))
|
||||
|
||||
354
litellm/llms/vertex_ai/vertex_gemma_models/transformation.py
Normal file
|
|
@ -0,0 +1,354 @@
|
|||
"""
|
||||
Transformation logic for Vertex AI Gemma Models
|
||||
|
||||
Handles the custom request/response format:
|
||||
- Request: Wraps messages in 'instances' with @requestFormat: "chatCompletions"
|
||||
- Response: Extracts data from 'predictions' wrapper
|
||||
|
||||
The actual message transformation reuses OpenAIGPTConfig since Gemma uses OpenAI-compatible format.
|
||||
"""
|
||||
|
||||
from typing import Any, Callable, Dict, List, Optional, Union, cast
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
class VertexGemmaConfig(OpenAIGPTConfig):
|
||||
"""
|
||||
Configuration and transformation class for Vertex AI Gemma models
|
||||
|
||||
Extends OpenAIGPTConfig to wrap/unwrap the instances/predictions format
|
||||
used by Vertex AI's Gemma deployment endpoint.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def should_fake_stream(
|
||||
self,
|
||||
model: Optional[str],
|
||||
stream: Optional[bool],
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Vertex AI Gemma models do not support streaming.
|
||||
Return True to enable fake streaming on the client side.
|
||||
"""
|
||||
return True
|
||||
|
||||
def _handle_fake_stream_response(
|
||||
self,
|
||||
model_response: ModelResponse,
|
||||
stream: bool,
|
||||
) -> Union[ModelResponse, Any]:
|
||||
"""
|
||||
Helper method to return fake stream iterator if streaming is requested.
|
||||
|
||||
Args:
|
||||
model_response: The completed model response
|
||||
stream: Whether streaming was requested
|
||||
|
||||
Returns:
|
||||
MockResponseIterator if stream=True, otherwise the model_response
|
||||
"""
|
||||
if stream:
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
return MockResponseIterator(model_response=model_response)
|
||||
return model_response
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform request to Vertex Gemma format.
|
||||
|
||||
Uses parent class to create OpenAI-compatible request, then wraps it
|
||||
in the Vertex Gemma instances format.
|
||||
"""
|
||||
# Get the base OpenAI request from parent class
|
||||
openai_request = super().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
# Remove params not needed/supported by Vertex Gemma
|
||||
openai_request.pop("model", None)
|
||||
openai_request.pop("stream", None) # Streaming not supported, will be faked client-side
|
||||
openai_request.pop("stream_options", None) # Stream options not supported
|
||||
|
||||
# Wrap in Vertex Gemma format
|
||||
return {
|
||||
"instances": [
|
||||
{
|
||||
"@requestFormat": "chatCompletions",
|
||||
**openai_request,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
def _unwrap_predictions_response(
|
||||
self,
|
||||
response_json: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Unwrap the Vertex Gemma predictions format to OpenAI format.
|
||||
|
||||
Vertex Gemma wraps the OpenAI-compatible response in a 'predictions' field.
|
||||
This method extracts it so the parent class can process it normally.
|
||||
"""
|
||||
if "predictions" not in response_json:
|
||||
raise BaseLLMException(
|
||||
status_code=422,
|
||||
message="Invalid response format: missing 'predictions' field",
|
||||
)
|
||||
|
||||
return response_json["predictions"]
|
||||
|
||||
def completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
api_key: str,
|
||||
custom_prompt_dict: dict,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
optional_params: dict,
|
||||
acompletion: bool,
|
||||
litellm_params: dict,
|
||||
logger_fn: Optional[Callable] = None,
|
||||
client: Optional[httpx.Client] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
encoding=None,
|
||||
custom_llm_provider: str = "vertex_ai",
|
||||
):
|
||||
"""
|
||||
Make completion request to Vertex Gemma endpoint.
|
||||
Supports both sync and async requests with fake streaming.
|
||||
"""
|
||||
if acompletion:
|
||||
return self._async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
timeout=timeout,
|
||||
encoding=encoding,
|
||||
)
|
||||
else:
|
||||
return self._sync_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
timeout=timeout,
|
||||
encoding=encoding,
|
||||
)
|
||||
|
||||
def _sync_completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
api_key: str,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding: Any,
|
||||
):
|
||||
"""Synchronous completion request"""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
# Check if streaming is requested (will be faked)
|
||||
stream = optional_params.get("stream", False)
|
||||
|
||||
# Transform the request using parent class methods
|
||||
request_data = self.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params.copy(),
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Set up headers
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
# Log the request
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": request_data,
|
||||
"api_base": api_base,
|
||||
},
|
||||
)
|
||||
|
||||
# Make the HTTP request
|
||||
http_handler = HTTPHandler(concurrent_limit=1)
|
||||
response = http_handler.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BaseLLMException(
|
||||
status_code=response.status_code,
|
||||
message=f"Request failed: {response.text}",
|
||||
)
|
||||
|
||||
response_json = response.json()
|
||||
|
||||
# Unwrap predictions to get OpenAI-compatible response
|
||||
openai_response = self._unwrap_predictions_response(response_json)
|
||||
|
||||
# Use litellm's standard response converter
|
||||
model_response = cast(
|
||||
ModelResponse,
|
||||
convert_to_model_response_object(
|
||||
response_object=openai_response,
|
||||
model_response_object=model_response,
|
||||
_response_headers={},
|
||||
),
|
||||
)
|
||||
|
||||
# Ensure model is set correctly
|
||||
model_response.model = model
|
||||
|
||||
# Log the response
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=response_json,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
)
|
||||
|
||||
# Return fake stream iterator if streaming was requested
|
||||
return self._handle_fake_stream_response(model_response=model_response, stream=stream)
|
||||
|
||||
async def _async_completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
api_key: str,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding: Any,
|
||||
):
|
||||
"""Asynchronous completion request"""
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
# Check if streaming is requested (will be faked)
|
||||
stream = optional_params.get("stream", False)
|
||||
|
||||
# Transform the request using parent class async methods
|
||||
request_data = await self.async_transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params.copy(),
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Set up headers
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
# Log the request
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": request_data,
|
||||
"api_base": api_base,
|
||||
},
|
||||
)
|
||||
|
||||
# Make the HTTP request
|
||||
http_handler = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.VERTEX_AI,
|
||||
)
|
||||
response = await http_handler.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BaseLLMException(
|
||||
status_code=response.status_code,
|
||||
message=f"Request failed: {response.text}",
|
||||
)
|
||||
|
||||
response_json = response.json()
|
||||
|
||||
# Unwrap predictions to get OpenAI-compatible response
|
||||
openai_response = self._unwrap_predictions_response(response_json)
|
||||
|
||||
# Use litellm's standard response converter
|
||||
model_response = cast(
|
||||
ModelResponse,
|
||||
convert_to_model_response_object(
|
||||
response_object=openai_response,
|
||||
model_response_object=model_response,
|
||||
_response_headers={},
|
||||
),
|
||||
)
|
||||
|
||||
# Ensure model is set correctly
|
||||
model_response.model = model
|
||||
|
||||
# Log the response
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=response_json,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
)
|
||||
|
||||
# Return fake stream iterator if streaming was requested
|
||||
return self._handle_fake_stream_response(model_response=model_response, stream=stream)
|
||||
|
||||
|
|
@ -4,10 +4,14 @@ Translation from OpenAI's `/chat/completions` endpoint to IBM WatsonX's `/text/c
|
|||
Docs: https://cloud.ibm.com/apidocs/watsonx-ai#text-chat
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import Dict, List, Optional, Tuple, Union
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.watsonx import WatsonXAIEndpoint, WatsonXAPIParams
|
||||
from litellm.types.llms.watsonx import (
|
||||
WatsonXAIEndpoint,
|
||||
WatsonXAPIParams,
|
||||
WatsonXModelPattern,
|
||||
)
|
||||
|
||||
from ....utils import _remove_additional_properties, _remove_strict_from_schema
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
|
@ -120,3 +124,95 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig):
|
|||
None if model.startswith("deployment/") else api_params["project_id"]
|
||||
)
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _apply_prompt_template_core(model: str, messages: List[Dict[str, str]], hf_template_fn) -> Optional[str]:
|
||||
"""Core logic for applying prompt templates"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
custom_prompt,
|
||||
ibm_granite_pt,
|
||||
mistral_instruct_pt,
|
||||
)
|
||||
|
||||
if WatsonXModelPattern.GRANITE_CHAT.value in model:
|
||||
return ibm_granite_pt(messages=messages)
|
||||
elif WatsonXModelPattern.IBM_MISTRAL.value in model:
|
||||
return mistral_instruct_pt(messages=messages)
|
||||
elif WatsonXModelPattern.GPT_OSS.value in model:
|
||||
hf_model = model.split("watsonx/")[-1] if "watsonx/" in model else model
|
||||
try:
|
||||
return hf_template_fn(model=hf_model, messages=messages)
|
||||
except Exception:
|
||||
pass
|
||||
elif WatsonXModelPattern.LLAMA3_INSTRUCT.value in model:
|
||||
return custom_prompt(
|
||||
role_dict={
|
||||
"system": {"pre_message": "<|start_header_id|>system<|end_header_id|>\n", "post_message": "<|eot_id|>"},
|
||||
"user": {"pre_message": "<|start_header_id|>user<|end_header_id|>\n", "post_message": "<|eot_id|>"},
|
||||
"assistant": {"pre_message": "<|start_header_id|>assistant<|end_header_id|>\n", "post_message": "<|eot_id|>"},
|
||||
},
|
||||
messages=messages,
|
||||
initial_prompt_value="<|begin_of_text|>",
|
||||
final_prompt_value="<|start_header_id|>assistant<|end_header_id|>\n",
|
||||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def aapply_prompt_template(model: str, messages: List[Dict[str, str]]) -> Optional[str]:
|
||||
"""Apply prompt template (async version)"""
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
ahf_chat_template,
|
||||
custom_prompt,
|
||||
hf_chat_template,
|
||||
ibm_granite_pt,
|
||||
mistral_instruct_pt,
|
||||
)
|
||||
|
||||
if WatsonXModelPattern.GRANITE_CHAT.value in model:
|
||||
return ibm_granite_pt(messages=messages)
|
||||
elif WatsonXModelPattern.IBM_MISTRAL.value in model:
|
||||
return mistral_instruct_pt(messages=messages)
|
||||
elif WatsonXModelPattern.GPT_OSS.value in model:
|
||||
hf_model = model.split("watsonx/")[-1] if "watsonx/" in model else model
|
||||
try:
|
||||
# Use sync if cached, async if not
|
||||
if hf_model in litellm.known_tokenizer_config:
|
||||
return hf_chat_template(model=hf_model, messages=messages)
|
||||
else:
|
||||
return await ahf_chat_template(model=hf_model, messages=messages)
|
||||
except Exception:
|
||||
pass
|
||||
elif WatsonXModelPattern.LLAMA3_INSTRUCT.value in model:
|
||||
return custom_prompt(
|
||||
role_dict={
|
||||
"system": {
|
||||
"pre_message": "<|start_header_id|>system<|end_header_id|>\n",
|
||||
"post_message": "<|eot_id|>",
|
||||
},
|
||||
"user": {
|
||||
"pre_message": "<|start_header_id|>user<|end_header_id|>\n",
|
||||
"post_message": "<|eot_id|>",
|
||||
},
|
||||
"assistant": {
|
||||
"pre_message": "<|start_header_id|>assistant<|end_header_id|>\n",
|
||||
"post_message": "<|eot_id|>",
|
||||
},
|
||||
},
|
||||
messages=messages,
|
||||
initial_prompt_value="<|begin_of_text|>",
|
||||
final_prompt_value="<|start_header_id|>assistant<|end_header_id|>\n",
|
||||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def apply_prompt_template(model: str, messages: List[Dict[str, str]]) -> Optional[str]:
|
||||
"""Apply prompt template (sync version)"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
hf_chat_template,
|
||||
)
|
||||
|
||||
return IBMWatsonXChatConfig._apply_prompt_template_core(
|
||||
model=model, messages=messages, hf_template_fn=hf_chat_template
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -131,34 +131,102 @@ def _get_api_params(
|
|||
)
|
||||
|
||||
|
||||
def convert_watsonx_messages_to_prompt(
|
||||
async def _aconvert_watsonx_messages_core(
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
provider: str,
|
||||
custom_prompt_dict: Dict,
|
||||
apply_template_fn,
|
||||
) -> str:
|
||||
"""Async core logic for converting watsonx messages to prompt"""
|
||||
from litellm.types.llms.watsonx import WatsonXModelPattern
|
||||
|
||||
# handle anthropic prompts and amazon titan prompts
|
||||
if model in custom_prompt_dict:
|
||||
# check if the model has a registered custom prompt
|
||||
model_prompt_dict = custom_prompt_dict[model]
|
||||
prompt = ptf.custom_prompt(
|
||||
return ptf.custom_prompt(
|
||||
messages=messages,
|
||||
role_dict=model_prompt_dict.get(
|
||||
"role_dict", model_prompt_dict.get("roles")
|
||||
),
|
||||
role_dict=model_prompt_dict.get("role_dict", model_prompt_dict.get("roles")),
|
||||
initial_prompt_value=model_prompt_dict.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_dict.get("final_prompt_value", ""),
|
||||
bos_token=model_prompt_dict.get("bos_token", ""),
|
||||
eos_token=model_prompt_dict.get("eos_token", ""),
|
||||
)
|
||||
return prompt
|
||||
elif provider == "ibm-mistralai":
|
||||
prompt = ptf.mistral_instruct_pt(messages=messages)
|
||||
elif provider == WatsonXModelPattern.IBM_MISTRALAI.value:
|
||||
return ptf.mistral_instruct_pt(messages=messages)
|
||||
else:
|
||||
prompt: str = ptf.prompt_factory( # type: ignore
|
||||
# Try applying specific template first
|
||||
result = await apply_template_fn(model=model, messages=messages)
|
||||
if result:
|
||||
return result
|
||||
# Fallback to default
|
||||
return ptf.prompt_factory(
|
||||
model=model, messages=messages, custom_llm_provider="watsonx"
|
||||
) # type: ignore
|
||||
|
||||
|
||||
def _convert_watsonx_messages_core(
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
provider: str,
|
||||
custom_prompt_dict: Dict,
|
||||
apply_template_fn,
|
||||
) -> str:
|
||||
"""Sync core logic for converting watsonx messages to prompt"""
|
||||
from litellm.types.llms.watsonx import WatsonXModelPattern
|
||||
|
||||
# handle anthropic prompts and amazon titan prompts
|
||||
if model in custom_prompt_dict:
|
||||
model_prompt_dict = custom_prompt_dict[model]
|
||||
return ptf.custom_prompt(
|
||||
messages=messages,
|
||||
role_dict=model_prompt_dict.get("role_dict", model_prompt_dict.get("roles")),
|
||||
initial_prompt_value=model_prompt_dict.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_dict.get("final_prompt_value", ""),
|
||||
bos_token=model_prompt_dict.get("bos_token", ""),
|
||||
eos_token=model_prompt_dict.get("eos_token", ""),
|
||||
)
|
||||
return prompt
|
||||
elif provider == WatsonXModelPattern.IBM_MISTRALAI.value:
|
||||
return ptf.mistral_instruct_pt(messages=messages)
|
||||
else:
|
||||
# Try applying specific template first
|
||||
result = apply_template_fn(model=model, messages=messages)
|
||||
if result:
|
||||
return result
|
||||
# Fallback to default
|
||||
return ptf.prompt_factory(
|
||||
model=model, messages=messages, custom_llm_provider="watsonx"
|
||||
) # type: ignore
|
||||
|
||||
|
||||
async def aconvert_watsonx_messages_to_prompt(
|
||||
model: str, messages: List[AllMessageValues], provider: str, custom_prompt_dict: Dict
|
||||
) -> str:
|
||||
"""Async version of convert_watsonx_messages_to_prompt"""
|
||||
from litellm.llms.watsonx.chat.transformation import IBMWatsonXChatConfig
|
||||
|
||||
return await _aconvert_watsonx_messages_core(
|
||||
model=model,
|
||||
messages=messages,
|
||||
provider=provider,
|
||||
custom_prompt_dict=custom_prompt_dict,
|
||||
apply_template_fn=IBMWatsonXChatConfig.aapply_prompt_template,
|
||||
)
|
||||
|
||||
|
||||
def convert_watsonx_messages_to_prompt(
|
||||
model: str, messages: List[AllMessageValues], provider: str, custom_prompt_dict: Dict
|
||||
) -> str:
|
||||
"""Sync version of convert_watsonx_messages_to_prompt"""
|
||||
from litellm.llms.watsonx.chat.transformation import IBMWatsonXChatConfig
|
||||
|
||||
return _convert_watsonx_messages_core(
|
||||
model=model,
|
||||
messages=messages,
|
||||
provider=provider,
|
||||
custom_prompt_dict=custom_prompt_dict,
|
||||
apply_template_fn=IBMWatsonXChatConfig.apply_prompt_template,
|
||||
)
|
||||
|
||||
|
||||
# Mixin class for shared IBM Watson X functionality
|
||||
|
|
|
|||
|
|
@ -228,39 +228,35 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig):
|
|||
"us-south",
|
||||
]
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: Dict,
|
||||
litellm_params: Dict,
|
||||
headers: Dict,
|
||||
) -> Dict:
|
||||
provider = model.split("/")[0]
|
||||
prompt = convert_watsonx_messages_to_prompt(
|
||||
model=model,
|
||||
messages=messages,
|
||||
provider=provider,
|
||||
custom_prompt_dict={},
|
||||
)
|
||||
def _build_request_payload(self, model: str, prompt: str, optional_params: Dict) -> Dict:
|
||||
"""Shared logic to build request payload"""
|
||||
extra_body_params = optional_params.pop("extra_body", {})
|
||||
optional_params.update(extra_body_params)
|
||||
watsonx_api_params = _get_api_params(params=optional_params)
|
||||
|
||||
watsonx_auth_payload = self._prepare_payload(
|
||||
model=model,
|
||||
api_params=watsonx_api_params,
|
||||
)
|
||||
|
||||
# init the payload to the text generation call
|
||||
payload = {
|
||||
watsonx_auth_payload = self._prepare_payload(model=model, api_params=watsonx_api_params)
|
||||
|
||||
return {
|
||||
"input": prompt,
|
||||
"moderations": optional_params.pop("moderations", {}),
|
||||
"parameters": optional_params,
|
||||
**watsonx_auth_payload,
|
||||
}
|
||||
|
||||
return payload
|
||||
async def atransform_request(self, model: str, messages: List[AllMessageValues], optional_params: Dict, litellm_params: Dict, headers: Dict) -> Dict:
|
||||
"""Async version of transform_request"""
|
||||
from litellm.llms.watsonx.common_utils import (
|
||||
aconvert_watsonx_messages_to_prompt,
|
||||
)
|
||||
|
||||
provider = model.split("/")[0]
|
||||
prompt = await aconvert_watsonx_messages_to_prompt(model=model, messages=messages, provider=provider, custom_prompt_dict={})
|
||||
return self._build_request_payload(model=model, prompt=prompt, optional_params=optional_params)
|
||||
|
||||
def transform_request(self, model: str, messages: List[AllMessageValues], optional_params: Dict, litellm_params: Dict, headers: Dict) -> Dict:
|
||||
"""Sync version of transform_request"""
|
||||
provider = model.split("/")[0]
|
||||
prompt = convert_watsonx_messages_to_prompt(model=model, messages=messages, provider=provider, custom_prompt_dict={})
|
||||
return self._build_request_payload(model=model, prompt=prompt, optional_params=optional_params)
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -85,6 +85,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
VertexAIModelRoute,
|
||||
get_vertex_ai_model_route,
|
||||
)
|
||||
from litellm.realtime_api.main import _realtime_health_check
|
||||
from litellm.secret_managers.main import get_secret_bool, get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -150,7 +154,6 @@ from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
|
|||
from .llms.bedrock.embed.embedding import BedrockEmbedding
|
||||
from .llms.bedrock.image.image_handler import BedrockImageGeneration
|
||||
from .llms.bytez.chat.transformation import BytezChatConfig
|
||||
from .llms.lemonade.chat.transformation import LemonadeChatConfig
|
||||
from .llms.codestral.completion.handler import CodestralTextCompletion
|
||||
from .llms.cohere.embed import handler as cohere_embed
|
||||
from .llms.custom_httpx.aiohttp_handler import BaseLLMAIOHTTPHandler
|
||||
|
|
@ -162,6 +165,7 @@ from .llms.gemini.common_utils import get_api_key_from_env
|
|||
from .llms.groq.chat.handler import GroqChatCompletion
|
||||
from .llms.heroku.chat.transformation import HerokuChatConfig
|
||||
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
|
||||
from .llms.lemonade.chat.transformation import LemonadeChatConfig
|
||||
from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion
|
||||
from .llms.oci.chat.transformation import OCIChatConfig
|
||||
from .llms.ollama.completion import handler as ollama
|
||||
|
|
@ -192,6 +196,7 @@ from .llms.vertex_ai.multimodal_embeddings.embedding_handler import (
|
|||
from .llms.vertex_ai.text_to_speech.text_to_speech_handler import VertexTextToSpeechAPI
|
||||
from .llms.vertex_ai.vertex_ai_partner_models.main import VertexAIPartnerModels
|
||||
from .llms.vertex_ai.vertex_embeddings.embedding_handler import VertexEmbedding
|
||||
from .llms.vertex_ai.vertex_gemma_models.main import VertexAIGemmaModels
|
||||
from .llms.vertex_ai.vertex_model_garden.main import VertexAIModelGardenModels
|
||||
from .llms.vllm.completion import handler as vllm_handler
|
||||
from .llms.watsonx.chat.handler import WatsonXChatHandler
|
||||
|
|
@ -255,6 +260,7 @@ vertex_multimodal_embedding = VertexMultimodalEmbedding()
|
|||
vertex_image_generation = VertexImageGeneration()
|
||||
google_batch_embeddings = GoogleBatchEmbeddings()
|
||||
vertex_partner_models_chat_completion = VertexAIPartnerModels()
|
||||
vertex_gemma_chat_completion = VertexAIGemmaModels()
|
||||
vertex_model_garden_chat_completion = VertexAIModelGardenModels()
|
||||
vertex_text_to_speech = VertexTextToSpeechAPI()
|
||||
sagemaker_llm = SagemakerLLM()
|
||||
|
|
@ -2875,7 +2881,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
extra_headers=headers,
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
vertex_ai_project = (
|
||||
optional_params.pop("vertex_project", None)
|
||||
or optional_params.pop("vertex_ai_project", None)
|
||||
|
|
@ -2897,7 +2903,9 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE")
|
||||
|
||||
new_params = safe_deep_copy(optional_params or {})
|
||||
if vertex_partner_models_chat_completion.is_vertex_partner_model(model):
|
||||
model_route = get_vertex_ai_model_route(model=model, litellm_params=litellm_params)
|
||||
|
||||
if model_route == VertexAIModelRoute.PARTNER_MODELS:
|
||||
model_response = vertex_partner_models_chat_completion.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
@ -2918,10 +2926,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
timeout=timeout,
|
||||
client=client,
|
||||
)
|
||||
elif "gemini" in model or (
|
||||
litellm_params.get("base_model") is not None
|
||||
and "gemini" in litellm_params["base_model"]
|
||||
):
|
||||
elif model_route == VertexAIModelRoute.GEMINI:
|
||||
model_response = vertex_chat_completion.completion( # type: ignore
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
@ -2943,7 +2948,29 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
api_base=api_base,
|
||||
extra_headers=headers,
|
||||
)
|
||||
elif "openai" in model:
|
||||
elif model_route == VertexAIModelRoute.GEMMA:
|
||||
# Vertex Gemma Models with custom prediction endpoint
|
||||
model_response = vertex_gemma_chat_completion.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
optional_params=new_params,
|
||||
litellm_params=litellm_params, # type: ignore
|
||||
logger_fn=logger_fn,
|
||||
encoding=encoding,
|
||||
api_base=api_base,
|
||||
vertex_location=vertex_ai_location,
|
||||
vertex_project=vertex_ai_project,
|
||||
vertex_credentials=vertex_credentials,
|
||||
logging_obj=logging,
|
||||
acompletion=acompletion,
|
||||
headers=headers,
|
||||
custom_prompt_dict=custom_prompt_dict,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
)
|
||||
elif model_route == VertexAIModelRoute.MODEL_GARDEN:
|
||||
# Vertex Model Garden - OpenAI compatible models
|
||||
model_response = vertex_model_garden_chat_completion.completion(
|
||||
model=model,
|
||||
|
|
@ -2965,7 +2992,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
timeout=timeout,
|
||||
client=client,
|
||||
)
|
||||
else:
|
||||
else: # VertexAIModelRoute.NON_GEMINI
|
||||
model_response = vertex_ai_non_gemini.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
@ -3969,6 +3996,7 @@ def embedding( # noqa: PLR0915
|
|||
"""
|
||||
azure = kwargs.get("azure", None)
|
||||
client = kwargs.pop("client", None)
|
||||
shared_session = kwargs.get("shared_session", None)
|
||||
max_retries = kwargs.get("max_retries", None)
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
mock_response: Optional[List[float]] = kwargs.get("mock_response", None) # type: ignore
|
||||
|
|
@ -4158,6 +4186,7 @@ def embedding( # noqa: PLR0915
|
|||
client=client,
|
||||
aembedding=aembedding,
|
||||
max_retries=max_retries,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
elif custom_llm_provider == "databricks":
|
||||
api_base = api_base or litellm.api_base or get_secret("DATABRICKS_API_BASE") # type: ignore
|
||||
|
|
|
|||
|
|
@ -866,6 +866,36 @@
|
|||
"mode": "audio_transcription",
|
||||
"output_cost_per_second": 0.0
|
||||
},
|
||||
"au.anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6.6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.475e-05,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 200000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 346
|
||||
},
|
||||
"azure/ada": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -3164,6 +3194,42 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure_ai/Phi-4-mini-reasoning": {
|
||||
"input_cost_per_token": 8e-08,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.2e-07,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/microsoft/",
|
||||
"supports_function_calling": true
|
||||
},
|
||||
"azure_ai/Phi-4-reasoning": {
|
||||
"input_cost_per_token": 1.25e-07,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-07,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/microsoft/",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_reasoning": true
|
||||
},
|
||||
"azure_ai/MAI-DS-R1": {
|
||||
"input_cost_per_token": 1.35e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5.4e-06,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/microsoft/",
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure_ai/cohere-rerank-v3-english": {
|
||||
"input_cost_per_query": 0.002,
|
||||
"input_cost_per_token": 0.0,
|
||||
|
|
@ -4775,7 +4841,7 @@
|
|||
"input_cost_per_token_above_200k_tokens": 6e-06,
|
||||
"litellm_provider": "anthropic",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_output_tokens": 64000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
|
|
@ -12975,16 +13041,18 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gpt-5-codex": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"gpt-5-pro-2025-10-06": {
|
||||
"input_cost_per_token": 1.5e-05,
|
||||
"input_cost_per_token_batches": 7.5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 400000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"max_output_tokens": 272000,
|
||||
"max_tokens": 272000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 1.2e-04,
|
||||
"output_cost_per_token_batches": 6e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
|
|
@ -12995,13 +13063,16 @@
|
|||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_native_streaming": false,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gpt-5-2025-08-07": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
|
|
@ -13303,8 +13374,7 @@
|
|||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true
|
||||
]
|
||||
},
|
||||
"gpt-realtime": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
|
|
@ -16565,6 +16635,42 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_response_schema": false
|
||||
},
|
||||
"oci/cohere.command-latest": {
|
||||
"input_cost_per_token": 1.56e-06,
|
||||
"litellm_provider": "oci",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.56e-06,
|
||||
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": false
|
||||
},
|
||||
"oci/cohere.command-a-03-2025": {
|
||||
"input_cost_per_token": 1.56e-06,
|
||||
"litellm_provider": "oci",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 4000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.56e-06,
|
||||
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": false
|
||||
},
|
||||
"oci/cohere.command-plus-latest": {
|
||||
"input_cost_per_token": 1.56e-06,
|
||||
"litellm_provider": "oci",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.56e-06,
|
||||
"source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/",
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": false
|
||||
},
|
||||
"ollama/codegeex4": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "ollama",
|
||||
|
|
@ -17098,6 +17204,25 @@
|
|||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 159
|
||||
},
|
||||
"openrouter/anthropic/claude-sonnet-4.5": {
|
||||
"input_cost_per_image": 0.0048,
|
||||
"input_cost_per_token": 3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.25e-05,
|
||||
"litellm_provider": "openrouter",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"tool_use_system_prompt_tokens": 159
|
||||
},
|
||||
"openrouter/bytedance/ui-tars-1.5-7b": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "openrouter",
|
||||
|
|
@ -19576,6 +19701,22 @@
|
|||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0
|
||||
},
|
||||
"together_ai/baai/bge-base-en-v1.5": {
|
||||
"input_cost_per_token": 8e-09,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 512,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 768
|
||||
},
|
||||
"together_ai/BAAI/bge-base-en-v1.5": {
|
||||
"input_cost_per_token": 8e-09,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 512,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 768
|
||||
},
|
||||
"together-ai-up-to-4b": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "together_ai",
|
||||
|
|
@ -19836,6 +19977,39 @@
|
|||
"supports_parallel_function_calling": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"together_ai/moonshotai/Kimi-K2-Instruct-0905": {
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-06,
|
||||
"source": "https://www.together.ai/models/kimi-k2-0905",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"together_ai/Qwen/Qwen3-Next-80B-A3B-Instruct": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-06,
|
||||
"source": "https://www.together.ai/models/qwen3-next-80b-a3b-instruct",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-06,
|
||||
"source": "https://www.together.ai/models/qwen3-next-80b-a3b-thinking",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"tts-1": {
|
||||
"input_cost_per_character": 1.5e-05,
|
||||
"litellm_provider": "openai",
|
||||
|
|
@ -20045,12 +20219,12 @@
|
|||
},
|
||||
"us.anthropic.claude-sonnet-4-5-20250929-v1:0": {
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 33e-07,
|
||||
"input_cost_per_token": 33e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 66e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 6.6e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 2.475e-05,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 66e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 64000,
|
||||
|
|
|
|||
|
|
@ -109,10 +109,21 @@ class MCPRequestHandler:
|
|||
request.body = mock_body # type: ignore
|
||||
if ".well-known" in str(request.url): # public routes
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
# elif litellm_api_key == "":
|
||||
# from fastapi import HTTPException
|
||||
|
||||
# raise HTTPException(
|
||||
# status_code=401,
|
||||
# detail="LiteLLM API key is missing. Please add it or use OAuth authentication.",
|
||||
# headers={
|
||||
# "WWW-Authenticate": f'Bearer resource_metadata=f"{request.base_url}/.well-known/oauth-protected-resource"',
|
||||
# },
|
||||
# )
|
||||
else:
|
||||
validated_user_api_key_auth = await user_api_key_auth(
|
||||
api_key=litellm_api_key, request=request
|
||||
)
|
||||
|
||||
return (
|
||||
validated_user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
|
|
@ -344,14 +355,14 @@ class MCPRequestHandler:
|
|||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
|
||||
# Already loaded
|
||||
if user_api_key_auth.object_permission:
|
||||
return user_api_key_auth.object_permission
|
||||
|
||||
|
||||
# Need to fetch from DB
|
||||
if user_api_key_auth.object_permission_id and prisma_client:
|
||||
return await get_object_permission(
|
||||
|
|
@ -361,7 +372,7 @@ class MCPRequestHandler:
|
|||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -369,16 +380,19 @@ class MCPRequestHandler:
|
|||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
"""Helper to get team object_permission from cache or DB."""
|
||||
from litellm.proxy.auth.auth_checks import get_object_permission, get_team_object
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
get_object_permission,
|
||||
get_team_object,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
|
||||
if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client:
|
||||
return None
|
||||
|
||||
|
||||
# First get the team object (which may have object_permission already loaded)
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
|
|
@ -387,14 +401,14 @@ class MCPRequestHandler:
|
|||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
if not team_obj:
|
||||
return None
|
||||
|
||||
|
||||
# Already loaded
|
||||
if team_obj.object_permission:
|
||||
return team_obj.object_permission
|
||||
|
||||
|
||||
# Need to fetch from DB using object_permission_id
|
||||
if team_obj.object_permission_id:
|
||||
return await get_object_permission(
|
||||
|
|
@ -404,7 +418,7 @@ class MCPRequestHandler:
|
|||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -415,26 +429,38 @@ class MCPRequestHandler:
|
|||
"""
|
||||
Get list of allowed tool names for a specific server based on key/team permissions.
|
||||
Follows same inheritance logic as get_allowed_mcp_servers.
|
||||
|
||||
|
||||
Args:
|
||||
server_id: Server ID to check permissions for
|
||||
user_api_key_auth: User auth
|
||||
|
||||
|
||||
Returns:
|
||||
List[str] if restrictions exist, None if no restrictions (allow all)
|
||||
"""
|
||||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
|
||||
try:
|
||||
# Get key and team object permissions
|
||||
key_obj_perm = await MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
team_obj_perm = await MCPRequestHandler._get_team_object_permission(user_api_key_auth)
|
||||
|
||||
key_obj_perm = await MCPRequestHandler._get_key_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
team_obj_perm = await MCPRequestHandler._get_team_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
||||
# Extract tool permissions for this server
|
||||
key_tools = key_obj_perm.mcp_tool_permissions.get(server_id) if key_obj_perm and key_obj_perm.mcp_tool_permissions else None
|
||||
team_tools = team_obj_perm.mcp_tool_permissions.get(server_id) if team_obj_perm and team_obj_perm.mcp_tool_permissions else None
|
||||
|
||||
key_tools = (
|
||||
key_obj_perm.mcp_tool_permissions.get(server_id)
|
||||
if key_obj_perm and key_obj_perm.mcp_tool_permissions
|
||||
else None
|
||||
)
|
||||
team_tools = (
|
||||
team_obj_perm.mcp_tool_permissions.get(server_id)
|
||||
if team_obj_perm and team_obj_perm.mcp_tool_permissions
|
||||
else None
|
||||
)
|
||||
|
||||
# Apply same inheritance logic as get_allowed_mcp_servers
|
||||
if team_tools:
|
||||
if key_tools:
|
||||
|
|
@ -446,7 +472,7 @@ class MCPRequestHandler:
|
|||
else:
|
||||
# No team restrictions → use key restrictions
|
||||
return key_tools
|
||||
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}")
|
||||
return None
|
||||
|
|
@ -459,12 +485,12 @@ class MCPRequestHandler:
|
|||
) -> bool:
|
||||
"""
|
||||
Check if a specific tool is allowed for a server based on key/team permissions.
|
||||
|
||||
|
||||
Args:
|
||||
tool_name: Name of the tool to check
|
||||
server_id: Server ID
|
||||
user_api_key_auth: User auth
|
||||
|
||||
|
||||
Returns:
|
||||
True if allowed, False if blocked
|
||||
"""
|
||||
|
|
@ -472,15 +498,15 @@ class MCPRequestHandler:
|
|||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
|
||||
# None means no restrictions (allow all)
|
||||
if allowed_tools is None:
|
||||
return True
|
||||
|
||||
|
||||
# Empty list means no tools allowed
|
||||
if not allowed_tools:
|
||||
return False
|
||||
|
||||
|
||||
# Check if tool is in allowed list
|
||||
return tool_name in allowed_tools
|
||||
|
||||
|
|
@ -555,7 +581,7 @@ class MCPRequestHandler:
|
|||
) -> List[str]:
|
||||
"""
|
||||
Get allowed MCP servers for a team.
|
||||
|
||||
|
||||
Uses the helper _get_team_object_permission which:
|
||||
1. First checks if object_permission is already loaded on the team
|
||||
2. If not, fetches from DB using object_permission_id if it exists
|
||||
|
|
@ -571,7 +597,7 @@ class MCPRequestHandler:
|
|||
object_permissions = await MCPRequestHandler._get_team_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
||||
|
||||
if object_permissions is None:
|
||||
return []
|
||||
|
||||
|
|
|
|||
|
|
@ -195,6 +195,7 @@ class MCPServerManager:
|
|||
name=name_for_prefix,
|
||||
alias=alias,
|
||||
server_name=server_name,
|
||||
spec_path=server_config.get("spec_path", None),
|
||||
url=server_config.get("url", None) or "",
|
||||
command=server_config.get("command", None) or "",
|
||||
args=server_config.get("args", None) or [],
|
||||
|
|
@ -215,15 +216,167 @@ class MCPServerManager:
|
|||
extra_headers=server_config.get("extra_headers", None),
|
||||
allowed_tools=server_config.get("allowed_tools", None),
|
||||
disallowed_tools=server_config.get("disallowed_tools", None),
|
||||
allowed_params=server_config.get("allowed_params", None),
|
||||
access_groups=server_config.get("access_groups", None),
|
||||
)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
|
||||
# Check if this is an OpenAPI-based server
|
||||
spec_path = server_config.get("spec_path", None)
|
||||
if spec_path:
|
||||
verbose_logger.info(
|
||||
f"Loading OpenAPI spec from {spec_path} for server {server_name}"
|
||||
)
|
||||
self._register_openapi_tools(
|
||||
spec_path=spec_path,
|
||||
server=new_server,
|
||||
base_url=server_config.get("url", ""),
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}"
|
||||
)
|
||||
|
||||
self.initialize_tool_name_to_mcp_server_name_mapping()
|
||||
|
||||
def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str):
|
||||
"""
|
||||
Register tools from an OpenAPI specification for a given server.
|
||||
|
||||
This creates "virtual" MCP tools from OpenAPI endpoints that are:
|
||||
1. Registered in the global tool registry with server prefix
|
||||
2. Mapped to the server for routing
|
||||
3. Executed via the local tool handler
|
||||
|
||||
Args:
|
||||
spec_path: Path to the OpenAPI specification file
|
||||
server: The MCPServer instance to register tools for
|
||||
base_url: Base URL for API calls
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
build_input_schema,
|
||||
create_tool_function,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
get_base_url as get_openapi_base_url,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
load_openapi_spec,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
try:
|
||||
# Load OpenAPI spec
|
||||
spec = load_openapi_spec(spec_path)
|
||||
|
||||
# Use base_url from config if provided, otherwise extract from spec
|
||||
if not base_url:
|
||||
base_url = get_openapi_base_url(spec)
|
||||
|
||||
verbose_logger.info(
|
||||
f"Registering OpenAPI tools for server {server.name} with base URL: {base_url}"
|
||||
)
|
||||
|
||||
# Get server prefix for tool naming
|
||||
server_prefix = get_server_prefix(server)
|
||||
|
||||
# Build headers from server configuration
|
||||
headers = {}
|
||||
|
||||
# Add authentication headers if configured
|
||||
if server.authentication_token:
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
if server.auth_type == MCPAuth.bearer_token:
|
||||
headers["Authorization"] = f"Bearer {server.authentication_token}"
|
||||
elif server.auth_type == MCPAuth.api_key:
|
||||
headers["Authorization"] = f"ApiKey {server.authentication_token}"
|
||||
elif server.auth_type == MCPAuth.basic:
|
||||
headers["Authorization"] = f"Basic {server.authentication_token}"
|
||||
|
||||
# Add any extra headers from server config
|
||||
# Note: extra_headers is a List[str] of header names to forward, not a dict
|
||||
# For OpenAPI tools, we'll just use the authentication headers
|
||||
# If extra_headers were needed, they would be processed separately
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Using headers for OpenAPI tools (excluding sensitive values): "
|
||||
f"{list(headers.keys())}"
|
||||
)
|
||||
|
||||
# Extract and register tools from OpenAPI paths
|
||||
paths = spec.get("paths", {})
|
||||
registered_count = 0
|
||||
|
||||
verbose_logger.debug(f"Processing {len(paths)} paths from OpenAPI spec")
|
||||
|
||||
for path, path_item in paths.items():
|
||||
for method in ["get", "post", "put", "delete", "patch"]:
|
||||
if method not in path_item:
|
||||
continue
|
||||
|
||||
operation = path_item[method]
|
||||
|
||||
# Generate tool name (without prefix initially)
|
||||
operation_id = operation.get(
|
||||
"operationId", f"{method}_{path.replace('/', '_')}"
|
||||
)
|
||||
base_tool_name = operation_id.replace(" ", "_").lower()
|
||||
|
||||
# Add server prefix to tool name
|
||||
prefixed_tool_name = add_server_prefix_to_tool_name(
|
||||
base_tool_name, server_prefix
|
||||
)
|
||||
|
||||
# Get description
|
||||
description = operation.get(
|
||||
"summary",
|
||||
operation.get("description", f"{method.upper()} {path}"),
|
||||
)
|
||||
|
||||
# Build input schema using imported function
|
||||
input_schema = build_input_schema(operation)
|
||||
|
||||
# Create tool function with headers using imported function
|
||||
tool_func = create_tool_function(
|
||||
path, method, operation, base_url, headers=headers
|
||||
)
|
||||
tool_func.__name__ = prefixed_tool_name
|
||||
tool_func.__doc__ = description
|
||||
|
||||
# Register tool with prefixed name in global registry
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=prefixed_tool_name,
|
||||
description=description,
|
||||
input_schema=input_schema,
|
||||
handler=tool_func,
|
||||
)
|
||||
|
||||
# Update tool name to server name mapping (for both prefixed and base names)
|
||||
self.tool_name_to_mcp_server_name_mapping[base_tool_name] = (
|
||||
server_prefix
|
||||
)
|
||||
self.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = (
|
||||
server_prefix
|
||||
)
|
||||
|
||||
registered_count += 1
|
||||
verbose_logger.debug(
|
||||
f"Registered OpenAPI tool: {prefixed_tool_name} for server {server.name}"
|
||||
)
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully registered {registered_count} OpenAPI tools for server {server.name}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Failed to register OpenAPI tools for server {server.name}: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
|
||||
def remove_server(self, mcp_server: LiteLLM_MCPServerTable):
|
||||
"""
|
||||
Remove a server from the registry
|
||||
|
|
@ -469,6 +622,10 @@ class MCPServerManager:
|
|||
Returns:
|
||||
List[MCPTool]: List of tools available on the server with prefixed names
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"Connecting to url: {server.url}")
|
||||
verbose_logger.info(f"_get_tools_from_server for {server.name}...")
|
||||
|
||||
|
|
@ -481,7 +638,14 @@ class MCPServerManager:
|
|||
extra_headers=extra_headers,
|
||||
)
|
||||
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
## HANDLE OPENAPI TOOLS
|
||||
if server.spec_path:
|
||||
_tools = global_mcp_tool_registry.list_tools(tool_prefix=server.name)
|
||||
tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(
|
||||
_tools
|
||||
)
|
||||
else:
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
|
||||
prefixed_or_original_tools = self._create_prefixed_tools(
|
||||
tools, server, add_prefix=add_prefix
|
||||
|
|
@ -597,11 +761,69 @@ class MCPServerManager:
|
|||
Check if the tool is allowed or banned for the given server
|
||||
"""
|
||||
if server.allowed_tools:
|
||||
return tool_name in server.allowed_tools
|
||||
return (
|
||||
tool_name in server.allowed_tools
|
||||
or f"{server.name}-{tool_name}" in server.allowed_tools
|
||||
)
|
||||
if server.disallowed_tools:
|
||||
return tool_name not in server.disallowed_tools
|
||||
return (
|
||||
tool_name not in server.disallowed_tools
|
||||
and f"{server.name}-{tool_name}" not in server.disallowed_tools
|
||||
)
|
||||
return True
|
||||
|
||||
def validate_allowed_params(
|
||||
self, tool_name: str, arguments: Dict[str, Any], server: MCPServer
|
||||
) -> None:
|
||||
"""
|
||||
Filter arguments to only include allowed parameters for the given tool.
|
||||
|
||||
Args:
|
||||
tool_name: Name of the tool (with or without prefix)
|
||||
arguments: Dictionary of arguments to filter
|
||||
server: MCPServer configuration
|
||||
|
||||
Returns:
|
||||
Filtered dictionary containing only allowed parameters
|
||||
|
||||
Raises:
|
||||
HTTPException: If allowed_params is configured for this tool but arguments contain disallowed params
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
get_server_name_prefix_tool_mcp,
|
||||
)
|
||||
|
||||
# If no allowed_params configured, return all arguments
|
||||
if not server.allowed_params:
|
||||
return
|
||||
|
||||
# Get the unprefixed tool name to match against config
|
||||
unprefixed_tool_name, _ = get_server_name_prefix_tool_mcp(tool_name)
|
||||
|
||||
# Check both prefixed and unprefixed tool names
|
||||
allowed_params_list = server.allowed_params.get(
|
||||
tool_name
|
||||
) or server.allowed_params.get(unprefixed_tool_name)
|
||||
|
||||
# If this tool doesn't have allowed_params specified, allow all params
|
||||
if allowed_params_list is None:
|
||||
return None
|
||||
|
||||
# Filter arguments to only include allowed parameters
|
||||
disallowed_params = [
|
||||
param for param in arguments.keys() if param not in allowed_params_list
|
||||
]
|
||||
|
||||
if disallowed_params:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"Parameters {disallowed_params} are not allowed for tool {tool_name}. "
|
||||
f"Allowed parameters: {allowed_params_list}. "
|
||||
f"Contact proxy admin to allow these parameters."
|
||||
},
|
||||
)
|
||||
|
||||
async def check_tool_permission_for_key_team(
|
||||
self,
|
||||
tool_name: str,
|
||||
|
|
@ -621,18 +843,20 @@ class MCPServerManager:
|
|||
Raises:
|
||||
HTTPException: If tool is not allowed for this key/team
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
|
||||
if not user_api_key_auth:
|
||||
return
|
||||
|
||||
|
||||
# Check if tool is allowed
|
||||
is_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
|
||||
tool_name=tool_name,
|
||||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
|
||||
if not is_allowed:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
@ -641,6 +865,64 @@ class MCPServerManager:
|
|||
},
|
||||
)
|
||||
|
||||
async def _call_openapi_tool_handler(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tool_name: str,
|
||||
arguments: Dict[str, Any],
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call an OpenAPI tool handler directly.
|
||||
|
||||
For OpenAPI servers, instead of using MCP protocol, we call the tool handler
|
||||
that was registered during OpenAPI spec parsing. This handler makes direct
|
||||
HTTP requests to the API.
|
||||
|
||||
Args:
|
||||
tool_name: The full tool name (with prefix) to call
|
||||
arguments: Tool arguments to pass to the handler
|
||||
|
||||
Returns:
|
||||
CallToolResult with the response from the API
|
||||
"""
|
||||
from mcp.types import TextContent
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
# Get the tool from the registry
|
||||
tool = global_mcp_tool_registry.get_tool(f"{server.name}-{tool_name}")
|
||||
if tool is None:
|
||||
# Tool not found in registry
|
||||
error_msg = f"OpenAPI tool {tool_name} not found in registry"
|
||||
verbose_logger.error(error_msg)
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text=error_msg)],
|
||||
isError=True,
|
||||
)
|
||||
|
||||
try:
|
||||
# Call the tool handler with the arguments
|
||||
# The handler is an async function that makes the HTTP request
|
||||
handler_result = await tool.handler(**arguments)
|
||||
|
||||
# Convert the handler result (string response) to CallToolResult format
|
||||
result = CallToolResult(
|
||||
content=[TextContent(type="text", text=str(handler_result))],
|
||||
isError=False,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"Error calling OpenAPI tool {tool_name}: {str(e)}"
|
||||
verbose_logger.error(error_msg)
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text=error_msg)],
|
||||
isError=True,
|
||||
)
|
||||
|
||||
async def pre_call_tool_check(
|
||||
self,
|
||||
name: str,
|
||||
|
|
@ -650,7 +932,6 @@ class MCPServerManager:
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
server: MCPServer,
|
||||
):
|
||||
|
||||
## check if the tool is allowed or banned for the given server
|
||||
if not self.check_allowed_or_banned_tools(name, server):
|
||||
raise HTTPException(
|
||||
|
|
@ -667,6 +948,13 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
## filter parameters based on allowed_params configuration
|
||||
self.validate_allowed_params(
|
||||
tool_name=name,
|
||||
arguments=arguments,
|
||||
server=server,
|
||||
)
|
||||
|
||||
pre_hook_kwargs = {
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
|
|
@ -730,6 +1018,143 @@ class MCPServerManager:
|
|||
verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {str(e)}")
|
||||
raise e
|
||||
|
||||
def _create_during_hook_task(
|
||||
self,
|
||||
name: str,
|
||||
arguments: Dict[str, Any],
|
||||
server_name_from_prefix: Optional[str],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
start_time: datetime.datetime,
|
||||
):
|
||||
"""Create and return a during hook task for MCP tool calls."""
|
||||
from litellm.types.llms.base import HiddenParams
|
||||
from litellm.types.mcp import MCPDuringCallRequestObject
|
||||
|
||||
request_obj = MCPDuringCallRequestObject(
|
||||
tool_name=name,
|
||||
arguments=arguments,
|
||||
server_name=server_name_from_prefix,
|
||||
start_time=start_time.timestamp() if start_time else None,
|
||||
hidden_params=HiddenParams(),
|
||||
)
|
||||
|
||||
during_hook_kwargs = {
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
"server_name": server_name_from_prefix,
|
||||
"user_api_key_auth": user_api_key_auth,
|
||||
}
|
||||
|
||||
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(
|
||||
request_obj, during_hook_kwargs
|
||||
)
|
||||
|
||||
return asyncio.create_task(
|
||||
proxy_logging_obj.during_call_hook(
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
data=synthetic_llm_data,
|
||||
call_type="mcp_call", # type: ignore
|
||||
)
|
||||
)
|
||||
|
||||
async def _call_regular_mcp_tool(
|
||||
self,
|
||||
mcp_server: MCPServer,
|
||||
original_tool_name: str,
|
||||
arguments: Dict[str, Any],
|
||||
tasks: List,
|
||||
mcp_auth_header: Optional[str],
|
||||
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
|
||||
oauth2_headers: Optional[Dict[str, str]],
|
||||
raw_headers: Optional[Dict[str, str]],
|
||||
proxy_logging_obj: Optional[ProxyLogging],
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call a regular MCP tool using the MCP client.
|
||||
|
||||
Args:
|
||||
mcp_server: The MCP server configuration
|
||||
original_tool_name: The original tool name (without prefix)
|
||||
arguments: Tool arguments
|
||||
tasks: List of async tasks to append to (for during hooks)
|
||||
mcp_auth_header: MCP auth header (deprecated)
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers
|
||||
oauth2_headers: Optional OAuth2 headers
|
||||
raw_headers: Optional raw headers from the request
|
||||
proxy_logging_obj: Optional ProxyLogging object for hook integration
|
||||
|
||||
Returns:
|
||||
CallToolResult from the MCP server
|
||||
|
||||
Raises:
|
||||
BlockedPiiEntityError: If PII is blocked by guardrails
|
||||
GuardrailRaisedException: If guardrails block the call
|
||||
HTTPException: If an HTTP error occurs
|
||||
"""
|
||||
# Get server-specific auth header if available
|
||||
server_auth_header: Optional[Union[Dict[str, str], str]] = None
|
||||
if mcp_server_auth_headers and mcp_server.alias:
|
||||
server_auth_header = mcp_server_auth_headers.get(mcp_server.alias)
|
||||
elif mcp_server_auth_headers and mcp_server.server_name:
|
||||
server_auth_header = mcp_server_auth_headers.get(mcp_server.server_name)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
|
||||
# oauth2 headers
|
||||
extra_headers: Optional[Dict[str, str]] = None
|
||||
if mcp_server.auth_type == MCPAuth.oauth2:
|
||||
extra_headers = oauth2_headers
|
||||
|
||||
if mcp_server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
extra_headers = {}
|
||||
for header in mcp_server.extra_headers:
|
||||
if header in raw_headers:
|
||||
extra_headers[header] = raw_headers[header]
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=mcp_server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
|
||||
call_tool_params = MCPCallToolRequestParams(
|
||||
name=original_tool_name,
|
||||
arguments=arguments,
|
||||
)
|
||||
|
||||
async def _call_tool_via_client(client, params):
|
||||
async with client:
|
||||
return await client.call_tool(params)
|
||||
|
||||
tasks.append(
|
||||
asyncio.create_task(_call_tool_via_client(client, call_tool_params))
|
||||
)
|
||||
|
||||
# IMPORTANT: Must await tasks INSIDE the context manager to keep connection alive
|
||||
try:
|
||||
mcp_responses = await asyncio.gather(*tasks)
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
# Re-raise guardrail exceptions to properly fail the MCP call
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
|
||||
# If proxy_logging_obj is None, the tool call result is at index 0
|
||||
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
|
||||
result_index = 1 if proxy_logging_obj else 0
|
||||
result = mcp_responses[result_index]
|
||||
|
||||
return cast(CallToolResult, result)
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
name: str,
|
||||
|
|
@ -793,95 +1218,63 @@ class MCPServerManager:
|
|||
server=mcp_server,
|
||||
)
|
||||
|
||||
# Get server-specific auth header if available
|
||||
server_auth_header: Optional[Union[Dict[str, str], str]] = None
|
||||
if mcp_server_auth_headers and mcp_server.alias:
|
||||
server_auth_header = mcp_server_auth_headers.get(mcp_server.alias)
|
||||
elif mcp_server_auth_headers and mcp_server.server_name:
|
||||
server_auth_header = mcp_server_auth_headers.get(mcp_server.server_name)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
|
||||
# oauth2 headers
|
||||
extra_headers: Optional[Dict[str, str]] = None
|
||||
if mcp_server.auth_type == MCPAuth.oauth2:
|
||||
extra_headers = oauth2_headers
|
||||
|
||||
if mcp_server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
extra_headers = {}
|
||||
for header in mcp_server.extra_headers:
|
||||
if header in raw_headers:
|
||||
extra_headers[header] = raw_headers[header]
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=mcp_server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
|
||||
async with client:
|
||||
# Use the original tool name (without prefix) for the actual call
|
||||
call_tool_params = MCPCallToolRequestParams(
|
||||
name=original_tool_name,
|
||||
# Prepare tasks for during hooks
|
||||
tasks = []
|
||||
if proxy_logging_obj:
|
||||
during_hook_task = self._create_during_hook_task(
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
server_name_from_prefix=server_name_from_prefix,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
start_time=start_time,
|
||||
)
|
||||
tasks = []
|
||||
if proxy_logging_obj:
|
||||
# Create synthetic LLM data for during hook processing
|
||||
from litellm.types.llms.base import HiddenParams
|
||||
from litellm.types.mcp import MCPDuringCallRequestObject
|
||||
tasks.append(during_hook_task)
|
||||
|
||||
request_obj = MCPDuringCallRequestObject(
|
||||
tool_name=name,
|
||||
arguments=arguments,
|
||||
server_name=server_name_from_prefix,
|
||||
start_time=start_time.timestamp() if start_time else None,
|
||||
hidden_params=HiddenParams(),
|
||||
# For OpenAPI servers, call the tool handler directly instead of via MCP client
|
||||
if mcp_server.spec_path:
|
||||
verbose_logger.debug(
|
||||
f"Calling OpenAPI tool {name} directly via HTTP handler"
|
||||
)
|
||||
tasks.append(
|
||||
asyncio.create_task(
|
||||
self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
)
|
||||
)
|
||||
else:
|
||||
# For regular MCP servers, use the MCP client
|
||||
return await self._call_regular_mcp_tool(
|
||||
mcp_server=mcp_server,
|
||||
original_tool_name=original_tool_name,
|
||||
arguments=arguments,
|
||||
tasks=tasks,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
during_hook_kwargs = {
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
"server_name": server_name_from_prefix,
|
||||
"user_api_key_auth": user_api_key_auth,
|
||||
}
|
||||
# For OpenAPI tools, await outside the client context
|
||||
try:
|
||||
mcp_responses = await asyncio.gather(*tasks)
|
||||
|
||||
synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format(
|
||||
request_obj, during_hook_kwargs
|
||||
)
|
||||
# If proxy_logging_obj is None, the tool call result is at index 0
|
||||
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
|
||||
result_index = 1 if proxy_logging_obj else 0
|
||||
result = mcp_responses[result_index]
|
||||
|
||||
during_hook_task = asyncio.create_task(
|
||||
proxy_logging_obj.during_call_hook(
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
data=synthetic_llm_data,
|
||||
call_type="mcp_call", # type: ignore
|
||||
)
|
||||
)
|
||||
tasks.append(during_hook_task)
|
||||
|
||||
tasks.append(asyncio.create_task(client.call_tool(call_tool_params)))
|
||||
try:
|
||||
mcp_responses = await asyncio.gather(*tasks)
|
||||
|
||||
# If proxy_logging_obj is None, the tool call result is at index 0
|
||||
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
|
||||
result_index = 1 if proxy_logging_obj else 0
|
||||
result = mcp_responses[result_index]
|
||||
|
||||
return cast(CallToolResult, result)
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
# Re-raise guardrail exceptions to properly fail the MCP call
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
return cast(CallToolResult, result)
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
# Re-raise guardrail exceptions to properly fail the MCP call
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
|
||||
#########################################################
|
||||
# End of Methods that call the upstream MCP servers
|
||||
|
|
@ -951,7 +1344,7 @@ class MCPServerManager:
|
|||
get_prisma_client_or_throw,
|
||||
)
|
||||
|
||||
verbose_logger.info("Loading MCP servers from database into registry...")
|
||||
verbose_logger.debug("Loading MCP servers from database into registry...")
|
||||
|
||||
# perform authz check to filter the mcp servers user has access to
|
||||
prisma_client = get_prisma_client_or_throw(
|
||||
|
|
@ -967,7 +1360,9 @@ class MCPServerManager:
|
|||
)
|
||||
self.add_update_server(server)
|
||||
|
||||
verbose_logger.info(f"Registry now contains {len(self.get_registry())} servers")
|
||||
verbose_logger.debug(
|
||||
f"Registry now contains {len(self.get_registry())} servers"
|
||||
)
|
||||
|
||||
def get_mcp_server_by_id(self, server_id: str) -> Optional[MCPServer]:
|
||||
"""
|
||||
|
|
@ -1055,6 +1450,7 @@ class MCPServerManager:
|
|||
if not server:
|
||||
return {
|
||||
"server_id": server_id,
|
||||
"server_name": None,
|
||||
"status": "unknown",
|
||||
"error": "Server not found",
|
||||
"last_health_check": datetime.now().isoformat(),
|
||||
|
|
@ -1069,6 +1465,7 @@ class MCPServerManager:
|
|||
|
||||
return {
|
||||
"server_id": server_id,
|
||||
"server_name": server.name,
|
||||
"status": "healthy",
|
||||
"tools_count": len(tools),
|
||||
"last_health_check": datetime.now().isoformat(),
|
||||
|
|
@ -1081,6 +1478,7 @@ class MCPServerManager:
|
|||
|
||||
return {
|
||||
"server_id": server_id,
|
||||
"server_name": server.name,
|
||||
"status": "unhealthy",
|
||||
"last_health_check": datetime.now().isoformat(),
|
||||
"response_time_ms": round(response_time, 2),
|
||||
|
|
|
|||
|
|
@ -0,0 +1,236 @@
|
|||
"""
|
||||
This module is used to generate MCP tools from OpenAPI specs.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
# Store the base URL and headers globally
|
||||
BASE_URL = ""
|
||||
HEADERS: Dict[str, str] = {}
|
||||
|
||||
|
||||
def load_openapi_spec(filepath: str) -> Dict[str, Any]:
|
||||
"""Load OpenAPI specification from JSON file."""
|
||||
with open(filepath, "r") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def get_base_url(spec: Dict[str, Any]) -> str:
|
||||
"""Extract base URL from OpenAPI spec."""
|
||||
# OpenAPI 3.x
|
||||
if "servers" in spec and spec["servers"]:
|
||||
return spec["servers"][0]["url"]
|
||||
# OpenAPI 2.x (Swagger)
|
||||
elif "host" in spec:
|
||||
scheme = spec.get("schemes", ["https"])[0]
|
||||
base_path = spec.get("basePath", "")
|
||||
return f"{scheme}://{spec['host']}{base_path}"
|
||||
return ""
|
||||
|
||||
|
||||
def extract_parameters(operation: Dict[str, Any]) -> tuple:
|
||||
"""Extract parameter names from OpenAPI operation."""
|
||||
path_params = []
|
||||
query_params = []
|
||||
body_params = []
|
||||
|
||||
# OpenAPI 3.x and 2.x parameters
|
||||
if "parameters" in operation:
|
||||
for param in operation["parameters"]:
|
||||
param_name = param["name"]
|
||||
if param.get("in") == "path":
|
||||
path_params.append(param_name)
|
||||
elif param.get("in") == "query":
|
||||
query_params.append(param_name)
|
||||
elif param.get("in") == "body":
|
||||
body_params.append(param_name)
|
||||
|
||||
# OpenAPI 3.x requestBody
|
||||
if "requestBody" in operation:
|
||||
body_params.append("body")
|
||||
|
||||
return path_params, query_params, body_params
|
||||
|
||||
|
||||
def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Build MCP input schema from OpenAPI operation."""
|
||||
properties = {}
|
||||
required = []
|
||||
|
||||
# Process parameters
|
||||
if "parameters" in operation:
|
||||
for param in operation["parameters"]:
|
||||
param_name = param["name"]
|
||||
param_schema = param.get("schema", {})
|
||||
param_type = param_schema.get("type", "string")
|
||||
|
||||
properties[param_name] = {
|
||||
"type": param_type,
|
||||
"description": param.get("description", ""),
|
||||
}
|
||||
|
||||
if param.get("required", False):
|
||||
required.append(param_name)
|
||||
|
||||
# Process requestBody (OpenAPI 3.x)
|
||||
if "requestBody" in operation:
|
||||
request_body = operation["requestBody"]
|
||||
content = request_body.get("content", {})
|
||||
|
||||
# Try to get JSON schema
|
||||
if "application/json" in content:
|
||||
schema = content["application/json"].get("schema", {})
|
||||
properties["body"] = {
|
||||
"type": "object",
|
||||
"description": request_body.get("description", "Request body"),
|
||||
"properties": schema.get("properties", {}),
|
||||
}
|
||||
if request_body.get("required", False):
|
||||
required.append("body")
|
||||
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": properties,
|
||||
"required": required if required else [],
|
||||
}
|
||||
|
||||
|
||||
def create_tool_function(
|
||||
path: str,
|
||||
method: str,
|
||||
operation: Dict[str, Any],
|
||||
base_url: str,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
"""Create a tool function for an OpenAPI operation.
|
||||
|
||||
Args:
|
||||
path: API endpoint path
|
||||
method: HTTP method (get, post, put, delete, patch)
|
||||
operation: OpenAPI operation object
|
||||
base_url: Base URL for the API
|
||||
headers: Optional headers to include in requests (e.g., authentication)
|
||||
"""
|
||||
if headers is None:
|
||||
headers = {}
|
||||
|
||||
path_params, query_params, body_params = extract_parameters(operation)
|
||||
all_params = path_params + query_params + body_params
|
||||
|
||||
# Build function signature dynamically
|
||||
if all_params:
|
||||
params_str = ", ".join(f"{p}: str = ''" for p in all_params)
|
||||
else:
|
||||
params_str = ""
|
||||
|
||||
# Create the function code as a string
|
||||
func_code = f'''
|
||||
async def tool_function({params_str}) -> str:
|
||||
"""Dynamically generated tool function."""
|
||||
url = base_url + path
|
||||
|
||||
# Replace path parameters
|
||||
path_param_names = {path_params}
|
||||
for param_name in path_param_names:
|
||||
param_value = locals().get(param_name, "")
|
||||
if param_value:
|
||||
url = url.replace("{{" + param_name + "}}", str(param_value))
|
||||
|
||||
# Build query params
|
||||
query_param_names = {query_params}
|
||||
params = {{}}
|
||||
for param_name in query_param_names:
|
||||
param_value = locals().get(param_name, "")
|
||||
if param_value:
|
||||
params[param_name] = param_value
|
||||
|
||||
# Build request body
|
||||
body_param_names = {body_params}
|
||||
json_body = None
|
||||
if body_param_names:
|
||||
body_value = locals().get("body", {{}})
|
||||
if isinstance(body_value, dict):
|
||||
json_body = body_value
|
||||
elif body_value:
|
||||
# If it's a string, try to parse as JSON
|
||||
import json as json_module
|
||||
try:
|
||||
json_body = json_module.loads(body_value) if isinstance(body_value, str) else {{"data": body_value}}
|
||||
except:
|
||||
json_body = {{"data": body_value}}
|
||||
|
||||
# Make HTTP request
|
||||
async with httpx.AsyncClient() as client:
|
||||
if "{method.lower()}" == "get":
|
||||
response = await client.get(url, params=params, headers=headers)
|
||||
elif "{method.lower()}" == "post":
|
||||
response = await client.post(url, params=params, json=json_body, headers=headers)
|
||||
elif "{method.lower()}" == "put":
|
||||
response = await client.put(url, params=params, json=json_body, headers=headers)
|
||||
elif "{method.lower()}" == "delete":
|
||||
response = await client.delete(url, params=params, headers=headers)
|
||||
elif "{method.lower()}" == "patch":
|
||||
response = await client.patch(url, params=params, json=json_body, headers=headers)
|
||||
else:
|
||||
return "Unsupported HTTP method: {method}"
|
||||
|
||||
return response.text
|
||||
'''
|
||||
|
||||
# Execute the function code to create the actual function
|
||||
local_vars = {
|
||||
"httpx": httpx,
|
||||
"headers": headers,
|
||||
"base_url": base_url,
|
||||
"path": path,
|
||||
"method": method,
|
||||
}
|
||||
exec(func_code, local_vars)
|
||||
|
||||
return local_vars["tool_function"]
|
||||
|
||||
|
||||
def register_tools_from_openapi(spec: Dict[str, Any], base_url: str):
|
||||
"""Register MCP tools from OpenAPI specification."""
|
||||
paths = spec.get("paths", {})
|
||||
|
||||
for path, path_item in paths.items():
|
||||
for method in ["get", "post", "put", "delete", "patch"]:
|
||||
if method in path_item:
|
||||
operation = path_item[method]
|
||||
|
||||
# Generate tool name
|
||||
operation_id = operation.get(
|
||||
"operationId", f"{method}_{path.replace('/', '_')}"
|
||||
)
|
||||
tool_name = operation_id.replace(" ", "_").lower()
|
||||
|
||||
# Get description
|
||||
description = operation.get(
|
||||
"summary", operation.get("description", f"{method.upper()} {path}")
|
||||
)
|
||||
|
||||
# Build input schema
|
||||
input_schema = build_input_schema(operation)
|
||||
|
||||
# Create tool function
|
||||
tool_func = create_tool_function(path, method, operation, base_url)
|
||||
tool_func.__name__ = tool_name
|
||||
tool_func.__doc__ = description
|
||||
|
||||
# Register tool with local registry
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=tool_name,
|
||||
description=description,
|
||||
input_schema=input_schema,
|
||||
handler=tool_func,
|
||||
)
|
||||
verbose_logger.debug(f"Registered tool: {tool_name}")
|
||||
|
|
@ -364,25 +364,25 @@ if MCP_AVAILABLE:
|
|||
def _tool_name_matches(tool_name: str, filter_list: List[str]) -> bool:
|
||||
"""
|
||||
Check if a tool name matches any name in the filter list.
|
||||
|
||||
|
||||
Checks both the full tool name and unprefixed version (without server prefix).
|
||||
This allows users to configure simple tool names regardless of prefixing.
|
||||
|
||||
|
||||
Args:
|
||||
tool_name: The tool name to check (may be prefixed like "server-tool_name")
|
||||
filter_list: List of tool names to match against
|
||||
|
||||
|
||||
Returns:
|
||||
True if the tool name (prefixed or unprefixed) is in the filter list
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
get_server_name_prefix_tool_mcp,
|
||||
)
|
||||
|
||||
|
||||
# Check if the full name is in the list
|
||||
if tool_name in filter_list:
|
||||
return True
|
||||
|
||||
|
||||
# Check if the unprefixed name is in the list
|
||||
unprefixed_name, _ = get_server_name_prefix_tool_mcp(tool_name)
|
||||
return unprefixed_name in filter_list
|
||||
|
|
@ -393,34 +393,36 @@ if MCP_AVAILABLE:
|
|||
) -> List[MCPTool]:
|
||||
"""
|
||||
Filter tools by allowed/disallowed tools configuration.
|
||||
|
||||
|
||||
If allowed_tools is set, only tools in that list are returned.
|
||||
If disallowed_tools is set, tools in that list are excluded.
|
||||
Tool names are matched with and without server prefixes for flexibility.
|
||||
|
||||
|
||||
Args:
|
||||
tools: List of tools to filter
|
||||
mcp_server: Server configuration with allowed_tools/disallowed_tools
|
||||
|
||||
|
||||
Returns:
|
||||
Filtered list of tools
|
||||
"""
|
||||
tools_to_return = tools
|
||||
|
||||
|
||||
# Filter by allowed_tools (whitelist)
|
||||
if mcp_server.allowed_tools:
|
||||
tools_to_return = [
|
||||
tool for tool in tools
|
||||
tool
|
||||
for tool in tools
|
||||
if _tool_name_matches(tool.name, mcp_server.allowed_tools)
|
||||
]
|
||||
|
||||
|
||||
# Filter by disallowed_tools (blacklist)
|
||||
if mcp_server.disallowed_tools:
|
||||
tools_to_return = [
|
||||
tool for tool in tools_to_return
|
||||
tool
|
||||
for tool in tools_to_return
|
||||
if not _tool_name_matches(tool.name, mcp_server.disallowed_tools)
|
||||
]
|
||||
|
||||
|
||||
return tools_to_return
|
||||
|
||||
async def _get_tools_from_mcp_servers(
|
||||
|
|
@ -497,17 +499,17 @@ if MCP_AVAILABLE:
|
|||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
)
|
||||
|
||||
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
||||
filtered_tools = await filter_tools_by_key_team_permissions(
|
||||
tools=filtered_tools,
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
|
||||
all_tools.extend(filtered_tools)
|
||||
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering"
|
||||
)
|
||||
|
|
@ -520,6 +522,7 @@ if MCP_AVAILABLE:
|
|||
verbose_logger.info(
|
||||
f"Successfully fetched {len(all_tools)} tools total from all MCP servers"
|
||||
)
|
||||
|
||||
return all_tools
|
||||
|
||||
async def filter_tools_by_key_team_permissions(
|
||||
|
|
@ -529,7 +532,7 @@ if MCP_AVAILABLE:
|
|||
) -> List[MCPTool]:
|
||||
"""
|
||||
Filter tools based on key/team mcp_tool_permissions.
|
||||
|
||||
|
||||
Note: Tool names in the DB are stored without server prefixes,
|
||||
but tool names from MCP servers are prefixed. We need to strip
|
||||
the prefix before comparing.
|
||||
|
|
@ -551,7 +554,7 @@ if MCP_AVAILABLE:
|
|||
else:
|
||||
# No restrictions, return all tools
|
||||
filtered_tools = tools
|
||||
|
||||
|
||||
return filtered_tools
|
||||
|
||||
async def _list_mcp_tools(
|
||||
|
|
@ -596,30 +599,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
# Continue with empty managed tools list instead of failing completely
|
||||
|
||||
# Get tools from local registry
|
||||
local_tools = []
|
||||
try:
|
||||
local_tools_raw = global_mcp_tool_registry.list_tools()
|
||||
|
||||
# Convert local tools to MCPTool format
|
||||
for tool in local_tools_raw:
|
||||
# Convert from litellm.types.mcp_server.tool_registry.MCPTool to mcp.types.Tool
|
||||
mcp_tool = MCPTool(
|
||||
name=tool.name,
|
||||
description=tool.description,
|
||||
inputSchema=tool.input_schema,
|
||||
)
|
||||
local_tools.append(mcp_tool)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
f"Error getting tools from local registry: {str(e)}"
|
||||
)
|
||||
# Continue with empty local tools list instead of failing completely
|
||||
|
||||
# Combine all tools
|
||||
all_tools = managed_tools + local_tools
|
||||
|
||||
return all_tools
|
||||
return managed_tools
|
||||
|
||||
@client
|
||||
async def call_mcp_tool(
|
||||
|
|
@ -680,33 +660,42 @@ if MCP_AVAILABLE:
|
|||
standard_logging_mcp_tool_call
|
||||
)
|
||||
litellm_logging_obj.model = f"MCP: {name}"
|
||||
# Try managed server tool first (pass the full prefixed name)
|
||||
# Primary and recommended way to use MCP servers
|
||||
# Check if tool exists in local registry first (for OpenAPI-based tools)
|
||||
# These tools are registered with their prefixed names
|
||||
#########################################################
|
||||
mcp_server: Optional[MCPServer] = (
|
||||
global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
)
|
||||
if mcp_server:
|
||||
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (
|
||||
mcp_server.mcp_info or {}
|
||||
).get("mcp_server_cost_info")
|
||||
response = await _handle_managed_mcp_tool(
|
||||
name=name, # Pass the full name (potentially prefixed)
|
||||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
local_tool = global_mcp_tool_registry.get_tool(name)
|
||||
if local_tool:
|
||||
verbose_logger.debug(f"Executing local registry tool: {name}")
|
||||
response = await _handle_local_mcp_tool(name, arguments)
|
||||
|
||||
# Fall back to local tool registry (use original name)
|
||||
#########################################################
|
||||
# Deprecated: Local MCP Server Tool
|
||||
# Try managed MCP server tool (pass the full prefixed name)
|
||||
# Primary and recommended way to use external MCP servers
|
||||
#########################################################
|
||||
else:
|
||||
response = await _handle_local_mcp_tool(original_tool_name, arguments)
|
||||
mcp_server: Optional[MCPServer] = (
|
||||
global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
)
|
||||
if mcp_server:
|
||||
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (
|
||||
mcp_server.mcp_info or {}
|
||||
).get("mcp_server_cost_info")
|
||||
response = await _handle_managed_mcp_tool(
|
||||
name=name, # Pass the full name (potentially prefixed)
|
||||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
# Fall back to local tool registry with original name (legacy support)
|
||||
#########################################################
|
||||
# Deprecated: Local MCP Server Tool
|
||||
#########################################################
|
||||
else:
|
||||
response = await _handle_local_mcp_tool(original_tool_name, arguments)
|
||||
|
||||
#########################################################
|
||||
# Post MCP Tool Call Hook
|
||||
|
|
@ -778,14 +767,21 @@ if MCP_AVAILABLE:
|
|||
Handle tool execution for local registry tools
|
||||
Note: Local tools don't use prefixes, so we use the original name
|
||||
"""
|
||||
import inspect
|
||||
|
||||
tool = global_mcp_tool_registry.get_tool(name)
|
||||
if not tool:
|
||||
raise HTTPException(status_code=404, detail=f"Tool '{name}' not found")
|
||||
|
||||
try:
|
||||
result = tool.handler(**arguments)
|
||||
# Check if handler is async or sync
|
||||
if inspect.iscoroutinefunction(tool.handler):
|
||||
result = await tool.handler(**arguments)
|
||||
else:
|
||||
result = tool.handler(**arguments)
|
||||
return [TextContent(text=str(result), type="text")]
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error executing local tool {name}: {str(e)}")
|
||||
return [TextContent(text=f"Error: {str(e)}", type="text")]
|
||||
|
||||
def _get_mcp_servers_in_path(path: str) -> Optional[List[str]]:
|
||||
|
|
@ -906,6 +902,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
await session_manager.handle_request(scope, receive, send)
|
||||
except Exception as e:
|
||||
raise e
|
||||
verbose_logger.exception(f"Error handling MCP request: {e}")
|
||||
# Instead of re-raising, try to send a graceful error response
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1,10 +1,18 @@
|
|||
import json
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
from litellm.types.mcp_server.tool_registry import MCPTool
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.types import Tool as MCPToolSDKTool
|
||||
else:
|
||||
try:
|
||||
from mcp.types import Tool as MCPToolSDKTool
|
||||
except ImportError:
|
||||
MCPToolSDKTool = None # type: ignore
|
||||
|
||||
|
||||
class MCPToolRegistry:
|
||||
"""
|
||||
|
|
@ -39,12 +47,34 @@ class MCPToolRegistry:
|
|||
"""
|
||||
return self.tools.get(name)
|
||||
|
||||
def list_tools(self) -> List[MCPTool]:
|
||||
def list_tools(self, tool_prefix: Optional[str] = None) -> List[MCPTool]:
|
||||
"""
|
||||
List all registered tools
|
||||
"""
|
||||
if tool_prefix:
|
||||
return [
|
||||
tool
|
||||
for tool in self.tools.values()
|
||||
if tool.name.startswith(tool_prefix)
|
||||
]
|
||||
return list(self.tools.values())
|
||||
|
||||
def convert_tools_to_mcp_sdk_tool_type(
|
||||
self, tools: List[MCPTool]
|
||||
) -> List["MCPToolSDKTool"]:
|
||||
if MCPToolSDKTool is None:
|
||||
raise ImportError(
|
||||
"MCP SDK is not installed. Please install it with: pip install 'litellm[proxy]'"
|
||||
)
|
||||
return [
|
||||
MCPToolSDKTool(
|
||||
name=tool.name,
|
||||
description=tool.description,
|
||||
inputSchema=tool.input_schema,
|
||||
)
|
||||
for tool in tools
|
||||
]
|
||||
|
||||
def load_tools_from_config(
|
||||
self, mcp_tools_config: Optional[Dict[str, Any]] = None
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
self.__BUILD_MANIFEST={__rewrites:{afterFiles:[],beforeFiles:[],fallback:[]},"/_error":["static/chunks/pages/_error-28b803cb2479b966.js"],sortedPages:["/_app","/_error"]},self.__BUILD_MANIFEST_CB&&self.__BUILD_MANIFEST_CB();
|
||||
|
|
@ -0,0 +1 @@
|
|||
self.__BUILD_MANIFEST={__rewrites:{afterFiles:[],beforeFiles:[],fallback:[]},"/_error":["static/chunks/pages/_error-cf5ca766ac8f493f.js"],sortedPages:["/_app","/_error"]},self.__BUILD_MANIFEST_CB&&self.__BUILD_MANIFEST_CB();
|
||||