Merge branch 'litellm_yj_may1' into codex/budget-race-enforcement

This commit is contained in:
yuneng-jiang 2026-05-01 14:32:18 -07:00 • committed by GitHub
commit c2cea58567
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
93 changed files with 4461 additions and 579 deletions

View file

@ -1,75 +0,0 @@
name: Check Lazy OpenAPI Snapshot
on:
pull_request:
branches:
- main
- litellm_internal_staging
- "litellm_**"
permissions:
contents: read
checks: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
verify:
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
with:
version: "0.10.9"
- name: Cache uv dependencies
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Install dependencies
run: uv sync --frozen --all-groups --all-extras
- name: Regenerate snapshot to /tmp
id: regen
run: |
cp litellm/proxy/_lazy_openapi_snapshot.json /tmp/snapshot.committed.json
uv run --no-sync python -m litellm.proxy._lazy_openapi_snapshot
mv litellm/proxy/_lazy_openapi_snapshot.json /tmp/snapshot.fresh.json
mv /tmp/snapshot.committed.json litellm/proxy/_lazy_openapi_snapshot.json
- name: Compare
id: diff
continue-on-error: true
run: |
diff -q /tmp/snapshot.fresh.json litellm/proxy/_lazy_openapi_snapshot.json
- name: Mark neutral if drift
if: steps.diff.outcome == 'failure'
uses: LouisBrunner/checks-action@6b626ffbad7cc56fd58627f774b9067e6118af23 # v2.0.0
with:
token: ${{ secrets.GITHUB_TOKEN }}
name: lazy-openapi-snapshot
conclusion: neutral
output: |
{
"title": "Lazy openapi snapshot is stale",
"summary": "Run `python -m litellm.proxy._lazy_openapi_snapshot` and commit the regenerated `litellm/proxy/_lazy_openapi_snapshot.json`. Not blocking — the snapshot will regenerate at release if not committed."
}

View file

@ -68,7 +68,7 @@ Managing LLM calls across providers gets complicated fast — different SDKs, au
<td><img height="60" alt="Stripe" src="https://github.com/user-attachments/assets/f7296d4f-9fbd-460d-9d05-e4df31697c4b" /></td>
<td><img height="60" alt="image" src="https://github.com/user-attachments/assets/436fca71-988b-40bb-b5fe-8450c80fdbd0" /></td>
<td><img height="60" alt="Google ADK" src="https://github.com/user-attachments/assets/caf270a2-5aee-45c4-8222-41a2070c4f19" /></td>
<td><img height="60" alt="Greptile" src="https://github.com/user-attachments/assets/0be4bd8a-7cfa-48d3-9090-f415fe948280" /></td>
<td><img height="60" alt="Greptile" src="https://github.com/user-attachments/assets/3db0ae72-0843-4005-a56d-bba1dde2193d" /></td>
<td><img height="60" alt="OpenHands" src="https://github.com/user-attachments/assets/a6150c4c-149e-4cae-888b-8b92be6e003f" /></td>
<td><h2>Netflix</h2></td>
<td><img height="60" alt="OpenAI Agents SDK" src="https://github.com/user-attachments/assets/c02f7be0-8c2e-4d27-aea7-7c024bfaebc0" /></td>

View file

@ -857,10 +857,16 @@ async def project_info(
where={"team_id": project.team_id}
)
if team:
is_team_member = (
user_api_key_dict.user_id in team.admins
or user_api_key_dict.user_id in team.members
)
caller_user_id = user_api_key_dict.user_id
for m in team.members_with_roles or []:
m_user_id = (
m.get("user_id")
if isinstance(m, dict)
else getattr(m, "user_id", None)
)
if m_user_id == caller_user_id:
is_team_member = True
break
if not (is_admin or is_team_member):
raise HTTPException(
@ -911,20 +917,20 @@ async def list_projects(
include={"litellm_budget_table": True, "object_permission": True}
)
else:
# Get projects for teams the user belongs to
user_teams = await prisma_client.db.litellm_teamtable.find_many(
where={
"OR": [
{"members": {"has": user_api_key_dict.user_id}},
{"admins": {"has": user_api_key_dict.user_id}},
]
}
# Look up the user's team memberships via the reverse-index on
# LiteLLM_UserTable.teams (maintained by team_member_add alongside
# members_with_roles). This avoids a full scan of all team rows.
user_record = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_api_key_dict.user_id},
)
user_team_ids = (
user_record.teams
if user_record is not None and user_record.teams
else []
)
team_ids = [team.team_id for team in user_teams]
projects = await prisma_client.db.litellm_projecttable.find_many(
where={"team_id": {"in": team_ids}},
where={"team_id": {"in": user_team_ids}},
include={"litellm_budget_table": True, "object_permission": True},
)

View file

@ -2,11 +2,23 @@
Arize Phoenix API client for fetching prompt versions from Arize Phoenix.
"""
import urllib.parse
from typing import Any, Dict, Optional
from litellm.llms.custom_httpx.http_handler import HTTPHandler
def _sanitize_id(identifier: str) -> str:
"""Reject path traversal characters and URL-encode the identifier."""
if any(c in identifier for c in ("/", "\\", "#", "?")):
raise ValueError(
f"Invalid identifier {identifier!r}: contains disallowed characters"
)
if ".." in identifier:
raise ValueError(f"Invalid identifier {identifier!r}: path traversal detected")
return urllib.parse.quote(identifier, safe="")
class ArizePhoenixClient:
"""
Client for interacting with Arize Phoenix API to fetch prompt versions.
@ -53,7 +65,8 @@ class ArizePhoenixClient:
Returns:
Dictionary containing prompt version data, or None if not found
"""
url = f"{self.api_base}/v1/prompt_versions/{prompt_version_id}"
safe_id = _sanitize_id(prompt_version_id)
url = f"{self.api_base}/v1/prompt_versions/{safe_id}"
try:
# Use the underlying httpx client directly to avoid query param extraction

View file

@ -3,11 +3,27 @@ BitBucket API client for fetching .prompt files from BitBucket repositories.
"""
import base64
import urllib.parse
from typing import Any, Dict, List, Optional
from litellm.llms.custom_httpx.http_handler import HTTPHandler
def _sanitize_file_path(file_path: str) -> str:
"""Reject path traversal and URL-encode each path segment."""
if "#" in file_path or "?" in file_path:
raise ValueError(
f"Invalid file path {file_path!r}: contains URL special characters"
)
parts = file_path.split("/")
for part in parts:
if part == "..":
raise ValueError(
f"Invalid file path {file_path!r}: path traversal detected"
)
return "/".join(urllib.parse.quote(part, safe="") for part in parts)
class BitBucketClient:
"""
Client for interacting with BitBucket API to fetch .prompt files.
@ -72,7 +88,8 @@ class BitBucketClient:
Returns:
File content as string, or None if file not found
"""
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{file_path}"
safe_path = _sanitize_file_path(file_path)
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}"
try:
response = self.http_handler.get(url, headers=self.headers)
@ -119,7 +136,8 @@ class BitBucketClient:
Returns:
List of file paths
"""
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{directory_path}"
safe_dir = _sanitize_file_path(directory_path) if directory_path else ""
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_dir}"
try:
response = self.http_handler.get(url, headers=self.headers)
@ -211,7 +229,8 @@ class BitBucketClient:
Returns:
Dictionary containing file metadata, or None if file not found
"""
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{file_path}"
safe_path = _sanitize_file_path(file_path)
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}"
try:
# Use GET with Range header to get just the headers (HEAD equivalent)

View file

@ -31,15 +31,23 @@ def load_cli_token() -> Optional[dict]:
return None
def get_litellm_gateway_api_key() -> Optional[str]:
def get_litellm_gateway_api_key(
expected_base_url: Optional[str] = None,
) -> Optional[str]:
"""
Get the stored CLI API key for use with LiteLLM SDK.
This function reads the token file created by `litellm-proxy login`
and returns the API key for use in Python scripts.
Args:
expected_base_url: When provided, the key is only returned if it was
originally issued for this URL. Pass the target server URL to
prevent credential leakage when the client is pointed at a
different (possibly malicious) server.
Returns:
str: The API key if found, None otherwise
str: The API key if found (and origin matches), None otherwise
Example:
>>> import litellm
@ -53,6 +61,10 @@ def get_litellm_gateway_api_key() -> Optional[str]:
>>> )
"""
token_data = load_cli_token()
if token_data and "key" in token_data:
return token_data["key"]
return None
if not token_data or "key" not in token_data:
return None
if expected_base_url is not None:
stored_url = token_data.get("base_url")
if stored_url != expected_base_url.rstrip("/"):
return None
return token_data["key"]

View file

@ -2244,7 +2244,7 @@ class CustomStreamWrapper:
asyncio.create_task(
self.logging_obj.async_failure_handler(e, traceback_exception)
)
raise e
self._handle_stream_fallback_error(e)
except Exception as e:
traceback_exception = traceback.format_exc()
if self.logging_obj is not None:

View file

@ -199,6 +199,47 @@ def validate_url(url: str) -> Tuple[str, str]:
return rewritten, host_header
def assert_same_origin(candidate_url: str, expected_url: str) -> None:
"""Verify ``candidate_url`` shares scheme, host, and port with ``expected_url``.
Use when an upstream API returns a URL meant for follow-up requests
(e.g. an async-job polling URL that will be hit with the operator's
API key in the headers). The upstream is trusted because the operator
configured ``api_base``, but the URL it hands back must actually point
back at the same origin or we'd be blindly forwarding credentials
wherever the upstream told us to.
Hostnames are compared case-insensitively. Default ports are made
explicit (HTTP→80, HTTPS→443) so ``https://api.example.com:443/...``
and ``https://api.example.com/...`` are treated as the same origin.
Error messages identify *which* component mismatched but never echo
the operator's ``expected`` host or the candidate's hostname back to
the caller — in the SSRF threat model the caller is the attacker,
and reflecting host info would be a secondary leak of operator
infrastructure details.
"""
candidate = urlparse(candidate_url)
expected = urlparse(expected_url)
if candidate.scheme not in _ALLOWED_SCHEMES:
raise SSRFError("URL scheme is not allowed")
if candidate.scheme != expected.scheme:
raise SSRFError("Origin mismatch on scheme")
candidate_host = _normalize_host(candidate.hostname or "")
expected_host = _normalize_host(expected.hostname or "")
if not candidate_host or candidate_host != expected_host:
raise SSRFError("Origin mismatch on host")
default_port = 443 if candidate.scheme == "https" else 80
candidate_port = candidate.port if candidate.port is not None else default_port
expected_port = expected.port if expected.port is not None else default_port
if candidate_port != expected_port:
raise SSRFError("Origin mismatch on port")
_MAX_REDIRECTS = 10

View file

@ -16,6 +16,7 @@ import litellm
from litellm.constants import AZURE_OPERATION_POLLING_TIMEOUT, DEFAULT_MAX_RETRIES
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -898,6 +899,17 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
operation_location_url = response.headers["operation-location"]
else:
raise AzureOpenAIError(status_code=500, message=response.text)
# Reject polling URLs that don't share an origin with ``api_base``.
# Without this an upstream-controlled or attacker-controlled
# value would receive the operator's Azure API key in the
# request headers below. VERIA-51.
try:
assert_same_origin(operation_location_url, api_base)
except SSRFError as ssrf_err:
raise AzureOpenAIError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
response = await async_handler.get(
url=operation_location_url,
headers=headers,
@ -908,8 +920,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
timeout_secs: int = AZURE_OPERATION_POLLING_TIMEOUT
start_time = time.time()
if "status" not in response.json():
raise Exception(
"Expected 'status' in response. Got={}".format(response.json())
# Don't reflect the raw response body — when the polling
# URL points at an internal JSON API (cloud metadata
# service etc.) reflecting it here turns Blind SSRF into
# Full-Read SSRF. VERIA-51.
raise AzureOpenAIError(
status_code=502,
message="Polling response missing 'status' field",
)
while response.json()["status"] not in ["succeeded", "failed"]:
if time.time() - start_time > timeout_secs:
@ -1009,6 +1026,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
operation_location_url = response.headers["operation-location"]
else:
raise AzureOpenAIError(status_code=500, message=response.text)
try:
assert_same_origin(operation_location_url, api_base)
except SSRFError as ssrf_err:
raise AzureOpenAIError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
response = sync_handler.get(
url=operation_location_url,
headers=headers,
@ -1019,8 +1043,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
timeout_secs: int = AZURE_OPERATION_POLLING_TIMEOUT
start_time = time.time()
if "status" not in response.json():
raise Exception(
"Expected 'status' in response. Got={}".format(response.json())
raise AzureOpenAIError(
status_code=502,
message="Polling response missing 'status' field",
)
while response.json()["status"] not in ["succeeded", "failed"]:
if time.time() - start_time > timeout_secs:

View file

@ -17,6 +17,7 @@ from urllib.parse import quote
import httpx
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.constants import (
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION,
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI,
@ -599,6 +600,16 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"Azure Document Intelligence returned 202 but no Operation-Location header found"
)
# Reject cross-origin polling URLs — the auth headers
# below would otherwise leak to whatever URL the upstream
# (or an attacker-controlled upstream) returns. VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(
f"Azure Document Intelligence: rejected polling URL ({ssrf_err})"
)
# Get headers for polling (need auth)
poll_headers = {
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get(
@ -711,6 +722,14 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"Azure Document Intelligence returned 202 but no Operation-Location header found"
)
# Reject cross-origin polling URLs (see sync path). VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(
f"Azure Document Intelligence: rejected polling URL ({ssrf_err})"
)
# Get headers for polling (need auth)
poll_headers = {
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get(

View file

@ -15,6 +15,7 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -331,6 +332,17 @@ class BlackForestLabsImageEdit:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}
@ -416,6 +428,17 @@ class BlackForestLabsImageEdit:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}

View file

@ -15,6 +15,7 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -317,6 +318,17 @@ class BlackForestLabsImageGeneration:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}
@ -402,6 +414,17 @@ class BlackForestLabsImageGeneration:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}

View file

@ -3,7 +3,7 @@ Google AI Studio /batchEmbedContents Embeddings Endpoint
"""
import json
from typing import Any, Dict, Literal, Optional, Union
from typing import Any, Dict, List, Literal, Optional, Tuple, Union
import httpx
@ -13,8 +13,8 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
)
from litellm.types.llms.openai import EmbeddingInput
from litellm.types.llms.vertex_ai import (
GeminiEmbeddingInput,
VertexAIBatchEmbeddingsRequestBody,
VertexAIBatchEmbeddingsResponseObject,
)
@ -23,7 +23,6 @@ from litellm.types.utils import EmbeddingResponse
from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from .batch_embed_content_transformation import (
_is_file_reference,
_is_multimodal_input,
process_embed_content_response,
process_response,
transform_openai_input_gemini_content,
@ -32,9 +31,24 @@ from .batch_embed_content_transformation import (
class GoogleBatchEmbeddings(VertexLLM):
@staticmethod
def _flatten_and_detect_file_refs(
input: GeminiEmbeddingInput,
) -> Tuple[List[str], bool]:
"""Flatten nested input lists and detect file references."""
input_list = [input] if isinstance(input, str) else input
flat_elements = [
e
for item in input_list
for e in (item if isinstance(item, list) else [item])
if isinstance(e, str)
]
has_file_refs = any(_is_file_reference(e) for e in flat_elements)
return flat_elements, has_file_refs
def _resolve_file_references(
self,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
api_key: str,
sync_handler: HTTPHandler,
) -> Dict[str, Dict[str, str]]:
@ -42,7 +56,7 @@ class GoogleBatchEmbeddings(VertexLLM):
Resolve Gemini file references (files/...) to get mime_type and uri.
Args:
input: EmbeddingInput that may contain file references
input: GeminiEmbeddingInput that may contain file references
api_key: Gemini API key
sync_handler: HTTP client
@ -73,7 +87,7 @@ class GoogleBatchEmbeddings(VertexLLM):
async def _async_resolve_file_references(
self,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
api_key: str,
async_handler: AsyncHTTPHandler,
) -> Dict[str, Dict[str, str]]:
@ -81,7 +95,7 @@ class GoogleBatchEmbeddings(VertexLLM):
Async version of _resolve_file_references.
Args:
input: EmbeddingInput that may contain file references
input: GeminiEmbeddingInput that may contain file references
api_key: Gemini API key
async_handler: Async HTTP client
@ -110,10 +124,10 @@ class GoogleBatchEmbeddings(VertexLLM):
return resolved_files
def batch_embeddings(
def batch_embeddings( # noqa: PLR0915
self,
model: str,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
print_verbose,
model_response: EmbeddingResponse,
custom_llm_provider: Literal["gemini", "vertex_ai"],
@ -151,8 +165,7 @@ class GoogleBatchEmbeddings(VertexLLM):
optional_params = optional_params or {}
is_multimodal = _is_multimodal_input(input)
use_embed_content = is_multimodal or (custom_llm_provider == "vertex_ai")
use_embed_content = custom_llm_provider == "vertex_ai"
mode: Literal["embedding", "batch_embedding"]
if use_embed_content:
mode = "embedding"
@ -215,8 +228,22 @@ class GoogleBatchEmbeddings(VertexLLM):
resolved_files=resolved_files,
)
else:
flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input)
if has_file_refs and not api_key:
raise ValueError(
"An API key is required to resolve Gemini file references (files/...). "
"Pass api_key= or set GEMINI_API_KEY."
)
resolved_files = {}
if api_key and has_file_refs:
resolved_files = self._resolve_file_references(
input=flat_elements, api_key=api_key, sync_handler=sync_handler
)
request_data = transform_openai_input_gemini_content(
input=input, model=model, optional_params=optional_params
input=input,
model=model,
optional_params=optional_params,
resolved_files=resolved_files,
)
## LOGGING
@ -264,7 +291,7 @@ class GoogleBatchEmbeddings(VertexLLM):
url: str,
data: Optional[Union[VertexAIBatchEmbeddingsRequestBody, dict]],
model_response: EmbeddingResponse,
input: EmbeddingInput,
input: GeminiEmbeddingInput,
timeout: Optional[Union[float, httpx.Timeout]],
headers={},
client: Optional[AsyncHTTPHandler] = None,
@ -303,8 +330,22 @@ class GoogleBatchEmbeddings(VertexLLM):
resolved_files=resolved_files,
)
else:
flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input)
if has_file_refs and not api_key:
raise ValueError(
"An API key is required to resolve Gemini file references (files/...). "
"Pass api_key= or set GEMINI_API_KEY."
)
resolved_files = {}
if api_key and has_file_refs:
resolved_files = await self._async_resolve_file_references(
input=flat_elements, api_key=api_key, async_handler=async_handler
)
data = transform_openai_input_gemini_content(
input=input, model=model, optional_params=optional_params or {}
input=input,
model=model,
optional_params=optional_params or {},
resolved_files=resolved_files,
)
## LOGGING

View file

@ -6,12 +6,12 @@ Why separate file? Make it easy to see how transformation works
from typing import Dict, List, Optional, Tuple
from litellm.types.llms.openai import EmbeddingInput
from litellm.types.llms.vertex_ai import (
BlobType,
ContentType,
EmbedContentRequest,
FileDataType,
GeminiEmbeddingInput,
PartType,
VertexAIBatchEmbeddingsRequestBody,
VertexAIBatchEmbeddingsResponseObject,
@ -114,33 +114,77 @@ def _parse_data_url(data_url: str) -> Tuple[str, str]:
return media_type, base64_data
def _is_multimodal_input(input: EmbeddingInput) -> bool:
def _is_multimodal_input(input: GeminiEmbeddingInput) -> bool:
"""
Check if the input contains multimodal data (data URIs, file references, or GCS URLs).
Check if the input contains multimodal data (data URIs, file references,
GCS URLs, or nested lists for combined embeddings).
Args:
input: EmbeddingInput (str or List[str])
input: GeminiEmbeddingInput — str, List[str], or List[List[str]] for combined embeddings
Returns:
bool: True if any element is a data URI, file reference, or GCS URL
bool: True if any element is multimodal or a nested list
"""
if isinstance(input, str):
input_list = [input]
else:
input_list = input
return _is_multimodal_element(input)
for element in input_list:
if isinstance(element, str):
if element.startswith("data:") and ";base64," in element:
return True
if _is_file_reference(element):
return True
if _is_gcs_url(element):
for element in input:
if isinstance(element, list):
if any(
_is_multimodal_element(sub) for sub in element if isinstance(sub, str)
):
return True
elif isinstance(element, str) and _is_multimodal_element(element):
return True
return False
def _is_multimodal_element(element: str) -> bool:
"""Check if a single string element is multimodal."""
if element.startswith("data:") and ";base64," in element:
return True
if _is_file_reference(element):
return True
if _is_gcs_url(element):
return True
return False
def _build_part_for_input(
element: str,
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
) -> PartType:
"""
Build a single PartType for an input element, handling text, data URIs,
file references, and GCS URLs.
"""
resolved_files = resolved_files or {}
if element.startswith("data:") and ";base64," in element:
mime_type, base64_data = _parse_data_url(element)
blob: BlobType = {"mime_type": mime_type, "data": base64_data}
return PartType(inline_data=blob)
elif _is_gcs_url(element):
mime_type = _infer_mime_type_from_gcs_url(element)
file_data: FileDataType = {
"mime_type": mime_type,
"file_uri": element,
}
return PartType(file_data=file_data)
elif _is_file_reference(element):
if element not in resolved_files:
raise ValueError(f"File reference {element} not resolved")
file_info = resolved_files[element]
file_data_ref: FileDataType = {
"mime_type": file_info["mime_type"],
"file_uri": file_info["uri"],
}
return PartType(file_data=file_data_ref)
else:
return PartType(text=element)
_SUPPORTED_EMBED_PARAMS = {"outputDimensionality", "taskType", "title"}
@ -155,37 +199,60 @@ def _filter_embed_params(optional_params: dict) -> dict:
def transform_openai_input_gemini_content(
input: EmbeddingInput, model: str, optional_params: dict
input: GeminiEmbeddingInput,
model: str,
optional_params: dict,
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
) -> VertexAIBatchEmbeddingsRequestBody:
"""
The content to embed. Only the parts.text fields will be counted.
Transform OpenAI embedding input to Gemini batchEmbedContents format.
Each input element becomes a separate EmbedContentRequest, supporting
text, data URIs, file references, and GCS URLs.
If an element is a list (nested input), all sub-elements are combined
into a single content with multiple parts, producing one combined
embedding for the group.
Examples:
input=["text", "image"] → 2 separate embeddings
input=[["text", "image"]] → 1 combined embedding
input=[["text", "image"], "x"] → 2 embeddings (1 combined + 1 separate)
"""
gemini_model_name = "models/{}".format(model)
gemini_params = _filter_embed_params(optional_params)
input_list = [input] if isinstance(input, str) else input
requests: List[EmbedContentRequest] = []
if isinstance(input, str):
for element in input_list:
if isinstance(element, list):
if not element:
raise ValueError("Nested input list must not be empty")
for sub in element:
if not isinstance(sub, str):
raise ValueError(
f"Elements inside a nested input list must be strings, got {type(sub)}"
)
parts = [
_build_part_for_input(sub, resolved_files=resolved_files)
for sub in element
]
else:
parts = [_build_part_for_input(element, resolved_files=resolved_files)]
request = EmbedContentRequest(
model=gemini_model_name,
content=ContentType(parts=[PartType(text=input)]),
content=ContentType(parts=parts),
**gemini_params,
)
requests.append(request)
else:
for i in input:
request = EmbedContentRequest(
model=gemini_model_name,
content=ContentType(parts=[PartType(text=i)]),
**gemini_params,
)
requests.append(request)
return VertexAIBatchEmbeddingsRequestBody(requests=requests)
def transform_openai_input_gemini_embed_content(
input: EmbeddingInput,
input: GeminiEmbeddingInput,
model: str,
optional_params: dict,
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
@ -194,7 +261,7 @@ def transform_openai_input_gemini_embed_content(
Transform OpenAI embedding input to Gemini embedContent format (multimodal).
Args:
input: EmbeddingInput (str or List[str]) with text, data URIs, or file references
input: GeminiEmbeddingInput with text, data URIs, or file references
model: Model name
optional_params: Additional parameters (taskType, outputDimensionality, etc.)
resolved_files: Dict mapping file names (files/abc) to {mime_type, uri}
@ -210,31 +277,14 @@ def transform_openai_input_gemini_embed_content(
parts: List[PartType] = []
for element in input_list:
if isinstance(element, list):
raise ValueError(
"Nested (combined) embeddings are not supported on the embedContent path. "
"Use the batchEmbedContents path or pass a flat list instead."
)
if not isinstance(element, str):
raise ValueError(f"Unsupported input type: {type(element)}")
if element.startswith("data:") and ";base64," in element:
mime_type, base64_data = _parse_data_url(element)
blob: BlobType = {"mime_type": mime_type, "data": base64_data}
parts.append(PartType(inline_data=blob))
elif _is_gcs_url(element):
mime_type = _infer_mime_type_from_gcs_url(element)
file_data: FileDataType = {
"mime_type": mime_type,
"file_uri": element,
}
parts.append(PartType(file_data=file_data))
elif _is_file_reference(element):
if element not in resolved_files:
raise ValueError(f"File reference {element} not resolved")
file_info = resolved_files[element]
file_data_ref: FileDataType = {
"mime_type": file_info["mime_type"],
"file_uri": file_info["uri"],
}
parts.append(PartType(file_data=file_data_ref))
else:
parts.append(PartType(text=element))
parts.append(_build_part_for_input(element, resolved_files=resolved_files))
request_body: dict = {
"content": ContentType(parts=parts),
@ -245,7 +295,7 @@ def transform_openai_input_gemini_embed_content(
def process_embed_content_response(
input: EmbeddingInput,
input: GeminiEmbeddingInput,
model_response: EmbeddingResponse,
model: str,
response_json: dict,
@ -291,7 +341,7 @@ def process_embed_content_response(
def process_response(
input: EmbeddingInput,
input: GeminiEmbeddingInput,
model_response: EmbeddingResponse,
model: str,
_predictions: VertexAIBatchEmbeddingsResponseObject,
@ -308,8 +358,29 @@ def process_response(
model_response.data = openai_embeddings
model_response.model = model
input_text = get_formatted_prompt(data={"input": input}, call_type="embedding")
prompt_tokens = token_counter(model=model, text=input_text)
has_nested = isinstance(input, list) and any(isinstance(e, list) for e in input)
if _is_multimodal_input(input) or has_nested:
input_list = input if isinstance(input, list) else [input]
text_elements: List[str] = []
for e in input_list:
if isinstance(e, list):
text_elements.extend(
sub
for sub in e
if isinstance(sub, str) and not _is_multimodal_element(sub)
)
elif isinstance(e, str) and not _is_multimodal_element(e):
text_elements.append(e)
if text_elements:
input_text = get_formatted_prompt(
data={"input": text_elements}, call_type="embedding"
)
prompt_tokens = token_counter(model=model, text=input_text)
else:
prompt_tokens = 0
else:
input_text = get_formatted_prompt(data={"input": input}, call_type="embedding")
prompt_tokens = token_counter(model=model, text=input_text)
model_response.usage = Usage(
prompt_tokens=prompt_tokens, total_tokens=prompt_tokens
)

View file

@ -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.",
@ -28329,7 +28329,7 @@
"/v1/vector_stores": {
"get": {
"description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list",
"operationId": "vector_store_list_v1_vector_stores_get",
"operationId": "vector_store_list_v1_vector_stores_get_2",
"parameters": [
{
"in": "query",
@ -28430,7 +28430,7 @@
},
"post": {
"description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```",
"operationId": "vector_store_create_v1_vector_stores_post",
"operationId": "vector_store_create_v1_vector_stores_post_2",
"responses": {
"200": {
"content": {
@ -28455,7 +28455,7 @@
"/v1/vector_stores/{vector_store_id}": {
"delete": {
"description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete",
"operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete",
"operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete_2",
"parameters": [
{
"in": "path",
@ -28499,7 +28499,7 @@
},
"get": {
"description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve",
"operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get",
"operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get_2",
"parameters": [
{
"in": "path",
@ -28543,7 +28543,7 @@
},
"post": {
"description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify",
"operationId": "vector_store_update_v1_vector_stores__vector_store_id__post",
"operationId": "vector_store_update_v1_vector_stores__vector_store_id__post_2",
"parameters": [
{
"in": "path",
@ -28588,7 +28588,7 @@
},
"/v1/vector_stores/{vector_store_id}/files": {
"get": {
"operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get",
"operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get_2",
"parameters": [
{
"in": "path",
@ -28631,7 +28631,7 @@
]
},
"post": {
"operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post",
"operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post_2",
"parameters": [
{
"in": "path",
@ -28676,7 +28676,7 @@
},
"/v1/vector_stores/{vector_store_id}/files/{file_id}": {
"delete": {
"operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete",
"operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete_2",
"parameters": [
{
"in": "path",
@ -28728,7 +28728,7 @@
]
},
"get": {
"operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get",
"operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get_2",
"parameters": [
{
"in": "path",
@ -28780,7 +28780,7 @@
]
},
"post": {
"operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post",
"operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post_2",
"parameters": [
{
"in": "path",
@ -28834,7 +28834,7 @@
},
"/v1/vector_stores/{vector_store_id}/files/{file_id}/content": {
"get": {
"operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get",
"operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get_2",
"parameters": [
{
"in": "path",
@ -28889,7 +28889,7 @@
"/v1/vector_stores/{vector_store_id}/search": {
"post": {
"description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search",
"operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post",
"operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post_2",
"parameters": [
{
"in": "path",
@ -28935,7 +28935,7 @@
"/vector_stores": {
"get": {
"description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list",
"operationId": "vector_store_list_vector_stores_get",
"operationId": "vector_store_list_vector_stores_get_2",
"parameters": [
{
"in": "query",
@ -29036,7 +29036,7 @@
},
"post": {
"description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```",
"operationId": "vector_store_create_vector_stores_post",
"operationId": "vector_store_create_vector_stores_post_2",
"responses": {
"200": {
"content": {
@ -29061,7 +29061,7 @@
"/vector_stores/{vector_store_id}": {
"delete": {
"description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete",
"operationId": "vector_store_delete_vector_stores__vector_store_id__delete",
"operationId": "vector_store_delete_vector_stores__vector_store_id__delete_2",
"parameters": [
{
"in": "path",
@ -29105,7 +29105,7 @@
},
"get": {
"description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve",
"operationId": "vector_store_retrieve_vector_stores__vector_store_id__get",
"operationId": "vector_store_retrieve_vector_stores__vector_store_id__get_2",
"parameters": [
{
"in": "path",
@ -29149,7 +29149,7 @@
},
"post": {
"description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify",
"operationId": "vector_store_update_vector_stores__vector_store_id__post",
"operationId": "vector_store_update_vector_stores__vector_store_id__post_2",
"parameters": [
{
"in": "path",
@ -29194,7 +29194,7 @@
},
"/vector_stores/{vector_store_id}/files": {
"get": {
"operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get",
"operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get_2",
"parameters": [
{
"in": "path",
@ -29237,7 +29237,7 @@
]
},
"post": {
"operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post",
"operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post_2",
"parameters": [
{
"in": "path",
@ -29282,7 +29282,7 @@
},
"/vector_stores/{vector_store_id}/files/{file_id}": {
"delete": {
"operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete",
"operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete_2",
"parameters": [
{
"in": "path",
@ -29334,7 +29334,7 @@
]
},
"get": {
"operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get",
"operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get_2",
"parameters": [
{
"in": "path",
@ -29386,7 +29386,7 @@
]
},
"post": {
"operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post",
"operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post_2",
"parameters": [
{
"in": "path",
@ -29440,7 +29440,7 @@
},
"/vector_stores/{vector_store_id}/files/{file_id}/content": {
"get": {
"operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get",
"operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get_2",
"parameters": [
{
"in": "path",
@ -29495,7 +29495,7 @@
"/vector_stores/{vector_store_id}/search": {
"post": {
"description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search",
"operationId": "vector_store_search_vector_stores__vector_store_id__search_post",
"operationId": "vector_store_search_vector_stores__vector_store_id__search_post_2",
"parameters": [
{
"in": "path",

View file

@ -11,9 +11,19 @@ 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_METHOD_SUFFIXES = {
"delete",
"get",
"head",
"options",
"patch",
"post",
"put",
"trace",
}
def _stabilize_multi_method_route_ids(routes) -> None:
@ -39,13 +49,46 @@ def load_snapshot() -> Optional[Dict[str, Dict]]:
return None
def _normalize_operation_ids(paths: Dict[str, Dict]) -> None:
"""Make FastAPI-generated operation IDs stable for multi-method routes.
FastAPI derives the default operation ID suffix from the first item in the
route's methods set. For routes registered with several HTTP methods, that
set iteration order can vary between processes, which makes the snapshot
drift even when no routes changed.
"""
for path_ops in paths.values():
if not isinstance(path_ops, dict):
continue
methods = {method for method in path_ops if method in HTTP_METHODS}
if not methods:
continue
for method, operation in path_ops.items():
if method not in HTTP_METHODS or not isinstance(operation, dict):
continue
operation_id = operation.get("operationId")
if not isinstance(operation_id, str):
continue
for suffix in methods:
suffix_token = f"_{suffix}"
if operation_id.endswith(suffix_token):
operation["operationId"] = (
operation_id[: -len(suffix_token)] + f"_{method}"
)
break
def generate_snapshot() -> Dict[str, Dict]:
import importlib
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:
@ -57,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
@ -67,13 +111,24 @@ def generate_snapshot() -> Dict[str, Dict]:
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 full.get("paths", {}).values():
for op in path_ops.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": full.get("paths", {}),
"paths": paths,
"components": {"schemas": full.get("components", {}).get("schemas", {})},
}
return fragments

View file

@ -668,6 +668,8 @@ class LiteLLMRoutes(enum.Enum):
"/models/{model_id}",
"/guardrails/list",
"/v2/guardrails/list",
"/project/list",
"/project/info",
]
+ spend_tracking_routes
+ key_management_routes
@ -692,6 +694,9 @@ class LiteLLMRoutes(enum.Enum):
"/model/{model_id}/update",
"/prompt/list",
"/prompt/info",
# Project read routes - endpoint scopes results to caller's teams (non-admin)
"/project/list",
"/project/info",
# Invitation routes - org/team admins checked in endpoint via _user_has_admin_privileges
"/invitation/new",
"/invitation/delete",

View file

@ -60,6 +60,10 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
_safe_get_request_query_params,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
TOOL_CAPABLE_CALL_TYPES,
@ -486,7 +490,10 @@ async def common_checks( # noqa: PLR0915
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
_model: Optional[Union[str, List[str]]] = get_model_from_request(
request_body, route
request_data=request_body,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
)
# 1. If team is blocked

View file

@ -2,7 +2,7 @@ import os
import re
import sys
from functools import lru_cache
from typing import Any, List, Optional, Tuple
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
from fastapi import HTTPException, Request, status
@ -167,6 +167,81 @@ def _allow_model_level_clientside_configurable_parameters(
)
# Config dicts whose entries are spread as ``**dict`` into outbound LLM
# API calls. ``litellm_embedding_config`` is consumed by the Milvus
# vector store transformer; future nested-config keys with the same
# threat shape should be added here.
_NESTED_CONFIG_KEYS: Tuple[str, ...] = ("litellm_embedding_config",)
# Banned root-level params. Same list applies to every entry in
# ``_NESTED_CONFIG_KEYS`` because those dicts get spread as ``**kwargs``
# into the same outbound calls.
_BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
"api_base",
"base_url",
"user_config",
"aws_sts_endpoint",
"aws_web_identity_token",
"aws_role_name",
"vertex_credentials",
# Endpoint-targeting fields that retarget the outbound request or
# an observability callback. An attacker-controlled value either
# exfiltrates the request payload (incl. messages + admin-set
# tokens) to the attacker's host, or coerces the proxy into
# authenticating against the attacker's host with admin secrets.
"aws_bedrock_runtime_endpoint",
"langsmith_base_url",
"langfuse_host",
"posthog_host",
"braintrust_host",
"slack_webhook_url",
# Provider-specific endpoint overrides that flow into the outbound
# request via ``optional_params``. Same threat as ``api_base``:
# ``s3_endpoint_url`` redirects Bedrock file uploads to attacker
# S3; ``sagemaker_base_url`` redirects all SageMaker traffic;
# ``deployment_url`` redirects SAP deployments.
"s3_endpoint_url",
"sagemaker_base_url",
"deployment_url",
)
def _check_banned_params(
body: dict,
general_settings: dict,
llm_router: Optional[Router],
model: str,
) -> None:
"""Raise ``ValueError`` if ``body`` carries a banned param without admin opt-in.
Shared between the root-level check and the nested-config check so a
new banned param only needs to be added in one place.
"""
for param in _BANNED_REQUEST_BODY_PARAMS:
if param not in body:
continue
if general_settings.get("allow_client_side_credentials") is True:
return
if (
_allow_model_level_clientside_configurable_parameters(
model=model,
param=param,
request_body_value=body[param],
llm_router=llm_router,
)
is True
):
return
raise ValueError(
f"Rejected Request: {param} is not allowed in request body. "
"Clientside passthrough requires explicit admin opt-in via "
"either `general_settings.allow_client_side_credentials = true` "
"(proxy-wide) or `configurable_clientside_auth_params` on the "
"deployment in your proxy config.yaml. "
"Relevant Issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997",
)
def is_request_body_safe(
request_body: dict, general_settings: dict, llm_router: Optional[Router], model: str
) -> bool:
@ -175,72 +250,31 @@ def is_request_body_safe(
A malicious user can set the api_base to their own domain and invoke POST /chat/completions to intercept and steal the OpenAI API key.
Relevant issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997
The blocklist is enforced unconditionally. Legitimate clientside
credential / endpoint passthrough goes through one of the two
explicit admin opt-ins (``general_settings.allow_client_side_credentials``
proxy-wide or ``configurable_clientside_auth_params`` per deployment).
Historically there was a third, *implicit*, *caller-controlled* path:
``check_complete_credentials`` returned True when the caller supplied
any non-empty ``api_key``, which made the entire blocklist a no-op.
That bypass turned every missing entry on the blocklist into an
exploitable SSRF / credential-exfil hole — see GHSA-jh89-88fc-qrfp,
GHSA-3frq-6r6h-7j64, and the chain of veria-admin findings (Dv_m860l,
b_yRJeQ5, stN90yjP, LBlyOAc8, U2TD78kg). Removed: the blocklist now
has a single, predictable failure mode for missing entries (a 400),
not a credential leak.
Iterative single-level descent into ``_NESTED_CONFIG_KEYS`` (rather
than recursion) covers nested-config attacks like Milvus's
``litellm_embedding_config.api_base`` (VERIA-6) without exposing a
recursion-depth DoS surface.
"""
banned_params = [
"api_base",
"base_url",
"user_config",
"aws_sts_endpoint",
"aws_web_identity_token",
"aws_role_name",
"vertex_credentials",
# Endpoint-targeting fields that retarget the outbound request or
# an observability callback. An attacker-controlled value either
# exfiltrates the request payload (incl. messages + admin-set
# tokens) to the attacker's host, or coerces the proxy into
# authenticating against the attacker's host with admin secrets.
"aws_bedrock_runtime_endpoint",
"langsmith_base_url",
"langfuse_host",
"posthog_host",
"braintrust_host",
"slack_webhook_url",
# Provider-specific endpoint overrides that flow into the outbound
# request via ``optional_params``. Same threat as ``api_base``:
# ``s3_endpoint_url`` redirects Bedrock file uploads to attacker
# S3; ``sagemaker_base_url`` redirects all SageMaker traffic;
# ``deployment_url`` redirects SAP deployments.
"s3_endpoint_url",
"sagemaker_base_url",
"deployment_url",
]
# The blocklist is enforced unconditionally. Legitimate clientside
# credential / endpoint passthrough goes through one of the two
# explicit admin opt-ins (``general_settings.allow_client_side_credentials``
# proxy-wide or ``configurable_clientside_auth_params`` per deployment).
# Historically there was a third, *implicit*, *caller-controlled* path:
# ``check_complete_credentials`` returned True when the caller supplied
# any non-empty ``api_key``, which made the entire blocklist a no-op.
# That bypass turned every missing entry on the blocklist into an
# exploitable SSRF / credential-exfil hole — see GHSA-jh89-88fc-qrfp,
# GHSA-3frq-6r6h-7j64, and the chain of veria-admin findings (Dv_m860l,
# b_yRJeQ5, stN90yjP, LBlyOAc8, U2TD78kg). Removed: the blocklist now
# has a single, predictable failure mode for missing entries (a 400),
# not a credential leak.
for param in banned_params:
if param in request_body:
if general_settings.get("allow_client_side_credentials") is True:
return True
elif (
_allow_model_level_clientside_configurable_parameters(
model=model,
param=param,
request_body_value=request_body[param],
llm_router=llm_router,
)
is True
):
return True
raise ValueError(
f"Rejected Request: {param} is not allowed in request body. "
"Clientside passthrough requires explicit admin opt-in via "
"either `general_settings.allow_client_side_credentials = true` "
"(proxy-wide) or `configurable_clientside_auth_params` on the "
"deployment in your proxy config.yaml. "
"Relevant Issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997",
)
_check_banned_params(request_body, general_settings, llm_router, model)
for nested_key in _NESTED_CONFIG_KEYS:
nested = request_body.get(nested_key)
if isinstance(nested, dict):
_check_banned_params(nested, general_settings, llm_router, model)
return True
@ -942,20 +976,257 @@ def get_end_user_id_from_request_body(
return None
def get_model_from_request(
request_data: dict, route: str
) -> Optional[Union[str, List[str]]]:
# First try to get model from request_data
model = request_data.get("model") or request_data.get("target_model_names")
MODEL_ROUTING_HEADER_NAME = "x-litellm-model"
_MODEL_ROUTING_ROUTE_MARKERS = (
"/files",
"/batches",
"/vector_stores",
"/skills",
"/evals",
"/fine_tuning",
"/videos",
)
_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS = (
"/files",
"/batches",
"/skills",
"/evals",
)
_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS = (
"/files",
"/batches",
"/fine_tuning",
)
_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS = (
"/files",
"/batches",
"/vector_stores",
)
_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS = ("/evals",)
_MODEL_ROUTING_ID_FIELDS = (
"file_id",
"input_file_id",
"output_file_id",
"error_file_id",
"batch_id",
"fine_tuning_job_id",
"training_file",
"validation_file",
"vector_store_id",
"video_id",
"character_id",
)
if model is not None:
model_names = model.split(",")
if len(model_names) == 1:
model = model_names[0].strip()
def _append_model_candidates(candidates: List[str], value: Any) -> None:
if value is None:
return
values = value if isinstance(value, (list, tuple, set)) else [value]
for item in values:
if item is None:
continue
if isinstance(item, str):
model_names = [model.strip() for model in item.split(",")]
else:
model = [m.strip() for m in model_names]
model_names = [str(item).strip()]
candidates.extend(model for model in model_names if model)
# If model not in request_data, try to extract from route
def _dedupe_model_candidates(candidates: List[str]) -> List[str]:
deduped: List[str] = []
for model in candidates:
if model not in deduped:
deduped.append(model)
return deduped
def _get_case_insensitive_mapping_value(
mapping: Optional[Mapping[str, Any]], key: str
) -> Any:
if not mapping:
return None
if key in mapping:
return mapping[key]
key_lower = key.lower()
for mapping_key, value in mapping.items():
if str(mapping_key).lower() == key_lower:
return value
return None
def _route_matches_any_marker(route: str, markers: Tuple[str, ...]) -> bool:
normalized_route = route.lower()
return any(marker in normalized_route for marker in markers)
def _route_uses_model_routing_sources(route: str) -> bool:
return _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_ROUTE_MARKERS)
def _extract_models_from_managed_resource_id(
resource_id: Any, resource_id_field: Optional[str] = None
) -> List[str]:
if not isinstance(resource_id, str) or not resource_id:
return []
candidates: List[str] = []
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
)
_append_model_candidates(
candidates=candidates, value=decode_model_from_file_id(resource_id)
)
unified_file_id = _is_base64_encoded_unified_file_id(resource_id)
if unified_file_id:
_append_model_candidates(
candidates=candidates,
value=get_models_from_unified_file_id(unified_file_id),
)
_append_model_candidates(
candidates=candidates,
value=get_model_id_from_unified_batch_id(unified_file_id),
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from managed file/batch ID: %s", str(e)
)
try:
from litellm.llms.base_llm.managed_resources.utils import parse_unified_id
parsed_id = parse_unified_id(resource_id)
if parsed_id:
_append_model_candidates(
candidates=candidates, value=parsed_id.get("model_id")
)
_append_model_candidates(
candidates=candidates, value=parsed_id.get("target_model_names")
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from unified managed resource ID: %s", str(e)
)
if resource_id_field in ("video_id", "character_id"):
try:
from litellm.types.videos.utils import (
decode_character_id_with_provider,
decode_video_id_with_provider,
)
if resource_id_field == "video_id":
_append_model_candidates(
candidates=candidates,
value=decode_video_id_with_provider(resource_id).get("model_id"),
)
else:
_append_model_candidates(
candidates=candidates,
value=decode_character_id_with_provider(resource_id).get(
"model_id"
),
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from managed video/character ID: %s", str(e)
)
return _dedupe_model_candidates(candidates)
def _extract_model_candidates_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
) -> List[str]:
candidates: List[str] = []
uses_model_routing_sources = _route_uses_model_routing_sources(route=route)
uses_header_or_query_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS
)
uses_query_target_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS
)
uses_body_target_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS
)
uses_completion_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS
)
body_model = request_data.get("model")
_append_model_candidates(candidates, body_model)
if uses_body_target_model_sources or not body_model:
_append_model_candidates(candidates, request_data.get("target_model_names"))
if uses_completion_model_sources and isinstance(
request_data.get("completion"), dict
):
_append_model_candidates(candidates, request_data["completion"].get("model"))
if uses_model_routing_sources:
if uses_header_or_query_model_sources:
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(request_query_params, "model"),
)
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(
request_headers, MODEL_ROUTING_HEADER_NAME
),
)
if uses_query_target_model_sources:
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(
request_query_params, "target_model_names"
),
)
for field in _MODEL_ROUTING_ID_FIELDS:
_append_model_candidates(
candidates,
_extract_models_from_managed_resource_id(
request_data.get(field), resource_id_field=field
),
)
return _dedupe_model_candidates(candidates)
def _format_model_candidates(
candidates: List[str],
) -> Optional[Union[str, List[str]]]:
if not candidates:
return None
if len(candidates) == 1:
return candidates[0]
return candidates
def get_model_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
) -> Optional[Union[str, List[str]]]:
candidates = _extract_model_candidates_from_request(
request_data=request_data,
route=route,
request_headers=request_headers,
request_query_params=request_query_params,
)
model = _format_model_candidates(candidates)
# If no explicit model was found, try to extract from route
if model is None:
# Parse model from route that follows the pattern /openai/deployments/{model}/*
match = re.match(r"/openai/deployments/([^/]+)", route)

