From 756a647d84d462735d6a4f04f45fa589db053240 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. --- .../send_emails/sendgrid_email.py | 1 - .../proxy/common_utils/check_batch_cost.py | 23 +++- .../guardrail_hooks/deepkeep/deepkeep.py | 17 ++- .../guardrail_hooks/test_deepkeep.py | 112 ++++++++++++++++++ 4 files changed, 142 insertions(+), 11 deletions(-) diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py index 2dc158a3cfb..a1e8def2bb2 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py @@ -15,7 +15,6 @@ from litellm.llms.custom_httpx.http_handler import ( from .base_email import BaseEmailLogger - SENDGRID_API_ENDPOINT = "https://api.sendgrid.com/v3/mail/send" diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 68531069a95..705a682f298 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -328,21 +328,32 @@ class CheckBatchCost: # CheckBatchCost bypasses async_post_call_success_hook, so convert raw # output/error file IDs to managed base64 IDs before the DB write here. - managed_files_hook = self.proxy_logging_obj.get_proxy_hook("managed_files") + managed_files_hook = self.proxy_logging_obj.get_proxy_hook( + "managed_files" + ) if managed_files_hook is not None: from litellm.proxy._types import UserAPIKeyAuth + _minimal_auth = UserAPIKeyAuth( user_id=job.created_by or "default-user-id", team_id=getattr(job, "team_id", None), ) for _file_attr in ["output_file_id", "error_file_id"]: _raw_file_id = getattr(response, _file_attr, None) - if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id): + if _raw_file_id and not _is_base64_encoded_unified_file_id( + _raw_file_id + ): try: - _unified_file_id = managed_files_hook.get_unified_output_file_id( - output_file_id=_raw_file_id, - model_id=model_id, - model_name=str(model_name) if model_name else deployment_info.model_name or None, + _unified_file_id = ( + managed_files_hook.get_unified_output_file_id( + output_file_id=_raw_file_id, + model_id=model_id, + model_name=( + str(model_name) + if model_name + else deployment_info.model_name or None + ), + ) ) await managed_files_hook.store_unified_file_id( file_id=_unified_file_id, diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index 3d0d34f0439..63dedcf9ac4 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,11 @@ 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 +183,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 +344,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."""