Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py: 16%
317 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# Pillar Security Guardrails
4# https://www.pillar.security/
5#
6# +-------------------------------------------------------------+
8# Standard library imports
9import json
10import os
11from typing import TYPE_CHECKING, Any, Final, Literal, Protocol
12from urllib.parse import quote
14# Third-party imports
15from fastapi import HTTPException
16from typing_extensions import NotRequired, ReadOnly, TypedDict
18# LiteLLM imports
19from litellm import DualCache
20from litellm._logging import verbose_proxy_logger
21from litellm._version import version as litellm_version
22from litellm.integrations.custom_guardrail import (
23 CustomGuardrail,
24 log_guardrail_information,
25)
26from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
27from litellm.llms.custom_httpx.http_handler import (
28 get_async_httpx_client,
29 httpxSpecialProvider,
30)
31from litellm.proxy._types import UserAPIKeyAuth
32from litellm.proxy.common_utils.callback_utils import (
33 TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY,
34 add_guardrail_to_applied_guardrails_header,
35 get_metadata_variable_name_from_kwargs,
36)
37from litellm.types.guardrails import GuardrailEventHooks
38from litellm.types.utils import LLMResponseTypes
40if TYPE_CHECKING: 40 ↛ 41line 40 didn't jump to line 41 because the condition on line 40 was never true
41 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
43MAX_PILLAR_HEADER_VALUE_BYTES: Final = 8 * 1024
46class _PillarProtectResponse(TypedDict):
47 """Body returned by Pillar's `/api/v1/protect` endpoint."""
49 flagged: ReadOnly[NotRequired[bool]]
50 session_id: ReadOnly[NotRequired[str]]
51 scanners: ReadOnly[NotRequired[dict[str, object]]]
52 evidence: ReadOnly[NotRequired[list[object]]]
53 masked_session_messages: ReadOnly[NotRequired[list[object]]]
56class _PillarProtectHTTPResponse(Protocol):
57 def raise_for_status(self) -> object: ... 57 ↛ exitline 57 didn't return from function 'raise_for_status' because
59 def json(self) -> _PillarProtectResponse: ... 59 ↛ exitline 59 didn't return from function 'json' because
62class _PillarProtectHTTPClient(Protocol):
63 async def post( 63 ↛ exitline 63 didn't return from function 'post' because
64 self,
65 *,
66 url: str,
67 headers: dict[str, str],
68 json: dict[str, object],
69 timeout: float,
70 ) -> _PillarProtectHTTPResponse: ...
73def _encode_json_for_header(data: object) -> str:
74 """
75 JSON-serialize and URL-encode data for safe header transmission.
76 """
77 json_payload: Final = json.dumps(data, ensure_ascii=False, separators=(",", ":"))
78 return quote(json_payload, safe="")
81def _truncate_evidence_payload(
82 evidence: object, max_bytes: int = MAX_PILLAR_HEADER_VALUE_BYTES
83) -> tuple[object, str, bool]:
84 """
85 Truncate evidence payload so the encoded header value stays within max_bytes.
87 Returns:
88 truncated_evidence: Evidence list/value after truncation
89 encoded_value: URL-encoded JSON string for header
90 was_truncated: Whether truncation occurred
91 """
92 if not isinstance(evidence, list):
93 encoded = _encode_json_for_header(evidence)
94 if len(encoded.encode("utf-8")) <= max_bytes:
95 return evidence, encoded, False
96 truncated_value: Final = "[truncated]"
97 return truncated_value, _encode_json_for_header(truncated_value), True
99 truncated: Final[list[object]] = []
100 encoded = _encode_json_for_header(truncated)
101 truncated_flag = False
103 for entry in evidence:
104 working_entry: object
105 if isinstance(entry, dict):
106 working_entry = dict(entry)
107 else:
108 working_entry = entry
110 truncated.append(working_entry)
111 encoded = _encode_json_for_header(truncated)
113 if len(encoded.encode("utf-8")) <= max_bytes:
114 continue
116 truncated_flag = True
117 if isinstance(working_entry, dict):
118 evidence_text = str(working_entry.get("evidence", ""))
119 if evidence_text:
120 step = max(1, len(evidence_text) // 2)
121 while len(encoded.encode("utf-8")) > max_bytes and evidence_text:
122 evidence_text = evidence_text[:-step] if len(evidence_text) > step else evidence_text[:-1]
123 step = max(1, step // 2)
124 truncated_text = f"{evidence_text}...[truncated]" if evidence_text else "[truncated]"
125 working_entry["evidence"] = truncated_text
126 working_entry["evidence_truncated"] = True
127 encoded = _encode_json_for_header(truncated)
129 if len(encoded.encode("utf-8")) <= max_bytes:
130 continue
132 truncated.pop()
133 encoded = _encode_json_for_header(truncated)
135 return truncated, encoded, truncated_flag
138def build_pillar_response_headers(metadata_store: dict[str, object]) -> dict[str, str]:
139 """
140 Create URL-safe Pillar response headers and apply truncation metadata.
141 """
142 headers: Final[dict[str, str]] = {}
144 if "pillar_flagged" in metadata_store:
145 headers["x-pillar-flagged"] = str(metadata_store["pillar_flagged"]).lower()
147 if "pillar_scanners" in metadata_store:
148 headers["x-pillar-scanners"] = _encode_json_for_header(metadata_store["pillar_scanners"])
150 if "pillar_evidence" in metadata_store:
151 truncated_evidence, encoded_value, truncated_flag = _truncate_evidence_payload(
152 metadata_store["pillar_evidence"]
153 )
154 metadata_store["pillar_evidence"] = truncated_evidence
155 if truncated_flag:
156 metadata_store["pillar_evidence_truncated"] = True
157 headers["x-pillar-evidence"] = encoded_value
159 if "pillar_session_id_response" in metadata_store:
160 headers["x-pillar-session-id"] = quote(str(metadata_store["pillar_session_id_response"]), safe="")
162 if headers:
163 metadata_store["pillar_response_headers"] = headers
164 metadata_store[TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY] = True
166 return headers
169# Exception classes
170class PillarGuardrailMissingSecrets(Exception):
171 """Exception raised when Pillar API key is missing."""
174class PillarGuardrailAPIError(Exception):
175 """Exception raised when there's an error calling the Pillar API."""
178# Main guardrail class
179class PillarGuardrail(CustomGuardrail):
180 """
181 Pillar Security Guardrail for LiteLLM.
183 Provides comprehensive AI security scanning for input prompts and output responses
184 using the Pillar Security API.
185 """
187 SUPPORTED_ON_FLAGGED_ACTIONS = ["block", "monitor", "mask"]
188 DEFAULT_ON_FLAGGED_ACTION = "monitor"
189 SUPPORTED_FALLBACK_ACTIONS = ["allow", "block"]
190 DEFAULT_FALLBACK_ACTION = "allow"
191 BASE_API_URL = "https://api.pillar.security"
192 DEFAULT_TIMEOUT = 5.0 # 5 seconds - fast failure detection with graceful degradation
194 def __init__(
195 self,
196 guardrail_name: str | None = "pillar-security",
197 api_key: str | None = None,
198 api_base: str | None = None,
199 on_flagged_action: str | None = None,
200 async_mode: bool | None = None,
201 persist_session: bool | None = None,
202 include_scanners: bool | None = None,
203 include_evidence: bool | None = None,
204 fallback_on_error: str | None = None,
205 timeout: float | None = None,
206 **kwargs,
207 ) -> None:
208 """
209 Initialize the Pillar guardrail.
211 Args:
212 guardrail_name: Name of the guardrail instance
213 api_key: Pillar API key
214 api_base: Pillar API base URL
215 on_flagged_action: Action to take when content is flagged ('block' or 'monitor')
216 fallback_on_error: Action when API errors occur ('allow' or 'block')
217 timeout: Timeout for API calls in seconds
218 **kwargs: Additional arguments passed to parent class
220 Note:
221 LiteLLM virtual key context (user_id, team_id, key_alias, etc.) is always
222 automatically passed as X-LiteLLM-* headers to enable application/user tracking.
223 """
224 self.async_handler: _PillarProtectHTTPClient = get_async_httpx_client(
225 llm_provider=httpxSpecialProvider.GuardrailCallback
226 )
227 self.api_key = api_key or os.environ.get("PILLAR_API_KEY")
229 if self.api_key is None:
230 msg: Final = (
231 "Couldn't get Pillar API key, either set the `PILLAR_API_KEY` in the environment or "
232 "pass it as a parameter to the guardrail in the config file"
233 )
234 raise PillarGuardrailMissingSecrets(msg)
236 self.api_base = api_base or os.getenv("PILLAR_API_BASE") or self.BASE_API_URL
238 # Validate and set on_flagged_action
239 action = on_flagged_action or os.environ.get("PILLAR_ON_FLAGGED_ACTION")
240 if action and action in self.SUPPORTED_ON_FLAGGED_ACTIONS:
241 self.on_flagged_action = action
242 else:
243 if action:
244 verbose_proxy_logger.warning("Invalid action '%s', using default", action)
245 self.on_flagged_action = self.DEFAULT_ON_FLAGGED_ACTION
247 verbose_proxy_logger.debug("Pillar Guardrail: Initialized with on_flagged_action: %s", self.on_flagged_action)
249 self.async_mode = self._resolve_bool_config(
250 provided_value=async_mode,
251 env_var="PILLAR_ASYNC",
252 default=None,
253 setting_name="async_mode",
254 )
255 self.persist_session = self._resolve_bool_config(
256 provided_value=persist_session,
257 env_var="PILLAR_PERSIST",
258 default=None,
259 setting_name="persist_session",
260 )
261 self.include_scanners = self._resolve_bool_config(
262 provided_value=include_scanners,
263 env_var="PILLAR_INCLUDE_SCANNERS",
264 default=True,
265 setting_name="include_scanners",
266 )
267 self.include_evidence = self._resolve_bool_config(
268 provided_value=include_evidence,
269 env_var="PILLAR_INCLUDE_EVIDENCE",
270 default=True,
271 setting_name="include_evidence",
272 )
274 # Validate and set fallback_on_error
275 action = fallback_on_error or os.environ.get("PILLAR_FALLBACK_ON_ERROR")
276 if action and action in self.SUPPORTED_FALLBACK_ACTIONS:
277 self.fallback_on_error = action
278 else:
279 if action:
280 verbose_proxy_logger.warning(
281 "Invalid fallback action '%s', using default '%s'", action, self.DEFAULT_FALLBACK_ACTION
282 )
283 self.fallback_on_error = self.DEFAULT_FALLBACK_ACTION
285 verbose_proxy_logger.debug("Pillar Guardrail: Initialized with fallback_on_error: %s", self.fallback_on_error)
287 # Set timeout with graceful fallback on invalid configuration
288 if timeout is not None:
289 self.timeout = timeout
290 else:
291 try:
292 self.timeout = float(os.environ.get("PILLAR_TIMEOUT", str(self.DEFAULT_TIMEOUT)))
293 except (ValueError, TypeError):
294 verbose_proxy_logger.warning(
295 "Pillar Guardrail: Invalid PILLAR_TIMEOUT value '%s', falling back to default %ss",
296 os.environ.get("PILLAR_TIMEOUT"),
297 self.DEFAULT_TIMEOUT,
298 )
299 self.timeout = self.DEFAULT_TIMEOUT
301 super().__init__(
302 guardrail_name=guardrail_name,
303 supported_event_hooks=list(self.get_supported_event_hooks()),
304 **kwargs,
305 )
307 # =========================================================================
308 # PUBLIC HOOK METHODS (Main Interface)
309 # =========================================================================
311 @log_guardrail_information
312 async def async_pre_call_hook(
313 self,
314 user_api_key_dict: UserAPIKeyAuth,
315 cache: DualCache,
316 data: dict,
317 call_type: Literal[
318 "completion",
319 "text_completion",
320 "embeddings",
321 "image_generation",
322 "moderation",
323 "audio_transcription",
324 "pass_through_endpoint",
325 "rerank",
326 "mcp_call",
327 "anthropic_messages",
328 ],
329 ) -> Exception | str | dict | None:
330 """
331 Pre-call hook to scan the request for security threats before sending to LLM.
333 Args:
334 user_api_key_dict: User API key authentication info
335 cache: LiteLLM cache instance
336 data: Request data
337 call_type: Type of LLM call
339 Returns:
340 Original data if safe, raises HTTPException if blocked
342 Raises:
343 HTTPException: If request should be blocked due to security threats
344 """
345 event_type: Final = GuardrailEventHooks.pre_call
346 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
347 verbose_proxy_logger.debug("Pillar Guardrail: Pre-call scanning disabled for %s", self.guardrail_name)
348 return data
350 verbose_proxy_logger.debug("Pillar Guardrail: Pre-call hook")
351 result: Final = await self.run_pillar_guardrail(data, user_api_key_dict)
353 # Add guardrail name to response headers
354 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
356 return result
358 @log_guardrail_information
359 async def async_moderation_hook(
360 self,
361 data: dict,
362 user_api_key_dict: UserAPIKeyAuth,
363 call_type: Literal[
364 "completion",
365 "embeddings",
366 "image_generation",
367 "moderation",
368 "audio_transcription",
369 "responses",
370 "mcp_call",
371 "anthropic_messages",
372 ],
373 ) -> Exception | str | dict | None:
374 """
375 During-call hook to scan the request in parallel with LLM processing.
377 Args:
378 data: Request data
379 user_api_key_dict: User API key authentication info
380 call_type: Type of LLM call
382 Returns:
383 Original data if safe, raises HTTPException if blocked
385 Raises:
386 HTTPException: If request should be blocked due to security threats
387 """
388 event_type: Final = GuardrailEventHooks.during_call
389 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
390 verbose_proxy_logger.debug("Pillar Guardrail: During-call scanning disabled for %s", self.guardrail_name)
391 return data
393 verbose_proxy_logger.debug("Pillar Guardrail: During-call moderation hook")
394 result: Final = await self.run_pillar_guardrail(data, user_api_key_dict)
396 # Add guardrail name to response headers
397 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
399 return result
401 @log_guardrail_information
402 async def async_post_call_success_hook(
403 self,
404 data: dict,
405 user_api_key_dict: UserAPIKeyAuth,
406 response: LLMResponseTypes,
407 ) -> LLMResponseTypes:
408 """
409 Post-call hook to scan LLM responses before returning to user.
411 Args:
412 data: Original request data
413 user_api_key_dict: User API key authentication info
414 response: LLM response to scan
416 Returns:
417 Original response if safe, raises HTTPException if blocked
419 Raises:
420 HTTPException: If response should be blocked due to security threats
421 """
422 event_type: Final = GuardrailEventHooks.post_call
423 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
424 verbose_proxy_logger.debug("Pillar Guardrail: Post-call scanning disabled for %s", self.guardrail_name)
425 return response
427 verbose_proxy_logger.debug("Pillar Guardrail: Post-call hook")
429 # Extract response messages in the format Pillar expects
430 response_dict = response.model_dump() if hasattr(response, "model_dump") else {}
431 response_messages: Final = [
432 choice.get("message") for choice in response_dict.get("choices", []) if choice.get("message")
433 ]
435 if not response_messages:
436 verbose_proxy_logger.debug("Pillar Guardrail: No response content to scan, skipping post-call analysis")
437 return response
439 # Create complete conversation: original messages + response messages
440 post_call_data: Final = data.copy()
441 post_call_data["messages"] = data.get("messages", []) + response_messages
443 # Reuse the existing guardrail logic - zero duplication!
444 await self.run_pillar_guardrail(post_call_data, user_api_key_dict)
446 # Add guardrail name to response headers
447 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
449 return response
451 # =========================================================================
452 # CORE LOGIC METHOD
453 # =========================================================================
455 async def run_pillar_guardrail(self, data: dict, user_api_key_dict: UserAPIKeyAuth) -> dict:
456 """
457 Core method to run the Pillar guardrail scan.
459 Args:
460 data: Request data containing messages and metadata
461 user_api_key_dict: User API key authentication info containing key context
463 Returns:
464 Original data if safe or in monitor mode
466 Raises:
467 HTTPException: If content is flagged and action is 'block', or if API fails and fallback_on_error is 'block'
468 """
469 # Check if messages are present
470 if not data.get("messages"):
471 verbose_proxy_logger.debug("Pillar Guardrail: No messages detected, bypassing security scan")
472 return data
474 try:
475 headers: Final = self._prepare_headers(user_api_key_dict)
476 payload: Final = self._prepare_payload(data)
478 response: Final = await self._call_pillar_api(
479 headers=headers,
480 payload=payload,
481 )
483 # Process the response - handles blocking or monitoring
484 self._process_pillar_response(response, data)
485 return data
487 except Exception as e:
488 # If it's already an HTTPException from content being flagged, re-raise it
489 if isinstance(e, HTTPException):
490 raise e
492 # Handle API communication errors based on fallback_on_error setting
493 verbose_proxy_logger.error("Pillar Guardrail: API communication failed - %s", e)
495 return self._handle_api_error(e, data)
497 # =========================================================================
498 # PRIVATE HELPER METHODS (In logical order of usage)
499 # =========================================================================
501 def _handle_api_error(self, error: Exception, data: dict) -> dict:
502 """
503 Handle API errors based on fallback_on_error configuration.
505 Args:
506 error: The exception that occurred during API communication
507 data: Original request data
509 Returns:
510 Original data if fallback_on_error is 'allow'
512 Raises:
513 HTTPException: If fallback_on_error is 'block'
514 """
515 if self.fallback_on_error == "allow":
516 verbose_proxy_logger.warning(
517 "Pillar Guardrail: API unavailable, proceeding without scanning (fallback_on_error=allow)"
518 )
519 return data
520 else: # fallback_on_error == "block"
521 verbose_proxy_logger.warning(
522 "Pillar Guardrail: API unavailable, blocking request (fallback_on_error=block)"
523 )
524 raise HTTPException(
525 status_code=503,
526 detail={
527 "error": "Pillar Security Guardrail Unavailable",
528 "message": "Security scanning service is temporarily unavailable and fallback is set to block",
529 "original_error": str(error),
530 },
531 )
533 def _prepare_headers(self, user_api_key_dict: UserAPIKeyAuth) -> dict[str, str]:
534 """
535 Prepare headers for the Pillar API request.
537 Args:
538 user_api_key_dict: User API key authentication info containing key context
540 Returns:
541 Dictionary of headers to send to Pillar API
542 """
543 if not self.api_key:
544 msg: Final = (
545 "Couldn't get Pillar API key, either set the `PILLAR_API_KEY` in the environment or "
546 "pass it as a parameter to the guardrail in the config file"
547 )
548 raise PillarGuardrailMissingSecrets(msg)
550 headers: Final[dict[str, str]] = {
551 "Authorization": f"Bearer {self.api_key}",
552 "Content-Type": "application/json",
553 }
555 # Add Pillar-specific headers based on configuration
556 self._set_bool_header(headers, "plr_scanners", self.include_scanners)
557 self._set_bool_header(headers, "plr_evidence", self.include_evidence)
558 self._set_bool_header(headers, "plr_async", self.async_mode)
559 self._set_bool_header(headers, "plr_persist", self.persist_session)
561 # Always add LiteLLM virtual key context headers (metadata excluded for security)
562 context_mapping: Final = {
563 "X-LiteLLM-Key-Name": user_api_key_dict.key_name,
564 "X-LiteLLM-Key-Alias": user_api_key_dict.key_alias,
565 "X-LiteLLM-User-Id": user_api_key_dict.user_id,
566 "X-LiteLLM-User-Email": user_api_key_dict.user_email,
567 "X-LiteLLM-Team-Id": user_api_key_dict.team_id,
568 "X-LiteLLM-Team-Name": user_api_key_dict.team_alias,
569 "X-LiteLLM-Org-Id": user_api_key_dict.org_id,
570 }
571 for header_name, value in context_mapping.items():
572 if value:
573 headers[header_name] = str(value)
575 return headers
577 def _set_bool_header(self, headers: dict[str, str], header_name: str, value: bool | None) -> None:
578 """Apply a boolean value as a lowercase string HTTP header when provided."""
580 if value is None:
581 return
582 headers[header_name] = "true" if value else "false"
584 def _resolve_bool_config(
585 self,
586 provided_value: bool | str | int | None,
587 env_var: str | None,
588 default: bool | None,
589 setting_name: str,
590 ) -> bool | None:
591 """Resolve configuration precedence: explicit value -> environment -> default."""
593 if provided_value is not None:
594 try:
595 return self._parse_bool_value(provided_value)
596 except ValueError:
597 verbose_proxy_logger.warning(
598 "Pillar Guardrail: Invalid boolean value '%s' for %s, falling back to default.",
599 provided_value,
600 setting_name,
601 )
602 return default
604 if env_var:
605 env_value: Final = os.getenv(env_var)
606 if env_value is not None:
607 try:
608 return self._parse_bool_value(env_value)
609 except ValueError:
610 verbose_proxy_logger.warning(
611 "Pillar Guardrail: Invalid boolean env value '%s' for %s, falling back to default.",
612 env_value,
613 env_var,
614 )
615 return default
617 return default
619 @staticmethod
620 def _parse_bool_value(value: bool | str | int) -> bool:
621 """Normalise various truthy/falsey inputs to a strict boolean."""
623 if isinstance(value, bool):
624 return value
625 if isinstance(value, int):
626 return bool(value)
628 value_str: Final = str(value).strip().lower()
629 if value_str in {"true", "1", "yes", "y", "on"}:
630 return True
631 if value_str in {"false", "0", "no", "n", "off"}:
632 return False
633 raise ValueError(f"Unrecognised boolean value: {value}")
635 def _extract_model_and_provider(self, data: dict) -> tuple[str, str]:
636 """
637 Extract the model and provider from the request data.
639 Args:
640 data: Request data
642 Returns:
643 Tuple of (model_name, provider_name)
644 """
645 model: Final = data.get("model")
646 if not model:
647 return "unknown", "unknown"
649 # Use LiteLLM's standard provider detection and model cleaning
650 try:
651 clean_model, provider, _, _ = get_llm_provider(
652 model=model,
653 custom_llm_provider=data.get("custom_llm_provider"),
654 api_base=data.get("api_base"),
655 api_key=data.get("api_key"),
656 )
657 return clean_model or "unknown", provider or "unknown"
658 except Exception:
659 # Fallback if get_llm_provider fails
660 return (
661 model or "unknown",
662 data.get("custom_llm_provider") or data.get("provider") or "unknown",
663 )
665 def _prepare_payload(self, data: dict) -> dict[str, Any]:
666 """
667 Prepare the payload for the Pillar API request following the /api/v1/protect contract.
669 This method supports multi-modal content (images, files, audio, video, etc.) as messages
670 are passed through without modification. The messages array can contain any OpenAI-compatible
671 message structure including:
672 - Text content (string)
673 - Multi-modal content blocks (image_url, image_file, audio, video, document, file)
674 - Attachments
675 - Tool calls
677 Args:
678 data: Request data
680 Returns:
681 Formatted payload for Pillar API
682 """
683 messages: Final = data.get("messages", [])
684 tools: Final = data.get("tools", [])
685 metadata: Final = {
686 "source": "litellm",
687 "version": litellm_version,
688 }
690 # Build payload following Pillar API format
691 payload: Final = {
692 "messages": messages,
693 "tools": tools,
694 "metadata": metadata,
695 }
697 # User ID: use LiteLLM user field
698 user_id: Final = data.get("user")
699 if user_id:
700 payload["user_id"] = user_id
702 # Session ID: use metadata.pillar_session_id if provided
703 session_id: Final = data.get("metadata", {}).get("pillar_session_id")
704 if session_id:
705 payload["session_id"] = session_id
707 # Extract model and provider from actual request data
708 model, provider = self._extract_model_and_provider(data)
709 payload["model"] = model
710 payload["provider"] = provider
712 verbose_proxy_logger.debug(
713 "Pillar Guardrail: Request context - user=%s, session=%s, model=%s, provider=%s",
714 user_id,
715 session_id,
716 model,
717 provider,
718 )
719 return payload
721 async def _call_pillar_api(self, headers: dict[str, str], payload: dict[str, Any]) -> _PillarProtectResponse:
722 """
723 Call the Pillar API and return the response.
725 Args:
726 headers: HTTP headers for the request
727 payload: Request payload
729 Returns:
730 Pillar API response as dictionary
731 """
732 verbose_proxy_logger.debug(
733 "Pillar Guardrail: Scanning %s messages for security threats", len(payload.get("messages", []))
734 )
735 response: Final = await self.async_handler.post(
736 url=f"{self.api_base}/api/v1/protect",
737 headers=headers,
738 json=payload,
739 timeout=self.timeout,
740 )
741 response.raise_for_status()
742 res: Final = response.json()
744 flagged: Final = res.get("flagged")
745 session_id: Final = res.get("session_id")
746 verbose_proxy_logger.debug("Pillar Guardrail: Analysis complete - flagged=%s, session=%s", flagged, session_id)
747 return res
749 def _process_pillar_response(self, pillar_response: _PillarProtectResponse, original_data: dict) -> None:
750 """
751 Process the Pillar API response and handle detections based on configuration.
753 Args:
754 pillar_response: Response from Pillar API
755 original_data: Original request data (modified in-place with session info)
757 Raises:
758 HTTPException: If content is flagged and action is 'block'
759 """
760 if not pillar_response:
761 return
763 flagged: Final = pillar_response.get("flagged", False)
765 metadata_field: Final = get_metadata_variable_name_from_kwargs(original_data)
766 if metadata_field not in original_data or not isinstance(original_data.get(metadata_field), dict):
767 original_data[metadata_field] = {}
768 metadata_store: Final = original_data[metadata_field]
770 # Backwards compatibility - ensure metadata alias exists when different key used
771 if metadata_field != "metadata":
772 if "metadata" not in original_data or not isinstance(original_data.get("metadata"), dict):
773 original_data["metadata"] = metadata_store
775 # Store session_id from Pillar response for potential reuse
776 pillar_session_id: Final = pillar_response.get("session_id")
777 if pillar_session_id:
778 verbose_proxy_logger.debug("Pillar Guardrail: Received session_id from server: %s", pillar_session_id)
779 # Store in request metadata for use in subsequent hooks
780 if "pillar_session_id" not in metadata_store:
781 metadata_store["pillar_session_id"] = pillar_session_id
782 metadata_store["pillar_session_id_response"] = pillar_session_id
784 # Always set flagged status and scanner/evidence data for monitor mode
785 metadata_store["pillar_flagged"] = flagged
786 if self.include_scanners:
787 metadata_store["pillar_scanners"] = pillar_response.get("scanners", {})
788 if self.include_evidence:
789 metadata_store["pillar_evidence"] = pillar_response.get("evidence", [])
791 if flagged:
792 verbose_proxy_logger.warning("Pillar Guardrail: Threat detected")
793 if self.on_flagged_action == "block":
794 self._raise_pillar_detection_exception(pillar_response)
795 elif self.on_flagged_action == "mask":
796 verbose_proxy_logger.info("Pillar Guardrail: Masking mode - masking flagged content")
797 masked_messages: Final = pillar_response.get("masked_session_messages", [])
798 if masked_messages:
799 original_data["messages"] = masked_messages
800 else:
801 verbose_proxy_logger.warning(
802 "Pillar Guardrail: Masking requested but no masked_session_messages in response"
803 )
804 elif self.on_flagged_action == "monitor":
805 verbose_proxy_logger.info("Pillar Guardrail: Monitoring mode - allowing flagged content to proceed")
807 build_pillar_response_headers(metadata_store)
809 def _raise_pillar_detection_exception(self, pillar_response: _PillarProtectResponse) -> None:
810 """
811 Raise an HTTPException for Pillar security detections.
813 Args:
814 pillar_response: Response from Pillar API containing detection details
816 Raises:
817 HTTPException: Always raises with security detection details
818 """
819 pillar_response_dict: Final[dict[str, object]] = {
820 "session_id": pillar_response.get("session_id"),
821 }
823 # Conditionally include scanners and evidence based on config
824 if self.include_scanners:
825 pillar_response_dict["scanners"] = pillar_response.get("scanners", {})
826 if self.include_evidence:
827 pillar_response_dict["evidence"] = pillar_response.get("evidence", [])
829 error_detail: Final = {
830 "error": "Blocked by Pillar Security Guardrail",
831 "detection_message": "Security threats detected",
832 "pillar_response": pillar_response_dict,
833 }
835 verbose_proxy_logger.warning("Pillar Guardrail: Request blocked - Security threats detected")
837 raise HTTPException(status_code=400, detail=error_detail)
839 # =========================================================================
840 # STATIC/CLASS METHODS
841 # =========================================================================
843 @staticmethod
844 def get_config_model() -> type["GuardrailConfigModel"] | None:
845 """
846 Get the configuration model for this guardrail.
848 Returns:
849 Pydantic model class for guardrail configuration
850 """
851 from litellm.types.proxy.guardrails.guardrail_hooks.pillar import (
852 PillarGuardrailConfigModel,
853 )
855 return PillarGuardrailConfigModel
857 @classmethod
858 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
859 return [
860 GuardrailEventHooks.pre_call,
861 GuardrailEventHooks.during_call,
862 GuardrailEventHooks.post_call,
863 GuardrailEventHooks.pre_mcp_call,
864 GuardrailEventHooks.during_mcp_call,
865 ]