mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
7c94149aeb
200 changed files with 18182 additions and 1436 deletions
6
.github/workflows/create-release.yml
vendored
6
.github/workflows/create-release.yml
vendored
|
|
@ -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`,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
196
docs/my-website/docs/providers/crusoe.md
Normal file
196
docs/my-website/docs/providers/crusoe.md
Normal 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`
|
||||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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="",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
81
litellm/litellm_core_utils/secret_redaction.py
Normal file
81
litellm/litellm_core_utils/secret_redaction.py
Normal 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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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", {})},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
)
|
||||
|
|
|
|||
118
litellm/proxy/auth/trusted_proxy_utils.py
Normal file
118
litellm/proxy/auth/trusted_proxy_utils.py
Normal 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."
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
35
litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py
Normal file
35
litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py
Normal 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,
|
||||
}
|
||||
81
litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py
Normal file
81
litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
1029
litellm/proxy/spend_tracking/budget_reservation.py
Normal file
1029
litellm/proxy/spend_tracking/budget_reservation.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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},
|
||||
|
|
|
|||
85
litellm/proxy/spend_tracking/spend_log_error_logger.py
Normal file
85
litellm/proxy/spend_tracking/spend_log_error_logger.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue