Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py: 18%
365 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
1import asyncio
2import base64
3import os
4from collections.abc import Mapping, Sequence
5from types import MappingProxyType
6from typing import TYPE_CHECKING, Final, Literal, Optional
8import httpx
9from fastapi import HTTPException
10from typing_extensions import ReadOnly, TypedDict
12from litellm._logging import verbose_proxy_logger
13from litellm.exceptions import Timeout as LiteLLMTimeout
14from litellm.integrations.custom_guardrail import (
15 CustomGuardrail,
16 log_guardrail_information,
17)
18from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts, message_with_slot_texts
19from litellm.llms.custom_httpx.http_handler import (
20 get_async_httpx_client,
21 httpxSpecialProvider,
22)
23from litellm.types.guardrails import GuardrailEventHooks
24from litellm.types.llms.openai import AllMessageValues
25from litellm.types.utils import GenericGuardrailAPIInputs
27if TYPE_CHECKING: 27 ↛ 28line 27 didn't jump to line 28 because the condition on line 27 was never true
28 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
29 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
32_SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS: Final = 30.0
33_SANITIZE_FILE_QUEUED_STATUSES: Final = frozenset({"created", "in progress"})
34_PROTECT_ROLES: Final = frozenset({"system", "user", "assistant"})
37class PromptSecurityGuardrailMissingSecrets(Exception):
38 pass
41def _modified_or_original(text: str, verdict: "_ProtectVerdict") -> str:
42 modified_text: Final = verdict.get("modified_text") if verdict.get("action") == "modify" else None
43 return text if modified_text is None else modified_text
46def _inputs_with_structured_messages(
47 inputs: GenericGuardrailAPIInputs, rewritten_messages: Sequence[AllMessageValues] | None
48) -> GenericGuardrailAPIInputs:
49 if rewritten_messages is None:
50 return inputs
51 patched: Final[GenericGuardrailAPIInputs] = {
52 **inputs,
53 "structured_messages": list(rewritten_messages), # mutable-ok: the TypedDict field is declared as a list
54 }
55 return patched
58def _inputs_with_modifications(
59 inputs: GenericGuardrailAPIInputs,
60 modified_texts: list[str],
61 rewritten_messages: Sequence[AllMessageValues] | None,
62) -> GenericGuardrailAPIInputs:
63 if not modified_texts:
64 return _inputs_with_structured_messages(inputs, rewritten_messages)
65 with_texts: Final[GenericGuardrailAPIInputs] = {**inputs, "texts": modified_texts}
66 return _inputs_with_structured_messages(with_texts, rewritten_messages)
69class _ProtectVerdict(TypedDict, total=False):
70 """One side (``prompt`` or ``response``) of an ``/api/protect`` verdict."""
72 action: ReadOnly[str]
73 violations: ReadOnly[Sequence[str]]
74 modified_messages: ReadOnly[Sequence[Mapping[str, object]]]
75 modified_text: ReadOnly[str]
78class _ProtectResult(TypedDict, total=False):
79 prompt: ReadOnly[_ProtectVerdict | None]
80 response: ReadOnly[_ProtectVerdict | None]
83class _ProtectResponse(TypedDict, total=False):
84 result: ReadOnly[_ProtectResult]
87class _SanitizeUploadResponse(TypedDict, total=False):
88 jobId: ReadOnly[str]
91class _SanitizeMetadata(TypedDict, total=False):
92 action: ReadOnly[str]
93 violations: ReadOnly[Sequence[str]]
96class _SanitizeStatusResponse(TypedDict, total=False):
97 """One poll of ``/api/sanitizeFile``."""
99 status: ReadOnly[str]
100 content: ReadOnly[str]
101 metadata: ReadOnly[_SanitizeMetadata]
104class _SanitizeResult(TypedDict):
105 action: ReadOnly[str]
106 content: ReadOnly[str | None]
107 metadata: ReadOnly[_SanitizeMetadata]
108 violations: ReadOnly[Sequence[str]]
111class PromptSecurityGuardrail(CustomGuardrail):
112 @classmethod
113 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
114 return [
115 GuardrailEventHooks.pre_call,
116 GuardrailEventHooks.during_call,
117 GuardrailEventHooks.post_call,
118 ]
120 def __init__(
121 self,
122 api_key: str | None = None,
123 api_base: str | None = None,
124 user: str | None = None,
125 system_prompt: str | None = None,
126 check_tool_results: bool | None = None,
127 streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
128 file_sanitization_timeout: float = _SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS,
129 file_sanitization_fail_open: bool | None = None,
130 block_on_file_modify: bool | None = None,
131 **kwargs,
132 ):
133 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
134 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
135 self.api_key = api_key or os.environ.get("PROMPT_SECURITY_API_KEY")
136 self.api_base = api_base or os.environ.get("PROMPT_SECURITY_API_BASE")
137 self.user = user or os.environ.get("PROMPT_SECURITY_USER")
138 self.system_prompt = system_prompt or os.environ.get("PROMPT_SECURITY_SYSTEM_PROMPT")
140 # Configure whether to check tool/function results for indirect prompt injection
141 # Default: False (Filter out tool/function messages)
142 # True: Transform to "other" role and send to API
143 if check_tool_results is None:
144 check_tool_results_env: Final = os.environ.get("PROMPT_SECURITY_CHECK_TOOL_RESULTS", "false").lower()
145 self.check_tool_results = check_tool_results_env in ("true", "1", "yes")
146 else:
147 self.check_tool_results = check_tool_results
149 if not self.api_key or not self.api_base:
150 msg: Final = (
151 "Couldn't get Prompt Security api base or key, "
152 "either set the `PROMPT_SECURITY_API_BASE` and `PROMPT_SECURITY_API_KEY` in the environment "
153 "or pass them as parameters to the guardrail in the config file"
154 )
155 raise PromptSecurityGuardrailMissingSecrets(msg)
157 self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = (
158 "block_only" if streaming_transform_mode is None else streaming_transform_mode
159 )
161 # Configuration for file sanitization
162 self.max_poll_attempts = 30 # Maximum number of polling attempts
163 self.poll_interval = 2 # Seconds between polling attempts
164 self.file_sanitization_timeout = file_sanitization_timeout
165 self.file_sanitization_fail_open = file_sanitization_fail_open is not False
166 self.block_on_file_modify = block_on_file_modify is not False
168 super().__init__(**kwargs)
170 def supports_scan_only_tool_results(self) -> bool:
171 return self.check_tool_results
173 @log_guardrail_information
174 async def apply_guardrail(
175 self,
176 inputs: GenericGuardrailAPIInputs,
177 request_data: dict,
178 input_type: Literal["request", "response"],
179 logging_obj: Optional["LiteLLMLoggingObj"] = None,
180 ) -> GenericGuardrailAPIInputs:
181 """
182 Apply Prompt Security guardrail to the given inputs.
184 This method is called by LiteLLM's guardrail framework for ALL endpoints:
185 - /chat/completions
186 - /responses
187 - /messages (Anthropic)
188 - /embeddings
189 - /image/generations
190 - /audio/transcriptions
191 - /rerank
192 - MCP server
193 - and more...
195 Args:
196 inputs: Dictionary containing:
197 - texts: List of texts to check
198 - images: Optional list of image URLs
199 - tool_calls: Optional list of tool calls
200 - structured_messages: Optional full message structure
201 request_data: The original request data
202 input_type: "request" for input checking, "response" for output checking
203 logging_obj: Optional logging object
205 Returns:
206 The inputs (potentially modified if action is "modify")
208 Raises:
209 HTTPException: If content is blocked by Prompt Security
210 """
211 texts: Final = inputs.get("texts", [])
212 images: Final = inputs.get("images", [])
213 structured_messages: Final = inputs.get("structured_messages", [])
215 # Resolve user API key alias from request metadata
216 user_api_key_alias: Final = self._resolve_key_alias_from_request_data(request_data)
218 verbose_proxy_logger.debug(
219 "Prompt Security Guardrail: apply_guardrail called with input_type=%s, "
220 "texts=%d, images=%d, structured_messages=%d",
221 input_type,
222 len(texts),
223 len(images),
224 len(structured_messages),
225 )
227 if input_type == "request":
228 return await self._apply_guardrail_on_request(
229 inputs=inputs,
230 texts=texts,
231 images=images,
232 structured_messages=structured_messages,
233 request_data=request_data,
234 user_api_key_alias=user_api_key_alias,
235 )
236 else: # response
237 return await self._apply_guardrail_on_response(
238 inputs=inputs,
239 texts=texts,
240 user_api_key_alias=user_api_key_alias,
241 )
243 async def _apply_guardrail_on_request(
244 self,
245 inputs: GenericGuardrailAPIInputs,
246 texts: list[str],
247 images: list[str],
248 structured_messages: list,
249 request_data: dict,
250 user_api_key_alias: str | None,
251 ) -> GenericGuardrailAPIInputs:
252 """Handle request-side guardrail checks."""
253 # If we have structured messages, use them (they contain role information)
254 # Otherwise, convert texts to simple user messages
255 if structured_messages:
256 messages = list(structured_messages)
257 else:
258 messages = [{"role": "user", "content": text} for text in texts]
260 # Process any embedded files/images in messages
261 messages = await self.process_message_files(messages, user_api_key_alias=user_api_key_alias)
263 # Also process standalone images from inputs
264 if images:
265 await self._process_standalone_images(images, user_api_key_alias)
267 # Filter messages by role for the API call
268 filtered_messages: Final = self.filter_messages_by_role(messages)
270 if not filtered_messages:
271 verbose_proxy_logger.debug("Prompt Security Guardrail: No messages to check after filtering")
272 return inputs
274 # Call Prompt Security API
275 headers: Final = self._build_headers(user_api_key_alias)
276 payload: Final = {
277 "messages": filtered_messages,
278 "user": user_api_key_alias or self.user,
279 "system_prompt": self.system_prompt,
280 }
282 self._log_api_request(
283 method="POST",
284 url=f"{self.api_base}/api/protect",
285 headers=headers,
286 payload={"messages_count": len(filtered_messages)},
287 )
289 response: Final = await self.async_handler.post(
290 f"{self.api_base}/api/protect",
291 headers=headers,
292 json=payload,
293 )
294 response.raise_for_status()
295 res: Final[_ProtectResponse] = response.json()
297 self._log_api_response(
298 url=f"{self.api_base}/api/protect",
299 status_code=response.status_code,
300 payload={"result": res.get("result")},
301 )
303 result: Final = res.get("result", {}).get("prompt", {})
304 if result is None:
305 return inputs
307 action: Final = result.get("action")
308 violations: Final = result.get("violations", [])
310 if action == "block":
311 raise HTTPException(
312 status_code=400,
313 detail="Blocked by Prompt Security, Violations: " + ", ".join(violations),
314 )
315 elif action == "modify":
316 modified_messages: Final = result.get("modified_messages", [])
317 return _inputs_with_modifications(
318 inputs,
319 self._extract_texts_from_messages(modified_messages),
320 self._structured_messages_with_modifications(structured_messages, modified_messages),
321 )
323 return inputs
325 def _is_sent_to_protect(self, message: Mapping[str, object]) -> bool:
326 return self.check_tool_results or message.get("role") in _PROTECT_ROLES
328 def _structured_messages_with_modifications(
329 self,
330 structured_messages: Sequence[AllMessageValues],
331 modified_messages: Sequence[Mapping[str, object]],
332 ) -> tuple[AllMessageValues, ...] | None:
333 sent_indices: Final = tuple(
334 index for index, message in enumerate(structured_messages) if self._is_sent_to_protect(message)
335 )
336 if not sent_indices or len(sent_indices) != len(modified_messages):
337 return None
338 rewritten: Final = tuple(
339 message_with_slot_texts(structured_messages[index], self._extract_texts_from_messages((modified,)))
340 for index, modified in zip(sent_indices, modified_messages)
341 )
342 replacements: Final = MappingProxyType(
343 {index: message for index, message in zip(sent_indices, rewritten) if message is not None}
344 )
345 if len(replacements) != len(sent_indices):
346 return None
347 return tuple(replacements.get(index, message) for index, message in enumerate(structured_messages))
349 async def _apply_guardrail_on_response(
350 self,
351 inputs: GenericGuardrailAPIInputs,
352 texts: list[str],
353 user_api_key_alias: str | None,
354 ) -> GenericGuardrailAPIInputs:
355 """Handle response-side guardrail checks, one protect verdict per text.
357 Prompt Security rewrites a single string, so texts from several choices must be scanned separately
358 or one ``modified_text`` cannot be mapped back onto the choice it came from. It also returns no span
359 offsets, so on a stream every text is held back in full until the final verdict: a value the vendor
360 redacts later may start anywhere in text that looked clean so far, and streamed bytes cannot be recalled.
361 """
362 if not texts:
363 return inputs
365 verdicts: Final = await asyncio.gather(
366 *(self._protect_response_text(text, user_api_key_alias) for text in texts)
367 )
368 violations: Final = tuple(
369 violation
370 for verdict in verdicts
371 if verdict.get("action") == "block"
372 for violation in verdict.get("violations", ())
373 )
374 if any(verdict.get("action") == "block" for verdict in verdicts):
375 raise HTTPException(
376 status_code=400,
377 detail="Blocked by Prompt Security, Violations: " + ", ".join(violations),
378 )
379 returned_texts: Final = [ # mutable-ok: GenericGuardrailAPIInputs.texts is list[str]
380 _modified_or_original(text, verdict) for text, verdict in zip(texts, verdicts, strict=True)
381 ]
382 patched: Final[GenericGuardrailAPIInputs] = {
383 **inputs,
384 "texts": returned_texts,
385 "stream_holdback_chars": [ # mutable-ok: GenericGuardrailAPIInputs.stream_holdback_chars is list[int]
386 len(text) for text in returned_texts
387 ],
388 }
389 return patched
391 async def _protect_response_text(self, text: str, user_api_key_alias: str | None) -> _ProtectVerdict:
392 headers: Final = self._build_headers(user_api_key_alias)
393 payload: Final = {
394 "response": text,
395 "user": user_api_key_alias or self.user,
396 "system_prompt": self.system_prompt,
397 }
399 self._log_api_request(
400 method="POST",
401 url=f"{self.api_base}/api/protect",
402 headers=headers,
403 payload={"response_length": len(text)},
404 )
406 response: Final = await self.async_handler.post(
407 f"{self.api_base}/api/protect",
408 headers=headers,
409 json=payload,
410 )
411 response.raise_for_status()
412 res: Final[_ProtectResponse] = response.json()
414 self._log_api_response(
415 url=f"{self.api_base}/api/protect",
416 status_code=response.status_code,
417 payload={"result": res.get("result")},
418 )
420 verdict: Final = res.get("result", {}).get("response", {})
421 return {} if verdict is None else verdict
423 def _extract_texts_from_messages(self, messages: Sequence[Mapping[str, object]]) -> list[str]:
424 return [text for message in messages for text in message_slot_texts(message)]
426 async def _process_standalone_images(self, images: list[str], user_api_key_alias: str | None) -> None:
427 """Process standalone images from inputs (data URLs)."""
428 for image_url in images:
429 if image_url.startswith("data:"):
430 try:
431 header, encoded = image_url.split(",", 1)
432 file_data = base64.b64decode(encoded)
433 mime_type = header.split(";")[0].split(":")[1]
434 extension = mime_type.split("/")[-1]
435 filename = f"image.{extension}"
437 result = await self.sanitize_file_content(
438 file_data, filename, user_api_key_alias=user_api_key_alias
439 )
440 self._raise_if_file_blocked(result, "Image")
441 except HTTPException:
442 raise
443 except Exception as e:
444 verbose_proxy_logger.error("Error processing image: %s", e)
446 @staticmethod
447 def _resolve_key_alias_from_request_data(request_data: dict) -> str | None:
448 """Resolve user API key alias from request_data metadata."""
449 # Check litellm_metadata first (set by guardrail framework)
450 litellm_metadata: Final = request_data.get("litellm_metadata", {})
451 if litellm_metadata:
452 alias = litellm_metadata.get("user_api_key_alias")
453 if alias:
454 return alias
456 # Then check regular metadata
457 metadata: Final = request_data.get("metadata", {})
458 if metadata:
459 alias = metadata.get("user_api_key_alias")
460 if alias:
461 return alias
463 return None
465 async def sanitize_file_content(
466 self,
467 file_data: bytes,
468 filename: str,
469 user_api_key_alias: str | None = None,
470 ) -> _SanitizeResult:
471 """
472 Sanitize file content using Prompt Security API.
473 Returns: dict with keys 'action', 'content', 'metadata'
474 """
475 try:
476 return await asyncio.wait_for(
477 self._sanitize_file_content(file_data, filename, user_api_key_alias),
478 timeout=self.file_sanitization_timeout,
479 )
480 except (asyncio.TimeoutError, httpx.TimeoutException, LiteLLMTimeout) as exc:
481 if not self.file_sanitization_fail_open:
482 verbose_proxy_logger.error(
483 "Prompt Security Guardrail: file sanitization for %s timed out with %s; failing closed",
484 filename,
485 type(exc).__name__,
486 )
487 raise HTTPException(status_code=408, detail="File sanitization timeout") from exc
489 verbose_proxy_logger.error(
490 "Prompt Security Guardrail: file sanitization for %s timed out with %s; failing open",
491 filename,
492 type(exc).__name__,
493 )
494 fail_open_result: Final[_SanitizeResult] = {
495 "action": "allow",
496 "content": None,
497 "metadata": {},
498 "violations": (),
499 }
500 return fail_open_result
502 async def _sanitize_file_content(
503 self,
504 file_data: bytes,
505 filename: str,
506 user_api_key_alias: str | None,
507 ) -> _SanitizeResult:
508 headers: Final = {"APP-ID": self.api_key}
509 if user_api_key_alias:
510 headers["X-LiteLLM-Key-Alias"] = user_api_key_alias
512 self._log_api_request(
513 method="POST",
514 url=f"{self.api_base}/api/sanitizeFile",
515 headers=headers,
516 payload=f"file upload: {filename}",
517 )
519 # Step 1: Upload file for sanitization
520 files: Final = {"file": (filename, file_data)}
521 upload_response: Final = await self.async_handler.post(
522 f"{self.api_base}/api/sanitizeFile",
523 headers=headers,
524 files=files,
525 )
526 upload_response.raise_for_status()
527 upload_result: Final[_SanitizeUploadResponse] = upload_response.json()
528 job_id: Final = upload_result.get("jobId")
530 self._log_api_response(
531 url=f"{self.api_base}/api/sanitizeFile",
532 status_code=upload_response.status_code,
533 payload={"jobId": job_id},
534 )
536 if not job_id:
537 raise HTTPException(status_code=500, detail="Failed to get jobId from Prompt Security")
539 verbose_proxy_logger.debug("Prompt Security Guardrail: File sanitization started with jobId=%s", job_id)
541 # Step 2: Poll for results
542 for attempt in range(self.max_poll_attempts):
543 await asyncio.sleep(self.poll_interval)
545 self._log_api_request(
546 method="GET",
547 url=f"{self.api_base}/api/sanitizeFile",
548 headers=headers,
549 payload={"jobId": job_id},
550 )
551 poll_response = await self.async_handler.get(
552 f"{self.api_base}/api/sanitizeFile",
553 headers=headers,
554 params={"jobId": job_id},
555 )
556 poll_response.raise_for_status()
557 result: _SanitizeStatusResponse = poll_response.json()
559 self._log_api_response(
560 url=f"{self.api_base}/api/sanitizeFile",
561 status_code=poll_response.status_code,
562 payload={"jobId": job_id, "status": result.get("status")},
563 )
565 status = result.get("status")
567 if status == "done":
568 verbose_proxy_logger.debug(
569 "Prompt Security Guardrail: File sanitization completed for jobId=%s",
570 job_id,
571 )
572 return {
573 "action": result.get("metadata", {}).get("action", "allow"),
574 "content": result.get("content"),
575 "metadata": result.get("metadata", {}),
576 "violations": result.get("metadata", {}).get("violations", []),
577 }
579 if status not in _SANITIZE_FILE_QUEUED_STATUSES:
580 raise HTTPException(status_code=500, detail=f"Unexpected sanitization status: {status}")
582 verbose_proxy_logger.debug(
583 "Prompt Security Guardrail: File sanitization status=%s for jobId=%s (attempt %d/%d)",
584 status,
585 job_id,
586 attempt + 1,
587 self.max_poll_attempts,
588 )
590 raise HTTPException(status_code=408, detail="File sanitization timeout")
592 def _raise_if_file_blocked(self, sanitization_result: _SanitizeResult, resource_name: str) -> None:
593 action: Final = sanitization_result.get("action")
594 if action != "block" and not (action == "modify" and self.block_on_file_modify):
595 return
597 violations: Final = sanitization_result.get("violations", ())
598 raise HTTPException(
599 status_code=400,
600 detail=f"{resource_name} blocked by Prompt Security. Violations: {', '.join(violations)}",
601 )
603 async def _process_image_url_item(self, item: dict, user_api_key_alias: str | None) -> dict:
604 """Process and sanitize image_url items."""
605 image_url_data: Final = item.get("image_url", {})
606 url: Final = image_url_data.get("url", "") if isinstance(image_url_data, dict) else image_url_data
608 if not url.startswith("data:"):
609 return item
611 try:
612 header, encoded = url.split(",", 1)
613 file_data: Final = base64.b64decode(encoded)
614 mime_type: Final = header.split(";")[0].split(":")[1]
615 extension: Final = mime_type.split("/")[-1]
616 filename: Final = f"image.{extension}"
618 sanitization_result: Final = await self.sanitize_file_content(
619 file_data, filename, user_api_key_alias=user_api_key_alias
620 )
621 action: Final = sanitization_result.get("action")
622 self._raise_if_file_blocked(sanitization_result, "File")
624 if action == "modify":
625 sanitized_content: Final = sanitization_result.get("content", "")
626 if sanitized_content:
627 sanitized_encoded: Final = base64.b64encode(sanitized_content.encode()).decode()
628 sanitized_url: Final = f"{header},{sanitized_encoded}"
629 if isinstance(image_url_data, dict):
630 image_url_data["url"] = sanitized_url
631 else:
632 item["image_url"] = sanitized_url
633 verbose_proxy_logger.info("File content modified by Prompt Security")
635 return item
636 except HTTPException:
637 raise
638 except Exception as e:
639 verbose_proxy_logger.error("Error sanitizing image file: %s", e)
640 raise HTTPException(status_code=500, detail=f"File sanitization failed: {e}")
642 async def _process_document_item(self, item: dict, user_api_key_alias: str | None) -> dict:
643 """Process and sanitize document/file items."""
644 doc_data: Final = item.get("document") or item.get("file") or item
646 if isinstance(doc_data, dict):
647 url = doc_data.get("url", "")
648 doc_content = doc_data.get("data", "")
649 else:
650 url = doc_data if isinstance(doc_data, str) else ""
651 doc_content = ""
653 if not (url.startswith("data:") or doc_content):
654 return item
656 try:
657 header = ""
658 if url.startswith("data:"):
659 header, encoded = url.split(",", 1)
660 file_data = base64.b64decode(encoded)
661 mime_type = header.split(";")[0].split(":")[1]
662 else:
663 file_data = base64.b64decode(doc_content)
664 mime_type = (
665 doc_data.get("mime_type", "application/pdf") if isinstance(doc_data, dict) else "application/pdf"
666 )
668 if "pdf" in mime_type:
669 filename = "document.pdf"
670 elif "word" in mime_type or "docx" in mime_type:
671 filename = "document.docx"
672 elif "excel" in mime_type or "xlsx" in mime_type:
673 filename = "document.xlsx"
674 else:
675 extension: Final = mime_type.split("/")[-1]
676 filename = f"document.{extension}"
678 verbose_proxy_logger.info("Sanitizing document: %s", filename)
680 sanitization_result: Final = await self.sanitize_file_content(
681 file_data, filename, user_api_key_alias=user_api_key_alias
682 )
683 action: Final = sanitization_result.get("action")
684 self._raise_if_file_blocked(sanitization_result, "Document")
686 if action == "modify":
687 sanitized_content: Final = sanitization_result.get("content", "")
688 if sanitized_content:
689 sanitized_encoded: Final = base64.b64encode(
690 sanitized_content if isinstance(sanitized_content, bytes) else sanitized_content.encode()
691 ).decode()
693 if url.startswith("data:") and header:
694 sanitized_url: Final = f"{header},{sanitized_encoded}"
695 if isinstance(doc_data, dict):
696 doc_data["url"] = sanitized_url
697 elif isinstance(doc_data, dict):
698 doc_data["data"] = sanitized_encoded
700 verbose_proxy_logger.info("Document content modified by Prompt Security")
702 return item
703 except HTTPException:
704 raise
705 except Exception as e:
706 verbose_proxy_logger.error("Error sanitizing document: %s", e)
707 raise HTTPException(status_code=500, detail=f"Document sanitization failed: {e}")
709 async def process_message_files(self, messages: list, user_api_key_alias: str | None = None) -> list:
710 """Process messages and sanitize any file content (images, documents, PDFs, etc.)."""
711 processed_messages: Final = []
713 for message in messages:
714 content = message.get("content")
716 if not isinstance(content, list):
717 processed_messages.append(message)
718 continue
720 processed_content = []
721 for item in content:
722 if isinstance(item, dict):
723 item_type = item.get("type")
724 if item_type == "image_url":
725 item = await self._process_image_url_item(item, user_api_key_alias)
726 elif item_type in ["document", "file"]:
727 item = await self._process_document_item(item, user_api_key_alias)
729 processed_content.append(item)
731 processed_message = message.copy()
732 processed_message["content"] = processed_content
733 processed_messages.append(processed_message)
735 return processed_messages
737 def filter_messages_by_role(self, messages: list) -> list:
738 """Filter messages to only include standard OpenAI/Anthropic roles.
740 Behavior depends on check_tool_results flag:
741 - False (default): Filters out tool/function roles completely
742 - True: Transforms tool/function to "other" role and includes them
744 This allows checking tool results for indirect prompt injection when enabled.
745 """
746 filtered_messages: Final = []
747 transformed_count = 0
748 filtered_count = 0
750 for message in messages:
751 role = message.get("role", "")
752 if role in _PROTECT_ROLES:
753 filtered_messages.append(message)
754 else:
755 if self.check_tool_results:
756 transformed_message = {
757 "role": "other",
758 **{key: value for key, value in message.items() if key != "role"},
759 }
760 filtered_messages.append(transformed_message)
761 transformed_count += 1
762 verbose_proxy_logger.debug(
763 "Prompt Security Guardrail: Transformed message from role '%s' to 'other'",
764 role,
765 )
766 else:
767 filtered_count += 1
768 verbose_proxy_logger.debug(
769 "Prompt Security Guardrail: Filtered message with role '%s'",
770 role,
771 )
773 if transformed_count > 0:
774 verbose_proxy_logger.debug(
775 "Prompt Security Guardrail: Transformed %d tool/function messages to 'other' role",
776 transformed_count,
777 )
779 if filtered_count > 0:
780 verbose_proxy_logger.debug(
781 "Prompt Security Guardrail: Filtered %d messages (%d -> %d messages)",
782 filtered_count,
783 len(messages),
784 len(filtered_messages),
785 )
787 return filtered_messages
789 def _build_headers(self, user_api_key_alias: str | None = None) -> dict:
790 headers: Final = {"APP-ID": self.api_key, "Content-Type": "application/json"}
791 if user_api_key_alias:
792 headers["X-LiteLLM-Key-Alias"] = user_api_key_alias
793 return headers
795 @staticmethod
796 def _redact_headers(headers: dict) -> dict:
797 return {name: ("REDACTED" if name.lower() == "app-id" else value) for name, value in headers.items()}
799 def _log_api_request(
800 self,
801 method: str,
802 url: str,
803 headers: dict,
804 payload: object,
805 ) -> None:
806 verbose_proxy_logger.debug(
807 "Prompt Security request %s %s headers=%s payload=%s",
808 method,
809 url,
810 self._redact_headers(headers),
811 payload,
812 )
814 def _log_api_response(
815 self,
816 url: str,
817 status_code: int,
818 payload: object,
819 ) -> None:
820 verbose_proxy_logger.debug(
821 "Prompt Security response %s status=%s payload=%s",
822 url,
823 status_code,
824 payload,
825 )
827 @staticmethod
828 def get_config_model() -> type["GuardrailConfigModel"] | None:
829 from litellm.types.proxy.guardrails.guardrail_hooks.prompt_security import (
830 PromptSecurityGuardrailConfigModel,
831 )
833 return PromptSecurityGuardrailConfigModel