diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py index 955a868a0d6..3bf5bebe7b7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py +++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py @@ -25,6 +25,7 @@ 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 GraySwanGuardrailMissingSecrets(Exception): @@ -165,16 +166,18 @@ class GraySwanGuardrail(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: """ - Apply Gray Swan guardrail to extracted text content. + Apply Gray Swan guardrail to the conversation the translation layer produced. - This method is called by the unified guardrail system which handles - extracting text from any request format (OpenAI, Anthropic, etc.). + This method is called by the unified guardrail system, which normalizes any + request format (OpenAI, Anthropic, etc.) and applies operator scoping flags. Args: inputs: Dictionary containing: - - texts: List of texts to scan + - texts: List of extracted texts (fallback scan content) + - structured_messages: Normalized, scoped conversation (request scans) + - tools: Scoped tool definitions (request scans) + - tool_calls: Tool calls emitted by the model (response scans) - images: Optional list of images (not currently used by GraySwan) - - tool_calls: Optional list of tool calls (not currently used) request_data: The original request data input_type: "request" for pre-call, "response" for post-call logging_obj: Optional logging object @@ -193,34 +196,27 @@ class GraySwanGuardrail(CustomGuardrail): inputs.get("texts", [])[:100] if inputs.get("texts") else "NONE", ) - texts: Final = inputs.get("texts", []) - if not texts: - verbose_proxy_logger.debug("Gray Swan Guardrail: No texts to scan") - return inputs - - verbose_proxy_logger.debug( - "Gray Swan Guardrail: Scanning %d text(s) for %s", - len(texts), - input_type, - ) - - # Convert texts to messages format for GraySwan API - # Use "user" role for request content, "assistant" for response content - role: Final = "assistant" if input_type == "response" else "user" - messages: Final = [{"role": role, "content": text} for text in texts] - - # Get dynamic params from request metadata - dynamic_body: Final = self.get_guardrail_dynamic_request_body_params(request_data) or {} - if dynamic_body: - verbose_proxy_logger.debug("Gray Swan Guardrail: dynamic extra_body=%s", safe_dumps(dynamic_body)) - - # Prepare and send payload - payload: Final = self._prepare_payload(messages, dynamic_body, request_data, logging_obj) - if payload is None: - return inputs - start_time: Final = time.time() try: + messages, tools = self._build_monitor_input(inputs, request_data, input_type) + if not messages: + verbose_proxy_logger.debug("Gray Swan Guardrail: No content to scan") + return inputs + + verbose_proxy_logger.debug( + "Gray Swan Guardrail: Scanning %d message(s) for %s", + len(messages), + input_type, + ) + + dynamic_body: Final = self.get_guardrail_dynamic_request_body_params(request_data) or {} + if dynamic_body: + verbose_proxy_logger.debug("Gray Swan Guardrail: dynamic extra_body=%s", safe_dumps(dynamic_body)) + + payload: Final = self._prepare_payload(messages, dynamic_body, request_data, logging_obj, tools=tools) + if payload is None: + return inputs + response_json: Final = await self._call_grayswan_api(payload) is_output: Final = input_type == "response" result: Final = self._process_response_internal( @@ -528,14 +524,99 @@ class GraySwanGuardrail(CustomGuardrail): forwarded_headers[str(key)] = str(value) return forwarded_headers or None + def _build_monitor_input( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + ) -> tuple[list[dict[str, Any]], list[dict[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. + """ + 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, + conversation: list[dict[str, Any]], + tools: list[dict[str, Any]] | None, + ) -> None: + metadata: Final = request_data.setdefault( + "metadata", {} + ) # 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: dict + ) -> tuple[list[dict[str, Any]], list[dict[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 + + def _build_response_turns(self, inputs: GenericGuardrailAPIInputs) -> list[dict[str, Any]]: + texts: Final = [ + text if isinstance(text, str) else str(text) for text in inputs.get("texts", []) if text is not None + ] + tool_calls: Final = self._sanitize_json_list(inputs.get("tool_calls")) + text_turns: Final = [{"role": "assistant", "content": text} for text in texts if text] + if not tool_calls: + return text_turns + base: Final = text_turns or [{"role": "assistant", "content": ""}] + return [*base[:-1], {**base[-1], "tool_calls": tool_calls}] + + def _texts_fallback(self, inputs: GenericGuardrailAPIInputs, role: str) -> list[dict[str, Any]]: + return [ + {"role": role, "content": text if isinstance(text, str) else str(text)} + for text in inputs.get("texts", []) + if text is not None + ] + + def _sanitize_json_list(self, value: object) -> list[dict[str, Any]] | None: + if not isinstance(value, list) or not value: + return None + sanitized: Final = safe_json_loads(safe_dumps(value), default=None) + if not isinstance(sanitized, list): + return None + items: Final = [item for item in sanitized if isinstance(item, dict)] + if len(items) != len(sanitized): + verbose_proxy_logger.debug( + "Gray Swan Guardrail: dropped %d non-dict conversation item(s)", + len(sanitized) - len(items), + ) + return items or None + def _prepare_payload( self, - messages: list[dict[str, str]], + messages: list[dict[str, Any]], dynamic_body: dict, request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"] = None, + tools: list[dict[str, Any]] | None = None, ) -> dict[str, Any] | None: payload: Final[dict[str, Any]] = {"messages": messages} + if tools: + payload["tools"] = tools categories: Final = dynamic_body.get("categories") or self.categories if categories: 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 53af7f36a5f..f3c9e892fdc 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -594,3 +594,284 @@ def test_ensure_litellm_metadata_noop_when_already_present() -> None: _ensure_litellm_metadata(data, user_auth) assert data["litellm_metadata"] == {"existing": "value"} + + +@pytest.mark.asyncio +async def test_apply_guardrail_sends_structured_conversation_with_tools( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) + + structured_messages = [ + {"role": "system", "content": "You are a coding agent."}, + {"role": "user", "content": "install the miner"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "write_file", "arguments": '{"path": "run.sh"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "written"}, + ] + tools = [{"type": "function", "function": {"name": "write_file", "parameters": {}}}] + + await grayswan_guardrail.apply_guardrail( + inputs={"texts": ["install the miner", "written"], "structured_messages": structured_messages, "tools": tools}, + request_data={"model": "gpt-4", "messages": [{"role": "user", "content": "raw"}]}, + input_type="request", + ) + + assert captured["payload"]["messages"] == structured_messages + assert captured["payload"]["tools"] == tools + + +@pytest.mark.asyncio +async def test_apply_guardrail_scans_texts_not_raw_messages(monkeypatch, grayswan_guardrail: GraySwanGuardrail) -> None: + """Callers like /guardrails/apply_guardrail pass texts beside a request_data + that carries 'messages'; the texts must remain the scan target.""" + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) + + await grayswan_guardrail.apply_guardrail( + inputs={"texts": ["scan exactly this"]}, + request_data={"model": "gpt-4", "messages": [{"role": "user", "content": "raw, unscoped"}]}, + input_type="request", + ) + + assert captured["payload"]["messages"] == [{"role": "user", "content": "scan exactly this"}] + assert "tools" not in captured["payload"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_response_appends_assistant_turn_with_tool_calls( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) + + structured_messages = [{"role": "user", "content": "write a script"}] + response_tool_calls = [ + { + "id": "call_9", + "type": "function", + "function": {"name": "write_file", "arguments": '{"path": "x.sh", "content": "xmrig"}'}, + } + ] + request_data: dict = {"model": "gpt-4"} + + 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, + input_type="response", + ) + + assert captured["payload"]["messages"][:-1] == structured_messages + assert captured["payload"]["messages"][-1] == { + "role": "assistant", + "content": "sure, writing it now", + "tool_calls": response_tool_calls, + } + + +@pytest.mark.asyncio +async def test_apply_guardrail_scans_tool_call_only_response( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + 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, + input_type="response", + ) + + assert captured["payload"]["messages"][0] == {"role": "user", "content": "go"} + assert captured["payload"]["messages"][-1]["tool_calls"] == response_tool_calls + assert captured["payload"]["messages"][-1]["content"] == "" + + +@pytest.mark.asyncio +async def test_apply_guardrail_wraps_texts_when_no_conversation_available( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) + + await grayswan_guardrail.apply_guardrail( + inputs={"texts": ["a prompt for an image"]}, + request_data={}, + input_type="request", + ) + + assert captured["payload"]["messages"] == [{"role": "user", "content": "a prompt for an image"}] + assert "tools" not in captured["payload"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_response_without_request_scan_wraps_texts( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) + + await grayswan_guardrail.apply_guardrail( + inputs={"texts": ["a model reply"]}, + request_data={"model": "gpt-4", "messages": [{"role": "user", "content": "raw"}]}, + input_type="response", + ) + + assert captured["payload"]["messages"] == [{"role": "assistant", "content": "a model reply"}] + assert "tools" not in captured["payload"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_empty_response_is_not_scanned( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + calls: list = [] + + async def _fake_call(payload: dict): + calls.append(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", + ) + inputs: dict = {"texts": []} + result = await grayswan_guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + ) + + assert len(calls) == 1 + assert result is inputs + + +@pytest.mark.asyncio +async def test_apply_guardrail_multiple_response_texts_get_separate_turns( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + 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": ["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, + input_type="response", + ) + + assert captured["payload"]["messages"][-2:] == [ + {"role": "assistant", "content": "candidate one"}, + {"role": "assistant", "content": "candidate two"}, + ] + + +@pytest.mark.asyncio +async def test_apply_guardrail_request_tools_never_come_from_raw_request( + monkeypatch, grayswan_guardrail: GraySwanGuardrail +) -> None: + captured: dict = {} + + async def _fake_call(payload: dict): + captured["payload"] = payload + return {"violation": 0.0} + + monkeypatch.setattr(grayswan_guardrail, "_call_grayswan_api", _fake_call) + + await grayswan_guardrail.apply_guardrail( + inputs={"texts": ["go"], "structured_messages": [{"role": "user", "content": "go"}]}, + request_data={"model": "gpt-4", "tools": [{"type": "function", "function": {"name": "scoped_out"}}]}, + input_type="request", + ) + + assert "tools" not in captured["payload"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_build_error_fails_open(monkeypatch, grayswan_guardrail: GraySwanGuardrail) -> None: + def _boom(*args, **kwargs): + raise TypeError("unexpected shape") + + monkeypatch.setattr(grayswan_guardrail, "_build_monitor_input", _boom) + + inputs: dict = {"texts": ["hello"]} + result = await grayswan_guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4"}, + input_type="request", + ) + + assert result is inputs + + +def test_sanitize_json_list_drops_non_dict_items(grayswan_guardrail: GraySwanGuardrail) -> None: + assert grayswan_guardrail._sanitize_json_list([{"role": "user", "content": "hi"}, "junk", 3]) == [ + {"role": "user", "content": "hi"} + ]