mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
5973d9fd2b
commit
528fa380f5
2 changed files with 159 additions and 36 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue