Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py: 15%

280 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2XecGuard guardrail integration for LiteLLM. 

3 

4Calls the CyCraft XecGuard API (https://api-xecguard.cycraft.ai) 

5to scan the full conversation history against configured policies 

6(prompt-injection, PII, harmful-content, custom rules) and, when 

7grounding documents are supplied via request metadata, also validates 

8the assistant response against those reference documents via the 

9/grounding endpoint. 

10 

11Design notes (intentional divergences from the framework defaults): 

12 * The full conversation history (system + user + assistant) is always 

13 forwarded to XecGuard regardless of ``scan_type``. This bypasses the 

14 framework's optional ``skip_system_message_in_guardrail`` behaviour 

15 on purpose - policy enforcement depends on system-prompt visibility. 

16 * ``apply_guardrail`` is defined directly on this class so the 

17 ``during_call`` dispatch (proxy/utils.py checks for the method on 

18 ``type(callback).__dict__``) reaches our implementation. 

19 * ``async_logging_hook`` is overridden because the framework calls it 

20 directly for ``logging_only`` mode - it does NOT bridge to 

21 ``apply_guardrail``. Our override runs the scan non-blockingly and 

22 swallows every exception. 

23""" 

24 

25import asyncio 

26import os 

27from datetime import datetime 

28from typing import TYPE_CHECKING, Any, Final, Literal, Optional 

29 

30from fastapi.exceptions import HTTPException 

31 

32from litellm._logging import verbose_proxy_logger 

33from litellm.integrations.custom_guardrail import ( 

34 CustomGuardrail, 

35 log_guardrail_information, 

36) 

37from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys 

38from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload 

39from litellm.llms.custom_httpx.http_handler import ( 

40 get_async_httpx_client, 

41 httpxSpecialProvider, 

42) 

43from litellm.types.guardrails import GuardrailEventHooks 

44from litellm.types.utils import ( 

45 GenericGuardrailAPIInputs, 

46 GuardrailStatus, 

47 StandardLoggingGuardrailInformation, 

48) 

49 

50if TYPE_CHECKING: 50 ↛ 51line 50 didn't jump to line 51 because the condition on line 50 was never true

51 from litellm.litellm_core_utils.litellm_logging import ( 

52 Logging as LiteLLMLoggingObj, 

53 ) 

54 from litellm.types.proxy.guardrails.guardrail_hooks.base import ( 

55 GuardrailConfigModel, 

56 ) 

57 

58 

59def _sanitize_scan_result_for_logging(scan_result: dict) -> dict: 

60 without_secrets: Final = {key: value for key, value in scan_result.items() if key != "secret_fields"} 

61 redacted: Final = redact_nested_match_and_regex_keys(without_secrets) 

62 masked: Final = mask_credentials_in_payload(redacted if isinstance(redacted, dict) else without_secrets) 

63 return masked if isinstance(masked, dict) else without_secrets 

64 

65 

66_DEFAULT_API_BASE: Final = "https://api-xecguard.cycraft.ai" 

67_SCAN_ENDPOINT: Final = "/xecguard/v1/scan" 

68_GROUNDING_ENDPOINT: Final = "/xecguard/v1/grounding" 

69_DEFAULT_MODEL: Final = "xecguard_v2" 

70_DEFAULT_GROUNDING_STRICTNESS: Final = "BALANCED" 

71_METADATA_GROUNDING_KEY: Final = "xecguard_grounding_documents" 

72_RATIONALE_TRUNCATE_CHARS: Final = 200 

73_DEFAULT_POLICIES: Final = [ 

74 "Default_Policy_SystemPromptEnforcement", 

75 "Default_Policy_HarmfulContentProtection", 

76 "Default_Policy_GeneralPromptAttackProtection", 

77] 

78 

79 

80class XecGuardMissingCredentials(Exception): 

81 pass 

82 

83 

84class XecGuardGuardrail(CustomGuardrail): 

85 def __init__( 

86 self, 

87 api_key: str | None = None, 

88 api_base: str | None = None, 

89 xecguard_model: str | None = None, 

90 policy_names: list[str] | None = None, 

91 block_on_error: bool | None = None, 

92 grounding_strictness: str | None = None, 

93 **kwargs: Any, 

94 ) -> None: 

95 self.api_key = api_key or os.environ.get("XECGUARD_API_KEY") 

96 if not self.api_key: 

97 raise XecGuardMissingCredentials( 

98 "XecGuard API key is required. " 

99 "Set XECGUARD_API_KEY in the " 

100 "environment or pass api_key in " 

101 "the guardrail config." 

102 ) 

103 

104 self.api_base = (api_base or os.environ.get("XECGUARD_API_BASE") or _DEFAULT_API_BASE).rstrip("/") 

105 

106 self.xecguard_model = xecguard_model or _DEFAULT_MODEL 

107 self.policy_names = policy_names 

108 

109 if block_on_error is None: 

110 env: Final = os.environ.get("XECGUARD_BLOCK_ON_ERROR", "true") 

111 self.block_on_error = env.lower() in ( 

112 "true", 

113 "1", 

114 "yes", 

115 ) 

116 else: 

117 self.block_on_error = block_on_error 

118 

119 self.grounding_strictness = grounding_strictness or _DEFAULT_GROUNDING_STRICTNESS 

120 

121 self.async_handler = get_async_httpx_client( 

122 llm_provider=httpxSpecialProvider.GuardrailCallback, 

123 ) 

124 

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

126 

127 super().__init__(**kwargs) 

128 

129 @staticmethod 

130 def get_config_model() -> type["GuardrailConfigModel"] | None: 

131 from litellm.types.proxy.guardrails.guardrail_hooks.xecguard import ( 

132 XecGuardConfigModel, 

133 ) 

134 

135 return XecGuardConfigModel 

136 

137 @classmethod 

138 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

139 return [ 

140 GuardrailEventHooks.pre_call, 

141 GuardrailEventHooks.during_call, 

142 GuardrailEventHooks.post_call, 

143 GuardrailEventHooks.logging_only, 

144 ] 

145 

146 @log_guardrail_information 

147 async def apply_guardrail( 

148 self, 

149 inputs: GenericGuardrailAPIInputs, 

150 request_data: dict, 

151 input_type: Literal["request", "response"], 

152 logging_obj: Optional["LiteLLMLoggingObj"] = None, 

153 ) -> GenericGuardrailAPIInputs: 

154 messages: Final = self._build_full_history( 

155 request_data=request_data, 

156 inputs=inputs, 

157 input_type=input_type, 

158 ) 

159 if not messages: 

160 return inputs 

161 

162 scan_type: Final = "input" if input_type == "request" else "response" 

163 scan_result: Final = await self._call_scan(messages=messages, scan_type=scan_type) 

164 if scan_result is None: 

165 return inputs 

166 

167 if scan_result.get("decision") == "UNSAFE": 

168 raise HTTPException( 

169 status_code=400, 

170 detail={ 

171 "error": self._format_scan_block_message(scan_result), 

172 "guardrail_name": self.guardrail_name or "xecguard", 

173 "xecguard_response": scan_result, 

174 }, 

175 ) 

176 

177 if input_type == "response": 

178 documents: Final = self._extract_grounding_documents(request_data) 

179 if documents: 

180 grounding_result: Final = await self._call_grounding( 

181 messages=messages, 

182 documents=documents, 

183 ) 

184 if grounding_result is not None and grounding_result.get("decision") == "UNSAFE": 

185 raise HTTPException( 

186 status_code=400, 

187 detail={ 

188 "error": self._format_grounding_block_message(grounding_result), 

189 "guardrail_name": self.guardrail_name or "xecguard", 

190 "xecguard_response": grounding_result, 

191 }, 

192 ) 

193 

194 return inputs 

195 

196 async def async_logging_hook( 

197 self, 

198 kwargs: dict, 

199 result: object, 

200 call_type: str, 

201 ) -> tuple[dict, object]: 

202 """Observe-only scan for logging_only mode. 

203 

204 Never blocks, never raises - all errors are swallowed. Records a 

205 StandardLoggingGuardrailInformation entry so the scan decision 

206 reaches downstream loggers (Langfuse, DataDog, etc.). 

207 """ 

208 if ( 

209 isinstance(kwargs, dict) 

210 and "litellm_params" in kwargs 

211 and "metadata" in kwargs["litellm_params"] 

212 and "standard_logging_guardrail_information" in kwargs["litellm_params"]["metadata"] 

213 and kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] 

214 ): 

215 return kwargs, result 

216 

217 start_time: Final = datetime.now() 

218 try: 

219 assistant_text: Final = self._extract_assistant_text_from_response(result) 

220 request_data: Final = {**kwargs} 

221 if assistant_text is not None: 

222 request_data["response"] = result 

223 messages = self._build_full_history( 

224 request_data=request_data, 

225 inputs={}, 

226 input_type="response", 

227 ) 

228 scan_type = "response" 

229 else: 

230 messages = self._build_full_history( 

231 request_data=request_data, 

232 inputs={}, 

233 input_type="request", 

234 ) 

235 scan_type = "input" 

236 

237 if not messages: 

238 return kwargs, result 

239 

240 scan_result: Final = await self._call_scan( 

241 messages=messages, 

242 scan_type=scan_type, 

243 suppress_errors=True, 

244 ) 

245 if scan_result is None: 

246 return kwargs, result 

247 

248 guardrail_status: Final[GuardrailStatus] = ( 

249 "guardrail_intervened" if scan_result.get("decision") == "UNSAFE" else "success" 

250 ) 

251 end_time: Final = datetime.now() 

252 slg: Final = StandardLoggingGuardrailInformation( 

253 guardrail_name=self.guardrail_name or "xecguard", 

254 guardrail_mode=GuardrailEventHooks.logging_only, 

255 guardrail_response=_sanitize_scan_result_for_logging(scan_result), 

256 guardrail_status=guardrail_status, 

257 start_time=start_time.timestamp(), 

258 end_time=end_time.timestamp(), 

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

260 masked_entity_count=None, 

261 ) 

262 existing: Final = kwargs["standard_logging_object"].get("guardrail_information") 

263 if isinstance(existing, list): 

264 existing.append(slg) 

265 else: 

266 kwargs["standard_logging_object"]["guardrail_information"] = [slg] 

267 

268 except Exception as exc: 

269 verbose_proxy_logger.debug( 

270 "XecGuard logging_only swallowed exception: %s", 

271 str(exc), 

272 ) 

273 return kwargs, result 

274 

275 def logging_hook( 

276 self, 

277 kwargs: dict, 

278 result: object, 

279 call_type: str, 

280 ) -> tuple[dict, object]: 

281 """Sync counterpart to ``async_logging_hook``. 

282 

283 Runs the async version on an available loop, swallowing every 

284 exception. Mirrors the pattern used by the Presidio guardrail 

285 for sync logging callbacks. 

286 """ 

287 try: 

288 try: 

289 loop = asyncio.get_event_loop() 

290 except RuntimeError: 

291 loop = asyncio.new_event_loop() 

292 asyncio.set_event_loop(loop) 

293 if loop.is_running(): 

294 return kwargs, result 

295 loop.run_until_complete(self.async_logging_hook(kwargs=kwargs, result=result, call_type=call_type)) 

296 except Exception as exc: 

297 verbose_proxy_logger.debug( 

298 "XecGuard sync logging_hook swallowed exception: %s", 

299 str(exc), 

300 ) 

301 return kwargs, result 

302 

303 # ------------------------------------------------------------------ 

304 # HTTP helpers 

305 # ------------------------------------------------------------------ 

306 

307 async def _call_scan( 

308 self, 

309 messages: list[dict], 

310 scan_type: str, 

311 suppress_errors: bool = False, 

312 ) -> dict | None: 

313 payload: Final[dict[str, object]] = { 

314 "model": self.xecguard_model, 

315 "scan_type": scan_type, 

316 "messages": messages, 

317 "policy_names": (self.policy_names if self.policy_names else _DEFAULT_POLICIES), 

318 } 

319 return await self._post( 

320 path=_SCAN_ENDPOINT, 

321 payload=payload, 

322 suppress_errors=suppress_errors, 

323 ) 

324 

325 async def _call_grounding( 

326 self, 

327 messages: list[dict], 

328 documents: list[dict], 

329 ) -> dict | None: 

330 prompt: Final = self._extract_last_text_by_role(messages, "user") 

331 response_text: Final = self._extract_last_text_by_role(messages, "assistant") 

332 if prompt is None or response_text is None: 

333 return None 

334 payload: Final = { 

335 "model": self.xecguard_model, 

336 "prompt": prompt, 

337 "response": response_text, 

338 "documents": documents, 

339 "strictness": self.grounding_strictness, 

340 } 

341 return await self._post(path=_GROUNDING_ENDPOINT, payload=payload) 

342 

343 async def _post( 

344 self, 

345 path: str, 

346 payload: dict, 

347 suppress_errors: bool = False, 

348 ) -> dict | None: 

349 endpoint: Final = f"{self.api_base}{path}" 

350 verbose_proxy_logger.debug( 

351 "XecGuard: POST %s payload_keys=%s", 

352 endpoint, 

353 list(payload.keys()), 

354 ) 

355 try: 

356 response: Final = await self.async_handler.post( 

357 url=endpoint, 

358 headers={ 

359 "Authorization": f"Bearer {self.api_key}", 

360 "Content-Type": "application/json", 

361 }, 

362 json=payload, 

363 timeout=10.0, 

364 ) 

365 response.raise_for_status() 

366 return response.json() 

367 except Exception as exc: 

368 verbose_proxy_logger.error("XecGuard API error: %s", str(exc)) 

369 if suppress_errors: 

370 return None 

371 if self.block_on_error: 

372 raise HTTPException( 

373 status_code=400, 

374 detail={ 

375 "error": (f"XecGuard API unreachable (block_on_error=True): {exc}"), 

376 "guardrail_name": self.guardrail_name or "xecguard", 

377 }, 

378 ) from exc 

379 return None 

380 

381 # ------------------------------------------------------------------ 

382 # Message-assembly helpers (respect the full-history requirement) 

383 # ------------------------------------------------------------------ 

384 

385 def _build_full_history( 

386 self, 

387 request_data: dict, 

388 inputs: GenericGuardrailAPIInputs, 

389 input_type: str, 

390 ) -> list[dict]: 

391 """Assemble the full message list that will be sent to XecGuard. 

392 

393 Always reads from ``request_data['messages']`` so the framework's 

394 optional ``skip_system_message_in_guardrail`` filter cannot strip 

395 system prompts. Synthesises a trailing user/assistant message when 

396 the request data is incomplete. 

397 """ 

398 raw_messages: Final = request_data.get("messages") or [] 

399 messages: Final[list[dict]] = [self._normalize_message(m) for m in raw_messages if isinstance(m, dict)] 

400 

401 if input_type == "request": 

402 if not messages: 

403 return [] 

404 if messages[-1].get("role") != "user": 

405 synthesized: Final = self._synthesize_user_from_inputs(inputs) 

406 if synthesized is None: 

407 return [] 

408 messages.append(synthesized) 

409 return messages 

410 

411 # input_type == "response" 

412 assistant_text: Final = self._extract_assistant_text_from_response(request_data.get("response")) 

413 if assistant_text is None: 

414 return [] 

415 messages.append({"role": "assistant", "content": assistant_text}) 

416 return messages 

417 

418 @staticmethod 

419 def _normalize_message(message: dict) -> dict: 

420 """Flatten multimodal content to a plain string for XecGuard.""" 

421 role: Final = message.get("role") or "user" 

422 content: Final = message.get("content") 

423 if isinstance(content, str): 

424 return {"role": role, "content": content} 

425 if isinstance(content, list): 

426 parts: Final[list[str]] = [] 

427 for item in content: 

428 if isinstance(item, dict) and item.get("type") == "text": 

429 text = item.get("text") 

430 if isinstance(text, str): 

431 parts.append(text) 

432 return {"role": role, "content": "\n".join(parts)} 

433 return {"role": role, "content": ""} 

434 

435 @staticmethod 

436 def _synthesize_user_from_inputs(inputs: object) -> dict | None: 

437 if not isinstance(inputs, dict): 

438 return None 

439 texts: Final = inputs.get("texts") 

440 if not texts: 

441 return None 

442 joined: Final = "\n".join(t for t in texts if isinstance(t, str) and t) 

443 if not joined: 

444 return None 

445 return {"role": "user", "content": joined} 

446 

447 @staticmethod 

448 def _extract_last_text_by_role(messages: list[dict], role: str) -> str | None: 

449 for message in reversed(messages): 

450 if message.get("role") == role: 

451 content = message.get("content") 

452 if isinstance(content, str) and content: 

453 return content 

454 return None 

455 return None 

456 

457 @staticmethod 

458 def _extract_assistant_text_from_response(response: Any) -> str | None: 

459 if response is None: 

460 return None 

461 choices = None 

462 if hasattr(response, "choices"): 

463 choices = response.choices 

464 elif isinstance(response, dict): 

465 choices = response.get("choices") 

466 if not choices: 

467 return None 

468 text_parts: Final[list[str]] = [] 

469 for choice in choices: 

470 content = XecGuardGuardrail._extract_choice_content(choice) 

471 text = XecGuardGuardrail._content_to_text(content) 

472 if text: 

473 text_parts.append(text) 

474 return "\n".join(text_parts) or None 

475 

476 @staticmethod 

477 def _extract_choice_content(choice: Any) -> Any: 

478 if hasattr(choice, "message"): 

479 message = choice.message 

480 elif isinstance(choice, dict): 

481 message = choice.get("message") 

482 else: 

483 return None 

484 if message is None: 

485 return None 

486 if hasattr(message, "content"): 

487 return message.content 

488 if isinstance(message, dict): 

489 return message.get("content") 

490 return None 

491 

492 @staticmethod 

493 def _content_to_text(content: object) -> str | None: 

494 if isinstance(content, str) and content: 

495 return content 

496 if isinstance(content, list): 

497 parts: Final = [ 

498 item.get("text") 

499 for item in content 

500 if isinstance(item, dict) and item.get("type") == "text" and isinstance(item.get("text"), str) 

501 ] 

502 joined: Final = "\n".join(p for p in parts if p) 

503 return joined or None 

504 return None 

505 

506 # ------------------------------------------------------------------ 

507 # Grounding document extraction 

508 # ------------------------------------------------------------------ 

509 

510 @staticmethod 

511 def _extract_grounding_documents(request_data: dict) -> list[dict]: 

512 metadata: Final = request_data.get("metadata") or request_data.get("litellm_metadata") 

513 if not isinstance(metadata, dict): 

514 return [] 

515 raw_docs: Final = metadata.get(_METADATA_GROUNDING_KEY) 

516 if not isinstance(raw_docs, list) or not raw_docs: 

517 return [] 

518 valid_docs: Final[list[dict]] = [] 

519 for doc in raw_docs: 

520 if ( 

521 isinstance(doc, dict) 

522 and isinstance(doc.get("document_id"), str) 

523 and isinstance(doc.get("context"), str) 

524 ): 

525 valid_docs.append( 

526 { 

527 "document_id": doc["document_id"], 

528 "context": doc["context"], 

529 } 

530 ) 

531 else: 

532 verbose_proxy_logger.debug( 

533 "XecGuard: dropping malformed grounding document: %r", 

534 doc, 

535 ) 

536 return valid_docs 

537 

538 # ------------------------------------------------------------------ 

539 # Error-message formatting 

540 # ------------------------------------------------------------------ 

541 

542 @staticmethod 

543 def _format_scan_block_message(result: dict) -> str: 

544 trace_id: Final = result.get("trace_id", "") 

545 violations = result.get("xecguard_result") 

546 if not isinstance(violations, list): 

547 violations = [] 

548 seen: Final[list[str]] = [] 

549 for v in violations: 

550 if not isinstance(v, dict): 

551 continue 

552 name = v.get("violated_policy_name") 

553 if isinstance(name, str) and name and name not in seen: 

554 seen.append(name) 

555 policies: Final = ",".join(seen) if seen else "unknown" 

556 rationale = "" 

557 for v in violations: 

558 if isinstance(v, dict): 

559 candidate = v.get("rationale") 

560 if isinstance(candidate, str) and candidate: 

561 rationale = candidate[:_RATIONALE_TRUNCATE_CHARS] 

562 break 

563 return f"Blocked by XecGuard: policies=[{policies}] trace_id={trace_id} rationale={rationale}" 

564 

565 @staticmethod 

566 def _format_grounding_block_message(result: dict) -> str: 

567 trace_id: Final = result.get("trace_id", "") 

568 detail: Final = result.get("xecguard_result") 

569 rules: list[str] = [] 

570 rationale = "" 

571 if isinstance(detail, dict): 

572 raw_rules: Final = detail.get("violated_rules_list") 

573 if isinstance(raw_rules, list): 

574 rules = [r for r in raw_rules if isinstance(r, str)] 

575 candidate: Final = detail.get("rationale") 

576 if isinstance(candidate, str): 

577 rationale = candidate[:_RATIONALE_TRUNCATE_CHARS] 

578 rules_str: Final = ",".join(rules) if rules else "unknown" 

579 return f"Blocked by XecGuard grounding: rules=[{rules_str}] trace_id={trace_id} rationale={rationale}"