Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_vcr-cassette-llm-tests-af37

# Conflicts:
#	litellm/llms/custom_httpx/llm_http_handler.py
This commit is contained in:
mateo-berri 2026-04-30 17:56:02 -07:00
commit 722a1a9f8f
128 changed files with 41497 additions and 1774 deletions

View file

@ -0,0 +1,75 @@
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."
}

2
.npmrc
View file

@ -2,4 +2,4 @@
# Packages needing lifecycle scripts: npm rebuild <pkg>
ignore-scripts=true
# Protects local npm install only — npm ci (used in CI) ignores this
min-release-age=3d
min-release-age=3

View file

@ -2,4 +2,4 @@
# Packages needing lifecycle scripts: npm rebuild <pkg>
ignore-scripts=true
# Protects local npm install only — npm ci (used in CI) ignores this
min-release-age=3d
min-release-age=3

View file

@ -2,4 +2,4 @@
# Packages needing lifecycle scripts: npm rebuild <pkg>
ignore-scripts=true
# Protects local npm install only — npm ci (used in CI) ignores this
min-release-age=3d
min-release-age=3

View file

@ -0,0 +1,2 @@
-- Search tool allowlists live on LiteLLM_ObjectPermissionTable (with agents, MCP, vector stores).
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN IF NOT EXISTS "search_tools" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -277,6 +277,7 @@ model LiteLLM_ObjectPermissionTable {
models String[] @default([])
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -543,15 +543,17 @@ def _handle_retrieve_batch_providers_without_provider_config(
)
else:
raise litellm.exceptions.BadRequestError(
message="LiteLLM doesn't support {} for 'create_batch'. Only 'openai' is supported.".format(
custom_llm_provider
),
message=(
"LiteLLM doesn't support custom_llm_provider={} for 'retrieve_batch' without a `model` kwarg. "
"Supported via this path: 'openai', 'azure', 'vertex_ai', 'anthropic'. "
"'bedrock' is supported but requires `model` to be passed so the provider config can be loaded."
).format(custom_llm_provider),
model="n/a",
llm_provider=custom_llm_provider,
response=httpx.Response(
status_code=400,
content="Unsupported provider",
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
request=httpx.Request(method="retrieve_batch", url="https://github.com/BerriAI/litellm"), # type: ignore
),
)
return response

View file

@ -1393,6 +1393,10 @@ except (ValueError, TypeError):
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME = "litellm-internal-health-check"
LITTELM_CLI_SERVICE_ACCOUNT_NAME = "litellm-cli"
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME = "litellm_internal_jobs"
# Stable identifier substituted in place of the master key on UserAPIKeyAuth
# objects so the master key (or its hash) never propagates to spend logs,
# Prometheus metrics, audit trails, or any other downstream consumer.
LITELLM_PROXY_MASTER_KEY_ALIAS = "litellm_proxy_master_key"
# Key Rotation Constants
LITELLM_KEY_ROTATION_ENABLED = os.getenv("LITELLM_KEY_ROTATION_ENABLED", "false")
@ -1421,6 +1425,7 @@ LITELLM_PROXY_ADMIN_NAME = "default_user_id"
LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli"
LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token"
CLI_SSO_SESSION_CACHE_KEY_PREFIX = "cli_sso_session"
CLI_SSO_SESSION_TTL_SECONDS = 600
CLI_JWT_TOKEN_NAME = "cli-jwt-token"
# Support both CLI_JWT_EXPIRATION_HOURS and LITELLM_CLI_JWT_EXPIRATION_HOURS for backwards compatibility
CLI_JWT_EXPIRATION_HOURS = int(

View file

@ -87,9 +87,7 @@ class PromptManagementBase(ABC):
try:
messages = compiled_prompt_client["prompt_template"] + client_messages
except Exception as e:
raise ValueError(
f"Error compiling prompt: {e}. Prompt id={prompt_id}, prompt_variables={prompt_variables}, client_messages={client_messages}, dynamic_callback_params={dynamic_callback_params}"
)
raise ValueError(f"Error compiling prompt: {e}. Prompt id={prompt_id}")
compiled_prompt_client["completed_messages"] = messages
return compiled_prompt_client
@ -116,9 +114,7 @@ class PromptManagementBase(ABC):
try:
messages = compiled_prompt_client["prompt_template"] + client_messages
except Exception as e:
raise ValueError(
f"Error compiling prompt: {e}. Prompt id={prompt_id}, prompt_variables={prompt_variables}, client_messages={client_messages}, dynamic_callback_params={dynamic_callback_params}"
)
raise ValueError(f"Error compiling prompt: {e}. Prompt id={prompt_id}")
compiled_prompt_client["completed_messages"] = messages
return compiled_prompt_client

View file

@ -23,6 +23,13 @@ def _raise_env_reference_error(param: str, *, source: str) -> None:
)
def validate_no_callback_env_reference(
param: str, value: object, *, source: str
) -> None:
if _is_env_reference(value):
_raise_env_reference_error(param, source=source)
# Hardcoded list of supported callback params to avoid runtime inspection issues with TypedDict
_supported_callback_params = [
"langfuse_public_key",
@ -66,8 +73,9 @@ def initialize_standard_callback_dynamic_params(
for param in _supported_callback_params:
if param in kwargs:
_param_value = kwargs.get(param)
if _is_env_reference(_param_value):
_raise_env_reference_error(param, source="request body")
validate_no_callback_env_reference(
param, _param_value, source="request body"
)
standard_callback_dynamic_params[param] = _param_value # type: ignore
# 2. Fallback: check "metadata" or "litellm_params" -> "metadata"
@ -80,8 +88,9 @@ def initialize_standard_callback_dynamic_params(
for param in _supported_callback_params:
if param not in standard_callback_dynamic_params and param in metadata:
_param_value = metadata.get(param)
if _is_env_reference(_param_value):
_raise_env_reference_error(param, source="metadata")
validate_no_callback_env_reference(
param, _param_value, source="metadata"
)
standard_callback_dynamic_params[param] = _param_value # type: ignore
return standard_callback_dynamic_params

View file

@ -824,8 +824,6 @@ def convert_to_model_response_object( # noqa: PLR0915
stream=stream,
start_time=start_time,
end_time=end_time,
hidden_params=hidden_params,
_response_headers=_response_headers,
convert_tool_call_to_json_mode=convert_tool_call_to_json_mode,
)
raise Exception(

View file

@ -1661,6 +1661,20 @@ def _sanitize_anthropic_tool_use_id(tool_use_id: str) -> str:
return sanitized
_ANTHROPIC_DOCUMENT_BASE64_MEDIA_TYPES = {"application/pdf", "text/plain"}
def _is_anthropic_document_data_uri(url: str) -> bool:
# Anthropic's base64 document source accepts only application/pdf and
# text/plain (see select_anthropic_content_block_type_for_file). Routing
# other mimes here would produce a document block the API rejects, so we
# leave them on the image code path.
match = re.match(r"data:([^;,]+)", url)
if not match:
return False
return match.group(1) in _ANTHROPIC_DOCUMENT_BASE64_MEDIA_TYPES
def convert_to_anthropic_tool_result(
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
force_base64: bool = False,
@ -1698,14 +1712,24 @@ def convert_to_anthropic_tool_result(
"""
anthropic_content: Union[
str,
List[Union[AnthropicMessagesToolResultContent, AnthropicMessagesImageParam]],
List[
Union[
AnthropicMessagesToolResultContent,
AnthropicMessagesImageParam,
AnthropicMessagesDocumentParam,
]
],
] = ""
if isinstance(message["content"], str):
anthropic_content = message["content"]
elif isinstance(message["content"], List):
content_list = message["content"]
anthropic_content_list: List[
Union[AnthropicMessagesToolResultContent, AnthropicMessagesImageParam]
Union[
AnthropicMessagesToolResultContent,
AnthropicMessagesImageParam,
AnthropicMessagesDocumentParam,
]
] = []
for content in content_list:
if content["type"] == "text":
@ -1720,21 +1744,62 @@ def convert_to_anthropic_tool_result(
text_content["cache_control"] = cache_control_value
anthropic_content_list.append(text_content)
elif content["type"] == "image_url":
image_url_value = content["image_url"]
format = (
content["image_url"].get("format")
if isinstance(content["image_url"], dict)
image_url_value.get("format")
if isinstance(image_url_value, dict)
else None
)
_anthropic_image_param = create_anthropic_image_param(
content["image_url"], format=format, is_bedrock_invoke=force_base64
url_str = (
image_url_value.get("url")
if isinstance(image_url_value, dict)
else image_url_value
)
_anthropic_image_param = add_cache_control_to_content(
anthropic_content_element=_anthropic_image_param,
# Data URIs with non-image mime types (e.g. application/pdf) must
# translate to Anthropic document blocks, not image blocks —
# wrapping a PDF in `type: "image"` is rejected by the API.
if isinstance(url_str, str) and _is_anthropic_document_data_uri(
url_str
):
synth_file_message: ChatCompletionFileObject = {
"type": "file",
"file": {"file_data": url_str},
}
_document_block = anthropic_process_openai_file_message(
synth_file_message
)
_document_block = add_cache_control_to_content(
anthropic_content_element=cast(
AnthropicMessagesDocumentParam, _document_block
),
original_content_element=content,
)
anthropic_content_list.append(
cast(AnthropicMessagesDocumentParam, _document_block)
)
else:
_anthropic_image_param = create_anthropic_image_param(
image_url_value,
format=format,
is_bedrock_invoke=force_base64,
)
_anthropic_image_param = add_cache_control_to_content(
anthropic_content_element=_anthropic_image_param,
original_content_element=content,
)
anthropic_content_list.append(
cast(AnthropicMessagesImageParam, _anthropic_image_param)
)
elif content["type"] == "file":
file_content = cast(ChatCompletionFileObject, content)
_file_block = anthropic_process_openai_file_message(file_content)
_file_block = add_cache_control_to_content(
anthropic_content_element=cast(
AnthropicMessagesDocumentParam, _file_block
),
original_content_element=content,
)
anthropic_content_list.append(
cast(AnthropicMessagesImageParam, _anthropic_image_param)
)
anthropic_content_list.append(_file_block)
anthropic_content = anthropic_content_list
anthropic_tool_result: Optional[AnthropicMessagesToolResultParam] = None
@ -3977,6 +4042,55 @@ def _convert_to_bedrock_tool_call_result(
tool_result_content_blocks.append(
BedrockToolResultContentBlock(image=_block["image"])
)
elif "document" in _block:
tool_result_content_blocks.append(
BedrockToolResultContentBlock(document=_block["document"])
)
else:
verbose_logger.warning(
"Bedrock Converse: unrecognized BedrockContentBlock keys "
"%s for image_url tool-result block %s; dropping.",
list(_block.keys()),
content,
)
elif content["type"] == "file":
# Match the user-message path (_process_file_message): accept
# either file_data (base64 data URI) or file_id (server-side
# reference / URL) and hand off to BedrockImageProcessor. Raise
# BadRequestError on both-None rather than silently dropping.
file_obj = content.get("file") or {}
file_data = file_obj.get("file_data")
file_id = file_obj.get("file_id")
if file_data is None and file_id is None:
raise litellm.BadRequestError(
message="file_data and file_id cannot both be None. Got={}".format(
content
),
model="",
llm_provider="bedrock",
)
file_format = file_obj.get("format")
_file_block: BedrockContentBlock = (
BedrockImageProcessor.process_image_sync(
image_url=cast(str, file_id or file_data),
format=file_format,
)
)
if "document" in _file_block:
tool_result_content_blocks.append(
BedrockToolResultContentBlock(document=_file_block["document"])
)
elif "image" in _file_block:
tool_result_content_blocks.append(
BedrockToolResultContentBlock(image=_file_block["image"])
)
else:
verbose_logger.warning(
"Bedrock Converse: unrecognized BedrockContentBlock keys "
"%s for file tool-result block %s; dropping.",
list(_file_block.keys()),
content,
)
message.get("name", "")
id = str(message.get("tool_call_id", str(uuid.uuid4())))

View file

@ -60,6 +60,9 @@ def _redact_choice_content(choice):
def _redact_responses_api_output(output_items):
"""Helper to redact ResponsesAPIResponse output items."""
for output_item in output_items:
if hasattr(output_item, "text"):
output_item.text = "redacted-by-litellm"
if hasattr(output_item, "content") and isinstance(output_item.content, list):
for content_part in output_item.content:
if hasattr(content_part, "text"):
@ -75,6 +78,28 @@ def _redact_responses_api_output(output_items):
summary_item.text = "redacted-by-litellm"
def _redact_responses_api_output_dict(output_items, redacted_str: str):
"""Helper to redact ResponsesAPIResponse output items in dict form."""
for output_item in output_items:
if not isinstance(output_item, dict):
continue
if "text" in output_item:
output_item["text"] = redacted_str
if isinstance(output_item.get("content"), list):
for content_item in output_item["content"]:
if isinstance(content_item, dict) and "text" in content_item:
content_item["text"] = redacted_str
if output_item.get("type") == "reasoning" and isinstance(
output_item.get("summary"), list
):
for summary_item in output_item["summary"]:
if isinstance(summary_item, dict) and "text" in summary_item:
summary_item["text"] = redacted_str
def _redact_standard_logging_object(model_call_details: dict):
"""Redact messages and response inside standard_logging_object if present."""
standard_logging_object = model_call_details.get("standard_logging_object")
@ -93,28 +118,11 @@ def _redact_standard_logging_object(model_call_details: dict):
if isinstance(response, dict) and "output" in response:
# ResponsesAPIResponse format - redact content in output items
if isinstance(response.get("output"), list):
for output_item in response["output"]:
if isinstance(output_item, dict) and "content" in output_item:
if isinstance(output_item["content"], list):
for content_item in output_item["content"]:
if (
isinstance(content_item, dict)
and "text" in content_item
):
content_item["text"] = redacted_str
_redact_responses_api_output_dict(response["output"], redacted_str)
elif isinstance(response, dict) and "choices" in response:
# ModelResponse dict format - redact content in choices
if isinstance(response.get("choices"), list):
for choice in response["choices"]:
if isinstance(choice, dict):
if "message" in choice and isinstance(choice["message"], dict):
choice["message"]["content"] = redacted_str
if "audio" in choice["message"]:
choice["message"]["audio"] = None
elif "delta" in choice and isinstance(choice["delta"], dict):
choice["delta"]["content"] = redacted_str
if "audio" in choice["delta"]:
choice["delta"]["audio"] = None
_redact_model_response_dict_choices(response["choices"], redacted_str)
elif isinstance(response, str):
standard_logging_object["response"] = redacted_str
else:
@ -122,6 +130,29 @@ def _redact_standard_logging_object(model_call_details: dict):
standard_logging_object["response"] = {"text": redacted_str}
def _redact_model_response_dict_choices(choices, redacted_str: str):
for choice in choices:
if isinstance(choice, dict):
if "message" in choice and isinstance(choice["message"], dict):
choice["message"]["content"] = redacted_str
if "reasoning_content" in choice["message"]:
choice["message"]["reasoning_content"] = redacted_str
if "thinking_blocks" in choice["message"]:
choice["message"]["thinking_blocks"] = None
if "audio" in choice["message"]:
choice["message"]["audio"] = None
elif "delta" in choice and isinstance(choice["delta"], dict):
choice["delta"]["content"] = redacted_str
if "reasoning_content" in choice["delta"]:
choice["delta"]["reasoning_content"] = redacted_str
if "thinking_blocks" in choice["delta"]:
choice["delta"]["thinking_blocks"] = None
if "audio" in choice["delta"]:
choice["delta"]["audio"] = None
else:
_redact_choice_content(choice)
def perform_redaction(model_call_details: dict, result):
"""
Performs the actual redaction on the logging object and result.
@ -132,6 +163,7 @@ def perform_redaction(model_call_details: dict, result):
]
model_call_details["prompt"] = ""
model_call_details["input"] = ""
_redact_standard_logging_object(model_call_details)
# Redact streaming response
if (
@ -171,30 +203,14 @@ def perform_redaction(model_call_details: dict, result):
elif isinstance(_result, dict) and "choices" in _result:
# Handle dict representation of ModelResponse (e.g., from model_dump())
if _result.get("choices") is not None:
for choice in _result["choices"]:
if isinstance(choice, dict):
if "message" in choice and isinstance(choice["message"], dict):
choice["message"]["content"] = "redacted-by-litellm"
if "reasoning_content" in choice["message"]:
choice["message"][
"reasoning_content"
] = "redacted-by-litellm"
if "thinking_blocks" in choice["message"]:
choice["message"]["thinking_blocks"] = None
if "audio" in choice["message"]:
choice["message"]["audio"] = None
elif "delta" in choice and isinstance(choice["delta"], dict):
choice["delta"]["content"] = "redacted-by-litellm"
if "reasoning_content" in choice["delta"]:
choice["delta"][
"reasoning_content"
] = "redacted-by-litellm"
if "thinking_blocks" in choice["delta"]:
choice["delta"]["thinking_blocks"] = None
if "audio" in choice["delta"]:
choice["delta"]["audio"] = None
else:
_redact_choice_content(choice)
_redact_model_response_dict_choices(
_result["choices"], "redacted-by-litellm"
)
elif isinstance(_result, dict) and "output" in _result:
if isinstance(_result.get("output"), list):
_redact_responses_api_output_dict(
_result["output"], "redacted-by-litellm"
)
elif isinstance(_result, litellm.ResponsesAPIResponse):
if hasattr(_result, "output"):
_redact_responses_api_output(_result.output)

View file

@ -21,6 +21,8 @@ class SensitiveDataMasker:
"auth",
"authorization",
"credential",
# Plural form: Vertex uses ``vertex_credentials``; segment-exact
# matching otherwise misses it because "credential" != "credentials".
"credentials",
"access",
"private",

View file

@ -1,6 +1,7 @@
import hashlib
import json
import os
import re
import urllib.parse
from datetime import datetime
from typing import (
@ -37,6 +38,11 @@ else:
AWSPreparedRequest = Any
# Real AWS region names are lowercase letters, digits, and hyphens
# (e.g. "us-east-1", "eu-west-2", "us-gov-west-1", "cn-north-1").
_VALID_AWS_REGION_PATTERN = re.compile(r"\A[a-z0-9-]+\Z")
class Boto3CredentialsInfo(BaseModel):
credentials: Credentials
aws_region_name: str
@ -284,6 +290,9 @@ class BaseAWSLLM:
if not region: # Check if region is empty
return None
if not _VALID_AWS_REGION_PATTERN.match(region):
return None
return region
except Exception:
# Catch any unexpected errors and return None
@ -481,6 +490,7 @@ class BaseAWSLLM:
str: The AWS region name
"""
aws_region_name = optional_params.get("aws_region_name", None)
self._validate_aws_region_name(aws_region_name)
### SET REGION NAME ###
if aws_region_name is None:
# check model arn #
@ -519,8 +529,25 @@ class BaseAWSLLM:
except Exception:
aws_region_name = "us-west-2"
self._validate_aws_region_name(aws_region_name)
return aws_region_name
@staticmethod
def _validate_aws_region_name(aws_region_name: Optional[str]) -> None:
"""
Validate that an AWS region name conforms to the expected format
(lowercase alphanumerics and hyphens). Raises ValueError otherwise.
"""
if aws_region_name is None:
return
if not isinstance(aws_region_name, str) or not _VALID_AWS_REGION_PATTERN.match(
aws_region_name
):
raise ValueError(
f"Invalid AWS region format: {aws_region_name!r}. "
"Region names must contain only lowercase letters, digits, and hyphens."
)
def get_aws_region_name_for_non_llm_api_calls(
self,
aws_region_name: Optional[str] = None,
@ -532,6 +559,7 @@ class BaseAWSLLM:
For non-llm api calls eg. Guardrails, Vector Stores we just need to check the dynamic param or env vars.
"""
self._validate_aws_region_name(aws_region_name)
if aws_region_name is None:
# check env #
litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
@ -549,6 +577,8 @@ class BaseAWSLLM:
if aws_region_name is None:
aws_region_name = "us-west-2"
self._validate_aws_region_name(aws_region_name)
return aws_region_name
@staticmethod

View file

@ -2535,16 +2535,15 @@ class BaseLLMHTTPHandler:
},
)
delete_kwargs: Dict[str, Any] = {
"url": url,
"headers": headers,
"timeout": timeout,
}
if data:
delete_kwargs["json"] = data
try:
# Only send a JSON body when the provider supplied request data;
# some providers (e.g. Azure OpenAI) reject DELETE with a body.
delete_kwargs: Dict[str, Any] = {
"url": url,
"headers": headers,
"timeout": timeout,
}
if data:
delete_kwargs["json"] = data
response = await async_httpx_client.delete(**delete_kwargs)
except Exception as e:
@ -2626,16 +2625,15 @@ class BaseLLMHTTPHandler:
},
)
delete_kwargs: Dict[str, Any] = {
"url": url,
"headers": headers,
"timeout": timeout,
}
if data:
delete_kwargs["json"] = data
try:
# Only send a JSON body when the provider supplied request data;
# some providers (e.g. Azure OpenAI) reject DELETE with a body.
delete_kwargs: Dict[str, Any] = {
"url": url,
"headers": headers,
"timeout": timeout,
}
if data:
delete_kwargs["json"] = data
response = sync_httpx_client.delete(**delete_kwargs)
except Exception as e:

View file

@ -25,7 +25,6 @@ else:
LiteLLMLoggingObj = Any
MILVUS_OPTIONAL_PARAMS = {
"dbName",
"annsField",
"limit",
"filter",
@ -33,7 +32,6 @@ MILVUS_OPTIONAL_PARAMS = {
"groupingField",
"outputFields",
"searchParams",
"partitionNames",
"consistencyLevel",
}
@ -173,13 +171,21 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig):
url = f"{api_base}/v2/vectordb/entities/search"
# Build the request body for Azure AI Search with vector search
request_body = {
request_body: Dict[str, Any] = {
"collectionName": index_name,
"data": [query_vector],
"annsField": "book_intro_vector",
**vector_store_search_optional_params,
}
db_name = litellm_params.get("milvus_db_name")
if db_name:
request_body["dbName"] = db_name
partition_names = litellm_params.get("milvus_partition_names")
if partition_names:
request_body["partitionNames"] = partition_names
#########################################################
# Update logging object with details of the request
#########################################################

View file

@ -712,6 +712,7 @@
},
"anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
@ -735,6 +736,7 @@
},
"anthropic.claude-haiku-4-5@20251001": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
@ -955,6 +957,7 @@
},
"anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -982,6 +985,7 @@
},
"anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1011,6 +1015,7 @@
},
"global.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1040,6 +1045,7 @@
},
"us.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1127,6 +1133,7 @@
},
"anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1157,6 +1164,7 @@
},
"global.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1187,6 +1195,7 @@
},
"us.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1277,6 +1286,7 @@
},
"anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
@ -1305,6 +1315,7 @@
},
"global.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
@ -1333,6 +1344,7 @@
},
"us.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
@ -1447,11 +1459,13 @@
},
"anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -17921,11 +17935,13 @@
},
"global.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -17982,6 +17998,7 @@
},
"global.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
@ -30116,6 +30133,7 @@
},
"us.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 1.1e-06,
"litellm_provider": "bedrock_converse",
@ -30267,11 +30285,13 @@
},
"us.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"input_cost_per_token_above_200k_tokens": 6.6e-06,
"output_cost_per_token_above_200k_tokens": 2.475e-05,
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -30372,6 +30392,7 @@
},
"us.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -30399,6 +30420,7 @@
},
"global.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",

View file

@ -21,6 +21,7 @@ import httpx
from httpx._types import CookieTypes, QueryParamTypes, RequestFiles
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
@ -390,19 +391,28 @@ def _sync_streaming(
):
from litellm.utils import executor
raw_bytes: List[bytes] = []
flush_scheduled = False
try:
raw_bytes: List[bytes] = []
for chunk in response.iter_bytes(): # type: ignore
raw_bytes.append(chunk)
yield chunk
executor.submit(
litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=raw_bytes,
provider_config=provider_config,
)
except Exception as e:
raise e
finally:
if not flush_scheduled and raw_bytes:
flush_scheduled = True
try:
executor.submit(
litellm_logging_obj.flush_passthrough_collected_chunks,
raw_bytes=raw_bytes,
provider_config=provider_config,
)
except Exception as e:
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush "
"in _sync_streaming; %d buffered chunks dropped: %s",
len(raw_bytes),
e,
)
async def _async_streaming(
@ -411,23 +421,45 @@ async def _async_streaming(
provider_config: "BasePassthroughConfig",
):
iter_response = await response
try:
iter_response.raise_for_status()
raw_bytes: List[bytes] = []
async for chunk in iter_response.aiter_bytes(): # type: ignore
raw_bytes.append(chunk)
yield chunk
asyncio.create_task(
litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
)
except Exception:
try:
await iter_response.aclose()
except Exception:
pass
raise
raw_bytes: List[bytes] = []
flush_scheduled = False
try:
async for chunk in iter_response.aiter_bytes(): # type: ignore
raw_bytes.append(chunk)
yield chunk
except Exception:
try:
await iter_response.aclose()
except Exception:
pass
raise
finally:
# GeneratorExit (raised on client disconnect) is not caught by
# `except Exception`; the finally block ensures partial usage
# still gets flushed for spend tracking. See LIT-2642.
if not flush_scheduled and raw_bytes:
flush_scheduled = True
try:
asyncio.create_task(
litellm_logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=raw_bytes,
provider_config=provider_config,
)
)
except Exception as e:
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush "
"in _async_streaming; %d buffered chunks dropped: %s",
len(raw_bytes),
e,
)

View file

@ -117,7 +117,10 @@ class MCPRequestHandler:
return b"{}"
request.body = mock_body # type: ignore
if ".well-known" in str(request.url): # public routes
# Only OAuth metadata routes registered under /.well-known/ are public.
# Match on request.url.path (path-only, exact prefix) so the substring
# cannot be smuggled via query string, hostname, or a deeper URL segment.
if request.url.path.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif has_explicit_litellm_key:
# Explicit x-litellm-api-key provided - always validate normally
@ -126,27 +129,37 @@ class MCPRequestHandler:
)
elif oauth2_headers:
# No x-litellm-api-key, but Authorization header present.
# Could be a LiteLLM key (backward compat) OR an OAuth2 token
# from an upstream MCP provider (e.g. Atlassian).
# Try LiteLLM auth first; on auth failure, treat as OAuth2 passthrough.
# Could be a LiteLLM key (backward compat) OR an opaque OAuth2 token
# the operator wants forwarded to an upstream OAuth2-mode MCP server.
# Try LiteLLM auth first; on auth failure, only fall back to anonymous
# passthrough when the request actually targets a server whose operator
# configured ``auth_type=oauth2``. For any other server (api_key,
# bearer_token, basic, etc.), a failed LiteLLM auth is a real failure
# and must propagate — otherwise an attacker can exchange any garbage
# bearer for an anonymous session.
try:
validated_user_api_key_auth = await user_api_key_auth(
api_key=litellm_api_key, request=request
)
except HTTPException as e:
if e.status_code in (401, 403):
except (HTTPException, ProxyException) as e:
# HTTPException.status_code is int; ProxyException.code is
# normalized to str in its __init__ but can be ``"None"`` or any
# non-numeric string when the caller didn't supply a numeric
# code, so we compare against both int and str forms rather
# than coercing (``int("None")`` would raise ValueError and
# rewrite the auth error as a 500).
status = e.status_code if isinstance(e, HTTPException) else e.code
if status in (
401,
403,
"401",
"403",
) and MCPRequestHandler._target_servers_use_oauth2(
path=request.url.path, mcp_servers=mcp_servers
):
verbose_logger.debug(
"MCP OAuth2: Authorization header is not a valid LiteLLM key, "
"treating as OAuth2 token passthrough"
)
validated_user_api_key_auth = UserAPIKeyAuth()
else:
raise
except ProxyException as e:
if str(e.code) in ("401", "403"):
verbose_logger.debug(
"MCP OAuth2: Authorization header is not a valid LiteLLM key, "
"treating as OAuth2 token passthrough"
"MCP OAuth2: target server is OAuth2-mode, treating "
"Authorization as upstream OAuth2 token passthrough"
)
validated_user_api_key_auth = UserAPIKeyAuth()
else:
@ -165,6 +178,62 @@ class MCPRequestHandler:
dict(headers),
)
@staticmethod
def _extract_target_server_names_from_path(path: str) -> List[str]:
"""
Extract the target MCP server name from the standard MCP transport
URL patterns: ``/mcp/{server_name}[/...]`` and
``/{server_name}/mcp[/...]``. Returns ``[]`` for any other path so
callers fail closed when the target cannot be resolved.
REST/admin endpoints, OAuth2 server endpoints
(``/{server_name}/authorize``, ``/token`` etc.), and ``.well-known``
discovery routes intentionally fall through — those flows do not need
OAuth2 token passthrough. Clients aggregating multiple servers should
use ``x-mcp-servers``, which takes precedence over path parsing.
"""
segments = [s for s in path.split("/") if s]
if len(segments) >= 2 and segments[0] == "mcp":
return [segments[1]]
if len(segments) >= 2 and segments[1] == "mcp":
return [segments[0]]
return []
@staticmethod
def _target_servers_use_oauth2(path: str, mcp_servers: Optional[List[str]]) -> bool:
"""
True only when EVERY MCP server the request targets is configured for
``auth_type == oauth2``. If any target is non-OAuth2 — or if the target
cannot be resolved at all — return False so the caller fails closed.
Used to gate the "treat Authorization as opaque OAuth2 token" fallback
in :meth:`process_mcp_request` so a failed LiteLLM-auth cannot be
exchanged for an anonymous session against a non-OAuth2 server.
"""
# Inline imports avoid a circular dependency: mcp_server_manager imports
# from this module.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.types.mcp import MCPAuth
# Use the x-mcp-servers header verbatim when present (including the
# explicitly-empty list, which means "no targets" → fail closed).
# Only fall back to path parsing when the header was absent entirely.
target_names = (
mcp_servers
if mcp_servers is not None
else MCPRequestHandler._extract_target_server_names_from_path(path)
)
if not target_names:
return False
for name in target_names:
server = global_mcp_server_manager.get_mcp_server_by_name(name)
if server is None or server.auth_type != MCPAuth.oauth2:
return False
return True
@staticmethod
def _get_mcp_auth_header_from_headers(headers: Headers) -> Optional[str]:
"""

View file

@ -1,4 +1,5 @@
import base64
import binascii
import json
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, Iterable, List, Optional, Set, Union, cast
@ -498,6 +499,82 @@ async def rotate_mcp_server_credentials_master_key(
)
def _decode_user_credential(stored: str) -> Optional[str]:
"""Read back a value persisted in ``LiteLLM_MCPUserCredentials.credential_b64``.
Tries nacl decryption first (current write format). Falls back to a
plain ``urlsafe_b64decode`` for rows persisted by older code that wrote
the credential without encryption. Returns ``None`` when neither path
yields a valid string.
"""
decrypted = decrypt_value_helper(
value=stored,
key="mcp_user_credential",
exception_type="debug",
return_original_value=False,
)
if decrypted is not None:
return decrypted
try:
return base64.urlsafe_b64decode(stored).decode()
except (binascii.Error, UnicodeDecodeError, ValueError, TypeError):
return None
def _decode_oauth_payload(stored: str) -> Optional[Dict[str, Any]]:
"""Return the OAuth2 payload dict if ``stored`` holds one, else ``None``.
A row is considered an OAuth2 credential iff its decoded value parses as
a JSON object with ``"type": "oauth2"``. Plain BYOK credentials (which
share the same column) decode to a non-JSON string and return ``None``.
"""
decoded = _decode_user_credential(stored)
if decoded is None:
return None
try:
parsed = json.loads(decoded)
except (ValueError, TypeError):
return None
if isinstance(parsed, dict) and parsed.get("type") == "oauth2":
return parsed
return None
async def rotate_mcp_user_credentials_master_key(
prisma_client: PrismaClient, new_master_key: str
):
"""Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``.
Reads each ``credential_b64`` with the current salt key (falling back to
legacy plain base64 for unmigrated rows) and writes it back encrypted
under the new master key. Rows that are unreadable under both paths
are logged and skipped so one corrupt row does not abort the rotation.
"""
rows = await prisma_client.db.litellm_mcpusercredentials.find_many()
for row in rows:
plaintext = _decode_user_credential(row.credential_b64)
if plaintext is None:
verbose_proxy_logger.warning(
"rotate_mcp_user_credentials_master_key: could not decode "
"credential for user_id=%s server_id=%s, skipping",
row.user_id,
row.server_id,
)
continue
re_encrypted = encrypt_value_helper(
plaintext, new_encryption_key=new_master_key
)
await prisma_client.db.litellm_mcpusercredentials.update(
where={
"user_id_server_id": {
"user_id": row.user_id,
"server_id": row.server_id,
}
},
data={"credential_b64": re_encrypted},
)
async def store_user_credential(
prisma_client: PrismaClient,
user_id: str,
@ -506,7 +583,7 @@ async def store_user_credential(
) -> None:
"""Store a user credential for a BYOK MCP server."""
encoded = base64.urlsafe_b64encode(credential.encode()).decode()
encoded = encrypt_value_helper(credential)
await prisma_client.db.litellm_mcpusercredentials.upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
@ -532,16 +609,7 @@ async def get_user_credential(
)
if row is None:
return None
try:
return base64.urlsafe_b64decode(row.credential_b64).decode()
except Exception:
# Fall back to nacl decryption for credentials stored by older code
return decrypt_value_helper(
value=row.credential_b64,
key="byok_credential",
exception_type="debug",
return_original_value=False,
)
return _decode_user_credential(row.credential_b64)
async def has_user_credential(
@ -582,7 +650,7 @@ async def store_user_oauth_credential(
) -> None:
"""Persist an OAuth2 access token for a user+server pair.
The payload is JSON-serialised and stored base64-encoded in the same
The payload is JSON-serialised and stored encrypted in the same
``credential_b64`` column used by BYOK. A ``"type": "oauth2"`` key
differentiates it from plain BYOK API keys.
"""
@ -606,29 +674,27 @@ async def store_user_oauth_credential(
payload["scopes"] = scopes
# Guard against silently overwriting a BYOK credential with an OAuth token.
# BYOK credentials lack a "type" field (or use a non-"oauth2" type).
# Skip the guard when the caller knows the row is already an OAuth2 credential
# (e.g. during token refresh), saving an extra DB round-trip.
if not skip_byok_guard:
existing = await prisma_client.db.litellm_mcpusercredentials.find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
if existing is not None:
_byok_error = ValueError(
f"A non-OAuth2 credential already exists for user {user_id} "
f"and server {server_id}. Refusing to overwrite."
if (
existing is not None
and _decode_oauth_payload(existing.credential_b64) is None
):
# Existing row is either a BYOK secret or an OAuth2 row that no
# longer decrypts (e.g. after a salt-key rotation). In either
# case, refuse to overwrite — the caller would clobber data
# that may still be recoverable.
raise ValueError(
f"Existing credential for user {user_id} and server "
f"{server_id} could not be verified as an OAuth2 token. "
f"Refusing to overwrite."
)
try:
raw = json.loads(
base64.urlsafe_b64decode(existing.credential_b64).decode()
)
except Exception:
# Credential is not base64+JSON — it's a plain-text BYOK key.
raise _byok_error
if raw.get("type") != "oauth2":
raise _byok_error
encoded = base64.urlsafe_b64encode(json.dumps(payload).encode()).decode()
encoded = encrypt_value_helper(json.dumps(payload))
await prisma_client.db.litellm_mcpusercredentials.upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
@ -672,15 +738,7 @@ async def get_user_oauth_credential(
)
if row is None:
return None
try:
decoded = base64.urlsafe_b64decode(row.credential_b64).decode()
parsed = json.loads(decoded)
if isinstance(parsed, dict) and parsed.get("type") == "oauth2":
return parsed
# Row exists but is a BYOK (plain string), not an OAuth token
return None
except Exception:
return None
return _decode_oauth_payload(row.credential_b64)
async def list_user_oauth_credentials(
@ -694,14 +752,11 @@ async def list_user_oauth_credentials(
)
results: List[Dict[str, Any]] = []
for row in rows:
try:
decoded = base64.urlsafe_b64decode(row.credential_b64).decode()
parsed = json.loads(decoded)
if isinstance(parsed, dict) and parsed.get("type") == "oauth2":
parsed["server_id"] = row.server_id
results.append(parsed)
except Exception:
pass # Skip non-OAuth rows (BYOK plain strings)
payload = _decode_oauth_payload(row.credential_b64)
if payload is None:
continue
payload["server_id"] = row.server_id
results.append(payload)
return results

View file

@ -131,6 +131,22 @@ def decode_state_hash(encrypted_state: str) -> dict:
return state_data
def _get_validated_client_redirect_uri(state_data: Dict[str, Any]) -> str:
"""Return a loopback client redirect URI from OAuth state."""
redirect_uri = state_data.get("client_redirect_uri") or state_data.get("base_url")
if not redirect_uri or not isinstance(redirect_uri, str):
raise HTTPException(status_code=400, detail="Invalid redirect URI")
validate_loopback_redirect_uri(redirect_uri)
return redirect_uri
def _append_query_params(url: str, params: Dict[str, str]) -> str:
parsed = urlparse(url)
query_params = parse_qsl(parsed.query, keep_blank_values=True)
query_params.extend(params.items())
return urlunparse(parsed._replace(query=urlencode(query_params)))
def _resolve_oauth2_server_for_root_endpoints(
client_ip: Optional[str] = None,
) -> Optional[MCPServer]:
@ -568,7 +584,7 @@ async def authorize(
else None
)
if mcp_server is None and mcp_server_name is None:
mcp_server = _resolve_oauth2_server_for_root_endpoints()
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
raise HTTPException(status_code=404, detail="MCP server not found")
# Use server's stored client_id when caller doesn't supply one.
@ -630,7 +646,7 @@ async def token_endpoint(
lookup_name, client_ip=client_ip
)
if mcp_server is None and mcp_server_name is None:
mcp_server = _resolve_oauth2_server_for_root_endpoints()
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
raise HTTPException(status_code=404, detail="MCP server not found")
return await exchange_token_with_server(
@ -651,7 +667,6 @@ async def token_endpoint(
async def callback(code: str, state: str):
try:
state_data = decode_state_hash(state)
base_url = state_data["base_url"]
original_state = state_data["original_state"]
# Re-validate loopback at the sink. /authorize rejects non-loopback
@ -659,10 +674,10 @@ async def callback(code: str, state: str):
# minted before that check was added have no expiry and remain
# valid indefinitely. Validating here blocks the open-redirect +
# code-theft primitive even for pre-fix states.
validate_loopback_redirect_uri(base_url)
redirect_uri = _get_validated_client_redirect_uri(state_data)
params = {"code": code, "state": original_state}
complete_returned_url = f"{base_url}?{urlencode(params)}"
complete_returned_url = _append_query_params(redirect_uri, params)
return RedirectResponse(url=complete_returned_url, status_code=302)
except HTTPException:
@ -719,16 +734,16 @@ def _build_oauth_protected_resource_response(
)
request_base_url = get_request_base_url(request)
client_ip = IPAddressUtils.get_mcp_client_ip(request)
# When no server name provided, try to resolve the single OAuth2 server
if mcp_server_name is None:
resolved = _resolve_oauth2_server_for_root_endpoints()
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
mcp_server_name = resolved.server_name or resolved.name
mcp_server: Optional[MCPServer] = None
if mcp_server_name:
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
mcp_server_name, client_ip=client_ip
)
@ -835,10 +850,11 @@ def _build_oauth_authorization_server_response(
)
request_base_url = get_request_base_url(request)
client_ip = IPAddressUtils.get_mcp_client_ip(request)
# When no server name provided, try to resolve the single OAuth2 server
if mcp_server_name is None:
resolved = _resolve_oauth2_server_for_root_endpoints()
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
mcp_server_name = resolved.server_name or resolved.name
@ -855,7 +871,6 @@ def _build_oauth_authorization_server_response(
mcp_server: Optional[MCPServer] = None
if mcp_server_name:
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
mcp_server_name, client_ip=client_ip
)
@ -1007,8 +1022,9 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
"client_secret": "dummy",
"redirect_uris": [f"{request_base_url}/callback"],
}
client_ip = IPAddressUtils.get_mcp_client_ip(request)
if not mcp_server_name:
resolved = _resolve_oauth2_server_for_root_endpoints()
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
return await register_client_with_server(
request=request,
@ -1021,7 +1037,6 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
)
return dummy_return
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
mcp_server_name, client_ip=client_ip
)

View file

@ -41,6 +41,7 @@ from litellm.constants import (
MCP_TOOL_LISTING_TIMEOUT,
)
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
@ -1499,6 +1500,47 @@ class MCPServerManager:
)
return await client.get_prompt(get_prompt_request_params)
@staticmethod
def _is_same_authority_metadata_url(url: str, server_url: str) -> bool:
"""
Whether ``url`` shares scheme, host, and port with ``server_url``.
Same-authority metadata URLs are produced by our well-known discovery
construction and by resource servers that publish protected-resource
metadata on the resource origin. These must keep working for
administrator-configured internal MCP servers, so they are fetched
directly. Cross-origin URLs are fetched through ``async_safe_get``.
"""
try:
target = urlparse(url)
base = urlparse(server_url)
except Exception:
return False
if target.scheme not in ("http", "https") or not target.hostname:
return False
target_port = target.port or (443 if target.scheme == "https" else 80)
base_port = base.port or (443 if base.scheme == "https" else 80)
return (
base.scheme == target.scheme
and (base.hostname or "").lower() == target.hostname.lower()
and base_port == target_port
)
async def _fetch_oauth_discovery_url(self, url: str, server_url: str) -> Any:
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.MCP,
params={"timeout": MCP_METADATA_TIMEOUT},
)
if self._is_same_authority_metadata_url(url, server_url):
# Same-authority URLs may point at administrator-configured
# internal MCP servers. Do not run them through user URL
# validation, but also do not follow redirects because the
# redirect target would not inherit the same-authority guarantee.
return await client.get(url, follow_redirects=False)
return await async_safe_get(client, url)
async def _descovery_metadata(
self,
server_url: str,
@ -1514,7 +1556,7 @@ class MCPServerManager:
resource_scopes,
) = await self._attempt_well_known_discovery(server_url)
metadata = await self._fetch_authorization_server_metadata(
authorization_servers
authorization_servers, server_url
)
if (
metadata is None
@ -1555,7 +1597,7 @@ class MCPServerManager:
authorization_servers,
resource_scopes,
) = await self._fetch_oauth_metadata_from_resource(
resource_metadata_url
resource_metadata_url, server_url
)
else:
(
@ -1576,7 +1618,7 @@ class MCPServerManager:
if authorization_servers:
metadata = await self._fetch_authorization_server_metadata(
authorization_servers
authorization_servers, server_url
)
preferred_scopes = scopes or resource_scopes
@ -1616,19 +1658,26 @@ class MCPServerManager:
return resource_metadata_url, scopes
async def _fetch_oauth_metadata_from_resource(
self, resource_metadata_url: str
self, resource_metadata_url: str, server_url: str
) -> Tuple[List[str], Optional[List[str]]]:
if not resource_metadata_url:
return [], None
try:
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.MCP,
params={"timeout": MCP_METADATA_TIMEOUT},
response = await self._fetch_oauth_discovery_url(
resource_metadata_url, server_url
)
response = await client.get(resource_metadata_url)
response.raise_for_status()
data = response.json()
except SSRFError as exc:
verbose_logger.warning(
"MCP OAuth discovery: refusing to fetch resource metadata from %s "
"(rejected by SSRF guard for server %s): %s",
resource_metadata_url,
server_url,
exc,
)
return [], None
except Exception as exc: # pragma: no cover - network issues
verbose_logger.debug(
"Failed to fetch MCP OAuth metadata from %s: %s",
@ -1677,23 +1726,25 @@ class MCPServerManager:
(
authorization_servers,
scopes,
) = await self._fetch_oauth_metadata_from_resource(url)
) = await self._fetch_oauth_metadata_from_resource(url, server_url)
if authorization_servers:
return authorization_servers, scopes
return [], None
async def _fetch_authorization_server_metadata(
self, authorization_servers: List[str]
self, authorization_servers: List[str], server_url: str
) -> Optional[MCPOAuthMetadata]:
for issuer in authorization_servers:
metadata = await self._fetch_single_authorization_server_metadata(issuer)
metadata = await self._fetch_single_authorization_server_metadata(
issuer, server_url
)
if metadata is not None:
return metadata
return None
async def _fetch_single_authorization_server_metadata(
self, issuer_url: str
self, issuer_url: str, server_url: str
) -> Optional[MCPOAuthMetadata]:
try:
parsed = urlparse(issuer_url)
@ -1721,13 +1772,18 @@ class MCPServerManager:
for url in candidate_urls:
try:
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.MCP,
params={"timeout": MCP_METADATA_TIMEOUT},
)
response = await client.get(url)
response = await self._fetch_oauth_discovery_url(url, server_url)
response.raise_for_status()
data = response.json()
except SSRFError as exc:
verbose_logger.warning(
"MCP OAuth discovery: refusing to fetch authorization-server "
"metadata from %s (rejected by SSRF guard for server %s): %s",
url,
server_url,
exc,
)
continue
except Exception as exc: # pragma: no cover - network issues
verbose_logger.debug(
"Failed to fetch authorization metadata from %s: %s",

View file

@ -0,0 +1,432 @@
"""
Lazy registration for optional feature routers. Each LAZY_FEATURES entry
imports its module only on the first request matching its path prefix,
saving ~700 MB at idle for deployments that don't use these features.
First hit pays the import cost (1-3 s for heavy modules); /openapi.json
omits each feature's routes until the feature is warmed.
"""
import asyncio
import importlib
import sys
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Callable, Dict, Tuple
from starlette.types import Receive, Scope, Send
from litellm._logging import verbose_proxy_logger
if TYPE_CHECKING:
from fastapi import APIRouter, FastAPI
def _include_router(attr_name: str = "router") -> Callable[["FastAPI", object], None]:
def _register(app: "FastAPI", module: object) -> None:
app.include_router(getattr(module, attr_name))
return _register
def _mount_app(
prefix: str, attr_name: str = "app"
) -> Callable[["FastAPI", object], None]:
def _register(app: "FastAPI", module: object) -> None:
app.mount(path=prefix, app=getattr(module, attr_name))
return _register
@dataclass(frozen=True)
class LazyFeature:
name: str
module_path: str
path_prefixes: Tuple[str, ...]
register_fn: Callable[["FastAPI", object], None] = field(
default_factory=lambda: _include_router("router")
)
# For routes whose path has a leading parameter (e.g. /{server}/authorize)
# — startswith can't match those, so the matcher also checks endswith.
path_suffixes: Tuple[str, ...] = ()
# Keep the stub injected even after load — for mounted ASGI sub-apps
# whose routes don't appear in the parent app's openapi spec.
persistent_swagger_stub: bool = False
LAZY_FEATURES: Tuple[LazyFeature, ...] = (
LazyFeature(
name="guardrails",
module_path="litellm.proxy.guardrails.guardrail_endpoints",
path_prefixes=(
"/guardrails",
"/v2/guardrails",
"/apply_guardrail",
"/policies/usage",
),
),
LazyFeature(
name="policies",
module_path="litellm.proxy.management_endpoints.policy_endpoints",
# Trailing slash to avoid matching /policies/... (policy_engine).
path_prefixes=("/policy/", "/utils/test_policies_and_guardrails"),
),
LazyFeature(
name="policy_engine",
module_path="litellm.proxy.policy_engine.policy_endpoints",
path_prefixes=("/policies",),
),
LazyFeature(
name="policy_resolve",
module_path="litellm.proxy.policy_engine.policy_resolve_endpoints",
path_prefixes=("/policies/resolve", "/policies/attachments/estimate-impact"),
),
LazyFeature(
name="agents",
module_path="litellm.proxy.agent_endpoints.endpoints",
path_prefixes=("/v1/agents", "/agents", "/agent/"),
),
LazyFeature(
name="a2a",
module_path="litellm.proxy.agent_endpoints.a2a_endpoints",
path_prefixes=("/a2a", "/v1/a2a"),
),
LazyFeature(
name="vector_stores",
module_path="litellm.proxy.vector_store_endpoints.endpoints",
path_prefixes=("/v1/vector_stores", "/vector_stores", "/v1/indexes"),
),
LazyFeature(
name="vector_store_management",
module_path="litellm.proxy.vector_store_endpoints.management_endpoints",
# Trailing slash to avoid matching /vector_stores/... (vector_stores).
path_prefixes=("/vector_store/", "/v1/vector_store/"),
),
LazyFeature(
name="vector_store_files",
# Routes appear under both /v1/vector_stores/{id}/files and the
# un-versioned form, so both prefixes must trigger the load.
module_path="litellm.proxy.vector_store_files_endpoints.endpoints",
path_prefixes=("/v1/vector_stores", "/vector_stores"),
),
LazyFeature(
name="tools",
module_path="litellm.proxy.management_endpoints.tool_management_endpoints",
path_prefixes=("/v1/tool", "/tool"),
),
LazyFeature(
name="search_tools",
module_path="litellm.proxy.search_endpoints.search_tool_management",
path_prefixes=("/search_tools",),
),
# mcp_management owns most /v1/mcp/* admin routes; mcp_app is the mounted
# streaming sub-app at /mcp.
LazyFeature(
name="mcp_management",
module_path="litellm.proxy.management_endpoints.mcp_management_endpoints",
path_prefixes=("/v1/mcp/",),
),
LazyFeature(
# Also serves /.well-known/oauth-* (OAuth metadata discovery).
# No /mcp/oauth prefix here: the mounted /mcp sub-app would
# shadow it, and there are no actual routes there anyway.
name="mcp_byok_oauth",
module_path="litellm.proxy._experimental.mcp_server.byok_oauth_endpoints",
path_prefixes=("/v1/mcp/oauth", "/.well-known/oauth-"),
),
LazyFeature(
# Serves OAuth dance endpoints (/authorize, /token, /callback,
# /register) plus several /.well-known/ discovery URLs at the proxy
# root — needed for MCP-over-OAuth flows even before /mcp is hit.
name="mcp_discoverable",
module_path="litellm.proxy._experimental.mcp_server.discoverable_endpoints",
path_prefixes=(
"/.well-known/oauth-",
"/.well-known/openid-configuration",
"/.well-known/jwks.json",
"/authorize",
"/token",
"/callback",
"/register",
),
# Catches the /{mcp_server_name}/authorize|token|register variants.
path_suffixes=("/authorize", "/token", "/register"),
),
LazyFeature(
name="mcp_rest",
module_path="litellm.proxy._experimental.mcp_server.rest_endpoints",
path_prefixes=("/mcp-rest",),
),
LazyFeature(
# Hardcoded /mcp matches BASE_MCP_ROUTE; importing the constant
# here would defeat lazy loading.
name="mcp_app",
module_path="litellm.proxy._experimental.mcp_server.server",
path_prefixes=("/mcp",),
register_fn=_mount_app("/mcp", attr_name="app"),
persistent_swagger_stub=True,
),
LazyFeature(
name="config_overrides",
module_path="litellm.proxy.management_endpoints.config_override_endpoints",
path_prefixes=("/config_overrides",),
),
LazyFeature(
name="realtime",
module_path="litellm.proxy.realtime_endpoints.endpoints",
path_prefixes=("/openai/v1/realtime", "/v1/realtime", "/realtime"),
),
LazyFeature(
name="anthropic_passthrough",
module_path="litellm.proxy.anthropic_endpoints.endpoints",
path_prefixes=("/v1/messages", "/anthropic", "/api/event_logging"),
),
LazyFeature(
name="anthropic_skills",
module_path="litellm.proxy.anthropic_endpoints.skills_endpoints",
path_prefixes=("/v1/skills", "/skills"),
),
LazyFeature(
name="langfuse_passthrough",
module_path="litellm.proxy.vertex_ai_endpoints.langfuse_endpoints",
path_prefixes=("/langfuse",),
),
LazyFeature(
name="evals",
module_path="litellm.proxy.openai_evals_endpoints.endpoints",
path_prefixes=("/v1/evals", "/evals"),
),
LazyFeature(
name="claude_code_marketplace",
module_path="litellm.proxy.anthropic_endpoints.claude_code_endpoints",
path_prefixes=("/claude-code",),
register_fn=_include_router("claude_code_marketplace_router"),
),
LazyFeature(
name="scim",
module_path="litellm.proxy.management_endpoints.scim.scim_v2",
path_prefixes=("/scim",),
register_fn=_include_router("scim_router"),
),
LazyFeature(
name="cloudzero",
module_path="litellm.proxy.spend_tracking.cloudzero_endpoints",
path_prefixes=("/cloudzero",),
),
LazyFeature(
name="vantage",
module_path="litellm.proxy.spend_tracking.vantage_endpoints",
path_prefixes=("/vantage",),
),
LazyFeature(
name="usage_ai",
module_path="litellm.proxy.management_endpoints.usage_endpoints",
path_prefixes=("/usage/ai",),
),
LazyFeature(
name="prompts",
module_path="litellm.proxy.prompts.prompt_endpoints",
path_prefixes=("/prompts", "/utils/dotprompt_json_converter"),
),
LazyFeature(
name="jwt_mappings",
module_path="litellm.proxy.management_endpoints.jwt_key_mapping_endpoints",
path_prefixes=("/jwt/key/mapping",),
),
LazyFeature(
name="compliance",
module_path="litellm.proxy.management_endpoints.compliance_endpoints",
path_prefixes=("/compliance",),
),
LazyFeature(
name="access_groups",
module_path="litellm.proxy.management_endpoints.access_group_endpoints",
path_prefixes=("/access_group", "/v1/access_group", "/v1/unified_access_group"),
),
)
class LazyFeatureMiddleware:
"""ASGI middleware that imports + registers a feature router on first
matching request. Idempotent; once loaded, subsequent requests skip."""
def __init__(
self,
app,
fastapi_app: "FastAPI",
features: Tuple[LazyFeature, ...] = LAZY_FEATURES,
):
self.app = app
self._fastapi_app = fastapi_app
self._features = features
# Loaded set / per-feature locks live on app.state so the warm endpoint
# and the middleware share them — preventing duplicate registrations
# when both paths fire for the same feature.
if not hasattr(fastapi_app.state, "lazy_loaded"):
fastapi_app.state.lazy_loaded = set()
fastapi_app.state.lazy_locks = {}
@property
def _loaded(self) -> set:
return self._fastapi_app.state.lazy_loaded
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
# Short-circuit once every feature has loaded.
if scope["type"] in ("http", "websocket") and len(self._loaded) < len(
self._features
):
path = scope.get("path", "")
for feat in self._features:
if feat.module_path in self._loaded:
continue
if any(path.startswith(p) for p in feat.path_prefixes) or any(
path.endswith(s) for s in feat.path_suffixes
):
await _force_load(self._fastapi_app, feat)
await self.app(scope, receive, send)
async def _force_load(app: "FastAPI", feat: LazyFeature) -> bool:
"""Import + register a lazy feature exactly once per (app, module).
Shared by the middleware and the /lazy/warm endpoint."""
if not hasattr(app.state, "lazy_loaded"):
app.state.lazy_loaded = set()
app.state.lazy_locks = {}
lock = app.state.lazy_locks.setdefault(feat.module_path, asyncio.Lock())
async with lock:
if feat.module_path in app.state.lazy_loaded:
return False
try:
# Import on a thread (heavy modules take 1-3 s). register_fn
# mutates app.router.routes, so it stays on the loop thread.
loop = asyncio.get_running_loop()
module = await loop.run_in_executor(
None, importlib.import_module, feat.module_path
)
feat.register_fn(app, module)
app.state.lazy_loaded.add(feat.module_path)
app.openapi_schema = None
verbose_proxy_logger.info(
"Lazy-loaded optional feature %r (module: %s)",
feat.name,
feat.module_path,
)
return True
except Exception as exc:
# Mark loaded anyway so we don't retry on every request.
app.state.lazy_loaded.add(feat.module_path)
verbose_proxy_logger.warning(
"Failed to lazy-load optional feature %r (module: %s): %s. "
"This feature's endpoints will return 404 until restart.",
feat.name,
feat.module_path,
exc,
)
return False
def attach_lazy_features(app: "FastAPI") -> None:
app.include_router(_make_warmup_router(app))
app.add_middleware(LazyFeatureMiddleware, fastapi_app=app)
def _make_warmup_router(app: "FastAPI") -> "APIRouter":
"""POST /lazy/warm/{name}: load a feature and return its partial openapi
so the Swagger plugin can merge in-place without a full /openapi.json refetch.
Requires auth — anyone who can hit the proxy can already trigger the same
imports by sending a real request to a feature's prefix, but gating this
debug endpoint avoids unauthenticated callers forcing the import chain."""
from fastapi import APIRouter, Depends, HTTPException
from fastapi.openapi.utils import get_openapi
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
@router.post(
"/lazy/warm/{name}",
include_in_schema=False,
dependencies=[Depends(user_api_key_auth)],
)
async def warm(name: str):
feat = next((f for f in LAZY_FEATURES if f.name == name), None)
if feat is None:
raise HTTPException(404, f"unknown lazy feature: {name}")
if feat.persistent_swagger_stub:
return {"stub_path": None, "paths": {}, "components": {"schemas": {}}}
await _force_load(app, feat)
feat_routes = [
r
for r in app.routes
if any(getattr(r, "path", "").startswith(p) for p in feat.path_prefixes)
]
full = get_openapi(title=app.title, version=app.version, routes=feat_routes)
# Force all operations under one tag so they group under a single Swagger
# section — many lazy modules tag routes inconsistently.
for path_ops in full.get("paths", {}).values():
for op in path_ops.values():
if isinstance(op, dict):
op["tags"] = [feat.name]
return {
"stub_path": feat.path_prefixes[0],
"paths": full.get("paths", {}),
"components": {"schemas": full.get("components", {}).get("schemas", {})},
}
return router
def inject_lazy_stubs(schema: Dict) -> Dict:
"""Inject openapi entries for unloaded features. Uses the snapshot file
when available (full route info), otherwise falls back to a single
placeholder per feature. Any failure logs and returns the schema unchanged
so /openapi.json never 500s on a cosmetic injection bug."""
try:
from litellm.proxy._lazy_openapi_snapshot import load_snapshot
snapshot = load_snapshot()
paths = schema.setdefault("paths", {})
schemas = schema.setdefault("components", {}).setdefault("schemas", {})
for feat in LAZY_FEATURES:
if feat.module_path in sys.modules and not feat.persistent_swagger_stub:
continue
fragment = (snapshot or {}).get(feat.name)
if fragment:
for p, ops in fragment.get("paths", {}).items():
paths.setdefault(p, ops)
for name, sch in (
fragment.get("components", {}).get("schemas", {}).items()
):
schemas.setdefault(name, sch)
continue
prefix = feat.path_prefixes[0]
if prefix in paths:
continue
paths[prefix] = {
"get": {
"tags": [feat.name],
"summary": feat.name,
"responses": {"200": {"description": "OK"}},
}
}
except Exception as exc:
verbose_proxy_logger.warning("inject_lazy_stubs failed: %s", exc)
return schema
def lazy_tag_to_prefix() -> Dict[str, str]:
"""feature.name -> first prefix, used by the Swagger warmup JS plugin.
Returns empty when the snapshot is loaded — the plugin is unnecessary
because /openapi.json already has full route info."""
from litellm.proxy._lazy_openapi_snapshot import load_snapshot
if load_snapshot():
return {}
return {
feat.name: feat.path_prefixes[0]
for feat in LAZY_FEATURES
if not feat.persistent_swagger_stub
}

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,70 @@
"""
Per-feature OpenAPI snapshot for lazy-loaded routers.
The committed JSON is generated by `python -m litellm.proxy._lazy_openapi_snapshot`
and consumed at runtime so /openapi.json can show full route info for unloaded
features without importing them. CI verifies the file is current and surfaces
any drift as a neutral check.
"""
import json
import sys
from pathlib import Path
from typing import Dict, Optional
SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json"
def load_snapshot() -> Optional[Dict[str, Dict]]:
if not SNAPSHOT_FILE.exists():
return None
try:
with SNAPSHOT_FILE.open() as f:
return json.load(f)
except (json.JSONDecodeError, OSError):
return None
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
for feat in LAZY_FEATURES:
if feat.module_path in sys.modules:
continue
try:
module = importlib.import_module(feat.module_path)
feat.register_fn(app, module)
except Exception as exc:
sys.stderr.write(f"warning: skip {feat.name}: {exc}\n")
fragments: Dict[str, Dict] = {}
for feat in LAZY_FEATURES:
feat_routes = [
r
for r in app.routes
if any(getattr(r, "path", "").startswith(p) for p in feat.path_prefixes)
]
if not feat_routes:
continue
full = get_openapi(title=app.title, version=app.version, routes=feat_routes)
# Group all of a feature's routes under one tag.
for path_ops in full.get("paths", {}).values():
for op in path_ops.values():
if isinstance(op, dict):
op["tags"] = [feat.name]
fragments[feat.name] = {
"paths": full.get("paths", {}),
"components": {"schemas": full.get("components", {}).get("schemas", {})},
}
return fragments
if __name__ == "__main__":
fragments = generate_snapshot()
SNAPSHOT_FILE.write_text(json.dumps(fragments, indent=2, sort_keys=True) + "\n")
sys.stdout.write(f"wrote {len(fragments)} feature fragments to {SNAPSHOT_FILE}\n")

View file

@ -17,6 +17,9 @@ from typing_extensions import Required, TypedDict
from litellm._uuid import uuid
from litellm.constants import MCP_STDIO_ALLOWED_COMMANDS
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
validate_no_callback_env_reference,
)
from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.openai import (
AllMessageValues,
@ -904,6 +907,7 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase):
agents: Optional[List[str]] = None
agent_access_groups: Optional[List[str]] = None
models: Optional[List[str]] = None
search_tools: Optional[List[str]] = None
class BudgetLimitEntry(LiteLLMPydanticObjectBase):
@ -1868,8 +1872,10 @@ class AddTeamCallback(LiteLLMPydanticObjectBase):
raise ValueError(
f"Invalid callback variable: {key}. Must be one of {valid_keys}"
)
if not isinstance(value, str):
callback_vars[key] = str(value)
callback_vars[key] = str(value)
validate_no_callback_env_reference(
key, callback_vars[key], source="key/team callback metadata"
)
return values
@ -1934,6 +1940,7 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
agent_access_groups: Optional[List[str]] = []
mcp_toolsets: Optional[List[str]] = None
blocked_tools: Optional[List[str]] = []
search_tools: Optional[List[str]] = []
class LiteLLM_TeamTable(TeamBase):
@ -2154,8 +2161,8 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase):
description="The USD cost per request to the target endpoint. This is used to calculate the cost of the request to the target endpoint.",
)
auth: bool = Field(
default=False,
description="Whether authentication is required for the pass-through endpoint. If True, requests to the endpoint will require a valid LiteLLM API key.",
default=True,
description="Whether authentication is required for the pass-through endpoint. Defaults to True so a pass-through silently created without an explicit value still requires a valid LiteLLM API key — set to False only if the endpoint is meant to be a public forwarder (e.g. an unauthenticated webhook target).",
)
guardrails: Optional[PassThroughGuardrailsConfig] = Field(
default=None,

View file

@ -374,7 +374,7 @@ def _guardrail_modification_check(
coerced = _coerce_to_dict(container)
if coerced is None:
return False
return any(coerced.get(key) for key in _GUARDRAIL_MODIFICATION_KEYS)
return any(key in coerced for key in _GUARDRAIL_MODIFICATION_KEYS)
# Check both metadata keys — callers can populate either depending on the
# endpoint. Cover the top-level too so root-level injection is rejected.
@ -915,7 +915,8 @@ async def get_team_member_default_budget(
Fetches the team-level default per-member budget referenced by team.metadata["team_member_budget_id"].
This budget is applied to team members whose TeamMembership row has no
linked budget. Results are cached for performance.
linked budget, or whose linked budget has max_budget=NULL. Results are
cached for performance.
Args:
budget_id: The budget_id pulled from team.metadata["team_member_budget_id"]
@ -2962,6 +2963,116 @@ async def can_user_call_model(
)
def _search_tool_names_from_object_permission(
object_permission: Optional[LiteLLM_ObjectPermissionTable],
) -> List[str]:
"""Return allowlisted search tool names from object_permission (empty = unrestricted)."""
if object_permission is None:
return []
raw = object_permission.search_tools
if not raw:
return []
return list(raw)
def _can_object_call_search_tool(
search_tool_name: str,
allowed_search_tools: List[str],
object_type: Literal["key", "team", "project"],
) -> Literal[True]:
"""
Check if an object (key/team/project) can access a specific search tool.
Similar to _can_object_call_model but for search tools.
Args:
search_tool_name: The search tool being requested
allowed_search_tools: List of allowed search tool names for this object
object_type: Type of object for error messaging
Returns:
True if access is allowed
Raises:
ProxyException if access is denied
"""
# Empty list means all search tools are allowed
if not allowed_search_tools:
return True
# Check if the search tool is in the allowlist
if search_tool_name in allowed_search_tools:
return True
# Access denied
raise ProxyException(
message=f"{object_type.capitalize()} not allowed to access search tool: {search_tool_name}. "
f"Allowed search tools: {allowed_search_tools}",
type=ProxyErrorTypes.key_model_access_denied,
param="search_tool_name",
code=status.HTTP_403_FORBIDDEN,
)
async def can_key_call_search_tool(
search_tool_name: str,
valid_token: UserAPIKeyAuth,
) -> Literal[True]:
"""
Check if a key can access a specific search tool.
Similar to can_key_call_model but for search tools.
Args:
search_tool_name: The search tool being requested
valid_token: The authenticated key
Returns:
True if access is allowed
Raises:
ProxyException if access is denied
"""
return _can_object_call_search_tool(
search_tool_name=search_tool_name,
allowed_search_tools=_search_tool_names_from_object_permission(
valid_token.object_permission
),
object_type="key",
)
async def can_team_call_search_tool(
search_tool_name: str,
team_object: Optional[LiteLLM_TeamTable],
) -> Literal[True]:
"""
Check if a team can access a specific search tool.
Similar to can_team_access_model but for search tools.
Args:
search_tool_name: The search tool being requested
team_object: The team object
Returns:
True if access is allowed
Raises:
ProxyException if access is denied
"""
if team_object is None:
return True
return _can_object_call_search_tool(
search_tool_name=search_tool_name,
allowed_search_tools=_search_tool_names_from_object_permission(
team_object.object_permission
),
object_type="team",
)
async def is_valid_fallback_model(
model: str,
llm_router: Optional[Router],
@ -3293,6 +3404,7 @@ async def _check_team_member_budget(
if (
team_membership is not None
and team_membership.litellm_budget_table is not None
and team_membership.litellm_budget_table.max_budget is not None
):
team_member_budget = team_membership.litellm_budget_table.max_budget
else:

View file

@ -21,6 +21,7 @@ import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm.caching import DualCache
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
from litellm.proxy._types import *
@ -472,7 +473,12 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints(
for endpoint in pass_through_endpoints:
if isinstance(endpoint, dict) and endpoint.get("path", "") == route:
## IF AUTH DISABLED
if endpoint.get("auth") is not True:
# Default to True: a config dict with no ``auth`` key
# otherwise produced an unauthenticated forwarder. The
# Pydantic ``PassThroughGenericEndpoint.auth`` default
# is also True, but raw config dicts skip that path —
# so this runtime check has to default to True too.
if endpoint.get("auth", True) is not True:
return UserAPIKeyAuth()
## IF AUTH ENABLED
### IF CUSTOM PARSER REQUIRED
@ -1119,10 +1125,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
if is_master_key_valid:
# Substitute a stable alias for the raw master key so neither the
# master key nor its hash propagates into spend logs, Prometheus
# /metrics labels, audit trails, rate-limit buckets, or any other
# downstream consumer of UserAPIKeyAuth.api_key.
_user_api_key_obj = await _return_user_api_key_auth_obj(
user_obj=None,
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key=master_key,
api_key=LITELLM_PROXY_MASTER_KEY_ALIAS,
parent_otel_span=parent_otel_span,
valid_token_dict={
**end_user_params,

View file

@ -474,6 +474,10 @@ async def retrieve_batch( # noqa: PLR0915
)
# Fix: The helper sets "file_id" but we need "batch_id"
data["batch_id"] = data.pop("file_id", original_batch_id)
# Provider-config providers (e.g. bedrock) require `model` in kwargs
# so litellm.aretrieve_batch can load BedrockBatchesConfig. Without
# it the call falls into the legacy provider switch and 400s.
data["model"] = model_from_id
# Retrieve batch using model credentials
response = await litellm.aretrieve_batch(

View file

@ -313,23 +313,24 @@ sequenceDiagram
participant Proxy as LiteLLM Proxy
participant SSO as SSO Provider
CLI->>CLI: Generate key ID (sk-uuid)
CLI->>Browser: Open /sso/key/generate?source=litellm-cli&key=sk-uuid
CLI->>Proxy: POST /sso/cli/start
Proxy->>CLI: Return login_id, poll_secret, user_code
CLI->>Browser: Open /sso/key/generate?source=litellm-cli&key=login_id
Browser->>Proxy: GET /sso/key/generate?source=litellm-cli&key=sk-uuid
Proxy->>Proxy: Set cli_state = litellm-session-token:sk-uuid
Proxy->>SSO: Redirect with state=litellm-session-token:sk-uuid
Browser->>Proxy: GET /sso/key/generate?source=litellm-cli&key=login_id
Proxy->>Proxy: Set cli_state = litellm-session-token:login_id
Proxy->>SSO: Redirect with state=litellm-session-token:login_id
SSO->>Browser: Show login page
Browser->>SSO: User authenticates
SSO->>Proxy: Redirect to /sso/callback?state=litellm-session-token:sk-uuid
SSO->>Proxy: Redirect to /sso/callback?state=litellm-session-token:login_id
Proxy->>Proxy: Check if state starts with "litellm-session-token:"
Proxy->>Proxy: Generate API key with ID=sk-uuid
Proxy->>Browser: Show success page
Proxy->>Browser: Prompt for user_code
Browser->>Proxy: POST /sso/cli/complete/login_id
CLI->>Proxy: Poll /sso/cli/poll/sk-uuid
Proxy->>CLI: Return {"status": "ready", "key": "sk-uuid"}
CLI->>Proxy: Poll /sso/cli/poll/login_id with poll_secret header
Proxy->>CLI: Return {"status": "ready", "key": "jwt"}
CLI->>CLI: Save key to ~/.litellm/token.json
```
@ -343,13 +344,13 @@ The CLI provides three authentication commands:
### Authentication Flow Steps
1. **Generate Session ID**: CLI generates a unique key ID (`sk-{uuid}`)
2. **Open Browser**: CLI opens browser to `/sso/key/generate` with CLI source and key parameters
3. **SSO Redirect**: Proxy sets the formatted state (`litellm-session-token:sk-uuid`) as OAuth state parameter and redirects to SSO provider
1. **Start Session**: CLI creates a short-lived login session with `/sso/cli/start`
2. **Open Browser**: CLI opens browser to `/sso/key/generate` with CLI source and login ID parameters
3. **SSO Redirect**: Proxy sets the formatted state (`litellm-session-token:{login_id}`) as OAuth state parameter and redirects to SSO provider
4. **User Authentication**: User completes SSO authentication in browser
5. **Callback Processing**: SSO provider redirects back to proxy with state parameter
6. **Key Generation**: Proxy detects CLI login (state starts with "litellm-session-token:") and generates API key with pre-specified ID
7. **Polling**: CLI polls `/sso/cli/poll/{key_id}` endpoint until key is ready
6. **User Code Verification**: Browser confirms the verification code shown in the CLI
7. **Polling**: CLI polls `/sso/cli/poll/{login_id}` with the polling secret header until the JWT is ready
8. **Token Storage**: CLI saves the authentication token to `~/.litellm/token.json`
### Benefits of This Approach
@ -357,7 +358,7 @@ The CLI provides three authentication commands:
- **No Local Server**: No need to run a local callback server
- **Standard OAuth**: Uses OAuth 2.0 state parameter correctly
- **Remote Compatible**: Works with remote proxy servers
- **Secure**: Uses UUID session identifiers
- **Secure**: Keeps the polling secret out of the browser handoff
- **Simple Setup**: No additional OAuth redirect URL configuration needed
### Token Storage

View file

@ -5,6 +5,7 @@ import time
import webbrowser
from pathlib import Path
from typing import Any, Dict, List, Optional
from urllib.parse import urlencode
import click
import requests
@ -241,7 +242,7 @@ def prompt_team_selection(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any
def prompt_team_selection_fallback(
teams: List[Dict[str, Any]]
teams: List[Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""Fallback team selection for non-interactive environments"""
if not teams:
@ -279,6 +280,7 @@ def prompt_team_selection_fallback(
def _poll_for_ready_data(
url: str,
*,
headers: Optional[Dict[str, str]] = None,
total_timeout: int = 300,
poll_interval: int = 2,
request_timeout: int = 10,
@ -291,7 +293,10 @@ def _poll_for_ready_data(
) -> Optional[Dict[str, Any]]:
for attempt in range(total_timeout // poll_interval):
try:
response = requests.get(url, timeout=request_timeout)
request_kwargs: Dict[str, Any] = {"timeout": request_timeout}
if headers is not None:
request_kwargs["headers"] = headers
response = requests.get(url, **request_kwargs)
if response.status_code == 200:
data = response.json()
status = data.get("status")
@ -346,7 +351,23 @@ def _normalize_teams(teams, team_details):
return []
def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
def _start_cli_sso_flow(base_url: str) -> Dict[str, Any]:
response = requests.post(f"{base_url}/sso/cli/start", timeout=10)
response.raise_for_status()
data = response.json()
required_fields = ("login_id", "poll_secret", "user_code")
if not all(isinstance(data.get(field), str) for field in required_fields):
raise ValueError("Invalid CLI SSO start response")
return data
def _get_cli_sso_poll_headers(poll_secret: str) -> Dict[str, str]:
return {"x-litellm-cli-poll-secret": poll_secret}
def _poll_for_authentication(
base_url: str, key_id: str, poll_secret: str
) -> Optional[dict]:
"""
Poll the server for authentication completion and handle team selection.
@ -356,6 +377,7 @@ def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
poll_url = f"{base_url}/sso/cli/poll/{key_id}"
data = _poll_for_ready_data(
poll_url,
headers=_get_cli_sso_poll_headers(poll_secret),
pending_message="Still waiting for authentication...",
)
if not data:
@ -373,6 +395,7 @@ def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
jwt_with_team = _handle_team_selection_during_polling(
base_url=base_url,
key_id=key_id,
poll_secret=poll_secret,
teams=normalized_teams,
)
@ -410,7 +433,7 @@ def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
def _handle_team_selection_during_polling(
base_url: str, key_id: str, teams: List[Dict[str, Any]]
base_url: str, key_id: str, poll_secret: str, teams: List[Dict[str, Any]]
) -> Optional[str]:
"""
Handle team selection and re-poll with selected team_id.
@ -441,6 +464,7 @@ def _handle_team_selection_during_polling(
poll_url = f"{base_url}/sso/cli/poll/{key_id}?team_id={team_id}"
data = _poll_for_ready_data(
poll_url,
headers=_get_cli_sso_poll_headers(poll_secret),
pending_message="Still waiting for team authentication...",
other_status_message="Waiting for team authentication to complete...",
http_error_log_every=10,
@ -514,29 +538,24 @@ def _render_and_prompt_for_team_selection(teams: List[Dict[str, Any]]) -> Option
@click.pass_context
def login(ctx: click.Context):
"""Login to LiteLLM proxy using SSO authentication"""
from litellm._uuid import uuid
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
from litellm.proxy.client.cli.interface import show_commands
base_url = ctx.obj["base_url"]
# Check if we have an existing key to regenerate
existing_key = get_stored_api_key()
# Generate unique key ID for this login session
key_id = f"sk-{str(uuid.uuid4())}"
try:
# Construct SSO login URL with CLI source and pre-generated key
sso_url = f"{base_url}/sso/key/generate?source={LITELLM_CLI_SOURCE_IDENTIFIER}&key={key_id}"
cli_sso_flow = _start_cli_sso_flow(base_url=base_url)
key_id = cli_sso_flow["login_id"]
poll_secret = cli_sso_flow["poll_secret"]
user_code = cli_sso_flow["user_code"]
# If we have an existing key, include it as a parameter to the login endpoint
# The server will encode it in the OAuth state parameter for the SSO flow
if existing_key:
sso_url += f"&existing_key={existing_key}"
sso_url = f"{base_url}/sso/key/generate?" + urlencode(
{"source": LITELLM_CLI_SOURCE_IDENTIFIER, "key": key_id}
)
click.echo(f"Opening browser to: {sso_url}")
click.echo("Please complete the SSO authentication in your browser...")
click.echo(f"Verification code: {user_code}")
click.echo(f"Session ID: {key_id}")
# Open browser
@ -545,7 +564,9 @@ def login(ctx: click.Context):
# Poll for authentication completion
click.echo("Waiting for authentication...")
auth_result = _poll_for_authentication(base_url=base_url, key_id=key_id)
auth_result = _poll_for_authentication(
base_url=base_url, key_id=key_id, poll_secret=poll_secret
)
if auth_result:
api_key = auth_result["api_key"]

View file

@ -14,6 +14,8 @@ from litellm.types.utils import (
blue_color_code = "\033[94m"
reset_color_code = "\033[0m"
TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY = "_pillar_response_headers_trusted"
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
@ -417,10 +419,19 @@ def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]:
if "semantic-similarity" in _metadata:
headers["x-litellm-semantic-similarity"] = str(_metadata["semantic-similarity"])
is_trusted_pillar_metadata = (
_metadata.get(TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY) is True
)
pillar_headers = _metadata.get("pillar_response_headers")
if isinstance(pillar_headers, dict):
headers.update(pillar_headers)
elif "pillar_flagged" in _metadata:
if is_trusted_pillar_metadata and isinstance(pillar_headers, dict):
headers.update(
{
key: str(value)
for key, value in pillar_headers.items()
if isinstance(key, str) and key.lower().startswith("x-pillar-")
}
)
elif is_trusted_pillar_metadata and "pillar_flagged" in _metadata:
headers["x-pillar-flagged"] = str(_metadata["pillar_flagged"]).lower()
return headers

View file

@ -0,0 +1,52 @@
"""Helpers for unauthenticated logo / favicon endpoints."""
import os
from typing import Optional, Tuple
from litellm._logging import verbose_proxy_logger
LOCAL_IMAGE_HEADER_BYTES = 512
def detect_local_image_media_type(header: bytes) -> Optional[str]:
"""Return a browser image media type for supported local image signatures."""
if header[0:8] == b"\x89PNG\r\n\x1a\n":
return "image/png"
if header[0:4] == b"GIF8" and header[5:6] == b"a":
return "image/gif"
if header[0:3] == b"\xff\xd8\xff":
return "image/jpeg"
if header[0:4] == b"RIFF" and header[8:12] == b"WEBP":
return "image/webp"
if header[0:4] in (b"\x00\x00\x01\x00", b"\x00\x00\x02\x00"):
return "image/x-icon"
return None
def resolve_validated_local_image_path(candidate: str) -> Optional[Tuple[str, str]]:
"""Resolve ``candidate`` only when it is an existing supported image file."""
if not candidate:
return None
try:
resolved = os.path.realpath(os.path.expanduser(candidate))
except (OSError, ValueError):
return None
if not os.path.isfile(resolved):
return None
try:
with open(resolved, "rb") as f:
header = f.read(LOCAL_IMAGE_HEADER_BYTES)
except OSError as exc:
verbose_proxy_logger.debug("Could not read local asset %r: %s", candidate, exc)
return None
media_type = detect_local_image_media_type(header)
if media_type is None:
verbose_proxy_logger.warning(
"Local asset %r is not a supported image file; falling back to default.",
candidate,
)
return None
return resolved, media_type

View file

@ -1,10 +1,6 @@
from datetime import datetime
from fastapi import APIRouter, Depends, Request, Response
from fastapi.responses import ORJSONResponse
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi.responses import ORJSONResponse, StreamingResponse
import litellm
from litellm._uuid import uuid
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -30,12 +26,17 @@ async def google_generate_content(
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
@ -43,48 +44,33 @@ async def google_generate_content(
if "model" not in data:
data["model"] = model_name
# Extract generationConfig and pass it as config parameter
generation_config = data.pop("generationConfig", None)
if generation_config:
data["config"] = generation_config
# Add user authentication metadata for cost tracking
data = await add_litellm_data_to_request(
data=data,
request=request,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
general_settings=general_settings,
version=version,
)
# Create logging object with full request metadata so callbacks (e.g. S3) get user/trace_id
data["litellm_call_id"] = request.headers.get(
"x-litellm-call-id", str(uuid.uuid4())
)
logging_obj, data = litellm.utils.function_setup(
original_function="agenerate_content",
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(),
**data,
)
data["litellm_logging_obj"] = logging_obj
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content(**data)
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
response=response,
request_data=data,
request=request,
user_api_key_dict=user_api_key_dict,
logging_obj=logging_obj,
version=version,
proxy_logging_obj=proxy_logging_obj,
)
fastapi_response.headers.update(success_headers)
return response
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="agenerate_content",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=model_name,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
@router.post(
@ -101,73 +87,52 @@ async def google_stream_generate_content(
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
data = await _read_request_body(request=request)
if "model" not in data:
data["model"] = model_name
data["stream"] = True
data["stream"] = True # enforce streaming for this endpoint
# Extract generationConfig and pass it as config parameter
generation_config = data.pop("generationConfig", None)
if generation_config:
data["config"] = generation_config
# Add user authentication metadata for cost tracking
data = await add_litellm_data_to_request(
data=data,
request=request,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
general_settings=general_settings,
version=version,
)
# Create logging object with full request metadata so streaming END callbacks (e.g. S3) get user/trace_id
data["litellm_call_id"] = request.headers.get(
"x-litellm-call-id", str(uuid.uuid4())
)
logging_obj, data = litellm.utils.function_setup(
original_function="agenerate_content_stream",
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(),
**data,
)
data["litellm_logging_obj"] = logging_obj
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content_stream(**data)
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
response=response,
request_data=data,
request=request,
user_api_key_dict=user_api_key_dict,
logging_obj=logging_obj,
version=version,
proxy_logging_obj=proxy_logging_obj,
)
# Check if response is an async iterator (streaming response)
if response is not None and hasattr(response, "__aiter__"):
return StreamingResponse(
content=response,
media_type="text/event-stream",
headers=success_headers,
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="agenerate_content_stream",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=model_name,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
fastapi_response.headers.update(success_headers)
return response
@router.post(

View file

@ -71,6 +71,7 @@ from litellm.types.utils import (
)
GUARDRAIL_NAME = "bedrock"
_BEDROCK_DYNAMIC_BODY_DENYLIST = frozenset({"content", "source"})
class GuardrailMessageFilterResult(NamedTuple):
@ -413,11 +414,18 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
api_key: Optional[str] = None
if request_data:
bedrock_request_data.update(
dynamic_request_body_params = (
self.get_guardrail_dynamic_request_body_params(
request_data=request_data
)
)
bedrock_request_data.update(
{
key: value
for key, value in dynamic_request_body_params.items()
if key not in _BEDROCK_DYNAMIC_BODY_DENYLIST
}
)
if request_data.get("api_key") is not None:
api_key = request_data["api_key"]

View file

@ -29,6 +29,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.callback_utils import (
TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY,
add_guardrail_to_applied_guardrails_header,
get_metadata_variable_name_from_kwargs,
)
@ -144,6 +145,7 @@ def build_pillar_response_headers(metadata_store: Dict[str, Any]) -> Dict[str, s
if headers:
metadata_store["pillar_response_headers"] = headers
metadata_store[TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY] = True
return headers

View file

@ -41,6 +41,7 @@ class KeyManagementEventHooks:
"""
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -61,9 +62,11 @@ class KeyManagementEventHooks:
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=response.token_id or "",
@ -102,6 +105,7 @@ class KeyManagementEventHooks:
"""
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -117,9 +121,11 @@ class KeyManagementEventHooks:
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=data.key,
@ -140,6 +146,7 @@ class KeyManagementEventHooks:
):
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -189,9 +196,11 @@ class KeyManagementEventHooks:
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.token,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=existing_key_row.token,
@ -220,6 +229,7 @@ class KeyManagementEventHooks:
"""
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -237,9 +247,11 @@ class KeyManagementEventHooks:
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.token,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=key.token,

View file

@ -192,13 +192,19 @@ class UserManagementEventHooks:
if not litellm.store_audit_logs:
return
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
)
await create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.USER_TABLE_NAME,
object_id=user_id,

View file

@ -6,6 +6,7 @@ from collections import OrderedDict
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from fastapi import Request
from pydantic import ValidationError as PydanticValidationError
from starlette.datastructures import Headers
import litellm
@ -104,6 +105,112 @@ LITELLM_METADATA_ROUTES = (
"files",
)
_UNTRUSTED_ROOT_CONTROL_FIELDS = (
"proxy_server_request",
"standard_logging_object",
"secret_fields",
"mock_response",
"mock_tool_calls",
"disable_global_guardrails",
"disable_global_guardrail",
"opted_out_global_guardrails",
"applied_guardrails",
"applied_policies",
"policy_sources",
"pillar_response_headers",
"_guardrail_pipelines",
"_pipeline_managed_guardrails",
)
_UNTRUSTED_METADATA_CONTROL_FIELDS = (
"disable_global_guardrails",
"disable_global_guardrail",
"opted_out_global_guardrails",
"pillar_response_headers",
"_pillar_response_headers_trusted",
"pillar_flagged",
"pillar_scanners",
"pillar_evidence",
"pillar_evidence_truncated",
"pillar_session_id_response",
"applied_guardrails",
"applied_policies",
"policy_sources",
"standard_logging_object",
"proxy_server_request",
"secret_fields",
"_guardrail_pipelines",
"_pipeline_managed_guardrails",
)
_UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS = frozenset(
{
"litellm-disable-message-redaction",
}
)
_CLIENT_MOCK_CONTROL_FIELDS = frozenset({"mock_response", "mock_tool_calls"})
_ALLOW_CLIENT_MOCK_RESPONSE_METADATA_KEY = "allow_client_mock_response"
_ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY = (
"allow_client_message_redaction_opt_out"
)
def _strip_untrusted_request_header_controls(
headers: Any,
*,
allow_client_message_redaction_opt_out: bool = False,
) -> None:
if not isinstance(headers, dict):
return
for header_name in list(headers.keys()):
if (
isinstance(header_name, str)
and header_name.lower() in _UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS
):
if allow_client_message_redaction_opt_out:
continue
headers.pop(header_name, None)
def _is_false_like(value: Any) -> bool:
if isinstance(value, bool):
return value is False
if isinstance(value, str):
return value.strip().lower() in {"false", "0", "no", "off"}
return False
def _key_or_team_metadata_flag_is_true(
user_api_key_dict: UserAPIKeyAuth,
metadata_key: str,
) -> bool:
for admin_metadata in (user_api_key_dict.metadata, user_api_key_dict.team_metadata):
if (
isinstance(admin_metadata, dict)
and admin_metadata.get(metadata_key) is True
):
return True
return False
def _key_or_team_allows_client_mock_response(
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
return _key_or_team_metadata_flag_is_true(
user_api_key_dict=user_api_key_dict,
metadata_key=_ALLOW_CLIENT_MOCK_RESPONSE_METADATA_KEY,
)
def _key_or_team_allows_client_message_redaction_opt_out(
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
return _key_or_team_metadata_flag_is_true(
user_api_key_dict=user_api_key_dict,
metadata_key=_ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY,
)
def _get_metadata_variable_name(request: Request) -> str:
"""
@ -228,13 +335,25 @@ def convert_key_logging_metadata_to_callback(
for var, value in data.callback_vars.items():
if team_callback_settings_obj.callback_vars is None:
team_callback_settings_obj.callback_vars = {}
team_callback_settings_obj.callback_vars[var] = str(
litellm.utils.get_secret(value, default_value=value) or value
)
team_callback_settings_obj.callback_vars[var] = str(value)
return team_callback_settings_obj
def _get_validated_callback_metadata(
item: dict, *, source: str
) -> Optional[AddTeamCallback]:
try:
return AddTeamCallback(**item)
except (PydanticValidationError, ValueError) as e:
verbose_proxy_logger.warning(
"Ignoring invalid %s callback metadata: %s",
source,
_sanitize_for_log(str(e)),
)
return None
class KeyAndTeamLoggingSettings:
"""
Helper class to get the dynamic logging settings for the key and team
@ -274,8 +393,11 @@ def _get_dynamic_logging_metadata(
#########################################################################################
if key_dynamic_logging_settings is not None:
for item in key_dynamic_logging_settings:
callback = _get_validated_callback_metadata(item=item, source="key-level")
if callback is None:
continue
callback_settings_obj = convert_key_logging_metadata_to_callback(
data=AddTeamCallback(**item),
data=callback,
team_callback_settings_obj=callback_settings_obj,
)
#########################################################################################
@ -283,8 +405,11 @@ def _get_dynamic_logging_metadata(
#########################################################################################
elif team_dynamic_logging_settings is not None:
for item in team_dynamic_logging_settings:
callback = _get_validated_callback_metadata(item=item, source="team-level")
if callback is None:
continue
callback_settings_obj = convert_key_logging_metadata_to_callback(
data=AddTeamCallback(**item),
data=callback,
team_callback_settings_obj=callback_settings_obj,
)
#########################################################################################
@ -904,6 +1029,14 @@ class LiteLLMProxyRequestSetup:
callback_vars_dict.pop("team_id", None)
callback_vars_dict.pop("success_callback", None)
callback_vars_dict.pop("failure_callback", None)
callback_vars_dict = {
key: (
litellm.utils.get_secret(value, default_value=value) or value
if isinstance(value, str)
else value
)
for key, value in callback_vars_dict.items()
}
return TeamCallbackMetadata(
success_callback=team_config.get("success_callback", None),
@ -962,11 +1095,15 @@ async def add_litellm_data_to_request( # noqa: PLR0915
# Strip internal-only keys from user input before the proxy sets its own.
# These keys are injected by the proxy itself below — user-supplied values
# must not be trusted.
for _internal_key in (
"proxy_server_request",
"standard_logging_object",
"secret_fields",
):
_allow_client_mock_response = _key_or_team_allows_client_mock_response(
user_api_key_dict
)
_allow_client_message_redaction_opt_out = (
_key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict)
)
for _internal_key in _UNTRUSTED_ROOT_CONTROL_FIELDS:
if _allow_client_mock_response and _internal_key in _CLIENT_MOCK_CONTROL_FIELDS:
continue
data.pop(_internal_key, None)
# Strip spoofable auth metadata from user-supplied metadata dict
_user_metadata = data.get("metadata")
@ -1007,6 +1144,17 @@ async def add_litellm_data_to_request( # noqa: PLR0915
forward_llm_provider_auth_headers=forward_llm_auth,
authenticated_with_header=authenticated_with_header,
)
_strip_untrusted_request_header_controls(
_headers,
allow_client_message_redaction_opt_out=_allow_client_message_redaction_opt_out,
)
if (
not _allow_client_message_redaction_opt_out
and litellm.turn_off_message_logging is True
and "turn_off_message_logging" in data
and _is_false_like(data["turn_off_message_logging"])
):
data.pop("turn_off_message_logging", None)
verbose_proxy_logger.debug(f"Request Headers: {_headers}")
verbose_proxy_logger.debug(f"Raw Headers: {_raw_headers}")
@ -1144,8 +1292,18 @@ async def add_litellm_data_to_request( # noqa: PLR0915
for _meta_key in ("metadata", "litellm_metadata"):
_user_meta = data.get(_meta_key)
if isinstance(_user_meta, dict):
_user_meta.pop("_pipeline_managed_guardrails", None)
for _k in [k for k in _user_meta if k.startswith("user_api_key_")]:
_strip_untrusted_request_header_controls(
_user_meta.get("headers"),
allow_client_message_redaction_opt_out=(
_allow_client_message_redaction_opt_out
),
)
for _k in [
k
for k in _user_meta
if k.startswith("user_api_key_")
or k in _UNTRUSTED_METADATA_CONTROL_FIELDS
]:
_user_meta.pop(_k, None)
# Strip caller-supplied routing/budget tags unless the admin has opted

View file

@ -2069,6 +2069,9 @@ async def delete_user(
litellm_proxy_admin_name,
prisma_client,
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
)
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
@ -2162,9 +2165,11 @@ async def delete_user(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.USER_TABLE_NAME,
object_id=user_id,

View file

@ -37,6 +37,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._experimental.mcp_server.db import (
rotate_mcp_server_credentials_master_key,
rotate_mcp_user_credentials_master_key,
)
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken
@ -65,6 +66,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
attach_object_permission_to_dict,
handle_update_object_permission_common,
validate_key_mcp_servers_against_team,
validate_key_search_tools_against_team,
)
from litellm.proxy.management_helpers.team_member_permission_checks import (
TeamMemberPermissionChecks,
@ -768,6 +770,10 @@ async def _common_key_generation_helper( # noqa: PLR0915
object_permission=data_json.get("object_permission"),
team_obj=team_table,
)
await validate_key_search_tools_against_team(
object_permission=data_json.get("object_permission"),
team_obj=team_table,
)
data_json = await _set_object_permission(
data_json=data_json,
@ -2010,6 +2016,10 @@ async def _validate_mcp_servers_for_key_update(
object_permission=object_permission_dict,
team_obj=effective_team_obj,
)
await validate_key_search_tools_against_team(
object_permission=object_permission_dict,
team_obj=effective_team_obj,
)
async def _validate_update_key_data(
@ -3709,6 +3719,17 @@ async def _rotate_master_key( # noqa: PLR0915
"Failed to rotate MCP server credentials: %s", str(e)
)
# 4b. process MCP user-scoped credentials table (BYOK + OAuth2 tokens)
try:
await rotate_mcp_user_credentials_master_key(
prisma_client=prisma_client,
new_master_key=new_master_key,
)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to rotate MCP user credentials: %s", str(e)
)
# 5. process credentials table
try:
credentials = await prisma_client.db.litellm_credentialstable.find_many()
@ -5245,6 +5266,9 @@ async def block_key(
proxy_logging_obj,
user_api_key_cache,
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
)
if prisma_client is None:
raise Exception("{}".format(CommonProxyErrors.db_not_connected_error.value))
@ -5288,9 +5312,11 @@ async def block_key(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=hashed_token,
@ -5354,6 +5380,9 @@ async def unblock_key(
proxy_logging_obj,
user_api_key_cache,
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
)
if prisma_client is None:
raise Exception("{}".format(CommonProxyErrors.db_not_connected_error.value))
@ -5397,9 +5426,11 @@ async def unblock_key(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=hashed_token,
@ -5580,7 +5611,6 @@ async def test_key_logging(
"content": "Hello, this is a test from litellm /key/health. No LLM API call was made for this",
}
],
"mock_response": "test response",
}
data = await add_litellm_data_to_request(
data=data,
@ -5589,6 +5619,7 @@ async def test_key_logging(
general_settings=general_settings,
request=request,
)
data["mock_response"] = "test response"
await litellm.acompletion(
**data
) # make mock completion call to trigger key based callbacks

View file

@ -56,6 +56,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.management_helpers.audit_logs import get_audit_log_changed_by
router = APIRouter(prefix="/v1/mcp", tags=["mcp"])
@ -2230,7 +2231,12 @@ if MCP_AVAILABLE:
detail={"error": "Only proxy admins can create MCP toolsets."},
)
touched_by = (
litellm_changed_by or user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME
get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME,
)
or LITELLM_PROXY_ADMIN_NAME
)
try:
result = await create_mcp_toolset(prisma_client, payload, touched_by)
@ -2321,7 +2327,12 @@ if MCP_AVAILABLE:
detail={"error": "Only proxy admins can update MCP toolsets."},
)
touched_by = (
litellm_changed_by or user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME
get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME,
)
or LITELLM_PROXY_ADMIN_NAME
)
try:
result = await update_mcp_toolset(prisma_client, payload, touched_by)

View file

@ -4,15 +4,22 @@ Endpoints to control callbacks per team
Use this when each team should control its own callbacks
"""
import asyncio
import copy
import json
import traceback
from typing import List, Optional
from datetime import datetime, timezone
from typing import Any, List, Optional
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy._types import (
AddTeamCallback,
LiteLLM_AuditLogs,
LitellmTableNames,
ProxyErrorTypes,
ProxyException,
TeamCallbackMetadata,
@ -24,6 +31,106 @@ from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
router = APIRouter()
_CALLBACK_VARS_REDACTED = "***REDACTED***"
def _redact_callback_secrets(metadata: Any) -> Any:
"""Strip secret values out of a team-metadata snapshot before audit logging.
Both ``team_metadata["logging"]`` (list of ``AddTeamCallback`` dicts) and
``team_metadata["callback_settings"]["callback_vars"]`` carry provider
credentials such as ``langfuse_secret_key``, ``langsmith_api_key``, and
``gcs_path_service_account``. Persisting them verbatim into
``LiteLLM_AuditLogs`` would let anyone with read access to the audit
table harvest team callback credentials, so we replace each value with
a fixed marker. The keys themselves are kept so the audit reader can
still see *which* fields changed.
"""
if not isinstance(metadata, dict):
return metadata
redacted = copy.deepcopy(metadata)
logging_entries = redacted.get("logging")
if isinstance(logging_entries, list):
for entry in logging_entries:
if isinstance(entry, dict) and isinstance(entry.get("callback_vars"), dict):
entry["callback_vars"] = {
k: _CALLBACK_VARS_REDACTED for k in entry["callback_vars"]
}
callback_settings = redacted.get("callback_settings")
if isinstance(callback_settings, dict) and isinstance(
callback_settings.get("callback_vars"), dict
):
callback_settings["callback_vars"] = {
k: _CALLBACK_VARS_REDACTED for k in callback_settings["callback_vars"]
}
return redacted
def _log_audit_task_exception(task: "asyncio.Task[None]") -> None:
"""Surface a fire-and-forget audit-log task failure.
``asyncio.create_task`` swallows exceptions silently — if the audit
write fails (transient DB error etc.) we'd otherwise lose the row
without any signal. Log at warning level so the operator sees there's
a gap in the audit trail.
"""
if task.cancelled():
return
exc = task.exception()
if exc is not None:
verbose_proxy_logger.warning("Failed to write team-callback audit log: %s", exc)
async def _emit_team_callback_audit_log(
*,
team_id: str,
before_metadata: Any,
after_metadata: Any,
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: Optional[str],
) -> None:
"""Emit an audit-log row for a team-callback mutation.
Mirrors the ``store_audit_logs``-gated pattern used in
``team_endpoints.py``: the call is async-fire-and-forget and is a no-op
when audit logging is not enabled on the proxy. Captured under
``LitellmTableNames.TEAM_TABLE_NAME`` so the row co-locates with other
team mutations in the audit table.
Callback secrets are redacted before serialization so the audit table
cannot itself become a credential-harvest sink.
"""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
redacted_before = _redact_callback_secrets(before_metadata)
redacted_after = _redact_callback_secrets(after_metadata)
task = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.TEAM_TABLE_NAME,
object_id=team_id,
action="updated",
updated_values=json.dumps({"metadata": redacted_after}, default=str),
before_value=json.dumps({"metadata": redacted_before}, default=str),
)
)
)
task.add_done_callback(_log_audit_task_exception)
@router.post(
"/team/{team_id:path}/callback",
tags=["team management"],
@ -123,6 +230,7 @@ async def add_team_callbacks(
param="callback_name",
)
before_metadata = copy.deepcopy(team_metadata)
team_callback_settings.append(data.model_dump())
team_metadata["logging"] = team_callback_settings
@ -132,6 +240,14 @@ async def add_team_callbacks(
where={"team_id": team_id}, data={"metadata": team_metadata_json} # type: ignore
)
await _emit_team_callback_audit_log(
team_id=team_id,
before_metadata=before_metadata,
after_metadata=team_metadata,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
return {
"status": "success",
"data": new_team_row,
@ -165,6 +281,10 @@ async def disable_team_logging(
http_request: Request,
team_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
),
):
"""
Disable all logging callbacks for a team
@ -198,6 +318,7 @@ async def disable_team_logging(
# Update team metadata to disable logging
team_metadata = _existing_team.metadata
before_metadata = copy.deepcopy(team_metadata)
team_callback_settings = team_metadata.get("callback_settings", {})
team_callback_settings_obj = TeamCallbackMetadata(**team_callback_settings)
@ -222,6 +343,17 @@ async def disable_team_logging(
},
)
# Disabling a team's logging callbacks is itself a logging-control
# action — emit an audit-log row so the action remains traceable
# even though the team's own observability is now off.
await _emit_team_callback_audit_log(
team_id=team_id,
before_metadata=before_metadata,
after_metadata=team_metadata,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
return {
"status": "success",
"message": f"Logging disabled for team {team_id}",

View file

@ -906,6 +906,9 @@ async def new_team( # noqa: PLR0915
prisma_client,
user_api_key_cache,
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
)
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
@ -1174,9 +1177,11 @@ async def new_team( # noqa: PLR0915
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.TEAM_TABLE_NAME,
object_id=data.team_id,
@ -1214,7 +1219,10 @@ async def _create_team_update_audit_log(
user_api_key_dict: User API key authentication details
litellm_proxy_admin_name: Name of the proxy admin
"""
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
)
_before_value = existing_team_row.json(exclude_none=True)
_before_value = json.dumps(_before_value, default=str)
@ -1225,9 +1233,11 @@ async def _create_team_update_audit_log(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.TEAM_TABLE_NAME,
object_id=team_id,
@ -2003,21 +2013,34 @@ def team_member_add_duplication_check(
async def _validate_team_member_add_permissions(
user_api_key_dict: UserAPIKeyAuth,
complete_team_data: LiteLLM_TeamTable,
data: TeamMemberAddRequest,
) -> None:
"""Validate if user has permission to add members to the team."""
"""Validate if user has permission to add members to the team.
Standard users can self-join an *available team*, but the bypass
must not be allowed to escalate them to ``role=admin`` or to add
other users into the team. When access is granted via the
available-team bypass we therefore enforce that every member in
the request matches the caller's own ``user_id`` and is being
added with ``role="user"``.
"""
if (
hasattr(user_api_key_dict, "user_role")
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
and not _is_user_team_admin(
user_api_key_dict=user_api_key_dict, team_obj=complete_team_data
)
and not await _is_user_org_admin_for_team(
user_api_key_dict=user_api_key_dict, team_obj=complete_team_data
)
and not _is_available_team(
team_id=complete_team_data.team_id,
user_api_key_dict=user_api_key_dict,
)
getattr(user_api_key_dict, "user_role", None)
== LitellmUserRoles.PROXY_ADMIN.value
):
return
if _is_user_team_admin(
user_api_key_dict=user_api_key_dict, team_obj=complete_team_data
):
return
if await _is_user_org_admin_for_team(
user_api_key_dict=user_api_key_dict, team_obj=complete_team_data
):
return
if not _is_available_team(
team_id=complete_team_data.team_id,
user_api_key_dict=user_api_key_dict,
):
raise HTTPException(
status_code=403,
@ -2029,6 +2052,34 @@ async def _validate_team_member_add_permissions(
},
)
# Available-team self-join: caller may add only themselves, only as a
# standard user. Enforce that here so the bypass cannot be used as a
# privilege-escalation or cross-user-injection primitive.
members = data.member if isinstance(data.member, list) else [data.member]
caller_user_id = getattr(user_api_key_dict, "user_id", None)
for member in members:
if getattr(member, "role", "user") != "user":
raise HTTPException(
status_code=403,
detail={
"error": (
"Available-team self-join cannot assign 'admin' role. "
"Only proxy/team/org admins can add admins to a team."
)
},
)
member_user_id = getattr(member, "user_id", None)
if not caller_user_id or not member_user_id or member_user_id != caller_user_id:
raise HTTPException(
status_code=403,
detail={
"error": (
"Available-team self-join can only add the caller "
"(user_id must match the authenticated user's user_id)."
)
},
)
async def _process_team_members(
data: TeamMemberAddRequest,
@ -2049,8 +2100,11 @@ async def _process_team_members(
# Resolve allowed_models: explicit request value, or fall back to team's default_team_member_models
member_allowed_models = data.allowed_models
if member_allowed_models is None and complete_team_data.default_team_member_models:
member_allowed_models = complete_team_data.default_team_member_models
team_default_member_models = getattr(
complete_team_data, "default_team_member_models", None
)
if member_allowed_models is None and team_default_member_models:
member_allowed_models = team_default_member_models
if isinstance(data.member, Member):
try:
@ -2381,6 +2435,7 @@ async def team_member_add(
await _validate_team_member_add_permissions(
user_api_key_dict=user_api_key_dict,
complete_team_data=complete_team_data,
data=data,
)
# Validate and populate user_email/user_id for members before processing
@ -2992,6 +3047,9 @@ async def delete_team(
litellm_proxy_admin_name,
prisma_client,
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
)
if prisma_client is None:
raise HTTPException(status_code=500, detail={"error": "No db connected"})
@ -3051,9 +3109,11 @@ async def delete_team(
request_data=LiteLLM_AuditLogs(
id=str(uuid.uuid4()),
updated_at=datetime.now(timezone.utc),
changed_by=litellm_changed_by
or user_api_key_dict.user_id
or litellm_proxy_admin_name,
changed_by=get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.TEAM_TABLE_NAME,
object_id=team_id,
@ -4694,6 +4754,8 @@ async def update_team_member_permissions(
complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump())
# Available-team self-join must NOT grant write access to team-wide
# permission policies; only proxy/team/org admins can update them.
if (
hasattr(user_api_key_dict, "user_role")
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
@ -4703,16 +4765,12 @@ async def update_team_member_permissions(
and not await _is_user_org_admin_for_team(
user_api_key_dict=user_api_key_dict, team_obj=complete_team_data
)
and not _is_available_team(
team_id=complete_team_data.team_id,
user_api_key_dict=user_api_key_dict,
)
):
raise HTTPException(
status_code=403,
detail={
"error": "Call not allowed. User not proxy admin OR team admin. route={}, team_id={}".format(
"/team/member_add",
"/team/permissions_update",
complete_team_data.team_id,
)
},

View file

@ -13,7 +13,9 @@ import base64
import hashlib
import inspect
import os
import re
import secrets
from html import escape
from copy import deepcopy
from typing import (
TYPE_CHECKING,
@ -27,13 +29,13 @@ from typing import (
Union,
cast,
)
from urllib.parse import urlencode, urlparse
from urllib.parse import parse_qs, urlencode, urlparse
if TYPE_CHECKING:
import httpx
import jwt
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi.responses import RedirectResponse
import litellm
@ -41,6 +43,9 @@ from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.caching import DualCache
from litellm.constants import (
CLI_SSO_SESSION_CACHE_KEY_PREFIX,
CLI_SSO_SESSION_TTL_SECONDS,
LITELLM_CLI_SOURCE_IDENTIFIER,
LITELLM_UI_SESSION_DURATION,
MAX_SPENDLOG_ROWS_TO_QUERY,
MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE,
@ -70,7 +75,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object
from litellm.proxy.auth.auth_utils import _has_user_setup_sso
from litellm.proxy.auth.auth_utils import _get_request_ip_address, _has_user_setup_sso
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.admin_ui_utils import (
@ -123,6 +128,250 @@ router = APIRouter()
# Metadata fields (token_type, expires_in, scope) are intentionally kept so
# response convertors see the same fields in the PKCE path as in the non-PKCE path.
_OAUTH_TOKEN_FIELDS = frozenset({"access_token", "id_token", "refresh_token"})
_CLI_SSO_FLOW_CACHE_KEY_PREFIX = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:flow"
_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX = (
f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:start_rate_limit"
)
_CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS = 60
_CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS = 30
_CLI_SSO_USER_CODE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
_CLI_SSO_LOGIN_ID_RE = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$")
def _hash_cli_sso_secret(secret: str) -> str:
return hashlib.sha256(secret.encode("utf-8")).hexdigest()
def _normalize_cli_sso_user_code(user_code: str) -> str:
return "".join(ch for ch in user_code.upper() if ch.isalnum())
def _generate_cli_sso_user_code() -> str:
user_code = "".join(secrets.choice(_CLI_SSO_USER_CODE_ALPHABET) for _ in range(8))
return f"{user_code[:4]}-{user_code[4:]}"
def _get_cli_sso_flow_cache_key(login_id: str) -> str:
return f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:{login_id}"
def _is_valid_cli_sso_login_id(login_id: Optional[str]) -> bool:
return isinstance(login_id, str) and bool(_CLI_SSO_LOGIN_ID_RE.fullmatch(login_id))
def _get_cli_sso_start_rate_limit_cache_key(
request: Request, use_x_forwarded_for: Optional[bool] = False
) -> str:
client_ip = (
_get_request_ip_address(
request=request, use_x_forwarded_for=use_x_forwarded_for
)
or "unknown"
)
client_ip_hash = _hash_cli_sso_secret(client_ip)
return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}"
def _check_cli_sso_start_rate_limit(
request: Request,
cache: DualCache,
use_x_forwarded_for: Optional[bool] = False,
) -> None:
rate_limit_cache_key = _get_cli_sso_start_rate_limit_cache_key(
request=request, use_x_forwarded_for=use_x_forwarded_for
)
current_attempts = cache.increment_cache(
key=rate_limit_cache_key,
value=1,
ttl=_CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS,
)
if current_attempts > _CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS:
raise HTTPException(
status_code=429,
detail="Too many CLI login attempts. Try again later.",
)
def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dict:
if not _is_valid_cli_sso_login_id(login_id):
raise HTTPException(status_code=400, detail="Invalid CLI login session")
cache_key = _get_cli_sso_flow_cache_key(cast(str, login_id))
flow = cache.get_cache(key=cache_key)
if not isinstance(flow, dict) or "poll_secret_hash" not in flow:
raise HTTPException(status_code=400, detail="Invalid CLI login session")
return flow
def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None:
cache.set_cache(
key=_get_cli_sso_flow_cache_key(login_id),
value=flow,
ttl=CLI_SSO_SESSION_TTL_SECONDS,
)
def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool:
expected_poll_secret_hash = flow.get("poll_secret_hash")
if not isinstance(expected_poll_secret_hash, str) or not isinstance(
poll_secret, str
):
return False
supplied_poll_secret_hash = _hash_cli_sso_secret(poll_secret)
return secrets.compare_digest(supplied_poll_secret_hash, expected_poll_secret_hash)
def _render_cli_sso_verification_page(
verify_url: str, browser_complete_token: str
) -> str:
escaped_verify_url = escape(verify_url, quote=True)
escaped_browser_complete_token = escape(browser_complete_token, quote=True)
return f"""
<!doctype html>
<html>
<head>
<title>LiteLLM CLI Login</title>
<style>
body {{
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;
margin: 0;
min-height: 100vh;
display: flex;
align-items: center;
justify-content: center;
background: #f8fafc;
color: #0f172a;
}}
main {{
width: min(420px, calc(100vw - 32px));
background: #ffffff;
border: 1px solid #e2e8f0;
border-radius: 8px;
padding: 28px;
box-shadow: 0 12px 32px rgba(15, 23, 42, 0.08);
}}
h1 {{ font-size: 22px; margin: 0 0 12px; }}
p {{ line-height: 1.5; margin: 0 0 18px; color: #334155; }}
label {{ display: block; font-weight: 600; margin-bottom: 8px; }}
input {{
box-sizing: border-box;
width: 100%;
padding: 12px;
border: 1px solid #cbd5e1;
border-radius: 6px;
font-size: 20px;
letter-spacing: 0.08em;
text-transform: uppercase;
}}
button {{
margin-top: 16px;
width: 100%;
padding: 12px;
border: 0;
border-radius: 6px;
background: #0f172a;
color: #ffffff;
font-weight: 600;
cursor: pointer;
}}
</style>
</head>
<body>
<main>
<h1>Complete CLI Login</h1>
<p>Enter the verification code shown in your terminal to finish this login.</p>
<form method="post" action="{escaped_verify_url}">
<input type="hidden" name="browser_complete_token" value="{escaped_browser_complete_token}" />
<label for="user_code">Verification code</label>
<input id="user_code" name="user_code" autocomplete="one-time-code" required autofocus />
<button type="submit">Continue</button>
</form>
</main>
</body>
</html>
"""
@router.post("/sso/cli/start", tags=["experimental"], include_in_schema=False)
async def cli_sso_start(request: Request):
from litellm.proxy.proxy_server import general_settings, user_api_key_cache
_check_cli_sso_start_rate_limit(
request=request,
cache=user_api_key_cache,
use_x_forwarded_for=bool(
(general_settings or {}).get("use_x_forwarded_for", False)
),
)
login_id = f"cli-{secrets.token_urlsafe(24)}"
poll_secret = secrets.token_urlsafe(32)
user_code = _generate_cli_sso_user_code()
flow = {
"poll_secret_hash": _hash_cli_sso_secret(poll_secret),
"user_code_hash": _hash_cli_sso_secret(_normalize_cli_sso_user_code(user_code)),
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
return {
"login_id": login_id,
"poll_secret": poll_secret,
"user_code": user_code,
"expires_in": CLI_SSO_SESSION_TTL_SECONDS,
}
@router.post(
"/sso/cli/complete/{login_id}", tags=["experimental"], include_in_schema=False
)
async def cli_sso_complete(request: Request, login_id: str):
from fastapi.responses import HTMLResponse
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
render_cli_sso_success_page,
)
from litellm.proxy.proxy_server import user_api_key_cache
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=user_api_key_cache)
if not flow.get("sso_complete") or not flow.get("session_data"):
raise HTTPException(status_code=400, detail="CLI login is not ready")
body = (await request.body()).decode("utf-8")
form_values = parse_qs(body)
supplied_user_code = (form_values.get("user_code") or [""])[0]
supplied_browser_complete_token = (
form_values.get("browser_complete_token") or [""]
)[0]
supplied_user_code_hash = _hash_cli_sso_secret(
_normalize_cli_sso_user_code(supplied_user_code)
)
supplied_browser_complete_token_hash = _hash_cli_sso_secret(
supplied_browser_complete_token
)
expected_user_code_hash = flow.get("user_code_hash")
if not isinstance(expected_user_code_hash, str) or not secrets.compare_digest(
supplied_user_code_hash, expected_user_code_hash
):
raise HTTPException(status_code=400, detail="Invalid verification code")
expected_browser_complete_token_hash = flow.get("browser_complete_token_hash")
if not isinstance(
expected_browser_complete_token_hash, str
) or not secrets.compare_digest(
supplied_browser_complete_token_hash, expected_browser_complete_token_hash
):
raise HTTPException(status_code=400, detail="Invalid verification code")
flow["user_code_verified"] = True
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
html_content = render_cli_sso_success_page()
return HTMLResponse(content=html_content, status_code=200)
def normalize_email(email: Optional[str]) -> Optional[str]:
@ -333,6 +582,7 @@ async def google_login(
from litellm.proxy.proxy_server import (
premium_user,
prisma_client,
user_api_key_cache,
user_custom_ui_sso_sign_in_handler,
)
@ -382,14 +632,15 @@ async def google_login(
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=request,
sso_callback_route="sso/callback",
existing_key=existing_key,
)
# Store CLI key in state for OAuth flow
if source == LITELLM_CLI_SOURCE_IDENTIFIER:
_get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
# Store CLI login handle in state for OAuth flow
cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state(
source=source,
key=key,
existing_key=existing_key,
)
# check if user defined a custom auth sso sign in handler, if yes, use it
@ -1392,18 +1643,12 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
)
if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
# Extract the key ID and existing_key from the state
# State format: {PREFIX}:{key}:{existing_key} or {PREFIX}:{key}
state_parts = state.split(":", 2) # Split into max 3 parts
# State format: {PREFIX}:{login_id}
state_parts = state.split(":", 1)
key_id = state_parts[1] if len(state_parts) > 1 else None
existing_key = state_parts[2] if len(state_parts) > 2 else None
verbose_proxy_logger.info(
f"CLI SSO callback detected for key: {key_id}, existing_key: {existing_key}"
)
return await cli_sso_callback(
request=request, key=key_id, existing_key=existing_key, result=result
)
verbose_proxy_logger.info("CLI SSO callback detected")
return await cli_sso_callback(request=request, key=key_id, result=result)
# Control-plane cross-origin: read return_to from cookie.
# Starlette's cookie_parser already handles RFC 2109 unquoting.
@ -1424,13 +1669,10 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
async def cli_sso_callback(
request: Request,
key: Optional[str] = None,
existing_key: Optional[str] = None,
result: Optional[Union[OpenID, dict]] = None,
):
"""CLI SSO callback - stores session info for JWT generation on polling"""
verbose_proxy_logger.info(
f"CLI SSO callback for key: {key}, existing_key: {existing_key}"
)
verbose_proxy_logger.info("CLI SSO callback")
from litellm.proxy.proxy_server import (
prisma_client,
@ -1438,11 +1680,7 @@ async def cli_sso_callback(
user_api_key_cache,
)
if not key or not key.startswith("sk-"):
raise HTTPException(
status_code=400,
detail="Invalid key parameter. Must be a valid key ID starting with 'sk-'",
)
flow = _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
if prisma_client is None:
raise HTTPException(
@ -1480,9 +1718,6 @@ async def cli_sso_callback(
status_code=500, detail="Failed to retrieve user information from SSO"
)
# Store session info in cache (10 min TTL)
from litellm.constants import CLI_SSO_SESSION_CACHE_KEY_PREFIX
# Get all teams from user_info - CLI will let user select which one
teams: List[str] = []
if hasattr(user_info, "teams") and user_info.teams:
@ -1523,21 +1758,25 @@ async def cli_sso_callback(
"team_details": team_details,
}
cache_key = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:{key}"
user_api_key_cache.set_cache(key=cache_key, value=session_data, ttl=600)
flow["session_data"] = session_data
flow["sso_complete"] = True
browser_complete_token = secrets.token_urlsafe(32)
flow["browser_complete_token_hash"] = _hash_cli_sso_secret(
browser_complete_token
)
_set_cli_sso_flow(login_id=cast(str, key), cache=user_api_key_cache, flow=flow)
verbose_proxy_logger.info(
f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}"
)
# Return success page
from fastapi.responses import HTMLResponse
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
render_cli_sso_success_page,
verify_url = str(request.url_for("cli_sso_complete", login_id=key))
html_content = _render_cli_sso_verification_page(
verify_url=verify_url,
browser_complete_token=browser_complete_token,
)
html_content = render_cli_sso_success_page()
return HTMLResponse(content=html_content, status_code=200)
except Exception as e:
@ -1548,7 +1787,11 @@ async def cli_sso_callback(
@router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False)
async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
async def cli_poll_key(
key_id: str,
team_id: Optional[str] = None,
x_litellm_cli_poll_secret: Optional[str] = Header(default=None),
):
"""
CLI polling endpoint - retrieves session from cache and generates JWT.
@ -1557,22 +1800,25 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
2. Second poll (with team_id): Generates JWT with selected team and deletes session
Args:
key_id: The session key ID
key_id: The CLI login session ID
team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams.
"""
from litellm.constants import CLI_SSO_SESSION_CACHE_KEY_PREFIX
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
from litellm.proxy.proxy_server import user_api_key_cache
if not key_id.startswith("sk-"):
raise HTTPException(status_code=400, detail="Invalid key ID format")
try:
# Look up session in cache
cache_key = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:{key_id}"
session_data = user_api_key_cache.get_cache(key=cache_key)
flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache)
if not _verify_cli_sso_poll_secret(
flow=flow, poll_secret=x_litellm_cli_poll_secret
):
raise HTTPException(status_code=403, detail="Invalid CLI polling secret")
if session_data:
if not flow.get("sso_complete") or not flow.get("user_code_verified"):
return {"status": "pending"}
session_data = flow.get("session_data")
if isinstance(session_data, dict):
user_teams = session_data.get("teams", [])
user_team_details = session_data.get("team_details")
user_id = session_data["user_id"]
@ -1632,7 +1878,7 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
)
# Delete cache entry (single-use)
user_api_key_cache.delete_cache(key=cache_key)
user_api_key_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id))
verbose_proxy_logger.info(
f"CLI JWT generated for user: {user_id}, team: {team_id}"
@ -1650,6 +1896,8 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
else:
return {"status": "pending"}
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.error(f"Error polling for CLI JWT: {e}")
raise HTTPException(
@ -2393,20 +2641,15 @@ class SSOAuthenticationHandler:
This is used to authenticate through the CLI login flow.
The state parameter format is: {PREFIX}:{key}:{existing_key}
- If existing_key is provided, it's included in the state
The state parameter format is: {PREFIX}:{login_id}
- The state parameter is used to pass data through the OAuth flow without changing the callback URL
"""
from litellm.constants import (
LITELLM_CLI_SESSION_TOKEN_PREFIX,
LITELLM_CLI_SOURCE_IDENTIFIER,
)
if source == LITELLM_CLI_SOURCE_IDENTIFIER and key:
if existing_key:
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}:{existing_key}"
else:
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
else:
return None

View file

@ -21,6 +21,28 @@ from litellm.proxy._types import (
from litellm.types.utils import StandardAuditLogPayload
_audit_log_callback_cache: Dict[str, CustomLogger] = {}
ALLOW_LITELLM_CHANGED_BY_HEADER_METADATA_KEY = "allow_litellm_changed_by_header"
def _allows_litellm_changed_by_header(user_api_key_dict: UserAPIKeyAuth) -> bool:
for admin_metadata in (user_api_key_dict.metadata, user_api_key_dict.team_metadata):
if (
isinstance(admin_metadata, dict)
and admin_metadata.get(ALLOW_LITELLM_CHANGED_BY_HEADER_METADATA_KEY) is True
):
return True
return False
def get_audit_log_changed_by(
*,
litellm_changed_by: Optional[str],
user_api_key_dict: UserAPIKeyAuth,
litellm_proxy_admin_name: Optional[str],
) -> Optional[str]:
if litellm_changed_by and _allows_litellm_changed_by_header(user_api_key_dict):
return litellm_changed_by
return user_api_key_dict.user_id or litellm_proxy_admin_name
def _resolve_audit_log_callback(name: str) -> Optional[CustomLogger]:
@ -143,8 +165,10 @@ async def create_object_audit_log(
if _store_audit_logs is not True:
return
_changed_by = (
litellm_changed_by or user_api_key_dict.user_id or litellm_proxy_admin_name
_changed_by = get_audit_log_changed_by(
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
)
await create_audit_log_for_update(

View file

@ -335,8 +335,9 @@ async def validate_key_mcp_servers_against_team(
disallowed_servers = requested_servers - all_allowed_servers
if disallowed_servers:
if team_obj is not None:
team_id = team_obj.team_id
detail = (
f"Key requests MCP servers not allowed by team '{team_obj.team_id}': "
f"Key requests MCP servers not allowed by team '{team_id}': "
f"{sorted(disallowed_servers)}. "
f"Team allows: {sorted(team_allowed_servers)}. "
f"Global (allow_all_keys) servers: {sorted(allow_all_keys_servers)}."
@ -365,8 +366,9 @@ async def validate_key_mcp_servers_against_team(
disallowed_groups = requested_access_groups - team_access_groups
if disallowed_groups:
if team_obj is not None:
team_id = team_obj.team_id
detail = (
f"Key requests MCP access groups not allowed by team '{team_obj.team_id}': "
f"Key requests MCP access groups not allowed by team '{team_id}': "
f"{sorted(disallowed_groups)}. "
f"Team allows: {sorted(team_access_groups)}."
)
@ -390,13 +392,60 @@ async def validate_key_mcp_servers_against_team(
if team_mcp_toolsets:
disallowed_toolsets = requested_toolsets - set(team_mcp_toolsets)
if disallowed_toolsets:
team_id = team_obj.team_id
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"error": (
f"Key requests MCP toolsets not allowed by team '{team_obj.team_id}': "
f"Key requests MCP toolsets not allowed by team '{team_id}': "
f"{sorted(disallowed_toolsets)}. "
f"Team allows: {sorted(team_mcp_toolsets)}."
)
},
)
def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[str]:
"""Return search_tool_name values from a key's object_permission dict."""
if not object_permission or not isinstance(object_permission, dict):
return []
raw = object_permission.get("search_tools")
if not isinstance(raw, list):
return []
return [str(x) for x in raw if x]
async def validate_key_search_tools_against_team(
object_permission: Optional[dict],
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
) -> None:
"""
Validate key object_permission.search_tools is a subset of the team's allowlist.
Empty team allowlist means no restriction at team layer (skip).
"""
requested = _extract_requested_search_tools(object_permission)
if not requested:
return
team_tools: List[str] = []
if team_obj is not None and team_obj.object_permission is not None:
st = team_obj.object_permission.search_tools
if st:
team_tools = list(st)
if not team_tools:
return
disallowed = set(requested) - set(team_tools)
if disallowed:
team_id = team_obj.team_id if team_obj is not None else "unknown"
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"error": (
f"Key requests search tools not allowed by team '{team_id}': "
f"{sorted(disallowed)}. Team allows: {sorted(team_tools)}."
)
},
)

View file

@ -549,10 +549,16 @@ class AnthropicPassthroughLoggingHandler:
# Create a mock user API key dict for the managed object storage
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
_request_metadata = (kwargs.get("litellm_params", {}) or {}).get(
"metadata", {}
) or {}
user_api_key_dict = UserAPIKeyAuth(
user_id=kwargs.get("user_id", "default-user"),
user_id=_request_metadata.get(
"user_api_key_user_id", "default-user"
),
api_key="",
team_id=None,
team_id=_request_metadata.get("user_api_key_team_id"),
team_alias=None,
user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value
user_email=None,

View file

@ -849,10 +849,16 @@ class VertexPassthroughLoggingHandler:
# Create a mock user API key dict for the managed object storage
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
_request_metadata = (kwargs.get("litellm_params", {}) or {}).get(
"metadata", {}
) or {}
user_api_key_dict = UserAPIKeyAuth(
user_id=kwargs.get("user_id", "default-user"),
user_id=_request_metadata.get(
"user_api_key_user_id", "default-user"
),
api_key="",
team_id=None,
team_id=_request_metadata.get("user_api_key_team_id"),
team_alias=None,
user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value
user_email=None,

View file

@ -41,7 +41,6 @@ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.passthrough import BasePassthroughUtils
from litellm.proxy._types import (
CommonProxyErrors,
ConfigFieldInfo,
ConfigFieldUpdate,
LiteLLMRoutes,
@ -2325,12 +2324,14 @@ async def _register_pass_through_endpoint(
dependencies = None
if auth is not None and str(auth).lower() == "true":
if premium_user is not True:
raise ValueError(
"Error Setting Authentication on Pass Through Endpoint: {}".format(
CommonProxyErrors.not_premium_user.value
)
)
# Authentication on a pass-through endpoint used to be enterprise-
# only — which left the OSS tier with no safe configuration: the
# default was ``auth=False`` (unauthenticated forwarder) and the
# safe ``auth=True`` raised at startup unless the operator had a
# license. The default is now ``True`` (safe-by-default), and
# turning it on no longer requires a license: an unauthenticated
# forwarder is a deployment choice the operator should be allowed
# to make explicitly, but the safe option must always be free.
dependencies = [Depends(user_api_key_auth)]
if path not in LiteLLMRoutes.openai_routes.value:
LiteLLMRoutes.openai_routes.value.append(path)

View file

@ -36,21 +36,16 @@ class PassThroughStreamingHandler:
passthrough_success_handler_obj: PassThroughEndpointLogging,
url_route: str,
):
"""
- Yields chunks from the response
- Collect non-empty chunks for post-processing (logging)
- Inject cost into chunks if include_cost_in_streaming_usage is enabled
"""
try:
raw_bytes: List[bytes] = []
# Extract model name for cost injection
model_name = PassThroughStreamingHandler._extract_model_for_cost_injection(
request_body=request_body,
url_route=url_route,
endpoint_type=endpoint_type,
litellm_logging_obj=litellm_logging_obj,
)
raw_bytes: List[bytes] = []
logging_scheduled = False
model_name = PassThroughStreamingHandler._extract_model_for_cost_injection(
request_body=request_body,
url_route=url_route,
endpoint_type=endpoint_type,
litellm_logging_obj=litellm_logging_obj,
)
try:
async for chunk in response.aiter_bytes():
raw_bytes.append(chunk)
if (
@ -58,7 +53,6 @@ class PassThroughStreamingHandler:
and model_name
):
if endpoint_type == EndpointType.VERTEX_AI:
# Only handle streamRawPredict (uses Anthropic format)
if "streamRawPredict" in url_route or "rawPredict" in url_route:
modified_chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
chunk, model_name
@ -73,25 +67,32 @@ class PassThroughStreamingHandler:
chunk = modified_chunk
yield chunk
# After all chunks are processed, handle post-processing
end_time = datetime.now()
asyncio.create_task(
PassThroughStreamingHandler._route_streaming_logging_to_handler(
litellm_logging_obj=litellm_logging_obj,
passthrough_success_handler_obj=passthrough_success_handler_obj,
url_route=url_route,
request_body=request_body or {},
endpoint_type=endpoint_type,
start_time=start_time,
raw_bytes=raw_bytes,
end_time=end_time,
)
)
except Exception as e:
verbose_proxy_logger.error(f"Error in chunk_processor: {str(e)}")
raise
finally:
# GeneratorExit (raised on client disconnect) is not caught by
# `except Exception`; the finally block ensures partial usage
# still gets logged for spend tracking. See LIT-2642.
if not logging_scheduled and raw_bytes:
logging_scheduled = True
try:
asyncio.create_task(
PassThroughStreamingHandler._route_streaming_logging_to_handler(
litellm_logging_obj=litellm_logging_obj,
passthrough_success_handler_obj=passthrough_success_handler_obj,
url_route=url_route,
request_body=request_body or {},
endpoint_type=endpoint_type,
start_time=start_time,
raw_bytes=raw_bytes,
end_time=datetime.now(),
)
)
except Exception as e:
verbose_proxy_logger.error(
f"Error scheduling chunk_processor logging: {str(e)}"
)
@staticmethod
async def _route_streaming_logging_to_handler(

View file

@ -91,6 +91,7 @@ from litellm.proxy._types import (
TeamDefaultSettings,
TokenCountRequest,
TransformRequestBody,
UI_TEAM_ID,
UserAPIKeyAuth,
)
from litellm.proxy.common_utils.callback_utils import (
@ -235,37 +236,11 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import (
router as mcp_byok_oauth_router,
)
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
router as mcp_discoverable_endpoints_router,
)
from litellm.proxy._experimental.mcp_server.rest_endpoints import (
router as mcp_rest_endpoints_router,
)
from litellm.proxy._experimental.mcp_server.server import app as mcp_app
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
from litellm.proxy._types import *
from litellm.proxy.agent_endpoints.a2a_endpoints import router as a2a_router
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.agent_endpoints.endpoints import router as agent_endpoints_router
from litellm.proxy.agent_endpoints.model_list_helpers import (
append_agents_to_model_group,
append_agents_to_model_info,
)
from litellm.proxy._lazy_features import attach_lazy_features
from litellm.proxy.analytics_endpoints.analytics_endpoints import (
router as analytics_router,
)
from litellm.proxy.anthropic_endpoints.claude_code_endpoints import (
claude_code_marketplace_router,
)
from litellm.proxy.anthropic_endpoints.endpoints import router as anthropic_router
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
router as anthropic_skills_router,
)
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
get_team_object,
@ -328,7 +303,6 @@ from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router
from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router
from litellm.proxy.fine_tuning_endpoints.endpoints import set_fine_tuning_config
from litellm.proxy.google_endpoints.endpoints import router as google_router
from litellm.proxy.guardrails.guardrail_endpoints import router as guardrails_router
from litellm.proxy.guardrails.init_guardrails import (
init_guardrails_v2,
initialize_guardrails,
@ -344,9 +318,6 @@ from litellm.proxy.hooks.prompt_injection_detection import (
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
from litellm.proxy.image_endpoints.endpoints import router as image_router
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.management_endpoints.access_group_endpoints import (
router as access_group_router,
)
from litellm.proxy.management_endpoints.budget_management_endpoints import (
router as budget_management_router,
)
@ -360,12 +331,6 @@ from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_privileges,
admin_can_invite_user,
)
from litellm.proxy.management_endpoints.compliance_endpoints import (
router as compliance_router,
)
from litellm.proxy.management_endpoints.config_override_endpoints import (
router as config_override_router,
)
from litellm.proxy.management_endpoints.cost_tracking_settings import (
router as cost_tracking_settings_router,
)
@ -379,9 +344,6 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import (
router as internal_user_router,
)
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import (
router as jwt_key_mapping_router,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
delete_verification_tokens,
duration_in_seconds,
@ -390,9 +352,6 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
from litellm.proxy.management_endpoints.key_management_endpoints import (
router as key_management_router,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
router as mcp_management_router,
)
from litellm.proxy.management_endpoints.model_access_group_management_endpoints import (
router as model_access_group_management_router,
)
@ -407,11 +366,9 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
from litellm.proxy.management_endpoints.organization_endpoints import (
router as organization_router,
)
from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router
from litellm.proxy.management_endpoints.router_settings_endpoints import (
router as router_settings_router,
)
from litellm.proxy.management_endpoints.scim.scim_v2 import scim_router
from litellm.proxy.management_endpoints.tag_management_endpoints import (
router as tag_management_router,
)
@ -423,9 +380,6 @@ from litellm.proxy.management_endpoints.team_endpoints import (
update_team,
validate_membership,
)
from litellm.proxy.management_endpoints.tool_management_endpoints import (
router as tool_management_router,
)
from litellm.proxy.management_endpoints.workflow_management_endpoints import (
router as workflow_management_router,
)
@ -434,7 +388,6 @@ from litellm.proxy.management_endpoints.ui_sso import (
get_disabled_non_admin_personal_key_creation,
)
from litellm.proxy.management_endpoints.ui_sso import router as ui_sso_router
from litellm.proxy.management_endpoints.usage_endpoints import router as usage_ai_router
from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import (
router as user_agent_analytics_router,
)
@ -444,7 +397,6 @@ from litellm.proxy.middleware.in_flight_requests_middleware import (
)
from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware
from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
from litellm.proxy.openai_evals_endpoints.endpoints import router as evals_router
from litellm.proxy.openai_files_endpoints.files_endpoints import (
router as openai_files_router,
)
@ -464,27 +416,16 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
router as pass_through_router,
)
from litellm.proxy.policy_engine.policy_endpoints import router as policy_crud_router
from litellm.proxy.policy_engine.policy_resolve_endpoints import (
router as policy_resolve_router,
)
from litellm.proxy.prompts.prompt_endpoints import router as prompts_router
from litellm.proxy.public_endpoints import router as public_endpoints_router
from litellm.proxy.rag_endpoints.endpoints import router as rag_router
from litellm.proxy.realtime_endpoints.endpoints import router as webrtc_router
from litellm.proxy.rerank_endpoints.endpoints import router as rerank_router
from litellm.proxy.response_api_endpoints.endpoints import router as response_router
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.search_endpoints.endpoints import router as search_router
from litellm.proxy.search_endpoints.search_tool_management import (
router as search_tool_management_router,
)
from litellm.proxy.spend_tracking.cloudzero_endpoints import router as cloudzero_router
from litellm.proxy.spend_tracking.spend_management_endpoints import (
router as spend_management_router,
)
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
from litellm.proxy.spend_tracking.vantage_endpoints import router as vantage_router
from litellm.proxy.types_utils.utils import get_instance_fn
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
router as ui_crud_endpoints_router,
@ -514,16 +455,6 @@ from litellm.proxy.utils import (
prefetch_config_params,
update_spend,
)
from litellm.proxy.vector_store_endpoints.endpoints import router as vector_store_router
from litellm.proxy.vector_store_endpoints.management_endpoints import (
router as vector_store_management_router,
)
from litellm.proxy.vector_store_files_endpoints.endpoints import (
router as vector_store_files_router,
)
from litellm.proxy.vertex_ai_endpoints.langfuse_endpoints import (
router as langfuse_router,
)
from litellm.proxy.video_endpoints.endpoints import router as video_router
from litellm.router import (
AssistantsTypedDict,
@ -1103,6 +1034,11 @@ def get_openapi_schema():
openapi_schema = CustomOpenAPISpec.add_llm_api_request_schema_body(openapi_schema)
# Stub unloaded lazy features so they appear as Swagger sections.
from litellm.proxy._lazy_features import inject_lazy_stubs
openapi_schema = inject_lazy_stubs(openapi_schema)
# Fix Swagger UI execute path error when server_root_path is set
if server_root_path:
openapi_schema["servers"] = [{"url": "/" + server_root_path.strip("/")}]
@ -1129,6 +1065,11 @@ def custom_openapi():
openapi_schema = CustomOpenAPISpec.add_llm_api_request_schema_body(openapi_schema)
# Stub unloaded lazy features so they appear as Swagger sections.
from litellm.proxy._lazy_features import inject_lazy_stubs
openapi_schema = inject_lazy_stubs(openapi_schema)
# Fix Swagger UI execute path error when server_root_path is set
if server_root_path:
openapi_schema["servers"] = [{"url": "/" + server_root_path.strip("/")}]
@ -1556,14 +1497,78 @@ def mount_swagger_ui():
app.mount("/swagger", StaticFiles(directory=swagger_directory), name="swagger")
# On dropdown expand: one-time fetch to the prefix (triggers lazy load),
# then spec re-download so real routes replace the stub. Raw JS (no
# <script> tag) since it's injected inside the existing inline script.
from fastapi.responses import HTMLResponse
from litellm.proxy._lazy_features import lazy_tag_to_prefix
_lazy_plugin_js = (
"const TAG_TO_PREFIX = " + json.dumps(lazy_tag_to_prefix()) + ";"
"const warmedTags = new Set();"
"const LAZY_TAGS = new Set(Object.keys(TAG_TO_PREFIX));"
"const hideStubRows = () => {"
"document.querySelectorAll('.opblock').forEach(op => {"
"const d = op.querySelector('.opblock-summary-description');"
"if (d && LAZY_TAGS.has(d.textContent.trim())) op.style.display = 'none';"
"});};"
"const annotateLazyHeaders = () => {"
"document.querySelectorAll('.opblock-tag').forEach(tagEl => {"
"const m = (tagEl.id || '').match(/^operations-tag-(.+)$/);"
"if (!m || !LAZY_TAGS.has(m[1])) return;"
"const existing = tagEl.querySelector('.lazy-load-hint');"
"if (warmedTags.has(m[1])) { if (existing) existing.remove(); return; }"
"if (existing) return;"
"const hint = document.createElement('small');"
"hint.className = 'lazy-load-hint';"
"hint.textContent = ' (expand to load routes)';"
"hint.style.opacity = '0.6';"
"hint.style.marginLeft = '6px';"
"const target = tagEl.querySelector('a span') || tagEl.querySelector('span') || tagEl;"
"target.appendChild(hint);"
"});};"
"setInterval(() => { hideStubRows(); annotateLazyHeaders(); }, 200);"
"const LazyLoadPlugin = () => ({"
"afterLoad:function(system){setTimeout(()=>{"
"for(const tag of LAZY_TAGS)system.layoutActions.show(['operations-tag',tag],false);"
"},200);},"
"statePlugins:{layout:{wrapActions:{show:(ori,sys)=>(...args)=>{"
"const thing=args[0];const shown=args[1];let tag=null;"
"if(Array.isArray(thing)){for(const t of thing)if(TAG_TO_PREFIX[t])tag=t;}"
"if(shown!==false&&tag&&!warmedTags.has(tag)){warmedTags.add(tag);"
"fetch('/lazy/warm/'+tag,{method:'POST',credentials:'include'}).then(r=>r.json()).then(d=>{"
"if(!d.paths||Object.keys(d.paths).length===0)return;"
"const cur=sys.specSelectors.specJson().toJS();"
"const merged={};let inserted=false;"
"for(const k in (cur.paths||{})){"
"if(k===d.stub_path){for(const nk in d.paths)merged[nk]=d.paths[nk];inserted=true;}"
"else{merged[k]=cur.paths[k];}}"
"if(!inserted)Object.assign(merged,d.paths);"
"cur.paths=merged;"
"cur.components=cur.components||{};"
"cur.components.schemas=Object.assign(cur.components.schemas||{},(d.components||{}).schemas||{});"
"sys.specActions.updateSpec(JSON.stringify(cur));"
"}).catch(()=>{});}"
"return ori(...args);}}}}});"
)
def swagger_monkey_patch(*args, **kwargs):
return get_swagger_ui_html(
response = get_swagger_ui_html(
*args,
**kwargs,
swagger_js_url=f"{custom_root_path_swagger_path}/swagger-ui-bundle.js",
swagger_css_url=f"{custom_root_path_swagger_path}/swagger-ui.css",
swagger_favicon_url=f"{custom_root_path_swagger_path}/favicon.png",
)
body = response.body.decode("utf-8")
body = body.replace(
"const ui = SwaggerUIBundle({",
_lazy_plugin_js
+ 'const ui = SwaggerUIBundle({plugins:[LazyLoadPlugin],tagsSorter:"alpha",',
1,
)
return HTMLResponse(content=body)
applications.get_swagger_ui_html = swagger_monkey_patch
@ -3864,11 +3869,19 @@ class ProxyConfig:
## MCP TOOLS
mcp_tools_config = config.get("mcp_tools", None)
if mcp_tools_config:
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
global_mcp_tool_registry.load_tools_from_config(mcp_tools_config)
## AGENTS
agent_config = config.get("agent_list", None)
if agent_config:
from litellm.proxy.agent_endpoints.agent_registry import (
global_agent_registry,
)
global_agent_registry.load_agents_from_config(agent_config) # type: ignore
mcp_servers_config = config.get("mcp_servers", None)
@ -10586,6 +10599,10 @@ async def model_info_v2(
verbose_proxy_logger.debug("all_models: %s", all_models)
# Append A2A agents to models list
from litellm.proxy.agent_endpoints.model_list_helpers import (
append_agents_to_model_info,
)
all_models = await append_agents_to_model_info(
models=all_models,
user_api_key_dict=user_api_key_dict,
@ -11435,6 +11452,10 @@ async def model_group_info(
)
# Append A2A agents to model groups
from litellm.proxy.agent_endpoints.model_list_helpers import (
append_agents_to_model_group,
)
model_groups = await append_agents_to_model_group(
model_groups=model_groups,
user_api_key_dict=user_api_key_dict,
@ -12032,7 +12053,7 @@ async def onboarding(invite_link: str, request: Request):
"""
- Get the invite link
- Validate it's still 'valid'
- Invalidate the link (prevents abuse)
- Return a short-lived onboarding token
- Get user from db
- Pass in user_email if set
"""
@ -12070,7 +12091,7 @@ async def onboarding(invite_link: str, request: Request):
)
#### CHECK IF ALREADY USED
if invite_obj.is_accepted is True:
if invite_obj.is_accepted is True or invite_obj.accepted_at is not None:
raise HTTPException(
status_code=401,
detail={"error": "Invitation link has already been used."},
@ -12086,24 +12107,6 @@ async def onboarding(invite_link: str, request: Request):
status_code=401, detail={"error": "User does not exist in db."}
)
user_email = user_obj.user_email
response = await generate_key_helper_fn(
request_type="key",
**{
"user_role": user_obj.user_role,
"duration": LITELLM_UI_SESSION_DURATION,
"key_max_budget": litellm.max_ui_session_budget,
"models": [],
"aliases": {},
"config": {},
"spend": 0,
"user_id": user_obj.user_id,
"team_id": "litellm-dashboard",
}, # type: ignore
)
key = response["token"] # type: ignore
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):
litellm_dashboard_ui += "ui/onboarding"
@ -12111,13 +12114,24 @@ async def onboarding(invite_link: str, request: Request):
litellm_dashboard_ui += "/ui/onboarding"
import jwt
user_email = user_obj.user_email
onboarding_token = jwt.encode( # type: ignore
{
"token_type": "litellm_onboarding",
"invitation_link": invite_link,
"user_id": user_obj.user_id,
"exp": litellm.utils.get_utc_datetime() + timedelta(minutes=15),
},
master_key,
algorithm="HS256",
)
disabled_non_admin_personal_key_creation = (
get_disabled_non_admin_personal_key_creation()
)
returned_ui_token_object = ReturnedUITokenObject(
user_id=user_obj.user_id,
key=key,
key=onboarding_token,
user_email=user_obj.user_email,
user_role=user_obj.user_role,
login_method="username_password",
@ -12142,8 +12156,117 @@ async def onboarding(invite_link: str, request: Request):
}
def _get_onboarding_claims_from_request(request: Request) -> dict:
global master_key, general_settings
if master_key is None:
raise ProxyException(
message="Master Key not set for Proxy. Please set Master Key to use Admin UI. Set `LITELLM_MASTER_KEY` in .env or set general_settings:master_key in config.yaml. https://docs.litellm.ai/docs/proxy/virtual_keys. If set, use `--detailed_debug` to debug issue.",
type=ProxyErrorTypes.auth_error,
param="master_key",
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
auth_header_name = general_settings.get("litellm_key_header_name", "Authorization")
onboarding_auth_header = request.headers.get(auth_header_name)
if onboarding_auth_header is None:
raise HTTPException(
status_code=401,
detail={"error": "Missing onboarding session for invitation link."},
)
onboarding_token = onboarding_auth_header
if onboarding_token.lower().startswith("bearer "):
onboarding_token = onboarding_token.split(" ", 1)[1]
import jwt
try:
return jwt.decode(
onboarding_token,
master_key,
algorithms=["HS256"],
)
except Exception:
raise HTTPException(
status_code=401,
detail={"error": "Invalid onboarding session for invitation link."},
)
async def _rollback_onboarding_invite_claim(
invitation_link: str,
user_id: str,
) -> None:
global prisma_client
if prisma_client is None:
return
try:
await prisma_client.db.litellm_invitationlink.update_many(
where={"id": invitation_link, "is_accepted": True},
data={
"accepted_at": None,
"is_accepted": False,
"updated_at": litellm.utils.get_utc_datetime(),
"updated_by": user_id,
},
)
except Exception:
verbose_proxy_logger.exception(
"Failed to roll back onboarding invitation after session key mint failed."
)
async def _generate_onboarding_ui_session_token(user_obj: Any) -> str:
global master_key, general_settings
response = await generate_key_helper_fn(
request_type="key",
**{
"user_role": user_obj.user_role,
"duration": LITELLM_UI_SESSION_DURATION,
"key_max_budget": litellm.max_ui_session_budget,
"models": [],
"aliases": {},
"config": {},
"spend": 0,
"user_id": user_obj.user_id,
"team_id": UI_TEAM_ID,
}, # type: ignore
)
key = response["token"] # type: ignore
from litellm.types.proxy.ui_sso import ReturnedUITokenObject
import jwt
disabled_non_admin_personal_key_creation = (
get_disabled_non_admin_personal_key_creation()
)
returned_ui_token_object = ReturnedUITokenObject(
user_id=user_obj.user_id,
key=key,
user_email=user_obj.user_email,
user_role=user_obj.user_role,
login_method="username_password",
premium_user=premium_user,
auth_header_name=general_settings.get(
"litellm_key_header_name", "Authorization"
),
disabled_non_admin_personal_key_creation=disabled_non_admin_personal_key_creation,
server_root_path=get_server_root_path(),
)
assert master_key is not None
return jwt.encode( # type: ignore
cast(dict, returned_ui_token_object),
master_key,
algorithm="HS256",
)
@app.post("/onboarding/claim_token", include_in_schema=False)
async def claim_onboarding_link(data: InvitationClaim):
async def claim_onboarding_link(data: InvitationClaim, request: Request):
"""
Special route. Allows UI link share user to update their password.
@ -12155,7 +12278,7 @@ async def claim_onboarding_link(data: InvitationClaim):
This route can only update user password.
"""
global prisma_client
global prisma_client, master_key, general_settings
### VALIDATE INVITE LINK ###
if prisma_client is None:
raise HTTPException(
@ -12180,7 +12303,7 @@ async def claim_onboarding_link(data: InvitationClaim):
)
#### CHECK IF ALREADY USED
if invite_obj.is_accepted is True:
if invite_obj.is_accepted is True or invite_obj.accepted_at is not None:
raise HTTPException(
status_code=401,
detail={"error": "Invitation link has already been used."},
@ -12196,39 +12319,105 @@ async def claim_onboarding_link(data: InvitationClaim):
)
},
)
### UPDATE USER OBJECT ###
hashed_pw = hash_password(data.password)
user_obj = await prisma_client.db.litellm_usertable.update(
where={"user_id": invite_obj.user_id}, data={"password": hashed_pw}
)
if user_obj is None:
onboarding_claims = _get_onboarding_claims_from_request(request=request)
if (
onboarding_claims.get("token_type") != "litellm_onboarding"
or onboarding_claims.get("invitation_link") != data.invitation_link
or onboarding_claims.get("user_id") != data.user_id
):
raise HTTPException(
status_code=401, detail={"error": "User does not exist in db."}
status_code=401,
detail={"error": "Invalid onboarding session for invitation link."},
)
#### MARK LINK AS USED
hashed_pw = hash_password(data.password)
current_time = litellm.utils.get_utc_datetime()
await prisma_client.db.litellm_invitationlink.update(
where={"id": data.invitation_link},
data={
"accepted_at": current_time,
"updated_at": current_time,
"is_accepted": True,
"updated_by": invite_obj.user_id, # type: ignore
},
)
async with prisma_client.db.tx() as tx:
updated_count = await tx.litellm_invitationlink.update_many(
where={"id": data.invitation_link, "is_accepted": False},
data={
"is_accepted": True,
"updated_at": current_time,
"updated_by": invite_obj.user_id, # type: ignore
},
)
if updated_count == 0:
raise HTTPException(
status_code=401,
detail={"error": "Invitation link has already been used."},
)
### UPDATE USER OBJECT ###
user_obj = await tx.litellm_usertable.update(
where={"user_id": invite_obj.user_id}, data={"password": hashed_pw}
)
if user_obj is None:
raise HTTPException(
status_code=401, detail={"error": "User does not exist in db."}
)
#### MARK LINK AS USED
current_time = litellm.utils.get_utc_datetime()
await tx.litellm_invitationlink.update(
where={"id": data.invitation_link},
data={
"accepted_at": current_time,
"updated_at": current_time,
"updated_by": invite_obj.user_id, # type: ignore
},
)
if user_obj and hasattr(user_obj, "__dict__"):
user_obj.__dict__.pop("password", None)
return user_obj
try:
jwt_token = await _generate_onboarding_ui_session_token(user_obj=user_obj)
except Exception as e:
await _rollback_onboarding_invite_claim(
invitation_link=data.invitation_link,
user_id=data.user_id,
)
if isinstance(e, HTTPException):
raise e
raise HTTPException(
status_code=500,
detail={
"error": "Failed to create onboarding session. Please retry the invitation link."
},
) from e
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):
litellm_dashboard_ui += "ui/"
else:
litellm_dashboard_ui += "/ui/"
litellm_dashboard_ui += "?login=success"
return {
"login_url": litellm_dashboard_ui,
"token": jwt_token,
"user_email": user_obj.user_email,
"user": user_obj,
}
@app.get("/get_logo_url", include_in_schema=False)
def get_logo_url():
"""Get the current logo URL from environment"""
"""Get the current logo URL from environment.
Only HTTP(S) URLs are returned — those are intended to be loaded
directly by the browser from a public/internal CDN. Local file
paths set via ``UI_LOGO_PATH`` are NOT returned: they are admin-
only filesystem details, the dashboard falls back to ``/get_image``
which serves the file only when it is a supported image. Without
this filter, the unauthenticated endpoint would disclose internal
hostnames or filesystem paths to any caller.
"""
logo_path = os.getenv("UI_LOGO_PATH", "")
return {"logo_url": logo_path}
if logo_path.startswith(("http://", "https://")):
return {"logo_url": logo_path}
return {"logo_url": ""}
@app.get("/get_image", include_in_schema=False)
@ -12267,61 +12456,44 @@ async def get_image():
if assets_dir != current_dir and not os.path.exists(default_logo):
default_logo = default_site_logo
cache_dir = assets_dir if os.access(assets_dir, os.W_OK) else current_dir
cache_path = os.path.join(cache_dir, "cached_logo.jpg")
logo_path = os.getenv("UI_LOGO_PATH", default_logo)
verbose_proxy_logger.debug("Reading logo from path: %s", logo_path)
# If UI_LOGO_PATH points to a local file, serve it directly (skip cache)
from litellm.proxy.common_utils.static_asset_utils import (
resolve_validated_local_image_path,
)
if logo_path != default_logo and not logo_path.startswith(("http://", "https://")):
if os.path.exists(logo_path):
return FileResponse(logo_path, media_type="image/jpeg")
# Custom path doesn't exist — fall back to default
safe_logo = resolve_validated_local_image_path(logo_path)
if safe_logo is not None:
safe_logo_path, media_type = safe_logo
return FileResponse(safe_logo_path, media_type=media_type)
verbose_proxy_logger.warning(
f"UI_LOGO_PATH '{logo_path}' does not exist, falling back to default logo"
"UI_LOGO_PATH %r is not a supported image file or does not exist, "
"falling back to default logo",
logo_path,
)
logo_path = default_logo
# [OPTIMIZATION] For HTTP URLs and default logo, check if the cached image exists
if os.path.exists(cache_path):
return FileResponse(cache_path, media_type="image/jpeg")
# Check if the logo path is an HTTP/HTTPS URL
# Remote logo URLs are loaded by the browser. The proxy should not fetch
# arbitrary admin-configured URLs server-side.
if logo_path.startswith(("http://", "https://")):
try:
# Download the image and cache it
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
return RedirectResponse(url=logo_path)
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.UI,
params={"timeout": 5.0},
)
response = await async_client.get(logo_path)
if response.status_code == 200:
# Save the image to a local file
with open(cache_path, "wb") as f:
f.write(response.content)
# Return the cached image as a FileResponse
return FileResponse(cache_path, media_type="image/jpeg")
else:
# Handle the case when the image cannot be downloaded
return FileResponse(default_logo, media_type="image/jpeg")
except Exception as e:
# Handle any exceptions during the download (e.g., timeout, connection error)
verbose_proxy_logger.debug(f"Error downloading logo from {logo_path}: {e}")
return FileResponse(default_logo, media_type="image/jpeg")
else:
# Return the local image file if the logo path is not an HTTP/HTTPS URL
return FileResponse(logo_path, media_type="image/jpeg")
# Default logo (resolved from the bundled asset, not user-controlled).
safe_logo = resolve_validated_local_image_path(logo_path)
if safe_logo is not None:
safe_logo_path, media_type = safe_logo
return FileResponse(safe_logo_path, media_type=media_type)
return FileResponse(default_site_logo, media_type="image/jpeg")
@app.get("/get_favicon", include_in_schema=False)
async def get_favicon():
"""Get custom favicon for the admin UI."""
from fastapi.responses import Response
from litellm.proxy.common_utils.static_asset_utils import (
resolve_validated_local_image_path,
)
current_dir = os.path.dirname(os.path.abspath(__file__))
default_favicon = os.path.join(current_dir, "_experimental", "out", "favicon.ico")
@ -12334,42 +12506,17 @@ async def get_favicon():
raise HTTPException(status_code=404, detail="Default favicon not found")
if favicon_url.startswith(("http://", "https://")):
try:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.UI,
params={"timeout": 5.0},
)
response = await async_client.get(favicon_url)
if response.status_code == 200:
content_type = response.headers.get("content-type", "image/x-icon")
return Response(
content=response.content,
media_type=content_type,
)
else:
verbose_proxy_logger.warning(
"Failed to fetch favicon from %s: status %s",
favicon_url,
response.status_code,
)
if os.path.exists(default_favicon):
return FileResponse(default_favicon, media_type="image/x-icon")
raise HTTPException(status_code=404, detail="Favicon not found")
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.debug(
"Error downloading favicon from %s: %s", favicon_url, e
)
if os.path.exists(default_favicon):
return FileResponse(default_favicon, media_type="image/x-icon")
raise HTTPException(status_code=404, detail="Favicon not found")
return RedirectResponse(url=favicon_url)
else:
if os.path.exists(favicon_url):
return FileResponse(favicon_url, media_type="image/x-icon")
safe_favicon = resolve_validated_local_image_path(favicon_url)
if safe_favicon is not None:
safe_favicon_path, media_type = safe_favicon
return FileResponse(safe_favicon_path, media_type=media_type)
verbose_proxy_logger.warning(
"LITELLM_FAVICON_URL %r is not a supported image file or does not "
"exist, falling back to default favicon",
favicon_url,
)
if os.path.exists(default_favicon):
return FileResponse(default_favicon, media_type="image/x-icon")
raise HTTPException(status_code=404, detail="Favicon not found")
@ -14240,66 +14387,41 @@ app.include_router(container_router)
app.include_router(search_router)
app.include_router(image_router)
app.include_router(fine_tuning_router)
app.include_router(vector_store_router)
app.include_router(vector_store_management_router)
app.include_router(vector_store_files_router)
app.include_router(credential_router)
app.include_router(llm_passthrough_router)
app.include_router(webrtc_router)
app.include_router(mcp_management_router)
app.include_router(mcp_byok_oauth_router)
app.include_router(anthropic_router)
app.include_router(anthropic_skills_router)
app.include_router(evals_router)
app.include_router(claude_code_marketplace_router)
app.include_router(google_router)
app.include_router(langfuse_router)
app.include_router(pass_through_router)
app.include_router(health_router)
app.include_router(key_management_router)
app.include_router(internal_user_router)
app.include_router(team_router)
app.include_router(ui_sso_router)
app.include_router(scim_router)
app.include_router(organization_router)
app.include_router(customer_router)
app.include_router(spend_management_router)
app.include_router(cloudzero_router)
app.include_router(vantage_router)
app.include_router(caching_router)
app.include_router(analytics_router)
app.include_router(guardrails_router)
app.include_router(policy_router)
app.include_router(usage_ai_router)
app.include_router(policy_crud_router)
app.include_router(policy_resolve_router)
app.include_router(search_tool_management_router)
app.include_router(prompts_router)
app.include_router(callback_management_endpoints_router)
app.include_router(debugging_endpoints_router)
app.include_router(ui_crud_endpoints_router)
app.include_router(openai_files_router)
app.include_router(team_callback_router)
app.include_router(jwt_key_mapping_router)
app.include_router(budget_management_router)
app.include_router(model_management_router)
app.include_router(model_access_group_management_router)
app.include_router(tag_management_router)
app.include_router(tool_management_router)
app.include_router(workflow_management_router)
app.include_router(memory_router)
app.include_router(cost_tracking_settings_router)
app.include_router(router_settings_router)
app.include_router(fallback_management_router)
app.include_router(cache_settings_router)
app.include_router(config_override_router)
app.include_router(user_agent_analytics_router)
app.include_router(enterprise_router)
app.include_router(ui_discovery_endpoints_router)
app.include_router(agent_endpoints_router)
app.include_router(compliance_router)
app.include_router(a2a_router)
app.include_router(access_group_router)
# Eager: /models/{name}:method overlaps with the OpenAI /models endpoint.
app.include_router(google_router)
attach_lazy_features(app)
async def _stream_mcp_asgi_response(
@ -14532,8 +14654,3 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request):
f"Error handling dynamic MCP route for {mcp_server_name}: {str(e)}"
)
raise HTTPException(status_code=500, detail=f"Internal server error: {str(e)}")
app.mount(path=BASE_MCP_ROUTE, app=mcp_app)
app.include_router(mcp_rest_endpoints_router)
app.include_router(mcp_discoverable_endpoints_router)

View file

@ -277,6 +277,7 @@ model LiteLLM_ObjectPermissionTable {
models String[] @default([])
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -134,10 +134,48 @@ async def search(
if "search_tool_name" in data and data["search_tool_name"]:
data["model"] = data["search_tool_name"]
search_tool_name_value = data["search_tool_name"]
# Authorization check: verify key can access this search tool
from litellm.proxy.auth.auth_checks import (
can_key_call_search_tool,
can_team_call_search_tool,
get_team_object,
)
try:
# Check key-level access
await can_key_call_search_tool(
search_tool_name=search_tool_name_value,
valid_token=user_api_key_dict,
)
# Check team-level access if key is associated with a team
if user_api_key_dict.team_id:
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
team_object = await get_team_object(
team_id=user_api_key_dict.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_dict.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
await can_team_call_search_tool(
search_tool_name=search_tool_name_value,
team_object=team_object,
)
except Exception as e:
verbose_proxy_logger.error(
f"Search tool authorization failed for {search_tool_name_value}: {str(e)}"
)
raise
if llm_router is not None and hasattr(llm_router, "search_tools"):
search_tool_name_value = data["search_tool_name"]
verbose_proxy_logger.debug(
f"Search endpoint - Looking for search_tool_name: {search_tool_name_value}. "
f"Available search tools in router: {[tool.get('search_tool_name') for tool in llm_router.search_tools]}. "
@ -163,6 +201,16 @@ async def search(
data["metadata"] = {}
data["metadata"]["model_group"] = search_tool_name_value
# Ensure team context is available to search router credential resolution.
# add_litellm_data_to_request() also injects these values, but this keeps
# search endpoint behavior explicit and resilient for direct router paths.
if "metadata" not in data or not isinstance(data.get("metadata"), dict):
data["metadata"] = {}
if getattr(user_api_key_dict, "team_metadata", None) is not None:
data["metadata"]["user_api_key_team_metadata"] = user_api_key_dict.team_metadata
if getattr(user_api_key_dict, "team_id", None) is not None:
data["metadata"]["user_api_key_team_id"] = user_api_key_dict.team_id
# Process request using ProxyBaseLLMRequestProcessing
processor = ProxyBaseLLMRequestProcessing(data=data)
try:

View file

@ -53,20 +53,13 @@ def _get_max_string_length_prompt_in_db() -> int:
def _is_master_key(api_key: Optional[str], _master_key: Optional[str]) -> bool:
"""
Raw-only constant-time master-key comparison. The hashed form is never
considered equivalent — only the raw master-key string matches.
"""
if _master_key is None or api_key is None:
return False
## string comparison
is_master_key = secrets.compare_digest(api_key, _master_key)
if is_master_key:
return True
## hash comparison
is_master_key = secrets.compare_digest(api_key, hash_token(_master_key))
if is_master_key:
return True
return False
return secrets.compare_digest(api_key, _master_key)
def _get_spend_logs_metadata(
@ -235,8 +228,6 @@ def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> d
def get_logging_payload( # noqa: PLR0915
kwargs, response_obj, start_time, end_time
) -> SpendLogsPayload:
from litellm.proxy.proxy_server import general_settings, master_key
if kwargs is None:
kwargs = {}
@ -295,11 +286,6 @@ def get_logging_payload( # noqa: PLR0915
if api_key.startswith("sk-"):
# hash the api_key
api_key = hash_token(api_key)
if (
_is_master_key(api_key=api_key, _master_key=master_key)
and general_settings.get("disable_adding_master_key_hash_to_db") is True
):
api_key = "litellm_proxy_master_key" # use a known alias, if the user disabled storing master key in db
if (
standard_logging_payload is not None
@ -324,11 +310,6 @@ def get_logging_payload( # noqa: PLR0915
and standard_logging_payload.get("request_tags") is not None
): # use 'tags' from standard logging payload instead
request_tags = json.dumps(standard_logging_payload["request_tags"])
if (
_is_master_key(api_key=api_key, _master_key=master_key)
and general_settings.get("disable_adding_master_key_hash_to_db") is True
):
api_key = "litellm_proxy_master_key" # use a known alias, if the user disabled storing master key in db
_model_id = metadata.get("model_info", {}).get("id", "")
_model_group = metadata.get("model_group", "")

View file

@ -16,7 +16,9 @@ from fastapi import APIRouter, Depends, HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy._types import (
LiteLLM_ManagedVectorStoresTable,
ResponseLiteLLM_ManagedVectorStore,
@ -38,6 +40,81 @@ from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
router = APIRouter()
_LITELLM_PARAMS_MASKER = SensitiveDataMasker()
_REDACT_LITELLM_PARAMS_MAX_DEPTH = 10
def _redact_sensitive_litellm_params(litellm_params: Any, _depth: int = 0) -> Any:
"""
Replace credential-bearing values in ``litellm_params`` with
``REDACTED_BY_LITELM`` while preserving non-secret keys (``api_base``,
``region``, ``model``, ``api_version``).
Handles three input shapes:
* ``dict`` — recurse into nested dicts (e.g. ``litellm_embedding_config``
which itself carries ``api_key`` / ``aws_*`` / ``vertex_credentials``).
* ``str`` — the in-memory registry occasionally holds the params as a
JSON-serialized string. Parse, redact, re-serialize. If parsing
fails, return the redaction sentinel rather than echo the value
back verbatim.
* Anything else, or ``None`` — passed through.
Recursion depth is bounded by ``_REDACT_LITELLM_PARAMS_MAX_DEPTH`` —
matching the convention of other allowlisted recursive helpers in the
repo (see ``tests/code_coverage_tests/recursive_detector.py``).
"""
if _depth >= _REDACT_LITELLM_PARAMS_MAX_DEPTH:
return REDACTED_BY_LITELM_STRING
if litellm_params is None:
return None
if isinstance(litellm_params, str):
try:
parsed = json.loads(litellm_params)
except (TypeError, ValueError):
return REDACTED_BY_LITELM_STRING
return json.dumps(_redact_sensitive_litellm_params(parsed, _depth + 1))
if not isinstance(litellm_params, dict):
return litellm_params
out: Dict[str, Any] = {}
for k, v in litellm_params.items():
if _LITELLM_PARAMS_MASKER.is_sensitive_key(k):
out[k] = REDACTED_BY_LITELM_STRING
elif isinstance(v, dict):
out[k] = _redact_sensitive_litellm_params(v, _depth + 1)
else:
out[k] = v
return out
async def _fetch_and_authorize_vector_store(
vector_store_id: str,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: Any,
) -> "LiteLLM_ManagedVectorStore":
"""
Look up a vector store by id and confirm the caller can access it.
Raises HTTPException(404) on miss and HTTPException(403) on access
denial.
"""
row = await prisma_client.db.litellm_managedvectorstorestable.find_unique(
where={"vector_store_id": vector_store_id}
)
if row is None:
raise HTTPException(
status_code=404,
detail=f"Vector store with ID {vector_store_id} not found",
)
typed = LiteLLM_ManagedVectorStore(**row.model_dump())
if not await _check_vector_store_access(typed, user_api_key_dict):
raise HTTPException(
status_code=403,
detail="Access denied: You do not have permission to access this vector store",
)
return typed
def _resolve_embedding_config_from_router(
embedding_model: str, llm_router
@ -555,7 +632,11 @@ async def list_vector_stores(
accessible_vector_stores = []
for vs in vector_store_map.values():
if await _check_vector_store_access(vs, user_api_key_dict):
accessible_vector_stores.append(vs)
redacted = LiteLLM_ManagedVectorStore(**vs)
redacted["litellm_params"] = _redact_sensitive_litellm_params(
vs.get("litellm_params")
)
accessible_vector_stores.append(redacted)
total_count = len(accessible_vector_stores)
total_pages = (total_count + page_size - 1) // page_size
@ -716,33 +797,29 @@ async def get_vector_store_info(
created_at=vector_store.get("created_at") or None,
updated_at=vector_store.get("updated_at") or None,
litellm_credential_name=vector_store.get("litellm_credential_name"),
litellm_params=vector_store.get("litellm_params") or None,
litellm_params=_redact_sensitive_litellm_params(
vector_store.get("litellm_params")
),
team_id=vector_store.get("team_id") or None,
user_id=vector_store.get("user_id") or None,
)
return {"vector_store": vector_store_pydantic_obj}
vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
where={"vector_store_id": data.vector_store_id}
)
vector_store_typed = await _fetch_and_authorize_vector_store(
vector_store_id=data.vector_store_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
if vector_store is None:
raise HTTPException(
status_code=404,
detail=f"Vector store with ID {data.vector_store_id} not found",
vector_store_dict = dict(vector_store_typed)
if "litellm_params" in vector_store_dict:
vector_store_dict["litellm_params"] = _redact_sensitive_litellm_params(
vector_store_dict["litellm_params"]
)
# Check access control for DB vector store
vector_store_dict = vector_store.model_dump() # type: ignore[attr-defined]
vector_store_typed = LiteLLM_ManagedVectorStore(**vector_store_dict)
if not await _check_vector_store_access(vector_store_typed, user_api_key_dict):
raise HTTPException(
status_code=403,
detail="Access denied: You do not have permission to access this vector store",
)
return {"vector_store": vector_store_dict}
except HTTPException:
# Preserve 403/404 from the access-control / not-found checks above;
# the catch-all below would otherwise rewrite them as 500.
raise
except Exception as e:
verbose_proxy_logger.exception(f"Error getting vector store info: {str(e)}")
raise HTTPException(status_code=500, detail=str(e))
@ -773,6 +850,15 @@ async def update_vector_store(
update_data = data.model_dump(exclude_unset=True)
vector_store_id = update_data.pop("vector_store_id")
# Per-store access control: anyone authenticated who passes the
# premium-feature gate could otherwise update *any* vector store —
# including stores belonging to other teams.
await _fetch_and_authorize_vector_store(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
# Handle metadata serialization
if update_data.get("vector_store_metadata") is not None:
update_data["vector_store_metadata"] = safe_dumps(
@ -820,11 +906,24 @@ async def update_vector_store(
f"Updated vector store {vector_store_id} in both database and in-memory registry"
)
# The DB row is returned in full, so the response would otherwise
# echo the persisted ``litellm_params`` (including provider
# credentials) back to the caller — even when the caller only
# changed unrelated fields like ``vector_store_description``.
response_vs = LiteLLM_ManagedVectorStore(**updated_vs)
response_vs["litellm_params"] = _redact_sensitive_litellm_params(
updated_vs.get("litellm_params")
)
return {
"status": "success",
"message": f"Vector store {vector_store_id} updated successfully",
"vector_store": updated_vs,
"vector_store": response_vs,
}
except HTTPException:
# Preserve 403/404 responses from the access-control / not-found
# checks above; the catch-all below would otherwise rewrite them
# as 500 with the original status code embedded in the detail.
raise
except Exception as e:
verbose_proxy_logger.exception(f"Error updating vector store: {str(e)}")
raise HTTPException(status_code=500, detail=str(e))

View file

@ -8,7 +8,7 @@ import asyncio
import random
import traceback
from functools import partial
from typing import Any, Callable
from typing import Any, Callable, Dict, Optional, Tuple
from litellm._logging import verbose_router_logger
@ -20,6 +20,28 @@ class SearchAPIRouter:
Provides methods for search tool selection, load balancing, and fallback handling.
"""
@staticmethod
def _resolve_search_provider_credentials(
*,
tool_litellm_params: Dict[str, Any],
) -> Tuple[Optional[str], Optional[str]]:
"""
Resolve search provider credentials from tool configuration ONLY.
Credentials are stored only in search_tool.litellm_params, never in team/key metadata.
This ensures secrets are not exposed in team/key API responses.
Args:
tool_litellm_params: Search tool litellm_params with credentials
Returns:
Tuple of (api_key, api_base) from tool configuration
"""
resolved_api_key: Optional[str] = tool_litellm_params.get("api_key")
resolved_api_base: Optional[str] = tool_litellm_params.get("api_base")
return resolved_api_key, resolved_api_base
@staticmethod
async def update_router_search_tools(router_instance: Any, search_tools: list):
"""
@ -198,14 +220,15 @@ class SearchAPIRouter:
# Extract search provider and other params from litellm_params
litellm_params = selected_tool.get("litellm_params", {})
search_provider = litellm_params.get("search_provider")
api_key = litellm_params.get("api_key")
api_base = litellm_params.get("api_base")
if not search_provider:
raise ValueError(
f"search_provider not found in litellm_params for search tool '{search_tool_name}'"
)
api_key, api_base = SearchAPIRouter._resolve_search_provider_credentials(
tool_litellm_params=litellm_params,
)
verbose_router_logger.debug(
f"Selected search tool with provider: {search_provider}"
)

View file

@ -330,7 +330,11 @@ class AnthropicMessagesToolResultParam(TypedDict, total=False):
content: Union[
str,
Iterable[
Union[AnthropicMessagesToolResultContent, AnthropicMessagesImageParam]
Union[
AnthropicMessagesToolResultContent,
AnthropicMessagesImageParam,
AnthropicMessagesDocumentParam,
]
],
]
cache_control: Optional[Union[dict, ChatCompletionCachedContent]]

View file

@ -237,6 +237,8 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
# Vector Store Params
vector_store_id: Optional[str] = None
milvus_text_field: Optional[str] = None
milvus_db_name: Optional[str] = None
milvus_partition_names: Optional[List[str]] = None
@model_validator(mode="before")
@classmethod

View file

@ -712,6 +712,7 @@
},
"anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
@ -735,6 +736,7 @@
},
"anthropic.claude-haiku-4-5@20251001": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
@ -955,6 +957,7 @@
},
"anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -982,6 +985,7 @@
},
"anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1011,6 +1015,7 @@
},
"global.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1040,6 +1045,7 @@
},
"us.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1127,6 +1133,7 @@
},
"anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1171,6 +1178,7 @@
},
"global.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",
@ -1201,6 +1209,7 @@
},
"us.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1291,6 +1300,7 @@
},
"anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
@ -1319,6 +1329,7 @@
},
"global.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
@ -1347,6 +1358,7 @@
},
"us.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
@ -1461,11 +1473,13 @@
},
"anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -17935,11 +17949,13 @@
},
"global.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -17996,6 +18012,7 @@
},
"global.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "bedrock_converse",
@ -30170,6 +30187,7 @@
},
"us.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 1.1e-06,
"litellm_provider": "bedrock_converse",
@ -30321,11 +30339,13 @@
},
"us.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"input_cost_per_token_above_200k_tokens": 6.6e-06,
"output_cost_per_token_above_200k_tokens": 2.475e-05,
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -30426,6 +30446,7 @@
},
"us.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -30453,6 +30474,7 @@
},
"global.anthropic.claude-opus-4-5-20251101-v1:0": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "bedrock_converse",

View file

@ -277,6 +277,7 @@ model LiteLLM_ObjectPermissionTable {
models String[] @default([])
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -557,6 +557,7 @@ async def test_avertex_batch_prediction(monkeypatch):
mock_get_response = MagicMock()
mock_get_response.json.return_value = mock_vertex_batch_response
mock_get_response.status_code = 200
mock_get_response.is_redirect = False
mock_get_response.raise_for_status.return_value = None
mock_get.return_value = mock_get_response

View file

@ -46,6 +46,7 @@ IGNORE_FUNCTIONS = [
"dict", # max depth set. _LiteLLMParamsDictView.dict() calls builtin dict(), not itself.
"_read_image_bytes", # max depth set.
"_get_masked_values", # max depth set (default 20) to prevent infinite recursion while masking nested sensitive config dicts.
"_redact_sensitive_litellm_params", # max depth set (default 10).
]

View file

@ -812,7 +812,7 @@ def test_redact_msgs_from_logs_with_dynamic_params():
# Assert redaction occurred
assert _redacted_response_obj.choices[0].message.content == "redacted-by-litellm"
# Test Case 3: standard_callback_dynamic_params does not override litellm.turn_off_message_logging
# Test Case 3: standard_callback_dynamic_params does not set turn_off_message_logging
# since litellm.turn_off_message_logging is True redaction should occur
standard_callback_dynamic_params = StandardCallbackDynamicParams()
litellm_logging_obj.model_call_details["standard_callback_dynamic_params"] = (

View file

@ -13,11 +13,13 @@ import logging
import time
from unittest.mock import AsyncMock, patch
import httpx
import pytest
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.responses.main import mock_responses_api_response
from litellm.types.utils import StandardLoggingPayload
@ -126,17 +128,10 @@ async def test_redaction_responses_api():
test_custom_logger = TestCustomLogger(turn_off_message_logging=True)
litellm.callbacks = [test_custom_logger]
# Mock a ResponsesAPIResponse-style response
mock_response = {
"output": [{"text": "This is a test response"}],
"model": "gpt-3.5-turbo",
"usage": {"input_tokens": 5, "output_tokens": 5, "total_tokens": 10},
}
response = await litellm.aresponses(
model="gpt-3.5-turbo",
input="hi",
mock_response=mock_response,
mock_response="This is a test response",
)
await asyncio.sleep(1)
@ -163,6 +158,7 @@ async def test_redaction_responses_api():
assert (
content_item["text"] == "redacted-by-litellm"
), f"Expected redacted text but got: {content_item['text']}"
assert "This is a test response" not in json.dumps(standard_logging_payload)
print(
"logged standard logging payload for ResponsesAPIResponse",
json.dumps(standard_logging_payload, indent=2),
@ -176,29 +172,36 @@ async def test_redaction_responses_api_stream():
test_custom_logger = TestCustomLogger(turn_off_message_logging=True)
litellm.callbacks = [test_custom_logger]
# Mock a ResponsesAPIResponse-style response with streaming chunks
mock_response = [
{
"output": [{"text": "This"}],
"model": "gpt-3.5-turbo",
},
{
"output": [{"text": " is"}],
"model": "gpt-3.5-turbo",
},
{
"output": [{"text": " a test response"}],
"model": "gpt-3.5-turbo",
"usage": {"input_tokens": 5, "output_tokens": 5, "total_tokens": 10},
},
]
mocked_response_payload = mock_responses_api_response(
"This is a test response"
).model_dump()
response = await litellm.aresponses(
model="gpt-3.5-turbo",
input="hi",
mock_response=mock_response,
stream=True,
)
async def mock_post(self, url, headers, timeout, stream=False, **kwargs):
stream_content = (
"data: "
+ json.dumps(
{
"type": "response.completed",
"response": mocked_response_payload,
}
)
+ "\n\ndata: [DONE]\n\n"
)
return httpx.Response(
status_code=200,
content=stream_content,
request=httpx.Request("POST", url),
)
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=mock_post,
):
response = await litellm.aresponses(
model="gpt-3.5-turbo",
input="hi",
stream=True,
)
# Consume the stream
chunks = []
@ -445,18 +448,11 @@ async def test_disable_redaction_header_responses_api():
test_custom_logger = TestCustomLogger()
litellm.callbacks = [test_custom_logger]
# Mock a ResponsesAPIResponse-style response
mock_response = {
"output": [{"text": "This is a test response"}],
"model": "gpt-3.5-turbo",
"usage": {"input_tokens": 5, "output_tokens": 5, "total_tokens": 10},
}
# Pass the header via litellm_metadata (as the proxy does for Responses API)
response = await litellm.aresponses(
model="gpt-3.5-turbo",
input="hi",
mock_response=mock_response,
mock_response="This is a test response",
litellm_metadata={"headers": {"litellm-disable-message-redaction": "true"}},
)
@ -464,14 +460,14 @@ async def test_disable_redaction_header_responses_api():
standard_logging_payload = test_custom_logger.logged_standard_logging_payload
assert standard_logging_payload is not None
# Verify that messages are NOT redacted because the header was set
# Verify that the direct SDK path still honors the explicit header.
print(
"logged standard logging payload for ResponsesAPI with disable header",
json.dumps(standard_logging_payload, indent=2, default=str),
)
# The content should NOT be redacted
assert standard_logging_payload["response"] != {"text": "redacted-by-litellm"}
response = standard_logging_payload["response"]
assert response["output"][0]["content"][0]["text"] == "This is a test response"
assert standard_logging_payload["messages"][0]["content"] == "hi"

View file

@ -6,13 +6,19 @@ from httpx import AsyncClient
from typing import Any, Optional, List, Literal
# The proxy strips client-supplied `mock_response` unless the calling key or
# team has this admin-metadata flag set. See `_UNTRUSTED_ROOT_CONTROL_FIELDS`
# in litellm/proxy/litellm_pre_call_utils.py.
_ALLOW_CLIENT_MOCK_METADATA = {"allow_client_mock_response": True}
async def generate_key(
session, models: Optional[List[str]] = None, team_id: Optional[str] = None
):
"""Helper function to generate a key with specific model access controls"""
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {}
data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)}
if models is not None:
data["models"] = models
if team_id is not None:
@ -25,7 +31,7 @@ async def generate_team(session, models: Optional[List[str]] = None):
"""Helper function to generate a team with specific model access"""
url = "http://0.0.0.0:4000/team/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {}
data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)}
if models is not None:
data["models"] = models
async with session.post(url, headers=headers, json=data) as response:
@ -111,7 +117,12 @@ async def test_model_access_update():
# Create initial key with restricted access
response = await client.post(
"/key/generate", json={"models": ["openai/gpt-4"]}, headers=headers
"/key/generate",
json={
"models": ["openai/gpt-4"],
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
key_data = response.json()
@ -214,7 +225,11 @@ async def test_team_model_access_update():
# Create initial team with restricted access
response = await client.post(
"/team/new",
json={"models": ["openai/gpt-4"], "name": "test-team"},
json={
"models": ["openai/gpt-4"],
"name": "test-team",
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
@ -223,7 +238,12 @@ async def test_team_model_access_update():
# Generate a key for this team
response = await client.post(
"/key/generate", json={"team_id": team_id}, headers=headers
"/key/generate",
json={
"team_id": team_id,
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
key = response.json()["key"]

View file

@ -115,6 +115,13 @@ async def test_proxy_failure_metrics():
"litellm_llm_api_failed_requests_metric_total{", # Deprecated but may still be used
]
# Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for
# hash_token(master_key) so the master key (or its hash) never
# propagates into metrics. See PR #26484.
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS
# Check if either pattern is in metrics and contains required fields
found_metric = False
for pattern in expected_patterns:
@ -125,8 +132,7 @@ async def test_proxy_failure_metrics():
'api_key_alias="None"' in line
and 'exception_class="Openai.RateLimitError"' in line
and 'exception_status="429"' in line
and 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"'
in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and 'route="/chat/completions"' in line
):
@ -135,8 +141,7 @@ async def test_proxy_failure_metrics():
# For deprecated llm_api metric, check llm-specific fields
elif "litellm_llm_api_failed_requests_metric_total{" in line:
if (
'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"'
in line
f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'model="429"' in line
): # The deprecated metric uses the actual model from the request
found_metric = True
@ -156,8 +161,7 @@ async def test_proxy_failure_metrics():
for line in metrics.split("\n"):
if (
total_requests_pattern in line
and 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"'
in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and 'status_code="429"' in line
):
@ -195,6 +199,12 @@ async def test_proxy_success_metrics():
assert END_USER_ID not in metrics
# Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for
# hash_token(master_key) (PR #26484).
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS
# Check if the success metric is present and correct - use flexible matching
# Check for request_total_latency_metric with required fields
# Note: The model can be "gpt-3.5-turbo-0301" or similar depending on what's returned
@ -203,8 +213,7 @@ async def test_proxy_success_metrics():
if (
"litellm_request_total_latency_metric_bucket{" in line
and 'api_key_alias="None"' in line
and 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"'
in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-openai-endpoint"' in line
and 'le="0.005"' in line
):
@ -221,8 +230,7 @@ async def test_proxy_success_metrics():
if (
"litellm_llm_api_latency_metric_bucket{" in line
and 'api_key_alias="None"' in line
and 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"'
in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-openai-endpoint"' in line
and 'le="0.005"' in line
):
@ -298,6 +306,12 @@ async def test_proxy_fallback_metrics():
print("/metrics", metrics)
# Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for
# hash_token(master_key) (PR #26484).
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS
# Check if successful fallback metric is incremented - use flexible matching
found_successful_fallback = False
for line in metrics.split("\n"):
@ -307,8 +321,7 @@ async def test_proxy_fallback_metrics():
and 'exception_class="Openai.RateLimitError"' in line
and 'exception_status="429"' in line
and 'fallback_model="fake-openai-endpoint"' in line
and 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"'
in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and "1.0" in line
):
@ -328,8 +341,7 @@ async def test_proxy_fallback_metrics():
and 'exception_class="Openai.RateLimitError"' in line
and 'exception_status="429"' in line
and 'fallback_model="unknown-model"' in line
and 'hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"'
in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and "1.0" in line
):

View file

@ -2,4 +2,4 @@
# Packages needing lifecycle scripts: npm rebuild <pkg>
ignore-scripts=true
# Protects local npm install only — npm ci (used in CI) ignores this
min-release-age=3d
min-release-age=3

View file

@ -2,4 +2,4 @@
# Packages needing lifecycle scripts: npm rebuild <pkg>
ignore-scripts=true
# Protects local npm install only — npm ci (used in CI) ignores this
min-release-age=3d
min-release-age=3

View file

@ -45,8 +45,11 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG)
from starlette.datastructures import URL
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
)
from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, UserAPIKeyAuth
from litellm.caching.caching import DualCache
from unittest.mock import patch, AsyncMock
@ -54,6 +57,119 @@ proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
import json
def test_get_audit_log_changed_by_prefers_authenticated_user():
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="authenticated-user",
)
assert (
get_audit_log_changed_by(
litellm_changed_by="spoofed-user",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="proxy-admin",
)
== "authenticated-user"
)
def test_get_audit_log_changed_by_honors_header_with_admin_opt_in():
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="service-account",
metadata={"allow_litellm_changed_by_header": True},
)
assert (
get_audit_log_changed_by(
litellm_changed_by="delegated-user",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="proxy-admin",
)
== "delegated-user"
)
def test_get_audit_log_changed_by_honors_header_with_team_opt_in():
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="service-account",
team_metadata={"allow_litellm_changed_by_header": True},
)
assert (
get_audit_log_changed_by(
litellm_changed_by="delegated-user",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="proxy-admin",
)
== "delegated-user"
)
def test_get_audit_log_changed_by_ignores_header_without_opt_in_when_user_id_missing():
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
assert (
get_audit_log_changed_by(
litellm_changed_by="spoofed-user",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="proxy-admin",
)
== "proxy-admin"
)
def test_get_audit_log_changed_by_honors_header_with_opt_in_when_user_id_missing():
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={"allow_litellm_changed_by_header": True},
)
assert (
get_audit_log_changed_by(
litellm_changed_by="delegated-user",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="proxy-admin",
)
== "delegated-user"
)
@pytest.mark.asyncio
async def test_create_internal_user_audit_log_uses_changed_by_helper():
from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
user_id="service-account",
metadata={"allow_litellm_changed_by_header": True},
)
with (
patch("litellm.store_audit_logs", True),
patch(
"litellm.proxy.hooks.user_management_event_hooks.create_audit_log_for_update",
new_callable=AsyncMock,
) as mock_create_audit_log_for_update,
):
await UserManagementEventHooks.create_internal_user_audit_log(
user_id="target-user",
action="updated",
litellm_changed_by="delegated-user",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="proxy-admin",
before_value='{"before": true}',
after_value='{"after": true}',
)
request_data = mock_create_audit_log_for_update.await_args.kwargs["request_data"]
assert request_data.changed_by == "delegated-user"
assert request_data.changed_by_api_key == "test-key"
assert request_data.object_id == "target-user"
assert request_data.action == "updated"
@pytest.mark.asyncio
async def test_create_audit_log_for_update_premium_user():
"""

View file

@ -1,6 +1,5 @@
import os
import sys
from unittest import mock
sys.path.insert(0, os.path.abspath("../.."))
@ -26,50 +25,30 @@ async def test_get_favicon_default():
@pytest.mark.asyncio
async def test_get_favicon_with_custom_url():
"""Test that get_favicon fetches from a custom URL."""
os.environ["LITELLM_FAVICON_URL"] = "https://example.com/favicon.ico"
async def test_get_favicon_with_custom_url(monkeypatch):
"""Test that get_favicon redirects browser-loaded custom URLs."""
monkeypatch.setenv("LITELLM_FAVICON_URL", "https://example.com/favicon.ico")
mock_response = mock.Mock()
mock_response.status_code = 200
mock_response.content = b"\x00\x00\x01\x00"
mock_response.headers = {"content-type": "image/x-icon"}
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
try:
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.return_value = mock_response
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
assert response.status_code == 200
assert response.headers["content-type"] == "image/x-icon"
finally:
os.environ.pop("LITELLM_FAVICON_URL", None)
assert response.status_code == 307
assert response.headers["location"] == "https://example.com/favicon.ico"
@pytest.mark.asyncio
async def test_get_favicon_url_error_fallback():
"""Test that get_favicon falls back to default on error."""
os.environ["LITELLM_FAVICON_URL"] = "https://invalid.com/favicon.ico"
async def test_get_favicon_remote_url_is_not_server_fetched(monkeypatch):
"""Test that get_favicon does not validate remote URLs server-side."""
monkeypatch.setenv("LITELLM_FAVICON_URL", "https://invalid.com/favicon.ico")
try:
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.side_effect = httpx.ConnectError("unreachable")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
assert response.status_code in [200, 404]
finally:
os.environ.pop("LITELLM_FAVICON_URL", None)
assert response.status_code == 307
assert response.headers["location"] == "https://invalid.com/favicon.ico"

View file

@ -5,85 +5,48 @@ from unittest import mock
# Standard path insertion
sys.path.insert(0, os.path.abspath("../.."))
import pytest
import httpx
import pytest
from litellm.proxy.proxy_server import app
@pytest.mark.asyncio
async def test_get_image_error_handling():
async def test_get_image_redirects_remote_logo_without_server_fetch(monkeypatch):
"""
Test that get_image handles network errors gracefully and doesn't hang.
Remote logo URLs should be loaded by the browser, not fetched by the proxy.
"""
# Set an unreachable URL
os.environ["UI_LOGO_PATH"] = "http://invalid-url-12345.com/logo.jpg"
monkeypatch.setenv("UI_LOGO_PATH", "http://invalid-url-12345.com/logo.jpg")
# Clear cache
parent_dir = os.path.dirname(
os.path.dirname(
app.__file__
if hasattr(app, "__file__")
else "litellm/proxy/proxy_server.py"
)
)
cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg")
if os.path.exists(cache_path):
os.remove(cache_path)
# Mock AsyncHTTPHandler to simulate a timeout or connection error
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.side_effect = httpx.ConnectError("Network is unreachable")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://testserver"
) as ac:
response = await ac.get("/get_image")
assert response.status_code == 200
assert response.headers["content-type"] == "image/jpeg"
assert response.status_code == 307
assert response.headers["location"] == "http://invalid-url-12345.com/logo.jpg"
mock_get.assert_not_called()
@pytest.mark.asyncio
async def test_get_image_cache_logic():
async def test_get_image_remote_logo_does_not_use_stale_cache(monkeypatch, tmp_path):
"""
Test that once cached, get_image doesn't hit the network.
A stale pre-fix cache file should not mask a configured remote logo URL.
"""
os.environ["UI_LOGO_PATH"] = "http://example.com/logo.jpg"
# Clear cache
parent_dir = os.path.dirname(
os.path.dirname(
app.__file__
if hasattr(app, "__file__")
else "litellm/proxy/proxy_server.py"
)
)
cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg")
if os.path.exists(cache_path):
os.remove(cache_path)
# Mock response
mock_response = mock.Mock()
mock_response.status_code = 200
mock_response.content = b"fake image data"
monkeypatch.setenv("UI_LOGO_PATH", "http://example.com/logo.jpg")
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
(tmp_path / "cached_logo.jpg").write_bytes(b"\xff\xd8\xff cached logo")
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.return_value = mock_response
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://testserver"
) as ac:
# First call - should hit download logic
response1 = await ac.get("/get_image")
assert response1.status_code == 200
assert mock_get.call_count == 1
response = await ac.get("/get_image")
# Second call - should hit cache
response2 = await ac.get("/get_image")
assert response2.status_code == 200
# If cache works, mock_get shouldn't be called again
assert mock_get.call_count == 1
assert response.status_code == 307
assert response.headers["location"] == "http://example.com/logo.jpg"
mock_get.assert_not_called()

View file

@ -2715,7 +2715,12 @@ async def test_master_key_hashing(prisma_client):
request=request, api_key=bearer_token
)
assert result.api_key == hash_token(master_key)
# Master-key auth substitutes a stable alias so the master key (or
# its hash) never propagates into spend logs / metrics / audit trails.
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
assert result.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS
assert result.api_key != hash_token(master_key)
except Exception as e:
print("Got Exception", e)

View file

@ -39,6 +39,23 @@ def test_routes_on_litellm_proxy():
this prevents accidentelly deleting /threads, or /batches etc
"""
# Force-load lazy features so the test sees the full route set. Continue
# on per-feature import failure — the assertion below still catches
# missing-route regressions.
import importlib
from litellm.proxy._lazy_features import LAZY_FEATURES
registered_paths = [getattr(r, "path", "") for r in app.routes]
for feat in LAZY_FEATURES:
if any(rp.startswith(p) for p in feat.path_prefixes for rp in registered_paths):
continue
try:
module = importlib.import_module(feat.module_path)
feat.register_fn(app, module)
except Exception as exc:
print(f"warning: failed to force-load {feat.name}: {exc}")
_all_routes = []
for route in app.routes:

View file

@ -1553,6 +1553,7 @@ async def test_add_callback_via_key(prisma_client):
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(
metadata={
"allow_client_mock_response": True,
"logging": [
{
"callback_name": "langfuse", # 'otel', 'langfuse', 'lunary'
@ -1563,7 +1564,7 @@ async def test_add_callback_via_key(prisma_client):
"langfuse_host": "https://us.cloud.langfuse.com",
},
}
]
],
}
),
)
@ -1657,6 +1658,7 @@ async def test_add_callback_via_key_litellm_pre_call_utils(
team_id=None,
max_parallel_requests=None,
metadata={
"allow_client_mock_response": True,
"logging": [
{
"callback_name": "langfuse",
@ -1667,7 +1669,7 @@ async def test_add_callback_via_key_litellm_pre_call_utils(
"langfuse_host": "https://us.cloud.langfuse.com",
},
}
]
],
},
tpm_limit=None,
rpm_limit=None,
@ -1813,6 +1815,7 @@ async def test_add_callback_via_key_litellm_pre_call_utils_gcs_bucket(
team_id=None,
max_parallel_requests=None,
metadata={
"allow_client_mock_response": True,
"logging": [
{
"callback_name": "gcs_bucket",
@ -1822,7 +1825,7 @@ async def test_add_callback_via_key_litellm_pre_call_utils_gcs_bucket(
"gcs_path_service_account": "pathrise-convert-1606954137718-a956eef1a2a8.json",
},
}
]
],
},
tpm_limit=None,
rpm_limit=None,
@ -1946,6 +1949,7 @@ async def test_add_callback_via_key_litellm_pre_call_utils_langsmith(
team_id=None,
max_parallel_requests=None,
metadata={
"allow_client_mock_response": True,
"logging": [
{
"callback_name": "langsmith",
@ -1956,7 +1960,7 @@ async def test_add_callback_via_key_litellm_pre_call_utils_langsmith(
"langsmith_base_url": "https://api.smith.langchain.com",
},
}
]
],
},
tpm_limit=None,
rpm_limit=None,

View file

@ -235,18 +235,10 @@ async def test_add_key_or_team_level_spend_logs_metadata_to_request(
"langfuse_host": "https://us.cloud.langfuse.com",
"langfuse_public_key": "pk-lf-9636b7a6-c066",
"langfuse_secret_key": "sk-lf-7cc8b620",
},
{
"langfuse_host": "os.environ/LANGFUSE_HOST_TEMP",
"langfuse_public_key": "os.environ/LANGFUSE_PUBLIC_KEY_TEMP",
"langfuse_secret_key": "os.environ/LANGFUSE_SECRET_KEY_TEMP",
},
}
],
)
def test_dynamic_logging_metadata_key_and_team_metadata(callback_vars):
os.environ["LANGFUSE_PUBLIC_KEY_TEMP"] = "pk-lf-9636b7a6-c066"
os.environ["LANGFUSE_SECRET_KEY_TEMP"] = "sk-lf-7cc8b620"
os.environ["LANGFUSE_HOST_TEMP"] = "https://us.cloud.langfuse.com"
from litellm.proxy.proxy_server import ProxyConfig
proxy_config = ProxyConfig()
@ -317,6 +309,41 @@ def test_dynamic_logging_metadata_key_and_team_metadata(callback_vars):
assert "os.environ" not in var
def test_dynamic_logging_metadata_ignores_env_references_from_key_metadata(
monkeypatch,
):
monkeypatch.setenv("LANGFUSE_SECRET_KEY_TEMP", "server-side-secret")
monkeypatch.setattr(
litellm.utils,
"get_secret",
lambda *args, **kwargs: pytest.fail("get_secret should not be called"),
)
from litellm.proxy.proxy_server import ProxyConfig
proxy_config = ProxyConfig()
user_api_key_dict = UserAPIKeyAuth(
api_key="test-key",
metadata={
"logging": [
{
"callback_name": "langfuse",
"callback_type": "success",
"callback_vars": {
"langfuse_secret_key": "os.environ/LANGFUSE_SECRET_KEY_TEMP",
},
}
]
},
team_metadata={},
)
callbacks = _get_dynamic_logging_metadata(
user_api_key_dict=user_api_key_dict, proxy_config=proxy_config
)
assert callbacks is None
@pytest.mark.parametrize(
"callback_vars",
[
@ -1263,11 +1290,16 @@ def test_proxy_config_state_post_init_callback_call(monkeypatch):
}
)
LiteLLMProxyRequestSetup.add_team_based_callbacks_from_config(
callback_metadata = LiteLLMProxyRequestSetup.add_team_based_callbacks_from_config(
team_id="test",
proxy_config=pc,
)
assert callback_metadata is not None
assert callback_metadata.callback_vars is not None
assert callback_metadata.callback_vars["langfuse_public_key"] == "test_public_key"
assert callback_metadata.callback_vars["langfuse_secret"] == "test_secret_key"
config = pc.get_config_state()
assert config["litellm_settings"]["default_team_settings"][0]["team_id"] == "test"

View file

@ -1118,11 +1118,17 @@ async def test_jwt_non_admin_team_route_access(monkeypatch):
@pytest.mark.asyncio
async def test_x_litellm_api_key():
"""
Check if auth can pick up x-litellm-api-key header, even if Bearer token is provided
Check if auth can pick up x-litellm-api-key header, even if Bearer token is provided.
On a master-key match, ``UserAPIKeyAuth.api_key`` (and the derived
``token``) are now the stable alias ``LITELLM_PROXY_MASTER_KEY_ALIAS``
rather than ``hash_token(master_key)`` — the master key (or its hash)
must not propagate into spend logs / metrics / audit trails.
"""
from fastapi import Request
from starlette.datastructures import URL
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy._types import (
LiteLLM_TeamTable,
LiteLLM_TeamTableCachedObj,
@ -1148,7 +1154,8 @@ async def test_x_litellm_api_key():
api_key="Bearer " + ignored_key,
custom_litellm_key_header=master_key,
)
assert valid_token.token == hash_token(master_key)
assert valid_token.token == LITELLM_PROXY_MASTER_KEY_ALIAS
assert valid_token.token != hash_token(master_key)
@pytest.mark.asyncio

View file

@ -2476,3 +2476,186 @@ def test_bedrock_tools_pt_passes_ttl_for_claude_4_5():
cache_blocks_old = [b for b in result_old if "cachePoint" in b]
assert len(cache_blocks_old) == 1
assert "ttl" not in cache_blocks_old[0]["cachePoint"]
def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document():
"""
OpenAI `{type: "file", file: {file_data: "data:application/pdf;..."}}` inside
a tool-message content list should translate to an Anthropic document block
inside the tool_result content. Reuses anthropic_process_openai_file_message,
which already handles this for user messages.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_tool_result,
)
pdf_b64 = "JVBERi0xLjQKJeLjz9MK"
message = {
"tool_call_id": "toolu_pdf_1",
"role": "tool",
"name": "fetch_document",
"content": [
{
"type": "file",
"file": {
"file_data": f"data:application/pdf;base64,{pdf_b64}",
"filename": "summary.pdf",
},
},
],
}
result = convert_to_anthropic_tool_result(message)
assert result["type"] == "tool_result"
assert result["tool_use_id"] == "toolu_pdf_1"
content = result["content"]
assert isinstance(content, list) and len(content) == 1
block = content[0]
assert block["type"] == "document"
assert block["source"]["type"] == "base64"
assert block["source"]["media_type"] == "application/pdf"
assert block["source"]["data"] == pdf_b64
def test_convert_to_anthropic_tool_result_image_url_pdf_data_uri_becomes_document():
"""
Regression: a PDF sent as an `image_url` data URI on the tool-result path
must translate to an Anthropic document block (not an image block — Anthropic
rejects image blocks whose media_type is a non-image like application/pdf).
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_tool_result,
)
pdf_b64 = "JVBERi0xLjQKJeLjz9MK"
message = {
"tool_call_id": "toolu_pdf_img_1",
"role": "tool",
"name": "fetch_document",
"content": [
{
"type": "image_url",
"image_url": {
"url": f"data:application/pdf;base64,{pdf_b64}",
},
},
],
}
result = convert_to_anthropic_tool_result(message)
content = result["content"]
assert isinstance(content, list) and len(content) == 1
block = content[0]
assert block["type"] == "document"
assert block["source"]["media_type"] == "application/pdf"
assert block["source"]["data"] == pdf_b64
def test_convert_to_anthropic_tool_result_image_url_unsupported_mime_stays_image_path():
"""
An `image_url` data URI whose mime is neither application/pdf nor text/plain
(e.g. application/json) must NOT be routed through the document path. Anthropic
only accepts application/pdf and text/plain as base64 document media_types —
anything else would produce a document block the API rejects. The old
(pre-fix) behavior was to wrap such data as an image block, which also
fails but stays on the image code path; preserve that failure mode rather
than switching to a document path that is equally broken.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_tool_result,
)
message = {
"tool_call_id": "toolu_json_1",
"role": "tool",
"name": "fetch_json",
"content": [
{
"type": "image_url",
"image_url": {
"url": "data:application/json;base64,eyJrIjoidiJ9",
},
},
],
}
result = convert_to_anthropic_tool_result(message)
content = result["content"]
assert isinstance(content, list) and len(content) == 1
block = content[0]
assert block["type"] == "image", (
f"unsupported mime {block.get('source', {}).get('media_type')!r} "
f"should not be routed to document path; got {block}"
)
def test_convert_to_anthropic_tool_result_image_url_text_plain_data_uri_becomes_document():
"""
text/plain is one of the two mimes Anthropic accepts as a base64 document
media_type. Confirm it routes through the document path so tightening the
gate to {application/pdf, text/plain} (not "application/*") covers both.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_tool_result,
)
txt_b64 = "aGVsbG8=" # "hello"
message = {
"tool_call_id": "toolu_txt_1",
"role": "tool",
"name": "fetch_text",
"content": [
{
"type": "image_url",
"image_url": {
"url": f"data:text/plain;base64,{txt_b64}",
},
},
],
}
result = convert_to_anthropic_tool_result(message)
content = result["content"]
assert isinstance(content, list) and len(content) == 1
block = content[0]
assert block["type"] == "document"
assert block["source"]["media_type"] == "text/plain"
assert block["source"]["data"] == txt_b64
def test_convert_to_anthropic_tool_result_image_url_png_still_becomes_image():
"""
Regression: image_url with a real image mime type must continue to translate
to an Anthropic image block. Locks in existing behavior after the
data-URI-mime-type branching for PDFs.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_tool_result,
)
png_b64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGNgYGBgAAAABQABXvMqOgAAAABJRU5ErkJggg=="
message = {
"tool_call_id": "toolu_png_1",
"role": "tool",
"name": "fetch_image",
"content": [
{
"type": "image_url",
"image_url": {
"url": f"data:image/png;base64,{png_b64}",
},
},
],
}
result = convert_to_anthropic_tool_result(message)
content = result["content"]
assert isinstance(content, list) and len(content) == 1
block = content[0]
assert block["type"] == "image"
assert block["source"]["media_type"] == "image/png"

View file

@ -5,10 +5,17 @@ Covers the proxy flow where headers arrive in litellm_params["metadata"]["header
but litellm_params["litellm_metadata"] is None.
"""
from types import SimpleNamespace
import pytest
import litellm
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
from litellm.litellm_core_utils.redact_messages import (
_redact_responses_api_output,
perform_redaction,
should_redact_message_logging,
)
from litellm.responses.main import mock_responses_api_response
@pytest.fixture(autouse=True)
@ -68,8 +75,7 @@ class TestShouldRedactMessageLogging:
assert should_redact_message_logging(details) is True
def test_disable_redaction_via_header_proxy_flow(self):
"""litellm-disable-message-redaction should suppress redaction
even when global setting is on, and litellm_metadata is None."""
"""Core helper still honors the explicit disable-redaction header."""
litellm.turn_off_message_logging = True
details = _make_model_call_details(
metadata_headers={"litellm-disable-message-redaction": "true"},
@ -77,6 +83,14 @@ class TestShouldRedactMessageLogging:
)
assert should_redact_message_logging(details) is False
def test_disable_redaction_via_header_when_global_off(self):
"""litellm-disable-message-redaction is still honored when global redaction is off."""
details = _make_model_call_details(
metadata_headers={"litellm-disable-message-redaction": "true"},
litellm_metadata=None,
)
assert should_redact_message_logging(details) is False
# ---- SDK direct-call flow: headers in litellm_metadata ----
def test_enable_redaction_via_header_in_litellm_metadata(self):
@ -127,6 +141,16 @@ class TestShouldRedactMessageLogging:
)
assert should_redact_message_logging(details) is False
def test_dynamic_param_false_overrides_global_redaction(self):
"""Dynamic turn_off_message_logging=False should take precedence."""
litellm.turn_off_message_logging = True
details = _make_model_call_details(
metadata_headers={},
litellm_metadata=None,
standard_callback_dynamic_params={"turn_off_message_logging": False},
)
assert should_redact_message_logging(details) is False
# ---- non-dict metadata safety ----
def test_both_metadata_fields_none(self):
@ -145,3 +169,183 @@ class TestShouldRedactMessageLogging:
litellm_metadata=None,
)
assert should_redact_message_logging(details) is True
class TestPerformRedaction:
def test_redacts_standard_logging_and_responses_api_dicts(self):
details = {
"messages": [{"role": "user", "content": "sensitive input"}],
"prompt": "sensitive prompt",
"input": "sensitive input",
"standard_logging_object": {
"messages": [{"role": "user", "content": "sensitive input"}],
"response": {
"output": [
{"text": "top-level text"},
{"content": [{"text": "nested text"}]},
{"type": "reasoning", "summary": [{"text": "reasoning"}]},
],
"usage": {"total_tokens": 1},
},
},
}
result = {
"output": [
{"text": "top-level result"},
{"content": [{"text": "nested result"}]},
{"type": "reasoning", "summary": [{"text": "reasoning result"}]},
],
"usage": {"total_tokens": 1},
}
redacted = perform_redaction(details, result)
assert details["messages"] == [
{"role": "user", "content": "redacted-by-litellm"}
]
assert details["prompt"] == ""
assert details["input"] == ""
logged_response = details["standard_logging_object"]["response"]
assert logged_response["usage"] == {"total_tokens": 1}
assert logged_response["output"][0]["text"] == "redacted-by-litellm"
assert logged_response["output"][1]["content"][0]["text"] == (
"redacted-by-litellm"
)
assert logged_response["output"][2]["summary"][0]["text"] == (
"redacted-by-litellm"
)
assert redacted["usage"] == {"total_tokens": 1}
assert redacted["output"][0]["text"] == "redacted-by-litellm"
assert redacted["output"][1]["content"][0]["text"] == "redacted-by-litellm"
assert redacted["output"][2]["summary"][0]["text"] == "redacted-by-litellm"
assert result["output"][0]["text"] == "top-level result"
def test_redacts_model_response_dict_choices(self):
result = {
"choices": [
{
"message": {
"content": "message content",
"reasoning_content": "message reasoning",
"thinking_blocks": ["thinking"],
"audio": {"data": "audio"},
}
},
{
"delta": {
"content": "delta content",
"reasoning_content": "delta reasoning",
"thinking_blocks": ["delta thinking"],
"audio": {"data": "audio"},
}
},
]
}
redacted = perform_redaction({}, result)
message = redacted["choices"][0]["message"]
assert message["content"] == "redacted-by-litellm"
assert message["reasoning_content"] == "redacted-by-litellm"
assert message["thinking_blocks"] is None
assert message["audio"] is None
delta = redacted["choices"][1]["delta"]
assert delta["content"] == "redacted-by-litellm"
assert delta["reasoning_content"] == "redacted-by-litellm"
assert delta["thinking_blocks"] is None
assert delta["audio"] is None
def test_redacts_standard_logging_model_response_dict_choices(self):
details = {
"standard_logging_object": {
"response": {
"choices": [
{
"message": {
"content": "message content",
"reasoning_content": "message reasoning",
"thinking_blocks": ["thinking"],
"audio": {"data": "audio"},
}
},
{
"delta": {
"content": "delta content",
"reasoning_content": "delta reasoning",
"thinking_blocks": ["delta thinking"],
"audio": {"data": "audio"},
}
},
]
}
}
}
perform_redaction(details, None)
choices = details["standard_logging_object"]["response"]["choices"]
message = choices[0]["message"]
assert message["content"] == "redacted-by-litellm"
assert message["reasoning_content"] == "redacted-by-litellm"
assert message["thinking_blocks"] is None
assert message["audio"] is None
delta = choices[1]["delta"]
assert delta["content"] == "redacted-by-litellm"
assert delta["reasoning_content"] == "redacted-by-litellm"
assert delta["thinking_blocks"] is None
assert delta["audio"] is None
def test_redacts_object_choices_inside_model_response_dict(self):
result = {
"choices": [
litellm.Choices(
message=litellm.Message(
content="message content",
role="assistant",
reasoning_content="message reasoning",
)
)
]
}
redacted = perform_redaction({}, result)
choice = redacted["choices"][0]
assert choice.message.content == "redacted-by-litellm"
assert choice.message.reasoning_content == "redacted-by-litellm"
def test_redacts_response_output_objects_with_top_level_text(self):
output_items = [
SimpleNamespace(text="top-level output"),
"non-dict output item",
]
_redact_responses_api_output(output_items)
assert output_items[0].text == "redacted-by-litellm"
assert output_items[1] == "non-dict output item"
def test_skips_non_dict_response_output_items(self):
result = {
"output": [
"non-dict output item",
{"content": [{"text": "nested result"}]},
]
}
redacted = perform_redaction({}, result)
assert redacted["output"][0] == "non-dict output item"
assert redacted["output"][1]["content"][0]["text"] == "redacted-by-litellm"
def test_redacts_responses_api_response_object(self):
response = mock_responses_api_response("sensitive output")
redacted = perform_redaction({}, response)
assert redacted.output[0].content[0].text == "redacted-by-litellm"
assert response.output[0].content[0].text == "sensitive output"

View file

@ -4148,6 +4148,224 @@ def test_transform_response_finish_reason_stop_when_json_mode_filters_all_tools(
assert result.choices[0].finish_reason == "stop"
def test_bedrock_tool_message_openai_file_pdf_becomes_document():
"""
OpenAI Chat Completions `{type: "file", file: {file_data: "data:application/pdf;...", filename}}`
inside a tool message content list should translate to a Bedrock
toolResult.content[].document block.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
_bedrock_converse_messages_pt,
)
pdf_b64 = "JVBERi0xLjQKJeLjz9MK" # tiny "%PDF-1.4\n" header
messages = [
{"role": "user", "content": "Summarize the attached PDF."},
{
"tool_call_id": "tooluse_pdf_1",
"role": "tool",
"name": "fetch_document",
"content": [
{
"type": "file",
"file": {
"file_data": f"data:application/pdf;base64,{pdf_b64}",
"filename": "summary.pdf",
},
},
],
},
]
translated_msg = _bedrock_converse_messages_pt(
messages=messages, model="", llm_provider=""
)
tool_result = translated_msg[-1]["content"][-1]["toolResult"]
assert tool_result["toolUseId"] == "tooluse_pdf_1"
assert len(tool_result["content"]) == 1
block = tool_result["content"][0]
assert "document" in block, f"expected document block, got {block}"
assert block["document"]["format"] == "pdf"
assert block["document"]["source"]["bytes"] == pdf_b64
assert block["document"]["name"].startswith("DocumentPDFmessages_")
assert block["document"]["name"].endswith("_pdf")
def test_bedrock_tool_message_image_url_pdf_data_uri_becomes_document():
"""
Regression for the processor-returns-document-but-wrapper-drops-it bug:
when a caller sends a PDF as an `image_url` data URI on the tool-result path,
BedrockImageProcessor correctly returns a {"document": ...} block, but the
tool-result wrapper only appended the "image" case, silently dropping documents.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
_bedrock_converse_messages_pt,
)
pdf_b64 = "JVBERi0xLjQKJeLjz9MK"
messages = [
{"role": "user", "content": "Summarize the attached PDF."},
{
"tool_call_id": "tooluse_pdf_img_1",
"role": "tool",
"name": "fetch_document",
"content": [
{
"type": "image_url",
"image_url": {
"url": f"data:application/pdf;base64,{pdf_b64}",
},
},
],
},
]
translated_msg = _bedrock_converse_messages_pt(
messages=messages, model="", llm_provider=""
)
tool_result = translated_msg[-1]["content"][-1]["toolResult"]
assert tool_result["toolUseId"] == "tooluse_pdf_img_1"
assert len(tool_result["content"]) == 1
block = tool_result["content"][0]
assert "document" in block, f"expected document block, got {block}"
assert block["document"]["format"] == "pdf"
assert block["document"]["source"]["bytes"] == pdf_b64
def test_bedrock_tool_message_file_id_http_url_becomes_document():
"""
OpenAI `file.file_id` is a server-side file reference. The Bedrock
user-message path (_process_file_message at factory.py:4796) accepts either
`file_data` or `file_id` and forwards to BedrockImageProcessor. The
tool-result path must match: when `file_id` is an http(s) PDF URL, it
should resolve to a Bedrock document block, not be silently dropped.
"""
from unittest.mock import patch
from litellm.litellm_core_utils.prompt_templates.factory import (
BedrockImageProcessor,
_bedrock_converse_messages_pt,
)
pdf_url = "https://example.com/whitepaper.pdf"
fake_document_block = {
"document": {
"format": "pdf",
"name": "fake_doc",
"source": {"bytes": "ZmFrZQ=="},
}
}
messages = [
{"role": "user", "content": "Summarize the attached PDF."},
{
"tool_call_id": "tooluse_fid_1",
"role": "tool",
"name": "fetch_document",
"content": [
{
"type": "file",
"file": {
"file_id": pdf_url,
"filename": "whitepaper.pdf",
},
},
],
},
]
with patch.object(
BedrockImageProcessor,
"process_image_sync",
return_value=fake_document_block,
) as mock_proc:
translated_msg = _bedrock_converse_messages_pt(
messages=messages, model="", llm_provider=""
)
mock_proc.assert_called_once()
assert mock_proc.call_args.kwargs["image_url"] == pdf_url
tool_result = translated_msg[-1]["content"][-1]["toolResult"]
assert len(tool_result["content"]) == 1
block = tool_result["content"][0]
assert "document" in block, f"expected document block, got {block}"
assert block["document"]["source"]["bytes"] == "ZmFrZQ=="
def test_bedrock_tool_message_file_without_data_or_id_raises():
"""
The user-message path raises BadRequestError when a `type: "file"` block
has neither `file_data` nor `file_id` (factory.py:4802-4809). The
tool-result path must match — silently dropping the block makes the model
see an empty tool result and obscures the caller bug.
"""
import litellm
from litellm.litellm_core_utils.prompt_templates.factory import (
_bedrock_converse_messages_pt,
)
messages = [
{"role": "user", "content": "Summarize."},
{
"tool_call_id": "tooluse_bad_1",
"role": "tool",
"name": "fetch_document",
"content": [
{
"type": "file",
"file": {"filename": "nothing.pdf"},
},
],
},
]
with pytest.raises(litellm.BadRequestError):
_bedrock_converse_messages_pt(messages=messages, model="", llm_provider="")
def test_bedrock_tool_message_image_url_png_still_becomes_image():
"""
Regression: image_url with an image mime type must continue to translate
to a Bedrock image block (not document). Locks in existing behavior after
the document-passthrough fix.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
_bedrock_converse_messages_pt,
)
png_b64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGNgYGBgAAAABQABXvMqOgAAAABJRU5ErkJggg=="
messages = [
{"role": "user", "content": "Describe the attached image."},
{
"tool_call_id": "tooluse_png_1",
"role": "tool",
"name": "fetch_image",
"content": [
{
"type": "image_url",
"image_url": {
"url": f"data:image/png;base64,{png_b64}",
},
},
],
},
]
translated_msg = _bedrock_converse_messages_pt(
messages=messages, model="", llm_provider=""
)
tool_result = translated_msg[-1]["content"][-1]["toolResult"]
assert len(tool_result["content"]) == 1
block = tool_result["content"][0]
assert "image" in block, f"expected image block, got {block}"
assert "document" not in block
assert block["image"]["format"] == "png"
assert block["image"]["source"]["bytes"] == png_b64
def test_transform_response_does_not_leak_body_on_parse_failure():
from litellm.llms.bedrock.common_utils import BedrockError

View file

@ -180,6 +180,140 @@ def test_get_aws_region_name_boto3_fallback():
mock_boto3_session.assert_not_called()
@pytest.mark.parametrize(
"bad_region",
[
"us-east-1@example.com/",
"us-east-1@example.com",
"us-east-1/path",
"us-east-1.example.com",
"us-east-1:8080",
"us-east-1#fragment",
"us-east-1?query=1",
"us-east-1\\path",
"US-EAST-1", # uppercase not allowed
"us east 1", # spaces not allowed
"", # empty string not allowed
"us-east-1\n", # trailing newline must not slip past $
],
)
def test_get_aws_region_name_rejects_malformed_region(bad_region):
"""
Region names are interpolated into endpoint URL templates, so any value
containing characters that would alter URL parsing must be rejected.
"""
base_aws_llm = BaseAWSLLM()
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm._get_aws_region_name(
optional_params={"aws_region_name": bad_region}
)
@pytest.mark.parametrize(
"valid_region",
[
"us-east-1",
"eu-west-2",
"ap-southeast-1",
"us-gov-west-1",
"cn-north-1",
"me-south-1",
],
)
def test_get_aws_region_name_accepts_valid_regions(valid_region):
"""Real AWS region formats must continue to work after the format guard."""
base_aws_llm = BaseAWSLLM()
result = base_aws_llm._get_aws_region_name(
optional_params={"aws_region_name": valid_region}
)
assert result == valid_region
def test_get_aws_region_name_rejects_malformed_region_from_env():
"""
A malformed AWS_REGION / AWS_REGION_NAME env value must also be rejected
before it can flow into a URL template.
"""
base_aws_llm = BaseAWSLLM()
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
def side_effect(key, default=None):
if key == "AWS_REGION_NAME":
return "us-east-1@example.com/"
return default
mock_get_secret.side_effect = side_effect
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm._get_aws_region_name(optional_params={})
def test_get_aws_region_name_for_non_llm_api_calls_rejects_malformed_param():
"""
The non-LLM helper (used by Guardrails, Vector Stores, etc.) must validate
a region passed in directly so it can't flow into a URL template.
"""
base_aws_llm = BaseAWSLLM()
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm.get_aws_region_name_for_non_llm_api_calls(
aws_region_name="us-east-1@example.com/"
)
def test_get_aws_region_name_for_non_llm_api_calls_rejects_malformed_env():
"""
A malformed AWS_REGION / AWS_REGION_NAME env value must be rejected on the
non-LLM path too — Guardrails and Vector Stores read the same env vars.
"""
base_aws_llm = BaseAWSLLM()
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
def side_effect(key, default=None):
if key == "AWS_REGION_NAME":
return "us-east-1@example.com/"
return default
mock_get_secret.side_effect = side_effect
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm.get_aws_region_name_for_non_llm_api_calls()
def test_get_aws_region_name_for_non_llm_api_calls_accepts_valid_region():
"""The non-LLM helper still returns valid regions unchanged."""
base_aws_llm = BaseAWSLLM()
assert (
base_aws_llm.get_aws_region_name_for_non_llm_api_calls(
aws_region_name="us-east-1"
)
== "us-east-1"
)
def test_get_aws_region_from_model_arn_rejects_malformed_region():
"""
If the region segment of a model ARN does not match the expected format,
the helper must return None so the caller falls back to env / default.
"""
base_aws_llm = BaseAWSLLM()
bad_arn = (
"arn:aws:bedrock:us-east-1@example.com:123456789012"
":foundation-model/anthropic.claude-3-sonnet"
)
assert base_aws_llm._get_aws_region_from_model_arn(bad_arn) is None
good_arn = (
"arn:aws:bedrock:us-east-1:123456789012"
":foundation-model/anthropic.claude-3-sonnet"
)
assert base_aws_llm._get_aws_region_from_model_arn(good_arn) == "us-east-1"
def test_sign_request_with_env_var_bearer_token():
# Create instance of actual class
llm = BaseAWSLLM()

View file

@ -1,3 +1,4 @@
import asyncio
import os
import sys
from unittest.mock import AsyncMock, Mock, patch
@ -8,6 +9,8 @@ import pytest
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
_google_genai_streaming_hidden_params,
@ -103,7 +106,9 @@ def test_fingerprint_agentic_tools_is_deterministic():
tools_a = {"tool_calls": [{"id": "1", "input": {"q": "abc"}, "name": "web_search"}]}
tools_b = {"tool_calls": [{"name": "web_search", "input": {"q": "abc"}, "id": "1"}]}
assert handler._fingerprint_agentic_tools(tools_a) == handler._fingerprint_agentic_tools(tools_b)
assert handler._fingerprint_agentic_tools(
tools_a
) == handler._fingerprint_agentic_tools(tools_b)
@pytest.mark.asyncio
@ -350,3 +355,70 @@ def test_google_genai_streaming_hidden_params_model_info_and_router_fallback():
response_headers=httpx.Headers({}),
)
assert from_router["model_id"] == "router-model-id"
def _build_delete_response_mock(captured: dict):
"""Returns a fake httpx delete that records its kwargs."""
def _response() -> httpx.Response:
return httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
content=b'{"id": "resp_x", "object": "response", "deleted": true}',
request=httpx.Request(method="DELETE", url="https://test.openai.azure.com"),
)
async def fake_async_delete(*args, **kwargs):
captured.update(kwargs)
return _response()
def fake_sync_delete(*args, **kwargs):
captured.update(kwargs)
return _response()
return fake_async_delete, fake_sync_delete
def test_async_delete_responses_omits_body_for_azure():
"""Azure responses DELETE rejects requests with any body. Verify the handler
does not pass `json=` to httpx when the transformer returns an empty dict."""
captured: dict = {}
fake_async_delete, _ = _build_delete_response_mock(captured)
async def run():
with patch.object(AsyncHTTPHandler, "delete", new=fake_async_delete):
await litellm.adelete_responses(
response_id="resp_xyz",
custom_llm_provider="azure",
api_base="https://test.openai.azure.com",
api_key="test-key",
api_version="2025-03-01-preview",
)
asyncio.run(run())
assert "json" not in captured
assert "data" not in captured
assert captured["url"].endswith(
"/openai/responses/resp_xyz?api-version=2025-03-01-preview"
)
def test_sync_delete_responses_omits_body_for_azure():
captured: dict = {}
_, fake_sync_delete = _build_delete_response_mock(captured)
with patch.object(HTTPHandler, "delete", new=fake_sync_delete):
litellm.delete_responses(
response_id="resp_xyz",
custom_llm_provider="azure",
api_base="https://test.openai.azure.com",
api_key="test-key",
api_version="2025-03-01-preview",
)
assert "json" not in captured
assert "data" not in captured
assert captured["url"].endswith(
"/openai/responses/resp_xyz?api-version=2025-03-01-preview"
)

View file

@ -0,0 +1,244 @@
"""Regression tests for LIT-2642 — interrupted streams must still flush usage."""
import asyncio
from typing import List
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
def _make_streaming_response(chunks: List[bytes]):
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
mock.headers = httpx.Headers({"content-type": "application/vnd.amazon.eventstream"})
mock.raise_for_status = MagicMock(return_value=None)
async def _aiter_bytes():
for chunk in chunks:
yield chunk
mock.aiter_bytes = _aiter_bytes
mock.aclose = AsyncMock()
return mock
def _make_logging_obj():
mock = MagicMock()
mock.async_flush_passthrough_collected_chunks = AsyncMock()
return mock
class _ImmediateExecutor:
def submit(self, fn, *args, **kwargs):
fn(*args, **kwargs)
@pytest.mark.asyncio
async def test_async_streaming_flushes_on_normal_completion():
from litellm.passthrough.main import _async_streaming
chunks = [b"chunk-1", b"chunk-2", b"chunk-3"]
mock_response = _make_streaming_response(chunks)
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
async for chunk in _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
):
received.append(chunk)
assert received == chunks
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == chunks
assert call_kwargs["provider_config"] is provider_config
@pytest.mark.asyncio
async def test_async_streaming_flushes_on_client_disconnect():
from litellm.passthrough.main import _async_streaming
chunks = [
b'{"chunk": 1, "outputTokens": 10}',
b'{"chunk": 2, "outputTokens": 12}',
b'{"chunk": 3, "outputTokens": 8}',
]
mock_response = _make_streaming_response(chunks)
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
gen = _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
)
received = [await gen.__anext__()]
await gen.aclose()
assert received == [chunks[0]]
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == [chunks[0]]
@pytest.mark.asyncio
async def test_async_streaming_does_not_flush_on_4xx():
from litellm.passthrough.main import _async_streaming
err_response = MagicMock(spec=httpx.Response)
err_response.status_code = 429
def _raise():
raise httpx.HTTPStatusError(
"429",
request=httpx.Request("POST", "https://example.com"),
response=httpx.Response(
429, request=httpx.Request("POST", "https://example.com")
),
)
err_response.raise_for_status = _raise
err_response.aclose = AsyncMock()
async def response_coro():
return err_response
mock_logging_obj = _make_logging_obj()
with pytest.raises(httpx.HTTPStatusError):
async for _ in _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=MagicMock(),
):
pass
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_not_called()
@pytest.mark.asyncio
async def test_async_streaming_flushes_on_upstream_exception_with_partial_data():
from litellm.passthrough.main import _async_streaming
partial_chunks = [b"partial-chunk-1", b"partial-chunk-2"]
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.raise_for_status = MagicMock(return_value=None)
mock_response.aclose = AsyncMock()
async def _aiter_bytes_then_raise():
for c in partial_chunks:
yield c
raise httpx.ReadError("upstream disconnected")
mock_response.aiter_bytes = _aiter_bytes_then_raise
async def response_coro():
return mock_response
mock_logging_obj = _make_logging_obj()
provider_config = MagicMock()
received = []
with pytest.raises(httpx.ReadError):
async for chunk in _async_streaming(
response=response_coro(),
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
):
received.append(chunk)
assert received == partial_chunks
await asyncio.sleep(0)
mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = (
mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs
)
assert call_kwargs["raw_bytes"] == partial_chunks
def test_sync_streaming_flushes_on_normal_completion():
from litellm.passthrough.main import _sync_streaming
chunks = [b"a", b"b", b"c"]
mock_response = MagicMock(spec=httpx.Response)
def _iter_bytes():
yield from chunks
mock_response.iter_bytes = _iter_bytes
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
with patch("litellm.utils.executor", _ImmediateExecutor()):
received = list(
_sync_streaming(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
)
)
assert received == chunks
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
def test_sync_streaming_flushes_on_early_close():
from litellm.passthrough.main import _sync_streaming
chunks = [b"first", b"second", b"third"]
mock_response = MagicMock(spec=httpx.Response)
def _iter_bytes():
yield from chunks
mock_response.iter_bytes = _iter_bytes
mock_logging_obj = MagicMock()
mock_logging_obj.flush_passthrough_collected_chunks = MagicMock()
provider_config = MagicMock()
with patch("litellm.utils.executor", _ImmediateExecutor()):
gen = _sync_streaming(
response=mock_response,
litellm_logging_obj=mock_logging_obj,
provider_config=provider_config,
)
first = next(gen)
gen.close()
assert first == chunks[0]
mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once()
call_kwargs = mock_logging_obj.flush_passthrough_collected_chunks.call_args.kwargs
assert call_kwargs["raw_bytes"] == [chunks[0]]

View file

@ -551,11 +551,14 @@ class TestMCPOAuth2AuthFlow:
async def test_oauth2_token_in_authorization_header_fallback(self):
"""
When only Authorization header is present with a non-LiteLLM OAuth2 token,
When only Authorization header is present with a non-LiteLLM OAuth2 token
AND the target server is operator-configured for ``auth_type=oauth2``,
auth should fall back to permissive mode (OAuth2 passthrough).
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
@ -568,10 +571,19 @@ class TestMCPOAuth2AuthFlow:
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
oauth2_server = MagicMock()
oauth2_server.auth_type = MCPAuth.oauth2
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = oauth2_server
(
auth_result,
mcp_auth_header,
@ -695,9 +707,11 @@ class TestMCPOAuth2AuthFlow:
async def test_proxy_exception_oauth2_fallback(self):
"""
user_api_key_auth raises ProxyException (not HTTPException) in production.
The OAuth2 fallback must catch ProxyException with code 401/403 too.
The OAuth2 fallback must catch ProxyException with code 401/403 too,
but only when the target server is operator-configured for ``auth_type=oauth2``.
"""
from litellm.proxy._types import ProxyException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
@ -716,10 +730,19 @@ class TestMCPOAuth2AuthFlow:
code=401,
)
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_proxy_exception,
oauth2_server = MagicMock()
oauth2_server.auth_type = MCPAuth.oauth2
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_proxy_exception,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = oauth2_server
(
auth_result,
mcp_auth_header,
@ -768,6 +791,290 @@ class TestMCPOAuth2AuthFlow:
await MCPRequestHandler.process_mcp_request(scope)
@pytest.mark.asyncio
class TestMCPPublicRouteGuard:
"""
Regression tests for GHSA-7cwm-3279-qf3c / HW6xR21d:
the public-route bypass at the top of process_mcp_request must match
the exact `/.well-known/` path prefix, not a substring of the URL.
"""
async def test_well_known_substring_in_query_does_not_bypass_auth(self):
"""
URL with `.well-known` smuggled into the query string must still
require valid LiteLLM auth.
"""
from fastapi import HTTPException
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/private_server",
"query_string": b"redirect=.well-known/oauth-protected-resource",
"headers": [(b"authorization", b"Bearer sk-bogus")],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
# Explicit unresolvable target — proves auth still fails even
# when the registry has no info to fall back to.
mock_mgr.get_mcp_server_by_name.return_value = None
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_well_known_segment_in_middle_of_path_does_not_bypass_auth(self):
"""
Path containing `.well-known` as a non-prefix component (e.g. a server
name or sub-path) must still require auth.
"""
from fastapi import HTTPException
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/.well-known-fake/tools",
"headers": [(b"authorization", b"Bearer sk-bogus")],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = None
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_legitimate_well_known_path_still_bypasses_auth(self):
"""
Real OAuth discovery routes registered under /.well-known/ must remain
public so unauthenticated clients can fetch them per RFC 8414/9728.
"""
scope = {
"type": "http",
"method": "GET",
"path": "/.well-known/oauth-protected-resource",
"headers": [],
}
# No mock needed — public path should not call user_api_key_auth at all
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
) as mock_auth:
(auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope)
mock_auth.assert_not_called()
assert isinstance(auth_result, UserAPIKeyAuth)
@pytest.mark.asyncio
class TestMCPOAuth2FallbackTargetGating:
"""
Regression tests for GHSA-h8fm-g6wc-j228 / HW6xR21d:
The OAuth2 passthrough fallback must only fire when the target MCP server
is operator-configured for ``auth_type=oauth2``. A failed LiteLLM-auth
against a non-OAuth2 server (api_key, bearer_token, basic, etc.) must
propagate as a real auth error, not be exchanged for an anonymous session.
"""
@staticmethod
def _make_server(auth_type):
server = MagicMock()
server.auth_type = auth_type
return server
async def test_fallback_blocked_when_target_is_not_oauth2(self):
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/api_key_server",
"headers": [(b"authorization", b"Bearer anything-at-all")],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPOAuth2FallbackTargetGating._make_server(MCPAuth.api_key)
)
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_fallback_blocked_when_target_unresolvable(self):
"""
If the target server cannot be resolved from path or x-mcp-servers,
we cannot prove it is OAuth2-mode, so we must fail closed.
"""
from fastapi import HTTPException
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/never_registered_server",
"headers": [(b"authorization", b"Bearer anything")],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = None
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_fallback_allowed_when_target_is_oauth2_mode(self):
"""
Operator-configured OAuth2 passthrough still works: target server has
``auth_type=oauth2`` → failed LiteLLM auth falls back to anonymous so
the bearer can be forwarded to upstream.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/atlassian_mcp",
"headers": [
(b"authorization", b"Bearer atlassian-oauth2-access-token-xyz"),
],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.return_value = (
TestMCPOAuth2FallbackTargetGating._make_server(MCPAuth.oauth2)
)
(auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope)
assert isinstance(auth_result, UserAPIKeyAuth)
async def test_fallback_blocked_when_any_target_in_header_is_not_oauth2(self):
"""
x-mcp-servers can list multiple targets. If ANY of them is non-OAuth2,
the fallback must be blocked — otherwise an attacker can mix one
OAuth2-mode server in to enable bypass against the others.
"""
from fastapi import HTTPException
from litellm.types.mcp import MCPAuth
scope = {
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [
(b"authorization", b"Bearer anything"),
(b"x-mcp-servers", b"oauth2_server,api_key_server"),
],
}
async def mock_user_api_key_auth_fails(api_key, request):
raise HTTPException(status_code=401, detail="Invalid API key")
def mock_lookup(name, client_ip=None):
if name == "oauth2_server":
return TestMCPOAuth2FallbackTargetGating._make_server(MCPAuth.oauth2)
return TestMCPOAuth2FallbackTargetGating._make_server(MCPAuth.api_key)
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_fails,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
):
mock_mgr.get_mcp_server_by_name.side_effect = mock_lookup
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_proxy_exception_with_non_numeric_code_propagates(self):
"""
``ProxyException`` normalises ``code`` via ``str()`` in its __init__,
so callers may produce ``"None"`` or any non-numeric string when no
explicit code was supplied. The exception handler must not coerce
with ``int(...)`` (which would raise ``ValueError`` and rewrite the
auth error as an unhandled 500); it must simply re-raise.
"""
from litellm.proxy._types import ProxyException
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/atlassian_mcp",
"headers": [(b"authorization", b"Bearer anything")],
}
async def mock_user_api_key_auth_no_code(api_key, request):
raise ProxyException(
message="Authentication Error",
type="auth_error",
param="api_key",
code=None,
)
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
side_effect=mock_user_api_key_auth_no_code,
):
with pytest.raises(ProxyException):
await MCPRequestHandler.process_mcp_request(scope)
class TestMCPCustomHeaderName:
"""Test suite for custom MCP authentication header name functionality"""

View file

@ -0,0 +1,402 @@
"""
Tests for the encrypted-at-rest persistence of MCP user credentials.
The ``LiteLLM_MCPUserCredentials.credential_b64`` column previously stored
both BYOK API keys and OAuth2 access tokens as plain ``urlsafe_b64encode``
of the raw value, leaving credentials readable from any DB read. The fix
runs every write through ``encrypt_value_helper`` (nacl SecretBox) and
keeps a plain-base64 fallback on read so existing rows continue to work.
"""
import base64
import json
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy._experimental.mcp_server.db import (
_decode_user_credential,
get_user_credential,
get_user_oauth_credential,
list_user_oauth_credentials,
rotate_mcp_user_credentials_master_key,
store_user_credential,
store_user_oauth_credential,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
SALT_KEY = "test-salt-key-for-byok-credential-tests-1234"
@pytest.fixture(autouse=True)
def _set_salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY)
def _make_prisma_with_existing(row):
"""Build a MagicMock prisma_client whose user-credentials table returns ``row``
for find_unique and behaves async-correctly for upsert/find_many."""
prisma = MagicMock()
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=row)
prisma.db.litellm_mcpusercredentials.upsert = AsyncMock()
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[])
return prisma
def _legacy_row(payload: str):
"""A row exactly as the pre-fix code would have written it: plain
``urlsafe_b64encode`` of the raw payload, no encryption."""
row = MagicMock()
row.credential_b64 = base64.urlsafe_b64encode(payload.encode()).decode()
row.user_id = "alice"
row.server_id = "srv-1"
return row
def _stored_value(prisma) -> str:
"""Pull the credential_b64 value passed to the most recent upsert call."""
call = prisma.db.litellm_mcpusercredentials.upsert.call_args
data = call.kwargs["data"]
create_value = data["create"]["credential_b64"]
update_value = data["update"]["credential_b64"]
assert create_value == update_value, "create/update must agree"
return create_value
# ── BYOK round-trip ───────────────────────────────────────────────────────────
@pytest.mark.asyncio
async def test_store_user_credential_does_not_persist_plaintext():
# Stored bytes must not just be base64 of the secret — that's the regression.
secret = "sk-proj-very-secret-byok-key"
prisma = _make_prisma_with_existing(row=None)
await store_user_credential(prisma, "alice", "srv-1", secret)
stored = _stored_value(prisma)
plain_b64 = base64.urlsafe_b64encode(secret.encode()).decode()
assert stored != plain_b64
# And the secret must not appear anywhere in a plain-b64 decode of the column.
try:
decoded_bytes = base64.urlsafe_b64decode(stored)
except Exception:
decoded_bytes = b""
assert secret.encode() not in decoded_bytes
@pytest.mark.asyncio
async def test_byok_round_trip_returns_plaintext():
secret = "sk-proj-very-secret-byok-key"
prisma = _make_prisma_with_existing(row=None)
await store_user_credential(prisma, "alice", "srv-1", secret)
stored = _stored_value(prisma)
row = MagicMock()
row.credential_b64 = stored
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=row)
result = await get_user_credential(prisma, "alice", "srv-1")
assert result == secret
@pytest.mark.asyncio
async def test_byok_get_returns_plaintext_for_legacy_row():
# Backward-compat: rows persisted by the pre-fix code (plain base64) must
# still decrypt-or-decode cleanly.
legacy_secret = "legacy-byok-key"
prisma = _make_prisma_with_existing(row=_legacy_row(legacy_secret))
result = await get_user_credential(prisma, "alice", "srv-1")
assert result == legacy_secret
@pytest.mark.asyncio
async def test_byok_get_returns_none_for_missing_row():
prisma = _make_prisma_with_existing(row=None)
result = await get_user_credential(prisma, "alice", "srv-1")
assert result is None
# ── OAuth2 round-trip ─────────────────────────────────────────────────────────
@pytest.mark.asyncio
async def test_store_user_oauth_credential_does_not_persist_plaintext():
access_token = "ya29.a0AfH6SMBverysecretaccesstoken"
prisma = _make_prisma_with_existing(row=None)
await store_user_oauth_credential(
prisma, "alice", "srv-1", access_token, refresh_token="rfr-xyz"
)
stored = _stored_value(prisma)
try:
decoded_bytes = base64.urlsafe_b64decode(stored)
except Exception:
decoded_bytes = b""
assert access_token.encode() not in decoded_bytes
assert b"rfr-xyz" not in decoded_bytes
@pytest.mark.asyncio
async def test_oauth_round_trip_returns_payload():
access_token = "ya29.a0AfH6SMBverysecretaccesstoken"
prisma = _make_prisma_with_existing(row=None)
await store_user_oauth_credential(
prisma,
"alice",
"srv-1",
access_token,
refresh_token="rfr-xyz",
scopes=["a", "b"],
)
stored = _stored_value(prisma)
row = MagicMock()
row.credential_b64 = stored
row.server_id = "srv-1"
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=row)
result = await get_user_oauth_credential(prisma, "alice", "srv-1")
assert result is not None
assert result["type"] == "oauth2"
assert result["access_token"] == access_token
assert result["refresh_token"] == "rfr-xyz"
assert result["scopes"] == ["a", "b"]
@pytest.mark.asyncio
async def test_oauth_get_returns_payload_for_legacy_row():
payload = {
"type": "oauth2",
"access_token": "legacy-token",
"connected_at": "2024-01-01T00:00:00Z",
}
legacy_b64 = base64.urlsafe_b64encode(json.dumps(payload).encode()).decode()
row = MagicMock()
row.credential_b64 = legacy_b64
row.server_id = "srv-1"
prisma = _make_prisma_with_existing(row=row)
result = await get_user_oauth_credential(prisma, "alice", "srv-1")
assert result is not None
assert result["access_token"] == "legacy-token"
@pytest.mark.asyncio
async def test_oauth_get_returns_none_for_byok_row():
# A row that holds a BYOK string must not leak as an OAuth payload.
prisma = _make_prisma_with_existing(row=_legacy_row("plain-byok-not-json"))
result = await get_user_oauth_credential(prisma, "alice", "srv-1")
assert result is None
# ── BYOK guard inside store_user_oauth_credential ─────────────────────────────
@pytest.mark.asyncio
async def test_byok_guard_rejects_overwriting_legacy_byok():
prisma = _make_prisma_with_existing(row=_legacy_row("plain-byok-key"))
with pytest.raises(ValueError, match="could not be verified as an OAuth2"):
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok")
@pytest.mark.asyncio
async def test_byok_guard_rejects_overwriting_encrypted_byok():
# Simulate a row written by the new (encrypted) code path: write a BYOK,
# then attempt to overwrite with an OAuth token.
prisma = _make_prisma_with_existing(row=None)
await store_user_credential(prisma, "alice", "srv-1", "sk-secret-byok")
encrypted_row = MagicMock()
encrypted_row.credential_b64 = _stored_value(prisma)
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(
return_value=encrypted_row
)
with pytest.raises(ValueError, match="could not be verified as an OAuth2"):
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok")
@pytest.mark.asyncio
async def test_byok_guard_allows_overwriting_existing_oauth():
# Refresh path: row already holds an OAuth payload, write must succeed.
prisma = _make_prisma_with_existing(row=None)
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-1")
oauth_row = MagicMock()
oauth_row.credential_b64 = _stored_value(prisma)
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=oauth_row)
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-2")
# Final upsert wrote a new payload (different from the first)
assert _stored_value(prisma) != oauth_row.credential_b64
# ── list_user_oauth_credentials ───────────────────────────────────────────────
@pytest.mark.asyncio
async def test_list_oauth_credentials_filters_byok_and_returns_payloads():
# Three rows: one encrypted OAuth, one legacy-plaintext OAuth, one BYOK.
# Only the two OAuth rows should come back.
prisma = _make_prisma_with_existing(row=None)
await store_user_oauth_credential(prisma, "alice", "srv-encrypted", "tok-enc")
encrypted_b64 = _stored_value(prisma)
encrypted_row = MagicMock()
encrypted_row.credential_b64 = encrypted_b64
encrypted_row.server_id = "srv-encrypted"
legacy_payload = {
"type": "oauth2",
"access_token": "tok-legacy",
"connected_at": "2024-01-01T00:00:00Z",
}
legacy_row = MagicMock()
legacy_row.credential_b64 = base64.urlsafe_b64encode(
json.dumps(legacy_payload).encode()
).decode()
legacy_row.server_id = "srv-legacy"
byok_row = MagicMock()
byok_row.credential_b64 = base64.urlsafe_b64encode(b"plain-byok-key").decode()
byok_row.server_id = "srv-byok"
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(
return_value=[encrypted_row, legacy_row, byok_row]
)
results = await list_user_oauth_credentials(prisma, "alice")
server_ids = {r["server_id"] for r in results}
assert server_ids == {"srv-encrypted", "srv-legacy"}
tokens = {r["access_token"] for r in results}
assert tokens == {"tok-enc", "tok-legacy"}
# ── _decode_user_credential helper ────────────────────────────────────────────
def test_decode_user_credential_handles_garbage():
# Malformed input must return None, not raise.
assert _decode_user_credential("not-base64-and-not-encrypted!!!") is None
def test_decode_user_credential_handles_none():
# Defensive: a null DB value must return None, not propagate TypeError.
assert _decode_user_credential(None) is None
def test_decode_user_credential_legacy_path():
plain = "legacy-secret"
stored = base64.urlsafe_b64encode(plain.encode()).decode()
assert _decode_user_credential(stored) == plain
# ── master-key rotation ───────────────────────────────────────────────────────
@pytest.mark.asyncio
async def test_rotate_re_encrypts_byok_with_new_key(monkeypatch):
# Encrypt a row under the current salt, then rotate to a new key, then
# confirm the stored ciphertext decrypts under the NEW key — and not under
# the old one.
prisma = _make_prisma_with_existing(row=None)
secret = "sk-original-byok-key"
await store_user_credential(prisma, "alice", "srv-1", secret)
encrypted_old = _stored_value(prisma)
row = MagicMock()
row.user_id = "alice"
row.server_id = "srv-1"
row.credential_b64 = encrypted_old
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_mcpusercredentials.update = AsyncMock()
new_master_key = "rotated-salt-key-9999-9999-9999-9999"
await rotate_mcp_user_credentials_master_key(
prisma_client=prisma, new_master_key=new_master_key
)
update_call = prisma.db.litellm_mcpusercredentials.update.call_args
new_stored = update_call.kwargs["data"]["credential_b64"]
assert new_stored != encrypted_old, "rotation must produce different ciphertext"
# Decrypt the rotated value under the NEW salt key — round-trips to plaintext.
monkeypatch.setenv("LITELLM_SALT_KEY", new_master_key)
assert (
decrypt_value_helper(
value=new_stored,
key="mcp_user_credential",
exception_type="debug",
return_original_value=False,
)
== secret
)
@pytest.mark.asyncio
async def test_rotate_migrates_legacy_plaintext_rows(monkeypatch):
# A legacy plain-base64 row must also get re-encrypted under the new key.
prisma = _make_prisma_with_existing(row=None)
legacy_row = MagicMock()
legacy_row.user_id = "alice"
legacy_row.server_id = "srv-legacy"
legacy_row.credential_b64 = base64.urlsafe_b64encode(b"legacy-plain").decode()
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(
return_value=[legacy_row]
)
prisma.db.litellm_mcpusercredentials.update = AsyncMock()
new_key = "another-rotation-key-aaaa-bbbb-cccc-dddd"
await rotate_mcp_user_credentials_master_key(
prisma_client=prisma, new_master_key=new_key
)
new_stored = prisma.db.litellm_mcpusercredentials.update.call_args.kwargs["data"][
"credential_b64"
]
monkeypatch.setenv("LITELLM_SALT_KEY", new_key)
assert (
decrypt_value_helper(
value=new_stored,
key="mcp_user_credential",
exception_type="debug",
return_original_value=False,
)
== "legacy-plain"
)
@pytest.mark.asyncio
async def test_rotate_skips_undecodable_rows():
# One bad row must not abort the rotation for the rest.
prisma = _make_prisma_with_existing(row=None)
bad_row = MagicMock()
bad_row.user_id = "alice"
bad_row.server_id = "srv-corrupt"
bad_row.credential_b64 = "!!! not base64 and not encrypted !!!"
good_row = MagicMock()
good_row.user_id = "bob"
good_row.server_id = "srv-ok"
good_row.credential_b64 = base64.urlsafe_b64encode(b"good-byok").decode()
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(
return_value=[bad_row, good_row]
)
prisma.db.litellm_mcpusercredentials.update = AsyncMock()
await rotate_mcp_user_credentials_master_key(
prisma_client=prisma, new_master_key="new-key-xxxx"
)
# Only one update call — the good row.
assert prisma.db.litellm_mcpusercredentials.update.call_count == 1
where = prisma.db.litellm_mcpusercredentials.update.call_args.kwargs["where"]
assert where["user_id_server_id"]["server_id"] == "srv-ok"

View file

@ -230,7 +230,6 @@ async def test_authorize_endpoint_forwards_pkce_parameters():
async def test_token_endpoint_forwards_code_verifier():
"""Test that token endpoint forwards code_verifier for PKCE flow"""
try:
import httpx
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
@ -632,8 +631,7 @@ async def test_token_endpoint_respects_x_forwarded_proto():
) as mock_get_client:
mock_get_client.return_value = mock_async_client
# Call token endpoint
response = await token_endpoint(
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="test_code",
@ -933,8 +931,7 @@ async def test_token_endpoint_respects_x_forwarded_host():
) as mock_get_client:
mock_get_client.return_value = mock_async_client
# Call token endpoint
response = await token_endpoint(
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="test_code",
@ -1240,6 +1237,7 @@ def _create_oauth2_server(
alias="test_oauth",
client_id="test_client_id",
client_secret="test_client_secret",
available_on_public_internet=True,
):
"""Helper to create a mock OAuth2 MCPServer."""
from litellm.proxy._types import MCPTransport
@ -1258,6 +1256,7 @@ def _create_oauth2_server(
authorization_url="https://provider.com/oauth/authorize",
token_url="https://provider.com/oauth/token",
scopes=["read", "write"],
available_on_public_internet=available_on_public_internet,
)
@ -1352,6 +1351,47 @@ async def test_authorize_root_fails_with_multiple_oauth2_servers():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_authorize_root_does_not_resolve_private_server_for_external_client():
"""Root /authorize must not auto-select an MCP server hidden from the caller IP."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
authorize,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
):
with pytest.raises(HTTPException) as exc_info:
await authorize(
request=mock_request,
client_id="dummy_client",
mcp_server_name=None,
redirect_uri="http://localhost:62646/callback",
state="test_state",
)
assert exc_info.value.status_code == 404
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_token_root_resolves_single_oauth2_server():
"""When /token is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
@ -1417,6 +1457,50 @@ async def test_token_root_resolves_single_oauth2_server():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_token_root_does_not_resolve_private_server_for_external_client():
"""Root /token must not exchange codes for a hidden MCP server."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
token_endpoint,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
):
with pytest.raises(HTTPException) as exc_info:
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="test_auth_code",
redirect_uri="http://localhost:62646/callback",
client_id="dummy_client",
mcp_server_name=None,
client_secret=None,
code_verifier="test_verifier",
)
assert exc_info.value.status_code == 404
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_register_root_resolves_single_oauth2_server():
"""When /register is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
@ -1454,6 +1538,48 @@ async def test_register_root_resolves_single_oauth2_server():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_register_root_does_not_resolve_private_server_for_external_client():
"""Root /register must not reveal or use a hidden MCP server."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value={}),
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
),
):
result = await register_client(request=mock_request, mcp_server_name=None)
assert result["client_id"] == "dummy_client"
assert result["redirect_uris"] == ["https://llm.example.com/callback"]
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_discovery_root_includes_server_name_prefix():
"""When root discovery is hit and exactly 1 OAuth2 server exists, include server name in URLs."""
@ -1493,6 +1619,54 @@ async def test_discovery_root_includes_server_name_prefix():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_discovery_root_does_not_expose_private_server_for_external_client():
"""Root discovery must use caller visibility before adding server-specific metadata."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_build_oauth_authorization_server_response,
_build_oauth_protected_resource_response,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
):
authorization_response = _build_oauth_authorization_server_response(
request=mock_request,
mcp_server_name=None,
)
resource_response = _build_oauth_protected_resource_response(
request=mock_request,
mcp_server_name=None,
use_standard_pattern=False,
)
assert "/test_oauth/" not in authorization_response["authorization_endpoint"]
assert "/test_oauth/" not in authorization_response["token_endpoint"]
assert authorization_response["scopes_supported"] == []
assert resource_response["authorization_servers"] == ["https://llm.example.com"]
assert resource_response["scopes_supported"] == []
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_oauth_callback_redirects_with_state():
"""Test OAuth callback endpoint properly decodes state and redirects to client callback URL."""
@ -1536,6 +1710,44 @@ async def test_oauth_callback_redirects_with_state():
mock_decode.assert_called_once_with("encrypted_state_value")
@pytest.mark.asyncio
async def test_oauth_callback_preserves_client_redirect_uri_query():
"""The callback should append code/state without dropping a client's existing query."""
try:
from urllib.parse import parse_qs, urlparse
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
) as mock_decode:
mock_decode.return_value = {
"base_url": "http://localhost:3000/ui/mcp/oauth/callback",
"original_state": "test-uuid-state-123",
"code_challenge": "test_challenge",
"code_challenge_method": "S256",
"client_redirect_uri": (
"http://localhost:3000/ui/mcp/oauth/callback?session=abc"
),
}
response = await callback(
code="test_authorization_code_12345",
state="encrypted_state_value",
)
assert response.status_code == 302
parsed_location = urlparse(response.headers["location"])
query_params = parse_qs(parsed_location.query)
assert query_params["session"] == ["abc"]
assert query_params["code"] == ["test_authorization_code_12345"]
assert query_params["state"] == ["test-uuid-state-123"]
@pytest.mark.asyncio
async def test_oauth_callback_handles_invalid_state():
"""Test OAuth callback returns error page when state decryption fails."""
@ -1948,6 +2160,48 @@ async def test_callback_revalidates_loopback_on_decoded_base_url():
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_callback_revalidates_loopback_on_decoded_client_redirect_uri():
"""If a state contains a full client_redirect_uri, validate that exact sink."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
) as mock_decode:
mock_decode.return_value = {
"base_url": "http://localhost:3000/cb",
"original_state": "s",
"code_challenge": None,
"code_challenge_method": None,
"client_redirect_uri": "https://attacker.example.com/cb",
}
with pytest.raises(HTTPException) as exc_info:
await callback(code="stolen_code", state="encrypted_stale_state")
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_callback_rejects_state_missing_redirect_uri():
"""Malformed state without a redirect target should fail with a structured 400."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
) as mock_decode:
mock_decode.return_value = {
"original_state": "s",
"code_challenge": None,
"code_challenge_method": None,
}
with pytest.raises(HTTPException) as exc_info:
await callback(code="code", state="encrypted_malformed_state")
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_token_endpoint_sets_no_store_cache_control():
"""RFC 6749 §5.1 / OAuth 2.1 draft-15 §4.1.3: the token response

View file

@ -4,6 +4,7 @@ import logging
import os
import sys
from datetime import datetime
from typing import Any, Dict
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -719,7 +720,8 @@ class TestMCPServerManager:
return_value=mock_client,
):
servers, scopes = await manager._fetch_oauth_metadata_from_resource(
"https://protected.example.com/.well-known/oauth"
"https://protected.example.com/.well-known/oauth",
"https://protected.example.com/mcp",
)
assert servers == [
@ -772,7 +774,8 @@ class TestMCPServerManager:
mock_well_known.assert_awaited_once_with("http://localhost:8001/mcp")
mock_fetch_auth.assert_awaited_once_with(
["https://login.microsoftonline.com/test-tenant-id/v2.0"]
["https://login.microsoftonline.com/test-tenant-id/v2.0"],
"http://localhost:8001/mcp",
)
assert result is mock_metadata
assert result.scopes == ["api://some-scope/.default"]
@ -784,7 +787,7 @@ class TestMCPServerManager:
manager = MCPServerManager()
issuer = "https://login.microsoftonline.com/test-tenant-id/v2.0"
def build_response(url: str):
def build_response(url: str, **kwargs):
mock_response = MagicMock()
if url == f"{issuer}/.well-known/openid-configuration":
mock_response.json.return_value = {
@ -810,7 +813,12 @@ class TestMCPServerManager:
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
result = await manager._fetch_single_authorization_server_metadata(issuer)
# The Azure issuer is cross-origin against the server_url — use
# the issuer itself as server_url so the test exercises the
# well-known fetch logic without needing real DNS.
result = await manager._fetch_single_authorization_server_metadata(
issuer, issuer
)
assert result is not None
assert (
@ -846,7 +854,9 @@ class TestMCPServerManager:
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
result = await manager._fetch_single_authorization_server_metadata(issuer)
result = await manager._fetch_single_authorization_server_metadata(
issuer, issuer
)
assert result is not None
assert (
@ -910,7 +920,7 @@ class TestMCPServerManager:
):
result = await manager._descovery_metadata(server_url)
mock_fetch_auth.assert_awaited_once_with(["https://example.com"])
mock_fetch_auth.assert_awaited_once_with(["https://example.com"], server_url)
assert result is mock_metadata
assert result.scopes == ["read"]
@ -2947,5 +2957,359 @@ class TestMCPServerManagerExpandToolPermissions:
assert sorted(result["uuid-a"]) == ["read_file", "write_file"]
class TestOAuthDiscoverySSRFGuard:
"""SSRF guard for the OAuth metadata discovery follow-up fetches.
The vulnerability: a malicious MCP server returns a ``WWW-Authenticate``
header pointing at an attacker-chosen ``resource_metadata`` URL, then a
PRM JSON whose ``authorization_servers[0]`` points at internal/loopback
addresses, coercing the proxy into making blind GETs to cloud-metadata
services, internal admin panels, or loopback debug endpoints.
"""
@staticmethod
def _patch_resolves(monkeypatch, mapping):
"""Patch ``socket.getaddrinfo`` for a deterministic SSRF-guard test.
``mapping`` is ``{hostname: [ip-string, ...]}``; unknown hosts raise
``gaierror`` (treated as "unresolvable" -> blocked by async_safe_get).
"""
import socket as _socket
def fake_getaddrinfo(host, port, *args, **kwargs):
if host not in mapping:
raise _socket.gaierror(f"unknown host {host}")
family = _socket.AF_INET
return [
(family, _socket.SOCK_STREAM, 0, "", (ip, port)) for ip in mapping[host]
]
monkeypatch.setattr(
"litellm.litellm_core_utils.url_utils.socket.getaddrinfo",
fake_getaddrinfo,
)
def test_same_authority_url_is_direct_fetch_eligible(self):
# Same scheme + host + port skips DNS entirely — the well-known
# endpoint construction in _attempt_well_known_discovery always
# produces same-authority URLs against the admin's server_url.
assert MCPServerManager._is_same_authority_metadata_url(
"https://example.com/.well-known/oauth-protected-resource",
"https://example.com/mcp",
)
def test_same_host_different_port_uses_safe_fetch_path(self):
assert not MCPServerManager._is_same_authority_metadata_url(
"https://example.com:9999/.well-known/oauth-protected-resource",
"https://example.com/mcp",
)
@pytest.mark.parametrize(
"ip",
[
"127.0.0.1", # loopback
"10.0.0.5", # RFC1918
"172.16.0.1", # RFC1918
"192.168.1.1", # RFC1918
"169.254.169.254", # AWS / Azure / GCP IMDS
"100.100.100.200", # Alibaba Cloud metadata
"0.0.0.0", # unspecified
"::1", # IPv6 loopback
"fe80::1", # IPv6 link-local
"fc00::1", # IPv6 ULA
],
)
@pytest.mark.asyncio
async def test_cross_origin_blocked_when_resolves_to_unsafe_ip(
self, monkeypatch, ip
):
self._patch_resolves(monkeypatch, {"attacker.example.com": [ip]})
manager = MCPServerManager()
mock_client = MagicMock()
mock_client.get = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
result = await manager._fetch_single_authorization_server_metadata(
"https://attacker.example.com",
"https://legit-mcp.example.com/mcp",
)
assert result is None
mock_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_cross_origin_allowed_when_resolves_to_public_ip(self, monkeypatch):
self._patch_resolves(
monkeypatch, {"login.microsoftonline.com": ["20.190.151.7"]}
)
manager = MCPServerManager()
mock_response = MagicMock()
mock_response.is_redirect = False
mock_response.raise_for_status = MagicMock()
mock_response.json.return_value = {
"authorization_servers": ["https://login.microsoftonline.com/tenant/v2.0"],
"scopes_supported": ["mcp.read"],
}
mock_client = MagicMock()
mock_client.get = AsyncMock(return_value=mock_response)
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
servers, scopes = await manager._fetch_oauth_metadata_from_resource(
"https://login.microsoftonline.com/tenant/v2.0/.well-known/openid-configuration",
"https://atlassian-mcp.example.com/mcp",
)
assert servers == ["https://login.microsoftonline.com/tenant/v2.0"]
assert scopes == ["mcp.read"]
mock_client.get.assert_awaited_once()
assert mock_client.get.await_args.kwargs["follow_redirects"] is False
assert (
mock_client.get.await_args.kwargs["headers"]["Host"]
== "login.microsoftonline.com"
)
@pytest.mark.asyncio
async def test_cross_origin_blocked_when_unresolvable(self, monkeypatch):
self._patch_resolves(monkeypatch, {})
manager = MCPServerManager()
mock_client = MagicMock()
mock_client.get = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
servers, scopes = await manager._fetch_oauth_metadata_from_resource(
"https://nope.example.invalid/.well-known/oauth-authorization-server",
"https://legit-mcp.example.com/mcp",
)
assert servers == []
assert scopes is None
mock_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_non_http_scheme_is_not_safe(self):
manager = MCPServerManager()
mock_client = MagicMock()
mock_client.get = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
servers, scopes = await manager._fetch_oauth_metadata_from_resource(
"file:///etc/passwd",
"https://example.com/mcp",
)
result = await manager._fetch_single_authorization_server_metadata(
"gopher://example.com/",
"https://example.com/mcp",
)
assert servers == []
assert scopes is None
assert result is None
mock_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_dual_resolution_blocked_if_any_ip_unsafe(self, monkeypatch):
# If the attacker controls a DNS record returning multiple A records,
# one of which is private, async_safe_get rejects before any network call.
self._patch_resolves(
monkeypatch, {"dual-stack.example.com": ["8.8.8.8", "127.0.0.1"]}
)
manager = MCPServerManager()
mock_client = MagicMock()
mock_client.get = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
servers, scopes = await manager._fetch_oauth_metadata_from_resource(
"https://dual-stack.example.com/.well-known/oauth-authorization-server",
"https://legit-mcp.example.com/mcp",
)
assert servers == []
assert scopes is None
mock_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_fetch_oauth_metadata_refuses_unsafe_url(self, monkeypatch):
# End-to-end: a malicious WWW-Authenticate redirecting to a loopback
# resource_metadata URL must not produce a network call.
self._patch_resolves(monkeypatch, {"attacker.example.com": ["127.0.0.1"]})
manager = MCPServerManager()
mock_client = MagicMock()
mock_client.get = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
servers, scopes = await manager._fetch_oauth_metadata_from_resource(
"https://attacker.example.com/meta",
"https://legit-mcp.example.com/mcp",
)
assert servers == []
assert scopes is None
mock_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_empty_getaddrinfo_result_blocks_url(self, monkeypatch):
# POSIX doesn't strictly forbid an empty success-list from getaddrinfo.
# async_safe_get must fail closed rather than making a network call.
monkeypatch.setattr(
"litellm.litellm_core_utils.url_utils.socket.getaddrinfo",
lambda *a, **k: [],
)
manager = MCPServerManager()
mock_client = MagicMock()
mock_client.get = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
servers, scopes = await manager._fetch_oauth_metadata_from_resource(
"https://no-records.example.com/.well-known/oauth-authorization-server",
"https://legit-mcp.example.com/mcp",
)
assert servers == []
assert scopes is None
mock_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_cross_origin_redirect_is_revalidated(self, monkeypatch):
self._patch_resolves(
monkeypatch,
{
"provider.example.com": ["8.8.8.8"],
"127.0.0.1": ["127.0.0.1"],
},
)
manager = MCPServerManager()
redirect_response = MagicMock()
redirect_response.is_redirect = True
redirect_response.headers = {"location": "http://127.0.0.1/admin"}
mock_client = MagicMock()
mock_client.get = AsyncMock(return_value=redirect_response)
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
result = await manager._fetch_single_authorization_server_metadata(
"https://provider.example.com",
"https://legit-mcp.example.com/mcp",
)
assert result is None
assert mock_client.get.await_count == 3
@pytest.mark.asyncio
async def test_same_authority_fetch_does_not_follow_redirects(self):
# Same-authority URLs may be internal admin-configured MCP servers, so
# they are fetched directly. Redirects are still disabled because a
# Location target would not inherit the same-authority guarantee.
manager = MCPServerManager()
mock_response = MagicMock()
mock_response.json.return_value = {
"authorization_servers": ["https://auth.example.com"],
}
mock_response.raise_for_status = MagicMock()
captured_kwargs: Dict[str, Any] = {}
async def fake_get(url, **kwargs):
captured_kwargs.update(kwargs)
return mock_response
mock_client = MagicMock()
mock_client.get = AsyncMock(side_effect=fake_get)
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
await manager._fetch_oauth_metadata_from_resource(
"https://protected.example.com/.well-known/oauth",
"https://protected.example.com/mcp",
)
assert captured_kwargs.get("follow_redirects") is False
@pytest.mark.asyncio
async def test_same_authority_auth_server_fetch_does_not_follow_redirects(self):
# Same redirect-bypass concern for the authorization-server fetch path.
manager = MCPServerManager()
mock_response = MagicMock()
mock_response.json.return_value = {
"authorization_endpoint": "https://provider.example.com/authorize",
"token_endpoint": "https://provider.example.com/token",
}
mock_response.raise_for_status = MagicMock()
captured_kwargs: Dict[str, Any] = {}
async def fake_get(url, **kwargs):
captured_kwargs.update(kwargs)
return mock_response
mock_client = MagicMock()
mock_client.get = AsyncMock(side_effect=fake_get)
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
await manager._fetch_single_authorization_server_metadata(
"https://provider.example.com",
"https://provider.example.com",
)
assert captured_kwargs.get("follow_redirects") is False
@pytest.mark.asyncio
async def test_fetch_authorization_server_refuses_unsafe_issuer(self, monkeypatch):
# Mirrors the GHSA-mrfv repro: PRM lists a loopback issuer URL.
self._patch_resolves(monkeypatch, {"attacker.example.com": ["127.0.0.1"]})
manager = MCPServerManager()
mock_client = MagicMock()
mock_client.get = AsyncMock()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
):
result = await manager._fetch_single_authorization_server_metadata(
"http://attacker.example.com:19999",
"https://legit-mcp.example.com/mcp",
)
assert result is None
mock_client.get.assert_not_called()
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -2078,6 +2078,33 @@ class TestGuardrailModificationCheck:
)
assert exc.value.status_code == 403
@pytest.mark.parametrize(
"key",
[
"guardrails",
"disable_global_guardrails",
"disable_global_guardrail",
"opted_out_global_guardrails",
],
)
@pytest.mark.parametrize("empty_value", [{}, [], "", 0, False])
def test_rejects_empty_value_modification(self, key, empty_value):
"""Regression: an explicitly-supplied empty/falsy value still expresses
intent to modify and must trigger the permission check. Truthiness-based
gating let callers bypass the check by sending e.g.
``metadata={"guardrails": {}}``, which downstream evaluation interpreted
as "disable all guardrails" while the auth layer treated it as no-op.
"""
from fastapi import HTTPException
with patch(
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
return_value=False,
):
with pytest.raises(HTTPException) as exc:
self._call({"metadata": {key: empty_value}})
assert exc.value.status_code == 403
def test_rejects_injection_via_litellm_metadata_key(self):
"""Caller can populate the OTHER metadata key; that must also 403."""
from fastapi import HTTPException
@ -2336,3 +2363,134 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau
)
assert exc_info.value.current_cost == 250.0
assert exc_info.value.max_budget == 200.0
@pytest.mark.asyncio
async def test_team_member_budget_check_null_clone_falls_back_to_team_default():
"""Per-member NULL max_budget falls through to the team default cap."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
from litellm.proxy.utils import ProxyLogging
team_object = LiteLLM_TeamTable(
team_id="test-team",
metadata={"team_member_budget_id": "budget-default"},
)
user_object = LiteLLM_UserTable(user_id="test-user")
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user",
team_id="test-team",
)
# Per-member row exists with NULL max_budget (the cloned-from-incomplete-default case).
team_membership = LiteLLM_TeamMembership(
user_id="test-user",
team_id="test-team",
spend=0.0,
budget_id="budget-clone",
litellm_budget_table=LiteLLM_BudgetTable(max_budget=None),
)
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
fake_default_row = MagicMock()
fake_default_row.max_budget = 65.0
fake_default_row.dict = MagicMock(
return_value={"budget_id": "budget-default", "max_budget": 65.0}
)
prisma_client = MagicMock()
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=fake_default_row
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == "spend:team_member:test-user:test-team":
return 500.0
return fallback_spend
with (
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new_callable=AsyncMock,
return_value=team_membership,
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=DualCache(),
proxy_logging_obj=proxy_logging_obj,
)
assert exc_info.value.current_cost == 500.0
assert exc_info.value.max_budget == 65.0
prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once()
@pytest.mark.asyncio
async def test_team_member_budget_check_null_clone_with_null_default_skips_enforcement():
"""When per-member and team default are both NULL, enforcement still skips."""
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
from litellm.proxy.utils import ProxyLogging
team_object = LiteLLM_TeamTable(
team_id="test-team",
metadata={"team_member_budget_id": "budget-default"},
)
user_object = LiteLLM_UserTable(user_id="test-user")
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user",
team_id="test-team",
)
team_membership = LiteLLM_TeamMembership(
user_id="test-user",
team_id="test-team",
spend=0.0,
budget_id="budget-clone",
litellm_budget_table=LiteLLM_BudgetTable(max_budget=None),
)
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
fake_default_row = MagicMock()
fake_default_row.max_budget = None
fake_default_row.dict = MagicMock(
return_value={"budget_id": "budget-default", "max_budget": None}
)
prisma_client = MagicMock()
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=fake_default_row
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == "spend:team_member:test-user:test-team":
return 1000.0
return fallback_spend
with (
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new_callable=AsyncMock,
return_value=team_membership,
),
):
# No raise: both rows are NULL, so enforcement is correctly skipped.
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=DualCache(),
proxy_logging_obj=proxy_logging_obj,
)

View file

@ -6,11 +6,12 @@ This module tests the auth commands and their associated functionality.
import pytest
import requests
from unittest.mock import AsyncMock, patch, Mock, call
from unittest.mock import patch, Mock, call
from litellm.proxy.client.cli.commands.auth import (
_normalize_teams,
_poll_for_ready_data,
_poll_for_authentication,
_start_cli_sso_flow,
)
@ -57,6 +58,18 @@ async def test_normalize_teams_with_details_with_aliases():
]
@patch("litellm.proxy.client.cli.commands.auth.requests.post")
def test_start_cli_sso_flow_rejects_invalid_response(request_mock):
"""Test CLI SSO start rejects malformed server responses"""
response = Mock()
response.raise_for_status = Mock()
response.json.return_value = {"login_id": "cli-session", "user_code": "ABCD-EFGH"}
request_mock.return_value = response
with pytest.raises(ValueError, match="Invalid CLI SSO start response"):
_start_cli_sso_flow("https://litellm.com")
@pytest.mark.asyncio
@patch(
"litellm.proxy.client.cli.commands.auth.requests.get",
@ -195,10 +208,11 @@ async def test_poll_for_ready_connection_failure(sleep_mock, click_mock, request
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
async def test_poll_for_authentication_no_data(click_mock, poll_mock, handle_mock):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual is None
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_not_called()
@ -214,10 +228,11 @@ async def test_poll_for_authentication_no_data(click_mock, poll_mock, handle_moc
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
async def test_poll_for_authentication_no_teams(click_mock, poll_mock, handle_mock):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual is None
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_not_called()
@ -243,7 +258,7 @@ async def test_poll_for_authentication_team_selection_success(
click_mock, poll_mock, handle_mock
):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual == {
"api_key": "jwt-123",
"user_id": "user-123",
@ -252,11 +267,13 @@ async def test_poll_for_authentication_team_selection_success(
}
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_called_once_with(
base_url="https://litellm.com",
key_id="key-123",
poll_secret="poll-secret",
teams=[
{"team_id": "1", "team_alias": None},
{"team_id": "2", "team_alias": None},
@ -283,15 +300,17 @@ async def test_poll_for_authentication_team_selection_cancelled(
click_mock, poll_mock, handle_mock
):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual is None
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_called_once_with(
base_url="https://litellm.com",
key_id="key-123",
poll_secret="poll-secret",
teams=[{"team_id": "team-1", "team_alias": None}],
)
click_mock.assert_called_once()
@ -314,7 +333,7 @@ async def test_poll_for_authentication_auto_assigned_team(
click_mock, poll_mock, handle_mock
):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual == {
"api_key": "jwt-456",
"user_id": "user-456",
@ -323,6 +342,7 @@ async def test_poll_for_authentication_auto_assigned_team(
}
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_not_called()

View file

@ -2,14 +2,16 @@
Tests for the invite-link onboarding endpoints.
Covers the security behavior of:
GET /onboarding/get_token – rejects already-used links before showing any user data
POST /onboarding/claim_token – rejects already-used links; marks is_accepted=True only
after the password is successfully written
GET /onboarding/get_token – rejects already-used links and returns only a
short-lived onboarding token, not a UI session key
POST /onboarding/claim_token – requires that onboarding token; mints the UI
session key only after the password is written
"""
from datetime import timedelta
from unittest.mock import AsyncMock, MagicMock, patch
import jwt
import pytest
from fastapi import HTTPException
@ -22,14 +24,27 @@ from litellm.proxy._types import InvitationClaim
# ---------------------------------------------------------------------------
def _make_invite(*, is_accepted: bool, expired: bool = False) -> MagicMock:
class _AsyncTx:
def __init__(self, db: MagicMock):
self.db = db
async def __aenter__(self) -> MagicMock:
return self.db
async def __aexit__(self, exc_type, exc, tb):
return False
def _make_invite(
*, is_accepted: bool, expired: bool = False, claimed: bool = False
) -> MagicMock:
now = litellm.utils.get_utc_datetime()
invite = MagicMock()
invite.id = "invite-abc"
invite.user_id = "user-123"
invite.is_accepted = is_accepted
invite.expires_at = now - timedelta(days=1) if expired else now + timedelta(days=6)
invite.accepted_at = None
invite.accepted_at = now if claimed else None
return invite
@ -45,11 +60,39 @@ def _make_prisma(invite: MagicMock, user: MagicMock | None = None) -> MagicMock:
prisma = MagicMock()
prisma.db.litellm_invitationlink.find_unique = AsyncMock(return_value=invite)
prisma.db.litellm_invitationlink.update = AsyncMock()
prisma.db.litellm_invitationlink.update_many = AsyncMock(return_value=1)
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user)
prisma.db.litellm_usertable.update = AsyncMock(return_value=user)
prisma.db.tx = MagicMock(return_value=_AsyncTx(prisma.db))
return prisma
def _make_onboarding_token(
*,
invitation_link: str = "invite-abc",
user_id: str = "user-123",
token_type: str = "litellm_onboarding",
master_key: str = "sk-test",
) -> str:
return jwt.encode(
{
"token_type": token_type,
"invitation_link": invitation_link,
"user_id": user_id,
"exp": litellm.utils.get_utc_datetime() + timedelta(minutes=15),
},
master_key,
algorithm="HS256",
)
def _make_claim_request(token: str | None = None) -> MagicMock:
request = MagicMock()
request.headers = {"Authorization": f"Bearer {token}"} if token is not None else {}
request.base_url = "http://localhost:4000/"
return request
# ---------------------------------------------------------------------------
# GET /onboarding/get_token
# ---------------------------------------------------------------------------
@ -120,10 +163,10 @@ async def test_get_token_rejects_missing_link():
@pytest.mark.asyncio
async def test_get_token_does_not_set_is_accepted():
async def test_get_token_returns_onboarding_token_without_minting_ui_key():
"""
A valid, unused link should succeed and must NOT flip is_accepted to True.
That flag is only written after the password is claimed.
A valid, unused link should return a short-lived onboarding token, but
must not reserve the invite or mint a usable UI/API key on GET.
"""
from litellm.proxy.proxy_server import onboarding
@ -133,6 +176,252 @@ async def test_get_token_does_not_set_is_accepted():
request = MagicMock()
request.base_url = "http://localhost:4000/"
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.premium_user", False),
patch(
"litellm.proxy.proxy_server.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_generate_key,
patch(
"litellm.proxy.proxy_server.get_custom_url",
return_value="http://localhost:4000/",
),
patch(
"litellm.proxy.proxy_server.get_disabled_non_admin_personal_key_creation",
return_value=False,
),
patch("litellm.proxy.proxy_server.get_server_root_path", return_value=""),
):
result = await onboarding(invite_link="invite-abc", request=request)
# Endpoint succeeded
assert "token" in result
assert "login_url" in result
outer_claims = jwt.decode(result["token"], "sk-test", algorithms=["HS256"])
onboarding_token = outer_claims["key"]
onboarding_claims = jwt.decode(onboarding_token, "sk-test", algorithms=["HS256"])
assert onboarding_claims["token_type"] == "litellm_onboarding"
assert onboarding_claims["invitation_link"] == "invite-abc"
assert onboarding_claims["user_id"] == "user-123"
assert not onboarding_token.startswith("sk-")
mock_generate_key.assert_not_called()
prisma.db.litellm_invitationlink.update_many.assert_not_called()
prisma.db.litellm_invitationlink.update.assert_not_called()
# ---------------------------------------------------------------------------
# POST /onboarding/claim_token
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_claim_token_rejects_already_used_link():
"""
If is_accepted is True, the password has already been set.
A second claim attempt must be rejected with 401.
"""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=True, claimed=True)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=_make_claim_request())
assert exc_info.value.status_code == 401
assert "already been used" in exc_info.value.detail["error"]
# Password must never have been written
prisma.db.litellm_usertable.update.assert_not_called()
@pytest.mark.asyncio
async def test_claim_token_rejects_expired_link():
"""An expired link must be rejected even if is_accepted is False."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False, expired=True)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=_make_claim_request())
assert exc_info.value.status_code == 401
assert "expired" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_claim_token_rejects_mismatched_user_id():
"""The user_id in the request must match the one on the invite."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="wrong-user",
password="NewP@ssw0rd",
)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=_make_claim_request())
assert exc_info.value.status_code == 401
assert "does not match" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_claim_token_rejects_missing_onboarding_token():
"""The password endpoint must require the onboarding token returned by get_token."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch("litellm.proxy.proxy_server.general_settings", {}),
):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=_make_claim_request())
assert exc_info.value.status_code == 401
assert "Missing onboarding session" in exc_info.value.detail["error"]
prisma.db.litellm_usertable.update.assert_not_called()
@pytest.mark.asyncio
async def test_claim_token_rejects_wrong_onboarding_session():
"""The onboarding token must be bound to the invite and user being claimed."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
request = _make_claim_request(
_make_onboarding_token(invitation_link="other-invite")
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch("litellm.proxy.proxy_server.general_settings", {}),
):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=request)
assert exc_info.value.status_code == 401
assert "Invalid onboarding session" in exc_info.value.detail["error"]
prisma.db.litellm_usertable.update.assert_not_called()
@pytest.mark.asyncio
async def test_claim_token_rejects_invalid_bearer_token():
"""A regular API key must not be accepted as an onboarding token."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
request = _make_claim_request("sk-regular-key")
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch("litellm.proxy.proxy_server.general_settings", {}),
):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=request)
assert exc_info.value.status_code == 401
assert "Invalid onboarding session" in exc_info.value.detail["error"]
prisma.db.litellm_usertable.update.assert_not_called()
@pytest.mark.asyncio
async def test_claim_token_rejects_concurrent_reuse_before_password_write():
"""Only the first valid claim may reserve the invitation."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
prisma = _make_prisma(invite)
prisma.db.litellm_invitationlink.update_many = AsyncMock(return_value=0)
request = _make_claim_request(_make_onboarding_token())
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch(
"litellm.proxy.proxy_server.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_generate_key,
):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=request)
assert exc_info.value.status_code == 401
assert "already been used" in exc_info.value.detail["error"]
prisma.db.litellm_usertable.update.assert_not_called()
mock_generate_key.assert_not_called()
@pytest.mark.asyncio
async def test_claim_token_sets_accepted_at_after_password_written():
"""
A valid first-time claim must:
1. Write the hashed password to the user table.
2. Set accepted_at on the invitation link after the password write succeeds.
"""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
user = _make_user()
prisma = _make_prisma(invite, user)
request = _make_claim_request(_make_onboarding_token())
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
mock_token_response = {"token": "sk-generated-key", "user_id": "user-123"}
with (
@ -155,113 +444,13 @@ async def test_get_token_does_not_set_is_accepted():
),
patch("litellm.proxy.proxy_server.get_server_root_path", return_value=""),
):
result = await onboarding(invite_link="invite-abc", request=request)
# Endpoint succeeded
assert "token" in result
assert "login_url" in result
# is_accepted must NOT have been updated here
prisma.db.litellm_invitationlink.update.assert_not_called()
# ---------------------------------------------------------------------------
# POST /onboarding/claim_token
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_claim_token_rejects_already_used_link():
"""
If is_accepted is True, the password has already been set.
A second claim attempt must be rejected with 401.
"""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=True)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data)
assert exc_info.value.status_code == 401
assert "already been used" in exc_info.value.detail["error"]
# Password must never have been written
prisma.db.litellm_usertable.update.assert_not_called()
@pytest.mark.asyncio
async def test_claim_token_rejects_expired_link():
"""An expired link must be rejected even if is_accepted is False."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False, expired=True)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data)
assert exc_info.value.status_code == 401
assert "expired" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_claim_token_rejects_mismatched_user_id():
"""The user_id in the request must match the one on the invite."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
prisma = _make_prisma(invite)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="wrong-user",
password="NewP@ssw0rd",
)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data)
assert exc_info.value.status_code == 401
assert "does not match" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_claim_token_sets_is_accepted_after_password_written():
"""
A valid first-time claim must:
1. Write the hashed password to the user table.
2. Flip is_accepted to True on the invitation link — and only after the
password write succeeds.
"""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
user = _make_user()
prisma = _make_prisma(invite, user)
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
result = await claim_onboarding_link(data=data)
result = await claim_onboarding_link(data=data, request=request)
# Password was written
prisma.db.litellm_invitationlink.update_many.assert_called_once()
reserve_kwargs = prisma.db.litellm_invitationlink.update_many.call_args.kwargs
assert reserve_kwargs["where"] == {"id": "invite-abc", "is_accepted": False}
assert reserve_kwargs["data"]["is_accepted"] is True
prisma.db.litellm_usertable.update.assert_called_once()
call_kwargs = prisma.db.litellm_usertable.update.call_args
assert call_kwargs.kwargs["where"] == {"user_id": "user-123"}
@ -270,5 +459,50 @@ async def test_claim_token_sets_is_accepted_after_password_written():
# is_accepted was flipped to True on the invitation link
prisma.db.litellm_invitationlink.update.assert_called_once()
link_update_data = prisma.db.litellm_invitationlink.update.call_args.kwargs["data"]
assert link_update_data["is_accepted"] is True
assert "is_accepted" not in link_update_data
assert link_update_data["accepted_at"] is not None
outer_claims = jwt.decode(result["token"], "sk-test", algorithms=["HS256"])
assert outer_claims["key"] == "sk-generated-key"
@pytest.mark.asyncio
async def test_claim_token_rolls_back_invite_when_session_key_mint_fails():
"""A session key failure must not leave the invite permanently consumed."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
user = _make_user()
prisma = _make_prisma(invite, user)
request = _make_claim_request(_make_onboarding_token())
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd",
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch(
"litellm.proxy.proxy_server.generate_key_helper_fn",
new_callable=AsyncMock,
side_effect=Exception("key mint failed"),
),
):
with pytest.raises(HTTPException) as exc_info:
await claim_onboarding_link(data=data, request=request)
assert exc_info.value.status_code == 500
assert "Failed to create onboarding session" in exc_info.value.detail["error"]
assert prisma.db.litellm_invitationlink.update_many.call_count == 2
rollback_kwargs = prisma.db.litellm_invitationlink.update_many.call_args_list[
1
].kwargs
assert rollback_kwargs["where"] == {
"id": "invite-abc",
"is_accepted": True,
}
assert rollback_kwargs["data"]["accepted_at"] is None
assert rollback_kwargs["data"]["is_accepted"] is False

View file

@ -2581,3 +2581,49 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
@pytest.mark.asyncio
async def test_master_key_auth_substitutes_alias_for_api_key():
"""
When the master key authenticates a request, the resulting
``UserAPIKeyAuth.api_key`` must be the stable alias
``LITELLM_PROXY_MASTER_KEY_ALIAS`` — never the raw master key (which
would propagate downstream and be hashed into spend logs, Prometheus
``/metrics`` labels, or audit trails) and never the master-key hash.
"""
from fastapi import Request
from starlette.datastructures import URL
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
from litellm.proxy.utils import hash_token
import litellm.proxy.proxy_server as _proxy_server_mod
attrs = _proxy_server_attrs_for_custom_auth(user_custom_auth=None)
master_key = attrs["master_key"]
_orig = {k: getattr(_proxy_server_mod, k, None) for k in attrs}
try:
for k, v in attrs.items():
setattr(_proxy_server_mod, k, v)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
result = await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {master_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
assert result.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS
assert result.api_key != master_key
assert result.api_key != hash_token(master_key)
finally:
for k, v in _orig.items():
setattr(_proxy_server_mod, k, v)

View file

@ -1,17 +1,15 @@
import json
import os
import sys
import tempfile
import time
from pathlib import Path
from unittest.mock import MagicMock, Mock, mock_open, patch
from unittest.mock import Mock, mock_open, patch
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import pytest
from click.testing import CliRunner
from litellm.proxy.client.cli.commands.auth import (
@ -26,6 +24,22 @@ from litellm.proxy.client.cli.commands.auth import (
)
def _mock_cli_sso_start_response(
login_id: str = "cli-session-uuid-456",
poll_secret: str = "poll-secret",
user_code: str = "ABCD-EFGH",
) -> Mock:
mock_response = Mock()
mock_response.status_code = 200
mock_response.json.return_value = {
"login_id": login_id,
"poll_secret": poll_secret,
"user_code": user_code,
}
mock_response.raise_for_status = Mock()
return mock_response
class TestTokenUtilities:
"""Test token file utility functions"""
@ -243,12 +257,15 @@ class TestLoginCommand:
with (
patch("webbrowser.open") as mock_browser,
patch(
"requests.post",
return_value=_mock_cli_sso_start_response(login_id="cli-test-uuid-123"),
) as mock_post,
patch("requests.get", return_value=mock_response) as mock_get,
patch("litellm.proxy.client.cli.commands.auth.save_token") as mock_save,
patch(
"litellm.proxy.client.cli.interface.show_commands"
) as mock_show_commands,
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -261,7 +278,13 @@ class TestLoginCommand:
mock_browser.assert_called_once()
call_args = mock_browser.call_args[0][0]
assert "https://test.example.com/sso/key/generate" in call_args
assert "sk-test-uuid-123" in call_args
assert "cli-test-uuid-123" in call_args
assert "Verification code: ABCD-EFGH" in result.output
mock_post.assert_called_once()
mock_get.assert_called()
assert mock_get.call_args.kwargs["headers"] == {
"x-litellm-cli-poll-secret": "poll-secret"
}
# Verify JWT was saved
mock_save.assert_called_once()
@ -284,9 +307,9 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", return_value=mock_response),
patch("time.sleep") as mock_sleep,
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
patch("time.sleep"),
):
# Mock time.sleep to avoid actual delays in tests
@ -306,9 +329,9 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", return_value=mock_response),
patch("time.sleep"),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -325,12 +348,12 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch(
"requests.get",
side_effect=requests.RequestException("Connection failed"),
),
patch("time.sleep"),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -345,8 +368,8 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", side_effect=KeyboardInterrupt),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -369,9 +392,9 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", return_value=mock_response),
patch("time.sleep"),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -386,8 +409,8 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", side_effect=ValueError("Invalid value")),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -556,6 +579,12 @@ class TestCLIKeyRegenerationFlow:
# Simulate user selecting team #2 (team-beta)
with (
patch("webbrowser.open") as mock_browser,
patch(
"requests.post",
return_value=_mock_cli_sso_start_response(
login_id="cli-session-uuid-456"
),
),
patch(
"requests.get", side_effect=[mock_first_response, mock_second_response]
) as mock_get,
@ -563,7 +592,6 @@ class TestCLIKeyRegenerationFlow:
patch(
"litellm.proxy.client.cli.interface.show_commands"
) as mock_show_commands,
patch("litellm._uuid.uuid.uuid4", return_value="session-uuid-456"),
patch("click.prompt", return_value="2"),
): # User selects index 2
@ -585,8 +613,11 @@ class TestCLIKeyRegenerationFlow:
# First poll should be without team_id
first_poll_url = mock_get.call_args_list[0][0][0]
assert "sk-session-uuid-456" in first_poll_url
assert "cli-session-uuid-456" in first_poll_url
assert "team_id=" not in first_poll_url
assert mock_get.call_args_list[0].kwargs["headers"] == {
"x-litellm-cli-poll-secret": "poll-secret"
}
# Second poll should include team_id=team-beta
second_poll_url = mock_get.call_args_list[1][0][0]
@ -621,10 +652,15 @@ class TestCLIKeyRegenerationFlow:
with (
patch("webbrowser.open") as mock_browser,
patch(
"requests.post",
return_value=_mock_cli_sso_start_response(
login_id="cli-session-uuid-solo"
),
),
patch("requests.get", return_value=mock_response),
patch("litellm.proxy.client.cli.commands.auth.save_token") as mock_save,
patch("litellm.proxy.client.cli.interface.show_commands"),
patch("litellm._uuid.uuid.uuid4", return_value="session-uuid-solo"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -637,7 +673,7 @@ class TestCLIKeyRegenerationFlow:
call_args = mock_browser.call_args[0][0]
assert "https://test.example.com/sso/key/generate" in call_args
assert "source=litellm-cli" in call_args
assert "key=sk-session-uuid-solo" in call_args
assert "key=cli-session-uuid-solo" in call_args
# Verify JWT was saved
mock_save.assert_called_once()

View file

@ -104,7 +104,9 @@ def test_remove_sensitive_info_from_deployment_with_excluded_keys():
assert sanitized_config["litellm_params"]["access_token"] != "token-12345"
assert "*" in sanitized_config["litellm_params"]["access_token"]
# With excluded_keys, litellm_credentials_name should NOT be masked (even if it would match patterns)
# With excluded_keys, litellm_credentials_name should NOT be masked.
# ``remove_sensitive_info_from_deployment`` mutates its input, so feed it
# a fresh copy rather than the already-sanitized one.
sanitized_config = remove_sensitive_info_from_deployment(
copy.deepcopy(base_config), excluded_keys={"litellm_credentials_name"}
)

View file

@ -0,0 +1,97 @@
"""
Unit tests for unauthenticated logo / favicon endpoint helpers.
Local image paths are an existing deployment workflow, so the helper keeps
arbitrary local image paths working while refusing non-image files like
``/etc/passwd`` or ``/proc/self/environ``.
"""
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.proxy.common_utils.static_asset_utils import (
detect_local_image_media_type,
resolve_validated_local_image_path,
)
@pytest.mark.parametrize(
("body", "media_type"),
[
(b"\x89PNG\r\n\x1a\nfake png body", "image/png"),
(b"GIF89a fake gif body", "image/gif"),
(b"\xff\xd8\xff fake jpeg body", "image/jpeg"),
(b"RIFF\x00\x00\x00\x00WEBP fake webp body", "image/webp"),
(b"\x00\x00\x01\x00 fake ico body", "image/x-icon"),
],
)
def test_detect_local_image_media_type_accepts_supported_images(body, media_type):
assert detect_local_image_media_type(body) == media_type
def test_detect_local_image_media_type_rejects_non_images():
assert detect_local_image_media_type(b"root:x:0:0:root:/root:/bin/bash") is None
class TestResolveValidatedLocalImagePath:
def test_returns_resolved_path_for_arbitrary_local_image(self, tmp_path):
logo = tmp_path / "logo.png"
logo.write_bytes(b"\x89PNG\r\n\x1a\nfake png body")
result = resolve_validated_local_image_path(str(logo))
assert result == (str(logo.resolve()), "image/png")
def test_rejects_etc_passwd(self):
result = resolve_validated_local_image_path("/etc/passwd")
assert result is None
def test_rejects_proc_self_environ(self):
result = resolve_validated_local_image_path("/proc/self/environ")
assert result is None
def test_rejects_symlink_pointing_to_non_image(self, tmp_path):
secret = tmp_path / "secret.txt"
secret.write_text("password=hunter2")
symlink = tmp_path / "logo.png"
os.symlink(str(secret), str(symlink))
result = resolve_validated_local_image_path(str(symlink))
assert result is None
def test_accepts_symlink_pointing_to_image(self, tmp_path):
logo = tmp_path / "real_logo.png"
logo.write_bytes(b"\x89PNG\r\n\x1a\nfake png body")
symlink = tmp_path / "logo.png"
os.symlink(str(logo), str(symlink))
result = resolve_validated_local_image_path(str(symlink))
assert result == (str(logo.resolve()), "image/png")
def test_rejects_path_traversal_to_non_image(self, tmp_path):
assets_dir = tmp_path / "assets"
assets_dir.mkdir()
secret = tmp_path / "secret.txt"
secret.write_text("nope")
traversal = str(assets_dir / ".." / "secret.txt")
result = resolve_validated_local_image_path(traversal)
assert result is None
def test_rejects_directory(self, tmp_path):
result = resolve_validated_local_image_path(str(tmp_path))
assert result is None
def test_rejects_nonexistent_file(self, tmp_path):
result = resolve_validated_local_image_path(str(tmp_path / "missing.jpg"))
assert result is None
def test_rejects_empty_path(self):
assert resolve_validated_local_image_path("") is None

View file

@ -4,7 +4,7 @@ Test to verify the Google GenAI proxy API endpoints
"""
import os
import sys
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -13,520 +13,171 @@ sys.path.insert(
) # Adds the parent directory to the system path
def test_google_generate_content_endpoint():
"""Test that the google_generate_content endpoint correctly routes requests"""
# Skip this test if we can't import the required modules due to missing dependencies
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
def _build_test_client():
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
from litellm.proxy.google_endpoints.endpoints import router as google_router
app = FastAPI()
app.include_router(google_router)
return TestClient(app)
def _patch_base_process(return_value=None):
"""Patch ProxyBaseLLMRequestProcessing.base_process_llm_request so endpoint
tests don't run the full pipeline. Returns the AsyncMock so callers can
inspect call args."""
if return_value is None:
return_value = {"test": "response"}
return patch(
"litellm.proxy.google_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new_callable=AsyncMock,
return_value=return_value,
)
def test_google_generate_content_endpoint():
"""generateContent routes through ProxyBaseLLMRequestProcessing with the
agenerate_content route_type — that pipeline runs pre_call_hook +
during_call_hook + post_call_success_hook for every guardrail callback."""
try:
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock the router's agenerate_content method
with patch("litellm.proxy.proxy_server.llm_router") as mock_router:
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Send a request to the endpoint
with _patch_base_process() as mock_base:
response = client.post(
"/v1beta/models/test-model:generateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
)
# Verify the response
assert response.status_code == 200
assert response.json() == {"test": "response"}
# Verify that agenerate_content was called
mock_router.agenerate_content.assert_called_once()
mock_base.assert_called_once()
kwargs = mock_base.call_args.kwargs
assert kwargs["route_type"] == "agenerate_content"
assert kwargs["model"] == "test-model"
def test_google_stream_generate_content_endpoint():
"""Test that the google_stream_generate_content endpoint correctly routes streaming requests"""
# Skip this test if we can't import the required modules due to missing dependencies
"""streamGenerateContent must route through the same processor with the
streaming route_type so the guardrail pipeline runs."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock the router's agenerate_content_stream method to return a stream
async def mock_stream_generator():
yield 'data: {"test": "stream_chunk_1"}\n\n'
yield 'data: {"test": "stream_chunk_2"}\n\n'
yield "data: [DONE]\n\n"
with patch("litellm.proxy.proxy_server.llm_router") as mock_router:
mock_router.agenerate_content_stream = AsyncMock(
return_value=mock_stream_generator()
)
# Send a request to the endpoint
with (
_patch_base_process() as mock_base,
patch(
"litellm.proxy.google_endpoints.endpoints.ProxyBaseLLMRequestProcessing.__init__",
return_value=None,
) as mock_init,
):
response = client.post(
"/v1beta/models/test-model:streamGenerateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
)
# Verify the response
assert response.status_code == 200
mock_base.assert_called_once()
kwargs = mock_base.call_args.kwargs
assert kwargs["route_type"] == "agenerate_content_stream"
assert kwargs["model"] == "test-model"
# Verify that agenerate_content_stream was called with correct parameters
mock_router.agenerate_content_stream.assert_called_once()
call_args = mock_router.agenerate_content_stream.call_args
assert call_args[1]["stream"] is True
assert call_args[1]["model"] == "test-model"
assert call_args[1]["contents"] == [
# stream=True must be forced into the data the processor receives.
init_kwargs = mock_init.call_args.kwargs
assert init_kwargs["data"]["stream"] is True
assert init_kwargs["data"]["model"] == "test-model"
assert init_kwargs["data"]["contents"] == [
{"role": "user", "parts": [{"text": "Hello"}]}
]
def test_google_generate_content_with_cost_tracking_metadata():
"""Test that the google_generate_content endpoint includes user metadata for cost tracking"""
def test_google_generate_content_data_flows_through_processor():
"""The body the client sends must reach ProxyBaseLLMRequestProcessing
intact so the pipeline can apply guardrails to it."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
_patch_base_process(),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
"litellm.proxy.google_endpoints.endpoints.ProxyBaseLLMRequestProcessing.__init__",
return_value=None,
) as mock_init,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to return data with metadata
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
# Simulate adding user metadata
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
"user_api_key_team_id": "test-team-id",
"user_api_key": "hashed-key",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request to the endpoint
response = client.post(
client.post(
"/v1beta/models/test-model:generateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that add_litellm_data_to_request was called
mock_add_data.assert_called_once()
# Verify that agenerate_content was called with metadata
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that litellm_metadata exists and contains user information
assert "litellm_metadata" in called_data
assert called_data["litellm_metadata"]["user_api_key_user_id"] == "test-user-id"
assert called_data["litellm_metadata"]["user_api_key_team_id"] == "test-team-id"
def test_google_stream_generate_content_with_cost_tracking_metadata():
"""Test that the google_stream_generate_content endpoint includes user metadata for cost tracking"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock the router's agenerate_content_stream method to return a stream
mock_stream = AsyncMock()
mock_stream.__aiter__ = lambda self: mock_stream
mock_stream.__anext__.side_effect = StopAsyncIteration
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content_stream = AsyncMock(return_value=mock_stream)
# Mock add_litellm_data_to_request to return data with metadata
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
# Simulate adding user metadata
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
"user_api_key_team_id": "test-team-id",
"user_api_key": "hashed-key",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request to the endpoint
response = client.post(
"/v1beta/models/test-model:streamGenerateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that add_litellm_data_to_request was called
mock_add_data.assert_called_once()
# Verify that agenerate_content_stream was called with metadata
mock_router.agenerate_content_stream.assert_called_once()
call_args = mock_router.agenerate_content_stream.call_args
called_data = call_args[1]
# Verify that litellm_metadata exists and contains user information
assert "litellm_metadata" in called_data
assert called_data["litellm_metadata"]["user_api_key_user_id"] == "test-user-id"
assert called_data["litellm_metadata"]["user_api_key_team_id"] == "test-team-id"
# Verify stream is set to True
assert called_data["stream"] is True
def test_google_generate_content_with_system_instruction():
"""
Test that systemInstruction is correctly passed through from the endpoint to the router.
This test verifies the fix for systemInstruction being dropped when forwarding
requests to Vertex AI through the Google GenAI endpoint.
"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to pass through data unchanged
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
return data
mock_add_data.side_effect = mock_add_litellm_data
# Define the systemInstruction to test
system_instruction = {"parts": [{"text": "Your name is Doodle."}]}
# Send a request with systemInstruction
response = client.post(
"/v1beta/models/gemini-2.5-pro:generateContent",
json={
"systemInstruction": system_instruction,
"contents": [
{"parts": [{"text": "What is your name?"}], "role": "user"}
],
},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that agenerate_content was called
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that systemInstruction is present in the call arguments
assert "systemInstruction" in called_data
assert called_data["systemInstruction"] == system_instruction
assert (
called_data["systemInstruction"]["parts"][0]["text"]
== "Your name is Doodle."
)
# Verify contents are also present
assert "contents" in called_data
assert len(called_data["contents"]) == 1
assert called_data["contents"][0]["role"] == "user"
def test_google_generate_content_with_image_config():
"""
Test that imageConfig is correctly passed through from generationConfig to the router.
This test verifies that imageConfig parameters (aspectRatio, imageSize) are preserved
when forwarding requests to Google GenAI through the endpoint.
"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to pass through data unchanged
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request with generationConfig containing imageConfig
response = client.post(
"/v1beta/models/gemini-3-pro-image-preview:generateContent",
json={
"contents": [
{
"role": "user",
"parts": [
{
"text": "Create a vibrant infographic about photosynthesis"
}
],
}
],
"contents": [{"role": "user", "parts": [{"text": "Hello"}]}],
"systemInstruction": {"parts": [{"text": "Your name is Doodle."}]},
"generationConfig": {
"responseModalities": ["TEXT", "IMAGE"],
"imageConfig": {"aspectRatio": "9:16", "imageSize": "4K"},
},
},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that agenerate_content was called
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that config is present in the call arguments
assert "config" in called_data
# Verify that imageConfig is preserved in the config
assert "imageConfig" in called_data["config"]
assert called_data["config"]["imageConfig"]["aspectRatio"] == "9:16"
assert called_data["config"]["imageConfig"]["imageSize"] == "4K"
# Verify that responseModalities is also preserved
assert "responseModalities" in called_data["config"]
assert called_data["config"]["responseModalities"] == ["TEXT", "IMAGE"]
# Verify contents are also present
assert "contents" in called_data
assert len(called_data["contents"]) == 1
assert called_data["contents"][0]["role"] == "user"
data = mock_init.call_args.kwargs["data"]
assert data["model"] == "test-model"
assert data["contents"] == [{"role": "user", "parts": [{"text": "Hello"}]}]
assert data["systemInstruction"] == {
"parts": [{"text": "Your name is Doodle."}]
}
# generationConfig arrives intact here; the rename to `config` is
# done downstream in route_request (see test_route_llm_request).
assert data["generationConfig"]["responseModalities"] == ["TEXT", "IMAGE"]
assert data["generationConfig"]["imageConfig"]["aspectRatio"] == "9:16"
def test_google_generate_content_metadata_and_trace_id_callbacks():
"""Test that google_generate_content sets litellm_call_id and logging_obj for callbacks (e.g. S3, Langfuse)"""
def test_google_generate_content_forwards_call_id_header():
"""The endpoint must forward the x-litellm-call-id header to the processor
so the helper can stamp it on the logging object. Trace continuity from
client → callbacks (S3, Langfuse, etc.) depends on this header surviving
the hop through these endpoints."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to return data with metadata
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
# Simulate adding user metadata
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request to the endpoint with x-litellm-call-id header
test_call_id = "test-custom-call-id"
response = client.post(
with _patch_base_process() as mock_base:
client.post(
"/v1beta/models/test-model:generateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
headers={
"Authorization": "Bearer sk-test-key",
"x-litellm-call-id": test_call_id,
},
headers={"x-litellm-call-id": "trace-abc-123"},
)
assert response.status_code == 200
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that the litellm_logging_obj got assigned in the final called_data to router
assert "litellm_logging_obj" in called_data
assert "litellm_call_id" in called_data
assert called_data["litellm_call_id"] == test_call_id
forwarded_request = mock_base.call_args.kwargs["request"]
assert forwarded_request.headers.get("x-litellm-call-id") == "trace-abc-123"
def test_google_stream_generate_content_metadata_and_trace_id_callbacks():
"""Test that google_stream_generate_content sets litellm_call_id and logging_obj for callbacks"""
def test_google_count_tokens_unchanged():
"""countTokens has its own path and isn't affected by the pipeline change."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
app = FastAPI()
app.include_router(google_router)
client = TestClient(app)
fake_response = MagicMock()
fake_response.original_response = {
"totalTokens": 7,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 7}],
}
fake_response.total_tokens = 7
mock_stream = AsyncMock()
mock_stream.__aiter__ = lambda self: mock_stream
mock_stream.__anext__.side_effect = StopAsyncIteration
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
with patch(
"litellm.proxy.proxy_server.token_counter",
new_callable=AsyncMock,
return_value=fake_response,
):
mock_router.agenerate_content_stream = AsyncMock(return_value=mock_stream)
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
test_call_id = "test-custom-stream-call-id"
response = client.post(
"/v1beta/models/test-model:streamGenerateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello stream"}]}]},
headers={
"Authorization": "Bearer sk-test-key",
"x-litellm-call-id": test_call_id,
},
"/v1beta/models/test-model:countTokens",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
)
assert response.status_code == 200
mock_router.agenerate_content_stream.assert_called_once()
call_args = mock_router.agenerate_content_stream.call_args
called_data = call_args[1]
assert "litellm_logging_obj" in called_data
assert "litellm_call_id" in called_data
assert called_data["litellm_call_id"] == test_call_id
body = response.json()
assert body["totalTokens"] == 7

View file

@ -1853,6 +1853,60 @@ async def test_make_bedrock_api_request_logging_event_type_for_spend_logs():
assert mock_log.call_args.kwargs["event_type"] == GuardrailEventHooks.pre_call
@pytest.mark.asyncio
async def test_make_bedrock_api_request_filters_dynamic_evaluation_overrides():
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT"
)
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = None
mock_bedrock_response = MagicMock()
mock_bedrock_response.status_code = 200
mock_bedrock_response.json.return_value = {"action": "NONE", "assessments": []}
prepared_request = MagicMock()
prepared_request.url = "https://bedrock.test/apply"
prepared_request.body = b"{}"
prepared_request.headers = {}
with (
patch.object(
guardrail.async_handler, "post", new_callable=AsyncMock
) as mock_post,
patch.object(
guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")
),
patch.object(
guardrail, "_prepare_request", return_value=prepared_request
) as mock_prepare_request,
patch.object(
guardrail,
"get_guardrail_dynamic_request_body_params",
return_value={
"content": [{"text": {"text": "benign replacement"}}],
"source": "OUTPUT",
"outputScope": "FULL",
},
),
):
mock_post.return_value = mock_bedrock_response
await guardrail.make_bedrock_api_request(
source="INPUT",
messages=[{"role": "user", "content": "actual prompt"}],
request_data={"model": "gpt-4o"},
)
prepared_data = mock_prepare_request.call_args.kwargs["data"]
assert prepared_data["source"] == "INPUT"
assert "actual prompt" in json.dumps(prepared_data["content"])
assert "benign replacement" not in json.dumps(prepared_data["content"])
assert prepared_data["outputScope"] == "FULL"
@pytest.mark.asyncio
async def test_during_call_hook_invokes_bedrock_async_moderation_hook():
"""

View file

@ -505,6 +505,38 @@ def test_get_logging_caching_headers_pillar_metadata():
)
def test_get_logging_caching_headers_ignores_untrusted_pillar_headers():
request_data = {
"metadata": {
"pillar_response_headers": {
"set-cookie": "session=evil",
"x-pillar-flagged": "true",
},
"pillar_flagged": True,
}
}
headers = get_logging_caching_headers(request_data)
assert "set-cookie" not in headers
assert "x-pillar-flagged" not in headers
def test_get_logging_caching_headers_filters_non_pillar_headers():
request_data = {
"metadata": {
"pillar_flagged": True,
}
}
build_pillar_response_headers(request_data["metadata"])
request_data["metadata"]["pillar_response_headers"]["set-cookie"] = "session=evil"
headers = get_logging_caching_headers(request_data)
assert headers["x-pillar-flagged"] == "true"
assert "set-cookie" not in headers
def test_get_logging_caching_headers_truncates_large_evidence():
long_text = "悪" * 6000 # multi-byte unicode to test URL encoding and truncation
request_data = {

View file

@ -0,0 +1,294 @@
"""
Audit-log emission for the team-callback admin endpoints.
The endpoints in ``team_callback_endpoints.py`` mutate a team's logging
callbacks (``add_team_callbacks``) or zero them out entirely
(``disable_team_logging``). Both are admin-only mutations, and the
disable variant is itself a logging-control action, so when the operator
has Enterprise audit logging enabled (``litellm.store_audit_logs = True``)
each call must emit a row that captures who did it and what the metadata
looked like before/after.
"""
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import Request
import litellm
from litellm.proxy._types import (
AddTeamCallback,
LitellmTableNames,
UserAPIKeyAuth,
)
from litellm.proxy.management_endpoints.team_callback_endpoints import (
add_team_callbacks,
disable_team_logging,
)
def _admin_auth() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="hashed",
user_id="admin-user",
user_role="proxy_admin",
)
def _existing_team_row(metadata: dict) -> MagicMock:
row = MagicMock()
row.team_id = "team-1"
row.metadata = metadata
return row
def _patch_prisma(existing_metadata: dict):
"""Build a context-manager that patches the proxy's ``prisma_client``
to return ``existing_metadata`` from ``get_data`` and a stub team row
from ``litellm_teamtable.update``."""
mock_prisma = MagicMock()
mock_prisma.get_data = AsyncMock(return_value=_existing_team_row(existing_metadata))
updated_row = MagicMock()
updated_row.team_id = "team-1"
mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=updated_row)
return mock_prisma
@pytest.mark.asyncio
async def test_disable_team_logging_emits_audit_log_when_enabled(monkeypatch):
monkeypatch.setattr(litellm, "store_audit_logs", True)
mock_prisma = _patch_prisma(
{
"callback_settings": {
"success_callback": ["langfuse"],
"failure_callback": [],
}
}
)
audit_calls = []
async def capture(request_data):
audit_calls.append(request_data)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new=capture,
),
):
await disable_team_logging(
http_request=MagicMock(spec=Request),
team_id="team-1",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
# asyncio.create_task fires the coroutine eagerly; await one tick to let
# the audit-log emit run before the test exits.
import asyncio
for _ in range(3):
await asyncio.sleep(0)
assert len(audit_calls) == 1
log = audit_calls[0]
assert log.table_name == LitellmTableNames.TEAM_TABLE_NAME
assert log.object_id == "team-1"
assert log.action == "updated"
assert log.changed_by == "admin-user"
before = json.loads(log.before_value)
after = json.loads(log.updated_values)
# Before: the team's pre-existing success_callback survives in the snapshot.
assert before["metadata"]["callback_settings"]["success_callback"] == ["langfuse"]
# After: callbacks zeroed out by the endpoint.
assert after["metadata"]["callback_settings"]["success_callback"] == []
assert after["metadata"]["callback_settings"]["failure_callback"] == []
@pytest.mark.asyncio
async def test_disable_team_logging_no_audit_when_disabled(monkeypatch):
monkeypatch.setattr(litellm, "store_audit_logs", False)
mock_prisma = _patch_prisma(
{
"callback_settings": {
"success_callback": ["langfuse"],
"failure_callback": [],
}
}
)
audit_calls = []
async def capture(request_data):
audit_calls.append(request_data)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new=capture,
),
):
await disable_team_logging(
http_request=MagicMock(spec=Request),
team_id="team-1",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
assert audit_calls == []
@pytest.mark.asyncio
async def test_add_team_callbacks_emits_audit_log_when_enabled(monkeypatch):
monkeypatch.setattr(litellm, "store_audit_logs", True)
mock_prisma = _patch_prisma({"logging": []})
audit_calls = []
async def capture(request_data):
audit_calls.append(request_data)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new=capture,
),
):
await add_team_callbacks(
data=AddTeamCallback(
callback_name="langfuse",
callback_type="success",
callback_vars={
"langfuse_public_key": "pk",
"langfuse_secret_key": "sk",
},
),
http_request=MagicMock(spec=Request),
team_id="team-1",
user_api_key_dict=_admin_auth(),
litellm_changed_by="ops-on-call",
)
import asyncio
for _ in range(3):
await asyncio.sleep(0)
assert len(audit_calls) == 1
log = audit_calls[0]
assert log.table_name == LitellmTableNames.TEAM_TABLE_NAME
assert log.object_id == "team-1"
assert log.action == "updated"
# ``litellm_changed_by`` header takes precedence over the auth user_id.
assert log.changed_by == "ops-on-call"
before = json.loads(log.before_value)
after = json.loads(log.updated_values)
assert before["metadata"]["logging"] == []
assert len(after["metadata"]["logging"]) == 1
assert after["metadata"]["logging"][0]["callback_name"] == "langfuse"
# Callback secrets MUST NOT leak into the audit log payload.
callback_vars = after["metadata"]["logging"][0]["callback_vars"]
assert callback_vars["langfuse_public_key"] != "pk"
assert callback_vars["langfuse_secret_key"] != "sk"
# Key names are preserved so the auditor can see which fields changed.
assert "langfuse_public_key" in callback_vars
assert "langfuse_secret_key" in callback_vars
# And no plaintext secret should appear anywhere in the serialized row.
assert "sk" not in log.updated_values.replace("sk-", "") # crude leak check
assert "pk" not in (log.updated_values.replace("pk-", "").replace("public_key", ""))
@pytest.mark.asyncio
async def test_disable_team_logging_redacts_existing_callback_secrets(monkeypatch):
monkeypatch.setattr(litellm, "store_audit_logs", True)
# Existing team has populated callback_vars containing secrets — redaction
# must apply to the BEFORE snapshot too.
mock_prisma = _patch_prisma(
{
"callback_settings": {
"success_callback": ["langfuse"],
"failure_callback": [],
"callback_vars": {
"langfuse_public_key": "pk-real",
"langfuse_secret_key": "sk-real-secret",
},
}
}
)
audit_calls = []
async def capture(request_data):
audit_calls.append(request_data)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new=capture,
),
):
await disable_team_logging(
http_request=MagicMock(spec=Request),
team_id="team-1",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
import asyncio
for _ in range(3):
await asyncio.sleep(0)
assert len(audit_calls) == 1
log = audit_calls[0]
# The pre-existing secret_key value must NOT appear in the serialized
# before_value or updated_values.
assert "sk-real-secret" not in log.before_value
assert "sk-real-secret" not in log.updated_values
assert "pk-real" not in log.before_value
assert "pk-real" not in log.updated_values
@pytest.mark.asyncio
async def test_add_team_callbacks_no_audit_when_disabled(monkeypatch):
monkeypatch.setattr(litellm, "store_audit_logs", False)
mock_prisma = _patch_prisma({"logging": []})
audit_calls = []
async def capture(request_data):
audit_calls.append(request_data)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new=capture,
),
):
await add_team_callbacks(
data=AddTeamCallback(
callback_name="langfuse",
callback_type="success",
callback_vars={
"langfuse_public_key": "pk",
"langfuse_secret_key": "sk",
},
),
http_request=MagicMock(spec=Request),
team_id="team-1",
user_api_key_dict=_admin_auth(),
litellm_changed_by=None,
)
assert audit_calls == []

View file

@ -968,6 +968,20 @@ def test_add_new_models_to_team():
)
def _make_team_member_add_request(
member_user_id: Optional[str] = "regular-user",
role: str = "user",
team_id: str = "test-team-123",
):
"""Build a TeamMemberAddRequest with one Member entry for tests below."""
from litellm.proxy._types import Member, TeamMemberAddRequest
return TeamMemberAddRequest(
team_id=team_id,
member=Member(role=role, user_id=member_user_id),
)
@pytest.mark.asyncio
async def test_validate_team_member_add_permissions_admin():
"""
@ -977,17 +991,15 @@ async def test_validate_team_member_add_permissions_admin():
_validate_team_member_add_permissions,
)
# Create admin user
admin_user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
# Create mock team
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "test-team-123"
# Should not raise any exception for admin
await _validate_team_member_add_permissions(
user_api_key_dict=admin_user,
complete_team_data=team,
data=_make_team_member_add_request(member_user_id="any-user", role="admin"),
)
@ -1000,20 +1012,17 @@ async def test_validate_team_member_add_permissions_non_admin():
_validate_team_member_add_permissions,
)
# Create non-admin user
regular_user = UserAPIKeyAuth(
user_id="regular-user",
user_role=LitellmUserRoles.INTERNAL_USER,
team_id="different-team",
)
# Create mock team
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "test-team-123"
team.members_with_roles = []
team.organization_id = None
# Mock the helper functions to return False
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
@ -1024,17 +1033,303 @@ async def test_validate_team_member_add_permissions_non_admin():
return_value=False,
),
):
# Should raise HTTPException for non-admin
with pytest.raises(HTTPException) as exc_info:
await _validate_team_member_add_permissions(
user_api_key_dict=regular_user,
complete_team_data=team,
data=_make_team_member_add_request(),
)
assert exc_info.value.status_code == 403
assert "not proxy admin OR team admin" in str(exc_info.value.detail)
# ── VERIA-56 regression tests for _is_available_team self-join enforcement ───
@pytest.mark.asyncio
async def test_available_team_self_join_with_caller_user_id_allowed():
"""A standard user adding themselves to an available team with role=user
is the only legitimate use of the available-team bypass."""
from litellm.proxy.management_endpoints.team_endpoints import (
_validate_team_member_add_permissions,
)
user = UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER,
)
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "public-team"
team.members_with_roles = []
team.organization_id = None
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
return_value=False,
),
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_available_team",
return_value=True,
),
):
await _validate_team_member_add_permissions(
user_api_key_dict=user,
complete_team_data=team,
data=_make_team_member_add_request(member_user_id="alice", role="user"),
)
@pytest.mark.asyncio
async def test_available_team_self_join_blocks_admin_role():
"""Privesc shape from VERIA-56: caller adds themselves with role=admin
via the available-team bypass. Must be rejected."""
from litellm.proxy.management_endpoints.team_endpoints import (
_validate_team_member_add_permissions,
)
user = UserAPIKeyAuth(user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER)
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "public-team"
team.members_with_roles = []
team.organization_id = None
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
return_value=False,
),
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_available_team",
return_value=True,
),
pytest.raises(HTTPException) as exc_info,
):
await _validate_team_member_add_permissions(
user_api_key_dict=user,
complete_team_data=team,
data=_make_team_member_add_request(member_user_id="alice", role="admin"),
)
assert exc_info.value.status_code == 403
assert "admin" in str(exc_info.value.detail).lower()
@pytest.mark.asyncio
async def test_available_team_self_join_blocks_other_user_id():
"""Cross-user-injection shape from VERIA-56: caller adds someone else
via the available-team bypass. Must be rejected."""
from litellm.proxy.management_endpoints.team_endpoints import (
_validate_team_member_add_permissions,
)
user = UserAPIKeyAuth(user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER)
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "public-team"
team.members_with_roles = []
team.organization_id = None
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
return_value=False,
),
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_available_team",
return_value=True,
),
pytest.raises(HTTPException) as exc_info,
):
await _validate_team_member_add_permissions(
user_api_key_dict=user,
complete_team_data=team,
data=_make_team_member_add_request(
member_user_id="bob-victim", role="user"
),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_available_team_self_join_blocks_when_caller_has_no_user_id():
"""If the auth context has no user_id we cannot prove self-join, so the
bypass must fail closed."""
from litellm.proxy.management_endpoints.team_endpoints import (
_validate_team_member_add_permissions,
)
user = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER) # no user_id
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "public-team"
team.members_with_roles = []
team.organization_id = None
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
return_value=False,
),
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_available_team",
return_value=True,
),
pytest.raises(HTTPException) as exc_info,
):
await _validate_team_member_add_permissions(
user_api_key_dict=user,
complete_team_data=team,
data=_make_team_member_add_request(member_user_id="alice", role="user"),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_available_team_self_join_blocks_email_only_member():
"""An email-only member entry can't be safely self-join-validated; the
caller must use their own user_id explicitly."""
from litellm.proxy._types import Member, TeamMemberAddRequest
from litellm.proxy.management_endpoints.team_endpoints import (
_validate_team_member_add_permissions,
)
user = UserAPIKeyAuth(user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER)
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "public-team"
team.members_with_roles = []
team.organization_id = None
data = TeamMemberAddRequest(
team_id="public-team",
member=Member(role="user", user_email="alice@example.com"),
)
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
return_value=False,
),
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_available_team",
return_value=True,
),
pytest.raises(HTTPException) as exc_info,
):
await _validate_team_member_add_permissions(
user_api_key_dict=user,
complete_team_data=team,
data=data,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_available_team_self_join_blocks_admin_role_in_member_list():
"""Bulk shape: list of members where one has role=admin must be rejected
even if the caller's own entry is correct."""
from litellm.proxy._types import Member, TeamMemberAddRequest
from litellm.proxy.management_endpoints.team_endpoints import (
_validate_team_member_add_permissions,
)
user = UserAPIKeyAuth(user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER)
team = MagicMock(spec=LiteLLM_TeamTable)
team.team_id = "public-team"
team.members_with_roles = []
team.organization_id = None
data = TeamMemberAddRequest(
team_id="public-team",
member=[
Member(role="user", user_id="alice"),
Member(role="admin", user_id="alice"),
],
)
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
return_value=False,
),
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_available_team",
return_value=True,
),
pytest.raises(HTTPException) as exc_info,
):
await _validate_team_member_add_permissions(
user_api_key_dict=user,
complete_team_data=team,
data=data,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_update_team_member_permissions_blocks_non_admin_via_available_team(
mock_db_client,
):
"""A non-admin caller invoking /team/permissions_update on an available
team must be rejected. The previous code path delegated to
``_is_available_team`` and accepted the write; this PR removes that
bypass entirely so the result is 403 even with the bypass mocked True."""
test_team_id = "public-team"
update_payload = {
"team_id": test_team_id,
"team_member_permissions": ["/key/generate"],
}
existing_row = MagicMock(spec=LiteLLM_TeamTable)
existing_row.model_dump.return_value = {
"team_id": test_team_id,
"team_alias": "Public Team",
"team_member_permissions": [],
"spend": 0.0,
"models": [],
}
existing_row.team_id = test_team_id
existing_row.members_with_roles = []
existing_row.organization_id = None
non_admin_auth = UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER,
)
with (
patch(
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
new_callable=AsyncMock,
return_value=existing_row,
),
patch(
"litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin",
return_value=False,
),
patch(
# Even with the available-team bypass mocked True, the endpoint
# must NOT consult it any more — the gate should reject the
# non-admin caller outright.
"litellm.proxy.management_endpoints.team_endpoints._is_available_team",
return_value=True,
),
):
app.dependency_overrides[user_api_key_auth] = lambda: non_admin_auth
try:
response = client.post("/team/permissions_update", json=update_payload)
finally:
app.dependency_overrides = {}
assert response.status_code == 403
body = response.json()
assert "permissions_update" in str(body) or "not proxy admin" in str(body)
@pytest.mark.asyncio
async def test_process_team_members_single_member():
"""

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