fix(guardrails): forward grayswan scan id header (#32544)

* fix(guardrails): forward grayswan scan id header

* test(guardrails): cover grayswan scan id forwarding

* fix(guardrails): prevent overwriting existing metadata headers when extracting scan id

* test(guardrails): cover header merging logic

* chore(guardrails): fix formatting

* test(guardrails): enforce case preservation

* chore(guardrails): corrected grayswan type annotations

* fix(guardrails): sanitized grayswan header metadata

* test(guardrails): covered grayswan logging headers

* fix(guardrails): guard grayswan header lookup against None and drop dead comment

- Fall back to {} when proxy_server_request is explicitly None so
  request_data.get(...).get('headers') never raises AttributeError.
- Remove the commented-out user_api_key_auth pop; it was inert and
  greptile called it out as ambiguous.

---------

Co-authored-by: Theodore Drzewinski <93957989+tediferJones@users.noreply.github.com>
This commit is contained in:
yucheng-berri 2026-07-08 15:05:27 -07:00 • committed by GitHub
parent 5973d9fd2b
commit 528fa380f5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 159 additions and 36 deletions

View file

@ -213,7 +213,7 @@ class GraySwanGuardrail(CustomGuardrail):
verbose_proxy_logger.debug("Gray Swan Guardrail: dynamic extra_body=%s", safe_dumps(dynamic_body))
# Prepare and send payload
payload = self._prepare_payload(messages, dynamic_body, request_data)
payload = self._prepare_payload(messages, dynamic_body, request_data, logging_obj)
if payload is None:
return inputs
@ -502,10 +502,38 @@ class GraySwanGuardrail(CustomGuardrail):
"grayswan-api-key": self.api_key,
}
def _extract_inbound_headers(
self,
request_data: dict,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> Optional[dict[str, str]]:
headers = (request_data.get("proxy_server_request") or {}).get("headers")
if not headers:
headers = request_data.get("headers")
if not headers:
headers = (request_data.get("metadata") or {}).get("headers")
if not headers and logging_obj and getattr(logging_obj, "model_call_details", None):
headers = (
(logging_obj.model_call_details or {}).get("litellm_params", {}).get("metadata", {}).get("headers")
)
if not isinstance(headers, dict):
return None
forwarded_header_names = ("shade_scan_id",)
forwarded_headers = {}
for key, value in headers.items():
if str(key).lower() in forwarded_header_names:
forwarded_headers[str(key)] = str(value)
return forwarded_headers or None
def _prepare_payload(
self, messages: List[Dict[str, str]], dynamic_body: dict, request_data: dict
) -> Optional[Dict[str, Any]]:
payload: Dict[str, Any] = {"messages": messages}
self,
messages: list[dict[str, str]],
dynamic_body: dict,
request_data: dict,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> Optional[dict[str, Any]]:
payload: dict[str, Any] = {"messages": messages}
categories = dynamic_body.get("categories") or self.categories
if categories:
@ -523,10 +551,16 @@ class GraySwanGuardrail(CustomGuardrail):
if "metadata" in dynamic_body:
payload["metadata"] = dynamic_body["metadata"]
inbound_headers = self._extract_inbound_headers(request_data, logging_obj)
litellm_metadata = request_data.get("litellm_metadata")
if isinstance(litellm_metadata, dict) and litellm_metadata:
cleaned_litellm_metadata = dict(litellm_metadata)
# cleaned_litellm_metadata.pop("user_api_key_auth", None)
cleaned_litellm_metadata = dict(litellm_metadata) if isinstance(litellm_metadata, dict) else {}
if inbound_headers:
existing_headers = cleaned_litellm_metadata.get("headers")
cleaned_litellm_metadata["headers"] = (
{**existing_headers, **inbound_headers} if isinstance(existing_headers, dict) else inbound_headers
)
if cleaned_litellm_metadata:
sanitized = safe_json_loads(safe_dumps(cleaned_litellm_metadata), default={})
if isinstance(sanitized, dict) and sanitized:
payload["litellm_metadata"] = sanitized

View file

@ -5,6 +5,7 @@ from fastapi import HTTPException
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.grayswan import grayswan as grayswan_module
from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import (
GraySwanGuardrail,
GraySwanGuardrailAPIError,
@ -70,12 +71,118 @@ def test_prepare_payload_includes_dynamic_metadata(
assert payload["metadata"] == dynamic_body["metadata"]
def test_prepare_payload_forwards_only_scan_id_header(
grayswan_guardrail: GraySwanGuardrail,
) -> None:
messages = [{"role": "user", "content": "hello"}]
request_data = {
"proxy_server_request": {
"headers": {
"SHADE_SCAN_ID": "scan-123",
"authorization": "Bearer secret",
}
},
"litellm_metadata": {"request_id": "request-123"},
}
payload = grayswan_guardrail._prepare_payload(messages, {}, request_data)
assert payload["litellm_metadata"] == {
"request_id": "request-123",
"headers": {"SHADE_SCAN_ID": "scan-123"},
}
def test_prepare_payload_merges_scan_id_with_existing_metadata_headers(
grayswan_guardrail: GraySwanGuardrail,
) -> None:
messages = [{"role": "user", "content": "hello"}]
request_data = {
"proxy_server_request": {
"headers": {
"shade_scan_id": "scan-123",
}
},
"litellm_metadata": {
"request_id": "request-123",
"headers": {"x-existing": "keep-me"},
},
}
payload = grayswan_guardrail._prepare_payload(messages, {}, request_data)
assert payload["litellm_metadata"] == {
"request_id": "request-123",
"headers": {
"x-existing": "keep-me",
"shade_scan_id": "scan-123",
},
}
def test_prepare_payload_sanitizes_headers_when_litellm_metadata_absent(
monkeypatch: pytest.MonkeyPatch,
grayswan_guardrail: GraySwanGuardrail,
) -> None:
messages = [{"role": "user", "content": "hello"}]
request_data = {
"proxy_server_request": {
"headers": {
"shade_scan_id": "scan-123",
}
}
}
monkeypatch.setattr(grayswan_module, "safe_dumps", lambda _data: "{}")
payload = grayswan_guardrail._prepare_payload(messages, {}, request_data)
assert "litellm_metadata" not in payload
def test_prepare_payload_extracts_headers_from_logging_obj(
grayswan_guardrail: GraySwanGuardrail,
) -> None:
messages = [{"role": "user", "content": "hello"}]
request_data = {}
logging_obj = type(
"LoggingObj",
(),
{
"model_call_details": {
"litellm_params": {
"metadata": {
"headers": {
"shade_scan_id": "scan-from-logging",
"authorization": "Bearer secret",
}
}
}
}
},
)()
payload = grayswan_guardrail._prepare_payload(messages, {}, request_data, logging_obj)
assert payload["litellm_metadata"] == {
"headers": {"shade_scan_id": "scan-from-logging"},
}
def test_prepare_payload_ignores_logging_obj_without_model_call_details(
grayswan_guardrail: GraySwanGuardrail,
) -> None:
messages = [{"role": "user", "content": "hello"}]
payload = grayswan_guardrail._prepare_payload(messages, {}, {}, object())
assert "litellm_metadata" not in payload
def test_process_response_does_not_block_under_threshold(
grayswan_guardrail: GraySwanGuardrail,
) -> None:
grayswan_guardrail._process_grayswan_response(
{"violation": 0.3, "violated_rules": []}
)
grayswan_guardrail._process_grayswan_response({"violation": 0.3, "violated_rules": []})
def test_process_response_blocks_when_threshold_exceeded() -> None:
@ -127,16 +234,12 @@ class _DummyClient:
self.calls: list[dict] = []
async def post(self, *, url: str, headers: dict, json: dict, timeout: float):
self.calls.append(
{"url": url, "headers": headers, "json": json, "timeout": timeout}
)
self.calls.append({"url": url, "headers": headers, "json": json, "timeout": timeout})
return _DummyResponse(self.payload)
@pytest.mark.asyncio
async def test_run_guardrail_posts_payload(
monkeypatch, grayswan_guardrail: GraySwanGuardrail
) -> None:
async def test_run_guardrail_posts_payload(monkeypatch, grayswan_guardrail: GraySwanGuardrail) -> None:
dummy_client = _DummyClient({"violation": 0.1})
grayswan_guardrail.async_handler = dummy_client
@ -308,9 +411,7 @@ def test_process_response_passthrough_raises_exception_in_pre_call() -> None:
# Should raise ModifyResponseException
with pytest.raises(ModifyResponseException) as exc:
guardrail._process_grayswan_response(
response_json, data, GuardrailEventHooks.pre_call
)
guardrail._process_grayswan_response(response_json, data, GuardrailEventHooks.pre_call)
assert "Gray Swan Cygnal Guardrail" in exc.value.message
assert exc.value.model == "gpt-4"
@ -338,9 +439,7 @@ def test_process_response_passthrough_raises_exception_in_during_call() -> None:
# Should raise ModifyResponseException
with pytest.raises(ModifyResponseException) as exc:
guardrail._process_grayswan_response(
response_json, data, GuardrailEventHooks.during_call
)
guardrail._process_grayswan_response(response_json, data, GuardrailEventHooks.during_call)
assert "Gray Swan Cygnal Guardrail" in exc.value.message
assert exc.value.model == "gpt-4"
@ -365,9 +464,7 @@ def test_process_response_passthrough_stores_detection_info_in_post_call() -> No
}
# Should NOT raise an exception in post_call
guardrail._process_grayswan_response(
response_json, data, GuardrailEventHooks.post_call
)
guardrail._process_grayswan_response(response_json, data, GuardrailEventHooks.post_call)
# Verify detection info was stored in metadata
assert "metadata" in data
@ -400,9 +497,7 @@ def test_process_response_passthrough_does_not_raise_if_under_threshold() -> Non
}
# Should not raise an exception since under threshold
guardrail._process_grayswan_response(
response_json, data, GuardrailEventHooks.pre_call
)
guardrail._process_grayswan_response(response_json, data, GuardrailEventHooks.pre_call)
# Should not have any detection info since it didn't exceed threshold
assert "guardrail_detections" not in data.get("metadata", {})
@ -436,10 +531,7 @@ def test_format_violation_message() -> None:
assert "Gray Swan Cygnal Guardrail" in message
assert "the input query has a violation score of 0.85" in message
assert "violating the rule(s): 1, 3, 5" in message
assert (
"Mutation effort to make the harmful intention disguised was DETECTED"
in message
)
assert "Mutation effort to make the harmful intention disguised was DETECTED" in message
# IPI should not be in message since it's False
assert "Indirect Prompt Injection was DETECTED" not in message
@ -450,10 +542,7 @@ def test_format_violation_message() -> None:
assert "Gray Swan Cygnal Guardrail" in message
assert "the model response has a violation score of 0.85" in message
assert "violating the rule(s): 1, 3, 5" in message
assert (
"Mutation effort to make the harmful intention disguised was DETECTED"
in message
)
assert "Mutation effort to make the harmful intention disguised was DETECTED" in message
def test_prepare_payload_includes_litellm_metadata(