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

1# +-------------------------------------------------------------+ 

2# 

3# Use IBM Guardrails Detector for your LLM calls 

4# Based on IBM's FMS Guardrails 

5# 

6# +-------------------------------------------------------------+ 

7 

8import os 

9from collections.abc import AsyncGenerator 

10from datetime import datetime 

11from typing import Any, Final 

12 

13import httpx 

14 

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 

31 

32GUARDRAIL_NAME: Final = "ibm_guardrails" 

33 

34 

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 ) 

54 

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 ) 

61 

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.") 

65 

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.") 

69 

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 

76 

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" 

82 

83 self.guardrail_name = guardrail_name 

84 self.guardrail_provider = "ibm_guardrails" 

85 

86 # store kwargs as optional_params 

87 self.optional_params = kwargs 

88 

89 # Set supported event hooks 

90 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) 

91 

92 super().__init__(guardrail_name=guardrail_name, **kwargs) 

93 

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 ) 

100 

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. 

109 

110 Args: 

111 contents: List of text strings to analyze 

112 request_data: Optional request data for logging purposes 

113 

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() 

119 

120 payload: Final = {"contents": contents, "detector_params": self.detector_params} 

121 

122 headers: Final = { 

123 "Authorization": f"Bearer {self.auth_token}", 

124 "content-type": "application/json", 

125 "detector-id": self.detector_id, 

126 } 

127 

128 # Add any extra headers to the request 

129 for header, value in self.extra_headers.items(): 

130 headers[header] = value 

131 

132 verbose_proxy_logger.debug( 

133 "IBM Detector Server request to %s with payload: %s", 

134 self.api_url, 

135 payload, 

136 ) 

137 

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() 

146 

147 end_time = datetime.now() 

148 duration = (end_time - start_time).total_seconds() 

149 

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 ) 

167 

168 return response_json 

169 

170 except httpx.HTTPError as e: 

171 end_time = datetime.now() 

172 duration = (end_time - start_time).total_seconds() 

173 

174 verbose_proxy_logger.error("IBM Detector Server request failed: %s", str(e)) 

175 

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 ) 

188 

189 raise 

190 

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. 

199 

200 Args: 

201 content: Text string to analyze 

202 request_data: Optional request data for logging purposes 

203 

204 Returns: 

205 List of detections 

206 """ 

207 start_time: Final = datetime.now() 

208 

209 payload: Final = { 

210 "content": content, 

211 "detectors": {self.detector_id: self.detector_params}, 

212 } 

213 

214 headers: Final = { 

215 "Authorization": f"Bearer {self.auth_token}", 

216 "content-type": "application/json", 

217 } 

218 

219 # Add any extra headers to the request 

220 for header, value in self.extra_headers.items(): 

221 headers[header] = value 

222 

223 verbose_proxy_logger.debug( 

224 "IBM Orchestrator request to %s with payload: %s", 

225 self.api_url, 

226 payload, 

227 ) 

228 

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() 

237 

238 end_time = datetime.now() 

239 duration = (end_time - start_time).total_seconds() 

240 

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 ) 

254 

255 return response_json.get("detections", []) 

256 

257 except httpx.HTTPError as e: 

258 end_time = datetime.now() 

259 duration = (end_time - start_time).total_seconds() 

260 

261 verbose_proxy_logger.error("IBM Orchestrator request failed: %s", str(e)) 

262 

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 ) 

275 

276 raise 

277 

278 def _filter_detections_by_threshold(self, detections: list[IBMDetectorDetection]) -> list[IBMDetectorDetection]: 

279 """ 

280 Filter detections based on score threshold. 

281 

282 Args: 

283 detections: List of detections 

284 

285 Returns: 

286 Filtered list of detections that meet the threshold 

287 """ 

288 if self.score_threshold is None: 

289 return detections 

290 

291 return [detection for detection in detections if detection.get("score", 0.0) >= self.score_threshold] 

292 

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. 

298 

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" 

307 

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 

317 

318 if has_detections: 

319 return "guardrail_intervened" 

320 

321 return "success" 

322 

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" 

326 

327 def _determine_guardrail_status_orchestrator( 

328 self, response_json: IBMDetectorResponseOrchestrator 

329 ) -> GuardrailStatus: 

330 """ 

331 Determine the guardrail status based on IBM Orchestrator response. 

332 

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" 

341 

342 detections: Final = response_json.get("detections", []) 

343 # Apply threshold filtering 

344 filtered: Final = self._filter_detections_by_threshold(detections) 

345 

346 if filtered: 

347 return "guardrail_intervened" 

348 

349 return "success" 

350 

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" 

354 

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. 

358 

359 Args: 

360 detections_list: List of lists of detections 

361 

362 Returns: 

363 Formatted error message string 

364 """ 

365 total_detections = 0 

366 error_message = "IBM Guardrail Detector failed:\n\n" 

367 

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) 

373 

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" 

380 

381 error_message += "\n" 

382 

383 error_message = f"IBM Guardrail Detector failed: {total_detections} violation(s) detected\n\n" + error_message 

384 return error_message.strip() 

385 

386 def _create_error_message_orchestrator(self, detections: list[IBMDetectorDetection]) -> str: 

387 """ 

388 Create a detailed error message from orchestrator response. 

389 

390 Args: 

391 detections: List of detections 

392 

393 Returns: 

394 Formatted error message string 

395 """ 

396 filtered_detections: Final = self._filter_detections_by_threshold(detections) 

397 

398 error_message = f"IBM Guardrail Detector failed: {len(filtered_detections)} violation(s) detected\n\n" 

399 

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", "") 

405 

406 error_message += f"- {detection_type.upper()} (detector: {detector_id}, score: {score:.3f})\n" 

407 error_message += f" Text: '{text}'\n\n" 

408 

409 return error_message.strip() 

410 

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") 

424 

425 from litellm.proxy.common_utils.callback_utils import ( 

426 add_guardrail_to_applied_guardrails_header, 

427 ) 

428 

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 

432 

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 ) 

443 

444 verbose_proxy_logger.debug("IBM Detector Server async_pre_call_hook result: %s", result) 

445 

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 

453 

454 if has_violations and self.block_on_detection: 

455 error_message = self._create_error_message_detector_server(result) 

456 raise ValueError(error_message) 

457 

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 ) 

466 

467 verbose_proxy_logger.debug( 

468 "IBM Orchestrator async_pre_call_hook result: %s", 

469 orchestrator_result, 

470 ) 

471 

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) 

476 

477 # Add guardrail to applied guardrails header 

478 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

479 

480 return data 

481 

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 

491 

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 ) 

497 

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 

501 

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 ) 

512 

513 verbose_proxy_logger.debug("IBM Detector Server async_moderation_hook result: %s", result) 

514 

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 

522 

523 if has_violations and self.block_on_detection: 

524 error_message = self._create_error_message_detector_server(result) 

525 raise ValueError(error_message) 

526 

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 ) 

535 

536 verbose_proxy_logger.debug( 

537 "IBM Orchestrator async_moderation_hook result: %s", 

538 orchestrator_result, 

539 ) 

540 

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) 

545 

546 # Add guardrail to applied guardrails header 

547 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

548 

549 return data 

550 

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 

559 

560 It can be used to reject a response 

561 

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 

568 

569 if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: 

570 return 

571 

572 verbose_proxy_logger.debug("async_post_call_success_hook response: %s", response) 

573 

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 

583 

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 

589 

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) 

596 

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 ) 

605 

606 verbose_proxy_logger.debug( 

607 "IBM Detector Server async_post_call_success_hook result: %s", 

608 result, 

609 ) 

610 

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 

618 

619 if has_violations and self.block_on_detection: 

620 error_message = self._create_error_message_detector_server(result) 

621 raise ValueError(error_message) 

622 

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 ) 

631 

632 verbose_proxy_logger.debug( 

633 "IBM Orchestrator async_post_call_success_hook result: %s", 

634 orchestrator_result, 

635 ) 

636 

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) 

641 

642 # Add guardrail to applied guardrails header 

643 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

644 

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 

653 

654 This is useful for guardrails that need to see the entire response, such as PII masking. 

655 

656 Triggered by mode: 'post_call' 

657 """ 

658 async for item in response: 

659 yield item 

660 

661 @staticmethod 

662 def get_config_model(): 

663 from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( 

664 IBMDetectorGuardrailConfigModel, 

665 ) 

666 

667 return IBMDetectorGuardrailConfigModel 

668 

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 ]