From c36d929dddb886d2f3c3b24c0335688d7343e1bf Mon Sep 17 00:00:00 2001 From: Yaniv Israel Date: Sun, 24 May 2026 21:33:54 +0300 Subject: [PATCH] fix(deepkeep): address greptile review comments - extra_headers: fix type annotation (list -> Dict[str, str]) and actually merge them into _build_request_headers() so user-configured headers reach the DeepKeep API - user_api_key_hash: only fall back to user_api_key_token when no explicit hash is already set, avoiding silent overwrite - apply_guardrail: preserve tool_calls and structured_messages in the return value so downstream callers don't lose that content Adds tests for all four fixes. --- .../guardrail_hooks/deepkeep/deepkeep.py | 14 ++- .../guardrail_hooks/test_deepkeep.py | 112 ++++++++++++++++++ 2 files changed, 122 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index 3d0d34f0439..4d55e31c7f6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py @@ -72,7 +72,7 @@ class DeepKeepGuardrail(CustomGuardrail): api_base: Optional[str] = None, firewall_id: Optional[str] = None, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", - extra_headers: Optional[list] = None, + extra_headers: Optional[Dict[str, str]] = None, **kwargs, ): self.async_handler = get_async_httpx_client( @@ -114,7 +114,7 @@ class DeepKeepGuardrail(CustomGuardrail): self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( unreachable_fallback ) - self.extra_headers = extra_headers or [] + self.extra_headers: Dict[str, str] = extra_headers or {} # Set supported event hooks if "supported_event_hooks" not in kwargs: @@ -168,8 +168,8 @@ class DeepKeepGuardrail(CustomGuardrail): if value is not None: result_metadata[key] = value - # Handle the token → hash alias - if metadata_dict.get("user_api_key_token") is not None: + # Handle the token → hash alias (only when no explicit hash was provided) + if metadata_dict.get("user_api_key_token") is not None and "user_api_key_hash" not in result_metadata: result_metadata["user_api_key_hash"] = metadata_dict["user_api_key_token"] return result_metadata @@ -180,6 +180,8 @@ class DeepKeepGuardrail(CustomGuardrail): "Content-Type": "application/json", "X-API-Key": self.deepkeep_api_key, } + if self.extra_headers: + headers.update(self.extra_headers) return headers def _fail_open_passthrough( @@ -339,6 +341,10 @@ class DeepKeepGuardrail(CustomGuardrail): 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 except GuardrailRaisedException: 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 d876df0e3a1..e01a0ccbdaf 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_deepkeep.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_deepkeep.py @@ -417,6 +417,118 @@ class TestDeepKeepGuardrail: assert config_model is not None assert config_model.ui_friendly_name() == "DeepKeep AI Firewall" + def test_build_request_headers_includes_extra_headers(self): + """should merge extra_headers into the request headers.""" + guardrail = DeepKeepGuardrail( + api_key="test-api-key-123", + api_base="https://test.deepkeep.ai", + firewall_id="fw-123", + extra_headers={"X-Custom-Header": "custom-value", "X-Tenant": "tenant-1"}, + guardrail_name="test", + event_hook="pre_call", + ) + + headers = guardrail._build_request_headers() + assert headers["X-API-Key"] == "test-api-key-123" + assert headers["Content-Type"] == "application/json" + assert headers["X-Custom-Header"] == "custom-value" + assert headers["X-Tenant"] == "tenant-1" + + def test_build_request_headers_no_extra_headers(self): + """should not fail and return only base headers when extra_headers is None.""" + guardrail = DeepKeepGuardrail( + api_key="test-api-key-123", + api_base="https://test.deepkeep.ai", + firewall_id="fw-123", + guardrail_name="test", + event_hook="pre_call", + ) + + headers = guardrail._build_request_headers() + assert set(headers.keys()) == {"Content-Type", "X-API-Key"} + + def test_extract_user_api_key_metadata_token_does_not_overwrite_hash(self): + """should not overwrite user_api_key_hash with user_api_key_token when hash is already set.""" + guardrail = DeepKeepGuardrail( + api_key="test-key", + api_base="https://test.deepkeep.ai", + firewall_id="fw-123", + guardrail_name="test", + event_hook="pre_call", + ) + + request_data = { + "metadata": { + "user_api_key_hash": "the-real-hash", + "user_api_key_token": "the-raw-token", + } + } + + metadata = guardrail._extract_user_api_key_metadata(request_data) + # hash was set explicitly, token alias must NOT overwrite it + assert metadata["user_api_key_hash"] == "the-real-hash" + + def test_extract_user_api_key_metadata_token_used_as_hash_fallback(self): + """should use user_api_key_token as hash alias only when no explicit hash is present.""" + guardrail = DeepKeepGuardrail( + api_key="test-key", + api_base="https://test.deepkeep.ai", + firewall_id="fw-123", + guardrail_name="test", + event_hook="pre_call", + ) + + request_data = { + "metadata": { + "user_api_key_token": "the-raw-token", + } + } + + metadata = guardrail._extract_user_api_key_metadata(request_data) + assert metadata["user_api_key_hash"] == "the-raw-token" + + @pytest.mark.asyncio + async def test_apply_guardrail_preserves_tool_calls_and_structured_messages(self): + """should include tool_calls and structured_messages in the return value.""" + guardrail = DeepKeepGuardrail( + api_key="test-key", + api_base="https://test.deepkeep.ai", + firewall_id="fw-123", + guardrail_name="test", + event_hook="pre_call", + ) + + mock_response = Response( + status_code=200, + json={"action": "NONE", "blocked_reason": None, "texts": None, "images": None}, + request=Request( + "POST", + "https://test.deepkeep.ai/v3/openai/beta/litellm_basic_guardrail_api", + ), + ) + + sample_tool_calls = [{"id": "call_1", "type": "function", "function": {"name": "get_weather"}}] + sample_structured = [{"role": "tool", "content": "sunny"}] + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await guardrail.apply_guardrail( + inputs={ + "texts": ["what's the weather?"], + "tool_calls": sample_tool_calls, + "structured_messages": sample_structured, + }, + request_data={"metadata": {}}, + input_type="request", + ) + + assert result["tool_calls"] == sample_tool_calls + assert result["structured_messages"] == sample_structured + @pytest.mark.asyncio async def test_firewall_id_in_payload(self): """should include firewall_id in additional_provider_specific_params."""