diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index 441bc789878..256ac753aa8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py @@ -238,6 +238,36 @@ class DeepKeepGuardrail(CustomGuardrail): verbose_proxy_logger.error("DeepKeep guardrail API error: %s", str(error)) raise DeepKeepGuardrailAPIError(f"DeepKeep guardrail API failed: {str(error)}") + @staticmethod + def _build_return_inputs( + *, + response_json: Dict[str, Any], + texts: list, + images: Optional[Any], + tools: Optional[Any], + tool_calls: Optional[Any], + structured_messages: Optional[Any], + ) -> GenericGuardrailAPIInputs: + """Merge original inputs with any guardrail-modified values from the API response.""" + return_inputs = GenericGuardrailAPIInputs(texts=texts) + if response_json.get("texts"): + return_inputs["texts"] = response_json["texts"] + if response_json.get("images"): + return_inputs["images"] = response_json["images"] + elif images: + return_inputs["images"] = images + if response_json.get("tools"): + return_inputs["tools"] = response_json["tools"] + elif tools: + return_inputs["tools"] = tools + if response_json.get("tool_calls"): + return_inputs["tool_calls"] = response_json["tool_calls"] + elif tool_calls: + return_inputs["tool_calls"] = tool_calls + if structured_messages: + return_inputs["structured_messages"] = structured_messages + return return_inputs + @log_guardrail_information async def apply_guardrail( self, @@ -336,21 +366,14 @@ class DeepKeepGuardrail(CustomGuardrail): should_wrap_with_default_message=False, ) - # Build return inputs – apply any modifications from GUARDRAIL_INTERVENED - return_inputs = GenericGuardrailAPIInputs(texts=texts) - if response_json.get("texts"): - return_inputs["texts"] = response_json["texts"] - if response_json.get("images"): - return_inputs["images"] = response_json["images"] - elif images: - return_inputs["images"] = images - if tools: - return_inputs["tools"] = tools - if tool_calls: - return_inputs["tool_calls"] = tool_calls - if structured_messages: - return_inputs["structured_messages"] = structured_messages - return return_inputs + return self._build_return_inputs( + response_json=response_json, + texts=texts, + images=images, + tools=tools, + tool_calls=tool_calls, + structured_messages=structured_messages, + ) except GuardrailRaisedException: raise diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_deepkeep.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_deepkeep.py index e01a0ccbdaf..0e8cdab6784 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_deepkeep.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_deepkeep.py @@ -529,6 +529,61 @@ class TestDeepKeepGuardrail: assert result["tool_calls"] == sample_tool_calls assert result["structured_messages"] == sample_structured + @pytest.mark.asyncio + async def test_apply_guardrail_applies_tool_redactions_from_response(self): + """should use redacted tools/tool_calls from response when GUARDRAIL_INTERVENED returns them.""" + guardrail = DeepKeepGuardrail( + api_key="test-key", + api_base="https://test.deepkeep.ai", + firewall_id="fw-123", + guardrail_name="test", + event_hook="pre_call", + ) + + redacted_tools = [{"type": "function", "function": {"name": "get_data", "description": "[REDACTED]"}}] + redacted_tool_calls = [{"id": "call_1", "type": "function", "function": {"name": "get_data", "arguments": "{}"}}] + + mock_response = Response( + status_code=200, + json={ + "action": "GUARDRAIL_INTERVENED", + "blocked_reason": None, + "texts": None, + "images": None, + "tools": redacted_tools, + "tool_calls": redacted_tool_calls, + }, + request=Request( + "POST", + "https://test.deepkeep.ai/v3/openai/beta/litellm_basic_guardrail_api", + ), + ) + + original_tools = [{"type": "function", "function": {"name": "get_data", "description": "sensitive info"}}] + original_tool_calls = [{"id": "call_1", "type": "function", "function": {"name": "get_data", "arguments": '{"secret": "value"}'}}] + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await guardrail.apply_guardrail( + inputs={ + "texts": ["run the tool"], + "tools": original_tools, + "tool_calls": original_tool_calls, + }, + request_data={"metadata": {}}, + input_type="request", + ) + + # Redacted versions from the API response must be used, not the originals + assert result["tools"] == redacted_tools + assert result["tool_calls"] == redacted_tool_calls + assert result["tools"] != original_tools + assert result["tool_calls"] != original_tool_calls + @pytest.mark.asyncio async def test_firewall_id_in_payload(self): """should include firewall_id in additional_provider_specific_params."""