mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
72682f4bd4
commit
126522ab91
4 changed files with 179 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue