From 4b63b53fe471fac94811218ce8ff6c762d1a8057 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Wed, 30 Sep 2026 17:42:36 +0300 Subject: [PATCH] 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):