mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
commit
722a1a9f8f
128 changed files with 41497 additions and 1774 deletions
75
.github/workflows/check-lazy-openapi-snapshot.yml
vendored
Normal file
75
.github/workflows/check-lazy-openapi-snapshot.yml
vendored
Normal 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
2
.npmrc
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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[];
|
||||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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())))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
432
litellm/proxy/_lazy_features.py
Normal file
432
litellm/proxy/_lazy_features.py
Normal 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
|
||||
}
|
||||
31651
litellm/proxy/_lazy_openapi_snapshot.json
Normal file
31651
litellm/proxy/_lazy_openapi_snapshot.json
Normal file
File diff suppressed because it is too large
Load diff
70
litellm/proxy/_lazy_openapi_snapshot.py
Normal file
70
litellm/proxy/_lazy_openapi_snapshot.py
Normal 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")
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
52
litellm/proxy/common_utils/static_asset_utils.py
Normal file
52
litellm/proxy/common_utils/static_asset_utils.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)}."
|
||||
)
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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[]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"] = (
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
@ -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
Loading…
Add table
Reference in a new issue