mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #25905 from stuxf/fix/metadata-security-controls
fix(proxy): read guardrail config from admin metadata, fix tag routing consistency
This commit is contained in:
commit
32714a4dc2
16 changed files with 1282 additions and 261 deletions
|
|
@ -255,26 +255,44 @@ class CustomGuardrail(CustomLogger):
|
|||
f"Event hook {event_hook} is not in the supported event hooks {supported_event_hooks}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_admin_metadata(data: dict) -> dict:
|
||||
"""Return merged admin-configured key and team metadata from the request data.
|
||||
|
||||
The proxy may inject admin metadata (user_api_key_metadata,
|
||||
user_api_key_team_metadata) into either ``metadata`` or
|
||||
``litellm_metadata`` depending on endpoint. Check both so a caller
|
||||
cannot shadow admin config by pre-populating the other key.
|
||||
Key-level settings override team-level.
|
||||
"""
|
||||
team_meta: dict = {}
|
||||
key_meta: dict = {}
|
||||
for key in ("metadata", "litellm_metadata"):
|
||||
# Defensive: an unparsed JSON-string metadata could leak past the
|
||||
# proxy's normal parse path; don't AttributeError on .get().
|
||||
meta = data.get(key)
|
||||
if not isinstance(meta, dict):
|
||||
continue
|
||||
team_meta = meta.get("user_api_key_team_metadata") or team_meta
|
||||
key_meta = meta.get("user_api_key_metadata") or key_meta
|
||||
return {**team_meta, **key_meta}
|
||||
|
||||
def get_disable_global_guardrail(self, data: dict) -> Optional[bool]:
|
||||
"""
|
||||
Returns True if the global guardrail should be disabled
|
||||
Returns True if the global guardrail should be disabled.
|
||||
|
||||
Reads from admin-configured key/team metadata only, not from
|
||||
the request body, to prevent callers from disabling guardrails.
|
||||
"""
|
||||
if "disable_global_guardrails" in data:
|
||||
return data["disable_global_guardrails"]
|
||||
metadata = data.get("litellm_metadata") or data.get("metadata", {})
|
||||
if "disable_global_guardrails" in metadata:
|
||||
return metadata["disable_global_guardrails"]
|
||||
return False
|
||||
return self._get_admin_metadata(data).get("disable_global_guardrails", False)
|
||||
|
||||
def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]:
|
||||
"""
|
||||
Returns the list of global guardrail names the team/key has opted out of.
|
||||
|
||||
Reads from admin-configured key/team metadata only.
|
||||
"""
|
||||
if "opted_out_global_guardrails" in data:
|
||||
value = data["opted_out_global_guardrails"]
|
||||
return value if isinstance(value, list) else []
|
||||
metadata = data.get("litellm_metadata") or data.get("metadata", {})
|
||||
value = metadata.get("opted_out_global_guardrails")
|
||||
value = self._get_admin_metadata(data).get("opted_out_global_guardrails")
|
||||
return value if isinstance(value, list) else []
|
||||
|
||||
def _is_valid_response_type(self, result: Any) -> bool:
|
||||
|
|
@ -417,7 +435,9 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
requested_guardrails = self.get_guardrail_from_metadata(data)
|
||||
disable_global_guardrail = self.get_disable_global_guardrail(data)
|
||||
opted_out_global_guardrails = self.get_opted_out_global_guardrails_from_metadata(data)
|
||||
opted_out_global_guardrails = (
|
||||
self.get_opted_out_global_guardrails_from_metadata(data)
|
||||
)
|
||||
verbose_logger.debug(
|
||||
"inside should_run_guardrail for guardrail=%s event_type= %s guardrail_supported_event_hooks= %s requested_guardrails= %s self.default_on= %s",
|
||||
self.guardrail_name,
|
||||
|
|
@ -426,7 +446,10 @@ class CustomGuardrail(CustomLogger):
|
|||
requested_guardrails,
|
||||
self.default_on,
|
||||
)
|
||||
if self.default_on is True and self.guardrail_name in opted_out_global_guardrails:
|
||||
if (
|
||||
self.default_on is True
|
||||
and self.guardrail_name in opted_out_global_guardrails
|
||||
):
|
||||
return False
|
||||
|
||||
if self.default_on is True and disable_global_guardrail is not True:
|
||||
|
|
|
|||
|
|
@ -48,7 +48,6 @@ _supported_callback_params = [
|
|||
"braintrust_host",
|
||||
"slack_webhook_url",
|
||||
"lunary_public_key",
|
||||
"turn_off_message_logging",
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy._types import (
|
||||
RBAC_ROLES,
|
||||
CallInfo,
|
||||
|
|
@ -327,15 +328,63 @@ def _global_proxy_budget_check(
|
|||
)
|
||||
|
||||
|
||||
_GUARDRAIL_MODIFICATION_KEYS: tuple = (
|
||||
"guardrails",
|
||||
"disable_global_guardrails",
|
||||
"disable_global_guardrail",
|
||||
"opted_out_global_guardrails",
|
||||
)
|
||||
|
||||
|
||||
def _guardrail_modification_check(
|
||||
request_body: dict, team_object: Optional[LiteLLM_TeamTable]
|
||||
) -> None:
|
||||
_request_metadata: dict = request_body.get("metadata", {}) or {}
|
||||
if not _request_metadata.get("guardrails"):
|
||||
return
|
||||
"""
|
||||
Reject user-supplied metadata flags that would modify guardrail behavior
|
||||
unless the team has explicit permission. Checked keys include the plural
|
||||
``guardrails`` list plus the per-request toggles that influence whether
|
||||
default-on guardrails run (``disable_global_guardrails``,
|
||||
``disable_global_guardrail`` singular, and ``opted_out_global_guardrails``).
|
||||
|
||||
User-supplied values for the bypass toggles are also silently ignored by
|
||||
``_get_admin_metadata`` at read time; this check adds defense in depth by
|
||||
failing loudly at the auth layer so operators see an explicit 403 instead
|
||||
of a confusing silent-ignore.
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails
|
||||
|
||||
def _coerce_to_dict(container: Any) -> Optional[dict]:
|
||||
"""Accept dict or JSON-string (from multipart/form-data or extra_body).
|
||||
|
||||
Without this, an attacker can smuggle guardrail keys past the check by
|
||||
sending ``{"metadata": "{\\"disable_global_guardrails\\": true}"}`` —
|
||||
``isinstance(dict)`` on the string returns False, the check returns
|
||||
no-modification, and ``add_litellm_data_to_request`` parses the string
|
||||
to a dict downstream.
|
||||
"""
|
||||
if isinstance(container, dict):
|
||||
return container
|
||||
if isinstance(container, str):
|
||||
parsed = safe_json_loads(container)
|
||||
return parsed if isinstance(parsed, dict) else None
|
||||
return None
|
||||
|
||||
def _user_requested_modification(container: Any) -> bool:
|
||||
coerced = _coerce_to_dict(container)
|
||||
if coerced is None:
|
||||
return False
|
||||
return any(coerced.get(key) for key in _GUARDRAIL_MODIFICATION_KEYS)
|
||||
|
||||
# Check both metadata keys — callers can populate either depending on the
|
||||
# endpoint. Cover the top-level too so root-level injection is rejected.
|
||||
modifies = (
|
||||
_user_requested_modification(request_body.get("metadata"))
|
||||
or _user_requested_modification(request_body.get("litellm_metadata"))
|
||||
or _user_requested_modification(request_body)
|
||||
)
|
||||
if not modifies:
|
||||
return
|
||||
|
||||
if not can_modify_guardrails(team_object):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
@ -451,9 +500,9 @@ async def common_checks( # noqa: PLR0915
|
|||
model=_model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=valid_token.team_model_aliases
|
||||
if valid_token
|
||||
else None,
|
||||
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}",
|
||||
|
|
|
|||
|
|
@ -823,19 +823,30 @@ def get_end_user_id_from_request_body(
|
|||
user_from_body_user_field = request_body["user"]
|
||||
return str(user_from_body_user_field)
|
||||
|
||||
def _as_dict(value: Any) -> dict:
|
||||
# metadata / litellm_metadata can arrive as JSON strings from
|
||||
# multipart/form-data or extra_body; coerce so string-encoded
|
||||
# payloads can't evade end-user attribution.
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
||||
parsed = safe_json_loads(value)
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
return {}
|
||||
|
||||
# Check 4: 'litellm_metadata.user' in request_body (commonly Anthropic)
|
||||
litellm_metadata = request_body.get("litellm_metadata")
|
||||
if isinstance(litellm_metadata, dict):
|
||||
user_from_litellm_metadata = litellm_metadata.get("user")
|
||||
if user_from_litellm_metadata is not None:
|
||||
return str(user_from_litellm_metadata)
|
||||
litellm_metadata = _as_dict(request_body.get("litellm_metadata"))
|
||||
user_from_litellm_metadata = litellm_metadata.get("user")
|
||||
if user_from_litellm_metadata is not None:
|
||||
return str(user_from_litellm_metadata)
|
||||
|
||||
# Check 5: 'metadata.user_id' in request_body (another common pattern)
|
||||
metadata_dict = request_body.get("metadata")
|
||||
if isinstance(metadata_dict, dict):
|
||||
user_id_from_metadata_field = metadata_dict.get("user_id")
|
||||
if user_id_from_metadata_field is not None:
|
||||
return str(user_id_from_metadata_field)
|
||||
metadata_dict = _as_dict(request_body.get("metadata"))
|
||||
user_id_from_metadata_field = metadata_dict.get("user_id")
|
||||
if user_id_from_metadata_field is not None:
|
||||
return str(user_id_from_metadata_field)
|
||||
|
||||
# Check 6: 'safety_identifier' in request body (OpenAI Responses API parameter)
|
||||
# SECURITY NOTE: safety_identifier can be set by any caller in the request body.
|
||||
|
|
|
|||
|
|
@ -197,10 +197,10 @@ def check_file_size_under_limit(
|
|||
|
||||
if llm_router is not None and request_data["model"] in router_model_names:
|
||||
try:
|
||||
deployment: Optional[
|
||||
Deployment
|
||||
] = llm_router.get_deployment_by_model_group_name(
|
||||
model_group_name=request_data["model"]
|
||||
deployment: Optional[Deployment] = (
|
||||
llm_router.get_deployment_by_model_group_name(
|
||||
model_group_name=request_data["model"]
|
||||
)
|
||||
)
|
||||
if (
|
||||
deployment
|
||||
|
|
@ -426,7 +426,16 @@ def get_tags_from_request_body(request_body: dict) -> List[str]:
|
|||
List of tag names (strings), empty list if no valid tags found
|
||||
"""
|
||||
metadata_variable_name = get_metadata_variable_name_from_kwargs(request_body)
|
||||
metadata = request_body.get(metadata_variable_name) or {}
|
||||
metadata = request_body.get(metadata_variable_name)
|
||||
# metadata can arrive as a JSON string from multipart/form-data or extra_body;
|
||||
# coerce defensively so .get() below never raises AttributeError.
|
||||
if isinstance(metadata, str):
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
||||
parsed = safe_json_loads(metadata)
|
||||
metadata = parsed if isinstance(parsed, dict) else {}
|
||||
elif not isinstance(metadata, dict):
|
||||
metadata = {}
|
||||
tags_in_metadata: Any = metadata.get("tags", [])
|
||||
tags_in_request_body: Any = request_body.get("tags", [])
|
||||
combined_tags: List[str] = []
|
||||
|
|
|
|||
|
|
@ -981,13 +981,16 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
# Init - Proxy Server Request
|
||||
# we do this as soon as entering so we track the original request
|
||||
##########################################################
|
||||
# Track arrival time for queue time metric
|
||||
# Track arrival time for queue time metric. The body snapshot is filled
|
||||
# in after the admin-injection strip below so the audit / spend-tracking
|
||||
# consumers of proxy_server_request["body"] see the cleaned metadata
|
||||
# rather than attacker-forged user_api_key_* fields.
|
||||
arrival_time = time.time()
|
||||
data["proxy_server_request"] = {
|
||||
"url": str(request.url),
|
||||
"method": request.method,
|
||||
"headers": _headers,
|
||||
"body": copy.copy(data), # use copy instead of deepcopy
|
||||
"body": None, # filled in post-strip; see below
|
||||
"arrival_time": arrival_time, # Track when request arrived at proxy
|
||||
}
|
||||
|
||||
|
|
@ -1069,9 +1072,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
verbose_proxy_logger.warning(
|
||||
f"Failed to parse 'metadata' as JSON dict. Received value: {data['metadata']}"
|
||||
)
|
||||
data[_metadata_variable_name]["requester_metadata"] = copy.deepcopy(
|
||||
data["metadata"]
|
||||
)
|
||||
# requester_metadata is snapshotted AFTER the strip below so
|
||||
# downstream consumers (e.g. PANW guardrail reading user_ip /
|
||||
# profile_id) don't see attacker-injected admin slots preserved in
|
||||
# the deepcopy.
|
||||
|
||||
# Parse litellm_metadata if it's a string (e.g., from multipart/form-data or extra_body)
|
||||
if "litellm_metadata" in data and data["litellm_metadata"] is not None:
|
||||
|
|
@ -1083,11 +1087,89 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
)
|
||||
else:
|
||||
data["litellm_metadata"] = parsed_litellm_metadata
|
||||
# Merge litellm_metadata into the metadata variable (preserving existing values)
|
||||
if isinstance(data["litellm_metadata"], dict):
|
||||
for key, value in data["litellm_metadata"].items():
|
||||
if key not in data[_metadata_variable_name]:
|
||||
data[_metadata_variable_name][key] = value
|
||||
|
||||
# Strip internal pipeline state and admin-injection slots from user input.
|
||||
# Runs AFTER the string-to-dict parse above so JSON-string metadata (sent
|
||||
# via multipart/form-data or extra_body) cannot smuggle admin fields past
|
||||
# the isinstance(dict) guard.
|
||||
#
|
||||
# The proxy populates a family of ``user_api_key_*`` fields below
|
||||
# (user_api_key_metadata, user_api_key_user_id, user_api_key_alias,
|
||||
# user_api_key_spend, user_api_key_team_metadata, …) into
|
||||
# data[_metadata_variable_name]. Because the proxy only writes to ONE of
|
||||
# the two metadata dicts, a caller pre-populating any of these keys on
|
||||
# the OTHER metadata dict would have their forged values surface in
|
||||
# guardrails, spend tracking, audit logs, and identity resolution. Strip
|
||||
# by prefix so new ``user_api_key_*`` fields added in the future are
|
||||
# covered without per-key maintenance.
|
||||
for _meta_key in ("metadata", "litellm_metadata"):
|
||||
_user_meta = data.get(_meta_key)
|
||||
if isinstance(_user_meta, dict):
|
||||
_user_meta.pop("_pipeline_managed_guardrails", None)
|
||||
for _k in [k for k in _user_meta if k.startswith("user_api_key_")]:
|
||||
_user_meta.pop(_k, None)
|
||||
|
||||
# Strip caller-supplied routing/budget tags unless the admin has opted
|
||||
# this key or team in via metadata.allow_client_tags=True. Tags drive
|
||||
# tag-based routing and tag budget attribution — accepting them from
|
||||
# untrusted callers lets an attacker reach restricted deployments or
|
||||
# misattribute spend to a victim team's tag.
|
||||
_admin_allow_client_tags = False
|
||||
for _admin_meta in (
|
||||
user_api_key_dict.metadata,
|
||||
user_api_key_dict.team_metadata,
|
||||
):
|
||||
if (
|
||||
isinstance(_admin_meta, dict)
|
||||
and _admin_meta.get("allow_client_tags") is True
|
||||
):
|
||||
_admin_allow_client_tags = True
|
||||
break
|
||||
if not _admin_allow_client_tags:
|
||||
_stripped_from: List[str] = []
|
||||
for _meta_key in ("metadata", "litellm_metadata"):
|
||||
_user_meta = data.get(_meta_key)
|
||||
if isinstance(_user_meta, dict) and "tags" in _user_meta:
|
||||
_user_meta.pop("tags", None)
|
||||
_stripped_from.append(_meta_key)
|
||||
# Also strip the root-level `tags` field. get_tags_from_request_body
|
||||
# reads request_body["tags"] directly and feeds it to the policy
|
||||
# engine, so leaving it in place here would let the strip-in-metadata
|
||||
# above be trivially bypassed by moving the tags to the body root.
|
||||
if "tags" in data:
|
||||
data.pop("tags", None)
|
||||
_stripped_from.append("tags (root)")
|
||||
if _stripped_from:
|
||||
verbose_proxy_logger.warning(
|
||||
"Stripped caller-supplied tags from %s: this key/team does "
|
||||
"not have `allow_client_tags: true` in its metadata. Set it "
|
||||
"to opt into client-supplied routing/budget tags.",
|
||||
", ".join(_stripped_from),
|
||||
)
|
||||
|
||||
# Fill in the proxy_server_request body snapshot now that metadata has
|
||||
# been parsed and stripped. Consumers (standard_logging_payload, lago,
|
||||
# spend_tracking_utils, streaming_iterator) read `body` to audit the
|
||||
# request; taking the snapshot here ensures they see cleaned metadata.
|
||||
data["proxy_server_request"]["body"] = copy.copy(data)
|
||||
|
||||
# Snapshot the (now-cleaned) requester-supplied metadata for downstream
|
||||
# consumers. Taking the deepcopy AFTER the strip prevents attacker-
|
||||
# injected admin slots (user_api_key_*, tags without opt-in,
|
||||
# _pipeline_managed_guardrails) from surviving in requester_metadata
|
||||
# where guardrails and audit paths may read from it.
|
||||
if "metadata" in data and isinstance(data["metadata"], dict):
|
||||
data[_metadata_variable_name]["requester_metadata"] = copy.deepcopy(
|
||||
data["metadata"]
|
||||
)
|
||||
|
||||
# Now merge litellm_metadata into the metadata variable (preserving existing
|
||||
# values) — runs AFTER the strip so attacker injections in litellm_metadata
|
||||
# cannot cross-contaminate the admin-authoritative metadata dict.
|
||||
if "litellm_metadata" in data and isinstance(data["litellm_metadata"], dict):
|
||||
for key, value in data["litellm_metadata"].items():
|
||||
if key not in data[_metadata_variable_name]:
|
||||
data[_metadata_variable_name][key] = value
|
||||
|
||||
data = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
|
||||
data=data,
|
||||
|
|
@ -1250,15 +1332,24 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
user_agent = request.headers["user-agent"]
|
||||
data[_metadata_variable_name]["user_agent"] = user_agent
|
||||
|
||||
# Check if using tag based routing
|
||||
# Check if using tag based routing. The helper reads caller-controlled
|
||||
# sources (x-litellm-tags header, data["tags"] root-level), so its result
|
||||
# is still gated by the same allow_client_tags flag that gated the
|
||||
# body-metadata tag strip above. Otherwise the strip is trivially
|
||||
# bypassed by sending tags via header or at the root of the body.
|
||||
tags = LiteLLMProxyRequestSetup.add_request_tag_to_metadata(
|
||||
llm_router=llm_router,
|
||||
headers=_headers,
|
||||
data=data,
|
||||
)
|
||||
|
||||
if tags is not None:
|
||||
if tags is not None and _admin_allow_client_tags:
|
||||
data[_metadata_variable_name]["tags"] = tags
|
||||
elif tags is not None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Ignored caller-supplied tags from header/root body: this "
|
||||
"key/team does not have `allow_client_tags: true` in its metadata."
|
||||
)
|
||||
|
||||
# Team Callbacks controls
|
||||
callback_settings_obj = _get_dynamic_logging_metadata(
|
||||
|
|
|
|||
|
|
@ -2201,9 +2201,11 @@ def run_ollama_serve():
|
|||
with open(os.devnull, "w") as devnull:
|
||||
subprocess.Popen(command, stdout=devnull, stderr=devnull)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(f"""
|
||||
verbose_proxy_logger.debug(
|
||||
f"""
|
||||
LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve`
|
||||
""")
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _get_process_rss_mb() -> Optional[float]:
|
||||
|
|
@ -7157,7 +7159,10 @@ async def chat_completion( # noqa: PLR0915
|
|||
global user_temperature, user_request_timeout, user_max_tokens, user_api_base
|
||||
data = await _read_request_body(request=request)
|
||||
if user_api_key_dict is not None:
|
||||
if data.get("metadata") is None:
|
||||
if not isinstance(data.get("metadata"), dict):
|
||||
# Covers both missing and JSON-string metadata (multipart /
|
||||
# extra_body); otherwise `data["metadata"][k] = v` below raises
|
||||
# TypeError on a string value and 500s the request.
|
||||
data["metadata"] = {}
|
||||
if (
|
||||
hasattr(user_api_key_dict, "user_id")
|
||||
|
|
@ -11372,7 +11377,9 @@ async def async_queue_request(
|
|||
# if users are using user_api_key_auth, set `user` in `data`
|
||||
data["user"] = user_api_key_dict.user_id
|
||||
|
||||
if "metadata" not in data:
|
||||
if not isinstance(data.get("metadata"), dict):
|
||||
# Covers both missing and JSON-string metadata (multipart /
|
||||
# extra_body); see above for the same guard upstream.
|
||||
data["metadata"] = {}
|
||||
data["metadata"]["user_api_key"] = user_api_key_dict.api_key
|
||||
data["metadata"]["user_api_key_metadata"] = user_api_key_dict.metadata
|
||||
|
|
|
|||
|
|
@ -29,6 +29,9 @@ from litellm.caching.redis_cache import RedisPipelineIncrementOperation
|
|||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
from litellm.router_utils.cooldown_callbacks import (
|
||||
_get_prometheus_logger_from_callbacks,
|
||||
)
|
||||
|
|
@ -100,9 +103,9 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
self.dual_cache = dual_cache
|
||||
self.redis_increment_operation_queue: List[RedisPipelineIncrementOperation] = []
|
||||
asyncio.create_task(self.periodic_sync_in_memory_spend_with_redis())
|
||||
self.provider_budget_config: Optional[
|
||||
GenericBudgetConfigType
|
||||
] = provider_budget_config
|
||||
self.provider_budget_config: Optional[GenericBudgetConfigType] = (
|
||||
provider_budget_config
|
||||
)
|
||||
self.deployment_budget_config: Optional[GenericBudgetConfigType] = None
|
||||
self.tag_budget_config: Optional[GenericBudgetConfigType] = None
|
||||
self._init_provider_budgets()
|
||||
|
|
@ -175,7 +178,10 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
spend_map=spend_map,
|
||||
potential_deployments=potential_deployments,
|
||||
request_tags=_get_tags_from_request_kwargs(
|
||||
request_kwargs=request_kwargs
|
||||
request_kwargs=request_kwargs,
|
||||
metadata_variable_name=get_metadata_variable_name_from_kwargs(
|
||||
request_kwargs or {}
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -304,6 +310,16 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
deployment_configs: Dict[str, GenericBudgetInfo] = {}
|
||||
deployment_providers: List[Optional[str]] = []
|
||||
|
||||
# Resolve tags once before the loop (loop-invariant)
|
||||
_request_tags: List[str] = []
|
||||
if self.tag_budget_config:
|
||||
_request_tags = _get_tags_from_request_kwargs(
|
||||
request_kwargs=request_kwargs,
|
||||
metadata_variable_name=get_metadata_variable_name_from_kwargs(
|
||||
request_kwargs or {}
|
||||
),
|
||||
)
|
||||
|
||||
for deployment in healthy_deployments:
|
||||
# Check provider budgets
|
||||
if self.provider_budget_config:
|
||||
|
|
@ -330,17 +346,14 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
cache_keys.append(
|
||||
f"deployment_spend:{model_id}:{budget_config.budget_duration}"
|
||||
)
|
||||
# Check tag budgets
|
||||
if self.tag_budget_config:
|
||||
request_tags = _get_tags_from_request_kwargs(
|
||||
request_kwargs=request_kwargs
|
||||
|
||||
# Check tag budgets (outside loop — tags are per-request, not per-deployment)
|
||||
for _tag in _request_tags:
|
||||
_tag_budget_config = self._get_budget_config_for_tag(_tag)
|
||||
if _tag_budget_config:
|
||||
cache_keys.append(
|
||||
f"tag_spend:{_tag}:{_tag_budget_config.budget_duration}"
|
||||
)
|
||||
for _tag in request_tags:
|
||||
_tag_budget_config = self._get_budget_config_for_tag(_tag)
|
||||
if _tag_budget_config:
|
||||
cache_keys.append(
|
||||
f"tag_spend:{_tag}:{_tag_budget_config.budget_duration}"
|
||||
)
|
||||
return (
|
||||
cache_keys,
|
||||
provider_configs,
|
||||
|
|
@ -459,7 +472,10 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
response_cost=response_cost,
|
||||
)
|
||||
|
||||
request_tags = _get_tags_from_request_kwargs(kwargs)
|
||||
request_tags = _get_tags_from_request_kwargs(
|
||||
kwargs,
|
||||
metadata_variable_name=get_metadata_variable_name_from_kwargs(kwargs or {}),
|
||||
)
|
||||
if len(request_tags) > 0:
|
||||
for _tag in request_tags:
|
||||
_tag_budget_config = self._get_budget_config_for_tag(_tag)
|
||||
|
|
|
|||
|
|
@ -102,13 +102,11 @@ def _match_deployment(
|
|||
return {"matched_via": "tags", "matched_value": matched_value}
|
||||
|
||||
# 2. Regex match against request headers.
|
||||
# When match_any=False and the deployment has both plain tags and tag_regex,
|
||||
# the strict tag check has already failed (step 1 returned None). Allow
|
||||
# the regex to fire only when the deployment has NO plain tags, so we never
|
||||
# use regex as a backdoor around the operator's strict-tag policy.
|
||||
strict_tag_check_failed = (
|
||||
not match_any and bool(deployment_tags) and bool(request_tags)
|
||||
)
|
||||
# When match_any=False and the deployment has plain tags, the strict tag
|
||||
# check either didn't run (no request tags) or failed (step 1 returned
|
||||
# None). Block the regex path so it cannot circumvent the operator's
|
||||
# strict-tag policy.
|
||||
strict_tag_check_failed = not match_any and bool(deployment_tags)
|
||||
if deployment_tag_regex and header_strings and not strict_tag_check_failed:
|
||||
regex_match = _is_valid_deployment_tag_regex(
|
||||
deployment_tag_regex, header_strings
|
||||
|
|
|
|||
|
|
@ -162,9 +162,13 @@ async def test_add_key_or_team_level_spend_logs_metadata_to_request(
|
|||
|
||||
print(f"team_sl_metadata: {team_sl_metadata}")
|
||||
mock_request.url.path = "/chat/completions"
|
||||
# Opt the key into client-supplied tags so request_tags are preserved
|
||||
# and merged with admin-configured key/team tags. Without this flag,
|
||||
# request_tags would be stripped by add_litellm_data_to_request.
|
||||
key_metadata = {
|
||||
"tags": key_tags,
|
||||
"spend_logs_metadata": key_sl_metadata,
|
||||
"allow_client_tags": True,
|
||||
}
|
||||
team_metadata = {
|
||||
"tags": team_tags,
|
||||
|
|
@ -838,12 +842,13 @@ async def test_add_litellm_data_to_request_duplicate_tags(
|
|||
mock_request.headers = {}
|
||||
mock_request.state = State()
|
||||
|
||||
# Setup key with tags in metadata
|
||||
# Setup key with tags in metadata. Opt into client-supplied tags so the
|
||||
# request_tags are preserved for the merge under test.
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test_api_key",
|
||||
user_id="test_user_id",
|
||||
org_id="test_org_id",
|
||||
metadata={"tags": key_tags},
|
||||
metadata={"tags": key_tags, "allow_client_tags": True},
|
||||
)
|
||||
|
||||
# Setup request data with tags
|
||||
|
|
|
|||
|
|
@ -173,17 +173,16 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
assert result is False
|
||||
|
||||
def test_should_run_guardrail_with_disable_global_guardrail(self):
|
||||
"""Test that disable_global_guardrail disables a global guardrail when set to True"""
|
||||
"""Test that disable_global_guardrails only works from admin metadata"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
# Create a guardrail with default_on=True (global guardrail)
|
||||
custom_guardrail = CustomGuardrail(
|
||||
guardrail_name="global_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
|
||||
# Test 1: Global guardrail runs by default when default_on=True
|
||||
# Test 1: Global guardrail runs by default
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
|
|
@ -193,7 +192,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
)
|
||||
assert result is True, "Global guardrail should run when default_on=True"
|
||||
|
||||
# Test 2: Global guardrail is disabled when disable_global_guardrail=True at root level
|
||||
# Test 2: User-injected disable at root level is IGNORED
|
||||
data_with_disable_root = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
|
|
@ -203,23 +202,10 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
data=data_with_disable_root, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), "Global guardrail should be disabled when disable_global_guardrail=True"
|
||||
result is True
|
||||
), "User-injected disable_global_guardrails should be ignored"
|
||||
|
||||
# Test 3: Global guardrail is disabled when disable_global_guardrail=True in litellm_metadata
|
||||
data_with_disable_litellm = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"litellm_metadata": {"disable_global_guardrails": True},
|
||||
}
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_with_disable_litellm, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), "Global guardrail should be disabled when disable_global_guardrail=True in litellm_metadata"
|
||||
|
||||
# Test 4: Global guardrail is disabled when disable_global_guardrail=True in metadata
|
||||
# Test 3: User-injected disable in metadata is IGNORED
|
||||
data_with_disable_metadata = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
|
|
@ -228,25 +214,51 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_with_disable_metadata, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), "Global guardrail should be disabled when disable_global_guardrail=True in metadata"
|
||||
assert result is True, "User-injected metadata disable should be ignored"
|
||||
|
||||
# Test 5: Global guardrail runs when disable_global_guardrail=False
|
||||
data_with_disable_false = {
|
||||
# Test 4: Admin-configured disable via user_api_key_metadata IS respected
|
||||
data_with_admin_disable = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"disable_global_guardrails": False,
|
||||
"metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}},
|
||||
}
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_with_disable_false, event_type=GuardrailEventHooks.pre_call
|
||||
data=data_with_admin_disable, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert result is False, "Admin-configured disable should be respected"
|
||||
|
||||
# Test 5: Admin config in metadata isn't shadowed by user-supplied litellm_metadata
|
||||
data_cross_key = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}},
|
||||
"litellm_metadata": {"request_tags": ["user-supplied"]},
|
||||
}
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_cross_key, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is True
|
||||
), "Global guardrail should still run when disable_global_guardrail=False"
|
||||
result is False
|
||||
), "Admin config in metadata must not be shadowed by user-supplied litellm_metadata"
|
||||
|
||||
# Test 6: After the pre-call strip runs, user-injected
|
||||
# user_api_key_metadata in the non-authoritative metadata key is gone.
|
||||
# _get_admin_metadata must then surface admin config unchanged.
|
||||
data_post_strip = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}},
|
||||
"litellm_metadata": {}, # post-strip: attacker payload removed
|
||||
}
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_post_strip, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), "Admin config in metadata must be respected when other metadata key is empty"
|
||||
|
||||
def test_should_run_guardrail_with_opted_out_global_guardrails(self):
|
||||
"""Test the per-guardrail opt-out list for global (default_on=True) guardrails"""
|
||||
"""Test that per-guardrail opt-out only works from admin metadata"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
custom_guardrail = CustomGuardrail(
|
||||
|
|
@ -255,7 +267,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
event_hook=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
|
||||
# Test 1: guardrail in the opt-out list at root level → skipped
|
||||
# Test 1: User-injected opt-out at root level is IGNORED
|
||||
data_root = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
|
|
@ -265,23 +277,10 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
custom_guardrail.should_run_guardrail(
|
||||
data=data_root, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is False
|
||||
is True
|
||||
)
|
||||
|
||||
# Test 2: guardrail in the opt-out list inside litellm_metadata → skipped
|
||||
data_litellm = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"litellm_metadata": {"opted_out_global_guardrails": ["global_guardrail"]},
|
||||
}
|
||||
assert (
|
||||
custom_guardrail.should_run_guardrail(
|
||||
data=data_litellm, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
# Test 3: guardrail in the opt-out list inside metadata → skipped
|
||||
# Test 2: User-injected opt-out in metadata is IGNORED
|
||||
data_metadata = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
|
|
@ -291,7 +290,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
custom_guardrail.should_run_guardrail(
|
||||
data=data_metadata, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is False
|
||||
is True
|
||||
)
|
||||
|
||||
# Test 4: a different guardrail in the opt-out list → still runs
|
||||
|
|
@ -588,7 +587,9 @@ class TestGuardrailSensitiveFieldStripping:
|
|||
duration=1.0,
|
||||
)
|
||||
|
||||
logged_response = request_data["metadata"]["standard_logging_guardrail_information"][0]["guardrail_response"]
|
||||
logged_response = request_data["metadata"][
|
||||
"standard_logging_guardrail_information"
|
||||
][0]["guardrail_response"]
|
||||
assert "secret_fields" not in logged_response
|
||||
assert "sk-live-SHOULD-NOT-APPEAR" not in json.dumps(logged_response)
|
||||
|
||||
|
|
@ -599,7 +600,12 @@ class TestGuardrailSensitiveFieldStripping:
|
|||
|
||||
guardrail.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=[
|
||||
{"result": "ok", "secret_fields": {"raw_headers": {"authorization": "Bearer sk-secret"}}},
|
||||
{
|
||||
"result": "ok",
|
||||
"secret_fields": {
|
||||
"raw_headers": {"authorization": "Bearer sk-secret"}
|
||||
},
|
||||
},
|
||||
{"result": "also_ok"},
|
||||
],
|
||||
request_data=request_data,
|
||||
|
|
@ -608,6 +614,7 @@ class TestGuardrailSensitiveFieldStripping:
|
|||
)
|
||||
|
||||
import json
|
||||
|
||||
serialized = json.dumps(request_data)
|
||||
assert "secret_fields" not in serialized
|
||||
assert "sk-secret" not in serialized
|
||||
|
|
@ -621,21 +628,21 @@ class TestCustomGuardrailPassthroughSupport:
|
|||
"""
|
||||
Test that async_post_call_success_deployment_hook handles raw httpx.Response objects
|
||||
from passthrough endpoints without crashing with TypeError.
|
||||
|
||||
|
||||
This tests Fix #3: TypeError: TypedDict does not support instance and class checks
|
||||
"""
|
||||
import httpx
|
||||
|
||||
custom_guardrail = CustomGuardrail()
|
||||
|
||||
|
||||
# Mock the async_post_call_success_hook to return None (guardrail didn't modify response)
|
||||
custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None)
|
||||
|
||||
|
||||
# Create a mock httpx.Response object (typical passthrough response)
|
||||
mock_response = AsyncMock(spec=httpx.Response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.text = "Mock response"
|
||||
|
||||
|
||||
request_data = {
|
||||
"guardrails": ["test_guardrail"],
|
||||
"user_api_key_user_id": "test_user",
|
||||
|
|
@ -644,14 +651,14 @@ class TestCustomGuardrailPassthroughSupport:
|
|||
"user_api_key_hash": "test_hash",
|
||||
"user_api_key_request_route": "passthrough_route",
|
||||
}
|
||||
|
||||
|
||||
# This should not raise TypeError: TypedDict does not support instance and class checks
|
||||
result = await custom_guardrail.async_post_call_success_deployment_hook(
|
||||
request_data=request_data,
|
||||
response=mock_response,
|
||||
call_type=CallTypes.allm_passthrough_route,
|
||||
)
|
||||
|
||||
|
||||
# When result is None, should return the original response
|
||||
assert result == mock_response
|
||||
|
||||
|
|
@ -659,53 +666,53 @@ class TestCustomGuardrailPassthroughSupport:
|
|||
async def test_async_post_call_success_deployment_hook_with_none_call_type(self):
|
||||
"""
|
||||
Test that async_post_call_success_deployment_hook handles None call_type gracefully.
|
||||
|
||||
|
||||
This ensures that even if call_type is None (before fix #1), the guardrail doesn't crash.
|
||||
"""
|
||||
custom_guardrail = CustomGuardrail()
|
||||
|
||||
|
||||
# Mock the async_post_call_success_hook to return None
|
||||
custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None)
|
||||
|
||||
|
||||
mock_response = AsyncMock()
|
||||
|
||||
|
||||
request_data = {
|
||||
"guardrails": ["test_guardrail"],
|
||||
"user_api_key_user_id": "test_user",
|
||||
}
|
||||
|
||||
|
||||
# Call with None call_type - should not crash
|
||||
result = await custom_guardrail.async_post_call_success_deployment_hook(
|
||||
request_data=request_data,
|
||||
response=mock_response,
|
||||
call_type=None,
|
||||
)
|
||||
|
||||
|
||||
# Should return the original response when result is None
|
||||
assert result == mock_response
|
||||
|
||||
def test_is_valid_response_type_with_none(self):
|
||||
"""
|
||||
Test _is_valid_response_type helper method correctly identifies None as invalid.
|
||||
|
||||
|
||||
This is part of Fix #3: Safely handling TypedDict types that don't support isinstance checks.
|
||||
"""
|
||||
custom_guardrail = CustomGuardrail()
|
||||
|
||||
|
||||
# None should be invalid
|
||||
assert custom_guardrail._is_valid_response_type(None) is False
|
||||
|
||||
def test_is_valid_response_type_with_typeddict_error(self):
|
||||
"""
|
||||
Test _is_valid_response_type gracefully handles TypeError from TypedDict.
|
||||
|
||||
|
||||
This tests Fix #3: When isinstance() is called with TypedDict types, it raises TypeError.
|
||||
The method should catch this and allow the response through.
|
||||
"""
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
custom_guardrail = CustomGuardrail()
|
||||
|
||||
|
||||
# Create a valid LiteLLM response object
|
||||
response = ModelResponse(
|
||||
id="test-id",
|
||||
|
|
@ -714,13 +721,12 @@ class TestCustomGuardrailPassthroughSupport:
|
|||
model="test-model",
|
||||
object="chat.completion",
|
||||
)
|
||||
|
||||
|
||||
# This should return True (it's a valid response type or TypeError is caught)
|
||||
result = custom_guardrail._is_valid_response_type(response)
|
||||
assert result is True
|
||||
|
||||
|
||||
|
||||
class TestEventTypeLogging:
|
||||
"""Tests for event_type logging in guardrail information."""
|
||||
|
||||
|
|
@ -1014,7 +1020,9 @@ class TestTracingFieldsPopulation:
|
|||
guardrail_json_response="blocked",
|
||||
request_data=request_data,
|
||||
guardrail_status="guardrail_intervened",
|
||||
tracing_detail=GuardrailTracingDetail(policy_template="EU AI Act Article 5"),
|
||||
tracing_detail=GuardrailTracingDetail(
|
||||
policy_template="EU AI Act Article 5"
|
||||
),
|
||||
)
|
||||
|
||||
slg_list = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
|
|
|
|||
|
|
@ -82,13 +82,18 @@ def test_env_reference_in_litellm_params_metadata_raises():
|
|||
def test_non_string_values_are_not_flagged():
|
||||
kwargs = {
|
||||
"langsmith_sampling_rate": 0.5,
|
||||
"turn_off_message_logging": True,
|
||||
}
|
||||
|
||||
params = initialize_standard_callback_dynamic_params(kwargs)
|
||||
|
||||
assert params.get("langsmith_sampling_rate") == 0.5
|
||||
assert params.get("turn_off_message_logging") is True
|
||||
|
||||
|
||||
def test_turn_off_message_logging_not_extracted_from_request():
|
||||
"""turn_off_message_logging is admin-only — must not be settable via request."""
|
||||
kwargs = {"turn_off_message_logging": True}
|
||||
params = initialize_standard_callback_dynamic_params(kwargs)
|
||||
assert params.get("turn_off_message_logging") is None
|
||||
|
||||
|
||||
def test_empty_kwargs_returns_empty_params():
|
||||
|
|
|
|||
|
|
@ -62,9 +62,9 @@ def reset_constants_module():
|
|||
# Reload modules before test
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
|
||||
yield
|
||||
|
||||
|
||||
# Reload modules after test to clean up
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
|
@ -157,9 +157,9 @@ def test_experimental_ui_token_ignores_litellm_ui_session_duration(
|
|||
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
|
||||
now = get_utc_datetime()
|
||||
# Must be ~10 min, NOT 24h. If LITELLM_UI_SESSION_DURATION were incorrectly used, this would fail.
|
||||
assert expires <= now + timedelta(minutes=11), (
|
||||
"Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION"
|
||||
)
|
||||
assert expires <= now + timedelta(
|
||||
minutes=11
|
||||
), "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION"
|
||||
|
||||
|
||||
def test_get_experimental_ui_login_jwt_auth_token_invalid(
|
||||
|
|
@ -293,13 +293,15 @@ def test_get_cli_jwt_auth_token_custom_expiration(
|
|||
|
||||
# Set custom expiration to 48 hours
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
|
||||
|
||||
|
||||
# Reload the constants module to pick up the new env var
|
||||
importlib.reload(constants)
|
||||
# Also reload auth_checks to pick up the new constant value
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
||||
|
||||
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
||||
valid_sso_user_defined_values
|
||||
)
|
||||
|
||||
# Decrypt and verify token contents
|
||||
decrypted_token = decrypt_value_helper(
|
||||
|
|
@ -315,7 +317,6 @@ def test_get_cli_jwt_auth_token_custom_expiration(
|
|||
assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1)
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_internal_user_params_with_get_user_object(monkeypatch):
|
||||
"""Test that default_internal_user_params is used when creating a new user via get_user_object"""
|
||||
|
|
@ -436,7 +437,9 @@ async def test_get_user_object_upsert_includes_user_email():
|
|||
mock_prisma_client.db.litellm_usertable.create.assert_called_once()
|
||||
creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"]
|
||||
|
||||
assert "user_email" in creation_args, "user_email should be included when upserting a new user"
|
||||
assert (
|
||||
"user_email" in creation_args
|
||||
), "user_email should be included when upserting a new user"
|
||||
assert creation_args["user_email"] == "test@example.com"
|
||||
assert creation_args["user_id"] == "new_test_user"
|
||||
|
||||
|
|
@ -463,7 +466,9 @@ def test_log_budget_lookup_failure_skips_user_not_found():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
|
||||
)
|
||||
async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch):
|
||||
"""
|
||||
Test that _get_team_db_check correctly calls the `new_team` function
|
||||
|
|
@ -497,8 +502,12 @@ async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeyp
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
|
||||
async def test_get_team_db_check_does_not_call_new_team_if_exists(mock_new_team, monkeypatch):
|
||||
@patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
|
||||
)
|
||||
async def test_get_team_db_check_does_not_call_new_team_if_exists(
|
||||
mock_new_team, monkeypatch
|
||||
):
|
||||
"""
|
||||
Test that _get_team_db_check does NOT call the `new_team` function
|
||||
if the team already exists in the database.
|
||||
|
|
@ -541,8 +550,9 @@ async def test_vector_store_access_check_early_returns(
|
|||
if vector_store_registry:
|
||||
vector_store_registry.get_vector_store_ids_to_run.return_value = None
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch(
|
||||
"litellm.vector_store_registry", vector_store_registry
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
patch("litellm.vector_store_registry", vector_store_registry),
|
||||
):
|
||||
result = await vector_store_access_check(
|
||||
request_body=request_body,
|
||||
|
|
@ -639,8 +649,9 @@ async def test_vector_store_access_check_with_permissions():
|
|||
mock_vector_store_registry = MagicMock()
|
||||
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-1"]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch(
|
||||
"litellm.vector_store_registry", mock_vector_store_registry
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
||||
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
||||
):
|
||||
result = await vector_store_access_check(
|
||||
request_body=request_body,
|
||||
|
|
@ -653,8 +664,9 @@ async def test_vector_store_access_check_with_permissions():
|
|||
# Test with denied access
|
||||
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-3"]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch(
|
||||
"litellm.vector_store_registry", mock_vector_store_registry
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
||||
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await vector_store_access_check(
|
||||
|
|
@ -687,8 +699,9 @@ async def test_vector_store_access_check_with_team_permissions():
|
|||
"team-store-allowed"
|
||||
]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch(
|
||||
"litellm.vector_store_registry", mock_vector_store_registry
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
||||
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
||||
):
|
||||
result = await vector_store_access_check(
|
||||
request_body=request_body,
|
||||
|
|
@ -702,8 +715,9 @@ async def test_vector_store_access_check_with_team_permissions():
|
|||
"team-store-denied"
|
||||
]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch(
|
||||
"litellm.vector_store_registry", mock_vector_store_registry
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
||||
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await vector_store_access_check(
|
||||
|
|
@ -1598,12 +1612,15 @@ async def test_custom_auth_common_checks_opt_in():
|
|||
mock_request = MagicMock()
|
||||
|
||||
# Default (no flag) — common_checks should NOT be called
|
||||
with patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_common, patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{},
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_common,
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{},
|
||||
),
|
||||
):
|
||||
mock_common.return_value = True
|
||||
result = await _run_post_custom_auth_checks(
|
||||
|
|
@ -1616,12 +1633,15 @@ async def test_custom_auth_common_checks_opt_in():
|
|||
mock_common.assert_not_called()
|
||||
|
||||
# With flag=True — common_checks SHOULD be called
|
||||
with patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_common, patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"custom_auth_run_common_checks": True},
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_common,
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"custom_auth_run_common_checks": True},
|
||||
),
|
||||
):
|
||||
mock_common.return_value = True
|
||||
result = await _run_post_custom_auth_checks(
|
||||
|
|
@ -1660,9 +1680,7 @@ async def test_virtual_key_budget_check_reads_from_spend_counter():
|
|||
return 1.5
|
||||
return fallback_spend
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
||||
await _virtual_key_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
|
|
@ -1692,9 +1710,7 @@ async def test_virtual_key_budget_check_fallback_no_counter():
|
|||
async def mock_get_current_spend(counter_key, fallback_spend):
|
||||
return fallback_spend
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
||||
await _virtual_key_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
|
|
@ -1723,9 +1739,7 @@ async def test_team_budget_check_reads_from_spend_counter():
|
|||
return 1.5
|
||||
return fallback_spend
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
||||
await _team_max_budget_check(
|
||||
team_object=team_object,
|
||||
|
|
@ -1763,12 +1777,13 @@ async def test_team_member_budget_check_reads_from_spend_counter():
|
|||
return 1.5
|
||||
return fallback_spend
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
|
||||
), patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_membership",
|
||||
new_callable=AsyncMock,
|
||||
return_value=team_membership,
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_membership",
|
||||
new_callable=AsyncMock,
|
||||
return_value=team_membership,
|
||||
),
|
||||
):
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
||||
await _check_team_member_budget(
|
||||
|
|
@ -1780,3 +1795,142 @@ async def test_team_member_budget_check_reads_from_spend_counter():
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
assert exc_info.value.current_cost == 1.5
|
||||
|
||||
|
||||
class TestGuardrailModificationCheck:
|
||||
"""Defense-in-depth: `_guardrail_modification_check` must 403 when the
|
||||
caller's metadata attempts to modify any guardrail-related key and the
|
||||
team lacks the `modify_guardrails` permission. Checks both the
|
||||
historically-covered `guardrails` list and the bypass toggles that
|
||||
`_get_admin_metadata` silently ignores at read time.
|
||||
"""
|
||||
|
||||
def _call(self, request_body):
|
||||
from litellm.proxy.auth.auth_checks import _guardrail_modification_check
|
||||
|
||||
team_object = MagicMock()
|
||||
team_object.metadata = {} # no permission
|
||||
return _guardrail_modification_check(
|
||||
request_body=request_body, team_object=team_object
|
||||
)
|
||||
|
||||
def test_noop_when_no_guardrail_keys_present(self):
|
||||
# no-op — should return silently
|
||||
self._call({"metadata": {"unrelated": "value"}})
|
||||
|
||||
def test_rejects_guardrails_list(self):
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call({"metadata": {"guardrails": ["custom"]}})
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_rejects_disable_global_guardrails_plural(self):
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call({"metadata": {"disable_global_guardrails": True}})
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_rejects_disable_global_guardrail_singular(self):
|
||||
"""VERIA-28's originally-reported singular-key typo variant."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call({"metadata": {"disable_global_guardrail": True}})
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_rejects_opted_out_global_guardrails(self):
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call(
|
||||
{"metadata": {"opted_out_global_guardrails": ["some_guardrail"]}}
|
||||
)
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_rejects_injection_via_litellm_metadata_key(self):
|
||||
"""Caller can populate the OTHER metadata key; that must also 403."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call({"litellm_metadata": {"disable_global_guardrails": True}})
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_rejects_root_level_injection(self):
|
||||
"""Top-level injection (`request_body["disable_global_guardrails"]`)
|
||||
was VERIA-28's easiest variant to hit — keep it rejected."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call({"disable_global_guardrails": True})
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_allows_when_team_has_permission(self):
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=True,
|
||||
):
|
||||
# no-op, should not raise
|
||||
self._call({"metadata": {"disable_global_guardrails": True}})
|
||||
|
||||
def test_rejects_string_encoded_metadata_bypass(self):
|
||||
"""Regression: attacker sends metadata as JSON string to bypass the
|
||||
isinstance(dict) guard. The check must coerce the string to dict
|
||||
and evaluate guardrail modification keys inside it."""
|
||||
import json as _json
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
attacker_payload = {"disable_global_guardrails": True}
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call({"metadata": _json.dumps(attacker_payload)})
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_rejects_string_encoded_litellm_metadata_bypass(self):
|
||||
"""Same bypass via the litellm_metadata key."""
|
||||
import json as _json
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
attacker_payload = {"guardrails": ["evaded"]}
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
||||
return_value=False,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
self._call({"litellm_metadata": _json.dumps(attacker_payload)})
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
def test_noop_when_string_is_not_json_object(self):
|
||||
"""Unparseable strings should not trigger a 403 — they have no keys."""
|
||||
self._call({"metadata": "not-json"})
|
||||
self._call({"metadata": '"just a string"'})
|
||||
|
|
|
|||
|
|
@ -835,3 +835,41 @@ def test_safe_get_request_headers_state_unavailable():
|
|||
|
||||
result = _safe_get_request_headers(mock_request)
|
||||
assert result == {"content-type": "application/json"}
|
||||
|
||||
|
||||
class TestGetTagsFromRequestBodyStringCoerce:
|
||||
"""Regression: the auth-time tag helper used `metadata.get("tags", ...)`
|
||||
directly, which raised AttributeError when metadata arrived as a JSON
|
||||
string (multipart/form-data or extra_body). That turned into a DoS at
|
||||
auth time and potentially bypassed tag-based RBAC if the caller caught
|
||||
the exception and fell through with empty tags.
|
||||
"""
|
||||
|
||||
def test_json_string_metadata_is_coerced_to_dict(self):
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
get_tags_from_request_body,
|
||||
)
|
||||
|
||||
metadata_json = json.dumps({"tags": ["a", "b"]})
|
||||
# Must not raise
|
||||
tags = get_tags_from_request_body({"metadata": metadata_json})
|
||||
assert tags == ["a", "b"]
|
||||
|
||||
def test_unparseable_string_metadata_is_ignored(self):
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
get_tags_from_request_body,
|
||||
)
|
||||
|
||||
# Must not raise; must yield no metadata tags but keep root tags
|
||||
tags = get_tags_from_request_body(
|
||||
{"metadata": "not-json", "tags": ["root-only"]}
|
||||
)
|
||||
assert tags == ["root-only"]
|
||||
|
||||
def test_dict_metadata_still_works(self):
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
get_tags_from_request_body,
|
||||
)
|
||||
|
||||
tags = get_tags_from_request_body({"metadata": {"tags": ["x"]}})
|
||||
assert tags == ["x"]
|
||||
|
|
|
|||
|
|
@ -1295,7 +1295,6 @@ def test_create_file_with_nested_litellm_metadata(
|
|||
"target_model_names": "gpt-3.5-turbo",
|
||||
"litellm_metadata[spend_logs_metadata][owner]": "john_doe",
|
||||
"litellm_metadata[spend_logs_metadata][team]": "engineering",
|
||||
"litellm_metadata[tags]": "production",
|
||||
"litellm_metadata[environment]": "prod",
|
||||
},
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
|
|
@ -1306,11 +1305,12 @@ def test_create_file_with_nested_litellm_metadata(
|
|||
result = response.json()
|
||||
assert result["id"] == "file-test-123"
|
||||
|
||||
# Verify nested metadata was correctly parsed
|
||||
# Verify nested metadata was correctly parsed.
|
||||
# Note: caller-supplied `tags` is stripped by default; test removed
|
||||
# to keep the parsing test focused on parser correctness.
|
||||
assert "spend_logs_metadata" in captured_litellm_metadata
|
||||
assert captured_litellm_metadata["spend_logs_metadata"]["owner"] == "john_doe"
|
||||
assert captured_litellm_metadata["spend_logs_metadata"]["team"] == "engineering"
|
||||
assert captured_litellm_metadata["tags"] == "production"
|
||||
assert captured_litellm_metadata["environment"] == "prod"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -207,6 +207,580 @@ async def test_add_litellm_data_to_request_parses_string_metadata():
|
|||
assert updated_data["metadata"]["generation_name"] == "gen123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_strips_admin_injection_slots():
|
||||
"""User-supplied user_api_key_metadata / user_api_key_team_metadata /
|
||||
_pipeline_managed_guardrails must be stripped from both metadata keys
|
||||
before the proxy writes its own admin-populated values. Otherwise a
|
||||
caller can shadow admin config via the non-`_metadata_variable_name`
|
||||
metadata key (e.g. litellm_metadata while the proxy writes to metadata).
|
||||
"""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
# Caller tries to inject admin config into BOTH metadata keys
|
||||
attacker_admin_payload = {"disable_global_guardrails": True}
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {
|
||||
"user_api_key_metadata": attacker_admin_payload,
|
||||
"user_api_key_team_metadata": attacker_admin_payload,
|
||||
"_pipeline_managed_guardrails": ["evaded"],
|
||||
},
|
||||
"litellm_metadata": {
|
||||
"user_api_key_metadata": attacker_admin_payload,
|
||||
"user_api_key_team_metadata": attacker_admin_payload,
|
||||
"_pipeline_managed_guardrails": ["evaded"],
|
||||
},
|
||||
}
|
||||
|
||||
real_admin_metadata = {"admin_flag": "from_proxy"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata=real_admin_metadata,
|
||||
team_metadata=real_admin_metadata,
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
# The key that matches `_metadata_variable_name` gets proxy-populated
|
||||
# with the real admin payload; the OTHER key must not retain the
|
||||
# attacker's injection.
|
||||
populated = updated["metadata"]
|
||||
assert populated["user_api_key_metadata"] == real_admin_metadata
|
||||
assert populated["user_api_key_team_metadata"] == real_admin_metadata
|
||||
assert "_pipeline_managed_guardrails" not in populated or populated[
|
||||
"_pipeline_managed_guardrails"
|
||||
] != ["evaded"]
|
||||
|
||||
other = updated.get("litellm_metadata") or {}
|
||||
assert other.get("user_api_key_metadata") in (None, {}, real_admin_metadata)
|
||||
assert other.get("user_api_key_team_metadata") in (None, {}, real_admin_metadata)
|
||||
assert "_pipeline_managed_guardrails" not in other
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_strips_all_user_api_key_prefix_keys():
|
||||
"""Strip must cover the full user_api_key_* family, not a hand-maintained
|
||||
list of 2-3 names. Proxy writes a dozen such fields (user_id, alias,
|
||||
spend, team_id, request_route, …) and an attacker populating any of them
|
||||
in the non-authoritative metadata key would otherwise forge identity /
|
||||
spend in audit logs and guardrails."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
attacker_injected = {
|
||||
"user_api_key_user_id": "victim",
|
||||
"user_api_key_alias": "admin-key",
|
||||
"user_api_key_spend": 0.0,
|
||||
"user_api_key_team_id": "victim-team",
|
||||
"user_api_key_end_user_id": "victim-user",
|
||||
"user_api_key_request_route": "/fake/route",
|
||||
"user_api_key_hash": "fake-hash",
|
||||
}
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {**attacker_injected},
|
||||
"litellm_metadata": {**attacker_injected},
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
user_id="real-user",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
spend=42.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
# The non-authoritative metadata dict must not retain ANY attacker-injected
|
||||
# user_api_key_* key.
|
||||
other = updated.get("litellm_metadata") or {}
|
||||
attacker_leaks = [k for k in other if k.startswith("user_api_key_")]
|
||||
assert attacker_leaks == [], f"Unexpected leaked keys: {attacker_leaks}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_string_metadata_does_not_crash():
|
||||
"""Regression: pre-strip code that pre-populated data['metadata'][k]=v
|
||||
before the string-to-dict parse would crash on JSON-string metadata.
|
||||
The snapshot / strip / admin-population pipeline must survive metadata
|
||||
arriving as a string."""
|
||||
import json as _json
|
||||
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "multipart/form-data"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": _json.dumps({"generation_name": "test"}),
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
# Must not raise TypeError / AttributeError.
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
# The parsed metadata should be a dict and the proxy snapshot body
|
||||
# should have been taken AFTER the strip (so no leaked user_api_key_*
|
||||
# from a raw string snapshot).
|
||||
assert isinstance(updated["metadata"], dict)
|
||||
assert updated["metadata"].get("generation_name") == "test"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_proxy_server_request_body_is_post_strip():
|
||||
"""Regression: proxy_server_request['body'] used to be snapshotted before
|
||||
the admin-slot strip, so standard_logging_object and spend-tracking
|
||||
readers saw attacker-injected payload. Snapshot must now be post-strip."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {"user_api_key_user_id": "victim"},
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
user_id="real-user",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
snapshot_body = updated["proxy_server_request"]["body"]
|
||||
assert snapshot_body is not None
|
||||
snapshot_metadata = snapshot_body.get("metadata") or {}
|
||||
assert "user_api_key_user_id" not in snapshot_metadata or (
|
||||
snapshot_metadata["user_api_key_user_id"] != "victim"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection():
|
||||
"""Regression: metadata arriving as a JSON string (multipart/form-data or
|
||||
extra_body) must not bypass the admin-injection strip. The parse happens
|
||||
AFTER receipt, so the strip has to run after the parse, not before.
|
||||
"""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "multipart/form-data"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
# Attacker encodes an admin-injection payload inside a JSON string.
|
||||
attacker_payload = {
|
||||
"user_api_key_metadata": {"disable_global_guardrails": True},
|
||||
"user_api_key_team_metadata": {"disable_global_guardrails": True},
|
||||
"_pipeline_managed_guardrails": ["evaded"],
|
||||
}
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": json.dumps(attacker_payload),
|
||||
"litellm_metadata": json.dumps(attacker_payload),
|
||||
}
|
||||
|
||||
real_admin_metadata = {"admin_flag": "from_proxy"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata=real_admin_metadata,
|
||||
team_metadata=real_admin_metadata,
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
populated = updated["metadata"]
|
||||
# The real admin payload from user_api_key_dict wins.
|
||||
assert populated["user_api_key_metadata"] == real_admin_metadata
|
||||
assert populated["user_api_key_team_metadata"] == real_admin_metadata
|
||||
assert populated.get("_pipeline_managed_guardrails") != ["evaded"]
|
||||
|
||||
other = updated.get("litellm_metadata") or {}
|
||||
# After the strip, litellm_metadata has no admin-injection slots.
|
||||
assert "user_api_key_metadata" not in other
|
||||
assert "user_api_key_team_metadata" not in other
|
||||
assert "_pipeline_managed_guardrails" not in other
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_ignores_x_litellm_tags_header_without_permission():
|
||||
"""Regression: the `x-litellm-tags` header bypassed the body-metadata
|
||||
tag strip. Header tags must also be gated by `allow_client_tags`."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {
|
||||
"Content-Type": "application/json",
|
||||
"x-litellm-tags": "restricted-tier,victim-team",
|
||||
}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert "tags" not in (updated.get("metadata") or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_ignores_root_level_tags_without_permission():
|
||||
"""Regression: root-level `data["tags"]` bypassed the body-metadata
|
||||
tag strip. Root-level tags must also be gated by `allow_client_tags`."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"tags": ["restricted-tier", "victim-team"],
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert "tags" not in (updated.get("metadata") or {})
|
||||
# Also ensure the root-level tags are removed. get_tags_from_request_body
|
||||
# reads request_body["tags"] directly, so leaving it in place would let
|
||||
# the policy engine see caller-supplied tags even after the metadata
|
||||
# strip.
|
||||
assert "tags" not in updated
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_honors_header_tags_when_opted_in():
|
||||
"""When allow_client_tags=True, header-supplied tags flow through."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {
|
||||
"Content-Type": "application/json",
|
||||
"x-litellm-tags": "production,ab-test",
|
||||
}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"allow_client_tags": True},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert updated["metadata"].get("tags") == ["production", "ab-test"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_strips_user_tags_without_permission():
|
||||
"""Caller-supplied metadata.tags must be stripped when the key/team
|
||||
metadata does not opt in via allow_client_tags=True. Otherwise an
|
||||
attacker can reach restricted tag-routed deployments or attribute
|
||||
spend to a victim team's tag."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {"tags": ["restricted-tier", "victim-team"]},
|
||||
"litellm_metadata": {"tags": ["also-stripped"]},
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert "tags" not in (updated.get("metadata") or {})
|
||||
assert "tags" not in (updated.get("litellm_metadata") or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_preserves_user_tags_when_key_opts_in():
|
||||
"""When key.metadata.allow_client_tags=True, caller-supplied tags are
|
||||
preserved and reach the router."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {"tags": ["opted-in-tag"]},
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"allow_client_tags": True},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert updated["metadata"].get("tags") == ["opted-in-tag"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_preserves_user_tags_when_team_opts_in():
|
||||
"""Team-level allow_client_tags is also honored (not just key-level)."""
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
request_mock.url = MagicMock()
|
||||
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
|
||||
request_mock.method = "POST"
|
||||
request_mock.query_params = {}
|
||||
request_mock.headers = {"Content-Type": "application/json"}
|
||||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {"tags": ["team-allowed"]},
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
team_metadata={"allow_client_tags": True},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
model_max_budget={},
|
||||
team_spend=0.0,
|
||||
team_max_budget=200.0,
|
||||
)
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request_mock,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert updated["metadata"].get("tags") == ["team-allowed"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_user_spend_and_budget():
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
|
|
@ -221,7 +795,10 @@ async def test_add_litellm_data_to_request_user_spend_and_budget():
|
|||
request_mock.client = MagicMock()
|
||||
request_mock.client.host = "127.0.0.1"
|
||||
|
||||
data = {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
|
|
@ -275,9 +852,11 @@ async def test_add_litellm_data_to_request_audio_transcription_multipart():
|
|||
"file": b"Fake audio bytes",
|
||||
}
|
||||
|
||||
# Opt the key in to client-supplied tags so the parsed tags from the
|
||||
# JSON-string multipart body aren't stripped by the admin-injection strip.
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
metadata={"allow_client_tags": True},
|
||||
team_metadata={},
|
||||
spend=0.0,
|
||||
max_budget=100.0,
|
||||
|
|
@ -1023,6 +1602,7 @@ def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data():
|
|||
# Restore original model_group_settings
|
||||
litellm.model_group_settings = original_model_group_settings
|
||||
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Optional
|
||||
|
|
@ -1040,15 +1620,16 @@ class TestCustomLogger(CustomLogger):
|
|||
def __init__(self):
|
||||
self.standard_logging_object: Optional[StandardLoggingPayload] = None
|
||||
super().__init__()
|
||||
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
print(f"SUCCESS CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}")
|
||||
self.standard_logging_object = kwargs.get("standard_logging_object")
|
||||
print(f"Captured standard_logging_object: {self.standard_logging_object}")
|
||||
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
print(f"FAILURE CALLBACK CALLED! kwargs keys: {list(kwargs.keys())}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_metadata_from_request_headers():
|
||||
"""
|
||||
|
|
@ -1065,8 +1646,16 @@ async def test_add_litellm_metadata_from_request_headers():
|
|||
|
||||
try:
|
||||
# Prepare test data (ensure no streaming, add mock_response and api_key to route to litellm.acompletion)
|
||||
headers = {"x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}'}
|
||||
data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "stream": False, "mock_response": "Hi", "api_key": "fake-key"}
|
||||
headers = {
|
||||
"x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}'
|
||||
}
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"stream": False,
|
||||
"mock_response": "Hi",
|
||||
"api_key": "fake-key",
|
||||
}
|
||||
|
||||
# Create mock request with headers
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
|
@ -1078,9 +1667,7 @@ async def test_add_litellm_metadata_from_request_headers():
|
|||
|
||||
# Create mock user API key dict
|
||||
mock_user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
org_id="test-org"
|
||||
api_key="test-key", user_id="test-user", org_id="test-org"
|
||||
)
|
||||
|
||||
# Create mock proxy logging object
|
||||
|
|
@ -1095,7 +1682,7 @@ async def test_add_litellm_metadata_from_request_headers():
|
|||
|
||||
async def mock_post_call_success_hook(*args, **kwargs):
|
||||
# Return the response unchanged
|
||||
return kwargs.get('response', args[2] if len(args) > 2 else None)
|
||||
return kwargs.get("response", args[2] if len(args) > 2 else None)
|
||||
|
||||
mock_proxy_logging_obj.during_call_hook = mock_during_call_hook
|
||||
mock_proxy_logging_obj.pre_call_hook = mock_pre_call_hook
|
||||
|
|
@ -1108,10 +1695,15 @@ async def test_add_litellm_metadata_from_request_headers():
|
|||
general_settings = {}
|
||||
|
||||
# Create mock select_data_generator with correct signature
|
||||
def mock_select_data_generator(response=None, user_api_key_dict=None, request_data=None):
|
||||
def mock_select_data_generator(
|
||||
response=None, user_api_key_dict=None, request_data=None
|
||||
):
|
||||
async def mock_generator():
|
||||
yield "data: " + json.dumps({"choices": [{"delta": {"content": "Hello"}}]}) + "\n\n"
|
||||
yield "data: " + json.dumps(
|
||||
{"choices": [{"delta": {"content": "Hello"}}]}
|
||||
) + "\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return mock_generator()
|
||||
|
||||
# Create the processor
|
||||
|
|
@ -1129,22 +1721,28 @@ async def test_add_litellm_metadata_from_request_headers():
|
|||
select_data_generator=mock_select_data_generator,
|
||||
llm_router=None,
|
||||
model="gpt-4",
|
||||
is_streaming_request=False
|
||||
is_streaming_request=False,
|
||||
)
|
||||
|
||||
# Sleep for 3 seconds to allow logging to complete
|
||||
await asyncio.sleep(3)
|
||||
|
||||
# Check if standard_logging_object was set
|
||||
assert test_logger.standard_logging_object is not None, "standard_logging_object should be populated after LLM request"
|
||||
assert (
|
||||
test_logger.standard_logging_object is not None
|
||||
), "standard_logging_object should be populated after LLM request"
|
||||
|
||||
# Verify the logging object contains expected metadata
|
||||
standard_logging_obj = test_logger.standard_logging_object
|
||||
|
||||
print(f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}")
|
||||
print(
|
||||
f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}"
|
||||
)
|
||||
|
||||
SPEND_LOGS_METADATA = standard_logging_obj["metadata"]["spend_logs_metadata"]
|
||||
assert SPEND_LOGS_METADATA == dict(json.loads(headers["x-litellm-spend-logs-metadata"])), "spend_logs_metadata should be the same as the headers"
|
||||
assert SPEND_LOGS_METADATA == dict(
|
||||
json.loads(headers["x-litellm-spend-logs-metadata"])
|
||||
), "spend_logs_metadata should be the same as the headers"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
|
@ -1197,7 +1795,9 @@ def test_get_internal_user_header_from_mapping_returns_expected_header():
|
|||
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
|
||||
]
|
||||
|
||||
header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings)
|
||||
header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(
|
||||
mappings
|
||||
)
|
||||
assert header_name == "X-OpenWebUI-User-Id"
|
||||
|
||||
|
||||
|
|
@ -1205,7 +1805,9 @@ def test_get_internal_user_header_from_mapping_none_when_absent():
|
|||
mappings = [
|
||||
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"}
|
||||
]
|
||||
header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(mappings)
|
||||
header_name = LiteLLMProxyRequestSetup.get_internal_user_header_from_mapping(
|
||||
mappings
|
||||
)
|
||||
assert header_name is None
|
||||
|
||||
single = {"header_name": "X-Only-Customer", "litellm_user_role": "customer"}
|
||||
|
|
@ -1218,7 +1820,10 @@ def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present():
|
|||
headers = {"X-OpenWebUI-User-Id": "internal-user-123"}
|
||||
general_settings = {
|
||||
"user_header_mappings": [
|
||||
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},
|
||||
{
|
||||
"header_name": "X-OpenWebUI-User-Id",
|
||||
"litellm_user_role": "internal_user",
|
||||
},
|
||||
{"header_name": "X-OpenWebUI-User-Email", "litellm_user_role": "customer"},
|
||||
]
|
||||
}
|
||||
|
|
@ -1312,7 +1917,7 @@ async def test_team_guardrails_append_to_key_guardrails():
|
|||
|
||||
metadata = updated_data.get("metadata", {})
|
||||
guardrails = metadata.get("guardrails", [])
|
||||
|
||||
|
||||
assert "key-guardrail-1" in guardrails
|
||||
assert "key-guardrail-2" in guardrails
|
||||
assert "team-guardrail-1" in guardrails
|
||||
|
|
@ -1341,7 +1946,7 @@ async def test_request_guardrails_do_not_override_key_guardrails():
|
|||
metadata={"guardrails": ["key-guardrail-1"]},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
|
||||
# Test case: Request with empty guardrails should not result in empty guardrails
|
||||
data_with_empty = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
|
|
@ -1361,7 +1966,7 @@ async def test_request_guardrails_do_not_override_key_guardrails():
|
|||
|
||||
_metadata = updated_data_empty.get("metadata", {})
|
||||
requested_guardrails = _metadata.get("guardrails", [])
|
||||
|
||||
|
||||
assert "guardrails" not in updated_data_empty
|
||||
assert "key-guardrail-1" in requested_guardrails
|
||||
assert len(requested_guardrails) == 1
|
||||
|
|
@ -1476,7 +2081,10 @@ def test_update_model_if_key_alias_exists():
|
|||
assert data["model"] == "xai/grok-4-fast-non-reasoning"
|
||||
|
||||
# Test case 2: Key alias doesn't exist
|
||||
data = {"model": "unknown-model", "messages": [{"role": "user", "content": "Hello"}]}
|
||||
data = {
|
||||
"model": "unknown-model",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
aliases={"modelAlias": "xai/grok-4-fast-non-reasoning"},
|
||||
|
|
@ -1594,16 +2202,22 @@ async def test_embedding_header_forwarding_with_model_group():
|
|||
|
||||
# Verify that only x- prefixed headers (except x-stainless) were forwarded
|
||||
forwarded_headers = updated_data["headers"]
|
||||
assert "X-Custom-Header" in forwarded_headers, "X-Custom-Header should be forwarded"
|
||||
assert (
|
||||
"X-Custom-Header" in forwarded_headers
|
||||
), "X-Custom-Header should be forwarded"
|
||||
assert forwarded_headers["X-Custom-Header"] == "custom-value"
|
||||
assert "X-Request-ID" in forwarded_headers, "X-Request-ID should be forwarded"
|
||||
assert forwarded_headers["X-Request-ID"] == "test-request-123"
|
||||
|
||||
# Verify that authorization header was NOT forwarded (sensitive header)
|
||||
assert "Authorization" not in forwarded_headers, "Authorization header should not be forwarded"
|
||||
assert (
|
||||
"Authorization" not in forwarded_headers
|
||||
), "Authorization header should not be forwarded"
|
||||
|
||||
# Verify that Content-Type was NOT forwarded (doesn't start with x-)
|
||||
assert "Content-Type" not in forwarded_headers, "Content-Type should not be forwarded"
|
||||
assert (
|
||||
"Content-Type" not in forwarded_headers
|
||||
), "Content-Type should not be forwarded"
|
||||
|
||||
# Verify original data fields are preserved
|
||||
assert updated_data["model"] == "local-openai/text-embedding-3-small"
|
||||
|
|
@ -1659,8 +2273,9 @@ async def test_embedding_header_forwarding_without_model_group_config():
|
|||
)
|
||||
|
||||
# Verify that headers were NOT added since model is not in forward list
|
||||
assert "headers" not in updated_data or updated_data.get("headers") is None, \
|
||||
"Headers should not be forwarded for models not in forward_client_headers_to_llm_api list"
|
||||
assert (
|
||||
"headers" not in updated_data or updated_data.get("headers") is None
|
||||
), "Headers should not be forwarded for models not in forward_client_headers_to_llm_api list"
|
||||
|
||||
# Verify original data fields are preserved
|
||||
assert updated_data["model"] == "text-embedding-ada-002"
|
||||
|
|
@ -1714,7 +2329,9 @@ async def test_add_guardrails_from_policy_engine():
|
|||
attachment_registry = get_attachment_registry()
|
||||
attachment_registry._attachments = [
|
||||
PolicyAttachment(policy="global-baseline", scope="*"), # applies to all
|
||||
PolicyAttachment(policy="healthcare", teams=["healthcare-team"]), # applies to healthcare team
|
||||
PolicyAttachment(
|
||||
policy="healthcare", teams=["healthcare-team"]
|
||||
), # applies to healthcare team
|
||||
]
|
||||
attachment_registry._initialized = True
|
||||
|
||||
|
|
@ -1757,7 +2374,10 @@ async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_po
|
|||
data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"policies": ["PII-POLICY-GLOBAL", "HIPAA-POLICY"], # Dynamic policies - should be accepted and removed
|
||||
"policies": [
|
||||
"PII-POLICY-GLOBAL",
|
||||
"HIPAA-POLICY",
|
||||
], # Dynamic policies - should be accepted and removed
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
|
|
@ -1780,7 +2400,9 @@ async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_po
|
|||
)
|
||||
|
||||
# Verify that 'policies' was removed from the request body
|
||||
assert "policies" not in data, "'policies' should be removed from request body to prevent forwarding to LLM provider"
|
||||
assert (
|
||||
"policies" not in data
|
||||
), "'policies' should be removed from request body to prevent forwarding to LLM provider"
|
||||
|
||||
# Verify that other fields are preserved
|
||||
assert "model" in data
|
||||
|
|
@ -1869,7 +2491,9 @@ async def test_bearer_token_not_in_debug_logs():
|
|||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
secret_token = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.fakesignature"
|
||||
secret_token = (
|
||||
"eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.fakesignature"
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.headers = {
|
||||
|
|
@ -1898,8 +2522,10 @@ async def test_bearer_token_not_in_debug_logs():
|
|||
logger.setLevel(logging.DEBUG)
|
||||
|
||||
try:
|
||||
with patch("litellm.proxy.proxy_server.llm_router", None), \
|
||||
patch("litellm.proxy.proxy_server.premium_user", True):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", None),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
):
|
||||
await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=mock_request,
|
||||
|
|
@ -2020,9 +2646,7 @@ def test_resolve_project_model_specific_wins():
|
|||
"gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}},
|
||||
"defaultconfig": {"azure": {"litellm_credentials": "team-default"}},
|
||||
}
|
||||
result = _resolve_credential_from_model_config(
|
||||
"gpt-4", project_config, team_config
|
||||
)
|
||||
result = _resolve_credential_from_model_config("gpt-4", project_config, team_config)
|
||||
assert result == "proj-gpt4"
|
||||
|
||||
|
||||
|
|
@ -2034,9 +2658,7 @@ def test_resolve_project_default_wins_over_team():
|
|||
"gpt-4": {"azure": {"litellm_credentials": "team-gpt4"}},
|
||||
"defaultconfig": {"azure": {"litellm_credentials": "team-default"}},
|
||||
}
|
||||
result = _resolve_credential_from_model_config(
|
||||
"gpt-4", project_config, team_config
|
||||
)
|
||||
result = _resolve_credential_from_model_config("gpt-4", project_config, team_config)
|
||||
assert result == "proj-default"
|
||||
|
||||
|
||||
|
|
@ -2091,12 +2713,8 @@ def test_apply_overrides_project_model_specific(setup_test_credentials):
|
|||
},
|
||||
project_metadata={
|
||||
"model_config": {
|
||||
"defaultconfig": {
|
||||
"azure": {"litellm_credentials": "hotel-rec-azure"}
|
||||
},
|
||||
"gpt-4-vision": {
|
||||
"azure": {"litellm_credentials": "hotel-rec-vision"}
|
||||
},
|
||||
"defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}},
|
||||
"gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}},
|
||||
}
|
||||
},
|
||||
)
|
||||
|
|
@ -2123,12 +2741,8 @@ def test_apply_overrides_project_default(setup_test_credentials):
|
|||
},
|
||||
project_metadata={
|
||||
"model_config": {
|
||||
"defaultconfig": {
|
||||
"azure": {"litellm_credentials": "hotel-rec-azure"}
|
||||
},
|
||||
"gpt-4-vision": {
|
||||
"azure": {"litellm_credentials": "hotel-rec-vision"}
|
||||
},
|
||||
"defaultconfig": {"azure": {"litellm_credentials": "hotel-rec-azure"}},
|
||||
"gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}},
|
||||
}
|
||||
},
|
||||
)
|
||||
|
|
@ -2231,9 +2845,7 @@ def test_apply_overrides_missing_credential_name(setup_test_credentials):
|
|||
api_key="test-key",
|
||||
team_metadata={
|
||||
"model_config": {
|
||||
"gpt-4": {
|
||||
"azure": {"litellm_credentials": "nonexistent-credential"}
|
||||
}
|
||||
"gpt-4": {"azure": {"litellm_credentials": "nonexistent-credential"}}
|
||||
}
|
||||
},
|
||||
)
|
||||
|
|
@ -2272,9 +2884,7 @@ def test_apply_overrides_no_model_in_data(setup_test_credentials):
|
|||
api_key="test-key",
|
||||
team_metadata={
|
||||
"model_config": {
|
||||
"defaultconfig": {
|
||||
"azure": {"litellm_credentials": "some-cred"}
|
||||
}
|
||||
"defaultconfig": {"azure": {"litellm_credentials": "some-cred"}}
|
||||
}
|
||||
},
|
||||
)
|
||||
|
|
@ -2305,9 +2915,7 @@ def test_apply_overrides_clientside_api_version_preserved(setup_test_credentials
|
|||
api_key="test-key",
|
||||
team_metadata={
|
||||
"model_config": {
|
||||
"gpt-4-vision": {
|
||||
"azure": {"litellm_credentials": "hotel-rec-vision"}
|
||||
}
|
||||
"gpt-4-vision": {"azure": {"litellm_credentials": "hotel-rec-vision"}}
|
||||
}
|
||||
},
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue