Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py: 23%
170 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1# +-------------------------------------------------------------+
2#
3# Use DeepKeep AI Firewall for your LLM calls
4# https://www.deepkeep.ai/
5#
6# +-------------------------------------------------------------+
8import os
9from collections.abc import Mapping
10from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol
12import httpx
13from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack
15from litellm._logging import verbose_proxy_logger
16from litellm._version import version as litellm_version
17from litellm.exceptions import GuardrailRaisedException, Timeout
18from litellm.integrations.custom_guardrail import (
19 CustomGuardrail,
20 log_guardrail_information,
21)
22from litellm.llms.custom_httpx.http_handler import (
23 get_async_httpx_client,
24 httpxSpecialProvider,
25)
26from litellm.types.guardrails import GuardrailEventHooks
27from litellm.types.utils import GenericGuardrailAPIInputs
29if TYPE_CHECKING: 29 ↛ 30line 29 didn't jump to line 30 because the condition on line 29 was never true
30 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
31 from litellm.types.llms.openai import (
32 AllMessageValues,
33 ChatCompletionToolCallChunk,
34 ChatCompletionToolParam,
35 )
36 from litellm.types.utils import ChatCompletionMessageToolCall
38GUARDRAIL_NAME: Final = "deepkeep"
40# Default DeepKeep API endpoint path
41_DEEPKEEP_GUARDRAIL_ENDPOINT: Final = "/v3/openai/beta/litellm_basic_guardrail_api"
44class DeepKeepFirewallResponse(TypedDict):
45 """Body returned by the DeepKeep firewall endpoint."""
47 action: ReadOnly[NotRequired[str]]
48 blocked_reason: ReadOnly[NotRequired[str]]
49 texts: ReadOnly[NotRequired["list[str]"]]
50 images: ReadOnly[NotRequired["list[str]"]]
51 tools: ReadOnly[NotRequired["list[ChatCompletionToolParam]"]]
52 tool_calls: ReadOnly[NotRequired["list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall]"]]
53 structured_messages: ReadOnly[NotRequired["list[AllMessageValues]"]]
56class _DeepKeepInitKwargsView(TypedDict):
57 """Typed read of the guardrail name carried in the untyped base-guardrail kwargs."""
59 guardrail_name: ReadOnly[str | None]
62class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
63 """Base-class constructor options this guardrail forwards untouched to CustomGuardrail."""
65 guardrail_name: ReadOnly[str | None]
68class _DeepKeepMetadataSource(TypedDict, total=False):
69 """Typed read of the two untyped metadata mappings this guardrail merges."""
71 litellm_metadata: ReadOnly[Mapping[str, object]]
72 metadata: ReadOnly[Mapping[str, object]]
75class _FirewallResponseBody(Protocol):
76 def json(self) -> DeepKeepFirewallResponse: ... 76 ↛ exitline 76 didn't return from function 'json' because
79def _firewall_response_body(response: _FirewallResponseBody) -> DeepKeepFirewallResponse:
80 return response.json()
83class DeepKeepGuardrailMissingSecrets(Exception):
84 """Exception raised when DeepKeep API key or firewall_id is missing."""
87class DeepKeepGuardrailAPIError(Exception):
88 """Exception raised when there's an error calling the DeepKeep API."""
91class DeepKeepGuardrail(CustomGuardrail):
92 """
93 DeepKeep AI Firewall integration for LiteLLM.
95 Provides content moderation, prompt injection detection, PII protection,
96 and policy enforcement through the DeepKeep AI Firewall API.
98 DeepKeep's firewall evaluates LLM inputs and outputs against a configurable
99 set of guardrails (detectors + actions) managed via the DeepKeep platform.
101 Configuration example (litellm config YAML):
102 guardrails:
103 - guardrail_name: deepkeep-firewall
104 litellm_params:
105 guardrail: deepkeep
106 mode: pre_call
107 api_key: os.environ/DEEPKEEP_API_KEY
108 api_base: https://your-deepkeep-instance.example.com
109 deepkeep_firewall_id: your-firewall-id
110 """
112 def __init__(
113 self,
114 api_key: str | None = None,
115 api_base: str | None = None,
116 firewall_id: str | None = None,
117 unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
118 extra_headers: Mapping[str, str] | list[str] | None = None,
119 **kwargs: Unpack[_CustomGuardrailOptions],
120 ):
121 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
123 # API key
124 deepkeep_api_key: Final = api_key or os.environ.get("DEEPKEEP_API_KEY")
125 if not deepkeep_api_key:
126 raise DeepKeepGuardrailMissingSecrets(
127 "DeepKeep API key is required. Set the `DEEPKEEP_API_KEY` environment "
128 "variable or pass `api_key` in the guardrail config."
129 )
130 self.deepkeep_api_key: str = deepkeep_api_key
132 # Firewall ID
133 self.firewall_id = firewall_id or os.environ.get("DEEPKEEP_FIREWALL_ID")
134 if not self.firewall_id:
135 raise DeepKeepGuardrailMissingSecrets(
136 "DeepKeep firewall_id is required. Set the `DEEPKEEP_FIREWALL_ID` environment "
137 "variable or pass `deepkeep_firewall_id` in the guardrail config."
138 )
140 # API base URL
141 base_url = api_base or os.environ.get("DEEPKEEP_API_BASE")
142 if not base_url:
143 raise DeepKeepGuardrailMissingSecrets(
144 "DeepKeep API base URL is required. Set the `DEEPKEEP_API_BASE` environment "
145 "variable or pass `api_base` in the guardrail config."
146 )
148 # Normalize the API base – ensure it ends with the guardrail endpoint
149 base_url = base_url.rstrip("/")
150 if base_url.endswith(_DEEPKEEP_GUARDRAIL_ENDPOINT.rstrip("/")):
151 self.api_base = base_url
152 else:
153 self.api_base = f"{base_url}{_DEEPKEEP_GUARDRAIL_ENDPOINT}"
155 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback
156 if extra_headers is not None and not isinstance(extra_headers, Mapping):
157 verbose_proxy_logger.warning(
158 "DeepKeep guardrail ignoring `extra_headers`: expected a mapping of header name to value, got %s. "
159 "`litellm_params.extra_headers` is a list of header names to forward and is not supported by this guardrail",
160 type(extra_headers).__name__,
161 )
162 self.extra_headers: dict[str, str] = dict(extra_headers) if isinstance(extra_headers, Mapping) else {}
164 # Set supported event hooks
165 if "supported_event_hooks" not in kwargs:
166 kwargs["supported_event_hooks"] = [
167 GuardrailEventHooks.pre_call,
168 GuardrailEventHooks.post_call,
169 GuardrailEventHooks.during_call,
170 ]
172 super().__init__(**kwargs)
174 init_view: Final[_DeepKeepInitKwargsView] = {"guardrail_name": kwargs.get("guardrail_name", "unknown")}
176 verbose_proxy_logger.debug(
177 "DeepKeep guardrail initialized: guardrail_name=%s, api_base=%s, firewall_id=%s",
178 init_view["guardrail_name"],
179 self.api_base,
180 self.firewall_id,
181 )
183 def _extract_user_api_key_metadata(self, request_data: _DeepKeepMetadataSource) -> dict[str, object]:
184 """
185 Extract user API key metadata from request_data for the DeepKeep API.
187 Args:
188 request_data: Request data dictionary containing metadata.
190 Returns:
191 Dictionary with user API key metadata fields.
192 """
193 result_metadata: Final[dict[str, object]] = {}
195 litellm_metadata: Final = request_data.get("litellm_metadata", {})
196 top_level_metadata: Final = request_data.get("metadata", {})
197 metadata_dict: Final[Mapping[str, object]] = {**top_level_metadata, **litellm_metadata}
199 if not metadata_dict:
200 return result_metadata
202 # Extract standard user API key fields
203 _METADATA_KEYS: Final = [
204 "user_api_key_hash",
205 "user_api_key_alias",
206 "user_api_key_user_id",
207 "user_api_key_user_email",
208 "user_api_key_team_id",
209 "user_api_key_team_alias",
210 "user_api_key_end_user_id",
211 "user_api_key_org_id",
212 ]
213 for key in _METADATA_KEYS:
214 value = metadata_dict.get(key)
215 if value is not None:
216 result_metadata[key] = value
218 # Handle the token → hash alias (only when no explicit hash was provided)
219 if metadata_dict.get("user_api_key_token") is not None and "user_api_key_hash" not in result_metadata:
220 result_metadata["user_api_key_hash"] = metadata_dict["user_api_key_token"]
222 return result_metadata
224 def _build_request_headers(self) -> dict[str, str]:
225 """Build HTTP headers for the DeepKeep API request."""
226 headers: Final[dict[str, str]] = {
227 "Content-Type": "application/json",
228 "X-API-Key": self.deepkeep_api_key,
229 }
230 if self.extra_headers:
231 headers.update(self.extra_headers)
232 return headers
234 def _fail_open_passthrough(
235 self,
236 *,
237 inputs: GenericGuardrailAPIInputs,
238 input_type: Literal["request", "response"],
239 logging_obj: Optional["LiteLLMLoggingObj"],
240 error: Exception,
241 http_status_code: int | None = None,
242 ) -> GenericGuardrailAPIInputs:
243 """Allow the request to proceed when the guardrail is unreachable (fail-open mode)."""
244 status_suffix: Final = f" http_status_code={http_status_code}" if http_status_code else ""
245 verbose_proxy_logger.critical(
246 "DeepKeep guardrail unreachable (fail-open). Proceeding without guardrail.%s "
247 "guardrail_name=%s api_base=%s input_type=%s litellm_call_id=%s litellm_trace_id=%s",
248 status_suffix,
249 getattr(self, "guardrail_name", None),
250 getattr(self, "api_base", None),
251 input_type,
252 getattr(logging_obj, "litellm_call_id", None) if logging_obj else None,
253 getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None,
254 exc_info=error,
255 )
256 return_inputs: Final[GenericGuardrailAPIInputs] = {}
257 return_inputs.update(inputs)
258 return return_inputs
260 def _handle_guardrail_request_error(
261 self,
262 error: Exception,
263 inputs: GenericGuardrailAPIInputs,
264 input_type: Literal["request", "response"],
265 logging_obj: Optional["LiteLLMLoggingObj"],
266 is_unreachable: bool = True,
267 ) -> GenericGuardrailAPIInputs:
268 """Handle errors from the DeepKeep API with fail-open/fail-closed logic."""
269 if is_unreachable and self.unreachable_fallback == "fail_open":
270 http_status_code: Final[int | None] = getattr(getattr(error, "response", None), "status_code", None)
271 return self._fail_open_passthrough(
272 inputs=inputs,
273 input_type=input_type,
274 logging_obj=logging_obj,
275 error=error,
276 **({"http_status_code": http_status_code} if http_status_code else {}),
277 )
278 verbose_proxy_logger.error("DeepKeep guardrail API error: %s", str(error))
279 raise DeepKeepGuardrailAPIError(f"DeepKeep guardrail API failed: {error}")
281 @staticmethod
282 def _build_return_inputs(
283 *,
284 response_json: DeepKeepFirewallResponse,
285 texts: list[str],
286 images: "list[str] | None",
287 tools: "list[ChatCompletionToolParam] | None",
288 tool_calls: "list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None",
289 structured_messages: "list[AllMessageValues] | None",
290 ) -> GenericGuardrailAPIInputs:
291 """Merge original inputs with any guardrail-modified values from the API response.
293 Presence is checked with ``is not None`` (not truthiness) so that an
294 intentional empty-list replacement such as ``texts: []`` or
295 ``tool_calls: []`` is honoured and forwarded downstream rather than
296 silently discarded in favour of the original content.
297 """
298 return_inputs: Final = GenericGuardrailAPIInputs(texts=texts)
299 texts_override: Final = response_json.get("texts")
300 if texts_override is not None:
301 return_inputs["texts"] = texts_override
302 images_override: Final = response_json.get("images")
303 if images_override is not None:
304 return_inputs["images"] = images_override
305 elif images is not None:
306 return_inputs["images"] = images
307 tools_override: Final = response_json.get("tools")
308 if tools_override is not None:
309 return_inputs["tools"] = tools_override
310 elif tools is not None:
311 return_inputs["tools"] = tools
312 tool_calls_override: Final = response_json.get("tool_calls")
313 if tool_calls_override is not None:
314 return_inputs["tool_calls"] = tool_calls_override
315 elif tool_calls is not None:
316 return_inputs["tool_calls"] = tool_calls
317 structured_messages_override: Final = response_json.get("structured_messages")
318 if structured_messages_override is not None:
319 return_inputs["structured_messages"] = structured_messages_override
320 elif structured_messages is not None:
321 return_inputs["structured_messages"] = structured_messages
322 return return_inputs
324 @log_guardrail_information
325 async def apply_guardrail(
326 self,
327 inputs: GenericGuardrailAPIInputs,
328 request_data: dict,
329 input_type: Literal["request", "response"],
330 logging_obj: Optional["LiteLLMLoggingObj"] = None,
331 ) -> GenericGuardrailAPIInputs:
332 """
333 Apply the DeepKeep AI Firewall guardrail to the given inputs.
335 This is the main method called by the LiteLLM framework for guardrail evaluation.
337 Args:
338 inputs: Dictionary containing texts, images, tools, tool_calls, structured_messages.
339 request_data: Request data dictionary containing metadata.
340 input_type: Whether this is a "request" (pre-call) or "response" (post-call) guardrail.
341 logging_obj: Optional logging object for tracking the guardrail execution.
343 Returns:
344 GenericGuardrailAPIInputs with original or modified content.
346 Raises:
347 GuardrailRaisedException: If the guardrail blocks the request.
348 DeepKeepGuardrailAPIError: If the API call fails (in fail-closed mode).
349 """
350 verbose_proxy_logger.debug("DeepKeep guardrail: applying guardrail, input_type=%s", input_type)
352 texts: Final = inputs.get("texts", [])
353 images: Final = inputs.get("images")
354 tools: Final = inputs.get("tools")
355 structured_messages: Final = inputs.get("structured_messages")
356 tool_calls: Final = inputs.get("tool_calls")
357 model: Final = inputs.get("model")
359 if request_data is None:
360 request_data = {}
362 request_body: Final = request_data.get("body") or {}
364 # Merge additional provider-specific params from config and dynamic params
365 additional_params: Final[dict[str, object]] = {"firewall_id": self.firewall_id}
366 dynamic_params: Final = self.get_guardrail_dynamic_request_body_params(request_body)
367 if dynamic_params:
368 additional_params.update({k: v for k, v in dynamic_params.items() if k != "firewall_id"})
370 # Extract user API key metadata
371 user_metadata: Final = self._extract_user_api_key_metadata(request_data)
373 # Build request payload
374 guardrail_request: Final[dict[str, object]] = {
375 "litellm_call_id": (logging_obj.litellm_call_id if logging_obj else None),
376 "litellm_trace_id": (logging_obj.litellm_trace_id if logging_obj else None),
377 "texts": texts,
378 "request_data": user_metadata,
379 "litellm_version": litellm_version,
380 "images": images,
381 "tools": tools,
382 "structured_messages": structured_messages,
383 "tool_calls": tool_calls,
384 "additional_provider_specific_params": additional_params,
385 "input_type": input_type,
386 "model": model,
387 }
389 headers: Final = self._build_request_headers()
391 try:
392 response: Final = await self.async_handler.post(
393 url=self.api_base,
394 json=guardrail_request,
395 headers=headers,
396 )
398 response.raise_for_status()
399 response_json: Final = _firewall_response_body(response)
401 verbose_proxy_logger.debug("DeepKeep guardrail response: %s", response_json)
403 action: Final = response_json.get("action", "NONE")
405 if action == "BLOCKED":
406 error_message: Final = response_json.get("blocked_reason") or "Content violates policy"
407 verbose_proxy_logger.warning("DeepKeep guardrail blocked request: %s", error_message)
408 raise GuardrailRaisedException(
409 guardrail_name=GUARDRAIL_NAME,
410 message=error_message,
411 should_wrap_with_default_message=False,
412 blocked_content=True,
413 )
415 return self._build_return_inputs(
416 response_json=response_json,
417 texts=texts,
418 images=images,
419 tools=tools,
420 tool_calls=tool_calls,
421 structured_messages=structured_messages,
422 )
424 except GuardrailRaisedException:
425 raise
426 except Timeout as e:
427 return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
428 except httpx.HTTPStatusError as e:
429 status_code: Final = getattr(getattr(e, "response", None), "status_code", None)
430 is_unreachable: Final = status_code in (502, 503, 504)
431 return self._handle_guardrail_request_error(
432 e, inputs, input_type, logging_obj, is_unreachable=is_unreachable
433 )
434 except httpx.RequestError as e:
435 return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
436 except Exception as e: # noqa: BLE001 # route unexpected errors through fail-open/closed handling
437 return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj, is_unreachable=False)
439 @staticmethod
440 def get_config_model() -> type | None:
441 from litellm.types.proxy.guardrails.guardrail_hooks.deepkeep import (
442 DeepKeepGuardrailConfigModel,
443 )
445 return DeepKeepGuardrailConfigModel