mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
e4520bf458
commit
c36d929ddd
2 changed files with 122 additions and 4 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue