Merge remote-tracking branch 'origin/litellm_internal_staging' into codex/cloud-storage-file-guard

# Conflicts:
#	litellm/llms/vertex_ai/files/handler.py
#	litellm/llms/vertex_ai/files/transformation.py
This commit is contained in:
user 2026-05-04 11:25:41 -07:00
commit 7c94149aeb
200 changed files with 18182 additions and 1436 deletions

View file

@ -4,7 +4,7 @@ on:
workflow_dispatch:
inputs:
tag:
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0-dev.2, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
required: true
type: string
commit_hash:
@ -46,9 +46,11 @@ jobs:
const commitHash = process.env.COMMIT_HASH;
// Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases.
// Accept both PEP 440 (`.dev`) and SemVer (`-dev`) separators so tags
// like `1.84.0.dev2` and `1.84.0-dev.2` are both detected.
// PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]`
// are stable maintenance releases, not pre-releases.
const isPrerelease = /(?:rc|nightly|alpha|beta|\.dev)/i.test(tag);
const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag);
const cosignSection = [
`## Verify Docker Image Signature`,

View file

@ -1,9 +1,9 @@
# Base image for building
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin

View file

@ -3,7 +3,7 @@ ARG LITELLM_BUILD_IMAGE=python:3.11-alpine@sha256:f07e2ace46f560f09a6eeec7b4913b
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=python:3.11-alpine@sha256:f07e2ace46f560f09a6eeec7b4913b80ee99546e749ef82342a419a326620856
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin

View file

@ -1,9 +1,9 @@
# Base image for building
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin

View file

@ -3,7 +3,7 @@ ARG LITELLM_BUILD_IMAGE=python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973a
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973afcbd606dbf021a589cab77d6b00b579d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin

View file

@ -1,4 +1,4 @@
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
FROM python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973afcbd606dbf021a589cab77d6b00b579d

View file

@ -1,8 +1,8 @@
# Base images
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
ARG PROXY_EXTRAS_SOURCE=published
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin

View file

@ -0,0 +1,196 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Crusoe
## Overview
| Property | Details |
|-------|-------|
| Description | Crusoe Cloud provides GPU-accelerated inference for open-source large language models, optimized for performance and cost efficiency. |
| Provider Route on LiteLLM | `crusoe/` |
| Link to Provider Doc | [Crusoe Managed Inference Documentation ↗](https://docs.crusoecloud.com/managed-inference/overview/index.html) |
| Base URL | `https://managed-inference-api-proxy.crusoecloud.com/v1` |
| Supported Operations | [`/chat/completions`](#sample-usage) |
<br />
<br />
**We support ALL Crusoe models, just set `crusoe/` as a prefix when sending completion requests**
## Available Models
| Model | Description | Context Window |
|-------|-------------|----------------|
| `crusoe/deepseek-ai/DeepSeek-R1-0528` | DeepSeek R1 reasoning model (May 2025) | 163,840 tokens |
| `crusoe/deepseek-ai/DeepSeek-V3-0324` | DeepSeek V3 chat model (March 2025) | 163,840 tokens |
| `crusoe/google/gemma-3-12b-it` | Google Gemma 3 12B instruction-tuned | 131,072 tokens |
| `crusoe/meta-llama/Llama-3.3-70B-Instruct` | Llama 3.3 70B instruction-tuned | 131,072 tokens |
| `crusoe/moonshotai/Kimi-K2-Thinking` | Kimi K2 extended thinking model | 262,144 tokens |
| `crusoe/openai/gpt-oss-120b` | OpenAI 120B open-source model | 131,072 tokens |
| `crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507` | Qwen3 235B MoE instruction-tuned | 262,144 tokens |
## Required Variables
```python showLineNumbers title="Environment Variables"
os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key
```
## Usage - LiteLLM Python SDK
### Non-streaming
```python showLineNumbers title="Crusoe Non-streaming Completion"
import os
import litellm
from litellm import completion
os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key
messages = [{"content": "Hello, how are you?", "role": "user"}]
# Crusoe call
response = completion(
model="crusoe/meta-llama/Llama-3.3-70B-Instruct",
messages=messages
)
print(response)
```
### Streaming
```python showLineNumbers title="Crusoe Streaming Completion"
import os
import litellm
from litellm import completion
os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key
messages = [{"content": "Write a short story about AI", "role": "user"}]
# Crusoe call with streaming
response = completion(
model="crusoe/meta-llama/Llama-3.3-70B-Instruct",
messages=messages,
stream=True
)
for chunk in response:
print(chunk)
```
### Function Calling
```python showLineNumbers title="Crusoe Function Calling"
import os
import litellm
from litellm import completion
os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key
tools = [{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather in a location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA"
}
},
"required": ["location"]
}
}
}]
messages = [{"role": "user", "content": "What's the weather in Boston?"}]
response = completion(
model="crusoe/meta-llama/Llama-3.3-70B-Instruct",
messages=messages,
tools=tools,
tool_choice="auto"
)
print(response)
```
## Usage - LiteLLM Proxy Server
```yaml showLineNumbers title="config.yaml"
model_list:
- model_name: llama-3.3-70b
litellm_params:
model: crusoe/meta-llama/Llama-3.3-70B-Instruct
api_key: os.environ/CRUSOE_API_KEY
- model_name: deepseek-r1
litellm_params:
model: crusoe/deepseek-ai/DeepSeek-R1-0528
api_key: os.environ/CRUSOE_API_KEY
- model_name: deepseek-v3
litellm_params:
model: crusoe/deepseek-ai/DeepSeek-V3-0324
api_key: os.environ/CRUSOE_API_KEY
- model_name: qwen3-235b
litellm_params:
model: crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507
api_key: os.environ/CRUSOE_API_KEY
- model_name: kimi-k2
litellm_params:
model: crusoe/moonshotai/Kimi-K2-Thinking
api_key: os.environ/CRUSOE_API_KEY
```
## Custom API Base
**Option 1: Environment variable**
```python showLineNumbers title="Custom API Base via env var"
import os
from litellm import completion
os.environ["CRUSOE_API_BASE"] = "https://custom.crusoecloud.com/v1"
os.environ["CRUSOE_API_KEY"] = "" # your API key
response = completion(
model="crusoe/meta-llama/Llama-3.3-70B-Instruct",
messages=[{"content": "Hello!", "role": "user"}],
)
```
**Option 2: Pass directly**
```python showLineNumbers title="Custom API Base via parameter"
from litellm import completion
response = completion(
model="crusoe/meta-llama/Llama-3.3-70B-Instruct",
messages=[{"content": "Hello!", "role": "user"}],
api_base="https://custom.crusoecloud.com/v1",
api_key="your-api-key",
)
```
## Supported OpenAI Parameters
- `temperature`
- `max_tokens`
- `max_completion_tokens`
- `top_p`
- `frequency_penalty`
- `presence_penalty`
- `stop`
- `n`
- `stream`
- `tools`
- `tool_choice`
- `response_format`
- `seed`
- `user`
- `logit_bias`
- `logprobs`
- `top_logprobs`

View file

@ -10,28 +10,21 @@ has already authenticated the user) and you need to extract user information fro
custom headers or other request attributes.
"""
from typing import TYPE_CHECKING, Dict, Optional, Union, cast
from typing import cast
from fastapi import Request
from fastapi.responses import RedirectResponse
if TYPE_CHECKING:
from fastapi_sso.sso.base import OpenID
else:
from typing import Any as OpenID
from litellm.proxy.management_endpoints.types import CustomOpenID
class EnterpriseCustomSSOHandler:
"""
Enterprise Custom SSO Handler for LiteLLM Proxy
This class provides methods for handling custom SSO authentication flows
where users can implement their own authentication logic by processing
request headers and returning user information in OpenID format.
"""
@staticmethod
async def handle_custom_ui_sso_sign_in(
request: Request,
@ -40,16 +33,16 @@ class EnterpriseCustomSSOHandler:
Allow a user to execute their custom code to parse incoming request headers and return a OpenID object
Use this when you have an OAuth proxy in front of LiteLLM (where the OAuth proxy has already authenticated the user)
Args:
request: The FastAPI request object containing headers and other request data
Returns:
RedirectResponse: Redirect response that sends the user to the LiteLLM UI with authentication token
Raises:
ValueError: If custom_ui_sso_sign_in_handler is not configured
Example:
This method is typically called when a user has already been authenticated by an
external OAuth proxy and the proxy has added custom headers containing user information.
@ -60,27 +53,44 @@ class EnterpriseCustomSSOHandler:
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
from litellm.proxy.proxy_server import (
CommonProxyErrors,
general_settings,
premium_user,
user_custom_ui_sso_sign_in_handler,
)
from litellm.proxy.auth.trusted_proxy_utils import (
require_trusted_proxy_request,
)
if premium_user is not True:
raise ValueError(CommonProxyErrors.not_premium_user.value)
if user_custom_ui_sso_sign_in_handler is None:
raise ValueError("custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings.")
custom_sso_login_handler = cast(CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler)
openid_response: OpenID = await custom_sso_login_handler.handle_custom_ui_sso_sign_in(
raise ValueError(
"custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings."
)
require_trusted_proxy_request(
request=request,
general_settings=general_settings,
feature_name="Custom UI SSO",
)
custom_sso_login_handler = cast(
CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler
)
openid_response: OpenID = (
await custom_sso_login_handler.handle_custom_ui_sso_sign_in(
request=request,
)
)
# Import here to avoid circular imports
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
return await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=openid_response,
request=request,
received_response=None,
generic_client_id=None,
ui_access_mode=None,
)
)

View file

@ -588,24 +588,21 @@ async def update_project( # noqa: PLR0915
param="project_id",
)
# Validate team exists and get team object for limit + permission checks
team_id_to_check = data.team_id or existing_project.team_id
team_obj_for_checks = None
if team_id_to_check is not None:
team_obj_for_checks = await _validate_team_exists(
team_id=team_id_to_check, prisma_client=prisma_client
# Permission to *edit* the project must be evaluated against the
# project's CURRENT team. Sourcing the team from `data.team_id`
# would let an admin of any team pass the check by supplying their
# own team_id, hijacking the project (VERIA-55).
target_team_id = data.team_id or existing_project.team_id
target_team_obj = None
if target_team_id is not None:
target_team_obj = await _validate_team_exists(
team_id=target_team_id, prisma_client=prisma_client
)
# Check if user has permission to update this project
has_permission = await _check_user_permission_for_project(
user_api_key_dict=user_api_key_dict,
team_id=existing_project.team_id,
prisma_client=prisma_client,
team_object=(
LiteLLM_TeamTable(**team_obj_for_checks.model_dump())
if team_obj_for_checks
else None
),
)
if not has_permission:
@ -614,10 +611,32 @@ async def update_project( # noqa: PLR0915
detail={"error": "Only admins or team admins can update projects"},
)
# Reassigning to a different team also requires admin rights on the
# destination team — otherwise a team admin could shed projects into
# an unsuspecting team's namespace.
if data.team_id is not None and data.team_id != existing_project.team_id:
can_assign_to_target = await _check_user_permission_for_project(
user_api_key_dict=user_api_key_dict,
team_id=data.team_id,
prisma_client=prisma_client,
team_object=(
LiteLLM_TeamTable(**target_team_obj.model_dump())
if target_team_obj
else None
),
)
if not can_assign_to_target:
raise HTTPException(
status_code=403,
detail={
"error": "Cannot reassign project to a team you are not an admin of"
},
)
# Validate project limits against team limits
if team_obj_for_checks is not None:
if target_team_obj is not None:
_check_team_project_limits(
team_object=LiteLLM_TeamTable(**team_obj_for_checks.model_dump()),
team_object=LiteLLM_TeamTable(**target_team_obj.model_dump()),
data=data,
)

View file

@ -288,6 +288,7 @@ disable_token_counter: bool = False
disable_add_transform_inline_image_block: bool = False
disable_add_user_agent_to_request_tags: bool = False
disable_anthropic_gemini_context_caching_transform: bool = False
disable_vertex_batch_output_transformation: bool = False
extra_spend_tag_headers: Optional[List[str]] = None
in_memory_llm_clients_cache: "LLMClientCache"
safe_memory_mode: bool = False
@ -330,6 +331,9 @@ enable_model_config_credential_overrides: bool = False
enable_key_alias_format_validation: bool = (
False # opt-in validation of key_alias format on /key/generate and /key/update
)
enable_gemini_default_thinking_level_low: bool = (
False # opt-in: force thinkingLevel low/minimal for Gemini 3 thinking param mapping
)
####################
logging: bool = True
enable_loadbalancing_on_batch_endpoints: Optional[bool] = None

View file

@ -1,12 +1,12 @@
import ast
import logging
import os
import re
import sys
from datetime import datetime
from logging import Formatter
from typing import Any, Dict, List, Optional
from typing import Any, Dict, Optional
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
@ -21,74 +21,11 @@ _ENABLE_SECRET_REDACTION = (
os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true"
)
_REDACTED = "REDACTED"
def _build_secret_patterns() -> re.Pattern:
patterns: List[str] = [
# ── PEM private key / certificate blocks ──
r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----",
# ── GCP OAuth2 access tokens (ya29.*) ──
r"\bya29\.[A-Za-z0-9_.~+/-]+",
# ── Credential %s formatting (space separator, no key= prefix) ──
r"(?:client_secret|azure_password|azure_username)\s+[^\s,'\"})\]{}>]+",
# AWS access key IDs
r"(?:AKIA|ASIA)[0-9A-Z]{16}",
# AWS secrets / session tokens / access key IDs (key=value)
r"(?:aws_secret_access_key|aws_session_token|aws_access_key_id)"
r"\s*[:=]\s*[A-Za-z0-9/+=]{20,}",
# Bearer tokens (OAuth, JWT, etc.)
r"Bearer\s+[A-Za-z0-9\-._~+/]{10,}=*",
# Basic auth headers
r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}",
# OpenAI / Anthropic sk- prefixed keys
r"sk-[A-Za-z0-9\-_]{20,}",
# Generic api_key / api-key / apikey (handles 'key': 'value' dict repr)
r"(?:api[_-]?key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]{8,}",
# x-api-key / api-key header values (handles 'key': 'value' dict repr)
r"(?:x-api-key|api-key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+",
# Anthropic internal header keys
r"x-ak-[A-Za-z0-9\-_]{20,}",
# Google API keys
r"AIza[0-9A-Za-z\-_]{35}",
# Password / secret params (handles key=value and 'key': 'value')
# Word boundary prevents O(n^2) backtracking on long word-char runs.
r"(?:^|(?<=\W))\w*(?:password|passwd|client_secret|secret_key|_secret)"
r"['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+",
# Database connection string credentials (scheme://user:pass@host)
r"(?<=://)[^\s'\"]*:[^\s'\"@]+(?=@)",
# Databricks personal access tokens
r"dapi[0-9a-f]{32}",
# ── Key-name-based redaction ──
# Catches secrets inside dicts/config dumps by matching on the KEY name
# regardless of what the value looks like.
# e.g. 'master_key': 'any-value-here', "database_url": "postgres://..."
# private_key with PEM-aware value capture
r"""private_key['\"]?\s*[:=]\s*['\"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'\"})\]{}>]+)""",
r"(?:master_key|database_url|db_url|connection_string|"
r"signing_key|encryption_key|"
r"auth_token|access_token|refresh_token|"
r"slack_webhook_url|webhook_url|"
r"database_connection_string|"
r"huggingface_token|jwt_secret)"
r"""['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+""",
# ── Raw JWTs (without Bearer prefix) ──
r"\beyJ[A-Za-z0-9_-]{10,}\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]*",
# ── Azure SAS tokens in URLs ──
r"[?&]sig=[A-Za-z0-9%+/=]+",
# ── Full JSON service-account blobs (single-line and multi-line) ──
r'\{[^{}]*"type"\s*:\s*"service_account"[^{}]*(?:\{[^{}]*\}[^{}]*)*\}',
]
return re.compile("|".join(patterns), re.IGNORECASE)
_SECRET_RE = _build_secret_patterns()
def _redact_string(value: str) -> str:
if not _ENABLE_SECRET_REDACTION:
return value
return _SECRET_RE.sub(_REDACTED, value)
return redact_string(value)
def redact_secrets(value: str) -> str:

View file

@ -387,6 +387,27 @@ def _get_batch_job_total_usage_from_file_content(
)
def _get_models_from_batch_input_file_content(
file_content_dictionary: List[dict],
) -> List[str]:
"""Extract the distinct ``body.model`` values from a batch *input* file.
Used by the proxy's batch pre-call hook to enforce that the caller is
authorized for every model named inside the JSONL — not just the one
on the outer request — so the proxy's per-key model allowlist isn't
bypassed by smuggling expensive models into the batch file.
"""
models: List[str] = []
seen: set = set()
for _item in file_content_dictionary:
body = _item.get("body") or {}
model = body.get("model")
if model and model not in seen:
seen.add(model)
models.append(model)
return models
def _get_batch_job_input_file_usage(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
@ -403,11 +424,25 @@ def _get_batch_job_input_file_usage(
for _item in file_content_dictionary:
body = _item.get("body", {})
model = body.get("model", model_name or "")
messages = body.get("messages", [])
# Chat completion payloads.
messages = body.get("messages")
if messages:
item_tokens = token_counter(model=model, messages=messages)
prompt_tokens += item_tokens
prompt_tokens += token_counter(model=model, messages=messages)
continue
# Text completion payloads (`prompt`).
prompt = body.get("prompt")
if prompt:
prompt_tokens += _count_prompt_or_input_tokens(model=model, value=prompt)
continue
# Embedding payloads (`input`).
input_data = body.get("input")
if input_data:
prompt_tokens += _count_prompt_or_input_tokens(
model=model, value=input_data
)
return Usage(
total_tokens=prompt_tokens + completion_tokens,
@ -416,6 +451,43 @@ def _get_batch_job_input_file_usage(
)
def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
"""Token-count a ``prompt`` / ``input`` field that the OpenAI batch
schema allows in four shapes:
- ``str``: a single text prompt.
- ``list[str]``: multiple text prompts.
- ``list[int]``: a pre-tokenized prompt (each int counts as 1 token).
- ``list[list[int]]``: multiple pre-tokenized prompts.
Pre-fix only the string shapes were counted, so a caller could send
a large ``list[list[int]]`` payload and slip past TPM rate limits
with a recorded cost of zero tokens.
"""
if isinstance(value, str):
return token_counter(model=model, text=value)
if isinstance(value, list):
total = 0
for chunk in value:
if isinstance(chunk, str):
total += token_counter(model=model, text=chunk)
elif isinstance(chunk, int):
# Single pre-tokenized prompt at the top level: each
# int counts as one token.
total += 1
elif isinstance(chunk, list):
# Nested pre-tokenized prompt: every int contributes a
# token. Mixed string/int items still count.
total += sum(1 if isinstance(t, int) else 0 for t in chunk)
total += sum(
token_counter(model=model, text=t)
for t in chunk
if isinstance(t, str)
)
return total
return 0
def _get_batch_job_usage_from_response_body(response_body: dict) -> Usage:
"""
Get the tokens of a batch job from the response body

View file

@ -419,9 +419,6 @@ CACHED_STREAMING_CHUNK_DELAY = float(os.getenv("CACHED_STREAMING_CHUNK_DELAY", 0
AUDIO_SPEECH_CHUNK_SIZE = int(
os.getenv("AUDIO_SPEECH_CHUNK_SIZE", 8192)
) # chunk_size for audio speech streaming. Balance between latency and memory usage
MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB = int(
os.getenv("MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB", 512)
)
DEFAULT_MAX_TOKENS_FOR_TRITON = int(os.getenv("DEFAULT_MAX_TOKENS_FOR_TRITON", 2000))
#### Networking settings ####
# Sentinel used when `REQUEST_TIMEOUT` is unset: `litellm.request_timeout` keeps this

View file

@ -513,7 +513,10 @@ def cost_per_token( # noqa: PLR0915
return fireworks_ai_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "azure":
return azure_openai_cost_per_token(
model=model, usage=usage_block, response_time_ms=response_time_ms
model=model,
usage=usage_block,
response_time_ms=response_time_ms,
service_tier=service_tier,
)
elif custom_llm_provider == "gemini":
return gemini_cost_per_token(
@ -539,6 +542,7 @@ def cost_per_token( # noqa: PLR0915
usage=usage_block,
response_time_ms=response_time_ms,
request_model=request_model,
service_tier=service_tier,
)
else:
model_info = _cached_get_model_info_helper(

View file

@ -220,23 +220,57 @@ def _set_structured_outputs(span: "Span", response_obj, msg_attrs, span_attrs):
safe_set_attribute(span, f"{prefix}.{msg_attrs.MESSAGE_ROLE}", message_role)
def _safe_get(obj, key, default=None):
"""Read ``key`` from a dict-like or Pydantic-model-like object.
The arize/langfuse_otel logger receives ``usage`` objects from many sources:
plain dicts, litellm ``Usage`` (which exposes ``.get``), and raw OpenAI
Pydantic models (e.g. ``openai.types.completion_usage.CompletionUsage`` and
nested ``CompletionTokensDetails`` / ``OutputTokensDetails``) which do NOT
expose ``.get``. Calling ``.get`` on the latter raised ``AttributeError`` —
see https://github.com/BerriAI/litellm/issues/13672.
"""
if obj is None:
return default
getter = getattr(obj, "get", None)
if callable(getter):
try:
return getter(key, default)
except TypeError:
# Some objects expose `.get` with a different signature
pass
return getattr(obj, key, default)
def _set_usage_outputs(span: "Span", response_obj, span_attrs):
usage = response_obj and response_obj.get("usage")
if not usage:
return
safe_set_attribute(
span, span_attrs.LLM_TOKEN_COUNT_TOTAL, usage.get("total_tokens")
span, span_attrs.LLM_TOKEN_COUNT_TOTAL, _safe_get(usage, "total_tokens")
)
completion_tokens = _safe_get(usage, "completion_tokens") or _safe_get(
usage, "output_tokens"
)
completion_tokens = usage.get("completion_tokens") or usage.get("output_tokens")
if completion_tokens:
safe_set_attribute(
span, span_attrs.LLM_TOKEN_COUNT_COMPLETION, completion_tokens
)
prompt_tokens = usage.get("prompt_tokens") or usage.get("input_tokens")
prompt_tokens = _safe_get(usage, "prompt_tokens") or _safe_get(
usage, "input_tokens"
)
if prompt_tokens:
safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_PROMPT, prompt_tokens)
reasoning_tokens = usage.get("output_tokens_details", {}).get("reasoning_tokens")
# Reasoning tokens live in `completion_tokens_details` for Chat Completions
# API (Usage) and in `output_tokens_details` for Responses API
# (ResponseAPIUsage). Both nested objects may be plain Pydantic models
# without `.get`.
token_details = _safe_get(usage, "completion_tokens_details") or _safe_get(
usage, "output_tokens_details"
)
reasoning_tokens = _safe_get(token_details, "reasoning_tokens")
if reasoning_tokens:
safe_set_attribute(
span,

View file

@ -18,6 +18,17 @@ class CustomSSOLoginHandler(CustomLogger):
self,
request: Request,
) -> OpenID:
from litellm.proxy.auth.trusted_proxy_utils import (
require_trusted_proxy_request,
)
from litellm.proxy.proxy_server import general_settings
require_trusted_proxy_request(
request=request,
general_settings=general_settings,
feature_name="Custom UI SSO",
)
request_headers_dict = dict(request.headers)
return OpenID(
id=request_headers_dict.get("x-litellm-user-id"),

View file

@ -90,6 +90,29 @@ def _extract_cache_read_input_tokens(usage_obj) -> int:
return cache_read_input_tokens
def resolve_langfuse_credentials(
langfuse_public_key=None,
langfuse_secret=None,
langfuse_secret_key=None,
langfuse_host=None,
allow_env_credentials: bool = True,
):
if allow_env_credentials is False and langfuse_host is not None:
secret_key = langfuse_secret or langfuse_secret_key
public_key = langfuse_public_key
else:
secret_key = (
langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY")
)
public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY")
resolved_host = langfuse_host or os.getenv(
"LANGFUSE_HOST", "https://cloud.langfuse.com"
)
return public_key, secret_key, resolved_host
class LangFuseLogger:
# Class variables or attributes
def __init__(
@ -98,6 +121,7 @@ class LangFuseLogger:
langfuse_secret=None,
langfuse_host=None,
flush_interval=1,
allow_env_credentials: bool = True,
):
try:
import langfuse
@ -106,11 +130,13 @@ class LangFuseLogger:
raise Exception(
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n{traceback.format_exc()}\033[0m"
)
# Instance variables
self.secret_key = langfuse_secret or os.getenv("LANGFUSE_SECRET_KEY")
self.public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY")
self.langfuse_host = langfuse_host or os.getenv(
"LANGFUSE_HOST", "https://cloud.langfuse.com"
self.public_key, self.secret_key, self.langfuse_host = (
resolve_langfuse_credentials(
langfuse_public_key=langfuse_public_key,
langfuse_secret=langfuse_secret,
langfuse_host=langfuse_host,
allow_env_credentials=allow_env_credentials,
)
)
if not (
self.langfuse_host.startswith("http://")
@ -160,9 +186,10 @@ class LangFuseLogger:
project_id = None
if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is not None:
upstream_langfuse_debug_env = os.getenv("UPSTREAM_LANGFUSE_DEBUG")
upstream_langfuse_debug = (
str_to_bool(self.upstream_langfuse_debug)
if self.upstream_langfuse_debug is not None
str_to_bool(upstream_langfuse_debug_env)
if upstream_langfuse_debug_env is not None
else None
)
self.upstream_langfuse_secret_key = os.getenv(
@ -173,7 +200,7 @@ class LangFuseLogger:
)
self.upstream_langfuse_host = os.getenv("UPSTREAM_LANGFUSE_HOST")
self.upstream_langfuse_release = os.getenv("UPSTREAM_LANGFUSE_RELEASE")
self.upstream_langfuse_debug = os.getenv("UPSTREAM_LANGFUSE_DEBUG")
self.upstream_langfuse_debug = upstream_langfuse_debug_env
self.upstream_langfuse = Langfuse(
public_key=self.upstream_langfuse_public_key,
secret_key=self.upstream_langfuse_secret_key,

View file

@ -115,8 +115,10 @@ class LangFuseHandler:
langfuse_logger = LangFuseLogger(
langfuse_public_key=credentials.get("langfuse_public_key"),
langfuse_secret=credentials.get("langfuse_secret"),
langfuse_secret=credentials.get("langfuse_secret")
or credentials.get("langfuse_secret_key"),
langfuse_host=credentials.get("langfuse_host"),
allow_env_credentials=credentials.get("langfuse_host") is None,
)
in_memory_dynamic_logger_cache.set_cache(
credentials=credentials,

View file

@ -20,7 +20,7 @@ from ...litellm_core_utils.specialty_caches.dynamic_logging_cache import (
DynamicLoggingCache,
)
from ..prompt_management_base import PromptManagementBase
from .langfuse import LangFuseLogger
from .langfuse import LangFuseLogger, resolve_langfuse_credentials
from .langfuse_handler import LangFuseHandler
if TYPE_CHECKING:
@ -46,6 +46,7 @@ def langfuse_client_init(
langfuse_secret_key=None,
langfuse_host=None,
flush_interval=1,
allow_env_credentials: bool = True,
) -> LangfuseClass:
"""
Initialize Langfuse client with caching to prevent multiple initializations.
@ -70,14 +71,12 @@ def langfuse_client_init(
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n\033[0m"
)
# Instance variables
secret_key = (
langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY")
)
public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY")
langfuse_host = langfuse_host or os.getenv(
"LANGFUSE_HOST", "https://cloud.langfuse.com"
public_key, secret_key, langfuse_host = resolve_langfuse_credentials(
langfuse_public_key=langfuse_public_key,
langfuse_secret=langfuse_secret,
langfuse_secret_key=langfuse_secret_key,
langfuse_host=langfuse_host,
allow_env_credentials=allow_env_credentials,
)
if not (
@ -222,6 +221,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
langfuse_secret=dynamic_callback_params.get("langfuse_secret"),
langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"),
langfuse_host=dynamic_callback_params.get("langfuse_host"),
allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None,
)
langfuse_prompt_client = self._get_prompt_from_id(
langfuse_prompt_id=prompt_id,
@ -246,6 +246,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
langfuse_secret=dynamic_callback_params.get("langfuse_secret"),
langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"),
langfuse_host=dynamic_callback_params.get("langfuse_host"),
allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None,
)
langfuse_prompt_client = self._get_prompt_from_id(
langfuse_prompt_id=prompt_id,

View file

@ -19,6 +19,7 @@ from litellm.integrations.langsmith_mock_client import (
create_mock_langsmith_client,
should_use_langsmith_mock,
)
from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -112,17 +113,28 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_project: Optional[str] = None,
langsmith_base_url: Optional[str] = None,
langsmith_tenant_id: Optional[str] = None,
allow_env_credentials: bool = True,
) -> LangsmithCredentialsObject:
_credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY")
_credentials_project = (
langsmith_project or os.getenv("LANGSMITH_PROJECT") or "litellm-completion"
)
_credentials_base_url = (
langsmith_base_url
or os.getenv("LANGSMITH_BASE_URL")
or "https://api.smith.langchain.com"
)
_credentials_tenant_id = langsmith_tenant_id or os.getenv("LANGSMITH_TENANT_ID")
if allow_env_credentials is False and langsmith_base_url is not None:
_credentials_api_key = langsmith_api_key
_credentials_project = langsmith_project or "litellm-completion"
_credentials_base_url = langsmith_base_url
_credentials_tenant_id = langsmith_tenant_id
else:
_credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY")
_credentials_project = (
langsmith_project
or os.getenv("LANGSMITH_PROJECT")
or "litellm-completion"
)
_credentials_base_url = (
langsmith_base_url
or os.getenv("LANGSMITH_BASE_URL")
or "https://api.smith.langchain.com"
)
_credentials_tenant_id = langsmith_tenant_id or os.getenv(
"LANGSMITH_TENANT_ID"
)
return LangsmithCredentialsObject(
LANGSMITH_API_KEY=_credentials_api_key,
@ -153,6 +165,15 @@ class LangsmithLogger(CustomBatchLogger):
for key in ("session_id", "thread_id", "conversation_id"):
if key in requester_metadata and key not in extra_metadata:
extra_metadata[key] = requester_metadata[key]
# helper is shallow; also scrub nested requester_metadata since
# LangSmith forwards the whole dict into `extra`
extra_metadata = redact_user_api_key_info(metadata=extra_metadata)
nested = extra_metadata.get("requester_metadata")
if isinstance(nested, dict):
extra_metadata["requester_metadata"] = redact_user_api_key_info(
metadata=nested
)
return extra_metadata
def _build_outputs_with_usage(
@ -540,6 +561,10 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_tenant_id=standard_callback_dynamic_params.get(
"langsmith_tenant_id", None
),
allow_env_credentials=standard_callback_dynamic_params.get(
"langsmith_base_url", None
)
is None,
)
else:
credentials = self.default_credentials

View file

@ -69,6 +69,8 @@ class OpenTelemetryConfig:
deployment_environment: Optional[str] = None
model_id: Optional[str] = None
ignore_context_propagation: Optional[bool] = None
# When True, create a private TracerProvider instead of reusing or setting the global one.
skip_set_global: bool = False
def __post_init__(self) -> None:
# If endpoint is specified but exporter is still the default "console",
@ -259,16 +261,21 @@ class OpenTelemetry(CustomLogger):
try:
existing_provider = get_existing_provider_fn()
# If a real SDK provider exists (set by another SDK like Langfuse), use it
# This uses a positive check for SDK providers instead of a negative check for proxy providers
if isinstance(existing_provider, sdk_provider_class):
verbose_logger.debug(
"OpenTelemetry: Using existing %s: %s",
provider_name,
type(existing_provider).__name__,
)
provider = existing_provider
# Don't call set_provider to preserve existing context
if skip_set_global:
verbose_logger.debug(
"OpenTelemetry: existing %s found but skip_set_global=True; creating private %s for isolation",
provider_name,
provider_name,
)
provider = create_new_provider_fn()
else:
verbose_logger.debug(
"OpenTelemetry: Using existing %s: %s",
provider_name,
type(existing_provider).__name__,
)
provider = existing_provider
else:
# Default proxy provider or unknown type, create our own
verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name)
@ -293,6 +300,12 @@ class OpenTelemetry(CustomLogger):
return provider
def _skip_set_global(self) -> bool:
# langfuse_otel relies on the Langfuse SDK's providers; don't overwrite them.
return self.config.skip_set_global or (
hasattr(self, "callback_name") and self.callback_name == "langfuse_otel"
)
def _init_tracing(self, tracer_provider):
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
@ -303,11 +316,6 @@ class OpenTelemetry(CustomLogger):
provider.add_span_processor(self._get_span_processor())
return provider
# CRITICAL FIX: For Langfuse OTEL, skip setting global provider to prevent interference
skip_global = (
hasattr(self, "callback_name") and self.callback_name == "langfuse_otel"
)
tracer_provider = self._get_or_create_provider(
provider=tracer_provider,
provider_name="TracerProvider",
@ -315,16 +323,18 @@ class OpenTelemetry(CustomLogger):
sdk_provider_class=TracerProvider,
create_new_provider_fn=create_tracer_provider,
set_provider_fn=trace.set_tracer_provider,
skip_set_global=skip_global,
skip_set_global=self._skip_set_global(),
)
# Grab our tracer from the TracerProvider (not from global context)
# This ensures we use the provided TracerProvider (e.g., for testing)
self.tracer = tracer_provider.get_tracer(LITELLM_TRACER_NAME)
self._tracer_provider = tracer_provider
self.span_kind = SpanKind
def _init_metrics(self, meter_provider):
if not self.config.enable_metrics:
self._meter_provider = None
self._operation_duration_histogram = None
self._token_usage_histogram = None
self._cost_histogram = None
@ -350,7 +360,9 @@ class OpenTelemetry(CustomLogger):
sdk_provider_class=MeterProvider,
create_new_provider_fn=create_meter_provider,
set_provider_fn=metrics.set_meter_provider,
skip_set_global=self._skip_set_global(),
)
self._meter_provider = meter_provider
meter = meter_provider.get_meter(__name__)
@ -388,6 +400,7 @@ class OpenTelemetry(CustomLogger):
def _init_logs(self, logger_provider):
# nothing to do if events disabled
if not self.config.enable_events:
self._logger_provider = None
return
from opentelemetry._logs import get_logger_provider, set_logger_provider
@ -404,13 +417,14 @@ class OpenTelemetry(CustomLogger):
)
return provider
self._get_or_create_provider(
self._logger_provider = self._get_or_create_provider(
provider=logger_provider,
provider_name="LoggerProvider",
get_existing_provider_fn=get_logger_provider,
sdk_provider_class=OTLoggerProvider,
create_new_provider_fn=create_logger_provider,
set_provider_fn=set_logger_provider,
skip_set_global=self._skip_set_global(),
)
def log_success_event(self, kwargs, response_obj, start_time, end_time):
@ -1073,7 +1087,7 @@ class OpenTelemetry(CustomLogger):
# See: https://github.com/open-telemetry/opentelemetry-python/pull/4676
# TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords
from opentelemetry._logs import SeverityNumber, get_logger
from opentelemetry._logs import SeverityNumber
try:
from opentelemetry.sdk._logs import ( # type: ignore[attr-defined] # OTEL < 1.39.0
@ -1084,7 +1098,10 @@ class OpenTelemetry(CustomLogger):
LogRecord as SdkLogRecord, # type: ignore[attr-defined] # OTEL >= 1.39.0
)
otel_logger = get_logger(LITELLM_LOGGER_NAME)
# Resolve through the handler's own LoggerProvider (which may be a
# private one when skip_set_global=True) rather than the module-level
# get_logger() which always goes through the global provider.
otel_logger = self._logger_provider.get_logger(LITELLM_LOGGER_NAME)
parent_ctx = span.get_span_context()
provider = (kwargs.get("litellm_params") or {}).get(

View file

@ -2,6 +2,7 @@
Helper functions to query prometheus API
"""
import json
import time
from datetime import datetime, timedelta
from typing import Optional
@ -81,6 +82,24 @@ def is_prometheus_connected() -> bool:
return False
def _quote_promql_string_literal(value: str) -> str:
"""Render ``value`` as a PromQL double-quoted string literal.
PromQL string literals follow Go's escape rules
(https://prometheus.io/docs/prometheus/latest/querying/basics/): a
backslash begins an escape sequence and a bare ``"`` ends the literal.
Without escaping, callers that accept arbitrary user-supplied values
(like the ``api_key`` filter on ``/global/spend/logs``) can inject extra
label matchers or selectors and read cross-tenant metrics.
JSON's quoting rules are a strict subset of Go's, so ``json.dumps`` of
a Python string produces a literal Prometheus accepts: ``\\``, ``\\"``,
and the standard ``\\n`` / ``\\t`` / ``\\uNNNN`` control-character
escapes. The returned value already includes the surrounding quotes.
"""
return json.dumps(value, ensure_ascii=False)
async def get_daily_spend_from_prometheus(api_key: Optional[str]):
"""
Expected Response Format:
@ -109,8 +128,11 @@ async def get_daily_spend_from_prometheus(api_key: Optional[str]):
if api_key is None:
query = "sum(delta(litellm_spend_metric_total[1d]))"
else:
quoted_api_key = _quote_promql_string_literal(api_key)
query = (
f'sum(delta(litellm_spend_metric_total{{hashed_api_key="{api_key}"}}[1d]))'
"sum(delta(litellm_spend_metric_total{"
f"hashed_api_key={quoted_api_key}"
"}[1d]))"
)
params = {

View file

@ -6,7 +6,8 @@ from typing import Any, Optional
import httpx
import litellm
from litellm._logging import _redact_string, verbose_logger
from litellm._logging import _ENABLE_SECRET_REDACTION, _redact_string, verbose_logger
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.types.utils import LlmProviders
from ..exceptions import (
@ -261,10 +262,18 @@ def exception_type( # type: ignore # noqa: PLR0915
original_exception=original_exception
)
try:
error_str = str(original_exception)
error_str = (
redact_string(str(original_exception))
if _ENABLE_SECRET_REDACTION
else str(original_exception)
)
if model:
if hasattr(original_exception, "message"):
error_str = str(original_exception.message)
error_str = (
redact_string(str(original_exception.message))
if _ENABLE_SECRET_REDACTION
else str(original_exception.message)
)
if isinstance(original_exception, BaseException):
exception_type = type(original_exception).__name__
else:
@ -2431,7 +2440,8 @@ def exception_type( # type: ignore # noqa: PLR0915
else:
raise APIConnectionError(
message="{}\n{}".format(
str(original_exception), _redact_string(traceback.format_exc())
str(original_exception),
_redact_string(traceback.format_exc()),
),
llm_provider=custom_llm_provider,
model=model,
@ -2461,7 +2471,8 @@ def exception_type( # type: ignore # noqa: PLR0915
raise e # it's already mapped
raised_exc = APIConnectionError(
message="{}\n{}".format(
original_exception, _redact_string(traceback.format_exc())
original_exception,
_redact_string(traceback.format_exc()),
),
llm_provider="",
model="",

View file

@ -3242,10 +3242,15 @@ class Logging(LiteLLMLoggingBaseClass):
),
langfuse_secret=self.standard_callback_dynamic_params.get(
"langfuse_secret"
),
)
or self.standard_callback_dynamic_params.get("langfuse_secret_key"),
langfuse_host=self.standard_callback_dynamic_params.get(
"langfuse_host"
),
allow_env_credentials=self.standard_callback_dynamic_params.get(
"langfuse_host"
)
is None,
)
return langFuseLogger
@ -4720,7 +4725,7 @@ class StandardLoggingPayloadSetup:
):
for key, value in litellm_params["metadata"].items():
# Skip non-serializable objects like UserAPIKeyAuth
if key == "user_api_key_auth":
if key in {"user_api_key_auth", "user_api_key_budget_reservation"}:
continue
merged_metadata[key] = value

View file

@ -0,0 +1,81 @@
"""
Credential/secret redaction utilities.
This module owns the compiled regex and the public `redact_string` helper so
that any part of the codebase (logging, exception mapping, etc.) can scrub
secrets from strings without depending on the logging-configuration module.
"""
import re
from typing import List
_REDACTED = "REDACTED"
def _build_secret_patterns() -> "re.Pattern[str]":
patterns: List[str] = [
# PEM private key / certificate blocks
r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----",
# GCP OAuth2 access tokens (ya29.*)
r"\bya29\.[A-Za-z0-9_.~+/-]+",
# Credential %s formatting (space separator, no key= prefix)
r"(?:client_secret|azure_password|azure_username)\s+[^\s,'\"})\]{}>]+",
# AWS access key IDs
r"(?:AKIA|ASIA)[0-9A-Z]{16}",
# AWS secrets / session tokens / access key IDs (key=value)
r"(?:aws_secret_access_key|aws_session_token|aws_access_key_id)"
r"\s*[:=]\s*[A-Za-z0-9/+=]{20,}",
# Bearer tokens (OAuth, JWT, etc.)
r"Bearer\s+[A-Za-z0-9\-._~+/]{10,}=*",
# Basic auth headers
r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}",
# OpenAI / Anthropic sk- prefixed keys
r"sk-[A-Za-z0-9\-_]{20,}",
# Generic api_key / api-key / apikey (handles 'key': 'value' dict repr)
r"(?:api[_-]?key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]{8,}",
# x-api-key / api-key header values (handles 'key': 'value' dict repr)
r"(?:x-api-key|api-key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+",
# Anthropic internal header keys
r"x-ak-[A-Za-z0-9\-_]{20,}",
# Google API keys (bare key value)
r"AIza[0-9A-Za-z\-_]{35}",
# URL query-param key=VALUE (e.g. ?key=AIza... or &key=...) — catches the
# full "key=<secret>" fragment so the value is redacted regardless of format.
r"(?<=[?&])key=[^\s&'\"]{8,}",
# Password / secret params (handles key=value and 'key': 'value')
# Word boundary prevents O(n^2) backtracking on long word-char runs.
r"(?:^|(?<=\W))\w*(?:password|passwd|client_secret|secret_key|_secret)"
r"['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+",
# Database connection string credentials (scheme://user:pass@host)
r"(?<=://)[^\s'\"]*:[^\s'\"@]+(?=@)",
# Databricks personal access tokens
r"dapi[0-9a-f]{32}",
# ── Key-name-based redaction ──
# Catches secrets inside dicts/config dumps by matching on the KEY name
# regardless of what the value looks like.
# e.g. 'master_key': 'any-value-here', "database_url": "postgres://..."
# private_key with PEM-aware value capture
r"""private_key['\"]?\s*[:=]\s*['\"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'\"})\]{}>]+)""",
r"(?:master_key|database_url|db_url|connection_string|"
r"signing_key|encryption_key|"
r"auth_token|access_token|refresh_token|"
r"slack_webhook_url|webhook_url|"
r"database_connection_string|"
r"huggingface_token|jwt_secret)"
r"""['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+""",
# Raw JWTs (without Bearer prefix)
r"\beyJ[A-Za-z0-9_-]{10,}\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]*",
# Azure SAS tokens in URLs
r"[?&]sig=[A-Za-z0-9%+/=]+",
# Full JSON service-account blobs (single-line and multi-line)
r'\{[^{}]*"type"\s*:\s*"service_account"[^{}]*(?:\{[^{}]*\}[^{}]*)*\}',
]
return re.compile("|".join(patterns), re.IGNORECASE)
_SECRET_RE = _build_secret_patterns()
def redact_string(value: str) -> str:
"""Scrub known secret/credential patterns from *value* and return the result."""
return _SECRET_RE.sub(_REDACTED, value)

View file

@ -1088,24 +1088,29 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
elif param == "thinking":
optional_params["thinking"] = value
elif param == "reasoning_effort" and isinstance(value, str):
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=value, model=model
)
# For Claude 4.6+ models, effort is controlled via output_config,
# not thinking budget_tokens. Map reasoning_effort to output_config.
if AnthropicConfig._is_claude_4_6_model(
model
) or AnthropicConfig._is_claude_4_7_model(model):
effort_map = {
"low": "low",
"minimal": "low",
"medium": "medium",
"high": "high",
"xhigh": "xhigh",
"max": "max",
}
mapped_effort = effort_map.get(value, value)
optional_params["output_config"] = {"effort": mapped_effort}
if mapped_thinking is None:
optional_params.pop("thinking", None)
optional_params.pop("output_config", None)
else:
optional_params["thinking"] = mapped_thinking
# For Claude 4.6+ models, effort is controlled via output_config,
# not thinking budget_tokens. Map reasoning_effort to output_config.
if AnthropicConfig._is_claude_4_6_model(
model
) or AnthropicConfig._is_claude_4_7_model(model):
effort_map = {
"low": "low",
"minimal": "low",
"medium": "medium",
"high": "high",
"xhigh": "xhigh",
"max": "max",
}
mapped_effort = effort_map.get(value, value)
optional_params["output_config"] = {"effort": mapped_effort}
elif param == "web_search_options" and isinstance(value, dict):
hosted_web_search_tool = self.map_web_search_tool(
cast(OpenAIWebSearchOptions, value)

View file

@ -27,6 +27,16 @@ from litellm.utils import get_model_info
if TYPE_CHECKING:
pass
# Anthropic-only fields that the translator above already maps into the
# OpenAI-format completion_kwargs (output_config → reasoning_effort /
# response_format, etc.). They must be filtered out of the raw
# extra_kwargs re-merge below or non-Anthropic backends reject the call
# with 400 "Extra inputs are not permitted". Add new entries here when
# extending AnthropicMessagesRequestOptionalParams with another Anthropic-
# specific key.
ANTHROPIC_ONLY_REQUEST_KEYS: frozenset[str] = frozenset({"output_config"})
########################################################
# init adapter
ANTHROPIC_ADAPTER = AnthropicAdapter()
@ -202,8 +212,12 @@ class LiteLLMMessagesToCompletionTransformationHandler:
request_data["output_format"] = output_format
# Extract output_config from extra_kwargs so the translator can use it
# (e.g. output_config.effort for adaptive thinking → reasoning_effort)
extra_kwargs = extra_kwargs or {}
# (e.g. output_config.effort for adaptive thinking → reasoning_effort,
# output_config.format → response_format for structured outputs).
# Use explicit None check rather than `or {}` so an explicit empty dict
# caller-passed argument is preserved (matters for tests that drive
# the fallback inference path).
extra_kwargs = extra_kwargs if extra_kwargs is not None else {}
if "output_config" in extra_kwargs:
request_data["output_config"] = extra_kwargs["output_config"]
@ -225,8 +239,23 @@ class LiteLLMMessagesToCompletionTransformationHandler:
"include_usage": True,
}
excluded_keys = {"anthropic_messages"}
extra_kwargs = extra_kwargs or {}
# Keys that must NOT be forwarded as raw extras into the OpenAI-format
# ``completion_kwargs`` after translation. The translator above has
# already consumed the meaningful parts of these inputs (e.g.
# ``output_config.format`` → ``response_format``, ``output_config.effort``
# → ``reasoning_effort`` for non-Claude targets). Re-adding the raw
# Anthropic-shaped key here causes 400 "Extra inputs are not permitted"
# on non-Anthropic backends (Azure OpenAI, Fireworks, Bedrock Nova,
# etc.) and is silently lossy on Anthropic-family targets, which would
# see the translated key ``response_format`` AND a duplicate, conflicting
# ``output_config``.
#
# Maintainability: when adding a new Anthropic-only request param to
# ``AnthropicMessagesRequestOptionalParams``, also extend
# ``ANTHROPIC_ONLY_REQUEST_KEYS`` here so it doesn't silently leak.
excluded_keys = ANTHROPIC_ONLY_REQUEST_KEYS | {"anthropic_messages"}
# NOTE: extra_kwargs was already coerced from None to {} at the top of
# this method (line ~220). It is guaranteed to be a dict here.
for key, value in extra_kwargs.items():
if (
key == "litellm_logging_obj"

View file

@ -667,7 +667,7 @@ class LiteLLMAnthropicMessagesAdapter:
@staticmethod
def translate_anthropic_thinking_to_reasoning_effort(
thinking: Dict[str, Any]
thinking: Dict[str, Any],
) -> Optional[str]:
"""
Translate Anthropic's thinking parameter to OpenAI's reasoning_effort.
@ -1084,10 +1084,23 @@ class LiteLLMAnthropicMessagesAdapter:
anthropic_message_request: AnthropicMessagesRequest,
new_kwargs: ChatCompletionRequest,
) -> None:
"""Translate output_format to response_format when applicable."""
if "output_format" not in anthropic_message_request:
return
output_format = anthropic_message_request["output_format"]
"""Translate Anthropic structured-output config to OpenAI ``response_format``.
Accepts either the legacy top-level ``output_format`` field OR the
newer ``output_config.format`` (sub-key on ``output_config``) so that
both shapes flow through to non-Anthropic backends as
``response_format``. Without the ``output_config.format`` branch,
callers using the new Anthropic Structured Outputs API would have
their schema silently dropped on the adapter path — only the legacy
top-level ``output_format`` was being mapped.
``output_format`` takes precedence when both are provided.
"""
output_format: Any = anthropic_message_request.get("output_format")
if not output_format:
output_config = anthropic_message_request.get("output_config")
if isinstance(output_config, dict):
output_format = output_config.get("format")
if not output_format:
return
response_format = self.translate_anthropic_output_format_to_openai(

View file

@ -793,6 +793,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
client=client,
litellm_params=litellm_params,
api_base=api_base,
api_version=api_version,
)
azure_client = self.get_azure_openai_client(
api_version=api_version,

View file

@ -12,7 +12,10 @@ from litellm.utils import get_model_info
def cost_per_token(
model: str, usage: Usage, response_time_ms: Optional[float] = 0.0
model: str,
usage: Usage,
response_time_ms: Optional[float] = 0.0,
service_tier: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -47,4 +50,5 @@ def cost_per_token(
model=model,
usage=usage,
custom_llm_provider="azure",
service_tier=service_tier,
)

View file

@ -65,6 +65,7 @@ def cost_per_token(
usage: Usage,
response_time_ms: Optional[float] = 0.0,
request_model: Optional[str] = None,
service_tier: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculate the cost per token for Azure AI models.
@ -102,6 +103,7 @@ def cost_per_token(
model=model,
usage=usage,
custom_llm_provider="azure_ai",
service_tier=service_tier,
)
except Exception as e:
# For Model Router, the model name (e.g., "azure-model-router") may not be in the cost map

View file

@ -449,9 +449,13 @@ class AmazonConverseConfig(BaseConfig):
optional_params.update(reasoning_config)
else:
# Anthropic and other models: convert to thinking parameter
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=reasoning_effort, model=model
)
if mapped_thinking is None:
optional_params.pop("thinking", None)
else:
optional_params["thinking"] = mapped_thinking
@staticmethod
def _clamp_thinking_budget_tokens(optional_params: dict) -> None:

View file

@ -149,9 +149,9 @@ class CloudflareChatConfig(BaseConfig):
) -> ModelResponse:
completion_response = raw_response.json()
model_response.choices[0].message.content = completion_response["result"][ # type: ignore
"response"
]
# Support both "response" and "response_text" keys (newer models like Nemotron use "response_text")
result = completion_response["result"]
model_response.choices[0].message.content = result.get("response") if result.get("response") is not None else result.get("response_text", "") # type: ignore
prompt_tokens = litellm.utils.get_token_count(messages=messages, model=model)
completion_tokens = len(
@ -201,8 +201,10 @@ class CloudflareChatResponseIterator(BaseModelResponseIterator):
index = int(chunk.get("index", 0))
if "response" in chunk:
if "response" in chunk and chunk["response"] is not None:
text = chunk["response"]
elif "response_text" in chunk and chunk["response_text"] is not None:
text = chunk["response_text"]
returned_chunk = GenericStreamingChunk(
text=text,

View file

@ -2,4 +2,15 @@ No transformation is required for hosted_vllm embedding.
VLLM is a superset of OpenAI's `embedding` endpoint.
To pass provider-specific parameters, see [this](https://docs.litellm.ai/docs/completion/provider_specific_params)
## `encoding_format`
For OpenAI-compatible embedding calls (including `openai/...` with a custom `api_base` pointing at vLLM), LiteLLM resolves `encoding_format` when it is not set on the request:
1. Explicit value on the embedding call (`encoding_format=...`).
2. Model config (`litellm_params.encoding_format` on the proxy `model_list` entry).
3. Environment variable `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT` (e.g. in `.env` or container env).
4. Default **`float`**.
That avoids forwarding `encoding_format=None` to the provider/SDK where some servers behave poorly.
To pass provider-specific parameters, see [provider-specific params](https://docs.litellm.ai/docs/completion/provider_specific_params).

View file

@ -244,9 +244,11 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
),
status_code=400,
)
elif effective_effort == "minimal":
# minimal is opt-out: unknown models pass through; only block when
# the model map explicitly sets supports_minimal_reasoning_effort=false.
elif effective_effort in ("minimal", "low"):
# minimal/low are opt-out: unknown models pass through; only block when
# the model map explicitly sets supports_{level}_reasoning_effort=false.
# Example: gpt-5.5-pro only accepts {medium, high, xhigh}, so it sets
# supports_low_reasoning_effort=false (and supports_minimal=false).
if self._is_reasoning_effort_level_explicitly_disabled(
model, effective_effort
):

View file

@ -106,5 +106,13 @@
"base_url": "https://aihubmix.com/v1",
"api_key_env": "AIHUBMIX_API_KEY",
"api_base_env": "AIHUBMIX_API_BASE"
},
"crusoe": {
"base_url": "https://managed-inference-api-proxy.crusoecloud.com/v1",
"api_key_env": "CRUSOE_API_KEY",
"api_base_env": "CRUSOE_API_BASE",
"param_mappings": {
"max_completion_tokens": "max_tokens"
}
}
}

View file

@ -656,6 +656,8 @@ def process_items(schema, depth=0):
and ("items" not in schema or schema.get("items") == {})
):
schema["items"] = {"type": "object"}
elif schema.get("type") == "array" and "items" not in schema:
schema["items"] = {"type": "object"}
for key, value in schema.items():
if isinstance(value, dict):
process_items(value, depth + 1)

View file

@ -1,4 +1,5 @@
import asyncio
import time
from urllib.parse import unquote
from typing import Any, Coroutine, Optional, Tuple, Union
@ -21,9 +22,10 @@ from litellm.types.llms.openai import (
HttpxBinaryResponseContent,
OpenAIFileObject,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
from .transformation import VertexAIJsonlFilesTransformation
from .transformation import VertexAIFilesConfig, VertexAIJsonlFilesTransformation
vertex_ai_files_transformation = VertexAIJsonlFilesTransformation()
@ -198,11 +200,30 @@ class VertexAIFilesHandler(GCSBucketBase):
mock_response = httpx.Response(
status_code=200,
content=file_content,
headers={"content-type": "application/octet-stream"},
headers={
"content-type": "application/octet-stream",
"content-length": str(len(file_content)),
},
request=httpx.Request(method="GET", url=decoded_file_id),
)
return HttpxBinaryResponseContent(response=mock_response)
# Apply transformation to convert Vertex AI batch outputs to OpenAI format
config = VertexAIFilesConfig()
# Create a logging object for transformation
logging_obj = Logging(
model="",
messages=[],
stream=False,
call_type="afile_content",
start_time=time.time(),
litellm_call_id="",
function_id="",
)
return config.transform_file_content_response(
raw_response=mock_response, logging_obj=logging_obj, litellm_params={}
)
def file_content(
self,

View file

@ -1,10 +1,15 @@
import base64
import json
import os
from typing import Any, Dict, List, Optional, Tuple, Union
import re
import time
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
import litellm
from litellm._uuid import uuid
from litellm.files.utils import FilesAPIUtils
from litellm.litellm_core_utils.cloud_storage_security import (
@ -16,6 +21,7 @@ from litellm.litellm_core_utils.cloud_storage_security import (
split_configured_cloud_bucket_name,
validate_managed_cloud_file_id,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.transformation import (
@ -39,11 +45,135 @@ from litellm.types.llms.openai import (
PathLike,
)
from litellm.types.llms.vertex_ai import GcsBucketResponse
from litellm.types.utils import ExtractedFileData, LlmProviders
from litellm.types.utils import ExtractedFileData, LlmProviders, ModelResponse
from ..common_utils import VertexAIError
from ..vertex_llm_base import VertexBase
_GCP_LABEL_VALUE_MAX_LEN = 63
_CUSTOM_ID_RAW_LABEL_PREFIX = "b32_"
def _sanitize_gcp_label_value(value: str) -> str:
"""
Sanitize a string to meet GCP label value constraints.
GCP label values must:
- Be lowercase
- Contain only letters, numbers, underscores, and hyphens
- Be max 63 characters
Args:
value: The string to sanitize
Returns:
A sanitized string that meets GCP label constraints
"""
sanitized = re.sub(r"[^a-z0-9_-]", "_", value.lower())
return sanitized[:_GCP_LABEL_VALUE_MAX_LEN]
def _encode_gcp_label_value_chunks(value: str) -> List[str]:
"""Encode arbitrary text across one or more GCP-label-safe values."""
max_encoded_len = _GCP_LABEL_VALUE_MAX_LEN - len(_CUSTOM_ID_RAW_LABEL_PREFIX)
encoded = (
base64.b32encode(value.encode("utf-8")).decode("ascii").rstrip("=").lower()
)
return [
f"{_CUSTOM_ID_RAW_LABEL_PREFIX}{encoded[i : i + max_encoded_len]}"
for i in range(0, len(encoded), max_encoded_len)
] or [_CUSTOM_ID_RAW_LABEL_PREFIX]
def _decode_gcp_label_value_chunks(values: List[str]) -> Optional[str]:
"""Decode values produced by _encode_gcp_label_value_chunks."""
encoded_parts = []
for value in values:
if not value.startswith(_CUSTOM_ID_RAW_LABEL_PREFIX):
return None
encoded_parts.append(value[len(_CUSTOM_ID_RAW_LABEL_PREFIX) :])
encoded = "".join(encoded_parts).upper()
padding = "=" * (-len(encoded) % 8)
try:
return base64.b32decode(encoded + padding).decode("utf-8")
except Exception:
return None
def _set_litellm_batch_custom_id_labels(labels: Dict[str, str], custom_id: Any) -> None:
"""
Store OpenAI batch custom_id for Vertex batch correlation.
``litellm_custom_id`` is GCP-label-safe (may alter casing and characters).
``litellm_custom_id_raw`` encodes the original string for
round-trip correlation in batch output transforms.
"""
custom_id_str = str(custom_id)
labels["litellm_custom_id"] = _sanitize_gcp_label_value(custom_id_str)
raw_label_chunks = _encode_gcp_label_value_chunks(custom_id_str)
labels["litellm_custom_id_raw"] = raw_label_chunks[0]
for index, raw_label_chunk in enumerate(raw_label_chunks[1:], start=1):
labels[f"litellm_custom_id_raw_{index}"] = raw_label_chunk
def _get_litellm_batch_custom_id_from_labels(labels: Dict[str, Any]) -> str:
"""Prefer encoded custom_id when present (see _set_litellm_batch_custom_id_labels)."""
raw = labels.get("litellm_custom_id_raw")
if raw:
raw_chunks = [str(raw)]
chunk_prefix = "litellm_custom_id_raw_"
indexed_chunks = []
for key, value in labels.items():
if key.startswith(chunk_prefix) and key[len(chunk_prefix) :].isdigit():
indexed_chunks.append((int(key[len(chunk_prefix) :]), str(value)))
raw_chunks.extend(
raw_label_chunk
for _, raw_label_chunk in sorted(indexed_chunks, key=lambda item: item[0])
)
decoded = _decode_gcp_label_value_chunks(raw_chunks)
if decoded is not None:
return decoded
return str(raw)
return str(labels.get("litellm_custom_id", "unknown"))
def _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
openai_jsonl_content: List[Dict[str, Any]],
map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""
Transforms OpenAI JSONL batch entries to Vertex AI JSONL lines.
jsonl body for vertex is {"request": <request_body>}
Example Vertex jsonl
{"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}}
{"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}}
"""
vertex_jsonl_content = []
for _openai_jsonl_content in openai_jsonl_content:
openai_request_body = _openai_jsonl_content.get("body") or {}
vertex_request_body = _transform_request_body(
messages=openai_request_body.get("messages", []),
model=openai_request_body.get("model", ""),
optional_params=map_openai_to_vertex_params(openai_request_body),
custom_llm_provider="vertex_ai",
litellm_params={},
cached_content=None,
)
# Add custom_id as a label for correlation in batch outputs
custom_id = _openai_jsonl_content.get("custom_id")
if custom_id is not None:
if "labels" not in vertex_request_body:
vertex_request_body["labels"] = {}
_set_litellm_batch_custom_id_labels(
vertex_request_body["labels"], custom_id
)
vertex_jsonl_content.append({"request": vertex_request_body})
return vertex_jsonl_content
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
@ -239,28 +369,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
self, openai_jsonl_content: List[Dict[str, Any]]
) -> List[Dict[str, Any]]:
"""
Transforms OpenAI JSONL content to VertexAI JSONL content
jsonl body for vertex is {"request": <request_body>}
Example Vertex jsonl
{"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}}
{"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}}
"""
vertex_jsonl_content = []
for _openai_jsonl_content in openai_jsonl_content:
openai_request_body = _openai_jsonl_content.get("body") or {}
vertex_request_body = _transform_request_body(
messages=openai_request_body.get("messages", []),
model=openai_request_body.get("model", ""),
optional_params=self._map_openai_to_vertex_params(openai_request_body),
custom_llm_provider="vertex_ai",
litellm_params={},
cached_content=None,
)
vertex_jsonl_content.append({"request": vertex_request_body})
return vertex_jsonl_content
return _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
openai_jsonl_content=openai_jsonl_content,
map_openai_to_vertex_params=self._map_openai_to_vertex_params,
)
def transform_create_file_request(
self,
@ -461,8 +573,229 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> HttpxBinaryResponseContent:
"""
Transform file content response, converting Vertex AI batch output to OpenAI format if applicable.
This method automatically detects and transforms Vertex AI batch prediction outputs
(predictions.jsonl files) into OpenAI-compatible batch response format.
If the file is not a batch output or transformation fails, the original content
is returned as-is to maintain backward compatibility.
"""
try:
# Allow users to opt out of automatic Vertex batch output -> OpenAI
# transformation, e.g. if they consume raw `predictions.jsonl` directly.
if getattr(litellm, "disable_vertex_batch_output_transformation", False):
return HttpxBinaryResponseContent(response=raw_response)
# Try to transform batch output if it's a JSONL file
content = raw_response.content
if content:
transformed_content = self._try_transform_vertex_batch_output_to_openai(
content=content,
logging_obj=logging_obj,
)
if transformed_content != content:
# Create a new response with transformed content and updated Content-Length
# Update headers with correct Content-Length
new_headers = dict(raw_response.headers)
new_headers["content-length"] = str(len(transformed_content))
mock_response = httpx.Response(
status_code=raw_response.status_code,
content=transformed_content,
headers=new_headers,
request=raw_response.request,
)
return HttpxBinaryResponseContent(response=mock_response)
except Exception:
# If transformation fails, return as-is
pass
return HttpxBinaryResponseContent(response=raw_response)
def _try_transform_vertex_batch_output_to_openai(
self, content: bytes, logging_obj: Optional[LiteLLMLoggingObj] = None
) -> bytes:
"""
Try to transform Vertex AI batch output to OpenAI format.
If conversion fails at any point, return the original content as-is.
Vertex AI batch output format (predictions.jsonl):
{
"request": {"contents": [...], "labels": {"litellm_custom_id": "request-1", "litellm_custom_id_raw": "..."}},
"status": "",
"response": {"candidates": [...], "modelVersion": "gemini-2.5-flash", ...},
"processed_time": "2026-04-13T10:18:18.102004+00:00"
}
OpenAI batch output format:
{
"id": "batch_req_...",
"custom_id": "request-1",
"response": {
"status_code": 200,
"request_id": "chatcmpl-...",
"body": {<OpenAI chat completion response>}
},
"error": null
}
"""
try:
# Decode content
content_str = content.decode("utf-8")
# Check if it's JSONL (multiple lines)
lines = content_str.strip().split("\n")
if not lines:
return content
# Try to parse the first line to see if it's Vertex AI batch output
first_line = json.loads(lines[0])
# Check if it has Vertex AI batch output structure with discriminating fields
# Must have request, response, and processed_time
# Plus either candidates (success) or status (error)
has_base_structure = (
"response" in first_line
and "request" in first_line
and "processed_time" in first_line
)
has_success_or_error = (
"candidates" in first_line.get("response", {})
or "promptFeedback" in first_line.get("response", {})
or bool(first_line.get("status"))
)
if not (has_base_structure and has_success_or_error):
# Not a Vertex AI batch output, return as-is
return content
vertex_gemini_config = VertexGeminiConfig()
# Always use a fresh local Logging object for the per-line transformation
# so we never mutate the caller's logging_obj (which already went through
# pre_call and has its own model/start_time/optional_params set).
batch_transform_logging_obj = Logging(
model="",
messages=[],
stream=False,
call_type="batch_transform",
start_time=time.time(),
litellm_call_id="",
function_id="",
)
batch_transform_logging_obj.optional_params = {}
mock_httpx_response = httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
request=httpx.Request(method="POST", url="https://example.com"),
)
# Transform all lines
transformed_lines = []
for line in lines:
if not line.strip():
continue
try:
vertex_output = json.loads(line)
openai_output = (
self._transform_single_vertex_batch_output_to_openai(
vertex_output=vertex_output,
vertex_gemini_config=vertex_gemini_config,
logging_obj=batch_transform_logging_obj,
mock_httpx_response=mock_httpx_response,
)
)
transformed_lines.append(json.dumps(openai_output))
except Exception:
# If any line fails, return original content
return content
# Return transformed content
return "\n".join(transformed_lines).encode("utf-8")
except Exception:
# If anything fails, return original content
return content
def _transform_single_vertex_batch_output_to_openai(
self,
vertex_output: Dict[str, Any],
vertex_gemini_config: VertexGeminiConfig,
logging_obj: Logging,
mock_httpx_response: httpx.Response,
) -> Dict[str, Any]:
"""
Transform a single Vertex AI batch output line to OpenAI format.
Uses the existing VertexGeminiConfig transformation for the response.
"""
# Extract custom_id from request labels (prefer raw for OpenAI round-trip)
request_data = vertex_output.get("request", {})
labels = request_data.get("labels", {}) or {}
custom_id = _get_litellm_batch_custom_id_from_labels(labels)
# Check if there's an error
status = vertex_output.get("status", "")
has_error = bool(status)
if has_error:
return {
"id": f"batch_req_{uuid.uuid4()}",
"custom_id": custom_id,
"response": None,
"error": {
"code": "vertex_ai_error",
"message": status,
},
}
# Transform successful response using existing transformation
vertex_response = vertex_output.get("response", {})
# Extract model from response
model = vertex_response.get("modelVersion", "gemini-1.5-flash-001")
if "@" in model:
model = model.split("@")[0]
try:
# Use existing VertexGeminiConfig transformation
model_response = ModelResponse()
transformed_response = vertex_gemini_config._transform_google_generate_content_to_openai_model_response(
completion_response=vertex_response,
model_response=model_response,
model=model,
logging_obj=logging_obj,
raw_response=mock_httpx_response,
)
# Convert ModelResponse to dict
response_dict = transformed_response.model_dump()
# Return in OpenAI batch format
return {
"id": f"batch_req_{uuid.uuid4()}",
"custom_id": custom_id,
"response": {
"status_code": 200,
"request_id": response_dict.get("id", ""),
"body": response_dict,
},
"error": None,
}
except Exception as e:
return {
"id": f"batch_req_{uuid.uuid4()}",
"custom_id": custom_id,
"response": None,
"error": {
"code": "transformation_error",
"message": f"Failed to transform response: {str(e)}",
},
}
class VertexAIJsonlFilesTransformation(VertexGeminiConfig):
"""
@ -500,29 +833,11 @@ class VertexAIJsonlFilesTransformation(VertexGeminiConfig):
def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content(
self, openai_jsonl_content: List[Dict[str, Any]]
):
"""
Transforms OpenAI JSONL content to VertexAI JSONL content
jsonl body for vertex is {"request": <request_body>}
Example Vertex jsonl
{"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}}
{"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}}
"""
vertex_jsonl_content = []
for _openai_jsonl_content in openai_jsonl_content:
openai_request_body = _openai_jsonl_content.get("body") or {}
vertex_request_body = _transform_request_body(
messages=openai_request_body.get("messages", []),
model=openai_request_body.get("model", ""),
optional_params=self._map_openai_to_vertex_params(openai_request_body),
custom_llm_provider="vertex_ai",
litellm_params={},
cached_content=None,
)
vertex_jsonl_content.append({"request": vertex_request_body})
return vertex_jsonl_content
) -> List[Dict[str, Any]]:
return _openai_batch_jsonl_entries_to_vertex_wrapped_requests(
openai_jsonl_content=openai_jsonl_content,
map_openai_to_vertex_params=self._map_openai_to_vertex_params,
)
def _get_gcs_object_name(
self,

View file

@ -212,6 +212,22 @@ def _process_gemini_media(
return _apply_gemini_metadata(
part, model, media_resolution_enum, video_metadata
)
elif image_url.startswith(
"https://generativelanguage.googleapis.com/v1beta/files/"
):
# Gemini Files API URIs — the file is already uploaded to Google's
# servers; pass the URI through as file_data without fetching it.
# These URLs return 403 when accessed directly, so we must not try
# to resolve their MIME type via HTTP.
if format:
file_data = FileDataType(mime_type=format, file_uri=image_url)
else:
# Gemini Files API references can be passed through as URI-only.
file_data = cast(FileDataType, {"file_uri": image_url})
part = {"file_data": file_data}
return _apply_gemini_metadata(
part, model, media_resolution_enum, video_metadata
)
elif (
"https://" in image_url
and (image_type := format or _get_image_mime_type_from_url(image_url))
@ -743,16 +759,22 @@ def _transform_request_body( # noqa: PLR0915
]
data = RequestBody(contents=content)
if system_instructions is not None:
data["system_instruction"] = system_instructions
if tools is not None:
data["tools"] = tools
if tool_choice is not None:
data["toolConfig"] = tool_choice
if include_server_side_tool_invocations:
if "toolConfig" not in data:
data["toolConfig"] = {}
data["toolConfig"]["includeServerSideToolInvocations"] = True
# Vertex rejects system_instruction/tools/toolConfig alongside cachedContent.
# Treat dropping these fields as a request mutation guarded by modify_params.
can_send_cache_incompatible_fields = (
cached_content is None or litellm.modify_params is False
)
if can_send_cache_incompatible_fields:
if system_instructions is not None:
data["system_instruction"] = system_instructions
if tools is not None:
data["tools"] = tools
if tool_choice is not None:
data["toolConfig"] = tool_choice
if include_server_side_tool_invocations:
if "toolConfig" not in data:
data["toolConfig"] = {}
data["toolConfig"]["includeServerSideToolInvocations"] = True
if safety_settings is not None:
data["safetySettings"] = safety_settings
if generation_config is not None and len(generation_config) > 0:

View file

@ -979,15 +979,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
params["includeThoughts"] = False
else:
params["includeThoughts"] = True
if thinking_budget >= 10000:
is_gemini3flash = (
"gemini-3-flash-preview" in model.lower()
or "gemini-3-flash" in model.lower()
)
params["thinkingLevel"] = (
"minimal" if is_gemini3flash else "low"
)
else:
# Follow provider defaults unless explicitly opted into legacy behavior.
if litellm.enable_gemini_default_thinking_level_low is True:
is_gemini3flash = (
"gemini-3-flash-preview" in model.lower()
or "gemini-3-flash" in model.lower()

View file

@ -13,6 +13,7 @@ from litellm.types.llms.vertex_ai import VertexPartnerProvider
from litellm.types.router import GenericLiteLLMParams
from ....vertex_llm_base import VertexBase
from ..output_params_utils import sanitize_vertex_anthropic_output_params
class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, VertexBase):
@ -158,12 +159,10 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert
"model", None
) # do not pass model in request body to vertex ai
anthropic_messages_request.pop(
"output_format", None
) # do not pass output_format in request body to vertex ai - vertex ai does not support output_format as yet
anthropic_messages_request.pop(
"output_config", None
) # do not pass output_config in request body to vertex ai - vertex ai does not support output_config
# Vertex AI Claude accepts ``output_config.format`` (structured outputs)
# and ``output_format``, but rejects ``output_config.effort`` with 400
# "Extra inputs are not permitted". Sanitize in place so the supported
# bits flow through.
sanitize_vertex_anthropic_output_params(anthropic_messages_request)
return anthropic_messages_request

View file

@ -0,0 +1,50 @@
"""
Shared sanitization for ``output_config`` / ``output_format`` on Vertex AI
Claude. Lives in its own module so both the chat-completion transformation
(``transformation.py``) and the Messages pass-through transformation
(``experimental_pass_through/transformation.py``) can import it without
forming a cycle through the parent module's heavier imports.
CodeQL flagged the ``..transformation`` import path as a potential cyclic
import; extracting the helper into a leaf module resolves the warning and
keeps the parent module's import surface narrow.
"""
# Keys inside ``output_config`` that Vertex AI Claude does not accept.
# Today only ``effort`` triggers "Extra inputs are not permitted"; add new
# entries here as Vertex parity drifts. Keep this list narrow — anything
# Vertex DOES accept (e.g. ``format`` for structured outputs) must be
# preserved so callers can rely on Anthropic-native features.
VERTEX_UNSUPPORTED_OUTPUT_CONFIG_KEYS: frozenset = frozenset({"effort"})
def sanitize_vertex_anthropic_output_params(data: dict) -> None:
"""
Strip Vertex-unsupported keys from ``output_config`` /
``output_format`` in-place; forward whatever remains.
Behavior:
* ``output_config`` containing only unsupported keys (e.g. ``effort``
alone) is removed entirely so the request body has no empty dict.
* ``output_config`` containing a mix of supported + unsupported keys
has the unsupported subset filtered out and the rest forwarded.
* ``output_config`` that is supported in full passes through unchanged.
* ``output_format`` is forwarded as-is (Vertex AI Claude accepts it).
* Non-dict values for ``output_config`` are dropped to avoid sending
malformed payloads downstream.
"""
output_config = data.get("output_config")
if output_config is None:
return
if not isinstance(output_config, dict):
data.pop("output_config", None)
return
sanitized = {
k: v
for k, v in output_config.items()
if k not in VERTEX_UNSUPPORTED_OUTPUT_CONFIG_KEYS
}
if sanitized:
data["output_config"] = sanitized
else:
data.pop("output_config", None)

View file

@ -10,6 +10,7 @@ from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
from ....anthropic.chat.transformation import AnthropicConfig
from .output_params_utils import sanitize_vertex_anthropic_output_params
class VertexAIError(Exception):
@ -105,11 +106,12 @@ class VertexAIAnthropicConfig(AnthropicConfig):
data.pop("model", None) # vertex anthropic doesn't accept 'model' parameter
# VertexAI doesn't support output_format parameter, remove it if present
data.pop("output_format", None)
# VertexAI doesn't support output_config parameter, remove it if present
data.pop("output_config", None)
# Vertex AI Claude accepts ``output_config.format`` (structured outputs /
# JSON Schema) but NOT ``output_config.effort`` — sending ``effort`` to
# Vertex returns 400 "Extra inputs are not permitted". Sanitize in place:
# forward the structured-output bits, drop the unsupported keys.
# Same treatment for the legacy top-level ``output_format`` field.
sanitize_vertex_anthropic_output_params(data)
tools = optional_params.get("tools")
tool_search_used = self.is_tool_search_used(tools)

View file

@ -4923,8 +4923,17 @@ def embedding( # noqa: PLR0915
if encoding_format is not None:
optional_params["encoding_format"] = encoding_format
else:
# Omiting causes openai sdk to add default value of "float"
optional_params["encoding_format"] = None
env_fmt = get_secret_str("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT")
if env_fmt is not None and env_fmt.strip().lower() == "none":
optional_params.pop("encoding_format", None)
else:
_default_fmt = (
optional_params.get("encoding_format") or env_fmt or "float"
)
if _default_fmt.strip().lower() == "none":
optional_params.pop("encoding_format", None)
else:
optional_params["encoding_format"] = _default_fmt
api_version = None

View file

@ -19928,7 +19928,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.5-2026-04-23": {
"cache_read_input_token_cost": 5e-07,
@ -19976,7 +19976,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.5-pro": {
"cache_read_input_token_cost": 3e-06,
@ -20019,7 +20019,8 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false,
"supports_low_reasoning_effort": false
},
"gpt-5.5-pro-2026-04-23": {
"cache_read_input_token_cost": 3e-06,
@ -20062,7 +20063,8 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false,
"supports_low_reasoning_effort": false
},
"gpt-5.4": {
"cache_read_input_token_cost": 2.5e-07,
@ -22061,6 +22063,98 @@
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true
},
"crusoe/deepseek-ai/DeepSeek-R1-0528": {
"input_cost_per_token": 3e-06,
"litellm_provider": "crusoe",
"max_input_tokens": 163840,
"max_output_tokens": 163840,
"max_tokens": 163840,
"mode": "chat",
"output_cost_per_token": 7e-06,
"supports_function_calling": false,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": false
},
"crusoe/deepseek-ai/DeepSeek-V3-0324": {
"input_cost_per_token": 1.5e-06,
"litellm_provider": "crusoe",
"max_input_tokens": 163840,
"max_output_tokens": 163840,
"max_tokens": 163840,
"mode": "chat",
"output_cost_per_token": 1.5e-06,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"crusoe/google/gemma-3-12b-it": {
"input_cost_per_token": 1e-07,
"litellm_provider": "crusoe",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1e-07,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true
},
"crusoe/meta-llama/Llama-3.3-70B-Instruct": {
"input_cost_per_token": 2e-07,
"litellm_provider": "crusoe",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2e-07,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"crusoe/moonshotai/Kimi-K2-Thinking": {
"input_cost_per_token": 2.5e-06,
"litellm_provider": "crusoe",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"supports_function_calling": false,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": false
},
"crusoe/openai/gpt-oss-120b": {
"input_cost_per_token": 8e-07,
"litellm_provider": "crusoe",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 8e-07,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507": {
"input_cost_per_token": 3e-06,
"litellm_provider": "crusoe",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 3e-06,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"lambda_ai/deepseek-llama3.3-70b": {
"input_cost_per_token": 2e-07,
"litellm_provider": "lambda_ai",

View file

@ -409,9 +409,12 @@ class MCPRequestHandler:
Permission hierarchy (all rules are intersections):
1. Get allowed servers from key permissions
2. Get allowed servers from team permissions
3. Get allowed servers from end_user permissions
4. Final result = intersection of key/team AND end_user (if end_user has permissions set)
2. Get allowed servers from team permissions (key inherits from team, or intersection)
3. Get allowed servers from end_user permissions (intersected if set)
4. Get allowed servers from agent permissions (intersected if set)
5. Get allowed servers from org permissions — org acts as a ceiling: if the org
has an explicit MCP server list, the combined key/team/end_user/agent result is
capped to that list. If the org has no list, no extra restriction is applied.
Returns:
List[str]: List of allowed MCP servers by server id
@ -435,6 +438,10 @@ class MCPRequestHandler:
# Calculate key/team allowed servers using inheritance and intersection logic
#########################################################
allowed_mcp_servers: List[str] = []
has_lower_level_mcp_restrictions = (
len(allowed_mcp_servers_for_key) > 0
or len(allowed_mcp_servers_for_team) > 0
)
if len(allowed_mcp_servers_for_team) > 0:
if len(allowed_mcp_servers_for_key) > 0:
# Key has its own MCP permissions - use intersection with team permissions
@ -459,6 +466,7 @@ class MCPRequestHandler:
# If end_user has explicit MCP server permissions, apply intersection
if len(allowed_mcp_servers_for_end_user) > 0:
has_lower_level_mcp_restrictions = True
verbose_logger.debug(
f"End user {user_api_key_auth.end_user_id} has explicit MCP permissions: {allowed_mcp_servers_for_end_user}"
)
@ -490,6 +498,7 @@ class MCPRequestHandler:
)
)
if len(allowed_mcp_servers_for_agent) > 0:
has_lower_level_mcp_restrictions = True
# Intersect: agent can only use servers allowed by BOTH key/team AND agent config
allowed_mcp_servers = [
s
@ -500,6 +509,30 @@ class MCPRequestHandler:
f"Applied agent intersection filter. Final allowed servers: {allowed_mcp_servers}"
)
#########################################################
# Apply org-level ceiling if org_id is set
#########################################################
if user_api_key_auth and user_api_key_auth.org_id:
allowed_mcp_servers_for_org = (
await MCPRequestHandler._get_allowed_mcp_servers_for_org(
user_api_key_auth
)
)
if len(allowed_mcp_servers_for_org) > 0:
if has_lower_level_mcp_restrictions:
# Lower-level restrictions exist, so org can only cap them.
allowed_mcp_servers = [
s
for s in allowed_mcp_servers
if s in allowed_mcp_servers_for_org
]
else:
# No lower-level restrictions → org list becomes the ceiling
allowed_mcp_servers = allowed_mcp_servers_for_org
verbose_logger.debug(
f"Applied org ceiling filter. Final allowed servers: {allowed_mcp_servers}"
)
return list(set(allowed_mcp_servers))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
@ -638,6 +671,27 @@ class MCPRequestHandler:
allowed_tools = list(set(allowed_tools) & set(agent_tools))
else:
allowed_tools = agent_tools
# Apply org-level tool ceiling if org_id is set
if user_api_key_auth.org_id:
# _get_org_object_permission uses user_api_key_cache, so this is not a
# fresh DB round-trip when get_allowed_mcp_servers was already called.
org_obj_perm = await MCPRequestHandler._get_org_object_permission(
user_api_key_auth
)
org_tools = (
global_mcp_server_manager.expand_tool_permissions(
org_obj_perm.mcp_tool_permissions
).get(server_id)
if org_obj_perm and org_obj_perm.mcp_tool_permissions
else None
)
if org_tools is not None:
if allowed_tools is not None:
allowed_tools = list(set(allowed_tools) & set(org_tools))
else:
allowed_tools = list(org_tools)
return allowed_tools
except Exception as e:
@ -805,6 +859,120 @@ class MCPRequestHandler:
)
return []
# Sentinel stored in cache when an org has no object_permission, so we
# don't re-query the DB on every MCP request for that org.
_ORG_NO_PERMISSION_SENTINEL = "__org_no_mcp_permission__"
@staticmethod
async def _get_org_object_permission(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
):
"""
Get org object_permission, using user_api_key_cache to avoid DB hits on every request.
Caches both positive results and the absence of an object_permission so that orgs
with no MCP permissions configured (the common default) do not trigger a DB query
on every request.
"""
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
if not user_api_key_auth or not user_api_key_auth.org_id:
return None
if prisma_client is None:
verbose_logger.debug("prisma_client is None")
return None
org_id = user_api_key_auth.org_id
cache_key = f"org_object_permission:{org_id}"
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
try:
cached = await user_api_key_cache.async_get_cache(key=cache_key)
if cached is not None:
# Sentinel means the DB confirmed no object_permission for this org
if cached == MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL:
return None
# Redis deserialises to a plain dict; reconstruct the Pydantic model
# so callers can access .mcp_servers / .mcp_tool_permissions as attrs.
if isinstance(cached, dict):
return LiteLLM_ObjectPermissionTable(**cached)
return cached
org_row = await prisma_client.db.litellm_organizationtable.find_unique(
where={"organization_id": org_id},
include={"object_permission": True},
)
if org_row is None or org_row.object_permission is None:
# Cache the negative result so subsequent calls skip the DB
await user_api_key_cache.async_set_cache(
key=cache_key,
value=MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL,
)
return None
# Convert raw Prisma model → Pydantic before caching. Caching the
# Pydantic .dict() ensures the value survives a Redis JSON round-trip
# as a plain dict that we can reconstruct above (same pattern used by
# get_end_user_object / get_team_object in auth_checks.py).
obj_perm = LiteLLM_ObjectPermissionTable(**org_row.object_permission.dict())
await user_api_key_cache.async_set_cache(
key=cache_key, value=obj_perm.dict()
)
return obj_perm
except Exception as e:
verbose_logger.warning(f"Failed to get org object permission: {str(e)}")
return None
@staticmethod
async def _get_allowed_mcp_servers_for_org(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
"""
Get allowed MCP servers for an organization.
Returns the MCP servers from the org's object_permission.
An empty result means the org places no restriction (allow-all from this level).
"""
try:
object_permissions = await MCPRequestHandler._get_org_object_permission(
user_api_key_auth
)
if object_permissions is None:
return []
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
# Expand names/aliases to canonical server IDs (consistent with key/team/end-user path)
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
object_permissions.mcp_servers or []
)
access_group_servers = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
)
)
tool_perm_servers = list(
global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).keys()
)
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(
f"Failed to get allowed MCP servers for org: {str(e)}"
)
return []
@staticmethod
async def _get_allowed_mcp_servers_for_end_user(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,

View file

@ -2138,6 +2138,47 @@ if MCP_AVAILABLE:
#########################################################
local_tool = global_mcp_tool_registry.get_tool(name)
if local_tool:
# OpenAPI-backed tools used to bypass `pre_call_tool_check` —
# only the managed path ran allowed/banned-tool checks, key/team
# tool permissions, and parameter validation. Run the same checks
# before dispatching to the local registry. Refuse the call if
# we cannot resolve a server: tools registered via
# openapi_to_mcp_generator are always tied to a server, so a
# missing mcp_server here means the tool->server mapping has
# not finished initializing or the registry entry is orphaned.
# Skipping the check would re-open the same authorization gap.
if mcp_server is None:
raise HTTPException(
status_code=503,
detail=(
f"MCP server for tool '{name}' is not available; "
"refusing to dispatch without authorization checks. "
"Retry once the server is registered."
),
)
# `pre_call_tool_check` calls into `proxy_logging_obj` for the
# pre-call guardrail hooks, so source it from the canonical
# `proxy_server` module the same way `_handle_managed_mcp_tool`
# does. `kwargs.get("proxy_logging_obj")` is None on the MCP
# entry path and would crash with AttributeError after the
# security checks pass.
from litellm.proxy.proxy_server import proxy_logging_obj
hook_result = await global_mcp_server_manager.pre_call_tool_check(
name=original_tool_name,
arguments=arguments or {},
server_name=server_name or mcp_server.name,
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
server=mcp_server,
raw_headers=raw_headers,
)
# `pre_call_tool_check` may return guardrail-modified
# arguments; honor them on the local path too.
if isinstance(hook_result, dict) and "arguments" in hook_result:
arguments = hook_result["arguments"]
verbose_logger.debug(f"Executing local registry tool: {name}")
# For BYOK servers the credential must be injected via a ContextVar
# because the tool function has headers baked into its closure.

View file

@ -3616,7 +3616,7 @@
},
"get": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__get",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -3660,7 +3660,7 @@
},
"patch": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__patch",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -3704,7 +3704,7 @@
},
"post": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__post",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -3748,7 +3748,7 @@
},
"put": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__put",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13299,7 +13299,7 @@
},
"get": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__get",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13338,7 +13338,7 @@
},
"patch": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__patch",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13377,7 +13377,7 @@
},
"post": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__post",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13416,7 +13416,7 @@
},
"put": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__put",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -14008,7 +14008,7 @@
"/mcp-rest/test/connection": {
"post": {
"description": "Test if we can connect to the provided MCP server before adding it",
"operationId": "test_connection_mcp_rest_test_connection_post",
"operationId": "test_connection_mcp_rest_test_connection_post_2",
"requestBody": {
"content": {
"application/json": {
@ -14053,7 +14053,7 @@
"/mcp-rest/test/tools/list": {
"post": {
"description": "Preview tools available from MCP server before adding it",
"operationId": "test_tools_list_mcp_rest_test_tools_list_post",
"operationId": "test_tools_list_mcp_rest_test_tools_list_post_2",
"requestBody": {
"content": {
"application/json": {
@ -14098,7 +14098,7 @@
"/mcp-rest/tools/call": {
"post": {
"description": "REST API to call a specific MCP tool with the provided arguments",
"operationId": "call_tool_rest_api_mcp_rest_tools_call_post",
"operationId": "call_tool_rest_api_mcp_rest_tools_call_post_2",
"responses": {
"200": {
"content": {
@ -14123,7 +14123,7 @@
"/mcp-rest/tools/list": {
"get": {
"description": "List all available tools with information about the server they belong to.\n\nExample response:\n{\n \"tools\": [\n {\n \"name\": \"create_zap\",\n \"description\": \"Create a new zap\",\n \"inputSchema\": \"tool_input_schema\",\n \"mcp_info\": {\n \"server_name\": \"zapier\",\n \"logo_url\": \"https://www.zapier.com/logo.png\",\n }\n }\n ],\n \"error\": null,\n \"message\": \"Successfully retrieved tools\"\n}",
"operationId": "list_tool_rest_api_mcp_rest_tools_list_get",
"operationId": "list_tool_rest_api_mcp_rest_tools_list_get_2",
"parameters": [
{
"description": "The server id to list tools for",
@ -21896,7 +21896,7 @@
"/policies/usage/overview": {
"get": {
"description": "Return policy performance overview for the dashboard.",
"operationId": "policies_usage_overview_policies_usage_overview_get",
"operationId": "policies_usage_overview_policies_usage_overview_get_2",
"parameters": [
{
"description": "YYYY-MM-DD",
@ -22521,7 +22521,7 @@
"/policies/attachments/estimate-impact": {
"post": {
"description": "Estimate how many keys and teams would be affected by a policy attachment.\n\nUse this before creating an attachment to preview the blast radius.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/attachments/estimate-impact\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"policy_name\": \"hipaa-compliance\",\n \"tags\": [\"healthcare\", \"health-*\"]\n }'\n```",
"operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post",
"operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post_2",
"requestBody": {
"content": {
"application/json": {
@ -22568,7 +22568,7 @@
"/policies/resolve": {
"post": {
"description": "Resolve which policies and guardrails apply for a given context.\n\nUse this endpoint to debug \"what guardrails would apply to a request\nwith this team/key/model/tags combination?\"\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/resolve\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"tags\": [\"healthcare\"],\n \"model\": \"gpt-4\"\n }'\n```",
"operationId": "resolve_policies_for_context_policies_resolve_post",
"operationId": "resolve_policies_for_context_policies_resolve_post_2",
"parameters": [
{
"description": "Force a DB sync before resolving. Default uses in-memory cache.",
@ -26922,7 +26922,7 @@
},
"get": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_get",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -26961,7 +26961,7 @@
},
"head": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_head",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27000,7 +27000,7 @@
},
"options": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27039,7 +27039,7 @@
},
"patch": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_patch",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27078,7 +27078,7 @@
},
"post": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_post",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27117,7 +27117,7 @@
},
"put": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -28329,7 +28329,7 @@
"/v1/vector_stores": {
"get": {
"description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list",
"operationId": "vector_store_list_v1_vector_stores_get",
"operationId": "vector_store_list_v1_vector_stores_get_2",
"parameters": [
{
"in": "query",
@ -28430,7 +28430,7 @@
},
"post": {
"description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```",
"operationId": "vector_store_create_v1_vector_stores_post",
"operationId": "vector_store_create_v1_vector_stores_post_2",
"responses": {
"200": {
"content": {
@ -28455,7 +28455,7 @@
"/v1/vector_stores/{vector_store_id}": {
"delete": {
"description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete",
"operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete",
"operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete_2",
"parameters": [
{
"in": "path",
@ -28499,7 +28499,7 @@
},
"get": {
"description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve",
"operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get",
"operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get_2",
"parameters": [
{
"in": "path",
@ -28543,7 +28543,7 @@
},
"post": {
"description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify",
"operationId": "vector_store_update_v1_vector_stores__vector_store_id__post",
"operationId": "vector_store_update_v1_vector_stores__vector_store_id__post_2",
"parameters": [
{
"in": "path",
@ -28588,7 +28588,7 @@
},
"/v1/vector_stores/{vector_store_id}/files": {
"get": {
"operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get",
"operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get_2",
"parameters": [
{
"in": "path",
@ -28631,7 +28631,7 @@
]
},
"post": {
"operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post",
"operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post_2",
"parameters": [
{
"in": "path",
@ -28676,7 +28676,7 @@
},
"/v1/vector_stores/{vector_store_id}/files/{file_id}": {
"delete": {
"operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete",
"operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete_2",
"parameters": [
{
"in": "path",
@ -28728,7 +28728,7 @@
]
},
"get": {
"operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get",
"operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get_2",
"parameters": [
{
"in": "path",
@ -28780,7 +28780,7 @@
]
},
"post": {
"operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post",
"operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post_2",
"parameters": [
{
"in": "path",
@ -28834,7 +28834,7 @@
},
"/v1/vector_stores/{vector_store_id}/files/{file_id}/content": {
"get": {
"operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get",
"operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get_2",
"parameters": [
{
"in": "path",
@ -28889,7 +28889,7 @@
"/v1/vector_stores/{vector_store_id}/search": {
"post": {
"description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search",
"operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post",
"operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post_2",
"parameters": [
{
"in": "path",
@ -28935,7 +28935,7 @@
"/vector_stores": {
"get": {
"description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list",
"operationId": "vector_store_list_vector_stores_get",
"operationId": "vector_store_list_vector_stores_get_2",
"parameters": [
{
"in": "query",
@ -29036,7 +29036,7 @@
},
"post": {
"description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```",
"operationId": "vector_store_create_vector_stores_post",
"operationId": "vector_store_create_vector_stores_post_2",
"responses": {
"200": {
"content": {
@ -29061,7 +29061,7 @@
"/vector_stores/{vector_store_id}": {
"delete": {
"description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete",
"operationId": "vector_store_delete_vector_stores__vector_store_id__delete",
"operationId": "vector_store_delete_vector_stores__vector_store_id__delete_2",
"parameters": [
{
"in": "path",
@ -29105,7 +29105,7 @@
},
"get": {
"description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve",
"operationId": "vector_store_retrieve_vector_stores__vector_store_id__get",
"operationId": "vector_store_retrieve_vector_stores__vector_store_id__get_2",
"parameters": [
{
"in": "path",
@ -29149,7 +29149,7 @@
},
"post": {
"description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify",
"operationId": "vector_store_update_vector_stores__vector_store_id__post",
"operationId": "vector_store_update_vector_stores__vector_store_id__post_2",
"parameters": [
{
"in": "path",
@ -29194,7 +29194,7 @@
},
"/vector_stores/{vector_store_id}/files": {
"get": {
"operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get",
"operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get_2",
"parameters": [
{
"in": "path",
@ -29237,7 +29237,7 @@
]
},
"post": {
"operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post",
"operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post_2",
"parameters": [
{
"in": "path",
@ -29282,7 +29282,7 @@
},
"/vector_stores/{vector_store_id}/files/{file_id}": {
"delete": {
"operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete",
"operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete_2",
"parameters": [
{
"in": "path",
@ -29334,7 +29334,7 @@
]
},
"get": {
"operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get",
"operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get_2",
"parameters": [
{
"in": "path",
@ -29386,7 +29386,7 @@
]
},
"post": {
"operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post",
"operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post_2",
"parameters": [
{
"in": "path",
@ -29440,7 +29440,7 @@
},
"/vector_stores/{vector_store_id}/files/{file_id}/content": {
"get": {
"operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get",
"operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get_2",
"parameters": [
{
"in": "path",
@ -29495,7 +29495,7 @@
"/vector_stores/{vector_store_id}/search": {
"post": {
"description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search",
"operationId": "vector_store_search_vector_stores__vector_store_id__search_post",
"operationId": "vector_store_search_vector_stores__vector_store_id__search_post_2",
"parameters": [
{
"in": "path",

View file

@ -8,12 +8,35 @@ any drift as a neutral check.
"""
import json
import re
import sys
from pathlib import Path
from typing import Dict, Optional
from typing import Dict, Optional, Set
SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json"
HTTP_METHODS = {"delete", "get", "head", "options", "patch", "post", "put"}
HTTP_METHOD_SUFFIXES = {
"delete",
"get",
"head",
"options",
"patch",
"post",
"put",
"trace",
}
def _stabilize_multi_method_route_ids(routes) -> None:
"""FastAPI derives route IDs from a set of methods; make snapshots stable."""
for route in routes:
methods = sorted(getattr(route, "methods", None) or [])
if len(methods) <= 1 or not getattr(route, "path_format", None):
continue
operation_id = f"{route.name}{route.path_format}"
operation_id = re.sub(r"\W", "_", operation_id)
route.unique_id = f"{operation_id}_{methods[0].lower()}"
def load_snapshot() -> Optional[Dict[str, Dict]]:
@ -38,12 +61,12 @@ def _normalize_operation_ids(paths: Dict[str, Dict]) -> None:
if not isinstance(path_ops, dict):
continue
methods = {method for method in path_ops if method in HTTP_METHODS}
methods = {method for method in path_ops if method in HTTP_METHOD_SUFFIXES}
if not methods:
continue
for method, operation in path_ops.items():
if method not in HTTP_METHODS or not isinstance(operation, dict):
if method not in HTTP_METHOD_SUFFIXES or not isinstance(operation, dict):
continue
operation_id = operation.get("operationId")
@ -65,7 +88,7 @@ def generate_snapshot() -> Dict[str, Dict]:
from fastapi.openapi.utils import get_openapi
from litellm.proxy._lazy_features import LAZY_FEATURES
from litellm.proxy.proxy_server import app
from litellm.proxy.proxy_server import app, ensure_unique_openapi_operation_ids
for feat in LAZY_FEATURES:
if feat.module_path in sys.modules:
@ -77,6 +100,7 @@ def generate_snapshot() -> Dict[str, Dict]:
sys.stderr.write(f"warning: skip {feat.name}: {exc}\n")
fragments: Dict[str, Dict] = {}
used_operation_ids: Set[str] = set()
for feat in LAZY_FEATURES:
feat_routes = [
r
@ -85,14 +109,24 @@ def generate_snapshot() -> Dict[str, Dict]:
]
if not feat_routes:
continue
_stabilize_multi_method_route_ids(feat_routes)
full = get_openapi(title=app.title, version=app.version, routes=feat_routes)
paths = full.get("paths", {})
_normalize_operation_ids(paths)
# Group all of a feature's routes under one tag.
for path_ops in paths.values():
for op in path_ops.values():
for path_ops in full.get("paths", {}).values():
for method, op in path_ops.items():
if isinstance(op, dict):
operation_id = op.get("operationId")
if isinstance(operation_id, str):
for suffix in HTTP_METHOD_SUFFIXES:
if operation_id.endswith(f"_{suffix}"):
op["operationId"] = (
operation_id[: -len(suffix)] + method
)
break
op["tags"] = [feat.name]
full = ensure_unique_openapi_operation_ids(full, used_operation_ids)
fragments[feat.name] = {
"paths": paths,
"components": {"schemas": full.get("components", {}).get("schemas", {})},

View file

@ -724,21 +724,73 @@ class LiteLLMRoutes(enum.Enum):
"/organization/member_delete",
]
# Routes accessible by Admin Viewer (read-only admin access)
admin_viewer_routes = [
"/user/list",
"/user/available_users",
"/user/available_roles",
"/user/daily/activity",
"/team/daily/activity",
"/tag/daily/activity",
"/tag/list",
"/audit",
"/audit/{id}",
"/global/activity",
"/global/activity/model",
"/global/activity/cache_hits",
] + info_routes
# Routes accessible by Admin Viewer (read-only admin access).
#
# Admin Viewer follows a read-parity-with-Proxy-Admin rule: anything Proxy
# Admin can read/list/get, Admin Viewer can too (no writes, no cost-incurring
# actions).
#
# NOTE: This list is no longer the primary mechanism for granting access —
# `_check_proxy_admin_viewer_access()` in route_checks.py default-allows
# any safe HTTP method (GET/HEAD/OPTIONS) on non-inference routes. This
# list now matters only for non-GET routes that are semantically reads
# (e.g. POST /spend/calculate). Adding a new GET endpoint does not require
# updating this list — the default-allow behavior covers it automatically.
admin_viewer_routes = (
[
"/user/list",
"/user/available_users",
"/user/available_roles",
"/user/daily/activity",
"/team/daily/activity",
"/tag/daily/activity",
"/tag/list",
"/audit",
"/audit/{id}",
"/global/activity",
"/global/activity/model",
"/global/activity/cache_hits",
# Customer / end-user listing (handlers already gate on
# PROXY_ADMIN_VIEW_ONLY — the route gate must match).
"/customer/list",
"/customer/info",
# UI Logs page detail drawer (single + session). The list endpoint
# `/spend/logs/ui` is covered via spend_tracking_routes below.
"/spend/logs/ui/{logId}",
"/spend/logs/session/ui",
# Settings / observability read endpoints exposed in admin-only
# sidebar groups (Logging & Alerts, Admin Settings, Budgets,
# Invitations).
"/callbacks/list",
"/callbacks/configs",
"/get/config/callbacks",
"/alerting/settings",
"/config/list",
"/config/field/info",
"/budget/list",
"/budget/settings",
# Invitation viewing (admin viewer cannot create/delete; can read).
"/invitation/info",
# Guardrails / Policies pages (read-only views).
"/guardrails/list",
"/v2/guardrails/list",
"/guardrails/submissions",
"/guardrails/submissions/{guardrail_id}",
"/guardrails/usage/overview",
"/policies/attachments/list",
# MCP semantic filter settings (read).
"/get/mcp_semantic_filter_settings",
# Model cost map maintenance views (read-only status / source).
"/schedule/model_cost_map_reload/status",
"/model/cost_map/source",
]
# Spend tracking reads (/spend/logs, /spend/logs/ui, /spend/keys,
# /spend/users, /spend/tags, /spend/calculate, /cost/estimate). Admin
# Viewer can already read /global/spend/* via global_spend_tracking_routes;
# the per-tenant /spend/* views were the missing peer.
+ spend_tracking_routes
+ info_routes
)
# All routes accesible by an Org Admin
org_admin_allowed_routes = (
@ -2386,6 +2438,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For headers are only trusted from these IPs.",
)
trusted_proxy_ranges: Optional[List[str]] = Field(
None,
description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.",
)
store_model_in_db: Optional[bool] = Field(
None,
description="If True, models and config are stored in and loaded from the database. Default is False.",
@ -2579,6 +2635,7 @@ class UserAPIKeyAuth(
user_spend: Optional[float] = None
user_max_budget: Optional[float] = None
request_route: Optional[str] = None
budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True)
user: Optional[Any] = None # Expanded user object when expand=user is used
created_by_user: Optional[Any] = (
None # Expanded created_by user when expand=user is used

View file

@ -60,6 +60,10 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
_safe_get_request_query_params,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
TOOL_CAPABLE_CALL_TYPES,
@ -486,7 +490,10 @@ async def common_checks( # noqa: PLR0915
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
_model: Optional[Union[str, List[str]]] = get_model_from_request(
request_body, route
request_data=request_body,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
)
# 1. If team is blocked
@ -495,23 +502,28 @@ async def common_checks( # noqa: PLR0915
f"Team={team_object.team_id} is blocked. Update via `/team/unblock` if you're an admin."
)
# 2. If team can call model
# 2. If team can call model (or key's access_group_ids grant it)
if _model and team_object:
with tracer.trace("litellm.proxy.auth.common_checks.can_team_access_model"):
if not await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=(
valid_token.team_model_aliases if valid_token else None
),
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
try:
await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=(
valid_token.team_model_aliases if valid_token else None
),
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
raise
if not await _key_access_group_grants_model(
model=_model,
valid_token=valid_token,
team_object=team_object,
llm_router=llm_router,
):
raise
# 2.2. If team member has per-member model scope, enforce it
if _model and team_object and valid_token and valid_token.user_id:
@ -656,13 +668,7 @@ async def common_checks( # noqa: PLR0915
end_user_object is not None
and end_user_object.litellm_budget_table is not None
):
end_user_budget = end_user_object.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_object.spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_object.spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
)
await _check_end_user_budget(end_user_obj=end_user_object, route=route)
_enforce_user_param_check(general_settings, request, request_body, route)
_reject_clientside_metadata_tags_check(general_settings, request_body, route)
@ -1012,7 +1018,7 @@ async def _apply_default_budget_to_end_user(
return end_user_obj
def _check_end_user_budget(
async def _check_end_user_budget(
end_user_obj: LiteLLM_EndUserTable,
route: str,
) -> None:
@ -1033,11 +1039,20 @@ def _check_end_user_budget(
return
end_user_budget = end_user_obj.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_obj.spend > end_user_budget:
if end_user_budget is None:
return
from litellm.proxy.proxy_server import get_current_spend
end_user_spend = await get_current_spend(
counter_key=f"spend:end_user:{end_user_obj.user_id}",
fallback_spend=end_user_obj.spend or 0.0,
)
if end_user_spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_obj.spend,
current_cost=end_user_spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_obj.spend}, Budget={end_user_budget}",
message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_spend}, Budget={end_user_budget}",
)
@ -1091,7 +1106,7 @@ async def get_end_user_object(
)
# Check budget limits
_check_end_user_budget(end_user_obj=return_obj, route=route)
await _check_end_user_budget(end_user_obj=return_obj, route=route)
return return_obj
@ -1124,7 +1139,7 @@ async def get_end_user_object(
)
# Check budget limits
_check_end_user_budget(end_user_obj=_response, route=route)
await _check_end_user_budget(end_user_obj=_response, route=route)
return _response
@ -1616,9 +1631,12 @@ async def _cache_key_object(
## CACHE REFRESH TIME
user_api_key_obj.last_refreshed_at = time.time()
cached_key_obj = _copy_user_api_key_auth_for_cache(
user_api_key_obj=user_api_key_obj
)
await _cache_management_object(
key=key,
value=user_api_key_obj,
value=cached_key_obj,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
model_type=UserAPIKeyAuth,
@ -2348,7 +2366,7 @@ async def get_key_object(
model_type=UserAPIKeyAuth,
)
if user_api_key_auth is not None:
return user_api_key_auth
return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth)
if check_cache_only:
raise Exception(
@ -2401,6 +2419,16 @@ async def get_key_object(
return _response
def _copy_user_api_key_auth_for_cache(
user_api_key_obj: UserAPIKeyAuth,
) -> UserAPIKeyAuth:
copied_key_obj = user_api_key_obj.model_copy()
copied_key_obj.budget_reservation = None
copied_key_obj.parent_otel_span = None
copied_key_obj.request_route = None
return copied_key_obj
@log_db_metrics
async def get_object_permission(
object_permission_id: str,
@ -2952,6 +2980,77 @@ async def can_team_access_model(
raise
async def _key_access_group_grants_model(
model: Union[str, List[str]],
valid_token: Optional[UserAPIKeyAuth],
team_object: Optional[LiteLLM_TeamTable],
llm_router: Optional[Router],
) -> bool:
"""
Returns True if the key's `access_group_ids` expand to models that grant
access to `model`. Used to let a key's access group override a team's
model restriction in `common_checks`.
A key's access group only counts if the access group itself authorizes the
caller as an owner — that is, the group's `assigned_team_ids` includes the
key's `team_id`, or the group's `assigned_key_ids` includes the key's
token. This preserves the team-as-owner boundary (a team member cannot
escalate by naming a group assigned to a different team) while still
letting a group reach the key without first being added to the team's
`access_group_ids` list.
"""
if valid_token is None:
return False
key_access_group_ids = list(valid_token.access_group_ids or [])
if not key_access_group_ids:
return False
from litellm.proxy.proxy_server import prisma_client as _prisma_client
from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging_obj
from litellm.proxy.proxy_server import user_api_key_cache as _user_api_key_cache
if _prisma_client is None or _user_api_key_cache is None:
return False
key_team_id = valid_token.team_id or (
team_object.team_id if team_object is not None else None
)
key_token = valid_token.token
authorized_models: List[str] = []
for ag_id in key_access_group_ids:
try:
ag = await get_access_object(
access_group_id=ag_id,
prisma_client=_prisma_client,
user_api_key_cache=_user_api_key_cache,
proxy_logging_obj=_proxy_logging_obj,
)
except Exception:
continue
team_authorized = bool(
key_team_id and key_team_id in (ag.assigned_team_ids or [])
)
key_authorized = bool(key_token and key_token in (ag.assigned_key_ids or []))
if team_authorized or key_authorized:
authorized_models.extend(ag.access_model_names or [])
if not authorized_models:
return False
try:
_can_object_call_model(
model=model,
llm_router=llm_router,
models=list(set(authorized_models)),
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
object_type="key",
)
return True
except ProxyException:
return False
def can_project_access_model(
model: Union[str, List[str]],
project_object: LiteLLM_ProjectTableCachedObj,
@ -3967,13 +4066,19 @@ async def _tag_max_budget_check(
if (
tag_object.litellm_budget_table is not None
and tag_object.litellm_budget_table.max_budget is not None
and tag_object.spend is not None
and tag_object.spend > tag_object.litellm_budget_table.max_budget
):
from litellm.proxy.proxy_server import get_current_spend
tag_spend = await get_current_spend(
counter_key=f"spend:tag:{tag_name}",
fallback_spend=tag_object.spend or 0.0,
)
if tag_spend <= tag_object.litellm_budget_table.max_budget:
continue
raise litellm.BudgetExceededError(
current_cost=tag_object.spend,
current_cost=tag_spend,
max_budget=tag_object.litellm_budget_table.max_budget,
message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_object.spend}, Max budget: {tag_object.litellm_budget_table.max_budget}",
message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_spend}, Max budget: {tag_object.litellm_budget_table.max_budget}",
)

View file

@ -2,7 +2,7 @@ import os
import re
import sys
from functools import lru_cache
from typing import Any, List, Optional, Tuple
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
from fastapi import HTTPException, Request, status
@ -976,20 +976,257 @@ def get_end_user_id_from_request_body(
return None
def get_model_from_request(
request_data: dict, route: str
) -> Optional[Union[str, List[str]]]:
# First try to get model from request_data
model = request_data.get("model") or request_data.get("target_model_names")
MODEL_ROUTING_HEADER_NAME = "x-litellm-model"
_MODEL_ROUTING_ROUTE_MARKERS = (
"/files",
"/batches",
"/vector_stores",
"/skills",
"/evals",
"/fine_tuning",
"/videos",
)
_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS = (
"/files",
"/batches",
"/skills",
"/evals",
)
_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS = (
"/files",
"/batches",
"/fine_tuning",
)
_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS = (
"/files",
"/batches",
"/vector_stores",
)
_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS = ("/evals",)
_MODEL_ROUTING_ID_FIELDS = (
"file_id",
"input_file_id",
"output_file_id",
"error_file_id",
"batch_id",
"fine_tuning_job_id",
"training_file",
"validation_file",
"vector_store_id",
"video_id",
"character_id",
)
if model is not None:
model_names = model.split(",")
if len(model_names) == 1:
model = model_names[0].strip()
def _append_model_candidates(candidates: List[str], value: Any) -> None:
if value is None:
return
values = value if isinstance(value, (list, tuple, set)) else [value]
for item in values:
if item is None:
continue
if isinstance(item, str):
model_names = [model.strip() for model in item.split(",")]
else:
model = [m.strip() for m in model_names]
model_names = [str(item).strip()]
candidates.extend(model for model in model_names if model)
# If model not in request_data, try to extract from route
def _dedupe_model_candidates(candidates: List[str]) -> List[str]:
deduped: List[str] = []
for model in candidates:
if model not in deduped:
deduped.append(model)
return deduped
def _get_case_insensitive_mapping_value(
mapping: Optional[Mapping[str, Any]], key: str
) -> Any:
if not mapping:
return None
if key in mapping:
return mapping[key]
key_lower = key.lower()
for mapping_key, value in mapping.items():
if str(mapping_key).lower() == key_lower:
return value
return None
def _route_matches_any_marker(route: str, markers: Tuple[str, ...]) -> bool:
normalized_route = route.lower()
return any(marker in normalized_route for marker in markers)
def _route_uses_model_routing_sources(route: str) -> bool:
return _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_ROUTE_MARKERS)
def _extract_models_from_managed_resource_id(
resource_id: Any, resource_id_field: Optional[str] = None
) -> List[str]:
if not isinstance(resource_id, str) or not resource_id:
return []
candidates: List[str] = []
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
)
_append_model_candidates(
candidates=candidates, value=decode_model_from_file_id(resource_id)
)
unified_file_id = _is_base64_encoded_unified_file_id(resource_id)
if unified_file_id:
_append_model_candidates(
candidates=candidates,
value=get_models_from_unified_file_id(unified_file_id),
)
_append_model_candidates(
candidates=candidates,
value=get_model_id_from_unified_batch_id(unified_file_id),
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from managed file/batch ID: %s", str(e)
)
try:
from litellm.llms.base_llm.managed_resources.utils import parse_unified_id
parsed_id = parse_unified_id(resource_id)
if parsed_id:
_append_model_candidates(
candidates=candidates, value=parsed_id.get("model_id")
)
_append_model_candidates(
candidates=candidates, value=parsed_id.get("target_model_names")
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from unified managed resource ID: %s", str(e)
)
if resource_id_field in ("video_id", "character_id"):
try:
from litellm.types.videos.utils import (
decode_character_id_with_provider,
decode_video_id_with_provider,
)
if resource_id_field == "video_id":
_append_model_candidates(
candidates=candidates,
value=decode_video_id_with_provider(resource_id).get("model_id"),
)
else:
_append_model_candidates(
candidates=candidates,
value=decode_character_id_with_provider(resource_id).get(
"model_id"
),
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from managed video/character ID: %s", str(e)
)
return _dedupe_model_candidates(candidates)
def _extract_model_candidates_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
) -> List[str]:
candidates: List[str] = []
uses_model_routing_sources = _route_uses_model_routing_sources(route=route)
uses_header_or_query_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS
)
uses_query_target_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS
)
uses_body_target_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS
)
uses_completion_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS
)
body_model = request_data.get("model")
_append_model_candidates(candidates, body_model)
if uses_body_target_model_sources or not body_model:
_append_model_candidates(candidates, request_data.get("target_model_names"))
if uses_completion_model_sources and isinstance(
request_data.get("completion"), dict
):
_append_model_candidates(candidates, request_data["completion"].get("model"))
if uses_model_routing_sources:
if uses_header_or_query_model_sources:
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(request_query_params, "model"),
)
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(
request_headers, MODEL_ROUTING_HEADER_NAME
),
)
if uses_query_target_model_sources:
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(
request_query_params, "target_model_names"
),
)
for field in _MODEL_ROUTING_ID_FIELDS:
_append_model_candidates(
candidates,
_extract_models_from_managed_resource_id(
request_data.get(field), resource_id_field=field
),
)
return _dedupe_model_candidates(candidates)
def _format_model_candidates(
candidates: List[str],
) -> Optional[Union[str, List[str]]]:
if not candidates:
return None
if len(candidates) == 1:
return candidates[0]
return candidates
def get_model_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
) -> Optional[Union[str, List[str]]]:
candidates = _extract_model_candidates_from_request(
request_data=request_data,
route=route,
request_headers=request_headers,
request_query_params=request_query_params,
)
model = _format_model_candidates(candidates)
# If no explicit model was found, try to extract from route
if model is None:
# Parse model from route that follows the pattern /openai/deployments/{model}/*
match = re.match(r"/openai/deployments/([^/]+)", route)

View file

@ -707,11 +707,48 @@ class JWTHandler:
verbose_proxy_logger.error(f"Error fetching OIDC UserInfo: {str(e)}")
raise Exception(f"Failed to fetch OIDC UserInfo: {str(e)}")
async def auth_jwt(self, token: str) -> dict:
_unscoped_jwt_warning_emitted = False
@classmethod
def _build_decode_kwargs(cls) -> dict:
"""Build the audience/issuer/options kwargs for ``jwt.decode``.
Setting ``JWT_AUDIENCE`` (and optionally ``JWT_ISSUER``) turns on the
corresponding PyJWT verifications, blocking cross-tenant tokens
minted by other applications that share the same IdP signing keys.
When both are unset PyJWT only checks the signature and expiry, which
is preserved for backward compatibility but logged once as a warning.
"""
audience = os.getenv("JWT_AUDIENCE")
decode_options = None
issuer = os.getenv("JWT_ISSUER")
if (
audience is None
and issuer is None
and not cls._unscoped_jwt_warning_emitted
):
verbose_proxy_logger.warning(
"JWT auth is enabled but neither JWT_AUDIENCE nor JWT_ISSUER "
"is configured. Tokens minted by any application that shares "
"the same IdP signing keys will be accepted. Set JWT_AUDIENCE "
"(and ideally JWT_ISSUER) to scope this proxy."
)
cls._unscoped_jwt_warning_emitted = True
options: dict = {}
if audience is None:
decode_options = {"verify_aud": False}
options["verify_aud"] = False
if issuer is None:
options["verify_iss"] = False
return {
"audience": audience,
"issuer": issuer,
"options": options or None,
}
async def auth_jwt(self, token: str) -> dict:
decode_kwargs = self._build_decode_kwargs()
header = jwt.get_unverified_header(token)
@ -747,9 +784,8 @@ class JWTHandler:
token,
public_key_obj, # type: ignore
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
options=decode_options, # type: ignore[arg-type]
audience=audience,
leeway=self.leeway, # allow testing of expired tokens
**decode_kwargs,
)
return payload
@ -775,8 +811,7 @@ class JWTHandler:
token,
key,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
audience=audience,
options=decode_options,
**decode_kwargs,
)
return payload

View file

@ -1,19 +1,69 @@
from typing import Any, Dict
from typing import Any, Dict, FrozenSet
from fastapi import Request
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.trusted_proxy_utils import require_trusted_proxy_request
# OAuth2-proxy header trust is for **identity assertion** from a trusted
# upstream auth proxy (oauth2-proxy, Authelia, etc.). The allowlist below
# is the only safe surface — anything else (``user_role``, ``api_key``,
# ``permissions``, ``max_budget``, ``user_max_budget``,
# ``team_tpm_limit``, ``end_user_max_budget``, ``allowed_model_region``,
# and dozens of similar policy fields scattered across the
# ``LiteLLM_VerificationTokenView`` hierarchy) is a privilege grant that
# would let a caller forge their own enforcement parameters by sending
# the matching header.
#
# A denylist of "privileged fields" is unmaintainable in this codebase:
# the auth model has ~50 budget/spend/limit/permission fields and gains
# more with each release. An allowlist scoped to identity assertion is
# default-secure — new fields are blocked automatically.
#
# Operators who need a trusted upstream to assert anything beyond
# identity should switch to JWT authentication, which validates a
# signature on the assertion rather than blindly trusting headers.
ALLOWED_OAUTH2_PROXY_FIELDS: FrozenSet[str] = frozenset(
{
"user_id",
"user_email",
"team_id",
"team_alias",
"org_id",
"models",
}
)
async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
"""
Handle request from oauth2 proxy.
Resolve a ``UserAPIKeyAuth`` from request headers per the admin-set
``oauth2_config_mappings``.
The auth model assumes the proxy is deployed behind a trusted OAuth2
reverse proxy that injects authenticated identity headers (e.g.
oauth2-proxy, Authelia).
**Identity-only allowlist.** ``oauth2_config_mappings`` maps header
names to ``UserAPIKeyAuth`` fields. Without an allowlist, an admin
who maps the wrong header to ``user_role`` lets any caller send
``X-User-Role: proxy_admin`` and gain full admin privileges
(Pydantic coerces the string into the enum). Only fields in
``ALLOWED_OAUTH2_PROXY_FIELDS`` (identity assertion only — see the
constant's comment) may be mapped; any other mapping is rejected at
request time so the misconfiguration surfaces loudly rather than as
a silent privesc.
"""
from litellm.proxy.proxy_server import general_settings
verbose_proxy_logger.debug("Handling oauth2 proxy request")
# Define the OAuth2 config mappings
require_trusted_proxy_request(
request=request,
general_settings=general_settings,
feature_name="OAuth2 proxy auth",
)
oauth2_config_mappings: Dict[str, str] = (
general_settings.get("oauth2_config_mappings") or {}
)
@ -21,21 +71,32 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
if not oauth2_config_mappings:
raise ValueError("Oauth2 config mappings not found in general_settings")
# Initialize a dictionary to store the mapped values
auth_data: Dict[str, Any] = {}
# Extract values from headers based on the mappings
disallowed = sorted(
set(oauth2_config_mappings.keys()) - ALLOWED_OAUTH2_PROXY_FIELDS
)
if disallowed:
raise ValueError(
"Oauth2 proxy auth refuses to map non-identity UserAPIKeyAuth "
f"fields from request headers: {disallowed}. Only identity "
f"fields are accepted ({sorted(ALLOWED_OAUTH2_PROXY_FIELDS)}); "
"anything else (privileges, budgets, rate limits, metadata) "
"would let a caller forge enforcement parameters by spoofing "
"the matching header. If you need a trusted upstream to "
"assert anything beyond identity, use JWT auth "
"(signature-validated) instead of header-trust."
)
auth_data: Dict[str, Any] = {}
for key, header in oauth2_config_mappings.items():
value = request.headers.get(header)
if value:
# Convert max_budget to float if present
if key == "max_budget":
auth_data[key] = float(value)
# Convert models to list if present
elif key == "models":
auth_data[key] = [model.strip() for model in value.split(",")]
else:
auth_data[key] = value
if not value:
continue
if key == "models":
auth_data[key] = [model.strip() for model in value.split(",")]
else:
auth_data[key] = value
verbose_proxy_logger.debug(
"Auth data before creating UserAPIKeyAuth object: keys=%s",
list(auth_data.keys()),
@ -45,5 +106,4 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
"UserAPIKeyAuth object created with keys: %s",
list(user_api_key_auth.__fields_set__),
)
# Create and return UserAPIKeyAuth object
return user_api_key_auth

View file

@ -202,6 +202,7 @@ class RouteChecks:
route=route,
_user_role=_user_role,
request_data=request_data,
request=request,
)
elif (
_user_role == LitellmUserRoles.INTERNAL_USER.value
@ -596,14 +597,66 @@ class RouteChecks:
return True
return False
# HTTP methods that are intrinsically read-only and therefore safe to
# default-allow for PROXY_ADMIN_VIEW_ONLY. Anything else (POST/PUT/PATCH/
# DELETE) is treated as a write attempt and goes through the explicit
# write-allowlist below.
_SAFE_HTTP_METHODS = frozenset({"GET", "HEAD", "OPTIONS"})
# Explicit write routes that PROXY_ADMIN_VIEW_ONLY must NEVER call. The
# role-principle is "no writes, ever" — the management_routes list is the
# authoritative source for which non-llm routes are writes; we just need
# to filter out the read endpoints (info / list) that share the prefix.
# A cleaner approach is to denylist by HTTP verb (POST/PUT/PATCH/DELETE);
# this block stays as a backstop in case a write is implemented as GET.
_ADMIN_VIEWER_BLOCKED_WRITE_ROUTES = frozenset(
[
"/user/new",
"/user/delete",
"/user/bulk_update",
"/team/new",
"/team/update",
"/team/delete",
"/model/new",
"/model/update",
"/model/delete",
"/key/generate",
"/key/delete",
"/key/update",
"/key/regenerate",
"/key/service-account/generate",
"/key/block",
"/key/unblock",
]
)
@staticmethod
def _check_proxy_admin_viewer_access(
route: str,
_user_role: str,
request_data: dict,
request: Optional[Request] = None,
) -> None:
"""
Check access for PROXY_ADMIN_VIEW_ONLY role
Check access for PROXY_ADMIN_VIEW_ONLY role.
Admin Viewer follows a read-parity-with-Proxy-Admin rule: anything Proxy
Admin can read/list/get, Admin Viewer can read/list/get. The only
exclusions are cost-incurring inference routes (Playground, /chat/
completions, etc.) and any state-mutating request.
Implementation:
1. LLM/inference routes → 403 (cost-incurring).
2. Safe HTTP method (GET/HEAD/OPTIONS) → allow by default. This is
the read-parity guarantee — every new GET endpoint added anywhere
in the codebase is automatically readable by Admin Viewer
without needing to remember to add it to an allowlist.
3. Unsafe HTTP method (POST/PUT/PATCH/DELETE):
- Allow `/user/update` only when restricted to user_email/password.
- Block all explicit writes in `_ADMIN_VIEWER_BLOCKED_WRITE_ROUTES`.
- Otherwise allow only if the route is in admin_viewer_routes /
global_spend_tracking_routes (legacy explicit-allow set).
- Else 403.
"""
if RouteChecks.is_llm_api_route(route=route):
raise HTTPException(
@ -611,65 +664,60 @@ class RouteChecks:
detail=f"user not allowed to access this OpenAI routes, role= {_user_role}",
)
# Check if this is a write operation on management routes
if RouteChecks.check_route_access(
route=route, allowed_routes=LiteLLMRoutes.management_routes.value
):
# For management routes, only allow read operations or specific allowed updates
if route == "/user/update":
# Check the Request params are valid for PROXY_ADMIN_VIEW_ONLY
if request_data is not None and isinstance(request_data, dict):
_params_updated = request_data.keys()
for param in _params_updated:
if param not in ["user_email", "password"]:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route} and updating invalid param: {param}. only user_email and password can be updated",
)
elif (
route
in [
"/user/new",
"/user/delete",
"/user/bulk_update",
"/team/new",
"/team/update",
"/team/delete",
"/model/new",
"/model/update",
"/model/delete",
"/key/generate",
"/key/delete",
"/key/update",
"/key/regenerate",
"/key/service-account/generate",
"/key/block",
"/key/unblock",
]
or route.startswith("/key/")
and route.endswith("/regenerate")
):
# Block write operations for PROXY_ADMIN_VIEW_ONLY
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}",
)
# Allow read operations on management routes (like /user/info, /team/info, /model/info)
method = request.method.upper() if request is not None else "GET"
is_safe_method = method in RouteChecks._SAFE_HTTP_METHODS
# ── Safe HTTP method: default-allow ──────────────────────────────
if is_safe_method:
return
elif RouteChecks.check_route_access(
route=route, allowed_routes=LiteLLMRoutes.admin_viewer_routes.value
):
# Allow access to admin viewer routes (read-only admin endpoints)
# ── Unsafe HTTP method: explicit checks ──────────────────────────
# Allow `/user/update` for self-service email / password change.
if route == "/user/update":
if request_data is not None and isinstance(request_data, dict):
for param in request_data.keys():
if param not in ["user_email", "password"]:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=(
f"user not allowed to access this route, role= {_user_role}. "
f"Trying to access: {route} and updating invalid param: {param}. "
"only user_email and password can be updated"
),
)
return
elif RouteChecks.check_route_access(
route=route, allowed_routes=LiteLLMRoutes.global_spend_tracking_routes.value
# Hard-block known write routes regardless of HTTP method (defensive
# — these are POSTs in practice, but pinning them here protects
# against future GET-shaped writes).
if route in RouteChecks._ADMIN_VIEWER_BLOCKED_WRITE_ROUTES or (
route.startswith("/key/") and route.endswith("/regenerate")
):
# Allow access to global spend tracking routes (read-only spend endpoints)
# proxy_admin_viewer role description: "view all keys, view all spend"
return
else:
# For other routes, block access
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}",
)
# Legacy explicit-allow sets (kept for routes that are POST but
# semantically read-only, e.g. /spend/calculate). Both admin_viewer_routes
# and global_spend_tracking_routes are reads/listings.
if RouteChecks.check_route_access(
route=route, allowed_routes=LiteLLMRoutes.admin_viewer_routes.value
):
return
if RouteChecks.check_route_access(
route=route, allowed_routes=LiteLLMRoutes.global_spend_tracking_routes.value
):
return
# NOTE: We intentionally do NOT fall back to allowing all
# `management_routes`. That set is a mix of reads (info/list — handled
# via the safe-method branch above) and writes (`/team/block`,
# `/team/permissions_update`, `/jwt/key/mapping/{new,update,delete}`,
# `/key/bulk_update`, `/key/{id}/reset_spend`). A blanket allow would
# let Admin Viewer POST these write endpoints — violating the
# "no writes, ever" rule. Default-deny instead.
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}",
)

View file

@ -0,0 +1,118 @@
import ipaddress
from typing import Any, Dict, List, Optional, Union
from fastapi import Request
from litellm._logging import verbose_proxy_logger
TRUSTED_PROXY_RANGES_KEY = "trusted_proxy_ranges"
TrustedProxyNetwork = Union[ipaddress.IPv4Network, ipaddress.IPv6Network]
def _get_proxy_general_settings() -> Dict[str, Any]:
try:
from litellm.proxy.proxy_server import general_settings
return general_settings or {}
except ImportError:
return {}
def _normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str) -> List[str]:
if not configured_ranges:
return []
if isinstance(configured_ranges, str):
return [
raw_range.strip()
for raw_range in configured_ranges.split(",")
if raw_range.strip()
]
if isinstance(configured_ranges, (list, tuple, set)):
return [
str(raw_range).strip()
for raw_range in configured_ranges
if str(raw_range).strip()
]
verbose_proxy_logger.warning(
"Invalid %s value: expected a list of CIDR ranges, got %s",
setting_name,
type(configured_ranges).__name__,
)
return []
def parse_trusted_proxy_ranges(
configured_ranges: Any,
*,
setting_name: str = TRUSTED_PROXY_RANGES_KEY,
) -> List[TrustedProxyNetwork]:
networks: List[TrustedProxyNetwork] = []
for cidr in _normalize_cidr_ranges(configured_ranges, setting_name=setting_name):
try:
networks.append(ipaddress.ip_network(cidr, strict=False))
except ValueError:
verbose_proxy_logger.warning(
"Invalid CIDR in %s: %s, skipping", setting_name, cidr
)
return networks
def _get_direct_client_ip(request: Request) -> Optional[str]:
client = getattr(request, "client", None)
client_host = getattr(client, "host", None)
if isinstance(client_host, str):
return client_host
return None
def _is_ip_in_networks(
client_ip: Optional[str], networks: List[TrustedProxyNetwork]
) -> bool:
if not client_ip or not networks:
return False
try:
addr = ipaddress.ip_address(client_ip.strip())
except ValueError:
return False
return any(addr in network for network in networks)
def require_trusted_proxy_request(
*,
request: Request,
general_settings: Optional[Dict[str, Any]] = None,
feature_name: str,
setting_name: str = TRUSTED_PROXY_RANGES_KEY,
) -> None:
"""
Fail closed unless the direct TCP peer is one of the configured
trusted reverse proxies.
Header-based auth paths must validate the direct peer, not
X-Forwarded-For, because the direct peer is the actor supplying the
identity headers.
"""
if general_settings is None:
general_settings = _get_proxy_general_settings()
trusted_networks = parse_trusted_proxy_ranges(
general_settings.get(setting_name), setting_name=setting_name
)
if not trusted_networks:
raise ValueError(
f"{feature_name} requires general_settings.{setting_name} before "
"trusting identity headers from an upstream proxy."
)
direct_client_ip = _get_direct_client_ip(request)
if not _is_ip_in_networks(direct_client_ip, trusted_networks):
verbose_proxy_logger.warning(
"%s rejected identity headers from untrusted direct client IP %r",
feature_name,
direct_client_ip,
)
raise ValueError(
f"{feature_name} only accepts identity headers from configured "
f"trusted proxy ranges. Direct client IP {direct_client_ip!r} "
"is not trusted."
)

