Generic Guardrails: Forward request headers + litellm_version to gene… (#20729)

* Generic Guardrails: Forward request headers + litellm_version to generic guardrail API

* Generic Guardrail: Change the request headers addition to be with allowlist instead denylist
This commit is contained in:
Itay Ovadia 2026-02-11 08:41:05 +02:00 • committed by GitHub
parent 72682f4bd4
commit 126522ab91
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 179 additions and 0 deletions

View file

@ -93,6 +93,12 @@ Implement `POST /beta/litellm_basic_guardrail_api`
"user_api_key_end_user_id": "end user id associated with the litellm virtual key used",
"user_api_key_org_id": "org id associated with the litellm virtual key used"
},
"request_headers": { // optional: inbound request headers (allowlist). Allowed headers show their value; all others show "[present]" to indicate the header existed.
"User-Agent": "OpenAI/Python 2.17.0",
"Content-Type": "application/json",
"X-Request-Id": "[present]"
},
"litellm_version": "1.x.y", // optional: LiteLLM library version running this proxy
"input_type": "request", // "request" or "response"
"litellm_call_id": "unique_call_id", // the call id of the individual LLM call
"litellm_trace_id": "trace_id", // the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation

View file

@ -5,10 +5,12 @@
# +-------------------------------------------------------------+
# Thank you users! We ❤️ you! - Krrish & Ishaan
import fnmatch
import os
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
@ -31,6 +33,110 @@ if TYPE_CHECKING:
GUARDRAIL_NAME = "generic_guardrail_api"
# Headers whose values are forwarded as-is (case-insensitive). Glob patterns supported (e.g. x-stainless-*, x-litellm*).
_HEADER_VALUE_ALLOWLIST = frozenset({
"host",
"accept-encoding",
"connection",
"accept",
"content-type",
"user-agent",
"x-stainless-*",
"x-litellm-*",
"content-length",
})
# Placeholder for headers that exist but are not on the allowlist (we don't expose their value).
_HEADER_PRESENT_PLACEHOLDER = "[present]"
def _header_value_allowed(header_name: str) -> bool:
"""Return True if this header's value may be forwarded (allowlist, including globs)."""
lower = header_name.lower()
if lower in _HEADER_VALUE_ALLOWLIST:
return True
for pattern in _HEADER_VALUE_ALLOWLIST:
if "*" in pattern and fnmatch.fnmatch(lower, pattern):
return True
return False
def _sanitize_inbound_headers(headers: Any) -> Optional[Dict[str, str]]:
"""
Sanitize inbound headers before passing them to a 3rd party guardrail service.
- Allowlist: only headers in the allowlist have their values forwarded (exact + glob: x-stainless-*, x-litellm-*).
- All other headers are included with value "[present]" so the guardrail knows the header existed.
- Coerces values to str (for JSON serialization).
"""
if not headers or not isinstance(headers, dict):
return None
sanitized: Dict[str, str] = {}
for k, v in headers.items():
if k is None:
continue
key = str(k)
if _header_value_allowed(key):
try:
sanitized[key] = str(v)
except Exception:
continue
else:
sanitized[key] = _HEADER_PRESENT_PLACEHOLDER
return sanitized or None
def _extract_inbound_headers(
request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]
) -> Optional[Dict[str, str]]:
"""
Extract inbound headers from available request context.
We try multiple locations to support different call paths:
- proxy endpoints: request_data["proxy_server_request"]["headers"]
- if the guardrail is passed the proxy_server_request object directly
- metadata headers captured in litellm_pre_call_utils
- response hooks: fallback to logging_obj.model_call_details
"""
# 1) Most common path (proxy): full request context in proxy_server_request
headers = request_data.get("proxy_server_request", {}).get("headers")
if headers:
return _sanitize_inbound_headers(headers)
# 2) Some guardrails pass proxy_server_request as request_data itself
headers = request_data.get("headers")
if headers:
return _sanitize_inbound_headers(headers)
# 3) Pre-call: headers stored in request metadata
metadata_headers = (request_data.get("metadata") or {}).get("headers")
if metadata_headers:
return _sanitize_inbound_headers(metadata_headers)
litellm_metadata_headers = (request_data.get("litellm_metadata") or {}).get(
"headers"
)
if litellm_metadata_headers:
return _sanitize_inbound_headers(litellm_metadata_headers)
# 4) Post-call: headers not present on response; fallback to logging object
if logging_obj and getattr(logging_obj, "model_call_details", None):
try:
details = logging_obj.model_call_details or {}
headers = (
details.get("litellm_params", {})
.get("metadata", {})
.get("headers", None)
)
if headers:
return _sanitize_inbound_headers(headers)
except Exception:
pass
return None
class GenericGuardrailAPI(CustomGuardrail):
"""
@ -207,6 +313,7 @@ class GenericGuardrailAPI(CustomGuardrail):
# Extract user API key metadata
user_metadata = self._extract_user_api_key_metadata(request_data)
inbound_headers = _extract_inbound_headers(request_data=request_data, logging_obj=logging_obj)
# Create request payload
guardrail_request = GenericGuardrailAPIRequest(
@ -214,6 +321,8 @@ class GenericGuardrailAPI(CustomGuardrail):
litellm_trace_id=logging_obj.litellm_trace_id if logging_obj else None,
texts=texts,
request_data=user_metadata,
request_headers=inbound_headers,
litellm_version=litellm_version,
images=images,
tools=tools,
structured_messages=structured_messages,

View file

@ -60,6 +60,14 @@ class GenericGuardrailAPIRequest(BaseModel):
tools: Optional[List[ChatCompletionToolParam]] = None
texts: Optional[List[str]] = None
request_data: GenericGuardrailAPIMetadata
request_headers: Optional[Dict[str, str]] = Field(
default=None,
description="Sanitized inbound request headers from the original proxy request.",
)
litellm_version: Optional[str] = Field(
default=None,
description="LiteLLM library version running this proxy.",
)
additional_provider_specific_params: Optional[Dict[str, Any]] = None
tool_calls: Optional[
Union[List[ChatCompletionToolCallChunk], List[ChatCompletionMessageToolCall]]

View file

@ -14,10 +14,14 @@ import pytest
import litellm
from litellm import ModelResponse
from litellm.exceptions import GuardrailRaisedException
from litellm._version import version as litellm_version
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPI,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import (
_HEADER_PRESENT_PLACEHOLDER,
)
from litellm.types.utils import Choices, Message
@ -351,6 +355,58 @@ class TestMetadataExtraction:
# Should be empty dict
assert request_metadata == {}
@pytest.mark.asyncio
async def test_inbound_headers_and_litellm_version_forwarded_and_sanitized(
self, generic_guardrail, mock_request_data_input
):
"""
Ensure inbound proxy request headers are forwarded in JSON payload with allowlist:
allowed headers show their value; all other headers show presence only ([present]).
"""
# Add proxy_server_request headers as they exist in proxy request context
request_data = dict(mock_request_data_input)
request_data["proxy_server_request"] = {
"headers": {
"User-Agent": "OpenAI/Python 2.17.0",
"Authorization": "Bearer should-not-forward",
"Cookie": "session=should-not-forward",
"X-Request-Id": "req_123",
}
}
mock_response = MagicMock()
mock_response.json.return_value = {
"action": "NONE",
"texts": ["test"],
}
mock_response.raise_for_status = MagicMock()
with patch.object(
generic_guardrail.async_handler, "post", return_value=mock_response
) as mock_post:
await generic_guardrail.apply_guardrail(
inputs={"texts": ["test"]},
request_data=request_data,
input_type="request",
)
call_args = mock_post.call_args
json_payload = call_args.kwargs["json"]
# New fields should exist
assert json_payload["litellm_version"] == litellm_version
assert "request_headers" in json_payload
assert isinstance(json_payload["request_headers"], dict)
req_headers = json_payload["request_headers"]
# Allowed: value forwarded
assert req_headers.get("User-Agent") == "OpenAI/Python 2.17.0"
# Not on allowlist: key present, value is placeholder only
assert req_headers.get("Authorization") == _HEADER_PRESENT_PLACEHOLDER
assert req_headers.get("Cookie") == _HEADER_PRESENT_PLACEHOLDER
assert req_headers.get("X-Request-Id") == _HEADER_PRESENT_PLACEHOLDER
class TestGuardrailActions:
"""Test different guardrail action responses"""