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
756a647d84
4 changed files with 142 additions and 11 deletions
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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