View file

@ -11,7 +11,7 @@ import asyncio
import re
import secrets
from datetime import datetime, timezone
from typing import Any, List, Optional, Tuple, cast
from typing import Any, Iterator, List, Optional, Tuple, Union, cast
import fastapi
from fastapi import HTTPException, Request, WebSocket, status
@ -63,6 +63,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
_safe_get_request_query_params,
populate_request_with_path_params,
)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
@ -118,6 +119,29 @@ azure_apim_header = APIKeyHeader(
)
def _get_model_from_request_context(
request_data: dict,
route: str,
request: Optional[Request],
) -> Optional[Union[str, List[str]]]:
return get_model_from_request(
request_data=request_data,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
)
def _get_model_names_for_budget_checks(
model: Optional[Union[str, List[str]]],
) -> List[str]:
if model is None:
return []
if isinstance(model, str):
return [model]
return model
def _get_bearer_token_or_received_api_key(api_key: str) -> str:
if api_key.startswith("Bearer "): # ensure Bearer token passed in
api_key = api_key.replace("Bearer ", "") # extract the token
@ -884,7 +908,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
# Check if model has zero cost - if so, skip all budget checks
model = get_model_from_request(request_data, route)
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
@ -1254,6 +1282,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
valid_token=valid_token,
request_data=request_data,
route=route,
request=request,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
@ -1279,7 +1308,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
user_obj = None
# Check 2a. Check if model has zero cost - if so, skip all budget checks
model = get_model_from_request(request_data, route)
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
@ -1403,21 +1436,29 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Check 5. Token Model Spend is under Model budget
max_budget_per_model = valid_token.model_max_budget
current_model = request_data.get("model", None)
current_model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
current_models = _get_model_names_for_budget_checks(
model=current_model
)
if (
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and prisma_client is not None
and current_model is not None
and current_models
and valid_token.token is not None
):
## GET THE SPEND FOR THIS MODEL
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=model_name,
)
# Check 5b. End-user model max budget
end_user_mmb = valid_token.end_user_model_max_budget
@ -1425,14 +1466,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and current_models
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=model_name,
)
# Check 6: Additional Common Checks across jwt + key auth
if valid_token.team_id is not None:
@ -1862,10 +1904,12 @@ async def _run_centralized_common_checks(
user_api_key_auth_obj.project_metadata = project_object.metadata
user_api_key_auth_obj.project_alias = project_object.project_alias
skip_budget_checks = False
model = get_model_from_request(request_data, route)
if model is not None and llm_router is not None:
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
skip_budget_checks = _should_skip_budget_checks(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
_ = await common_checks(
request=request,
@ -1883,6 +1927,21 @@ async def _run_centralized_common_checks(
project_object=project_object,
)
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request_data=request_data,
route=route,
llm_router=llm_router,
team_object=team_object,
user_object=user_object,
end_user_id=end_user_id,
end_user_object=end_user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
skip_budget_checks=skip_budget_checks,
)
async def _noop_none() -> None:
"""Sentinel coroutine for asyncio.gather when a fetch is unnecessary
@ -1890,6 +1949,59 @@ async def _noop_none() -> None:
return None
async def _reserve_budget_after_common_checks(
user_api_key_auth_obj: UserAPIKeyAuth,
request_data: dict,
route: str,
llm_router: Optional[Any],
team_object: Optional[LiteLLM_TeamTableCachedObj],
user_object: Optional[LiteLLM_UserTable],
prisma_client: Optional[PrismaClient],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
skip_budget_checks: bool,
end_user_id: Optional[str] = None,
end_user_object: Optional[LiteLLM_EndUserTable] = None,
) -> None:
user_api_key_auth_obj.budget_reservation = None
if skip_budget_checks:
return
from litellm.proxy.spend_tracking.budget_reservation import (
reserve_budget_for_request,
)
user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request(
request_body=request_data,
route=route,
llm_router=llm_router,
valid_token=user_api_key_auth_obj,
team_object=team_object,
user_object=user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
end_user_id=end_user_id,
end_user_object=end_user_object,
)
def _should_skip_budget_checks(
request_data: dict,
route: str,
request: Optional[Request],
llm_router: Optional[Any],
) -> bool:
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
if model is not None and llm_router is not None:
return _is_model_cost_zero(model=model, llm_router=llm_router)
return False
@tracer.wrap()
async def user_api_key_auth(
request: Request,
@ -1927,6 +2039,7 @@ async def user_api_key_auth(
request_data=request_data,
custom_litellm_key_header=custom_litellm_key_header,
)
user_api_key_auth_obj.budget_reservation = None
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj)
@ -2134,6 +2247,7 @@ async def _enforce_key_and_fallback_model_access(
valid_token: UserAPIKeyAuth,
request_data: dict,
route: str,
request: Optional[Request],
llm_model_list: Optional[list],
llm_router: Optional[Any],
) -> None:
@ -2152,10 +2266,10 @@ async def _enforce_key_and_fallback_model_access(
):
pass
else:
model = get_model_from_request(request_data, route)
fallback_models = cast(
Optional[List[ALL_FALLBACK_MODEL_VALUES]],
request_data.get("fallbacks", None),
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
if model is not None:
@ -2166,20 +2280,69 @@ async def _enforce_key_and_fallback_model_access(
llm_router=llm_router,
)
if fallback_models is not None:
for m in fallback_models:
await can_key_call_model(
model=m["model"] if isinstance(m, dict) else m,
llm_model_list=llm_model_list,
valid_token=valid_token,
llm_router=llm_router,
)
await is_valid_fallback_model(
model=m["model"] if isinstance(m, dict) else m,
llm_router=llm_router,
user_model=None,
# Validate every fallback model name reachable by this request.
# All three fields (``fallbacks``, ``context_window_fallbacks``,
# ``content_policy_fallbacks``) are forwarded to the router as
# per-request kwargs whether they appear at the top level of
# ``request_data`` or nested under ``router_settings_override``.
# Both surfaces must be validated against the API key's model
# allowlist or a caller can smuggle a restricted model. VERIA-44.
fallback_names: List[str] = []
override_settings = request_data.get("router_settings_override")
for _fb_key in ROUTER_FALLBACK_FIELDS:
fallback_names.extend(
iter_router_fallback_model_names(request_data.get(_fb_key))
)
if isinstance(override_settings, dict):
fallback_names.extend(
iter_router_fallback_model_names(override_settings.get(_fb_key))
)
for _name in dict.fromkeys(fallback_names): # dedupe, preserve order
await can_key_call_model(
model=_name,
llm_model_list=llm_model_list,
valid_token=valid_token,
llm_router=llm_router,
)
await is_valid_fallback_model(
model=_name,
llm_router=llm_router,
user_model=None,
)
ROUTER_FALLBACK_FIELDS: Tuple[str, ...] = (
"fallbacks",
"context_window_fallbacks",
"content_policy_fallbacks",
)
def iter_router_fallback_model_names(fallbacks: Any) -> Iterator[str]:
"""Yield leaf model names from any of the supported fallbacks shapes.
Handles the simple top-level shape (``str`` or ``{"model": str}``) and
the nested router-config shape (``[{primary: [fallback_list]}]``).
"""
if not isinstance(fallbacks, list):
return
for entry in fallbacks:
if isinstance(entry, str):
yield entry
elif isinstance(entry, dict):
if isinstance(entry.get("model"), str):
yield entry["model"]
continue
for fallback_list in entry.values():
if not isinstance(fallback_list, list):
continue
for m in fallback_list:
if isinstance(m, str):
yield m
elif isinstance(m, dict) and isinstance(m.get("model"), str):
yield m["model"]
async def _run_post_custom_auth_checks(
valid_token: UserAPIKeyAuth,
@ -2239,11 +2402,17 @@ async def _run_post_custom_auth_checks(
valid_token=valid_token,
request_data=request_data,
route=route,
request=request,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
current_model = request_data.get("model", None)
current_model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
current_models = _get_model_names_for_budget_checks(model=current_model)
# 3. Check key-level model_max_budget
max_budget_per_model = valid_token.model_max_budget
@ -2251,13 +2420,14 @@ async def _run_post_custom_auth_checks(
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and current_model is not None
and current_models
and valid_token.token is not None
):
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=model_name,
)
# 4. Check end-user model_max_budget
end_user_mmb = valid_token.end_user_model_max_budget
@ -2265,14 +2435,15 @@ async def _run_post_custom_auth_checks(
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and current_models
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=model_name,
)
# team / user / end_user / project context objects are fetched by
# the centralized common_checks gate in user_api_key_auth after

View file

@ -97,6 +97,55 @@ def _serialize_http_exception_detail(
return str(detail), None
def _collect_response_file_search_vector_store_ids(data: Dict[str, Any]) -> set[str]:
vector_store_ids: set[str] = set()
tools = data.get("tools")
if not isinstance(tools, list):
return vector_store_ids
for tool in tools:
if not isinstance(tool, dict) or tool.get("type") != "file_search":
continue
ids = tool.get("vector_store_ids") or []
if not isinstance(ids, list):
raise HTTPException(
status_code=400,
detail={
"error": "file_search.vector_store_ids must be a list of strings"
},
)
for vector_store_id in ids:
if not isinstance(vector_store_id, str) or not vector_store_id:
raise HTTPException(
status_code=400,
detail={
"error": "file_search.vector_store_ids must be a list of strings"
},
)
vector_store_ids.add(vector_store_id)
return vector_store_ids
async def _authorize_response_file_search_vector_stores(
data: Dict[str, Any],
user_api_key_dict: UserAPIKeyAuth,
) -> None:
vector_store_ids = _collect_response_file_search_vector_store_ids(data)
if not vector_store_ids:
return
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
)
for vector_store_id in sorted(vector_store_ids):
await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]:
"""Parses an event line and returns an error code if present, else None."""
event_line = (
@ -791,6 +840,11 @@ class ProxyBaseLLMRequestProcessing:
version=version,
proxy_config=proxy_config,
)
if route_type in {"aresponses", "_aresponses_websocket"}:
await _authorize_response_file_search_vector_stores(
data=self.data,
user_api_key_dict=user_api_key_dict,
)
# Calculate request queue time after add_litellm_data_to_request
# which sets arrival_time in proxy_server_request
@ -1604,6 +1658,12 @@ class ProxyBaseLLMRequestProcessing:
# here would duplicate the guardrail API call
# (e.g. double OpenAI Moderation charges).
continue
if "async_post_call_streaming_iterator_hook" in type(cb).__dict__:
# Skip — the guardrail already scanned the assembled
# response via its own streaming iterator hook in the
# streaming pipeline. re running this function async_post_call_success_hook
# here would duplicate the scan and can spuriously block the guardrail that already passed / failed.
continue
else:
guardrail_result = await cb.async_post_call_success_hook(
user_api_key_dict=captured_user_api_key_dict,

View file

@ -60,6 +60,7 @@ from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import (
ToolDiscoveryQueue,
)
from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient, ProxyLogging
@ -192,17 +193,16 @@ class DBSpendUpdateWriter:
verbose_proxy_logger.debug("Runs spend update on all tables")
except Exception:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue "
"may not have completed for this request. "
"response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s - %s",
"response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s",
response_cost,
token,
user_id,
team_id,
org_id,
end_user_id,
traceback.format_exc(),
)
def _enqueue_tool_registry_upsert(
@ -491,9 +491,7 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.exception(
f"Update Key DB Call failed to execute - {str(e)}"
)
spend_log_error("Update Key DB Call failed to execute - %s", str(e), exc=e)
raise e
async def _update_user_db(
@ -540,14 +538,14 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to enqueue user spend update. "
"user_id=%s, end_user_id=%s, response_cost=%s - %s\n%s",
"user_id=%s, end_user_id=%s, response_cost=%s - %s",
user_id,
end_user_id,
response_cost,
str(e),
traceback.format_exc(),
exc=e,
)
async def _update_team_db(
@ -585,23 +583,23 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to enqueue team member spend update. "
"team_id=%s, user_id=%s, response_cost=%s - %s\n%s",
"team_id=%s, user_id=%s, response_cost=%s - %s",
team_id,
user_id,
response_cost,
str(e),
traceback.format_exc(),
exc=e,
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to enqueue team spend update. "
"team_id=%s, response_cost=%s - %s\n%s",
"team_id=%s, response_cost=%s - %s",
team_id,
response_cost,
str(e),
traceback.format_exc(),
exc=e,
)
raise e
@ -626,13 +624,13 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to enqueue org spend update. "
"org_id=%s, response_cost=%s - %s\n%s",
"org_id=%s, response_cost=%s - %s",
org_id,
response_cost,
str(e),
traceback.format_exc(),
exc=e,
)
raise e
@ -654,13 +652,13 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to enqueue agent spend update. "
"agent_id=%s, response_cost=%s - %s\n%s",
"agent_id=%s, response_cost=%s - %s",
agent_id,
response_cost,
str(e),
traceback.format_exc(),
exc=e,
)
raise e
@ -707,13 +705,13 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to enqueue tag spend update. "
"request_tags=%s, response_cost=%s - %s\n%s",
"request_tags=%s, response_cost=%s - %s",
request_tags,
response_cost,
str(e),
traceback.format_exc(),
exc=e,
)
raise e
@ -906,11 +904,11 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_agent_spend_update_transactions,
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to commit spend updates from Redis to DB. "
"Data already popped from Redis may be lost. Error: %s\n%s",
"Data already popped from Redis may be lost. Error: %s",
str(e),
traceback.format_exc(),
exc=e,
)
finally:
await self.pod_lock_manager.release_lock(
@ -1074,11 +1072,11 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_tag_spend_update_transactions,
)
except Exception as e:
verbose_proxy_logger.error(
spend_log_error(
"Spend tracking - failed to commit daily tag spend updates from Redis to DB. "
"Data already popped from Redis may be lost. Error: %s\n%s",
"Data already popped from Redis may be lost. Error: %s",
str(e),
traceback.format_exc(),
exc=e,
)
finally:
await self.pod_lock_manager.release_lock(
@ -1736,11 +1734,15 @@ class DBSpendUpdateWriter:
except Exception as batch_error:
# Log detailed error information for debugging batch upsert failures
# This helps diagnose issues like unique constraint violations
verbose_proxy_logger.exception(
f"Daily {entity_type} spend batch upsert failed. "
f"Table: {table_name}, Constraint: {unique_constraint_name}, "
f"Batch size: {len(transactions_to_process)}, "
f"Error: {str(batch_error)}"
spend_log_error(
"Daily %s spend batch upsert failed. "
"Table: %s, Constraint: %s, Batch size: %d, Error: %s",
entity_type,
table_name,
unique_constraint_name,
len(transactions_to_process),
str(batch_error),
exc=batch_error,
)
raise

View file

@ -14,10 +14,12 @@ memory in long-lived deployments.
import asyncio
from collections import OrderedDict
from datetime import datetime
from typing import TYPE_CHECKING, ClassVar, Optional
from litellm._logging import verbose_proxy_logger
from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
if TYPE_CHECKING:
from litellm.caching.dual_cache import DualCache
@ -35,6 +37,10 @@ class SpendCounterReseed:
spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend
spend:user:{user_id} -> LiteLLM_UserTable.spend
spend:org:{org_id} -> LiteLLM_OrganizationTable.spend
End-user and tag spend counters intentionally do not reseed here. Their
auth paths already load the corresponding objects via get_end_user_object()
and get_tag_objects_batch(); callers pass those values as fallback_spend.
"""
_locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict()
@ -69,9 +75,10 @@ class SpendCounterReseed:
"""
if prisma_client is None:
return None
# Per-window counters share prefixes with primary counters but
# don't correspond to a DB row.
if ":window:" in counter_key:
# Per-window key/team counters share prefixes with primary counters
# but don't correspond to a DB row. Do not reject arbitrary entity IDs
# or tag names that merely contain ":window:".
if SpendCounterReseed._is_key_or_team_window_counter(counter_key):
return None
try:
if counter_key.startswith("spend:key:"):
@ -97,6 +104,10 @@ class SpendCounterReseed:
row = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}
)
elif counter_key.startswith("spend:end_user:"):
return None
elif counter_key.startswith("spend:tag:"):
return None
elif counter_key.startswith("spend:org:"):
org_id = counter_key[len("spend:org:") :]
row = await prisma_client.db.litellm_organizationtable.find_unique(
@ -113,11 +124,27 @@ class SpendCounterReseed:
return None
return float(getattr(row, "spend", 0.0) or 0.0)
@staticmethod
def _is_key_or_team_window_counter(counter_key: str) -> bool:
for prefix in ("spend:key:", "spend:team:"):
if not counter_key.startswith(prefix):
continue
_, separator, duration = counter_key.rpartition(":window:")
if not separator or not duration:
return False
try:
duration_in_seconds(duration)
except Exception:
return False
return True
return False
@staticmethod
async def coalesced(
prisma_client: Optional["PrismaClient"],
spend_counter_cache: "DualCache",
counter_key: str,
require_cache_warm: bool = False,
) -> Optional[float]:
"""
Reseed a cold spend counter from the DB and warm the cache,
@ -152,12 +179,156 @@ class SpendCounterReseed:
return None
# Warm even when 0 so subsequent reads hit cache, not DB.
try:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=db_spend, refresh_ttl=True
)
if spend_counter_cache.redis_cache is not None:
current_value = (
await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=db_spend,
refresh_ttl=True,
)
)
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
else:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=db_spend, refresh_ttl=True
)
except Exception:
verbose_proxy_logger.exception(
"SpendCounterReseed.coalesced: failed to warm counter %s",
counter_key,
)
if require_cache_warm:
raise
return db_spend
@staticmethod
async def window_from_spend_logs(
prisma_client: Optional["PrismaClient"],
entity_type: str,
entity_id: str,
window_start: datetime,
) -> Optional[float]:
if prisma_client is None:
return None
if entity_type == "Key":
group_field = "api_key"
where = {
"api_key": entity_id,
"startTime": {"gte": window_start},
}
elif entity_type == "Team":
group_field = "team_id"
where = {
"team_id": entity_id,
"startTime": {"gte": window_start},
}
else:
return None
try:
response = await prisma_client.db.litellm_spendlogs.group_by(
by=[group_field],
where=where, # type: ignore[arg-type]
sum={"spend": True},
)
except Exception:
verbose_proxy_logger.exception(
"SpendCounterReseed.window_from_spend_logs: failed for %s=%s",
entity_type,
entity_id,
)
return None
if not response:
return 0.0
first_row = response[0]
sum_row = (
first_row.get("_sum")
if isinstance(first_row, dict)
else getattr(first_row, "_sum", None)
)
spend = (
sum_row.get("spend")
if isinstance(sum_row, dict)
else getattr(sum_row, "spend", None)
)
return float(spend or 0.0)
@staticmethod
async def coalesced_window(
prisma_client: Optional["PrismaClient"],
spend_counter_cache: "DualCache",
counter_key: str,
entity_type: str,
entity_id: str,
window_start: datetime,
) -> Optional[float]:
lock = await SpendCounterReseed._get_lock(counter_key)
async with lock:
redis_clean_miss = False
if spend_counter_cache.redis_cache is not None:
try:
val = await spend_counter_cache.redis_cache.async_get_cache(
key=counter_key
)
if val is not None:
return float(val)
redis_clean_miss = True
except Exception:
pass
if not redis_clean_miss:
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
if val is not None:
return float(val)
window_spend = await SpendCounterReseed.window_from_spend_logs(
prisma_client=prisma_client,
entity_type=entity_type,
entity_id=entity_id,
window_start=window_start,
)
if window_spend is None:
return None
try:
if spend_counter_cache.redis_cache is not None:
seeded = await spend_counter_cache.redis_cache.async_set_cache(
key=counter_key,
value=window_spend,
nx=True,
)
if seeded:
current_value = window_spend
else:
current_cached_value = (
await spend_counter_cache.redis_cache.async_get_cache(
key=counter_key
)
)
if current_cached_value is None:
current_value = (
await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=window_spend,
)
)
else:
current_value = float(current_cached_value)
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
else:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=window_spend
)
except Exception:
verbose_proxy_logger.exception(
"SpendCounterReseed.coalesced_window: failed to warm counter %s",
counter_key,
)
raise
return window_spend

View file

@ -21,6 +21,7 @@ from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
build_sandbox_globals,
compile_sandboxed,
@ -842,7 +843,10 @@ async def list_guardrail_submissions(
if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
# Admin Viewer follows the read-parity rule: see all submissions like a
# Proxy Admin would (no writes — registration / approval still gated
# elsewhere by their own per-action checks).
is_admin = _user_has_admin_view(user_api_key_dict)
visible_team_ids: Optional[List[str]] = None
if not is_admin:
visible_team_ids = await _get_user_team_ids(user_api_key_dict)

View file

@ -1160,14 +1160,38 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
from litellm.types.utils import ModelResponse
all_chunks: List[ModelResponseStream] = []
passthrough_due_to_unknown_stream_shape = False
try:
async for chunk in response:
if isinstance(chunk, ModelResponseStream):
all_chunks.append(chunk)
if passthrough_due_to_unknown_stream_shape:
yield chunk
else:
all_chunks.append(chunk)
elif isinstance(chunk, bytes):
yield chunk # type: ignore[misc]
continue
else:
if all_chunks:
# Flush buffered chunks and switch to transparent passthrough for this stream shape.
# NOTE: these buffered chunks are emitted unmasked because this
# stream mixed chunk types and cannot be safely reconstructed.
verbose_proxy_logger.warning(
"Presidio apply_to_output: mixed stream detected (ModelResponseStream + unknown event). "
"Flushing %d buffered chunks without PII masking and switching to transparent passthrough.",
len(all_chunks),
)
for buffered_chunk in all_chunks:
yield buffered_chunk
all_chunks = []
passthrough_due_to_unknown_stream_shape = True
yield chunk
if passthrough_due_to_unknown_stream_shape:
verbose_proxy_logger.warning(
"Presidio apply_to_output: streaming response contained unknown event objects "
"(e.g. /v1/responses events). Output PII masking was skipped for this response."
)
return
if not all_chunks:
verbose_proxy_logger.warning(
"Presidio apply_to_output: streaming response contained only "

View file

@ -0,0 +1,35 @@
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .qohash import QostodianNexus
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
import litellm
_instance = QostodianNexus(
api_base=litellm_params.api_base,
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
additional_provider_specific_params=litellm_params.additional_provider_specific_params,
extra_headers=getattr(litellm_params, "extra_headers", None),
)
litellm.logging_callback_manager.add_litellm_callback(_instance)
return _instance
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.QOSTODIAN_NEXUS.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.QOSTODIAN_NEXUS.value: QostodianNexus,
}

View file

@ -0,0 +1,81 @@
"""
Qostodian Nexus (by Qohash) — LiteLLM guardrail integration.
"""
import os
from typing import TYPE_CHECKING, Literal, Optional, Type
from litellm.integrations.custom_guardrail import log_guardrail_information
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import (
GenericGuardrailAPI,
)
from litellm.types.proxy.guardrails.guardrail_hooks.qohash import (
QostodianNexusConfigModel,
)
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
GUARDRAIL_NAME = "qostodian_nexus"
class QostodianNexus(GenericGuardrailAPI):
def __init__(
self,
api_base: Optional[str] = None,
**kwargs,
):
api_base = api_base or os.environ.get(
"QOSTODIAN_NEXUS_API_BASE", "http://nexus:8800"
)
kwargs["guardrail_name"] = kwargs.get("guardrail_name", GUARDRAIL_NAME)
# Merge built-in Qostodian Nexus identifier headers with any caller-supplied extra_headers
nexus_headers = [
"x-qostodian-nexus-identifiers-trace",
"x-qostodian-nexus-identifiers-source",
"x-qostodian-nexus-identifiers-container",
"x-qostodian-nexus-identifiers-identity",
]
existing = kwargs.get("extra_headers") or []
kwargs["extra_headers"] = nexus_headers + [
h for h in existing if h not in nexus_headers
]
super().__init__(
api_base=api_base,
**kwargs,
)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
"""
Apply Qostodian Nexus to the given inputs.
NOTE: This override is intentionally a pass-through. It must be present
directly in this class's __dict__ so that LiteLLM's unified guardrail
routing check (`"apply_guardrail" in type(callback).__dict__` in
litellm/proxy/utils.py) routes calls correctly. Do not remove.
"""
return await super().apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type=input_type,
logging_obj=logging_obj,
)
@classmethod
def get_config_model(cls) -> Optional[Type[QostodianNexusConfigModel]]:
"""
Returns the config model for Qostodian Nexus.
"""
return QostodianNexusConfigModel

View file

@ -27,6 +27,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import (
_get_batch_job_input_file_usage,
_get_file_content_as_dictionary,
_get_models_from_batch_input_file_content,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
@ -246,6 +247,17 @@ class _PROXY_BatchRateLimiter(CustomLogger):
file_content_as_dict = _get_file_content_as_dictionary(file_content.content)
# Validate every model named in the batch JSONL against the
# caller's per-key model allowlist. Without this, a caller
# could smuggle restricted/expensive models inside the file
# and the upstream provider would execute the batch under
# the proxy's shared API key.
if user_api_key_dict is not None:
await self._enforce_batch_file_model_access(
user_api_key_dict=user_api_key_dict,
file_content_as_dict=file_content_as_dict,
)
input_file_usage = _get_batch_job_input_file_usage(
file_content_dictionary=file_content_as_dict,
custom_llm_provider=custom_llm_provider,
@ -256,12 +268,69 @@ class _PROXY_BatchRateLimiter(CustomLogger):
request_count=request_count,
)
except HTTPException as e:
# Distinguish intentional 403s from `_enforce_batch_file_model_access`
# from genuine I/O failures so security-relevant rejections show up
# in the access log instead of getting buried in error noise.
if e.status_code == 403:
verbose_proxy_logger.warning(
f"Batch rejected: caller not authorized for a model named in {file_id}: {e.detail}"
)
else:
verbose_proxy_logger.error(
f"Batch input file rejected for {file_id}: status={e.status_code} detail={e.detail}"
)
raise
except Exception as e:
verbose_proxy_logger.error(
f"Error counting input file usage for {file_id}: {str(e)}"
)
raise
async def _enforce_batch_file_model_access(
self,
user_api_key_dict: UserAPIKeyAuth,
file_content_as_dict: List[dict],
) -> None:
"""Reject the batch if the caller is not authorized for every
``body.model`` named inside the JSONL.
Reuses ``can_key_call_model`` so the same allowlist semantics
(wildcards, access groups, ``all-proxy-models``, team aliases)
the proxy enforces on `/chat/completions` apply here.
"""
from litellm.proxy.auth.auth_checks import can_key_call_model
from litellm.proxy.proxy_server import llm_router
models = _get_models_from_batch_input_file_content(file_content_as_dict)
if not models:
return
llm_model_list = llm_router.model_list if llm_router is not None else None
for model in models:
try:
await can_key_call_model(
model=model,
llm_model_list=llm_model_list,
valid_token=user_api_key_dict,
llm_router=llm_router,
)
except HTTPException:
raise
except Exception as e:
# `can_key_call_model` raises ProxyException on denial;
# re-shape to a 403 so the batch endpoint returns a
# consistent rejection without leaking internal types.
raise HTTPException(
status_code=403,
detail={
"error": (
"Batch input file references a model the caller is "
f"not authorized to use: model={model}, reason={str(e)}"
)
},
)
async def _fetch_managed_file_content(
self,
file_id: str,

View file

@ -32,10 +32,25 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
if user_api_key_dict.team_id is not None:
return
# The reservation path admits at the strict-`<` boundary and
# atomically pre-fills the same counter we'd read here. Re-checking
# with `>=` would reject a request the reservation already admitted
# when the reservation fills the counter to exactly max_budget.
# Imported lazily to avoid a circular import via proxy.utils.
from litellm.proxy.spend_tracking.budget_reservation import (
get_reserved_counter_keys,
)
user_counter_key = f"spend:user:{user_id}"
if user_counter_key in get_reserved_counter_keys(
user_api_key_dict.budget_reservation
):
return
from litellm.proxy.proxy_server import get_current_spend
curr_spend = await get_current_spend(
counter_key=f"spend:user:{user_id}",
counter_key=user_counter_key,
fallback_spend=user_api_key_dict.user_spend or 0.0,
)

View file

@ -19,6 +19,10 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.spend_tracking.spend_log_error_logger import (
should_suppress_spend_log_tracebacks,
spend_log_error,
)
from litellm.proxy.utils import ProxyUpdateSpend
from litellm.types.utils import StandardLoggingPayload
from litellm.utils import get_end_user_id_for_cost_tracking
@ -30,16 +34,35 @@ class _ProxyDBLogger(CustomLogger):
kwargs, response_obj, start_time, end_time
)
async def async_post_call_failure_hook(
self,
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
request_route = user_api_key_dict.request_route
if _ProxyDBLogger._should_track_errors_in_db() is False:
return
async def async_post_call_failure_hook(
self,
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
try:
await _release_budget_reservation(
budget_reservation=user_api_key_dict.budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to release budget reservation during failure handling"
)
try:
await _invalidate_budget_reservation_counters(
budget_reservation=user_api_key_dict.budget_reservation
)
if user_api_key_dict.budget_reservation is not None:
user_api_key_dict.budget_reservation["finalized"] = True
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate budget reservation counters after failure release failed"
)
request_route = user_api_key_dict.request_route
if _ProxyDBLogger._should_track_errors_in_db() is False:
return
elif request_route is not None and not (
RouteChecks.is_llm_api_route(route=request_route)
or RouteChecks.is_info_route(route=request_route)
@ -55,12 +78,18 @@ class _ProxyDBLogger(CustomLogger):
)
_metadata["user_api_key"] = user_api_key_dict.api_key
_metadata["status"] = "failure"
_metadata["error_information"] = (
StandardLoggingPayloadSetup.get_error_information(
original_exception=original_exception,
traceback_str=traceback_str,
)
_error_information = StandardLoggingPayloadSetup.get_error_information(
original_exception=original_exception,
traceback_str=traceback_str,
)
if should_suppress_spend_log_tracebacks():
# Drop the traceback key entirely so the per-row Metadata pane in
# the UI (which renders the JSON blob verbatim) doesn't show a
# noisy ``"traceback": ""`` line. Downstream consumers all use
# ``.get("traceback")`` / truthy checks, and the TypedDict marks
# the field as optional, so omitting is type-safe.
_error_information.pop("traceback", None)
_metadata["error_information"] = _error_information
_metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(
metadata=_metadata,
@ -155,66 +184,64 @@ class _ProxyDBLogger(CustomLogger):
f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}"
)
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
litellm_params = kwargs.get("litellm_params", {}) or {}
end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
litellm_params = kwargs.get("litellm_params", {}) or {}
end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
budget_reservation = _get_budget_reservation_from_metadata(
metadata=metadata
)
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None))
end_user_max_budget = metadata.get("user_api_end_user_max_budget", None)
sl_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object", None
)
response_cost = (
sl_object.get("response_cost", None)
if sl_object is not None
else kwargs.get("response_cost", None)
)
tags: Optional[List[str]] = (
sl_object.get("request_tags", None) if sl_object is not None else None
)
if response_cost is not None:
user_api_key = metadata.get("user_api_key", None)
response_cost = (
sl_object.get("response_cost", None)
if sl_object is not None
else kwargs.get("response_cost", None)
)
tags = _get_request_tags_for_cost_tracking(
sl_object=sl_object,
metadata=metadata,
)
if response_cost is not None:
user_api_key = metadata.get("user_api_key", None)
if kwargs.get("cache_hit", False) is True:
response_cost = 0.0
verbose_proxy_logger.debug(
f"Cache Hit: response_cost {response_cost}, for user_id {user_id}"
)
verbose_proxy_logger.debug(
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}"
)
if _should_track_cost_callback(
user_api_key=user_api_key,
verbose_proxy_logger.debug(
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}"
)
if _should_track_cost_callback(
user_api_key=user_api_key,
user_id=user_id,
team_id=team_id,
end_user_id=end_user_id,
):
## UPDATE DATABASE
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
org_id=org_id,
)
# Atomically update spend counters (in-memory + Redis)
# for cross-pod budget enforcement.
await increment_spend_counters(
token=user_api_key,
team_id=team_id,
user_id=user_id,
response_cost=response_cost,
org_id=org_id,
)
end_user_id=end_user_id,
):
## UPDATE DATABASE
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key=user_api_key,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
org_id=org_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
response_cost=response_cost,
budget_reservation=budget_reservation,
request_tags=tags,
)
# update cache (fire-and-forget for backward compat:
# cached object fields, soft budget alerts, etc.)
@ -234,10 +261,15 @@ class _ProxyDBLogger(CustomLogger):
token=user_api_key,
key_alias=key_alias,
end_user_id=end_user_id,
response_cost=response_cost,
max_budget=end_user_max_budget,
)
response_cost=response_cost,
max_budget=end_user_max_budget,
)
elif budget_reservation is not None:
await _release_budget_reservation(
budget_reservation=budget_reservation
)
else:
await _release_budget_reservation(budget_reservation=budget_reservation)
# Non-model call types (health checks, afile_delete) have no model or standard_logging_object.
# Use .get() for "stream" to avoid KeyError on health checks.
if sl_object is None and not kwargs.get("model"):
@ -280,9 +312,7 @@ class _ProxyDBLogger(CustomLogger):
)
)
verbose_proxy_logger.exception(
"Error in tracking cost callback - %s", str(e)
)
spend_log_error("Error in tracking cost callback - %s", str(e), exc=e)
@staticmethod
async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict:
@ -366,7 +396,7 @@ class _ProxyDBLogger(CustomLogger):
return
def _should_track_cost_callback(
def _should_track_cost_callback(
user_api_key: Optional[str],
user_id: Optional[str],
team_id: Optional[str],
@ -387,4 +417,135 @@ def _should_track_cost_callback(
or end_user_id is not None
):
return True
return False
return False
def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]:
metadata_budget_reservation = metadata.get("user_api_key_budget_reservation")
if isinstance(metadata_budget_reservation, dict):
return metadata_budget_reservation
user_api_key_auth_obj = metadata.get("user_api_key_auth")
if user_api_key_auth_obj is None:
return None
if isinstance(user_api_key_auth_obj, dict):
budget_reservation = user_api_key_auth_obj.get("budget_reservation")
return budget_reservation if isinstance(budget_reservation, dict) else None
return getattr(user_api_key_auth_obj, "budget_reservation", None)
def _get_request_tags_for_cost_tracking(
sl_object: Optional[StandardLoggingPayload],
metadata: dict,
) -> Optional[List[str]]:
if sl_object is not None:
request_tags = sl_object.get("request_tags", None)
if isinstance(request_tags, list):
return request_tags
metadata_tags = metadata.get("tags", None)
if isinstance(metadata_tags, list):
return metadata_tags
return None
async def _update_database_and_spend_counters(
proxy_logging_obj: Any,
increment_spend_counters: Any,
user_api_key: Optional[str],
user_id: Optional[str],
end_user_id: Optional[str],
team_id: Optional[str],
org_id: Optional[str],
kwargs: dict,
completion_response: Optional[Union[litellm.ModelResponse, Any]],
start_time: Any,
end_time: Any,
response_cost: float,
budget_reservation: Optional[dict],
request_tags: Optional[List[str]] = None,
) -> None:
try:
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
org_id=org_id,
)
except Exception:
if budget_reservation is not None:
try:
await _release_budget_reservation(budget_reservation=budget_reservation)
except Exception:
verbose_proxy_logger.exception(
"Failed to release budget reservation after database update failed"
)
try:
await _invalidate_budget_reservation_counters(
budget_reservation=budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate budget reservation counters after release failed"
)
raise
try:
await increment_spend_counters(
token=user_api_key,
team_id=team_id,
user_id=user_id,
response_cost=response_cost,
org_id=org_id,
budget_reservation=budget_reservation,
end_user_id=end_user_id,
tags=request_tags,
)
except Exception:
if budget_reservation is not None:
try:
await _invalidate_budget_reservation_counters(
budget_reservation=budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate budget reservation counters after spend counter update failed"
)
finally:
budget_reservation["finalized"] = True
raise
async def _release_budget_reservation(budget_reservation: Optional[dict]) -> None:
if budget_reservation is None:
return
from litellm.proxy.spend_tracking.budget_reservation import (
release_budget_reservation,
)
await release_budget_reservation(
budget_reservation=budget_reservation,
)
async def _invalidate_budget_reservation_counters(
budget_reservation: Optional[dict],
) -> None:
if budget_reservation is None:
return
from litellm.proxy.spend_tracking.budget_reservation import (
invalidate_budget_reservation_counters,
)
await invalidate_budget_reservation_counters(
budget_reservation=budget_reservation,
)

View file

@ -893,6 +893,10 @@ class LiteLLMProxyRequestSetup:
data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr(
user_api_key_dict, "end_user_max_budget", None
)
if user_api_key_dict.budget_reservation is not None:
data[_metadata_variable_name][
"user_api_key_budget_reservation"
] = user_api_key_dict.budget_reservation
# Add the full UserAPIKeyAuth object for MCP server access control
data[_metadata_variable_name]["user_api_key_auth"] = user_api_key_dict
return data

View file

@ -38,6 +38,17 @@ def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
)
def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None:
"""Admin Viewer parity: PROXY_ADMIN or PROXY_ADMIN_VIEW_ONLY may read."""
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
if not _user_has_admin_view(user_api_key_dict):
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,
@ -370,7 +381,7 @@ async def create_access_group(
async def list_access_groups(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> List[AccessGroupResponse]:
_require_proxy_admin(user_api_key_dict)
_require_admin_view(user_api_key_dict)
prisma_client = get_prisma_client_or_throw(
CommonProxyErrors.db_not_connected_error.value
)
@ -389,7 +400,7 @@ 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)
_require_admin_view(user_api_key_dict)
prisma_client = get_prisma_client_or_throw(
CommonProxyErrors.db_not_connected_error.value
)

View file

@ -17,6 +17,7 @@ from fastapi import APIRouter, Depends, HTTPException
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.utils import jsonify_object
router = APIRouter()
@ -238,7 +239,7 @@ async def budget_settings(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=400,
detail={
@ -305,7 +306,7 @@ async def list_budget(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=400,
detail={

View file

@ -1,5 +1,7 @@
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from fastapi import HTTPException, status
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.proxy._types import (
@ -29,6 +31,34 @@ def _user_has_admin_view(user_api_key_dict: UserAPIKeyAuth) -> bool:
)
def require_caller_user_id_for_non_admin(
user_api_key_dict: UserAPIKeyAuth,
) -> str:
"""Return the caller's user_id, or raise 403 if missing.
Non-admin analytics endpoints scope queries by the caller's own user_id.
Service-account keys are deliberately created with user_id=None
(key_management_endpoints.py forces ``data.user_id = None`` at key
creation). Without this guard, that None value flows through to the
daily-activity builder, which treats ``entity_id is None`` as "no filter"
and returns every tenant's data.
Callers must check is_admin first; this helper is only valid on the
non-admin scoping branch.
"""
if user_api_key_dict.user_id is None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"error": (
"Service-account keys cannot query user analytics. "
"Use a user-bound key, or call as a proxy admin."
)
},
)
return user_api_key_dict.user_id
def _is_user_team_admin(
user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable
) -> bool:

View file

@ -267,9 +267,11 @@ async def get_hashicorp_vault_config(
Get current Hashicorp Vault configuration.
Returns decrypted values from DB, or falls back to current env vars.
"""
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.proxy_server import prisma_client, proxy_config
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Admin Viewer follows the read-parity rule.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail="Only admin users can view config overrides",

View file

@ -35,6 +35,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
_user_has_admin_view,
require_caller_user_id_for_non_admin,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_helper_fn,
@ -618,6 +619,40 @@ def _normalize_user_info_user_id(
return user_id
def _enforce_user_info_access(
user_id: Optional[str], user_api_key_dict: UserAPIKeyAuth
) -> None:
"""Re-validate that the caller may read the resolved ``user_id`` after
URL-decoding has been finalized.
The route-level check in ``RouteChecks.non_proxy_admin_allowed_routes_check``
runs against ``request.query_params``, which decodes a literal ``+`` to a
space. ``_normalize_user_info_user_id`` then re-parses the raw query with
``unquote`` so the endpoint can return rows for user_ids that contain ``+``
(e.g. plus-addressed emails). That asymmetry let an attacker who registered
a username with a literal space pass the route check and then read another
user's row by sending the encoded ``+`` form. Re-checking ownership here
closes the gap without changing the supported user_id grammar.
"""
if user_id is None:
return
# Only true proxy admin bypasses ownership. PROXY_ADMIN_VIEW_ONLY is
# subject to the same `user_id == valid_token.user_id` rule that
# `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream
# for the `/user/info` route.
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return
if user_id == user_api_key_dict.user_id:
return
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=(
f"key not allowed to access this user's info. user_id={user_id}, "
f"key's user_id={user_api_key_dict.user_id}"
),
)
async def _get_user_info_teams(
prisma_client: Any,
user_id: Optional[str],
@ -732,6 +767,7 @@ async def user_info( # noqa: PLR0915
try:
user_id = _normalize_user_info_user_id(request=request, user_id=user_id)
_enforce_user_info_access(user_id=user_id, user_api_key_dict=user_api_key_dict)
if prisma_client is None:
raise Exception(
@ -2587,9 +2623,10 @@ async def get_user_daily_activity(
if is_admin:
entity_id = user_id # None means global view, otherwise filter by user
else:
caller_user_id = require_caller_user_id_for_non_admin(user_api_key_dict)
if user_id is None:
user_id = user_api_key_dict.user_id
if user_id != user_api_key_dict.user_id:
user_id = caller_user_id
if user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
@ -2684,9 +2721,10 @@ async def get_user_daily_activity_aggregated(
if is_admin:
entity_id = user_id # None means global view, otherwise filter by user
else:
caller_user_id = require_caller_user_id_for_non_admin(user_api_key_dict)
if user_id is None:
user_id = user_api_key_dict.user_id
if user_id != user_api_key_dict.user_id:
user_id = caller_user_id
if user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={

View file

@ -10,6 +10,7 @@ from litellm.proxy._types import (
hash_token,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
router = APIRouter()
@ -194,7 +195,8 @@ async def list_jwt_key_mappings(
):
from litellm.proxy.proxy_server import prisma_client
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Admin Viewer follows the read-parity rule.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403, detail="Only proxy admins can list JWT key mappings"
)
@ -233,7 +235,8 @@ async def info_jwt_key_mapping(
):
from litellm.proxy.proxy_server import prisma_client
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Admin Viewer follows the read-parity rule.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403, detail="Only proxy admins can get JWT key mapping info"
)

View file

@ -57,6 +57,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_is_user_org_admin_for_team,
_is_user_team_admin,
_set_object_metadata_field,
_team_member_has_permission,
)
from litellm.proxy.management_endpoints.model_management_endpoints import (
_add_model_to_db,
@ -809,6 +810,20 @@ async def _common_key_generation_helper( # noqa: PLR0915
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
if prisma_client:
# Mirror the membership rule applied to /key/update: when the
# caller specifies an organization_id, require that they are a
# member of (or proxy admin over) the target organization.
_is_proxy_admin = (
user_api_key_dict.user_role is not None
and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if not _is_proxy_admin:
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
prisma_client=prisma_client,
)
org_table = await get_org_object(
org_id=data.organization_id,
user_api_key_cache=user_api_key_cache,
@ -1168,6 +1183,42 @@ def check_org_key_rpm_tpm_limits(
)
async def _validate_caller_can_assign_key_org(
user_api_key_dict: UserAPIKeyAuth,
organization_id: str,
prisma_client: PrismaClient,
) -> None:
"""Reject ``/key/update`` requests that point a key at an organization
the caller does not belong to.
Mirrors the org-membership rule already enforced on ``/key/list`` in
``validate_key_list_check``. Proxy admins are checked at the call site.
"""
if user_api_key_dict.user_id is None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Cannot assign a key to an organization without a user_id on the caller's token",
)
user_row = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_api_key_dict.user_id},
include={"organization_memberships": True},
)
memberships = (
getattr(user_row, "organization_memberships", None) if user_row else None
)
member_org_ids = {
membership.organization_id
for membership in (memberships or [])
if membership.organization_id is not None
}
if organization_id not in member_org_ids:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"Caller is not a member of organization_id={organization_id}",
)
async def _check_org_key_limits(
org_table: LiteLLM_OrganizationTable,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
@ -2168,10 +2219,26 @@ async def _validate_update_key_data(
user_api_key_cache=user_api_key_cache,
)
# When the caller asks to change the key's organization_id, require that
# they are a member of (or a proxy admin over) the target organization.
# Without this gate, any caller could assign their key to an arbitrary
# organization_id by passing it in the request body — VERIA-55 secondary
# IDOR. The check mirrors the membership rule already used on the
# `/key/list` filter path in `validate_key_list_check`.
_existing_org_id = getattr(existing_key_row, "organization_id", None)
if (
data.organization_id is not None
and data.organization_id != _existing_org_id
and not _is_proxy_admin
):
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
prisma_client=prisma_client,
)
# Check org key limits only when throughput-related fields or organization_id change
_org_id_to_check = data.organization_id or getattr(
existing_key_row, "organization_id", None
)
_org_id_to_check = data.organization_id or _existing_org_id
_throughput_fields_changed = (
data.organization_id is not None
or data.tpm_limit is not None
@ -3868,6 +3935,22 @@ async def _execute_virtual_key_regeneration(
"""Generate new token, update DB, invalidate cache, and return response."""
from litellm.proxy.proxy_server import hash_token
# Apply the same membership rule used on /key/update: when the caller
# asks to point the regenerated key at a different organization_id,
# require they are a member of (or proxy admin over) the target org.
if data is not None and data.organization_id is not None:
_existing_org_id = getattr(key_in_db, "organization_id", None)
_is_proxy_admin = (
user_api_key_dict.user_role is not None
and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if data.organization_id != _existing_org_id and not _is_proxy_admin:
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
prisma_client=prisma_client,
)
new_token = await get_new_token(data=data)
new_token_hash = hash_token(new_token)
new_token_key_name = f"sk-...{new_token[-4:]}"
@ -4436,6 +4519,26 @@ def _get_admin_team_ids_from_objects(
]
def _get_team_ids_with_key_list_permission_from_objects(
user_api_key_dict: UserAPIKeyAuth,
team_objects: List[LiteLLM_TeamTable],
) -> List[str]:
"""Filter team objects to non-admin teams where the caller has /key/list
permission via team_member_permissions. These teams should grant the
caller full key visibility (same as a team admin), so other members'
keys and service account keys (user_id=NULL) are returned."""
return [
team.team_id
for team in team_objects
if not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team)
and _team_member_has_permission(
user_api_key_dict=user_api_key_dict,
team_obj=team,
permission=KeyManagementRoutes.KEY_LIST.value,
)
]
def _get_member_team_ids_from_objects(
user_api_key_dict: UserAPIKeyAuth,
team_objects: List[LiteLLM_TeamTable],
@ -4589,6 +4692,17 @@ async def list_keys(
user_api_key_dict=user_api_key_dict,
team_objects=team_objects,
)
# Non-admin members with /key/list permission get full team-key
# visibility for that team — matching the UI contract that
# granting this permission lets them see all keys within the team.
list_permission_team_ids = (
_get_team_ids_with_key_list_permission_from_objects(
user_api_key_dict=user_api_key_dict,
team_objects=team_objects,
)
)
if list_permission_team_ids:
admin_team_ids = list({*admin_team_ids, *list_permission_team_ids})
else:
admin_team_ids = None

View file

@ -2120,7 +2120,8 @@ if MCP_AVAILABLE:
Used by the UI to show a discovery grid when adding new MCP servers.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Admin Viewer follows the read-parity rule.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail={
@ -2177,7 +2178,8 @@ if MCP_AVAILABLE:
async def get_openapi_registry(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Admin Viewer follows the read-parity rule.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail={

View file

@ -4683,9 +4683,11 @@ async def team_member_permissions(
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
# Admin Viewer follows the read-parity rule: see team permissions like
# a Proxy Admin would. Team / org admins keep their existing scope.
if (
hasattr(user_api_key_dict, "user_role")
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
and not _user_has_admin_view(user_api_key_dict)
and not _is_user_team_admin(
user_api_key_dict=user_api_key_dict, team_obj=complete_team_data
)

View file

@ -678,6 +678,7 @@ async def google_login(
google_client_id=google_client_id,
generic_client_id=generic_client_id,
state=cli_state,
request=request,
)
if return_to is not None and sso_redirect is not None:
if SSOAuthenticationHandler._validate_return_to(return_to):
@ -1159,6 +1160,30 @@ async def get_generic_sso_response(
authorization_code = request.query_params.get("code")
if code_verifier:
# State-to-session-cookie binding. The non-PKCE branch below
# delegates to fastapi-sso's ``verify_and_process``, which
# performs its own session-cookie check. The PKCE branch
# bypasses that helper, so we validate the URL ``state``
# against the ``litellm_oauth_state`` cookie set on the
# redirect response — without this an attacker can pre-mint
# a state + cached PKCE verifier and hijack a victim's auth
# code (Login-CSRF / token theft).
url_state = request.query_params.get("state")
cookie_state = request.cookies.get("litellm_oauth_state")
if (
not url_state
or not cookie_state
or not secrets.compare_digest(url_state, cookie_state)
):
raise ProxyException(
message=(
"Invalid OAuth state parameter — does not match "
"the browser-bound state cookie."
),
type=ProxyErrorTypes.auth_error,
param="state",
code=status.HTTP_400_BAD_REQUEST,
)
if not authorization_code:
raise ProxyException(
message="Missing authorization code in callback",
@ -2147,6 +2172,7 @@ class SSOAuthenticationHandler:
microsoft_client_id: Optional[str] = None,
generic_client_id: Optional[str] = None,
state: Optional[str] = None,
request: Optional[Request] = None,
) -> Optional[RedirectResponse]:
"""
Step 1. Call Get Login Redirect for the SSO provider. Send the redirect response to `redirect_url`
@ -2156,6 +2182,8 @@ class SSOAuthenticationHandler:
google_client_id (Optional[str], optional): The Google Client ID. Defaults to None.
microsoft_client_id (Optional[str], optional): The Microsoft Client ID. Defaults to None.
generic_client_id (Optional[str], optional): The Generic Client ID. Defaults to None.
request: Optional FastAPI request, used to drive the ``Secure``
attribute on the ``litellm_oauth_state`` CSRF cookie.
Returns:
RedirectResponse: The redirect response from the SSO provider.
@ -2266,6 +2294,7 @@ class SSOAuthenticationHandler:
generic_sso=generic_sso,
state=state,
generic_authorization_endpoint=generic_authorization_endpoint,
request=request,
)
raise ValueError(
"Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso"
@ -2276,6 +2305,7 @@ class SSOAuthenticationHandler:
generic_sso: Any,
state: Optional[str] = None,
generic_authorization_endpoint: Optional[str] = None,
request: Optional[Request] = None,
) -> Optional[RedirectResponse]:
"""
Get the redirect response for Generic SSO
@ -2285,10 +2315,13 @@ class SSOAuthenticationHandler:
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
with generic_sso:
# TODO: state should be a random string and added to the user session with cookie
# or a cryptographicly signed state that we can verify stateless
# For simplification we are using a static state, this is not perfect but some
# SSO providers do not allow stateless verification
# State is bound to the caller's browser via a ``litellm_oauth_state``
# HttpOnly cookie set on the redirect response below; the SSO
# callback validates the URL ``state`` against that cookie before
# completing the PKCE token exchange. Without this binding, an
# attacker who pre-mints a state + a cached PKCE verifier can hand
# the link to a victim and capture the resulting access token
# (Login CSRF / token theft).
(
redirect_params,
code_verifier,
@ -2355,6 +2388,31 @@ class SSOAuthenticationHandler:
# Update the redirect response
redirect_response.headers["location"] = new_url
# Bind state to the user's browser session. The /callback
# handler validates the URL ``state`` against this cookie via
# ``secrets.compare_digest`` before exchanging the PKCE
# code_verifier. Only set the cookie when PKCE is in use
# (i.e. inside this ``code_verifier`` branch) so two
# concurrent SSO sessions — one PKCE, one plain — cannot
# overwrite each other's state cookie.
state_value = redirect_params.get("state")
if state_value and redirect_response is not None:
# Production-safe default: require HTTPS for the
# CSRF-protection cookie unless we can prove the
# incoming request is HTTP (local dev). Without
# ``Secure`` the cookie is sent over plain HTTP,
# letting a network observer read and replay the
# state value and bypass this protection.
secure_flag = request is None or request.url.scheme == "https"
redirect_response.set_cookie(
key="litellm_oauth_state",
value=state_value,
max_age=600,
httponly=True,
samesite="lax",
secure=secure_flag,
)
return redirect_response
@staticmethod
@ -3972,6 +4030,7 @@ async def debug_sso_login(request: Request):
microsoft_client_id=microsoft_client_id,
google_client_id=google_client_id,
generic_client_id=generic_client_id,
request=request,
)

View file

@ -440,6 +440,15 @@ def _resolve_fetch_kwargs(
kwargs: Dict[str, Any] = {"start_date": start_date, "end_date": end_date}
if fn_name == "get_usage_data":
if not is_admin:
if user_id is None:
# Defense-in-depth: the endpoint guard in usage_endpoints/endpoints.py
# should have already rejected this. If we ever reach here it means
# a future caller invoked the helper without scoping — fail loudly
# rather than issuing an unfiltered global query.
raise ValueError(
"Non-admin caller has user_id=None; refusing to issue an "
"unscoped query. Endpoint-level guard missing."
)
kwargs["user_id"] = user_id
elif fn_args.get("user_id"):
kwargs["user_id"] = fn_args["user_id"]

View file

@ -44,13 +44,17 @@ async def usage_ai_chat(
"""
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_view,
require_caller_user_id_for_non_admin,
)
from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import (
stream_usage_ai_chat,
)
is_admin = _user_has_admin_view(user_api_key_dict)
user_id = user_api_key_dict.user_id
if is_admin:
user_id = user_api_key_dict.user_id
else:
user_id = require_caller_user_id_for_non_admin(user_api_key_dict)
messages = [{"role": m.role, "content": m.content} for m in data.messages]
return StreamingResponse(

View file

@ -47,6 +47,8 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
)
from litellm.proxy.utils import is_known_model
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store,
get_litellm_managed_vector_store,
is_allowed_to_call_vector_store_endpoint,
)
from litellm.secret_managers.main import get_secret_str
@ -533,6 +535,10 @@ async def milvus_proxy_route(
)
if vector_store is None:
raise Exception(f"Vector store not found for {vector_store_name}")
await assert_user_can_access_vector_store(
vector_store=vector_store,
user_api_key_dict=user_api_key_dict,
)
litellm_params = vector_store.get("litellm_params") or {}
auth_credentials = provider_config.get_auth_credentials(
litellm_params=litellm_params
@ -1438,6 +1444,10 @@ async def azure_proxy_route(
)
if vector_store is None:
raise Exception(f"Vector store not found for {vector_store_name}")
await assert_user_can_access_vector_store(
vector_store=vector_store,
user_api_key_dict=user_api_key_dict,
)
litellm_params = vector_store.get("litellm_params") or {}
auth_credentials = provider_config.get_auth_credentials(
litellm_params=litellm_params
@ -1777,6 +1787,11 @@ async def _base_vertex_proxy_route(
request=request,
api_key=api_key_to_use,
)
if router_credentials is not None:
await assert_user_can_access_vector_store(
vector_store=router_credentials,
user_api_key_dict=user_api_key_dict,
)
vertex_project: Optional[str] = get_vertex_project_id_from_url(endpoint)
vertex_location: Optional[str] = get_vertex_location_from_url(endpoint)
@ -1913,11 +1928,11 @@ async def vertex_discovery_proxy_route(
"Extracted vector store ID from endpoint: %s", vector_store_id
)
# Retrieve vector store credentials from the registry
vector_store_credentials = (
passthrough_endpoint_router.get_vector_store_credentials(
vector_store_id=vector_store_id
)
# Retrieve LiteLLM-managed vector store credentials if the datastore id
# is registered with LiteLLM. Unknown datastore ids keep the existing
# direct Vertex pass-through behavior.
vector_store_credentials = await get_litellm_managed_vector_store(
vector_store_id=vector_store_id
)
if vector_store_credentials:
@ -1925,7 +1940,7 @@ async def vertex_discovery_proxy_route(
"Found vector store credentials for ID: %s", vector_store_id
)
else:
verbose_proxy_logger.warning(
verbose_proxy_logger.debug(
"Vector store ID %s found in endpoint but no credentials found in registry",
vector_store_id,
)

View file

@ -2324,14 +2324,10 @@ async def _register_pass_through_endpoint(
dependencies = None
if auth is not None and str(auth).lower() == "true":
# Authentication on a pass-through endpoint used to be enterprise-
# only — which left the OSS tier with no safe configuration: the
# default was ``auth=False`` (unauthenticated forwarder) and the
# safe ``auth=True`` raised at startup unless the operator had a
# license. The default is now ``True`` (safe-by-default), and
# turning it on no longer requires a license: an unauthenticated
# forwarder is a deployment choice the operator should be allowed
# to make explicitly, but the safe option must always be free.
# Authentication on a pass-through endpoint used to be enterprise-only.
# That left OSS with no safe configuration: auth=True raised at startup
# unless the operator had a license. The safe option must always be free,
# and unauthenticated forwarding should require explicit opt-in.
dependencies = [Depends(user_api_key_auth)]
if path not in LiteLLMRoutes.openai_routes.value:
LiteLLMRoutes.openai_routes.value.append(path)

View file

@ -220,6 +220,7 @@ class AttachmentRegistry:
attachment: PolicyAttachment object to add
"""
self._attachments.append(attachment)
self._initialized = True
verbose_proxy_logger.debug(f"Added attachment for policy: {attachment.policy}")
def remove_attachments_for_policy(self, policy_name: str) -> int:

View file

@ -226,6 +226,7 @@ class PolicyRegistry:
policy: Policy object to add
"""
self._policies[policy_name] = policy
self._initialized = True
verbose_proxy_logger.debug(f"Added/updated policy: {policy_name}")
def remove_policy(self, policy_name: str) -> bool:

View file

@ -6,6 +6,7 @@ import inspect
import io
import os
import random
import re
import secrets
import shutil
import subprocess
@ -334,6 +335,7 @@ from litellm.proxy.management_endpoints.callback_management_endpoints import (
)
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_privileges,
_user_has_admin_view,
admin_can_invite_user,
)
from litellm.proxy.management_endpoints.cost_tracking_settings import (
@ -955,6 +957,85 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues]
def _generate_stable_operation_id(route: Any) -> str:
operation_id = re.sub(r"\W", "_", f"{route.name}{route.path_format}")
route_methods = sorted(route.methods or [])
if len(route_methods) == 1:
operation_id = f"{operation_id}_{route_methods[0].lower()}"
return operation_id
_OPENAPI_HTTP_METHODS = {
"delete",
"get",
"head",
"options",
"patch",
"post",
"put",
"trace",
}
def _strip_operation_id_method_suffix(operation_id: str) -> str:
base, separator, suffix = operation_id.rpartition("_")
if separator and suffix in _OPENAPI_HTTP_METHODS:
return base
return operation_id
def ensure_unique_openapi_operation_ids(
openapi_schema: Dict[str, Any],
reserved_operation_ids: Optional[Set[str]] = None,
) -> Dict[str, Any]:
operation_entries = []
operation_id_counts: Dict[str, int] = {}
for path_item in openapi_schema.get("paths", {}).values():
if not isinstance(path_item, dict):
continue
for method, operation in path_item.items():
if method not in _OPENAPI_HTTP_METHODS or not isinstance(operation, dict):
continue
operation_id = operation.get("operationId")
if not isinstance(operation_id, str):
continue
operation_entries.append((method, operation, operation_id))
operation_id_counts[operation_id] = (
operation_id_counts.get(operation_id, 0) + 1
)
used_operation_ids = set(reserved_operation_ids or set())
seen_operation_ids: Set[str] = set()
for method, operation, operation_id in operation_entries:
should_rewrite = (
operation_id_counts[operation_id] > 1
or operation_id in used_operation_ids
or operation_id in seen_operation_ids
)
if not should_rewrite:
seen_operation_ids.add(operation_id)
used_operation_ids.add(operation_id)
continue
base_operation_id = _strip_operation_id_method_suffix(operation_id)
new_operation_id = f"{base_operation_id}_{method}"
suffix = 2
while (
new_operation_id in used_operation_ids
or new_operation_id in seen_operation_ids
):
new_operation_id = f"{base_operation_id}_{method}_{suffix}"
suffix += 1
operation["operationId"] = new_operation_id
seen_operation_ids.add(new_operation_id)
used_operation_ids.add(new_operation_id)
if reserved_operation_ids is not None:
reserved_operation_ids.update(used_operation_ids)
return openapi_schema
app = FastAPI(
docs_url=_get_docs_url(),
redoc_url=_get_redoc_url(),
@ -964,6 +1045,7 @@ app = FastAPI(
version=version,
root_path=server_root_path,
lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues]
generate_unique_id_function=_generate_stable_operation_id,
)
vertex_live_passthrough_vertex_base = VertexBase()
@ -1043,6 +1125,7 @@ def get_openapi_schema():
from litellm.proxy._lazy_features import inject_lazy_stubs
openapi_schema = inject_lazy_stubs(openapi_schema)
openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema)
# Fix Swagger UI execute path error when server_root_path is set
if server_root_path:
@ -1074,6 +1157,7 @@ def custom_openapi():
from litellm.proxy._lazy_features import inject_lazy_stubs
openapi_schema = inject_lazy_stubs(openapi_schema)
openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema)
# Fix Swagger UI execute path error when server_root_path is set
if server_root_path:
@ -1845,6 +1929,9 @@ async def increment_spend_counters(
user_id: Optional[str],
response_cost: Optional[float],
org_id: Optional[str] = None,
budget_reservation: Optional[dict] = None,
end_user_id: Optional[str] = None,
tags: Optional[List[str]] = None,
):
"""
Atomically increment spend counters for budget enforcement.
@ -1856,7 +1943,14 @@ async def increment_spend_counters(
Awaited (not create_task) in the cost callback, so the counter is
updated before the next request's auth check runs.
"""
reserved_counter_keys = await _reconcile_budget_reservation_for_counter_update(
budget_reservation=budget_reservation,
response_cost=response_cost,
)
if response_cost is None or response_cost == 0:
if budget_reservation is not None:
budget_reservation["finalized"] = True
return
if token is not None:
@ -1871,11 +1965,13 @@ async def increment_spend_counters(
if isinstance(token, str) and token.startswith("sk-")
else token
)
await _init_and_increment_spend_counter(
counter_key=f"spend:key:{hashed_token}",
source_cache_key=hashed_token,
increment=response_cost,
)
key_counter_key = f"spend:key:{hashed_token}"
if key_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=key_counter_key,
source_cache_key=hashed_token,
increment=response_cost,
)
# Increment per-window budget counters for multi-budget keys
key_obj = await user_api_key_cache.async_get_cache(key=hashed_token)
@ -1892,17 +1988,28 @@ async def increment_spend_counters(
if isinstance(window, dict)
else window.budget_duration
)
await spend_counter_cache.async_increment_cache(
key=f"spend:key:{hashed_token}:window:{duration}",
value=response_cost,
)
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
if key_window_counter not in reserved_counter_keys:
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
)
await _init_and_increment_window_spend_counter(
counter_key=key_window_counter,
entity_type="Key",
entity_id=hashed_token,
window_start=get_budget_window_start(window),
increment=response_cost,
)
if team_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team:{team_id}",
source_cache_key=f"team_id:{team_id}",
increment=response_cost,
)
team_counter_key = f"spend:team:{team_id}"
if team_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=team_counter_key,
source_cache_key=f"team_id:{team_id}",
increment=response_cost,
)
# Increment per-window budget counters for multi-budget teams
team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{team_id}")
@ -1919,36 +2026,157 @@ async def increment_spend_counters(
if isinstance(window, dict)
else window.budget_duration
)
await spend_counter_cache.async_increment_cache(
key=f"spend:team:{team_id}:window:{duration}",
value=response_cost,
)
team_window_counter = f"spend:team:{team_id}:window:{duration}"
if team_window_counter not in reserved_counter_keys:
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
)
await _init_and_increment_window_spend_counter(
counter_key=team_window_counter,
entity_type="Team",
entity_id=team_id,
window_start=get_budget_window_start(window),
increment=response_cost,
)
if user_id is not None and team_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team_member:{user_id}:{team_id}",
source_cache_key=f"team_membership:{user_id}:{team_id}",
increment=response_cost,
)
team_member_counter_key = f"spend:team_member:{user_id}:{team_id}"
if team_member_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=team_member_counter_key,
source_cache_key=f"team_membership:{user_id}:{team_id}",
increment=response_cost,
)
if user_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:user:{user_id}",
source_cache_key=user_id,
user_counter_key = f"spend:user:{user_id}"
if user_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=user_counter_key,
source_cache_key=user_id,
increment=response_cost,
)
await _increment_end_user_and_tag_spend_counters(
end_user_id=end_user_id,
tags=tags,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
await _increment_org_spend_counter(
org_id=org_id,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
if budget_reservation is not None:
budget_reservation["finalized"] = True
async def _reconcile_budget_reservation_for_counter_update(
budget_reservation: Optional[dict],
response_cost: Optional[float],
) -> Set[str]:
if budget_reservation is None:
return set()
from litellm.proxy.spend_tracking.budget_reservation import (
get_reserved_counter_keys,
invalidate_budget_reservation_counters,
reconcile_budget_reservation,
)
reserved_counter_keys = get_reserved_counter_keys(
budget_reservation=budget_reservation
)
try:
await reconcile_budget_reservation(
budget_reservation=budget_reservation,
actual_cost=response_cost or 0.0,
finalize=False,
)
except Exception:
verbose_proxy_logger.warning(
"Failed to reconcile budget reservation after persisted spend; invalidating reserved counters and continuing",
exc_info=True,
)
try:
await invalidate_budget_reservation_counters(
budget_reservation=budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate reserved counters after reservation reconciliation failed"
)
return reserved_counter_keys
async def _increment_end_user_and_tag_spend_counters(
end_user_id: Optional[str],
tags: Optional[List[str]],
response_cost: float,
reserved_counter_keys: Set[str],
) -> None:
if end_user_id is not None:
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:end_user:{end_user_id}",
source_cache_key=f"end_user_id:{end_user_id}",
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
if org_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:org:{org_id}",
source_cache_key=f"org_id:{org_id}",
if tags is None:
return
seen_tags: Set[str] = set()
for tag_name in tags:
if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags:
continue
seen_tags.add(tag_name)
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:tag:{tag_name}",
source_cache_key=f"tag:{tag_name}",
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
async def _increment_org_spend_counter(
org_id: Optional[str],
response_cost: float,
reserved_counter_keys: Set[str],
) -> None:
if org_id is None:
return
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:org:{org_id}",
source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"],
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
async def _init_and_increment_unreserved_spend_counter(
counter_key: str,
source_cache_key: Union[str, List[str]],
increment: float,
reserved_counter_keys: Set[str],
) -> None:
if counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=counter_key,
source_cache_key=source_cache_key,
increment=increment,
)
async def _init_and_increment_spend_counter(
counter_key: str,
source_cache_key: str,
source_cache_key: Union[str, List[str]],
increment: float,
):
"""
@ -1967,31 +2195,163 @@ async def _init_and_increment_spend_counter(
under-counting (would allow overspend).
4. Increment atomically (both in-memory + Redis)
"""
current = await spend_counter_cache.async_get_cache(key=counter_key)
if current is None:
await _ensure_spend_counter_initialized(
counter_key=counter_key,
source_cache_key=source_cache_key,
)
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
async def _init_and_increment_window_spend_counter(
counter_key: str,
entity_type: str,
entity_id: str,
window_start: Optional[datetime],
increment: float,
):
if window_start is None:
verbose_proxy_logger.warning(
"Skipping spend counter increment for invalid budget window %s",
counter_key,
)
return
initialized = await _ensure_window_spend_counter_initialized(
counter_key=counter_key,
entity_type=entity_type,
entity_id=entity_id,
window_start=window_start,
)
if initialized is False:
return
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
async def _ensure_spend_counter_initialized(
counter_key: str,
source_cache_key: Union[str, List[str]],
):
is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key)
if is_warm is False:
# Shares the per-counter lock with get_current_spend.
db_spend = await SpendCounterReseed.coalesced(
prisma_client=prisma_client,
spend_counter_cache=spend_counter_cache,
counter_key=counter_key,
require_cache_warm=True,
)
if db_spend is None:
# DB unavailable - fall back to in-process cache (may be stale).
source = await user_api_key_cache.async_get_cache(key=source_cache_key)
base_spend: float = 0.0
if source is not None:
if isinstance(source, dict):
base_spend = source.get("spend", 0.0) or 0.0
else:
base_spend = getattr(source, "spend", 0.0) or 0.0
base_spend = await _get_source_cache_base_spend(
source_cache_key=source_cache_key
)
if base_spend > 0:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=base_spend, refresh_ttl=True
await _increment_spend_counter_cache(
counter_key=counter_key, increment=base_spend
)
await spend_counter_cache.async_increment_cache(
key=counter_key, value=increment, refresh_ttl=True
async def _get_source_cache_base_spend(
source_cache_key: Union[str, List[str]],
) -> float:
source_cache_keys = (
[source_cache_key] if isinstance(source_cache_key, str) else source_cache_key
)
for cache_key in source_cache_keys:
source = await user_api_key_cache.async_get_cache(key=cache_key)
if source is None:
continue
if isinstance(source, dict):
return float(source.get("spend", 0.0) or 0.0)
return float(getattr(source, "spend", 0.0) or 0.0)
return 0.0
async def _ensure_window_spend_counter_initialized(
counter_key: str,
entity_type: str,
entity_id: str,
window_start: datetime,
) -> bool:
is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key)
if is_warm is True:
return True
window_spend = await SpendCounterReseed.coalesced_window(
prisma_client=prisma_client,
spend_counter_cache=spend_counter_cache,
counter_key=counter_key,
entity_type=entity_type,
entity_id=entity_id,
window_start=window_start,
)
if window_spend is None:
verbose_proxy_logger.warning(
"Skipping cold spend counter seed for %s because window spend could not be loaded",
counter_key,
)
return False
return True
async def _is_spend_counter_cache_warm(counter_key: str) -> bool:
if spend_counter_cache.redis_cache is not None:
try:
current_value = await spend_counter_cache.redis_cache.async_get_cache(
key=counter_key,
)
if current_value is None:
return False
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
return True
except Exception as e:
verbose_proxy_logger.debug(
"Unable to read Redis spend counter %s before initialization, falling back to in-memory: %s",
counter_key,
e,
)
return spend_counter_cache.in_memory_cache.get_cache(key=counter_key) is not None
async def _increment_spend_counter_cache(counter_key: str, increment: float):
if spend_counter_cache.redis_cache is not None:
try:
current_value = await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=increment,
refresh_ttl=True,
)
except Exception:
await _invalidate_spend_counter(counter_key=counter_key)
raise
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
return current_value
return await spend_counter_cache.async_increment_cache(
key=counter_key,
value=increment,
refresh_ttl=True,
)
async def _invalidate_spend_counter(counter_key: str):
spend_counter_cache.in_memory_cache.delete_cache(key=counter_key)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_delete_cache(key=counter_key)
except Exception:
verbose_proxy_logger.debug(
"Unable to delete stale spend counter %s after increment failure",
counter_key,
exc_info=True,
)
async def update_cache( # noqa: PLR0915
@ -5889,10 +6249,15 @@ async def initialize( # noqa: PLR0915
if litellm_log_setting.upper() == "INFO":
import logging
from litellm._logging import verbose_proxy_logger, verbose_router_logger
from litellm._logging import (
verbose_logger,
verbose_proxy_logger,
verbose_router_logger,
)
# this must ALWAYS remain logging.INFO, DO NOT MODIFY THIS
verbose_logger.setLevel(level=logging.INFO) # set package log to info
verbose_router_logger.setLevel(
level=logging.INFO
) # set router logs to info
@ -11603,7 +11968,7 @@ async def alerting_settings(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=400,
detail={
@ -12715,7 +13080,7 @@ async def invitation_info(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=400,
detail={
@ -13137,7 +13502,7 @@ async def get_config_general_settings(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=400,
detail={"error": CommonProxyErrors.not_allowed_access.value},
@ -13201,7 +13566,7 @@ async def get_config_list(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=400,
detail={
@ -13913,8 +14278,8 @@ async def get_model_cost_map_reload_status(
Get the status of the scheduled model cost map reload job.
"""
# Check if user is admin
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Read-only status check — admin viewers can read.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}",
@ -14016,7 +14381,8 @@ async def get_model_cost_map_source(
- fallback_reason: human-readable reason why remote failed (null on success)
- model_count: number of models in the currently loaded cost map
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Read-only source info — admin viewers can read.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}",
@ -14273,8 +14639,8 @@ async def get_anthropic_beta_headers_reload_status(
Get the status of the scheduled Anthropic beta headers reload job.
"""
# Check if user is admin
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Read-only status — admin viewers can read.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}",
@ -14380,7 +14746,8 @@ async def get_adaptive_router_state(
adaptive-router deployment. Each snapshot's `router_name` field identifies
which deployment it came from.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Read-only state — admin viewers can read.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail={"error": CommonProxyErrors.not_allowed_access.value},

View file

@ -15,6 +15,7 @@ from fastapi.responses import ORJSONResponse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import (
@ -22,10 +23,88 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
get_form_data,
)
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
)
router = APIRouter()
def _raise_vector_store_scan_depth_exceeded() -> None:
raise HTTPException(
status_code=400,
detail={
"error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values"
},
)
def _append_payload_to_scan_stack(
payload_stack: list[tuple[Any, int]],
value: Any,
next_depth: int,
) -> None:
if isinstance(value, dict):
if next_depth > DEFAULT_MAX_RECURSE_DEPTH:
_raise_vector_store_scan_depth_exceeded()
payload_stack.append((value, next_depth))
elif isinstance(value, list):
if next_depth > DEFAULT_MAX_RECURSE_DEPTH:
if any(isinstance(item, (dict, list)) for item in value):
_raise_vector_store_scan_depth_exceeded()
return
payload_stack.append((value, next_depth))
def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]:
vector_store_ids: set[str] = set()
payload_stack = [(payload, 0)]
while payload_stack:
current_payload, depth = payload_stack.pop()
if depth > DEFAULT_MAX_RECURSE_DEPTH:
_raise_vector_store_scan_depth_exceeded()
if isinstance(current_payload, dict):
for key, value in current_payload.items():
if key == "vector_store_id":
if not isinstance(value, str) or not value:
raise HTTPException(
status_code=400,
detail={
"error": "vector_store_id must be a non-empty string"
},
)
vector_store_ids.add(value)
continue
if isinstance(value, (dict, list)):
_append_payload_to_scan_stack(
payload_stack=payload_stack,
value=value,
next_depth=depth + 1,
)
elif isinstance(current_payload, list):
for item in current_payload:
_append_payload_to_scan_stack(
payload_stack=payload_stack,
value=item,
next_depth=depth + 1,
)
return vector_store_ids
async def _authorize_nested_vector_store_ids(
payload: Any,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
for vector_store_id in sorted(_collect_vector_store_ids_from_payload(payload)):
await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
def _build_file_metadata_entry(
response: Any,
file_data: Optional[Tuple[str, bytes, str]] = None,
@ -385,6 +464,11 @@ async def rag_ingest(
},
)
await _authorize_nested_vector_store_ids(
payload=ingest_options,
user_api_key_dict=user_api_key_dict,
)
# Add litellm data
request_data: Dict[str, Any] = {}
request_data = await add_litellm_data_to_request(
@ -537,11 +621,20 @@ async def rag_query(
status_code=400,
detail={"error": "retrieval_config is required"},
)
if not isinstance(retrieval_config, dict):
raise HTTPException(
status_code=400,
detail={"error": "retrieval_config must be an object"},
)
if "vector_store_id" not in retrieval_config:
raise HTTPException(
status_code=400,
detail={"error": "retrieval_config must contain 'vector_store_id'"},
)
await _authorize_nested_vector_store_ids(
payload=retrieval_config,
user_api_key_dict=user_api_key_dict,
)
# Add litellm data
request_data: Dict[str, Any] = {}

View file

@ -6,6 +6,18 @@ from fastapi import HTTPException, status
import litellm
from litellm.proxy._types import UserAPIKeyAuth
# Router-internal mock_testing_* flag names — kept in sync with
# ``litellm.types.router.MockRouterTestingParams`` by the test
# ``test_mock_testing_kwarg_names_matches_dataclass``. Hardcoding (rather
# than deriving via ``dataclasses.fields(MockRouterTestingParams)`` at
# import time) avoids a cyclic import: ``litellm.types.router`` imports
# back into proxy modules before this module finishes loading.
_MOCK_TESTING_KWARG_NAMES: tuple = (
"mock_testing_fallbacks",
"mock_testing_context_fallbacks",
"mock_testing_content_policy_fallbacks",
)
if TYPE_CHECKING:
from litellm.router import Router as _Router
@ -322,6 +334,13 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
"""
await add_shared_session_to_data(data)
# Strip router-internal mock_testing_* flags. Combined with an
# unauthorized fallback in ``router_settings_override`` they let a
# caller deterministically execute requests against restricted
# models. VERIA-44.
for _key in _MOCK_TESTING_KWARG_NAMES:
data.pop(_key, None)
team_id = get_team_id_from_data(data)
router_model_names = llm_router.model_names if llm_router is not None else []

File diff suppressed because it is too large Load diff

View file

@ -6,6 +6,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
@ -127,10 +128,10 @@ async def get_cloudzero_settings(
Only the first 4 and last 4 characters of the API key are shown.
Returns null/empty values when settings are not configured (consistent with other settings endpoints).
Only admin users can view CloudZero settings.
Only admin users (Proxy Admin or Admin Viewer) can view CloudZero settings.
"""
# Validation
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Validation — Admin Viewer follows the read-parity rule.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail={"error": CommonProxyErrors.not_allowed_access.value},

View file

@ -0,0 +1,85 @@
"""
Logging helpers for spend-tracking error paths.
Proxy operators have asked for a way to keep both their downstream log sinks
and the SpendLogs UI free of the stack traces that the spend-tracking
machinery emits when it hits 4xx/5xx or transient DB errors. The errors still
need to be logged (and still flow to Sentry via
``proxy_logging_obj.failure_handler``), but the multi-line stack traces
dominate log volume and clutter the per-row Metadata pane in the UI.
The opt-in is a single env var, ``LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS=true``,
gated by ``should_suppress_spend_log_tracebacks``. When it returns ``True``:
* ``spend_log_error`` drops the traceback from the console / structured log
record (this module), and
* the failure callback in ``proxy_track_cost_callback`` drops the
``error_information.traceback`` field from the SpendLogs row before it is
persisted, so the UI's per-row Metadata pane (which renders the metadata
JSON verbatim) stays clean. The key is omitted entirely rather than set
to ``""`` — ``StandardLoggingPayloadErrorInformation`` marks the field
optional and every downstream consumer uses ``.get("traceback")``.
At DEBUG the full traceback is always preserved so operators can still
troubleshoot. The UI suppression follows the same gate.
"""
import logging
import os
from typing import Any, Optional
from litellm._logging import verbose_proxy_logger
from litellm.secret_managers.main import str_to_bool
SUPPRESS_SPEND_LOG_TRACEBACKS_ENV = "LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS"
def _is_suppression_env_enabled() -> bool:
"""Read the opt-in env var fresh each call so dynamic flips are honored.
Kept separate from ``should_suppress_spend_log_tracebacks`` so tests and
other call sites can introspect just the env-var state without also
consulting the live logger level.
"""
return str_to_bool(os.getenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV)) is True
def should_suppress_spend_log_tracebacks() -> bool:
"""Return ``True`` when spend-log traceback suppression should apply.
Suppression only kicks in when both:
* the operator opted in via the env var, and
* the proxy logger is at INFO or above (i.e. not DEBUG) — at DEBUG we
still want full tracebacks for troubleshooting.
"""
if not _is_suppression_env_enabled():
return False
return not verbose_proxy_logger.isEnabledFor(logging.DEBUG)
def spend_log_error(
message: str,
*args: Any,
exc: Optional[BaseException] = None,
) -> None:
"""Log a spend-tracking error, with the traceback gated on the env var.
By default this behaves like ``verbose_proxy_logger.exception`` — the
active exception (or ``exc`` if supplied) is attached so the formatter
renders its traceback. When ``LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS`` is
truthy and the logger is at INFO or above, the traceback is dropped and
only ``message % args`` is emitted.
Sentry / ``proxy_logging_obj.failure_handler`` is NOT invoked here — call
sites still own the alerting path. This helper is purely about console /
structured-log output volume.
"""
if should_suppress_spend_log_tracebacks():
verbose_proxy_logger.error(message, *args)
return
if exc is not None:
verbose_proxy_logger.error(
message, *args, exc_info=(type(exc), exc, exc.__traceback__)
)
else:
verbose_proxy_logger.error(message, *args, exc_info=True)

View file

@ -25,6 +25,7 @@ from litellm.litellm_core_utils.core_helpers import (
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.proxy.utils import PrismaClient, hash_token
from litellm.types.utils import (
CostBreakdown,
@ -471,9 +472,7 @@ def get_logging_payload( # noqa: PLR0915
return payload
except Exception as e:
verbose_proxy_logger.exception(
"Error creating spendlogs object - {}".format(str(e))
)
spend_log_error("Error creating spendlogs object - %s", str(e), exc=e)
raise e

View file

@ -7,6 +7,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
@ -140,9 +141,10 @@ async def get_vantage_settings(
View current Vantage settings.
Returns the current Vantage configuration with the API key masked for security.
Only admin users can view Vantage settings.
Only admin users (Proxy Admin or Admin Viewer) can view Vantage settings.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Admin Viewer follows the read-parity rule.
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail={"error": CommonProxyErrors.not_allowed_access.value},

View file

@ -36,6 +36,7 @@ from litellm.proxy._types import (
SpendLogsMetadata,
SpendLogsPayload,
)
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypes, CallTypesLiteral
@ -3188,6 +3189,8 @@ class PrismaClient:
t.organization_id as org_id,
p.project_alias AS project_alias,
tm.spend AS team_member_spend,
b_tm.tpm_limit AS team_member_tpm_limit,
b_tm.rpm_limit AS team_member_rpm_limit,
m.aliases AS team_model_aliases,
-- Added comma to separate b.* columns
b.max_budget AS litellm_budget_table_max_budget,
@ -3203,6 +3206,7 @@ class PrismaClient:
FROM "LiteLLM_VerificationToken" AS v
LEFT JOIN "LiteLLM_TeamTable" AS t ON v.team_id = t.team_id
LEFT JOIN "LiteLLM_TeamMembership" AS tm ON v.team_id = tm.team_id AND tm.user_id = v.user_id
LEFT JOIN "LiteLLM_BudgetTable" AS b_tm ON tm.budget_id = b_tm.budget_id
LEFT JOIN "LiteLLM_ModelTable" m ON t.model_id = m.id
LEFT JOIN "LiteLLM_BudgetTable" AS b ON v.budget_id = b.budget_id
LEFT JOIN "LiteLLM_ProjectTable" AS p ON v.project_id = p.project_id
@ -5103,6 +5107,11 @@ async def update_daily_tag_spend(
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e:
# NOTE: keep this as a plain ``error`` (no traceback) to match the
# historical behavior of this site. ``spend_log_error`` would attach
# the active exception's traceback whenever the suppression env var
# is unset, which would be a regression for operators who never saw
# one here before.
verbose_proxy_logger.error(f"Error updating daily tag spend: {e}")
@ -5235,9 +5244,7 @@ async def _monitor_spend_logs_queue(
await asyncio.sleep(current_interval)
except Exception as e:
verbose_proxy_logger.error(
f"Error in spend logs queue monitor: {str(e)}\n{traceback.format_exc()}"
)
spend_log_error("Error in spend logs queue monitor: %s", str(e), exc=e)
# Continue monitoring even if there's an error, with exponential backoff
current_interval = min(current_interval * backoff_multiplier, max_backoff)
await asyncio.sleep(current_interval)

View file

@ -1,8 +1,6 @@
from typing import Any, Dict, Optional
from fastapi import APIRouter, Depends, HTTPException, Request, Response
import litellm
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
LiteLLM_ManagedVectorStore,
)
@ -10,7 +8,10 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.utils import jsonify_object
from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store,
get_litellm_managed_vector_store,
)
from litellm.types.vector_stores import IndexCreateRequest
router = APIRouter()
@ -19,24 +20,6 @@ router = APIRouter()
########################################################
async def _check_vector_store_access(
vector_store: LiteLLM_ManagedVectorStore,
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
"""
Check if the user has access to the vector store.
Delegates to :func:`can_user_access_vector_store`, which honors:
- PROXY_ADMIN bypass
- legacy vector stores with no team_id
- key-level and team-level ``object_permission.vector_stores`` allowlists
- team_id match between key and store
"""
return await can_user_access_vector_store(
vector_store=vector_store, user_api_key_dict=user_api_key_dict
)
async def _update_request_data_with_litellm_managed_vector_store_registry(
data: Dict,
vector_store_id: str,
@ -53,35 +36,27 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
Raises:
HTTPException: If user doesn't have access to the vector store
"""
if litellm.vector_store_registry is not None:
vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = (
litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
vector_store_id=vector_store_id
vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = (
await get_litellm_managed_vector_store(vector_store_id=vector_store_id)
)
if vector_store_to_run is not None:
if user_api_key_dict is not None:
await assert_user_can_access_vector_store(
vector_store=vector_store_to_run,
user_api_key_dict=user_api_key_dict,
)
)
if vector_store_to_run is not None:
if user_api_key_dict is not None:
if not await _check_vector_store_access(
vector_store_to_run, user_api_key_dict
):
raise HTTPException(
status_code=403,
detail="Access denied: You do not have permission to access this vector store",
)
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get(
"custom_llm_provider"
)
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider")
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
return data
@ -121,8 +96,7 @@ async def vector_store_search(
)
data = await _read_request_body(request=request)
if "vector_store_id" not in data:
data["vector_store_id"] = vector_store_id
data["vector_store_id"] = vector_store_id
# Check for legacy vector store registry (non-managed vector stores)
data = await _update_request_data_with_litellm_managed_vector_store_registry(

View file

@ -1,7 +1,9 @@
import json
from typing import Any, Dict, Literal, Optional
from fastapi import HTTPException, Request
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
LiteLLM_ObjectPermissionTable,
@ -13,6 +15,21 @@ from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.utils import ProviderConfigManager
def _normalize_litellm_params(
vector_store: LiteLLM_ManagedVectorStore,
) -> LiteLLM_ManagedVectorStore:
litellm_params = vector_store.get("litellm_params")
if isinstance(litellm_params, str):
normalized = LiteLLM_ManagedVectorStore(**dict(vector_store))
try:
parsed = json.loads(litellm_params)
normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {}
except (TypeError, ValueError):
normalized["litellm_params"] = {}
return normalized
return vector_store
def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
return (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
@ -120,6 +137,104 @@ async def can_user_access_vector_store(
return False
async def get_litellm_managed_vector_store(
vector_store_id: str,
) -> Optional[LiteLLM_ManagedVectorStore]:
"""
Resolve a LiteLLM-managed vector store from the registry or shared cache.
Provider-native vector store IDs will not be present in either location and
return None, preserving direct provider behavior while still protecting
LiteLLM-managed multi-tenant stores.
"""
if not vector_store_id:
return None
if litellm.vector_store_registry is not None:
try:
vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
vector_store_id=vector_store_id
)
if vector_store is not None:
return _normalize_litellm_params(vector_store)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to resolve vector store id=%s from registry: %s",
vector_store_id,
e,
)
raise HTTPException(
status_code=500,
detail="Unable to validate vector store access",
) from e
try:
from litellm.proxy.auth.auth_checks import (
get_managed_vector_store_rows_by_uuids,
)
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
return None
rows = await get_managed_vector_store_rows_by_uuids(
uuids=[vector_store_id],
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if not rows:
return None
return _normalize_litellm_params(
LiteLLM_ManagedVectorStore(**rows[0].model_dump())
)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to resolve vector store id=%s from shared cache: %s",
vector_store_id,
e,
)
raise HTTPException(
status_code=500,
detail="Unable to validate vector store access",
) from e
async def assert_user_can_access_vector_store(
vector_store: LiteLLM_ManagedVectorStore,
user_api_key_dict: UserAPIKeyAuth,
detail: str = "Access denied: You do not have permission to access this vector store",
) -> None:
"""Raise 403 unless the caller can access the resolved vector store."""
if not await can_user_access_vector_store(vector_store, user_api_key_dict):
raise HTTPException(status_code=403, detail=detail)
async def assert_user_can_access_vector_store_id(
vector_store_id: str,
user_api_key_dict: UserAPIKeyAuth,
detail: str = "Access denied: You do not have permission to access this vector store",
) -> Optional[LiteLLM_ManagedVectorStore]:
"""
Resolve a managed vector store id and enforce ownership if it exists.
Unknown ids are treated as provider-native ids and are not rejected here.
"""
vector_store = await get_litellm_managed_vector_store(
vector_store_id=vector_store_id
)
if vector_store is not None:
await assert_user_can_access_vector_store(
vector_store=vector_store,
user_api_key_dict=user_api_key_dict,
detail=detail,
)
return vector_store
def _does_endpoint_match(endpoint_path: str, request_path: str) -> bool:
if endpoint_path in request_path:
return True

View file

@ -17,9 +17,11 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
prepare_data_with_credentials,
)
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
is_allowed_to_call_vector_store_files_endpoint,
)
from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
if TYPE_CHECKING:
from litellm.router import Router
@ -193,6 +195,8 @@ def _update_request_data_with_litellm_managed_vector_store_registry(
data: Dict,
vector_store_id: str,
llm_router: Optional["Router"] = None,
managed_vector_store: Optional[LiteLLM_ManagedVectorStore] = None,
should_lookup_registry: bool = True,
) -> Dict:
"""
Update request data with model routing information from managed vector store.
@ -262,23 +266,27 @@ def _update_request_data_with_litellm_managed_vector_store_registry(
return data
# Legacy path: Check vector store registry for non-managed vector stores
if litellm.vector_store_registry is not None:
# Legacy path: Check vector store registry for non-managed vector stores.
vector_store_to_run = managed_vector_store
if (
vector_store_to_run is None
and should_lookup_registry
and litellm.vector_store_registry is not None
):
vector_store_to_run = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
vector_store_id=vector_store_id
)
if vector_store_to_run is not None:
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get(
"custom_llm_provider"
)
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
if vector_store_to_run is not None:
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider")
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
return data
@ -363,8 +371,11 @@ async def vector_store_file_create(
)
data = await _read_request_body(request=request)
if "vector_store_id" not in data:
data["vector_store_id"] = vector_store_id
data["vector_store_id"] = vector_store_id
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs if present in request body
original_managed_file_id = None
@ -375,7 +386,11 @@ async def vector_store_file_create(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -459,9 +474,18 @@ async def vector_store_file_list(
query_params = dict(request.query_params)
data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id}
data.update(query_params)
data["vector_store_id"] = vector_store_id
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -541,6 +565,10 @@ async def vector_store_file_retrieve(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -549,7 +577,11 @@ async def vector_store_file_retrieve(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -635,6 +667,10 @@ async def vector_store_file_content(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -643,7 +679,11 @@ async def vector_store_file_content(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -729,6 +769,10 @@ async def vector_store_file_update(
data = await _read_request_body(request=request)
data["vector_store_id"] = vector_store_id
data["file_id"] = file_id
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -737,7 +781,11 @@ async def vector_store_file_update(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -823,6 +871,10 @@ async def vector_store_file_delete(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -831,7 +883,11 @@ async def vector_store_file_delete(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)

View file

@ -12,11 +12,13 @@ import base64
import os
from base64 import b64encode
from typing import Optional
from urllib.parse import unquote
import httpx
from fastapi import APIRouter, Request, Response
from fastapi import APIRouter, HTTPException, Request, Response, status
import litellm
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
@ -27,6 +29,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
router = APIRouter()
default_vertex_config = None
_DEFAULT_LANGFUSE_HOST = "https://cloud.langfuse.com"
def create_request_copy(request: Request):
@ -39,6 +42,116 @@ def create_request_copy(request: Request):
}
def _decode_to_convergence(value: str) -> str:
previous = value
while True:
decoded = unquote(previous)
if decoded == previous:
return decoded
previous = decoded
def _normalize_langfuse_base_url(base_target_url: str) -> str:
if not (
base_target_url.startswith("http://") or base_target_url.startswith("https://")
):
# Existing behavior allows host-only Langfuse settings.
base_target_url = "http://" + base_target_url
try:
base_url = httpx.URL(base_target_url)
except Exception as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"Invalid Langfuse host: {str(e)}"},
)
if base_url.scheme not in ("http", "https") or not base_url.host:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse host"},
)
if base_url.userinfo:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Langfuse host must not include credentials"},
)
return str(base_url)
def _validate_langfuse_proxy_path(endpoint: str) -> str:
decoded_endpoint = _decode_to_convergence(endpoint)
if any(ord(char) < 32 for char in decoded_endpoint):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse endpoint path"},
)
if "\\" in decoded_endpoint or decoded_endpoint.startswith("//"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse endpoint path"},
)
endpoint_path = "/" + decoded_endpoint.lstrip("/")
if any(segment in (".", "..") for segment in endpoint_path.split("/")):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse endpoint path"},
)
return endpoint_path
def _get_langfuse_proxy_credentials(
*,
dynamic_host_supplied: bool,
dynamic_langfuse_public_key: Optional[str],
dynamic_langfuse_secret_key: Optional[str],
):
if dynamic_host_supplied:
if not dynamic_langfuse_public_key or not dynamic_langfuse_secret_key:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "Dynamic Langfuse hosts must include dynamic Langfuse credentials"
},
)
return dynamic_langfuse_public_key, dynamic_langfuse_secret_key
return (
dynamic_langfuse_public_key
or litellm.utils.get_secret(secret_name="LANGFUSE_PUBLIC_KEY"),
dynamic_langfuse_secret_key
or litellm.utils.get_secret(secret_name="LANGFUSE_SECRET_KEY"),
)
def _build_langfuse_proxy_target(
*,
endpoint: str,
base_target_url: str,
dynamic_host_supplied: bool,
):
endpoint_path = _validate_langfuse_proxy_path(endpoint)
base_url = httpx.URL(_normalize_langfuse_base_url(base_target_url))
updated_url = base_url.copy_with(path=endpoint_path)
custom_headers = {}
if dynamic_host_supplied and getattr(litellm, "user_url_validation", True):
try:
target_url, host_header = validate_url(str(updated_url))
except SSRFError as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"Invalid Langfuse host: {str(e)}"},
)
custom_headers["Host"] = host_header
return target_url, custom_headers
return str(updated_url), custom_headers
@router.api_route(
"/langfuse/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
@ -91,44 +204,33 @@ async def langfuse_proxy_route(
elif k == "langfuse_host":
dynamic_langfuse_host = v
dynamic_host_supplied = dynamic_langfuse_host is not None
base_target_url: str = (
dynamic_langfuse_host
or os.getenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
or "https://cloud.langfuse.com"
or os.getenv("LANGFUSE_HOST", _DEFAULT_LANGFUSE_HOST)
or _DEFAULT_LANGFUSE_HOST
)
if not (
base_target_url.startswith("http://") or base_target_url.startswith("https://")
):
# add http:// if unset, assume communicating over private network - e.g. render
base_target_url = "http://" + base_target_url
encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction
if not encoded_endpoint.startswith("/"):
encoded_endpoint = "/" + encoded_endpoint
# Construct the full target URL using httpx
base_url = httpx.URL(base_target_url)
updated_url = base_url.copy_with(path=encoded_endpoint)
# Add or update query parameters
langfuse_public_key = dynamic_langfuse_public_key or litellm.utils.get_secret(
secret_name="LANGFUSE_PUBLIC_KEY"
langfuse_public_key, langfuse_secret_key = _get_langfuse_proxy_credentials(
dynamic_host_supplied=dynamic_host_supplied,
dynamic_langfuse_public_key=dynamic_langfuse_public_key,
dynamic_langfuse_secret_key=dynamic_langfuse_secret_key,
)
langfuse_secret_key = dynamic_langfuse_secret_key or litellm.utils.get_secret(
secret_name="LANGFUSE_SECRET_KEY"
target_url, target_headers = _build_langfuse_proxy_target(
endpoint=endpoint,
base_target_url=base_target_url,
dynamic_host_supplied=dynamic_host_supplied,
)
langfuse_combined_key = "Basic " + b64encode(
f"{langfuse_public_key}:{langfuse_secret_key}".encode("utf-8")
).decode("ascii")
target_headers["Authorization"] = langfuse_combined_key
## CREATE PASS-THROUGH
endpoint_func = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers={"Authorization": langfuse_combined_key},
target=target_url,
custom_headers=target_headers,
query_params=dict(request.query_params), # type: ignore
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(

View file

@ -9351,6 +9351,52 @@ class Router:
"""
return [m for m in self.model_list if m["litellm_params"]["model"] == model]
def _try_early_resolve_deployments_for_model_not_in_names(
self, model: str, request_team_id: Optional[str]
) -> Optional[Tuple[str, Union[List, Dict]]]:
"""
When ``model`` is not in ``self.model_names``, try team routes, pattern routes,
team pattern routers, then default deployment. Returns None if none apply.
"""
if model in self.model_names:
return None
# Check for team-specific deployments by team_public_model_name.
# This intentionally takes priority over team pattern routers below,
# so that named team deployments shadow wildcard/pattern routes.
if request_team_id is not None:
team_deployments = self._get_all_deployments(
model_name=model, team_id=request_team_id
)
if team_deployments:
return model, team_deployments
pattern_deployments = self.pattern_router.get_deployments_by_pattern(
model=model,
)
if pattern_deployments:
return model, pattern_deployments
if request_team_id is not None and request_team_id in self.team_pattern_routers:
pattern_deployments = self.team_pattern_routers[
request_team_id
].get_deployments_by_pattern(
model=model,
)
if pattern_deployments:
return model, pattern_deployments
if self.default_deployment is not None:
# Shallow copy with nested litellm_params copy (100x+ faster than deepcopy)
updated_deployment = self.default_deployment.copy()
updated_deployment["litellm_params"] = self.default_deployment[
"litellm_params"
].copy()
updated_deployment["litellm_params"]["model"] = model
return model, updated_deployment
return None
def _common_checks_available_deployment(
self,
model: str,
@ -9393,56 +9439,52 @@ class Router:
if _model_from_alias is not None:
model = _model_from_alias
if model not in self.model_names:
# Check for team-specific deployments by team_public_model_name.
# This intentionally takes priority over team pattern routers below,
# so that named team deployments shadow wildcard/pattern routes.
if request_team_id is not None:
team_deployments = self._get_all_deployments(
model_name=model, team_id=request_team_id
)
if team_deployments:
return model, team_deployments
# check if provider/ specific wildcard routing use pattern matching
pattern_deployments = self.pattern_router.get_deployments_by_pattern(
model=model,
)
if pattern_deployments:
return model, pattern_deployments
if (
request_team_id is not None
and request_team_id in self.team_pattern_routers
):
pattern_deployments = self.team_pattern_routers[
request_team_id
].get_deployments_by_pattern(
model=model,
)
if pattern_deployments:
return model, pattern_deployments
# check if default deployment is set
if self.default_deployment is not None:
# Shallow copy with nested litellm_params copy (100x+ faster than deepcopy)
updated_deployment = self.default_deployment.copy()
updated_deployment["litellm_params"] = self.default_deployment[
"litellm_params"
].copy()
updated_deployment["litellm_params"]["model"] = model
return model, updated_deployment
early = self._try_early_resolve_deployments_for_model_not_in_names(
model=model, request_team_id=request_team_id
)
if early is not None:
return early
## get healthy deployments
### get all deployments
healthy_deployments = self._get_all_deployments(
model_name=model, team_id=request_team_id
)
_pre_model_access_group_filter_len = len(healthy_deployments)
healthy_deployments = self._filter_deployments_by_model_access_groups(
model=model,
healthy_deployments=healthy_deployments,
request_kwargs=request_kwargs,
request_team_id=request_team_id,
)
_access_group_filter_emptied_candidates = (
_pre_model_access_group_filter_len > 0 and len(healthy_deployments) == 0
)
if len(healthy_deployments) == 0:
# check if the user sent in a deployment name instead
healthy_deployments = self._get_deployment_by_litellm_model(model=model)
# Do not fall back when access-group filtering removed every candidate;
# _get_deployment_by_litellm_model does not re-apply that filter.
if _pre_model_access_group_filter_len == 0:
_litellm_model_deployments = self._get_deployment_by_litellm_model(
model=model
)
healthy_deployments = self._filter_deployments_by_model_access_groups(
model=model,
healthy_deployments=_litellm_model_deployments,
request_kwargs=request_kwargs,
request_team_id=request_team_id,
)
# If the litellm-model lookup produced candidates that access-group
# filtering then removed, treat this the same as the by-name path
# being emptied: prevent default-model fallback from bypassing the
# restriction (the fallback model may have no access_groups and
# would short-circuit the filter).
if (
len(_litellm_model_deployments) > 0
and len(healthy_deployments) == 0
):
_access_group_filter_emptied_candidates = True
if verbose_router_logger.isEnabledFor(logging.DEBUG):
verbose_router_logger.debug(
@ -9451,7 +9493,13 @@ class Router:
if len(healthy_deployments) == 0:
# Check for default fallbacks if no deployments are found for the requested model
if self._has_default_fallbacks():
# Do not fall back to another model when access-group filtering removed every
# candidate for the requested name: re-filtering the fallback model can be a
# no-op when it has no access_groups, incorrectly serving a different model.
if (
self._has_default_fallbacks()
and not _access_group_filter_emptied_candidates
):
fallback_model = self._get_first_default_fallback()
if fallback_model:
verbose_router_logger.info(
@ -9462,6 +9510,14 @@ class Router:
healthy_deployments = self._get_all_deployments(
model_name=model, team_id=request_team_id
)
healthy_deployments = (
self._filter_deployments_by_model_access_groups(
model=model,
healthy_deployments=healthy_deployments,
request_kwargs=request_kwargs,
request_team_id=request_team_id,
)
)
# If still no deployments after checking for fallbacks, raise an error
if len(healthy_deployments) == 0:
@ -9487,6 +9543,70 @@ class Router:
return model, healthy_deployments
def _filter_deployments_by_model_access_groups(
self,
model: str,
healthy_deployments: List,
request_kwargs: Optional[Dict],
request_team_id: Optional[str],
) -> List:
"""
Restrict candidate deployments to caller-authorized model access groups.
This is only applied when:
- request metadata includes `user_api_key_auth`, and
- caller permissions for this model are access-group-only
(no explicit model, wildcard, or all-proxy grants).
"""
if not healthy_deployments or request_kwargs is None:
return healthy_deployments
metadata = request_kwargs.get("metadata") or {}
litellm_metadata = request_kwargs.get("litellm_metadata") or {}
user_api_key_auth = metadata.get("user_api_key_auth") or litellm_metadata.get(
"user_api_key_auth"
)
if user_api_key_auth is None:
return healthy_deployments
object_models = set(getattr(user_api_key_auth, "models", []) or [])
object_team_models = set(getattr(user_api_key_auth, "team_models", []) or [])
allowed_models = object_models | object_team_models
if not allowed_models:
return healthy_deployments
# If caller has direct model/wildcard/all-proxy access, do not constrain
# deployment choice by access group.
if (
model in allowed_models
or "*" in allowed_models
or "all-proxy-models" in allowed_models
):
return healthy_deployments
access_groups_for_model = self.get_model_access_groups(
model_name=model, team_id=request_team_id
)
if len(access_groups_for_model) == 0:
return healthy_deployments
allowed_access_groups = set(access_groups_for_model.keys()) & allowed_models
if not allowed_access_groups:
# No overlap means this request was not authorized via model access
# group membership for this model, so do not force group filtering.
return healthy_deployments
filtered_deployments = []
for deployment in healthy_deployments:
deployment_model_info = deployment.get("model_info") or {}
deployment_access_groups = set(
deployment_model_info.get("access_groups", []) or []
)
if deployment_access_groups & allowed_access_groups:
filtered_deployments.append(deployment)
return filtered_deployments
async def async_get_healthy_deployments(
self,
model: str,
@ -10007,6 +10127,7 @@ class Router:
messages=messages,
input=input,
specific_deployment=specific_deployment,
request_kwargs=request_kwargs,
)
if isinstance(healthy_deployments, dict):

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