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
This commit is contained in:
Caduri Katzav 2026-09-30 12:55:06 +03:00
parent 88e0c43e7c
commit 8fe7b4f00f
4 changed files with 133 additions and 15 deletions

View file

@ -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):

View file

@ -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,

View file

@ -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:

View file

@ -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"