From 8fe7b4f00f9f6a2545265a8c0dafeaccfcb7cbb1 Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Wed, 30 Sep 2026 12:55:06 +0300 Subject: [PATCH] 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"