View file

@ -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, 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:
@ -2199,6 +2241,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:
@ -2217,7 +2260,11 @@ async def _enforce_key_and_fallback_model_access(
):
pass
else:
model = get_model_from_request(request_data, route)
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
fallback_models = cast(
Optional[List[ALL_FALLBACK_MODEL_VALUES]],
request_data.get("fallbacks", None),
@ -2304,11 +2351,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
@ -2316,13 +2369,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
@ -2330,14 +2384,15 @@ async def _run_post_custom_auth_checks(
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and current_models
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=model_name,
)
# team / user / end_user / project context objects are fetched by
# the centralized common_checks gate in user_api_key_auth after

View file

@ -53,12 +53,16 @@ def clear_token() -> None:
os.remove(token_file)
def get_stored_api_key() -> Optional[str]:
"""Get the stored API key from token file"""
# Use the SDK-level utility
def get_stored_api_key(expected_base_url: Optional[str] = None) -> Optional[str]:
"""Get the stored API key from token file.
If expected_base_url is provided, the key is only returned when it was
originally issued for that URL. This prevents credential leakage when the
CLI is pointed at a different (possibly malicious) server.
"""
from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key
return get_litellm_gateway_api_key()
return get_litellm_gateway_api_key(expected_base_url=expected_base_url)
# Team selection utilities
@ -572,9 +576,11 @@ def login(ctx: click.Context):
api_key = auth_result["api_key"]
user_id = auth_result["user_id"]
# Save token data (simplified for CLI - we just need the key)
# Save token data. base_url is stored so we can verify origin
# before reusing the key on a subsequent CLI invocation.
save_token(
{
"base_url": base_url.rstrip("/"),
"key": api_key,
"user_id": user_id or "cli-user",
"user_email": "unknown",

View file

@ -74,9 +74,10 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
"""LiteLLM Proxy CLI - Manage your LiteLLM proxy server"""
ctx.ensure_object(dict)
# If no API key provided via flag or environment variable, try to load from saved token
# If no API key provided via flag or environment variable, try to load from saved token.
# Pass base_url so we only use the stored key when it was issued for this server.
if api_key is None:
api_key = get_stored_api_key()
api_key = get_stored_api_key(expected_base_url=base_url)
ctx.obj["base_url"] = base_url
ctx.obj["api_key"] = api_key

View file

@ -28,12 +28,17 @@ class Client:
api_key (Optional[str]): API key for authentication. If provided, it will be sent as a Bearer token.
timeout: Request timeout in seconds (default: 30)
"""
self._base_url = base_url.rstrip("/") # Remove trailing slash if present
self._api_key = get_litellm_gateway_api_key() or api_key
self._base_url = base_url.rstrip("/")
# Only use the stored CLI key when it was issued for this server.
self._api_key = api_key or get_litellm_gateway_api_key(
expected_base_url=self._base_url
)
# Initialize resource clients
self.http = HTTPClient(base_url=base_url, api_key=api_key, timeout=timeout)
self.http = HTTPClient(
base_url=base_url, api_key=self._api_key, timeout=timeout
)
self.models = ModelsManagementClient(
base_url=self._base_url, api_key=self._api_key
)

View file

@ -97,6 +97,55 @@ def _serialize_http_exception_detail(
return str(detail), None
def _collect_response_file_search_vector_store_ids(data: Dict[str, Any]) -> set[str]:
vector_store_ids: set[str] = set()
tools = data.get("tools")
if not isinstance(tools, list):
return vector_store_ids
for tool in tools:
if not isinstance(tool, dict) or tool.get("type") != "file_search":
continue
ids = tool.get("vector_store_ids") or []
if not isinstance(ids, list):
raise HTTPException(
status_code=400,
detail={
"error": "file_search.vector_store_ids must be a list of strings"
},
)
for vector_store_id in ids:
if not isinstance(vector_store_id, str) or not vector_store_id:
raise HTTPException(
status_code=400,
detail={
"error": "file_search.vector_store_ids must be a list of strings"
},
)
vector_store_ids.add(vector_store_id)
return vector_store_ids
async def _authorize_response_file_search_vector_stores(
data: Dict[str, Any],
user_api_key_dict: UserAPIKeyAuth,
) -> None:
vector_store_ids = _collect_response_file_search_vector_store_ids(data)
if not vector_store_ids:
return
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
)
for vector_store_id in sorted(vector_store_ids):
await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]:
"""Parses an event line and returns an error code if present, else None."""
event_line = (
@ -791,6 +840,11 @@ class ProxyBaseLLMRequestProcessing:
version=version,
proxy_config=proxy_config,
)
if route_type in {"aresponses", "_aresponses_websocket"}:
await _authorize_response_file_search_vector_stores(
data=self.data,
user_api_key_dict=user_api_key_dict,
)
# Calculate request queue time after add_litellm_data_to_request
# which sets arrival_time in proxy_server_request

View file

@ -225,10 +225,10 @@ class ToolPermissionGuardrail(CustomGuardrail):
def _parse_tool_call_arguments(
self, tool_call: ChatCompletionMessageToolCall
) -> Dict[str, Any]:
) -> tuple[Optional[Dict[str, Any]], Optional[str]]:
arguments = getattr(tool_call.function, "arguments", None)
if not arguments:
return {}
return None, "missing arguments"
parsed_arguments: Any = {}
try:
@ -236,22 +236,24 @@ class ToolPermissionGuardrail(CustomGuardrail):
parsed_arguments = json.loads(arguments)
elif isinstance(arguments, dict):
parsed_arguments = arguments
except json.JSONDecodeError as exc:
else:
return None, "arguments must be a JSON object"
except (json.JSONDecodeError, TypeError) as exc:
verbose_proxy_logger.warning(
"Tool Permission Guardrail: Failed to decode arguments for tool %s: %s",
tool_call.function.name,
exc,
)
return {}
return None, "arguments could not be parsed"
if isinstance(parsed_arguments, dict):
return parsed_arguments
return parsed_arguments, None
verbose_proxy_logger.debug(
"Tool Permission Guardrail: Ignoring non-dict arguments for tool %s",
"Tool Permission Guardrail: Rejecting non-dict arguments for tool %s",
tool_call.function.name,
)
return {}
return None, "arguments must be a JSON object"
def _collect_argument_paths(
self,
@ -331,10 +333,21 @@ class ToolPermissionGuardrail(CustomGuardrail):
continue
if rule.allowed_param_patterns and should_check_params:
arguments = self._parse_tool_call_arguments(tool_call)
arguments, parse_error = self._parse_tool_call_arguments(tool_call)
if parse_error:
default_message = f"Tool '{tool_identifier}' {parse_error} required by rule '{rule.id}'"
message = self.render_violation_message(
default=default_message,
context={"tool_name": tool_identifier, "rule_id": rule.id},
)
return False, rule.id, message
if not arguments:
last_pattern_failure_msg = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'"
continue
default_message = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'"
message = self.render_violation_message(
default=default_message,
context={"tool_name": tool_identifier, "rule_id": rule.id},
)
return False, rule.id, message
patterns_match, failure_message = self._patterns_match_for_rule(
arguments=arguments,
@ -365,6 +378,33 @@ class ToolPermissionGuardrail(CustomGuardrail):
)
return is_allowed, None, message
@staticmethod
def _get_mapping_value(item: Any, key: str) -> Any:
if isinstance(item, dict):
return item.get(key)
return getattr(item, key, None)
@staticmethod
def _legacy_function_call_id(choice_index: int) -> str:
return f"legacy_function_call_{choice_index}"
def _legacy_function_call_to_tool_call(
self, function_call: Any, choice_index: int
) -> Optional[ChatCompletionMessageToolCall]:
if function_call is None:
return None
function_name = self._get_mapping_value(function_call, "name")
arguments = self._get_mapping_value(function_call, "arguments") or ""
if not function_name:
return None
return ChatCompletionMessageToolCall(
id=self._legacy_function_call_id(choice_index),
type="function",
function={"name": function_name, "arguments": arguments},
)
def _extract_tool_calls_from_response(
self, response: ModelResponse
) -> List[ChatCompletionMessageToolCall]:
@ -379,13 +419,72 @@ class ToolPermissionGuardrail(CustomGuardrail):
"""
tool_calls = []
for choice in response.choices:
for choice_index, choice in enumerate(response.choices):
if isinstance(choice, Choices):
for tool in choice.message.tool_calls or []:
tool_calls.append(tool)
legacy_tool_call = self._legacy_function_call_to_tool_call(
getattr(choice.message, "function_call", None), choice_index
)
if legacy_tool_call is not None:
tool_calls.append(legacy_tool_call)
return tool_calls
def _get_request_tool_name(self, tool: Any) -> tuple[Optional[str], Optional[str]]:
tool_type = self._get_mapping_value(tool, "type")
if tool_type != "function":
return None, tool_type
function = self._get_mapping_value(tool, "function")
tool_name = self._get_mapping_value(function, "name")
return tool_name, tool_type
def _get_legacy_function_name(self, function: Any) -> Optional[str]:
return self._get_mapping_value(function, "name")
def _get_named_tool_choice(self, data: dict) -> Optional[str]:
tool_choice = data.get("tool_choice")
if not tool_choice or tool_choice in ("auto", "none", "required"):
return None
if isinstance(tool_choice, str):
return tool_choice
if self._get_mapping_value(tool_choice, "type") != "function":
return None
return self._get_mapping_value(
self._get_mapping_value(tool_choice, "function"), "name"
)
def _get_named_function_call(self, data: dict) -> Optional[str]:
function_call = data.get("function_call")
if not function_call or function_call in ("auto", "none"):
return None
if isinstance(function_call, str):
return function_call
return self._get_mapping_value(function_call, "name")
def _collect_request_tools(self, data: dict) -> List[tuple[str, Optional[str]]]:
request_tools: List[tuple[str, Optional[str]]] = []
for tool in data.get("tools") or []:
tool_name, tool_type = self._get_request_tool_name(tool)
if tool_name is not None:
request_tools.append((tool_name, tool_type))
for function in data.get("functions") or []:
function_name = self._get_legacy_function_name(function)
if function_name is not None:
request_tools.append((function_name, "function"))
for forced_tool_name in (
self._get_named_tool_choice(data),
self._get_named_function_call(data),
):
if forced_tool_name is not None:
request_tools.append((forced_tool_name, "function"))
return request_tools
def _modify_request_with_permission_errors(
self,
data: dict,
@ -410,19 +509,32 @@ class ToolPermissionGuardrail(CustomGuardrail):
for tool_use in denied_tool_names:
error_tool_names.add(tool_use)
# Modify the tools
tools: Optional[List[ChatCompletionToolParam]] = data.get("tools")
if tools is None:
return data
new_tools = []
for tool in tools:
if tool["type"] != "function":
continue
tool_name: str = tool["function"]["name"]
if tool_name not in error_tool_names:
if tools is not None:
new_tools = []
for tool in tools:
tool_name, tool_type = self._get_request_tool_name(tool)
if tool_type == "function" and tool_name in error_tool_names:
continue
new_tools.append(tool)
data["tools"] = new_tools
data["tools"] = new_tools
functions = data.get("functions")
if functions is not None:
data["functions"] = [
function
for function in functions
if self._get_legacy_function_name(function) not in error_tool_names
]
named_tool_choice = self._get_named_tool_choice(data)
if named_tool_choice in error_tool_names:
data["tool_choice"] = "none"
named_function_call = self._get_named_function_call(data)
if named_function_call in error_tool_names:
data["function_call"] = "none"
return data
def _create_permission_error_result(
@ -472,7 +584,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
error_results[tool_use.id] = error_result
# Modify the response content
for choice in response.choices:
for choice_index, choice in enumerate(response.choices):
if isinstance(choice, Choices):
filtered_tool_calls = []
error_messages = []
@ -490,6 +602,15 @@ class ToolPermissionGuardrail(CustomGuardrail):
filtered_tool_calls if filtered_tool_calls else None
)
legacy_tool_call = self._legacy_function_call_to_tool_call(
getattr(choice.message, "function_call", None), choice_index
)
if legacy_tool_call is not None:
legacy_error_result = error_results.get(legacy_tool_call.id)
if legacy_error_result is not None:
choice.message.function_call = None
error_messages.append(legacy_error_result.content)
# Add error messages to content
if error_messages:
existing_content = choice.message.content
@ -519,21 +640,16 @@ class ToolPermissionGuardrail(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return data
new_tools: Optional[List[ChatCompletionToolParam]] = data.get("tools")
if new_tools is None:
new_tools = self._collect_request_tools(data)
if not new_tools:
verbose_proxy_logger.warning(
"Tool Permission Guardrail: not running guardrail. No tools in data"
"Tool Permission Guardrail: not running guardrail. No tools or functions in data"
)
return data
# Check permissions for each tool
denied_tool_names = []
for tool in new_tools:
if tool["type"] != "function":
continue
tool_name: str = tool["function"]["name"]
tool_type: Optional[str] = tool.get("type")
for tool_name, tool_type in new_tools:
is_allowed, _, message = self._check_tool_permission(tool_name, tool_type)
if not is_allowed and message is not None:

View file

@ -29,6 +29,10 @@ ILLEGAL_DISPLAY_PARAMS = [
"exception", # internal; not JSON-serializable, never for display
"litellm_metadata", # internal tracking metadata with auth objects; not for display
]
# Provider routing fields. Allowed for proxy admins so they can see which
# region/version a deployment is checking; gated at the endpoint layer for
# non-admin callers (see _strip_admin_only_fields_from_health_result).
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS = ("api_base", "api_version")
MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"]

View file

@ -20,6 +20,7 @@ from litellm.proxy._types import (
CallInfo,
EnterpriseLicenseData,
Litellm_EntityType,
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
@ -28,6 +29,7 @@ from litellm.proxy._types import (
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.health_check import (
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS,
_clean_endpoint_data,
_update_litellm_params_for_health_check,
perform_health_check,
@ -723,6 +725,90 @@ async def _save_background_health_checks_to_db(
# Continue execution - don't let database save failure break health checks
_PROXY_ADMIN_ROLES = frozenset(
{
LitellmUserRoles.PROXY_ADMIN.value,
# View-only admins are operators (oncall, support); they need the
# routing fields (api_base, api_version) to diagnose health and tell
# which provider region a check is hitting. They cannot mutate config
# so granting them the read-only view is safe.
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
}
)
def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
"""
Return True if the caller has a proxy-admin role (full or view-only).
user_role on UserAPIKeyAuth can be either a LitellmUserRoles enum or its
string value depending on how the auth path constructed the object, so we
compare against the raw value rather than the enum identity.
"""
role = user_api_key_dict.user_role
if role is None:
return False
role_value = role.value if hasattr(role, "value") else role
return role_value in _PROXY_ADMIN_ROLES
def _strip_admin_only_fields_from_health_result(result: dict) -> dict:
"""
Return a copy of the /health response with provider routing fields
(``api_base``, ``api_version``) removed from each healthy/unhealthy
endpoint entry. Used to hide those fields from non-admin callers while
still showing them which deployments they own and whether each one is
healthy. Proxy admins receive the unmodified result.
"""
out = dict(result)
drop = set(ADMIN_ONLY_HEALTH_DISPLAY_PARAMS)
for key in ("healthy_endpoints", "unhealthy_endpoints"):
eps = out.get(key)
if isinstance(eps, list):
out[key] = [
(
{k: v for k, v in ep.items() if k not in drop}
if isinstance(ep, dict)
else ep
)
for ep in eps
]
return out
def _filter_health_check_results_by_model_ids(
results: dict, allowed_model_ids: set
) -> dict:
"""
Restrict a cached background health-check result dict to endpoints whose
model_id is in ``allowed_model_ids``.
Endpoints without a model_id (e.g. CLI-model entries that predate the
model_id wiring) are dropped conservatively — we cannot prove they belong
to the caller, so they are excluded rather than leaked.
Each retained endpoint is shallow-copied before being returned, so any
downstream transform (e.g. _strip_admin_only_fields_from_health_result)
cannot accidentally mutate the shared ``health_check_results`` cache.
"""
healthy = [
dict(ep)
for ep in (results.get("healthy_endpoints") or [])
if ep.get("model_id") in allowed_model_ids
]
unhealthy = [
dict(ep)
for ep in (results.get("unhealthy_endpoints") or [])
if ep.get("model_id") in allowed_model_ids
]
return {
"healthy_endpoints": healthy,
"unhealthy_endpoints": unhealthy,
"healthy_count": len(healthy),
"unhealthy_count": len(unhealthy),
}
async def _perform_health_check_and_save(
model_list,
target_model,
@ -771,6 +857,7 @@ async def _perform_health_check_and_save(
@router.get("/health", tags=["health"], dependencies=[Depends(user_api_key_auth)])
async def health_endpoint(
response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
model: Optional[str] = fastapi.Query(
None, description="Specify the model name (optional)"
@ -838,11 +925,26 @@ async def health_endpoint(
detail={"error": f"Model with ID {model_id} not found"},
)
is_admin = _is_proxy_admin(user_api_key_dict)
def _post_process(result: dict) -> dict:
# api_base / api_version reveal which provider/region/internal host the
# deployment talks to; only proxy admins receive them. Non-admin keys
# still see model/model_id and the healthy/unhealthy status. We also
# set a header so non-admin clients that previously parsed those
# fields can detect the change programmatically.
if is_admin:
return result
response.headers["Litellm-Health-Field-Notice"] = (
"api_base and api_version are admin-only on this endpoint"
)
return _strip_admin_only_fields_from_health_result(result)
try:
if llm_model_list is None:
# if no router set, check if user set a model using litellm --model ollama/llama2
if user_model is not None:
return await _perform_health_check_and_save(
cli_result = await _perform_health_check_and_save(
model_list=[],
target_model=None,
cli_model=user_model,
@ -853,20 +955,59 @@ async def health_endpoint(
model_id=None, # CLI model doesn't have model_id
max_concurrency=health_check_concurrency,
)
return _post_process(cli_result)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": "Model list not initialized"},
)
_llm_model_list = copy.deepcopy(llm_model_list)
### FILTER MODELS FOR ONLY THOSE USER HAS ACCESS TO ###
# Live path: scope by model_name (every deployment has one).
# Cache path: scope by model_id (the cache is keyed on model_id).
# Consequence: a deployment whose model_name the caller can access
# but which lacks model_info.id will appear in the live /health
# response but NOT in the background-cache /health response. This is
# surfaced via the "warnings" field below so operators can fix the
# missing model_info.id rather than guess at the discrepancy.
if len(user_api_key_dict.models) > 0:
pass
else:
pass #
allowed_models = set(user_api_key_dict.models)
_llm_model_list = [
m for m in _llm_model_list if m.get("model_name") in allowed_models
]
if use_background_health_checks:
return health_check_results
if len(user_api_key_dict.models) > 0:
allowed_model_ids = {
(m.get("model_info") or {}).get("id")
for m in _llm_model_list
if (m.get("model_info") or {}).get("id")
}
filtered = _filter_health_check_results_by_model_ids(
health_check_results, allowed_model_ids
)
if not allowed_model_ids:
# Caller has accessible model_names but none of the
# matching deployments expose a model_info.id, so the
# cache filter (which keys on model_id) drops every
# entry. Surface this both as a warning log and a
# structured "warnings" field on the response so the
# caller can distinguish "no deployments found" from
# "deployments excluded due to missing model_info.id".
verbose_proxy_logger.warning(
"health_endpoint: scoped key %s has accessible models %s "
"but none of the matching deployments carry a model_info.id; "
"background health-check cache will return an empty result.",
user_api_key_dict.user_id,
list(user_api_key_dict.models),
)
filtered["warnings"] = [
"Some accessible deployments are missing model_info.id "
"and were excluded from this response. Ask a proxy admin "
"to populate model_info.id for these models."
]
return _post_process(filtered)
return _post_process(health_check_results)
else:
return await _perform_health_check_and_save(
router_result = await _perform_health_check_and_save(
model_list=_llm_model_list,
target_model=target_model,
cli_model=None,
@ -877,6 +1018,7 @@ async def health_endpoint(
model_id=model_id,
max_concurrency=health_check_concurrency,
)
return _post_process(router_result)
except Exception as e:
verbose_proxy_logger.error(
"litellm.proxy.proxy_server.py::health_endpoint(): Exception occured - {}".format(

View file

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

View file

@ -1,6 +1,7 @@
import asyncio
import json
import time
import urllib.parse
from datetime import datetime
from typing import Literal, Optional
from urllib.parse import urlparse
@ -203,8 +204,16 @@ class AssemblyAIPassthroughLoggingHandler:
)
if _api_key is None:
raise ValueError("AssemblyAI API key not found")
if (
any(c in transcript_id for c in ("/", "\\", "#", "?"))
or ".." in transcript_id
):
raise ValueError(
f"Invalid transcript_id {transcript_id!r}: contains disallowed characters"
)
safe_transcript_id = urllib.parse.quote(transcript_id, safe="")
try:
url = f"{_base_url}/v2/transcript/{transcript_id}"
url = f"{_base_url}/v2/transcript/{safe_transcript_id}"
headers = {
"Authorization": f"Bearer {_api_key}",
"Content-Type": "application/json",

View file

@ -6,6 +6,7 @@ import inspect
import io
import os
import random
import re
import secrets
import shutil
import subprocess
@ -955,6 +956,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 +1044,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 +1124,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 +1156,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:
@ -13167,9 +13250,12 @@ async def update_config( # noqa: PLR0915
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
For Admin UI - allows admin to update config via UI
For Admin UI - allows admin to update config via UI.
Currently supports modifying General Settings + LiteLLM settings
Writes only the sections present in the request body to LiteLLM_Config rows
(one row per top-level section). Sections the caller did not send are left
untouched — this endpoint never persists pre-existing YAML values to DB as
a side effect of an unrelated update.
"""
global llm_router, llm_model_list, general_settings, proxy_config, proxy_logging_obj, master_key, prisma_client
try:
@ -13177,109 +13263,96 @@ async def update_config( # noqa: PLR0915
raise HTTPException(
status_code=403, detail="Only proxy admins can update config"
)
import base64
"""
- Update the ConfigTable DB
- Run 'add_deployment'
"""
if prisma_client is None:
raise Exception("No DB Connected")
if store_model_in_db is not True:
raise HTTPException(
status_code=500,
detail={
"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."
async def _read_section(param_name: str) -> dict:
row = await prisma_client.db.litellm_config.find_first(
where={"param_name": param_name}
)
if row is None or row.param_value is None:
return {}
return dict(row.param_value)
async def _upsert_section(param_name: str, value: dict) -> None:
serialized = json.dumps(value)
await prisma_client.db.litellm_config.upsert(
where={"param_name": param_name},
data={
"create": {"param_name": param_name, "param_value": serialized},
"update": {"param_value": serialized},
},
)
# invalidate the DualCache entry so the next reader (this process
# or any other proxy in the cluster) goes to DB.
await invalidate_config_param(param_name)
updated_settings = config_info.json(exclude_none=True)
updated_settings = prisma_client.jsonify_object(updated_settings)
for k, v in updated_settings.items():
if k == "router_settings":
await prisma_client.db.litellm_config.upsert(
where={"param_name": k},
data={
"create": {"param_name": k, "param_value": v},
"update": {"param_value": v},
},
)
await invalidate_config_param(k)
### OLD LOGIC [TODO] MOVE TO DB ###
# Load existing config
config = await proxy_config.get_config()
verbose_proxy_logger.debug("Loaded config: %s", config)
# update the general settings
# general_settings: merge per-key, with the alert_to_webhook_url side
# effect of auto-enabling slack alerting.
if config_info.general_settings is not None:
config.setdefault("general_settings", {})
updated_general_settings = config_info.general_settings.dict(
exclude_none=True
)
_existing_settings = config["general_settings"]
for k, v in updated_general_settings.items():
# overwrite existing settings with updated values
existing = await _read_section("general_settings")
updates = config_info.general_settings.dict(exclude_none=True)
for k, v in updates.items():
if k == "alert_to_webhook_url":
# check if slack is already enabled. if not, enable it
if "alerting" not in _existing_settings:
_existing_settings = {"alerting": ["slack"]}
elif isinstance(_existing_settings["alerting"], list):
if "slack" not in _existing_settings["alerting"]:
_existing_settings["alerting"].append("slack")
_existing_settings[k] = v
config["general_settings"] = _existing_settings
if "alerting" not in existing:
existing["alerting"] = ["slack"]
elif (
isinstance(existing["alerting"], list)
and "slack" not in existing["alerting"]
):
existing["alerting"].append("slack")
existing[k] = v
await _upsert_section("general_settings", existing)
# environment_variables: encrypt request values, then merge into existing.
if config_info.environment_variables is not None:
config.setdefault("environment_variables", {})
_updated_environment_variables = config_info.environment_variables
existing = await _read_section("environment_variables")
for k, v in config_info.environment_variables.items():
existing[k] = encrypt_value_helper(value=v)
await _upsert_section("environment_variables", existing)
# encrypt updated_environment_variables #
for k, v in _updated_environment_variables.items():
encrypted_value = encrypt_value_helper(value=v)
_updated_environment_variables[k] = encrypted_value
_existing_env_variables = config["environment_variables"]
for k, v in _updated_environment_variables.items():
# overwrite existing env variables with updated values
_existing_env_variables[k] = _updated_environment_variables[k]
# update the litellm settings
# litellm_settings: merge existing + request, request wins (matching
# router_settings semantics — the caller's value for any given key is
# what gets persisted). success_callback is special-cased: it is
# always normalized + deduped, and unioned with any existing list,
# because callbacks are additive (callers send the new entry, not
# the full set). Normalizing on every write — not only when an
# existing entry is present — keeps the DB free of mixed-case
# entries that delete_callback (lowercase lookup) cannot find.
if config_info.litellm_settings is not None:
config.setdefault("litellm_settings", {})
updated_litellm_settings = config_info.litellm_settings
config["litellm_settings"] = {
**updated_litellm_settings,
**config["litellm_settings"],
}
existing = await _read_section("litellm_settings")
updated_litellm_settings = dict(config_info.litellm_settings)
# if litellm.success_callback in updated_litellm_settings and config["litellm_settings"]
if (
"success_callback" in updated_litellm_settings
and "success_callback" in config["litellm_settings"]
):
# check both success callback are lists
if isinstance(
config["litellm_settings"]["success_callback"], list
) and isinstance(updated_litellm_settings["success_callback"], list):
updated_success_callbacks_normalized = normalize_callback_names(
updated_litellm_settings["success_callback"]
)
combined_success_callback = (
config["litellm_settings"]["success_callback"]
+ updated_success_callbacks_normalized
)
combined_success_callback = list(set(combined_success_callback))
config["litellm_settings"][
"success_callback"
] = combined_success_callback
incoming_cb = updated_litellm_settings.get("success_callback")
if isinstance(incoming_cb, list):
updated_litellm_settings["success_callback"] = normalize_callback_names(
incoming_cb
)
# Save the updated config
await proxy_config.save_config(new_config=config)
merged = {**existing, **updated_litellm_settings}
incoming_cb = updated_litellm_settings.get("success_callback")
existing_cb = existing.get("success_callback")
if isinstance(incoming_cb, list):
if isinstance(existing_cb, list):
# Normalize the existing list too — a row written by a
# different code path may still hold mixed-case names,
# which would otherwise dedup-miss against the lowercase
# incoming entries.
merged["success_callback"] = list(
set(normalize_callback_names(existing_cb) + incoming_cb)
)
else:
merged["success_callback"] = list(set(incoming_cb))
await _upsert_section("litellm_settings", merged)
# router_settings: merge existing + request, request wins.
if config_info.router_settings is not None:
existing = await _read_section("router_settings")
updates = config_info.router_settings.dict(exclude_none=True)
await _upsert_section("router_settings", {**existing, **updates})
await proxy_config.add_deployment(
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj

View file

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

View file

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

View file

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

View file

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

View file

@ -6,6 +6,13 @@ from typing_extensions import (
TypedDict,
)
from litellm.types.llms.openai import EmbeddingInput
# Gemini supports nested-list inputs (e.g. [["text", "image"]]) as an explicit
# opt-in for combined embeddings — a provider-specific extension of the
# OpenAI-faithful EmbeddingInput shape.
GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]]
class FunctionResponse(TypedDict):
name: str

View file

@ -377,15 +377,11 @@ def search(
_is_async = kwargs.pop("asearch", False) is True
# pull credentials from registry if available
vector_store_id_for_credentials = kwargs.get("vector_store_id", vector_store_id)
if (
litellm.vector_store_registry is not None
and vector_store_id_for_credentials is not None
):
if litellm.vector_store_registry is not None and vector_store_id is not None:
try:
registry_credentials = (
litellm.vector_store_registry.get_credentials_for_vector_store(
vector_store_id_for_credentials
vector_store_id
)
)
kwargs.update(registry_credentials)

View file

@ -34894,6 +34894,20 @@
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"zai.glm-5": {
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3.2e-06,
"supports_function_calling": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"zai.glm-4.7-flash": {
"input_cost_per_token": 7e-08,
"litellm_provider": "bedrock_converse",

View file

@ -477,3 +477,4 @@ def test_get_llm_provider_use_proxy_arg_true_with_direct_args():
assert provider == "litellm_proxy"
assert key == arg_api_key # Should use the argument key
assert base == arg_api_base # Should use the argument base

View file

@ -134,3 +134,62 @@ def test_is_assemblyai_route():
== False
)
assert handler.is_assemblyai_route("") == False
# --- Security: SSRF via transcript_id path traversal ---
def test_get_assembly_transcript_rejects_slash_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("../../admin/credentials")
def test_get_assembly_transcript_rejects_dotdot_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("..evil")
def test_get_assembly_transcript_rejects_fragment_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("abc#suffix")
def test_get_assembly_transcript_rejects_query_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("abc?x=1")
def test_get_assembly_transcript_allows_valid_id(
assembly_handler, mock_transcript_response
):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with patch("httpx.get") as mock_get:
mock_get.return_value.json.return_value = mock_transcript_response
mock_get.return_value.raise_for_status.return_value = None
transcript = assembly_handler._get_assembly_transcript(
"abc123-valid-id_xyz"
)
assert transcript == mock_transcript_response
called_url = mock_get.call_args[0][0]
assert "abc123-valid-id_xyz" in called_url
assert ".." not in called_url

View file

@ -2768,40 +2768,40 @@ async def test_update_config_success_callback_normalization():
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy._types import ConfigYAML
# Ensure feature is enabled and prisma_client is set
setattr(proxy_server, "store_model_in_db", True)
setattr(proxy_server, "proxy_logging_obj", MagicMock())
existing_litellm_settings = {"success_callback": ["langfuse"]}
class FakeRow:
def __init__(self, name, value):
self.param_name = name
self.param_value = value
upserted = {}
async def fake_find_first(where=None):
if where and where.get("param_name") == "litellm_settings":
return FakeRow("litellm_settings", existing_litellm_settings)
return None
async def fake_upsert(where=None, data=None):
upserted[where["param_name"]] = json.loads(data["update"]["param_value"])
class MockPrisma:
def __init__(self):
self.db = MagicMock()
self.db.litellm_config = MagicMock()
self.db.litellm_config.upsert = AsyncMock()
# proxy_server.update_config expects this to be sync returning a dict
def jsonify_object(self, obj):
return obj
self.db.litellm_config.find_first = AsyncMock(side_effect=fake_find_first)
self.db.litellm_config.upsert = AsyncMock(side_effect=fake_upsert)
setattr(proxy_server, "prisma_client", MockPrisma())
class MockProxyConfig:
def __init__(self):
self.saved_config = None
async def get_config(self):
# Existing config has one lowercase callback already
return {"litellm_settings": {"success_callback": ["langfuse"]}}
async def save_config(self, new_config: dict):
self.saved_config = new_config
async def add_deployment(self, prisma_client=None, proxy_logging_obj=None):
return None
mock_proxy_config = MockProxyConfig()
setattr(proxy_server, "proxy_config", mock_proxy_config)
setattr(proxy_server, "proxy_config", MockProxyConfig())
# Update config with mixed-case callbacks - expect normalization to lowercase
config_update = ConfigYAML(litellm_settings={"success_callback": ["SQS", "sQs"]})
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
@ -2810,9 +2810,10 @@ async def test_update_config_success_callback_normalization():
)
await proxy_server.update_config(config_update, user_api_key_dict=admin_user)
saved = mock_proxy_config.saved_config
assert saved is not None, "save_config was not called"
callbacks = saved["litellm_settings"]["success_callback"]
assert (
"litellm_settings" in upserted
), "litellm_config.upsert was not called for litellm_settings"
callbacks = upserted["litellm_settings"]["success_callback"]
# Deduped and normalized
assert "sqs" in callbacks

View file

@ -280,3 +280,55 @@ class TestDynamicProjectNameOnSpan:
if __name__ == "__main__":
unittest.main()
# --- Security: SSRF via prompt_version_id path traversal ---
def test_arize_phoenix_client_sanitize_id_rejects_traversal():
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
# dotdot without slashes
with pytest.raises(ValueError, match="path traversal"):
_sanitize_id("..something")
# full traversal (slash caught first)
with pytest.raises(ValueError, match="disallowed characters"):
_sanitize_id("../../projects")
def test_arize_phoenix_client_sanitize_id_rejects_slash():
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
with pytest.raises(ValueError, match="disallowed characters"):
_sanitize_id("valid/extra")
def test_arize_phoenix_client_sanitize_id_rejects_fragment():
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
with pytest.raises(ValueError, match="disallowed characters"):
_sanitize_id("abc#suffix")
def test_arize_phoenix_client_sanitize_id_rejects_query():
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
with pytest.raises(ValueError, match="disallowed characters"):
_sanitize_id("abc?x=1")
def test_arize_phoenix_client_sanitize_id_allows_uuid():
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
uid = "550e8400-e29b-41d4-a716-446655440000"
assert _sanitize_id(uid) == uid
def test_arize_phoenix_client_get_prompt_version_rejects_traversal():
from litellm.integrations.arize.arize_phoenix_client import ArizePhoenixClient
client = ArizePhoenixClient(
api_key="test-key", api_base="https://app.phoenix.arize.com"
)
with pytest.raises(ValueError, match="disallowed characters"):
client.get_prompt_version("../../projects")

View file

@ -11,6 +11,7 @@ sys.path.insert(
import litellm
from litellm.integrations.bitbucket import BitBucketPromptManager
from litellm.integrations.bitbucket.bitbucket_client import _sanitize_file_path
@patch("litellm.integrations.bitbucket.bitbucket_prompt_manager.BitBucketClient")
@ -370,3 +371,45 @@ def test_bitbucket_prompt_manager_list_templates(mock_client_class):
templates = manager.prompt_manager.list_templates()
assert isinstance(templates, list)
assert "test_prompt" in templates
# --- Security: path traversal / SSRF ---
def test_sanitize_file_path_rejects_traversal():
with pytest.raises(ValueError, match="path traversal"):
_sanitize_file_path("../../etc/passwd")
def test_sanitize_file_path_rejects_fragment():
with pytest.raises(ValueError, match="URL special characters"):
_sanitize_file_path("secret#.prompt")
def test_sanitize_file_path_rejects_query():
with pytest.raises(ValueError, match="URL special characters"):
_sanitize_file_path("secret?.prompt")
def test_sanitize_file_path_encodes_special_chars():
result = _sanitize_file_path("prompts/my prompt.prompt")
assert result == "prompts/my%20prompt.prompt"
def test_sanitize_file_path_allows_normal_paths():
assert _sanitize_file_path("prompts/my-prompt") == "prompts/my-prompt"
assert _sanitize_file_path("simple") == "simple"
def test_bitbucket_client_rejects_traversal_in_get_file_content():
from litellm.integrations.bitbucket.bitbucket_client import BitBucketClient
client = BitBucketClient(
{
"workspace": "ws",
"repository": "repo",
"access_token": "tok",
}
)
with pytest.raises(ValueError, match="path traversal"):
client.get_file_content("../../admin/credentials")

View file

@ -878,6 +878,39 @@ def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging):
assert "invalid maxOutputTokens" in str(excinfo.value)
@pytest.mark.asyncio
async def test_async_streaming_read_timeout_triggers_midstream_fallback(
logging_obj: Logging,
):
"""A mid-stream httpx.ReadTimeout must wrap into MidStreamFallbackError so
the Router's FallbackStreamWrapper can switch to a fallback model.
Previously __anext__ caught httpx.TimeoutException and re-raised it raw,
which bypassed _handle_stream_fallback_error and prevented stream_timeout
from triggering fallbacks the way connection-phase timeout does.
"""
import httpx
from litellm.exceptions import MidStreamFallbackError
async def _raise_read_timeout(**kwargs):
raise httpx.ReadTimeout("Timeout on reading data from socket")
response = CustomStreamWrapper(
completion_stream=None,
model="gpt-4",
logging_obj=logging_obj,
custom_llm_provider="openai",
make_call=_raise_read_timeout,
)
with pytest.raises(MidStreamFallbackError) as excinfo:
await response.__anext__()
assert excinfo.value.is_pre_first_chunk is True
assert isinstance(excinfo.value.original_exception, Exception)
def test_streaming_handler_with_created_time_propagation(
initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging
):

View file

@ -394,3 +394,76 @@ class TestHostAllowlist:
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
validate_url("http://internal.corp/")
# ── assert_same_origin ────────────────────────────────────────────────────────
from litellm.litellm_core_utils.url_utils import assert_same_origin
def test_assert_same_origin_matches_scheme_host_port():
"""A polling URL on the same scheme + host + port as the api_base
passes — the upstream is trusted; the URL it returned points back at
the same upstream."""
assert_same_origin(
"https://api.example.com/v1/operations/abc",
"https://api.example.com/v1/generate",
)
def test_assert_same_origin_treats_default_ports_as_explicit():
"""``https://x/`` and ``https://x:443/`` are the same origin."""
assert_same_origin("https://api.example.com/poll", "https://api.example.com:443/")
assert_same_origin("https://api.example.com:443/poll", "https://api.example.com/")
assert_same_origin("http://api.example.com/poll", "http://api.example.com:80/")
def test_assert_same_origin_rejects_different_host():
with pytest.raises(SSRFError, match="host"):
assert_same_origin(
"https://attacker.example.com/poll",
"https://api.example.com/generate",
)
def test_assert_same_origin_rejects_different_scheme():
with pytest.raises(SSRFError, match="scheme"):
assert_same_origin(
"http://api.example.com/poll", "https://api.example.com/generate"
)
def test_assert_same_origin_rejects_different_port():
with pytest.raises(SSRFError, match="port"):
assert_same_origin(
"https://api.example.com:8443/poll", "https://api.example.com/generate"
)
def test_assert_same_origin_rejects_non_http_scheme():
"""``file://`` polling URLs are rejected outright — the upstream
should never return a non-HTTP scheme."""
with pytest.raises(SSRFError, match="scheme"):
assert_same_origin("file:///etc/passwd", "https://api.example.com/")
def test_assert_same_origin_case_insensitive_host():
assert_same_origin(
"https://API.example.com/poll", "https://api.example.com/generate"
)
def test_assert_same_origin_error_message_does_not_leak_hostnames():
"""Greptile P2: in the SSRF threat model the caller is the attacker.
The error message must not echo the operator's expected host or the
attacker-supplied candidate host back to the caller — only identify
*which* component mismatched."""
with pytest.raises(SSRFError) as exc:
assert_same_origin(
"https://attacker.example.com:1234/poll",
"https://api.internal-corp.example/generate",
)
detail = str(exc.value)
assert "attacker.example.com" not in detail
assert "api.internal-corp.example" not in detail

View file

@ -0,0 +1,177 @@
"""
VERIA-51: polling URLs returned by upstream APIs (Azure DALL-E,
Azure Document Intelligence, Black Forest Labs) used to be followed
without origin validation. The handlers attached the operator's API
key to the polling request, so an attacker who could influence the
upstream response (or a compromised upstream) could redirect the proxy
to send credentials anywhere.
These tests assert each handler now rejects polling URLs that don't
share an origin with the original request URL.
"""
from unittest.mock import MagicMock, patch
import httpx
import pytest
# Azure DALL-E sync + async paths route through ``assert_same_origin``
# the same way as the cases below. The helper itself is unit-tested in
# ``tests/test_litellm/litellm_core_utils/test_url_utils.py``; the
# tests here exercise the wiring at sites with simpler signatures.
# ── Azure Document Intelligence polling ───────────────────────────────────────
def test_azure_di_sync_rejects_cross_origin_polling():
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
config = AzureDocumentIntelligenceOCRConfig()
raw_response = MagicMock()
raw_response.status_code = 202
raw_response.headers = {
"Operation-Location": "https://attacker.example.com/results/xyz",
}
raw_response.request = MagicMock()
raw_response.request.url = (
"https://eastus.cognitiveservices.azure.com/documentintelligence/.../analyze"
)
raw_response.request.headers = {"Ocp-Apim-Subscription-Key": "leak-me"}
with pytest.raises(ValueError, match="rejected polling URL"):
config.transform_ocr_response(
model="azure-doc-intel",
raw_response=raw_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
response={},
)
# ── Black Forest Labs polling ─────────────────────────────────────────────────
def test_bfl_image_generation_sync_rejects_cross_origin_polling():
from litellm.llms.black_forest_labs.image_generation.handler import (
BlackForestLabsImageGeneration,
)
handler = BlackForestLabsImageGeneration()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://attacker.example.com/get_result"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
sync_client = MagicMock()
sync_client.get = MagicMock()
with pytest.raises(Exception, match="Rejected polling URL"):
handler._poll_for_result_sync(
initial_response=initial_response,
headers={"x-key": "secret"},
sync_client=sync_client,
)
sync_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_bfl_image_generation_async_rejects_cross_origin_polling():
from litellm.llms.black_forest_labs.image_generation.handler import (
BlackForestLabsImageGeneration,
)
handler = BlackForestLabsImageGeneration()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://attacker.example.com/get_result"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
async_client = MagicMock()
async_client.get = MagicMock()
with pytest.raises(Exception, match="Rejected polling URL"):
await handler._poll_for_result_async(
initial_response=initial_response,
headers={"x-key": "secret"},
async_client=async_client,
)
async_client.get.assert_not_called()
def test_bfl_image_edit_sync_rejects_cross_origin_polling():
from litellm.llms.black_forest_labs.image_edit.handler import (
BlackForestLabsImageEdit,
)
handler = BlackForestLabsImageEdit()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://attacker.example.com/get_result"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro/edit"
sync_client = MagicMock()
sync_client.get = MagicMock()
with pytest.raises(Exception, match="Rejected polling URL"):
handler._poll_for_result_sync(
initial_response=initial_response,
headers={"x-key": "secret"},
sync_client=sync_client,
)
sync_client.get.assert_not_called()
def test_bfl_image_generation_same_origin_polling_passes():
"""Sanity check: when the polling URL shares origin with the original
request, the origin check passes and polling proceeds."""
from litellm.llms.black_forest_labs.image_generation.handler import (
BlackForestLabsImageGeneration,
)
handler = BlackForestLabsImageGeneration()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://api.bfl.ai/v1/get_result?id=abc"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
sync_client = MagicMock()
poll_response = MagicMock()
poll_response.status_code = 200
poll_response.json = MagicMock(return_value={"status": "Ready"})
sync_client.get = MagicMock(return_value=poll_response)
result = handler._poll_for_result_sync(
initial_response=initial_response,
headers={"x-key": "secret"},
sync_client=sync_client,
)
sync_client.get.assert_called_once()
assert result is poll_response

View file

@ -0,0 +1,290 @@
"""
Tests for Gemini batchEmbedContents transformation logic.
Covers:
- Text-only inputs (single and batch)
- Multimodal inputs (data URIs, GCS URLs, file references)
- Mixed text + multimodal inputs
- Response processing with correct indices
"""
import pytest
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
_build_part_for_input,
_is_multimodal_input,
process_response,
transform_openai_input_gemini_content,
transform_openai_input_gemini_embed_content,
)
from litellm.types.llms.vertex_ai import VertexAIBatchEmbeddingsResponseObject
from litellm.types.utils import EmbeddingResponse
IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
GCS_URL = "gs://my-bucket/image.png"
class TestIsMultimodalInput:
def test_text_only_string(self):
assert _is_multimodal_input("hello world") is False
def test_text_only_list(self):
assert _is_multimodal_input(["hello", "world"]) is False
def test_data_uri(self):
assert _is_multimodal_input([IMAGE_DATA_URI]) is True
def test_gcs_url(self):
assert _is_multimodal_input([GCS_URL]) is True
def test_file_reference(self):
assert _is_multimodal_input(["files/abc123"]) is True
def test_mixed_text_and_image(self):
assert _is_multimodal_input(["hello", IMAGE_DATA_URI]) is True
def test_nested_text_is_not_multimodal(self):
"""Nested list with text is not multimodal."""
assert _is_multimodal_input([["text_a", "text_b"]]) is False
def test_nested_list_with_image_is_multimodal(self):
assert _is_multimodal_input([["a red shoe", IMAGE_DATA_URI]]) is True
class TestBuildPartForInput:
def test_text_input(self):
part = _build_part_for_input("hello")
assert part["text"] == "hello"
assert part.get("inline_data") is None
def test_data_uri_input(self):
part = _build_part_for_input(IMAGE_DATA_URI)
assert part.get("text") is None
assert part["inline_data"] is not None
assert part["inline_data"]["mime_type"] == "image/png"
def test_gcs_url_input(self):
part = _build_part_for_input(GCS_URL)
assert part.get("text") is None
assert part["file_data"] is not None
assert part["file_data"]["mime_type"] == "image/png"
assert part["file_data"]["file_uri"] == GCS_URL
def test_file_reference_resolved(self):
resolved = {"files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"}}
part = _build_part_for_input("files/abc", resolved_files=resolved)
assert part["file_data"] is not None
assert part["file_data"]["mime_type"] == "image/jpeg"
def test_file_reference_unresolved_raises(self):
with pytest.raises(ValueError, match="not resolved"):
_build_part_for_input("files/abc")
class TestTransformOpenaiInputGeminiContent:
"""Test that transform_openai_input_gemini_content creates separate requests per input."""
def test_single_text(self):
result = transform_openai_input_gemini_content(
input="hello", model="gemini-embedding-2-preview", optional_params={}
)
assert len(result["requests"]) == 1
assert result["requests"][0]["content"]["parts"][0]["text"] == "hello"
def test_multiple_texts(self):
result = transform_openai_input_gemini_content(
input=["hello", "world"], model="gemini-embedding-2-preview", optional_params={}
)
assert len(result["requests"]) == 2
assert result["requests"][0]["content"]["parts"][0]["text"] == "hello"
assert result["requests"][1]["content"]["parts"][0]["text"] == "world"
def test_multimodal_inputs_are_separate_requests(self):
"""Key regression test for #24209: each input becomes its own request."""
result = transform_openai_input_gemini_content(
input=["The food was delicious", IMAGE_DATA_URI],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 2
# First request is text
assert result["requests"][0]["content"]["parts"][0]["text"] == "The food was delicious"
# Second request is image
assert result["requests"][1]["content"]["parts"][0]["inline_data"] is not None
def test_dimensions_mapped_to_output_dimensionality(self):
result = transform_openai_input_gemini_content(
input="hello",
model="gemini-embedding-2-preview",
optional_params={"dimensions": 256},
)
assert result["requests"][0]["outputDimensionality"] == 256
def test_model_name_prefixed(self):
result = transform_openai_input_gemini_content(
input="hello", model="gemini-embedding-2-preview", optional_params={}
)
assert result["requests"][0]["model"] == "models/gemini-embedding-2-preview"
def test_gcs_url_input(self):
result = transform_openai_input_gemini_content(
input=[GCS_URL], model="gemini-embedding-2-preview", optional_params={}
)
assert len(result["requests"]) == 1
assert result["requests"][0]["content"]["parts"][0]["file_data"] is not None
def test_mixed_text_image_gcs(self):
result = transform_openai_input_gemini_content(
input=["hello", IMAGE_DATA_URI, GCS_URL],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 3
def test_nested_input_combined_embedding(self):
"""Nested list produces one request with multiple parts (combined embedding)."""
result = transform_openai_input_gemini_content(
input=[["a red shoe", IMAGE_DATA_URI]],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 1
parts = result["requests"][0]["content"]["parts"]
assert len(parts) == 2
assert parts[0]["text"] == "a red shoe"
assert parts[1]["inline_data"] is not None
def test_mixed_nested_and_flat(self):
"""Mixed nested + flat produces correct number of requests."""
result = transform_openai_input_gemini_content(
input=[["text", IMAGE_DATA_URI], "standalone"],
model="gemini-embedding-2-preview",
optional_params={},
)
assert len(result["requests"]) == 2
# First: combined (2 parts)
assert len(result["requests"][0]["content"]["parts"]) == 2
# Second: standalone (1 part)
assert len(result["requests"][1]["content"]["parts"]) == 1
assert result["requests"][1]["content"]["parts"][0]["text"] == "standalone"
class TestTransformOpenaiInputGeminiEmbedContent:
"""Test transform_openai_input_gemini_embed_content (vertex_ai / embedContent path)."""
def test_text_and_image_combined(self):
result = transform_openai_input_gemini_embed_content(
input=["hello", IMAGE_DATA_URI],
model="gemini-embedding-2-preview",
optional_params={},
)
assert "content" in result
parts = result["content"]["parts"]
assert len(parts) == 2
assert parts[0]["text"] == "hello"
assert parts[1]["inline_data"] is not None
def test_gcs_url(self):
result = transform_openai_input_gemini_embed_content(
input=[GCS_URL],
model="gemini-embedding-2-preview",
optional_params={},
)
parts = result["content"]["parts"]
assert len(parts) == 1
assert parts[0]["file_data"]["file_uri"] == GCS_URL
def test_dimensions_mapped(self):
result = transform_openai_input_gemini_embed_content(
input="hello",
model="gemini-embedding-2-preview",
optional_params={"dimensions": 256},
)
assert result["outputDimensionality"] == 256
class TestProcessResponse:
"""Test that process_response sets correct indices."""
def test_single_embedding_index(self):
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}]
}
model_response = EmbeddingResponse()
result = process_response(
input="hello",
model_response=model_response,
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 1
assert result.data[0]["index"] == 0
def test_multiple_embeddings_have_correct_indices(self):
"""Regression test: indices should be 0, 1, 2... not all 0."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [
{"values": [0.1, 0.2]},
{"values": [0.3, 0.4]},
{"values": [0.5, 0.6]},
]
}
model_response = EmbeddingResponse()
result = process_response(
input=["a", "b", "c"],
model_response=model_response,
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 3
assert result.data[0]["index"] == 0
assert result.data[1]["index"] == 1
assert result.data[2]["index"] == 2
def test_multimodal_mixed_input(self):
"""process_response works with mixed text + multimodal inputs."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]
}
result = process_response(
input=["hello", IMAGE_DATA_URI],
model_response=EmbeddingResponse(),
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 2
assert result.data[0]["index"] == 0
assert result.data[1]["index"] == 1
# Should count tokens only for the text element, not the image
assert result.usage.prompt_tokens > 0
def test_nested_input_token_counting(self):
"""Nested list: only plain-text sub-elements should be counted."""
predictions: VertexAIBatchEmbeddingsResponseObject = {
"embeddings": [{"values": [0.1, 0.2]}]
}
result = process_response(
input=[["a red shoe", IMAGE_DATA_URI]],
model_response=EmbeddingResponse(),
model="gemini-embedding-2-preview",
_predictions=predictions,
)
assert len(result.data) == 1
assert result.usage.prompt_tokens > 0
def test_nested_empty_list_raises(self):
with pytest.raises(ValueError, match="must not be empty"):
transform_openai_input_gemini_content(
input=[[]],
model="gemini-embedding-2-preview",
optional_params={},
)
def test_nested_non_string_element_raises(self):
with pytest.raises(ValueError, match="must be strings"):
transform_openai_input_gemini_content(
input=[[["doubly", "nested"]]],
model="gemini-embedding-2-preview",
optional_params={},
)

View file

@ -2,6 +2,7 @@
Unit tests for auth_utils functions related to rate limiting and customer ID extraction.
"""
import base64
from typing import Optional
from unittest.mock import MagicMock, patch
@ -10,11 +11,12 @@ import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
_get_customer_id_from_standard_headers,
abbreviate_api_key,
check_complete_credentials,
get_end_user_id_from_request_body,
get_model_from_request,
get_key_model_rpm_limit,
get_key_model_tpm_limit,
get_model_from_request,
get_project_model_rpm_limit,
get_project_model_tpm_limit,
is_request_body_safe,
@ -258,6 +260,206 @@ def test_get_model_from_request_vertex_passthrough_still_works():
assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro"
def test_get_model_from_request_openai_deployment_route_still_works():
assert (
get_model_from_request(
request_data={},
route="/openai/deployments/my-azure-deployment/chat/completions",
)
== "my-azure-deployment"
)
def test_get_model_from_request_includes_file_endpoint_header_model():
assert (
get_model_from_request(
request_data={},
route="/v1/files",
request_headers={"X-LiteLLM-Model": "restricted-model"},
)
== "restricted-model"
)
def test_get_model_from_request_ignores_routing_header_on_standard_llm_routes():
assert (
get_model_from_request(
request_data={"model": "allowed-model"},
route="/v1/chat/completions",
request_headers={"x-litellm-model": "restricted-model"},
)
== "allowed-model"
)
def test_get_model_from_request_authorizes_all_file_routing_model_sources():
models = get_model_from_request(
request_data={"model": "body-model"},
route="/v1/files",
request_headers={"x-litellm-model": "header-model"},
request_query_params={"target_model_names": "query-model-a,query-model-b"},
)
assert isinstance(models, list)
assert set(models) == {
"body-model",
"query-model-a",
"query-model-b",
"header-model",
}
def test_get_model_from_request_extracts_simple_encoded_file_id_model():
from litellm.proxy.openai_files_endpoints.common_utils import (
encode_file_id_with_model,
)
file_id = encode_file_id_with_model(
file_id="file-provider-id",
model="restricted-model",
)
assert (
get_model_from_request(
request_data={"file_id": file_id},
route="/v1/files/{file_id}",
)
== "restricted-model"
)
def test_get_model_from_request_extracts_unified_file_id_models():
raw_unified_file_id = (
"litellm_proxy:application/octet-stream;unified_id,test-id;"
"target_model_names,model-a,model-b;llm_output_file_id,file-provider-id"
)
encoded_unified_file_id = (
base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=")
)
assert get_model_from_request(
request_data={"file_id": encoded_unified_file_id},
route="/v1/files/{file_id}",
) == ["model-a", "model-b"]
def test_get_model_from_request_extracts_eval_completion_model():
assert (
get_model_from_request(
request_data={"completion": {"model": "judge-model"}},
route="/v1/evals/{eval_id}/runs",
)
== "judge-model"
)
def test_get_model_from_request_includes_fine_tuning_target_model_query():
assert (
get_model_from_request(
request_data={},
route="/v1/fine_tuning/jobs",
request_query_params={"target_model_names": "fine-tune-model"},
)
== "fine-tune-model"
)
def test_get_model_from_request_extracts_video_id_model():
from litellm.types.videos.utils import encode_video_id_with_provider
video_id = encode_video_id_with_provider(
video_id="video-provider-id",
provider="openai",
model_id="video-model",
)
assert (
get_model_from_request(
request_data={"video_id": video_id},
route="/v1/videos/{video_id}",
)
== "video-model"
)
def test_get_model_from_request_only_runs_media_decoders_for_matching_fields():
with (
patch(
"litellm.types.videos.utils.decode_video_id_with_provider",
return_value={"model_id": "video-model"},
) as video_decoder,
patch(
"litellm.types.videos.utils.decode_character_id_with_provider",
return_value={"model_id": "character-model"},
) as character_decoder,
):
assert (
get_model_from_request(
request_data={"file_id": "file-provider-id"},
route="/v1/files/{file_id}",
)
is None
)
video_decoder.assert_not_called()
character_decoder.assert_not_called()
assert (
get_model_from_request(
request_data={"video_id": "video-provider-id"},
route="/v1/videos/{video_id}",
)
== "video-model"
)
video_decoder.assert_called_once_with("video-provider-id")
character_decoder.assert_not_called()
video_decoder.reset_mock()
character_decoder.reset_mock()
assert (
get_model_from_request(
request_data={"character_id": "character-provider-id"},
route="/v1/videos/{character_id}",
)
== "character-model"
)
video_decoder.assert_not_called()
character_decoder.assert_called_once_with("character-provider-id")
def test_get_model_from_request_handles_managed_id_decoder_failures():
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils.decode_model_from_file_id",
side_effect=Exception("decode failed"),
),
patch(
"litellm.llms.base_llm.managed_resources.utils.parse_unified_id",
side_effect=Exception("parse failed"),
),
patch(
"litellm.types.videos.utils.decode_video_id_with_provider",
side_effect=Exception("video decode failed"),
),
):
assert (
get_model_from_request(
request_data={"file_id": "not-a-managed-resource-id"},
route="/v1/files/{file_id}",
)
is None
)
assert (
get_model_from_request(
request_data={"video_id": "not-a-managed-resource-id"},
route="/v1/videos/{video_id}",
)
is None
)
def test_abbreviate_api_key():
assert abbreviate_api_key("sk-test-1234") == "sk-...1234"
def test_get_customer_user_header_returns_none_when_no_customer_role():
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping
@ -964,3 +1166,129 @@ class TestIsRequestBodySafeBlocksEndpointTargetingFields:
)
is True
)
# ── is_request_body_safe nested-config recursion (VERIA-6) ────────────────────
class TestIsRequestBodySafeNestedConfig:
"""The Milvus vector store transformer unpacks
``litellm_embedding_config`` as ``**kwargs`` into ``litellm.embedding(...)``
— same SSRF / credential-exfil surface as a top-level ``api_base`` in
the request body. ``is_request_body_safe`` must recurse into this
nested dict so a banned param can't be smuggled in via nesting."""
def test_root_level_api_base_blocked_when_no_opt_in(self):
"""Sanity check: pre-existing root-level enforcement still works."""
with pytest.raises(ValueError, match="api_base"):
is_request_body_safe(
request_body={"api_base": "https://attacker.example.com"},
general_settings={},
llm_router=None,
model="gpt-4",
)
def test_nested_api_base_in_embedding_config_blocked(self):
"""Smuggling ``api_base`` inside ``litellm_embedding_config`` is
the VERIA-6 bypass — must be blocked by the recursive check."""
with pytest.raises(ValueError, match="api_base"):
is_request_body_safe(
request_body={
"litellm_embedding_config": {
"api_base": "https://attacker.example.com",
"api_key": "leaked-key",
}
},
general_settings={},
llm_router=None,
model="milvus-store",
)
def test_nested_langfuse_host_in_embedding_config_blocked(self):
"""The recursion uses the *full* banned-param list, not a special
subset — so any flag that's banned at the root is also banned
when nested."""
with pytest.raises(ValueError, match="langfuse_host"):
is_request_body_safe(
request_body={
"litellm_embedding_config": {
"langfuse_host": "https://attacker.example.com"
}
},
general_settings={},
llm_router=None,
model="milvus-store",
)
def test_nested_api_base_allowed_when_admin_opts_in(self):
"""Admins who explicitly enable client-side credential passthrough
keep the existing escape hatch — same UX as for root-level."""
assert (
is_request_body_safe(
request_body={
"litellm_embedding_config": {
"api_base": "https://my-azure.example.com"
}
},
general_settings={"allow_client_side_credentials": True},
llm_router=None,
model="milvus-store",
)
is True
)
def test_safe_nested_config_accepted(self):
"""A nested config without any banned params passes — there's no
false-positive on legitimate ``api_version`` / model params."""
assert (
is_request_body_safe(
request_body={
"litellm_embedding_config": {
"api_version": "2024-02-15-preview",
}
},
general_settings={},
llm_router=None,
model="milvus-store",
)
is True
)
def test_non_dict_nested_config_does_not_break_check(self):
"""A bogus type for ``litellm_embedding_config`` (string, list,
None) must not crash the validator — it should just fall through."""
assert (
is_request_body_safe(
request_body={"litellm_embedding_config": "not-a-dict"},
general_settings={},
llm_router=None,
model="x",
)
is True
)
def test_deeply_nested_config_does_not_recurse(self):
"""Greptile P1: ``is_request_body_safe`` is iterative single-level —
a deeply-nested ``litellm_embedding_config`` cannot exhaust the
Python call stack to trigger a 500 ``RecursionError``. Build a
body 1000 levels deep; the validator must complete in O(1)
descent."""
body = {"litellm_embedding_config": {}}
cur = body["litellm_embedding_config"]
for _ in range(1000):
cur["litellm_embedding_config"] = {}
cur = cur["litellm_embedding_config"]
# Banned param at the deepest level shouldn't be reached — single
# level only.
cur["api_base"] = "https://attacker.example.com"
# No exception raised: deeper levels aren't checked.
assert (
is_request_body_safe(
request_body=body,
general_settings={},
llm_router=None,
model="x",
)
is True
)

View file

@ -1,6 +1,7 @@
import json
import os
import sys
from types import SimpleNamespace
from unittest.mock import ANY, AsyncMock, MagicMock, patch
sys.path.insert(
@ -34,6 +35,13 @@ from litellm.proxy.auth.user_api_key_auth import (
)
class _RoutingRequest:
def __init__(self, headers=None, query_params=None):
self.headers = headers or {}
self.query_params = query_params or {}
self.state = SimpleNamespace()
def test_get_api_key():
bearer_token = "Bearer sk-12345678"
api_key = "sk-12345678"
@ -177,6 +185,39 @@ async def test_custom_auth_honors_key_level_model_access_restriction_allowed_wit
)
@pytest.mark.asyncio
async def test_custom_auth_enforces_key_model_access_from_file_route_header_with_opt_in():
valid_token = UserAPIKeyAuth(token="test_token", models=["allowed-model"])
request = _RoutingRequest(headers={"x-litellm-model": "restricted-model"})
with (
patch(
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
new_callable=AsyncMock,
) as mock_can_key,
patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
),
patch(
"litellm.proxy.proxy_server.general_settings",
{"custom_auth_run_common_checks": True},
),
):
await _run_post_custom_auth_checks(
valid_token=valid_token,
request=request,
request_data={},
route="/v1/files",
parent_otel_span=None,
)
mock_can_key.assert_awaited_once_with(
model="restricted-model",
llm_model_list=ANY,
valid_token=valid_token,
llm_router=ANY,
)
@pytest.mark.asyncio
async def test_custom_auth_honors_key_level_model_access_restriction_denied_with_opt_in():
valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"])

View file

@ -231,6 +231,50 @@ class TestTokenUtilities:
result = get_stored_api_key()
assert result is None
def test_get_stored_api_key_base_url_match(self):
"""Stored key is returned when expected_base_url matches stored origin"""
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
return_value=token_data,
):
assert (
get_stored_api_key(expected_base_url="https://real-proxy.com")
== "sk-prod"
)
def test_get_stored_api_key_base_url_match_trailing_slash(self):
"""Trailing slash on expected_base_url is normalised before comparison"""
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
return_value=token_data,
):
assert (
get_stored_api_key(expected_base_url="https://real-proxy.com/")
== "sk-prod"
)
def test_get_stored_api_key_base_url_mismatch(self):
"""Stored key is NOT returned when expected_base_url differs from stored origin"""
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
return_value=token_data,
):
assert get_stored_api_key(expected_base_url="https://evil.com") is None
def test_get_stored_api_key_old_token_no_base_url(self):
"""Old tokens without a base_url field are rejected when origin check is requested"""
token_data = {"key": "sk-old-token"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
return_value=token_data,
):
assert (
get_stored_api_key(expected_base_url="https://real-proxy.com") is None
)
class TestLoginCommand:
"""Test login CLI command"""

View file

@ -220,6 +220,27 @@ class TestToolPermissionGuardrail:
assert tool_calls[0].id == "call_123"
assert tool_calls[0].function.name == "Read"
def test_extract_tool_calls_legacy_function_call_format(self):
response = ModelResponse(
choices=[
Choices(
message={
"function_call": {
"name": "Read",
"arguments": '{"file_path": "/test/file.txt"}',
},
}
)
]
)
tool_calls = self.guardrail._extract_tool_calls_from_response(response)
assert len(tool_calls) == 1
assert isinstance(tool_calls[0], ChatCompletionMessageToolCall)
assert tool_calls[0].id == "legacy_function_call_0"
assert tool_calls[0].function.name == "Read"
assert tool_calls[0].function.arguments == '{"file_path": "/test/file.txt"}'
def test_extract_tool_calls_empty_response(self):
response = ModelResponse(choices=[])
tool_calls = self.guardrail._extract_tool_calls_from_response(response)
@ -271,6 +292,31 @@ class TestToolPermissionGuardrail:
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
async def test_async_post_call_success_hook_with_denied_legacy_function_call_raises(
self,
):
response = ModelResponse(
choices=[
Choices(
message={
"function_call": {
"name": "Read",
"arguments": "{}",
},
}
)
]
)
user_api_key_dict = UserAPIKeyAuth()
data = {"guardrails": ["test-tool-permission"]}
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(GuardrailRaisedException):
await self.guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
async def test_async_post_call_success_hook_param_patterns_allow(self):
guardrail = ToolPermissionGuardrail(
@ -379,7 +425,9 @@ class TestToolPermissionGuardrail:
assert "berri" in choice.message.content
@pytest.mark.asyncio
async def test_async_post_call_success_hook_missing_arguments_default_allows(self):
async def test_async_post_call_success_hook_missing_arguments_blocks_param_rule(
self,
):
guardrail = ToolPermissionGuardrail(
guardrail_name="mail-guardrail",
rules=[
@ -405,9 +453,52 @@ class TestToolPermissionGuardrail:
data = {"guardrails": ["mail-guardrail"]}
with patch.object(guardrail, "should_run_guardrail", return_value=True):
await guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
with pytest.raises(GuardrailRaisedException):
await guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"arguments",
[
"{not-json",
'["owner@berri.ai"]',
],
)
async def test_async_post_call_success_hook_malformed_arguments_blocks_param_rule(
self, arguments
):
guardrail = ToolPermissionGuardrail(
guardrail_name="mail-guardrail",
rules=[
{
"id": "deny_gmail",
"tool_name": r"^mail_mcp-send_email$",
"decision": "deny",
"allowed_param_patterns": {"to[]": r"^.+@gmail\.com$"},
}
],
default_action="allow",
on_disallowed_action="block",
)
tool_call = {
"function": {
"name": "mail_mcp-send_email",
"arguments": arguments,
},
"type": "function",
}
response = ModelResponse(choices=[Choices(message={"tool_calls": [tool_call]})])
user_api_key_dict = UserAPIKeyAuth()
data = {"guardrails": ["mail-guardrail"]}
with patch.object(guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(GuardrailRaisedException):
await guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
async def test_async_pre_call_hook_block_mode(self):
@ -430,6 +521,65 @@ class TestToolPermissionGuardrail:
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_blocks_legacy_functions(self):
data = {
"functions": [
{"name": "Bash", "description": "allowed"},
{"name": "Read", "description": "denied"},
]
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(HTTPException) as excinfo:
await self.guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_blocks_named_legacy_function_call(self):
data = {
"functions": [{"name": "Bash"}],
"function_call": {"name": "Read"},
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(HTTPException) as excinfo:
await self.guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_blocks_named_tool_choice(self):
data = {
"tools": [{"type": "function", "function": {"name": "Bash"}}],
"tool_choice": {"type": "function", "function": {"name": "Read"}},
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(HTTPException) as excinfo:
await self.guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_uses_custom_template(self):
guardrail = ToolPermissionGuardrail(
@ -491,6 +641,41 @@ class TestToolPermissionGuardrail:
assert "Bash" in tool_names
assert "Read" not in tool_names
@pytest.mark.asyncio
async def test_async_pre_call_hook_rewrite_mode_filters_legacy_functions(self):
guardrail = ToolPermissionGuardrail(
guardrail_name="test-tool-permission",
rules=self.test_rules,
default_action="deny",
on_disallowed_action="rewrite",
)
data = {
"functions": [
{"name": "Bash", "description": "allowed"},
{"name": "Read", "description": "denied"},
],
"function_call": {"name": "Read"},
"tools": [
{"type": "function", "function": {"name": "Bash"}},
],
"tool_choice": {"type": "function", "function": {"name": "Read"}},
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(guardrail, "should_run_guardrail", return_value=True):
new_data = await guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert isinstance(new_data, dict)
assert [function["name"] for function in new_data["functions"]] == ["Bash"]
assert new_data["function_call"] == "none"
assert new_data["tool_choice"] == "none"
def test_modify_response_with_permission_errors(self):
# Setup a response with one tool_call
tool_call = ChatCompletionMessageToolCall(
@ -522,6 +707,40 @@ class TestToolPermissionGuardrail:
assert isinstance(choice.message.content, str)
assert "Permission denied" in choice.message.content
def test_modify_response_with_permission_errors_filters_legacy_function_call(self):
response = ModelResponse(
choices=[
Choices(
message={
"function_call": {
"name": "Read",
"arguments": "{}",
},
"content": "",
}
)
]
)
tool_call = self.guardrail._extract_tool_calls_from_response(response)[0]
denied_tools = [
(
tool_call,
PermissionError(
tool_name="Read",
rule_id="deny_read",
message="Tool 'Read' denied by rule 'deny_read'",
),
)
]
self.guardrail._modify_response_with_permission_errors(response, denied_tools)
choice = response.choices[0]
assert isinstance(choice, Choices)
assert choice.message.function_call is None
assert isinstance(choice.message.content, str)
assert "Permission denied" in choice.message.content
class TestToolPermissionGuardrailIntegration:
"""Integration tests for Tool Permission Guardrail"""

View file

@ -778,3 +778,373 @@ def test_get_callback_identifier_custom_logger_registry_and_fallback():
result = get_callback_identifier(my_callback_function)
# Should fall back to callback_name() which returns __name__
assert result == "my_callback_function"
# ---------------------------------------------------------------------------
# /health response shape: model-access scoping and display-field allowlist
# ---------------------------------------------------------------------------
# These tests pin the contract that the /health response (a) only includes
# deployments the calling key is allowed to see, and (b) does not return
# provider routing fields like api_base / api_version. They guard against
# regressions that would widen the response shape.
@pytest.mark.asyncio
async def test_health_endpoint_filters_model_list_by_user_access():
"""
health_endpoint() should restrict _llm_model_list to deployments whose
model_name appears in user_api_key_dict.models before running the health
check. A key scoped to ["model-a"] should only see model-a in the result,
not other deployments configured on the proxy.
"""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
full_model_list = [
{
"model_name": "model-a",
"litellm_params": {
"model": "openai/gpt-4o",
"api_base": "https://example-a.test",
},
"model_info": {"id": "id-a"},
},
{
"model_name": "model-b",
"litellm_params": {
"model": "openai/gpt-4o",
"api_base": "https://example-b.test",
"api_version": "2024-10-21",
},
"model_info": {"id": "id-b"},
},
]
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-test-key",
models=["model-a"],
)
captured: dict = {}
async def fake_perform(**kwargs):
captured["model_list"] = kwargs["model_list"]
return {
"healthy_endpoints": [],
"unhealthy_endpoints": [],
"healthy_count": 0,
"unhealthy_count": 0,
}
with (
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.use_background_health_checks", False),
patch("litellm.proxy.proxy_server.user_model", None),
patch("litellm.proxy.proxy_server.health_check_results", {}),
patch("litellm.proxy.proxy_server.health_check_details", True),
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
patch(
"litellm.proxy.health_endpoints._health_endpoints._perform_health_check_and_save",
side_effect=fake_perform,
),
):
from fastapi import Response
await health_endpoint(response=Response(), user_api_key_dict=user_api_key_dict)
assert (
"model_list" in captured
), "health_endpoint did not call _perform_health_check_and_save"
returned_names = {m["model_name"] for m in captured["model_list"]}
assert returned_names == {
"model-a"
}, f"health_endpoint did not scope model_list to caller access: {returned_names}"
@pytest.mark.asyncio
async def test_health_endpoint_filters_background_cache_by_user_access():
"""
When background_health_checks is enabled, health_endpoint() should also
scope the cached result to the caller's allowed models rather than
returning the cache verbatim.
"""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
full_model_list = [
{
"model_name": "model-a",
"litellm_params": {
"model": "openai/gpt-4o",
"api_base": "https://example-a.test",
},
"model_info": {"id": "id-a"},
},
{
"model_name": "model-b",
"litellm_params": {
"model": "openai/gpt-4o",
"api_base": "https://example-b.test",
},
"model_info": {"id": "id-b"},
},
]
cached_results = {
"healthy_endpoints": [
{
"model": "openai/gpt-4o",
"model_id": "id-a",
"api_base": "https://example-a.test",
},
{
"model": "openai/gpt-4o",
"model_id": "id-b",
"api_base": "https://example-b.test",
},
],
"unhealthy_endpoints": [],
"healthy_count": 2,
"unhealthy_count": 0,
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-test-key",
models=["model-a"],
)
with (
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.use_background_health_checks", True),
patch("litellm.proxy.proxy_server.user_model", None),
patch("litellm.proxy.proxy_server.health_check_results", cached_results),
patch("litellm.proxy.proxy_server.health_check_details", True),
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
):
from fastapi import Response
result = await health_endpoint(
response=Response(), user_api_key_dict=user_api_key_dict
)
# Sanity: the source cache had two entries before scoping; the scoping
# step is what reduces it to one. (This guards against the test passing
# vacuously when the cache filter drops everything because cached
# entries lack the model_id key — both entries carry model_id above.)
assert len(cached_results["healthy_endpoints"]) == 2
assert all(
ep.get("model_id") for ep in cached_results["healthy_endpoints"]
), "test fixture invariant: every cached entry must carry a model_id"
# The non-admin caller must not see api_base on the returned cache entries.
returned = result.get("healthy_endpoints", [])
assert (
len(returned) == 1
), f"expected exactly one cached entry after scoping, got {len(returned)}"
assert returned[0]["model_id"] == "id-a"
assert "api_base" not in returned[0]
assert result["healthy_count"] == 1
assert result["unhealthy_count"] == 0
@pytest.mark.asyncio
async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not():
"""
A proxy admin should still see ``api_base`` and ``api_version`` in the
/health response so they can tell which Vertex region / Azure resource
+ API version is healthy. A non-admin caller must not — both fields
should be stripped, and the response should carry a notice header so
non-admin clients can detect the change programmatically.
"""
from fastapi import Response
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
full_model_list = [
{
"model_name": "model-a",
"litellm_params": {
"model": "openai/gpt-4o",
"api_base": "https://example-a.test",
},
"model_info": {"id": "id-a"},
},
]
cached_results = {
"healthy_endpoints": [
{
"model": "openai/gpt-4o",
"model_id": "id-a",
"api_base": "https://us-central1-aiplatform.googleapis.com/v1/projects/p",
"api_version": "2024-10-21",
},
],
"unhealthy_endpoints": [],
"healthy_count": 1,
"unhealthy_count": 0,
}
admin_key = UserAPIKeyAuth(
api_key="hashed-admin-key",
models=["model-a"],
user_role=LitellmUserRoles.PROXY_ADMIN,
)
non_admin_key = UserAPIKeyAuth(
api_key="hashed-user-key",
models=["model-a"],
)
common_patches = [
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.use_background_health_checks", True),
patch("litellm.proxy.proxy_server.user_model", None),
patch("litellm.proxy.proxy_server.health_check_results", cached_results),
patch("litellm.proxy.proxy_server.health_check_details", True),
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
]
for p in common_patches:
p.start()
try:
admin_response = Response()
non_admin_response = Response()
admin_result = await health_endpoint(
response=admin_response, user_api_key_dict=admin_key
)
non_admin_result = await health_endpoint(
response=non_admin_response, user_api_key_dict=non_admin_key
)
finally:
for p in common_patches:
p.stop()
admin_eps = admin_result.get("healthy_endpoints", [])
non_admin_eps = non_admin_result.get("healthy_endpoints", [])
assert len(admin_eps) == 1
assert (
admin_eps[0]["api_base"]
== "https://us-central1-aiplatform.googleapis.com/v1/projects/p"
), "admin must see the full api_base so they can identify the region"
assert (
admin_eps[0]["api_version"] == "2024-10-21"
), "admin must see api_version so they can distinguish provider deployments"
assert len(non_admin_eps) == 1
assert "api_base" not in non_admin_eps[0]
assert "api_version" not in non_admin_eps[0]
# Non-admin response must advertise that api_base/api_version were
# withheld so clients that previously parsed them can detect the change.
assert (
non_admin_response.headers.get("Litellm-Health-Field-Notice")
== "api_base and api_version are admin-only on this endpoint"
)
assert "Litellm-Health-Field-Notice" not in admin_response.headers
# Stripping must produce a copy — the shared cache must still carry the
# routing fields so the next admin caller can read them.
cached_first = cached_results["healthy_endpoints"][0]
assert (
cached_first["api_base"]
== "https://us-central1-aiplatform.googleapis.com/v1/projects/p"
)
assert cached_first["api_version"] == "2024-10-21"
@pytest.mark.asyncio
async def test_health_endpoint_warns_when_scoped_models_lack_model_id():
"""
When a scoped key's accessible models exist on the proxy but none of the
matching deployments expose a ``model_info.id``, the cache filter drops
everything. The response should include a structured ``warnings`` field
so the caller can distinguish "no deployments configured" from
"deployments excluded due to missing model_info.id".
"""
from fastapi import Response
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
full_model_list = [
{
"model_name": "model-a",
"litellm_params": {
"model": "openai/gpt-4o",
"api_base": "https://example-a.test",
},
# Intentionally no model_info.id — this is the misconfiguration
# the warnings field is meant to flag.
"model_info": {},
},
]
cached_results = {
"healthy_endpoints": [
{
"model": "openai/gpt-4o",
"model_id": "id-a",
"api_base": "https://example-a.test",
},
],
"unhealthy_endpoints": [],
"healthy_count": 1,
"unhealthy_count": 0,
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-user-key",
models=["model-a"],
)
with (
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch("litellm.proxy.proxy_server.use_background_health_checks", True),
patch("litellm.proxy.proxy_server.user_model", None),
patch("litellm.proxy.proxy_server.health_check_results", cached_results),
patch("litellm.proxy.proxy_server.health_check_details", True),
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
):
result = await health_endpoint(
response=Response(), user_api_key_dict=user_api_key_dict
)
assert result["healthy_count"] == 0
assert result["unhealthy_count"] == 0
assert "warnings" in result, (
"empty cache result must surface a warnings field so the caller "
"can distinguish 'no deployments' from 'deployments excluded'"
)
assert any("model_info.id" in w for w in result["warnings"])
def test_clean_endpoint_data_strips_credentials_keeps_routing_fields():
"""
_clean_endpoint_data() drops credentials but leaves api_base /
api_version intact — the per-caller hide/show happens in the endpoint
layer based on user role, not in the cleaning helper. This guarantees
proxy admins continue to see those fields in the /health response.
"""
from litellm.proxy.health_check import _clean_endpoint_data
raw = {
"model": "openai/gpt-4o",
"api_key": "sk-test",
"api_base": "https://example.test/v1",
"api_version": "2024-10-21",
"aws_access_key_id": "AKIAEXAMPLE",
}
cleaned = _clean_endpoint_data(raw, details=True)
assert "api_key" not in cleaned
assert "aws_access_key_id" not in cleaned
assert cleaned.get("api_base") == "https://example.test/v1"
assert cleaned.get("api_version") == "2024-10-21"

View file

@ -0,0 +1,118 @@
import sys
from types import ModuleType, SimpleNamespace
from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids
def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch):
from litellm.proxy import _lazy_openapi_snapshot
route_a = SimpleNamespace(path="/feature-a/items")
route_b = SimpleNamespace(path="/feature-b/items")
fake_app = SimpleNamespace(
title="LiteLLM test",
version="0.0.0",
routes=[route_a, route_b],
)
fake_feature_a_module = ModuleType("fake_feature_a")
fake_feature_b_module = ModuleType("fake_feature_b")
monkeypatch.setitem(sys.modules, "fake_feature_a", fake_feature_a_module)
monkeypatch.setitem(sys.modules, "fake_feature_b", fake_feature_b_module)
fake_lazy_features_module = ModuleType("litellm.proxy._lazy_features")
fake_lazy_features_module.LAZY_FEATURES = [
SimpleNamespace(
name="feature-a",
module_path="fake_feature_a",
path_prefixes=("/feature-a",),
register_fn=lambda app, module: None,
),
SimpleNamespace(
name="feature-b",
module_path="fake_feature_b",
path_prefixes=("/feature-b",),
register_fn=lambda app, module: None,
),
]
monkeypatch.setitem(
sys.modules, "litellm.proxy._lazy_features", fake_lazy_features_module
)
def fake_get_openapi(title, version, routes):
path = routes[0].path
return {
"paths": {path: {"get": {"operationId": "shared_operation_id_get"}}},
"components": {"schemas": {"Example": {"type": "object"}}},
}
def fake_ensure_unique_openapi_operation_ids(schema, reserved_operation_ids):
for path_item in schema["paths"].values():
operation = path_item["get"]
operation_id = operation["operationId"]
if operation_id in reserved_operation_ids:
operation_id = f"{operation_id}_2"
operation["operationId"] = operation_id
reserved_operation_ids.add(operation_id)
return schema
fake_proxy_server_module = ModuleType("litellm.proxy.proxy_server")
fake_proxy_server_module.app = fake_app
fake_proxy_server_module.ensure_unique_openapi_operation_ids = (
fake_ensure_unique_openapi_operation_ids
)
monkeypatch.setitem(
sys.modules, "litellm.proxy.proxy_server", fake_proxy_server_module
)
monkeypatch.setattr("fastapi.openapi.utils.get_openapi", fake_get_openapi)
fragments = _lazy_openapi_snapshot.generate_snapshot()
assert (
fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["operationId"]
== "shared_operation_id_get"
)
assert (
fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["operationId"]
== "shared_operation_id_get_2"
)
assert fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["tags"] == [
"feature-a"
]
assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["tags"] == [
"feature-b"
]
def test_normalize_operation_ids_uses_each_http_method():
paths = {
"/proxy/{endpoint}": {
"delete": {"operationId": "proxy_route_proxy__endpoint__put"},
"get": {"operationId": "proxy_route_proxy__endpoint__put"},
"post": {"operationId": "proxy_route_proxy__endpoint__put"},
"put": {"operationId": "proxy_route_proxy__endpoint__put"},
}
}
_normalize_operation_ids(paths)
operations = paths["/proxy/{endpoint}"]
assert operations["delete"]["operationId"] == "proxy_route_proxy__endpoint__delete"
assert operations["get"]["operationId"] == "proxy_route_proxy__endpoint__get"
assert operations["post"]["operationId"] == "proxy_route_proxy__endpoint__post"
assert operations["put"]["operationId"] == "proxy_route_proxy__endpoint__put"
def test_normalize_operation_ids_preserves_custom_ids():
paths = {
"/proxy/{endpoint}": {
"get": {"operationId": "custom_operation"},
"post": {"operationId": "custom_operation"},
}
}
_normalize_operation_ids(paths)
operations = paths["/proxy/{endpoint}"]
assert operations["get"]["operationId"] == "custom_operation"
assert operations["post"]["operationId"] == "custom_operation"

View file

@ -5988,6 +5988,202 @@ async def test_reseed_warms_cache_even_on_zero_db_spend():
ps.prisma_client = orig_prisma
# -----------------------------------------------------------------------------
# /config/update — critical paths only.
#
# These exercise the four behaviors that broke or changed in the rewrite of
# update_config (litellm/proxy/proxy_server.py): targeted per-section writes,
# the removal of the store_model_in_db gate, env var encryption, and the
# success_callback / litellm_settings merge semantics. All other branches
# (auth, missing-DB, slack auto-enable, router_settings merge) are covered
# implicitly or by upstream tests.
# -----------------------------------------------------------------------------
class _FakeRow:
def __init__(self, param_name, param_value):
self.param_name = param_name
self.param_value = param_value
class _FakeLitellmConfig:
def __init__(self, initial_rows=None):
self.rows = dict(initial_rows or {})
self.upsert_calls: list = []
self.find_first = AsyncMock(side_effect=self._find_first)
self.upsert = AsyncMock(side_effect=self._upsert)
async def _find_first(self, where=None):
if where and "param_name" in where:
name = where["param_name"]
if name in self.rows:
return _FakeRow(name, self.rows[name])
return None
async def _upsert(self, where=None, data=None):
name = where["param_name"]
raw = data["update"]["param_value"]
value = json.loads(raw) if isinstance(raw, str) else raw
self.rows[name] = value
self.upsert_calls.append((name, value))
class _FakePrismaClient:
def __init__(self, initial_rows=None):
self.db = mock.MagicMock()
self.db.litellm_config = _FakeLitellmConfig(initial_rows=initial_rows)
self.jsonify_object = lambda obj: obj
@pytest.fixture
def _update_config_setup(monkeypatch):
"""Install fakes for the /config/update endpoint and return (client, prisma)."""
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth as auth_dep
def _install(initial_rows=None, store_model_in_db=True):
prisma = _FakePrismaClient(initial_rows=initial_rows)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
monkeypatch.setattr(
"litellm.proxy.proxy_server.store_model_in_db", store_model_in_db
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.encrypt_value_helper",
lambda value, **_: f"enc:{value}",
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.invalidate_config_param",
AsyncMock(return_value=None),
)
from litellm.proxy.proxy_server import proxy_config as real_proxy_config
monkeypatch.setattr(
real_proxy_config, "add_deployment", AsyncMock(return_value=None)
)
original_overrides = app.dependency_overrides.copy()
app.dependency_overrides[auth_dep] = lambda: UserAPIKeyAuth(
user_id="test_admin",
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
)
client = TestClient(app)
def _restore():
app.dependency_overrides = original_overrides
return client, prisma, _restore
return _install
def test_update_config_writes_only_sent_section(_update_config_setup):
"""A request that only touches general_settings must not write any other
section row, and must leave previously-written rows byte-identical."""
client, prisma, restore = _update_config_setup(
initial_rows={
"litellm_settings": {"drop_params": True},
"environment_variables": {"FOO": "enc:bar"},
}
)
try:
resp = client.post(
"/config/update",
json={"general_settings": {"store_prompts_in_spend_logs": True}},
)
assert resp.status_code == 200
written = {name for name, _ in prisma.db.litellm_config.upsert_calls}
assert written == {"general_settings"}
assert prisma.db.litellm_config.rows["litellm_settings"] == {
"drop_params": True
}
assert prisma.db.litellm_config.rows["environment_variables"] == {
"FOO": "enc:bar"
}
finally:
restore()
def test_update_config_can_flip_store_model_in_db_when_currently_false(
_update_config_setup,
):
"""The endpoint used to refuse all writes when store_model_in_db was
False, blocking the very request that would flip it to True."""
client, prisma, restore = _update_config_setup(store_model_in_db=False)
try:
resp = client.post(
"/config/update", json={"general_settings": {"store_model_in_db": True}}
)
assert resp.status_code == 200
assert (
prisma.db.litellm_config.rows["general_settings"]["store_model_in_db"]
is True
)
finally:
restore()
def test_update_config_environment_variables_encrypted_before_write(
_update_config_setup,
):
"""env var values must be encrypted before they hit the DB row."""
client, prisma, restore = _update_config_setup()
try:
resp = client.post(
"/config/update",
json={"environment_variables": {"OPENAI_API_KEY": "sk-secret"}},
)
assert resp.status_code == 200
stored = prisma.db.litellm_config.rows["environment_variables"]
assert stored == {"OPENAI_API_KEY": "enc:sk-secret"}
finally:
restore()
def test_update_config_litellm_settings_request_wins_for_non_callback_keys(
_update_config_setup,
):
"""Sending {"drop_params": False} when the row holds drop_params: True
must persist False (request wins). Untouched keys preserved."""
client, prisma, restore = _update_config_setup(
initial_rows={
"litellm_settings": {"drop_params": True, "set_verbose": True},
}
)
try:
resp = client.post(
"/config/update", json={"litellm_settings": {"drop_params": False}}
)
assert resp.status_code == 200
stored = prisma.db.litellm_config.rows["litellm_settings"]
assert stored["drop_params"] is False
assert stored["set_verbose"] is True
finally:
restore()
def test_update_config_success_callback_normalizes_existing_mixed_case(
_update_config_setup,
):
"""Existing mixed-case callback names (written elsewhere) must be
normalized to lowercase before union, otherwise the union dedup misses
against the lowercase incoming entry and delete_callback (lowercase
lookup) cannot find the original."""
client, prisma, restore = _update_config_setup(
initial_rows={"litellm_settings": {"success_callback": ["Langfuse", "SQS"]}}
)
try:
resp = client.post(
"/config/update",
json={"litellm_settings": {"success_callback": ["langfuse"]}},
)
assert resp.status_code == 200
stored = prisma.db.litellm_config.rows["litellm_settings"]["success_callback"]
assert set(stored) == {"langfuse", "sqs"}
finally:
restore()
# ---------------------------------------------------------------------------
# Lazy feature loading (LazyFeatureMiddleware) — verifies that optional
# routers are NOT imported at module load and ARE imported on first request
@ -5996,9 +6192,6 @@ async def test_reseed_warms_cache_even_on_zero_db_spend():
# ---------------------------------------------------------------------------
import sys
class TestLazyFeatureRegistry:
"""Sanity checks on the registry shape — guards against accidental edits."""

View file

@ -156,8 +156,11 @@ async def test_update_request_data_with_litellm_managed_vector_store_registry():
vector_store_id="test_store_id"
)
# Test with no vector store registry
with patch.object(litellm, "vector_store_registry", None):
# Test with no vector store registry or DB fallback
with (
patch.object(litellm, "vector_store_registry", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
original_data = {"existing_key": "existing_value"}
result = await _update_request_data_with_litellm_managed_vector_store_registry(
data=original_data, vector_store_id=vector_store_id

View file

@ -0,0 +1,540 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException, Request, Response
import litellm
from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth
def _mock_request() -> MagicMock:
request = MagicMock(spec=Request)
request.headers = {}
request.method = "POST"
request.query_params = {}
request.url.path = "/v1/vector_stores/vs_path/search"
return request
@pytest.mark.asyncio
async def test_vector_store_search_forces_path_id_over_body_id():
from litellm.proxy.vector_store_endpoints.endpoints import vector_store_search
captured_data = {}
async def fake_base_process(self, **kwargs):
captured_data.update(self.data)
return {"ok": True}
request = _mock_request()
with (
patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(
return_value={
"vector_store_id": "vs_body_victim",
"query": "test",
}
),
),
patch.object(litellm, "vector_store_registry", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch(
"litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=fake_base_process,
),
):
response = await vector_store_search(
request=request,
vector_store_id="vs_path_allowed",
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert response == {"ok": True}
assert captured_data["vector_store_id"] == "vs_path_allowed"
@pytest.mark.asyncio
async def test_vector_store_file_create_forces_path_id_over_body_id():
from litellm.proxy.vector_store_files_endpoints.endpoints import (
vector_store_file_create,
)
captured_data = {}
async def fake_base_process(self, **kwargs):
captured_data.update(self.data)
return {"ok": True}
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_path_allowed",
"custom_llm_provider": "openai",
"team_id": "team-a",
}
request = _mock_request()
with (
patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(
return_value={
"vector_store_id": "vs_body_victim",
"file_id": "file_123",
}
),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=fake_base_process,
),
):
response = await vector_store_file_create(
vector_store_id="vs_path_allowed",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert response == {"ok": True}
assert captured_data["vector_store_id"] == "vs_path_allowed"
assert captured_data["custom_llm_provider"] == "openai"
mock_registry.get_litellm_managed_vector_store_from_registry.assert_called_once_with(
vector_store_id="vs_path_allowed"
)
@pytest.mark.asyncio
async def test_vector_store_file_create_denies_other_team_path_store():
from litellm.proxy.vector_store_files_endpoints.endpoints import (
vector_store_file_create,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
request = _mock_request()
with (
patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(return_value={"file_id": "file_123"}),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=AsyncMock(),
) as mock_base_process,
):
with pytest.raises(HTTPException) as exc_info:
await vector_store_file_create(
vector_store_id="vs_other_team",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
mock_base_process.assert_not_called()
@pytest.mark.asyncio
async def test_rag_query_denies_nested_other_team_vector_store():
from litellm.proxy.rag_endpoints.endpoints import rag_query
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
request = _mock_request()
with (
patch(
"litellm.proxy.rag_endpoints.endpoints._read_request_body",
new=AsyncMock(
return_value={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"retrieval_config": {"vector_store_id": "vs_other_team"},
}
),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new=AsyncMock(),
) as mock_aquery,
):
with pytest.raises(HTTPException) as exc_info:
await rag_query(
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
mock_aquery.assert_not_called()
@pytest.mark.asyncio
async def test_rag_ingest_denies_nested_other_team_vector_store():
from litellm.proxy.rag_endpoints.endpoints import rag_ingest
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
request = _mock_request()
with (
patch(
"litellm.proxy.rag_endpoints.endpoints.parse_rag_ingest_request",
new=AsyncMock(
return_value=(
{
"vector_store": {
"custom_llm_provider": "openai",
"vector_store_id": "vs_other_team",
}
},
None,
"https://example.com/file.txt",
None,
)
),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aingest",
new=AsyncMock(),
) as mock_aingest,
):
with pytest.raises(HTTPException) as exc_info:
await rag_ingest(
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
mock_aingest.assert_not_called()
def test_rag_payload_scan_rejects_excessive_nesting():
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy.rag_endpoints.endpoints import (
_collect_vector_store_ids_from_payload,
)
payload = {}
current = payload
for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 1):
current["nested"] = {}
current = current["nested"]
current["vector_store_id"] = "vs_too_deep"
with pytest.raises(HTTPException) as exc_info:
_collect_vector_store_ids_from_payload(payload)
assert exc_info.value.status_code == 400
def test_rag_payload_scan_accepts_vector_store_id_at_depth_limit():
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy.rag_endpoints.endpoints import (
_collect_vector_store_ids_from_payload,
)
payload = {}
current = payload
for _ in range(DEFAULT_MAX_RECURSE_DEPTH):
current["nested"] = {}
current = current["nested"]
current["vector_store_id"] = "vs_at_limit"
assert _collect_vector_store_ids_from_payload(payload) == {"vs_at_limit"}
def test_rag_payload_scan_ignores_primitive_list_beyond_depth_limit():
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy.rag_endpoints.endpoints import (
_collect_vector_store_ids_from_payload,
)
payload = {}
current = payload
for _ in range(DEFAULT_MAX_RECURSE_DEPTH):
current["nested"] = {}
current = current["nested"]
current["labels"] = ["alpha", "beta"]
assert _collect_vector_store_ids_from_payload(payload) == set()
@pytest.mark.asyncio
async def test_responses_file_search_denies_other_team_vector_store():
from litellm.proxy.common_request_processing import (
_authorize_response_file_search_vector_stores,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
with patch.object(litellm, "vector_store_registry", mock_registry):
with pytest.raises(HTTPException) as exc_info:
await _authorize_response_file_search_vector_stores(
data={
"tools": [
{
"type": "file_search",
"vector_store_ids": ["vs_other_team"],
}
]
},
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_vertex_discovery_denies_other_team_vector_store_credentials():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_base_vertex_proxy_route,
)
request = _mock_request()
request.method = "GET"
vector_store_credentials = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "vertex_ai",
"team_id": "team-b",
}
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth",
new=AsyncMock(return_value=UserAPIKeyAuth(team_id="team-a")),
):
with pytest.raises(HTTPException) as exc_info:
await _base_vertex_proxy_route(
endpoint="projects/p/locations/us-central1/dataStores/vs_other_team",
request=request,
fastapi_response=Response(),
get_vertex_pass_through_handler=MagicMock(),
router_credentials=vector_store_credentials,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback():
from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = None
cache_helper = AsyncMock(
return_value=[
LiteLLM_ManagedVectorStoresTable(
vector_store_id="vs_cached",
custom_llm_provider="openai",
vector_store_name=None,
vector_store_description=None,
vector_store_metadata=None,
created_at=None,
updated_at=None,
litellm_credential_name=None,
litellm_params={"api_base": "https://example.com"},
team_id="team-a",
user_id=None,
)
]
)
with (
patch.object(litellm, "vector_store_registry", mock_registry),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids",
new=cache_helper,
),
):
vector_store = await get_litellm_managed_vector_store(
vector_store_id="vs_cached"
)
assert vector_store is not None
assert vector_store["vector_store_id"] == "vs_cached"
assert vector_store["team_id"] == "team-a"
cache_helper.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_managed_vector_store_fails_closed_on_lookup_error():
from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = (
RuntimeError("registry unavailable")
)
with patch.object(litellm, "vector_store_registry", mock_registry):
with pytest.raises(HTTPException) as exc_info:
await get_litellm_managed_vector_store(vector_store_id="vs_registry_only")
assert exc_info.value.status_code == 500
@pytest.mark.asyncio
async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
vertex_discovery_proxy_route,
)
request = _mock_request()
request.method = "GET"
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_litellm_managed_vector_store",
new=AsyncMock(return_value=None),
) as mock_lookup,
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._base_vertex_proxy_route",
new=AsyncMock(return_value={"ok": True}),
) as mock_base_route,
):
response = await vertex_discovery_proxy_route(
endpoint="projects/p/locations/us-central1/dataStores/vs_unknown",
request=request,
fastapi_response=Response(),
)
assert response == {"ok": True}
mock_lookup.assert_awaited_once_with(vector_store_id="vs_unknown")
assert mock_base_route.call_args.kwargs["router_credentials"] is None
@pytest.mark.asyncio
async def test_milvus_passthrough_denies_other_team_vector_store_index():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
milvus_proxy_route,
)
request = _mock_request()
request.url.path = "/milvus/v2/vectordb/entities/search"
index_object = MagicMock()
index_object.litellm_params.vector_store_name = "tenant-b-store"
index_object.litellm_params.vector_store_index = "tenant_b_collection"
mock_index_registry = MagicMock()
mock_index_registry.is_vector_store_index.return_value = True
mock_index_registry.get_vector_store_index_by_name.return_value = index_object
mock_vector_registry = MagicMock()
mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "milvus",
"team_id": "team-b",
"litellm_params": {"api_base": "https://milvus.example.com"},
}
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(),
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
new=AsyncMock(return_value={"collectionName": "managed_index"}),
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint",
return_value=True,
),
patch.object(litellm, "vector_store_index_registry", mock_index_registry),
patch.object(litellm, "vector_store_registry", mock_vector_registry),
):
with pytest.raises(HTTPException) as exc_info:
await milvus_proxy_route(
endpoint="v2/vectordb/entities/search",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_azure_passthrough_denies_other_team_vector_store_index():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
azure_proxy_route,
)
request = _mock_request()
request.url.path = "/azure/indexes/managed_index/docs/search"
index_object = MagicMock()
index_object.litellm_params.vector_store_name = "tenant-b-store"
mock_index_registry = MagicMock()
mock_index_registry.is_vector_store_index.side_effect = (
lambda vector_store_index_name: vector_store_index_name == "managed_index"
)
mock_index_registry.get_vector_store_index_by_name.return_value = index_object
mock_vector_registry = MagicMock()
mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "azure_ai",
"team_id": "team-b",
"litellm_params": {"api_base": "https://azure.example.com"},
}
with (
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
return_value=False,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(),
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint",
return_value=True,
),
patch.object(litellm, "vector_store_index_registry", mock_index_registry),
patch.object(litellm, "vector_store_registry", mock_vector_registry),
):
with pytest.raises(HTTPException) as exc_info:
await azure_proxy_route(
endpoint="indexes/managed_index/docs/search",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403

View file

@ -6,8 +6,8 @@ import {
deriveErrorMessage,
handleError,
} from "@/components/networking";
import { all_admin_roles } from "@/utils/roles";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { all_admin_roles } from "@/utils/roles";
// ── Types ────────────────────────────────────────────────────────────────────
@ -81,7 +81,6 @@ export const useProjects = () => {
return useQuery<ProjectResponse[]>({
queryKey: projectKeys.list({}),
queryFn: async () => fetchProjects(accessToken!),
enabled:
Boolean(accessToken) && all_admin_roles.includes(userRole || ""),
enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!),
});
};

View file

@ -169,8 +169,8 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
}
}, [isAdmin, userID]);
// For non-admins, always pass their own user_id
const effectiveUserId = isAdmin ? selectedUserId : userID || null;
// For non-admins or "my-usage" view, always pass their own user_id
const effectiveUserId = usageView === "my-usage" || !isAdmin ? userID || null : selectedUserId;
const startTime = useMemo(() => (dateValue.from ? new Date(dateValue.from) : null), [dateValue.from]);
const endTime = useMemo(() => (dateValue.to ? new Date(dateValue.to) : null), [dateValue.to]);
@ -477,10 +477,10 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
}
/>
)}
{/* Your Usage Panel */}
{usageView === "global" && (
{/* Your Usage / Global Usage Panel */}
{(usageView === "global" || usageView === "my-usage") && (
<>
{isAdmin && (
{isAdmin && usageView === "global" && (
<div className="mb-4">
<Text className="mb-2">Filter by user</Text>
<Select

View file

@ -11,7 +11,7 @@ import {
} from "@ant-design/icons";
import { Badge, Select } from "antd";
import React from "react";
export type UsageOption = "global" | "organization" | "team" | "customer" | "tag" | "agent" | "user" | "user-agent-activity";
export type UsageOption = "global" | "my-usage" | "organization" | "team" | "customer" | "tag" | "agent" | "user" | "user-agent-activity";
export interface UsageViewSelectProps {
value: UsageOption;
onChange: (value: UsageOption) => void;
@ -43,6 +43,13 @@ const OPTIONS: OptionConfig[] = [
descriptionForNonAdmin: "View your usage",
icon: <GlobalOutlined style={{ fontSize: "16px" }} />,
},
{
value: "my-usage",
label: "Your Usage",
description: "View your own usage",
icon: <UserOutlined style={{ fontSize: "16px" }} />,
adminOnly: true,
},
{
value: "organization",
label: "Organization Usage",