Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py: 10%
262 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 IBM Guardrails Detector for your LLM calls
4# Based on IBM's FMS Guardrails
5#
6# +-------------------------------------------------------------+
8import os
9from collections.abc import AsyncGenerator
10from datetime import datetime
11from typing import Any, Final
13import httpx
15import litellm
16from litellm._logging import verbose_proxy_logger
17from litellm.caching.caching import DualCache
18from litellm.integrations.custom_guardrail import CustomGuardrail
19from litellm.llms.custom_httpx.http_handler import (
20 get_async_httpx_client,
21 httpxSpecialProvider,
22)
23from litellm.proxy._types import UserAPIKeyAuth
24from litellm.proxy.guardrails._content_utils import iter_message_text
25from litellm.types.guardrails import GuardrailEventHooks
26from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
27 IBMDetectorDetection,
28 IBMDetectorResponseOrchestrator,
29)
30from litellm.types.utils import CallTypesLiteral, GuardrailStatus, ModelResponseStream
32GUARDRAIL_NAME: Final = "ibm_guardrails"
35class IBMGuardrailDetector(CustomGuardrail):
36 def __init__(
37 self,
38 guardrail_name: str = "ibm_detector",
39 auth_token: str | None = None,
40 base_url: str | None = None,
41 detector_id: str | None = None,
42 is_detector_server: bool = True,
43 detector_params: dict[str, Any] | None = None,
44 extra_headers: dict[str, str] | None = None,
45 score_threshold: float | None = None,
46 block_on_detection: bool = True,
47 verify_ssl: bool = True,
48 **kwargs,
49 ):
50 self.async_handler = get_async_httpx_client(
51 llm_provider=httpxSpecialProvider.GuardrailCallback,
52 params={"ssl_verify": verify_ssl},
53 )
55 # Set API configuration
56 self.auth_token = auth_token or os.getenv("IBM_GUARDRAILS_AUTH_TOKEN")
57 if not self.auth_token:
58 raise ValueError(
59 "IBM Guardrails auth token is required. Set IBM_GUARDRAILS_AUTH_TOKEN environment variable or pass auth_token parameter."
60 )
62 self.base_url = base_url
63 if not self.base_url:
64 raise ValueError("IBM Guardrails base_url is required. Pass base_url parameter.")
66 self.detector_id = detector_id
67 if not self.detector_id:
68 raise ValueError("IBM Guardrails detector_id is required. Pass detector_id parameter.")
70 self.is_detector_server = is_detector_server
71 self.detector_params = detector_params or {}
72 self.extra_headers = extra_headers or {}
73 self.score_threshold = score_threshold
74 self.block_on_detection = block_on_detection
75 self.verify_ssl = verify_ssl
77 # Construct API URL based on server type
78 if self.is_detector_server:
79 self.api_url = f"{self.base_url}/api/v1/text/contents"
80 else:
81 self.api_url = f"{self.base_url}/api/v2/text/detection/content"
83 self.guardrail_name = guardrail_name
84 self.guardrail_provider = "ibm_guardrails"
86 # store kwargs as optional_params
87 self.optional_params = kwargs
89 # Set supported event hooks
90 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
92 super().__init__(guardrail_name=guardrail_name, **kwargs)
94 verbose_proxy_logger.debug(
95 "IBM Guardrail Detector initialized with guardrail_name: %s, detector_id: %s, is_detector_server: %s",
96 self.guardrail_name,
97 self.detector_id,
98 self.is_detector_server,
99 )
101 async def _call_detector_server(
102 self,
103 contents: list[str],
104 event_type: GuardrailEventHooks,
105 request_data: dict | None = None,
106 ) -> list[list[IBMDetectorDetection]]:
107 """
108 Call IBM Detector Server directly.
110 Args:
111 contents: List of text strings to analyze
112 request_data: Optional request data for logging purposes
114 Returns:
115 List of lists where top-level list is per message in contents,
116 sublists are individual detections on that message
117 """
118 start_time: Final = datetime.now()
120 payload: Final = {"contents": contents, "detector_params": self.detector_params}
122 headers: Final = {
123 "Authorization": f"Bearer {self.auth_token}",
124 "content-type": "application/json",
125 "detector-id": self.detector_id,
126 }
128 # Add any extra headers to the request
129 for header, value in self.extra_headers.items():
130 headers[header] = value
132 verbose_proxy_logger.debug(
133 "IBM Detector Server request to %s with payload: %s",
134 self.api_url,
135 payload,
136 )
138 try:
139 response: Final = await self.async_handler.post(
140 url=self.api_url,
141 json=payload,
142 headers=headers,
143 )
144 response.raise_for_status()
145 response_json: Final[list[list[IBMDetectorDetection]]] = response.json()
147 end_time = datetime.now()
148 duration = (end_time - start_time).total_seconds()
150 # Add guardrail information to request trace
151 if request_data:
152 guardrail_status: Final = self._determine_guardrail_status_detector_server(response_json)
153 self.add_standard_logging_guardrail_information_to_request_data(
154 guardrail_provider=self.guardrail_provider,
155 guardrail_json_response={
156 "detections": [
157 [detection for detection in message_detections] for message_detections in response_json
158 ]
159 },
160 request_data=request_data,
161 guardrail_status=guardrail_status,
162 start_time=start_time.timestamp(),
163 end_time=end_time.timestamp(),
164 duration=duration,
165 event_type=event_type,
166 )
168 return response_json
170 except httpx.HTTPError as e:
171 end_time = datetime.now()
172 duration = (end_time - start_time).total_seconds()
174 verbose_proxy_logger.error("IBM Detector Server request failed: %s", str(e))
176 # Add guardrail information with failure status
177 if request_data:
178 self.add_standard_logging_guardrail_information_to_request_data(
179 guardrail_provider=self.guardrail_provider,
180 guardrail_json_response={"error": str(e)},
181 request_data=request_data,
182 guardrail_status="guardrail_failed_to_respond",
183 start_time=start_time.timestamp(),
184 end_time=end_time.timestamp(),
185 duration=duration,
186 event_type=event_type,
187 )
189 raise
191 async def _call_orchestrator(
192 self,
193 content: str,
194 event_type: GuardrailEventHooks,
195 request_data: dict | None = None,
196 ) -> list[IBMDetectorDetection]:
197 """
198 Call IBM FMS Guardrails Orchestrator.
200 Args:
201 content: Text string to analyze
202 request_data: Optional request data for logging purposes
204 Returns:
205 List of detections
206 """
207 start_time: Final = datetime.now()
209 payload: Final = {
210 "content": content,
211 "detectors": {self.detector_id: self.detector_params},
212 }
214 headers: Final = {
215 "Authorization": f"Bearer {self.auth_token}",
216 "content-type": "application/json",
217 }
219 # Add any extra headers to the request
220 for header, value in self.extra_headers.items():
221 headers[header] = value
223 verbose_proxy_logger.debug(
224 "IBM Orchestrator request to %s with payload: %s",
225 self.api_url,
226 payload,
227 )
229 try:
230 response: Final = await self.async_handler.post(
231 url=self.api_url,
232 json=payload,
233 headers=headers,
234 )
235 response.raise_for_status()
236 response_json: Final[IBMDetectorResponseOrchestrator] = response.json()
238 end_time = datetime.now()
239 duration = (end_time - start_time).total_seconds()
241 # Add guardrail information to request trace
242 if request_data:
243 guardrail_status: Final = self._determine_guardrail_status_orchestrator(response_json)
244 self.add_standard_logging_guardrail_information_to_request_data(
245 guardrail_provider=self.guardrail_provider,
246 guardrail_json_response=dict(response_json),
247 request_data=request_data,
248 guardrail_status=guardrail_status,
249 start_time=start_time.timestamp(),
250 end_time=end_time.timestamp(),
251 duration=duration,
252 event_type=event_type,
253 )
255 return response_json.get("detections", [])
257 except httpx.HTTPError as e:
258 end_time = datetime.now()
259 duration = (end_time - start_time).total_seconds()
261 verbose_proxy_logger.error("IBM Orchestrator request failed: %s", str(e))
263 # Add guardrail information with failure status
264 if request_data:
265 self.add_standard_logging_guardrail_information_to_request_data(
266 guardrail_provider=self.guardrail_provider,
267 guardrail_json_response={"error": str(e)},
268 request_data=request_data,
269 guardrail_status="guardrail_failed_to_respond",
270 start_time=start_time.timestamp(),
271 end_time=end_time.timestamp(),
272 duration=duration,
273 event_type=event_type,
274 )
276 raise
278 def _filter_detections_by_threshold(self, detections: list[IBMDetectorDetection]) -> list[IBMDetectorDetection]:
279 """
280 Filter detections based on score threshold.
282 Args:
283 detections: List of detections
285 Returns:
286 Filtered list of detections that meet the threshold
287 """
288 if self.score_threshold is None:
289 return detections
291 return [detection for detection in detections if detection.get("score", 0.0) >= self.score_threshold]
293 def _determine_guardrail_status_detector_server(
294 self, response_json: list[list[IBMDetectorDetection]]
295 ) -> GuardrailStatus:
296 """
297 Determine the guardrail status based on IBM Detector Server response.
299 Returns:
300 "success": Content allowed through with no violations
301 "guardrail_intervened": Content blocked due to detections
302 "guardrail_failed_to_respond": Technical error or API failure
303 """
304 try:
305 if not isinstance(response_json, list):
306 return "guardrail_failed_to_respond"
308 # Check if any detections were found
309 has_detections = False
310 for message_detections in response_json:
311 if message_detections:
312 # Apply threshold filtering
313 filtered = self._filter_detections_by_threshold(message_detections)
314 if filtered:
315 has_detections = True
316 break
318 if has_detections:
319 return "guardrail_intervened"
321 return "success"
323 except Exception as e:
324 verbose_proxy_logger.error("Error determining IBM Detector Server guardrail status: %s", str(e))
325 return "guardrail_failed_to_respond"
327 def _determine_guardrail_status_orchestrator(
328 self, response_json: IBMDetectorResponseOrchestrator
329 ) -> GuardrailStatus:
330 """
331 Determine the guardrail status based on IBM Orchestrator response.
333 Returns:
334 "success": Content allowed through with no violations
335 "guardrail_intervened": Content blocked due to detections
336 "guardrail_failed_to_respond": Technical error or API failure
337 """
338 try:
339 if not isinstance(response_json, dict):
340 return "guardrail_failed_to_respond"
342 detections: Final = response_json.get("detections", [])
343 # Apply threshold filtering
344 filtered: Final = self._filter_detections_by_threshold(detections)
346 if filtered:
347 return "guardrail_intervened"
349 return "success"
351 except Exception as e:
352 verbose_proxy_logger.error("Error determining IBM Orchestrator guardrail status: %s", str(e))
353 return "guardrail_failed_to_respond"
355 def _create_error_message_detector_server(self, detections_list: list[list[IBMDetectorDetection]]) -> str:
356 """
357 Create a detailed error message from detector server response.
359 Args:
360 detections_list: List of lists of detections
362 Returns:
363 Formatted error message string
364 """
365 total_detections = 0
366 error_message = "IBM Guardrail Detector failed:\n\n"
368 for idx, message_detections in enumerate(detections_list):
369 filtered_detections = self._filter_detections_by_threshold(message_detections)
370 if filtered_detections:
371 error_message += f"Message {idx + 1}:\n"
372 total_detections += len(filtered_detections)
374 for detection in filtered_detections:
375 detection_type = detection.get("detection_type", "unknown")
376 score = detection.get("score", 0.0)
377 text = detection.get("text", "")
378 error_message += f" - {detection_type.upper()} (score: {score:.3f})\n"
379 error_message += f" Text: '{text}'\n"
381 error_message += "\n"
383 error_message = f"IBM Guardrail Detector failed: {total_detections} violation(s) detected\n\n" + error_message
384 return error_message.strip()
386 def _create_error_message_orchestrator(self, detections: list[IBMDetectorDetection]) -> str:
387 """
388 Create a detailed error message from orchestrator response.
390 Args:
391 detections: List of detections
393 Returns:
394 Formatted error message string
395 """
396 filtered_detections: Final = self._filter_detections_by_threshold(detections)
398 error_message = f"IBM Guardrail Detector failed: {len(filtered_detections)} violation(s) detected\n\n"
400 for detection in filtered_detections:
401 detection_type = detection.get("detection_type", "unknown")
402 detector_id = detection.get("detector_id", self.detector_id)
403 score = detection.get("score", 0.0)
404 text = detection.get("text", "")
406 error_message += f"- {detection_type.upper()} (detector: {detector_id}, score: {score:.3f})\n"
407 error_message += f" Text: '{text}'\n\n"
409 return error_message.strip()
411 async def async_pre_call_hook(
412 self,
413 user_api_key_dict: UserAPIKeyAuth,
414 cache: DualCache,
415 data: dict,
416 call_type: CallTypesLiteral,
417 ) -> Exception | str | dict | None:
418 """
419 Runs before the LLM API call
420 Runs on only Input
421 Use this if you want to MODIFY the input
422 """
423 verbose_proxy_logger.debug("Running IBM Guardrail Detector pre-call hook")
425 from litellm.proxy.common_utils.callback_utils import (
426 add_guardrail_to_applied_guardrails_header,
427 )
429 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.pre_call
430 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
431 return data
433 # Covers multimodal list content + Responses-API input.
434 contents_to_check: Final[list[str]] = list(iter_message_text(data))
435 if contents_to_check:
436 if self.is_detector_server:
437 # Call detector server with all contents at once
438 result: Final = await self._call_detector_server(
439 contents=contents_to_check,
440 request_data=data,
441 event_type=GuardrailEventHooks.pre_call,
442 )
444 verbose_proxy_logger.debug("IBM Detector Server async_pre_call_hook result: %s", result)
446 # Check if any detections were found
447 has_violations = False
448 for message_detections in result:
449 filtered = self._filter_detections_by_threshold(message_detections)
450 if filtered:
451 has_violations = True
452 break
454 if has_violations and self.block_on_detection:
455 error_message = self._create_error_message_detector_server(result)
456 raise ValueError(error_message)
458 else:
459 # Call orchestrator for each content separately
460 for content in contents_to_check:
461 orchestrator_result = await self._call_orchestrator(
462 content=content,
463 request_data=data,
464 event_type=GuardrailEventHooks.pre_call,
465 )
467 verbose_proxy_logger.debug(
468 "IBM Orchestrator async_pre_call_hook result: %s",
469 orchestrator_result,
470 )
472 filtered = self._filter_detections_by_threshold(orchestrator_result)
473 if filtered and self.block_on_detection:
474 error_message = self._create_error_message_orchestrator(orchestrator_result)
475 raise ValueError(error_message)
477 # Add guardrail to applied guardrails header
478 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
480 return data
482 async def async_moderation_hook(
483 self,
484 data: dict,
485 user_api_key_dict: UserAPIKeyAuth,
486 call_type: CallTypesLiteral,
487 ):
488 """
489 Runs in parallel to LLM API call
490 Runs on only Input
492 This can NOT modify the input, only used to reject or accept a call before going to LLM API
493 """
494 from litellm.proxy.common_utils.callback_utils import (
495 add_guardrail_to_applied_guardrails_header,
496 )
498 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.during_call
499 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
500 return
502 # Covers multimodal list content + Responses-API input.
503 contents_to_check: Final[list[str]] = list(iter_message_text(data))
504 if contents_to_check:
505 if self.is_detector_server:
506 # Call detector server with all contents at once
507 result: Final = await self._call_detector_server(
508 contents=contents_to_check,
509 request_data=data,
510 event_type=GuardrailEventHooks.during_call,
511 )
513 verbose_proxy_logger.debug("IBM Detector Server async_moderation_hook result: %s", result)
515 # Check if any detections were found
516 has_violations = False
517 for message_detections in result:
518 filtered = self._filter_detections_by_threshold(message_detections)
519 if filtered:
520 has_violations = True
521 break
523 if has_violations and self.block_on_detection:
524 error_message = self._create_error_message_detector_server(result)
525 raise ValueError(error_message)
527 else:
528 # Call orchestrator for each content separately
529 for content in contents_to_check:
530 orchestrator_result = await self._call_orchestrator(
531 content=content,
532 request_data=data,
533 event_type=GuardrailEventHooks.during_call,
534 )
536 verbose_proxy_logger.debug(
537 "IBM Orchestrator async_moderation_hook result: %s",
538 orchestrator_result,
539 )
541 filtered = self._filter_detections_by_threshold(orchestrator_result)
542 if filtered and self.block_on_detection:
543 error_message = self._create_error_message_orchestrator(orchestrator_result)
544 raise ValueError(error_message)
546 # Add guardrail to applied guardrails header
547 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
549 return data
551 async def async_post_call_success_hook(
552 self,
553 data: dict,
554 user_api_key_dict: UserAPIKeyAuth,
555 response,
556 ):
557 """
558 Runs on response from LLM API call
560 It can be used to reject a response
562 Uses IBM Guardrails Detector to check the response for violations
563 """
564 from litellm.proxy.common_utils.callback_utils import (
565 add_guardrail_to_applied_guardrails_header,
566 )
567 from litellm.types.guardrails import GuardrailEventHooks
569 if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True:
570 return
572 verbose_proxy_logger.debug("async_post_call_success_hook response: %s", response)
574 # Check if the ModelResponse has text content in its choices
575 # to avoid sending empty content to IBM Detector (e.g., during tool calls)
576 if isinstance(response, litellm.ModelResponse):
577 has_text_content = False
578 for choice in response.choices:
579 if isinstance(choice, litellm.Choices):
580 if choice.message.content and isinstance(choice.message.content, str):
581 has_text_content = True
582 break
584 if not has_text_content:
585 verbose_proxy_logger.warning(
586 "IBM Guardrail Detector: not running guardrail. No output text in response"
587 )
588 return
590 contents_to_check: Final[list[str]] = []
591 for choice in response.choices:
592 if isinstance(choice, litellm.Choices):
593 verbose_proxy_logger.debug("async_post_call_success_hook choice: %s", choice)
594 if choice.message.content and isinstance(choice.message.content, str):
595 contents_to_check.append(choice.message.content)
597 if contents_to_check:
598 if self.is_detector_server:
599 # Call detector server with all contents at once
600 result: Final = await self._call_detector_server(
601 contents=contents_to_check,
602 request_data=data,
603 event_type=GuardrailEventHooks.post_call,
604 )
606 verbose_proxy_logger.debug(
607 "IBM Detector Server async_post_call_success_hook result: %s",
608 result,
609 )
611 # Check if any detections were found
612 has_violations = False
613 for message_detections in result:
614 filtered = self._filter_detections_by_threshold(message_detections)
615 if filtered:
616 has_violations = True
617 break
619 if has_violations and self.block_on_detection:
620 error_message = self._create_error_message_detector_server(result)
621 raise ValueError(error_message)
623 else:
624 # Call orchestrator for each content separately
625 for content in contents_to_check:
626 orchestrator_result = await self._call_orchestrator(
627 content=content,
628 request_data=data,
629 event_type=GuardrailEventHooks.post_call,
630 )
632 verbose_proxy_logger.debug(
633 "IBM Orchestrator async_post_call_success_hook result: %s",
634 orchestrator_result,
635 )
637 filtered = self._filter_detections_by_threshold(orchestrator_result)
638 if filtered and self.block_on_detection:
639 error_message = self._create_error_message_orchestrator(orchestrator_result)
640 raise ValueError(error_message)
642 # Add guardrail to applied guardrails header
643 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
645 async def async_post_call_streaming_iterator_hook(
646 self,
647 user_api_key_dict: UserAPIKeyAuth,
648 response: Any,
649 request_data: dict,
650 ) -> AsyncGenerator[ModelResponseStream, None]:
651 """
652 Passes the entire stream to the guardrail
654 This is useful for guardrails that need to see the entire response, such as PII masking.
656 Triggered by mode: 'post_call'
657 """
658 async for item in response:
659 yield item
661 @staticmethod
662 def get_config_model():
663 from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
664 IBMDetectorGuardrailConfigModel,
665 )
667 return IBMDetectorGuardrailConfigModel
669 @classmethod
670 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
671 return [
672 GuardrailEventHooks.pre_call,
673 GuardrailEventHooks.post_call,
674 GuardrailEventHooks.during_call,
675 ]