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.
This commit is contained in:
Yaniv Israel 2026-05-24 21:33:54 +03:00
parent e4520bf458
commit c36d929ddd
2 changed files with 122 additions and 4 deletions

View file

@ -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:

View file

@ -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."""