diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py index e47dea8c0ec..6a3ce00336f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py +++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py @@ -27,7 +27,6 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj GRAYSWAN_BLOCK_ERROR_MSG: Final = "Blocked by Gray Swan Guardrail" -GRAYSWAN_CONVERSATION_CACHE_KEY: Final = "_grayswan_request_conversation" class MonitorTurn(TypedDict): @@ -159,8 +158,9 @@ class GraySwanGuardrail(CustomGuardrail): Args: inputs: Dictionary containing: - texts: List of extracted texts (fallback scan content) - - structured_messages: Normalized, scoped conversation (request scans) - - tools: Scoped tool definitions (request scans) + - structured_messages: Normalized, scoped conversation (response + scans carry the request conversation plus the response turns) + - tools: Scoped tool definitions - tool_calls: Tool calls emitted by the model (response scans) - images: Optional list of images (not currently used by GraySwan) request_data: The original request data @@ -176,7 +176,7 @@ class GraySwanGuardrail(CustomGuardrail): """ start_time: Final = time.time() try: - messages, tools = self._build_monitor_input(inputs, request_data, input_type) + messages, tools = self._build_monitor_input(inputs, input_type) if not messages: verbose_proxy_logger.debug("Gray Swan Guardrail: No content to scan") return inputs @@ -505,53 +505,24 @@ class GraySwanGuardrail(CustomGuardrail): def _build_monitor_input( self, inputs: GenericGuardrailAPIInputs, - request_data: dict, # mutable-ok: the shared per-request state dict every hook receives; the request scan caches on it input_type: Literal["request", "response"], ) -> tuple[tuple[Mapping[str, Any], ...], tuple[Mapping[str, Any], ...] | None]: """Build the monitor conversation from the translation layer's scoped view. - Request scans send `structured_messages` and `tools` exactly as the unified - guardrail system produced them (normalized per API surface, operator scoping - flags applied) and cache them on the request. Response scans replay the - cached conversation and append the response turns. Without a structured - view, the pre-existing texts-only wrapping is kept. + Both scan types send `structured_messages` and `tools` exactly as the + unified guardrail system produced them (normalized per API surface, + operator scoping flags applied; response scans carry the request + conversation with the response turns appended). Without a structured + view, request scans keep the pre-existing texts-only wrapping and + response scans send context-free assistant turns (the SDK/direct-call + path, where no request conversation exists). """ + conversation: Final = self._sanitize_json_list(inputs.get("structured_messages")) + if conversation: + return conversation, self._sanitize_json_list(inputs.get("tools")) if input_type == "request": - conversation: Final = self._sanitize_json_list(inputs.get("structured_messages")) - if not conversation: - return self._texts_fallback(inputs, "user"), None - tools: Final = self._sanitize_json_list(inputs.get("tools")) - self._cache_request_conversation(request_data, conversation, tools) - return conversation, tools - response_turns: Final = self._build_response_turns(inputs) - if not response_turns: - return (), None - cached: Final = self._cached_request_conversation(request_data) - if cached is None: - return self._texts_fallback(inputs, "assistant"), None - return (*cached[0], *response_turns), cached[1] - - def _cache_request_conversation( - self, - request_data: dict, # mutable-ok: the shared per-request state dict every hook receives; caching on it is the point - conversation: tuple[Mapping[str, Any], ...], - tools: tuple[Mapping[str, Any], ...] | None, - ) -> None: - metadata: Final = request_data.setdefault( - "metadata", - {}, # mutable-ok: request metadata is shared mutable state other hooks also write to - ) # rebind-ok: response scans replay the request-time scoped conversation and request_data is the only object shared across hooks - if isinstance(metadata, dict): - metadata[GRAYSWAN_CONVERSATION_CACHE_KEY] = (conversation, tools) - - def _cached_request_conversation( - self, request_data: Mapping[str, Any] - ) -> tuple[tuple[Mapping[str, Any], ...], tuple[Mapping[str, Any], ...] | None] | None: - metadata: Final = request_data.get("metadata") - cached: Final = metadata.get(GRAYSWAN_CONVERSATION_CACHE_KEY) if isinstance(metadata, dict) else None - if isinstance(cached, tuple) and len(cached) == 2: - return cached - return None + return self._texts_fallback(inputs, "user"), None + return self._build_response_turns(inputs), None def _build_response_turns(self, inputs: GenericGuardrailAPIInputs) -> tuple[Mapping[str, Any], ...]: tool_calls: Final = self._sanitize_json_list(inputs.get("tool_calls")) 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 e2243d44993..d807dee7712 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -666,9 +666,11 @@ async def test_apply_guardrail_scans_texts_not_raw_messages(monkeypatch, grayswa @pytest.mark.asyncio -async def test_apply_guardrail_response_appends_assistant_turn_with_tool_calls( +async def test_apply_guardrail_response_uses_structured_conversation_with_tools( monkeypatch, grayswan_guardrail: GraySwanGuardrail ) -> None: + """The translation layer hands response scans the request conversation with + the response turns already appended; the payload sends it verbatim.""" captured: dict = {} async def _fake_call(payload: dict): @@ -677,34 +679,65 @@ async def test_apply_guardrail_response_appends_assistant_turn_with_tool_calls( monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) - structured_messages = [{"role": "user", "content": "write a script"}] - response_tool_calls = [ + structured_messages = [ + {"role": "user", "content": "write a script"}, { - "id": "call_9", - "type": "function", - "function": {"name": "write_file", "arguments": '{"path": "x.sh", "content": "xmrig"}'}, - } + "role": "assistant", + "content": "sure, writing it now", + "tool_calls": [ + { + "id": "call_9", + "type": "function", + "function": {"name": "write_file", "arguments": '{"path": "x.sh", "content": "xmrig"}'}, + } + ], + }, ] - request_data: dict = {"model": "gpt-4"} + tools = [{"type": "function", "function": {"name": "write_file", "parameters": {}}}] await grayswan_guardrail.apply_guardrail( - inputs={"texts": ["write a script"], "structured_messages": structured_messages}, - request_data=request_data, - input_type="request", - ) - await grayswan_guardrail.apply_guardrail( - inputs={"texts": ["sure, writing it now"], "tool_calls": response_tool_calls}, - request_data=request_data, + inputs={ + "texts": ["sure, writing it now"], + "structured_messages": structured_messages, + "tools": tools, + }, + request_data={"model": "gpt-4"}, input_type="response", ) wire = _wire(captured["payload"]) - assert wire["messages"][:-1] == structured_messages - assert wire["messages"][-1] == { - "role": "assistant", - "content": "sure, writing it now", - "tool_calls": response_tool_calls, - } + assert wire["messages"] == structured_messages + assert wire["tools"] == tools + + +@pytest.mark.asyncio +async def test_apply_guardrail_response_does_not_read_request_metadata( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + """Request scans no longer stash the conversation on request_data, and + response scans build the payload from inputs alone.""" + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) + + request_data: dict = {"model": "gpt-4"} + await grayswan_guardrail.apply_guardrail( + inputs={"texts": ["go"], "structured_messages": [{"role": "user", "content": "go"}]}, + request_data=request_data, + input_type="request", + ) + assert "_grayswan_request_conversation" not in (request_data.get("metadata") or {}) + + await grayswan_guardrail.apply_guardrail( + inputs={"texts": ["done"]}, + request_data=request_data, + input_type="response", + ) + assert _wire(captured["payload"])["messages"] == [{"role": "assistant", "content": "done"}] @pytest.mark.asyncio @@ -720,23 +753,15 @@ async def test_apply_guardrail_scans_tool_call_only_response( monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) response_tool_calls = [{"id": "call_2", "type": "function", "function": {"name": "run", "arguments": "{}"}}] - request_data: dict = {"model": "gpt-4"} - await grayswan_guardrail.apply_guardrail( - inputs={"texts": ["go"], "structured_messages": [{"role": "user", "content": "go"}]}, - request_data=request_data, - input_type="request", - ) await grayswan_guardrail.apply_guardrail( inputs={"texts": [], "tool_calls": response_tool_calls}, - request_data=request_data, + request_data={"model": "gpt-4"}, input_type="response", ) wire = _wire(captured["payload"]) - assert wire["messages"][0] == {"role": "user", "content": "go"} - assert wire["messages"][-1]["tool_calls"] == response_tool_calls - assert wire["messages"][-1]["content"] == "" + assert wire["messages"] == [{"role": "assistant", "content": "", "tool_calls": response_tool_calls}] @pytest.mark.asyncio @@ -795,20 +820,14 @@ async def test_apply_guardrail_empty_response_is_not_scanned( monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) - request_data: dict = {"model": "gpt-4"} - await grayswan_guardrail.apply_guardrail( - inputs={"texts": ["go"], "structured_messages": [{"role": "user", "content": "go"}]}, - request_data=request_data, - input_type="request", - ) inputs: dict = {"texts": []} result = await grayswan_guardrail.apply_guardrail( inputs=inputs, - request_data=request_data, + request_data={"model": "gpt-4"}, input_type="response", ) - assert len(calls) == 1 + assert len(calls) == 0 assert result is inputs @@ -824,19 +843,13 @@ async def test_apply_guardrail_multiple_response_texts_get_separate_turns( monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) - request_data: dict = {"model": "gpt-4"} - await grayswan_guardrail.apply_guardrail( - inputs={"texts": ["pick one"], "structured_messages": [{"role": "user", "content": "pick one"}]}, - request_data=request_data, - input_type="request", - ) await grayswan_guardrail.apply_guardrail( inputs={"texts": ["candidate one", "candidate two"]}, - request_data=request_data, + request_data={"model": "gpt-4"}, input_type="response", ) - assert _wire(captured["payload"])["messages"][-2:] == [ + assert _wire(captured["payload"])["messages"] == [ {"role": "assistant", "content": "candidate one"}, {"role": "assistant", "content": "candidate two"}, ]