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

1import asyncio 

2import base64 

3import os 

4from collections.abc import Mapping, Sequence 

5from types import MappingProxyType 

6from typing import TYPE_CHECKING, Final, Literal, Optional 

7 

8import httpx 

9from fastapi import HTTPException 

10from typing_extensions import ReadOnly, TypedDict 

11 

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 

26 

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 

30 

31 

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

35 

36 

37class PromptSecurityGuardrailMissingSecrets(Exception): 

38 pass 

39 

40 

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 

44 

45 

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 

56 

57 

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) 

67 

68 

69class _ProtectVerdict(TypedDict, total=False): 

70 """One side (``prompt`` or ``response``) of an ``/api/protect`` verdict.""" 

71 

72 action: ReadOnly[str] 

73 violations: ReadOnly[Sequence[str]] 

74 modified_messages: ReadOnly[Sequence[Mapping[str, object]]] 

75 modified_text: ReadOnly[str] 

76 

77 

78class _ProtectResult(TypedDict, total=False): 

79 prompt: ReadOnly[_ProtectVerdict | None] 

80 response: ReadOnly[_ProtectVerdict | None] 

81 

82 

83class _ProtectResponse(TypedDict, total=False): 

84 result: ReadOnly[_ProtectResult] 

85 

86 

87class _SanitizeUploadResponse(TypedDict, total=False): 

88 jobId: ReadOnly[str] 

89 

90 

91class _SanitizeMetadata(TypedDict, total=False): 

92 action: ReadOnly[str] 

93 violations: ReadOnly[Sequence[str]] 

94 

95 

96class _SanitizeStatusResponse(TypedDict, total=False): 

97 """One poll of ``/api/sanitizeFile``.""" 

98 

99 status: ReadOnly[str] 

100 content: ReadOnly[str] 

101 metadata: ReadOnly[_SanitizeMetadata] 

102 

103 

104class _SanitizeResult(TypedDict): 

105 action: ReadOnly[str] 

106 content: ReadOnly[str | None] 

107 metadata: ReadOnly[_SanitizeMetadata] 

108 violations: ReadOnly[Sequence[str]] 

109 

110 

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 ] 

119 

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

139 

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 

148 

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) 

156 

157 self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = ( 

158 "block_only" if streaming_transform_mode is None else streaming_transform_mode 

159 ) 

160 

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 

167 

168 super().__init__(**kwargs) 

169 

170 def supports_scan_only_tool_results(self) -> bool: 

171 return self.check_tool_results 

172 

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. 

183 

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

194 

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 

204 

205 Returns: 

206 The inputs (potentially modified if action is "modify") 

207 

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

214 

215 # Resolve user API key alias from request metadata 

216 user_api_key_alias: Final = self._resolve_key_alias_from_request_data(request_data) 

217 

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 ) 

226 

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 ) 

242 

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] 

259 

260 # Process any embedded files/images in messages 

261 messages = await self.process_message_files(messages, user_api_key_alias=user_api_key_alias) 

262 

263 # Also process standalone images from inputs 

264 if images: 

265 await self._process_standalone_images(images, user_api_key_alias) 

266 

267 # Filter messages by role for the API call 

268 filtered_messages: Final = self.filter_messages_by_role(messages) 

269 

270 if not filtered_messages: 

271 verbose_proxy_logger.debug("Prompt Security Guardrail: No messages to check after filtering") 

272 return inputs 

273 

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 } 

281 

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 ) 

288 

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

296 

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 ) 

302 

303 result: Final = res.get("result", {}).get("prompt", {}) 

304 if result is None: 

305 return inputs 

306 

307 action: Final = result.get("action") 

308 violations: Final = result.get("violations", []) 

309 

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 ) 

322 

323 return inputs 

324 

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 

327 

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

348 

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. 

356 

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 

364 

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 

390 

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 } 

398 

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 ) 

405 

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

413 

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 ) 

419 

420 verdict: Final = res.get("result", {}).get("response", {}) 

421 return {} if verdict is None else verdict 

422 

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

425 

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

436 

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) 

445 

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 

455 

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 

462 

463 return None 

464 

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 

488 

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 

501 

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 

511 

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 ) 

518 

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

529 

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 ) 

535 

536 if not job_id: 

537 raise HTTPException(status_code=500, detail="Failed to get jobId from Prompt Security") 

538 

539 verbose_proxy_logger.debug("Prompt Security Guardrail: File sanitization started with jobId=%s", job_id) 

540 

541 # Step 2: Poll for results 

542 for attempt in range(self.max_poll_attempts): 

543 await asyncio.sleep(self.poll_interval) 

544 

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

558 

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 ) 

564 

565 status = result.get("status") 

566 

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 } 

578 

579 if status not in _SANITIZE_FILE_QUEUED_STATUSES: 

580 raise HTTPException(status_code=500, detail=f"Unexpected sanitization status: {status}") 

581 

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 ) 

589 

590 raise HTTPException(status_code=408, detail="File sanitization timeout") 

591 

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 

596 

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 ) 

602 

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 

607 

608 if not url.startswith("data:"): 

609 return item 

610 

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

617 

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

623 

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

634 

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

641 

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 

645 

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

652 

653 if not (url.startswith("data:") or doc_content): 

654 return item 

655 

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 ) 

667 

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

677 

678 verbose_proxy_logger.info("Sanitizing document: %s", filename) 

679 

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

685 

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

692 

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 

699 

700 verbose_proxy_logger.info("Document content modified by Prompt Security") 

701 

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

708 

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 = [] 

712 

713 for message in messages: 

714 content = message.get("content") 

715 

716 if not isinstance(content, list): 

717 processed_messages.append(message) 

718 continue 

719 

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) 

728 

729 processed_content.append(item) 

730 

731 processed_message = message.copy() 

732 processed_message["content"] = processed_content 

733 processed_messages.append(processed_message) 

734 

735 return processed_messages 

736 

737 def filter_messages_by_role(self, messages: list) -> list: 

738 """Filter messages to only include standard OpenAI/Anthropic roles. 

739 

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 

743 

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 

749 

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 ) 

772 

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 ) 

778 

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 ) 

786 

787 return filtered_messages 

788 

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 

794 

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

798 

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 ) 

813 

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 ) 

826 

827 @staticmethod 

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

829 from litellm.types.proxy.guardrails.guardrail_hooks.prompt_security import ( 

830 PromptSecurityGuardrailConfigModel, 

831 ) 

832 

833 return PromptSecurityGuardrailConfigModel