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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2XecGuard guardrail integration for LiteLLM.
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.
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"""
25import asyncio
26import os
27from datetime import datetime
28from typing import TYPE_CHECKING, Any, Final, Literal, Optional
30from fastapi.exceptions import HTTPException
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)
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 )
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
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]
80class XecGuardMissingCredentials(Exception):
81 pass
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 )
104 self.api_base = (api_base or os.environ.get("XECGUARD_API_BASE") or _DEFAULT_API_BASE).rstrip("/")
106 self.xecguard_model = xecguard_model or _DEFAULT_MODEL
107 self.policy_names = policy_names
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
119 self.grounding_strictness = grounding_strictness or _DEFAULT_GROUNDING_STRICTNESS
121 self.async_handler = get_async_httpx_client(
122 llm_provider=httpxSpecialProvider.GuardrailCallback,
123 )
125 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
127 super().__init__(**kwargs)
129 @staticmethod
130 def get_config_model() -> type["GuardrailConfigModel"] | None:
131 from litellm.types.proxy.guardrails.guardrail_hooks.xecguard import (
132 XecGuardConfigModel,
133 )
135 return XecGuardConfigModel
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 ]
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
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
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 )
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 )
194 return inputs
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.
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
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"
237 if not messages:
238 return kwargs, result
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
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]
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
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``.
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
303 # ------------------------------------------------------------------
304 # HTTP helpers
305 # ------------------------------------------------------------------
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 )
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)
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
381 # ------------------------------------------------------------------
382 # Message-assembly helpers (respect the full-history requirement)
383 # ------------------------------------------------------------------
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.
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)]
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
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
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": ""}
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}
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
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
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
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
506 # ------------------------------------------------------------------
507 # Grounding document extraction
508 # ------------------------------------------------------------------
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
538 # ------------------------------------------------------------------
539 # Error-message formatting
540 # ------------------------------------------------------------------
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}"
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}"