From 528fa380f5a271865af9f85228148131063bdf2b Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Wed, 8 Jul 2026 15:05:27 -0700 Subject: [PATCH] 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> --- .../guardrail_hooks/grayswan/grayswan.py | 48 +++++- .../guardrail_hooks/test_grayswan.py | 147 ++++++++++++++---- 2 files changed, 159 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py index a14d2fc8608..9805b1a9117 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py +++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py index f2e7447239f..53af7f36a5f 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -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(