Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_suppress-spend-log-tracebacks-2208

This commit is contained in:
mateo-berri
2026-05-01 20:54:52 -07:00
29 changed files with 2568 additions and 174 deletions
+4 -2
View File
@@ -4,7 +4,7 @@ on:
workflow_dispatch:
inputs:
tag:
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0-dev.2, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
required: true
type: string
commit_hash:
@@ -46,9 +46,11 @@ jobs:
const commitHash = process.env.COMMIT_HASH;
// Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases.
// Accept both PEP 440 (`.dev`) and SemVer (`-dev`) separators so tags
// like `1.84.0.dev2` and `1.84.0-dev.2` are both detected.
// PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]`
// are stable maintenance releases, not pre-releases.
const isPrerelease = /(?:rc|nightly|alpha|beta|\.dev)/i.test(tag);
const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag);
const cosignSection = [
`## Verify Docker Image Signature`,
@@ -10,28 +10,21 @@ has already authenticated the user) and you need to extract user information fro
custom headers or other request attributes.
"""
from typing import TYPE_CHECKING, Dict, Optional, Union, cast
from typing import cast
from fastapi import Request
from fastapi.responses import RedirectResponse
if TYPE_CHECKING:
from fastapi_sso.sso.base import OpenID
else:
from typing import Any as OpenID
from litellm.proxy.management_endpoints.types import CustomOpenID
class EnterpriseCustomSSOHandler:
"""
Enterprise Custom SSO Handler for LiteLLM Proxy
This class provides methods for handling custom SSO authentication flows
where users can implement their own authentication logic by processing
request headers and returning user information in OpenID format.
"""
@staticmethod
async def handle_custom_ui_sso_sign_in(
request: Request,
@@ -40,16 +33,16 @@ class EnterpriseCustomSSOHandler:
Allow a user to execute their custom code to parse incoming request headers and return a OpenID object
Use this when you have an OAuth proxy in front of LiteLLM (where the OAuth proxy has already authenticated the user)
Args:
request: The FastAPI request object containing headers and other request data
Returns:
RedirectResponse: Redirect response that sends the user to the LiteLLM UI with authentication token
Raises:
ValueError: If custom_ui_sso_sign_in_handler is not configured
Example:
This method is typically called when a user has already been authenticated by an
external OAuth proxy and the proxy has added custom headers containing user information.
@@ -60,27 +53,44 @@ class EnterpriseCustomSSOHandler:
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
from litellm.proxy.proxy_server import (
CommonProxyErrors,
general_settings,
premium_user,
user_custom_ui_sso_sign_in_handler,
)
from litellm.proxy.auth.trusted_proxy_utils import (
require_trusted_proxy_request,
)
if premium_user is not True:
raise ValueError(CommonProxyErrors.not_premium_user.value)
if user_custom_ui_sso_sign_in_handler is None:
raise ValueError("custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings.")
custom_sso_login_handler = cast(CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler)
openid_response: OpenID = await custom_sso_login_handler.handle_custom_ui_sso_sign_in(
raise ValueError(
"custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings."
)
require_trusted_proxy_request(
request=request,
general_settings=general_settings,
feature_name="Custom UI SSO",
)
custom_sso_login_handler = cast(
CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler
)
openid_response: OpenID = (
await custom_sso_login_handler.handle_custom_ui_sso_sign_in(
request=request,
)
)
# Import here to avoid circular imports
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
return await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=openid_response,
request=request,
received_response=None,
generic_client_id=None,
ui_access_mode=None,
)
)
@@ -588,24 +588,21 @@ async def update_project( # noqa: PLR0915
param="project_id",
)
# Validate team exists and get team object for limit + permission checks
team_id_to_check = data.team_id or existing_project.team_id
team_obj_for_checks = None
if team_id_to_check is not None:
team_obj_for_checks = await _validate_team_exists(
team_id=team_id_to_check, prisma_client=prisma_client
# Permission to *edit* the project must be evaluated against the
# project's CURRENT team. Sourcing the team from `data.team_id`
# would let an admin of any team pass the check by supplying their
# own team_id, hijacking the project (VERIA-55).
target_team_id = data.team_id or existing_project.team_id
target_team_obj = None
if target_team_id is not None:
target_team_obj = await _validate_team_exists(
team_id=target_team_id, prisma_client=prisma_client
)
# Check if user has permission to update this project
has_permission = await _check_user_permission_for_project(
user_api_key_dict=user_api_key_dict,
team_id=existing_project.team_id,
prisma_client=prisma_client,
team_object=(
LiteLLM_TeamTable(**team_obj_for_checks.model_dump())
if team_obj_for_checks
else None
),
)
if not has_permission:
@@ -614,10 +611,32 @@ async def update_project( # noqa: PLR0915
detail={"error": "Only admins or team admins can update projects"},
)
# Reassigning to a different team also requires admin rights on the
# destination team — otherwise a team admin could shed projects into
# an unsuspecting team's namespace.
if data.team_id is not None and data.team_id != existing_project.team_id:
can_assign_to_target = await _check_user_permission_for_project(
user_api_key_dict=user_api_key_dict,
team_id=data.team_id,
prisma_client=prisma_client,
team_object=(
LiteLLM_TeamTable(**target_team_obj.model_dump())
if target_team_obj
else None
),
)
if not can_assign_to_target:
raise HTTPException(
status_code=403,
detail={
"error": "Cannot reassign project to a team you are not an admin of"
},
)
# Validate project limits against team limits
if team_obj_for_checks is not None:
if target_team_obj is not None:
_check_team_project_limits(
team_object=LiteLLM_TeamTable(**team_obj_for_checks.model_dump()),
team_object=LiteLLM_TeamTable(**target_team_obj.model_dump()),
data=data,
)
@@ -18,6 +18,17 @@ class CustomSSOLoginHandler(CustomLogger):
self,
request: Request,
) -> OpenID:
from litellm.proxy.auth.trusted_proxy_utils import (
require_trusted_proxy_request,
)
from litellm.proxy.proxy_server import general_settings
require_trusted_proxy_request(
request=request,
general_settings=general_settings,
feature_name="Custom UI SSO",
)
request_headers_dict = dict(request.headers)
return OpenID(
id=request_headers_dict.get("x-litellm-user-id"),
+35 -18
View File
@@ -69,6 +69,8 @@ class OpenTelemetryConfig:
deployment_environment: Optional[str] = None
model_id: Optional[str] = None
ignore_context_propagation: Optional[bool] = None
# When True, create a private TracerProvider instead of reusing or setting the global one.
skip_set_global: bool = False
def __post_init__(self) -> None:
# If endpoint is specified but exporter is still the default "console",
@@ -259,16 +261,21 @@ class OpenTelemetry(CustomLogger):
try:
existing_provider = get_existing_provider_fn()
# If a real SDK provider exists (set by another SDK like Langfuse), use it
# This uses a positive check for SDK providers instead of a negative check for proxy providers
if isinstance(existing_provider, sdk_provider_class):
verbose_logger.debug(
"OpenTelemetry: Using existing %s: %s",
provider_name,
type(existing_provider).__name__,
)
provider = existing_provider
# Don't call set_provider to preserve existing context
if skip_set_global:
verbose_logger.debug(
"OpenTelemetry: existing %s found but skip_set_global=True; creating private %s for isolation",
provider_name,
provider_name,
)
provider = create_new_provider_fn()
else:
verbose_logger.debug(
"OpenTelemetry: Using existing %s: %s",
provider_name,
type(existing_provider).__name__,
)
provider = existing_provider
else:
# Default proxy provider or unknown type, create our own
verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name)
@@ -293,6 +300,12 @@ class OpenTelemetry(CustomLogger):
return provider
def _skip_set_global(self) -> bool:
# langfuse_otel relies on the Langfuse SDK's providers; don't overwrite them.
return self.config.skip_set_global or (
hasattr(self, "callback_name") and self.callback_name == "langfuse_otel"
)
def _init_tracing(self, tracer_provider):
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
@@ -303,11 +316,6 @@ class OpenTelemetry(CustomLogger):
provider.add_span_processor(self._get_span_processor())
return provider
# CRITICAL FIX: For Langfuse OTEL, skip setting global provider to prevent interference
skip_global = (
hasattr(self, "callback_name") and self.callback_name == "langfuse_otel"
)
tracer_provider = self._get_or_create_provider(
provider=tracer_provider,
provider_name="TracerProvider",
@@ -315,16 +323,18 @@ class OpenTelemetry(CustomLogger):
sdk_provider_class=TracerProvider,
create_new_provider_fn=create_tracer_provider,
set_provider_fn=trace.set_tracer_provider,
skip_set_global=skip_global,
skip_set_global=self._skip_set_global(),
)
# Grab our tracer from the TracerProvider (not from global context)
# This ensures we use the provided TracerProvider (e.g., for testing)
self.tracer = tracer_provider.get_tracer(LITELLM_TRACER_NAME)
self._tracer_provider = tracer_provider
self.span_kind = SpanKind
def _init_metrics(self, meter_provider):
if not self.config.enable_metrics:
self._meter_provider = None
self._operation_duration_histogram = None
self._token_usage_histogram = None
self._cost_histogram = None
@@ -350,7 +360,9 @@ class OpenTelemetry(CustomLogger):
sdk_provider_class=MeterProvider,
create_new_provider_fn=create_meter_provider,
set_provider_fn=metrics.set_meter_provider,
skip_set_global=self._skip_set_global(),
)
self._meter_provider = meter_provider
meter = meter_provider.get_meter(__name__)
@@ -388,6 +400,7 @@ class OpenTelemetry(CustomLogger):
def _init_logs(self, logger_provider):
# nothing to do if events disabled
if not self.config.enable_events:
self._logger_provider = None
return
from opentelemetry._logs import get_logger_provider, set_logger_provider
@@ -404,13 +417,14 @@ class OpenTelemetry(CustomLogger):
)
return provider
self._get_or_create_provider(
self._logger_provider = self._get_or_create_provider(
provider=logger_provider,
provider_name="LoggerProvider",
get_existing_provider_fn=get_logger_provider,
sdk_provider_class=OTLoggerProvider,
create_new_provider_fn=create_logger_provider,
set_provider_fn=set_logger_provider,
skip_set_global=self._skip_set_global(),
)
def log_success_event(self, kwargs, response_obj, start_time, end_time):
@@ -1073,7 +1087,7 @@ class OpenTelemetry(CustomLogger):
# See: https://github.com/open-telemetry/opentelemetry-python/pull/4676
# TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords
from opentelemetry._logs import SeverityNumber, get_logger
from opentelemetry._logs import SeverityNumber
try:
from opentelemetry.sdk._logs import ( # type: ignore[attr-defined] # OTEL < 1.39.0
@@ -1084,7 +1098,10 @@ class OpenTelemetry(CustomLogger):
LogRecord as SdkLogRecord, # type: ignore[attr-defined] # OTEL >= 1.39.0
)
otel_logger = get_logger(LITELLM_LOGGER_NAME)
# Resolve through the handler's own LoggerProvider (which may be a
# private one when skip_set_global=True) rather than the module-level
# get_logger() which always goes through the global provider.
otel_logger = self._logger_provider.get_logger(LITELLM_LOGGER_NAME)
parent_ctx = span.get_span_context()
provider = (kwargs.get("litellm_params") or {}).get(
@@ -2,6 +2,7 @@
Helper functions to query prometheus API
"""
import json
import time
from datetime import datetime, timedelta
from typing import Optional
@@ -81,6 +82,24 @@ def is_prometheus_connected() -> bool:
return False
def _quote_promql_string_literal(value: str) -> str:
"""Render ``value`` as a PromQL double-quoted string literal.
PromQL string literals follow Go's escape rules
(https://prometheus.io/docs/prometheus/latest/querying/basics/): a
backslash begins an escape sequence and a bare ``"`` ends the literal.
Without escaping, callers that accept arbitrary user-supplied values
(like the ``api_key`` filter on ``/global/spend/logs``) can inject extra
label matchers or selectors and read cross-tenant metrics.
JSON's quoting rules are a strict subset of Go's, so ``json.dumps`` of
a Python string produces a literal Prometheus accepts: ``\\``, ``\\"``,
and the standard ``\\n`` / ``\\t`` / ``\\uNNNN`` control-character
escapes. The returned value already includes the surrounding quotes.
"""
return json.dumps(value, ensure_ascii=False)
async def get_daily_spend_from_prometheus(api_key: Optional[str]):
"""
Expected Response Format:
@@ -109,8 +128,11 @@ async def get_daily_spend_from_prometheus(api_key: Optional[str]):
if api_key is None:
query = "sum(delta(litellm_spend_metric_total[1d]))"
else:
quoted_api_key = _quote_promql_string_literal(api_key)
query = (
f'sum(delta(litellm_spend_metric_total{{hashed_api_key="{api_key}"}}[1d]))'
"sum(delta(litellm_spend_metric_total{"
f"hashed_api_key={quoted_api_key}"
"}[1d]))"
)
params = {
+4
View File
@@ -2438,6 +2438,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For headers are only trusted from these IPs.",
)
trusted_proxy_ranges: Optional[List[str]] = Field(
None,
description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.",
)
store_model_in_db: Optional[bool] = Field(
None,
description="If True, models and config are stored in and loaded from the database. Default is False.",
+90 -14
View File
@@ -502,23 +502,28 @@ async def common_checks( # noqa: PLR0915
f"Team={team_object.team_id} is blocked. Update via `/team/unblock` if you're an admin."
)
# 2. If team can call model
# 2. If team can call model (or key's access_group_ids grant it)
if _model and team_object:
with tracer.trace("litellm.proxy.auth.common_checks.can_team_access_model"):
if not await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=(
valid_token.team_model_aliases if valid_token else None
),
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
try:
await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=(
valid_token.team_model_aliases if valid_token else None
),
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
raise
if not await _key_access_group_grants_model(
model=_model,
valid_token=valid_token,
team_object=team_object,
llm_router=llm_router,
):
raise
# 2.2. If team member has per-member model scope, enforce it
if _model and team_object and valid_token and valid_token.user_id:
@@ -2975,6 +2980,77 @@ async def can_team_access_model(
raise
async def _key_access_group_grants_model(
model: Union[str, List[str]],
valid_token: Optional[UserAPIKeyAuth],
team_object: Optional[LiteLLM_TeamTable],
llm_router: Optional[Router],
) -> bool:
"""
Returns True if the key's `access_group_ids` expand to models that grant
access to `model`. Used to let a key's access group override a team's
model restriction in `common_checks`.
A key's access group only counts if the access group itself authorizes the
caller as an owner that is, the group's `assigned_team_ids` includes the
key's `team_id`, or the group's `assigned_key_ids` includes the key's
token. This preserves the team-as-owner boundary (a team member cannot
escalate by naming a group assigned to a different team) while still
letting a group reach the key without first being added to the team's
`access_group_ids` list.
"""
if valid_token is None:
return False
key_access_group_ids = list(valid_token.access_group_ids or [])
if not key_access_group_ids:
return False
from litellm.proxy.proxy_server import prisma_client as _prisma_client
from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging_obj
from litellm.proxy.proxy_server import user_api_key_cache as _user_api_key_cache
if _prisma_client is None or _user_api_key_cache is None:
return False
key_team_id = valid_token.team_id or (
team_object.team_id if team_object is not None else None
)
key_token = valid_token.token
authorized_models: List[str] = []
for ag_id in key_access_group_ids:
try:
ag = await get_access_object(
access_group_id=ag_id,
prisma_client=_prisma_client,
user_api_key_cache=_user_api_key_cache,
proxy_logging_obj=_proxy_logging_obj,
)
except Exception:
continue
team_authorized = bool(
key_team_id and key_team_id in (ag.assigned_team_ids or [])
)
key_authorized = bool(key_token and key_token in (ag.assigned_key_ids or []))
if team_authorized or key_authorized:
authorized_models.extend(ag.access_model_names or [])
if not authorized_models:
return False
try:
_can_object_call_model(
model=model,
llm_router=llm_router,
models=list(set(authorized_models)),
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
object_type="key",
)
return True
except ProxyException:
return False
def can_project_access_model(
model: Union[str, List[str]],
project_object: LiteLLM_ProjectTableCachedObj,
+42 -7
View File
@@ -707,11 +707,48 @@ class JWTHandler:
verbose_proxy_logger.error(f"Error fetching OIDC UserInfo: {str(e)}")
raise Exception(f"Failed to fetch OIDC UserInfo: {str(e)}")
async def auth_jwt(self, token: str) -> dict:
_unscoped_jwt_warning_emitted = False
@classmethod
def _build_decode_kwargs(cls) -> dict:
"""Build the audience/issuer/options kwargs for ``jwt.decode``.
Setting ``JWT_AUDIENCE`` (and optionally ``JWT_ISSUER``) turns on the
corresponding PyJWT verifications, blocking cross-tenant tokens
minted by other applications that share the same IdP signing keys.
When both are unset PyJWT only checks the signature and expiry, which
is preserved for backward compatibility but logged once as a warning.
"""
audience = os.getenv("JWT_AUDIENCE")
decode_options = None
issuer = os.getenv("JWT_ISSUER")
if (
audience is None
and issuer is None
and not cls._unscoped_jwt_warning_emitted
):
verbose_proxy_logger.warning(
"JWT auth is enabled but neither JWT_AUDIENCE nor JWT_ISSUER "
"is configured. Tokens minted by any application that shares "
"the same IdP signing keys will be accepted. Set JWT_AUDIENCE "
"(and ideally JWT_ISSUER) to scope this proxy."
)
cls._unscoped_jwt_warning_emitted = True
options: dict = {}
if audience is None:
decode_options = {"verify_aud": False}
options["verify_aud"] = False
if issuer is None:
options["verify_iss"] = False
return {
"audience": audience,
"issuer": issuer,
"options": options or None,
}
async def auth_jwt(self, token: str) -> dict:
decode_kwargs = self._build_decode_kwargs()
header = jwt.get_unverified_header(token)
@@ -747,9 +784,8 @@ class JWTHandler:
token,
public_key_obj, # type: ignore
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
options=decode_options, # type: ignore[arg-type]
audience=audience,
leeway=self.leeway, # allow testing of expired tokens
**decode_kwargs,
)
return payload
@@ -775,8 +811,7 @@ class JWTHandler:
token,
key,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
audience=audience,
options=decode_options,
**decode_kwargs,
)
return payload
+76 -16
View File
@@ -1,19 +1,69 @@
from typing import Any, Dict
from typing import Any, Dict, FrozenSet
from fastapi import Request
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.trusted_proxy_utils import require_trusted_proxy_request
# OAuth2-proxy header trust is for **identity assertion** from a trusted
# upstream auth proxy (oauth2-proxy, Authelia, etc.). The allowlist below
# is the only safe surface — anything else (``user_role``, ``api_key``,
# ``permissions``, ``max_budget``, ``user_max_budget``,
# ``team_tpm_limit``, ``end_user_max_budget``, ``allowed_model_region``,
# and dozens of similar policy fields scattered across the
# ``LiteLLM_VerificationTokenView`` hierarchy) is a privilege grant that
# would let a caller forge their own enforcement parameters by sending
# the matching header.
#
# A denylist of "privileged fields" is unmaintainable in this codebase:
# the auth model has ~50 budget/spend/limit/permission fields and gains
# more with each release. An allowlist scoped to identity assertion is
# default-secure — new fields are blocked automatically.
#
# Operators who need a trusted upstream to assert anything beyond
# identity should switch to JWT authentication, which validates a
# signature on the assertion rather than blindly trusting headers.
ALLOWED_OAUTH2_PROXY_FIELDS: FrozenSet[str] = frozenset(
{
"user_id",
"user_email",
"team_id",
"team_alias",
"org_id",
"models",
}
)
async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
"""
Handle request from oauth2 proxy.
Resolve a ``UserAPIKeyAuth`` from request headers per the admin-set
``oauth2_config_mappings``.
The auth model assumes the proxy is deployed behind a trusted OAuth2
reverse proxy that injects authenticated identity headers (e.g.
oauth2-proxy, Authelia).
**Identity-only allowlist.** ``oauth2_config_mappings`` maps header
names to ``UserAPIKeyAuth`` fields. Without an allowlist, an admin
who maps the wrong header to ``user_role`` lets any caller send
``X-User-Role: proxy_admin`` and gain full admin privileges
(Pydantic coerces the string into the enum). Only fields in
``ALLOWED_OAUTH2_PROXY_FIELDS`` (identity assertion only see the
constant's comment) may be mapped; any other mapping is rejected at
request time so the misconfiguration surfaces loudly rather than as
a silent privesc.
"""
from litellm.proxy.proxy_server import general_settings
verbose_proxy_logger.debug("Handling oauth2 proxy request")
# Define the OAuth2 config mappings
require_trusted_proxy_request(
request=request,
general_settings=general_settings,
feature_name="OAuth2 proxy auth",
)
oauth2_config_mappings: Dict[str, str] = (
general_settings.get("oauth2_config_mappings") or {}
)
@@ -21,21 +71,32 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
if not oauth2_config_mappings:
raise ValueError("Oauth2 config mappings not found in general_settings")
# Initialize a dictionary to store the mapped values
auth_data: Dict[str, Any] = {}
# Extract values from headers based on the mappings
disallowed = sorted(
set(oauth2_config_mappings.keys()) - ALLOWED_OAUTH2_PROXY_FIELDS
)
if disallowed:
raise ValueError(
"Oauth2 proxy auth refuses to map non-identity UserAPIKeyAuth "
f"fields from request headers: {disallowed}. Only identity "
f"fields are accepted ({sorted(ALLOWED_OAUTH2_PROXY_FIELDS)}); "
"anything else (privileges, budgets, rate limits, metadata) "
"would let a caller forge enforcement parameters by spoofing "
"the matching header. If you need a trusted upstream to "
"assert anything beyond identity, use JWT auth "
"(signature-validated) instead of header-trust."
)
auth_data: Dict[str, Any] = {}
for key, header in oauth2_config_mappings.items():
value = request.headers.get(header)
if value:
# Convert max_budget to float if present
if key == "max_budget":
auth_data[key] = float(value)
# Convert models to list if present
elif key == "models":
auth_data[key] = [model.strip() for model in value.split(",")]
else:
auth_data[key] = value
if not value:
continue
if key == "models":
auth_data[key] = [model.strip() for model in value.split(",")]
else:
auth_data[key] = value
verbose_proxy_logger.debug(
"Auth data before creating UserAPIKeyAuth object: keys=%s",
list(auth_data.keys()),
@@ -45,5 +106,4 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
"UserAPIKeyAuth object created with keys: %s",
list(user_api_key_auth.__fields_set__),
)
# Create and return UserAPIKeyAuth object
return user_api_key_auth
+118
View File
@@ -0,0 +1,118 @@
import ipaddress
from typing import Any, Dict, List, Optional, Union
from fastapi import Request
from litellm._logging import verbose_proxy_logger
TRUSTED_PROXY_RANGES_KEY = "trusted_proxy_ranges"
TrustedProxyNetwork = Union[ipaddress.IPv4Network, ipaddress.IPv6Network]
def _get_proxy_general_settings() -> Dict[str, Any]:
try:
from litellm.proxy.proxy_server import general_settings
return general_settings or {}
except ImportError:
return {}
def _normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str) -> List[str]:
if not configured_ranges:
return []
if isinstance(configured_ranges, str):
return [
raw_range.strip()
for raw_range in configured_ranges.split(",")
if raw_range.strip()
]
if isinstance(configured_ranges, (list, tuple, set)):
return [
str(raw_range).strip()
for raw_range in configured_ranges
if str(raw_range).strip()
]
verbose_proxy_logger.warning(
"Invalid %s value: expected a list of CIDR ranges, got %s",
setting_name,
type(configured_ranges).__name__,
)
return []
def parse_trusted_proxy_ranges(
configured_ranges: Any,
*,
setting_name: str = TRUSTED_PROXY_RANGES_KEY,
) -> List[TrustedProxyNetwork]:
networks: List[TrustedProxyNetwork] = []
for cidr in _normalize_cidr_ranges(configured_ranges, setting_name=setting_name):
try:
networks.append(ipaddress.ip_network(cidr, strict=False))
except ValueError:
verbose_proxy_logger.warning(
"Invalid CIDR in %s: %s, skipping", setting_name, cidr
)
return networks
def _get_direct_client_ip(request: Request) -> Optional[str]:
client = getattr(request, "client", None)
client_host = getattr(client, "host", None)
if isinstance(client_host, str):
return client_host
return None
def _is_ip_in_networks(
client_ip: Optional[str], networks: List[TrustedProxyNetwork]
) -> bool:
if not client_ip or not networks:
return False
try:
addr = ipaddress.ip_address(client_ip.strip())
except ValueError:
return False
return any(addr in network for network in networks)
def require_trusted_proxy_request(
*,
request: Request,
general_settings: Optional[Dict[str, Any]] = None,
feature_name: str,
setting_name: str = TRUSTED_PROXY_RANGES_KEY,
) -> None:
"""
Fail closed unless the direct TCP peer is one of the configured
trusted reverse proxies.
Header-based auth paths must validate the direct peer, not
X-Forwarded-For, because the direct peer is the actor supplying the
identity headers.
"""
if general_settings is None:
general_settings = _get_proxy_general_settings()
trusted_networks = parse_trusted_proxy_ranges(
general_settings.get(setting_name), setting_name=setting_name
)
if not trusted_networks:
raise ValueError(
f"{feature_name} requires general_settings.{setting_name} before "
"trusting identity headers from an upstream proxy."
)
direct_client_ip = _get_direct_client_ip(request)
if not _is_ip_in_networks(direct_client_ip, trusted_networks):
verbose_proxy_logger.warning(
"%s rejected identity headers from untrusted direct client IP %r",
feature_name,
direct_client_ip,
)
raise ValueError(
f"{feature_name} only accepts identity headers from configured "
f"trusted proxy ranges. Direct client IP {direct_client_ip!r} "
"is not trusted."
)
+62 -17
View File
@@ -11,7 +11,7 @@ import asyncio
import re
import secrets
from datetime import datetime, timezone
from typing import Any, List, Optional, Tuple, Union, cast
from typing import Any, Iterator, List, Optional, Tuple, Union, cast
import fastapi
from fastapi import HTTPException, Request, WebSocket, status
@@ -2271,10 +2271,6 @@ async def _enforce_key_and_fallback_model_access(
route=route,
request=request,
)
fallback_models = cast(
Optional[List[ALL_FALLBACK_MODEL_VALUES]],
request_data.get("fallbacks", None),
)
if model is not None:
await can_key_call_model(
@@ -2284,20 +2280,69 @@ async def _enforce_key_and_fallback_model_access(
llm_router=llm_router,
)
if fallback_models is not None:
for m in fallback_models:
await can_key_call_model(
model=m["model"] if isinstance(m, dict) else m,
llm_model_list=llm_model_list,
valid_token=valid_token,
llm_router=llm_router,
)
await is_valid_fallback_model(
model=m["model"] if isinstance(m, dict) else m,
llm_router=llm_router,
user_model=None,
# Validate every fallback model name reachable by this request.
# All three fields (``fallbacks``, ``context_window_fallbacks``,
# ``content_policy_fallbacks``) are forwarded to the router as
# per-request kwargs whether they appear at the top level of
# ``request_data`` or nested under ``router_settings_override``.
# Both surfaces must be validated against the API key's model
# allowlist or a caller can smuggle a restricted model. VERIA-44.
fallback_names: List[str] = []
override_settings = request_data.get("router_settings_override")
for _fb_key in ROUTER_FALLBACK_FIELDS:
fallback_names.extend(
iter_router_fallback_model_names(request_data.get(_fb_key))
)
if isinstance(override_settings, dict):
fallback_names.extend(
iter_router_fallback_model_names(override_settings.get(_fb_key))
)
for _name in dict.fromkeys(fallback_names): # dedupe, preserve order
await can_key_call_model(
model=_name,
llm_model_list=llm_model_list,
valid_token=valid_token,
llm_router=llm_router,
)
await is_valid_fallback_model(
model=_name,
llm_router=llm_router,
user_model=None,
)
ROUTER_FALLBACK_FIELDS: Tuple[str, ...] = (
"fallbacks",
"context_window_fallbacks",
"content_policy_fallbacks",
)
def iter_router_fallback_model_names(fallbacks: Any) -> Iterator[str]:
"""Yield leaf model names from any of the supported fallbacks shapes.
Handles the simple top-level shape (``str`` or ``{"model": str}``) and
the nested router-config shape (``[{primary: [fallback_list]}]``).
"""
if not isinstance(fallbacks, list):
return
for entry in fallbacks:
if isinstance(entry, str):
yield entry
elif isinstance(entry, dict):
if isinstance(entry.get("model"), str):
yield entry["model"]
continue
for fallback_list in entry.values():
if not isinstance(fallback_list, list):
continue
for m in fallback_list:
if isinstance(m, str):
yield m
elif isinstance(m, dict) and isinstance(m.get("model"), str):
yield m["model"]
async def _run_post_custom_auth_checks(
valid_token: UserAPIKeyAuth,
@@ -810,6 +810,20 @@ async def _common_key_generation_helper( # noqa: PLR0915
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
if prisma_client:
# Mirror the membership rule applied to /key/update: when the
# caller specifies an organization_id, require that they are a
# member of (or proxy admin over) the target organization.
_is_proxy_admin = (
user_api_key_dict.user_role is not None
and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if not _is_proxy_admin:
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
prisma_client=prisma_client,
)
org_table = await get_org_object(
org_id=data.organization_id,
user_api_key_cache=user_api_key_cache,
@@ -1169,6 +1183,42 @@ def check_org_key_rpm_tpm_limits(
)
async def _validate_caller_can_assign_key_org(
user_api_key_dict: UserAPIKeyAuth,
organization_id: str,
prisma_client: PrismaClient,
) -> None:
"""Reject ``/key/update`` requests that point a key at an organization
the caller does not belong to.
Mirrors the org-membership rule already enforced on ``/key/list`` in
``validate_key_list_check``. Proxy admins are checked at the call site.
"""
if user_api_key_dict.user_id is None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Cannot assign a key to an organization without a user_id on the caller's token",
)
user_row = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_api_key_dict.user_id},
include={"organization_memberships": True},
)
memberships = (
getattr(user_row, "organization_memberships", None) if user_row else None
)
member_org_ids = {
membership.organization_id
for membership in (memberships or [])
if membership.organization_id is not None
}
if organization_id not in member_org_ids:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"Caller is not a member of organization_id={organization_id}",
)
async def _check_org_key_limits(
org_table: LiteLLM_OrganizationTable,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
@@ -2169,10 +2219,26 @@ async def _validate_update_key_data(
user_api_key_cache=user_api_key_cache,
)
# When the caller asks to change the key's organization_id, require that
# they are a member of (or a proxy admin over) the target organization.
# Without this gate, any caller could assign their key to an arbitrary
# organization_id by passing it in the request body — VERIA-55 secondary
# IDOR. The check mirrors the membership rule already used on the
# `/key/list` filter path in `validate_key_list_check`.
_existing_org_id = getattr(existing_key_row, "organization_id", None)
if (
data.organization_id is not None
and data.organization_id != _existing_org_id
and not _is_proxy_admin
):
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
prisma_client=prisma_client,
)
# Check org key limits only when throughput-related fields or organization_id change
_org_id_to_check = data.organization_id or getattr(
existing_key_row, "organization_id", None
)
_org_id_to_check = data.organization_id or _existing_org_id
_throughput_fields_changed = (
data.organization_id is not None
or data.tpm_limit is not None
@@ -3869,6 +3935,22 @@ async def _execute_virtual_key_regeneration(
"""Generate new token, update DB, invalidate cache, and return response."""
from litellm.proxy.proxy_server import hash_token
# Apply the same membership rule used on /key/update: when the caller
# asks to point the regenerated key at a different organization_id,
# require they are a member of (or proxy admin over) the target org.
if data is not None and data.organization_id is not None:
_existing_org_id = getattr(key_in_db, "organization_id", None)
_is_proxy_admin = (
user_api_key_dict.user_role is not None
and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if data.organization_id != _existing_org_id and not _is_proxy_admin:
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
prisma_client=prisma_client,
)
new_token = await get_new_token(data=data)
new_token_hash = hash_token(new_token)
new_token_key_name = f"sk-...{new_token[-4:]}"
+63 -4
View File
@@ -678,6 +678,7 @@ async def google_login(
google_client_id=google_client_id,
generic_client_id=generic_client_id,
state=cli_state,
request=request,
)
if return_to is not None and sso_redirect is not None:
if SSOAuthenticationHandler._validate_return_to(return_to):
@@ -1159,6 +1160,30 @@ async def get_generic_sso_response(
authorization_code = request.query_params.get("code")
if code_verifier:
# State-to-session-cookie binding. The non-PKCE branch below
# delegates to fastapi-sso's ``verify_and_process``, which
# performs its own session-cookie check. The PKCE branch
# bypasses that helper, so we validate the URL ``state``
# against the ``litellm_oauth_state`` cookie set on the
# redirect response — without this an attacker can pre-mint
# a state + cached PKCE verifier and hijack a victim's auth
# code (Login-CSRF / token theft).
url_state = request.query_params.get("state")
cookie_state = request.cookies.get("litellm_oauth_state")
if (
not url_state
or not cookie_state
or not secrets.compare_digest(url_state, cookie_state)
):
raise ProxyException(
message=(
"Invalid OAuth state parameter — does not match "
"the browser-bound state cookie."
),
type=ProxyErrorTypes.auth_error,
param="state",
code=status.HTTP_400_BAD_REQUEST,
)
if not authorization_code:
raise ProxyException(
message="Missing authorization code in callback",
@@ -2147,6 +2172,7 @@ class SSOAuthenticationHandler:
microsoft_client_id: Optional[str] = None,
generic_client_id: Optional[str] = None,
state: Optional[str] = None,
request: Optional[Request] = None,
) -> Optional[RedirectResponse]:
"""
Step 1. Call Get Login Redirect for the SSO provider. Send the redirect response to `redirect_url`
@@ -2156,6 +2182,8 @@ class SSOAuthenticationHandler:
google_client_id (Optional[str], optional): The Google Client ID. Defaults to None.
microsoft_client_id (Optional[str], optional): The Microsoft Client ID. Defaults to None.
generic_client_id (Optional[str], optional): The Generic Client ID. Defaults to None.
request: Optional FastAPI request, used to drive the ``Secure``
attribute on the ``litellm_oauth_state`` CSRF cookie.
Returns:
RedirectResponse: The redirect response from the SSO provider.
@@ -2266,6 +2294,7 @@ class SSOAuthenticationHandler:
generic_sso=generic_sso,
state=state,
generic_authorization_endpoint=generic_authorization_endpoint,
request=request,
)
raise ValueError(
"Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso"
@@ -2276,6 +2305,7 @@ class SSOAuthenticationHandler:
generic_sso: Any,
state: Optional[str] = None,
generic_authorization_endpoint: Optional[str] = None,
request: Optional[Request] = None,
) -> Optional[RedirectResponse]:
"""
Get the redirect response for Generic SSO
@@ -2285,10 +2315,13 @@ class SSOAuthenticationHandler:
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
with generic_sso:
# TODO: state should be a random string and added to the user session with cookie
# or a cryptographicly signed state that we can verify stateless
# For simplification we are using a static state, this is not perfect but some
# SSO providers do not allow stateless verification
# State is bound to the caller's browser via a ``litellm_oauth_state``
# HttpOnly cookie set on the redirect response below; the SSO
# callback validates the URL ``state`` against that cookie before
# completing the PKCE token exchange. Without this binding, an
# attacker who pre-mints a state + a cached PKCE verifier can hand
# the link to a victim and capture the resulting access token
# (Login CSRF / token theft).
(
redirect_params,
code_verifier,
@@ -2355,6 +2388,31 @@ class SSOAuthenticationHandler:
# Update the redirect response
redirect_response.headers["location"] = new_url
# Bind state to the user's browser session. The /callback
# handler validates the URL ``state`` against this cookie via
# ``secrets.compare_digest`` before exchanging the PKCE
# code_verifier. Only set the cookie when PKCE is in use
# (i.e. inside this ``code_verifier`` branch) so two
# concurrent SSO sessions — one PKCE, one plain — cannot
# overwrite each other's state cookie.
state_value = redirect_params.get("state")
if state_value and redirect_response is not None:
# Production-safe default: require HTTPS for the
# CSRF-protection cookie unless we can prove the
# incoming request is HTTP (local dev). Without
# ``Secure`` the cookie is sent over plain HTTP,
# letting a network observer read and replay the
# state value and bypass this protection.
secure_flag = request is None or request.url.scheme == "https"
redirect_response.set_cookie(
key="litellm_oauth_state",
value=state_value,
max_age=600,
httponly=True,
samesite="lax",
secure=secure_flag,
)
return redirect_response
@staticmethod
@@ -3972,6 +4030,7 @@ async def debug_sso_login(request: Request):
microsoft_client_id=microsoft_client_id,
google_client_id=google_client_id,
generic_client_id=generic_client_id,
request=request,
)
+19
View File
@@ -6,6 +6,18 @@ from fastapi import HTTPException, status
import litellm
from litellm.proxy._types import UserAPIKeyAuth
# Router-internal mock_testing_* flag names — kept in sync with
# ``litellm.types.router.MockRouterTestingParams`` by the test
# ``test_mock_testing_kwarg_names_matches_dataclass``. Hardcoding (rather
# than deriving via ``dataclasses.fields(MockRouterTestingParams)`` at
# import time) avoids a cyclic import: ``litellm.types.router`` imports
# back into proxy modules before this module finishes loading.
_MOCK_TESTING_KWARG_NAMES: tuple = (
"mock_testing_fallbacks",
"mock_testing_context_fallbacks",
"mock_testing_content_policy_fallbacks",
)
if TYPE_CHECKING:
from litellm.router import Router as _Router
@@ -322,6 +334,13 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
"""
await add_shared_session_to_data(data)
# Strip router-internal mock_testing_* flags. Combined with an
# unauthorized fallback in ``router_settings_override`` they let a
# caller deterministically execute requests against restricted
# models. VERIA-44.
for _key in _MOCK_TESTING_KWARG_NAMES:
data.pop(_key, None)
team_id = get_team_id_from_data(data)
router_model_names = llm_router.model_names if llm_router is not None else []
@@ -26,7 +26,7 @@ class TestAnthropicTokenCounter(BaseTokenCounterTest):
return AnthropicTokenCounter()
def get_test_model(self) -> str:
return "claude-sonnet-4-20250514"
return "claude-haiku-4-5-20251001"
def get_test_messages(self) -> List[Dict[str, Any]]:
return [{"role": "user", "content": "Hello, how are you today?"}]
+1 -1
View File
@@ -158,7 +158,7 @@ def test_aaparallel_function_call(model):
@pytest.mark.parametrize(
"model",
[
"anthropic/claude-4-sonnet-20250514",
"anthropic/claude-haiku-4-5-20251001",
"bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
],
)
+2 -2
View File
@@ -1270,7 +1270,7 @@ def test_bedrock_claude_3_streaming():
@pytest.mark.parametrize(
"model",
[
"claude-4-sonnet-20250514",
"claude-haiku-4-5-20251001",
"cohere.command-r-plus-v1:0", # bedrock
"gpt-3.5-turbo",
],
@@ -2696,7 +2696,7 @@ def test_completion_claude_3_function_call_with_streaming():
try:
# test without max tokens
response = completion(
model="claude-4-sonnet-20250514",
model="claude-haiku-4-5-20251001",
messages=messages,
tools=tools,
tool_choice="required",
@@ -54,7 +54,7 @@ class BaseAnthropicMessagesTest(ABC):
print("making request to anthropic passthrough with thinking")
client = self.get_client()
response = client.messages.create(
model="claude-4-sonnet-20250514",
model="claude-haiku-4-5-20251001",
max_tokens=20000,
thinking={"type": "enabled", "budget_tokens": 16000},
messages=[
@@ -75,7 +75,7 @@ class BaseAnthropicMessagesTest(ABC):
collected_response = []
client = self.get_client()
with client.messages.stream(
model="claude-4-sonnet-20250514",
model="claude-haiku-4-5-20251001",
max_tokens=20000,
thinking={"type": "enabled", "budget_tokens": 16000},
messages=[
@@ -333,7 +333,7 @@ async def test_anthropic_messages_streaming_cost_injection():
}
payload = {
"model": "claude-4-sonnet-20250514",
"model": "claude-haiku-4-5-20251001",
"max_tokens": 10,
"stream": True,
"messages": [{"role": "user", "content": "Say 'Hi'"}],
+317
View File
@@ -1153,3 +1153,320 @@ async def test_can_key_call_model_via_access_group_ids():
valid_token=user_api_key_object,
llm_router=router,
)
# ---------------------------------------------------------------------------
# _key_access_group_grants_model (key access group overriding team restriction)
# ---------------------------------------------------------------------------
def _patch_proxy_server_globals():
"""Patch proxy_server's prisma_client and user_api_key_cache to non-None mocks
so the helper's None-guard doesn't short-circuit. The actual values don't
matter because get_access_object is patched separately to return fixtures."""
from unittest.mock import MagicMock, patch
return [
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
]
def _fake_access_group(
access_group_id: str,
access_model_names=None,
assigned_team_ids=None,
assigned_key_ids=None,
):
from litellm.proxy._types import LiteLLM_AccessGroupTable
return LiteLLM_AccessGroupTable(
access_group_id=access_group_id,
access_group_name=access_group_id,
access_model_names=access_model_names or [],
assigned_team_ids=assigned_team_ids or [],
assigned_key_ids=assigned_key_ids or [],
)
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_team_authorized():
"""Group's assigned_team_ids includes the key's team and grants the model → True.
This is the happy path equivalent of Andres's report: admin creates an
access group with assigned_team_ids=[team-a], grants claude-haiku-4-5,
attaches it to a key on team-a. Override fires.
"""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="test-token",
models=[],
access_group_ids=["premium-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[], # deliberately not synced — the access group itself authorizes
)
fake_ag = _fake_access_group(
access_group_id="premium-group",
access_model_names=["claude-haiku-4-5"],
assigned_team_ids=["team-a"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is True
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_key_directly_authorized():
"""Group's assigned_key_ids includes the key's token and grants the model → True.
Per-key authorization path: an admin scopes a group directly to a key
(assigned_key_ids) without listing the team.
"""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="test-token-hashed",
models=[],
access_group_ids=["per-key-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
fake_ag = _fake_access_group(
access_group_id="per-key-group",
access_model_names=["claude-haiku-4-5"],
assigned_team_ids=[],
assigned_key_ids=["test-token-hashed"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is True
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_key_has_no_groups():
"""Key with no access_group_ids → False (early return, no DB read)."""
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="test-token",
models=[],
access_group_ids=[],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=["any-group"],
)
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_group_does_not_cover_model():
"""Group authorizes the team but does not grant the requested model → False."""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="test-token",
models=[],
access_group_ids=["basic-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
fake_ag = _fake_access_group(
access_group_id="basic-group",
access_model_names=["gpt-4o-mini"],
assigned_team_ids=["team-a"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_group_authorizes_neither():
"""
Bypass regression test: a team member sets a foreign access group on their
key. The group grants the requested model but its assigned_team_ids /
assigned_key_ids do not include this caller's team or token. Override is
denied the team's 401 propagates.
"""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="team-a-token",
models=[],
access_group_ids=["team-b-premium"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
fake_ag = _fake_access_group(
access_group_id="team-b-premium",
access_model_names=["claude-opus-4-5"],
assigned_team_ids=["team-b"],
assigned_key_ids=["team-b-token"],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-opus-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
finally:
for p in patches:
p.stop()
@pytest.mark.asyncio
async def test_key_access_group_grants_model_when_get_access_object_raises():
"""Group lookup failure (404, network, etc.) is treated as no authorization."""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import _key_access_group_grants_model
valid_token = UserAPIKeyAuth(
token="test-token",
models=[],
access_group_ids=["missing-group"],
team_id="team-a",
)
team_object = LiteLLM_TeamTable(
team_id="team-a",
models=["mock-success"],
access_group_ids=[],
)
patches = _patch_proxy_server_globals() + [
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
side_effect=Exception("not found"),
),
]
for p in patches:
p.start()
try:
assert (
await _key_access_group_grants_model(
model="claude-haiku-4-5",
valid_token=valid_token,
team_object=team_object,
llm_router=None,
)
is False
)
finally:
for p in patches:
p.stop()
@@ -259,6 +259,189 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase):
), "Existing LoggerProvider should be respected and not overridden"
class TestOpenTelemetryDualHandlerIsolation(unittest.TestCase):
"""Two OpenTelemetry handlers coexisting via skip_set_global=True
must each get their own provider for every signal (tracer/meter/logger)."""
@staticmethod
def _wire_span_processor(exporter):
"""Context manager: while active, the next OpenTelemetry instance
wires its TracerProvider to `exporter`."""
return patch.object(
OpenTelemetry,
"_get_span_processor",
lambda self, dynamic_headers=None: SimpleSpanProcessor(exporter),
)
def test_skip_set_global_creates_isolated_tracer_provider(self):
from opentelemetry.sdk.trace import TracerProvider as SDKTracerProvider
fake_existing = SDKTracerProvider()
own_exporter = InMemorySpanExporter()
cfg = OpenTelemetryConfig(
exporter="console", service_name="iso-test", skip_set_global=True
)
with (
patch.object(trace, "get_tracer_provider", return_value=fake_existing),
patch.object(trace, "set_tracer_provider") as mock_set,
self._wire_span_processor(own_exporter),
):
handler = OpenTelemetry(config=cfg)
self.assertIsNot(handler._tracer_provider, fake_existing)
mock_set.assert_not_called()
handler.tracer.start_span("isolation_check").end()
handler._tracer_provider.force_flush(2000)
self.assertEqual(
[s.name for s in own_exporter.get_finished_spans()],
["isolation_check"],
)
def test_skip_set_global_via_callback_name_back_compat(self):
from opentelemetry.sdk.trace import TracerProvider as SDKTracerProvider
fake_existing = SDKTracerProvider()
cfg = OpenTelemetryConfig(exporter="console", service_name="lf-back-compat")
with (
patch.object(trace, "get_tracer_provider", return_value=fake_existing),
patch.object(trace, "set_tracer_provider"),
self._wire_span_processor(InMemorySpanExporter()),
):
handler = OpenTelemetry(config=cfg, callback_name="langfuse_otel")
self.assertIsNot(handler._tracer_provider, fake_existing)
def test_default_behavior_reuses_existing_sdk_tracer_provider(self):
from opentelemetry.sdk.trace import TracerProvider as SDKTracerProvider
fake_existing = SDKTracerProvider()
with patch.object(trace, "get_tracer_provider", return_value=fake_existing):
handler = OpenTelemetry(config=OpenTelemetryConfig(service_name="shared"))
self.assertIs(handler._tracer_provider, fake_existing)
def test_skip_set_global_creates_isolated_meter_provider(self):
from opentelemetry import metrics
from opentelemetry.sdk.metrics import MeterProvider as SDKMeterProvider
fake_existing = SDKMeterProvider()
cfg = OpenTelemetryConfig(
exporter="console",
service_name="meter-iso-test",
enable_metrics=True,
skip_set_global=True,
)
with (
patch.object(metrics, "get_meter_provider", return_value=fake_existing),
patch.object(metrics, "set_meter_provider") as mock_set,
self._wire_span_processor(InMemorySpanExporter()),
):
handler = OpenTelemetry(config=cfg)
self.assertIsNot(handler._meter_provider, fake_existing)
mock_set.assert_not_called()
def test_skip_set_global_creates_isolated_logger_provider(self):
from opentelemetry import _logs
from opentelemetry.sdk._logs import LoggerProvider as SDKLoggerProvider
fake_existing = SDKLoggerProvider()
cfg = OpenTelemetryConfig(
exporter="console",
service_name="logger-iso-test",
enable_events=True,
skip_set_global=True,
)
with (
patch.object(_logs, "get_logger_provider", return_value=fake_existing),
patch.object(_logs, "set_logger_provider") as mock_set,
self._wire_span_processor(InMemorySpanExporter()),
):
handler = OpenTelemetry(config=cfg)
self.assertIsNot(handler._logger_provider, fake_existing)
mock_set.assert_not_called()
def test_emitted_logs_route_to_isolated_logger_provider(self):
# End-to-end: emitted logs land in the handler's private LoggerProvider,
# not the global one. Guards against get_logger() bypassing self._logger_provider.
from opentelemetry import _logs
from opentelemetry.sdk._logs import LoggerProvider as SDKLoggerProvider
global_exporter = InMemoryLogExporter()
fake_existing = SDKLoggerProvider()
fake_existing.add_log_record_processor(
SimpleLogRecordProcessor(global_exporter)
)
private_exporter = InMemoryLogExporter()
cfg = OpenTelemetryConfig(
exporter="console",
service_name="logger-emit-test",
enable_events=True,
skip_set_global=True,
)
with (
patch.object(_logs, "get_logger_provider", return_value=fake_existing),
patch.object(_logs, "set_logger_provider"),
patch.object(
OpenTelemetry, "_get_log_exporter", return_value=private_exporter
),
self._wire_span_processor(InMemorySpanExporter()),
):
handler = OpenTelemetry(config=cfg)
span = handler.tracer.start_span("emit-test")
handler._emit_semantic_logs(
kwargs={"messages": [{"role": "user", "content": "hi"}]},
response_obj={"choices": []},
span=span,
)
span.end()
handler._logger_provider.force_flush(2000)
self.assertGreater(len(private_exporter.get_finished_logs()), 0)
self.assertEqual(len(global_exporter.get_finished_logs()), 0)
def test_two_handlers_each_receive_their_own_spans(self):
# Handler A gets explicit injection (production-ish: claims the global).
exporter_a = InMemorySpanExporter()
provider_a = TracerProvider()
provider_a.add_span_processor(SimpleSpanProcessor(exporter_a))
handler_a = OpenTelemetry(
config=OpenTelemetryConfig(service_name="handler-a"),
tracer_provider=provider_a,
)
# Handler B comes along with the global appearing to be A's provider.
exporter_b = InMemorySpanExporter()
cfg_b = OpenTelemetryConfig(
exporter="console", service_name="handler-b", skip_set_global=True
)
with (
patch.object(trace, "get_tracer_provider", return_value=provider_a),
patch.object(trace, "set_tracer_provider"),
self._wire_span_processor(exporter_b),
):
handler_b = OpenTelemetry(config=cfg_b)
self.assertIsNot(handler_a._tracer_provider, handler_b._tracer_provider)
handler_a.tracer.start_span("from_handler_a").end()
handler_b.tracer.start_span("from_handler_b").end()
provider_a.force_flush(2000)
handler_b._tracer_provider.force_flush(2000)
self.assertEqual(
sorted(s.name for s in exporter_a.get_finished_spans()),
["from_handler_a"],
)
self.assertEqual(
sorted(s.name for s in exporter_b.get_finished_spans()),
["from_handler_b"],
)
class TestOpenTelemetry(unittest.TestCase):
POLL_INTERVAL = 0.05
POLL_TIMEOUT = 2.0
@@ -0,0 +1,150 @@
"""
Tests for VERIA-53: PromQL string-literal quoting in
``get_daily_spend_from_prometheus``.
PromQL string literals follow Go's escape rules
(https://prometheus.io/docs/prometheus/latest/querying/basics/). JSON's
quoting is a strict subset of Go's, so ``json.dumps`` produces a literal
Prometheus parses identically.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
def test_quote_safe_input_round_trips():
from litellm.integrations.prometheus_helpers.prometheus_api import (
_quote_promql_string_literal,
)
assert _quote_promql_string_literal("sk-abc123") == '"sk-abc123"'
assert _quote_promql_string_literal("hash:deadbeef") == '"hash:deadbeef"'
def test_quote_escapes_double_quote():
from litellm.integrations.prometheus_helpers.prometheus_api import (
_quote_promql_string_literal,
)
# A bare double quote would otherwise terminate the label matcher and
# let the attacker append `, foo="..."} or sum(...)`.
assert _quote_promql_string_literal('hello"injected') == '"hello\\"injected"'
def test_quote_escapes_backslash():
from litellm.integrations.prometheus_helpers.prometheus_api import (
_quote_promql_string_literal,
)
assert _quote_promql_string_literal('a\\"b') == '"a\\\\\\"b"'
def test_quote_escapes_newlines_and_control_chars():
"""Beyond the security minimum, the canonical Go/JSON escape also
handles control characters that would otherwise produce an invalid
PromQL string literal."""
from litellm.integrations.prometheus_helpers.prometheus_api import (
_quote_promql_string_literal,
)
assert _quote_promql_string_literal("a\nb") == '"a\\nb"'
assert _quote_promql_string_literal("a\tb") == '"a\\tb"'
assert _quote_promql_string_literal("a\rb") == '"a\\rb"'
@pytest.mark.asyncio
async def test_get_daily_spend_does_not_pass_raw_quote_into_query():
from litellm.integrations.prometheus_helpers import prometheus_api
captured = {}
class _FakeResponse:
def json(self):
return {"data": {"result": []}}
async def _capture(url, params):
captured["url"] = url
captured["params"] = params
return _FakeResponse()
fake_client = MagicMock()
fake_client.get = AsyncMock(side_effect=_capture)
with patch.object(prometheus_api, "PROMETHEUS_URL", "http://prom.example"):
with patch.object(prometheus_api, "async_http_handler", fake_client):
await prometheus_api.get_daily_spend_from_prometheus(
api_key='sk-victim"} or sum(other_metric{a="b'
)
rendered_query = captured["params"]["query"]
# The legitimate matcher framing must still be intact: one outer
# `delta()` window, one inner `hashed_api_key="..."` matcher.
assert rendered_query.startswith(
'sum(delta(litellm_spend_metric_total{hashed_api_key="'
)
assert rendered_query.endswith('"}[1d]))')
# Every injected `"` from the attacker payload appears as `\"` so the
# PromQL parser treats them as literal characters inside the matcher
# value, never as the terminator that would let the rest parse as
# PromQL syntax.
inner = rendered_query[
len('sum(delta(litellm_spend_metric_total{hashed_api_key="') : -len('"}[1d]))')
]
assert '"' not in inner.replace('\\"', "")
@pytest.mark.asyncio
async def test_get_daily_spend_with_no_api_key_uses_unfiltered_query():
from litellm.integrations.prometheus_helpers import prometheus_api
captured = {}
class _FakeResponse:
def json(self):
return {"data": {"result": []}}
async def _capture(url, params):
captured["params"] = params
return _FakeResponse()
fake_client = MagicMock()
fake_client.get = AsyncMock(side_effect=_capture)
with patch.object(prometheus_api, "PROMETHEUS_URL", "http://prom.example"):
with patch.object(prometheus_api, "async_http_handler", fake_client):
await prometheus_api.get_daily_spend_from_prometheus(api_key=None)
assert captured["params"]["query"] == "sum(delta(litellm_spend_metric_total[1d]))"
@pytest.mark.asyncio
async def test_get_daily_spend_legitimate_hashed_key_unchanged():
"""A normal hex hashed_api_key flows through `json.dumps` as itself
plus the surrounding quotes no spurious escaping that would break
real lookups."""
from litellm.integrations.prometheus_helpers import prometheus_api
captured = {}
class _FakeResponse:
def json(self):
return {"data": {"result": []}}
async def _capture(url, params):
captured["params"] = params
return _FakeResponse()
fake_client = MagicMock()
fake_client.get = AsyncMock(side_effect=_capture)
legit_key = "a" * 64 # 64-char hex-ish hashed key
with patch.object(prometheus_api, "PROMETHEUS_URL", "http://prom.example"):
with patch.object(prometheus_api, "async_http_handler", fake_client):
await prometheus_api.get_daily_spend_from_prometheus(api_key=legit_key)
assert (
captured["params"]["query"]
== f'sum(delta(litellm_spend_metric_total{{hashed_api_key="{legit_key}"}}[1d]))'
)
@@ -2567,3 +2567,116 @@ async def test_auth_builder_single_team_fallback_membership_error_skips_no_raise
assert result["team_membership"] is None
mock_get_team.assert_called()
mock_get_membership.assert_called_once()
# ---------------------------------------------------------------------------
# JWTHandler._build_decode_kwargs — VERIA-27 (audience + issuer verification)
# ---------------------------------------------------------------------------
@pytest.fixture(autouse=False)
def _reset_unscoped_warning_flag():
"""Reset the once-per-process warning sentinel so each test sees a fresh
state."""
JWTHandler._unscoped_jwt_warning_emitted = False
yield
JWTHandler._unscoped_jwt_warning_emitted = False
def test_build_decode_kwargs_no_env_disables_both_verifications(
monkeypatch, _reset_unscoped_warning_flag
):
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_ISSUER", raising=False)
kwargs = JWTHandler._build_decode_kwargs()
assert kwargs["audience"] is None
assert kwargs["issuer"] is None
assert kwargs["options"] == {"verify_aud": False, "verify_iss": False}
def test_build_decode_kwargs_audience_only_enables_aud_verification(
monkeypatch, _reset_unscoped_warning_flag
):
monkeypatch.setenv("JWT_AUDIENCE", "my-proxy")
monkeypatch.delenv("JWT_ISSUER", raising=False)
kwargs = JWTHandler._build_decode_kwargs()
assert kwargs["audience"] == "my-proxy"
assert kwargs["issuer"] is None
# verify_aud not in options means PyJWT will verify audience
assert kwargs["options"] == {"verify_iss": False}
def test_build_decode_kwargs_issuer_only_enables_iss_verification(
monkeypatch, _reset_unscoped_warning_flag
):
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/")
kwargs = JWTHandler._build_decode_kwargs()
assert kwargs["audience"] is None
assert kwargs["issuer"] == "https://idp.example.com/"
assert kwargs["options"] == {"verify_aud": False}
def test_build_decode_kwargs_both_set_enables_full_verification(
monkeypatch, _reset_unscoped_warning_flag
):
monkeypatch.setenv("JWT_AUDIENCE", "my-proxy")
monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/")
kwargs = JWTHandler._build_decode_kwargs()
assert kwargs["audience"] == "my-proxy"
assert kwargs["issuer"] == "https://idp.example.com/"
# No verification opt-outs — PyJWT verifies both claims by default.
assert kwargs["options"] is None
def test_build_decode_kwargs_warns_once_when_unscoped(
monkeypatch, _reset_unscoped_warning_flag, caplog
):
"""The warning about unscoped JWT auth should fire on the first call but
not on every subsequent decode."""
import logging
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_ISSUER", raising=False)
caplog.set_level(logging.WARNING)
JWTHandler._build_decode_kwargs()
JWTHandler._build_decode_kwargs()
JWTHandler._build_decode_kwargs()
matching = [
r
for r in caplog.records
if "JWT auth is enabled" in r.getMessage()
and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()
]
assert (
len(matching) == 1
), f"Expected exactly one warning across 3 calls, got {len(matching)}"
def test_build_decode_kwargs_no_warning_when_scoped(
monkeypatch, _reset_unscoped_warning_flag, caplog
):
import logging
monkeypatch.setenv("JWT_AUDIENCE", "my-proxy")
monkeypatch.delenv("JWT_ISSUER", raising=False)
caplog.set_level(logging.WARNING)
JWTHandler._build_decode_kwargs()
matching = [
r
for r in caplog.records
if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()
]
assert matching == []
@@ -0,0 +1,211 @@
"""
Regression tests for the OAuth2-proxy header-forgery fix
(GHSA-5c3m-qffq-4r9m).
The hook reads HTTP request headers per ``oauth2_config_mappings`` and
constructs a ``UserAPIKeyAuth`` from them. The fix has two parts:
1. Only requests from configured trusted proxy CIDR ranges may provide
identity headers.
2. Only identity fields may be mapped from those headers. Without the
identity-only allowlist any field could be mapped including
``user_role``, which Pydantic coerces from the string
``"proxy_admin"`` into ``LitellmUserRoles.PROXY_ADMIN``.
"""
import os
import sys
import pytest
from fastapi import Request
from starlette.datastructures import Headers
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.oauth2_proxy_hook import (
ALLOWED_OAUTH2_PROXY_FIELDS,
handle_oauth2_proxy_request,
)
def _request_with_headers(headers: dict, *, client_host: str = "127.0.0.1") -> Request:
scope = {
"type": "http",
"client": (client_host, 12345),
"headers": [(k.lower().encode(), v.encode()) for k, v in headers.items()],
}
request = Request(scope=scope)
request._headers = Headers(headers)
return request
@pytest.fixture
def configure_proxy(monkeypatch):
"""
Yields a callable that sets ``oauth2_config_mappings`` and
``trusted_proxy_ranges`` on the proxy_server module for the duration
of one test. Defaults to a single identity mapping and localhost as
a trusted proxy.
"""
import litellm.proxy.proxy_server as proxy_server
def _configure(*, mappings=None, trusted_proxy_ranges=("127.0.0.1/32",)):
if mappings is None:
mappings = {"user_id": "x-user-id"}
settings = {
"oauth2_config_mappings": mappings,
"trusted_proxy_ranges": trusted_proxy_ranges,
}
monkeypatch.setattr(
proxy_server,
"general_settings",
settings,
raising=False,
)
return _configure
@pytest.mark.asyncio
async def test_returns_auth_for_simple_user_id_mapping(configure_proxy):
configure_proxy()
request = _request_with_headers({"x-user-id": "alice"})
auth = await handle_oauth2_proxy_request(request)
assert auth.user_id == "alice"
assert auth.user_role is None
@pytest.mark.asyncio
async def test_rejects_identity_headers_without_trusted_proxy_ranges(configure_proxy):
configure_proxy(trusted_proxy_ranges=None)
request = _request_with_headers({"x-user-id": "alice"})
with pytest.raises(ValueError, match="trusted_proxy_ranges"):
await handle_oauth2_proxy_request(request)
@pytest.mark.asyncio
async def test_rejects_identity_headers_from_untrusted_direct_client(configure_proxy):
configure_proxy(trusted_proxy_ranges=["10.0.0.0/24"])
request = _request_with_headers({"x-user-id": "alice"}, client_host="203.0.113.10")
with pytest.raises(ValueError, match="not trusted"):
await handle_oauth2_proxy_request(request)
@pytest.mark.parametrize(
"privileged_field",
[
# The GHSA-5c3m-qffq-4r9m primary privesc field.
"user_role",
# Key-level enforcement bypass shapes.
"api_key",
"token",
"permissions",
"allowed_routes",
"max_budget",
"spend",
"tpm_limit",
"rpm_limit",
"model_max_budget",
"metadata",
# User-level enforcement bypass — flagged by Greptile as a denylist gap.
"user_max_budget",
"user_tpm_limit",
"user_rpm_limit",
"user_spend",
# Team / org / end-user / region — same class, all denied by the
# identity-only allowlist.
"team_max_budget",
"team_spend",
"team_member_tpm_limit",
"organization_max_budget",
"organization_tpm_limit",
"end_user_max_budget",
"allowed_model_region",
# Anything not on ALLOWED_OAUTH2_PROXY_FIELDS is blocked, even
# fabricated field names admins might try.
"definitely_not_a_real_field",
],
)
@pytest.mark.asyncio
async def test_refuses_to_map_non_identity_fields(configure_proxy, privileged_field):
# GHSA-5c3m-qffq-4r9m attack shape: admin maps a privileged field
# to a header and a caller forges the value. The allowlist rejects
# any non-identity mapping at request time, regardless of whether
# the field ever appeared on a denylist — which is the whole reason
# we use an allowlist instead.
configure_proxy(mappings={privileged_field: f"x-{privileged_field}"})
request = _request_with_headers({f"x-{privileged_field}": "proxy_admin"})
with pytest.raises(ValueError) as exc:
await handle_oauth2_proxy_request(request)
assert privileged_field in str(exc.value)
@pytest.mark.parametrize("identity_field", sorted(ALLOWED_OAUTH2_PROXY_FIELDS))
def test_allowlist_is_identity_only(identity_field):
# Lock in the allowlist's intent: only identity-assertion fields are
# safe to populate from a header. If anyone proposes adding budget /
# spend / role / permission to ``ALLOWED_OAUTH2_PROXY_FIELDS``, this
# assertion forces them to update the test deliberately.
assert identity_field in {
"user_id",
"user_email",
"team_id",
"team_alias",
"org_id",
"models",
}
@pytest.mark.asyncio
async def test_user_role_header_forgery_attack_is_blocked(configure_proxy):
# End-to-end form of the privesc: with ``user_role`` mapped, the
# forged ``X-User-Role: proxy_admin`` header would have produced
# a ``UserAPIKeyAuth(user_role=PROXY_ADMIN)``. Now rejected before
# any auth object is constructed.
configure_proxy(
mappings={"user_id": "x-user-id", "user_role": "x-user-role"},
)
request = _request_with_headers(
{
"x-user-id": "attacker",
"x-user-role": LitellmUserRoles.PROXY_ADMIN.value,
}
)
with pytest.raises(ValueError, match="user_role"):
await handle_oauth2_proxy_request(request)
@pytest.mark.asyncio
async def test_safe_fields_still_pass_through(configure_proxy):
# The documented use case for OAuth2 proxy auth: identity assertion
# from a trusted upstream. Must remain unaffected by the denylist.
configure_proxy(
mappings={
"user_id": "x-user-id",
"user_email": "x-user-email",
"team_id": "x-team-id",
"models": "x-models",
},
)
request = _request_with_headers(
{
"x-user-id": "alice",
"x-user-email": "alice@example.com",
"x-team-id": "team-corp",
"x-models": "gpt-4, gpt-3.5-turbo",
}
)
auth = await handle_oauth2_proxy_request(request)
assert auth.user_id == "alice"
assert auth.user_email == "alice@example.com"
assert auth.team_id == "team-corp"
assert auth.models == ["gpt-4", "gpt-3.5-turbo"]
@@ -0,0 +1,237 @@
"""
VERIA-44: ``router_settings_override.fallbacks`` must be validated
against the API key's model allowlist at auth time. Without this, the
override is promoted to per-request kwargs after auth and lets a caller
execute requests against models their API key cannot call.
"""
from typing import List
from unittest.mock import AsyncMock, patch
import pytest
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import (
_enforce_key_and_fallback_model_access,
iter_router_fallback_model_names,
)
def _key_with_models(models: List[str]) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="hashed",
user_id="caller",
user_role=LitellmUserRoles.INTERNAL_USER,
models=models,
)
# ── iter_router_fallback_model_names ─────────────────────────────────────────
def testiter_router_fallback_model_names_router_config_shape():
"""Router-config shape: ``[{primary: [fallback_list]}]``."""
assert list(
iter_router_fallback_model_names(
[{"gpt-3.5-turbo": ["gpt-4", "claude-3"]}, {"gpt-4o": ["o1"]}]
)
) == ["gpt-4", "claude-3", "o1"]
def testiter_router_fallback_model_names_simple_string_shape():
"""Simple top-level shape: list of strings."""
assert list(iter_router_fallback_model_names(["gpt-4", "claude-3"])) == [
"gpt-4",
"claude-3",
]
def testiter_router_fallback_model_names_client_side_shape():
"""ClientSideFallbackModel shape: ``[{"model": "..."}]``."""
assert list(
iter_router_fallback_model_names([{"model": "gpt-4"}, {"model": "claude-3"}])
) == ["gpt-4", "claude-3"]
def testiter_router_fallback_model_names_empty_or_none():
assert list(iter_router_fallback_model_names(None)) == []
assert list(iter_router_fallback_model_names([])) == []
assert list(iter_router_fallback_model_names("not a list")) == []
# ── _enforce_key_and_fallback_model_access ────────────────────────────────────
@pytest.mark.asyncio
async def test_router_override_fallbacks_validated_against_key_allowlist():
"""A fallback nested inside ``router_settings_override`` is validated
against the API key's allowed models — not just the top-level
``fallbacks`` field."""
valid_token = _key_with_models(["gpt-3.5-turbo"])
request_data = {
"model": "gpt-3.5-turbo",
"router_settings_override": {
"fallbacks": [{"gpt-3.5-turbo": ["unauthorized-model"]}],
},
}
seen_models: List[str] = []
async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router):
seen_models.append(model)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
side_effect=fake_can_key_call_model,
),
patch(
"litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model",
new=AsyncMock(),
),
):
await _enforce_key_and_fallback_model_access(
valid_token=valid_token,
request_data=request_data,
route="/v1/chat/completions",
request=None,
llm_model_list=None,
llm_router=None,
)
# Both the primary model and the override-nested fallback must be
# checked against the API key's allowlist.
assert "gpt-3.5-turbo" in seen_models
assert "unauthorized-model" in seen_models
@pytest.mark.asyncio
@pytest.mark.parametrize(
"fallback_field",
[
"fallbacks",
"context_window_fallbacks",
"content_policy_fallbacks",
],
)
async def test_router_override_all_fallback_fields_validated(fallback_field):
"""All three fallback fields the router accepts as per-request kwargs
are validated context_window_fallbacks and content_policy_fallbacks
are promoted in route_llm_request.py too."""
valid_token = _key_with_models(["gpt-3.5-turbo"])
request_data = {
"model": "gpt-3.5-turbo",
"router_settings_override": {
fallback_field: [{"gpt-3.5-turbo": ["smuggled-model"]}],
},
}
seen: List[str] = []
async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router):
seen.append(model)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
side_effect=fake_can_key_call_model,
),
patch(
"litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model",
new=AsyncMock(),
),
):
await _enforce_key_and_fallback_model_access(
valid_token=valid_token,
request_data=request_data,
route="/v1/chat/completions",
request=None,
llm_model_list=None,
llm_router=None,
)
assert "smuggled-model" in seen
@pytest.mark.asyncio
@pytest.mark.parametrize(
"fallback_field",
[
"fallbacks",
"context_window_fallbacks",
"content_policy_fallbacks",
],
)
async def test_top_level_fallback_fields_validated(fallback_field):
"""All three top-level fallback fields are forwarded to the router as
per-request kwargs, so all three must be validated against the API
key's allowlist. Greptile P1 follow-up: previously only the
``fallbacks`` field was walked at the top level."""
valid_token = _key_with_models(["gpt-3.5-turbo"])
request_data = {
"model": "gpt-3.5-turbo",
fallback_field: [{"gpt-3.5-turbo": ["top-level-smuggled"]}],
}
seen: List[str] = []
async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router):
seen.append(model)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
side_effect=fake_can_key_call_model,
),
patch(
"litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model",
new=AsyncMock(),
),
):
await _enforce_key_and_fallback_model_access(
valid_token=valid_token,
request_data=request_data,
route="/v1/chat/completions",
request=None,
llm_model_list=None,
llm_router=None,
)
assert "top-level-smuggled" in seen
@pytest.mark.asyncio
async def test_router_override_without_fallbacks_does_not_break_auth():
"""``router_settings_override`` set without any fallback fields is a
no-op for the auth check only the primary model is validated."""
valid_token = _key_with_models(["gpt-3.5-turbo"])
request_data = {
"model": "gpt-3.5-turbo",
"router_settings_override": {"num_retries": 3, "timeout": 30},
}
seen: List[str] = []
async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router):
seen.append(model)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
side_effect=fake_can_key_call_model,
),
patch(
"litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model",
new=AsyncMock(),
),
):
await _enforce_key_and_fallback_model_access(
valid_token=valid_token,
request_data=request_data,
route="/v1/chat/completions",
request=None,
llm_model_list=None,
llm_router=None,
)
assert seen == ["gpt-3.5-turbo"]
@@ -0,0 +1,195 @@
"""
Unit tests for the VERIA-55 fixes:
- Project update permission must be evaluated against the project's *current*
team, not a team supplied in the request body.
- Key update may not assign a key to an organization the caller is not a
member of.
"""
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
# ---------------------------------------------------------------------------
# /project/update — _check_user_permission_for_project
# ---------------------------------------------------------------------------
def _make_prisma_with_team(team_id: str, admins: list):
prisma = MagicMock()
team_row = MagicMock()
team_row.team_id = team_id
team_row.admins = admins
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
return prisma
@pytest.mark.asyncio
async def test_project_perm_check_uses_current_team_not_caller_supplied():
"""The permission check must look at the project's existing team. Even
if the caller is admin of an unrelated team, they must not pass when no
explicit team_object is forced through."""
from enterprise.litellm_enterprise.proxy.management_endpoints.project_endpoints import (
_check_user_permission_for_project,
)
# Project lives on team-A, caller is admin only of team-B.
prisma = _make_prisma_with_team(team_id="team-A", admins=["alice"])
caller = UserAPIKeyAuth(
user_id="bob",
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
has_perm = await _check_user_permission_for_project(
user_api_key_dict=caller,
team_id="team-A",
prisma_client=prisma,
)
assert has_perm is False
prisma.db.litellm_teamtable.find_unique.assert_awaited_once()
@pytest.mark.asyncio
async def test_project_perm_check_allows_team_admin_of_existing_team():
from enterprise.litellm_enterprise.proxy.management_endpoints.project_endpoints import (
_check_user_permission_for_project,
)
prisma = _make_prisma_with_team(team_id="team-A", admins=["alice"])
alice = UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
has_perm = await _check_user_permission_for_project(
user_api_key_dict=alice,
team_id="team-A",
prisma_client=prisma,
)
assert has_perm is True
@pytest.mark.asyncio
async def test_project_perm_check_proxy_admin_always_allowed():
from enterprise.litellm_enterprise.proxy.management_endpoints.project_endpoints import (
_check_user_permission_for_project,
)
prisma = MagicMock()
admin = UserAPIKeyAuth(
user_id="root",
user_role=LitellmUserRoles.PROXY_ADMIN.value,
)
has_perm = await _check_user_permission_for_project(
user_api_key_dict=admin,
team_id="team-A",
prisma_client=prisma,
)
assert has_perm is True
# Admin shortcut should not even hit the DB.
prisma.db.litellm_teamtable.find_unique.assert_not_called()
# ---------------------------------------------------------------------------
# /key/update — _validate_caller_can_assign_key_org
# ---------------------------------------------------------------------------
def _make_prisma_with_user_orgs(user_id: str, org_ids: list):
prisma = MagicMock()
user_row = MagicMock()
user_row.organization_memberships = [
MagicMock(organization_id=org_id) for org_id in org_ids
]
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
return prisma
@pytest.mark.asyncio
async def test_assign_key_org_allows_member():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_caller_can_assign_key_org,
)
prisma = _make_prisma_with_user_orgs("alice", ["org-1", "org-2"])
caller = UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
# Should not raise.
await _validate_caller_can_assign_key_org(
user_api_key_dict=caller,
organization_id="org-2",
prisma_client=prisma,
)
@pytest.mark.asyncio
async def test_assign_key_org_blocks_non_member():
"""The IDOR: caller asks to point a key at an org they don't belong to."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_caller_can_assign_key_org,
)
prisma = _make_prisma_with_user_orgs("alice", ["org-1"])
caller = UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
with pytest.raises(HTTPException) as exc_info:
await _validate_caller_can_assign_key_org(
user_api_key_dict=caller,
organization_id="someone-elses-org",
prisma_client=prisma,
)
assert exc_info.value.status_code == 403
assert "someone-elses-org" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_assign_key_org_blocks_caller_without_user_id():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_caller_can_assign_key_org,
)
prisma = MagicMock()
caller = UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
with pytest.raises(HTTPException) as exc_info:
await _validate_caller_can_assign_key_org(
user_api_key_dict=caller,
organization_id="org-1",
prisma_client=prisma,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_assign_key_org_blocks_caller_with_no_memberships():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_caller_can_assign_key_org,
)
prisma = MagicMock()
user_row = MagicMock()
user_row.organization_memberships = None
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
caller = UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
with pytest.raises(HTTPException) as exc_info:
await _validate_caller_can_assign_key_org(
user_api_key_dict=caller,
organization_id="org-1",
prisma_client=prisma,
)
assert exc_info.value.status_code == 403
@@ -1841,6 +1841,7 @@ class TestCustomUISSO:
"x-forwarded-for": "192.168.1.1",
}
mock_request.base_url = "https://test.litellm.ai/"
mock_request.client.host = "10.0.0.10"
# Mock the custom handler
mock_custom_handler = MagicMock(spec=CustomSSOLoginHandler)
@@ -1866,36 +1867,73 @@ class TestCustomUISSO:
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
mock_custom_handler,
):
with patch.object(
SSOAuthenticationHandler,
"get_redirect_response_from_openid",
return_value=mock_redirect_response,
) as mock_get_redirect:
# Act
result = (
with patch(
"litellm.proxy.proxy_server.general_settings",
{"trusted_proxy_ranges": ["10.0.0.0/24"]},
):
with patch.object(
SSOAuthenticationHandler,
"get_redirect_response_from_openid",
return_value=mock_redirect_response,
) as mock_get_redirect:
# Act
result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
request=mock_request
)
# Assert
# Verify the custom handler was called with the request
mock_custom_handler.handle_custom_ui_sso_sign_in.assert_called_once_with(
request=mock_request
)
# Verify the redirect response was generated with correct OpenID
mock_get_redirect.assert_called_once_with(
result=expected_openid,
request=mock_request,
received_response=None,
generic_client_id=None,
ui_access_mode=None,
)
# Verify the result is the redirect response
assert result == mock_redirect_response
assert result.status_code == 303
@pytest.mark.asyncio
async def test_handle_custom_ui_sso_sign_in_rejects_untrusted_proxy(self):
"""Custom UI SSO rejects spoofed identity headers from direct clients."""
from enterprise.litellm_enterprise.proxy.auth.custom_sso_handler import (
EnterpriseCustomSSOHandler,
)
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
mock_request = MagicMock(spec=Request)
mock_request.headers = {
"x-litellm-user-id": "admin",
"x-litellm-user-email": "admin@example.com",
}
mock_request.base_url = "https://test.litellm.ai/"
mock_request.client.host = "203.0.113.10"
mock_custom_handler = MagicMock(spec=CustomSSOLoginHandler)
mock_custom_handler.handle_custom_ui_sso_sign_in = AsyncMock()
with patch("litellm.proxy.proxy_server.premium_user", True):
with patch(
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
mock_custom_handler,
):
with patch(
"litellm.proxy.proxy_server.general_settings",
{"trusted_proxy_ranges": ["10.0.0.0/24"]},
):
with pytest.raises(ValueError, match="not trusted"):
await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
request=mock_request
)
)
# Assert
# Verify the custom handler was called with the request
mock_custom_handler.handle_custom_ui_sso_sign_in.assert_called_once_with(
request=mock_request
)
# Verify the redirect response was generated with correct OpenID
mock_get_redirect.assert_called_once_with(
result=expected_openid,
request=mock_request,
received_response=None,
generic_client_id=None,
ui_access_mode=None,
)
# Verify the result is the redirect response
assert result == mock_redirect_response
assert result.status_code == 303
mock_custom_handler.handle_custom_ui_sso_sign_in.assert_not_called()
@pytest.mark.asyncio
async def test_custom_ui_sso_handler_execution_with_real_class(self):
@@ -1946,6 +1984,7 @@ class TestCustomUISSO:
"x-forwarded-for": "10.0.0.1",
}
mock_request.base_url = "https://custom.litellm.ai/"
mock_request.client.host = "10.0.0.20"
# Mock the redirect response method
mock_redirect_response = MagicMock()
@@ -1956,34 +1995,36 @@ class TestCustomUISSO:
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
test_handler_instance,
):
with patch.object(
SSOAuthenticationHandler,
"get_redirect_response_from_openid",
return_value=mock_redirect_response,
) as mock_get_redirect:
# Act
result = (
await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
with patch(
"litellm.proxy.proxy_server.general_settings",
{"trusted_proxy_ranges": ["10.0.0.0/24"]},
):
with patch.object(
SSOAuthenticationHandler,
"get_redirect_response_from_openid",
return_value=mock_redirect_response,
) as mock_get_redirect:
# Act
result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
request=mock_request
)
)
# Assert that our custom handler was executed
assert test_handler_instance.method_called is True
assert test_handler_instance.received_request == mock_request
# Assert that our custom handler was executed
assert test_handler_instance.method_called is True
assert test_handler_instance.received_request == mock_request
# Verify the redirect response was called with the OpenID from our custom handler
mock_get_redirect.assert_called_once()
call_args = mock_get_redirect.call_args.kwargs
# Verify the redirect response was called with the OpenID from our custom handler
mock_get_redirect.assert_called_once()
call_args = mock_get_redirect.call_args.kwargs
# Verify the OpenID object has the expected values from our custom handler
openid_result = call_args["result"]
assert openid_result.id == "custom_test_user_456"
assert openid_result.email == "custom@example.com"
assert openid_result.first_name == "Custom"
assert openid_result.last_name == "Handler"
assert openid_result.display_name == "Custom Handler Test"
assert openid_result.provider == "custom"
# Verify the OpenID object has the expected values from our custom handler
openid_result = call_args["result"]
assert openid_result.id == "custom_test_user_456"
assert openid_result.email == "custom@example.com"
assert openid_result.first_name == "Custom"
assert openid_result.last_name == "Handler"
assert openid_result.display_name == "Custom Handler Test"
assert openid_result.provider == "custom"
# Verify the request and other parameters were passed correctly
assert call_args["request"] == mock_request
@@ -5767,3 +5808,324 @@ class TestSyncUserRoleFromJwtRoleMap:
)
prisma.db.litellm_usertable.update.assert_not_called()
# ── VERIA-34 regression: PKCE state-to-session-cookie binding ───────────────
class TestPKCEStateCookieBinding:
"""The Generic SSO PKCE flow used the URL ``state`` parameter as a
cache-key for the PKCE ``code_verifier`` without binding the state to
the caller's browser. An attacker who pre-mints a state + cached
verifier could hand the link to a victim and capture the resulting
access token. Fix: set ``litellm_oauth_state`` HttpOnly cookie on
the redirect; verify the URL state matches the cookie before doing
the PKCE token exchange."""
@pytest.mark.asyncio
async def test_redirect_response_sets_oauth_state_cookie_when_pkce_enabled(self):
"""``get_generic_sso_redirect_response`` must set
``litellm_oauth_state`` on the redirect response when PKCE is on so
the callback can verify it later. The cookie must carry HttpOnly,
SameSite=Lax, and (because no http request was supplied to the
helper) the production-safe ``Secure`` default."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="https://idp.example.com/authorize?state=test-state-xyz"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "test-state-xyz",
"GENERIC_CLIENT_USE_PKCE": "true",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="https://idp.example.com/authorize",
)
assert response is not None
cookie_headers = response.headers.getlist("set-cookie")
cookie_str = next(
(c for c in cookie_headers if "litellm_oauth_state=" in c), None
)
assert (
cookie_str is not None
), f"litellm_oauth_state cookie not set; got: {cookie_headers}"
assert "test-state-xyz" in cookie_str
assert "HttpOnly" in cookie_str
assert "SameSite=lax" in cookie_str
# No incoming Request supplied → ``Secure`` defaults to True so a
# network observer on plain HTTP cannot read the state value.
assert "Secure" in cookie_str
@pytest.mark.asyncio
async def test_redirect_response_omits_oauth_state_cookie_when_pkce_disabled(
self,
):
"""Non-PKCE flows delegate to fastapi-sso's own session-cookie
binding; we do not set our cookie there because it would never be
validated (and could collide with a concurrent PKCE session in
the same browser)."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="https://idp.example.com/authorize?state=test-state-xyz"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "test-state-xyz",
"GENERIC_CLIENT_USE_PKCE": "false",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="https://idp.example.com/authorize",
)
assert response is not None
cookie_headers = response.headers.getlist("set-cookie")
assert not any(
"litellm_oauth_state=" in c for c in cookie_headers
), f"litellm_oauth_state cookie set on non-PKCE flow; got: {cookie_headers}"
@pytest.mark.asyncio
async def test_redirect_response_drops_secure_flag_for_http_dev(self):
"""When the incoming request is plain HTTP (local dev), ``Secure``
must be dropped so the browser will actually attach the cookie on
the callback hop."""
from fastapi.responses import RedirectResponse
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
)
mock_redirect = RedirectResponse(
url="http://idp.local/authorize?state=local-dev-state"
)
mock_generic_sso = MagicMock()
mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso)
mock_generic_sso.__exit__ = MagicMock(return_value=None)
mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect)
http_request = MagicMock(spec=Request)
http_request.url.scheme = "http"
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_STATE": "local-dev-state",
"GENERIC_CLIENT_USE_PKCE": "true",
},
):
response = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_generic_sso,
state=None,
generic_authorization_endpoint="http://idp.local/authorize",
request=http_request,
)
cookie_headers = response.headers.getlist("set-cookie")
cookie_str = next(
(c for c in cookie_headers if "litellm_oauth_state=" in c), None
)
assert cookie_str is not None
assert "Secure" not in cookie_str
@pytest.mark.asyncio
async def test_pkce_callback_rejects_missing_cookie(self):
"""When PKCE is enabled and a code_verifier is in the cache, the
callback must reject a request that has no ``litellm_oauth_state``
cookie (browser-to-server binding missing)."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {
"state": "attacker-minted-state",
"code": "auth-code",
}
# No oauth_state cookie set → request.cookies.get returns None.
mock_request.cookies = {}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "attacker-cached-verifier",
"_pkce_cache_key": "pkce_verifier:attacker-minted-state",
}
),
),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
pytest.raises(ProxyException) as exc_info,
):
await get_generic_sso_response(
request=mock_request,
jwt_handler=MagicMock(spec=JWTHandler),
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
assert "state" in str(exc_info.value.message).lower()
@pytest.mark.asyncio
async def test_pkce_callback_rejects_state_cookie_mismatch(self):
"""The Login-CSRF shape: attacker mints state ``A``, victim's browser
carries cookie state ``B``. The callback must reject."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {
"state": "attacker-minted-state",
"code": "auth-code",
}
mock_request.cookies = {"litellm_oauth_state": "victim-browser-state"}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "verifier",
"_pkce_cache_key": "pkce_verifier:attacker-minted-state",
}
),
),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
pytest.raises(ProxyException) as exc_info,
):
await get_generic_sso_response(
request=mock_request,
jwt_handler=MagicMock(spec=JWTHandler),
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
assert "state" in str(exc_info.value.message).lower()
@pytest.mark.asyncio
async def test_pkce_callback_accepts_matching_state_cookie(self):
"""Happy path: URL state and cookie state match (the legitimate
flow where the same browser that started the redirect lands on
the callback) the PKCE token exchange proceeds."""
from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
get_generic_sso_response,
)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "matched-state", "code": "auth-code"}
mock_request.cookies = {"litellm_oauth_state": "matched-state"}
with (
patch.object(
SSOAuthenticationHandler,
"prepare_token_exchange_parameters",
AsyncMock(
return_value={
"code_verifier": "verifier",
"_pkce_cache_key": "pkce_verifier:matched-state",
}
),
),
patch.object(
SSOAuthenticationHandler,
"_pkce_token_exchange",
AsyncMock(
return_value={
"access_token": "tok",
"id_token": "id",
"sub": "user@example.com",
"email": "user@example.com",
}
),
),
patch.object(
SSOAuthenticationHandler,
"_delete_pkce_verifier",
AsyncMock(),
),
patch("fastapi_sso.sso.base.DiscoveryDocument"),
patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()),
patch.dict(
os.environ,
{
"GENERIC_CLIENT_SECRET": "x",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth",
"GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo",
"GENERIC_CLIENT_USE_PKCE": "true",
},
),
):
jwt_handler = MagicMock(spec=JWTHandler)
jwt_handler.get_team_ids_from_jwt.return_value = []
result, _, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=jwt_handler,
generic_client_id="cid",
redirect_url="https://proxy.example.com/sso/callback",
sso_jwt_handler=None,
)
# State-cookie check passed, so the function got past the early
# ProxyException raise and produced an SSO result object.
assert result is not None
@@ -241,6 +241,53 @@ async def test_route_request_with_router_settings_override_preserves_existing():
assert call_kwargs["timeout"] == 30
def test_mock_testing_kwarg_names_matches_dataclass():
"""``_MOCK_TESTING_KWARG_NAMES`` is hardcoded to avoid a cyclic import
against ``litellm.types.router``. This test guards against drift
if a new ``mock_testing_*`` field is added to ``MockRouterTestingParams``
the strip list must be updated to keep covering it."""
from dataclasses import fields
from litellm.proxy.route_llm_request import _MOCK_TESTING_KWARG_NAMES
from litellm.types.router import MockRouterTestingParams
assert set(_MOCK_TESTING_KWARG_NAMES) == {
f.name for f in fields(MockRouterTestingParams)
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"mock_flag",
[
"mock_testing_fallbacks",
"mock_testing_context_fallbacks",
"mock_testing_content_policy_fallbacks",
],
)
async def test_route_request_strips_mock_testing_flags(mock_flag):
"""VERIA-44: router-internal testing flags must not survive a
user-supplied request body. Without this strip, an attacker can
combine ``mock_testing_fallbacks=true`` with an unauthorized fallback
in ``router_settings_override`` to deterministically execute requests
against restricted models."""
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
mock_flag: True,
}
llm_router = MagicMock()
llm_router.acompletion.return_value = "ok"
await route_request(data, llm_router, None, "acompletion")
call_kwargs = llm_router.acompletion.call_args[1]
assert mock_flag not in call_kwargs
# The flag is also gone from the original data dict so any subsequent
# processing (e.g. logging) doesn't see it either.
assert mock_flag not in data
@pytest.mark.parametrize(
"route_type", ["agenerate_content", "agenerate_content_stream"]
)