Merge branch 'BerriAI:litellm_internal_staging' into pgvector-changes

This commit is contained in:
Sujal Bhavsar 2026-05-26 09:46:00 +05:30 • committed by GitHub
commit 1eb0e4b7b4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
122 changed files with 4134 additions and 401 deletions

View file

@ -0,0 +1,47 @@
name: Create Daily oss-agent-shin Branch
on:
schedule:
- cron: "0 0 * * *" # Runs every day at midnight UTC
workflow_dispatch: # Allow manual trigger
jobs:
create-oss-agent-shin-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
persist-credentials: false
- name: Create daily oss-agent-shin branch
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
# Configure Git user
git config user.name "github-actions[bot]"
git config user.email "github-actions[bot]@users.noreply.github.com"
# Generate branch name with MM_DD_YYYY format
BRANCH_NAME="litellm_oss_agent_shin_$(date +'%m_%d_%Y')"
echo "Creating branch: $BRANCH_NAME"
# Fetch all branches
git fetch --all
# Check if the branch already exists
if git show-ref --verify --quiet refs/remotes/origin/$BRANCH_NAME; then
echo "Branch $BRANCH_NAME already exists. Skipping creation."
else
echo "Creating new branch: $BRANCH_NAME"
# Create the new branch from main
git checkout -b $BRANCH_NAME origin/main
# Push the new branch
git push origin $BRANCH_NAME
echo "Successfully created and pushed branch: $BRANCH_NAME"
fi

View file

@ -7,6 +7,7 @@ on:
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
workflow_dispatch:
permissions:
contents: read
@ -42,3 +43,16 @@ jobs:
workers: 2
reruns: 2
artifact-name: proxy-endpoints
# Behavior-pinning tests for litellm/proxy/proxy_server.py. Owns its
# own job (not a path on the proxy-endpoints job above) so its budget
# is independent and its coverage artifact is uploaded separately.
# See: https://www.notion.so/36c43b8acdab81ee845fd5365128a2fc
proxy-server:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: tests/test_litellm/proxy/proxy_server
workers: 4
reruns: 2
timeout-minutes: 60
artifact-name: proxy-server

View file

@ -55,22 +55,10 @@ COPY . .
# Set non-root flag for build time consistency
ENV LITELLM_NON_ROOT=true
# Stage the pre-built Admin UI from the checked-in Next.js static export.
# _experimental/out/ is regenerated as part of the release runbook.
# Restructure extensionless routes (foo.html -> foo/index.html) to match the layout
# proxy_server.py expects, and drop a readiness marker.
RUN mkdir -p /var/lib/litellm/ui /var/lib/litellm/assets && \
cp -r /app/litellm/proxy/_experimental/out/. /var/lib/litellm/ui/ && \
cp /app/litellm/proxy/logo.jpg /var/lib/litellm/assets/logo.jpg && \
( cd /var/lib/litellm/ui && \
for html_file in *.html; do \
if [ "$html_file" != "index.html" ] && [ -f "$html_file" ]; then \
folder_name="${html_file%.html}" && \
mkdir -p "$folder_name" && \
mv "$html_file" "$folder_name/index.html"; \
fi; \
done && \
touch .litellm_ui_ready )
touch /var/lib/litellm/ui/.litellm_ui_ready
RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
if [ "$PROXY_EXTRAS_SOURCE" = "published" ]; then \

View file

@ -24,6 +24,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
from litellm.litellm_core_utils.llm_cost_calc.utils import (
CostCalculatorUtils,
_generic_cost_per_character,
_get_regional_uplift_multiplier,
_get_service_tier_cost_key,
_parse_prompt_tokens_details,
calculate_cost_component,
@ -312,6 +313,10 @@ def cost_per_token( # noqa: PLR0915
audio_transcription_file_duration: float = 0.0, # for audio transcription calls - the file time in seconds
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: Optional[
str
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
response: Optional[Any] = None,
### REQUEST MODEL ###
request_model: Optional[str] = None, # original request model for router detection
@ -493,6 +498,7 @@ def cost_per_token( # noqa: PLR0915
usage=usage_block,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
data_residency=data_residency,
)
return prompt_cost, completion_cost
@ -521,7 +527,10 @@ def cost_per_token( # noqa: PLR0915
or call_type == CallTypes.retrieve_batch
):
return batch_cost_calculator(
usage=usage_block, model=model, custom_llm_provider=custom_llm_provider
usage=usage_block,
model=model,
custom_llm_provider=custom_llm_provider,
data_residency=data_residency,
)
elif call_type == "atranscription" or call_type == "transcription":
if _transcription_usage_has_token_details(usage_block):
@ -529,6 +538,7 @@ def cost_per_token( # noqa: PLR0915
model=model_without_prefix,
usage=usage_block,
service_tier=service_tier,
data_residency=data_residency,
)
return openai_cost_per_second(
@ -579,7 +589,10 @@ def cost_per_token( # noqa: PLR0915
)
elif custom_llm_provider == "openai":
return openai_cost_per_token(
model=model, usage=usage_block, service_tier=service_tier
model=model,
usage=usage_block,
service_tier=service_tier,
data_residency=data_residency,
)
elif custom_llm_provider == "databricks":
return databricks_cost_per_token(model=model, usage=usage_block)
@ -631,6 +644,7 @@ def cost_per_token( # noqa: PLR0915
usage=usage_block,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
data_residency=data_residency,
)
if (
@ -1117,6 +1131,10 @@ def completion_cost( # noqa: PLR0915
litellm_logging_obj: Optional[LitellmLoggingObject] = None,
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: Optional[
str
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
) -> float:
"""
Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm.
@ -1516,6 +1534,7 @@ def completion_cost( # noqa: PLR0915
combined_usage_object=cost_per_token_usage_object,
custom_llm_provider=custom_llm_provider,
litellm_model_name=model,
data_residency=data_residency,
)
elif call_type == _MCP_CALL_TYPE:
from litellm.proxy._experimental.mcp_server.cost_calculator import (
@ -1600,6 +1619,7 @@ def completion_cost( # noqa: PLR0915
audio_transcription_file_duration=audio_transcription_file_duration,
rerank_billed_units=rerank_billed_units,
service_tier=service_tier,
data_residency=data_residency,
response=completion_response,
request_model=request_model_for_cost,
)
@ -1811,6 +1831,10 @@ def response_cost_calculator(
litellm_logging_obj: Optional[LitellmLoggingObject] = None,
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: Optional[
str
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
) -> float:
"""
Returns
@ -1844,6 +1868,7 @@ def response_cost_calculator(
router_model_id=router_model_id,
litellm_logging_obj=litellm_logging_obj,
service_tier=service_tier,
data_residency=data_residency,
)
return response_cost
except Exception as e:
@ -2202,6 +2227,7 @@ def batch_cost_calculator(
model: str,
custom_llm_provider: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
data_residency: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculate the cost of a batch job.
@ -2286,6 +2312,11 @@ def batch_cost_calculator(
usage.completion_tokens * (output_cost_per_token) / 2
) # batch cost is usually half of the regular token cost
uplift = _get_regional_uplift_multiplier(model_info, data_residency)
if uplift != 1.0:
total_prompt_cost *= uplift
total_completion_cost *= uplift
return total_prompt_cost, total_completion_cost
@ -2431,6 +2462,7 @@ def handle_realtime_stream_cost_calculation(
combined_usage_object: Usage,
custom_llm_provider: str,
litellm_model_name: str,
data_residency: Optional[str] = None,
) -> float:
"""
Handles the cost calculation for realtime stream responses.
@ -2461,6 +2493,7 @@ def handle_realtime_stream_cost_calculation(
model=model_name,
usage=combined_usage_object,
custom_llm_provider=custom_llm_provider,
data_residency=data_residency,
)
except Exception:
continue

View file

@ -1,5 +1,7 @@
from typing import Optional
from litellm.llms.openai.data_residency import infer_openai_data_residency
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
_OPTIONAL_KWARGS_KEYS = frozenset(
@ -103,6 +105,10 @@ def get_litellm_params(
if litellm_trace_id is None:
litellm_trace_id = _meta.get("trace_id") or _meta.get("session_id")
data_residency: Optional[str] = infer_openai_data_residency(
custom_llm_provider, api_base
)
# Build base dict with explicit parameters (always included)
litellm_params = {
"acompletion": acompletion,
@ -112,6 +118,7 @@ def get_litellm_params(
"verbose": verbose,
"custom_llm_provider": custom_llm_provider,
"api_base": api_base,
"data_residency": data_residency,
"litellm_call_id": litellm_call_id,
"model_alias_map": model_alias_map,
"completion_call_id": completion_call_id,

View file

@ -1546,6 +1546,11 @@ class Logging(LiteLLMLoggingBaseClass):
if self.optional_params
else None
),
"data_residency": (
self.litellm_params.get("data_residency")
if hasattr(self, "litellm_params") and self.litellm_params
else None
),
}
except Exception as e: # error creating kwargs for cost calculation
debug_info = StandardLoggingModelCostFailureDebugInformation(

View file

@ -9,6 +9,7 @@ from litellm.types.utils import (
CacheCreationTokenDetails,
CallTypes,
CompletionTokensDetailsWrapper,
DataResidency,
ImageResponse,
ModelInfo,
PassthroughCallTypes,
@ -617,11 +618,46 @@ def _calculate_input_cost(
return prompt_cost
def _get_regional_uplift_multiplier(
model_info: ModelInfo, data_residency: Optional[str]
) -> float:
"""
Resolve the per-model regional-processing uplift multiplier for a given
data-residency region.
OpenAI applies a flat percentage uplift (e.g. +10%) on all token costs for
requests served from a regionalized hostname (eu./us.api.openai.com). The
multiplier is stored on the model entry as
``regional_processing_uplift_multiplier_<region>`` (e.g. 1.10).
Returns 1.0 (no uplift) when ``data_residency`` is ``None`` or when the
model has no multiplier configured for the given region.
"""
if data_residency is None:
return 1.0
residency = data_residency.lower()
if residency not in {r.value for r in DataResidency}:
return 1.0
multiplier = model_info.get(f"regional_processing_uplift_multiplier_{residency}")
if multiplier is None:
return 1.0
try:
return float(cast(float, multiplier))
except (TypeError, ValueError):
verbose_logger.exception(
"Invalid regional_processing_uplift_multiplier_%s for model; "
"defaulting to 1.0",
residency,
)
return 1.0
def generic_cost_per_token( # noqa: PLR0915
model: str,
usage: Usage,
custom_llm_provider: str,
service_tier: Optional[str] = None,
data_residency: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -631,6 +667,8 @@ def generic_cost_per_token( # noqa: PLR0915
Input:
- model: str, the model name without provider prefix
- usage: LiteLLM Usage block, containing anthropic caching information
- data_residency: optional OpenAI data-residency region (e.g. "eu", "us"),
used to apply the per-model regional-processing uplift multiplier.
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -781,6 +819,14 @@ def generic_cost_per_token( # noqa: PLR0915
)
completion_cost += float(image_tokens) * _output_cost_per_image_token
## REGIONAL DATA-RESIDENCY UPLIFT
# Applied as a flat multiplier across all token costs for the request
# when the upstream is a regionalized OpenAI host (eu./us.api.openai.com).
uplift = _get_regional_uplift_multiplier(model_info, data_residency)
if uplift != 1.0:
prompt_cost *= uplift
completion_cost *= uplift
return prompt_cost, completion_cost

View file

@ -146,6 +146,37 @@ class SensitiveDataMasker:
return masked_data
_default_masker = SensitiveDataMasker()
def mask_sensitive_keys(
data: Dict[str, Any], sensitive_fields: Set[str]
) -> Dict[str, Any]:
"""Return a new dict with values masked for keys listed in ``sensitive_fields``.
Unlike :meth:`SensitiveDataMasker.mask_dict`, this does exact key-name
matching (not segment matching), so callers explicitly enumerate which
fields to mask. Non-string and None values are passed through unchanged.
Values shorter than ``visible_prefix + visible_suffix`` (8 by default)
fall outside :meth:`SensitiveDataMasker._mask_value`'s partial-reveal
range and are replaced with a fixed-length all-mask string, so a short
credential is never returned verbatim.
"""
masked: Dict[str, Any] = {}
mask_char = _default_masker.mask_char
min_visible = _default_masker.visible_prefix + _default_masker.visible_suffix
for key, value in data.items():
if value is not None and key in sensitive_fields and isinstance(value, str):
if len(value) < min_visible:
masked[key] = mask_char * len(value) if value else value
else:
masked[key] = _default_masker._mask_value(value)
else:
masked[key] = value
return masked
# Usage example:
"""
masker = SensitiveDataMasker()

View file

@ -1,3 +1,5 @@
import asyncio
import hashlib
import json
import os
from typing import Any, Callable, Dict, Literal, NamedTuple, Optional, Union, cast
@ -449,6 +451,25 @@ class BaseAzureLLM(BaseOpenAILLM):
] = None
client_initialization_params: dict = locals()
client_initialization_params["is_async"] = _is_async
_lp = litellm_params or {}
_ad_provider = _lp.get("azure_ad_token_provider")
_ad_token = _lp.get("azure_ad_token")
_client_secret = _lp.get("client_secret")
_azure_password = _lp.get("azure_password")
client_initialization_params["azure_ad_token"] = (
hashlib.sha256(_ad_token.encode()).hexdigest()
if isinstance(_ad_token, str)
else None
)
client_initialization_params["azure_ad_token_provider"] = (
f"provider_id={id(_ad_provider) if callable(_ad_provider) else None}"
f"|tenant_id={_lp.get('tenant_id')}"
f"|client_id={_lp.get('client_id')}"
f"|client_secret={hashlib.sha256(_client_secret.encode()).hexdigest() if isinstance(_client_secret, str) else None}"
f"|azure_username={_lp.get('azure_username')}"
f"|azure_password={hashlib.sha256(_azure_password.encode()).hexdigest() if isinstance(_azure_password, str) else None}"
f"|azure_scope={_lp.get('azure_scope')}"
)
if client is None:
cached_client = self.get_cached_openai_client(
client_initialization_params=client_initialization_params,
@ -474,8 +495,29 @@ class BaseAzureLLM(BaseOpenAILLM):
if self._is_azure_v1_api_version(api_version):
# Extract only params that OpenAI client accepts
# Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview"
v1_params = {
"api_key": azure_client_params.get("api_key"),
# The OpenAI client accepts a callable for `api_key` and re-invokes it
# on every request (via `_refresh_api_key`), so passing
# `azure_ad_token_provider` directly preserves Azure AD token refresh
# behavior that the regular AzureOpenAI client provides.
v1_api_key: Optional[Union[str, Callable[[], Any]]] = (
azure_client_params.get("api_key")
or azure_client_params.get("azure_ad_token_provider")
or azure_client_params.get("azure_ad_token")
)
if _is_async is True and callable(v1_api_key):
# AsyncOpenAI expects an async provider; wrap the sync provider
# returned by azure-identity. Offload to a thread so a token
# refresh (blocking HTTP call to AAD on cache miss) does not
# stall the event loop.
_sync_provider = v1_api_key
async def _async_v1_api_key() -> str:
return await asyncio.to_thread(_sync_provider)
v1_api_key = _async_v1_api_key
v1_params: Dict[str, Any] = {
"api_key": v1_api_key,
"base_url": f"{api_base}/openai/v1/",
}
if "timeout" in azure_client_params:

View file

@ -177,8 +177,14 @@ def extract_model_id_from_unified_id(
if decoded_id:
unified_id = decoded_id
# Extract model ID
match = re.search(r"model_id,([^;]+)", unified_id)
# Extract model ID. Anchor to a field boundary (start of string or
# after `;`) so this regex doesn't substring-match the `model_id,`
# inside file_id encodings' `llm_output_file_model_id,<deployment_uuid>`
# field — that would feed the deployment UUID as a model candidate
# into the team-access check and 403 every team-BYOK file attach
# with `Tried to access <uuid>` (LIT-3244 patch/1.86.0 second-order
# finding).
match = re.search(r"(?:^|;)model_id,([^;]+)", unified_id)
if match:
return match.group(1).strip()

View file

@ -157,8 +157,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
def _get_agent_runtime_arn(self, model: str) -> str:
"""
Extract ARN from model string
model = "agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
returns: "arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
model = "agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
returns: "arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
"""
parts = model.split("/", 1)
if len(parts) != 2 or parts[0] != "agentcore":
@ -170,7 +170,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
def _extract_region_from_arn(self, arn: str) -> str:
"""
Extract region from ARN
arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC
arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp
returns: us-west-2
"""
parts = arn.split(":")

View file

@ -19,7 +19,10 @@ def cost_router(call_type: CallTypes) -> Literal["cost_per_token", "cost_per_sec
def cost_per_token(
model: str, usage: Usage, service_tier: Optional[str] = None
model: str,
usage: Usage,
service_tier: Optional[str] = None,
data_residency: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -27,6 +30,9 @@ def cost_per_token(
Input:
- model: str, the model name without provider prefix
- usage: LiteLLM Usage block, containing anthropic caching information
- data_residency: optional OpenAI data-residency region (e.g. "eu", "us"),
inferred from api_base. Applies the model's regional-processing
uplift multiplier when set.
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -37,6 +43,7 @@ def cost_per_token(
usage=usage,
custom_llm_provider="openai",
service_tier=service_tier,
data_residency=data_residency,
)
# ### Non-cached text tokens
# non_cached_text_tokens = usage.prompt_tokens

View file

@ -0,0 +1,41 @@
"""
Helpers for resolving OpenAI data-residency (regional processing) from an
api_base URL.
OpenAI enforces hostname-per-region for projects with geography restrictions
enabled and rejects requests sent to the wrong host, so the api_base hostname
is the authoritative signal of which region a request was processed in.
"""
from typing import Dict, Optional
from urllib.parse import urlparse
# Mapping of OpenAI regional hostnames to the corresponding data-residency
# value used by the cost calculator. See
# https://developers.openai.com/api/docs/pricing for the regional-processing
# uplift these hostnames trigger.
_OPENAI_REGIONAL_HOSTS: Dict[str, str] = {
"eu.api.openai.com": "eu",
"us.api.openai.com": "us",
}
def infer_openai_data_residency(
custom_llm_provider: Optional[str], api_base: Optional[str]
) -> Optional[str]:
"""
Derive the OpenAI data-residency region from an api_base URL.
Returns ``"eu"`` for the EU regional host, ``"us"`` for the US regional
host, and ``None`` for the default global host, any non-OpenAI provider,
or any non-OpenAI URL.
"""
if custom_llm_provider != "openai" or not api_base:
return None
try:
host = urlparse(api_base).hostname
except (TypeError, ValueError):
return None
if not host:
return None
return _OPENAI_REGIONAL_HOSTS.get(host.lower())

View file

@ -1011,6 +1011,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -1041,6 +1042,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -1071,6 +1073,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -1100,6 +1103,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -1129,6 +1133,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -1328,6 +1333,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"global.anthropic.claude-sonnet-4-6": {
@ -1358,6 +1364,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"us.anthropic.claude-sonnet-4-6": {
@ -1388,6 +1395,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"eu.anthropic.claude-sonnet-4-6": {
@ -1417,6 +1425,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"au.anthropic.claude-sonnet-4-6": {
@ -1446,6 +1455,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"jp.anthropic.claude-sonnet-4-6": {
@ -1475,6 +1485,7 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"anthropic.claude-sonnet-4-20250514-v1:0": {
@ -1996,6 +2007,7 @@
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 159,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -2093,6 +2105,7 @@
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"azure/computer-use-preview": {
@ -9654,6 +9667,7 @@
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"claude-sonnet-4-5-20250929-v1:0": {
@ -9851,6 +9865,7 @@
"us": 1.1,
"fast": 6.0
},
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -9886,7 +9901,8 @@
"fast": 6.0
},
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": true,
"supports_output_config": true
},
"claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
@ -9921,7 +9937,8 @@
"us": 1.1,
"fast": 6.0
},
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": true,
"supports_output_config": true
},
"claude-opus-4-7-20260416": {
"cache_creation_input_token_cost": 6.25e-06,
@ -9956,7 +9973,8 @@
"us": 1.1,
"fast": 6.0
},
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": true,
"supports_output_config": true
},
"claude-sonnet-4-20250514": {
"deprecation_date": "2026-05-14",
@ -14958,7 +14976,7 @@
"mode": "chat",
"output_cost_per_reasoning_token": 1.5e-06,
"output_cost_per_token": 1.5e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"source": "https://ai.google.dev/gemini-api/docs/models",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/completions",
@ -19014,6 +19032,8 @@
"output_cost_per_token": 8e-06,
"output_cost_per_token_batches": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -19087,6 +19107,8 @@
"output_cost_per_token": 1.6e-06,
"output_cost_per_token_batches": 8e-07,
"output_cost_per_token_priority": 2.8e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -19160,6 +19182,8 @@
"output_cost_per_token": 4e-07,
"output_cost_per_token_batches": 2e-07,
"output_cost_per_token_priority": 8e-07,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -19231,6 +19255,8 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"output_cost_per_token_priority": 1.7e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -19272,6 +19298,8 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -19293,6 +19321,8 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -19581,6 +19611,8 @@
"output_cost_per_token": 6e-07,
"output_cost_per_token_batches": 3e-07,
"output_cost_per_token_priority": 1e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -20284,6 +20316,8 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -21206,6 +21240,8 @@
"mode": "responses",
"output_cost_per_token": 0.00012,
"output_cost_per_token_batches": 6e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/batch",
"/v1/responses"
@ -21612,6 +21648,8 @@
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -21693,6 +21731,8 @@
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,
@ -28243,10 +28283,10 @@
"supports_tool_choice": true
},
"openrouter/xiaomi/mimo-v2-flash": {
"input_cost_per_token": 9e-08,
"output_cost_per_token": 2.9e-07,
"input_cost_per_token": 1e-07,
"output_cost_per_token": 3e-07,
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 0.0,
"cache_read_input_token_cost": 1e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 16384,
@ -28256,7 +28296,43 @@
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": false,
"supports_prompt_caching": false
"supports_prompt_caching": true
},
"openrouter/xiaomi/mimo-v2.5-pro": {
"input_cost_per_token": 1e-06,
"output_cost_per_token": 3e-06,
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": false,
"supports_response_schema": true,
"supports_prompt_caching": true
},
"openrouter/xiaomi/mimo-v2.5": {
"input_cost_per_token": 4e-07,
"output_cost_per_token": 2e-06,
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 8e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true,
"supports_response_schema": true,
"supports_prompt_caching": true
},
"openrouter/z-ai/glm-4.7": {
"input_cost_per_token": 4e-07,
@ -28987,14 +29063,16 @@
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
"supports_function_calling": true,
"supports_output_config": true
},
"perplexity/anthropic/claude-opus-4-7": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
"supports_function_calling": true,
"supports_output_config": true
},
"perplexity/anthropic/claude-opus-4-5": {
"litellm_provider": "perplexity",
@ -33405,6 +33483,7 @@
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -33433,6 +33512,7 @@
"supports_tool_choice": true,
"supports_vision": true,
"tool_use_system_prompt_tokens": 346,
"supports_output_config": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
@ -33546,6 +33626,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"vertex_ai/claude-sonnet-4-5@20250929": {
@ -40658,6 +40739,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_output_config": true,
"supports_minimal_reasoning_effort": true
},
"duckduckgo/search": {

View file

@ -118,15 +118,19 @@ class MCPRequestHandler:
return b"{}"
request.body = mock_body # type: ignore
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
request_route = get_request_route(request)
# 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/"):
if request_route.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
not litellm_api_key
and MCPRequestHandler._target_servers_delegate_auth_to_upstream( # noqa: E501
path=request.url.path, mcp_servers=mcp_servers
path=request_route, mcp_servers=mcp_servers
)
):
# Operator opted this oauth2 server into upstream-delegated auth
@ -174,7 +178,7 @@ class MCPRequestHandler:
"401",
"403",
) and MCPRequestHandler._target_servers_use_oauth2(
path=request.url.path, mcp_servers=mcp_servers
path=request_route, mcp_servers=mcp_servers
):
verbose_logger.debug(
"MCP OAuth2: target server is OAuth2-mode, treating "

View file

@ -1765,19 +1765,39 @@ async def _cache_team_object(
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: Optional[ProxyLogging],
):
key = "team_id:{}".format(team_id)
## CACHE REFRESH TIME!
team_table.last_refreshed_at = time.time()
# team_id is the table primary key — guaranteed unique, safe to write.
await _cache_management_object(
key=key,
key="team_id:{}".format(team_id),
value=team_table,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
model_type=LiteLLM_TeamTableCachedObj,
)
# Invalidate the alias-keyed cache so the JWT auth path with
# `team_alias_jwt_field` (which reads via `get_team_object_by_alias`)
# doesn't keep serving the pre-mutation team after every team-write
# endpoint (team_model_add, team_model_delete, update_team, etc.).
#
# Why DELETE and not WRITE: `team_alias` has no UNIQUE constraint in
# schema.prisma. Writing this cache from the generic refresh path
# would let a team admin who renamed their team to collide with
# another team's alias silently overwrite the cached team for
# JWT-by-alias auth (veria-ai review on #28739). Deleting forces the
# next reader through `get_team_object_by_alias`, which DOES enforce
# uniqueness (len(teams) > 1 raises HTTPException) before populating
# the cache from a verified single row.
if team_table.team_alias:
alias_key = "team_alias:{}".format(team_table.team_alias)
user_api_key_cache.delete_cache(key=alias_key)
if proxy_logging_obj is not None:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(
key=alias_key
)
async def _cache_key_object(
hashed_token: str,

View file

@ -498,9 +498,18 @@ def route_in_additonal_public_routes(current_route: str):
def get_request_route(request: Request) -> str:
"""
Helper to get the route from the request
Resolve the request route from the ASGI scope, with ``root_path`` stripped.
remove base url from path if set e.g. `/genai/chat/completions` -> `/chat/completions
Prefer this over ``request.url.path`` for any auth, ACL, routing, or
audit-log decision: Starlette reconstructs ``url.path`` by interpolating
the Host header into a URL string and re-parsing with ``urlsplit``, so a
malformed Host (e.g. ``localhost/?x=1``) collapses ``url.path`` to ``"/"``
while FastAPI continues to dispatch on ``scope["path"]``. ``scope["path"]``
is uvicorn's parse of the HTTP request line and matches the actual
handler, so it's the authoritative route.
Also normalizes sub-path deployments by stripping ``scope["root_path"]``
e.g. ``/genai/chat/completions`` -> ``/chat/completions``.
"""
try:
scope = request.scope

View file

@ -14,6 +14,9 @@ from litellm.utils import get_valid_models
_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields)
_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields)
def _check_wildcard_routing(model: str) -> bool:
"""
Returns True if a model is a provider wildcard.

View file

@ -62,7 +62,11 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_KEY_SUFFIXES = ("/regenerate", "/reset_spend")
class RouteChecks:
@staticmethod
def should_call_route(route: str, valid_token: UserAPIKeyAuth):
def should_call_route(
route: str,
valid_token: UserAPIKeyAuth,
request: Optional[Request] = None,
):
"""
Check if management route is disabled and raise exception
"""
@ -77,13 +81,15 @@ class RouteChecks:
# Check if Virtual Key is allowed to call the route - Applies to all Roles
RouteChecks.is_virtual_key_allowed_to_call_route(
route=route, valid_token=valid_token
route=route, valid_token=valid_token, request=request
)
return True
@staticmethod
def is_virtual_key_allowed_to_call_route(
route: str, valid_token: UserAPIKeyAuth
route: str,
valid_token: UserAPIKeyAuth,
request: Optional[Request] = None,
) -> bool:
"""
Raises Exception if Virtual Key is not allowed to call the route
@ -130,6 +136,21 @@ class RouteChecks:
):
return True
# Method-aware carve-out: allow GET on the two
# read-only MCP-server discovery endpoints
# (`/v1/mcp/server` and `/v1/mcp/server/{server_id}`)
# so virtual keys with allowed_routes=["llm_api_routes"]
# can list/inspect MCP servers. The GET handlers in
# mcp_management_endpoints.py sanitize the response
# for restricted virtual keys (stripping url,
# headers, env, credentials). POST/PUT/DELETE on
# these paths are admin-only management writes and
# are intentionally not covered.
if RouteChecks._is_get_mcp_server_discovery_route(
route=route, request=request
):
return True
# check if wildcard pattern is allowed
for allowed_route in valid_token.allowed_routes:
if RouteChecks._route_matches_wildcard_pattern(
@ -401,6 +422,31 @@ class RouteChecks:
return True
return False
@staticmethod
def _is_get_mcp_server_discovery_route(
route: str, request: Optional[Request]
) -> bool:
"""
Returns True if `request` is a GET against one of the two read-only
MCP-server discovery paths:
- GET `/v1/mcp/server` (list)
- GET `/v1/mcp/server/{server_id}` (single server, single segment)
Multi-segment paths (`/v1/mcp/server/{id}/approve`, etc.) and any
non-GET method return False, so admin-only management writes on the
same path prefix are not reachable through this carve-out.
"""
if request is None or request.method.upper() != "GET":
return False
if route == "/v1/mcp/server":
return True
prefix = "/v1/mcp/server/"
if not route.startswith(prefix):
return False
remainder = route[len(prefix) :]
return bool(remainder) and "/" not in remainder
@staticmethod
def is_management_route(route: str) -> bool:
"""
@ -627,7 +673,11 @@ class RouteChecks:
Returns:
bool: True if `thread` or `assistant` is in the request path, False otherwise
"""
if "thread" in request.url.path or "assistant" in request.url.path:
# Inline import — auth_utils participates in a proxy import cycle.
from .auth_utils import get_request_route # noqa: PLC0415
route = get_request_route(request)
if "thread" in route or "assistant" in route:
return True
return False

View file

@ -2200,7 +2200,9 @@ async def user_api_key_auth(
user_api_key_auth_obj.budget_reservation = None
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj)
RouteChecks.should_call_route(
route=route, valid_token=user_api_key_auth_obj, request=request
)
# Single authorization point. Builder paths MUST NOT call common_checks.
# Route through the same exception handler the builder uses so

View file

@ -546,7 +546,10 @@ def _add_vector_store_id_from_path(request_data: dict, request: Request) -> None
request_data: The request data dictionary to populate
request: The FastAPI Request object
"""
path = request.url.path
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
path = get_request_route(request)
vector_store_match = re.search(r"/vector_stores/([^/]+)/", path)
if vector_store_match:
vector_store_id = vector_store_match.group(1)

View file

@ -23,11 +23,11 @@ model_list:
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
#########################################################
########## batch specific params ########################
s3_bucket_name: litellm-proxy
s3_bucket_name: litellm-proxy-941277531214
s3_region_name: us-west-2
s3_access_key_id: os.environ/AWS_ACCESS_KEY_ID
s3_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_batch_role_arn: arn:aws:iam::888602223428:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV
aws_batch_role_arn: arn:aws:iam::941277531214:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV
model_info:
mode: batch

View file

@ -55,7 +55,7 @@ guardrails:
litellm_params:
guardrail: bedrock # supported values: "bedrock", "lakera"
mode: "during_call"
guardrailIdentifier: ff6ujrregl1q
guardrailIdentifier: 4w3d1di3snt5
guardrailVersion: "DRAFT"
- guardrail_name: "custom-pre-guard"
litellm_params:

View file

@ -151,7 +151,10 @@ async def test_endpoint(request: Request):
dict: A dictionary containing the route of the request URL.
"""
# ping the proxy server to check if its healthy
return {"route": request.url.path}
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
return {"route": get_request_route(request)}
@router.get(

View file

@ -333,8 +333,10 @@ def _get_metadata_variable_name(request: Request) -> str:
For ALL other endpoints we call this "metadata"
"""
path = request.url.path
# Inline imports — auth_utils/route_checks participate in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
path = get_request_route(request)
if "thread" in path or "assistant" in path:
return "litellm_metadata"

View file

@ -19,6 +19,7 @@ from pydantic import BaseModel, Field
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
from litellm.proxy._types import (
AUDIT_ACTIONS,
LiteLLM_AuditLogs,
@ -34,6 +35,10 @@ from litellm.types.management_endpoints import (
router = APIRouter()
# Cache fields holding credentials. Masked on read so plaintext Redis /
# Sentinel passwords never leave the server in a GET response.
_CACHE_SENSITIVE_FIELDS: set = {"password", "sentinel_password"}
_REDACTED_VALUE = "***REDACTED***"
@ -295,7 +300,11 @@ async def get_cache_settings(
else:
decrypted_settings["redis_type"] = "node"
current_values = decrypted_settings
# Mask credential fields so the GET response never carries
# plaintext Redis / Sentinel passwords off the server.
current_values = mask_sensitive_keys(
decrypted_settings, _CACHE_SENSITIVE_FIELDS
)
# Update field values with current values
for field in cache_fields:

View file

@ -1568,6 +1568,9 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
server_id = request.path_params.get("server_id", "")
if server_id:
@ -1584,7 +1587,7 @@ if MCP_AVAILABLE:
):
# For /token, require PKCE authorization_code; refresh_token
# grants must NOT bypass auth (see comment above).
path_lower = (request.url.path or "").rstrip("/").lower()
path_lower = get_request_route(request).rstrip("/").lower()
if path_lower.endswith("/token"):
body_data = await _read_request_body(request=request)
grant_type = (body_data or {}).get("grant_type", "")

View file

@ -2,7 +2,7 @@
## Helper utils for the management endpoints (keys/users/teams)
from datetime import datetime
from functools import wraps
from typing import List, Optional, Tuple
from typing import Any, Callable, List, Optional, Tuple
from fastapi import HTTPException, Request
@ -435,6 +435,63 @@ async def send_management_endpoint_alert(
)
async def _emit_management_endpoint_otel_span(
func: Callable,
kwargs: dict,
parent_otel_span: Any,
start_time: datetime,
end_time: datetime,
result: Any = None,
exception: Optional[Exception] = None,
) -> None:
"""Stamp + end the parent OTEL SERVER span for a management endpoint.
Routes the request/response (or exception) through the OTEL success/failure
hook. Falls back to ``func.__name__`` for the route when the handler has no
``http_request`` param — endpoints like ``/key/generate`` never receive one,
and gating the hook on it leaked their SERVER span (created in auth, never
ended → never exported). Always emitting keeps both success and failure
paths consistent.
"""
from litellm.proxy.proxy_server import open_telemetry_logger
if open_telemetry_logger is None:
return
http_request: Optional[Request] = kwargs.get("http_request")
if http_request is not None:
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
route = get_request_route(http_request)
request_body: dict = await _read_request_body(request=http_request)
else:
route = func.__name__
request_body = {}
logging_payload = ManagementEndpointLoggingPayload(
route=route,
request_data=request_body,
response=None,
start_time=start_time,
end_time=end_time,
exception=exception,
)
if exception is None:
await open_telemetry_logger.async_management_endpoint_success_hook(
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
else:
await open_telemetry_logger.async_management_endpoint_failure_hook(
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
def management_endpoint_wrapper(func):
"""
This wrapper does the following:
@ -446,13 +503,10 @@ def management_endpoint_wrapper(func):
@wraps(func)
async def wrapper(*args, **kwargs):
start_time = datetime.now()
_http_request: Optional[Request] = None
try:
result = await func(*args, **kwargs)
end_time = datetime.now()
try:
if kwargs is None:
kwargs = {}
user_api_key_dict: UserAPIKeyAuth = (
kwargs.get("user_api_key_dict") or UserAPIKeyAuth()
)
@ -462,31 +516,16 @@ def management_endpoint_wrapper(func):
user_api_key_dict=user_api_key_dict,
function_name=func.__name__,
)
_http_request = kwargs.get("http_request", None)
parent_otel_span = getattr(user_api_key_dict, "parent_otel_span", None)
if parent_otel_span is not None:
from litellm.proxy.proxy_server import open_telemetry_logger
if open_telemetry_logger is not None:
if _http_request:
_route = _http_request.url.path
_request_body: dict = await _read_request_body(
request=_http_request
)
_response = dict(result) if result is not None else None
logging_payload = ManagementEndpointLoggingPayload(
route=_route,
request_data=_request_body,
response=_response,
start_time=start_time,
end_time=end_time,
)
await open_telemetry_logger.async_management_endpoint_success_hook( # type: ignore
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
await _emit_management_endpoint_otel_span(
func=func,
kwargs=kwargs,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
result=result,
)
# Delete updated/deleted info from cache
_delete_api_key_from_cache(kwargs=kwargs)
@ -502,39 +541,19 @@ def management_endpoint_wrapper(func):
except Exception as e:
end_time = datetime.now()
if kwargs is None:
kwargs = {}
user_api_key_dict: UserAPIKeyAuth = (
kwargs.get("user_api_key_dict") or UserAPIKeyAuth()
)
parent_otel_span = getattr(user_api_key_dict, "parent_otel_span", None)
if parent_otel_span is not None:
from litellm.proxy.proxy_server import open_telemetry_logger
if open_telemetry_logger is not None:
_http_request = kwargs.get("http_request")
if _http_request:
_route = _http_request.url.path
_request_body: dict = await _read_request_body(
request=_http_request
)
else:
_route = func.__name__
_request_body = {}
logging_payload = ManagementEndpointLoggingPayload(
route=_route,
request_data=_request_body,
response=None,
start_time=start_time,
end_time=end_time,
exception=e,
)
await open_telemetry_logger.async_management_endpoint_failure_hook( # type: ignore
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
await _emit_management_endpoint_otel_span(
func=func,
kwargs=kwargs,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
exception=e,
)
raise e

View file

@ -1307,11 +1307,14 @@ def create_pass_through_route(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
subpath: str = "", # captures sub-paths when include_subpath=True
):
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
)
path = request.url.path
path = get_request_route(request)
# Parse request data based on content type
(

View file

@ -241,7 +241,10 @@ from litellm.litellm_core_utils.core_helpers import (
)
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.litellm_core_utils.sensitive_data_masker import (
SensitiveDataMasker,
mask_sensitive_keys,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
@ -990,6 +993,15 @@ _OPENAPI_HTTP_METHODS = {
}
# Credentials surfaced by `/get/config/callbacks` in the alerting block: the
# full Slack incoming-webhook URL is itself a credential, and the SMTP
# password is a service password. Masked on read so plaintext never reaches
# the UI. Kept here at module scope to match the analogous
# `_SSO_SENSITIVE_FIELDS` / `_CACHE_SENSITIVE_FIELDS` constants in the SSO
# and cache endpoint files.
_ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"}
def _strip_operation_id_method_suffix(operation_id: str) -> str:
base, separator, suffix = operation_id.rpartition("_")
if separator and suffix in _OPENAPI_HTTP_METHODS:
@ -14708,6 +14720,9 @@ async def get_config(): # noqa: PLR0915
value=env_variable, key=_var
)
_slack_env_vars[_var] = _decrypted_value
_slack_env_vars = mask_sensitive_keys(
_slack_env_vars, _ALERTING_SENSITIVE_VARS
)
_alerting_types = proxy_logging_obj.slack_alerting_instance.alert_types
_all_alert_types = (
@ -14744,6 +14759,7 @@ async def get_config(): # noqa: PLR0915
# decode + decrypt the value
_decrypted_value = decrypt_value_helper(value=env_variable, key=_var)
_email_env_vars[_var] = _decrypted_value
_email_env_vars = mask_sensitive_keys(_email_env_vars, _ALERTING_SENSITIVE_VARS)
alerting_data.append(
{

View file

@ -1817,7 +1817,10 @@ async def ui_view_spend_logs( # noqa: PLR0915
)
try:
is_v2 = "/spend/logs/v2" in request.url.path
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
is_v2 = "/spend/logs/v2" in get_request_route(request)
formats = ["%Y-%m-%d %H:%M:%S", "%Y-%m-%d"] if is_v2 else ["%Y-%m-%d %H:%M:%S"]
def parse_date(date_str: str) -> datetime:

View file

@ -9,6 +9,7 @@ from pydantic.fields import FieldInfo
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.proxy.management_endpoints.ui_sso import (
@ -19,6 +20,16 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
router = APIRouter()
# SSO secret fields returned by /get/sso_settings. These are masked on read so
# the UI can show "(set)" without ever transporting the plaintext OAuth secret
# off the server, matching the write-once + masked-on-read contract used for
# the HashiCorp Vault config override.
_SSO_SENSITIVE_FIELDS: Set[str] = {
"google_client_secret",
"microsoft_client_secret",
"generic_client_secret",
}
class IPAddress(BaseModel):
ip: str
@ -728,8 +739,9 @@ async def get_sso_settings():
schema = TypeAdapter(SSOConfig).json_schema(by_alias=True)
# Convert to dict for response
sso_dict = sso_config.model_dump()
# Convert to dict for response, masking OAuth client secrets so plaintext
# is never sent to the UI.
sso_dict = mask_sensitive_keys(sso_config.model_dump(), _SSO_SENSITIVE_FIELDS)
# Add descriptions to the response
result = {

View file

@ -330,11 +330,16 @@ def is_allowed_to_call_vector_store_endpoint(
provider_config.get_vector_store_endpoints_by_type()
)
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
request_route = get_request_route(request)
# Determine the permission type based on the request
permission_type = None
for endpoint in provider_vector_store_endpoints["read"]:
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "read"
break
@ -342,7 +347,7 @@ def is_allowed_to_call_vector_store_endpoint(
if permission_type is None:
for endpoint in provider_vector_store_endpoints["write"]:
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "write"
break
@ -392,10 +397,15 @@ def is_allowed_to_call_vector_store_files_endpoint(
provider_config.get_vector_store_file_endpoints_by_type()
)
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
request_route = get_request_route(request)
permission_type: Optional[str] = None
for endpoint in provider_vector_store_endpoints.get("read", ()):
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "read"
break
@ -403,7 +413,7 @@ def is_allowed_to_call_vector_store_files_endpoint(
if permission_type is None:
for endpoint in provider_vector_store_endpoints.get("write", ()):
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "write"
break

View file

@ -54,6 +54,7 @@ if TYPE_CHECKING:
else:
ResponseText = str # Fallback for ResponseText import
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.llms.openai.data_residency import infer_openai_data_residency
from litellm.secret_managers.main import get_secret_str
from litellm.types.responses.main import *
from litellm.types.router import GenericLiteLLMParams
@ -1139,6 +1140,9 @@ def responses(
"aresponses": _is_async,
"litellm_call_id": litellm_call_id,
"model_info": kwargs.get("model_info"),
"data_residency": infer_openai_data_residency(
custom_llm_provider, litellm_params.api_base
),
"metadata": (
kwargs["litellm_metadata"]
if "litellm_metadata" in kwargs
@ -2032,6 +2036,9 @@ def compact_responses(
litellm_params={
**responses_api_request_params,
"litellm_call_id": litellm_call_id,
"data_residency": infer_openai_data_residency(
custom_llm_provider, litellm_params.api_base
),
},
custom_llm_provider=custom_llm_provider,
)
@ -2129,6 +2136,11 @@ async def _aresponses_websocket(
api_key=api_key,
)
litellm_params_dict["data_residency"] = infer_openai_data_residency(
_custom_llm_provider,
dynamic_api_base or litellm_params.api_base or litellm.api_base,
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,

View file

@ -219,6 +219,12 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
output_cost_per_token_priority: Optional[
float
] # OpenAI priority service tier pricing
regional_processing_uplift_multiplier_eu: Optional[
float
] # OpenAI EU data-residency uplift multiplier applied to all token costs (e.g. 1.10 = +10%)
regional_processing_uplift_multiplier_us: Optional[
float
] # OpenAI US data-residency uplift multiplier applied to all token costs (e.g. 1.10 = +10%)
output_cost_per_character: Optional[float] # only for vertex ai models
output_cost_per_audio_token: Optional[float]
output_cost_per_token_above_128k_tokens: Optional[
@ -3602,6 +3608,20 @@ class ServiceTier(Enum):
PRIORITY = "priority"
class DataResidency(Enum):
"""
OpenAI data-residency / regional-processing regions.
Inferred from the OpenAI api_base host (eu.api.openai.com -> EU,
us.api.openai.com -> US). Used to apply the regional-processing
cost uplift (see ``regional_processing_uplift_multiplier_<region>``
on ModelInfo).
"""
US = "us"
EU = "eu"
LLMResponseTypes = Union[
ModelResponse,
EmbeddingResponse,

View file

@ -5942,6 +5942,12 @@ def _get_model_info_helper( # noqa: PLR0915
output_cost_per_token_priority=_model_info.get(
"output_cost_per_token_priority", None
),
regional_processing_uplift_multiplier_eu=_model_info.get(
"regional_processing_uplift_multiplier_eu", None
),
regional_processing_uplift_multiplier_us=_model_info.get(
"regional_processing_uplift_multiplier_us", None
),
output_cost_per_audio_token=_model_info.get(
"output_cost_per_audio_token", None
),

View file

@ -19050,6 +19050,8 @@
"output_cost_per_token": 8e-06,
"output_cost_per_token_batches": 4e-06,
"output_cost_per_token_priority": 1.4e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -19123,6 +19125,8 @@
"output_cost_per_token": 1.6e-06,
"output_cost_per_token_batches": 8e-07,
"output_cost_per_token_priority": 2.8e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -19196,6 +19200,8 @@
"output_cost_per_token": 4e-07,
"output_cost_per_token_batches": 2e-07,
"output_cost_per_token_priority": 8e-07,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -19267,6 +19273,8 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"output_cost_per_token_priority": 1.7e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -19308,6 +19316,8 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -19329,6 +19339,8 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_batches": 5e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -19617,6 +19629,8 @@
"output_cost_per_token": 6e-07,
"output_cost_per_token_batches": 3e-07,
"output_cost_per_token_priority": 1e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
@ -20320,6 +20334,8 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_flex": 5e-06,
"output_cost_per_token_priority": 2e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -21242,6 +21258,8 @@
"mode": "responses",
"output_cost_per_token": 0.00012,
"output_cost_per_token_batches": 6e-05,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/batch",
"/v1/responses"
@ -21648,6 +21666,8 @@
"output_cost_per_token": 2e-06,
"output_cost_per_token_flex": 1e-06,
"output_cost_per_token_priority": 3.6e-06,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@ -21729,6 +21749,8 @@
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"regional_processing_uplift_multiplier_eu": 1.10,
"regional_processing_uplift_multiplier_us": 1.10,
"mode": "chat",
"output_cost_per_token": 4e-07,
"output_cost_per_token_flex": 2e-07,

View file

@ -168,7 +168,7 @@ async def test_a2a_completion_bridge_bedrock_agentcore():
litellm._turn_on_debug()
# Bedrock AgentCore ARN (streaming-capable runtime)
agentcore_arn = "arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
agentcore_arn = "arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
send_message_payload = {
"message": {

View file

@ -145,6 +145,37 @@ def test_batch_cost_calculator_func_uses_custom_model_info():
), f"Expected total cost {expected}, got {cost}"
@pytest.mark.parametrize("data_residency", ["eu", "us"])
def test_batch_cost_calculator_applies_data_residency_uplift(
data_residency, monkeypatch
):
"""batch_cost_calculator should apply the regional uplift multiplier when
data_residency is set and the model carries a configured multiplier."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
prev_model_cost = litellm.model_cost
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base_prompt, base_completion = batch_cost_calculator(
usage=usage,
model="gpt-5",
custom_llm_provider="openai",
)
regional_prompt, regional_completion = batch_cost_calculator(
usage=usage,
model="gpt-5",
custom_llm_provider="openai",
data_residency=data_residency,
)
assert base_prompt > 0 and base_completion > 0
assert regional_prompt == pytest.approx(base_prompt * 1.10, rel=1e-9)
assert regional_completion == pytest.approx(base_completion * 1.10, rel=1e-9)
finally:
litellm.model_cost = prev_model_cost
@pytest.mark.asyncio
async def test_calculate_batch_cost_and_usage_uses_custom_model_info():
"""calculate_batch_cost_and_usage should thread model_info."""

View file

@ -38,7 +38,7 @@ async def test_async_create_file():
file=open(file_path, "rb"),
purpose="batch",
custom_llm_provider="bedrock",
s3_bucket_name="litellm-proxy",
s3_bucket_name="litellm-proxy-941277531214",
)
@ -55,7 +55,7 @@ async def test_async_file_and_batch():
file=open(file_path, "rb"),
purpose="batch",
custom_llm_provider="bedrock",
s3_bucket_name="litellm-proxy",
s3_bucket_name="litellm-proxy-941277531214",
)
print("CREATED FILE RESPONSE=", file_obj)
@ -70,7 +70,7 @@ async def test_async_file_and_batch():
# bedrock specific params
#########################################################
model="us.anthropic.claude-haiku-4-5-20251001-v1:0",
aws_batch_role_arn="arn:aws:iam::888602223428:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV",
aws_batch_role_arn="arn:aws:iam::941277531214:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV",
)
print("CREATED BATCH RESPONSE=", create_batch_response)
@ -129,7 +129,7 @@ async def test_mock_bedrock_file_url_mapping():
),
purpose="batch",
custom_llm_provider="bedrock",
s3_bucket_name="litellm-proxy",
s3_bucket_name="litellm-proxy-941277531214",
)
print(f"PUT URL: {captured_put_url}")

View file

@ -20,7 +20,7 @@ async def test_bedrock_guardrails_pii_masking():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
)
@ -60,7 +60,7 @@ async def test_bedrock_guardrails_pii_masking_content_list():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
)
@ -115,7 +115,7 @@ async def test_bedrock_guardrails_block_messages_api():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
)
@ -166,7 +166,7 @@ async def test_bedrock_guardrails_block_responses_api():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
)
@ -211,7 +211,7 @@ async def test_bedrock_guardrails_with_streaming():
)
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
supported_event_hooks=[GuardrailEventHooks.post_call],
guardrail_name="bedrock-post-guard",
@ -255,7 +255,7 @@ async def test_bedrock_guardrails_with_streaming_no_violation():
)
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
supported_event_hooks=[GuardrailEventHooks.post_call],
guardrail_name="bedrock-post-guard",
@ -299,7 +299,7 @@ async def test_bedrock_guardrails_streaming_request_body_mock():
# Create the guardrail
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
supported_event_hooks=[GuardrailEventHooks.post_call],
guardrail_name="bedrock-post-guard",
@ -382,7 +382,7 @@ async def test_bedrock_guardrail_aws_param_persistence():
from litellm.types.guardrails import GuardrailEventHooks
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
aws_access_key_id="test-access-key",
aws_secret_access_key="test-secret-key",

View file

@ -1,3 +1,4 @@
import json
import logging
import os
import sys
@ -44,6 +45,9 @@ from litellm.llms.bedrock.image_generation.image_handler import (
)
from litellm.llms.bedrock.common_utils import BedrockError
# Base64 placeholder used for mocked Bedrock image responses (a 1x1 PNG).
_MOCK_BEDROCK_IMAGE_B64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
@pytest.mark.parametrize(
"model,expected",
@ -528,17 +532,34 @@ def test_backward_compatibility_regular_nova_model():
def test_amazon_titan_image_gen():
"""Test Amazon Titan image generation with cost tracking."""
from litellm import image_generation
"""Test Amazon Titan image generation with cost tracking.
The Bedrock CI account is not entitled to amazon.titan-image-generator, so
the network call is mocked and only the transform + cost-tracking path is
exercised.
"""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
# Use v2 as v1 has reached end of life
model_id = "bedrock/amazon.titan-image-generator-v2:0"
response = litellm.image_generation(
model=model_id,
prompt="A serene mountain landscape at sunset with a lake reflection",
aws_region_name="us-east-1",
)
mock_payload = {"images": [_MOCK_BEDROCK_IMAGE_B64]}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = mock_payload
mock_response.text = json.dumps(mock_payload)
mock_response.headers = {}
client = HTTPHandler()
with patch.object(client, "post", return_value=mock_response):
response = litellm.image_generation(
model=model_id,
prompt="A serene mountain landscape at sunset with a lake reflection",
aws_region_name="us-east-1",
aws_access_key_id="fake-access-key-id",
aws_secret_access_key="fake-secret-access-key",
client=client,
)
print(f"response cost: {response._hidden_params['response_cost']}")

View file

@ -7,7 +7,6 @@ import sys
import traceback
from unittest.mock import AsyncMock, MagicMock, patch
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
@ -136,6 +135,51 @@ class TestVertexAIGeminiImageGeneration(BaseImageGenTest):
}
# Base64 placeholder used for mocked Bedrock image responses (a 1x1 PNG).
_MOCK_BEDROCK_IMAGE_B64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
async def _assert_mocked_bedrock_image_generation(call_args: dict) -> None:
"""Run ``aimage_generation`` with the Bedrock HTTP call mocked.
The CI account is not entitled to Nova Canvas, so the network call is
replaced with a canned Bedrock response. This keeps the request transform,
response transform, and cost-tracking path under test without live access.
"""
mock_payload = {"images": [_MOCK_BEDROCK_IMAGE_B64]}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = mock_payload
mock_response.text = json.dumps(mock_payload)
mock_response.headers = {}
custom_logger = TestCustomLogger()
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [custom_logger]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
return_value=mock_response,
):
response = await litellm.aimage_generation(
**call_args,
prompt="A image of a otter",
aws_access_key_id="fake-access-key-id",
aws_secret_access_key="fake-secret-access-key",
)
await asyncio.sleep(1)
assert custom_logger.standard_logging_payload is not None
assert custom_logger.standard_logging_payload["response_cost"] is not None
assert custom_logger.standard_logging_payload["response_cost"] > 0
assert response.data is not None
for d in response.data:
assert isinstance(d, Image)
assert d.b64_json is not None or d.url is not None
class TestBedrockNovaCanvasTextToImage(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
litellm.in_memory_llm_clients_cache = InMemoryCache()
@ -148,6 +192,12 @@ class TestBedrockNovaCanvasTextToImage(BaseImageGenTest):
"aws_region_name": "us-east-1",
}
@pytest.mark.asyncio(scope="module")
async def test_basic_image_generation(self):
await _assert_mocked_bedrock_image_generation(
self.get_base_image_generation_call_args()
)
class TestBedrockNovaCanvasColorGuidedGeneration(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
@ -162,6 +212,12 @@ class TestBedrockNovaCanvasColorGuidedGeneration(BaseImageGenTest):
"aws_region_name": "us-east-1",
}
@pytest.mark.asyncio(scope="module")
async def test_basic_image_generation(self):
await _assert_mocked_bedrock_image_generation(
self.get_base_image_generation_call_args()
)
class TestOpenAIGPTImage1(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:

View file

@ -82,7 +82,7 @@ async def _vertex_ai_mocks():
"bedrock/mistral.mistral-7b-instruct-v0:2",
"openai/gpt-4o",
"openai/self_hosted",
"bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"vertex_ai/gemini-1.5-flash",
],
)
@ -147,7 +147,7 @@ async def test_litellm_overhead_non_streaming(model):
[
"bedrock/mistral.mistral-7b-instruct-v0:2",
"openai/gpt-4o",
"bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"openai/self_hosted",
],
)

View file

@ -1,7 +1,6 @@
from dataclasses import dataclass, field
from typing import Dict, FrozenSet, List, Optional, Tuple
OMIT = object()
@ -22,6 +21,7 @@ class ModelEntry:
extra_params: Tuple[Tuple[str, str], ...] = field(default_factory=tuple)
required_env: FrozenSet[str] = field(default_factory=frozenset)
caps: FrozenSet[str] = field(default_factory=frozenset)
fail_reason: Optional[str] = None
def params(self) -> Dict[str, str]:
return dict(self.extra_params)
@ -205,6 +205,12 @@ BEDROCK_CONVERSE_MODELS: Tuple[ModelEntry, ...] = (
extra_params=(("aws_region_name", "us-east-1"),),
required_env=_BEDROCK_REQ,
caps=_CAPS_OPUS_4_7,
fail_reason=(
"claude-opus-4-7 is not entitled on the Bedrock CI account "
"941277531214 (model access requires an AWS Sales request, not "
"self-serve); this cell fails on purpose so it stays loud in CI — "
"remove this fail_reason once access is granted"
),
),
ModelEntry(
alias="bedrock-claude-opus-4-6",

View file

@ -15,7 +15,6 @@ from .grid_spec import (
all_cells,
)
_PROMPT_MESSAGES: List[Dict[str, str]] = [
{"role": "user", "content": "Step by step, calculate 47 * 53. Show your work."}
]
@ -168,6 +167,9 @@ async def test_reasoning_effort_grid(
if skip_reason:
pytest.skip(skip_reason)
if model.fail_reason:
pytest.xfail(model.fail_reason)
if route_name == "bedrock_invoke_messages":
status, exc = await _call_messages(model, effort)
else:

View file

@ -19,8 +19,8 @@ import httpx
@pytest.mark.parametrize(
"model",
[
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_13sf6-cALnp38iZD", # non-streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC", # streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_13sf6-4046UzHSwy", # non-streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp", # streaming invocation
],
)
def test_bedrock_agentcore_basic(model):
@ -44,7 +44,7 @@ def test_bedrock_agentcore_basic(model):
@pytest.mark.parametrize(
"model",
[
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_13sf6-cALnp38iZD", # streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_13sf6-4046UzHSwy", # streaming invocation
],
)
async def test_bedrock_agentcore_with_streaming(model):
@ -54,7 +54,7 @@ async def test_bedrock_agentcore_with_streaming(model):
print("running streming test for model=", model)
# litellm._turn_on_debug()
response = await litellm.acompletion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -82,7 +82,7 @@ def test_bedrock_agentcore_with_custom_params():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -105,7 +105,7 @@ def test_bedrock_agentcore_with_custom_params():
url = call_kwargs["url"]
print(f"URL: {url}")
assert (
"/runtimes/arn%3Aaws%3Abedrock-agentcore%3Aus-west-2%3A888602223428%3Aruntime%2Fhosted_agent_r9jvp-3ySZuRHjLC/invocations"
"/runtimes/arn%3Aaws%3Abedrock-agentcore%3Aus-west-2%3A941277531214%3Aruntime%2Fhosted_agent_r9jvp-Rq79QFC2fp/invocations"
in url
)
assert "qualifier=DEFAULT" in url
@ -150,7 +150,7 @@ def test_bedrock_agentcore_with_runtime_user_id():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -189,7 +189,7 @@ def test_bedrock_agentcore_with_session_and_user():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -234,7 +234,7 @@ def test_bedrock_agentcore_with_api_key_bearer_token():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -282,7 +282,7 @@ def test_bedrock_agentcore_with_all_parameters():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -350,7 +350,7 @@ def test_bedrock_agentcore_without_api_key_uses_sigv4():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -625,7 +625,7 @@ def test_agentcore_synchronous_non_streaming_response():
with patch.object(client, "post", return_value=mock_response) as mock_post:
# Make a synchronous (non-streaming) completion call
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",

View file

@ -115,7 +115,7 @@ def test_completion_bedrock_guardrails(streaming):
],
max_tokens=10,
guardrailConfig={
"guardrailIdentifier": "ff6ujrregl1q",
"guardrailIdentifier": "4w3d1di3snt5",
"guardrailVersion": "DRAFT",
"trace": "enabled",
},
@ -144,7 +144,7 @@ def test_completion_bedrock_guardrails(streaming):
stream=True,
max_tokens=10,
guardrailConfig={
"guardrailIdentifier": "ff6ujrregl1q",
"guardrailIdentifier": "4w3d1di3snt5",
"guardrailVersion": "DRAFT",
"trace": "enabled",
},
@ -475,7 +475,7 @@ def test_bedrock_claude_3(image_url):
],
}
response: ModelResponse = completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
num_retries=3,
**data,
) # type: ignore
@ -498,7 +498,7 @@ def test_bedrock_claude_3(image_url):
@pytest.mark.parametrize(
"model",
[
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
# "meta.llama3-70b-instruct-v1:0",
# "anthropic.claude-v2",
# "mistral.mixtral-8x7b-instruct-v0:1",
@ -537,7 +537,7 @@ def test_bedrock_stop_value(stop, model):
@pytest.mark.parametrize(
"model",
[
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
"mistral.mixtral-8x7b-instruct-v0:1",
],
)
@ -602,7 +602,7 @@ def test_bedrock_claude_3_tool_calling():
}
]
response: ModelResponse = completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
tools=tools,
tool_choice="auto",
@ -630,7 +630,7 @@ def test_bedrock_claude_3_tool_calling():
)
# In the second response, Claude should deduce answer from tool results
second_response = completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
tools=tools,
tool_choice="auto",
@ -737,7 +737,7 @@ def test_bedrock_ptu():
from openai.types.chat import ChatCompletion
model_id = (
"arn:aws:bedrock:us-west-2:888602223428:provisioned-model/8fxff74qyhs3"
"arn:aws:bedrock:us-west-2:941277531214:provisioned-model/8fxff74qyhs3"
)
try:
response = litellm.completion(
@ -752,7 +752,7 @@ def test_bedrock_ptu():
assert "url" in mock_client_post.call_args.kwargs
assert (
mock_client_post.call_args.kwargs["url"]
== "https://bedrock-runtime.us-west-2.amazonaws.com/model/arn%3Aaws%3Abedrock%3Aus-west-2%3A888602223428%3Aprovisioned-model%2F8fxff74qyhs3/converse"
== "https://bedrock-runtime.us-west-2.amazonaws.com/model/arn%3Aaws%3Abedrock%3Aus-west-2%3A941277531214%3Aprovisioned-model%2F8fxff74qyhs3/converse"
)
mock_client_post.assert_called_once()
@ -2327,7 +2327,7 @@ def test_bedrock_cross_region_inference(monkeypatch):
def test_bedrock_empty_content_real_call():
completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=[
{
"role": "user",

View file

@ -299,7 +299,10 @@ def test_completion_claude_3():
@pytest.mark.parametrize(
"model",
["anthropic/claude-sonnet-4-5-20250929", "anthropic.claude-3-sonnet-20240229-v1:0"],
[
"anthropic/claude-sonnet-4-5-20250929",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
],
)
def test_completion_claude_3_function_call(model):
litellm.set_verbose = True
@ -385,7 +388,7 @@ def test_completion_claude_3_function_call(model):
[
("gpt-3.5-turbo", None, None),
("claude-sonnet-4-5-20250929", None, None),
("anthropic.claude-3-sonnet-20240229-v1:0", None, None),
("us.anthropic.claude-sonnet-4-5-20250929-v1:0", None, None),
# (
# "azure_ai/command-r-plus",
# os.getenv("AZURE_COHERE_API_KEY"),
@ -1578,7 +1581,7 @@ def test_completion_openai():
[
# ("gpt-4o-2024-08-06", None),
# ("azure/gpt-4.1-mini", None),
("bedrock/anthropic.claude-3-sonnet-20240229-v1:0", None),
("bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", None),
# ("azure/gpt-4o-new-test", "2024-08-01-preview"),
],
)
@ -1666,15 +1669,13 @@ def custom_callback(
#################################################
print(
f"""
print(f"""
Model: {model},
Messages: {messages},
User: {user},
Seed: {kwargs["seed"]},
temperature: {kwargs["temperature"]},
"""
)
""")
assert kwargs["user"] == "ishaans app"
assert kwargs["model"] == "gpt-3.5-turbo-1106"
@ -2699,7 +2700,7 @@ def test_bedrock_deepseek_custom_prompt_dict():
def test_bedrock_deepseek_known_tokenizer_config(monkeypatch):
model = (
"deepseek_r1/arn:aws:bedrock:us-west-2:888602223428:imported-model/bnnr6463ejgf"
"deepseek_r1/arn:aws:bedrock:us-west-2:941277531214:imported-model/bnnr6463ejgf"
)
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from unittest.mock import Mock
@ -2914,8 +2915,8 @@ def response_format_tests(response: litellm.ModelResponse):
"model",
[
"bedrock/mistral.mistral-large-2407-v1:0",
"bedrock/cohere.command-r-plus-v1:0",
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-haiku-4-5-20251001-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
"mistral.mistral-7b-instruct-v0:2",
"meta.llama3-8b-instruct-v1:0",
],

View file

@ -142,7 +142,8 @@ def trade(model_name: str) -> List[Trade]: # type: ignore
@pytest.mark.parametrize(
"model", ["claude-haiku-4-5-20251001", "anthropic.claude-3-haiku-20240307-v1:0"]
"model",
["claude-haiku-4-5-20251001", "us.anthropic.claude-haiku-4-5-20251001-v1:0"],
)
@pytest.mark.flaky(retries=6, delay=10)
def test_function_call_parsing(model):

View file

@ -49,7 +49,7 @@ def get_current_weather(location, unit="fahrenheit"):
"mistral/mistral-large-latest",
"claude-haiku-4-5-20251001",
"gemini/gemini-2.5-flash-lite",
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
],
)
@pytest.mark.flaky(retries=3, delay=1)
@ -267,7 +267,6 @@ def test_aaparallel_function_call_with_anthropic_thinking(model):
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
_PARALLEL_TOOL_HISTORY_MESSAGES = [
{
"role": "user",
@ -303,7 +302,7 @@ _PARALLEL_TOOL_HISTORY_MESSAGES = [
[
# Bedrock Converse still requires modify_params to inject the dummy tool.
(
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
_PARALLEL_TOOL_HISTORY_MESSAGES,
True,
),
@ -314,7 +313,7 @@ _PARALLEL_TOOL_HISTORY_MESSAGES = [
False,
),
(
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
[
{
"role": "user",
@ -579,7 +578,7 @@ def test_groq_parallel_function_call():
@pytest.mark.parametrize(
"model",
[
"bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
],
)
def test_passing_tool_result_as_list(model):

View file

@ -57,7 +57,7 @@ async def test_completion_sagemaker(sync_mode):
print("testing sagemaker")
if sync_mode is True:
response = litellm.completion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -67,7 +67,7 @@ async def test_completion_sagemaker(sync_mode):
)
else:
response = await litellm.acompletion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -158,7 +158,7 @@ async def test_completion_sagemaker_messages_api(sync_mode):
"model",
[
# "sagemaker_chat/huggingface-pytorch-tgi-inference-2024-08-23-15-48-59-245",
"sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"sagemaker/litellm-ci-textgen",
],
)
# @pytest.mark.flaky(retries=3, delay=1)
@ -218,7 +218,7 @@ async def test_completion_sagemaker_stream(sync_mode, model):
"model",
[
# "sagemaker_chat/huggingface-pytorch-tgi-inference-2024-08-23-15-48-59-245",
"sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"sagemaker/litellm-ci-textgen",
],
)
async def test_completion_sagemaker_streaming_bad_request(sync_mode, model):
@ -256,7 +256,7 @@ async def test_acompletion_sagemaker_non_stream():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -282,7 +282,7 @@ async def test_acompletion_sagemaker_non_stream():
) as mock_post:
# Act: Call the litellm.acompletion function
response = await litellm.acompletion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -302,7 +302,7 @@ async def test_acompletion_sagemaker_non_stream():
assert args_to_sagemaker == expected_payload
assert (
kwargs["url"]
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations"
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/litellm-ci-textgen/invocations"
)
@ -316,7 +316,7 @@ async def test_completion_sagemaker_non_stream():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -342,7 +342,7 @@ async def test_completion_sagemaker_non_stream():
) as mock_post:
# Act: Call the litellm.acompletion function
response = litellm.completion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -362,7 +362,7 @@ async def test_completion_sagemaker_non_stream():
assert args_to_sagemaker == expected_payload
assert (
kwargs["url"]
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations"
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/litellm-ci-textgen/invocations"
)
@ -377,7 +377,7 @@ async def test_completion_sagemaker_prompt_template_non_stream():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -433,7 +433,7 @@ async def test_completion_sagemaker_non_stream_with_aws_params():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -459,7 +459,7 @@ async def test_completion_sagemaker_non_stream_with_aws_params():
) as mock_post:
# Act: Call the litellm.acompletion function
response = litellm.completion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -482,5 +482,5 @@ async def test_completion_sagemaker_non_stream_with_aws_params():
assert args_to_sagemaker == expected_payload
assert (
kwargs["url"]
== "https://runtime.sagemaker.us-west-5.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations"
== "https://runtime.sagemaker.us-west-5.amazonaws.com/endpoints/litellm-ci-textgen/invocations"
)

View file

@ -1174,7 +1174,7 @@ async def test_completion_replicate_llama3_streaming(sync_mode):
[
# ["bedrock/ai21.jamba-instruct-v1:0", "us-east-1"],
# ["bedrock/cohere.command-r-plus-v1:0", None],
["anthropic.claude-3-sonnet-20240229-v1:0", None],
["us.anthropic.claude-sonnet-4-5-20250929-v1:0", None],
# ["mistral.mistral-7b-instruct-v0:2", None],
# ["meta.llama3-8b-instruct-v1:0", None],
],
@ -1246,7 +1246,7 @@ def test_bedrock_claude_3_streaming():
try:
litellm.set_verbose = True
response: ModelResponse = completion( # type: ignore
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
max_tokens=10, # type: ignore
stream=True,
@ -1276,7 +1276,7 @@ def test_bedrock_claude_3_streaming():
"model",
[
"claude-haiku-4-5-20251001",
"cohere.command-r-plus-v1:0", # bedrock
"us.anthropic.claude-haiku-4-5-20251001-v1:0", # bedrock
"gpt-3.5-turbo",
],
)
@ -3500,7 +3500,7 @@ def test_unit_test_perplexity_citations_chunk():
[
"gpt-3.5-turbo",
"claude-sonnet-4-5-20250929",
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
# "vertex_ai/claude-3-5-sonnet@20240620",
],
)

View file

@ -27,7 +27,7 @@ async def test_basic_s3_logging(sync_mode, streaming):
verbose_logger.setLevel(level=logging.DEBUG)
litellm.success_callback = ["s3"]
litellm.s3_callback_params = {
"s3_bucket_name": "load-testing-oct",
"s3_bucket_name": "load-testing-oct-941277531214",
"s3_aws_secret_access_key": "os.environ/AWS_SECRET_ACCESS_KEY",
"s3_aws_access_key_id": "os.environ/AWS_ACCESS_KEY_ID",
"s3_region_name": "us-west-2",
@ -64,14 +64,14 @@ async def test_basic_s3_logging(sync_mode, streaming):
await asyncio.sleep(2)
print(f"response: {response}")
total_objects, all_s3_keys = list_all_s3_objects("load-testing-oct")
total_objects, all_s3_keys = list_all_s3_objects("load-testing-oct-941277531214")
# assert that atlest one key has response.id in it
assert any(response_id in key for key in all_s3_keys)
s3 = boto3.client("s3")
# delete all objects
for key in all_s3_keys:
s3.delete_object(Bucket="load-testing-oct", Key=key)
s3.delete_object(Bucket="load-testing-oct-941277531214", Key=key)
@pytest.mark.asyncio
@ -82,7 +82,7 @@ async def test_basic_s3_v2_logging(streaming):
from litellm.integrations.s3_v2 import S3Logger
litellm.s3_callback_params = {
"s3_bucket_name": "load-testing-oct",
"s3_bucket_name": "load-testing-oct-941277531214",
"s3_aws_secret_access_key": "test-secret",
"s3_aws_access_key_id": "test-key",
"s3_region_name": "us-west-2",

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
@ -67,7 +66,7 @@ def setup_vector_store_registry():
litellm.vector_store_registry = VectorStoreRegistry(
vector_stores=[
LiteLLM_ManagedVectorStore(
vector_store_id="T37J8R4WTM", custom_llm_provider="bedrock"
vector_store_id="LCYXFBR2TU", custom_llm_provider="bedrock"
)
]
)
@ -111,7 +110,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_completion(
response = await litellm.acompletion(
model="anthropic/claude-3.5-sonnet",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=client,
)
except Exception as e:
@ -152,7 +151,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call(
response = await litellm.acompletion(
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=async_client,
)
print("OPENAI RESPONSE:", json.dumps(dict(response), indent=4, default=str))
@ -196,7 +195,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_streaming(
response = await litellm.acompletion(
model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
stream=True,
client=async_client,
)
@ -255,7 +254,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools(
model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}",
messages=[{"role": "user", "content": "what is litellm?"}],
max_tokens=10,
tools=[{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}],
tools=[{"type": "file_search", "vector_store_ids": ["LCYXFBR2TU"]}],
)
assert response is not None
@ -279,7 +278,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools_
tools=[
{
"type": "file_search",
"vector_store_ids": ["T37J8R4WTM"],
"vector_store_ids": ["LCYXFBR2TU"],
"filters": {
"key": "user_id",
"value": "fake-user-id",
@ -387,7 +386,7 @@ async def test_bedrock_kb_request_body_has_transformed_filters(
tools=[
{
"type": "file_search",
"vector_store_ids": ["T37J8R4WTM"],
"vector_store_ids": ["LCYXFBR2TU"],
"filters": {
"key": "user_id",
"value": "fake-user-id",
@ -461,7 +460,7 @@ async def test_openai_with_knowledge_base_mock_openai(setup_vector_store_registr
await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=client,
)
except Exception as e:
@ -537,7 +536,7 @@ async def test_openai_with_vector_store_ids_in_tool_call_mock_openai(
await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}],
tools=[{"type": "file_search", "vector_store_ids": ["LCYXFBR2TU"]}],
client=client,
)
except Exception as e:
@ -611,7 +610,7 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[
{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]},
{"type": "file_search", "vector_store_ids": ["LCYXFBR2TU"]},
{"type": "file_search", "vector_store_ids": ["unknownVS"]},
],
client=client,
@ -645,7 +644,7 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist
# model="gpt-5.5",
# messages=[{"role": "user", "content": "what is litellm?"}],
# vector_store_ids = [
# "T37J8R4WTM"
# "LCYXFBR2TU"
# ],
# )
@ -667,7 +666,7 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist
# # expect the vector store request metadata object to have the correct values
# vector_store_request_metadata = standard_logging_vector_store_request_metadata[0]
# assert vector_store_request_metadata.get("vector_store_id") == "T37J8R4WTM"
# assert vector_store_request_metadata.get("vector_store_id") == "LCYXFBR2TU"
# assert vector_store_request_metadata.get("query") == "what is litellm?"
# assert vector_store_request_metadata.get("custom_llm_provider") == "bedrock"
@ -723,7 +722,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_without_vector_store_registry
response = await litellm.acompletion(
model="anthropic/claude-3.5-sonnet",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=client,
)
except Exception as e:

View file

@ -217,9 +217,120 @@ def _create_request_with_host_header(path: str, host_header: str) -> Request:
],
)
def test_get_request_route_not_bypassed_by_malformed_host(host_header: str):
for protected_path in ["/health", "/user/new", "/key/generate", "/get/internal_user_settings"]:
request = _create_request_with_host_header(path=protected_path, host_header=host_header)
result = get_request_route(request)
assert result == protected_path, (
f"Host: {host_header!r} caused route {protected_path!r} to resolve as {result!r}"
for protected_path in [
"/health",
"/user/new",
"/key/generate",
"/get/internal_user_settings",
]:
request = _create_request_with_host_header(
path=protected_path, host_header=host_header
)
result = get_request_route(request)
assert (
result == protected_path
), f"Host: {host_header!r} caused route {protected_path!r} to resolve as {result!r}"
# ---------------------------------------------------------------------------
# Regression tests for variant call sites that previously read request.url.path
# (Host-derived) instead of the ASGI scope path. Each test sends a Host header
# crafted to collapse url.path to a substring the call site's decision logic
# would match on, while scope["path"] is the real (unmatching) route.
# ---------------------------------------------------------------------------
_BYPASS_HOSTS = [
"localhost/?x=1",
"localhost:4000/?x=1",
"localhost/#test",
"localhost:4000/#test",
]
def _is_assistants(req):
return RouteChecks._is_assistants_api_request(req)
def _metadata_var_name(req):
from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name
return _get_metadata_variable_name(req)
def _vector_store_id_in_path(req):
from litellm.proxy.common_utils.http_parsing_utils import (
_add_vector_store_id_from_path,
)
data: dict = {}
_add_vector_store_id_from_path(request_data=data, request=req)
return "vector_store_id" in data
# (label, scope_path, host_suffix_template, predicate, expected) — host_suffix_template
# receives the host_header via %s substitution. The predicate is invoked on a Request
# whose scope["path"] is scope_path and whose Host header is the formatted suffix.
#
# The MCP entries (well_known_mcp_bypass, pkce_token_suffix) call
# get_request_route directly rather than the surrounding production handler
# (MCPRequestHandler.process_mcp_request / _mcp_oauth_user_api_key_auth) —
# those handlers require an ASGI scope plus MCP state to invoke, and the call
# sites do nothing with the path except feed it to this helper. The helper-
# level assertion is the relevant signal.
_CALL_SITES = [
("assistants_classification", "/key/generate", "%s/thread", _is_assistants, False),
(
"metadata_variable_name",
"/chat/completions",
"%s/thread",
_metadata_var_name,
"metadata",
),
(
"vector_store_id_extraction",
"/key/generate",
"%s/vector_stores/x/files",
_vector_store_id_in_path,
False,
),
(
"well_known_mcp_bypass",
"/mcp/tools/call",
"/.well-known/%s",
lambda r: get_request_route(r).startswith("/.well-known/"),
False,
),
(
"pkce_token_suffix",
"/mcp/server-id/token",
"%s",
lambda r: get_request_route(r).rstrip("/").lower().endswith("/token"),
True,
),
(
"spend_logs_v2_classification",
"/spend/logs",
"%s/spend/logs/v2",
lambda r: "/spend/logs/v2" in get_request_route(r),
False,
),
("health_route_echo", "/test", "%s", lambda r: get_request_route(r), "/test"),
]
@pytest.mark.parametrize("host_header", _BYPASS_HOSTS)
@pytest.mark.parametrize(
"label,scope_path,host_suffix_template,predicate,expected",
_CALL_SITES,
ids=[c[0] for c in _CALL_SITES],
)
def test_call_site_uses_scope_path(
label, scope_path, host_suffix_template, predicate, expected, host_header
):
"""Each call site that previously read request.url.path must now make its
decision against scope["path"]. The Host header is crafted so url.path
would resolve to a value that flips the decision under the old code."""
request = _create_request_with_host_header(
path=scope_path, host_header=host_suffix_template % host_header
)
assert predicate(request) == expected

View file

@ -3,6 +3,7 @@ async_management_endpoint_{success,failure}_hook integration points."""
import asyncio
from datetime import datetime
from unittest.mock import MagicMock
import pytest
@ -14,6 +15,7 @@ from litellm.proxy._types import (
from ._helpers import (
HttpStatusException,
assert_server_span_attrs,
get_server_span,
make_fastapi_http_exception,
make_httpx_status_error,
)
@ -28,6 +30,10 @@ def _real_user_api_key_dict(parent_span):
)
async def _noop_alert(*args, **kwargs):
return None
async def _drive_admin_failure(*, otel, exception, parent_span, route):
payload = ManagementEndpointLoggingPayload(
route=route,
@ -180,3 +186,173 @@ def test_admin_endpoint_failure_stamps_server_span(
expected_url_path=path,
where=f"{path} {expected_status}",
)
def test_management_wrapper_success_ends_server_span_without_http_request(
server_span_factory, otel_with_exporter, monkeypatch
):
"""Regression: management endpoints whose handler does not declare an
``http_request`` parameter (``/key/generate``, ``/user/new``, ``/mcp/*``,
...) must still get their parent SERVER span stamped + ended on success.
The success hook itself stamps 200 and ``end()``s the parent, but the
wrapper only invoked it when ``http_request`` was present — so on success
the span (created in auth) was never ended and never exported. This drives
the real wrapper around an ``http_request``-less handler and asserts the
SERVER span reaches the exporter with status 200.
"""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def fake_generate_key_fn(data=None, user_api_key_dict=None):
# No ``http_request`` parameter — mirrors generate_key_fn et al.
return {"key": "sk-xyz", "key_name": "k"}
asyncio.run(
fake_generate_key_fn(
data={},
user_api_key_dict=_real_user_api_key_dict(server_span),
)
)
assert_server_span_attrs(
exporter,
expected_status=200,
expected_url_path=KEY_GENERATE_PATH,
where="management wrapper success without http_request",
)
def test_management_wrapper_failure_ends_server_span(
server_span_factory, otel_with_exporter, monkeypatch
):
"""When the handler raises, the wrapper must route through the failure hook
and stamp + end the parent SERVER span with the error status — even for an
``http_request``-less handler (route falls back to ``func.__name__``)."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def failing_fn(data=None, user_api_key_dict=None):
raise HttpStatusException(500, "boom")
with pytest.raises(HttpStatusException):
asyncio.run(
failing_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))
)
assert_server_span_attrs(
exporter,
expected_status=500,
expected_url_path=KEY_GENERATE_PATH,
where="management wrapper failure",
)
def test_management_wrapper_success_with_http_request(
server_span_factory, otel_with_exporter, monkeypatch
):
"""Cover the branch where the handler DOES declare ``http_request``: the
route comes from ``http_request.url.path`` and the body is read from it."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
async def _fake_body(request=None):
return {"team_alias": "t"}
monkeypatch.setattr(mgmt_utils, "_read_request_body", _fake_body)
server_span = server_span_factory("/team/new")
http_request = MagicMock()
http_request.url.path = "/team/new"
@mgmt_utils.management_endpoint_wrapper
async def fake_new_team(data=None, http_request=None, user_api_key_dict=None):
return {"team_id": "t-1"}
asyncio.run(
fake_new_team(
data={},
http_request=http_request,
user_api_key_dict=_real_user_api_key_dict(server_span),
)
)
assert_server_span_attrs(
exporter,
expected_status=200,
expected_url_path="/team/new",
where="management wrapper success with http_request",
)
def test_management_wrapper_noop_when_otel_logger_absent(
server_span_factory, otel_with_exporter, monkeypatch
):
"""When no OTEL logger is registered, the helper early-returns and no SERVER
span is exported — and the handler result is still returned unchanged."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
_otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", None, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def fake_fn(data=None, user_api_key_dict=None):
return {"ok": True}
result = asyncio.run(
fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))
)
assert result == {"ok": True}
assert get_server_span(exporter) is None
def test_management_wrapper_swallows_post_success_errors(
server_span_factory, otel_with_exporter, monkeypatch
):
"""A failure in post-success bookkeeping (cache invalidation, alerting) must
not propagate — the handler result is returned regardless (non-blocking)."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, _exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
def _boom(*args, **kwargs):
raise RuntimeError("cache backend down")
monkeypatch.setattr(mgmt_utils, "_delete_api_key_from_cache", _boom)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def fake_fn(data=None, user_api_key_dict=None):
return {"ok": True}
result = asyncio.run(
fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))
)
assert result == {"ok": True}

View file

@ -1418,3 +1418,123 @@ def test_image_count_prevents_text_tokens_fallback():
f"got {prompt_cost}. text_tokens fallback may be double-charging."
)
assert completion_cost == 0.0
# ---------------------------------------------------------------------------
# Data-residency (OpenAI regional processing) tests
# ---------------------------------------------------------------------------
@pytest.fixture
def _local_model_cost_map():
prev_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
prev_model_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
yield
finally:
litellm.model_cost = prev_model_cost
if prev_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = prev_env
@pytest.mark.parametrize("data_residency", ["eu", "us"])
def test_data_residency_applies_uplift(data_residency, _local_model_cost_map):
"""gpt-5 should apply the regional processing uplift multiplier when
data_residency is set."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(
model="gpt-5",
usage=usage,
custom_llm_provider="openai",
)
regional = generic_cost_per_token(
model="gpt-5",
usage=usage,
custom_llm_provider="openai",
data_residency=data_residency,
)
base_total = base[0] + base[1]
regional_total = regional[0] + regional[1]
assert base_total > 0
assert regional_total == pytest.approx(base_total * 1.10, rel=1e-9)
assert regional[0] == pytest.approx(base[0] * 1.10, rel=1e-9)
assert regional[1] == pytest.approx(base[1] * 1.10, rel=1e-9)
def test_data_residency_no_uplift_for_unmarked_model(_local_model_cost_map):
"""A model without a regional_processing_uplift_multiplier_* entry should
fall back to base pricing, not error."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(
model="gpt-3.5-turbo",
usage=usage,
custom_llm_provider="openai",
)
with_residency = generic_cost_per_token(
model="gpt-3.5-turbo",
usage=usage,
custom_llm_provider="openai",
data_residency="eu",
)
assert base == with_residency
def test_data_residency_none_no_uplift(_local_model_cost_map):
"""data_residency=None should be a no-op even for models with a multiplier."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(
model="gpt-5",
usage=usage,
custom_llm_provider="openai",
)
explicit_none = generic_cost_per_token(
model="gpt-5",
usage=usage,
custom_llm_provider="openai",
data_residency=None,
)
assert base == explicit_none
def test_data_residency_composes_with_service_tier(_local_model_cost_map):
"""The uplift multiplies the priority-tier cost, not the standard one."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
priority_base = generic_cost_per_token(
model="gpt-5",
usage=usage,
custom_llm_provider="openai",
service_tier="priority",
)
priority_eu = generic_cost_per_token(
model="gpt-5",
usage=usage,
custom_llm_provider="openai",
service_tier="priority",
data_residency="eu",
)
priority_base_total = priority_base[0] + priority_base[1]
priority_eu_total = priority_eu[0] + priority_eu[1]
assert priority_base_total > 0
assert priority_eu_total == pytest.approx(priority_base_total * 1.10, rel=1e-9)

View file

@ -125,3 +125,40 @@ class TestGetLitellmParamsExplicitFields:
def test_no_log_from_explicit_param(self):
result = get_litellm_params(no_log=True)
assert result["no-log"] is True
class TestGetLitellmParamsDataResidency:
"""Verify that data_residency is inferred from OpenAI regional api_base."""
def test_eu_host_resolves_to_eu(self):
result = get_litellm_params(
custom_llm_provider="openai",
api_base="https://eu.api.openai.com/v1",
)
assert result["data_residency"] == "eu"
def test_us_host_resolves_to_us(self):
result = get_litellm_params(
custom_llm_provider="openai",
api_base="https://us.api.openai.com/v1",
)
assert result["data_residency"] == "us"
def test_global_host_resolves_to_none(self):
result = get_litellm_params(
custom_llm_provider="openai",
api_base="https://api.openai.com/v1",
)
assert result["data_residency"] is None
def test_no_api_base_is_none(self):
result = get_litellm_params(custom_llm_provider="openai")
assert result["data_residency"] is None
def test_non_openai_provider_does_not_resolve(self):
"""Regional OpenAI host doesn't apply to other providers."""
result = get_litellm_params(
custom_llm_provider="anthropic",
api_base="https://eu.api.openai.com/v1",
)
assert result["data_residency"] is None

View file

@ -1646,6 +1646,336 @@ def test_azure_v1_api_uses_openai_client(api_version):
), f"base_url should contain /openai/v1/, got {async_client.base_url}"
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_v1_api_with_azure_ad_token_provider(api_version):
"""
The v1 OpenAI client path must forward `azure_ad_token_provider` so Azure AD
auth works for `api_version` in {"v1", "latest", "preview"}.
Regression: https://github.com/BerriAI/litellm/issues/27945 — before the fix
the v1 branch only forwarded `api_key`, so AD-only configs raised
"The api_key client option must be set" on every request.
The OpenAI SDK accepts a callable for `api_key` and re-invokes it on every
request, so passing the provider directly preserves token refresh.
"""
from openai import AsyncOpenAI, OpenAI
base_llm = BaseAzureLLM()
api_base = "https://test.openai.azure.com"
token_value = "mock-azure-ad-token-from-provider"
def token_provider():
return token_value
init_return = {
"api_key": None,
"azure_endpoint": api_base,
"api_version": api_version,
"azure_ad_token": None,
"azure_ad_token_provider": token_provider,
}
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = init_return
client = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
_is_async=False,
)
assert isinstance(client, OpenAI)
# The SDK stores callables as `_api_key_provider` and refreshes
# `self.api_key` before each request.
assert client._api_key_provider is token_provider
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = init_return
async_client = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
_is_async=True,
)
assert isinstance(async_client, AsyncOpenAI)
# Async client requires an async provider; we wrap the sync provider
# so the SDK can `await` it.
assert async_client._api_key_provider is not None
assert async_client._api_key_provider is not token_provider
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_v1_api_async_token_provider_resolves_to_current_token(api_version):
"""
The async wrapper must call the underlying sync provider on each invocation
(not cache its first return value), so token rotation is honored.
"""
import asyncio
from openai import AsyncOpenAI
base_llm = BaseAzureLLM()
api_base = "https://test.openai.azure.com"
tokens = iter(["token-1", "token-2", "token-3"])
def rotating_provider():
return next(tokens)
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = {
"api_key": None,
"azure_endpoint": api_base,
"api_version": api_version,
"azure_ad_token": None,
"azure_ad_token_provider": rotating_provider,
}
async_client = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
_is_async=True,
)
assert isinstance(async_client, AsyncOpenAI)
loop = asyncio.new_event_loop()
try:
first = loop.run_until_complete(async_client._api_key_provider())
second = loop.run_until_complete(async_client._api_key_provider())
finally:
loop.close()
assert first == "token-1"
assert second == "token-2"
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_v1_api_with_static_azure_ad_token(api_version):
"""
When only `azure_ad_token` (a static string) is set, the v1 client should
receive it as `api_key`.
"""
from openai import OpenAI
base_llm = BaseAzureLLM()
api_base = "https://test.openai.azure.com"
token_value = "static-azure-ad-token"
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = {
"api_key": None,
"azure_endpoint": api_base,
"api_version": api_version,
"azure_ad_token": token_value,
"azure_ad_token_provider": None,
}
client = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
_is_async=False,
)
assert isinstance(client, OpenAI)
assert client.api_key == token_value
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_v1_api_key_wins_over_ad_token(api_version):
"""
Explicit `api_key` takes precedence over `azure_ad_token_provider` /
`azure_ad_token`, matching the priority documented in
`initialize_azure_sdk_client`.
"""
from openai import OpenAI
base_llm = BaseAzureLLM()
api_base = "https://test.openai.azure.com"
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = {
"api_key": "explicit-key",
"azure_endpoint": api_base,
"api_version": api_version,
"azure_ad_token": "should-be-ignored",
"azure_ad_token_provider": lambda: "also-ignored",
}
client = base_llm.get_azure_openai_client(
api_key="explicit-key",
api_base=api_base,
api_version=api_version,
_is_async=False,
)
assert isinstance(client, OpenAI)
assert client.api_key == "explicit-key"
assert client._api_key_provider is None
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_v1_client_cache_separates_distinct_ad_providers(api_version):
"""
Two configs sharing api_base/api_version but with different AD token
providers must not share a cached OpenAI client, otherwise requests for
one config would be sent with another config's AD credentials.
"""
from openai import AsyncOpenAI
litellm.in_memory_llm_clients_cache._cache = {}
base_llm = BaseAzureLLM()
api_base = "https://test.openai.azure.com"
def provider_a():
return "token-a"
def provider_b():
return "token-b"
def _init_for(provider):
return {
"api_key": None,
"azure_endpoint": api_base,
"api_version": api_version,
"azure_ad_token": None,
"azure_ad_token_provider": provider,
}
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = _init_for(provider_a)
client_a = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
litellm_params={"azure_ad_token_provider": provider_a},
_is_async=True,
)
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = _init_for(provider_b)
client_b = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
litellm_params={"azure_ad_token_provider": provider_b},
_is_async=True,
)
assert isinstance(client_a, AsyncOpenAI)
assert isinstance(client_b, AsyncOpenAI)
assert client_a is not client_b
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_v1_client_cache_separates_distinct_entra_credentials(api_version):
"""
Configs that synthesize an AD provider from tenant_id/client_id/client_secret
must not share a cached client when those inputs differ.
"""
from openai import AsyncOpenAI
litellm.in_memory_llm_clients_cache._cache = {}
base_llm = BaseAzureLLM()
api_base = "https://test.openai.azure.com"
def synth_provider():
return "synthesized-token"
def _init_synth():
return {
"api_key": None,
"azure_endpoint": api_base,
"api_version": api_version,
"azure_ad_token": None,
"azure_ad_token_provider": synth_provider,
}
common = {
"api_key": None,
"api_base": api_base,
"api_version": api_version,
"_is_async": True,
}
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = _init_synth()
client_a = base_llm.get_azure_openai_client(
litellm_params={
"tenant_id": "tenant-a",
"client_id": "client-a",
"client_secret": "secret-a",
},
**common,
)
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = _init_synth()
client_b = base_llm.get_azure_openai_client(
litellm_params={
"tenant_id": "tenant-b",
"client_id": "client-b",
"client_secret": "secret-b",
},
**common,
)
assert isinstance(client_a, AsyncOpenAI)
assert isinstance(client_b, AsyncOpenAI)
assert client_a is not client_b
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_v1_client_cache_reuses_for_identical_ad_config(api_version):
"""
Identical AD configs should still share a cached client (regression guard
so the cache-key change doesn't accidentally disable caching).
"""
from openai import AsyncOpenAI
litellm.in_memory_llm_clients_cache._cache = {}
base_llm = BaseAzureLLM()
api_base = "https://test.openai.azure.com"
def provider():
return "tok"
init_return = {
"api_key": None,
"azure_endpoint": api_base,
"api_version": api_version,
"azure_ad_token": None,
"azure_ad_token_provider": provider,
}
with patch.object(base_llm, "initialize_azure_sdk_client") as mock_init:
mock_init.return_value = init_return
client_a = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
litellm_params={"azure_ad_token_provider": provider},
_is_async=True,
)
client_b = base_llm.get_azure_openai_client(
api_key=None,
api_base=api_base,
api_version=api_version,
litellm_params={"azure_ad_token_provider": provider},
_is_async=True,
)
assert isinstance(client_a, AsyncOpenAI)
assert client_a is client_b
def test_azure_traditional_api_uses_azure_openai_client():
"""
Test that traditional Azure API versions still use AzureOpenAI client.

View file

@ -0,0 +1,134 @@
"""
Tests for `litellm.llms.base_llm.managed_resources.utils.extract_model_id_from_unified_id`.
The regex inside this helper is shared by both the vector-store unified-ID
format (`...;model_id,<value>;...`) and the file-ID format (`...;llm_output_file_model_id,<uuid>`).
A naive regex (`r"model_id,([^;]+)"`) substring-matches the latter and
returns the deployment UUID, which then gets fed as a model candidate
into the team-access check and 403s every team-BYOK file attach
(LIT-3244 patch/1.86.0 second-order finding). These tests pin the
field-boundary anchor that prevents that.
"""
import pytest
from litellm.llms.base_llm.managed_resources.utils import (
encode_unified_id,
extract_model_id_from_unified_id,
)
# ---------------------------------------------------------------------------
# Vector-store unified-ID shape — has a top-level `model_id,<value>` field.
# Existing behavior must be preserved: returns the value.
# ---------------------------------------------------------------------------
def test_extract_model_id_returns_value_for_vector_store_unified_id():
unified_id = (
"litellm_proxy:vector_store"
";unified_id,abc-123"
";target_model_names,gpt-4,gemini"
";resource_id,vs_xyz"
";model_id,deployment-uuid-456"
)
assert extract_model_id_from_unified_id(unified_id) == "deployment-uuid-456"
def test_extract_model_id_returns_value_when_field_is_first():
"""`model_id` is the very first field after the prefix (anchor must accept start-of-string)."""
unified_id = "litellm_proxy:vector_store;model_id,first-field-value;unified_id,abc"
# First field after the prefix is preceded by `;`, so it matches via the
# `;model_id,` branch. Pin that the anchor isn't accidentally too strict.
assert extract_model_id_from_unified_id(unified_id) == "first-field-value"
# ---------------------------------------------------------------------------
# File-ID shape — has `llm_output_file_model_id,<uuid>` but no top-level
# `model_id,` field. Must return None (the previous regex would have
# substring-matched and returned the deployment UUID).
# ---------------------------------------------------------------------------
def test_extract_model_id_returns_none_for_file_id_without_model_id_field():
"""Regression pin for LIT-3244 patch/1.86.0.
File-IDs constructed via `LITELLM_MANAGED_FILE_COMPLETE_STR` have
`llm_output_file_model_id,<deployment_uuid>` but no top-level
`model_id,` field. The previous regex matched the substring and
returned the UUID, which then 403'd team-BYOK file attaches with
`Tried to access <uuid>`.
"""
file_id = (
"litellm_proxy:text/plain"
";unified_id,file-uuid-123"
";target_model_names,openai/gpt-4o"
";llm_output_file_id,file-OpenAIReturnedId"
";llm_output_file_model_id,813bf25f-e5a7-4658-8253-a6f677be8eb5"
)
assert extract_model_id_from_unified_id(file_id) is None, (
"File-ID has no top-level `model_id,` field — the deployment UUID "
"in `llm_output_file_model_id,` must NOT be returned. Returning it "
"feeds the UUID as a model candidate into the team-access check "
"and 403s every team-BYOK file attach (LIT-3244 patch/1.86.0)."
)
def test_extract_model_id_returns_none_for_file_id_with_model_id_value_null():
"""The current file-ID builder writes `llm_output_file_model_id,None`
(the Python `None` stringified) when the upstream model_id isn't known.
Still no top-level `model_id,` field → must return None.
"""
file_id = (
"litellm_proxy:text/plain"
";unified_id,uuid"
";target_model_names,openai/gpt-4o"
";llm_output_file_id,file-Y"
";llm_output_file_model_id,None"
)
assert extract_model_id_from_unified_id(file_id) is None
# ---------------------------------------------------------------------------
# Base64-encoded inputs must decode and apply the same anchor.
# ---------------------------------------------------------------------------
def test_extract_model_id_decodes_base64_then_anchors():
file_id_plain = (
"litellm_proxy:text/plain"
";unified_id,uuid"
";target_model_names,openai/gpt-4o"
";llm_output_file_id,file-Y"
";llm_output_file_model_id,813bf25f-e5a7-4658-8253-a6f677be8eb5"
)
encoded = encode_unified_id(file_id_plain)
assert extract_model_id_from_unified_id(encoded) is None
vector_store_plain = (
"litellm_proxy:vector_store"
";unified_id,abc"
";target_model_names,gpt-4"
";resource_id,vs_xyz"
";model_id,real-model-id"
)
encoded_vs = encode_unified_id(vector_store_plain)
assert extract_model_id_from_unified_id(encoded_vs) == "real-model-id"
# ---------------------------------------------------------------------------
# Defensive: malformed / non-string inputs must not raise.
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("bad_input", [None, 42, b"bytes-not-str", []])
def test_extract_model_id_returns_none_for_non_string_input(bad_input):
assert extract_model_id_from_unified_id(bad_input) is None # type: ignore[arg-type]
def test_extract_model_id_returns_none_when_field_absent():
assert (
extract_model_id_from_unified_id(
"litellm_proxy:other;unified_id,abc;some_field,whatever"
)
is None
)

View file

@ -76,7 +76,7 @@ class TestAgentCoreAcceptHeader:
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_runtime",
messages=[{"role": "user", "content": "test"}],
api_key="test-jwt-token",
client=client,
@ -281,7 +281,7 @@ class TestAgentCoreStreamingJsonFallback:
with patch.object(client, "post", return_value=mock_response):
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,
@ -318,7 +318,7 @@ class TestAgentCoreStreamingJsonFallback:
client, "post", new_callable=AsyncMock, return_value=mock_response
):
response = await litellm.acompletion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,
@ -353,7 +353,7 @@ class TestAgentCoreStreamingJsonFallback:
Exception, match="Failed to read/parse JSON response body"
):
litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,
@ -383,7 +383,7 @@ class TestAgentCoreStreamingJsonFallback:
Exception, match="Failed to read/parse JSON response body"
):
await litellm.acompletion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,

View file

@ -0,0 +1,134 @@
"""
Tests that data_residency is correctly populated on the litellm logging
object's litellm_params for OpenAI Responses paths, even when
custom_llm_provider is resolved from the model string inside responses()
rather than passed explicitly.
"""
import json
from unittest.mock import MagicMock, patch
import litellm
def _make_responses_api_response_body() -> dict:
return {
"id": "resp-test",
"object": "response",
"created_at": 1234567890,
"model": "gpt-4.1",
"output": [
{
"type": "message",
"id": "msg-test",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "ok",
"annotations": [],
}
],
}
],
"status": "completed",
"usage": {
"input_tokens": 1,
"output_tokens": 1,
"total_tokens": 2,
},
}
def _make_mock_http_client(response_body: dict) -> MagicMock:
mock_client = MagicMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = response_body
mock_response.text = json.dumps(response_body)
mock_client.post.return_value = mock_response
return mock_client
def _capture_logging_obj():
captured = {}
real_init = litellm.Logging.__init__
def init_spy(self, *args, **kwargs):
real_init(self, *args, **kwargs)
captured["logging_obj"] = self
return captured, init_spy
def test_responses_eu_api_base_sets_data_residency():
"""When api_base is a regional OpenAI host and custom_llm_provider is
inferred from the model (not passed explicitly), data_residency must end
up on the logging object's litellm_params so the cost calculator can apply
the regional uplift."""
mock_client = _make_mock_http_client(_make_responses_api_response_body())
captured, init_spy = _capture_logging_obj()
with (
patch(
"litellm.llms.custom_httpx.llm_http_handler._get_httpx_client",
return_value=mock_client,
),
patch.object(litellm.Logging, "__init__", init_spy),
):
litellm.responses(
model="gpt-4.1",
input="hi",
api_base="https://eu.api.openai.com/v1",
api_key="test-key",
)
logging_obj = captured["logging_obj"]
assert logging_obj.litellm_params.get("data_residency") == "eu"
def test_responses_us_api_base_sets_data_residency():
mock_client = _make_mock_http_client(_make_responses_api_response_body())
captured, init_spy = _capture_logging_obj()
with (
patch(
"litellm.llms.custom_httpx.llm_http_handler._get_httpx_client",
return_value=mock_client,
),
patch.object(litellm.Logging, "__init__", init_spy),
):
litellm.responses(
model="gpt-4.1",
input="hi",
api_base="https://us.api.openai.com/v1",
api_key="test-key",
)
logging_obj = captured["logging_obj"]
assert logging_obj.litellm_params.get("data_residency") == "us"
def test_responses_global_api_base_leaves_data_residency_none():
mock_client = _make_mock_http_client(_make_responses_api_response_body())
captured, init_spy = _capture_logging_obj()
with (
patch(
"litellm.llms.custom_httpx.llm_http_handler._get_httpx_client",
return_value=mock_client,
),
patch.object(litellm.Logging, "__init__", init_spy),
):
litellm.responses(
model="gpt-4.1",
input="hi",
api_base="https://api.openai.com/v1",
api_key="test-key",
)
logging_obj = captured["logging_obj"]
assert logging_obj.litellm_params.get("data_residency") is None

View file

@ -0,0 +1,34 @@
"""Tests for the OpenAI data-residency inference helper."""
import pytest
from litellm.llms.openai.data_residency import infer_openai_data_residency
@pytest.mark.parametrize(
"api_base, expected",
[
("https://eu.api.openai.com/v1", "eu"),
("https://eu.api.openai.com", "eu"),
("https://us.api.openai.com/v1", "us"),
("https://us.api.openai.com", "us"),
("https://EU.api.openai.com/v1", "eu"),
("https://api.openai.com/v1", None),
("https://api.openai.com", None),
("https://example.com/v1", None),
("https://my-azure-endpoint.openai.azure.com/openai/deployments/foo", None),
("", None),
(None, None),
("not a url", None),
],
)
def test_infer_openai_data_residency(api_base, expected):
assert infer_openai_data_residency("openai", api_base) == expected
@pytest.mark.parametrize("custom_llm_provider", [None, "anthropic", "azure", "bedrock"])
def test_infer_openai_data_residency_non_openai_provider(custom_llm_provider):
assert (
infer_openai_data_residency(custom_llm_provider, "https://eu.api.openai.com/v1")
is None
)

View file

@ -3370,3 +3370,102 @@ async def test_resolve_end_user_reraises_budget_exceeded(
prisma_client=MagicMock(),
user_api_key_cache=cache,
)
@pytest.mark.asyncio
async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
"""
Regression pin for LIT-3244 patch/1.86.0 follow-up.
`_cache_team_object` is the canonical "refresh this team" primitive.
Two cache keys are in play:
- "team_id:<id>" — used by `get_team_object(team_id=...)`,
i.e. API-key auth and JWT-with-team_id_jwt_field
- "team_alias:<alias>" — used by `get_team_object_by_alias(team_alias=...)`,
i.e. JWT-with-team_alias_jwt_field
Invariants this test pins:
1. Writes the team_id-keyed entry with the refreshed object (team_id
is the table PK — guaranteed unique, safe to write).
2. DELETES (does NOT write) the team_alias-keyed entry. `team_alias`
has no UNIQUE constraint in schema.prisma, so writing it from
this generic refresh path would let a team admin who renames
their team to collide with another team's alias silently
overwrite the cached team for JWT-by-alias auth (veria-ai
review on #28739). Deleting forces the next JWT-by-alias
reader through `get_team_object_by_alias`, which enforces
len(teams)==1 before populating the cache.
3. When team_alias is None, NO alias-key operation happens (no
delete of an empty-keyed entry, no spurious write).
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
from litellm.proxy.auth.auth_checks import _cache_team_object
base_team_row = {
"team_id": "team-1234",
"team_alias": "H-Capacity",
"models": ["openai/*", "bedrock-claude-sonnet-4"],
}
# ===== team_alias is set =====
team_table = LiteLLM_TeamTableCachedObj(**base_team_row)
cache = MagicMock()
cache.async_set_cache = AsyncMock()
cache.delete_cache = MagicMock()
logging_obj = MagicMock()
logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
await _cache_team_object(
team_id="team-1234",
team_table=team_table,
user_api_key_cache=cache,
proxy_logging_obj=logging_obj,
)
# (1) team_id-keyed write fires with the refreshed object
written_keys = [
(c.kwargs.get("key") or c.args[0])
for c in cache.async_set_cache.await_args_list
]
assert written_keys == ["team_id:team-1234"], (
"Only the team_id-keyed write should fire; the alias key must be "
"deleted, NOT written. "
f"Got writes: {written_keys}"
)
written_value = (
cache.async_set_cache.await_args.kwargs.get("value")
or cache.async_set_cache.await_args.args[1]
)
assert written_value is team_table
# (2) team_alias-keyed entry is deleted in BOTH the in-memory cache
# and the Redis dual cache (mirrors _delete_cache_key_object pattern).
cache.delete_cache.assert_called_once_with(key="team_alias:H-Capacity")
logging_obj.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with(
key="team_alias:H-Capacity"
)
# ===== team_alias is None: no alias-key operation =====
aliasless = LiteLLM_TeamTableCachedObj(**{**base_team_row, "team_alias": None})
cache2 = MagicMock()
cache2.async_set_cache = AsyncMock()
cache2.delete_cache = MagicMock()
logging_obj2 = MagicMock()
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
await _cache_team_object(
team_id="team-no-alias",
team_table=aliasless,
user_api_key_cache=cache2,
proxy_logging_obj=logging_obj2,
)
cache2.delete_cache.assert_not_called()
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_not_awaited()
written_keys_aliasless = [
(c.kwargs.get("key") or c.args[0])
for c in cache2.async_set_cache.await_args_list
]
assert written_keys_aliasless == ["team_id:team-no-alias"]

View file

@ -262,12 +262,125 @@ def test_virtual_key_mcp_routes_allows_v1_mcp_server_subpaths(route):
)
def test_mcp_management_routes_classified_as_management_not_llm_api(route):
"""MCP server CRUD must be management routes, not llm_api routes, so
DISABLE_LLM_API_ENDPOINTS on admin nodes does not block the Admin UI."""
DISABLE_LLM_API_ENDPOINTS on admin nodes does not block the Admin UI.
Note: virtual keys with allowed_routes=["llm_api_routes"] can still call
*GET* `/v1/mcp/server` and *GET* `/v1/mcp/server/{server_id}` — that
carve-out is enforced method-aware inside
`is_virtual_key_allowed_to_call_route`, not by adding the paths to
`llm_api_routes`. So `is_llm_api_route()` still returns False here and
`DISABLE_LLM_API_ENDPOINTS` still does not block these paths.
"""
assert RouteChecks.is_llm_api_route(route=route) is False
assert RouteChecks.is_management_route(route=route) is True
def _mock_request(method: str) -> Request:
request = MagicMock(spec=Request)
request.method = method
return request
@pytest.mark.parametrize(
"route",
[
"/v1/mcp/server",
"/v1/mcp/server/abc-123",
],
)
def test_virtual_key_llm_api_routes_allows_get_mcp_server_discovery(route):
"""
Regression test: virtual keys with allowed_routes=["llm_api_routes"] must
be able to list/inspect MCP servers via GET /v1/mcp/server[/{server_id}].
The handlers strip credential-bearing fields via
`_sanitize_mcp_server_list_for_virtual_key` when the caller is a
restricted virtual key, so GET is safe to expose. The carve-out is
method-aware (see below) — non-GET requests to the same paths are
rejected at this layer, so admin-only writes remain gated.
"""
valid_token = UserAPIKeyAuth(
user_id="test_user",
allowed_routes=["llm_api_routes"],
)
result = RouteChecks.is_virtual_key_allowed_to_call_route(
route=route,
valid_token=valid_token,
request=_mock_request("GET"),
)
assert result is True
@pytest.mark.parametrize(
"route",
[
"/v1/mcp/server",
"/v1/mcp/server/abc-123",
],
)
@pytest.mark.parametrize("method", ["POST", "PUT", "PATCH", "DELETE"])
def test_virtual_key_llm_api_routes_rejects_non_get_mcp_server_discovery(route, method):
"""Method-aware: the MCP server discovery carve-out is GET-only.
POST/PUT/PATCH/DELETE on `/v1/mcp/server[/{server_id}]` are admin-only
management writes and must not be reachable via llm_api_routes.
"""
valid_token = UserAPIKeyAuth(
user_id="test_user",
allowed_routes=["llm_api_routes"],
)
with pytest.raises(HTTPException) as exc_info:
RouteChecks.is_virtual_key_allowed_to_call_route(
route=route,
valid_token=valid_token,
request=_mock_request(method),
)
assert exc_info.value.status_code == 403
@pytest.mark.parametrize(
"route",
[
# Multi-segment admin-only sub-paths must NOT be reachable via
# llm_api_routes, even on GET.
"/v1/mcp/server/abc-123/approve",
"/v1/mcp/server/abc-123/reject",
"/v1/mcp/server/oauth/session",
"/v1/mcp/server/abc-123/user-credential",
],
)
def test_virtual_key_llm_api_routes_rejects_mcp_multi_segment_admin_subpaths(
route,
):
"""Multi-segment admin-only MCP sub-paths are not reachable via llm_api_routes.
The discovery carve-out only matches `/v1/mcp/server` and
`/v1/mcp/server/{server_id}` (single segment after `/server/`), so any
path with additional segments is rejected even when the request is GET.
"""
valid_token = UserAPIKeyAuth(
user_id="test_user",
allowed_routes=["llm_api_routes"],
)
with pytest.raises(HTTPException) as exc_info:
RouteChecks.is_virtual_key_allowed_to_call_route(
route=route,
valid_token=valid_token,
request=_mock_request("GET"),
)
assert exc_info.value.status_code == 403
def test_spend_logs_v2_classified_as_management_not_llm_api():
"""Paginated spend logs are a management/spend read route, not an LLM API."""

View file

@ -0,0 +1 @@
line:0.0 branch:0.0

View file

@ -0,0 +1,201 @@
#!/usr/bin/env python3
"""Coverage gate for the proxy_server.py behavior-pinning project.
Reads a coverage XML report (produced by ``pytest --cov-branch
--cov-report=xml:<path>``) and asserts that line + branch coverage on
``litellm/proxy/proxy_server.py`` meets the per-PR target.
Target selection:
--pr-target {1|2|3} explicit target
(none) self-selected by inspecting which placeholder
test files have been filled (PR1 fills before
PR2, PR2 before PR3). With nothing filled, the
target is "PR0" (baseline, no minimum).
Exits 0 on PASS, non-zero on FAIL.
"""
from __future__ import annotations
import argparse
import ast
import sys
import xml.etree.ElementTree as ET
from pathlib import Path
from typing import Dict, List, Tuple
HERE = Path(__file__).resolve().parent
SOURCE_FILE = "litellm/proxy/proxy_server.py"
# PR target gates: (line%, branch%)
TARGETS: Dict[str, Tuple[float, float]] = {
"PR0": (0.0, 0.0),
"PR1": (25.0, 18.0),
"PR2": (50.0, 38.0),
"PR3": (70.0, 55.0),
}
# Which placeholder files each PR is expected to fill (see Notion plan).
PR1_FILES: List[str] = [
"test_lifecycle.py",
"test_proxy_config.py",
"test_spend_counters.py",
"test_background_health.py",
"test_openapi_customization.py",
"test_exception_handlers.py",
"test_streaming_helpers.py",
]
PR2_FILES: List[str] = [
"test_routes_models.py",
"test_routes_chat_completions.py",
"test_routes_completions.py",
"test_routes_embeddings.py",
"test_routes_moderations.py",
"test_routes_audio.py",
"test_routes_assistants.py",
"test_routes_threads.py",
"test_routes_utils.py",
"test_routes_model_info.py",
"test_routes_model_metrics.py",
"test_routes_queue.py",
]
PR3_FILES: List[str] = [
"test_routes_login_sso.py",
"test_routes_onboarding.py",
"test_routes_invitation.py",
"test_routes_config.py",
"test_routes_model_cost_map.py",
"test_routes_anthropic_beta.py",
"test_routes_misc.py",
]
def file_has_tests(path: Path) -> bool:
"""A test file is considered filled if it defines at least one ``test_*``."""
if not path.is_file():
return False
try:
tree = ast.parse(path.read_text())
except SyntaxError:
return False
for node in ast.walk(tree):
if isinstance(
node, (ast.FunctionDef, ast.AsyncFunctionDef)
) and node.name.startswith("test_"):
return True
return False
def detect_pr_target(dir_path: Path) -> str:
"""Pick the strictest PR whose files are fully filled in this directory."""
pr3_filled = all(file_has_tests(dir_path / f) for f in PR3_FILES)
pr2_filled = all(file_has_tests(dir_path / f) for f in PR2_FILES)
pr1_filled = all(file_has_tests(dir_path / f) for f in PR1_FILES)
if pr3_filled and pr2_filled and pr1_filled:
return "PR3"
if pr2_filled and pr1_filled:
return "PR2"
if pr1_filled:
return "PR1"
return "PR0"
def parse_coverage_xml(xml_path: Path) -> Tuple[float, float]:
"""Extract (line%, branch%) for proxy_server.py from a coverage XML report.
Returns (0.0, 0.0) if the file isn't found in the report.
"""
if not xml_path.is_file():
raise FileNotFoundError(f"Coverage XML not found at {xml_path}")
tree = ET.parse(xml_path)
root = tree.getroot()
for class_elem in root.iter("class"):
filename = class_elem.get("filename", "")
# Coverage tools emit either a repo-relative path or just the basename
# depending on configuration. Match by suffix.
if filename.endswith("proxy/proxy_server.py") or filename.endswith(
"proxy_server.py"
):
line_rate = float(class_elem.get("line-rate", "0"))
branch_rate = float(class_elem.get("branch-rate", "0"))
return line_rate * 100.0, branch_rate * 100.0
return 0.0, 0.0
def parse_baseline(baseline_path: Path) -> Tuple[float, float]:
"""Parse ``line:<float> branch:<float>`` baseline; missing file -> (0, 0)."""
if not baseline_path.is_file():
return 0.0, 0.0
line_pct = 0.0
branch_pct = 0.0
for token in baseline_path.read_text().split():
if ":" not in token:
continue
key, _, value = token.partition(":")
try:
num = float(value)
except ValueError:
continue
if key == "line":
line_pct = num
elif key == "branch":
branch_pct = num
return line_pct, branch_pct
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--pr-target",
choices=["1", "2", "3"],
default=None,
help="Explicit PR target (1, 2, or 3). If omitted, self-selected.",
)
parser.add_argument(
"--coverage-xml",
default=str(HERE.parent.parent.parent.parent / ".cov_new.xml"),
help="Path to coverage XML (default: <repo>/.cov_new.xml)",
)
args = parser.parse_args()
if args.pr_target:
target = f"PR{args.pr_target}"
else:
target = detect_pr_target(HERE)
target_line, target_branch = TARGETS[target]
# The effective floor is the max of the PR target and the committed
# baseline. The baseline is updated as each PR lands so a future
# regression (e.g. a test deletion) trips this gate even if the
# static PR target is already met.
baseline_line, baseline_branch = parse_baseline(HERE / ".coverage_baseline")
line_min = max(target_line, baseline_line)
branch_min = max(target_branch, baseline_branch)
xml_path = Path(args.coverage_xml)
try:
line_pct, branch_pct = parse_coverage_xml(xml_path)
except FileNotFoundError as exc:
print(f"FAIL: {exc}", file=sys.stderr)
return 2
line_ok = line_pct >= line_min
branch_ok = branch_pct >= branch_min
status = "PASS" if (line_ok and branch_ok) else "FAIL"
print(
f"target={target} baseline=(line:{baseline_line:.2f} branch:{baseline_branch:.2f})"
)
print(
f"line: {line_pct:6.2f}% / {line_min:6.2f}% " f"{'OK' if line_ok else 'MISS'}"
)
print(
f"branch: {branch_pct:6.2f}% / {branch_min:6.2f}% "
f"{'OK' if branch_ok else 'MISS'}"
)
print(status)
return 0 if status == "PASS" else 1
if __name__ == "__main__":
sys.exit(main())

View file

@ -0,0 +1,249 @@
#!/usr/bin/env python3
"""Pin-list gate for the proxy_server.py behavior-pinning project.
For each identifier in a pin list, asserts that the test directory contains:
1. At least one happy-path test that references the identifier and uses
a real assertion (normalize(response.json()) == {...}, .model_validate,
or a dict-equality with >= 3 keys).
2. At least one error-path test (name hints at error OR asserts a 4xx/5xx
status OR uses pytest.raises).
3. No test that is "status-only" (its sole assert is on response.status_code).
``test_harness_smoke.py`` is ignored (harness self-tests don't count toward
behavior pinning).
Exits 0 on PASS, non-zero on FAIL.
"""
from __future__ import annotations
import argparse
import ast
import re
import sys
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, List, Optional, Set, Tuple
HERE = Path(__file__).resolve().parent
PIN_LINE_RE = re.compile(r"^- `([^`]+)`\s*$")
ERROR_NAME_HINTS = (
"error",
"fail",
"invalid",
"unauthorized",
"forbidden",
"missing",
"denied",
"rejected",
"bad",
"raises",
"exception",
"404",
"401",
"403",
"422",
"500",
)
ERROR_STATUS_CODES = frozenset({400, 401, 402, 403, 404, 405, 409, 422, 500, 502, 503})
@dataclass
class TestFunction:
name: str
file: Path
source: str
asserts: List[ast.Assert] = field(default_factory=list)
raises_calls: int = 0
status_code_asserts: List[int] = field(default_factory=list)
has_strong_assertion: bool = (
False # normalize() or .model_validate() or large dict-eq
)
def parse_pin_list(path: Path) -> List[str]:
items: List[str] = []
for line in path.read_text().splitlines():
m = PIN_LINE_RE.match(line)
if m:
items.append(m.group(1).strip())
return items
def _has_strong_assertion(node: ast.AST) -> bool:
"""True if an assert subtree contains normalize(), .model_validate(), or dict-eq with >=3 keys."""
for sub in ast.walk(node):
if isinstance(sub, ast.Call):
func = sub.func
if isinstance(func, ast.Name) and func.id == "normalize":
return True
if isinstance(func, ast.Attribute) and func.attr == "model_validate":
return True
if (
isinstance(sub, ast.Compare)
and len(sub.ops) == 1
and isinstance(sub.ops[0], ast.Eq)
):
# response.json() == {<dict literal with >= 3 keys>}
rhs = sub.comparators[0]
if isinstance(rhs, ast.Dict) and len(rhs.keys) >= 3:
return True
return False
def _extract_status_code(node: ast.Assert) -> Optional[int]:
"""If this assert is exactly ``X.status_code == <int>``, return the int."""
test = node.test
if not isinstance(test, ast.Compare):
return None
if len(test.ops) != 1 or not isinstance(test.ops[0], ast.Eq):
return None
left = test.left
if not (isinstance(left, ast.Attribute) and left.attr == "status_code"):
return None
right = test.comparators[0]
if isinstance(right, ast.Constant) and isinstance(right.value, int):
return right.value
return None
def collect_test_functions(test_dir: Path) -> List[TestFunction]:
funcs: List[TestFunction] = []
for path in sorted(test_dir.glob("test_*.py")):
# Skip the harness's own smoke tests — they don't count toward
# behavior pinning.
if path.name == "test_harness_smoke.py":
continue
source = path.read_text()
try:
tree = ast.parse(source)
except SyntaxError:
continue
for node in ast.walk(tree):
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
if not node.name.startswith("test_"):
continue
tf = TestFunction(name=node.name, file=path, source=source)
for sub in ast.walk(node):
if isinstance(sub, ast.Assert):
tf.asserts.append(sub)
sc = _extract_status_code(sub)
if sc is not None:
tf.status_code_asserts.append(sc)
if _has_strong_assertion(sub):
tf.has_strong_assertion = True
if isinstance(sub, ast.With):
for item in sub.items:
ctx = item.context_expr
if isinstance(ctx, ast.Call) and isinstance(
ctx.func, ast.Attribute
):
if ctx.func.attr == "raises":
tf.raises_calls += 1
funcs.append(tf)
return funcs
def _is_status_only(tf: TestFunction) -> bool:
"""A test that has >=1 status_code assert and ALL its asserts are status_code."""
return len(tf.asserts) >= 1 and len(tf.status_code_asserts) == len(tf.asserts)
def _looks_like_error_test(tf: TestFunction) -> bool:
name_lower = tf.name.lower()
if any(hint in name_lower for hint in ERROR_NAME_HINTS):
return True
if tf.raises_calls > 0:
return True
if any(sc in ERROR_STATUS_CODES for sc in tf.status_code_asserts):
return True
return False
def _references_pin(tf: TestFunction, pin: str) -> bool:
"""Cheap string-contains check against the test function's source.
This is intentionally permissive — if the pin identifier (e.g.
``update_cache`` or ``POST /chat/completions``) appears anywhere in
the test file we count it. Aliased route paths or parametrize
cases trigger the same reference.
"""
return pin in tf.source
def check(pin_list: List[str], funcs: List[TestFunction]) -> Tuple[bool, List[str]]:
failures: List[str] = []
status_only = [tf for tf in funcs if _is_status_only(tf)]
for tf in status_only:
failures.append(
f"status-only test (only asserts response.status_code): "
f"{tf.file.name}::{tf.name}"
)
by_pin: Dict[str, List[TestFunction]] = {pin: [] for pin in pin_list}
for tf in funcs:
for pin in pin_list:
if _references_pin(tf, pin):
by_pin[pin].append(tf)
for pin, matches in by_pin.items():
if not matches:
failures.append(f"no tests reference pin: {pin}")
continue
has_happy = any(
tf.has_strong_assertion and not _looks_like_error_test(tf) for tf in matches
)
has_error = any(_looks_like_error_test(tf) for tf in matches)
if not has_happy:
failures.append(
f"no happy-path test with strong assertion (normalize/model_validate/dict-eq>=3) "
f"for pin: {pin}"
)
if not has_error:
failures.append(f"no error-path test for pin: {pin}")
return (not failures), failures
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--list",
required=True,
help="Path to pin list file (markdown bullets in `- ` + backtick + symbol + backtick format)",
)
parser.add_argument(
"--test-dir",
default=str(HERE),
help="Test directory to scan (default: this directory)",
)
args = parser.parse_args()
pin_path = Path(args.list)
if not pin_path.is_file():
print(f"FAIL: pin list not found at {pin_path}", file=sys.stderr)
return 2
pin_list = parse_pin_list(pin_path)
if not pin_list:
print(f"FAIL: pin list at {pin_path} contained zero items", file=sys.stderr)
return 2
test_dir = Path(args.test_dir)
funcs = collect_test_functions(test_dir)
ok, failures = check(pin_list, funcs)
print(f"pins: {len(pin_list)}")
print(f"tests: {len(funcs)}")
if failures:
for f in failures:
print(f" - {f}")
print("PASS" if ok else "FAIL")
return 0 if ok else 1
if __name__ == "__main__":
sys.exit(main())

View file

@ -0,0 +1,513 @@
"""Shared fixtures for tests/test_litellm/proxy/proxy_server/.
All fixtures and helpers used by PR1/PR2/PR3 test files live here. Do NOT
add fixtures inside individual test files. If a fixture is missing, add it
here and update the Notion plan.
"""
from __future__ import annotations
import contextlib
import os
import sys
from pathlib import Path
from typing import Any, AsyncIterator, Callable, Dict, Iterator, List, Optional
from unittest.mock import AsyncMock, MagicMock
import pytest
# Repo root, anchored to this file (not CWD) so the path is correct no
# matter where pytest is invoked from. With the project installed via
# uv this is defensive — `litellm` already resolves through site-packages
# — but it lets the harness work in editable-source layouts too.
sys.path.insert(0, str(Path(__file__).resolve().parents[4]))
# ---------------------------------------------------------------------------
# normalize() — used by every dict-equality assertion to scrub volatile fields
# ---------------------------------------------------------------------------
VOLATILE_KEYS = frozenset(
{
"created_at",
"updated_at",
"key",
"token",
"id",
"request_id",
"expires",
"expires_at",
"litellm_call_id",
"key_alias",
"created",
}
)
def normalize(data: Any, volatile: frozenset[str] = VOLATILE_KEYS) -> Any:
"""Replace volatile field values with "<VOLATILE>" so dict equality works.
Recursive over dicts and lists. Pass an explicit ``volatile`` set to
extend or override the default.
"""
if isinstance(data, dict):
return {
k: ("<VOLATILE>" if k in volatile else normalize(v, volatile))
for k, v in data.items()
}
if isinstance(data, list):
return [normalize(v, volatile) for v in data]
return data
# ---------------------------------------------------------------------------
# app + client — session-scoped so app import + TestClient setup amortize
# ---------------------------------------------------------------------------
@pytest.fixture(scope="session")
def app():
"""Return the proxy_server FastAPI app with lifespan effectively disabled.
TestClient used WITHOUT the ``with`` context manager skips the lifespan,
so the startup event (DB connect, Router init, OTEL setup) never fires.
Module import still runs once; module-level globals are harmless.
"""
os.environ.setdefault("LITELLM_LOG", "ERROR")
from litellm.proxy.proxy_server import app as _app
return _app
@pytest.fixture(scope="session")
def client(app):
"""TestClient wrapping the session app.
NOT entered as a context manager — lifespan does not fire. Tests that
require a real lifespan should use a function-scoped TestClient with
a ``with`` block locally and accept the per-test cost.
"""
from fastapi.testclient import TestClient
return TestClient(app, raise_server_exceptions=False)
# ---------------------------------------------------------------------------
# mock_prisma — function-scoped MagicMock with the common table methods stubbed
# ---------------------------------------------------------------------------
# Tables most-touched by proxy_server.py routes. Add to this list if a
# test discovers a missing table.
_PRISMA_TABLES: List[str] = [
"litellm_verificationtoken",
"litellm_teamtable",
"litellm_usertable",
"litellm_endusertable",
"litellm_organizationtable",
"litellm_organizationmembership",
"litellm_proxymodeltable",
"litellm_modeltable",
"litellm_budgettable",
"litellm_spendlogs",
"litellm_invitationlink",
"litellm_credentialstable",
"litellm_mcpservertable",
"litellm_objectpermissiontable",
"litellm_configtable",
"litellm_audit_log",
"litellm_dailyuserspend",
"litellm_dailyteamspend",
"litellm_dailytagspend",
"litellm_managed_object_table",
"litellm_managed_vector_stores_table",
"litellm_promptstable",
"litellm_guardrailstable",
"litellm_managed_files",
"litellm_session_token_table",
"litellm_passthrough_endpoint_table",
"litellm_cron_job",
"litellm_passthrough_logs",
"litellm_health_check_table",
"litellm_mcpusercredentials",
]
def _make_table_mock() -> MagicMock:
table = MagicMock()
table.find_unique = AsyncMock(return_value=None)
table.find_many = AsyncMock(return_value=[])
table.find_first = AsyncMock(return_value=None)
table.create = AsyncMock()
table.create_many = AsyncMock()
table.update = AsyncMock()
table.update_many = AsyncMock()
table.upsert = AsyncMock()
table.delete = AsyncMock()
table.delete_many = AsyncMock()
table.count = AsyncMock(return_value=0)
table.group_by = AsyncMock(return_value=[])
table.aggregate = AsyncMock(return_value={})
return table
@pytest.fixture
def mock_prisma() -> MagicMock:
"""MagicMock prisma_client with .db.<table> methods stubbed.
Default returns: find_unique/find_first -> None, find_many/group_by -> [],
count -> 0. Override in a test with::
mock_prisma.db.litellm_teamtable.find_unique.return_value = ...
"""
client_mock = MagicMock()
client_mock.db = MagicMock()
client_mock.connect = AsyncMock()
client_mock.disconnect = AsyncMock()
client_mock.health_check = AsyncMock(return_value=True)
for table_name in _PRISMA_TABLES:
setattr(client_mock.db, table_name, _make_table_mock())
return client_mock
# ---------------------------------------------------------------------------
# auth_as — context manager that overrides user_api_key_auth dependency
# ---------------------------------------------------------------------------
@pytest.fixture
def auth_as(app) -> Callable[..., contextlib.AbstractContextManager]:
"""Context manager that overrides ``user_api_key_auth`` for a role.
Usage::
def test_admin_only(client, auth_as):
from litellm.proxy._types import LitellmUserRoles
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.get("/some/admin/route")
assert response.status_code == 200
Outside the ``with`` block the override is removed so other tests see
the real dependency.
"""
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@contextlib.contextmanager
def _auth_as(
role: Any = None,
user_id: str = "test-user-id",
team_id: Optional[str] = None,
api_key: str = "sk-test-key",
**kwargs: Any,
) -> Iterator[Any]:
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
if role is None:
role = LitellmUserRoles.PROXY_ADMIN
fake_auth = UserAPIKeyAuth(
api_key=api_key,
user_id=user_id,
team_id=team_id,
user_role=role,
**kwargs,
)
async def _override() -> UserAPIKeyAuth:
return fake_auth
previous = app.dependency_overrides.get(user_api_key_auth)
app.dependency_overrides[user_api_key_auth] = _override
try:
yield fake_auth
finally:
if previous is None:
app.dependency_overrides.pop(user_api_key_auth, None)
else:
app.dependency_overrides[user_api_key_auth] = previous
return _auth_as
# ---------------------------------------------------------------------------
# Response builders — used by mock_router for parametrized responses
# ---------------------------------------------------------------------------
def make_acompletion_response(
model: str = "gpt-4",
messages: Optional[List[Dict[str, Any]]] = None,
stream: bool = False,
tools: Optional[List[Dict[str, Any]]] = None,
content: str = "Hello from mock",
**kwargs: Any,
) -> Any:
"""Build a deterministic chat-completion response.
Returns:
- An async generator when ``stream=True``
- A tool-call shape when ``tools`` is non-empty
- A plain text response otherwise
"""
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Choices,
Function,
Message,
ModelResponse,
Usage,
)
if stream:
return _stream_chunks(model=model, content=content)
if tools:
tool_name = tools[0].get("function", {}).get("name", "fake_tool")
message = Message(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
id="call_test",
type="function",
function=Function(name=tool_name, arguments="{}"),
)
],
)
else:
message = Message(role="assistant", content=content)
return ModelResponse(
id="chatcmpl-test",
choices=[Choices(finish_reason="stop", index=0, message=message)],
created=0,
model=model,
object="chat.completion",
usage=Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2),
)
async def _stream_chunks(
model: str = "gpt-4", content: str = "Hi"
) -> AsyncIterator[Any]:
from litellm.types.utils import (
Delta,
ModelResponseStream,
StreamingChoices,
)
for piece in [content, ""]:
yield ModelResponseStream(
id="chatcmpl-test",
choices=[
StreamingChoices(
finish_reason=None if piece else "stop",
index=0,
delta=Delta(content=piece or None, role="assistant"),
)
],
created=0,
model=model,
object="chat.completion.chunk",
)
def make_embedding_response(
model: str = "text-embedding-ada-002",
input: Any = None,
dimensions: int = 8,
**kwargs: Any,
) -> Any:
from litellm.types.utils import EmbeddingResponse
if isinstance(input, list):
n = len(input)
elif input is None:
n = 1
else:
n = 1
return EmbeddingResponse(
model=model,
data=[
{"embedding": [0.0] * dimensions, "index": i, "object": "embedding"}
for i in range(n)
],
object="list",
usage={"prompt_tokens": n, "total_tokens": n},
)
def make_image_response(model: str = "dall-e-3", **kwargs: Any) -> Any:
from litellm.types.utils import ImageResponse
return ImageResponse(
created=0,
data=[{"url": "https://example.invalid/image.png"}],
)
def make_speech_response(**kwargs: Any) -> bytes:
"""Return a fake audio blob. The route serializes bytes to a streaming response."""
return b"\x00" * 128
def make_transcription_response(**kwargs: Any) -> Any:
from litellm.types.utils import TranscriptionResponse
return TranscriptionResponse(text="hello world")
def make_moderation_response(**kwargs: Any) -> Dict[str, Any]:
return {
"id": "modr-test",
"model": "text-moderation-latest",
"results": [
{
"flagged": False,
"categories": {},
"category_scores": {},
}
],
}
# ---------------------------------------------------------------------------
# mock_router — fake Router with all the *async* call surfaces stubbed
# ---------------------------------------------------------------------------
@pytest.fixture
def mock_router() -> MagicMock:
"""A MagicMock standing in for ``llm_router`` with parametrized responses."""
async def _acompletion(model: str = "gpt-4", messages=None, **kwargs):
return make_acompletion_response(model=model, messages=messages, **kwargs)
async def _aembedding(model: str = "text-embedding-ada-002", input=None, **kwargs):
return make_embedding_response(model=model, input=input, **kwargs)
async def _aimage_generation(**kwargs):
return make_image_response(**kwargs)
async def _aspeech(**kwargs):
return make_speech_response(**kwargs)
async def _atranscription(**kwargs):
return make_transcription_response(**kwargs)
async def _amoderation(**kwargs):
return make_moderation_response(**kwargs)
router = MagicMock()
router.acompletion = AsyncMock(side_effect=_acompletion)
router.aembedding = AsyncMock(side_effect=_aembedding)
router.aimage_generation = AsyncMock(side_effect=_aimage_generation)
router.aspeech = AsyncMock(side_effect=_aspeech)
router.atranscription = AsyncMock(side_effect=_atranscription)
router.amoderation = AsyncMock(side_effect=_amoderation)
router.model_list = [
{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}},
{
"model_name": "claude-sonnet",
"litellm_params": {"model": "anthropic/claude-3-5-sonnet-latest"},
},
{
"model_name": "bedrock-claude",
"litellm_params": {"model": "bedrock/anthropic.claude-3-5-sonnet"},
},
]
router.model_names = ["gpt-4", "claude-sonnet", "bedrock-claude"]
router.get_model_list = MagicMock(return_value=router.model_list)
return router
# ---------------------------------------------------------------------------
# mock_callbacks_disabled — autouse: zero out global callbacks per test
# ---------------------------------------------------------------------------
@pytest.fixture(autouse=True)
def mock_callbacks_disabled(monkeypatch) -> None:
"""Wipe ``litellm.callbacks`` and friends so tests don't leak side effects."""
import litellm
for attr in (
"callbacks",
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
"input_callback",
"service_callback",
):
if hasattr(litellm, attr):
monkeypatch.setattr(litellm, attr, [], raising=False)
# ---------------------------------------------------------------------------
# Builders for DB-like objects (used by routes that load from DB)
# ---------------------------------------------------------------------------
def make_user(
user_id: str = "user-test",
role: Any = None,
teams: Optional[List[str]] = None,
max_budget: Optional[float] = None,
spend: float = 0.0,
**kwargs: Any,
) -> Any:
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles
if role is None:
role = LitellmUserRoles.INTERNAL_USER
return LiteLLM_UserTable(
user_id=user_id,
user_role=role,
teams=teams or [],
max_budget=max_budget,
spend=spend,
**kwargs,
)
def make_team(
team_id: str = "team-test",
team_alias: str = "Test Team",
max_budget: Optional[float] = None,
spend: float = 0.0,
members_with_roles: Optional[List[Dict[str, Any]]] = None,
**kwargs: Any,
) -> Any:
from litellm.proxy._types import LiteLLM_TeamTable
return LiteLLM_TeamTable(
team_id=team_id,
team_alias=team_alias,
max_budget=max_budget,
spend=spend,
members_with_roles=members_with_roles or [],
**kwargs,
)
def make_key(
token: str = "hashed-test-key",
key_alias: Optional[str] = None,
team_id: Optional[str] = None,
user_id: str = "user-test",
spend: float = 0.0,
max_budget: Optional[float] = None,
**kwargs: Any,
) -> Any:
from litellm.proxy._types import LiteLLM_VerificationToken
return LiteLLM_VerificationToken(
token=token,
key_alias=key_alias,
team_id=team_id,
user_id=user_id,
spend=spend,
max_budget=max_budget,
**kwargs,
)

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1,283 @@
"""Smoke tests for the proxy_server/ test harness.
Validates that fixtures + scripts work end-to-end before PR1/PR2/PR3 depend
on them. ``_pin_check.py`` skips this file explicitly so it doesn't count
toward behavior pinning.
"""
from __future__ import annotations
import importlib.util
import sys
import textwrap
from pathlib import Path
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from .conftest import ( # type: ignore[import-not-found]
make_acompletion_response,
make_embedding_response,
normalize,
)
HERE = Path(__file__).resolve().parent
# ---------------------------------------------------------------------------
# Fixture smoke tests
# ---------------------------------------------------------------------------
def test_app_fixture_returns_fastapi_app(app):
assert isinstance(app, FastAPI)
assert app.router is not None
def test_client_fixture_returns_testclient(client):
assert isinstance(client, TestClient)
assert hasattr(client, "post")
assert hasattr(client, "get")
def test_mock_prisma_has_team_table(mock_prisma):
assert hasattr(mock_prisma.db, "litellm_teamtable")
assert callable(mock_prisma.db.litellm_teamtable.find_unique)
assert callable(mock_prisma.db.litellm_teamtable.find_many)
def test_mock_prisma_has_key_table(mock_prisma):
assert hasattr(mock_prisma.db, "litellm_verificationtoken")
assert callable(mock_prisma.db.litellm_verificationtoken.find_unique)
def test_auth_as_admin_overrides_dependency(app, auth_as):
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
with auth_as(LitellmUserRoles.PROXY_ADMIN):
assert user_api_key_auth in app.dependency_overrides
def test_auth_as_internal_user_overrides_dependency(app, auth_as):
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
with auth_as(LitellmUserRoles.INTERNAL_USER) as fake_auth:
assert user_api_key_auth in app.dependency_overrides
assert fake_auth.user_role == LitellmUserRoles.INTERNAL_USER
def test_auth_as_cleans_up_on_exit(app, auth_as):
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
assert user_api_key_auth not in app.dependency_overrides
with auth_as(LitellmUserRoles.PROXY_ADMIN):
pass
assert user_api_key_auth not in app.dependency_overrides
def test_mock_router_acompletion_callable(mock_router):
from unittest.mock import AsyncMock
assert isinstance(mock_router.acompletion, AsyncMock)
assert isinstance(mock_router.aembedding, AsyncMock)
assert isinstance(mock_router.aimage_generation, AsyncMock)
@pytest.mark.asyncio
async def test_make_acompletion_response_stream():
gen = make_acompletion_response(model="gpt-4", stream=True)
chunks = [chunk async for chunk in gen]
assert len(chunks) >= 1
# Last chunk should have finish_reason set
assert chunks[-1].choices[0].finish_reason == "stop"
def test_make_acompletion_response_tools():
resp = make_acompletion_response(
model="gpt-4",
tools=[{"type": "function", "function": {"name": "fake_tool"}}],
)
assert resp.choices[0].message.tool_calls is not None
assert resp.choices[0].message.tool_calls[0].function.name == "fake_tool"
def test_make_embedding_response_shape():
resp = make_embedding_response(input=["a", "b", "c"], dimensions=4)
data = resp.data
assert len(data) == 3
assert len(data[0]["embedding"]) == 4
def test_normalize_replaces_volatile_keys():
out = normalize({"key": "abc", "spend": 0, "nested": {"id": "x", "value": 5}})
assert out == {
"key": "<VOLATILE>",
"spend": 0,
"nested": {"id": "<VOLATILE>", "value": 5},
}
def test_normalize_handles_lists():
out = normalize([{"key": "a"}, {"key": "b"}])
assert out == [{"key": "<VOLATILE>"}, {"key": "<VOLATILE>"}]
# ---------------------------------------------------------------------------
# Script smoke tests — _coverage_check.py
# ---------------------------------------------------------------------------
def _load_script(name: str):
spec = importlib.util.spec_from_file_location(name, HERE / f"{name}.py")
assert spec is not None and spec.loader is not None
mod = importlib.util.module_from_spec(spec)
# Register in sys.modules so dataclasses can resolve cls.__module__.
sys.modules[name] = mod
spec.loader.exec_module(mod)
return mod
def _write_cov_xml(tmp_path: Path, line_rate: float, branch_rate: float) -> Path:
xml = textwrap.dedent(f"""\
<?xml version="1.0" ?>
<coverage version="7.0">
<packages>
<package name="litellm.proxy">
<classes>
<class filename="litellm/proxy/proxy_server.py"
line-rate="{line_rate}" branch-rate="{branch_rate}"/>
</classes>
</package>
</packages>
</coverage>
""")
path = tmp_path / "cov.xml"
path.write_text(xml)
return path
def test_coverage_check_pass_on_synthetic_xml(tmp_path):
cov_check = _load_script("_coverage_check")
xml = _write_cov_xml(tmp_path, line_rate=0.75, branch_rate=0.60)
line_pct, branch_pct = cov_check.parse_coverage_xml(xml)
assert line_pct == pytest.approx(75.0)
assert branch_pct == pytest.approx(60.0)
def test_coverage_check_fail_on_low_coverage(tmp_path, monkeypatch, capsys):
cov_check = _load_script("_coverage_check")
xml = _write_cov_xml(tmp_path, line_rate=0.10, branch_rate=0.05)
monkeypatch.setattr(
sys,
"argv",
["_coverage_check.py", "--pr-target", "3", "--coverage-xml", str(xml)],
)
rc = cov_check.main()
assert rc == 1
out = capsys.readouterr().out
assert "FAIL" in out
def test_coverage_check_pass_on_high_coverage(tmp_path, monkeypatch, capsys):
cov_check = _load_script("_coverage_check")
xml = _write_cov_xml(tmp_path, line_rate=0.75, branch_rate=0.60)
monkeypatch.setattr(
sys,
"argv",
["_coverage_check.py", "--pr-target", "3", "--coverage-xml", str(xml)],
)
rc = cov_check.main()
assert rc == 0
out = capsys.readouterr().out
assert "PASS" in out
# ---------------------------------------------------------------------------
# Script smoke tests — _pin_check.py
# ---------------------------------------------------------------------------
def _write_pin_list(tmp_path: Path, items: list) -> Path:
path = tmp_path / "pins.txt"
path.write_text("\n".join(f"- `{item}`" for item in items) + "\n")
return path
def _write_test_file(tmp_path: Path, name: str, body: str) -> Path:
path = tmp_path / name
path.write_text(textwrap.dedent(body))
return path
def test_pin_check_pass_on_complete_pins(tmp_path):
pin_check = _load_script("_pin_check")
_write_pin_list(tmp_path, ["update_cache"])
_write_test_file(
tmp_path,
"test_thing.py",
"""\
def test_update_cache_happy():
data = update_cache(value=1)
assert data == {"key1": 1, "key2": 2, "key3": 3}
def test_update_cache_error():
import pytest
with pytest.raises(ValueError):
update_cache(value=None)
""",
)
pin_list = pin_check.parse_pin_list(tmp_path / "pins.txt")
funcs = pin_check.collect_test_functions(tmp_path)
ok, failures = pin_check.check(pin_list, funcs)
assert ok, failures
def test_pin_check_fail_on_missing_pin(tmp_path):
pin_check = _load_script("_pin_check")
_write_pin_list(tmp_path, ["update_cache", "never_referenced_symbol"])
_write_test_file(
tmp_path,
"test_thing.py",
"""\
def test_update_cache_happy():
data = update_cache(value=1)
assert data == {"key1": 1, "key2": 2, "key3": 3}
def test_update_cache_error():
import pytest
with pytest.raises(ValueError):
update_cache(value=None)
""",
)
pin_list = pin_check.parse_pin_list(tmp_path / "pins.txt")
funcs = pin_check.collect_test_functions(tmp_path)
ok, failures = pin_check.check(pin_list, funcs)
assert not ok
assert any("never_referenced_symbol" in f for f in failures)
def test_pin_check_fail_on_status_only_test(tmp_path):
pin_check = _load_script("_pin_check")
_write_pin_list(tmp_path, ["some_route"])
_write_test_file(
tmp_path,
"test_thing.py",
"""\
def test_some_route_happy():
response = client.get("/some_route")
assert response.status_code == 200
def test_some_route_error():
response = client.get("/some_route")
assert response.status_code == 404
""",
)
pin_list = pin_check.parse_pin_list(tmp_path / "pins.txt")
funcs = pin_check.collect_test_functions(tmp_path)
ok, failures = pin_check.check(pin_list, funcs)
assert not ok
assert any("status-only" in f for f in failures)

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -0,0 +1 @@
"""Placeholder. Filled by a follow-up PR per the Notion plan."""

View file

@ -604,6 +604,49 @@ def test_ui_extensionless_route_requires_restructure(tmp_path):
assert "login" in response.text
def test_admin_ui_export_serves_nested_extensionless_routes():
out_dir = (
Path(litellm.__file__).parent / "proxy" / "_experimental" / "out"
)
assert out_dir.is_dir(), f"missing UI export at {out_dir}"
nested_html_offenders = [
path.relative_to(out_dir).as_posix()
for path in out_dir.rglob("*.html")
if path.parent != out_dir
and path.name != "index.html"
and "_next" not in path.parts
and "litellm-asset-prefix" not in path.parts
]
assert not nested_html_offenders, (
"Nested routes must be named index.html. Offenders: "
f"{nested_html_offenders}"
)
callback_index = out_dir / "mcp" / "oauth" / "callback" / "index.html"
assert callback_index.is_file(), (
f"MCP OAuth callback page must exist at {callback_index}; "
"without it /ui/mcp/oauth/callback 404s after Linear redirects back."
)
fastapi_app = FastAPI()
fastapi_app.mount(
"/ui", StaticFiles(directory=str(out_dir), html=True), name="ui"
)
client = TestClient(fastapi_app)
redirect = client.get(
"/ui/mcp/oauth/callback?code=abc&state=xyz",
follow_redirects=False,
)
assert redirect.status_code == 307
assert redirect.headers["location"].endswith("/ui/mcp/oauth/callback/?code=abc&state=xyz")
landed = client.get("/ui/mcp/oauth/callback?code=abc&state=xyz")
assert landed.status_code == 200
assert "<html" in landed.text.lower()
def test_restructure_always_happens(monkeypatch):
"""
Test that restructuring logic always executes regardless of LITELLM_NON_ROOT setting.

View file

@ -360,11 +360,15 @@ class TestProxySettingEndpoints:
assert "proxy_base_url" in values
assert "user_email" in values
# Verify values match our mock config
# Verify non-secret values match our mock config. OAuth client
# secrets are masked on read so the GET response never carries
# plaintext credentials.
assert values["google_client_id"] == "test_google_client_id"
assert values["google_client_secret"] == "test_google_client_secret"
assert values["google_client_secret"] != "test_google_client_secret"
assert "*" in values["google_client_secret"]
assert values["microsoft_client_id"] == "test_microsoft_client_id"
assert values["microsoft_client_secret"] == "test_microsoft_client_secret"
assert values["microsoft_client_secret"] != "test_microsoft_client_secret"
assert "*" in values["microsoft_client_secret"]
assert values["proxy_base_url"] == "https://example.com"
assert values["user_email"] == "admin@example.com"
@ -1321,10 +1325,12 @@ class TestProxySettingEndpoints:
assert "values" in data
assert "field_schema" in data
# Verify decrypted values are returned
# Verify decrypted values are returned. OAuth client secrets are
# masked on read so plaintext is never sent to the UI.
values = data["values"]
assert values["google_client_id"] == "decrypted_google_id"
assert values["google_client_secret"] == "decrypted_google_secret"
assert values["google_client_secret"] != "decrypted_google_secret"
assert "*" in values["google_client_secret"]
assert values["microsoft_client_id"] == "decrypted_microsoft_id"
assert values["proxy_base_url"] == "https://decrypted.example.com"

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