Merge branch 'main' into litellm_dev_02_11_2026_p1

This commit is contained in:
Krish Dholakia 2026-02-13 09:02:24 -08:00 • committed by GitHub
commit c55469241c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
196 changed files with 17050 additions and 3456 deletions

View file

@ -40,38 +40,33 @@ outputs:
runs:
using: composite
steps:
- name: Helm | Setup
uses: azure/setup-helm@v4
with:
version: v3.20.0
- name: Helm | Login
shell: bash
run: echo ${{ inputs.registry_password }} | helm registry login -u ${{ inputs.registry_username }} --password-stdin ${{ inputs.registry }}
env:
HELM_EXPERIMENTAL_OCI: '1'
- name: Helm | Dependency
if: inputs.update_dependencies == 'true'
shell: bash
run: helm dependency update ${{ inputs.path == null && format('{0}/{1}', 'charts', inputs.name) || inputs.path }}
env:
HELM_EXPERIMENTAL_OCI: '1'
- name: Helm | Package
shell: bash
run: helm package ${{ inputs.path == null && format('{0}/{1}', 'charts', inputs.name) || inputs.path }} --version ${{ inputs.tag }} --app-version ${{ inputs.app_version }}
env:
HELM_EXPERIMENTAL_OCI: '1'
- name: Helm | Push
shell: bash
run: helm push ${{ inputs.name }}-${{ inputs.tag }}.tgz oci://${{ inputs.registry }}/${{ inputs.repository }}
env:
HELM_EXPERIMENTAL_OCI: '1'
- name: Helm | Logout
shell: bash
run: helm registry logout ${{ inputs.registry }}
env:
HELM_EXPERIMENTAL_OCI: '1'
- name: Helm | Output
id: output
shell: bash
run: echo "image=${{ inputs.registry }}/${{ inputs.repository }}/${{ inputs.name }}:${{ inputs.tag }}" >> $GITHUB_OUTPUT
run: echo "image=${{ inputs.registry }}/${{ inputs.repository }}/${{ inputs.name }}:${{ inputs.tag }}" >> $GITHUB_OUTPUT

View file

@ -0,0 +1,95 @@
name: LiteLLM Unit Tests (Matrix)
on:
pull_request:
branches: [main]
# Cancel in-progress runs for the same PR
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
test:
runs-on: ubuntu-latest
timeout-minutes: 15
strategy:
fail-fast: false
matrix:
test-group:
# tests/test_litellm split by subdirectory (~560 files total)
- name: "llms"
path: "tests/test_litellm/llms"
workers: 4
# tests/test_litellm/proxy split by subdirectory (~180 files total)
- name: "proxy-guardrails"
path: "tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers"
workers: 4
- name: "proxy-core"
path: "tests/test_litellm/proxy/auth tests/test_litellm/proxy/client tests/test_litellm/proxy/db tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine"
workers: 4
- name: "proxy-misc"
path: "tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py"
workers: 4
- name: "integrations"
path: "tests/test_litellm/integrations"
workers: 4
- name: "core-utils"
path: "tests/test_litellm/litellm_core_utils"
workers: 2
- name: "other"
path: "tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/test_litellm/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types"
workers: 4
- name: "root"
path: "tests/test_litellm/test_*.py"
workers: 4
# tests/proxy_unit_tests split alphabetically (~48 files total)
- name: "proxy-unit-a"
path: "tests/proxy_unit_tests/test_[a-o]*.py"
workers: 2
- name: "proxy-unit-b"
path: "tests/proxy_unit_tests/test_[p-z]*.py"
workers: 2
name: test (${{ matrix.test-group.name }})
steps:
- uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Install Poetry
uses: snok/install-poetry@v1
- name: Cache Poetry dependencies
uses: actions/cache@v4
with:
path: |
~/.cache/pypoetry
~/.cache/pip
.venv
key: ${{ runner.os }}-poetry-${{ hashFiles('poetry.lock') }}
restore-keys: |
${{ runner.os }}-poetry-
- name: Install dependencies
run: |
poetry config virtualenvs.in-project true
poetry install --with dev,proxy-dev --extras "proxy semantic-router"
poetry run pip install pytest-retry==1.6.3 pytest-xdist google-genai==1.22.0 \
google-cloud-aiplatform>=1.38 fastapi-offline==1.7.3 python-multipart==0.0.22 openapi-core
- name: Setup litellm-enterprise
run: |
cd enterprise && poetry run pip install -e . && cd ..
- name: Run tests - ${{ matrix.test-group.name }}
run: |
poetry run pytest ${{ matrix.test-group.path }} \
--tb=short -vv \
--maxfail=10 \
-n ${{ matrix.test-group.workers }} \
--durations=20

View file

@ -1,8 +1,12 @@
name: LiteLLM Mock Tests (folder - tests/test_litellm)
# DEPRECATED: This workflow is replaced by test-litellm-matrix.yml which runs
# the same tests in parallel across 10 jobs for faster CI times.
# Kept for manual debugging only.
on:
pull_request:
branches: [ main ]
workflow_dispatch: # Manual trigger only
# pull_request:
# branches: [ main ]
jobs:
test:

View file

@ -1,52 +1,22 @@
# Custom Semgrep Rules
# Custom Semgrep rules for LiteLLM
All `.yml` files under `.semgrep/rules/` run in CI (CircleCI `semgrep` job).
Add custom rule YAML files here. Semgrep loads all `.yml`/`.yaml` files under this directory.
## Add a Rule
* Add a `.yml` file under `.semgrep/rules/<language>/<domain>/`
[Rule syntax →](https://semgrep.dev/docs/writing-rules/rule-syntax/)
## Organizing Rules
### Structure: language → domain
```
.semgrep/rules/<language>/<domain>/<rule-name>.yml
```
Examples:
- `python/security/unsafe-yaml-load.yml`
- `python/reliability/missing-timeout-http.yml`
- `python/performance/blocking-io-in-async.yml`
### Rule metadata
Match tags to the folder for consistent filtering:
```yaml
metadata:
tags: [python, security]
```
### Severity expectations
All rules must fail CI on findings. No warn-only rules.
- Use `severity: ERROR` in rule metadata
- If a rule is noisy → refine until low false positives before adding
## Run Locally
**Run only custom rules (CI / fail on findings):**
```bash
semgrep scan --config .semgrep/rules . --error
```
With Semgrep registry:
**Run with registry + custom rules:**
```bash
semgrep scan --config auto --config .semgrep/rules .
```
**Layout:**
- `python/` – Python-specific rules (security, patterns)
- Add more subdirs as needed (e.g. `generic/` for language-agnostic rules)
See [Semgrep rule syntax](https://semgrep.dev/docs/writing-rules/rule-syntax/).

View file

@ -0,0 +1,14 @@
# Unbounded memory growth – data structures without a clear max limit
# Can lead to OOM under load.
rules:
- id: unbounded-asyncio-queue
message: asyncio.Queue() with no maxsize can grow unbounded. Use asyncio.Queue(maxsize=N) for integrations (e.g. log queues).
severity: ERROR
languages: [python]
pattern-either:
- pattern: asyncio.Queue()
- pattern: asyncio.Queue(maxsize=0)
metadata:
category: correctness
cwe: "CWE-400: Uncontrolled Resource Consumption"

View file

@ -1,7 +1,9 @@
# LiteLLM Makefile
# Simple Makefile for running tests and basic development tasks
.PHONY: help test test-unit test-integration test-unit-helm \
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
info lint lint-dev format \
install-dev install-proxy-dev install-test-deps \
install-helm-unittest check-circular-imports check-import-safety
@ -25,6 +27,16 @@ help:
@echo " make check-import-safety - Check import safety"
@echo " make test - Run all tests"
@echo " make test-unit - Run unit tests (tests/test_litellm)"
@echo " make test-unit-llms - Run LLM provider tests (~225 files)"
@echo " make test-unit-proxy-guardrails - Run proxy guardrails+mgmt tests (~51 files)"
@echo " make test-unit-proxy-core - Run proxy auth+client+db+hooks tests (~52 files)"
@echo " make test-unit-proxy-misc - Run proxy misc tests (~77 files)"
@echo " make test-unit-integrations - Run integration tests (~60 files)"
@echo " make test-unit-core-utils - Run core utils tests (~32 files)"
@echo " make test-unit-other - Run other tests (caching, responses, etc., ~69 files)"
@echo " make test-unit-root - Run root-level tests (~34 files)"
@echo " make test-proxy-unit-a - Run proxy_unit_tests (a-o, ~20 files)"
@echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)"
@echo " make test-integration - Run integration tests"
@echo " make test-unit-helm - Run helm unit tests"
@ -129,6 +141,38 @@ test:
test-unit: install-test-deps
poetry run pytest tests/test_litellm -x -vv -n 4
# Matrix test targets (matching CI workflow groups)
test-unit-llms: install-test-deps
poetry run pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20
test-unit-proxy-guardrails: install-test-deps
poetry run pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20
test-unit-proxy-core: install-test-deps
poetry run pytest tests/test_litellm/proxy/auth tests/test_litellm/proxy/client tests/test_litellm/proxy/db tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine --tb=short -vv -n 4 --durations=20
test-unit-proxy-misc: install-test-deps
poetry run pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
test-unit-integrations: install-test-deps
poetry run pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20
test-unit-core-utils: install-test-deps
poetry run pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
test-unit-other: install-test-deps
poetry run pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/test_litellm/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
test-unit-root: install-test-deps
poetry run pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
# Proxy unit tests (tests/proxy_unit_tests split alphabetically)
test-proxy-unit-a: install-test-deps
poetry run pytest tests/proxy_unit_tests/test_[a-o]*.py --tb=short -vv -n 2 --durations=20
test-proxy-unit-b: install-test-deps
poetry run pytest tests/proxy_unit_tests/test_[p-z]*.py --tb=short -vv -n 2 --durations=20
test-integration:
poetry run pytest tests/ -k "not test_litellm"

View file

@ -26,6 +26,10 @@ version: 1.1.0
# It is recommended to use it with quotes.
appVersion: v1.80.12
annotations:
org.opencontainers.image.source: "https://github.com/BerriAI/litellm"
org.opencontainers.image.url: "https://docs.litellm.ai/"
dependencies:
- name: "postgresql"
version: ">=13.3.0"

View file

@ -59,7 +59,8 @@ RUN mkdir -p /var/lib/litellm/ui && \
mkdir -p "$folder_name" && \
mv "$html_file" "$folder_name/index.html"; \
fi; \
done ) && \
done && \
touch .litellm_ui_ready ) && \
cd /app/ui/litellm-dashboard && rm -rf ./out
# Build litellm wheel and place it in wheels dir (replace any PyPI wheels)

View file

@ -70,9 +70,12 @@ docker compose -f docker-compose.yml -f docker-compose.hardened.yml up -d
This setup:
- Builds from `docker/Dockerfile.non_root` with Prisma engines and Node toolchain baked into the image.
- Runs the proxy as a non-root user with a read-only rootfs and only two writable tmpfs mounts:
- Runs the proxy as a non-root user with a read-only rootfs and only writable tmpfs mounts:
- `/app/cache` (Prisma/NPM cache; backing `PRISMA_BINARY_CACHE_DIR`, `NPM_CONFIG_CACHE`, `XDG_CACHE_HOME`)
- `/app/migrations` (Prisma migration workspace; backing `LITELLM_MIGRATION_DIR`)
- Pre-builds and serves the admin UI from read-only paths:
- `/var/lib/litellm/ui` (pre-restructured Next.js UI with `.litellm_ui_ready` marker)
- `/var/lib/litellm/assets` (UI logos and assets)
- Routes all outbound traffic through a local Squid proxy that denies egress, so Prisma migrations must use the cached CLI and engines.
You should also verify offline Prisma behaviour with:

View file

@ -389,6 +389,10 @@ Compaction blocks are also supported in streaming mode. You'll receive:
### Adaptive Thinking
:::note
When using `reasoning_effort` with Claude Opus 4.6, all values (`low`, `medium`, `high`) are mapped to `thinking: {type: "adaptive"}`. To use explicit thinking budgets with `type: "enabled"`, pass the native `thinking` parameter directly (see "Native thinking param" tab below).
:::
<Tabs>
<TabItem value="completions" label="/chat/completions">
@ -434,6 +438,21 @@ curl --location 'http://0.0.0.0:4000/v1/messages' \
}'
```
</TabItem>
<TabItem value="native" label="Native thinking param">
Use the `thinking` parameter directly for adaptive thinking via the SDK:
```python
import litellm
response = litellm.completion(
model="anthropic/claude-opus-4-6",
messages=[{"role": "user", "content": "Solve this complex problem: What is the optimal strategy for..."}],
thinking={"type": "adaptive"},
)
```
</TabItem>
</Tabs>

View file

@ -0,0 +1,394 @@
---
slug: minimax_m2_5
title: "Day 0 Support: MiniMax-M2.5"
date: 2026-02-12T10:00:00
authors:
- name: Sameer Kankute
title: SWE @ LiteLLM (LLM Translation)
url: https://www.linkedin.com/in/sameer-kankute/
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
- name: Krrish Dholakia
title: "CEO, LiteLLM"
url: https://www.linkedin.com/in/krish-d/
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
- name: Ishaan Jaff
title: "CTO, LiteLLM"
url: https://www.linkedin.com/in/reffajnaahsi/
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
description: "Day 0 support for MiniMax-M2.5 on LiteLLM"
tags: [minimax, M2.5, llm]
hide_table_of_contents: false
---
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
LiteLLM now supports MiniMax-M2.5 on Day 0. Use it across OpenAI-compatible and Anthropic-compatible APIs through the LiteLLM AI Gateway.
## Supported Models
LiteLLM supports the following MiniMax models:
| Model | Description | Input Cost | Output Cost | Context Window |
|-------|-------------|------------|-------------|----------------|
| **MiniMax-M2.5** | Advanced reasoning, Agentic capabilities | $0.3/M tokens | $1.2/M tokens | 1M tokens |
| **MiniMax-M2.5-lightning** | Faster and More Agile (~100 tps) | $0.3/M tokens | $2.4/M tokens | 1M tokens |
## Features Supported
- **Prompt Caching**: Reduce costs with cached prompts ($0.03/M tokens for cache read, $0.375/M tokens for cache write)
- **Function Calling**: Built-in tool calling support
- **Reasoning**: Advanced reasoning capabilities with thinking support
- **System Messages**: Full system message support
- **Cost Tracking**: Automatic cost calculation for all requests
## Docker Image
```bash
docker pull litellm/litellm:v1.81.3-stable
```
## Usage - OpenAI Compatible API (/v1/chat/completions)
<Tabs>
<TabItem value="proxy" label="LiteLLM Proxy">
**1. Setup config.yaml**
```yaml
model_list:
- model_name: minimax-m2-5
litellm_params:
model: minimax/MiniMax-M2.5
api_key: os.environ/MINIMAX_API_KEY
api_base: https://api.minimax.io/v1
```
**2. Start the proxy**
```bash
docker run -d \
-p 4000:4000 \
-e MINIMAX_API_KEY=$MINIMAX_API_KEY \
-v $(pwd)/config.yaml:/app/config.yaml \
ghcr.io/berriai/litellm:v1.81.3-stable \
--config /app/config.yaml
```
**3. Test it!**
```bash
curl --location 'http://0.0.0.0:4000/chat/completions' \
--header 'Content-Type: application/json' \
--header 'Authorization: Bearer $LITELLM_KEY' \
--data '{
"model": "minimax-m2-5",
"messages": [
{
"role": "user",
"content": "what llm are you"
}
]
}'
```
</TabItem>
</Tabs>
### With Reasoning Split
```bash
curl --location 'http://0.0.0.0:4000/chat/completions' \
--header 'Content-Type: application/json' \
--header 'Authorization: Bearer $LITELLM_KEY' \
--data '{
"model": "minimax-m2-5",
"messages": [
{
"role": "user",
"content": "Solve: 2+2=?"
}
],
"extra_body": {
"reasoning_split": true
}
}'
```
## Usage - Anthropic Compatible API (/v1/messages)
<Tabs>
<TabItem value="proxy" label="LiteLLM Proxy">
**1. Setup config.yaml**
```yaml
model_list:
- model_name: minimax-m2-5
litellm_params:
model: minimax/MiniMax-M2.5
api_key: os.environ/MINIMAX_API_KEY
api_base: https://api.minimax.io/anthropic/v1/messages
```
**2. Start the proxy**
```bash
docker run -d \
-p 4000:4000 \
-e MINIMAX_API_KEY=$MINIMAX_API_KEY \
-v $(pwd)/config.yaml:/app/config.yaml \
ghcr.io/berriai/litellm:v1.81.3-stable \
--config /app/config.yaml
```
**3. Test it!**
```bash
curl --location 'http://0.0.0.0:4000/v1/messages' \
--header 'Content-Type: application/json' \
--header 'Authorization: Bearer $LITELLM_KEY' \
--data '{
"model": "minimax-m2-5",
"max_tokens": 1000,
"messages": [
{
"role": "user",
"content": "what llm are you"
}
]
}'
```
</TabItem>
</Tabs>
### With Thinking
```bash
curl --location 'http://0.0.0.0:4000/v1/messages' \
--header 'Content-Type: application/json' \
--header 'Authorization: Bearer $LITELLM_KEY' \
--data '{
"model": "minimax-m2-5",
"max_tokens": 1000,
"thinking": {
"type": "enabled",
"budget_tokens": 1000
},
"messages": [
{
"role": "user",
"content": "Solve: 2+2=?"
}
]
}'
```
## Usage - LiteLLM SDK
### OpenAI-compatible API
```python
import litellm
response = litellm.completion(
model="minimax/MiniMax-M2.5",
messages=[
{"role": "user", "content": "Hello, how are you?"}
],
api_key="your-minimax-api-key",
api_base="https://api.minimax.io/v1"
)
print(response.choices[0].message.content)
```
### Anthropic-compatible API
```python
import litellm
response = litellm.anthropic.messages.acreate(
model="minimax/MiniMax-M2.5",
messages=[{"role": "user", "content": "Hello, how are you?"}],
api_key="your-minimax-api-key",
api_base="https://api.minimax.io/anthropic/v1/messages",
max_tokens=1000
)
print(response.choices[0].message.content)
```
### With Thinking
```python
response = litellm.anthropic.messages.acreate(
model="minimax/MiniMax-M2.5",
messages=[{"role": "user", "content": "Solve: 2+2=?"}],
thinking={"type": "enabled", "budget_tokens": 1000},
api_key="your-minimax-api-key"
)
# Access thinking content
for block in response.choices[0].message.content:
if hasattr(block, 'type') and block.type == 'thinking':
print(f"Thinking: {block.thinking}")
```
### With Reasoning Split (OpenAI API)
```python
response = litellm.completion(
model="minimax/MiniMax-M2.5",
messages=[
{"role": "user", "content": "Solve: 2+2=?"}
],
extra_body={"reasoning_split": True},
api_key="your-minimax-api-key",
api_base="https://api.minimax.io/v1"
)
# Access thinking and response
if hasattr(response.choices[0].message, 'reasoning_details'):
print(f"Thinking: {response.choices[0].message.reasoning_details}")
print(f"Response: {response.choices[0].message.content}")
```
## Cost Tracking
LiteLLM automatically tracks costs for MiniMax-M2.5 requests. The pricing is:
- **Input**: $0.3 per 1M tokens
- **Output**: $1.2 per 1M tokens
- **Cache Read**: $0.03 per 1M tokens
- **Cache Write**: $0.375 per 1M tokens
### Accessing Cost Information
```python
response = litellm.completion(
model="minimax/MiniMax-M2.5",
messages=[{"role": "user", "content": "Hello!"}],
api_key="your-minimax-api-key"
)
# Access cost information
print(f"Cost: ${response._hidden_params.get('response_cost', 0)}")
```
## Streaming Support
### OpenAI API
```python
response = litellm.completion(
model="minimax/MiniMax-M2.5",
messages=[{"role": "user", "content": "Tell me a story"}],
stream=True,
api_key="your-minimax-api-key",
api_base="https://api.minimax.io/v1"
)
for chunk in response:
if chunk.choices[0].delta.content:
print(chunk.choices[0].delta.content, end="")
```
### Streaming with Reasoning Split
```python
stream = litellm.completion(
model="minimax/MiniMax-M2.5",
messages=[
{"role": "user", "content": "Tell me a story"},
],
extra_body={"reasoning_split": True},
stream=True,
api_key="your-minimax-api-key",
api_base="https://api.minimax.io/v1"
)
reasoning_buffer = ""
text_buffer = ""
for chunk in stream:
if hasattr(chunk.choices[0].delta, "reasoning_details") and chunk.choices[0].delta.reasoning_details:
for detail in chunk.choices[0].delta.reasoning_details:
if "text" in detail:
reasoning_text = detail["text"]
new_reasoning = reasoning_text[len(reasoning_buffer):]
if new_reasoning:
print(new_reasoning, end="", flush=True)
reasoning_buffer = reasoning_text
if chunk.choices[0].delta.content:
content_text = chunk.choices[0].delta.content
new_text = content_text[len(text_buffer):] if text_buffer else content_text
if new_text:
print(new_text, end="", flush=True)
text_buffer = content_text
```
## Using with Native SDKs
### Anthropic SDK via LiteLLM Proxy
```python
import os
os.environ["ANTHROPIC_BASE_URL"] = "http://localhost:4000"
os.environ["ANTHROPIC_API_KEY"] = "sk-1234" # Your LiteLLM proxy key
import anthropic
client = anthropic.Anthropic()
message = client.messages.create(
model="minimax-m2-5",
max_tokens=1000,
system="You are a helpful assistant.",
messages=[
{
"role": "user",
"content": [
{
"type": "text",
"text": "Hi, how are you?"
}
]
}
]
)
for block in message.content:
if block.type == "thinking":
print(f"Thinking:\n{block.thinking}\n")
elif block.type == "text":
print(f"Text:\n{block.text}\n")
```
### OpenAI SDK via LiteLLM Proxy
```python
import os
os.environ["OPENAI_BASE_URL"] = "http://localhost:4000"
os.environ["OPENAI_API_KEY"] = "sk-1234" # Your LiteLLM proxy key
from openai import OpenAI
client = OpenAI()
response = client.chat.completions.create(
model="minimax-m2-5",
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hi, how are you?"},
],
extra_body={"reasoning_split": True},
)
# Access thinking and response
if hasattr(response.choices[0].message, 'reasoning_details'):
print(f"Thinking:\n{response.choices[0].message.reasoning_details[0]['text']}\n")
print(f"Text:\n{response.choices[0].message.content}\n")
```

View file

@ -93,6 +93,12 @@ Implement `POST /beta/litellm_basic_guardrail_api`
"user_api_key_end_user_id": "end user id associated with the litellm virtual key used",
"user_api_key_org_id": "org id associated with the litellm virtual key used"
},
"request_headers": { // optional: inbound request headers (allowlist). Allowed headers show their value; all others show "[present]" to indicate the header existed.
"User-Agent": "OpenAI/Python 2.17.0",
"Content-Type": "application/json",
"X-Request-Id": "[present]"
},
"litellm_version": "1.x.y", // optional: LiteLLM library version running this proxy
"input_type": "request", // "request" or "response"
"litellm_call_id": "unique_call_id", // the call id of the individual LLM call
"litellm_trace_id": "trace_id", // the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation

View file

@ -18,7 +18,7 @@ Each provider uses their own search backend:
| Provider | Search Engine | Notes |
|----------|---------------|-------|
| **OpenAI** (`gpt-4o-search-preview`, `gpt-4o-mini-search-preview`, `gpt-5-search-api`) | OpenAI's internal search | Real-time web data |
| **OpenAI** (`gpt-5-search-api`, `gpt-4o-search-preview`, `gpt-4o-mini-search-preview`) | OpenAI's internal search | Real-time web data |
| **xAI** (`grok-3`) | xAI's search + X/Twitter | Real-time social media data |
| **Google AI/Vertex** (`gemini-2.0-flash`) | **Google Search** | Uses actual Google search results |
| **Anthropic** (`claude-3-5-sonnet`) | Anthropic's web search | Real-time web data |
@ -45,6 +45,19 @@ Use `web_search_options` when you need to:
**Anthropic Web Search Models**: Claude models that support web search: `claude-3-5-sonnet-latest`, `claude-3-5-sonnet-20241022`, `claude-3-5-haiku-latest`, `claude-3-5-haiku-20241022`, `claude-3-7-sonnet-20250219`
:::
## OpenAI Web Search: Two Approaches
OpenAI offers two distinct ways to use web search depending on the endpoint and model:
| Approach | Endpoint | Models | How to enable |
|----------|----------|--------|---------------|
| **Search Models** | `/chat/completions` | `gpt-5-search-api`, `gpt-4o-search-preview`, `gpt-4o-mini-search-preview` | Pass `web_search_options` parameter |
| **Web Search Tool** | `/responses` | `gpt-5`, `gpt-4.1`, `gpt-4o`, and other regular models | Pass `web_search_preview` tool |
:::tip Search models search automatically
Search models like `gpt-5-search-api` **automatically search the web** even without the `web_search_options` parameter. Use `web_search_options` to set `search_context_size` (`"low"`, `"medium"`, `"high"`) or specify `user_location` for localized results.
:::
## `/chat/completions` (litellm.completion)
### Quick Start
@ -56,7 +69,7 @@ Use `web_search_options` when you need to:
from litellm import completion
response = completion(
model="openai/gpt-4o-search-preview",
model="openai/gpt-5-search-api",
messages=[
{
"role": "user",
@ -76,31 +89,36 @@ response = completion(
```yaml
model_list:
# OpenAI
# OpenAI search models
- model_name: gpt-5-search-api
litellm_params:
model: openai/gpt-5-search-api
api_key: os.environ/OPENAI_API_KEY
- model_name: gpt-4o-search-preview
litellm_params:
model: openai/gpt-4o-search-preview
api_key: os.environ/OPENAI_API_KEY
# xAI
- model_name: grok-3
litellm_params:
model: xai/grok-3
api_key: os.environ/XAI_API_KEY
# Anthropic
- model_name: claude-3-5-sonnet-latest
litellm_params:
model: anthropic/claude-3-5-sonnet-latest
api_key: os.environ/ANTHROPIC_API_KEY
# VertexAI
- model_name: gemini-2-flash
litellm_params:
model: gemini-2.0-flash
vertex_project: your-project-id
vertex_location: us-central1
# Google AI Studio
- model_name: gemini-2-flash-studio
litellm_params:
@ -108,13 +126,13 @@ model_list:
api_key: os.environ/GOOGLE_API_KEY
```
2. Start the proxy
2. Start the proxy
```bash
litellm --config /path/to/config.yaml
```
3. Test it!
3. Test it!
```python showLineNumbers
from openai import OpenAI
@ -126,13 +144,18 @@ client = OpenAI(
)
response = client.chat.completions.create(
model="grok-3", # or any other web search enabled model
model="gpt-5-search-api", # or any other web search enabled model
messages=[
{
"role": "user",
"content": "What was a positive news story from today?"
}
]
],
extra_body={
"web_search_options": {
"search_context_size": "medium"
}
}
)
```
</TabItem>
@ -149,7 +172,7 @@ from litellm import completion
# Customize search context size
response = completion(
model="openai/gpt-4o-search-preview",
model="openai/gpt-5-search-api",
messages=[
{
"role": "user",
@ -257,6 +280,12 @@ response = client.chat.completions.create(
## `/responses` (litellm.responses)
Use the `web_search_preview` tool with models like `gpt-5`, `gpt-4.1`, `gpt-4o`, etc.
:::info
Search-dedicated models like `gpt-5-search-api` and `gpt-4o-search-preview` do **not** support the `/responses` endpoint. Use them with `/chat/completions` + `web_search_options` instead (see above).
:::
### Quick Start
<Tabs>
@ -266,18 +295,14 @@ response = client.chat.completions.create(
from litellm import responses
response = responses(
model="openai/gpt-4o",
input=[
{
"role": "user",
"content": "What was a positive news story from today?"
}
],
model="openai/gpt-5",
input="What is the capital of France?",
tools=[{
"type": "web_search_preview" # enables web search with default medium context size
}]
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
@ -285,19 +310,24 @@ response = responses(
```yaml
model_list:
- model_name: gpt-4o
- model_name: gpt-5
litellm_params:
model: openai/gpt-4o
model: openai/gpt-5
api_key: os.environ/OPENAI_API_KEY
- model_name: gpt-4.1
litellm_params:
model: openai/gpt-4.1
api_key: os.environ/OPENAI_API_KEY
```
2. Start the proxy
2. Start the proxy
```bash
litellm --config /path/to/config.yaml
```
3. Test it!
3. Test it!
```python showLineNumbers
from openai import OpenAI
@ -309,11 +339,11 @@ client = OpenAI(
)
response = client.responses.create(
model="gpt-4o",
model="gpt-5",
tools=[{
"type": "web_search_preview"
}],
input="What was a positive news story from today?",
input="What is the capital of France?",
)
print(response.output_text)
@ -331,13 +361,8 @@ from litellm import responses
# Customize search context size
response = responses(
model="openai/gpt-4o",
input=[
{
"role": "user",
"content": "What was a positive news story from today?"
}
],
model="openai/gpt-5",
input="What is the capital of France?",
tools=[{
"type": "web_search_preview",
"search_context_size": "low" # Options: "low", "medium" (default), "high"
@ -358,12 +383,12 @@ client = OpenAI(
# Customize search context size
response = client.responses.create(
model="gpt-4o",
model="gpt-5",
tools=[{
"type": "web_search_preview",
"search_context_size": "low" # Options: "low", "medium" (default), "high"
}],
input="What was a positive news story from today?",
input="What is the capital of France?",
)
print(response.output_text)
@ -417,14 +442,14 @@ model_list:
web_search_options:
search_context_size: "high" # Options: "low", "medium", "high"
# Different context size for different models
- model_name: gpt-4o-search-preview
# OpenAI search model with custom context size
- model_name: gpt-5-search-api
litellm_params:
model: openai/gpt-4o-search-preview
model: openai/gpt-5-search-api
api_key: os.environ/OPENAI_API_KEY
web_search_options:
search_context_size: "low"
# Gemini with medium context (default)
- model_name: gemini-2-flash
litellm_params:
@ -449,6 +474,7 @@ Use `litellm.supports_web_search(model="model_name")` -> returns `True` if model
```python showLineNumbers
# Check OpenAI models
assert litellm.supports_web_search(model="openai/gpt-5-search-api") == True
assert litellm.supports_web_search(model="openai/gpt-4o-search-preview") == True
# Check xAI models
@ -472,13 +498,20 @@ assert litellm.supports_web_search(model="gemini/gemini-2.0-flash") == True
```yaml
model_list:
# OpenAI
- model_name: gpt-5-search-api
litellm_params:
model: openai/gpt-5-search-api
api_key: os.environ/OPENAI_API_KEY
model_info:
supports_web_search: True
- model_name: gpt-4o-search-preview
litellm_params:
model: openai/gpt-4o-search-preview
api_key: os.environ/OPENAI_API_KEY
model_info:
supports_web_search: True
# xAI
- model_name: grok-3
litellm_params:
@ -533,6 +566,12 @@ Expected Response
```json showLineNumbers
{
"data": [
{
"model_group": "gpt-5-search-api",
"providers": ["openai"],
"max_tokens": 128000,
"supports_web_search": true
},
{
"model_group": "gpt-4o-search-preview",
"providers": ["openai"],

View file

@ -1473,6 +1473,20 @@ LiteLLM translates OpenAI's `reasoning_effort` to Anthropic's `thinking` paramet
| "medium" | "budget_tokens": 2048 |
| "high" | "budget_tokens": 4096 |
:::note
For Claude Opus 4.6, all `reasoning_effort` values (`low`, `medium`, `high`) are mapped to `thinking: {type: "adaptive"}`. To use explicit thinking budgets, pass the native `thinking` parameter directly:
```python
from litellm import completion
resp = completion(
model="anthropic/claude-opus-4-6",
messages=[{"role": "user", "content": "What is the capital of France?"}],
thinking={"type": "enabled", "budget_tokens": 1024},
)
```
:::
<Tabs>
<TabItem value="sdk" label="SDK">
@ -1614,8 +1628,65 @@ curl http://0.0.0.0:4000/v1/chat/completions \
</TabItem>
</Tabs>
#### Adaptive Thinking (Claude Opus 4.6)
<Tabs>
<TabItem value="sdk" label="SDK">
```python
response = litellm.completion(
model="anthropic/claude-opus-4-6",
messages=[{"role": "user", "content": "What is the optimal strategy for solving this problem?"}],
thinking={"type": "adaptive"},
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```bash
curl http://0.0.0.0:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "anthropic/claude-opus-4-6",
"messages": [{"role": "user", "content": "What is the optimal strategy for solving this problem?"}],
"thinking": {"type": "adaptive"}
}'
```
</TabItem>
</Tabs>
#### Enabled Thinking with Budget
<Tabs>
<TabItem value="sdk" label="SDK">
```python
response = litellm.completion(
model="anthropic/claude-opus-4-6",
messages=[{"role": "user", "content": "What is the capital of France?"}],
thinking={"type": "enabled", "budget_tokens": 5000},
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```bash
curl http://0.0.0.0:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "anthropic/claude-opus-4-6",
"messages": [{"role": "user", "content": "What is the capital of France?"}],
"thinking": {"type": "enabled", "budget_tokens": 5000}
}'
```
</TabItem>
</Tabs>
## **Passing Extra Headers to Anthropic API**

View file

@ -1,7 +1,7 @@
# Dashscope (Qwen API)
# Dashscope API (Qwen models)
https://dashscope.console.aliyun.com/
**We support ALL Qwen models, just set `dashscope/` as a prefix when sending completion requests**
**We support ALL Qwen models (from Alibaba Cloud), just set `dashscope/` as a prefix when sending completion requests**
## API Key
```python
@ -9,6 +9,26 @@ https://dashscope.console.aliyun.com/
os.environ['DASHSCOPE_API_KEY']
```
## API Base
You can optionally specify the API base URL depending on your region:
| Region | API Base |
|--------|----------|
| **International** | `https://dashscope-intl.aliyuncs.com/compatible-mode/v1` |
| **China/Beijing** | `https://dashscope.aliyuncs.com/compatible-mode/v1` |
```python
# Set via environment variable
os.environ['DASHSCOPE_API_BASE'] = "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"
# Or pass directly in the completion call
response = completion(
model="dashscope/qwen-turbo",
messages=[{"role": "user", "content": "hello"}],
api_base="https://dashscope-intl.aliyuncs.com/compatible-mode/v1"
)
```
## Sample Usage
```python
from litellm import completion
@ -43,9 +63,7 @@ for chunk in response:
```
## Supported Models - ALL Qwen Models Supported!
We support ALL Qwen models, just set `dashscope/` as a prefix when sending completion requests
## All supported Models
[DashScope Model List](https://help.aliyun.com/zh/model-studio/compatibility-of-openai-with-dashscope?spm=a2c4g.11186623.help-menu-2400256.d_2_8_0.1efd516e2tTXBn&scm=20140722.H_2833609._.OR_help-T_cn~zh-V_1#7f9c78ae99pwz)

View file

@ -230,7 +230,70 @@ os.environ["OPENAI_BASE_URL"] = "https://your_host/v1" # OPTIONAL
These also support the `OPENAI_BASE_URL` environment variable, which can be used to specify a custom API endpoint.
## OpenAI Vision Models
### OpenAI Web Search Models
OpenAI has two ways to use web search, depending on the endpoint:
| Approach | Endpoint | Models | How to enable |
|----------|----------|--------|---------------|
| **Search Models** | `/chat/completions` | `gpt-5-search-api`, `gpt-4o-search-preview`, `gpt-4o-mini-search-preview` | Pass `web_search_options` parameter |
| **Web Search Tool** | `/responses` | `gpt-5`, `gpt-4.1`, `gpt-4o`, and other regular models | Pass `web_search_preview` tool |
<Tabs>
<TabItem value="sdk-completion" label="SDK - /chat/completions">
```python showLineNumbers
from litellm import completion
response = completion(
model="openai/gpt-5-search-api",
messages=[{"role": "user", "content": "What is the capital of France?"}],
web_search_options={
"search_context_size": "medium" # Options: "low", "medium", "high"
}
)
```
</TabItem>
<TabItem value="sdk-responses" label="SDK - /responses">
```python showLineNumbers
from litellm import responses
response = responses(
model="openai/gpt-5",
input="What is the capital of France?",
tools=[{
"type": "web_search_preview",
"search_context_size": "low"
}]
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```yaml
model_list:
# Search model for /chat/completions
- model_name: gpt-5-search-api
litellm_params:
model: openai/gpt-5-search-api
api_key: os.environ/OPENAI_API_KEY
# Regular model for /responses with web_search_preview tool
- model_name: gpt-5
litellm_params:
model: openai/gpt-5
api_key: os.environ/OPENAI_API_KEY
```
</TabItem>
</Tabs>
For full details, see the [Web Search guide](../completion/web_search.md).
## OpenAI Vision Models
| Model Name | Function Call |
|-----------------------|-----------------------------------------------------------------|
| gpt-4o | `response = completion(model="gpt-4o", messages=messages)` |

View file

@ -37,6 +37,24 @@ for event in response:
print(event)
```
#### Web Search
```python showLineNumbers title="OpenAI Responses with Web Search"
import litellm
response = litellm.responses(
model="openai/gpt-5",
input="What is the capital of France?",
tools=[{
"type": "web_search_preview",
"search_context_size": "medium" # Options: "low", "medium", "high"
}]
)
print(response)
```
For full details, see the [Web Search guide](../../completion/web_search.md).
#### Image Generation with Streaming
```python showLineNumbers title="OpenAI Streaming Image Generation"
import litellm

View file

@ -0,0 +1,62 @@
# Scaleway
LiteLLM supports all [models available on Scaleway Generative APIs ↗](https://www.scaleway.com/en/docs/generative-apis/reference-content/supported-models/).
## Usage with LiteLLM Python SDK
```python
import os
from litellm import completion
os.environ["SCW_SECRET_KEY"] = "your-scaleway-secret-key"
messages = [{"role": "user", "content": "Write a short poem"}]
response = completion(model="scaleway/qwen3-235b-a22b-instruct-2507", messages=messages)
print(response)
```
## Usage with LiteLLM Proxy
### 1. Set Scaleway models in config.yaml
```yaml
model_list:
- model_name: scaleway-model
litellm_params:
model: scaleway/qwen3-235b-a22b-instruct-2507
api_key: "os.environ/SCW_SECRET_KEY" # ensure you have `SCW_SECRET_KEY` in your .env
```
### 2. Start proxy
```bash
litellm --config config.yaml
```
### 3. Query proxy
Assuming the proxy is running on [http://localhost:4000](http://localhost:4000):
```bash
curl http://localhost:4000/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer YOUR_LITELLM_MASTER_KEY" \
-d '{
"model": "scaleway-model",
"messages": [
{
"role": "system",
"content": "You are a helpful assistant."
},
{
"role": "user",
"content": "Write a short poem"
}
]
}'
```
`-H "Authorization: Bearer YOUR_LITELLM_MASTER_KEY" ` is only required if you have set a LiteLLM master key
## Supported features
Scaleway provider supports all features in [Generative APIs reference documentation ↗](https://www.scaleway.com/en/developers/api/generative-apis/), such as streaming, structured outputs and tool calling.

View file

@ -555,7 +555,7 @@ router_settings:
| DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT | Default token count for mock response completions. Default is 20
| DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT | Default token count for mock response prompts. Default is 10
| DEFAULT_MODEL_CREATED_AT_TIME | Default creation timestamp for models. Default is 1677610602
| DEFAULT_NUM_WORKERS_LITELLM_PROXY | Default number of workers for LiteLLM proxy. Default is 4. **We strongly recommend setting NUM Workers to Number of vCPUs available**
| DEFAULT_NUM_WORKERS_LITELLM_PROXY | Default number of workers for LiteLLM proxy when `NUM_WORKERS` is not set. Default is 1. **We strongly recommend setting NUM_WORKERS to the number of vCPUs available** (e.g. `NUM_WORKERS=8` or `--num_workers 8`).
| DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD | Default threshold for prompt injection similarity. Default is 0.7
| DEFAULT_POLLING_INTERVAL | Default polling interval for schedulers in seconds. Default is 0.03
| DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET | Default reasoning effort disable thinking budget. Default is 0
@ -746,6 +746,7 @@ router_settings:
| LITERAL_API_URL | API URL for Literal service
| LITERAL_BATCH_SIZE | Batch size for Literal operations
| LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX | Disable automatic URL suffix appending for Anthropic API base URLs. When set to `true`, prevents LiteLLM from automatically adding `/v1/messages` or `/v1/complete` to custom Anthropic API endpoints
| LITELLM_ASSETS_PATH | Path to directory for UI assets and logos. Used when running with read-only filesystem (e.g., Kubernetes). Default is `/var/lib/litellm/assets` in Docker.
| LITELLM_CLI_JWT_EXPIRATION_HOURS | Expiration time in hours for CLI-generated JWT tokens. Default is 24 hours
| LITELLM_DD_AGENT_HOST | Hostname or IP of DataDog agent for LiteLLM-specific logging. When set, logs are sent to agent instead of direct API
| LITELLM_DD_AGENT_PORT | Port of DataDog agent for LiteLLM-specific log intake. Default is 10518
@ -760,6 +761,7 @@ router_settings:
| LITELLM_MIGRATION_DIR | Custom migrations directory for prisma migrations, used for baselining db in read-only file systems.
| LITELLM_HOSTED_UI | URL of the hosted UI for LiteLLM
| LITELLM_UI_API_DOC_BASE_URL | Optional override for the API Reference base URL (used in sample code/docs) when the admin UI runs on a different host than the proxy. Defaults to `PROXY_BASE_URL` when unset.
| LITELLM_UI_PATH | Path to directory for Admin UI files. Used when running with read-only filesystem (e.g., Kubernetes). Default is `/var/lib/litellm/ui` in Docker.
| LITELM_ENVIRONMENT | Environment of LiteLLM Instance, used by logging services. Currently only used by DeepEval.
| LITELLM_KEY_ROTATION_ENABLED | Enable auto-key rotation for LiteLLM (boolean). Default is false.
| LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS | Interval in seconds for how often to run job that auto-rotates keys. Default is 86400 (24 hours).

View file

@ -469,6 +469,7 @@ credential_list:
api_version: "2023-05-15"
credential_info:
description: "Production credentials for EU region"
custom_llm_provider: "azure"
```
#### Key Parameters

View file

@ -250,11 +250,133 @@ The migrate deploy command:
### Read-only File System
If you see a `Permission denied` error, it means the LiteLLM pod is running with a read-only file system.
Running LiteLLM with `readOnlyRootFilesystem: true` is a Kubernetes security best practice that prevents container processes from writing to the root filesystem. LiteLLM fully supports this configuration.
To fix this, just set `LITELLM_MIGRATION_DIR="/path/to/writeable/directory"` in your environment.
#### Quick Fix for Permission Errors
LiteLLM will use this directory to write migration files.
If you see a `Permission denied` error, it means the LiteLLM pod is running with a read-only file system. LiteLLM needs writable directories for:
- **Database migrations**: Set `LITELLM_MIGRATION_DIR="/path/to/writable/directory"`
- **Admin UI**: Set `LITELLM_UI_PATH="/path/to/writable/directory"`
- **UI assets/logos**: Set `LITELLM_ASSETS_PATH="/path/to/writable/directory"`
#### Complete Read-Only Filesystem Setup (Kubernetes)
For production deployments with enhanced security, use this configuration:
**Option 1: Using EmptyDir Volumes with InitContainer (Recommended)**
This approach copies the pre-built UI from the Docker image to writable emptyDir volumes at pod startup.
```yaml
apiVersion: apps/v1
kind: Deployment
metadata:
name: litellm-proxy
spec:
template:
spec:
initContainers:
- name: setup-ui
image: ghcr.io/berriai/litellm:main-stable
command:
- sh
- -c
- |
cp -r /var/lib/litellm/ui/* /app/var/litellm/ui/ && \
cp -r /var/lib/litellm/assets/* /app/var/litellm/assets/
volumeMounts:
- name: ui-volume
mountPath: /app/var/litellm/ui
- name: assets-volume
mountPath: /app/var/litellm/assets
containers:
- name: litellm
image: ghcr.io/berriai/litellm:main-stable
env:
- name: LITELLM_NON_ROOT
value: "true"
- name: LITELLM_UI_PATH
value: "/app/var/litellm/ui"
- name: LITELLM_ASSETS_PATH
value: "/app/var/litellm/assets"
- name: LITELLM_MIGRATION_DIR
value: "/app/migrations"
- name: PRISMA_BINARY_CACHE_DIR
value: "/app/cache/prisma-python/binaries"
- name: XDG_CACHE_HOME
value: "/app/cache"
securityContext:
readOnlyRootFilesystem: true
runAsNonRoot: true
runAsUser: 101
capabilities:
drop:
- ALL
volumeMounts:
- name: config
mountPath: /app/config.yaml
subPath: config.yaml
readOnly: true
- name: ui-volume
mountPath: /app/var/litellm/ui
- name: assets-volume
mountPath: /app/var/litellm/assets
- name: cache
mountPath: /app/cache
- name: migrations
mountPath: /app/migrations
volumes:
- name: config
configMap:
name: litellm-config
- name: ui-volume
emptyDir:
sizeLimit: 100Mi
- name: assets-volume
emptyDir:
sizeLimit: 10Mi
- name: cache
emptyDir:
sizeLimit: 500Mi
- name: migrations
emptyDir:
sizeLimit: 64Mi
```
**Option 2: Without UI (API-only deployment)**
If you don't need the admin UI, you can run with minimal configuration:
```yaml
env:
- name: LITELLM_NON_ROOT
value: "true"
- name: LITELLM_MIGRATION_DIR
value: "/app/migrations"
securityContext:
readOnlyRootFilesystem: true
```
The proxy will log a warning about the UI but API endpoints will work normally.
#### Environment Variables for Read-Only Filesystems
| Variable | Purpose | Default |
|----------|---------|---------|
| `LITELLM_UI_PATH` | Admin UI directory | `/var/lib/litellm/ui` (Docker) |
| `LITELLM_ASSETS_PATH` | UI assets/logos | `/var/lib/litellm/assets` (Docker) |
| `LITELLM_MIGRATION_DIR` | Database migrations | Package directory |
| `PRISMA_BINARY_CACHE_DIR` | Prisma binary cache | System default |
| `XDG_CACHE_HOME` | General cache directory | System default |
#### Important Notes
1. **Migrations**: Always set `LITELLM_MIGRATION_DIR` to a writable emptyDir path
2. **Prisma Cache**: Set `PRISMA_BINARY_CACHE_DIR` and `XDG_CACHE_HOME` to writable paths
3. **Server Root Path**: If using a custom `server_root_path`, you must pre-process UI files in your Dockerfile as the proxy cannot modify files at runtime with read-only filesystem
4. **Automatic Detection**: The UI is automatically detected as pre-restructured if it contains a `.litellm_ui_ready` marker file (created by the official Docker images)
## 10. Use a Separate Health Check App
:::info

View file

@ -1023,6 +1023,134 @@ curl http://localhost:4000/v1/responses \
## Server-side compaction
For long-running conversations, you can enable **server-side compaction** so that when the rendered context size crosses a threshold, the server automatically runs compaction in-stream and emits a compaction item—no separate `POST /v1/responses/compact` call is required.
Supported on the OpenAI Responses API when using the `openai` or `azure` provider. Pass `context_management` with a compaction entry and `compact_threshold` (token count; minimum 1000). When the context crosses the threshold, the server compacts in-stream and continues. Chain turns with `previous_response_id` or by appending output items to your next input array. See [OpenAI Compaction guide](https://developers.openai.com/api/docs/guides/compaction) for details.
For explicit control over when compaction runs, use the standalone compact endpoint (`POST /v1/responses/compact`) instead.
### Python SDK
```python showLineNumbers title="Server-side compaction with LiteLLM Python SDK"
import litellm
# Non-streaming: enable compaction when context exceeds 200k tokens
response = litellm.responses(
model="openai/gpt-4o",
input="Your conversation input...",
context_management=[{"type": "compaction", "compact_threshold": 200000}],
max_output_tokens=1024,
)
print(response)
# Streaming: same context_management, compaction runs in-stream if threshold is crossed
stream = litellm.responses(
model="openai/gpt-4o",
input="Your conversation input...",
context_management=[{"type": "compaction", "compact_threshold": 200000}],
stream=True,
)
for event in stream:
print(event)
```
### LiteLLM Proxy (AI Gateway)
Use the OpenAI SDK with your proxy as `base_url`, or call the proxy with curl. The proxy forwards `context_management` to the provider.
**OpenAI Python SDK (proxy as base_url):**
```python showLineNumbers title="Server-side compaction via LiteLLM Proxy"
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:4000", # LiteLLM Proxy (AI Gateway)
api_key="your-proxy-api-key",
)
response = client.responses.create(
model="openai/gpt-4o",
input="Your conversation input...",
context_management=[{"type": "compaction", "compact_threshold": 200000}],
max_output_tokens=1024,
)
print(response)
```
**curl (proxy):**
```bash title="Server-side compaction via curl to LiteLLM Proxy"
curl -X POST "http://localhost:4000/v1/responses" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer your-proxy-api-key" \
-d '{
"model": "openai/gpt-4o",
"input": "Your conversation input...",
"context_management": [{"type": "compaction", "compact_threshold": 200000}],
"max_output_tokens": 1024
}'
```
## Shell tool
The **Shell tool** lets the model run commands in a hosted container or local runtime (OpenAI Responses API). You pass `tools=[{"type": "shell", "environment": {...}}]`; the `environment` object configures the runtime (e.g. `type: "container_auto"` for auto-provisioned containers). See [OpenAI Shell tool guide](https://developers.openai.com/api/docs/guides/tools-shell) for full options.
Supported when using the `openai` or `azure` provider with a model that supports the Shell tool.
### Python SDK
```python showLineNumbers title="Shell tool with LiteLLM Python SDK"
import litellm
response = litellm.responses(
model="openai/gpt-5.2",
input="List files in /mnt/data and run python --version.",
tools=[{"type": "shell", "environment": {"type": "container_auto"}}],
tool_choice="auto",
max_output_tokens=1024,
)
```
### LiteLLM Proxy (AI Gateway)
Use the OpenAI SDK with your proxy as `base_url`, or call the proxy with curl. The proxy forwards `tools` (including `type: "shell"`) to the provider.
**OpenAI Python SDK (proxy as base_url):**
```python showLineNumbers title="Shell tool via LiteLLM Proxy"
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:4000",
api_key="your-proxy-api-key",
)
response = client.responses.create(
model="openai/gpt-5.2",
input="List files in /mnt/data.",
tools=[{"type": "shell", "environment": {"type": "container_auto"}}],
tool_choice="auto",
max_output_tokens=1024,
)
```
**curl:**
```bash title="Shell tool via curl to LiteLLM Proxy"
curl -X POST "http://localhost:4000/v1/responses" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer your-proxy-api-key" \
-d '{
"model": "openai/gpt-5.2",
"input": "List files in /mnt/data.",
"tools": [{"type": "shell", "environment": {"type": "container_auto"}}],
"tool_choice": "auto",
"max_output_tokens": 1024
}'
```
## Session Management
LiteLLM Proxy supports session management for all supported models. This allows you to store and fetch conversation history (state) in LiteLLM Proxy.

View file

@ -874,6 +874,7 @@ const sidebars = {
},
"providers/sambanova",
"providers/sap",
"providers/scaleway",
"providers/stability",
"providers/synthetic",
"providers/snowflake",

View file

@ -1,11 +1,15 @@
from typing import Dict, Literal, Type, Union
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
from litellm_enterprise.proxy.hooks.managed_vector_stores import (
_PROXY_LiteLLMManagedVectorStores,
)
from litellm.integrations.custom_logger import CustomLogger
ENTERPRISE_PROXY_HOOKS: Dict[str, Type[CustomLogger]] = {
"managed_files": _PROXY_LiteLLMManagedFiles,
"managed_vector_stores": _PROXY_LiteLLMManagedVectorStores,
}
@ -13,6 +17,7 @@ def get_enterprise_proxy_hook(
hook_name: Union[
Literal[
"managed_files",
"managed_vector_stores",
"max_parallel_requests",
],
str,

View file

@ -41,6 +41,10 @@ class EnterpriseRouteChecks:
return get_secret_bool("DISABLE_ADMIN_ENDPOINTS") is True
# Routes that should remain accessible even when LLM API endpoints are disabled.
# These are read-only model listing routes needed by the Admin UI.
LLM_API_EXEMPT_ROUTES = ["/models", "/v1/models"]
@staticmethod
def should_call_route(route: str):
"""
@ -58,6 +62,7 @@ class EnterpriseRouteChecks:
)
elif (
RouteChecks.is_llm_api_route(route=route)
and route not in EnterpriseRouteChecks.LLM_API_EXEMPT_ROUTES
and EnterpriseRouteChecks.is_llm_api_route_disabled()
):
raise HTTPException(

View file

@ -0,0 +1,464 @@
# What is this?
## This hook is used to manage vector stores with target_model_names support
## It allows creating vector stores across multiple models and managing them with unified IDs
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
from fastapi import HTTPException
import litellm
from litellm import Router, verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.managed_resources import BaseManagedResource
from litellm.llms.base_llm.managed_resources.utils import (
generate_unified_id_string,
is_base64_encoded_unified_id,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.vector_stores import (
VectorStoreCreateOptionalRequestParams,
VectorStoreCreateResponse,
)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
from litellm.proxy.utils import PrismaClient as _PrismaClient
Span = Union[_Span, Any]
InternalUsageCache = _InternalUsageCache
PrismaClient = _PrismaClient
else:
Span = Any
InternalUsageCache = Any
PrismaClient = Any
class _PROXY_LiteLLMManagedVectorStores(
CustomLogger, BaseManagedResource[VectorStoreCreateResponse]
):
"""
Managed vector stores with target_model_names support.
This class provides functionality to:
- Create vector stores across multiple models
- Retrieve vector stores by unified ID
- Delete vector stores from all models
- List vector stores created by a user
"""
def __init__(
self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient
):
CustomLogger.__init__(self)
BaseManagedResource.__init__(self, internal_usage_cache, prisma_client)
# ============================================================================
# ABSTRACT METHOD IMPLEMENTATIONS
# ============================================================================
@property
def resource_type(self) -> str:
"""Return the resource type identifier."""
return "vector_store"
@property
def table_name(self) -> str:
"""Return the database table name for vector stores."""
# Prisma converts model name LiteLLM_ManagedVectorStoreTable to litellm_managedvectorstoretable
return "litellm_managedvectorstoretable"
def get_unified_resource_id_format(
self,
resource_object: VectorStoreCreateResponse,
target_model_names_list: List[str],
) -> str:
"""
Generate the format string for the unified vector store ID.
Format:
litellm_proxy:vector_store;unified_id,<uuid>;target_model_names,<models>;resource_id,<vs_id>;model_id,<model_id>
"""
# VectorStoreCreateResponse is a TypedDict, so resource_object is a dictionary
# Extract provider resource ID from the response
provider_resource_id = resource_object.get("id", "")
# Model ID is stored in hidden params if the response object supports it
# For TypedDict responses, we need to check if _hidden_params was added
hidden_params: Dict[str, Any] = {}
if hasattr(resource_object, "_hidden_params"):
hidden_params = getattr(resource_object, "_hidden_params", {}) or {}
model_id = hidden_params.get("model_id", "")
return generate_unified_id_string(
resource_type=self.resource_type,
unified_uuid=str(uuid.uuid4()),
target_model_names=target_model_names_list,
provider_resource_id=provider_resource_id,
model_id=model_id,
)
async def create_resource_for_model(
self,
llm_router: Router,
model: str,
request_data: Dict[str, Any],
litellm_parent_otel_span: Span,
) -> VectorStoreCreateResponse:
"""
Create a vector store for a specific model.
Args:
llm_router: LiteLLM router instance
model: Model name to create vector store for
request_data: Request data for vector store creation
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
VectorStoreCreateResponse from the provider
"""
# Use the router to create the vector store
response = await llm_router.avector_store_create(
model=model, **request_data
)
return response
# ============================================================================
# VECTOR STORE CRUD OPERATIONS
# ============================================================================
async def acreate_vector_store(
self,
create_request: VectorStoreCreateOptionalRequestParams,
llm_router: Router,
target_model_names_list: List[str],
litellm_parent_otel_span: Span,
user_api_key_dict: UserAPIKeyAuth,
) -> VectorStoreCreateResponse:
"""
Create a vector store across multiple models.
Args:
create_request: Vector store creation request parameters
llm_router: LiteLLM router instance
target_model_names_list: List of target model names
litellm_parent_otel_span: OpenTelemetry span for tracing
user_api_key_dict: User API key authentication details
Returns:
VectorStoreCreateResponse with unified ID
"""
verbose_logger.info(
f"Creating managed vector store for models: {target_model_names_list}"
)
# Create vector store for each model
# Convert TypedDict to Dict[str, Any] for base class compatibility
request_data_dict: Dict[str, Any] = dict(create_request)
responses = await self.create_resource_for_each_model(
llm_router=llm_router,
request_data=request_data_dict,
target_model_names_list=target_model_names_list,
litellm_parent_otel_span=litellm_parent_otel_span,
)
# Generate unified ID
unified_id = self.generate_unified_resource_id(
resource_objects=responses,
target_model_names_list=target_model_names_list,
)
# Extract model mappings from responses
model_mappings: Dict[str, str] = {}
for response in responses:
hidden_params = getattr(response, "_hidden_params", {}) or {}
model_id = hidden_params.get("model_id")
if model_id:
# VectorStoreCreateResponse is a TypedDict, use dict access
model_mappings[model_id] = response["id"]
verbose_logger.debug(
f"Created vector stores with model mappings: {model_mappings}"
)
# Store in database
await self.store_unified_resource_id(
unified_resource_id=unified_id,
resource_object=responses[0], # Store first response as template
litellm_parent_otel_span=litellm_parent_otel_span,
model_mappings=model_mappings,
user_api_key_dict=user_api_key_dict,
)
# Return response with unified ID
# VectorStoreCreateResponse is a TypedDict, so we need to create a new dict with the unified ID
response = responses[0].copy()
response["id"] = unified_id
verbose_logger.info(
f"Successfully created managed vector store with unified ID: {unified_id}"
)
return response
async def alist_vector_stores(
self,
user_api_key_dict: UserAPIKeyAuth,
limit: Optional[int] = None,
after: Optional[str] = None,
order: Optional[str] = None,
) -> Dict[str, Any]:
"""
List vector stores created by a user.
Args:
user_api_key_dict: User API key authentication details
limit: Maximum number of vector stores to return
after: Cursor for pagination
order: Sort order ('asc' or 'desc')
Returns:
Dictionary with list of vector stores and pagination info
"""
# Use the base class method
return await self.list_user_resources(
user_api_key_dict=user_api_key_dict,
limit=limit,
after=after,
)
# ============================================================================
# ACCESS CONTROL
# ============================================================================
async def check_vector_store_access(
self, vector_store_id: str, user_api_key_dict: UserAPIKeyAuth
) -> bool:
"""
Check if user has access to a vector store.
Args:
vector_store_id: The unified vector store ID
user_api_key_dict: User API key authentication details
Returns:
True if user has access, False otherwise
"""
is_unified_id = is_base64_encoded_unified_id(vector_store_id)
if is_unified_id:
# Check access for managed vector store
return await self.can_user_access_unified_resource_id(
vector_store_id,
user_api_key_dict,
)
# Not a managed vector store, allow access
return True
async def check_managed_vector_store_access(
self, data: Dict, user_api_key_dict: UserAPIKeyAuth
) -> bool:
"""
Check if user has access to a managed vector store in request data.
Args:
data: Request data containing vector_store_id
user_api_key_dict: User API key authentication details
Returns:
True if this is a managed vector store and user has access
Raises:
HTTPException: If user doesn't have access
"""
vector_store_id = cast(Optional[str], data.get("vector_store_id"))
is_unified_id = (
is_base64_encoded_unified_id(vector_store_id)
if vector_store_id
else False
)
if is_unified_id and vector_store_id:
if await self.can_user_access_unified_resource_id(
vector_store_id, user_api_key_dict
):
return True
else:
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
)
return False
# ============================================================================
# PRE-CALL HOOK (For Router Integration)
# ============================================================================
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: Any,
data: Dict,
call_type: str,
) -> Union[Exception, str, Dict, None]:
"""
Pre-call hook to handle vector store operations.
This hook intercepts vector store requests and:
- Validates access for managed vector stores
- Transforms unified IDs to provider-specific IDs
- Adds model routing information
Args:
user_api_key_dict: User API key authentication details
cache: Cache instance
data: Request data
call_type: Type of call being made
Returns:
Modified request data or None
"""
from litellm.llms.base_llm.managed_resources.utils import (
is_base64_encoded_unified_id,
parse_unified_id,
)
# Handle vector store search operations
if call_type == "avector_store_search":
vector_store_id = data.get("vector_store_id")
if vector_store_id:
# Check if it's a managed vector store ID
decoded_id = is_base64_encoded_unified_id(vector_store_id)
if decoded_id:
verbose_logger.debug(
f"Processing managed vector store search: {vector_store_id}"
)
# Check access
has_access = await self.can_user_access_unified_resource_id(
vector_store_id, user_api_key_dict
)
if not has_access:
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}",
)
# Parse the unified ID to extract components
parsed_id = parse_unified_id(vector_store_id)
if parsed_id:
# Extract the model ID and provider resource ID
model_id = parsed_id.get("model_id")
provider_resource_id = parsed_id.get("provider_resource_id")
target_model_names = parsed_id.get("target_model_names", [])
verbose_logger.debug(
f"Decoded vector store - model_id: {model_id}, provider_resource_id: {provider_resource_id}, target_model_names: {target_model_names}"
)
# Determine which model to use for routing
# Priority: model_id (deployment ID) > first target_model_name
routing_model = None
if model_id:
routing_model = model_id
elif target_model_names and len(target_model_names) > 0:
routing_model = target_model_names[0]
# Set the model for routing
if routing_model:
data["model"] = routing_model
verbose_logger.info(
f"Routing vector store search to model: {routing_model}"
)
# Replace the unified ID with the provider-specific ID
if provider_resource_id:
data["vector_store_id"] = provider_resource_id
verbose_logger.debug(
f"Replaced unified ID with provider resource ID: {provider_resource_id}"
)
# Handle vector store retrieve/delete operations
elif call_type in ("avector_store_retrieve", "avector_store_delete"):
await self.check_managed_vector_store_access(data, user_api_key_dict)
# If it's a managed vector store, we'll handle it in the endpoint
# No need to transform here as the endpoint will route to the hook
return data
# ============================================================================
# POST-CALL HOOK (For Response Transformation)
# ============================================================================
async def async_post_call_success_hook(
self,
data: Dict,
user_api_key_dict: UserAPIKeyAuth,
response: Any,
) -> Any:
"""
Post-call hook to transform responses.
This hook can be used to transform responses if needed.
For now, it just passes through the response.
Args:
data: Request data
user_api_key_dict: User API key authentication details
response: Response from the provider
Returns:
Potentially modified response
"""
# Currently no transformation needed
return response
# ============================================================================
# DEPLOYMENT FILTERING
# ============================================================================
async def async_filter_deployments( # type: ignore[override]
self,
model: str,
healthy_deployments: List,
messages: Optional[List] = None,
request_kwargs: Optional[Dict] = None,
parent_otel_span: Optional[Span] = None,
) -> List[Dict]:
"""
Filter deployments based on vector store availability.
This is used by the router to select only deployments that have
the vector store available.
Note: This method signature is a compromise between CustomLogger and BaseManagedResource
parent classes which have incompatible signatures. The type: ignore[override] is necessary
due to this multiple inheritance conflict.
Args:
model: Model name
healthy_deployments: List of healthy deployments
messages: Messages (unused for vector stores, required by CustomLogger interface)
request_kwargs: Request kwargs containing vector_store_id and mappings
parent_otel_span: OpenTelemetry span for tracing
Returns:
Filtered list of deployments
"""
return await BaseManagedResource.async_filter_deployments(
self,
model=model,
healthy_deployments=healthy_deployments,
request_kwargs=request_kwargs,
parent_otel_span=parent_otel_span,
resource_id_key="vector_store_id",
)

Binary file not shown.

Binary file not shown.

View file

@ -0,0 +1,3 @@
-- AlterTable
ALTER TABLE "LiteLLM_PolicyAttachmentTable" ADD COLUMN "tags" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -0,0 +1,33 @@
-- AlterTable
ALTER TABLE "LiteLLM_DeletedTeamTable" ADD COLUMN "access_group_ids" TEXT[] DEFAULT ARRAY[]::TEXT[];
-- AlterTable
ALTER TABLE "LiteLLM_DeletedVerificationToken" ADD COLUMN "access_group_ids" TEXT[] DEFAULT ARRAY[]::TEXT[];
-- AlterTable
ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN "access_group_ids" TEXT[] DEFAULT ARRAY[]::TEXT[];
-- AlterTable
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "access_group_ids" TEXT[] DEFAULT ARRAY[]::TEXT[];
-- CreateTable
CREATE TABLE "LiteLLM_AccessGroupTable" (
"access_group_id" TEXT NOT NULL,
"access_group_name" TEXT NOT NULL,
"description" TEXT,
"access_model_ids" TEXT[] DEFAULT ARRAY[]::TEXT[],
"access_mcp_server_ids" TEXT[] DEFAULT ARRAY[]::TEXT[],
"access_agent_ids" TEXT[] DEFAULT ARRAY[]::TEXT[],
"assigned_team_ids" TEXT[] DEFAULT ARRAY[]::TEXT[],
"assigned_key_ids" TEXT[] DEFAULT ARRAY[]::TEXT[],
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"created_by" TEXT,
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_by" TEXT,
CONSTRAINT "LiteLLM_AccessGroupTable_pkey" PRIMARY KEY ("access_group_id")
);
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_AccessGroupTable_access_group_name_key" ON "LiteLLM_AccessGroupTable"("access_group_name");

View file

@ -0,0 +1,22 @@
-- CreateTable
CREATE TABLE "LiteLLM_ManagedVectorStoreTable" (
"id" TEXT NOT NULL,
"unified_resource_id" TEXT NOT NULL,
"resource_object" JSONB,
"model_mappings" JSONB NOT NULL,
"flat_model_resource_ids" TEXT[] DEFAULT ARRAY[]::TEXT[],
"storage_backend" TEXT,
"storage_url" TEXT,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"created_by" TEXT,
"updated_at" TIMESTAMP(3) NOT NULL,
"updated_by" TEXT,
CONSTRAINT "LiteLLM_ManagedVectorStoreTable_pkey" PRIMARY KEY ("id")
);
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_ManagedVectorStoreTable_unified_resource_id_key" ON "LiteLLM_ManagedVectorStoreTable"("unified_resource_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedVectorStoreTable_unified_resource_id_idx" ON "LiteLLM_ManagedVectorStoreTable"("unified_resource_id");

View file

@ -128,6 +128,7 @@ model LiteLLM_TeamTable {
model_max_budget Json @default("{}")
router_settings Json? @default("{}")
team_member_permissions String[] @default([])
access_group_ids String[] @default([])
policies String[] @default([])
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team
@ -161,6 +162,7 @@ model LiteLLM_DeletedTeamTable {
model_max_budget Json @default("{}")
router_settings Json? @default("{}")
team_member_permissions String[] @default([])
access_group_ids String[] @default([])
policies String[] @default([])
model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false)
@ -293,6 +295,7 @@ model LiteLLM_VerificationToken {
allowed_cache_controls String[] @default([])
allowed_routes String[] @default([])
policies String[] @default([])
access_group_ids String[] @default([])
model_spend Json @default("{}")
model_max_budget Json @default("{}")
budget_id String?
@ -348,6 +351,7 @@ model LiteLLM_DeletedVerificationToken {
allowed_cache_controls String[] @default([])
allowed_routes String[] @default([])
policies String[] @default([])
access_group_ids String[] @default([])
model_spend Json @default("{}")
model_max_budget Json @default("{}")
router_settings Json? @default("{}")
@ -766,6 +770,22 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t
@@index([model_object_id])
}
model LiteLLM_ManagedVectorStoreTable {
id String @id @default(uuid())
unified_resource_id String @unique // The base64 encoded unified vector store ID
resource_object Json? // Stores the VectorStoreCreateResponse
model_mappings Json // Maps model_id -> provider_vector_store_id
flat_model_resource_ids String[] @default([]) // Flat list of provider vector store IDs for faster querying
storage_backend String? // Storage backend name (if applicable)
storage_url String? // Storage URL (if applicable)
created_at DateTime @default(now())
created_by String?
updated_at DateTime @updatedAt
updated_by String?
@@index([unified_resource_id])
}
model LiteLLM_ManagedVectorStoresTable {
vector_store_id String @id
custom_llm_provider String
@ -920,3 +940,23 @@ model LiteLLM_PolicyAttachmentTable {
updated_at DateTime @default(now()) @updatedAt
updated_by String?
}
//Unified Access Groups table for storing unified access groups
model LiteLLM_AccessGroupTable {
access_group_id String @id @default(uuid())
access_group_name String @unique
description String?
// Resource memberships - explicit arrays per type
access_model_ids String[] @default([])
access_mcp_server_ids String[] @default([])
access_agent_ids String[] @default([])
assigned_team_ids String[] @default([])
assigned_key_ids String[] @default([])
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
}

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-proxy-extras"
version = "0.4.34"
version = "0.4.36"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
authors = ["BerriAI"]
readme = "README.md"
@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "0.4.34"
version = "0.4.36"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-proxy-extras==",

View file

@ -175,6 +175,7 @@ _async_failure_callback: List[Union[str, Callable, "CustomLogger"]] = ( # Custo
pre_call_rules: List[Callable] = []
post_call_rules: List[Callable] = []
turn_off_message_logging: Optional[bool] = False
standard_logging_payload_excluded_fields: Optional[List[str]] = None # Fields to exclude from StandardLoggingPayload before callbacks receive it
log_raw_request_response: bool = False
redact_messages_in_exceptions: Optional[bool] = False
redact_user_api_key_info: Optional[bool] = False

View file

@ -275,7 +275,6 @@ LLM_CONFIG_NAMES = (
"LmStudioEmbeddingConfig",
"NscaleConfig",
"PerplexityChatConfig",
"PerplexityResponsesConfig",
"AzureOpenAIO1Config",
"IBMWatsonXAIConfig",
"IBMWatsonXChatConfig",

View file

@ -19,6 +19,7 @@
"mcp-client-2025-11-20": "mcp-client-2025-11-20",
"mcp-client-2025-04-04": "mcp-client-2025-04-04",
"mcp-servers-2025-12-04": "mcp-servers-2025-12-04",
"oauth-2025-04-20": "oauth-2025-04-20",
"output-128k-2025-02-19": "output-128k-2025-02-19",
"prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05",
"skills-2025-10-02": "skills-2025-10-02",

View file

@ -39,11 +39,19 @@ async def _handle_completed_batch(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
model_name: Optional[str] = None,
litellm_params: Optional[dict] = None,
) -> Tuple[float, Usage, List[str]]:
"""Helper function to process a completed batch and handle logging"""
"""Helper function to process a completed batch and handle logging
Args:
batch: The batch object
custom_llm_provider: The LLM provider
model_name: Optional model name
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
"""
# Get batch results
file_content_dictionary = await _get_batch_output_file_content_as_dictionary(
batch, custom_llm_provider
batch, custom_llm_provider, litellm_params=litellm_params
)
# Calculate costs and usage
@ -187,9 +195,16 @@ def calculate_vertex_ai_batch_cost_and_usage(
async def _get_batch_output_file_content_as_dictionary(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
litellm_params: Optional[dict] = None,
) -> List[dict]:
"""
Get the batch output file content as a list of dictionaries
Args:
batch: The batch object
custom_llm_provider: The LLM provider
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
Required for Azure and other providers that need authentication
"""
from litellm.files.main import afile_content
from litellm.proxy.openai_files_endpoints.common_utils import (
@ -211,13 +226,50 @@ async def _get_batch_output_file_content_as_dictionary(
except (IndexError, AttributeError) as e:
verbose_logger.error(f"Failed to extract LLM output file ID from unified file ID: {batch.output_file_id}, error: {e}")
_file_content = await afile_content(
file_id=file_id,
custom_llm_provider=custom_llm_provider,
)
# Build kwargs for afile_content with credentials from litellm_params
file_content_kwargs = {
"file_id": file_id,
"custom_llm_provider": custom_llm_provider,
}
# Extract and add credentials for file access
credentials = _extract_file_access_credentials(litellm_params)
file_content_kwargs.update(credentials)
_file_content = await afile_content(**file_content_kwargs)
return _get_file_content_as_dictionary(_file_content.content)
def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict:
"""
Extract credentials from litellm_params for file access operations.
This method extracts relevant authentication and configuration parameters
needed for accessing files across different providers (Azure, Vertex AI, etc.).
Args:
litellm_params: Dictionary containing litellm parameters with credentials
Returns:
Dictionary containing only the credentials needed for file access
"""
credentials = {}
if litellm_params:
# List of credential keys that should be passed to file operations
credential_keys = [
"api_key", "api_base", "api_version", "organization",
"azure_ad_token", "azure_ad_token_provider",
"vertex_project", "vertex_location", "vertex_credentials",
"timeout", "max_retries"
]
for key in credential_keys:
if key in litellm_params:
credentials[key] = litellm_params[key]
return credentials
def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]:
"""
Get the file content as a list of dictionaries from JSON Lines format

View file

@ -12,7 +12,8 @@ import asyncio
import time
import traceback
from concurrent.futures import ThreadPoolExecutor
from typing import TYPE_CHECKING, Any, List, Optional, Union
from threading import Lock
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
if TYPE_CHECKING:
from litellm.types.caching import RedisPipelineIncrementOperation
@ -71,6 +72,7 @@ class DualCache(BaseCache):
self.last_redis_batch_access_time = LimitedSizeOrderedDict(
max_size=default_max_redis_batch_cache_size
)
self._last_redis_batch_access_time_lock = Lock()
self.redis_batch_cache_expiry = (
default_redis_batch_cache_expiry
or litellm.default_redis_batch_cache_expiry
@ -236,22 +238,46 @@ class DualCache(BaseCache):
except Exception:
verbose_logger.error(traceback.format_exc())
def get_redis_batch_keys(
def _reserve_redis_batch_keys(
self,
current_time: float,
keys: List[str],
result: List[Any],
) -> List[str]:
sublist_keys = []
for key, value in zip(keys, result):
if value is None:
) -> Tuple[List[str], Dict[str, Optional[float]]]:
"""
Atomically choose keys to fetch from Redis and reserve their access time.
This prevents check-then-act races under concurrent async callers.
"""
sublist_keys: List[str] = []
previous_access_times: Dict[str, Optional[float]] = {}
with self._last_redis_batch_access_time_lock:
for key, value in zip(keys, result):
if value is not None:
continue
if (
key not in self.last_redis_batch_access_time
or current_time - self.last_redis_batch_access_time[key]
>= self.redis_batch_cache_expiry
):
sublist_keys.append(key)
return sublist_keys
previous_access_times[key] = self.last_redis_batch_access_time.get(
key
)
self.last_redis_batch_access_time[key] = current_time
return sublist_keys, previous_access_times
def _rollback_redis_batch_key_reservations(
self, previous_access_times: Dict[str, Optional[float]]
) -> None:
with self._last_redis_batch_access_time_lock:
for key, previous_time in previous_access_times.items():
if previous_time is None:
self.last_redis_batch_access_time.pop(key, None)
else:
self.last_redis_batch_access_time[key] = previous_time
async def async_batch_get_cache(
self,
@ -276,19 +302,23 @@ class DualCache(BaseCache):
- check the redis cache
"""
current_time = time.time()
sublist_keys = self.get_redis_batch_keys(current_time, keys, result)
sublist_keys, previous_access_times = self._reserve_redis_batch_keys(
current_time, keys, result
)
# Only hit Redis if the last access time was more than 5 seconds ago
# Only hit Redis if enough time has passed since last access.
if len(sublist_keys) > 0:
# If not found in in-memory cache, try fetching from Redis
redis_result = await self.redis_cache.async_batch_get_cache(
sublist_keys, parent_otel_span=parent_otel_span
)
# Update the last access time for ALL queried keys
# This includes keys with None values to throttle repeated Redis queries
for key in sublist_keys:
self.last_redis_batch_access_time[key] = current_time
try:
# If not found in in-memory cache, try fetching from Redis
redis_result = await self.redis_cache.async_batch_get_cache(
sublist_keys, parent_otel_span=parent_otel_span
)
except Exception:
# Do not throttle subsequent callers if the Redis read fails.
self._rollback_redis_batch_key_reservations(
previous_access_times
)
raise
# Short-circuit if redis_result is None or contains only None values
if redis_result is None or all(v is None for v in redis_result.values()):

View file

@ -227,6 +227,84 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return input_items, instructions
def _map_optional_params_to_responses_api_request(
self,
optional_params: dict,
responses_api_request: "ResponsesAPIOptionalRequestParams",
) -> None:
"""Map optional_params into responses_api_request (mutates in place)."""
for key, value in optional_params.items():
if value is None:
continue
if key in ("max_tokens", "max_completion_tokens"):
responses_api_request["max_output_tokens"] = value
elif key == "tools" and value is not None:
responses_api_request["tools"] = (
self._convert_tools_to_responses_format(
cast(List[Dict[str, Any]], value)
)
)
elif key == "response_format":
text_format = self._transform_response_format_to_text_format(value)
if text_format:
responses_api_request["text"] = text_format # type: ignore
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
responses_api_request[key] = value # type: ignore
elif key == "previous_response_id":
responses_api_request["previous_response_id"] = value
elif key == "reasoning_effort":
responses_api_request["reasoning"] = self._map_reasoning_effort(value)
elif key == "web_search_options":
self._add_web_search_tool(responses_api_request, value)
def _build_sanitized_litellm_params(
self, litellm_params: dict
) -> Dict[str, Any]:
"""Build sanitized litellm_params with merged metadata."""
responses_optional_param_keys = set(
ResponsesAPIOptionalRequestParams.__annotations__.keys()
)
sanitized: Dict[str, Any] = {
key: value
for key, value in litellm_params.items()
if key not in responses_optional_param_keys
}
legacy_metadata = litellm_params.get("metadata")
existing_litellm_metadata = litellm_params.get("litellm_metadata")
merged_litellm_metadata: Dict[str, Any] = {}
if isinstance(legacy_metadata, dict):
merged_litellm_metadata.update(legacy_metadata)
if isinstance(existing_litellm_metadata, dict):
merged_litellm_metadata.update(existing_litellm_metadata)
if merged_litellm_metadata:
sanitized["litellm_metadata"] = merged_litellm_metadata
else:
sanitized.pop("litellm_metadata", None)
return sanitized
def _merge_responses_api_request_into_request_data(
self,
request_data: Dict[str, Any],
responses_api_request: "ResponsesAPIOptionalRequestParams",
instructions: Optional[str],
) -> None:
"""Add non-None values from responses_api_request into request_data."""
for key, value in responses_api_request.items():
if value is None:
continue
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" and isinstance(value, str):
# OpenAI API requires user param to be max 64 chars - truncate if longer
if len(value) <= 64:
request_data["user"] = value
else:
request_data["user"] = value[:64]
else:
request_data[key] = value
def transform_request(
self,
model: str,
@ -251,36 +329,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
if instructions:
responses_api_request["instructions"] = instructions
# Map optional parameters
for key, value in optional_params.items():
if value is None:
continue
if key in ("max_tokens", "max_completion_tokens"):
responses_api_request["max_output_tokens"] = value
elif key == "tools" and value is not None:
# Convert chat completion tools to responses API tools format
responses_api_request["tools"] = (
self._convert_tools_to_responses_format(
cast(List[Dict[str, Any]], value)
)
)
elif key == "response_format":
# Convert response_format to text.format
text_format = self._transform_response_format_to_text_format(value)
if text_format:
responses_api_request["text"] = text_format # type: ignore
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
responses_api_request[key] = value # type: ignore
elif key == "metadata":
responses_api_request["metadata"] = value
elif key == "previous_response_id":
responses_api_request["previous_response_id"] = value
elif key == "reasoning_effort":
responses_api_request["reasoning"] = self._map_reasoning_effort(value)
elif key == "web_search_options":
self._add_web_search_tool(responses_api_request, value)
self._map_optional_params_to_responses_api_request(
optional_params, responses_api_request
)
# Get stream parameter from litellm_params if not in optional_params
stream = optional_params.get("stream") or litellm_params.get("stream", False)
verbose_logger.debug(f"Chat provider: Stream parameter: {stream}")
@ -304,11 +356,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
setattr(litellm_logging_obj, "call_type", CallTypes.responses.value)
sanitized_litellm_params = self._build_sanitized_litellm_params(
litellm_params
)
request_data = {
"model": api_model,
"input": input_items,
"litellm_logging_obj": litellm_logging_obj,
**litellm_params,
**sanitized_litellm_params,
"client": client,
}
@ -316,18 +372,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
f"Chat provider: Final request model={api_model}, input_items={len(input_items)}"
)
# Add non-None values from responses_api_request
for key, value in responses_api_request.items():
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
self._merge_responses_api_request_into_request_data(
request_data, responses_api_request, instructions
)
if headers:
request_data["extra_headers"] = headers

View file

@ -101,6 +101,11 @@ MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE = int(
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL = int(
os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")
)
# Default npm cache directory for STDIO MCP servers.
# npm/npx needs a writable cache dir; in containers the default (~/.npm)
# may not exist or be read-only. /tmp is always writable.
MCP_NPM_CACHE_DIR = os.getenv("MCP_NPM_CACHE_DIR", "/tmp/.npm_mcp_cache")
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL = int(
os.getenv("MCP_OAUTH2_TOKEN_CACHE_MIN_TTL", "10")
)
@ -1011,10 +1016,12 @@ BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[
BEDROCK_CONVERSE_MODELS = [
"qwen.qwen3-coder-480b-a35b-v1:0",
"qwen.qwen3-coder-next",
"qwen.qwen3-235b-a22b-2507-v1:0",
"qwen.qwen3-coder-30b-a3b-v1:0",
"qwen.qwen3-32b-v1:0",
"deepseek.v3-v1:0",
"deepseek.v3.2",
"openai.gpt-oss-20b-1:0",
"openai.gpt-oss-120b-1:0",
"anthropic.claude-haiku-4-5-20251001-v1:0",
@ -1057,6 +1064,8 @@ BEDROCK_CONVERSE_MODELS = [
"amazon.nova-pro-v1:0",
"writer.palmyra-x4-v1:0",
"writer.palmyra-x5-v1:0",
"minimax.minimax-m2.1",
"moonshotai.kimi-k2.5",
]

View file

@ -141,7 +141,7 @@ class CBFTransformer:
# Required CBF fields
'time/usage_start': usage_date.isoformat() if usage_date else None, # Required: ISO-formatted UTC datetime
'cost/cost': float(row.get('spend', 0.0)), # Required: billed cost
'resource/id': model, # Send model name
'resource/id': resource_id, # CZRN (CloudZero Resource Name)
# Usage metrics for token consumption
'usage/amount': total_tokens, # Numeric value of tokens consumed

View file

@ -29,6 +29,7 @@ from litellm.types.utils import (
LLMResponseTypes,
StandardLoggingGuardrailInformation,
)
from fastapi.exceptions import HTTPException
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -624,7 +625,9 @@ class CustomGuardrail(CustomLogger):
This gets logged on downsteam Langfuse, DataDog, etc.
"""
# Convert None to empty dict to satisfy type requirements
guardrail_response = {} if response is None else response
guardrail_response: Union[Dict[str, Any], str] = (
{} if response is None else response
)
# For apply_guardrail functions in custom_code_guardrail scenario,
# simplify the logged response to "allow", "deny", or "mask"
@ -648,6 +651,23 @@ class CustomGuardrail(CustomLogger):
)
return response
@staticmethod
def _is_guardrail_intervention(e: Exception) -> bool:
"""
Returns True if the exception represents an intentional guardrail block
(this was logged previously as an API failure - guardrail_failed_to_respond).
Guardrails signal intentional blocks by raising:
- HTTPException with status 400 (content policy violation)
- ModifyResponseException (passthrough mode violation)
"""
if isinstance(e, ModifyResponseException):
return True
if isinstance(e, HTTPException) and e.status_code == 400:
return True
return False
def _process_error(
self,
e: Exception,
@ -662,6 +682,11 @@ class CustomGuardrail(CustomLogger):
This gets logged on downsteam Langfuse, DataDog, etc.
"""
guardrail_status: GuardrailStatus = (
"guardrail_intervened"
if self._is_guardrail_intervention(e)
else "guardrail_failed_to_respond"
)
# For custom_code_guardrail scenario, log as "deny" instead of full exception
# Check if this is from custom_code_guardrail by checking the class name
guardrail_response: Union[Exception, str] = e
@ -671,7 +696,7 @@ class CustomGuardrail(CustomLogger):
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=guardrail_response,
request_data=request_data,
guardrail_status="guardrail_failed_to_respond",
guardrail_status=guardrail_status,
duration=duration,
start_time=start_time,
end_time=end_time,

View file

@ -774,15 +774,17 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
self, model_call_details: Dict
) -> Dict:
"""
Only redacts messages and responses when self.turn_off_message_logging is True
Redacts or excludes fields from StandardLoggingPayload before callbacks receive it.
This method handles two features:
1. turn_off_message_logging: When True, redacts messages and responses
2. standard_logging_payload_excluded_fields: Removes specified fields entirely
By default, self.turn_off_message_logging is False and this does nothing.
Return a redacted deepcopy of the provided logging payload.
Return a modified copy of the provided logging payload.
This is useful for logging payloads that contain sensitive information.
"""
import litellm
from copy import copy
from litellm import Choices, Message, ModelResponse
@ -790,14 +792,17 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
turn_off_message_logging: bool = getattr(
self, "turn_off_message_logging", False
)
excluded_fields: Optional[List[str]] = getattr(
litellm, "standard_logging_payload_excluded_fields", None
)
if turn_off_message_logging is False:
# Early return if no processing needed
if turn_off_message_logging is False and not excluded_fields:
return model_call_details
# Only make a shallow copy of the top-level dict to avoid deepcopy issues
# with complex objects like AuthenticationError that may be present
model_call_details_copy = copy(model_call_details)
redacted_str = "redacted-by-litellm"
standard_logging_object = model_call_details.get("standard_logging_object")
if standard_logging_object is None:
return model_call_details_copy
@ -805,39 +810,58 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
# Make a copy of just the standard_logging_object to avoid modifying the original
standard_logging_object_copy = copy(standard_logging_object)
if standard_logging_object_copy.get("messages") is not None:
standard_logging_object_copy["messages"] = [
Message(content=redacted_str).model_dump()
]
# Handle excluded fields - remove them entirely from the payload
if excluded_fields:
for field in excluded_fields:
if field in standard_logging_object_copy:
del standard_logging_object_copy[field]
if standard_logging_object_copy.get("response") is not None:
response = standard_logging_object_copy["response"]
# Check if this is a ResponsesAPIResponse (has "output" field)
if isinstance(response, dict) and "output" in response:
# Make a copy to avoid modifying the original
from copy import deepcopy
# Handle turn_off_message_logging - redact messages and responses (if not already excluded)
if turn_off_message_logging:
redacted_str = "redacted-by-litellm"
response_copy = deepcopy(response)
# Redact content in output array
if isinstance(response_copy.get("output"), list):
for output_item in response_copy["output"]:
if isinstance(output_item, dict) and "content" in output_item:
if isinstance(output_item["content"], list):
# Redact text in content items
for content_item in output_item["content"]:
if (
isinstance(content_item, dict)
and "text" in content_item
):
content_item["text"] = redacted_str
standard_logging_object_copy["response"] = response_copy
else:
# Standard ModelResponse format
model_response = ModelResponse(
choices=[Choices(message=Message(content=redacted_str))]
)
model_response_dict = model_response.model_dump()
standard_logging_object_copy["response"] = model_response_dict
if (
"messages" not in (excluded_fields or [])
and standard_logging_object_copy.get("messages") is not None
):
standard_logging_object_copy["messages"] = [
Message(content=redacted_str).model_dump()
]
if (
"response" not in (excluded_fields or [])
and standard_logging_object_copy.get("response") is not None
):
response = standard_logging_object_copy["response"]
# Check if this is a ResponsesAPIResponse (has "output" field)
if isinstance(response, dict) and "output" in response:
# Make a copy to avoid modifying the original
from copy import deepcopy
response_copy = deepcopy(response)
# Redact content in output array
if isinstance(response_copy.get("output"), list):
for output_item in response_copy["output"]:
if (
isinstance(output_item, dict)
and "content" in output_item
):
if isinstance(output_item["content"], list):
# Redact text in content items
for content_item in output_item["content"]:
if (
isinstance(content_item, dict)
and "text" in content_item
):
content_item["text"] = redacted_str
standard_logging_object_copy["response"] = response_copy
else:
# Standard ModelResponse format
model_response = ModelResponse(
choices=[Choices(message=Message(content=redacted_str))]
)
model_response_dict = model_response.model_dump()
standard_logging_object_copy["response"] = model_response_dict
model_call_details_copy["standard_logging_object"] = (
standard_logging_object_copy

View file

@ -70,6 +70,11 @@ class ExceptionCheckers:
Check if an error string indicates a context window exceeded error.
"""
_error_str_lowercase = error_str.lower()
# Exclude param validation errors (e.g. OpenAI "user" param max 64 chars)
if "string_above_max_length" in _error_str_lowercase:
return False
if "invalid 'user'" in _error_str_lowercase and "string too long" in _error_str_lowercase:
return False
known_exception_substrings = [
"exceed context limit",
"this model's maximum context length is",
@ -98,16 +103,18 @@ class ExceptionCheckers:
"""
Check if an error string indicates a content policy violation error.
"""
_lower = error_str.lower()
known_exception_substrings = [
"invalid_request_error",
"content_policy_violation",
"responsibleaipolicyviolation",
"the response was filtered due to the prompt triggering azure openai's content management",
"your task failed as a result of our safety system",
"the model produced invalid content",
"content_filter_policy",
"your request was rejected as a result of our safety system",
]
for substring in known_exception_substrings:
if substring in error_str.lower():
if substring in _lower:
return True
return False
@ -2060,6 +2067,19 @@ def exception_type( # type: ignore # noqa: PLR0915
if isinstance(body_dict, dict):
if isinstance(body_dict.get("error"), dict):
azure_error_code = body_dict["error"].get("code") # type: ignore[index]
# Also check inner_error for
# ResponsibleAIPolicyViolation which indicates a
# content policy violation even when the top-level
# code is generic (e.g. "invalid_request_error").
if azure_error_code != "content_policy_violation":
_inner = (
body_dict["error"].get("inner_error") # type: ignore[index]
or body_dict["error"].get("innererror") # type: ignore[index]
)
if isinstance(_inner, dict) and _inner.get(
"code"
) == "ResponsibleAIPolicyViolation":
azure_error_code = "content_policy_violation"
else:
azure_error_code = body_dict.get("code")
except Exception:

View file

@ -51,7 +51,7 @@ def handle_cohere_chat_model_custom_llm_provider(
if custom_llm_provider == "cohere" and model in litellm.cohere_chat_models:
return model, "cohere_chat"
if "/" in model:
if model and "/" in model:
_custom_llm_provider, _model = model.split("/", 1)
if (
_custom_llm_provider
@ -84,7 +84,7 @@ def handle_anthropic_text_model_custom_llm_provider(
):
return model, "anthropic_text"
if "/" in model:
if model and "/" in model:
_custom_llm_provider, _model = model.split("/", 1)
if (
_custom_llm_provider
@ -113,6 +113,12 @@ def get_llm_provider( # noqa: PLR0915
Return model, custom_llm_provider, dynamic_api_key, api_base
"""
try:
# Early validation - model is required
if model is None:
raise ValueError(
"model parameter is required but was None. Please provide a valid model name."
)
if litellm.LiteLLMProxyChatConfig._should_use_litellm_proxy_by_default(
litellm_params=litellm_params
):

View file

@ -2331,7 +2331,7 @@ class Logging(LiteLLMLoggingBaseClass):
result, LiteLLMBatch
):
litellm_params = self.litellm_params or {}
litellm_metadata = litellm_params.get("litellm_metadata", {})
litellm_metadata = litellm_params.get("litellm_metadata") or {}
if (
litellm_metadata.get("batch_ignore_default_logging", False) is True
): # polling job will query these frequently, don't spam db logs
@ -2369,6 +2369,7 @@ class Logging(LiteLLMLoggingBaseClass):
) = await _handle_completed_batch(
batch=result,
custom_llm_provider=self.custom_llm_provider,
litellm_params=self.litellm_params,
)
result._hidden_params["response_cost"] = response_cost
@ -3127,7 +3128,7 @@ class Logging(LiteLLMLoggingBaseClass):
self, dynamic_success_callbacks: Optional[List], global_callbacks: List
) -> List:
if dynamic_success_callbacks is None:
return global_callbacks
return list(global_callbacks)
return list(set(dynamic_success_callbacks + global_callbacks))
def _remove_internal_litellm_callbacks(self, callbacks: List) -> List:

View file

@ -664,35 +664,34 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
reasoning_effort: Optional[Union[REASONING_EFFORT, str]],
model: str,
) -> Optional[AnthropicThinkingParam]:
if reasoning_effort is None or reasoning_effort == "none":
return None
if AnthropicConfig._is_claude_opus_4_6(model):
return AnthropicThinkingParam(
type="adaptive",
)
elif reasoning_effort == "low":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
)
elif reasoning_effort == "medium":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
)
elif reasoning_effort == "high":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
)
elif reasoning_effort == "minimal":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
)
else:
if reasoning_effort is None:
return None
elif reasoning_effort == "low":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
)
elif reasoning_effort == "medium":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
)
elif reasoning_effort == "high":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
)
elif reasoning_effort == "minimal":
return AnthropicThinkingParam(
type="enabled",
budget_tokens=DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET,
)
else:
raise ValueError(f"Unmapped reasoning effort: {reasoning_effort}")
raise ValueError(f"Unmapped reasoning effort: {reasoning_effort}")
def _extract_json_schema_from_response_format(
self, value: Optional[dict]

View file

@ -38,9 +38,18 @@ def optionally_handle_anthropic_oauth(
Returns:
Tuple of (updated headers, api_key)
"""
# Check Authorization header (passthrough / forwarded requests)
auth_header = headers.get("authorization", "")
if auth_header and auth_header.startswith(f"Bearer {ANTHROPIC_OAUTH_TOKEN_PREFIX}"):
api_key = auth_header.replace("Bearer ", "")
headers.pop("x-api-key", None)
headers["anthropic-beta"] = ANTHROPIC_OAUTH_BETA_HEADER
headers["anthropic-dangerous-direct-browser-access"] = "true"
return headers, api_key
# Check api_key directly (standard chat/completion flow)
if api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX):
headers.pop("x-api-key", None)
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-beta"] = ANTHROPIC_OAUTH_BETA_HEADER
headers["anthropic-dangerous-direct-browser-access"] = "true"
return headers, api_key
@ -108,7 +117,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
if tools is None:
return False
for tool in tools:
if "type" in tool and tool["type"].startswith(ANTHROPIC_HOSTED_TOOLS.WEB_SEARCH.value):
if "type" in tool and tool["type"].startswith(
ANTHROPIC_HOSTED_TOOLS.WEB_SEARCH.value
):
return True
return False
@ -134,111 +145,126 @@ class AnthropicModelInfo(BaseLLMModelInfo):
"""
if not tools:
return False
for tool in tools:
tool_type = tool.get("type", "")
if tool_type in ["tool_search_tool_regex_20251119", "tool_search_tool_bm25_20251119"]:
if tool_type in [
"tool_search_tool_regex_20251119",
"tool_search_tool_bm25_20251119",
]:
return True
return False
def is_programmatic_tool_calling_used(self, tools: Optional[List]) -> bool:
"""
Check if programmatic tool calling is being used (tools with allowed_callers field).
Returns True if any tool has allowed_callers containing 'code_execution_20250825'.
"""
if not tools:
return False
for tool in tools:
# Check top-level allowed_callers
allowed_callers = tool.get("allowed_callers", None)
if allowed_callers and isinstance(allowed_callers, list):
if "code_execution_20250825" in allowed_callers:
return True
# Check function.allowed_callers for OpenAI format tools
function = tool.get("function", {})
if isinstance(function, dict):
function_allowed_callers = function.get("allowed_callers", None)
if function_allowed_callers and isinstance(function_allowed_callers, list):
if function_allowed_callers and isinstance(
function_allowed_callers, list
):
if "code_execution_20250825" in function_allowed_callers:
return True
return False
def is_input_examples_used(self, tools: Optional[List]) -> bool:
"""
Check if input_examples is being used in any tools.
Returns True if any tool has input_examples field.
"""
if not tools:
return False
for tool in tools:
# Check top-level input_examples
input_examples = tool.get("input_examples", None)
if input_examples and isinstance(input_examples, list) and len(input_examples) > 0:
if (
input_examples
and isinstance(input_examples, list)
and len(input_examples) > 0
):
return True
# Check function.input_examples for OpenAI format tools
function = tool.get("function", {})
if isinstance(function, dict):
function_input_examples = function.get("input_examples", None)
if function_input_examples and isinstance(function_input_examples, list) and len(function_input_examples) > 0:
if (
function_input_examples
and isinstance(function_input_examples, list)
and len(function_input_examples) > 0
):
return True
return False
def is_effort_used(self, optional_params: Optional[dict], model: Optional[str] = None) -> bool:
def is_effort_used(
self, optional_params: Optional[dict], model: Optional[str] = None
) -> bool:
"""
Check if effort parameter is being used.
Returns True if effort-related parameters are present.
"""
if not optional_params:
return False
# Check if reasoning_effort is provided for Claude Opus 4.5
if model and ("opus-4-5" in model.lower() or "opus_4_5" in model.lower()):
reasoning_effort = optional_params.get("reasoning_effort")
if reasoning_effort and isinstance(reasoning_effort, str):
return True
# Check if output_config is directly provided
output_config = optional_params.get("output_config")
if output_config and isinstance(output_config, dict):
effort = output_config.get("effort")
if effort and isinstance(effort, str):
return True
return False
def is_code_execution_tool_used(self, tools: Optional[List]) -> bool:
"""
Check if code execution tool is being used.
Returns True if any tool has type "code_execution_20250825".
"""
if not tools:
return False
for tool in tools:
tool_type = tool.get("type", "")
if tool_type == "code_execution_20250825":
return True
return False
def is_container_with_skills_used(self, optional_params: Optional[dict]) -> bool:
"""
Check if container with skills is being used.
Returns True if optional_params contains container with skills.
"""
if not optional_params:
return False
container = optional_params.get("container")
if container and isinstance(container, dict):
skills = container.get("skills")
@ -256,10 +282,10 @@ class AnthropicModelInfo(BaseLLMModelInfo):
def get_computer_tool_beta_header(self, computer_tool_version: str) -> str:
"""
Get the appropriate beta header for a given computer tool version.
Args:
computer_tool_version: The computer tool version (e.g., 'computer_20250124', 'computer_20241022')
Returns:
The corresponding beta header string
"""
@ -282,37 +308,37 @@ class AnthropicModelInfo(BaseLLMModelInfo):
) -> List[str]:
"""
Get list of common beta headers based on the features that are active.
Returns:
List of beta header strings
"""
from litellm.types.llms.anthropic import (
ANTHROPIC_EFFORT_BETA_HEADER,
)
betas = []
# Detect features
effort_used = self.is_effort_used(optional_params, model)
if effort_used:
betas.append(ANTHROPIC_EFFORT_BETA_HEADER) # effort-2025-11-24
if computer_tool_used:
beta_header = self.get_computer_tool_beta_header(computer_tool_used)
betas.append(beta_header)
# Anthropic no longer requires the prompt-caching beta header
# Prompt caching now works automatically when cache_control is used in messages
# Reference: https://docs.anthropic.com/en/docs/build-with-claude/prompt-caching
if file_id_used:
betas.append("files-api-2025-04-14")
betas.append("code-execution-2025-05-22")
if mcp_server_used:
betas.append("mcp-client-2025-04-04")
return list(set(betas))
def get_anthropic_headers(
@ -351,27 +377,35 @@ class AnthropicModelInfo(BaseLLMModelInfo):
# Tool search, programmatic tool calling, and input_examples all use the same beta header
if tool_search_used or programmatic_tool_calling_used or input_examples_used:
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
betas.add(ANTHROPIC_TOOL_SEARCH_BETA_HEADER)
# Effort parameter uses a separate beta header
if effort_used:
from litellm.types.llms.anthropic import ANTHROPIC_EFFORT_BETA_HEADER
betas.add(ANTHROPIC_EFFORT_BETA_HEADER)
# Code execution tool uses a separate beta header
if code_execution_tool_used:
betas.add("code-execution-2025-08-25")
# Container with skills uses a separate beta header
if container_with_skills_used:
betas.add("skills-2025-10-02")
_is_oauth = api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX)
headers = {
"anthropic-version": anthropic_version or "2023-06-01",
"x-api-key": api_key,
"accept": "application/json",
"content-type": "application/json",
}
if _is_oauth:
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-dangerous-direct-browser-access"] = "true"
betas.add(ANTHROPIC_OAUTH_BETA_HEADER)
else:
headers["x-api-key"] = api_key
if user_anthropic_beta_headers is not None:
betas.update(user_anthropic_beta_headers)
@ -381,7 +415,10 @@ class AnthropicModelInfo(BaseLLMModelInfo):
# Vertex AI requires web search beta header for web search to work
if web_search_tool_used:
from litellm.types.llms.anthropic import ANTHROPIC_BETA_HEADER_VALUES
headers["anthropic-beta"] = ANTHROPIC_BETA_HEADER_VALUES.WEB_SEARCH_2025_03_05.value
headers[
"anthropic-beta"
] = ANTHROPIC_BETA_HEADER_VALUES.WEB_SEARCH_2025_03_05.value
elif len(betas) > 0:
headers["anthropic-beta"] = ",".join(betas)
@ -398,7 +435,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
api_base: Optional[str] = None,
) -> Dict:
# Check for Anthropic OAuth token in headers
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
headers, api_key = optionally_handle_anthropic_oauth(
headers=headers, api_key=api_key
)
if api_key is None:
raise litellm.AuthenticationError(
message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` in your environment vars",
@ -416,11 +455,15 @@ class AnthropicModelInfo(BaseLLMModelInfo):
file_id_used = self.is_file_id_used(messages=messages)
web_search_tool_used = self.is_web_search_tool_used(tools=tools)
tool_search_used = self.is_tool_search_used(tools=tools)
programmatic_tool_calling_used = self.is_programmatic_tool_calling_used(tools=tools)
programmatic_tool_calling_used = self.is_programmatic_tool_calling_used(
tools=tools
)
input_examples_used = self.is_input_examples_used(tools=tools)
effort_used = self.is_effort_used(optional_params=optional_params, model=model)
code_execution_tool_used = self.is_code_execution_tool_used(tools=tools)
container_with_skills_used = self.is_container_with_skills_used(optional_params=optional_params)
container_with_skills_used = self.is_container_with_skills_used(
optional_params=optional_params
)
user_anthropic_beta_headers = self._get_user_anthropic_beta_headers(
anthropic_beta_header=headers.get("anthropic-beta")
)
@ -499,7 +542,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
def get_token_counter(self) -> Optional[BaseTokenCounter]:
"""
Factory method to create an Anthropic token counter.
Returns:
AnthropicTokenCounter instance for this provider.
"""

View file

@ -49,15 +49,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
# TODO: Add Anthropic `metadata` support
# "metadata",
]
@staticmethod
def _filter_billing_headers_from_system(system_param):
"""
Filter out x-anthropic-billing-header metadata from system parameter.
Args:
system_param: Can be a string or a list of system message content blocks
Returns:
Filtered system parameter (string or list), or None if all content was filtered
"""
@ -74,7 +74,9 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
text = content_block.get("text", "")
content_type = content_block.get("type", "")
# Skip text blocks that start with billing header
if content_type == "text" and text.startswith("x-anthropic-billing-header:"):
if content_type == "text" and text.startswith(
"x-anthropic-billing-header:"
):
continue
filtered_list.append(content_block)
else:
@ -111,11 +113,13 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
import os
# Check for Anthropic OAuth token in Authorization header
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
headers, api_key = optionally_handle_anthropic_oauth(
headers=headers, api_key=api_key
)
if api_key is None:
api_key = os.getenv("ANTHROPIC_API_KEY")
if "x-api-key" not in headers and api_key:
if "x-api-key" not in headers and "authorization" not in headers and api_key:
headers["x-api-key"] = api_key
if "anthropic-version" not in headers:
headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION
@ -149,7 +153,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
message="max_tokens is required for Anthropic /v1/messages API",
status_code=400,
)
# Filter out x-anthropic-billing-header from system messages
system_param = anthropic_messages_optional_request_params.get("system")
if system_param is not None:
@ -159,7 +163,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
else:
# Remove system parameter if all content was filtered out
anthropic_messages_optional_request_params.pop("system", None)
####### get required params for all anthropic messages requests ######
verbose_logger.debug(f"TRANSFORMATION DEBUG - Messages: {messages}")
anthropic_messages_request: AnthropicMessagesRequest = AnthropicMessagesRequest(
@ -244,25 +248,29 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
edits = context_management_param.get("edits", [])
has_compact = False
has_other = False
for edit in edits:
edit_type = edit.get("type", "")
if edit_type == "compact_20260112":
has_compact = True
else:
has_other = True
# Add compact header if any compact edits exist
if has_compact:
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.COMPACT_2026_01_12.value)
# Add context management header if any other edits exist
if has_other:
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value)
beta_values.add(
ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
)
# Check for structured outputs
if optional_params.get("output_format") is not None:
beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value)
beta_values.add(
ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value
)
# Check for fast mode
if optional_params.get("speed") == "fast":

View file

@ -901,7 +901,20 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
if response.json()["status"] == "failed":
error_data = response.json()
raise AzureOpenAIError(status_code=400, message=json.dumps(error_data))
# Preserve Azure error details (e.g. content_policy_violation,
# inner_error, content_filter_results) as structured body so
# exception_type() can route them correctly.
_error_body = error_data.get("error", error_data)
_error_msg = (
_error_body.get("message", "Image generation failed")
if isinstance(_error_body, dict)
else json.dumps(error_data)
)
raise AzureOpenAIError(
status_code=400,
message=_error_msg,
body=error_data,
)
result = response.json()["result"]
return httpx.Response(
@ -999,7 +1012,20 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
if response.json()["status"] == "failed":
error_data = response.json()
raise AzureOpenAIError(status_code=400, message=json.dumps(error_data))
# Preserve Azure error details (e.g. content_policy_violation,
# inner_error, content_filter_results) as structured body so
# exception_type() can route them correctly.
_error_body = error_data.get("error", error_data)
_error_msg = (
_error_body.get("message", "Image generation failed")
if isinstance(_error_body, dict)
else json.dumps(error_data)
)
raise AzureOpenAIError(
status_code=400,
message=_error_msg,
body=error_data,
)
result = response.json()["result"]
return httpx.Response(

View file

@ -0,0 +1,41 @@
"""
Managed Resources Module
This module provides base classes and utilities for managing resources
(files, vector stores, etc.) with target_model_names support.
The BaseManagedResource class provides common functionality for:
- Storing unified resource IDs with model mappings
- Retrieving resources by unified ID
- Deleting resources across multiple models
- Creating resources for multiple models
- Filtering deployments based on model mappings
"""
from .base_managed_resource import BaseManagedResource
from .utils import (
decode_unified_id,
encode_unified_id,
extract_model_id_from_unified_id,
extract_provider_resource_id_from_unified_id,
extract_resource_type_from_unified_id,
extract_target_model_names_from_unified_id,
extract_unified_uuid_from_unified_id,
generate_unified_id_string,
is_base64_encoded_unified_id,
parse_unified_id,
)
__all__ = [
"BaseManagedResource",
"is_base64_encoded_unified_id",
"extract_target_model_names_from_unified_id",
"extract_resource_type_from_unified_id",
"extract_unified_uuid_from_unified_id",
"extract_model_id_from_unified_id",
"extract_provider_resource_id_from_unified_id",
"generate_unified_id_string",
"encode_unified_id",
"decode_unified_id",
"parse_unified_id",
]

View file

@ -0,0 +1,605 @@
# What is this?
## Base class for managing resources (files, vector stores, etc.) with target_model_names support
## This provides common functionality for creating, retrieving, and managing resources across multiple models
import base64
import json
from abc import ABC, abstractmethod
from typing import (
TYPE_CHECKING,
Any,
Dict,
Generic,
List,
Optional,
TypeVar,
Union,
cast,
)
from litellm import verbose_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import SpecialEnums
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
from litellm.proxy.utils import PrismaClient as _PrismaClient
from litellm.router import Router as _Router
Span = Union[_Span, Any]
InternalUsageCache = _InternalUsageCache
PrismaClient = _PrismaClient
Router = _Router
else:
Span = Any
InternalUsageCache = Any
PrismaClient = Any
Router = Any
# Generic type for resource objects
ResourceObjectType = TypeVar('ResourceObjectType')
class BaseManagedResource(ABC, Generic[ResourceObjectType]):
"""
Base class for managing resources with target_model_names support.
This class provides common functionality for:
- Storing unified resource IDs with model mappings
- Retrieving resources by unified ID
- Deleting resources across multiple models
- Creating resources for multiple models
- Filtering deployments based on model mappings
Subclasses should implement:
- resource_type: str property
- table_name: str property
- create_resource_for_model: method to create resource on a specific model
- get_unified_resource_id_format: method to generate unified ID format
"""
def __init__(
self,
internal_usage_cache: InternalUsageCache,
prisma_client: PrismaClient,
):
self.internal_usage_cache = internal_usage_cache
self.prisma_client = prisma_client
# ============================================================================
# ABSTRACT METHODS
# ============================================================================
@property
@abstractmethod
def resource_type(self) -> str:
"""
Return the resource type identifier (e.g., 'file', 'vector_store', 'vector_store_file').
Used for logging and unified ID generation.
"""
pass
@property
@abstractmethod
def table_name(self) -> str:
"""
Return the database table name for this resource type.
Example: 'litellm_managedfiletable', 'litellm_managedvectorstoretable'
"""
pass
@abstractmethod
def get_unified_resource_id_format(
self,
resource_object: ResourceObjectType,
target_model_names_list: List[str],
) -> str:
"""
Generate the format string for the unified resource ID.
This should return a string that will be base64 encoded.
Example for files:
"litellm_proxy:application/json;unified_id,{uuid};target_model_names,{models};..."
Args:
resource_object: The resource object returned from the provider
target_model_names_list: List of target model names
Returns:
Format string to be base64 encoded
"""
pass
@abstractmethod
async def create_resource_for_model(
self,
llm_router: Router,
model: str,
request_data: Dict[str, Any],
litellm_parent_otel_span: Span,
) -> ResourceObjectType:
"""
Create a resource for a specific model.
Args:
llm_router: LiteLLM router instance
model: Model name to create resource for
request_data: Request data for resource creation
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
Resource object from the provider
"""
pass
# ============================================================================
# COMMON STORAGE OPERATIONS
# ============================================================================
async def store_unified_resource_id(
self,
unified_resource_id: str,
resource_object: Optional[ResourceObjectType],
litellm_parent_otel_span: Optional[Span],
model_mappings: Dict[str, str],
user_api_key_dict: UserAPIKeyAuth,
additional_db_fields: Optional[Dict[str, Any]] = None,
) -> None:
"""
Store unified resource ID with model mappings in cache and database.
Args:
unified_resource_id: The unified resource ID (base64 encoded)
resource_object: The resource object to store (can be None)
litellm_parent_otel_span: OpenTelemetry span for tracing
model_mappings: Dictionary mapping model_id -> provider_resource_id
user_api_key_dict: User API key authentication details
additional_db_fields: Additional fields to store in database
"""
verbose_logger.info(
f"Storing LiteLLM Managed {self.resource_type} with id={unified_resource_id} in cache"
)
# Prepare cache data
cache_data = {
"unified_resource_id": unified_resource_id,
"resource_object": resource_object,
"model_mappings": model_mappings,
"flat_model_resource_ids": list(model_mappings.values()),
"created_by": user_api_key_dict.user_id,
"updated_by": user_api_key_dict.user_id,
}
# Add additional fields if provided
if additional_db_fields:
cache_data.update(additional_db_fields)
# Store in cache
if resource_object is not None:
await self.internal_usage_cache.async_set_cache(
key=unified_resource_id,
value=cache_data,
litellm_parent_otel_span=litellm_parent_otel_span,
)
# Prepare database data
db_data = {
"unified_resource_id": unified_resource_id,
"model_mappings": json.dumps(model_mappings),
"flat_model_resource_ids": list(model_mappings.values()),
"created_by": user_api_key_dict.user_id,
"updated_by": user_api_key_dict.user_id,
}
# Add resource object if available
if resource_object is not None:
# Handle both dict and Pydantic models
if hasattr(resource_object, "model_dump_json"):
db_data["resource_object"] = resource_object.model_dump_json() # type: ignore
elif isinstance(resource_object, dict):
db_data["resource_object"] = json.dumps(resource_object)
# Extract storage metadata from hidden params if present
hidden_params = getattr(resource_object, "_hidden_params", {}) or {}
if "storage_backend" in hidden_params:
db_data["storage_backend"] = hidden_params["storage_backend"]
if "storage_url" in hidden_params:
db_data["storage_url"] = hidden_params["storage_url"]
# Add additional fields to database
if additional_db_fields:
db_data.update(additional_db_fields)
# Store in database
table = getattr(self.prisma_client.db, self.table_name)
result = await table.create(data=db_data)
verbose_logger.debug(
f"LiteLLM Managed {self.resource_type} with id={unified_resource_id} stored in db: {result}"
)
async def get_unified_resource_id(
self,
unified_resource_id: str,
litellm_parent_otel_span: Optional[Span] = None,
) -> Optional[Dict[str, Any]]:
"""
Retrieve unified resource by ID from cache or database.
Args:
unified_resource_id: The unified resource ID to retrieve
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
Dictionary containing resource data or None if not found
"""
# Check cache first
result = cast(
Optional[dict],
await self.internal_usage_cache.async_get_cache(
key=unified_resource_id,
litellm_parent_otel_span=litellm_parent_otel_span,
),
)
if result:
return result
# Check database
table = getattr(self.prisma_client.db, self.table_name)
db_object = await table.find_first(
where={"unified_resource_id": unified_resource_id}
)
if db_object:
return db_object.model_dump()
return None
async def delete_unified_resource_id(
self,
unified_resource_id: str,
litellm_parent_otel_span: Optional[Span] = None,
) -> Optional[ResourceObjectType]:
"""
Delete unified resource from cache and database.
Args:
unified_resource_id: The unified resource ID to delete
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
The deleted resource object or None if not found
"""
# Get old value from database
table = getattr(self.prisma_client.db, self.table_name)
initial_value = await table.find_first(
where={"unified_resource_id": unified_resource_id}
)
if initial_value is None:
raise Exception(
f"LiteLLM Managed {self.resource_type} with id={unified_resource_id} not found"
)
# Delete from cache
await self.internal_usage_cache.async_set_cache(
key=unified_resource_id,
value=None,
litellm_parent_otel_span=litellm_parent_otel_span,
)
# Delete from database
await table.delete(where={"unified_resource_id": unified_resource_id})
return initial_value.resource_object
async def can_user_access_unified_resource_id(
self,
unified_resource_id: str,
user_api_key_dict: UserAPIKeyAuth,
litellm_parent_otel_span: Optional[Span] = None,
) -> bool:
"""
Check if user has access to the unified resource ID.
Uses get_unified_resource_id() which checks cache first before hitting the database,
avoiding direct DB queries in the critical request path.
Args:
unified_resource_id: The unified resource ID to check
user_api_key_dict: User API key authentication details
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
True if user has access, False otherwise
"""
user_id = user_api_key_dict.user_id
# Use cached method instead of direct DB query
resource = await self.get_unified_resource_id(
unified_resource_id, litellm_parent_otel_span
)
if resource:
return resource.get("created_by") == user_id
return False
# ============================================================================
# MODEL MAPPING OPERATIONS
# ============================================================================
async def get_model_resource_id_mapping(
self,
resource_ids: List[str],
litellm_parent_otel_span: Span,
) -> Dict[str, Dict[str, str]]:
"""
Get model-specific resource IDs for a list of unified resource IDs.
Args:
resource_ids: List of unified resource IDs
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
Dictionary mapping unified_resource_id -> model_id -> provider_resource_id
Example:
{
"unified_resource_id_1": {
"model_id_1": "provider_resource_id_1",
"model_id_2": "provider_resource_id_2"
}
}
"""
resource_id_mapping: Dict[str, Dict[str, str]] = {}
for resource_id in resource_ids:
# Get unified resource from cache/db
unified_resource_object = await self.get_unified_resource_id(
resource_id, litellm_parent_otel_span
)
if unified_resource_object:
model_mappings = unified_resource_object.get("model_mappings", {})
# Handle both JSON string and dict
if isinstance(model_mappings, str):
model_mappings = json.loads(model_mappings)
resource_id_mapping[resource_id] = model_mappings
return resource_id_mapping
# ============================================================================
# RESOURCE CREATION OPERATIONS
# ============================================================================
async def create_resource_for_each_model(
self,
llm_router: Router,
request_data: Dict[str, Any],
target_model_names_list: List[str],
litellm_parent_otel_span: Span,
) -> List[ResourceObjectType]:
"""
Create a resource for each model in the target list.
Args:
llm_router: LiteLLM router instance
request_data: Request data for resource creation
target_model_names_list: List of target model names
litellm_parent_otel_span: OpenTelemetry span for tracing
Returns:
List of resource objects created for each model
"""
if llm_router is None:
raise Exception("LLM Router not initialized. Ensure models added to proxy.")
responses = []
for model in target_model_names_list:
individual_response = await self.create_resource_for_model(
llm_router=llm_router,
model=model,
request_data=request_data,
litellm_parent_otel_span=litellm_parent_otel_span,
)
responses.append(individual_response)
return responses
def generate_unified_resource_id(
self,
resource_objects: List[ResourceObjectType],
target_model_names_list: List[str],
) -> str:
"""
Generate a unified resource ID from multiple resource objects.
Args:
resource_objects: List of resource objects from different models
target_model_names_list: List of target model names
Returns:
Base64 encoded unified resource ID
"""
# Use the first resource object to generate the format
unified_id_format = self.get_unified_resource_id_format(
resource_object=resource_objects[0],
target_model_names_list=target_model_names_list,
)
# Convert to URL-safe base64 and strip padding
base64_unified_id = (
base64.urlsafe_b64encode(unified_id_format.encode()).decode().rstrip("=")
)
return base64_unified_id
def extract_model_mappings_from_responses(
self,
resource_objects: List[ResourceObjectType],
) -> Dict[str, str]:
"""
Extract model mappings from resource objects.
Args:
resource_objects: List of resource objects from different models
Returns:
Dictionary mapping model_id -> provider_resource_id
"""
model_mappings: Dict[str, str] = {}
for resource_object in resource_objects:
# Get hidden params if available
hidden_params = getattr(resource_object, "_hidden_params", {}) or {}
model_resource_id_mapping = hidden_params.get("model_resource_id_mapping")
if model_resource_id_mapping and isinstance(model_resource_id_mapping, dict):
model_mappings.update(model_resource_id_mapping)
return model_mappings
# ============================================================================
# DEPLOYMENT FILTERING
# ============================================================================
async def async_filter_deployments(
self,
model: str,
healthy_deployments: List,
request_kwargs: Optional[Dict] = None,
parent_otel_span: Optional[Span] = None,
resource_id_key: str = "resource_id",
) -> List[Dict]:
"""
Filter deployments based on model mappings for a resource.
This is used by the router to select only deployments that have
the resource available.
Args:
model: Model name
healthy_deployments: List of healthy deployments
request_kwargs: Request kwargs containing resource_id and mappings
parent_otel_span: OpenTelemetry span for tracing
resource_id_key: Key to use for resource ID in request_kwargs
Returns:
Filtered list of deployments
"""
if request_kwargs is None:
return healthy_deployments
resource_id = cast(Optional[str], request_kwargs.get(resource_id_key))
model_resource_id_mapping = cast(
Optional[Dict[str, Dict[str, str]]],
request_kwargs.get("model_resource_id_mapping"),
)
allowed_model_ids = []
if resource_id and model_resource_id_mapping:
model_id_dict = model_resource_id_mapping.get(resource_id, {})
allowed_model_ids = list(model_id_dict.keys())
if len(allowed_model_ids) == 0:
return healthy_deployments
return [
deployment
for deployment in healthy_deployments
if deployment.get("model_info", {}).get("id") in allowed_model_ids
]
# ============================================================================
# UTILITY METHODS
# ============================================================================
def get_unified_id_prefix(self) -> str:
"""
Get the prefix for unified IDs for this resource type.
Returns:
Prefix string (e.g., "litellm_proxy:")
"""
return SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value
async def list_user_resources(
self,
user_api_key_dict: UserAPIKeyAuth,
limit: Optional[int] = None,
after: Optional[str] = None,
additional_filters: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
List resources created by a user.
Args:
user_api_key_dict: User API key authentication details
limit: Maximum number of resources to return
after: Cursor for pagination
additional_filters: Additional filters to apply
Returns:
Dictionary with list of resources and pagination info
"""
where_clause: Dict[str, Any] = {}
# Filter by user who created the resource
if user_api_key_dict.user_id:
where_clause["created_by"] = user_api_key_dict.user_id
if after:
where_clause["id"] = {"gt": after}
# Add additional filters
if additional_filters:
where_clause.update(additional_filters)
# Fetch resources
fetch_limit = limit or 20
table = getattr(self.prisma_client.db, self.table_name)
resources = await table.find_many(
where=where_clause,
take=fetch_limit,
order={"created_at": "desc"},
)
resource_objects: List[Any] = []
for resource in resources:
try:
# Stop once we have enough
if len(resource_objects) >= (limit or 20):
break
# Parse resource object
resource_data = resource.resource_object
if isinstance(resource_data, str):
resource_data = json.loads(resource_data)
# Set unified ID
if hasattr(resource_data, "id"):
resource_data.id = resource.unified_resource_id
elif isinstance(resource_data, dict):
resource_data["id"] = resource.unified_resource_id
resource_objects.append(resource_data)
except Exception as e:
verbose_logger.warning(
f"Failed to parse {self.resource_type} object "
f"{resource.unified_resource_id}: {e}"
)
continue
return {
"object": "list",
"data": resource_objects,
"first_id": resource_objects[0].id if resource_objects else None,
"last_id": resource_objects[-1].id if resource_objects else None,
"has_more": len(resource_objects) == (limit or 20),
}

View file

@ -0,0 +1,364 @@
"""
Utility functions for managed resources.
This module provides common utility functions that can be used across
different managed resource types (files, vector stores, etc.).
"""
import base64
import re
from typing import List, Optional, Union, Literal
def is_base64_encoded_unified_id(
resource_id: str,
prefix: str = "litellm_proxy:",
) -> Union[str, Literal[False]]:
"""
Check if a resource ID is a base64 encoded unified ID.
Args:
resource_id: The resource ID to check
prefix: The expected prefix for unified IDs
Returns:
Decoded string if valid unified ID, False otherwise
"""
# Ensure resource_id is a string
if not isinstance(resource_id, str):
return False
# Add padding back if needed
padded = resource_id + "=" * (-len(resource_id) % 4)
# Decode from base64
try:
decoded = base64.urlsafe_b64decode(padded).decode()
if decoded.startswith(prefix):
return decoded
else:
return False
except Exception:
return False
def extract_target_model_names_from_unified_id(
unified_id: str,
) -> List[str]:
"""
Extract target model names from a unified resource ID.
Args:
unified_id: The unified resource ID (decoded or encoded)
Returns:
List of target model names
Example:
unified_id = "litellm_proxy:vector_store;unified_id,uuid;target_model_names,gpt-4,gemini-2.0"
returns: ["gpt-4", "gemini-2.0"]
"""
try:
# Ensure unified_id is a string
if not isinstance(unified_id, str):
return []
# Decode if it's base64 encoded
decoded_id = is_base64_encoded_unified_id(unified_id)
if decoded_id:
unified_id = decoded_id
# Extract model names using regex
match = re.search(r"target_model_names,([^;]+)", unified_id)
if match:
# Split on comma and strip whitespace from each model name
return [model.strip() for model in match.group(1).split(",")]
return []
except Exception:
return []
def extract_resource_type_from_unified_id(
unified_id: str,
) -> Optional[str]:
"""
Extract resource type from a unified resource ID.
Args:
unified_id: The unified resource ID (decoded or encoded)
Returns:
Resource type string or None
Example:
unified_id = "litellm_proxy:vector_store;unified_id,uuid;..."
returns: "vector_store"
"""
try:
# Ensure unified_id is a string
if not isinstance(unified_id, str):
return None
# Decode if it's base64 encoded
decoded_id = is_base64_encoded_unified_id(unified_id)
if decoded_id:
unified_id = decoded_id
# Extract resource type (comes after prefix and before first semicolon)
match = re.search(r"litellm_proxy:([^;]+)", unified_id)
if match:
return match.group(1).strip()
return None
except Exception:
return None
def extract_unified_uuid_from_unified_id(
unified_id: str,
) -> Optional[str]:
"""
Extract the UUID from a unified resource ID.
Args:
unified_id: The unified resource ID (decoded or encoded)
Returns:
UUID string or None
Example:
unified_id = "litellm_proxy:vector_store;unified_id,abc-123;..."
returns: "abc-123"
"""
try:
# Ensure unified_id is a string
if not isinstance(unified_id, str):
return None
# Decode if it's base64 encoded
decoded_id = is_base64_encoded_unified_id(unified_id)
if decoded_id:
unified_id = decoded_id
# Extract UUID
match = re.search(r"unified_id,([^;]+)", unified_id)
if match:
return match.group(1).strip()
return None
except Exception:
return None
def extract_model_id_from_unified_id(
unified_id: str,
) -> Optional[str]:
"""
Extract model ID from a unified resource ID.
Args:
unified_id: The unified resource ID (decoded or encoded)
Returns:
Model ID string or None
Example:
unified_id = "litellm_proxy:vector_store;...;model_id,gpt-4-model-id;..."
returns: "gpt-4-model-id"
"""
try:
# Ensure unified_id is a string
if not isinstance(unified_id, str):
return None
# Decode if it's base64 encoded
decoded_id = is_base64_encoded_unified_id(unified_id)
if decoded_id:
unified_id = decoded_id
# Extract model ID
match = re.search(r"model_id,([^;]+)", unified_id)
if match:
return match.group(1).strip()
return None
except Exception:
return None
def extract_provider_resource_id_from_unified_id(
unified_id: str,
) -> Optional[str]:
"""
Extract provider resource ID from a unified resource ID.
Args:
unified_id: The unified resource ID (decoded or encoded)
Returns:
Provider resource ID string or None
Example:
unified_id = "litellm_proxy:vector_store;...;resource_id,vs_abc123;..."
returns: "vs_abc123"
"""
try:
# Ensure unified_id is a string
if not isinstance(unified_id, str):
return None
# Decode if it's base64 encoded
decoded_id = is_base64_encoded_unified_id(unified_id)
if decoded_id:
unified_id = decoded_id
# Extract resource ID (try multiple patterns for different resource types)
patterns = [
r"resource_id,([^;]+)",
r"vector_store_id,([^;]+)",
r"file_id,([^;]+)",
]
for pattern in patterns:
match = re.search(pattern, unified_id)
if match:
return match.group(1).strip()
return None
except Exception:
return None
def generate_unified_id_string(
resource_type: str,
unified_uuid: str,
target_model_names: List[str],
provider_resource_id: str,
model_id: str,
additional_fields: Optional[dict] = None,
) -> str:
"""
Generate a unified ID string (before base64 encoding).
Args:
resource_type: Type of resource (e.g., "vector_store", "file")
unified_uuid: UUID for this unified resource
target_model_names: List of target model names
provider_resource_id: Resource ID from the provider
model_id: Model ID from the router
additional_fields: Additional fields to include in the ID
Returns:
Unified ID string (not yet base64 encoded)
Example:
generate_unified_id_string(
resource_type="vector_store",
unified_uuid="abc-123",
target_model_names=["gpt-4", "gemini"],
provider_resource_id="vs_xyz",
model_id="model-id-123",
)
returns: "litellm_proxy:vector_store;unified_id,abc-123;target_model_names,gpt-4,gemini;resource_id,vs_xyz;model_id,model-id-123"
"""
# Build the unified ID string
parts = [
f"litellm_proxy:{resource_type}",
f"unified_id,{unified_uuid}",
f"target_model_names,{','.join(target_model_names)}",
f"resource_id,{provider_resource_id}",
f"model_id,{model_id}",
]
# Add additional fields if provided
if additional_fields:
for key, value in additional_fields.items():
parts.append(f"{key},{value}")
return ";".join(parts)
def encode_unified_id(unified_id_string: str) -> str:
"""
Encode a unified ID string to base64.
Args:
unified_id_string: The unified ID string to encode
Returns:
Base64 encoded unified ID (URL-safe, padding stripped)
"""
return (
base64.urlsafe_b64encode(unified_id_string.encode())
.decode()
.rstrip("=")
)
def decode_unified_id(encoded_unified_id: str) -> Optional[str]:
"""
Decode a base64 encoded unified ID.
Args:
encoded_unified_id: The base64 encoded unified ID
Returns:
Decoded unified ID string or None if invalid
"""
try:
# Add padding back if needed
padded = encoded_unified_id + "=" * (-len(encoded_unified_id) % 4)
# Decode from base64
decoded = base64.urlsafe_b64decode(padded).decode()
# Verify it starts with the expected prefix
if decoded.startswith("litellm_proxy:"):
return decoded
return None
except Exception:
return None
def parse_unified_id(
unified_id: str,
) -> Optional[dict]:
"""
Parse a unified ID into its components.
Args:
unified_id: The unified ID (encoded or decoded)
Returns:
Dictionary with parsed components or None if invalid
Example:
{
"resource_type": "vector_store",
"unified_uuid": "abc-123",
"target_model_names": ["gpt-4", "gemini"],
"provider_resource_id": "vs_xyz",
"model_id": "model-id-123"
}
"""
try:
# Decode if needed
decoded_id = decode_unified_id(unified_id)
if not decoded_id:
# Maybe it's already decoded
if unified_id.startswith("litellm_proxy:"):
decoded_id = unified_id
else:
return None
return {
"resource_type": extract_resource_type_from_unified_id(decoded_id),
"unified_uuid": extract_unified_uuid_from_unified_id(decoded_id),
"target_model_names": extract_target_model_names_from_unified_id(decoded_id),
"provider_resource_id": extract_provider_resource_id_from_unified_id(decoded_id),
"model_id": extract_model_id_from_unified_id(decoded_id),
}
except Exception:
return None

View file

@ -246,6 +246,11 @@ class AmazonAnthropicClaudeMessagesConfig(
"sonnet_4.5",
"sonnet-4-5",
"sonnet_4_5",
# Opus 4.6
"opus-4.6",
"opus_4.6",
"opus-4-6",
"opus_4_6",
]
return any(pattern in model_lower for pattern in supported_patterns)

View file

@ -4,7 +4,7 @@ import os
import ssl
import typing
import urllib.request
from typing import Callable, Dict, Optional, Union
from typing import Any, Callable, Dict, Optional, Union
import aiohttp
import aiohttp.client_exceptions
@ -248,26 +248,25 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
# Only pass ssl kwarg when explicitly configured, to avoid
# overriding the session/connector defaults with None (which is
# not a valid value for aiohttp's ssl parameter).
ssl_kwargs: Dict[str, Union[bool, ssl.SSLContext]] = {}
if ssl_verify is not None:
ssl_kwargs["ssl"] = ssl_verify
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(
request_kwargs: Dict[str, Any] = {
"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,
**ssl_kwargs,
).__aenter__()
"proxy": proxy,
"server_hostname": sni_hostname,
}
if ssl_verify is not None:
request_kwargs["ssl"] = ssl_verify
response = await client_session.request(**request_kwargs).__aenter__()
return response

View file

@ -1206,7 +1206,28 @@ def get_async_httpx_client(
If not present, creates a new client
Caches the new client and returns it.
Note: When shared_session is provided, the cache is bypassed to ensure
the user's session (with its trace_configs, connector settings, etc.)
is used for the request.
"""
# When shared_session is provided, bypass cache and create a new handler
# that uses the user's session directly. This preserves the user's
# session configuration including trace_configs for aiohttp tracing.
if shared_session is not None:
verbose_logger.debug(
f"shared_session provided (ID: {id(shared_session)}), bypassing client cache"
)
if params is not None:
handler_params = {k: v for k, v in params.items() if k != "disable_aiohttp_transport"}
handler_params["shared_session"] = shared_session
return AsyncHTTPHandler(**handler_params)
else:
return AsyncHTTPHandler(
timeout=httpx.Timeout(timeout=600.0, connect=5.0),
shared_session=shared_session,
)
_params_key_name = ""
if params is not None:
for key, value in params.items():
@ -1233,12 +1254,10 @@ def get_async_httpx_client(
if params is not None:
# Filter out params that are only used for cache key, not for AsyncHTTPHandler.__init__
handler_params = {k: v for k, v in params.items() if k != "disable_aiohttp_transport"}
handler_params["shared_session"] = shared_session
_new_client = AsyncHTTPHandler(**handler_params)
else:
_new_client = AsyncHTTPHandler(
timeout=httpx.Timeout(timeout=600.0, connect=5.0),
shared_session=shared_session,
)
cache.set_cache(

View file

@ -772,12 +772,15 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
try:
return ModelResponseStream(
id=chunk["id"],
object="chat.completion.chunk",
created=chunk.get("created"),
model=chunk.get("model"),
choices=chunk.get("choices", []),
)
kwargs = {
"id": chunk["id"],
"object": "chat.completion.chunk",
"created": chunk.get("created"),
"model": chunk.get("model"),
"choices": chunk.get("choices", []),
}
if "usage" in chunk and chunk["usage"] is not None:
kwargs["usage"] = chunk["usage"]
return ModelResponseStream(**kwargs)
except Exception as e:
raise e

View file

@ -2,7 +2,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Union, cast, get_type_hin
import httpx
from openai.types.responses import ResponseReasoningItem
from pydantic import BaseModel
from pydantic import BaseModel, ValidationError
import litellm
from litellm._logging import verbose_logger
@ -240,25 +240,26 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
event_pydantic_model = OpenAIResponsesAPIConfig.get_event_model_class(
event_type=event_type
)
# Defensive: Some OpenAI-compatible providers may send `error.code: null`.
# Pydantic will raise a ValidationError when it expects a string but gets None.
# Coalesce a None `error.code` to a stable default string so streaming
# iteration does not crash (see issue report). This keeps behavior similar
# to previous fixes (coalesce before validation) and lets higher-level
# handlers still receive an `ErrorEvent` object.
# Some OpenAI-compatible providers send error.code: null; coalesce so validation succeeds.
try:
error_obj = parsed_chunk.get("error")
if isinstance(error_obj, dict) and error_obj.get("code") is None:
# Preserve other fields, but ensure `code` is a non-null string
parsed_chunk = dict(parsed_chunk)
parsed_chunk["error"] = dict(error_obj)
parsed_chunk["error"]["code"] = "unknown_error"
except Exception:
# If anything unexpected happens here, fall back to attempting
# instantiation and let higher-level handlers manage errors.
verbose_logger.debug("Failed to coalesce error.code in parsed_chunk")
return event_pydantic_model(**parsed_chunk)
try:
return event_pydantic_model(**parsed_chunk)
except ValidationError:
verbose_logger.debug(
"Pydantic validation failed for %s with chunk %s, "
"falling back to model_construct",
event_pydantic_model.__name__,
parsed_chunk,
)
return event_pydantic_model.model_construct(**parsed_chunk)
@staticmethod
def get_event_model_class(event_type: str) -> Any:
@ -307,6 +308,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
ResponsesAPIStreamEvents.MCP_CALL_FAILED: MCPCallFailedEvent,
ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE: ImageGenerationPartialImageEvent,
ResponsesAPIStreamEvents.ERROR: ErrorEvent,
# Shell tool events: passthrough as GenericEvent so payload is preserved
ResponsesAPIStreamEvents.SHELL_CALL_IN_PROGRESS: GenericEvent,
ResponsesAPIStreamEvents.SHELL_CALL_COMPLETED: GenericEvent,
ResponsesAPIStreamEvents.SHELL_CALL_OUTPUT: GenericEvent,
}
model_class = event_models.get(cast(ResponsesAPIStreamEvents, event_type))

View file

@ -26,6 +26,10 @@
"max_completion_tokens": "max_tokens"
}
},
"scaleway": {
"base_url": "https://api.scaleway.ai/v1",
"api_key_env": "SCW_SECRET_KEY"
},
"synthetic": {
"base_url": "https://api.synthetic.new/openai/v1",
"api_key_env": "SYNTHETIC_API_KEY",

View file

@ -102,11 +102,18 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig):
status_code=raw_response.status_code
)
if "embedding" not in response_data:
# Handle both raw array format (TEI) and wrapped format (standard HF)
if isinstance(response_data, list):
# TEI and some HF models return raw embedding arrays directly
embeddings = response_data
elif isinstance(response_data, dict) and "embedding" in response_data:
# Standard HF format with "embedding" key
embeddings = response_data["embedding"]
else:
raise SagemakerError(
status_code=500, message="HF response missing 'embedding' field"
status_code=500,
message=f"Unexpected response format. Expected list or dict with 'embedding' key, got: {type(response_data).__name__}",
)
embeddings = response_data["embedding"]
if not isinstance(embeddings, list):
raise SagemakerError(

View file

@ -529,6 +529,17 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
raise e
def _pop_and_merge_extra_body(data: RequestBody, optional_params: dict) -> None:
"""Pop extra_body from optional_params and shallow-merge into data, deep-merging dict values."""
extra_body: Optional[dict] = optional_params.pop("extra_body", None)
if extra_body is not None:
for k, v in extra_body.items():
if k in data and isinstance(data[k], dict) and isinstance(v, dict):
data[k].update(v)
else:
data[k] = v
def _transform_request_body(
messages: List[AllMessageValues],
model: str,
@ -619,6 +630,7 @@ def _transform_request_body(
# Only add labels for Vertex AI endpoints (not Google GenAI/AI Studio) and only if non-empty
if labels and custom_llm_provider != LlmProviders.GEMINI:
data["labels"] = labels
_pop_and_merge_extra_body(data, optional_params)
except Exception as e:
raise e

View file

@ -480,7 +480,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
tool = {VertexToolName.COMPUTER_USE.value: computer_use_config}
# Handle OpenAI-style web_search and web_search_preview tools
# Transform them to Gemini's googleSearch tool
elif "type" in tool and tool["type"] in ("web_search", "web_search_preview"):
elif "type" in tool and tool["type"] in (
"web_search",
"web_search_preview",
):
verbose_logger.info(
f"Gemini: Transforming OpenAI-style '{tool['type']}' tool to googleSearch"
)
@ -1196,6 +1199,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for the prohibited contents.",
"SPII": "The token generation was stopped as the response was flagged for Sensitive Personally Identifiable Information (SPII) contents.",
"IMAGE_SAFETY": "The token generation was stopped as the response was flagged for image safety reasons.",
"IMAGE_PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for prohibited image content.",
}
@staticmethod
@ -1218,6 +1222,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"SPII": "content_filter",
"MALFORMED_FUNCTION_CALL": "malformed_function_call", # openai doesn't have a way of representing this
"IMAGE_SAFETY": "content_filter",
"IMAGE_PROHIBITED_CONTENT": "content_filter",
}
def translate_exception_str(self, exception_string: str):
@ -1630,7 +1635,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
completion_image_tokens = response_tokens_details.image_tokens or 0
completion_audio_tokens = response_tokens_details.audio_tokens or 0
calculated_text_tokens = (
candidates_token_count - completion_image_tokens - completion_audio_tokens
candidates_token_count
- completion_image_tokens
- completion_audio_tokens
)
response_tokens_details.text_tokens = calculated_text_tokens
#########################################################
@ -2248,6 +2255,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
citation_metadata # older approach - maintaining to prevent regressions
)
## ADD TRAFFIC TYPE ##
traffic_type = completion_response.get("usageMetadata", {}).get(
"trafficType"
)
if traffic_type:
model_response._hidden_params.setdefault("provider_specific_fields", {})["traffic_type"] = traffic_type
except Exception as e:
raise VertexAIError(
message="Received={}, Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
@ -2906,6 +2920,12 @@ class ModelResponseIterator:
PromptTokensDetailsWrapper, usage.prompt_tokens_details
).web_search_requests = web_search_requests
traffic_type = processed_chunk.get("usageMetadata", {}).get(
"trafficType"
)
if traffic_type:
model_response._hidden_params.setdefault("provider_specific_fields", {})["traffic_type"] = traffic_type
setattr(model_response, "usage", usage) # type: ignore
model_response._hidden_params["is_finished"] = False

View file

@ -115,8 +115,13 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
vertex_project = self.get_vertex_ai_project(litellm_params)
vertex_location = self.get_vertex_ai_location(litellm_params)
# Construct full rag corpus path
full_rag_corpus = f"projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{vector_store_id}"
# Handle both full corpus path and just corpus ID
if vector_store_id.startswith("projects/"):
# Already a full path
full_rag_corpus = vector_store_id
else:
# Just the corpus ID, construct full path
full_rag_corpus = f"projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{vector_store_id}"
# Build the request body for Vertex AI RAG API
request_body: Dict[str, Any] = {

View file

@ -7383,6 +7383,16 @@ def stream_chunk_builder( # noqa: PLR0915
setattr(response, "usage", usage)
# Propagate provider_specific_fields from the last chunk (contains provider
# metadata like traffic_type set during streaming)
for chunk in reversed(chunks):
hidden = getattr(chunk, "_hidden_params", None)
if hidden and "provider_specific_fields" in hidden:
response._hidden_params.setdefault(
"provider_specific_fields", {}
).update(hidden["provider_specific_fields"])
break
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
setattr(

File diff suppressed because it is too large Load diff

View file

@ -536,10 +536,25 @@ class MCPRequestHandler:
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
try:
# Get key object permission (already loaded in main auth flow)
# Get key object permission (already loaded in main auth flow, or fetch from DB)
key_object_permission = MCPRequestHandler._get_key_object_permission(
user_api_key_auth
)
if key_object_permission is None and user_api_key_auth and user_api_key_auth.object_permission_id:
from litellm.proxy.auth.auth_checks import get_object_permission
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is not None:
key_object_permission = await get_object_permission(
object_permission_id=user_api_key_auth.object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if key_object_permission is None:
return []

View file

@ -14,6 +14,7 @@ import re
from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from urllib.parse import urlparse
import anyio
from fastapi import HTTPException
from httpx import HTTPStatusError
from mcp import ReadResourceResult, Resource
@ -887,8 +888,15 @@ class MCPServerManager:
# Handle stdio transport
if transport == MCPTransport.stdio:
# For stdio, we need to get the stdio config from the server
resolved_env = stdio_env if stdio_env is not None else server.env or {}
resolved_env = stdio_env if stdio_env is not None else dict(server.env or {})
# Ensure npm-based STDIO MCP servers have a writable cache dir.
# In containers the default (~/.npm or /app/.npm) may not exist
# or be read-only, causing npx to fail with ENOENT.
if "NPM_CONFIG_CACHE" not in resolved_env:
from litellm.constants import MCP_NPM_CACHE_DIR
resolved_env["NPM_CONFIG_CACHE"] = MCP_NPM_CACHE_DIR
stdio_config: Optional[MCPStdioConfig] = None
if server.command and server.args is not None:
stdio_config = MCPStdioConfig(
@ -1437,6 +1445,9 @@ class MCPServerManager:
"""
Fetch tools from MCP client with timeout and error handling.
Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts
with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details.
Args:
client: MCP client instance
server_name: Name of the server for logging
@ -1444,24 +1455,12 @@ class MCPServerManager:
Returns:
List of tools from the server
"""
async def _list_tools_task():
try:
try:
with anyio.fail_after(30.0):
tools = await client.list_tools()
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
except asyncio.CancelledError:
verbose_logger.warning(f"Client operation cancelled for {server_name}")
return []
except Exception as e:
verbose_logger.warning(
f"Client operation failed for {server_name}: {str(e)}"
)
return []
try:
return await asyncio.wait_for(_list_tools_task(), timeout=30.0)
except asyncio.TimeoutError:
except TimeoutError:
verbose_logger.warning(f"Timeout while listing tools from {server_name}")
return []
except asyncio.CancelledError:
@ -2481,6 +2480,9 @@ class MCPServerManager:
except asyncio.TimeoutError:
health_check_error = "Health check timed out after 10 seconds"
status = "unhealthy"
except asyncio.CancelledError:
health_check_error = "Health check was cancelled"
status = "unknown"
except Exception as e:
health_check_error = str(e)
status = "unhealthy"

View file

@ -24,6 +24,7 @@ from fastapi import FastAPI, HTTPException
from pydantic import AnyUrl, ConfigDict
from starlette.requests import Request as StarletteRequest
from starlette.types import Receive, Scope, Send
from starlette.responses import JSONResponse
from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
@ -41,6 +42,9 @@ from litellm.proxy._experimental.mcp_server.utils import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
@ -842,6 +846,7 @@ if MCP_AVAILABLE:
raw_headers: Optional[Dict[str, str]] = None,
log_list_tools_to_spendlogs: bool = False,
list_tools_log_source: Optional[str] = None,
litellm_trace_id: Optional[str] = None,
) -> List[MCPTool]:
"""
Helper method to fetch tools from MCP servers based on server filtering criteria.
@ -879,6 +884,7 @@ if MCP_AVAILABLE:
"model": "MCP: list_tools",
"call_type": CallTypes.list_mcp_tools.value,
"litellm_call_id": list_tools_call_id,
"litellm_trace_id": litellm_trace_id,
"metadata": {
"spend_logs_metadata": spend_logs_metadata,
},
@ -894,13 +900,14 @@ if MCP_AVAILABLE:
],
}
# Attach user identifiers when available (matches call_mcp_tool style)
# Attach user identifiers using the standard helper
if user_api_key_auth is not None:
user_api_key = getattr(user_api_key_auth, "api_key", None)
if user_api_key:
cast(dict, list_tools_request_data["metadata"])[
"user_api_key"
] = user_api_key
LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
data=list_tools_request_data,
user_api_key_dict=user_api_key_auth,
_metadata_variable_name="metadata",
)
user_identifier = getattr(
user_api_key_auth, "end_user_id", None
@ -1907,18 +1914,27 @@ if MCP_AVAILABLE:
raw_headers,
)
def _strip_stale_mcp_session_header(
async def _handle_stale_mcp_session(
scope: Scope,
receive: Receive,
send: Send,
mgr: "StreamableHTTPSessionManager",
) -> None:
) -> bool:
"""
Strip stale ``mcp-session-id`` headers so the session manager
creates a fresh session instead of returning 404 "Session not found".
Handle stale MCP session IDs to prevent "Session not found" errors.
When clients like VSCode reconnect after a reload they may resend a
session id that has already been cleaned up. Rather than letting the
SDK return a 404 error loop, we detect the stale id and remove the
header so a brand-new session is created transparently.
When clients reconnect after a server restart or session cleanup, they may
send a session ID that no longer exists. This function handles two scenarios:
1. Non-DELETE requests: Strip the stale session ID header so the session
manager creates a fresh session transparently.
2. DELETE requests: Return success (200) immediately for idempotent behavior,
since the desired state (session doesn't exist) is already achieved.
Returns:
True if the request was handled (DELETE on non-existent session)
False if the request should continue to the session manager
Fixes https://github.com/BerriAI/litellm/issues/20292
"""
@ -1930,10 +1946,30 @@ if MCP_AVAILABLE:
break
if _session_id is None:
return
return False
known_sessions = getattr(mgr, "_server_instances", None)
if known_sessions is not None and _session_id not in known_sessions:
if known_sessions is None or _session_id in known_sessions:
# Session exists or we can't check - let the session manager handle it
return False
# Session doesn't exist - handle based on request method
method = scope.get("method", "").upper()
if method == "DELETE":
# Idempotent DELETE: session doesn't exist, return success
verbose_logger.info(
f"DELETE request for non-existent MCP session '{_session_id}'. "
"Returning success (idempotent DELETE)."
)
success_response = JSONResponse(
status_code=200,
content={"message": "Session terminated successfully"}
)
await success_response(scope, receive, send)
return True
else:
# Non-DELETE: strip stale session ID to allow new session creation
verbose_logger.warning(
"MCP session ID '%s' not found in active sessions. "
"Stripping stale header to force new session creation.",
@ -1943,6 +1979,7 @@ if MCP_AVAILABLE:
(k, v) for k, v in scope["headers"]
if k != _mcp_session_header
]
return False
async def handle_streamable_http_mcp(
scope: Scope, receive: Receive, send: Send
@ -2005,7 +2042,12 @@ if MCP_AVAILABLE:
# Give it a moment to start up
await asyncio.sleep(0.1)
_strip_stale_mcp_session_header(scope, session_manager)
# Handle stale session IDs - either strip them for reconnection
# or return success for idempotent DELETE operations
handled = await _handle_stale_mcp_session(scope, receive, send, session_manager)
if handled:
# Request was fully handled (e.g., DELETE on non-existent session)
return
await session_manager.handle_request(scope, receive, send)
except HTTPException:

View file

@ -514,6 +514,8 @@ class LiteLLMRoutes(enum.Enum):
"/user/delete",
"/user/info",
"/user/list",
"/user/daily/activity",
"/user/daily/activity/aggregated",
# team
"/team/new",
"/team/update",
@ -526,6 +528,7 @@ class LiteLLMRoutes(enum.Enum):
"/team/available",
"/team/permissions_list",
"/team/permissions_update",
"/team/daily/activity",
# model
"/model/new",
"/model/update",
@ -893,6 +896,7 @@ class KeyRequestBase(GenerateRequestBase):
Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"]
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
router_settings: Optional[UpdateRouterConfig] = None
access_group_ids: Optional[List[str]] = None
class LiteLLMKeyType(str, enum.Enum):
@ -1502,6 +1506,7 @@ class TeamBase(LiteLLMPydanticObjectBase):
models: list = []
blocked: bool = False
router_settings: Optional[dict] = None
access_group_ids: Optional[List[str]] = None
class NewTeamRequest(TeamBase):
@ -1589,6 +1594,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
model_tpm_limit: Optional[Dict[str, int]] = None
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
router_settings: Optional[dict] = None
access_group_ids: Optional[List[str]] = None
class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase):
@ -2177,6 +2183,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase):
updated_by: Optional[str] = None
object_permission_id: Optional[str] = None
object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
access_group_ids: Optional[List[str]] = None
rotation_count: Optional[int] = 0 # Number of times key has been rotated
auto_rotate: Optional[bool] = False # Whether this key should be auto-rotated
rotation_interval: Optional[str] = None # How often to rotate (e.g., "30d", "90d")
@ -3945,6 +3952,18 @@ class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase):
file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, ResponsesAPIResponse]
class LiteLLM_ManagedVectorStoreTable(LiteLLMPydanticObjectBase):
"""Table for managing vector stores with target_model_names support."""
unified_resource_id: str
resource_object: Optional[Any] = None # VectorStoreCreateResponse
model_mappings: Dict[str, str]
flat_model_resource_ids: List[str]
created_by: Optional[str]
updated_by: Optional[str]
storage_backend: Optional[str] = None
storage_url: Optional[str] = None
class EnterpriseLicenseData(TypedDict, total=False):
expiration_date: str
user_id: str

View file

@ -925,10 +925,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if isinstance(
api_key, str
): # if generated token, make sure it starts with sk-.
_masked_key = "{}****{}".format(api_key[:4], api_key[-4:]) if len(api_key) > 8 else "****"
assert api_key.startswith(
"sk-"
), "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format(
api_key
_masked_key
) # prevent token hashes from being used
else:
verbose_logger.warning(

View file

@ -282,7 +282,7 @@ def _override_openai_response_model(
if isinstance(response_obj, dict):
downstream_model = response_obj.get("model")
if downstream_model != requested_model:
verbose_proxy_logger.warning(
verbose_proxy_logger.debug(
"%s: response model mismatch - requested=%r downstream=%r. Overriding response['model'] to requested model.",
log_context,
requested_model,
@ -301,7 +301,7 @@ def _override_openai_response_model(
downstream_model = getattr(response_obj, "model", None)
if downstream_model != requested_model:
verbose_proxy_logger.warning(
verbose_proxy_logger.debug(
"%s: response model mismatch - requested=%r downstream=%r. Overriding response.model to requested model.",
log_context,
requested_model,

View file

@ -112,7 +112,8 @@ class SpendUpdateQueue(BaseUpdateQueue):
for update in updates:
_key = f"{update.get('entity_type')}:{update.get('entity_id')}"
if _key not in _in_memory_map:
_in_memory_map[_key] = update
# avoid mutating caller-owned dicts while aggregating queue entries
_in_memory_map[_key] = update.copy()
else:
current_cost = _in_memory_map[_key].get("response_cost", 0) or 0
update_cost = update.get("response_cost", 0) or 0

View file

@ -5,10 +5,12 @@
# +-------------------------------------------------------------+
# Thank you users! We ❤️ you! - Krrish & Ishaan
import fnmatch
import os
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
@ -31,6 +33,110 @@ if TYPE_CHECKING:
GUARDRAIL_NAME = "generic_guardrail_api"
# Headers whose values are forwarded as-is (case-insensitive). Glob patterns supported (e.g. x-stainless-*, x-litellm*).
_HEADER_VALUE_ALLOWLIST = frozenset({
"host",
"accept-encoding",
"connection",
"accept",
"content-type",
"user-agent",
"x-stainless-*",
"x-litellm-*",
"content-length",
})
# Placeholder for headers that exist but are not on the allowlist (we don't expose their value).
_HEADER_PRESENT_PLACEHOLDER = "[present]"
def _header_value_allowed(header_name: str) -> bool:
"""Return True if this header's value may be forwarded (allowlist, including globs)."""
lower = header_name.lower()
if lower in _HEADER_VALUE_ALLOWLIST:
return True
for pattern in _HEADER_VALUE_ALLOWLIST:
if "*" in pattern and fnmatch.fnmatch(lower, pattern):
return True
return False
def _sanitize_inbound_headers(headers: Any) -> Optional[Dict[str, str]]:
"""
Sanitize inbound headers before passing them to a 3rd party guardrail service.
- Allowlist: only headers in the allowlist have their values forwarded (exact + glob: x-stainless-*, x-litellm-*).
- All other headers are included with value "[present]" so the guardrail knows the header existed.
- Coerces values to str (for JSON serialization).
"""
if not headers or not isinstance(headers, dict):
return None
sanitized: Dict[str, str] = {}
for k, v in headers.items():
if k is None:
continue
key = str(k)
if _header_value_allowed(key):
try:
sanitized[key] = str(v)
except Exception:
continue
else:
sanitized[key] = _HEADER_PRESENT_PLACEHOLDER
return sanitized or None
def _extract_inbound_headers(
request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]
) -> Optional[Dict[str, str]]:
"""
Extract inbound headers from available request context.
We try multiple locations to support different call paths:
- proxy endpoints: request_data["proxy_server_request"]["headers"]
- if the guardrail is passed the proxy_server_request object directly
- metadata headers captured in litellm_pre_call_utils
- response hooks: fallback to logging_obj.model_call_details
"""
# 1) Most common path (proxy): full request context in proxy_server_request
headers = request_data.get("proxy_server_request", {}).get("headers")
if headers:
return _sanitize_inbound_headers(headers)
# 2) Some guardrails pass proxy_server_request as request_data itself
headers = request_data.get("headers")
if headers:
return _sanitize_inbound_headers(headers)
# 3) Pre-call: headers stored in request metadata
metadata_headers = (request_data.get("metadata") or {}).get("headers")
if metadata_headers:
return _sanitize_inbound_headers(metadata_headers)
litellm_metadata_headers = (request_data.get("litellm_metadata") or {}).get(
"headers"
)
if litellm_metadata_headers:
return _sanitize_inbound_headers(litellm_metadata_headers)
# 4) Post-call: headers not present on response; fallback to logging object
if logging_obj and getattr(logging_obj, "model_call_details", None):
try:
details = logging_obj.model_call_details or {}
headers = (
details.get("litellm_params", {})
.get("metadata", {})
.get("headers", None)
)
if headers:
return _sanitize_inbound_headers(headers)
except Exception:
pass
return None
class GenericGuardrailAPI(CustomGuardrail):
"""
@ -207,6 +313,7 @@ class GenericGuardrailAPI(CustomGuardrail):
# Extract user API key metadata
user_metadata = self._extract_user_api_key_metadata(request_data)
inbound_headers = _extract_inbound_headers(request_data=request_data, logging_obj=logging_obj)
# Create request payload
guardrail_request = GenericGuardrailAPIRequest(
@ -214,6 +321,8 @@ class GenericGuardrailAPI(CustomGuardrail):
litellm_trace_id=logging_obj.litellm_trace_id if logging_obj else None,
texts=texts,
request_data=user_metadata,
request_headers=inbound_headers,
litellm_version=litellm_version,
images=images,
tools=tools,
structured_messages=structured_messages,

View file

@ -29,10 +29,7 @@ from fastapi import HTTPException
from litellm import Router
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import ModelResponseStream
@ -1056,7 +1053,6 @@ class ContentFilterGuardrail(CustomGuardrail):
masked_entity_count=masked_entity_count,
)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: "GenericGuardrailAPIInputs",

View file

@ -330,6 +330,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
original_inputs: Optional[dict] = None,
):
"""
Override to store only the Model Armor API response, not the entire data dict.

View file

@ -5,14 +5,9 @@ OpenAI Moderation Guardrail Integration for LiteLLM
from typing import (
TYPE_CHECKING,
Any,
AsyncGenerator,
Dict,
List,
Literal,
Optional,
Type,
Union,
)
from fastapi import HTTPException
@ -20,7 +15,7 @@ from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
log_guardrail_information
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
@ -32,10 +27,8 @@ from litellm.types.utils import GenericGuardrailAPIInputs
from .base import OpenAIGuardrailBase
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import OpenAIModerationResponse
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
from litellm.types.utils import ModelResponse, ModelResponseStream
class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
@ -236,108 +229,6 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
# Moderation doesn't modify content, just blocks - return inputs unchanged
return inputs
@log_guardrail_information
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: "UserAPIKeyAuth",
response: Any,
request_data: Dict[str, Any],
) -> AsyncGenerator["ModelResponseStream", None]:
"""
Process streaming response chunks for OpenAI moderation.
Collects all chunks from the stream, assembles them into a complete response,
and applies moderation check. If content violates moderation policy, raises HTTPException.
"""
# Import here to avoid circular imports
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.main import stream_chunk_builder
from litellm.types.utils import TextCompletionResponse
verbose_proxy_logger.debug("OpenAI Moderation: Running streaming response scan")
# Collect all chunks to process them together
all_chunks: List["ModelResponseStream"] = []
async for chunk in response:
all_chunks.append(chunk)
# Assemble the complete response from chunks
assembled_model_response: Optional[
Union["ModelResponse", TextCompletionResponse]
] = stream_chunk_builder(
chunks=all_chunks,
)
if isinstance(assembled_model_response, (type(None), TextCompletionResponse)):
# If we can't assemble a ModelResponse or it's a text completion,
# just yield the original chunks without moderation
verbose_proxy_logger.warning(
"OpenAI Moderation: Could not assemble ModelResponse from chunks, skipping moderation"
)
for chunk in all_chunks:
yield chunk
return
# Extract response text for moderation
response_text = self._extract_response_text(assembled_model_response)
if response_text:
verbose_proxy_logger.debug(
f"OpenAI Moderation: Streaming response text: {response_text[:100]}..." # Log first 100 chars
)
# Make moderation request - this will raise HTTPException if content is flagged
moderation_response = await self.async_make_request(
input_text=response_text,
)
# Check if content is flagged and raise exception if needed
self._check_moderation_result(moderation_response)
# If we reach here, content passed moderation - yield the original chunks
mock_response = MockResponseIterator(model_response=assembled_model_response)
# Return the reconstructed stream
async for chunk in mock_response:
yield chunk
def _extract_response_text(self, response: "ModelResponse") -> Optional[str]:
"""
Extract text content from the model response for moderation.
"""
if not hasattr(response, "choices") or not response.choices:
return None
response_texts = []
for choice in response.choices:
try:
# Try to get content from message (chat completion)
message = getattr(choice, "message", None)
if message:
content = getattr(message, "content", None)
if content and isinstance(content, str):
response_texts.append(content)
continue
# Try to get text (text completion)
text = getattr(choice, "text", None)
if text and isinstance(text, str):
response_texts.append(text)
continue
# Try to get content from delta (streaming)
delta = getattr(choice, "delta", None)
if delta:
content = getattr(delta, "content", None)
if content and isinstance(content, str):
response_texts.append(content)
continue
except (AttributeError, TypeError):
# Skip choices that don't have expected attributes
continue
return "\n".join(response_texts) if response_texts else None
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
"""

View file

@ -386,7 +386,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
continue
return final_results
except Exception as e:
raise e
# Sanitize exception to avoid leaking the original text (which may
# contain API keys or other secrets) in error responses.
raise Exception(
f"Presidio PII analysis failed: {type(e).__name__}"
) from e
async def anonymize_text(
self,
@ -443,9 +447,15 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
)
return redacted_text["text"]
else:
raise Exception(f"Invalid anonymizer response: {redacted_text}")
raise Exception("Invalid anonymizer response: received None")
except Exception as e:
raise e
# Sanitize exception to avoid leaking the original text (which may
# contain API keys or other secrets) in error responses.
if "Invalid anonymizer response" in str(e):
raise
raise Exception(
f"Presidio PII anonymization failed: {type(e).__name__}"
) from e
def filter_analyze_results_by_score(
self, analyze_results: Union[List[PresidioAnalyzeResponseItem], Dict]

View file

@ -21,6 +21,7 @@ from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
GUARDRAIL_TIMEOUT = 5
@ -334,3 +335,11 @@ class ZscalerAIGuard(CustomGuardrail):
user_facing_error = self._create_user_facing_error(f"{str(e)})")
# This exception will be caught by the proxy and returned to the user
raise HTTPException(status_code=500, detail=user_facing_error)
@staticmethod
def get_config_model() -> Optional[type["GuardrailConfigModel"]]:
from litellm.types.proxy.guardrails.guardrail_hooks.zscaler_ai_guard import (
ZscalerAIGuardConfigModel,
)
return ZscalerAIGuardConfigModel

View file

@ -12,6 +12,7 @@ from litellm._uuid import uuid
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy.utils import PrismaClient
from litellm.proxy.types_utils.utils import get_instance_fn
from litellm.secret_managers.main import get_secret
from litellm.types.guardrails import (
Guardrail,
@ -489,7 +490,7 @@ class InMemoryGuardrailHandler:
config_file_path: Optional[str] = None,
) -> Optional[CustomGuardrail]:
"""
Initialize a Custom Guardrail from a python file
Initialize a Custom Guardrail from a python file or module path
This initializes it by adding it to the litellm callback manager
"""
@ -498,26 +499,12 @@ class InMemoryGuardrailHandler:
"GuardrailsAIException - Please pass the config_file_path to initialize_guardrails_v2"
)
_file_name, _class_name = guardrail_type.split(".")
verbose_proxy_logger.debug(
"Initializing custom guardrail: %s, file_name: %s, class_name: %s",
"Initializing custom guardrail: %s",
guardrail_type,
_file_name,
_class_name,
)
directory = os.path.dirname(config_file_path)
module_file_path = os.path.join(directory, _file_name) + ".py"
spec = importlib.util.spec_from_file_location(_class_name, module_file_path) # type: ignore
if not spec:
raise ImportError(
f"Could not find a module specification for {module_file_path}"
)
module = importlib.util.module_from_spec(spec) # type: ignore
spec.loader.exec_module(module) # type: ignore
_guardrail_class = getattr(module, _class_name)
_guardrail_class = get_instance_fn(guardrail_type, config_file_path=config_file_path)
mode = litellm_params.mode
if mode is None:

View file

@ -255,12 +255,24 @@ class _PROXY_BatchRateLimiter(CustomLogger):
BatchFileUsage with total_tokens and request_count
"""
try:
# Read file content
file_content = await litellm.afile_content(
file_id=file_id,
custom_llm_provider=custom_llm_provider,
user_api_key_dict=user_api_key_dict,
# Check if this is a managed file (base64 encoded unified file ID)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
is_managed_file = _is_base64_encoded_unified_file_id(file_id)
if is_managed_file and user_api_key_dict is not None:
# For managed files, use the managed files hook directly
file_content = await self._fetch_managed_file_content(
file_id=file_id,
user_api_key_dict=user_api_key_dict,
)
else:
# For non-managed files, use the standard litellm.afile_content
file_content = await litellm.afile_content(
file_id=file_id,
custom_llm_provider=custom_llm_provider,
user_api_key_dict=user_api_key_dict,
)
file_content_as_dict = _get_file_content_as_dictionary(
file_content.content
@ -282,6 +294,67 @@ class _PROXY_BatchRateLimiter(CustomLogger):
)
raise
async def _fetch_managed_file_content(
self,
file_id: str,
user_api_key_dict: UserAPIKeyAuth,
) -> Any:
"""
Fetch file content from managed files hook.
This is needed for managed files because they require proper user context
to verify file ownership and access permissions.
Args:
file_id: The managed file ID (base64 encoded)
user_api_key_dict: User authentication information
Returns:
HttpxBinaryResponseContent with the file content
"""
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
# Import proxy_server dependencies at runtime to avoid circular imports
try:
from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
except ImportError as e:
raise ValueError(
f"Cannot import proxy_server dependencies: {str(e)}. "
"Managed files require proxy_server to be initialized."
)
# Get the managed files hook
if proxy_logging_obj is None:
raise ValueError(
"proxy_logging_obj not available. Cannot access managed files hook."
)
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
if managed_files_obj is None:
raise ValueError(
"Managed files hook not found. Cannot access managed file."
)
if not isinstance(managed_files_obj, BaseFileEndpoints):
raise ValueError(
"Managed files hook is not a BaseFileEndpoints instance."
)
if llm_router is None:
raise ValueError(
"llm_router not available. Cannot access managed files."
)
# Use the managed files hook to get file content
# This properly handles user permissions and file ownership
file_content = await managed_files_obj.afile_content(
file_id=file_id,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
llm_router=llm_router,
)
return file_content
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,

View file

@ -150,15 +150,26 @@ class KeyManagementEventHooks:
existing_key_row.key_alias
or f"virtual-key-{existing_key_row.token}"
)
new_secret_name = (
response.key_alias
or data.key_alias
or f"virtual-key-{response.token_id}"
)
verbose_proxy_logger.info(
"Updating secret in secret manager: secret_name=%s",
new_secret_name,
)
team_id = getattr(existing_key_row, "team_id", None)
await KeyManagementEventHooks._rotate_virtual_key_in_secret_manager(
current_secret_name=initial_secret_name,
new_secret_name=response.key_alias
or data.key_alias
or f"virtual-key-{response.token_id}",
new_secret_name=new_secret_name,
new_secret_value=response.key,
team_id=team_id,
)
verbose_proxy_logger.info(
"Secret updated in secret manager: secret_name=%s",
new_secret_name,
)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to rotate virtual key in secret manager: {e}"

View file

@ -202,8 +202,8 @@ class _ProxyDBLogger(CustomLogger):
max_budget=end_user_max_budget,
)
else:
if kwargs["stream"] is not True or (
kwargs["stream"] is True and "complete_streaming_response" in kwargs
if kwargs.get("stream") is not True or (
kwargs.get("stream") is True and "complete_streaming_response" in kwargs
):
if sl_object is not None:
cost_tracking_failure_debug_info: Union[dict, str] = (

View file

@ -0,0 +1,264 @@
from typing import List
from fastapi import APIRouter, Depends, HTTPException, status
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.utils import get_prisma_client_or_throw
from litellm.types.access_group import (
AccessGroupCreateRequest,
AccessGroupResponse,
AccessGroupUpdateRequest,
)
router = APIRouter(
tags=["access group management"],
)
def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": CommonProxyErrors.not_allowed_access.value},
)
def _record_to_response(record) -> AccessGroupResponse:
return AccessGroupResponse(
access_group_id=record.access_group_id,
access_group_name=record.access_group_name,
description=record.description,
access_model_ids=record.access_model_ids,
access_mcp_server_ids=record.access_mcp_server_ids,
access_agent_ids=record.access_agent_ids,
assigned_team_ids=record.assigned_team_ids,
assigned_key_ids=record.assigned_key_ids,
created_at=record.created_at,
created_by=record.created_by,
updated_at=record.updated_at,
updated_by=record.updated_by,
)
@router.post(
"/v1/access_group",
response_model=AccessGroupResponse,
status_code=status.HTTP_201_CREATED,
)
async def create_access_group(
data: AccessGroupCreateRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> AccessGroupResponse:
_require_proxy_admin(user_api_key_dict)
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
existing = await prisma_client.db.litellm_accessgrouptable.find_unique(
where={"access_group_name": data.access_group_name}
)
if existing is not None:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=f"Access group '{data.access_group_name}' already exists",
)
try:
record = await prisma_client.db.litellm_accessgrouptable.create(
data={
"access_group_name": data.access_group_name,
"description": data.description,
"access_model_ids": data.access_model_ids or [],
"access_mcp_server_ids": data.access_mcp_server_ids or [],
"access_agent_ids": data.access_agent_ids or [],
"assigned_team_ids": data.assigned_team_ids or [],
"assigned_key_ids": data.assigned_key_ids or [],
"created_by": user_api_key_dict.user_id,
"updated_by": user_api_key_dict.user_id,
}
)
except Exception as e:
# Race condition: another request created the same name between find_unique and create.
# Prisma raises UniqueViolationError (P2002) or similar for unique constraint.
if "unique constraint" in str(e).lower() or "P2002" in str(e):
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=f"Access group '{data.access_group_name}' already exists",
)
raise
return _record_to_response(record)
@router.get(
"/v1/access_group",
response_model=List[AccessGroupResponse],
)
async def list_access_groups(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> List[AccessGroupResponse]:
_require_proxy_admin(user_api_key_dict)
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
records = await prisma_client.db.litellm_accessgrouptable.find_many(
order={"created_at": "desc"}
)
return [_record_to_response(r) for r in records]
@router.get(
"/v1/access_group/{access_group_id}",
response_model=AccessGroupResponse,
)
async def get_access_group(
access_group_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> AccessGroupResponse:
_require_proxy_admin(user_api_key_dict)
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
record = await prisma_client.db.litellm_accessgrouptable.find_unique(
where={"access_group_id": access_group_id}
)
if record is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
return _record_to_response(record)
@router.put(
"/v1/access_group/{access_group_id}",
response_model=AccessGroupResponse,
)
async def update_access_group(
access_group_id: str,
data: AccessGroupUpdateRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> AccessGroupResponse:
_require_proxy_admin(user_api_key_dict)
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
existing = await prisma_client.db.litellm_accessgrouptable.find_unique(
where={"access_group_id": access_group_id}
)
if existing is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
update_data: dict = {"updated_by": user_api_key_dict.user_id}
for field, value in data.model_dump(exclude_unset=True).items():
update_data[field] = value
record = await prisma_client.db.litellm_accessgrouptable.update(
where={"access_group_id": access_group_id},
data=update_data,
)
return _record_to_response(record)
@router.delete(
"/v1/access_group/{access_group_id}",
status_code=status.HTTP_204_NO_CONTENT,
)
async def delete_access_group(
access_group_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> None:
_require_proxy_admin(user_api_key_dict)
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
try:
async with prisma_client.db.tx() as tx:
existing = await tx.litellm_accessgrouptable.find_unique(
where={"access_group_id": access_group_id}
)
if existing is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
# Remove access_group_id from teams and keys that reference it
teams_with_group = await tx.litellm_teamtable.find_many(
where={"access_group_ids": {"hasSome": [access_group_id]}}
)
for team in teams_with_group:
updated_ids = [tid for tid in (team.access_group_ids or []) if tid != access_group_id]
await tx.litellm_teamtable.update(
where={"team_id": team.team_id},
data={"access_group_ids": updated_ids},
)
keys_with_group = await tx.litellm_verificationtoken.find_many(
where={"access_group_ids": {"hasSome": [access_group_id]}}
)
for key in keys_with_group:
updated_ids = [kid for kid in (key.access_group_ids or []) if kid != access_group_id]
await tx.litellm_verificationtoken.update(
where={"token": key.token},
data={"access_group_ids": updated_ids},
)
await tx.litellm_accessgrouptable.delete(
where={"access_group_id": access_group_id}
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception(
"delete_access_group failed: access_group_id=%s error=%s",
access_group_id,
e,
)
if PrismaDBExceptionHandler.is_database_connection_error(e):
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=CommonProxyErrors.db_not_connected_error.value,
)
if "P2025" in str(e) or ("record" in str(e).lower() and "not found" in str(e).lower()):
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to delete access group. Please try again.",
)
# Alias routes for /v1/unified_access_group
router.add_api_route(
"/v1/unified_access_group",
create_access_group,
methods=["POST"],
response_model=AccessGroupResponse,
status_code=status.HTTP_201_CREATED,
)
router.add_api_route(
"/v1/unified_access_group",
list_access_groups,
methods=["GET"],
response_model=List[AccessGroupResponse],
)
router.add_api_route(
"/v1/unified_access_group/{access_group_id}",
get_access_group,
methods=["GET"],
response_model=AccessGroupResponse,
)
router.add_api_route(
"/v1/unified_access_group/{access_group_id}",
update_access_group,
methods=["PUT"],
response_model=AccessGroupResponse,
)
router.add_api_route(
"/v1/unified_access_group/{access_group_id}",
delete_access_group,
methods=["DELETE"],
status_code=status.HTTP_204_NO_CONTENT,
)

View file

@ -383,6 +383,17 @@ def _update_metadata_field(updated_kv: dict, field_name: str) -> None:
updated_kv["metadata"] = {field_name: _value}
def _has_non_empty_value(value: Any) -> bool:
"""Check if a value has real content (not None, not empty list, not blank string)."""
if value is None:
return False
if isinstance(value, list) and len(value) == 0:
return False
if isinstance(value, str) and value.strip() == "":
return False
return True
def _update_metadata_fields(updated_kv: dict) -> None:
"""
Helper function to update all metadata fields (both premium and standard).
@ -391,7 +402,7 @@ def _update_metadata_fields(updated_kv: dict) -> None:
updated_kv: The key-value dict being used for the update
"""
for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
if field in updated_kv and updated_kv[field] is not None:
if field in updated_kv and _has_non_empty_value(updated_kv[field]):
_update_metadata_field(updated_kv=updated_kv, field_name=field)
for field in LiteLLM_ManagementEndpoint_MetadataFields:

View file

@ -628,10 +628,11 @@ async def _common_key_generation_helper( # noqa: PLR0915
# Validate user-provided key format
if data.key is not None and not data.key.startswith("sk-"):
_masked = "{}****{}".format(data.key[:4], data.key[-4:]) if len(data.key) > 8 else "****"
raise HTTPException(
status_code=400,
detail={
"error": f"Invalid key format. LiteLLM Virtual Key must start with 'sk-'. Received: {data.key}"
"error": f"Invalid key format. LiteLLM Virtual Key must start with 'sk-'. Received: {_masked}"
},
)
@ -2770,6 +2771,7 @@ async def can_modify_verification_token(
Rules:
- Proxy admin can modify any key
- Internal jobs service account can modify any key (for auto-rotation)
- For team keys: only team admin or key owner can modify
- For personal keys: only key owner can modify
@ -2782,13 +2784,19 @@ async def can_modify_verification_token(
Returns:
True if user can modify the key, False otherwise
"""
from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
is_team_key = _is_team_key(data=key_info)
# 1. Proxy admin can modify any key
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return True
# 2. For team keys: only team admin or key owner can modify
# 2. Internal jobs service account can modify any key (for auto-rotation)
if user_api_key_dict.api_key == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME:
return True
# 3. For team keys: only team admin or key owner can modify
if is_team_key and key_info.team_id is not None:
# Get team object to check if user is team admin
team_table = await get_team_object(
@ -2818,7 +2826,7 @@ async def can_modify_verification_token(
# Not team admin and doesn't own the key
return False
# 3. For personal keys: only key owner can modify
# 4. For personal keys: only key owner can modify
if key_info.user_id is not None and key_info.user_id == user_api_key_dict.user_id:
return True
@ -3179,7 +3187,7 @@ def get_new_token(data: Optional[RegenerateKeyRequest]) -> str:
dependencies=[Depends(user_api_key_auth)],
)
@management_endpoint_wrapper
async def regenerate_key_fn(
async def regenerate_key_fn( # noqa: PLR0915
key: Optional[str] = None,
data: Optional[RegenerateKeyRequest] = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -3330,6 +3338,10 @@ async def regenerate_key_fn(
detail={"error": "You are not authorized to regenerate this key"},
)
verbose_proxy_logger.info(
"Key regeneration requested: key_alias=%s",
getattr(_key_in_db, "key_alias", None),
)
verbose_proxy_logger.debug("key_in_db: %s", _key_in_db)
new_token = get_new_token(data=data)
@ -3380,6 +3392,10 @@ async def regenerate_key_fn(
**updated_token_dict,
)
verbose_proxy_logger.info(
"Key regeneration completed: key_alias=%s",
getattr(_key_in_db, "key_alias", None),
)
asyncio.create_task(
KeyManagementEventHooks.async_key_rotated_hook(
data=data,

View file

@ -10,10 +10,13 @@ Endpoints here:
- DELETE `/v1/mcp/server/{server_id}` - Deletes the mcp server given `server_id`.
- GET `/v1/mcp/tools - lists all the tools available for a key
- GET `/v1/mcp/access_groups` - lists all available MCP access groups
- GET `/v1/mcp/discover` - Returns curated list of well-known MCP servers for discovery UI
"""
import importlib
import json
import os
from dataclasses import dataclass
from datetime import datetime, timedelta
from typing import Any, Dict, Iterable, List, Literal, Optional
@ -1176,3 +1179,88 @@ if MCP_AVAILABLE:
except Exception as e:
verbose_proxy_logger.exception(f"Error making agent public: {e}")
raise HTTPException(status_code=500, detail=str(e))
# --- MCP Discovery ---
_MCP_REGISTRY_PATH = os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"mcp_registry.json",
)
_mcp_registry_cache: Optional[Dict[str, Any]] = None
def _load_mcp_registry() -> Dict[str, Any]:
"""Load the curated MCP registry from disk. Cached after first read."""
global _mcp_registry_cache
if _mcp_registry_cache is not None:
return _mcp_registry_cache
try:
with open(_MCP_REGISTRY_PATH, "r") as f:
data: Dict[str, Any] = json.load(f)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to load MCP registry from {_MCP_REGISTRY_PATH}: {e}"
)
data = {"servers": []}
_mcp_registry_cache = data
return data
@router.get(
"/discover",
description="Returns a curated list of well-known MCP servers for discovery UI",
dependencies=[Depends(user_api_key_auth)],
)
async def discover_mcp_servers(
query: Optional[str] = Query(
None, description="Search filter for server names and descriptions"
),
category: Optional[str] = Query(
None, description="Filter by category"
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Returns a curated list of well-known MCP servers that can be added to the proxy.
Used by the UI to show a discovery grid when adding new MCP servers.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403,
detail={
"error": "Only proxy admins can access MCP discovery. Your role={}".format(
user_api_key_dict.user_role
)
},
)
registry = _load_mcp_registry()
servers = registry.get("servers", [])
# Apply query filter
if query:
query_lower = query.lower()
servers = [
s
for s in servers
if query_lower in s.get("name", "").lower()
or query_lower in s.get("title", "").lower()
or query_lower in s.get("description", "").lower()
]
# Apply category filter
if category:
servers = [
s for s in servers if s.get("category", "") == category
]
# Extract unique categories from the full list (before filtering)
all_servers = registry.get("servers", [])
categories = sorted(
set(s.get("category", "Other") for s in all_servers)
)
return {
"servers": servers,
"categories": categories,
}

View file

@ -0,0 +1,426 @@
{
"servers": [
{
"name": "github",
"title": "GitHub",
"description": "Manage repos, issues, PRs, and workflows through natural language",
"icon_url": "https://cdn.simpleicons.org/github",
"category": "Developer Tools",
"registry_url": "https://registry.modelcontextprotocol.io/servers/io.github.github%2Fgithub-mcp-server",
"transport": "http",
"url": "https://api.githubcopilot.com/mcp/",
"env_vars": [
{"name": "GITHUB_PERSONAL_ACCESS_TOKEN", "description": "GitHub Personal Access Token", "secret": true}
]
},
{
"name": "gitlab",
"title": "GitLab",
"description": "Official GitLab MCP Server for project and repository management",
"icon_url": "https://cdn.simpleicons.org/gitlab",
"category": "Developer Tools",
"registry_url": "https://registry.modelcontextprotocol.io/servers/com.gitlab%2Fmcp",
"transport": "http",
"url": "https://gitlab.com/api/v4/mcp",
"env_vars": [
{"name": "GITLAB_PERSONAL_ACCESS_TOKEN", "description": "GitLab Personal Access Token", "secret": true}
]
},
{
"name": "atlassian",
"title": "Atlassian (Jira & Confluence)",
"description": "Jira issues, Confluence pages, and Atlassian product integration",
"icon_url": "https://cdn.simpleicons.org/atlassian",
"category": "Developer Tools",
"registry_url": "https://registry.modelcontextprotocol.io/servers/com.atlassian%2Fatlassian-mcp-server",
"transport": "sse",
"url": "https://mcp.atlassian.com/v1/sse",
"env_vars": []
},
{
"name": "linear",
"title": "Linear",
"description": "Issue tracking, project management, and team workflow automation",
"icon_url": "https://cdn.simpleicons.org/linear",
"category": "Developer Tools",
"registry_url": "https://registry.modelcontextprotocol.io/servers/app.linear%2Flinear",
"transport": "sse",
"url": "https://mcp.linear.app/sse",
"env_vars": []
},
{
"name": "sentry",
"title": "Sentry",
"description": "Error monitoring, issue tracking, and debugging for AI assistants",
"icon_url": "https://cdn.simpleicons.org/sentry",
"category": "Developer Tools",
"registry_url": "https://registry.modelcontextprotocol.io/servers/io.github.getsentry%2Fsentry-mcp",
"transport": "stdio",
"command": "npx",
"args": ["-y", "@sentry/mcp-server"],
"env_vars": [
{"name": "SENTRY_ACCESS_TOKEN", "description": "Sentry Access Token", "secret": true}
]
},
{
"name": "slack",
"title": "Slack",
"description": "Channel management, messaging, and Slack workspace integration",
"icon_url": "https://cdn.simpleicons.org/slack",
"category": "Communication",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-slack"],
"env_vars": [
{"name": "SLACK_BOT_TOKEN", "description": "Slack Bot User OAuth Token", "secret": true},
{"name": "SLACK_TEAM_ID", "description": "Slack Team/Workspace ID", "secret": false}
]
},
{
"name": "discord",
"title": "Discord",
"description": "Discord server management, messaging, and bot integration",
"icon_url": "https://cdn.simpleicons.org/discord",
"category": "Communication",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-discord"],
"env_vars": [
{"name": "DISCORD_BOT_TOKEN", "description": "Discord Bot Token", "secret": true}
]
},
{
"name": "postgresql",
"title": "PostgreSQL",
"description": "Query and manage PostgreSQL databases with read-only access",
"icon_url": "https://cdn.simpleicons.org/postgresql",
"category": "Databases",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-postgres"],
"env_vars": [
{"name": "POSTGRES_CONNECTION_STRING", "description": "PostgreSQL connection string (e.g., postgresql://user:pass@host:5432/db)", "secret": true}
]
},
{
"name": "sqlite",
"title": "SQLite",
"description": "Query and manage SQLite databases",
"icon_url": "https://cdn.simpleicons.org/sqlite",
"category": "Databases",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-sqlite"],
"env_vars": [
{"name": "SQLITE_DB_PATH", "description": "Path to SQLite database file", "secret": false}
]
},
{
"name": "mysql",
"title": "MySQL",
"description": "Query and manage MySQL databases",
"icon_url": "https://cdn.simpleicons.org/mysql",
"category": "Databases",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-mysql"],
"env_vars": [
{"name": "MYSQL_HOST", "description": "MySQL host", "secret": false},
{"name": "MYSQL_USER", "description": "MySQL username", "secret": false},
{"name": "MYSQL_PASSWORD", "description": "MySQL password", "secret": true},
{"name": "MYSQL_DATABASE", "description": "MySQL database name", "secret": false}
]
},
{
"name": "mongodb",
"title": "MongoDB",
"description": "Query and manage MongoDB databases and collections",
"icon_url": "https://cdn.simpleicons.org/mongodb",
"category": "Databases",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-mongodb"],
"env_vars": [
{"name": "MONGODB_CONNECTION_STRING", "description": "MongoDB connection string", "secret": true}
]
},
{
"name": "redis",
"title": "Redis",
"description": "Interact with Redis key-value stores",
"icon_url": "https://cdn.simpleicons.org/redis",
"category": "Databases",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-redis"],
"env_vars": [
{"name": "REDIS_URL", "description": "Redis connection URL (e.g., redis://localhost:6379)", "secret": true}
]
},
{
"name": "snowflake",
"title": "Snowflake",
"description": "MCP Server for Snowflake from Snowflake Labs",
"icon_url": "https://cdn.simpleicons.org/snowflake",
"category": "Databases",
"registry_url": "https://registry.modelcontextprotocol.io/servers/io.github.Snowflake-Labs%2Fmcp",
"transport": "stdio",
"command": "uvx",
"args": ["snowflake-labs-mcp"],
"env_vars": [
{"name": "SNOWFLAKE_ACCOUNT", "description": "Snowflake account identifier (e.g., xy12345.us-east-1)", "secret": false},
{"name": "SNOWFLAKE_USER", "description": "Snowflake username", "secret": false},
{"name": "SNOWFLAKE_PASSWORD", "description": "Snowflake password", "secret": true}
]
},
{
"name": "notion",
"title": "Notion",
"description": "Official Notion MCP server for pages and databases",
"icon_url": "https://cdn.simpleicons.org/notion",
"category": "Productivity",
"registry_url": "https://registry.modelcontextprotocol.io/servers/com.notion%2Fmcp",
"transport": "sse",
"url": "https://mcp.notion.com/sse",
"env_vars": []
},
{
"name": "google_drive",
"title": "Google Drive",
"description": "Search and access files in Google Drive",
"icon_url": "https://cdn.simpleicons.org/googledrive",
"category": "Productivity",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-gdrive"],
"env_vars": [
{"name": "GOOGLE_CLIENT_ID", "description": "Google OAuth Client ID", "secret": false},
{"name": "GOOGLE_CLIENT_SECRET", "description": "Google OAuth Client Secret", "secret": true}
]
},
{
"name": "google_calendar",
"title": "Google Calendar",
"description": "Manage events and calendars in Google Calendar",
"icon_url": "https://cdn.simpleicons.org/googlecalendar",
"category": "Productivity",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-google-calendar"],
"env_vars": [
{"name": "GOOGLE_CLIENT_ID", "description": "Google OAuth Client ID", "secret": false},
{"name": "GOOGLE_CLIENT_SECRET", "description": "Google OAuth Client Secret", "secret": true}
]
},
{
"name": "obsidian",
"title": "Obsidian",
"description": "Read, search, and manage Obsidian vault notes and files",
"icon_url": "https://cdn.simpleicons.org/obsidian",
"category": "Productivity",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-obsidian"],
"env_vars": [
{"name": "OBSIDIAN_VAULT_PATH", "description": "Path to Obsidian vault directory", "secret": false}
]
},
{
"name": "brave_search",
"title": "Brave Search",
"description": "Web results, images, videos, and AI summaries via Brave Search API",
"icon_url": "https://cdn.simpleicons.org/brave",
"category": "Search",
"registry_url": "https://registry.modelcontextprotocol.io/servers/io.github.brave%2Fbrave-search-mcp-server",
"transport": "stdio",
"command": "npx",
"args": ["-y", "@brave/brave-search-mcp-server"],
"env_vars": [
{"name": "BRAVE_API_KEY", "description": "Brave Search API Key", "secret": true}
]
},
{
"name": "exa",
"title": "Exa",
"description": "Fast, intelligent web search and web crawling",
"icon_url": "https://cdn.simpleicons.org/exa",
"category": "Search",
"registry_url": "https://registry.modelcontextprotocol.io/servers/ai.exa%2Fexa",
"transport": "http",
"url": "https://mcp.exa.ai/mcp",
"env_vars": [
{"name": "EXA_API_KEY", "description": "Exa API Key", "secret": true}
]
},
{
"name": "tavily",
"title": "Tavily",
"description": "AI-optimized search engine for research and retrieval",
"icon_url": "https://cdn.simpleicons.org/tavily",
"category": "Search",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-tavily"],
"env_vars": [
{"name": "TAVILY_API_KEY", "description": "Tavily API Key", "secret": true}
]
},
{
"name": "puppeteer",
"title": "Puppeteer",
"description": "Browser automation, web scraping, and screenshot capture",
"icon_url": "https://cdn.simpleicons.org/puppeteer",
"category": "Web & Browser",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-puppeteer"],
"env_vars": []
},
{
"name": "playwright",
"title": "Playwright",
"description": "Browser automation and testing with Playwright",
"icon_url": "https://cdn.simpleicons.org/playwright",
"category": "Web & Browser",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-playwright"],
"env_vars": []
},
{
"name": "browserbase",
"title": "Browserbase",
"description": "Cloud browser automation and session management",
"icon_url": "https://cdn.simpleicons.org/browserbase",
"category": "Web & Browser",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-browserbase"],
"env_vars": [
{"name": "BROWSERBASE_API_KEY", "description": "Browserbase API Key", "secret": true},
{"name": "BROWSERBASE_PROJECT_ID", "description": "Browserbase Project ID", "secret": false}
]
},
{
"name": "aws",
"title": "AWS",
"description": "Interact with Amazon Web Services resources and APIs",
"icon_url": "https://cdn.simpleicons.org/amazonaws",
"category": "Cloud",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-aws"],
"env_vars": [
{"name": "AWS_ACCESS_KEY_ID", "description": "AWS Access Key ID", "secret": true},
{"name": "AWS_SECRET_ACCESS_KEY", "description": "AWS Secret Access Key", "secret": true},
{"name": "AWS_REGION", "description": "AWS Region (e.g., us-east-1)", "secret": false}
]
},
{
"name": "cloudflare",
"title": "Cloudflare",
"description": "Manage Cloudflare Workers, KV, R2, D1, and more",
"icon_url": "https://cdn.simpleicons.org/cloudflare",
"category": "Cloud",
"registry_url": "https://registry.modelcontextprotocol.io/servers/com.cloudflare.mcp%2Fmcp",
"transport": "sse",
"url": "https://bindings.mcp.cloudflare.com/sse",
"env_vars": []
},
{
"name": "filesystem",
"title": "Filesystem",
"description": "Read, write, and manage files and directories on disk",
"icon_url": "https://cdn.simpleicons.org/files",
"category": "System",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-filesystem"],
"env_vars": []
},
{
"name": "docker",
"title": "Docker",
"description": "Manage Docker containers, images, and networks",
"icon_url": "https://cdn.simpleicons.org/docker",
"category": "System",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-docker"],
"env_vars": []
},
{
"name": "stripe",
"title": "Stripe",
"description": "Manage payments, customers, and subscriptions via Stripe",
"icon_url": "https://cdn.simpleicons.org/stripe",
"category": "Finance",
"registry_url": "https://registry.modelcontextprotocol.io/servers/com.stripe%2Fmcp",
"transport": "http",
"url": "https://mcp.stripe.com",
"env_vars": []
},
{
"name": "shopify",
"title": "Shopify",
"description": "Manage Shopify stores, products, orders, and customers",
"icon_url": "https://cdn.simpleicons.org/shopify",
"category": "E-Commerce",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-shopify"],
"env_vars": [
{"name": "SHOPIFY_ACCESS_TOKEN", "description": "Shopify Admin API Access Token", "secret": true},
{"name": "SHOPIFY_STORE_URL", "description": "Shopify Store URL (e.g., mystore.myshopify.com)", "secret": false}
]
},
{
"name": "twilio",
"title": "Twilio",
"description": "Send SMS, make calls, and manage communication via Twilio",
"icon_url": "https://cdn.simpleicons.org/twilio",
"category": "Communication",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-twilio"],
"env_vars": [
{"name": "TWILIO_ACCOUNT_SID", "description": "Twilio Account SID", "secret": false},
{"name": "TWILIO_AUTH_TOKEN", "description": "Twilio Auth Token", "secret": true}
]
},
{
"name": "supabase",
"title": "Supabase",
"description": "Manage Supabase projects, databases, and storage",
"icon_url": "https://cdn.simpleicons.org/supabase",
"category": "Databases",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@anthropic/mcp-server-supabase"],
"env_vars": [
{"name": "SUPABASE_URL", "description": "Supabase Project URL", "secret": false},
{"name": "SUPABASE_SERVICE_ROLE_KEY", "description": "Supabase Service Role Key", "secret": true}
]
}
]
}

View file

@ -12,6 +12,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAX_POLICY_ESTIMATE_IMPACT_ROWS
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
@ -85,7 +86,6 @@ def _filter_keys_by_tags(keys: list, tag_patterns: list) -> tuple:
Returns (named_aliases, unnamed_count).
"""
from litellm.proxy.auth.route_checks import RouteChecks
affected: list = []
unnamed_count = 0
@ -111,7 +111,6 @@ def _filter_teams_by_tags(teams: list, tag_patterns: list) -> tuple:
Returns (named_aliases, unnamed_count).
"""
from litellm.proxy.auth.route_checks import RouteChecks
affected: list = []
unnamed_count = 0
@ -141,7 +140,6 @@ async def _find_affected_by_team_patterns(
Returns (new_teams, new_keys, unnamed_keys_count).
"""
from litellm.proxy.auth.route_checks import RouteChecks
new_teams: list = []
matched_team_ids: list = []
@ -178,7 +176,6 @@ async def _find_affected_keys_by_alias(
prisma_client: object, key_patterns: list, existing_keys: list
) -> list:
"""Find keys whose alias matches the given patterns."""
from litellm.proxy.auth.route_checks import RouteChecks
affected: list = []

View file

@ -318,7 +318,7 @@ class ProxyInitializationHelpers:
@click.option(
"--num_workers",
default=DEFAULT_NUM_WORKERS_LITELLM_PROXY,
help="Number of uvicorn / gunicorn workers to spin up. By default, it equals the number of logical CPUs in the system, or 4 workers if that cannot be determined.",
help="Number of uvicorn / gunicorn workers to spin up. Default is 1 (from DEFAULT_NUM_WORKERS_LITELLM_PROXY)",
envvar="NUM_WORKERS",
)
@click.option("--api_base", default=None, help="API base URL.")

View file

@ -393,6 +393,9 @@ from litellm.proxy.management_endpoints.tag_management_endpoints import (
from litellm.proxy.management_endpoints.team_callback_endpoints import (
router as team_callback_router,
)
from litellm.proxy.management_endpoints.access_group_endpoints import (
router as access_group_router,
)
from litellm.proxy.management_endpoints.team_endpoints import router as team_router
from litellm.proxy.management_endpoints.team_endpoints import (
update_team,
@ -1051,98 +1054,236 @@ try:
except FileNotFoundError:
return False
def _validate_ui_directory(ui_path: str) -> bool:
"""
Verify UI directory has minimum required structure.
Checks for:
- Directory exists
- Has index.html (main entry point)
- Has _next directory (Next.js assets)
Returns True if UI directory appears valid and servable.
"""
if not os.path.isdir(ui_path):
return False
# Must have main index.html
if not os.path.exists(os.path.join(ui_path, "index.html")):
return False
# Must have _next directory with Next.js assets
next_dir = os.path.join(ui_path, "_next")
if not os.path.isdir(next_dir):
return False
return True
def _is_ui_pre_restructured(ui_dir: str) -> bool:
"""
Detect if UI directory is already pre-restructured and ready to serve.
Returns True if:
1. Marker file .litellm_ui_ready exists (created by Dockerfile), OR
2. Restructuring pattern detected (subdirectories with index.html inside)
This allows skipping copy/restructure operations on read-only filesystems.
"""
if not os.path.isdir(ui_dir):
return False
# Primary signal: marker file created by Dockerfile
marker_file = os.path.join(ui_dir, ".litellm_ui_ready")
if os.path.exists(marker_file):
verbose_proxy_logger.debug(f"Found UI ready marker: {marker_file}")
return True
# Fallback signal: Detect restructuring pattern
# After restructuring, routes exist as directories with index.html inside
# (e.g., login/index.html instead of login.html)
# Check for main index.html first (basic UI structure requirement)
if not os.path.exists(os.path.join(ui_dir, "index.html")):
return False
# Look for ANY subdirectory with index.html (proves restructuring happened)
# Ignore directories starting with _ (Next.js internals like _next)
try:
for entry in os.scandir(ui_dir):
if entry.is_dir() and not entry.name.startswith("_"):
index_path = os.path.join(entry.path, "index.html")
if os.path.exists(index_path):
# Found at least one restructured route - this proves the pattern
verbose_proxy_logger.debug(
f"Detected restructured UI via pattern: found {entry.name}/index.html"
)
return True
except (PermissionError, OSError) as e:
verbose_proxy_logger.debug(
f"Could not scan {ui_dir} for restructuring detection: {e}"
)
return False
# No restructured routes found
return False
def _try_populate_ui_directory(
source_path: str, target_path: str
) -> tuple[bool, str]:
"""
Attempt to populate target UI directory from source.
Returns: (success: bool, error_message: str)
"""
try:
os.makedirs(target_path, exist_ok=True)
if not _dir_has_content(target_path) and _dir_has_content(source_path):
shutil.copytree(
source_path,
target_path,
dirs_exist_ok=True,
)
verbose_proxy_logger.info(f"Successfully populated UI at {target_path}")
return True, ""
else:
return False, "Source or target directory state invalid"
except (PermissionError, OSError) as e:
return False, str(e)
# Use a writable runtime UI directory whenever possible.
# This prevents mutating the packaged UI directory (e.g. site-packages or the repo checkout)
# and ensures extensionless routes like /ui/login work via <route>/index.html.
is_non_root = os.getenv("LITELLM_NON_ROOT", "").lower() == "true"
# Only use runtime UI path in Docker/non-root environments
# In local development, use the packaged UI directly
# Determine runtime UI path
# Priority: LITELLM_UI_PATH env var > default path based on is_non_root
if is_non_root:
# Use /var/lib/litellm/ui for Docker (more secure than /tmp)
runtime_ui_path = "/var/lib/litellm/ui"
default_runtime_ui_path = "/var/lib/litellm/ui"
else:
default_runtime_ui_path = packaged_ui_path
if _dir_has_content(runtime_ui_path):
runtime_ui_path = os.getenv("LITELLM_UI_PATH", default_runtime_ui_path)
# Validate packaged UI before proceeding
if not _validate_ui_directory(packaged_ui_path):
verbose_proxy_logger.error(
f"Packaged UI at {packaged_ui_path} is invalid or incomplete. "
f"UI may not function correctly."
)
# Decision tree for UI path selection:
# 1. If runtime path == packaged path: use packaged UI directly
# 2. If runtime UI exists and is pre-restructured: use it
# 3. If runtime UI exists but not restructured: use it (will restructure later)
# 4. If runtime UI missing: try to populate from packaged UI
# 4a. If population succeeds: use runtime UI
# 4b. If population fails: fall back to packaged UI
should_use_runtime_path = runtime_ui_path != packaged_ui_path
if should_use_runtime_path:
is_pre_restructured = _is_ui_pre_restructured(runtime_ui_path)
has_content = _dir_has_content(runtime_ui_path)
# Case 2: Runtime UI exists and is ready
if has_content and is_pre_restructured:
verbose_proxy_logger.info(
f"Using pre-built UI for non-root Docker: {runtime_ui_path}"
f"Using pre-restructured UI at {runtime_ui_path}"
)
ui_path = runtime_ui_path
# Case 3: Runtime UI exists but needs restructuring
elif has_content and not is_pre_restructured:
verbose_proxy_logger.warning(
f"UI at {runtime_ui_path} has content but is not properly restructured. "
f"Will attempt to restructure in place."
)
ui_path = runtime_ui_path
# Case 4: Runtime UI missing - try to populate
else:
verbose_proxy_logger.error(
f"UI not found at {runtime_ui_path}. Attempting to populate it from packaged UI."
)
verbose_proxy_logger.error(
f"Path exists: {os.path.exists(runtime_ui_path)}, Has content: {_dir_has_content(runtime_ui_path)}"
verbose_proxy_logger.info(
f"UI not found at {runtime_ui_path}. Attempting to populate from packaged UI."
)
try:
os.makedirs(runtime_ui_path, exist_ok=True)
if not _dir_has_content(runtime_ui_path) and _dir_has_content(
packaged_ui_path
):
shutil.copytree(
packaged_ui_path,
runtime_ui_path,
dirs_exist_ok=True,
)
except Exception as e:
verbose_proxy_logger.exception(
f"Failed to populate runtime UI directory {runtime_ui_path} from {packaged_ui_path}: {e}"
)
success, error = _try_populate_ui_directory(
packaged_ui_path, runtime_ui_path
)
if success:
# Case 4a: Population succeeded
ui_path = runtime_ui_path
else:
if _dir_has_content(runtime_ui_path):
verbose_proxy_logger.info(
f"Using populated UI for non-root Docker: {runtime_ui_path}"
)
ui_path = runtime_ui_path
# Case 4b: Population failed - fall back to packaged UI
verbose_proxy_logger.warning(
f"Failed to populate UI at {runtime_ui_path}: {error}. "
f"Falling back to packaged UI at {packaged_ui_path}. "
f"For read-only deployments, pre-build UI in Dockerfile "
f"or set LITELLM_UI_PATH to a writable emptyDir volume."
)
ui_path = packaged_ui_path
else:
# Local development: use packaged UI directly, no runtime copy needed
verbose_proxy_logger.info(
f"Using packaged UI directory for local development: {packaged_ui_path}"
)
# Case 1: Using packaged UI directly (local development)
verbose_proxy_logger.info(f"Using packaged UI directory: {packaged_ui_path}")
ui_path = packaged_ui_path
# Only modify files if a custom server root path is set
# Validate final UI path
if not _validate_ui_directory(ui_path):
verbose_proxy_logger.error(
f"Selected UI path {ui_path} is invalid or incomplete. UI may not work correctly."
)
# Only modify files if a custom server root path is set AND filesystem is writable
if server_root_path and server_root_path != "/":
# Iterate through files in the UI directory
for root, dirs, files in os.walk(ui_path):
for filename in files:
file_path = os.path.join(root, filename)
# Skip binary files and files that don't need path replacement
if filename.endswith(
(
".png",
".jpg",
".jpeg",
".gif",
".ico",
".woff",
".woff2",
".ttf",
".eot",
)
):
continue
try:
with open(file_path, "r", encoding="utf-8") as f:
content = f.read()
# Check if UI path is writable
is_writable = os.access(ui_path, os.W_OK)
# Replace the asset prefix with the server root path
modified_content = content.replace(
f"{litellm_asset_prefix}",
f"{server_root_path}",
)
if not is_writable:
verbose_proxy_logger.warning(
f"Cannot apply server_root_path replacements to UI at {ui_path}: "
f"path is not writable. Ensure server_root_path is '/' or pre-process "
f"UI files in Dockerfile with custom server_root_path."
)
else:
# Iterate through files in the UI directory
for root, dirs, files in os.walk(ui_path):
for filename in files:
file_path = os.path.join(root, filename)
# Skip binary files and files that don't need path replacement
if filename.endswith(
(
".png",
".jpg",
".jpeg",
".gif",
".ico",
".woff",
".woff2",
".ttf",
".eot",
)
):
continue
try:
with open(file_path, "r", encoding="utf-8") as f:
content = f.read()
# Replace the /.well-known/litellm-ui-config with the server root path
modified_content = modified_content.replace(
"/litellm/.well-known/litellm-ui-config",
f"{server_root_path}/.well-known/litellm-ui-config",
)
# Replace the asset prefix with the server root path
modified_content = content.replace(
f"{litellm_asset_prefix}",
f"{server_root_path}",
)
with open(file_path, "w", encoding="utf-8") as f:
f.write(modified_content)
except UnicodeDecodeError:
# Skip binary files that can't be decoded
continue
# Replace the /.well-known/litellm-ui-config with the server root path
modified_content = modified_content.replace(
"/litellm/.well-known/litellm-ui-config",
f"{server_root_path}/.well-known/litellm-ui-config",
)
with open(file_path, "w", encoding="utf-8") as f:
f.write(modified_content)
except (UnicodeDecodeError, PermissionError, OSError):
# Skip binary files or files we can't write to
continue
# # Mount the _next directory at the root level
app.mount(
@ -1186,14 +1327,22 @@ try:
continue
# Handle HTML file restructuring
# Always restructure the directory we actually serve.
# This is critical for extensionless routes like /ui/login (expects login/index.html).
# In development, we restructure directly in _experimental/out.
# In non-root Docker, we restructure in /var/lib/litellm/ui.
# Only restructure if:
# 1. UI is not already pre-restructured
# 2. Filesystem is writable
try:
if is_non_root and ui_path == "/var/lib/litellm/ui":
is_pre_restructured = _is_ui_pre_restructured(ui_path)
is_writable = os.access(ui_path, os.W_OK)
if is_pre_restructured:
verbose_proxy_logger.info(
f"Skipping runtime UI restructuring for non-root Docker. UI at {ui_path} is pre-restructured."
f"Skipping UI restructuring: {ui_path} is already pre-restructured"
)
elif not is_writable:
verbose_proxy_logger.warning(
f"Cannot restructure UI at {ui_path}: path is not writable. "
f"UI may not work correctly for extensionless routes. "
f"Pre-build and restructure UI in Dockerfile for read-only deployments."
)
else:
_restructure_ui_html_files(ui_path)
@ -4716,8 +4865,10 @@ async def async_assistants_data_generator(
if isinstance(e, HTTPException):
raise e
else:
error_traceback = traceback.format_exc()
error_msg = f"{str(e)}\n\n{error_traceback}"
# Only include the error message, not the traceback.
# The traceback is already logged above via verbose_proxy_logger.exception().
# Including it in the SSE response leaks internal details to clients.
error_msg = str(e)
proxy_exception = ProxyException(
message=getattr(e, "message", error_msg),
@ -4764,7 +4915,7 @@ def _restamp_streaming_chunk_model(
chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None)
)
if not model_mismatch_logged and downstream_model != requested_model_from_client:
verbose_proxy_logger.warning(
verbose_proxy_logger.debug(
"litellm_call_id=%s: streaming chunk model mismatch - requested=%r downstream=%r. Overriding model to requested.",
request_data.get("litellm_call_id"),
requested_model_from_client,
@ -4867,8 +5018,10 @@ async def async_data_generator(
elif isinstance(e, StreamingCallbackError):
error_msg = str(e)
else:
error_traceback = traceback.format_exc()
error_msg = f"{str(e)}\n\n{error_traceback}"
# Only include the error message, not the traceback.
# The traceback is already logged above via verbose_proxy_logger.exception().
# Including it in the SSE response leaks internal details to clients.
error_msg = str(e)
proxy_exception = ProxyException(
message=getattr(e, "message", error_msg),
@ -10294,18 +10447,34 @@ async def get_image():
default_site_logo = os.path.join(current_dir, "logo.jpg")
is_non_root = os.getenv("LITELLM_NON_ROOT", "").lower() == "true"
assets_dir = "/var/lib/litellm/assets" if is_non_root else current_dir
if is_non_root:
os.makedirs(assets_dir, exist_ok=True)
# Determine assets directory
# Priority: LITELLM_ASSETS_PATH env var > default based on is_non_root
default_assets_dir = "/var/lib/litellm/assets" if is_non_root else current_dir
assets_dir = os.getenv("LITELLM_ASSETS_PATH", default_assets_dir)
# Try to create assets_dir if it doesn't exist (simple try/except approach)
if not os.path.exists(assets_dir):
try:
os.makedirs(assets_dir, exist_ok=True)
verbose_proxy_logger.debug(f"Created assets directory at {assets_dir}")
except (PermissionError, OSError) as e:
verbose_proxy_logger.warning(
f"Cannot create assets directory at {assets_dir}: {e}. "
f"Logo caching may not work. Using current directory for assets."
)
assets_dir = current_dir
# Determine default logo path
default_logo = (
os.path.join(assets_dir, "logo.jpg") if is_non_root else default_site_logo
os.path.join(assets_dir, "logo.jpg")
if assets_dir != current_dir
else default_site_logo
)
if is_non_root and not os.path.exists(default_logo):
if assets_dir != current_dir and not os.path.exists(default_logo):
default_logo = default_site_logo
cache_dir = assets_dir if is_non_root else current_dir
cache_dir = assets_dir if os.access(assets_dir, os.W_OK) else current_dir
cache_path = os.path.join(cache_dir, "cached_logo.jpg")
# [OPTIMIZATION] Check if the cached image exists first
@ -11832,6 +12001,7 @@ app.include_router(enterprise_router)
app.include_router(ui_discovery_endpoints_router)
app.include_router(agent_endpoints_router)
app.include_router(a2a_router)
app.include_router(access_group_router)
########################################################
# MCP Server
########################################################

View file

@ -128,6 +128,7 @@ model LiteLLM_TeamTable {
model_max_budget Json @default("{}")
router_settings Json? @default("{}")
team_member_permissions String[] @default([])
access_group_ids String[] @default([])
policies String[] @default([])
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
@ -160,6 +161,7 @@ model LiteLLM_DeletedTeamTable {
model_max_budget Json @default("{}")
router_settings Json? @default("{}")
team_member_permissions String[] @default([])
access_group_ids String[] @default([])
policies String[] @default([])
model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases
@ -291,6 +293,7 @@ model LiteLLM_VerificationToken {
allowed_cache_controls String[] @default([])
allowed_routes String[] @default([])
policies String[] @default([])
access_group_ids String[] @default([])
model_spend Json @default("{}")
model_max_budget Json @default("{}")
budget_id String?
@ -346,6 +349,7 @@ model LiteLLM_DeletedVerificationToken {
allowed_cache_controls String[] @default([])
allowed_routes String[] @default([])
policies String[] @default([])
access_group_ids String[] @default([])
model_spend Json @default("{}")
model_max_budget Json @default("{}")
router_settings Json? @default("{}")
@ -764,6 +768,22 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t
@@index([model_object_id])
}
model LiteLLM_ManagedVectorStoreTable {
id String @id @default(uuid())
unified_resource_id String @unique // The base64 encoded unified vector store ID
resource_object Json? // Stores the VectorStoreCreateResponse
model_mappings Json // Maps model_id -> provider_vector_store_id
flat_model_resource_ids String[] @default([]) // Flat list of provider vector store IDs for faster querying
storage_backend String? // Storage backend name (if applicable)
storage_url String? // Storage URL (if applicable)
created_at DateTime @default(now())
created_by String?
updated_at DateTime @updatedAt
updated_by String?
@@index([unified_resource_id])
}
model LiteLLM_ManagedVectorStoresTable {
vector_store_id String @id
custom_llm_provider String
@ -917,3 +937,23 @@ model LiteLLM_PolicyAttachmentTable {
updated_at DateTime @default(now()) @updatedAt
updated_by String?
}
//Unified Access Groups table for storing unified access groups
model LiteLLM_AccessGroupTable {
access_group_id String @id @default(uuid())
access_group_name String @unique
description String?
// Resource memberships - explicit arrays per type
access_model_ids String[] @default([])
access_mcp_server_ids String[] @default([])
access_agent_ids String[] @default([])
assigned_team_ids String[] @default([])
assigned_key_ids String[] @default([])
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
}

View file

@ -1878,13 +1878,15 @@ async def ui_view_spend_logs( # noqa: PLR0915
verbose_proxy_logger.debug("data= %s", json.dumps(data, indent=4, default=str))
return {
"data": data,
"total": total_records,
"page": page,
"page_size": page_size,
"total_pages": total_pages,
}
return await _build_ui_spend_logs_response(
prisma_client,
data,
total_records,
page,
page_size,
total_pages,
enrich_session_counts=not is_v2,
)
except Exception as e:
verbose_proxy_logger.exception(f"Error in ui_view_spend_logs: {e}")
raise handle_exception_on_proxy(e)
@ -3129,6 +3131,91 @@ async def ui_view_session_spend_logs(
)
async def _build_ui_spend_logs_response(
prisma_client: "PrismaClient",
data: list,
total_records: int,
page: int,
page_size: int,
total_pages: int,
enrich_session_counts: bool = True,
) -> dict:
"""
Build the paginated response for the UI spend-logs endpoint.
When ``enrich_session_counts`` is ``True`` (the default for the v1/UI
endpoint), each row is enriched with ``session_total_count`` so the
frontend knows which sessions are expandable (multi-call sessions).
For every row that carries a ``session_id``, a single ``GROUP BY`` query
fetches the total number of logs in each referenced session. Rows without
a ``session_id`` default to ``1``.
When ``enrich_session_counts`` is ``False`` (v2 endpoint), rows are
serialised without the extra query.
Args:
prisma_client: The connected Prisma client instance.
data: A list of Prisma model instances (must support ``.model_dump()``
and have a ``session_id`` attribute).
total_records: Total number of matching records (for pagination).
page: Current page number.
page_size: Number of items per page.
total_pages: Total number of pages.
enrich_session_counts: Whether to add ``session_total_count`` to each
row. Defaults to ``True``.
Returns:
A dict with ``data`` (enriched rows), ``total``, ``page``,
``page_size``, and ``total_pages``.
"""
count_map: dict[str, int] = {}
if enrich_session_counts:
session_ids = list(
{row.session_id for row in data if getattr(row, "session_id", None)}
)
if session_ids:
# NOTE: This GROUP BY runs on every v1/UI page load. The IN clause
# is bounded by page_size (typically 25-50 distinct session IDs).
# If performance degrades at scale, consider short-lived caching or
# folding the count into the main query via a window function.
counts = await prisma_client.db.litellm_spendlogs.group_by(
by=["session_id"],
where={"session_id": {"in": session_ids}},
count={"session_id": True},
)
count_map = {
r["session_id"]: r["_count"]["session_id"]
for r in counts
if r.get("session_id")
}
if enrich_session_counts:
enriched: List[dict] = []
for row in data:
row_dict = (
dict(row)
if isinstance(row, dict)
else row.model_dump()
)
sid = row_dict.get("session_id")
row_dict["session_total_count"] = count_map.get(sid, 1) if sid else 1
enriched.append(row_dict)
response_data: list = enriched
else:
# v2 path: return raw Prisma model instances so FastAPI applies its
# own Pydantic-aware serialisation (preserves alias handling, custom
# serializers, etc.).
response_data = data # type: ignore[assignment]
return {
"data": response_data,
"total": total_records,
"page": page,
"page_size": page_size,
"total_pages": total_pages,
}
def _build_status_filter_condition(status_filter: Optional[str]) -> Dict[str, Any]:
"""
Helper function to build the status filter condition for database queries.

View file

@ -740,7 +740,7 @@ class ProxyLogging:
self, dynamic_success_callbacks: Optional[List], global_callbacks: List
) -> List:
if dynamic_success_callbacks is None:
return global_callbacks
return list(global_callbacks)
return list(set(dynamic_success_callbacks + global_callbacks))
def _parse_pre_mcp_call_hook_response(

View file

@ -136,10 +136,17 @@ async def vector_store_search(
if "vector_store_id" not in data:
data["vector_store_id"] = vector_store_id
# Check for legacy vector store registry (non-managed vector stores)
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, user_api_key_dict=user_api_key_dict
)
# The managed_vector_stores pre-call hook will handle:
# 1. Decoding managed vector store IDs
# 2. Extracting model and provider resource ID
# 3. Setting up proper routing
# 4. Authentication checks
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
@ -181,6 +188,14 @@ async def vector_store_create(
API Reference:
https://platform.openai.com/docs/api-reference/vector-stores/create
Supports target_model_names parameter for creating vector stores across multiple models:
```json
{
"name": "my-vector-store",
"target_model_names": "gpt-4,gemini-2.0"
}
```
"""
from litellm.proxy.proxy_server import (
_read_request_body,
@ -198,6 +213,47 @@ async def vector_store_create(
)
data = await _read_request_body(request=request)
# Check for target_model_names parameter
target_model_names = data.pop("target_model_names", None)
if target_model_names:
# Use managed vector stores for multi-model support
if isinstance(target_model_names, str):
target_model_names_list = [m.strip() for m in target_model_names.split(",")]
elif isinstance(target_model_names, list):
target_model_names_list = target_model_names
else:
raise HTTPException(
status_code=400,
detail="target_model_names must be a comma-separated string or list of model names",
)
# Get managed vector stores hook
managed_vector_stores = proxy_logging_obj.get_proxy_hook("managed_vector_stores")
if managed_vector_stores is None:
raise HTTPException(
status_code=500,
detail="Managed vector stores not configured. Please ensure the proxy is initialized with database support.",
)
if llm_router is None:
raise HTTPException(
status_code=500,
detail="LLM Router not initialized. Ensure models are added to proxy.",
)
# Create vector store across multiple models
response = await managed_vector_stores.acreate_vector_store(
create_request=data,
llm_router=llm_router,
target_model_names_list=target_model_names_list,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
user_api_key_dict=user_api_key_dict,
)
return response
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(

View file

@ -1,9 +1,9 @@
from typing import Dict, Optional
from typing import TYPE_CHECKING, Dict, Optional
import litellm
from fastapi import APIRouter, Depends, Request, Response
from fastapi.responses import ORJSONResponse
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -12,18 +12,252 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
handle_model_based_routing,
prepare_data_with_credentials,
)
from litellm.proxy.vector_store_endpoints.utils import (
is_allowed_to_call_vector_store_files_endpoint,
)
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.router import Router
router = APIRouter()
def _update_request_data_with_managed_file_id(
data: Dict,
file_id: str,
request: Request,
llm_router: Optional["Router"] = None,
) -> tuple[Dict, Optional[str]]:
"""
Update request data with model routing information from managed file ID.
This function handles two types of file IDs:
1. Simple encoded file IDs (format: litellm:{file_id};model,{model})
2. Unified managed file IDs (format: litellm_proxy:{mime};unified_id,{uuid};...;llm_output_file_id,{file_id};...)
For unified managed file IDs, it:
- Decodes the unified ID to extract the actual provider file ID (llm_output_file_id)
- Extracts the model routing information (target_model_names)
- Updates data with credentials for the correct deployment
Args:
data: Request data to update
file_id: File ID (can be managed/encoded or regular)
request: FastAPI request object
llm_router: LiteLLM router for credential lookup (required for managed files)
Returns:
Tuple of (updated request data, original_managed_file_id)
- original_managed_file_id is the original file_id if it was managed/encoded, None otherwise
"""
import re
from litellm import verbose_logger
from litellm.llms.base_llm.managed_resources.utils import (
is_base64_encoded_unified_id,
parse_unified_id,
)
# First, check if this is a unified managed file ID (base64 encoded)
decoded_id = is_base64_encoded_unified_id(file_id)
if decoded_id:
# This is a unified managed file ID
verbose_logger.debug(
f"Processing unified managed file ID: {file_id}"
)
# Parse the unified ID to extract components
parsed_id = parse_unified_id(file_id)
if parsed_id:
target_model_names = parsed_id.get("target_model_names", [])
# Extract the actual provider file ID from llm_output_file_id field
# Format: litellm_proxy:...;llm_output_file_id,{actual_file_id};...
llm_output_file_id = None
try:
match = re.search(r"llm_output_file_id,([^;]+)", decoded_id)
if match:
llm_output_file_id = match.group(1).strip()
except Exception:
pass
verbose_logger.debug(
f"Decoded unified file ID - target_model_names: {target_model_names}, llm_output_file_id: {llm_output_file_id}"
)
# Set the model for routing
if target_model_names and len(target_model_names) > 0:
routing_model = target_model_names[0]
data["model"] = routing_model
# Get credentials for the model
if llm_router:
credentials = llm_router.get_deployment_credentials_with_provider(
model_id=routing_model
)
if credentials:
prepare_data_with_credentials(
data=data,
credentials=credentials,
file_id=llm_output_file_id, # Use the actual provider file ID
)
verbose_logger.info(
f"Routing vector store file operation to model: {routing_model}, file_id: {file_id} -> {llm_output_file_id}"
)
return data, file_id # Return original managed file ID
# If we extracted the provider file ID but no routing, still use it
if llm_output_file_id:
data["file_id"] = llm_output_file_id
verbose_logger.debug(
f"Replaced unified file ID with provider file ID: {llm_output_file_id}"
)
return data, file_id # Return original managed file ID
return data, file_id if decoded_id else None
# Fall back to simple encoded file ID handling (format: litellm:{file_id};model,{model})
should_route, model_used, original_file_id, credentials = handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
check_file_id_encoding=True,
)
if should_route:
# Use model-based routing with credentials from config
prepare_data_with_credentials(
data=data,
credentials=credentials, # type: ignore
file_id=original_file_id, # Use decoded file ID if from encoded ID
)
verbose_logger.debug(
f"Routing vector store file operation using model: {model_used}"
+ (f", file_id: {file_id} -> {original_file_id}" if original_file_id else "")
)
return data, file_id # Return original file ID for response replacement
return data, None
def _replace_file_id_in_response(response, original_file_id: str):
"""
Replace the provider file ID in the response with the original managed file ID.
This ensures that when a user sends a managed file ID, they get back the same
managed file ID in the response, not the decoded provider file ID.
Args:
response: The response object from the provider
original_file_id: The original managed file ID to restore
Returns:
Modified response with original file ID
"""
if response is None:
return response
# Handle different response types
if isinstance(response, dict):
# For dict responses (e.g., VectorStoreFileDeleteResponse)
if "id" in response:
response["id"] = original_file_id
if "file_id" in response:
response["file_id"] = original_file_id
elif hasattr(response, "id"):
# For object responses (e.g., VectorStoreFileObject)
response.id = original_file_id
elif hasattr(response, "file_id"):
response.file_id = original_file_id
return response
def _update_request_data_with_litellm_managed_vector_store_registry(
data: Dict,
vector_store_id: str,
llm_router: Optional["Router"] = None,
) -> Dict:
"""
Update request data with model routing information from managed vector store.
This function handles two types of vector stores:
1. Legacy vector stores from registry (non-managed)
2. Managed vector stores with unified IDs (requires decoding)
For managed vector stores, this function:
- Decodes the unified vector store ID
- Extracts the model_id and provider resource ID
- Sets data["model"] so the router can use the correct deployment credentials
- Replaces the unified ID with the provider-specific ID
Args:
data: Request data to update
vector_store_id: Vector store ID (can be unified or legacy)
llm_router: LiteLLM router for credential lookup (required for managed vector stores)
Returns:
Updated request data with model routing information
"""
from litellm import verbose_logger
from litellm.llms.base_llm.managed_resources.utils import (
is_base64_encoded_unified_id,
parse_unified_id,
)
# Check if this is a managed vector store ID (base64 encoded unified ID)
decoded_id = is_base64_encoded_unified_id(vector_store_id)
if decoded_id:
# This is a managed vector store - decode and extract routing information
verbose_logger.debug(
f"Processing managed vector store ID: {vector_store_id}"
)
parsed_id = parse_unified_id(vector_store_id)
if parsed_id:
model_id = parsed_id.get("model_id")
provider_resource_id = parsed_id.get("provider_resource_id")
target_model_names = parsed_id.get("target_model_names", [])
verbose_logger.debug(
f"Decoded vector store - model_id: {model_id}, provider_resource_id: {provider_resource_id}, target_model_names: {target_model_names}"
)
# Set the model for routing - this tells the router which deployment to use
# The router will automatically get the credentials from the deployment
routing_model = None
if model_id:
routing_model = model_id
elif target_model_names and len(target_model_names) > 0:
routing_model = target_model_names[0]
if routing_model:
data["model"] = routing_model
verbose_logger.info(
f"Routing vector store files operation to model: {routing_model}"
)
# Replace unified vector store ID with provider resource ID
if provider_resource_id:
data["vector_store_id"] = provider_resource_id
verbose_logger.debug(
f"Replaced unified vector store ID with provider resource ID: {provider_resource_id}"
)
return data
# Legacy path: Check vector store registry for non-managed vector stores
if litellm.vector_store_registry is not None:
vector_store_to_run = (
litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
@ -42,6 +276,7 @@ def _update_request_data_with_litellm_managed_vector_store_registry(
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
return data
@ -128,8 +363,16 @@ async def vector_store_file_create(
if "vector_store_id" not in data:
data["vector_store_id"] = vector_store_id
# Handle managed file IDs if present in request body
original_managed_file_id = None
if "file_id" in data:
data, original_managed_file_id = _update_request_data_with_managed_file_id(
data=data, file_id=data["file_id"], request=request, llm_router=llm_router
)
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id
data=data, vector_store_id=vector_store_id, llm_router=llm_router
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -145,7 +388,7 @@ async def vector_store_file_create(
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -163,6 +406,12 @@ async def vector_store_file_create(
user_api_base=user_api_base,
version=version,
)
# Replace provider file ID with original managed file ID in response
if original_managed_file_id:
response = _replace_file_id_in_response(response, original_managed_file_id)
return response
except Exception as e: # noqa: BLE001
raise await processor._handle_llm_api_exception(
e=e,
@ -209,7 +458,7 @@ async def vector_store_file_list(
data.update(query_params)
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id
data=data, vector_store_id=vector_store_id, llm_router=llm_router
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -290,8 +539,14 @@ async def vector_store_file_retrieve(
"file_id": file_id,
}
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
data=data, file_id=file_id, request=request, llm_router=llm_router
)
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id
data=data, vector_store_id=vector_store_id, llm_router=llm_router
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -307,7 +562,7 @@ async def vector_store_file_retrieve(
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -325,6 +580,12 @@ async def vector_store_file_retrieve(
user_api_base=user_api_base,
version=version,
)
# Replace provider file ID with original managed file ID in response
if original_managed_file_id:
response = _replace_file_id_in_response(response, original_managed_file_id)
return response
except Exception as e: # noqa: BLE001
raise await processor._handle_llm_api_exception(
e=e,
@ -372,8 +633,14 @@ async def vector_store_file_content(
"file_id": file_id,
}
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
data=data, file_id=file_id, request=request, llm_router=llm_router
)
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id
data=data, vector_store_id=vector_store_id, llm_router=llm_router
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -389,7 +656,7 @@ async def vector_store_file_content(
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -407,6 +674,12 @@ async def vector_store_file_content(
user_api_base=user_api_base,
version=version,
)
# Replace provider file ID with original managed file ID in response
if original_managed_file_id:
response = _replace_file_id_in_response(response, original_managed_file_id)
return response
except Exception as e: # noqa: BLE001
raise await processor._handle_llm_api_exception(
e=e,
@ -454,8 +727,14 @@ async def vector_store_file_update(
data["vector_store_id"] = vector_store_id
data["file_id"] = file_id
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
data=data, file_id=file_id, request=request, llm_router=llm_router
)
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id
data=data, vector_store_id=vector_store_id, llm_router=llm_router
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -471,7 +750,7 @@ async def vector_store_file_update(
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -489,6 +768,12 @@ async def vector_store_file_update(
user_api_base=user_api_base,
version=version,
)
# Replace provider file ID with original managed file ID in response
if original_managed_file_id:
response = _replace_file_id_in_response(response, original_managed_file_id)
return response
except Exception as e: # noqa: BLE001
raise await processor._handle_llm_api_exception(
e=e,
@ -536,8 +821,14 @@ async def vector_store_file_delete(
"file_id": file_id,
}
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
data=data, file_id=file_id, request=request, llm_router=llm_router
)
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id
data=data, vector_store_id=vector_store_id, llm_router=llm_router
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -553,7 +844,7 @@ async def vector_store_file_delete(
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -571,6 +862,12 @@ async def vector_store_file_delete(
user_api_base=user_api_base,
version=version,
)
# Replace provider file ID with original managed file ID in response
if original_managed_file_id:
response = _replace_file_id_in_response(response, original_managed_file_id)
return response
except Exception as e: # noqa: BLE001
raise await processor._handle_llm_api_exception(
e=e,

View file

@ -4,11 +4,17 @@ RAG Ingestion classes for different providers.
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion
from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion
from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion
from litellm.rag.ingestion.s3_vectors_ingestion import S3VectorsRAGIngestion
from litellm.rag.ingestion.vertex_ai_ingestion import VertexAIRAGIngestion
__all__ = [
"BaseRAGIngestion",
"BedrockRAGIngestion",
"GeminiRAGIngestion",
"OpenAIRAGIngestion",
"S3VectorsRAGIngestion",
"VertexAIRAGIngestion",
]

View file

@ -0,0 +1,478 @@
"""
Vertex AI-specific RAG Ingestion implementation.
Vertex AI RAG Engine handles embedding and chunking internally when files are uploaded,
so this implementation skips the embedding step and directly uploads files to RAG corpora.
Based on: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/model-reference/rag-api-v1
"""
from __future__ import annotations
import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
if TYPE_CHECKING:
from litellm import Router
from litellm.types.rag import RAGIngestOptions
class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
"""
Vertex AI RAG Engine ingestion implementation.
Key differences from base:
- Embedding is handled by Vertex AI RAG Engine when files are uploaded
- Files are uploaded using the RAG API (import or upload)
- Chunking is done by Vertex AI RAG Engine (supports custom chunking config)
- Supports Google Cloud Storage (GCS) and Google Drive sources
- Supports custom parsing configurations (layout parser, LLM parser)
"""
def __init__(
self,
ingest_options: "RAGIngestOptions",
router: Optional["Router"] = None,
):
BaseRAGIngestion.__init__(self, ingest_options=ingest_options, router=router)
VertexBase.__init__(self)
# Extract Vertex AI specific configs from vector_store_config
litellm_params = dict(self.vector_store_config)
# Get project, location, and credentials using VertexBase methods
self.project_id = self.safe_get_vertex_ai_project(litellm_params)
self.location = self.get_vertex_ai_location(litellm_params) or "us-central1"
self.vertex_credentials = self.safe_get_vertex_ai_credentials(litellm_params)
async def embed(
self,
chunks: List[str],
) -> Optional[List[List[float]]]:
"""
Vertex AI RAG Engine handles embedding internally - skip this step.
Returns:
None (Vertex AI embeds when files are uploaded to RAG corpus)
"""
# Vertex AI RAG Engine handles embedding when files are uploaded
return None
async def store(
self,
file_content: Optional[bytes],
filename: Optional[str],
content_type: Optional[str],
chunks: List[str],
embeddings: Optional[List[List[float]]],
) -> Tuple[Optional[str], Optional[str]]:
"""
Store content in Vertex AI RAG corpus.
Vertex AI workflow:
1. Create RAG corpus (if not provided)
2. Upload file using RAG API (Vertex AI handles chunking/embedding)
Args:
file_content: Raw file bytes
filename: Name of the file
content_type: MIME type
chunks: Ignored - Vertex AI handles chunking
embeddings: Ignored - Vertex AI handles embedding
Returns:
Tuple of (rag_corpus_id, file_id)
"""
if not self.project_id:
raise ValueError(
"vertex_project is required for Vertex AI RAG ingestion. "
"Set it in vector_store config."
)
# Get or create RAG corpus
rag_corpus_id = self.vector_store_config.get("vector_store_id")
if not rag_corpus_id:
rag_corpus_id = await self._create_rag_corpus(
display_name=self.ingest_name or "litellm-rag-corpus",
description=self.vector_store_config.get("description"),
)
# Upload file to RAG corpus
result_file_id = None
if file_content and filename and rag_corpus_id:
result_file_id = await self._upload_file_to_corpus(
rag_corpus_id=rag_corpus_id,
filename=filename,
file_content=file_content,
content_type=content_type,
)
return rag_corpus_id, result_file_id
async def _create_rag_corpus(
self,
display_name: str,
description: Optional[str] = None,
) -> str:
"""
Create a Vertex AI RAG corpus.
Args:
display_name: Display name for the corpus
description: Optional description
Returns:
RAG corpus ID (format: projects/{project}/locations/{location}/ragCorpora/{corpus_id})
"""
# Get access token using VertexBase method
access_token, project_id = self._ensure_access_token(
credentials=self.vertex_credentials,
project_id=self.project_id,
custom_llm_provider="vertex_ai",
)
# Use the project_id from token if not set
if not self.project_id:
self.project_id = project_id
# Construct URL using vertex base URL helper
base_url = get_vertex_base_url(self.location)
url = (
f"{base_url}/v1beta1/"
f"projects/{self.project_id}/locations/{self.location}/ragCorpora"
)
# Build request body with camelCase keys (Vertex AI API format)
request_body: Dict[str, Any] = {
"displayName": display_name,
}
if description:
request_body["description"] = description
# Add vector database config if specified
vector_db_config = self.vector_store_config.get("vector_db_config")
if vector_db_config:
request_body["vectorDbConfig"] = vector_db_config
# Add embedding model config if specified
embedding_model = self.vector_store_config.get("embedding_model")
if embedding_model:
if "vectorDbConfig" not in request_body:
request_body["vectorDbConfig"] = {}
request_body["vectorDbConfig"]["ragEmbeddingModelConfig"] = {
"vertexPredictionEndpoint": {
"endpoint": embedding_model
}
}
verbose_logger.debug(f"Creating RAG corpus: {url}")
verbose_logger.debug(f"Request body: {json.dumps(request_body, indent=2)}")
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.RAG,
params={"timeout": 60.0},
)
response = await client.post(
url,
json=request_body,
headers={
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
},
)
if response.status_code not in [200, 201]:
error_msg = f"Failed to create RAG corpus: {response.text}"
verbose_logger.error(error_msg)
raise Exception(error_msg)
response_data = response.json()
verbose_logger.debug(f"Create corpus response: {json.dumps(response_data, indent=2)}")
# The response is a long-running operation
# Check if it's already done or if we need to poll
if response_data.get("done"):
# Operation completed immediately
corpus_name = response_data.get("response", {}).get("name", "")
else:
# Need to poll the operation
operation_name = response_data.get("name", "")
verbose_logger.debug(f"Polling operation: {operation_name}")
corpus_name = await self._poll_operation(
operation_name=operation_name,
access_token=access_token,
)
verbose_logger.debug(f"Created RAG corpus: {corpus_name}")
return corpus_name
async def _poll_operation(
self,
operation_name: str,
access_token: str,
max_retries: int = 30,
retry_delay: float = 2.0,
) -> str:
"""
Poll a long-running operation until it completes.
Args:
operation_name: The operation name (e.g., "operations/123456")
access_token: Access token for authentication
max_retries: Maximum number of polling attempts
retry_delay: Delay between polling attempts in seconds
Returns:
The corpus name from the completed operation
Raises:
Exception: If operation fails or times out
"""
import asyncio
base_url = get_vertex_base_url(self.location)
# Operation name is like: projects/{project}/locations/{location}/operations/{operation_id}
# We need to construct the full URL
url = f"{base_url}/v1beta1/{operation_name}"
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.RAG,
params={"timeout": 60.0},
)
for attempt in range(max_retries):
response = await client.get(
url,
headers={
"Authorization": f"Bearer {access_token}",
},
)
if response.status_code != 200:
error_msg = f"Failed to poll operation: {response.text}"
verbose_logger.error(error_msg)
raise Exception(error_msg)
operation_data = response.json()
if operation_data.get("done"):
# Check for errors
if "error" in operation_data:
error = operation_data["error"]
raise Exception(f"Operation failed: {error}")
# Extract corpus name from response
corpus_name = operation_data.get("response", {}).get("name", "")
if corpus_name:
return corpus_name
else:
raise Exception(f"No corpus name in operation response: {operation_data}")
verbose_logger.debug(f"Operation not done yet, attempt {attempt + 1}/{max_retries}")
await asyncio.sleep(retry_delay)
raise Exception(f"Operation timed out after {max_retries} attempts")
async def _upload_file_to_corpus(
self,
rag_corpus_id: str,
filename: str,
file_content: bytes,
content_type: Optional[str],
) -> str:
"""
Upload a file to Vertex AI RAG corpus using multipart upload.
Args:
rag_corpus_id: RAG corpus resource name
filename: Name of the file
file_content: File content bytes
content_type: MIME type
Returns:
File ID or resource name
"""
# Get access token using VertexBase method
access_token, _ = self._ensure_access_token(
credentials=self.vertex_credentials,
project_id=self.project_id,
custom_llm_provider="vertex_ai",
)
# Construct upload URL using vertex base URL helper
base_url = get_vertex_base_url(self.location)
url = (
f"{base_url}/upload/v1beta1/"
f"{rag_corpus_id}/ragFiles:upload"
)
# Build metadata for the file with snake_case keys (as per upload API docs)
metadata: Dict[str, Any] = {
"rag_file": {
"display_name": filename,
}
}
# Add description if provided
description = self.vector_store_config.get("file_description")
if description:
metadata["rag_file"]["description"] = description
# Add chunking configuration if provided
chunking_strategy = self.chunking_strategy
if chunking_strategy and isinstance(chunking_strategy, dict):
chunk_size = chunking_strategy.get("chunk_size")
chunk_overlap = chunking_strategy.get("chunk_overlap")
if chunk_size or chunk_overlap:
if "upload_rag_file_config" not in metadata:
metadata["upload_rag_file_config"] = {}
metadata["upload_rag_file_config"]["rag_file_transformation_config"] = {
"rag_file_chunking_config": {
"fixed_length_chunking": {}
}
}
chunking_config = metadata["upload_rag_file_config"][
"rag_file_transformation_config"
]["rag_file_chunking_config"]["fixed_length_chunking"]
if chunk_size:
chunking_config["chunk_size"] = chunk_size
if chunk_overlap:
chunking_config["chunk_overlap"] = chunk_overlap
verbose_logger.debug(f"Uploading file to RAG corpus: {url}")
verbose_logger.debug(f"Metadata: {json.dumps(metadata, indent=2)}")
# Prepare multipart form data
files = {
"metadata": (None, json.dumps(metadata), "application/json"),
"file": (filename, file_content, content_type or "application/octet-stream"),
}
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.RAG,
params={"timeout": 300.0}, # Longer timeout for large files
)
response = await client.post(
url,
files=files,
headers={
"Authorization": f"Bearer {access_token}",
"X-Goog-Upload-Protocol": "multipart",
},
)
if response.status_code not in [200, 201]:
error_msg = f"Failed to upload file: {response.text}"
verbose_logger.error(error_msg)
raise Exception(error_msg)
# Parse response to get file ID
try:
response_data = response.json()
# The response should contain the rag_file resource name
file_id = response_data.get("ragFile", {}).get("name", "")
if not file_id:
file_id = response_data.get("name", "")
verbose_logger.debug(f"Upload complete. File ID: {file_id}")
return file_id
except Exception as e:
verbose_logger.warning(f"Could not parse upload response: {e}")
return "uploaded"
async def _import_files_from_gcs(
self,
rag_corpus_id: str,
gcs_uris: List[str],
) -> str:
"""
Import files from Google Cloud Storage into RAG corpus.
Args:
rag_corpus_id: RAG corpus resource name
gcs_uris: List of GCS URIs (e.g., ["gs://bucket/file.pdf"])
Returns:
Operation name for tracking import progress
"""
# Get access token using VertexBase method
access_token, _ = self._ensure_access_token(
credentials=self.vertex_credentials,
project_id=self.project_id,
custom_llm_provider="vertex_ai",
)
# Construct import URL using vertex base URL helper
base_url = get_vertex_base_url(self.location)
url = (
f"{base_url}/v1beta1/"
f"{rag_corpus_id}/ragFiles:import"
)
# Build request body with camelCase keys (Vertex AI API format)
request_body: Dict[str, Any] = {
"importRagFilesConfig": {
"gcsSource": {
"uris": gcs_uris
}
}
}
# Add chunking configuration if provided
chunking_strategy = self.chunking_strategy
if chunking_strategy and isinstance(chunking_strategy, dict):
chunk_size = chunking_strategy.get("chunk_size")
chunk_overlap = chunking_strategy.get("chunk_overlap")
if chunk_size or chunk_overlap:
request_body["importRagFilesConfig"]["ragFileChunkingConfig"] = {
"chunkSize": chunk_size or 1024,
"chunkOverlap": chunk_overlap or 200,
}
# Add max embedding requests per minute if specified
max_embedding_qpm = self.vector_store_config.get("max_embedding_requests_per_min")
if max_embedding_qpm:
request_body["importRagFilesConfig"]["maxEmbeddingRequestsPerMin"] = max_embedding_qpm
verbose_logger.debug(f"Importing files from GCS: {url}")
verbose_logger.debug(f"Request body: {json.dumps(request_body, indent=2)}")
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.RAG,
params={"timeout": 60.0},
)
response = await client.post(
url,
json=request_body,
headers={
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
},
)
if response.status_code not in [200, 201]:
error_msg = f"Failed to import files: {response.text}"
verbose_logger.error(error_msg)
raise Exception(error_msg)
response_data = response.json()
operation_name = response_data.get("name", "")
verbose_logger.debug(f"Import operation started: {operation_name}")
return operation_name

View file

@ -32,6 +32,7 @@ from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion
from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion
from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion
from litellm.rag.ingestion.s3_vectors_ingestion import S3VectorsRAGIngestion
from litellm.rag.ingestion.vertex_ai_ingestion import VertexAIRAGIngestion
from litellm.rag.rag_query import RAGQuery
from litellm.types.rag import (
RAGIngestOptions,
@ -50,6 +51,7 @@ INGESTION_REGISTRY: Dict[str, Type[BaseRAGIngestion]] = {
"bedrock": BedrockRAGIngestion,
"gemini": GeminiRAGIngestion,
"s3_vectors": S3VectorsRAGIngestion,
"vertex_ai": VertexAIRAGIngestion,
}

Some files were not shown because too many files have changed in this diff Show more