From 1dbd73d5129454dd19b676afba284265233d147f Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Mon, 28 Sep 2026 18:02:45 +0300 Subject: [PATCH 1/4] fix(guardrails): make guardrails see the authenticated identity on every path Guardrails read user_api_key_alias, user_api_key_team_id, the key hash and the request route from request metadata, and the generic guardrail API forwards them to the vendor. The chat path strips caller copies of those fields and writes the real ones, but three other paths did not, so a caller could claim another key, team or request route (which call-type lookups key on) /guardrails/apply_guardrail passed the body metadata straight to the guardrail. It now drops the fields the chat path treats as untrusted, plus the bare user_api_key and the caller's headers (a slightly wider strip than the chat path, on purpose). Then it adds the authenticated identity and the proxy's real request headers Pass-through handed the raw body to pre_call_hook. It now runs the chat path's strip on both metadata buckets first, which also stops a body from switching off global guardrails, and drops the body's headers and proxy_server_request so it cannot pick the inbound headers a vendor sees. Those keys were already popped before the upstream send, so the forwarded body does not change. The pass-through guardrail text no longer includes litellm_metadata, which carried the key's identity dump into the scanned payload The unified guardrail only filled litellm_metadata when it was missing, so a caller-supplied bucket won. It now overwrites the identity fields of an existing bucket and drops any user_api_key_token there, since the proxy never writes one. It leaves user_api_key_auth_metadata alone because on litellm_metadata routes the proxy has already merged team metadata into it transform_user_api_key_dict_to_metadata used to dump every UserAPIKeyAuth field, including the raw token of non-sk keys, JWT claims, team membership, proxy config and org and project metadata with callback secrets. It now returns only the chat path's identity fields plus user_api_key_key_alias for existing readers, so MCP, pass-through, output-side handlers and realtime transcript guardrails all get that same identity allowlist and see the real user_api_key_alias. The generic guardrail uses user_api_key_token only from litellm_metadata and only when user_api_key_hash is missing, so a CLI session key's raw per-login token never reaches the vendor The strip and the identity fields live in one place in litellm_pre_call_utils, and the chat path calls them too --- .../litellm_core_utils/realtime_streaming.py | 15 +- .../guardrail_translation/base_translation.py | 40 +-- .../guardrail_translation/handler.py | 4 +- .../proxy/guardrails/guardrail_endpoints.py | 8 +- .../generic_guardrail_api.py | 5 +- .../panw_prisma_airs/panw_prisma_airs.py | 10 - .../unified_guardrail/unified_guardrail.py | 27 +- litellm/proxy/litellm_pre_call_utils.py | 90 ++++--- .../pass_through_endpoints.py | 14 ++ .../test_generic_guardrail_api.py | 16 ++ .../guardrail_hooks/test_grayswan.py | 15 +- .../test_unified_guardrail.py | 226 +++++++++++++++++ .../guardrails/test_guardrail_endpoints.py | 235 +++++++++++++++--- .../test_pass_through_endpoints.py | 78 ++++++ .../test_apply_guardrail_endpoint.py | 18 +- .../test_realtime_streaming.py | 131 ++++++++++ .../guardrail_translation/__init__.py | 0 .../test_base_translation.py | 53 ++++ 18 files changed, 850 insertions(+), 135 deletions(-) create mode 100644 tests/unit/llms/base_llm/guardrail_translation/__init__.py create mode 100644 tests/unit/llms/base_llm/guardrail_translation/test_base_translation.py diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index d2fbb26bb02..300834073f7 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -29,6 +29,7 @@ if TYPE_CHECKING: from websockets.asyncio.client import ClientConnection from websockets.exceptions import ConnectionClosed + from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks CLIENT_CONNECTION_CLASS = ClientConnection @@ -123,6 +124,12 @@ DefaultLoggedRealTimeEventTypes: Final = [ ] +def _as_user_api_key_auth(user_api_key_dict: object) -> "UserAPIKeyAuth | None": + from litellm.proxy._types import UserAPIKeyAuth + + return user_api_key_dict if isinstance(user_api_key_dict, UserAPIKeyAuth) else None + + class RealTimeStreaming: def __init__( self, @@ -831,6 +838,7 @@ class RealTimeStreaming: typed user messages and tool outputs use ``pre_call``. """ from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.types.guardrails import GuardrailEventHooks if event_hooks is None: @@ -852,7 +860,12 @@ class RealTimeStreaming: try: await callback.apply_guardrail( inputs={"texts": [transcript], "images": []}, - request_data={"user_api_key_dict": self.user_api_key_dict}, + request_data={ + "user_api_key_dict": self.user_api_key_dict, + "litellm_metadata": BaseTranslation.transform_user_api_key_dict_to_metadata( + _as_user_api_key_auth(self.user_api_key_dict) + ), + }, input_type="request", ) except Exception as e: diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 89ad67f0485..1a772ca8ee3 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -81,43 +81,17 @@ class BaseTranslation(ABC): @staticmethod def transform_user_api_key_dict_to_metadata( - user_api_key_dict: Any | None, + user_api_key_dict: "UserAPIKeyAuth | None", ) -> dict[str, object]: - """ - Transform user_api_key_dict to a metadata dict with prefixed keys. - - Converts keys like 'user_id' to 'user_api_key_user_id' to clearly indicate - the source of the metadata. - - Args: - user_api_key_dict: UserAPIKeyAuth object or dict with user information - - Returns: - Dict with keys prefixed with 'user_api_key_' - """ + """The authenticated key's identity as prefixed metadata, an allowlist safe to hand to guardrail vendors.""" if user_api_key_dict is None: return {} + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup - # Convert to dict if it's a Pydantic object - user_dict = user_api_key_dict.model_dump() if hasattr(user_api_key_dict, "model_dump") else user_api_key_dict - - if not isinstance(user_dict, dict): - return {} - - # Transform keys to be prefixed with 'user_api_key_' - transformed: Final[dict[str, object]] = {} - for key, value in user_dict.items(): - # Skip None values and internal fields - if value is None or key.startswith("_"): - continue - - # If key already has the prefix, use as-is, otherwise add prefix - if key.startswith("user_api_key_"): - transformed[key] = value - else: - transformed[f"user_api_key_{key}"] = value - - return transformed + return { + **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict), + "user_api_key_key_alias": user_api_key_dict.key_alias, + } @staticmethod def merge_user_api_key_metadata_into_request( diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index 1f295a6e656..ac7dcfd4f39 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -20,6 +20,8 @@ if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging +_PROXY_OWNED_PAYLOAD_KEYS: Final = frozenset({"metadata", "litellm_metadata", "litellm_logging_obj"}) + class PassThroughEndpointHandler(BaseTranslation): """ @@ -80,7 +82,7 @@ class PassThroughEndpointHandler(BaseTranslation): from litellm.litellm_core_utils.safe_json_dumps import safe_dumps payload_to_check: Final = { - k: v for k, v in data.items() if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj") + k: v for k, v in data.items() if not k.startswith("_") and k not in _PROXY_OWNED_PAYLOAD_KEYS } verbose_proxy_logger.debug("PassThroughEndpointHandler: Using full payload for guardrail") return safe_dumps(payload_to_check) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 6053ab26726..18d62142503 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -34,6 +34,7 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( ) from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router +from litellm.proxy.litellm_pre_call_utils import caller_metadata_with_authenticated_identity from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import GuardrailsRepository @@ -2404,9 +2405,14 @@ async def apply_guardrail( if litellm_logging_obj is not None: _patch_logging_obj_for_guardrail(litellm_logging_obj, request) + processed_metadata: Final = data.get("metadata") + inbound_headers: Final = processed_metadata.get("headers") if isinstance(processed_metadata, dict) else None request_data: Final[dict] = { **({"messages": request.messages} if request.messages is not None else {}), - **({"metadata": request.metadata} if request.metadata is not None else {}), + "metadata": { + **caller_metadata_with_authenticated_identity(request.metadata, user_api_key_dict), + **({"headers": inbound_headers} if inbound_headers is not None else {}), + }, } _input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type) guardrailed_inputs: Final = await active_guardrail.apply_guardrail( diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 3d1a173635e..320682c89b7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -292,9 +292,8 @@ class GenericGuardrailAPI(CustomGuardrail): if value is not None: result_metadata[field_name] = value - # handle user_api_key_token = user_api_key_hash - if metadata_dict.get("user_api_key_token") is not None: - result_metadata["user_api_key_hash"] = metadata_dict.get("user_api_key_token") + if litellm_metadata.get("user_api_key_token") is not None and "user_api_key_hash" not in result_metadata: + result_metadata["user_api_key_hash"] = litellm_metadata["user_api_key_token"] verbose_proxy_logger.debug( "Generic Guardrail API: Extracted user metadata: %s", diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index df5a265bb72..143e5818ee8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -1740,16 +1740,6 @@ class PanwPrismaAirsHandler(CustomGuardrail): call_id, _mcp_tool, ) - elif not request_data and logging_obj is None and input_type == "request": - # Direct /apply_guardrail endpoint — empty request_data, no - # logging_obj. Existing behavior: synthesize UUID. - call_id = str(uuid.uuid4()) - request_data["litellm_call_id"] = call_id - verbose_proxy_logger.warning( - "PANW Prisma AIRS: litellm_call_id missing from empty " - "request_data, synthesized %s (direct /apply_guardrail?)", - call_id, - ) else: call_id = str(uuid.uuid4()) request_data["litellm_call_id"] = call_id diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index d68a55f9a88..eeb604e98af 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -155,16 +155,25 @@ def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapp } -def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None: - """Populate data['litellm_metadata'] from user_api_key_dict if absent.""" - if "litellm_metadata" not in data: - from litellm.llms.base_llm.guardrail_translation.base_translation import ( - BaseTranslation, - ) +_PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_metadata"}) - user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict) - if user_metadata: - data["litellm_metadata"] = user_metadata + +def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None: + """Overwrite the identity fields of data['litellm_metadata'] from the authenticated key, in place.""" + from litellm.llms.base_llm.guardrail_translation.base_translation import ( + BaseTranslation, + ) + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup + + existing: Final = data.get("litellm_metadata") + if isinstance(existing, dict): + identity: Final = LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict) + existing.update({key: value for key, value in identity.items() if key not in _PROXY_ENRICHED_IDENTITY_FIELDS}) + existing.pop("user_api_key_token", None) + return + user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict) + if user_metadata: + data["litellm_metadata"] = user_metadata class UnifiedLLMGuardrails(CustomLogger): diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 866d84ca8f2..9201fcd7ca7 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -493,6 +493,40 @@ def _strip_untrusted_request_header_controls( headers.pop(header_name, None) +def is_untrusted_caller_metadata_key(key: str) -> bool: + return key.startswith("user_api_key_") or key in _UNTRUSTED_METADATA_CONTROL_FIELDS + + +def strip_untrusted_caller_metadata( + data: MutableMapping[str, object], *, allow_client_message_redaction_opt_out: bool +) -> None: + """Remove, in place, the proxy-owned slots a caller put in either metadata bucket of a request body.""" + for user_meta in (data.get("metadata"), data.get("litellm_metadata")): + if not isinstance(user_meta, dict): + continue + _strip_untrusted_request_header_controls( + user_meta.get("headers"), + allow_client_message_redaction_opt_out=allow_client_message_redaction_opt_out, + ) + for untrusted_key in tuple(key for key in user_meta if is_untrusted_caller_metadata_key(key)): + user_meta.pop(untrusted_key, None) + + +_GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS: Final = frozenset({"user_api_key", "headers"}) + + +def caller_metadata_with_authenticated_identity( + caller_metadata: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth +) -> dict[str, object]: + """Caller metadata minus proxy-owned slots, bare user_api_key and headers, with the key's identity on top.""" + caller_fields: Final = { + key: value + for key, value in (caller_metadata or {}).items() + if not (is_untrusted_caller_metadata_key(key) or key in _GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS) + } + return {**caller_fields, **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)} + + def _is_false_like(value: object) -> bool: if isinstance(value, bool): return value is False @@ -520,7 +554,7 @@ def _key_or_team_allows_client_mock_response( ) -def _key_or_team_allows_client_message_redaction_opt_out( +def key_or_team_allows_client_message_redaction_opt_out( user_api_key_dict: UserAPIKeyAuth, ) -> bool: return _key_or_team_metadata_flag_is_true( @@ -1642,6 +1676,24 @@ class LiteLLMProxyRequestSetup: ) return user_api_key_logged_metadata + @staticmethod + def get_key_scoped_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: + return { + "user_api_key_metadata": strip_callback_config(user_api_key_dict.metadata), + "user_api_key_team_metadata": strip_callback_config(user_api_key_dict.team_metadata), + "user_api_key_object_permission_id": user_api_key_dict.object_permission_id, + "user_api_key_team_object_permission_id": user_api_key_dict.team_object_permission_id, + } + + @staticmethod + def get_authenticated_identity_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: + """Identity fields derived from the authenticated key alone, for paths that skip the chat-path build.""" + return { + **LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict), + "user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict), + **LiteLLMProxyRequestSetup.get_key_scoped_metadata(user_api_key_dict), + } + @staticmethod def add_user_api_key_auth_to_request_metadata( data: dict, @@ -2001,7 +2053,7 @@ async def add_litellm_data_to_request( # These keys are injected by the proxy itself below — user-supplied values # must not be trusted. _allow_client_mock_response: Final = _key_or_team_allows_client_mock_response(user_api_key_dict) - _allow_client_message_redaction_opt_out = _key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict) + _allow_client_message_redaction_opt_out = key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict) for _internal_key in _UNTRUSTED_ROOT_CONTROL_FIELDS: if _allow_client_mock_response and _internal_key in _CLIENT_MOCK_CONTROL_FIELDS: continue @@ -2183,31 +2235,12 @@ async def add_litellm_data_to_request( # profile_id) don't see attacker-injected admin slots preserved in # the deepcopy. - # 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): - _strip_untrusted_request_header_controls( - _user_meta.get("headers"), - allow_client_message_redaction_opt_out=(_allow_client_message_redaction_opt_out), - ) - for _k in [ - k for k in _user_meta if k.startswith("user_api_key_") or k in _UNTRUSTED_METADATA_CONTROL_FIELDS - ]: - _user_meta.pop(_k, None) + strip_untrusted_caller_metadata( + data, allow_client_message_redaction_opt_out=_allow_client_message_redaction_opt_out + ) # Strip pricing overrides AFTER the litellm_metadata string-to-dict parse # above, for the same reason as the user_api_key_* strip — JSON-string @@ -2384,14 +2417,7 @@ async def add_litellm_data_to_request( data[_metadata_variable_name]["user_api_key_user_model_max_budget"] = user_model_budget # rebind-ok: out-param data[_metadata_variable_name].update(carried_budget_metadata(user_api_key_dict)) - data[_metadata_variable_name]["user_api_key_metadata"] = strip_callback_config(user_api_key_dict.metadata) - data[_metadata_variable_name]["user_api_key_team_metadata"] = strip_callback_config(user_api_key_dict.team_metadata) - data[_metadata_variable_name]["user_api_key_object_permission_id"] = getattr( - user_api_key_dict, "object_permission_id", None - ) - data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = getattr( - user_api_key_dict, "team_object_permission_id", None - ) + data[_metadata_variable_name].update(LiteLLMProxyRequestSetup.get_key_scoped_metadata(user_api_key_dict)) data[_metadata_variable_name]["headers"] = _logging_safe_headers data[_metadata_variable_name]["endpoint"] = str(request.url) # Carry the proxy-receive instant via metadata (like `endpoint`) so the diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..c1cd652ceba 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -100,6 +100,8 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + key_or_team_allows_client_message_redaction_opt_out, + strip_untrusted_caller_metadata, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path @@ -1120,6 +1122,18 @@ async def pass_through_request( _parsed_body = {} else: _parsed_body = await _read_request_body(request) + strip_untrusted_caller_metadata( + _parsed_body, + allow_client_message_redaction_opt_out=key_or_team_allows_client_message_redaction_opt_out( + user_api_key_dict + ), + ) + # Guardrails forward these to vendors as the inbound request headers; all are popped before the upstream send. + _parsed_body.pop("proxy_server_request", None) + _parsed_body.pop("headers", None) + for _caller_bucket in (_parsed_body.get("metadata"), _parsed_body.get("litellm_metadata")): + if isinstance(_caller_bucket, dict): + _caller_bucket.pop("headers", None) verbose_proxy_logger.debug( "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n", url, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index a5e79f84ef1..3f77a24b707 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -416,6 +416,22 @@ class TestMetadataExtraction: assert request_metadata["user_api_key_hash"] == "hashed-token-value" assert request_metadata["user_api_key_user_id"] == "test-user" + @pytest.mark.parametrize( + "request_data, expected_hash", + [ + pytest.param({"metadata": {"user_api_key_token": "caller-token"}}, None, id="caller-bucket-token-ignored"), + pytest.param( + {"litellm_metadata": {"user_api_key_token": "proxy-token", "user_api_key_hash": "logged-key"}}, + "logged-key", + id="hash-wins-over-token", + ), + ], + ) + def test_token_fallback_only_from_litellm_metadata_and_only_without_hash( + self, generic_guardrail, request_data, expected_hash + ): + assert generic_guardrail._extract_user_api_key_metadata(request_data).get("user_api_key_hash") == expected_hash + @pytest.mark.asyncio async def test_metadata_extraction_empty_when_no_metadata(self, generic_guardrail): """Test metadata extraction returns empty dict when no metadata available""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py index 53af7f36a5f..cce13844201 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -582,15 +582,20 @@ def test_ensure_litellm_metadata_populates_from_user_api_key_dict() -> None: assert data["litellm_metadata"]["user_api_key_team_id"] == "t1" -def test_ensure_litellm_metadata_noop_when_already_present() -> None: - """Verify _ensure_litellm_metadata does not overwrite existing litellm_metadata.""" +def test_ensure_litellm_metadata_overrides_caller_identity_in_existing_bucket() -> None: + """An existing litellm_metadata keeps its other keys, but its identity comes from the authenticated key.""" from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( _ensure_litellm_metadata, ) - user_auth = UserAPIKeyAuth(user_id="should-not-appear") - data: dict = {"litellm_metadata": {"existing": "value"}} + user_auth = UserAPIKeyAuth(user_id="auth-user", key_alias="auth-alias", team_id="auth-team") + bucket: dict = {"existing": "value", "user_api_key_alias": "batch-worker", "user_api_key_team_id": "team-exempt"} + data: dict = {"litellm_metadata": bucket} _ensure_litellm_metadata(data, user_auth) - assert data["litellm_metadata"] == {"existing": "value"} + assert data["litellm_metadata"] is bucket + assert bucket["existing"] == "value" + assert bucket["user_api_key_alias"] == "auth-alias" + assert bucket["user_api_key_team_id"] == "auth-team" + assert bucket["user_api_key_user_id"] == "auth-user" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index c90f88ec110..9b9fa0b89fb 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1,6 +1,7 @@ """Tests for unified guardrail.""" import logging +from collections.abc import Callable from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal @@ -2395,3 +2396,228 @@ class TestTranslationMappingsAreReadLive: assert not [ name for name, value in vars(unified_module).items() if isinstance(value, dict) and CallTypes.aocr in value ] + + +_RAW_CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" + + +def _sk_key(route: str) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-real-caller-key", + key_alias="prod-app", + team_id="team-prod", + metadata={"key_label": "k1"}, + team_metadata={"phoenix_project_name": "team-proj", "priority": "high"}, + request_route=route, + ) + + +def _cli_session_key(route: str) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + token=_RAW_CLI_SESSION_TOKEN, + key_alias="cli-session-alice", + user_id="alice", + is_session_token=True, + team_id="team-prod", + team_metadata={"phoenix_project_name": "team-proj", "priority": "high"}, + request_route=route, + ) + + +class TestGuardrailsSeeAuthenticatedIdentity: + """A request body cannot make a guardrail vendor see another key's identity, and the real one reaches it.""" + + @staticmethod + def _generic_guardrail(vendor_payloads: list[dict[str, object]]) -> CustomGuardrail: + import json + + import httpx + + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI + + def vendor(request: httpx.Request) -> httpx.Response: + vendor_payloads.append(json.loads(request.content)) + return httpx.Response(200, json={"action": "NONE"}) + + guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic") + guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor)) + return guardrail + + @pytest.mark.asyncio + @pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata"]) + async def test_pass_through_body_cannot_forge_identity(self, monkeypatch, bucket: str) -> None: + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + vendor_payloads: list[dict[str, object]] = [] + key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod") + data = { + "guardrail_to_apply": self._generic_guardrail(vendor_payloads), + "prompt": "hello", + bucket: { + "user_api_key_alias": "batch-worker", + "user_api_key_team_id": "team-exempt", + "user_api_key_token": "forged-hash", + }, + } + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert len(vendor_payloads) == 1 + identity = vendor_payloads[0]["request_data"] + assert identity["user_api_key_alias"] == "prod-app" + assert identity["user_api_key_team_id"] == "team-prod" + assert identity["user_api_key_hash"] == key.api_key + + @pytest.mark.asyncio + async def test_mcp_tool_call_reaches_vendor_with_key_alias(self) -> None: + from litellm.proxy.utils import ProxyLogging + + vendor_payloads: list[dict[str, object]] = [] + key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod") + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + mcp_kwargs = { + "name": "search", + "arguments": {"query": "hello"}, + "server_name": "docs", + "user_api_key_auth": key, + "user_api_key_user_id": key.user_id, + "user_api_key_team_id": key.team_id, + "user_api_key_end_user_id": None, + "user_api_key_hash": key.api_key, + "headers": {}, + } + data = proxy_logging._convert_mcp_to_llm_format( + proxy_logging._create_mcp_request_object_from_kwargs(mcp_kwargs), mcp_kwargs + ) + data["guardrail_to_apply"] = self._generic_guardrail(vendor_payloads) + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.call_mcp_tool.value + ) + + assert len(vendor_payloads) == 1 + identity = vendor_payloads[0]["request_data"] + assert identity["user_api_key_alias"] == "prod-app" + assert identity["user_api_key_team_id"] == "team-prod" + assert identity["user_api_key_hash"] == key.api_key + + @pytest.mark.asyncio + async def test_pass_through_body_cannot_forge_request_route(self, monkeypatch) -> None: + """Guardrails and call-type lookups key on user_api_key_request_route, so it must be the key's own route.""" + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + key = UserAPIKeyAuth(api_key="sk-real-caller-key", request_route="/openai/v1/chat/completions") + data = { + "guardrail_to_apply": RecordingGuardrail(), + "prompt": "hello", + "litellm_metadata": {"user_api_key_request_route": "/v1/embeddings"}, + } + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert data["litellm_metadata"]["user_api_key_request_route"] == "/openai/v1/chat/completions" + + @pytest.mark.asyncio + @pytest.mark.parametrize("route", ["/v1/messages", "/v1/responses", "/v1/chat/completions"]) + @pytest.mark.parametrize( + "make_key", [pytest.param(_sk_key, id="sk-key"), pytest.param(_cli_session_key, id="cli-session-key")] + ) + async def test_chat_path_request_keeps_proxy_metadata_and_sends_stable_hash( + self, monkeypatch, route: str, make_key: Callable[[str], UserAPIKeyAuth] + ) -> None: + """After the chat-path metadata build, the guardrail hook leaves the proxy's bucket as it was (team metadata + in user_api_key_auth_metadata included) and the vendor gets the logged key, never a raw CLI session token.""" + import copy + import json + from unittest.mock import MagicMock + + from fastapi import Request + + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request + + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + request = MagicMock(spec=Request) + request.url = MagicMock() + request.url.path = route + request.url.__str__.return_value = "http://localhost" + route + request.method = "POST" + request.query_params = {} + request.headers = {"Content-Type": "application/json"} + request.client = MagicMock() + request.client.host = "127.0.0.1" + request.state = MagicMock() + key = make_key(route) + data = await add_litellm_data_to_request( + data={"model": "m", "messages": [{"role": "user", "content": "hi"}]}, + request=request, + user_api_key_dict=key, + proxy_config=MagicMock(), + general_settings={}, + version="v", + ) + proxy_bucket = data.get("litellm_metadata") + unshared = ("litellm_parent_otel_span", "user_api_key_auth") + bucket_before = ( + copy.deepcopy({k: v for k, v in proxy_bucket.items() if k not in unshared}) if proxy_bucket else None + ) + vendor_payloads: list[dict[str, object]] = [] + data["guardrail_to_apply"] = self._generic_guardrail(vendor_payloads) + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.acompletion.value + ) + + assert len(vendor_payloads) == 1 + assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == LiteLLMProxyRequestSetup.get_logged_api_key( + key + ) + assert _RAW_CLI_SESSION_TOKEN not in json.dumps(vendor_payloads[0]) + if bucket_before is not None: + assert data["litellm_metadata"] is proxy_bucket + assert {k: v for k, v in proxy_bucket.items() if k in bucket_before} == bucket_before + assert bucket_before["user_api_key_auth_metadata"]["priority"] == "high" + assert "user_api_key_token" not in proxy_bucket + + @pytest.mark.asyncio + @pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata", None]) + async def test_pass_through_cli_session_key_sends_stable_hash(self, monkeypatch, bucket: str | None) -> None: + import json + + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + vendor_payloads: list[dict[str, object]] = [] + key = _cli_session_key("/anthropic/v1/messages") + forged_bucket = {bucket: {"user_api_key_token": "forged-hash"}} if bucket else {} + data = {"guardrail_to_apply": self._generic_guardrail(vendor_payloads), "prompt": "hello", **forged_bucket} + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert len(vendor_payloads) == 1 + assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == "cli-session-alice" + assert _RAW_CLI_SESSION_TOKEN not in json.dumps(vendor_payloads[0]) + assert vendor_payloads[0]["texts"] == ['{"prompt": "hello"}'] + + @pytest.mark.asyncio + async def test_token_only_key_drops_forged_token_already_in_proxy_bucket(self, monkeypatch) -> None: + """A key with no api_key logs no hash, so a user_api_key_token left in litellm_metadata would become the + vendor's hash if the hook kept it.""" + _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) + vendor_payloads: list[dict[str, object]] = [] + key = UserAPIKeyAuth(token="abc123hashed", key_alias="prod-app") + data = { + "guardrail_to_apply": self._generic_guardrail(vendor_payloads), + "prompt": "hello", + "litellm_metadata": {"user_api_key_token": "forged-hash"}, + } + + await UnifiedLLMGuardrails().async_pre_call_hook( + user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value + ) + + assert len(vendor_payloads) == 1 + assert "user_api_key_token" not in data["litellm_metadata"] + assert vendor_payloads[0]["request_data"].get("user_api_key_hash") is None diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 508736fb78e..8a82fbdc377 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -1487,12 +1487,12 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker): } -def _patch_apply_guardrail_env(mocker, guardrail_result): +def _patch_apply_guardrail_env(mocker, guardrail_result, processed_data=None, guardrail=None): mock_guardrail = mocker.Mock() mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result) mock_registry = mocker.Mock() - mock_registry.get_initialized_guardrail_callback.return_value = mock_guardrail + mock_registry.get_initialized_guardrail_callback.return_value = guardrail or mock_guardrail mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry) mock_logging_obj = mocker.Mock() @@ -1500,7 +1500,7 @@ def _patch_apply_guardrail_env(mocker, guardrail_result): mock_logging_obj.model_call_details = {} mock_processor = mocker.Mock() mock_processor.common_processing_pre_call_logic = AsyncMock( - return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj) + return_value=(processed_data or {"guardrail_name": "test-guardrail"}, mock_logging_obj) ) mocker.patch( "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", @@ -1535,11 +1535,12 @@ async def test_apply_guardrail_forwards_metadata_to_guardrail(mocker): user_api_key_dict=UserAPIKeyAuth(), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["What are tax loopholes?"]}, - request_data={"metadata": {"forbidden_topics": ["tax"]}}, - input_type="request", - ) + mock_guardrail.apply_guardrail.assert_awaited_once() + call = mock_guardrail.apply_guardrail.await_args.kwargs + assert call["inputs"] == {"texts": ["What are tax loopholes?"]} + assert call["input_type"] == "request" + assert "messages" not in call["request_data"] + assert call["request_data"]["metadata"]["forbidden_topics"] == ["tax"] @pytest.mark.asyncio @@ -1561,39 +1562,211 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker): user_api_key_dict=UserAPIKeyAuth(), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["What are tax loopholes?"]}, - request_data={ - "messages": messages, - "metadata": {"forbidden_topics": ["tax"]}, - }, - input_type="request", - ) + request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] + assert request_data["messages"] == messages + assert request_data["metadata"]["forbidden_topics"] == ["tax"] @pytest.mark.asyncio -async def test_apply_guardrail_omits_metadata_when_not_sent(mocker): - """Without metadata, request_data stays empty (backward-compatible).""" +async def test_apply_guardrail_authenticated_identity_overrides_client_metadata(mocker): + """A caller must not be able to claim another key's or team's identity in the body metadata.""" + mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) + caller = UserAPIKeyAuth( + api_key="sk-real-caller-key", + key_alias="real-caller", + team_id="real-team", + user_id="real-user", + ) + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={ + "user_api_key_alias": "exempt-batch-worker", + "user_api_key_team_id": "exempt-team", + "user_api_key_user_id": "someone-else", + "user_api_key_hash": "forged-hash", + "forbidden_topics": ["tax"], + }, + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["user_api_key_alias"] == "real-caller" + assert metadata["user_api_key_team_id"] == "real-team" + assert metadata["user_api_key_user_id"] == "real-user" + assert metadata["user_api_key_hash"] == caller.api_key + assert metadata["user_api_key_hash"] != "forged-hash" + assert metadata["forbidden_topics"] == ["tax"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_drops_client_identity_fields_the_key_does_not_set(mocker): + """Proxy-owned slots in the body never reach the guardrail, including user_api_key_token, which guardrails + map onto the key hash, and control fields the chat path also strips.""" + mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) + caller = UserAPIKeyAuth(metadata={"zguard_policy_id": "strict"}, object_permission_id="perm-real") + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={ + "user_api_key_alias": "exempt-batch-worker", + "user_api_key_team_id": "exempt-team", + "user_api_key_token": "forged-hash", + "user_api_key_metadata": {"zguard_policy_id": "permissive"}, + "user_api_key_object_permission_id": "perm-forged", + "user_api_key": "forged-key", + "applied_guardrails": ["already-ran"], + "headers": {"x-end-user": "someone-else"}, + "trace_label": "nightly", + }, + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["user_api_key_alias"] is None + assert metadata["user_api_key_team_id"] is None + assert "user_api_key_token" not in metadata + assert metadata["user_api_key_metadata"] == {"zguard_policy_id": "strict"} + assert metadata["user_api_key_object_permission_id"] == "perm-real" + assert metadata["user_api_key"] is None + assert "applied_guardrails" not in metadata + assert "headers" not in metadata + assert metadata["trace_label"] == "nightly" + + +@pytest.mark.asyncio +async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied_ones(mocker): + """Guardrails forward metadata headers to vendors, so they must be the proxy's view of the request.""" + real_headers = {"user-agent": "real-client/1.0"} + mock_guardrail = _patch_apply_guardrail_env( + mocker, + {"texts": ["ok"]}, + processed_data={"guardrail_name": "test-guardrail", "metadata": {"headers": real_headers}}, + ) + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={"headers": {"user-agent": "forged/1.0", "x-end-user": "someone-else"}}, + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=UserAPIKeyAuth()) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["headers"] == real_headers + + +@pytest.mark.asyncio +async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker): + """End to end through a real GenericGuardrailAPI: the vendor payload names the authenticated key even when + the body forges user_api_key_alias and user_api_key_token, which the generic guardrail maps onto the hash.""" + import httpx + + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI + + vendor_payloads = [] + + def vendor(request: httpx.Request) -> httpx.Response: + vendor_payloads.append(json.loads(request.content)) + return httpx.Response(200, json={"action": "NONE"}) + + generic_guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic") + generic_guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor)) + _patch_apply_guardrail_env(mocker, {"texts": ["unused"]}, guardrail=generic_guardrail) + caller = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="real-caller", team_id="real-team") + + request = ApplyGuardrailRequest( + guardrail_name="generic", + text="hello", + metadata={"user_api_key_alias": "exempt-batch-worker", "user_api_key_token": "forged-hash"}, + ) + response = await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + assert response.response_text == "hello" + assert len(vendor_payloads) == 1 + identity = vendor_payloads[0]["request_data"] + assert identity["user_api_key_alias"] == "real-caller" + assert identity["user_api_key_team_id"] == "real-team" + assert identity["user_api_key_hash"] == caller.api_key + + +@pytest.mark.asyncio +async def test_apply_guardrail_request_route_comes_from_the_key(mocker): + """Guardrails pick call-type behavior from user_api_key_request_route, so the body cannot choose it.""" + mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) + + request = ApplyGuardrailRequest( + guardrail_name="test-guardrail", + text="hello", + metadata={"user_api_key_request_route": "/v1/embeddings"}, + ) + await apply_guardrail( + fastapi_request=mocker.Mock(), + request=request, + user_api_key_dict=UserAPIKeyAuth(request_route="/guardrails/apply_guardrail"), + ) + + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert metadata["user_api_key_request_route"] == "/guardrails/apply_guardrail" + + +@pytest.mark.asyncio +async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocker): + """A CLI session key's raw per-login token must never reach the vendor; it gets the stable logged key.""" + import httpx + + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI + + raw_session_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" + vendor_payloads = [] + + def vendor(request: httpx.Request) -> httpx.Response: + vendor_payloads.append(json.loads(request.content)) + return httpx.Response(200, json={"action": "NONE"}) + + generic_guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic") + generic_guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor)) + _patch_apply_guardrail_env(mocker, {"texts": ["unused"]}, guardrail=generic_guardrail) + caller = UserAPIKeyAuth( + token=raw_session_token, key_alias="cli-session-alice", user_id="alice", is_session_token=True + ) + + request = ApplyGuardrailRequest( + guardrail_name="generic", text="hello", metadata={"user_api_key_token": raw_session_token} + ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) + + assert len(vendor_payloads) == 1 + assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == "cli-session-alice" + assert raw_session_token not in json.dumps(vendor_payloads[0]) + + +@pytest.mark.asyncio +async def test_apply_guardrail_carries_authenticated_identity_when_no_metadata_sent(mocker): + """request_data always carries the authenticated identity, even when the body has no metadata.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello") await apply_guardrail( fastapi_request=mocker.Mock(), request=request, - user_api_key_dict=UserAPIKeyAuth(), + user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller", team_id="known-team"), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["hello"]}, - request_data={}, - input_type="request", - ) + call = mock_guardrail.apply_guardrail.await_args.kwargs + assert call["inputs"] == {"texts": ["hello"]} + assert "messages" not in call["request_data"] + assert call["request_data"]["metadata"]["user_api_key_alias"] == "known-caller" + assert call["request_data"]["metadata"]["user_api_key_team_id"] == "known-team" @pytest.mark.asyncio async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(mocker): - """Explicitly-sent empty messages/metadata must be forwarded, not dropped; - only omitted fields stay out of request_data.""" + """Explicitly-sent empty messages must be forwarded, not dropped, and empty + metadata still carries the authenticated identity.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) request = ApplyGuardrailRequest( @@ -1605,14 +1778,12 @@ async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(moc await apply_guardrail( fastapi_request=mocker.Mock(), request=request, - user_api_key_dict=UserAPIKeyAuth(), + user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller"), ) - mock_guardrail.apply_guardrail.assert_awaited_once_with( - inputs={"texts": ["hello"]}, - request_data={"messages": [], "metadata": {}}, - input_type="request", - ) + request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] + assert request_data["messages"] == [] + assert request_data["metadata"]["user_api_key_alias"] == "known-caller" @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..652acfbe5b9 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -7622,3 +7622,81 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token() metadata = kwargs["litellm_params"]["metadata"] assert metadata["user_api_key"] == "cli-session-alice" assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice" + + +@pytest.mark.asyncio +async def test_pass_through_request_strips_caller_identity_before_guardrail_hooks(): + """ + Regression: a pass-through body skips add_litellm_data_to_request, so forged user_api_key_* fields, guardrail + control fields and inbound headers reached pre_call_hook guardrails as the caller's identity. The upstream body + is unchanged because these keys never reach it. + """ + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.types.llms.custom_http import httpxSpecialProvider + + upstream_bodies = [] + + def transport_handler(upstream_request: httpx.Request) -> httpx.Response: + upstream_bodies.append(json.loads(upstream_request.content)) + return httpx.Response(200, json={"ok": True}) + + real_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.PassThroughEndpoint, + params={"timeout": resolve_pass_through_request_timeout(None)}, + ) + cache_dict = litellm.in_memory_llm_clients_cache.cache_dict + cache_key = next(key for key, cached in cache_dict.items() if cached is real_handler) + cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler))) + + hook_data = [] + + def record_hook_data(user_api_key_dict, data, call_type): + hook_data.append({key: value for key, value in data.items() if key != "litellm_logging_obj"}) + return data + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=record_hook_data) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + forged = { + "user_api_key_alias": "batch-worker", + "user_api_key_team_id": "team-exempt", + "user_api_key_token": "forged-hash", + "user_api_key_request_route": "/v1/embeddings", + "disable_global_guardrails": True, + "headers": {"x-authenticated-user": "admin@corp"}, + "trace_label": "nightly", + } + forged_headers = {"x-authenticated-user": "admin@corp", "x-litellm-end-user-id": "victim"} + body = { + "prompt": "hello", + "metadata": forged, + "litellm_metadata": forged, + "headers": forged_headers, + "proxy_server_request": {"headers": forged_headers}, + } + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.query_params = QueryParams({}) + mock_request.body = AsyncMock(return_value=json.dumps(body).encode()) + + try: + with patch( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging + ): # test-quality-ok: read at call time + response = await pass_through_request( + request=mock_request, + target="https://upstream.test/v1/generate", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app"), + ) + finally: + cache_dict[cache_key] = real_handler + + assert response.status_code == 200 + assert hook_data == [ + {"prompt": "hello", "metadata": {"trace_label": "nightly"}, "litellm_metadata": {"trace_label": "nightly"}} + ] + assert upstream_bodies == [{"prompt": "hello"}] diff --git a/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py b/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py index 4f44a4adeed..bab9958dc5d 100644 --- a/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py +++ b/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py @@ -61,11 +61,11 @@ async def test_apply_guardrail_endpoint_returns_correct_response( assert response.response_text == "Redacted text: [REDACTED] and [REDACTED]" # Verify the guardrail was called with correct parameters - mock_guardrail.apply_guardrail.assert_called_once_with( - inputs={"texts": ["Test text with PII"]}, - request_data={}, - input_type="request", - ) + mock_guardrail.apply_guardrail.assert_called_once() + call = mock_guardrail.apply_guardrail.call_args.kwargs + assert call["inputs"] == {"texts": ["Test text with PII"]} + assert call["input_type"] == "request" + assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key @pytest.mark.asyncio @@ -197,6 +197,8 @@ async def test_apply_guardrail_endpoint_without_optional_params(mock_proxy_loggi assert response.response_text == "Processed text" # Verify the guardrail was called with correct parameters - mock_guardrail.apply_guardrail.assert_called_once_with( - inputs={"texts": ["Test text"]}, request_data={}, input_type="request" - ) + mock_guardrail.apply_guardrail.assert_called_once() + call = mock_guardrail.apply_guardrail.call_args.kwargs + assert call["inputs"] == {"texts": ["Test text"]} + assert call["input_type"] == "request" + assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 7e6d4d24905..ed9c666ac80 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -3550,3 +3550,134 @@ async def test_provider_bytes_are_sent_raw_after_pacing(): assert [call.args[0] for call in backend_ws.send.await_args_list] == [b"\x00\x01", '{"type":"endStream"}'] provider_config.pace_backend_send.assert_awaited_once_with(b"\x00\x01") + + +@pytest.mark.asyncio +async def test_realtime_transcript_guardrail_receives_authenticated_identity(monkeypatch: pytest.MonkeyPatch): + """Transcript guardrails get the session key's identity in litellm_metadata, as the chat path provides it.""" + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.guardrails import GuardrailEventHooks + + received_request_data = [] + + class IdentityRecordingGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + received_request_data.append(request_data) + return inputs + + monkeypatch.setattr( + litellm, + "callbacks", + [ + IdentityRecordingGuardrail( + guardrail_name="identity_recorder", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + ], + ) + key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod") + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=key) + + blocked = await streaming.run_realtime_guardrails("hello there") + + assert blocked is False + assert len(received_request_data) == 1 + identity = received_request_data[0]["litellm_metadata"] + assert identity["user_api_key_alias"] == "prod-app" + assert identity["user_api_key_team_id"] == "team-prod" + assert identity["user_api_key_hash"] == key.api_key + + +@pytest.mark.asyncio +async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pytest.MonkeyPatch): + """Gray Swan forwards litellm_metadata verbatim, so realtime must hand it identity and no key secrets.""" + import litellm + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + vendor_payloads: list[dict[str, object]] = [] + + class RecordingGraySwan(GraySwanGuardrail): + async def _call_grayswan_api(self, payload): + vendor_payloads.append(payload) + return {"violation": 0.0, "violated_rules": []} + + monkeypatch.setattr( + litellm, + "callbacks", + [ + RecordingGraySwan( + guardrail_name="grayswan", + api_key="test-key", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + ], + ) + key = UserAPIKeyAuth( + api_key="sk-real-caller-key", + key_alias="prod-app", + team_id="team-prod", + organization_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-ORG"}}]}, + jwt_claims={"email": "alice@corp.example", "name": "Alice Smith"}, + team_member={"user_id": "alice", "user_email": "alice@corp.example", "role": "admin"}, + ) + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=key) + + await streaming.run_realtime_guardrails("hello there", event_hooks=[GuardrailEventHooks.pre_call]) + + assert len(vendor_payloads) == 1 + vendor_metadata = vendor_payloads[0]["litellm_metadata"] + assert vendor_metadata["user_api_key_alias"] == "prod-app" + assert vendor_metadata["user_api_key_team_id"] == "team-prod" + serialized = json.dumps(vendor_metadata) + for leaked in ("SECRET-ORG", "Alice Smith", "alice@corp.example"): + assert leaked not in serialized + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "sdk_value", + [ + pytest.param({"key_alias": "forged", "team_id": "team-exempt"}, id="dict"), + pytest.param({"spend": "not-a-number"}, id="malformed-dict"), + pytest.param("sk-raw-string", id="string"), + ], +) +async def test_realtime_guardrail_gets_no_identity_from_non_auth_sdk_value( + monkeypatch: pytest.MonkeyPatch, sdk_value: object +): + """Only a proxy-authenticated UserAPIKeyAuth yields identity; an SDK-supplied value never raises or fakes one.""" + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + received_request_data: list[dict[str, object]] = [] + + class IdentityRecordingGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + received_request_data.append(request_data) + return inputs + + monkeypatch.setattr( + litellm, + "callbacks", + [ + IdentityRecordingGuardrail( + guardrail_name="identity_recorder", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + ], + ) + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=sdk_value) + + blocked = await streaming.run_realtime_guardrails("hello there") + + assert blocked is False + assert len(received_request_data) == 1 + assert not [key for key in received_request_data[0]["litellm_metadata"] if key.startswith("user_api_key")] diff --git a/tests/unit/llms/base_llm/guardrail_translation/__init__.py b/tests/unit/llms/base_llm/guardrail_translation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/guardrail_translation/test_base_translation.py b/tests/unit/llms/base_llm/guardrail_translation/test_base_translation.py new file mode 100644 index 00000000000..06745c5a15f --- /dev/null +++ b/tests/unit/llms/base_llm/guardrail_translation/test_base_translation.py @@ -0,0 +1,53 @@ +import json +from typing import Final + +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.proxy._types import UserAPIKeyAuth + +RAW_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" + + +def _fully_populated_session_key() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + token=RAW_SESSION_TOKEN, + is_session_token=True, + key_alias="cli-session-alice", + user_id="alice", + team_id="team-prod", + org_id="org-1", + metadata={"logging": [{"callback_name": "langfuse", "callback_vars": {"langfuse_secret_key": "SECRET-KEY"}}]}, + team_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-TEAM"}}]}, + organization_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-ORG"}}]}, + project_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-PROJECT"}}]}, + jwt_claims={"sub": "alice", "email": "alice@corp.example", "name": "Alice Smith"}, + team_member={"user_id": "alice", "user_email": "alice@corp.example", "role": "admin"}, + config={"internal": "proxy-config"}, + ) + + +def test_transform_emits_only_identity_never_credentials_or_callback_secrets(): + metadata = BaseTranslation.transform_user_api_key_dict_to_metadata(_fully_populated_session_key()) + serialized = json.dumps(metadata, default=str) + + assert metadata["user_api_key_alias"] == "cli-session-alice" + assert metadata["user_api_key_key_alias"] == "cli-session-alice" + assert metadata["user_api_key_hash"] == "cli-session-alice" + assert metadata["user_api_key_team_id"] == "team-prod" + assert RAW_SESSION_TOKEN not in serialized + assert "callback_vars" not in serialized + assert "SECRET" not in serialized + assert "Alice Smith" not in serialized + assert "proxy-config" not in serialized + for dropped in ( + "user_api_key_token", + "user_api_key_jwt_claims", + "user_api_key_team_member", + "user_api_key_organization_metadata", + "user_api_key_project_metadata", + "user_api_key_config", + ): + assert dropped not in metadata + + +def test_transform_of_no_key_is_empty(): + assert BaseTranslation.transform_user_api_key_dict_to_metadata(None) == {} From 88e0c43e7c4201b92d75c29c42aecdcea46e1d0c Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Tue, 29 Sep 2026 23:59:33 +0300 Subject: [PATCH 2/4] refactor(guardrails): hoist function-local imports to module top Moves the imports this PR added inside functions to the top of their modules: UserAPIKeyAuth and BaseTranslation in realtime_streaming, BaseTranslation, StreamingScanKey and LiteLLMProxyRequestSetup in unified_guardrail (so its annotations no longer need quotes), and the test-only imports in the unified guardrail, guardrail endpoint, pass-through and realtime tests. None of them creates a cycle, and import litellm still does not load the proxy request layer or fastapi One import stays inside its function: base_translation's LiteLLMProxyRequestSetup. import litellm loads base_translation before litellm.Router exists, and litellm_pre_call_utils imports Router, so hoisting it makes import litellm fail with "cannot import name 'Router' from 'litellm'". It carries a one-line comment saying so --- .../litellm_core_utils/realtime_streaming.py | 8 ++---- .../guardrail_translation/base_translation.py | 1 + .../unified_guardrail/unified_guardrail.py | 21 +++++--------- .../test_unified_guardrail.py | 28 ++++++------------- .../guardrails/test_guardrail_endpoints.py | 13 ++------- .../test_pass_through_endpoints.py | 5 ++-- .../test_realtime_streaming.py | 16 ++--------- 7 files changed, 27 insertions(+), 65 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 300834073f7..d1b962d1226 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -12,7 +12,9 @@ import litellm from litellm._logging import redact_internal_details_from_client_message, verbose_logger from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend +from litellm.proxy._types import UserAPIKeyAuth from litellm.types.llms.openai import ( OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, @@ -29,7 +31,6 @@ if TYPE_CHECKING: from websockets.asyncio.client import ClientConnection from websockets.exceptions import ConnectionClosed - from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks CLIENT_CONNECTION_CLASS = ClientConnection @@ -124,9 +125,7 @@ DefaultLoggedRealTimeEventTypes: Final = [ ] -def _as_user_api_key_auth(user_api_key_dict: object) -> "UserAPIKeyAuth | None": - from litellm.proxy._types import UserAPIKeyAuth - +def _as_user_api_key_auth(user_api_key_dict: object) -> UserAPIKeyAuth | None: return user_api_key_dict if isinstance(user_api_key_dict, UserAPIKeyAuth) else None @@ -838,7 +837,6 @@ class RealTimeStreaming: typed user messages and tool outputs use ``pre_call``. """ from litellm.integrations.custom_guardrail import CustomGuardrail - from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.types.guardrails import GuardrailEventHooks if event_hooks is None: diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 1a772ca8ee3..074d1b275ef 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -86,6 +86,7 @@ class BaseTranslation(ABC): """The authenticated key's identity as prefixed metadata, an allowlist safe to hand to guardrail vendors.""" if user_api_key_dict is None: return {} + # Lazy: `import litellm` loads this module before litellm.Router exists, and litellm_pre_call_utils imports it from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup return { diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index eeb604e98af..0743af68e18 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -20,7 +20,9 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings +from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation, StreamingScanKey from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import ( CallTypes, @@ -34,10 +36,6 @@ if TYPE_CHECKING: # Imported lazily at runtime (inside the streaming hook) to avoid a # module-level cyclic import with litellm.integrations.custom_guardrail. from litellm.integrations.custom_guardrail import ModifyResponseException - from litellm.llms.base_llm.guardrail_translation.base_translation import ( - BaseTranslation, - StreamingScanKey, - ) # Call types that stream JSON-RPC events (A2A); guardrail HTTPException is emitted as in-stream error A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message) @@ -56,7 +54,7 @@ class _EndpointTranslation(Protocol): def process_output_streaming_response(self) -> "Callable[..., Awaitable[object]]": ... @property - def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ... + def get_streaming_scan_key(self) -> Callable[[Sequence[object]], StreamingScanKey | None]: ... @property def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ... @@ -71,7 +69,7 @@ def _as_endpoint_translation(translation: _EndpointTranslation) -> _EndpointTran def resolve_endpoint_translation( user_api_key_dict: UserAPIKeyAuth, first_response_item: object | None -) -> "tuple[str, BaseTranslation] | None": +) -> tuple[str, BaseTranslation] | None: """ Resolve the endpoint guardrail translation for a streamed response: the request route wins, falling back to inferring the call type from the first @@ -108,7 +106,7 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]: return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0) -def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool: +def _is_redundant_scan(scan_key: StreamingScanKey | None, last_scan_key: StreamingScanKey | None) -> bool: if scan_key is None: return False return scan_key == last_scan_key or scan_key.has_nothing_to_scan @@ -160,11 +158,6 @@ _PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_metadata" def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None: """Overwrite the identity fields of data['litellm_metadata'] from the authenticated key, in place.""" - from litellm.llms.base_llm.guardrail_translation.base_translation import ( - BaseTranslation, - ) - from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup - existing: Final = data.get("litellm_metadata") if isinstance(existing, dict): identity: Final = LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict) @@ -414,7 +407,7 @@ class UnifiedLLMGuardrails(CustomLogger): @staticmethod def _resolve_transform_call_type( user_api_key_dict: UserAPIKeyAuth, - mappings: Mapping[CallTypes, type["BaseTranslation"]], + mappings: Mapping[CallTypes, type[BaseTranslation]], ) -> str | None: """Resolve the call type for the incremental_diff path, or None if the route is unresolvable / unsupported. @@ -677,7 +670,7 @@ class UnifiedLLMGuardrails(CustomLogger): call_type: str, sampling_rate: int, end_of_stream_only: bool, - mappings: Mapping[CallTypes, type["BaseTranslation"]], + mappings: Mapping[CallTypes, type[BaseTranslation]], ) -> AsyncGenerator[object, None]: """Emit guardrail text transformations as new deltas on the stream. diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index 9b9fa0b89fb..e945145f01a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1,11 +1,16 @@ """Tests for unified guardrail.""" +import copy +import json import logging from collections.abc import Callable from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal +from unittest.mock import MagicMock +import httpx import pytest +from fastapi import Request import litellm from litellm.caching import DualCache @@ -23,6 +28,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( openai_messages_without_tool, ) from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.llms.mistral.ocr.guardrail_translation.handler import OCRHandler from litellm.llms.openai.chat.guardrail_translation.handler import ( OpenAIChatCompletionsHandler, @@ -34,12 +40,15 @@ from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import MCPGuardrailTranslationHandler, ) from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import ( unified_guardrail as unified_module, ) from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request +from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices @@ -2429,13 +2438,6 @@ class TestGuardrailsSeeAuthenticatedIdentity: @staticmethod def _generic_guardrail(vendor_payloads: list[dict[str, object]]) -> CustomGuardrail: - import json - - import httpx - - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI - def vendor(request: httpx.Request) -> httpx.Response: vendor_payloads.append(json.loads(request.content)) return httpx.Response(200, json={"action": "NONE"}) @@ -2472,8 +2474,6 @@ class TestGuardrailsSeeAuthenticatedIdentity: @pytest.mark.asyncio async def test_mcp_tool_call_reaches_vendor_with_key_alias(self) -> None: - from litellm.proxy.utils import ProxyLogging - vendor_payloads: list[dict[str, object]] = [] key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod") proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) @@ -2530,14 +2530,6 @@ class TestGuardrailsSeeAuthenticatedIdentity: ) -> None: """After the chat-path metadata build, the guardrail hook leaves the proxy's bucket as it was (team metadata in user_api_key_auth_metadata included) and the vendor gets the logged key, never a raw CLI session token.""" - import copy - import json - from unittest.mock import MagicMock - - from fastapi import Request - - from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request - _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) request = MagicMock(spec=Request) request.url = MagicMock() @@ -2584,8 +2576,6 @@ class TestGuardrailsSeeAuthenticatedIdentity: @pytest.mark.asyncio @pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata", None]) async def test_pass_through_cli_session_key_sends_stable_hash(self, monkeypatch, bucket: str | None) -> None: - import json - _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) vendor_payloads: list[dict[str, object]] = [] key = _cli_session_key("/anthropic/v1/messages") diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 8a82fbdc377..0d6b9b9d13f 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -4,11 +4,13 @@ from datetime import datetime from typing import Dict, List, Optional from unittest.mock import AsyncMock +import httpx import pytest from fastapi import HTTPException +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_endpoints import ( CreateGuardrailRequest, @@ -33,6 +35,7 @@ from litellm.proxy.guardrails.guardrail_endpoints import ( from litellm.proxy.guardrails.guardrail_endpoints import ( test_custom_code_guardrail as run_custom_code_test_endpoint, ) +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) from litellm.proxy.guardrails.guardrail_registry import ( @@ -1661,11 +1664,6 @@ async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker): """End to end through a real GenericGuardrailAPI: the vendor payload names the authenticated key even when the body forges user_api_key_alias and user_api_key_token, which the generic guardrail maps onto the hash.""" - import httpx - - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI - vendor_payloads = [] def vendor(request: httpx.Request) -> httpx.Response: @@ -1715,11 +1713,6 @@ async def test_apply_guardrail_request_route_comes_from_the_key(mocker): @pytest.mark.asyncio async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocker): """A CLI session key's raw per-login token must never reach the vendor; it gets the stable logged key.""" - import httpx - - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI - raw_session_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" vendor_payloads = [] diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 652acfbe5b9..fe5dc05b477 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -26,6 +26,7 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, @@ -48,6 +49,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.types import utils as types_utils +from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -7631,9 +7633,6 @@ async def test_pass_through_request_strips_caller_identity_before_guardrail_hook control fields and inbound headers reached pre_call_hook guardrails as the caller's identity. The upstream body is unchanged because these keys never reach it. """ - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client - from litellm.types.llms.custom_http import httpxSpecialProvider - upstream_bodies = [] def transport_handler(upstream_request: httpx.Request) -> httpx.Response: diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index ed9c666ac80..a5713b355af 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -17,6 +17,8 @@ from litellm.litellm_core_utils.realtime_streaming import ( client_sent_openai_beta_realtime_header, ) from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail from litellm.types.guardrails import GuardrailEventHooks @@ -3555,11 +3557,6 @@ async def test_provider_bytes_are_sent_raw_after_pacing(): @pytest.mark.asyncio async def test_realtime_transcript_guardrail_receives_authenticated_identity(monkeypatch: pytest.MonkeyPatch): """Transcript guardrails get the session key's identity in litellm_metadata, as the chat path provides it.""" - import litellm - from litellm.integrations.custom_guardrail import CustomGuardrail - from litellm.proxy._types import UserAPIKeyAuth - from litellm.types.guardrails import GuardrailEventHooks - received_request_data = [] class IdentityRecordingGuardrail(CustomGuardrail): @@ -3594,11 +3591,6 @@ async def test_realtime_transcript_guardrail_receives_authenticated_identity(mon @pytest.mark.asyncio async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pytest.MonkeyPatch): """Gray Swan forwards litellm_metadata verbatim, so realtime must hand it identity and no key secrets.""" - import litellm - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail - from litellm.types.guardrails import GuardrailEventHooks - vendor_payloads: list[dict[str, object]] = [] class RecordingGraySwan(GraySwanGuardrail): @@ -3652,10 +3644,6 @@ async def test_realtime_guardrail_gets_no_identity_from_non_auth_sdk_value( monkeypatch: pytest.MonkeyPatch, sdk_value: object ): """Only a proxy-authenticated UserAPIKeyAuth yields identity; an SDK-supplied value never raises or fakes one.""" - import litellm - from litellm.integrations.custom_guardrail import CustomGuardrail - from litellm.types.guardrails import GuardrailEventHooks - received_request_data: list[dict[str, object]] = [] class IdentityRecordingGuardrail(CustomGuardrail): From 8fe7b4f00f9f6a2545265a8c0dafeaccfcb7cbb1 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Wed, 30 Sep 2026 12:55:06 +0300 Subject: [PATCH 3/4] fix(guardrails): give pass-through guardrails the real inbound headers Pass-through guardrails lost their inbound headers once the caller's body copies were stripped, so an operator's extra_headers allowlist forwarded nothing to the vendor. Both the pre-call and the post-call guardrail hooks now get the real request headers in proxy_server_request, built once the way the chat path builds them: clean_headers drops the proxy's auth header and a custom litellm_key_header_name, MCP upstream credential headers are dropped, and redact_credential_headers masks cookies and other credentials On the success path the headers are attached after the logging object is created, so they do not land in the logged request. When a pre-call guardrail blocks, the failure log now receives these cleaned and redacted headers, which matches what the chat path logs. The pass-through guardrail leaves proxy_server_request out of the text it scans, and the upstream body is unchanged because proxy_server_request is popped with the other litellm params before the send --- .../guardrail_translation/handler.py | 4 +- .../pass_through_endpoints.py | 32 ++++-- .../test_unified_guardrail.py | 8 +- .../test_pass_through_endpoints.py | 104 +++++++++++++++++- 4 files changed, 133 insertions(+), 15 deletions(-) diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index ac7dcfd4f39..6d8f13e1044 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -20,7 +20,9 @@ if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging -_PROXY_OWNED_PAYLOAD_KEYS: Final = frozenset({"metadata", "litellm_metadata", "litellm_logging_obj"}) +_PROXY_OWNED_PAYLOAD_KEYS: Final = frozenset( + {"metadata", "litellm_metadata", "litellm_logging_obj", "proxy_server_request"} +) class PassThroughEndpointHandler(BaseTranslation): diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index c1cd652ceba..7bd1bb7871f 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -26,6 +26,7 @@ from fastapi import ( status, ) from fastapi.responses import StreamingResponse +from starlette.datastructures import Headers from starlette.datastructures import UploadFile as StarletteUploadFile from starlette.websockets import WebSocketState from websockets.asyncio.client import connect @@ -65,6 +66,7 @@ from litellm.llms.base_llm.managed_resources.utils import ( ) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.passthrough import BasePassthroughUtils +from litellm.proxy._experimental.mcp_server.utils import upstream_credential_headers from litellm.proxy._types import ( ConfigFieldInfo, ConfigFieldUpdate, @@ -100,7 +102,9 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + clean_headers, key_or_team_allows_client_message_redaction_opt_out, + redact_credential_headers, strip_untrusted_caller_metadata, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError @@ -1024,6 +1028,14 @@ from litellm.passthrough.timeout_utils import ( ) +def _guardrail_request_headers(headers: Headers, litellm_key_header_name: str | None) -> Mapping[str, str]: + cleaned: Final = clean_headers(headers, litellm_key_header_name=litellm_key_header_name) + mcp_credential_headers: Final = upstream_credential_headers(cleaned) + return redact_credential_headers( + {name: value for name, value in cleaned.items() if name.lower() not in mcp_credential_headers} + ) + + async def pass_through_request( request: Request, target: str, @@ -1128,6 +1140,14 @@ async def pass_through_request( user_api_key_dict ), ) + # Lazy: proxy_server imports this module + from litellm.proxy.proxy_server import ( + general_settings as proxy_general_settings, + ) + from litellm.proxy.proxy_server import ( + general_settings_view, + ) + # Guardrails forward these to vendors as the inbound request headers; all are popped before the upstream send. _parsed_body.pop("proxy_server_request", None) _parsed_body.pop("headers", None) @@ -1188,6 +1208,10 @@ async def pass_through_request( if _parsed_body is None: _parsed_body = {} _parsed_body["litellm_logging_obj"] = logging_obj + guardrail_headers: Final = _guardrail_request_headers( + request.headers, litellm_key_header_name=proxy_general_settings.get("litellm_key_header_name") + ) + _parsed_body["proxy_server_request"] = {"headers": guardrail_headers} ### CALL HOOKS ### - modify incoming data / reject request before calling the model _parsed_body = await proxy_logging_obj.pre_call_hook( @@ -1236,13 +1260,6 @@ async def pass_through_request( # provider IDs before forwarding upstream. Gated by feature flag and # enterprise managed-files hook. Runs after pre_call_hook so # guardrails have already seen the managed IDs. - from litellm.proxy.proxy_server import ( - general_settings as proxy_general_settings, - ) - from litellm.proxy.proxy_server import ( - general_settings_view, - ) - _managed_id_provider: Final = resolve_passthrough_managed_id_provider(custom_llm_provider) if proxy_general_settings.get("passthrough_managed_object_ids", False) and _managed_id_provider is not None: @@ -1650,6 +1667,7 @@ async def pass_through_request( **existing_metadata, "guardrails": guardrails_to_run, } + hook_data["proxy_server_request"] = {"headers": guardrail_headers} post_call_guardrail_data = hook_data response_body = await proxy_logging_obj.post_call_success_hook( data=hook_data, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e945145f01a..afebcbdf8a6 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -2580,7 +2580,12 @@ class TestGuardrailsSeeAuthenticatedIdentity: vendor_payloads: list[dict[str, object]] = [] key = _cli_session_key("/anthropic/v1/messages") forged_bucket = {bucket: {"user_api_key_token": "forged-hash"}} if bucket else {} - data = {"guardrail_to_apply": self._generic_guardrail(vendor_payloads), "prompt": "hello", **forged_bucket} + data = { + "guardrail_to_apply": self._generic_guardrail(vendor_payloads), + "prompt": "hello", + "proxy_server_request": {"headers": {"x-tenant": "tenant-real"}}, + **forged_bucket, + } await UnifiedLLMGuardrails().async_pre_call_hook( user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value @@ -2590,6 +2595,7 @@ class TestGuardrailsSeeAuthenticatedIdentity: assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == "cli-session-alice" assert _RAW_CLI_SESSION_TOKEN not in json.dumps(vendor_payloads[0]) assert vendor_payloads[0]["texts"] == ['{"prompt": "hello"}'] + assert vendor_payloads[0]["request_headers"] == {"x-tenant": "[present]"} @pytest.mark.asyncio async def test_token_only_key_drops_forged_token_already_in_proxy_bucket(self, monkeypatch) -> None: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index fe5dc05b477..1b397ea3286 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -28,6 +28,10 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import ( + _extract_inbound_headers, +) +from litellm.proxy.litellm_pre_call_utils import _REDACTED_HEADER_VALUE from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, @@ -7667,7 +7671,7 @@ async def test_pass_through_request_strips_caller_identity_before_guardrail_hook "headers": {"x-authenticated-user": "admin@corp"}, "trace_label": "nightly", } - forged_headers = {"x-authenticated-user": "admin@corp", "x-litellm-end-user-id": "victim"} + forged_headers = {"x-authenticated-user": "admin@corp", "x-litellm-end-user-id": "victim", "x-tenant": "forged"} body = { "prompt": "hello", "metadata": forged, @@ -7677,14 +7681,26 @@ async def test_pass_through_request_strips_caller_identity_before_guardrail_hook } mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.headers = Headers( + { + "content-type": "application/json", + "x-tenant": "tenant-real", + "authorization": "Bearer sk-real-caller-key", + "x-my-key": "sk-custom-header-key", + "x-mcp-github-authorization": "Bearer mcp-upstream-token", + "cookie": "session=secret", + } + ) mock_request.query_params = QueryParams({}) mock_request.body = AsyncMock(return_value=json.dumps(body).encode()) try: - with patch( - "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging - ): # test-quality-ok: read at call time + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging + ), # test-quality-ok: read at call time + patch("litellm.proxy.proxy_server.general_settings", {"litellm_key_header_name": "x-my-key"}), + ): response = await pass_through_request( request=mock_request, target="https://upstream.test/v1/generate", @@ -7696,6 +7712,82 @@ async def test_pass_through_request_strips_caller_identity_before_guardrail_hook assert response.status_code == 200 assert hook_data == [ - {"prompt": "hello", "metadata": {"trace_label": "nightly"}, "litellm_metadata": {"trace_label": "nightly"}} + { + "prompt": "hello", + "metadata": {"trace_label": "nightly"}, + "litellm_metadata": {"trace_label": "nightly"}, + "proxy_server_request": { + "headers": { + "content-type": "application/json", + "x-tenant": "tenant-real", + "cookie": _REDACTED_HEADER_VALUE, + } + }, + } ] + vendor_headers = _extract_inbound_headers(request_data=hook_data[0], logging_obj=None, extra_allowlist={"x-tenant"}) + assert vendor_headers is not None and vendor_headers["x-tenant"] == "tenant-real" assert upstream_bodies == [{"prompt": "hello"}] + + +@pytest.mark.asyncio +async def test_pass_through_post_call_guardrails_receive_real_inbound_headers(): + """Post-call guardrails run on a copy of the body the litellm-param pop already stripped, so without an explicit + re-attach an operator's extra_headers allowlist forwarded nothing on the response side.""" + + def transport_handler(upstream_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"completion": "hi"}) + + real_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.PassThroughEndpoint, + params={"timeout": resolve_pass_through_request_timeout(None)}, + ) + cache_dict = litellm.in_memory_llm_clients_cache.cache_dict + cache_key = next(key for key, cached in cache_dict.items() if cached is real_handler) + cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler))) + + post_call_data = [] + + def record_post_call(data, user_api_key_dict, response): + post_call_data.append(data) + return response + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data) + mock_proxy_logging.post_call_success_hook = AsyncMock(side_effect=record_post_call) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.headers = Headers( + {"content-type": "application/json", "x-tenant": "tenant-real", "authorization": "Bearer sk-real-caller-key"} + ) + mock_request.query_params = QueryParams({}) + mock_request.body = AsyncMock( + return_value=json.dumps({"prompt": "hello", "headers": {"x-tenant": "forged"}}).encode() + ) + + try: + with patch( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging + ): # test-quality-ok: read at call time + response = await pass_through_request( + request=mock_request, + target="https://upstream.test/v1/generate", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app"), + guardrails_config=["gg"], + ) + finally: + cache_dict[cache_key] = real_handler + + assert response.status_code == 200 + assert len(post_call_data) == 1, "the post-call guardrail hook did not run" + assert post_call_data[0]["proxy_server_request"] == { + "headers": {"content-type": "application/json", "x-tenant": "tenant-real"} + } + vendor_headers = _extract_inbound_headers( + request_data=post_call_data[0], logging_obj=None, extra_allowlist={"x-tenant"} + ) + assert vendor_headers is not None and vendor_headers["x-tenant"] == "tenant-real" From 4b63b53fe471fac94811218ce8ff6c762d1a8057 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Wed, 30 Sep 2026 17:42:36 +0300 Subject: [PATCH 4/4] test(guardrails): assert whole guardrail payloads and cover the blocked-request log The apply_guardrail tests compare the whole request_data again, with the expected identity built from get_authenticated_identity_metadata and the key hash pinned, so an extra leaked key or a wrong hash fails them. A new pass-through test checks that a request blocked by a pre-call guardrail logs the cleaned and redacted inbound headers. The realtime Gray Swan test now injects an HTTP transport instead of overriding a private method, so the real request path runs _ensure_litellm_metadata is renamed to _apply_authenticated_identity_to_litellm_metadata, and docstrings that only repeated the code are gone. The apply_guardrail strip no longer lists user_api_key, because the authenticated identity always overwrites it. The identity helpers now return Mapping, since no caller mutates what they return --- .../unified_guardrail/unified_guardrail.py | 7 +- litellm/proxy/litellm_pre_call_utils.py | 18 ++- .../pass_through_endpoints.py | 2 +- .../guardrail_hooks/test_grayswan.py | 21 ++-- .../test_unified_guardrail.py | 9 +- .../guardrails/test_guardrail_endpoints.py | 106 +++++++----------- .../test_pass_through_endpoints.py | 51 +++++++-- .../test_apply_guardrail_endpoint.py | 28 +++-- .../test_realtime_streaming.py | 30 ++--- 9 files changed, 134 insertions(+), 138 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 0743af68e18..ba0e97ea65d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -156,8 +156,7 @@ def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapp _PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_metadata"}) -def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None: - """Overwrite the identity fields of data['litellm_metadata'] from the authenticated key, in place.""" +def _apply_authenticated_identity_to_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None: existing: Final = data.get("litellm_metadata") if isinstance(existing, dict): identity: Final = LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict) @@ -228,7 +227,7 @@ class UnifiedLLMGuardrails(CustomLogger): endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]()) - _ensure_litellm_metadata(data, user_api_key_dict) + _apply_authenticated_identity_to_litellm_metadata(data, user_api_key_dict) data = await endpoint_translation.process_input_messages( data=data, @@ -274,7 +273,7 @@ class UnifiedLLMGuardrails(CustomLogger): endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]()) - _ensure_litellm_metadata(data, user_api_key_dict) + _apply_authenticated_identity_to_litellm_metadata(data, user_api_key_dict) return await endpoint_translation.process_input_messages( data=data, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 9201fcd7ca7..8a7b79f69d4 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -500,7 +500,6 @@ def is_untrusted_caller_metadata_key(key: str) -> bool: def strip_untrusted_caller_metadata( data: MutableMapping[str, object], *, allow_client_message_redaction_opt_out: bool ) -> None: - """Remove, in place, the proxy-owned slots a caller put in either metadata bucket of a request body.""" for user_meta in (data.get("metadata"), data.get("litellm_metadata")): if not isinstance(user_meta, dict): continue @@ -512,17 +511,13 @@ def strip_untrusted_caller_metadata( user_meta.pop(untrusted_key, None) -_GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS: Final = frozenset({"user_api_key", "headers"}) - - def caller_metadata_with_authenticated_identity( caller_metadata: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth -) -> dict[str, object]: - """Caller metadata minus proxy-owned slots, bare user_api_key and headers, with the key's identity on top.""" +) -> Mapping[str, object]: caller_fields: Final = { key: value for key, value in (caller_metadata or {}).items() - if not (is_untrusted_caller_metadata_key(key) or key in _GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS) + if not (is_untrusted_caller_metadata_key(key) or key == "headers") } return {**caller_fields, **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)} @@ -1677,7 +1672,7 @@ class LiteLLMProxyRequestSetup: return user_api_key_logged_metadata @staticmethod - def get_key_scoped_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: + def get_key_scoped_metadata(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object]: return { "user_api_key_metadata": strip_callback_config(user_api_key_dict.metadata), "user_api_key_team_metadata": strip_callback_config(user_api_key_dict.team_metadata), @@ -1686,8 +1681,7 @@ class LiteLLMProxyRequestSetup: } @staticmethod - def get_authenticated_identity_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: - """Identity fields derived from the authenticated key alone, for paths that skip the chat-path build.""" + def get_authenticated_identity_metadata(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object]: return { **LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict), "user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict), @@ -2053,7 +2047,9 @@ async def add_litellm_data_to_request( # These keys are injected by the proxy itself below — user-supplied values # must not be trusted. _allow_client_mock_response: Final = _key_or_team_allows_client_mock_response(user_api_key_dict) - _allow_client_message_redaction_opt_out = key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict) + _allow_client_message_redaction_opt_out: Final = key_or_team_allows_client_message_redaction_opt_out( + user_api_key_dict + ) for _internal_key in _UNTRUSTED_ROOT_CONTROL_FIELDS: if _allow_client_mock_response and _internal_key in _CLIENT_MOCK_CONTROL_FIELDS: continue diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 7bd1bb7871f..9acc29b0b44 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -1148,7 +1148,7 @@ async def pass_through_request( general_settings_view, ) - # Guardrails forward these to vendors as the inbound request headers; all are popped before the upstream send. + # Only the proxy's own view of the inbound headers may reach guardrail vendors _parsed_body.pop("proxy_server_request", None) _parsed_body.pop("headers", None) for _caller_bucket in (_parsed_body.get("metadata"), _parsed_body.get("litellm_metadata")): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py index cce13844201..4b668c0b5bb 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -10,6 +10,9 @@ from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import ( GraySwanGuardrail, GraySwanGuardrailAPIError, ) +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + _apply_authenticated_identity_to_litellm_metadata, +) from litellm.types.guardrails import GuardrailEventHooks @@ -566,33 +569,23 @@ def test_prepare_payload_includes_litellm_metadata( assert payload["litellm_metadata"]["user_api_key_team_id"] == "team-456" -def test_ensure_litellm_metadata_populates_from_user_api_key_dict() -> None: - """Verify _ensure_litellm_metadata populates litellm_metadata.""" - from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( - _ensure_litellm_metadata, - ) - +def test_missing_litellm_metadata_is_populated_from_user_api_key_dict() -> None: user_auth = UserAPIKeyAuth(user_id="u1", team_id="t1", api_key="sk-test-hashed") data: dict = {} - _ensure_litellm_metadata(data, user_auth) + _apply_authenticated_identity_to_litellm_metadata(data, user_auth) assert "litellm_metadata" in data assert data["litellm_metadata"]["user_api_key_user_id"] == "u1" assert data["litellm_metadata"]["user_api_key_team_id"] == "t1" -def test_ensure_litellm_metadata_overrides_caller_identity_in_existing_bucket() -> None: - """An existing litellm_metadata keeps its other keys, but its identity comes from the authenticated key.""" - from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( - _ensure_litellm_metadata, - ) - +def test_existing_litellm_metadata_keeps_its_keys_but_takes_the_authenticated_identity() -> None: user_auth = UserAPIKeyAuth(user_id="auth-user", key_alias="auth-alias", team_id="auth-team") bucket: dict = {"existing": "value", "user_api_key_alias": "batch-worker", "user_api_key_team_id": "team-exempt"} data: dict = {"litellm_metadata": bucket} - _ensure_litellm_metadata(data, user_auth) + _apply_authenticated_identity_to_litellm_metadata(data, user_auth) assert data["litellm_metadata"] is bucket assert bucket["existing"] == "value" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index afebcbdf8a6..be2ec795e41 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -2434,8 +2434,6 @@ def _cli_session_key(route: str) -> UserAPIKeyAuth: class TestGuardrailsSeeAuthenticatedIdentity: - """A request body cannot make a guardrail vendor see another key's identity, and the real one reaches it.""" - @staticmethod def _generic_guardrail(vendor_payloads: list[dict[str, object]]) -> CustomGuardrail: def vendor(request: httpx.Request) -> httpx.Response: @@ -2505,7 +2503,6 @@ class TestGuardrailsSeeAuthenticatedIdentity: @pytest.mark.asyncio async def test_pass_through_body_cannot_forge_request_route(self, monkeypatch) -> None: - """Guardrails and call-type lookups key on user_api_key_request_route, so it must be the key's own route.""" _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) key = UserAPIKeyAuth(api_key="sk-real-caller-key", request_route="/openai/v1/chat/completions") data = { @@ -2528,8 +2525,6 @@ class TestGuardrailsSeeAuthenticatedIdentity: async def test_chat_path_request_keeps_proxy_metadata_and_sends_stable_hash( self, monkeypatch, route: str, make_key: Callable[[str], UserAPIKeyAuth] ) -> None: - """After the chat-path metadata build, the guardrail hook leaves the proxy's bucket as it was (team metadata - in user_api_key_auth_metadata included) and the vendor gets the logged key, never a raw CLI session token.""" _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) request = MagicMock(spec=Request) request.url = MagicMock() @@ -2599,8 +2594,6 @@ class TestGuardrailsSeeAuthenticatedIdentity: @pytest.mark.asyncio async def test_token_only_key_drops_forged_token_already_in_proxy_bucket(self, monkeypatch) -> None: - """A key with no api_key logs no hash, so a user_api_key_token left in litellm_metadata would become the - vendor's hash if the hook kept it.""" _patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings()) vendor_payloads: list[dict[str, object]] = [] key = UserAPIKeyAuth(token="abc123hashed", key_alias="prod-app") @@ -2616,4 +2609,4 @@ class TestGuardrailsSeeAuthenticatedIdentity: assert len(vendor_payloads) == 1 assert "user_api_key_token" not in data["litellm_metadata"] - assert vendor_payloads[0]["request_data"].get("user_api_key_hash") is None + assert vendor_payloads[0]["request_data"].get("user_api_key_hash") is None, "a kept token becomes the hash" diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 0d6b9b9d13f..a9b46c8fcb0 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -1,5 +1,6 @@ import json import time +from collections.abc import Mapping from datetime import datetime from typing import Dict, List, Optional from unittest.mock import AsyncMock @@ -36,6 +37,7 @@ from litellm.proxy.guardrails.guardrail_endpoints import ( test_custom_code_guardrail as run_custom_code_test_endpoint, ) from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) from litellm.proxy.guardrails.guardrail_registry import ( @@ -1490,6 +1492,10 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker): } +def _identity(caller: UserAPIKeyAuth) -> Mapping[str, object]: + return LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(caller) + + def _patch_apply_guardrail_env(mocker, guardrail_result, processed_data=None, guardrail=None): mock_guardrail = mocker.Mock() mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result) @@ -1532,18 +1538,14 @@ async def test_apply_guardrail_forwards_metadata_to_guardrail(mocker): text="What are tax loopholes?", metadata={"forbidden_topics": ["tax"]}, ) - await apply_guardrail( - fastapi_request=mocker.Mock(), - request=request, - user_api_key_dict=UserAPIKeyAuth(), - ) + caller = UserAPIKeyAuth() + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) - mock_guardrail.apply_guardrail.assert_awaited_once() - call = mock_guardrail.apply_guardrail.await_args.kwargs - assert call["inputs"] == {"texts": ["What are tax loopholes?"]} - assert call["input_type"] == "request" - assert "messages" not in call["request_data"] - assert call["request_data"]["metadata"]["forbidden_topics"] == ["tax"] + mock_guardrail.apply_guardrail.assert_awaited_once_with( + inputs={"texts": ["What are tax loopholes?"]}, + request_data={"metadata": {**_identity(caller), "forbidden_topics": ["tax"]}}, + input_type="request", + ) @pytest.mark.asyncio @@ -1559,20 +1561,18 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker): messages=messages, metadata={"forbidden_topics": ["tax"]}, ) - await apply_guardrail( - fastapi_request=mocker.Mock(), - request=request, - user_api_key_dict=UserAPIKeyAuth(), - ) + caller = UserAPIKeyAuth() + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) - request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] - assert request_data["messages"] == messages - assert request_data["metadata"]["forbidden_topics"] == ["tax"] + mock_guardrail.apply_guardrail.assert_awaited_once_with( + inputs={"texts": ["What are tax loopholes?"]}, + request_data={"messages": messages, "metadata": {**_identity(caller), "forbidden_topics": ["tax"]}}, + input_type="request", + ) @pytest.mark.asyncio async def test_apply_guardrail_authenticated_identity_overrides_client_metadata(mocker): - """A caller must not be able to claim another key's or team's identity in the body metadata.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) caller = UserAPIKeyAuth( api_key="sk-real-caller-key", @@ -1595,18 +1595,17 @@ async def test_apply_guardrail_authenticated_identity_overrides_client_metadata( await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] - assert metadata["user_api_key_alias"] == "real-caller" - assert metadata["user_api_key_team_id"] == "real-team" - assert metadata["user_api_key_user_id"] == "real-user" - assert metadata["user_api_key_hash"] == caller.api_key - assert metadata["user_api_key_hash"] != "forged-hash" - assert metadata["forbidden_topics"] == ["tax"] + assert metadata == {**_identity(caller), "forbidden_topics": ["tax"]} + assert ( + metadata["user_api_key_alias"], + metadata["user_api_key_team_id"], + metadata["user_api_key_user_id"], + metadata["user_api_key_hash"], + ) == ("real-caller", "real-team", "real-user", caller.api_key) @pytest.mark.asyncio async def test_apply_guardrail_drops_client_identity_fields_the_key_does_not_set(mocker): - """Proxy-owned slots in the body never reach the guardrail, including user_api_key_token, which guardrails - map onto the key hash, and control fields the chat path also strips.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) caller = UserAPIKeyAuth(metadata={"zguard_policy_id": "strict"}, object_permission_id="perm-real") @@ -1628,20 +1627,13 @@ async def test_apply_guardrail_drops_client_identity_fields_the_key_does_not_set await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] - assert metadata["user_api_key_alias"] is None - assert metadata["user_api_key_team_id"] is None - assert "user_api_key_token" not in metadata + assert metadata == {**_identity(caller), "trace_label": "nightly"}, "proxy-owned slots must not reach it" assert metadata["user_api_key_metadata"] == {"zguard_policy_id": "strict"} assert metadata["user_api_key_object_permission_id"] == "perm-real" - assert metadata["user_api_key"] is None - assert "applied_guardrails" not in metadata - assert "headers" not in metadata - assert metadata["trace_label"] == "nightly" @pytest.mark.asyncio async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied_ones(mocker): - """Guardrails forward metadata headers to vendors, so they must be the proxy's view of the request.""" real_headers = {"user-agent": "real-client/1.0"} mock_guardrail = _patch_apply_guardrail_env( mocker, @@ -1654,16 +1646,15 @@ async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied text="hello", metadata={"headers": {"user-agent": "forged/1.0", "x-end-user": "someone-else"}}, ) - await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=UserAPIKeyAuth()) + caller = UserAPIKeyAuth() + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] - assert metadata["headers"] == real_headers + assert metadata == {**_identity(caller), "headers": real_headers} @pytest.mark.asyncio async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker): - """End to end through a real GenericGuardrailAPI: the vendor payload names the authenticated key even when - the body forges user_api_key_alias and user_api_key_token, which the generic guardrail maps onto the hash.""" vendor_payloads = [] def vendor(request: httpx.Request) -> httpx.Response: @@ -1692,7 +1683,6 @@ async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identit @pytest.mark.asyncio async def test_apply_guardrail_request_route_comes_from_the_key(mocker): - """Guardrails pick call-type behavior from user_api_key_request_route, so the body cannot choose it.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) request = ApplyGuardrailRequest( @@ -1712,7 +1702,6 @@ async def test_apply_guardrail_request_route_comes_from_the_key(mocker): @pytest.mark.asyncio async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocker): - """A CLI session key's raw per-login token must never reach the vendor; it gets the stable logged key.""" raw_session_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" vendor_payloads = [] @@ -1739,27 +1728,21 @@ async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocke @pytest.mark.asyncio async def test_apply_guardrail_carries_authenticated_identity_when_no_metadata_sent(mocker): - """request_data always carries the authenticated identity, even when the body has no metadata.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) + caller = UserAPIKeyAuth(key_alias="known-caller", team_id="known-team") request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello") - await apply_guardrail( - fastapi_request=mocker.Mock(), - request=request, - user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller", team_id="known-team"), - ) + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) - call = mock_guardrail.apply_guardrail.await_args.kwargs - assert call["inputs"] == {"texts": ["hello"]} - assert "messages" not in call["request_data"] - assert call["request_data"]["metadata"]["user_api_key_alias"] == "known-caller" - assert call["request_data"]["metadata"]["user_api_key_team_id"] == "known-team" + mock_guardrail.apply_guardrail.assert_awaited_once_with( + inputs={"texts": ["hello"]}, request_data={"metadata": _identity(caller)}, input_type="request" + ) + metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"] + assert (metadata["user_api_key_alias"], metadata["user_api_key_team_id"]) == ("known-caller", "known-team") @pytest.mark.asyncio async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(mocker): - """Explicitly-sent empty messages must be forwarded, not dropped, and empty - metadata still carries the authenticated identity.""" mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]}) request = ApplyGuardrailRequest( @@ -1768,15 +1751,12 @@ async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(moc messages=[], metadata={}, ) - await apply_guardrail( - fastapi_request=mocker.Mock(), - request=request, - user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller"), - ) + caller = UserAPIKeyAuth(key_alias="known-caller") + await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller) - request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] - assert request_data["messages"] == [] - assert request_data["metadata"]["user_api_key_alias"] == "known-caller" + mock_guardrail.apply_guardrail.assert_awaited_once_with( + inputs={"texts": ["hello"]}, request_data={"messages": [], "metadata": _identity(caller)}, input_type="request" + ) @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 1b397ea3286..0384b47ad7e 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -7632,11 +7632,6 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token() @pytest.mark.asyncio async def test_pass_through_request_strips_caller_identity_before_guardrail_hooks(): - """ - Regression: a pass-through body skips add_litellm_data_to_request, so forged user_api_key_* fields, guardrail - control fields and inbound headers reached pre_call_hook guardrails as the caller's identity. The upstream body - is unchanged because these keys never reach it. - """ upstream_bodies = [] def transport_handler(upstream_request: httpx.Request) -> httpx.Response: @@ -7731,10 +7726,48 @@ async def test_pass_through_request_strips_caller_identity_before_guardrail_hook @pytest.mark.asyncio -async def test_pass_through_post_call_guardrails_receive_real_inbound_headers(): - """Post-call guardrails run on a copy of the body the litellm-param pop already stripped, so without an explicit - re-attach an operator's extra_headers allowlist forwarded nothing on the response side.""" +async def test_pass_through_pre_call_block_logs_cleaned_inbound_headers(): + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=HTTPException(status_code=400, detail="blocked")) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.headers = Headers( + { + "content-type": "application/json", + "x-tenant": "tenant-real", + "authorization": "Bearer sk-real-caller-key", + "cookie": "session=secret", + } + ) + mock_request.query_params = QueryParams({}) + forged_headers = {"x-tenant": "forged"} + mock_request.body = AsyncMock( + return_value=json.dumps( + {"prompt": "hello", "headers": forged_headers, "proxy_server_request": {"headers": forged_headers}} + ).encode() + ) + with ( + patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), # test-quality-ok: read at call time + pytest.raises(ProxyException), + ): + await pass_through_request( + request=mock_request, + target="https://upstream.test/v1/generate", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app"), + ) + + mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + logged_request = mock_proxy_logging.post_call_failure_hook.await_args.kwargs["request_data"] + assert logged_request["proxy_server_request"] == { + "headers": {"content-type": "application/json", "x-tenant": "tenant-real", "cookie": _REDACTED_HEADER_VALUE} + } + + +@pytest.mark.asyncio +async def test_pass_through_post_call_guardrails_receive_real_inbound_headers(): def transport_handler(upstream_request: httpx.Request) -> httpx.Response: return httpx.Response(200, json={"completion": "hi"}) @@ -7786,7 +7819,7 @@ async def test_pass_through_post_call_guardrails_receive_real_inbound_headers(): assert len(post_call_data) == 1, "the post-call guardrail hook did not run" assert post_call_data[0]["proxy_server_request"] == { "headers": {"content-type": "application/json", "x-tenant": "tenant-real"} - } + }, "the post-call body copy is already stripped, so the headers must be re-attached" vendor_headers = _extract_inbound_headers( request_data=post_call_data[0], logging_obj=None, extra_allowlist={"x-tenant"} ) diff --git a/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py b/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py index bab9958dc5d..316deb20350 100644 --- a/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py +++ b/tests/unit/enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py @@ -11,9 +11,17 @@ from fastapi import HTTPException from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.guardrails import ApplyGuardrailRequest, ApplyGuardrailResponse +def _identity_with_hash(caller: UserAPIKeyAuth) -> dict[str, object]: + return { + **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(caller), + "user_api_key_hash": caller.api_key, + } + + @pytest.mark.asyncio async def test_apply_guardrail_endpoint_returns_correct_response( mock_proxy_logging_ctx, @@ -61,11 +69,11 @@ async def test_apply_guardrail_endpoint_returns_correct_response( assert response.response_text == "Redacted text: [REDACTED] and [REDACTED]" # Verify the guardrail was called with correct parameters - mock_guardrail.apply_guardrail.assert_called_once() - call = mock_guardrail.apply_guardrail.call_args.kwargs - assert call["inputs"] == {"texts": ["Test text with PII"]} - assert call["input_type"] == "request" - assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key + mock_guardrail.apply_guardrail.assert_called_once_with( + inputs={"texts": ["Test text with PII"]}, + request_data={"metadata": _identity_with_hash(user_api_key_dict)}, + input_type="request", + ) @pytest.mark.asyncio @@ -197,8 +205,8 @@ async def test_apply_guardrail_endpoint_without_optional_params(mock_proxy_loggi assert response.response_text == "Processed text" # Verify the guardrail was called with correct parameters - mock_guardrail.apply_guardrail.assert_called_once() - call = mock_guardrail.apply_guardrail.call_args.kwargs - assert call["inputs"] == {"texts": ["Test text"]} - assert call["input_type"] == "request" - assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key + mock_guardrail.apply_guardrail.assert_called_once_with( + inputs={"texts": ["Test text"]}, + request_data={"metadata": _identity_with_hash(user_api_key_dict)}, + input_type="request", + ) diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index a5713b355af..fdd0f34fbf6 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -5,6 +5,7 @@ from dataclasses import dataclass from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from websockets.exceptions import ConnectionClosed from websockets.frames import Close @@ -16,6 +17,7 @@ from litellm.litellm_core_utils.realtime_streaming import ( RealTimeStreaming, client_sent_openai_beta_realtime_header, ) +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail @@ -3556,7 +3558,6 @@ async def test_provider_bytes_are_sent_raw_after_pacing(): @pytest.mark.asyncio async def test_realtime_transcript_guardrail_receives_authenticated_identity(monkeypatch: pytest.MonkeyPatch): - """Transcript guardrails get the session key's identity in litellm_metadata, as the chat path provides it.""" received_request_data = [] class IdentityRecordingGuardrail(CustomGuardrail): @@ -3590,26 +3591,20 @@ async def test_realtime_transcript_guardrail_receives_authenticated_identity(mon @pytest.mark.asyncio async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pytest.MonkeyPatch): - """Gray Swan forwards litellm_metadata verbatim, so realtime must hand it identity and no key secrets.""" vendor_payloads: list[dict[str, object]] = [] - class RecordingGraySwan(GraySwanGuardrail): - async def _call_grayswan_api(self, payload): - vendor_payloads.append(payload) - return {"violation": 0.0, "violated_rules": []} + def vendor(request: httpx.Request) -> httpx.Response: + vendor_payloads.append(json.loads(request.content)) + return httpx.Response(200, json={"violation": 0.0, "violated_rules": []}) - monkeypatch.setattr( - litellm, - "callbacks", - [ - RecordingGraySwan( - guardrail_name="grayswan", - api_key="test-key", - event_hook=GuardrailEventHooks.pre_call, - default_on=True, - ) - ], + grayswan = GraySwanGuardrail( + guardrail_name="grayswan", + api_key="test-key", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, ) + grayswan.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor)) + monkeypatch.setattr(litellm, "callbacks", [grayswan]) key = UserAPIKeyAuth( api_key="sk-real-caller-key", key_alias="prod-app", @@ -3643,7 +3638,6 @@ async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pyte async def test_realtime_guardrail_gets_no_identity_from_non_auth_sdk_value( monkeypatch: pytest.MonkeyPatch, sdk_value: object ): - """Only a proxy-authenticated UserAPIKeyAuth yields identity; an SDK-supplied value never raises or fakes one.""" received_request_data: list[dict[str, object]] = [] class IdentityRecordingGuardrail(CustomGuardrail):