Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/presidio.py: 8%
848 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# | PII Masking |
4# | with Microsoft Presidio |
5# | https://github.com/BerriAI/litellm/issues/ |
6# +-----------------------------------------------+
7#
8# Tell us how we can improve! - Krrish & Ishaan
11import asyncio
12import json
13import re
14import threading
15from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Iterator, Sequence
16from contextlib import asynccontextmanager
17from dataclasses import dataclass
18from datetime import datetime
19from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast
21import aiohttp
22from typing_extensions import NotRequired, ReadOnly
24import litellm
25from litellm import get_secret
26from litellm._logging import verbose_proxy_logger
27from litellm.constants import (
28 DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES,
29 PRESIDIO_ANALYZE_CHUNK_CONCURRENCY,
30 PRESIDIO_ANALYZE_CHUNK_OVERLAP_CHARS,
31)
32from litellm.types.utils import GenericGuardrailAPIInputs
34if TYPE_CHECKING: 34 ↛ 35line 34 didn't jump to line 35 because the condition on line 34 was never true
35 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
37from litellm.caching.caching import DualCache
38from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
39from litellm.integrations.custom_guardrail import (
40 CustomGuardrail,
41 log_guardrail_information,
42)
43from litellm.proxy._types import UserAPIKeyAuth
44from litellm.proxy.common_utils.sse_keepalive import split_complete_sse_frames
45from litellm.proxy.guardrails.anthropic_sse import (
46 anthropic_sse_chunks_from_response,
47 assemble_anthropic_sse_stream,
48 is_anthropic_sse_stream,
49 model_response_text,
50)
51from litellm.types.guardrails import (
52 GuardrailEventHooks,
53 LitellmParams,
54 PiiAction,
55 PiiEntityType,
56 PresidioPerRequestConfig,
57)
58from litellm.types.proxy.guardrails.guardrail_hooks.presidio import (
59 PresidioAnalyzeRequest,
60 PresidioAnalyzeResponseItem,
61)
62from litellm.types.utils import GuardrailStatus, StreamingChoices
63from litellm.utils import (
64 EmbeddingResponse,
65 ImageResponse,
66 ModelResponse,
67 ModelResponseStream,
68)
71class _PresidioAnonymizeItem(TypedDict, total=False):
72 entity_type: ReadOnly[str | None]
75class _PresidioAnonymizeResponse(TypedDict):
76 text: ReadOnly[str]
77 items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]]
80class _JsonResponse(Protocol):
81 def json(self) -> Awaitable[object]: ... 81 ↛ exitline 81 didn't return from function 'json' because
84async def _json_body(response: _JsonResponse) -> object:
85 return await response.json()
88_LoopSemaphores = dict[asyncio.AbstractEventLoop, asyncio.Semaphore]
91def _json_escaped_len(text: str) -> int:
92 """
93 Byte length of ``text`` as it appears serialized inside the JSON request
94 body sent to Presidio (``json.dumps`` escapes non-ASCII characters, so a
95 3-byte UTF-8 character can occupy 6+ bytes on the wire).
96 """
97 return len(json.dumps(text).encode("utf-8")) - 2 # strip the surrounding quotes
100_MAX_FIRST_SSE_FRAME_BYTES: Final = 64 * 1024
103@dataclass(frozen=True, slots=True)
104class _SsePreface:
105 """Complete leading SSE frames with no ``data:`` line, relayed verbatim before the stream shape is decided."""
107 raw: bytes
110_SSE_FRAME_END: Final = re.compile(rb"\r\n\r\n|\n\n|\r\r")
113def _split_sse_preface(complete_frames: bytes) -> tuple[bytes, bytes]:
114 """Split complete frames into ``(frames before the first data-bearing frame, that frame and everything after)``."""
115 start = 0
116 for end in _SSE_FRAME_END.finditer(complete_frames):
117 frame = complete_frames[start : end.end()]
118 if any(line.startswith(b"data:") for line in frame.splitlines()):
119 return complete_frames[:start], complete_frames[start:]
120 start = end.end()
121 return complete_frames, b""
124def _flush_unmaskable_buffer(all_chunks: list[ModelResponseStream]) -> Iterator[ModelResponseStream]:
125 """Buffered chunks flushed unmasked when a mixed stream shape makes reconstruction impossible."""
126 if not all_chunks:
127 return
128 verbose_proxy_logger.warning(
129 "Presidio apply_to_output: mixed stream detected (ModelResponseStream + unknown event). "
130 "Flushing %d buffered chunks without PII masking and switching to transparent passthrough.",
131 len(all_chunks),
132 )
133 yield from all_chunks
136async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGenerator[object, None]:
137 """
138 Relay leading data-less SSE frames (comment keepalives, events without a
139 ``data:`` line) as they complete, and join raw ``bytes`` chunks until they
140 hold one complete SSE event with a data line, so the stream shape is
141 decided on a whole frame rather than a transport fragment. Everything
142 after that first frame is forwarded untouched. The byte cap can only be
143 reached by a single unterminated frame.
144 """
145 pending = b""
146 try:
147 async for chunk in stream:
148 if not isinstance(chunk, bytes):
149 yield chunk
150 continue
151 pending += chunk
152 complete_frames, tail = split_complete_sse_frames(pending)
153 preface, classifiable = _split_sse_preface(complete_frames)
154 if preface:
155 yield _SsePreface(preface)
156 pending = classifiable + tail
157 if classifiable or len(pending) >= _MAX_FIRST_SSE_FRAME_BYTES:
158 break
159 else:
160 if pending:
161 yield pending
162 return
163 except Exception:
164 if pending:
165 yield pending
166 raise
167 yield pending
168 async for chunk in stream:
169 yield chunk
172class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
173 user_api_key_cache = None
174 ad_hoc_recognizers: list[str] | None = None
176 @classmethod
177 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
178 return [
179 GuardrailEventHooks.pre_call,
180 GuardrailEventHooks.during_call,
181 GuardrailEventHooks.post_call,
182 GuardrailEventHooks.logging_only,
183 GuardrailEventHooks.pre_mcp_call,
184 GuardrailEventHooks.post_mcp_call,
185 ]
187 # Class variables or attributes
188 def __init__(
189 self,
190 mock_testing: bool = False,
191 mock_redacted_text: _PresidioAnonymizeResponse | None = None,
192 presidio_analyzer_api_base: str | None = None,
193 presidio_anonymizer_api_base: str | None = None,
194 output_parse_pii: bool | None = False,
195 apply_to_output: bool = False,
196 presidio_ad_hoc_recognizers: str | None = None,
197 logging_only: bool | None = None,
198 pii_entities_config: dict[PiiEntityType | str, PiiAction] | None = None,
199 presidio_language: str | None = None,
200 presidio_score_thresholds: dict[PiiEntityType | str, float] | None = None,
201 presidio_entities_deny_list: list[PiiEntityType | str] | None = None,
202 presidio_analyze_chunk_size_bytes: int | None = None,
203 **kwargs,
204 ):
205 if logging_only is True:
206 self.logging_only = True
207 kwargs["event_hook"] = GuardrailEventHooks.logging_only
208 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
209 super().__init__(**kwargs)
210 self.guardrail_provider = "presidio"
211 self.pii_tokens: dict[
212 str, str
213 ] = {} # mapping of PII token to original text - only used with Presidio `replace` operation
214 self.mock_redacted_text = mock_redacted_text
215 self.output_parse_pii = output_parse_pii or False
216 self.apply_to_output = apply_to_output
218 # When output_parse_pii or apply_to_output is enabled, the guardrail must
219 # also run on post_call to unmask/mask the response. Expand the event_hook
220 # so should_run_guardrail returns True for both pre_call and post_call.
221 if (self.output_parse_pii or self.apply_to_output) and not logging_only:
222 current_hook: Final = self.event_hook
223 if isinstance(current_hook, str) and current_hook != "post_call":
224 self.event_hook = cast(list[GuardrailEventHooks], [current_hook, "post_call"])
225 elif isinstance(current_hook, list) and "post_call" not in current_hook:
226 self.event_hook = cast(list[GuardrailEventHooks], current_hook + ["post_call"])
227 self.pii_entities_config: dict[PiiEntityType | str, PiiAction] = pii_entities_config or {}
228 self.presidio_score_thresholds: dict[PiiEntityType | str, float] = presidio_score_thresholds or {}
229 self.presidio_entities_deny_list: list[PiiEntityType | str] = presidio_entities_deny_list or []
230 self.presidio_language = presidio_language or "en"
231 self.presidio_analyze_chunk_size_bytes: int = self._coerce_analyze_chunk_size(presidio_analyze_chunk_size_bytes)
232 # Shared HTTP session to prevent memory leaks (issue #14540)
233 self._http_session: aiohttp.ClientSession | None = None
234 # Lock to prevent race conditions when creating session under concurrent load
235 # Note: asyncio.Lock() can be created without an event loop; it only needs one when awaited
236 self._session_lock: asyncio.Lock = asyncio.Lock()
238 # Track main thread ID to safely identity when we are running in main loop vs background thread
240 self._main_thread_id = threading.get_ident()
242 # Loop-bound session cache for background threads
243 self._loop_sessions: dict[asyncio.AbstractEventLoop, aiohttp.ClientSession] = {}
245 # Per-loop semaphores bounding chunked-analyze fan-out across ALL
246 # concurrent oversized blocks/requests on this instance, not per call
247 self._loop_chunk_semaphores: _LoopSemaphores = {}
249 if mock_testing is True: # for testing purposes only
250 return
252 ad_hoc_recognizers: Final = presidio_ad_hoc_recognizers
253 if ad_hoc_recognizers is not None:
254 try:
255 with open(ad_hoc_recognizers, "r") as file:
256 self.ad_hoc_recognizers = json.load(file)
257 except FileNotFoundError:
258 raise Exception(f"File not found. file_path={ad_hoc_recognizers}")
259 except json.JSONDecodeError as e:
260 raise Exception(f"Error decoding JSON file: {e}, file_path={ad_hoc_recognizers}")
261 except Exception as e:
262 raise Exception(f"An error occurred: {e}, file_path={ad_hoc_recognizers}")
263 self.validate_environment(
264 presidio_analyzer_api_base=presidio_analyzer_api_base,
265 presidio_anonymizer_api_base=presidio_anonymizer_api_base,
266 )
268 def validate_environment(
269 self,
270 presidio_analyzer_api_base: str | None = None,
271 presidio_anonymizer_api_base: str | None = None,
272 ):
273 self.presidio_analyzer_api_base: str | None = presidio_analyzer_api_base or get_secret(
274 "PRESIDIO_ANALYZER_API_BASE", None
275 )
276 self.presidio_anonymizer_api_base: str | None = presidio_anonymizer_api_base or litellm.get_secret(
277 "PRESIDIO_ANONYMIZER_API_BASE", None
278 )
280 if self.presidio_analyzer_api_base is None:
281 raise Exception("Missing `PRESIDIO_ANALYZER_API_BASE` from environment")
282 if not self.presidio_analyzer_api_base.endswith("/"):
283 self.presidio_analyzer_api_base += "/"
284 if not (
285 self.presidio_analyzer_api_base.startswith("http://")
286 or self.presidio_analyzer_api_base.startswith("https://")
287 ):
288 # add http:// if unset, assume communicating over private network - e.g. render
289 self.presidio_analyzer_api_base = "http://" + self.presidio_analyzer_api_base
291 if self.presidio_anonymizer_api_base is None:
292 raise Exception("Missing `PRESIDIO_ANONYMIZER_API_BASE` from environment")
293 if not self.presidio_anonymizer_api_base.endswith("/"):
294 self.presidio_anonymizer_api_base += "/"
295 if not (
296 self.presidio_anonymizer_api_base.startswith("http://")
297 or self.presidio_anonymizer_api_base.startswith("https://")
298 ):
299 # add http:// if unset, assume communicating over private network - e.g. render
300 self.presidio_anonymizer_api_base = "http://" + self.presidio_anonymizer_api_base
302 @asynccontextmanager
303 async def _get_session_iterator(
304 self,
305 ) -> AsyncGenerator[aiohttp.ClientSession, None]:
306 """
307 Async context manager for yielding an HTTP session.
309 Logic:
310 1. If running in the main thread (where the object was initialized/destined to live normally),
311 use the shared `self._http_session` (protected by a lock).
312 2. If running in a background thread (e.g. logging hook), use a cached session for that loop.
313 """
314 current_loop: Final = asyncio.get_running_loop()
316 # Check if we are in the stored main thread
317 if threading.get_ident() == self._main_thread_id:
318 # Main thread -> use shared session
319 async with self._session_lock:
320 if self._http_session is None or self._http_session.closed:
321 self._http_session = aiohttp.ClientSession()
322 yield self._http_session
323 else:
324 # Background thread/loop -> use loop-bound session cache
325 # This avoids "attached to a different loop" or "no running event loop" errors
326 # when accessing the shared session created in the main loop
327 if current_loop not in self._loop_sessions or self._loop_sessions[current_loop].closed:
328 self._loop_sessions[current_loop] = aiohttp.ClientSession()
329 yield self._loop_sessions[current_loop]
331 async def _close_http_session(self) -> None:
332 """Close all cached HTTP sessions."""
333 if self._http_session is not None and not self._http_session.closed:
334 await self._http_session.close()
335 self._http_session = None
337 for session in self._loop_sessions.values():
338 if not session.closed:
339 await session.close()
340 self._loop_sessions.clear()
342 def __del__(self):
343 """Cleanup: we try to close, but doing async cleanup in __del__ is risky."""
345 def _has_block_action(self) -> bool:
346 """Return True if pii_entities_config has any BLOCK action (fail-closed on analyzer errors)."""
347 if not self.pii_entities_config:
348 return False
349 return any(action == PiiAction.BLOCK for action in self.pii_entities_config.values())
351 def _get_presidio_analyze_request_payload(
352 self,
353 text: str,
354 presidio_config: PresidioPerRequestConfig | None,
355 request_data: dict,
356 ) -> PresidioAnalyzeRequest:
357 """
358 Construct the payload for the Presidio analyze request
360 API Ref: https://microsoft.github.io/presidio/api-docs/api-docs.html#tag/Analyzer/paths/~1analyze/post
361 """
362 analyze_payload: Final[PresidioAnalyzeRequest] = PresidioAnalyzeRequest(
363 text=text,
364 language=self.presidio_language,
365 )
366 ##################################################################
367 ###### Check if user has configured any params for this guardrail
368 ################################################################
369 if self.ad_hoc_recognizers is not None:
370 analyze_payload["ad_hoc_recognizers"] = self.ad_hoc_recognizers
372 if self.pii_entities_config:
373 analyze_payload["entities"] = list(self.pii_entities_config.keys())
375 ##################################################################
376 ######### End of adding config params
377 ##################################################################
379 # Check if client side request passed any dynamic params
380 if presidio_config and presidio_config.language:
381 analyze_payload["language"] = presidio_config.language
383 casted_analyze_payload: Final[dict] = cast(dict, analyze_payload)
384 casted_analyze_payload.update(self.get_guardrail_dynamic_request_body_params(request_data=request_data))
385 return cast(PresidioAnalyzeRequest, casted_analyze_payload)
387 async def analyze_text(
388 self,
389 text: str,
390 presidio_config: PresidioPerRequestConfig | None,
391 request_data: dict,
392 ) -> list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse:
393 """
394 Send text to the Presidio analyzer endpoint and get analysis results
396 Texts larger than ``presidio_analyze_chunk_size_bytes`` (UTF-8) are split
397 into overlapping chunks, analyzed per chunk, and the per-chunk results
398 are remapped onto the original text. Presidio analyzer deployments
399 commonly cap the /analyze request body size (e.g. at 1 MB), and analyzer
400 latency grows with payload size.
401 """
402 # Chunk oversized texts before the try block so that a failing chunk
403 # keeps the same sanitized error message a single call would produce.
404 # A single-character text can never be split further, so it always
405 # takes the single-call path regardless of its encoded width.
406 if (
407 text
408 and len(text) > 1
409 and self.mock_redacted_text is None
410 and _json_escaped_len(text) > self.presidio_analyze_chunk_size_bytes
411 ):
412 return await self._analyze_text_chunked(
413 text=text,
414 presidio_config=presidio_config,
415 request_data=request_data,
416 )
417 try:
418 # Skip empty or whitespace-only text to avoid Presidio errors
419 # Common in tool/function calling where assistant content is empty
420 if not text or len(text.strip()) == 0:
421 verbose_proxy_logger.debug("Skipping Presidio analysis for empty/whitespace-only text")
422 return []
424 if self.mock_redacted_text is not None:
425 return self.mock_redacted_text
427 # Use shared session to prevent memory leak (issue #14540)
428 async with self._get_session_iterator() as session:
429 # Make the request to /analyze
430 analyze_url: Final = f"{self.presidio_analyzer_api_base}analyze"
432 analyze_payload: Final[PresidioAnalyzeRequest] = self._get_presidio_analyze_request_payload(
433 text=text,
434 presidio_config=presidio_config,
435 request_data=request_data,
436 )
438 verbose_proxy_logger.debug(
439 "Making request to: %s with payload: %s",
440 analyze_url,
441 analyze_payload,
442 )
444 def _fail_on_invalid_response(
445 reason: str,
446 ) -> list[PresidioAnalyzeResponseItem]:
447 should_fail_closed = bool(self.pii_entities_config) or self.output_parse_pii or self.apply_to_output
448 if should_fail_closed:
449 raise GuardrailRaisedException(
450 guardrail_name=self.guardrail_name,
451 message=f"Presidio analyzer returned invalid response; cannot verify PII when PII protection is configured: {reason}",
452 should_wrap_with_default_message=False,
453 )
454 verbose_proxy_logger.warning("Presidio analyzer %s, returning empty list", reason)
455 return []
457 async with session.post(
458 analyze_url,
459 json=analyze_payload,
460 headers={"Accept": "application/json"},
461 ) as response:
462 # Validate HTTP status
463 if response.status >= 400:
464 error_body = await response.text()
465 return _fail_on_invalid_response(
466 f"HTTP {response.status} from Presidio analyzer: {error_body[:200]}"
467 )
469 # Validate Content-Type is JSON
470 content_type: Final = getattr(
471 response,
472 "content_type",
473 response.headers.get("Content-Type", ""),
474 )
475 if "application/json" not in content_type:
476 error_body = await response.text()
477 return _fail_on_invalid_response(
478 f"expected application/json Content-Type but received '{content_type}'; body: '{error_body[:200]}'"
479 )
481 analyze_results: Final = await _json_body(response)
482 verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
484 # Handle error responses from Presidio (e.g., {'error': 'No text provided'})
485 # Presidio may return a dict instead of a list when errors occur
487 if isinstance(analyze_results, dict):
488 if "error" in analyze_results:
489 return _fail_on_invalid_response(f"error: {analyze_results.get('error')}")
490 # If it's a dict but not an error, try to process it as a single item
491 verbose_proxy_logger.debug(
492 "Presidio returned dict (not list), attempting to process as single item"
493 )
494 try:
495 return [PresidioAnalyzeResponseItem(**analyze_results)]
496 except Exception as e:
497 return _fail_on_invalid_response(f"failed to parse dict response: {e}")
499 # Handle unexpected types (str, None, etc.) - e.g. from malformed/error
500 if not isinstance(analyze_results, list):
501 return _fail_on_invalid_response(
502 f"unexpected type {type(analyze_results).__name__} (expected list or dict), response: {str(analyze_results)[:200]}"
503 )
505 # Normal case: list of results
506 final_results: Final = []
507 for item in analyze_results:
508 if not isinstance(item, dict):
509 verbose_proxy_logger.warning(
510 "Skipping invalid Presidio result item (expected dict, got %s): %s",
511 type(item).__name__,
512 str(item)[:100],
513 )
514 continue
515 try:
516 final_results.append(PresidioAnalyzeResponseItem(**item))
517 except Exception as e:
518 verbose_proxy_logger.warning(
519 "Failed to parse Presidio result item: %s (error: %s)",
520 item,
521 e,
522 )
523 continue
524 return final_results
525 except GuardrailRaisedException:
526 # Re-raise GuardrailRaisedException without wrapping
527 raise
528 except Exception as e:
529 # Sanitize exception to avoid leaking the original text (which may
530 # contain API keys or other secrets) in error responses.
531 raise Exception(f"Presidio PII analysis failed: {type(e).__name__}") from e
533 async def _analyze_text_chunked(
534 self,
535 text: str,
536 presidio_config: PresidioPerRequestConfig | None,
537 request_data: dict, # mutable-ok: shared per-request state dict, matching analyze_text's parameter
538 ) -> list[PresidioAnalyzeResponseItem]: # mutable-ok: analyze_text's declared return type requires list
539 """
540 Analyze an oversized text by splitting it into overlapping chunks.
542 Each chunk serializes to at most ``presidio_analyze_chunk_size_bytes``
543 bytes inside the JSON request body, so every /analyze call stays below
544 the analyzer deployment's request body limit; per-chunk results are remapped onto the original text and
545 merged. Raises exactly like a single ``analyze_text`` call if any chunk
546 fails.
548 Only the analyzer side is chunked: the later anonymize call still
549 receives the full original text, so texts above the anonymizer's own
550 body limit that contain detections keep failing there.
551 """
552 text_chunks: Final = self._split_text_for_analysis(
553 text=text,
554 chunk_size_bytes=self.presidio_analyze_chunk_size_bytes,
555 overlap_chars=PRESIDIO_ANALYZE_CHUNK_OVERLAP_CHARS,
556 )
557 verbose_proxy_logger.debug(
558 "Presidio analyze: text exceeds %s bytes, analyzing in %s overlapping chunks",
559 self.presidio_analyze_chunk_size_bytes,
560 len(text_chunks),
561 )
562 # Bound the fan-out so oversized requests cannot saturate the analyzer.
563 # The semaphore is shared per event loop across every chunked call on
564 # this instance, so many oversized blocks in one request (or many
565 # concurrent requests) still hold at most this many analyzer calls in
566 # flight. On the proxy's main thread the shared-session lock in
567 # _get_session_iterator additionally serializes the HTTP calls; the
568 # bound matters for loop-bound sessions (background threads).
569 analyze_semaphore: Final = self._get_chunk_semaphore()
571 async def _analyze_chunk_bounded(
572 chunk_text: str,
573 ) -> Sequence[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse:
574 async with analyze_semaphore:
575 return await self.analyze_text(
576 text=chunk_text,
577 presidio_config=presidio_config,
578 request_data=request_data,
579 )
581 gathered: Final = await asyncio.gather(
582 *(_analyze_chunk_bounded(chunk_text) for _, chunk_text in text_chunks),
583 return_exceptions=True,
584 )
585 chunk_results: Final = []
586 for result in gathered:
587 if isinstance(result, BaseException):
588 raise result
589 # analyze_text only returns a non-list shape when mock_redacted_text
590 # is set, and the chunked path is never entered in that case.
591 typed_result = cast("list[PresidioAnalyzeResponseItem]", result) # cast-ok: gather() erases element type
592 # Apply the configured score thresholds and deny list BEFORE the
593 # overlap merge: a below-threshold detection must not win overlap
594 # resolution against one the thresholds would keep. The same filter
595 # runs again downstream in check_pii, where it is a no-op for the
596 # already-filtered items.
597 filtered_result = self.filter_analyze_results_by_score(analyze_results=typed_result)
598 chunk_results.append(
599 cast("list[PresidioAnalyzeResponseItem]", filtered_result) # cast-ok: list input yields list
600 )
601 return self._merge_chunked_analyze_results(text_chunks=text_chunks, chunk_results=chunk_results)
603 def _get_chunk_semaphore(self) -> asyncio.Semaphore:
604 """Per-event-loop semaphore shared by all chunked analyze calls on this instance."""
605 loop: Final = asyncio.get_running_loop()
606 existing: Final = self._loop_chunk_semaphores.get(loop)
607 if existing is not None:
608 return existing
609 created: Final = asyncio.Semaphore(PRESIDIO_ANALYZE_CHUNK_CONCURRENCY)
610 self._loop_chunk_semaphores[loop] = created
611 return created
613 @staticmethod
614 def _coerce_analyze_chunk_size(value: int | None) -> int:
615 """
616 Validate a configured chunk size, falling back to the default.
618 Non-positive values would either bypass chunking entirely or degenerate
619 it into per-character splits (silently disabling detection), so they are
620 replaced by the default; values below 4 bytes are floored to 4 and the
621 splitter always emits at least one character per chunk, so the chunked
622 path can never re-enter itself.
623 """
624 if not value or value <= 0:
625 return DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES
626 return max(value, 4)
628 @staticmethod
629 def _split_text_for_analysis(
630 text: str,
631 chunk_size_bytes: int,
632 overlap_chars: int,
633 ) -> Sequence[tuple[int, str]]:
634 """
635 Split ``text`` into chunks whose JSON-serialized form is at most
636 ``chunk_size_bytes`` bytes (the analyzer body limit applies to the
637 JSON request body, where non-ASCII characters are escaped and larger
638 than their raw UTF-8 encoding).
640 Consecutive chunks overlap by up to ``overlap_chars`` characters so a
641 PII entity up to that length lying across a chunk boundary is still
642 seen whole by one of the chunks (longer boundary-straddling entities
643 may be seen only truncated); ``_merge_chunked_analyze_results`` resolves
644 the duplicate and truncated detections this produces. Returns
645 ``(char_offset, chunk_text)`` pairs where ``char_offset`` is the
646 chunk's start position in the original text.
647 """
648 chunks: Final = []
649 text_len: Final = len(text)
650 start = 0 # rebind-ok: chunk cursor advances across the loop
651 while start < text_len:
652 # Serialized length of a character is at least 1 byte, so a slice
653 # of chunk_size_bytes characters is a sufficient search window.
654 candidate = text[start : start + chunk_size_bytes]
655 if _json_escaped_len(candidate) <= chunk_size_bytes:
656 chunk = candidate
657 else:
658 # Largest prefix whose serialized form fits the budget.
659 low, high = 1, len(candidate)
660 while low < high:
661 mid = (low + high + 1) // 2
662 if _json_escaped_len(candidate[:mid]) <= chunk_size_bytes:
663 low = mid
664 else:
665 high = mid - 1
666 # low >= 1 keeps the loop advancing even when a single
667 # character serializes over a (floored, tiny) budget.
668 chunk = candidate[:low]
669 end = start + len(chunk)
670 chunks.append((start, chunk))
671 if end >= text_len:
672 break
673 # Cap the overlap so the next chunk always makes forward progress.
674 effective_overlap = min(overlap_chars, len(chunk) // 2)
675 start = max(start + 1, end - effective_overlap)
676 return chunks
678 @staticmethod
679 def _merge_chunked_analyze_results(
680 text_chunks: Sequence[tuple[int, str]],
681 chunk_results: Sequence[Sequence[PresidioAnalyzeResponseItem]],
682 ) -> list[PresidioAnalyzeResponseItem]: # mutable-ok: analyze_text's declared return type requires list
683 """
684 Remap per-chunk analyzer offsets onto the original text and merge.
686 A detection in an overlap region is reported by both neighbouring
687 chunks, and a boundary entity can additionally be reported truncated by
688 the chunk that saw only its head or tail. Same-entity-type detections
689 with overlapping remapped spans are therefore resolved by keeping the
690 longest span (highest score on ties) — mirroring the same-type conflict
691 removal Presidio's AnalyzerEngine applies within a single call, and
692 keeping overlapping spans from corrupting the numbered-token rewriter.
693 Detections of DIFFERENT entity types may still overlap, exactly as in a
694 single-call response. The merged list is sorted by position.
695 """
696 remapped: Final = []
697 for (char_offset, _), results in zip(text_chunks, chunk_results, strict=True):
698 for item in results:
699 item_start = item.get("start")
700 item_end = item.get("end")
701 if item_start is not None:
702 item["start"] = item_start + char_offset
703 if item_end is not None:
704 item["end"] = item_end + char_offset
705 remapped.append(item)
707 def _priority(item: PresidioAnalyzeResponseItem) -> tuple[int, float]:
708 span_start: Final = item.get("start") or 0
709 span_end: Final = item.get("end") or 0
710 return (-(span_end - span_start), -(item.get("score") or 0.0))
712 merged: Final = []
713 kept_spans_by_type: Final = {}
714 for item in sorted(remapped, key=_priority):
715 item_start = item.get("start")
716 item_end = item.get("end")
717 if item_start is None or item_end is None:
718 merged.append(item)
719 continue
720 kept_spans = kept_spans_by_type.setdefault(str(item.get("entity_type")), [])
721 if any(item_start < kept_end and kept_start < item_end for kept_start, kept_end in kept_spans):
722 continue
723 kept_spans.append((item_start, item_end))
724 merged.append(item)
725 merged.sort(key=lambda r: (r.get("start") or 0, r.get("end") or 0))
726 return merged
728 async def _post_presidio_anonymize(
729 self,
730 text: str,
731 analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse,
732 ) -> _PresidioAnonymizeResponse | None:
733 """POST to Presidio anonymize; returns parsed JSON body."""
734 # Use shared session to prevent memory leak (issue #14540)
735 async with self._get_session_iterator() as session:
736 anonymize_url: Final = f"{self.presidio_anonymizer_api_base}anonymize"
737 verbose_proxy_logger.debug("Making request to: %s", anonymize_url)
738 anonymize_payload: Final = {
739 "text": text,
740 "analyzer_results": analyze_results,
741 }
742 async with session.post(
743 anonymize_url,
744 json=anonymize_payload,
745 headers={"Accept": "application/json"},
746 ) as response:
747 if response.status >= 400:
748 error_body = await response.text()
749 raise Exception(f"Presidio anonymizer returned HTTP {response.status}: {error_body[:200]}")
750 content_type: Final = getattr(
751 response,
752 "content_type",
753 response.headers.get("Content-Type", ""),
754 )
755 if "application/json" not in content_type:
756 error_body = await response.text()
757 raise Exception(
758 f"Presidio anonymizer returned non-JSON Content-Type '{content_type}'; body: '{error_body[:200]}'"
759 )
760 return await response.json()
762 def _finalize_presidio_anonymize_simple(
763 self,
764 redacted_text: _PresidioAnonymizeResponse,
765 masked_entity_count: dict[str, int],
766 ) -> str:
767 # No need to build numbered tokens — just use Presidio's
768 # already-anonymized text directly. The old code incorrectly
769 # applied anonymizer item positions (which reference the
770 # *output* text) to the *original* text, causing offset errors.
771 for item in redacted_text.get("items", []):
772 entity_type = item.get("entity_type", None)
773 if entity_type is not None:
774 masked_entity_count[entity_type] = masked_entity_count.get(entity_type, 0) + 1
775 return redacted_text["text"]
777 def _finalize_presidio_anonymize_numbered_tokens(
778 self,
779 text: str,
780 analyze_results: Any,
781 request_data: dict | None,
782 masked_entity_count: dict[str, int],
783 ) -> str:
784 # output_parse_pii is True — we need sequentially numbered
785 # tokens and a pii_tokens mapping for later unmasking.
786 # Use analyze_results positions (which reference the ORIGINAL
787 # text) instead of anonymizer items (which reference the output).
788 new_text = text
789 if request_data is None:
790 verbose_proxy_logger.warning(
791 "Presidio anonymize_text called without request_data — "
792 "PII tokens cannot be stored per-request. "
793 "This may indicate a missing caller update."
794 )
795 request_data = {}
796 if not request_data.get("metadata"):
797 request_data["metadata"] = {}
798 if "pii_tokens" not in request_data["metadata"]:
799 request_data["metadata"]["pii_tokens"] = {}
800 pii_tokens: Final = request_data["metadata"]["pii_tokens"]
802 # Assign sequence numbers in forward (left-to-right) order so
803 # that <PERSON_1> is the first entity in the text, etc.
804 sorted_forward: Final = sorted(analyze_results, key=lambda x: x["start"])
805 seq_map: Final = {}
806 for idx, ar in enumerate(sorted_forward, start=1):
807 seq_map[(ar["start"], ar["end"])] = idx
809 # Apply replacements in reverse order by start position so
810 # that replacing later spans first does not shift earlier
811 # coordinates in the original text.
812 for ar in reversed(sorted_forward):
813 start = ar["start"]
814 end = ar["end"]
815 entity_type = ar["entity_type"]
816 replacement = f"<{entity_type}>"
817 seq = seq_map[(start, end)]
818 if replacement.endswith(">"):
819 replacement = f"{replacement[:-1]}_{seq}>"
820 else:
821 replacement = f"{replacement}_{seq}"
822 pii_tokens[replacement] = text[start:end]
823 new_text = new_text[:start] + replacement + new_text[end:]
824 masked_entity_count[entity_type] = masked_entity_count.get(entity_type, 0) + 1
825 return new_text
827 async def anonymize_text(
828 self,
829 text: str,
830 analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse,
831 output_parse_pii: bool,
832 masked_entity_count: dict[str, int],
833 request_data: dict | None = None,
834 ) -> str:
835 """
836 Send analysis results to the Presidio anonymizer endpoint to get redacted text
837 """
838 try:
839 # If there are no detections after filtering, return the original text
840 if isinstance(analyze_results, list) and len(analyze_results) == 0:
841 return text
843 redacted_text: Final = await self._post_presidio_anonymize(text, analyze_results)
844 if redacted_text is None:
845 raise Exception("Invalid anonymizer response: received None")
847 verbose_proxy_logger.debug("redacted_text: %s", redacted_text)
849 if not output_parse_pii:
850 return self._finalize_presidio_anonymize_simple(redacted_text, masked_entity_count)
852 return self._finalize_presidio_anonymize_numbered_tokens(
853 text, analyze_results, request_data, masked_entity_count
854 )
855 except Exception as e:
856 # Sanitize exception to avoid leaking the original text (which may
857 # contain API keys or other secrets) in error responses.
858 error_str: Final = str(e)
859 if "Invalid anonymizer response" in error_str or "Presidio anonymizer returned" in error_str:
860 raise
861 raise Exception(f"Presidio PII anonymization failed: {type(e).__name__}") from e
863 def filter_analyze_results_by_score(
864 self, analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse
865 ) -> list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse:
866 """
867 Drop detections that fall below configured per-entity score thresholds
868 or match an entity type in the deny list.
869 """
870 if not self.presidio_score_thresholds and not self.presidio_entities_deny_list:
871 return analyze_results
873 if not isinstance(analyze_results, list):
874 return analyze_results
876 filtered_results: Final[list[PresidioAnalyzeResponseItem]] = []
877 deny_list_strings: Final = [getattr(x, "value", str(x)) for x in self.presidio_entities_deny_list]
878 for item in analyze_results:
879 entity_type = item.get("entity_type")
881 str_entity_type = str(
882 getattr(entity_type, "value", entity_type) if entity_type is not None else entity_type
883 )
884 if entity_type and str_entity_type in deny_list_strings:
885 continue
887 if self.presidio_score_thresholds:
888 score = item.get("score")
889 threshold = None
890 if entity_type is not None:
891 threshold = self.presidio_score_thresholds.get(entity_type)
892 if threshold is None:
893 threshold = self.presidio_score_thresholds.get("ALL")
895 if threshold is not None:
896 if score is None or score < threshold:
897 continue
899 filtered_results.append(item)
901 return filtered_results
903 def raise_exception_if_blocked_entities_detected(
904 self, analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse
905 ):
906 """
907 Raise an exception if blocked entities are detected
908 """
909 if self.pii_entities_config is None:
910 return
912 if isinstance(analyze_results, dict):
913 # if mock testing is enabled, analyze_results is a dict
914 # we don't need to raise an exception in this case
915 return
917 for result in analyze_results:
918 entity_type = result.get("entity_type")
920 if entity_type:
921 # Check if entity_type is in config (supports both enum and string)
922 if entity_type in self.pii_entities_config and self.pii_entities_config[entity_type] == PiiAction.BLOCK:
923 raise BlockedPiiEntityError(
924 entity_type=entity_type,
925 guardrail_name=self.guardrail_name,
926 )
928 async def check_pii(
929 self,
930 text: str,
931 output_parse_pii: bool,
932 presidio_config: PresidioPerRequestConfig | None,
933 request_data: dict,
934 ) -> str:
935 """
936 Calls Presidio Analyze + Anonymize endpoints for PII Analysis + Masking
937 """
938 start_time: Final = datetime.now()
939 analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse | None = None
940 status: GuardrailStatus = "success"
941 masked_entity_count: Final[dict[str, int]] = {}
942 exception_str: str = ""
943 try:
944 if self.mock_redacted_text is not None:
945 redacted_text: Final = self.mock_redacted_text
946 else:
947 # First get analysis results
948 analyze_results = await self.analyze_text(
949 text=text,
950 presidio_config=presidio_config,
951 request_data=request_data,
952 )
954 verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
956 # Apply score threshold filtering if configured
957 analyze_results = self.filter_analyze_results_by_score(analyze_results=analyze_results)
959 ####################################################
960 # Blocked Entities check
961 ####################################################
962 self.raise_exception_if_blocked_entities_detected(analyze_results=analyze_results)
964 # Then anonymize the text using the analysis results
965 anonymized_text: Final = await self.anonymize_text(
966 text=text,
967 analyze_results=analyze_results,
968 output_parse_pii=output_parse_pii,
969 masked_entity_count=masked_entity_count,
970 request_data=request_data,
971 )
972 return anonymized_text
973 return redacted_text["text"]
974 except Exception as e:
975 status = "guardrail_failed_to_respond"
976 exception_str = str(e)
977 raise e
978 finally:
979 ####################################################
980 # Create Guardrail Trace for logging on Langfuse, Datadog, etc.
981 ####################################################
982 guardrail_json_response: Exception | str | dict | list[dict] = {}
983 if status == "success":
984 if isinstance(analyze_results, list):
985 guardrail_json_response = [dict(item) for item in analyze_results]
986 else:
987 guardrail_json_response = exception_str
988 self.add_standard_logging_guardrail_information_to_request_data(
989 guardrail_provider=self.guardrail_provider,
990 guardrail_json_response=guardrail_json_response,
991 request_data=request_data,
992 guardrail_status=status,
993 start_time=start_time.timestamp(),
994 end_time=datetime.now().timestamp(),
995 duration=(datetime.now() - start_time).total_seconds(),
996 masked_entity_count=masked_entity_count,
997 )
999 async def async_pre_call_hook(
1000 self,
1001 user_api_key_dict: UserAPIKeyAuth,
1002 cache: DualCache,
1003 data: dict,
1004 call_type: str,
1005 ):
1006 """
1007 - Check if request turned off pii
1008 - Check if user allowed to turn off pii (key permissions -> 'allow_pii_controls')
1010 - Take the request data
1011 - Call /analyze -> get the results
1012 - Call /anonymize w/ the analyze results -> get the redacted text
1014 For multiple messages in /chat/completions, we'll need to call them in parallel.
1015 """
1016 # Respect the configured event hook. In `logging_only` mode (and any config that
1017 # excludes pre_call) the live request must not be masked - masking is applied to a
1018 # copy at logging time via `async_logging_hook`. Without this gate the request sent
1019 # to the model would carry anonymization tokens and the response would echo them.
1020 if (
1021 self.should_run_guardrail(
1022 data=data,
1023 event_type=GuardrailEventHooks.pre_call,
1024 )
1025 is not True
1026 ):
1027 return data
1029 try:
1030 content_safety: Final = data.get("content_safety", None)
1031 verbose_proxy_logger.debug("content_safety: %s", content_safety)
1032 presidio_config: Final = self.get_presidio_settings_from_request_data(data)
1033 messages: Final = data.get("messages", None)
1034 if messages is None:
1035 return data
1036 tasks: Final = []
1037 task_mappings: list[tuple[int, int | None]] = [] # Track (message_index, content_index) for each task
1039 for msg_idx, m in enumerate(messages):
1040 content = m.get("content", None)
1041 if content is None:
1042 continue
1043 if isinstance(content, str):
1044 tasks.append(
1045 self.check_pii(
1046 text=content,
1047 output_parse_pii=self.output_parse_pii,
1048 presidio_config=presidio_config,
1049 request_data=data,
1050 )
1051 )
1052 task_mappings.append((msg_idx, None)) # None indicates string content
1053 elif isinstance(content, list):
1054 for content_idx, c in enumerate(content):
1055 text_str = c.get("text", None)
1056 if text_str is None:
1057 continue
1058 tasks.append(
1059 self.check_pii(
1060 text=text_str,
1061 output_parse_pii=self.output_parse_pii,
1062 presidio_config=presidio_config,
1063 request_data=data,
1064 )
1065 )
1066 task_mappings.append((msg_idx, int(content_idx)))
1068 responses: Final = await asyncio.gather(*tasks)
1070 # Map responses back to the correct message and content item
1071 for task_idx, r in enumerate(responses):
1072 mapping = task_mappings[task_idx]
1073 msg_idx = cast(int, mapping[0])
1074 content_idx_optional = cast(int | None, mapping[1])
1075 content = messages[msg_idx].get("content", None)
1076 if content is None:
1077 continue
1078 if isinstance(content, str) and content_idx_optional is None:
1079 messages[msg_idx]["content"] = r # replace content with redacted string
1080 elif isinstance(content, list) and content_idx_optional is not None:
1081 messages[msg_idx]["content"][content_idx_optional]["text"] = r
1083 verbose_proxy_logger.debug("Presidio PII Masking: Redacted pii message: %s", data["messages"])
1084 data["messages"] = messages
1085 return data
1086 except Exception as e:
1087 raise e
1089 def logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
1090 from concurrent.futures import ThreadPoolExecutor
1092 def run_in_new_loop():
1093 """Run the coroutine in a new event loop within this thread."""
1094 new_loop: Final = asyncio.new_event_loop()
1095 try:
1096 asyncio.set_event_loop(new_loop)
1097 return new_loop.run_until_complete(
1098 self.async_logging_hook(kwargs=kwargs, result=result, call_type=call_type)
1099 )
1100 finally:
1101 new_loop.close()
1102 asyncio.set_event_loop(None)
1104 try:
1105 # First, try to get the current event loop
1106 _ = asyncio.get_running_loop()
1107 # If we're already in an event loop, run in a separate thread
1108 # to avoid nested event loop issues
1109 with ThreadPoolExecutor(max_workers=1) as executor:
1110 future: Final = executor.submit(run_in_new_loop)
1111 return future.result()
1113 except RuntimeError:
1114 # No running event loop, we can safely run in this thread
1115 return run_in_new_loop()
1117 async def async_logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
1118 """
1119 Masks the input and output before logging to langfuse, datadog, etc.
1120 """
1121 if call_type == "completion" or call_type == "acompletion": # /chat/completions requests
1122 messages: Final[list | None] = kwargs.get("messages", None)
1123 tasks: Final = []
1124 task_mappings: list[tuple[int, int | None]] = [] # Track (message_index, content_index) for each task
1126 if messages is None:
1127 return kwargs, result
1129 presidio_config: Final = self.get_presidio_settings_from_request_data(kwargs)
1131 for msg_idx, m in enumerate(messages):
1132 content = m.get("content", None)
1133 if content is None:
1134 continue
1135 if isinstance(content, str):
1136 tasks.append(
1137 self.check_pii(
1138 text=content,
1139 output_parse_pii=False,
1140 presidio_config=presidio_config,
1141 request_data=kwargs,
1142 )
1143 ) # need to pass separately b/c presidio has context window limits
1144 task_mappings.append((msg_idx, None)) # None indicates string content
1145 elif isinstance(content, list):
1146 for content_idx, c in enumerate(content):
1147 text_str = c.get("text", None)
1148 if text_str is None:
1149 continue
1150 tasks.append(
1151 self.check_pii(
1152 text=text_str,
1153 output_parse_pii=False,
1154 presidio_config=presidio_config,
1155 request_data=kwargs,
1156 )
1157 )
1158 task_mappings.append((msg_idx, int(content_idx)))
1160 responses: Final = await asyncio.gather(*tasks)
1162 # Map responses back to the correct message and content item
1163 for task_idx, r in enumerate(responses):
1164 mapping = task_mappings[task_idx]
1165 msg_idx = cast(int, mapping[0])
1166 content_idx_optional = cast(int | None, mapping[1])
1167 content = messages[msg_idx].get("content", None)
1168 if content is None:
1169 continue
1170 if isinstance(content, str) and content_idx_optional is None:
1171 messages[msg_idx]["content"] = r # replace content with redacted string
1172 elif isinstance(content, list) and content_idx_optional is not None:
1173 messages[msg_idx]["content"][content_idx_optional]["text"] = r
1175 verbose_proxy_logger.debug("Presidio PII Masking: Redacted pii message: %s", messages)
1176 kwargs["messages"] = messages
1178 if (
1179 isinstance(result, ModelResponse)
1180 and result.choices
1181 and not isinstance(result.choices[0], StreamingChoices)
1182 ):
1183 await self._process_response_for_pii(response=result, request_data=kwargs, mode="mask")
1184 elif isinstance(result, dict) and self._is_anthropic_message_response(result):
1185 await self._process_anthropic_response_for_pii(
1186 response=result,
1187 request_data=kwargs,
1188 mode="mask",
1189 )
1191 return kwargs, result
1193 async def async_post_call_success_hook(
1194 self,
1195 data: dict,
1196 user_api_key_dict: UserAPIKeyAuth,
1197 response: ModelResponse | EmbeddingResponse | ImageResponse,
1198 ):
1199 """
1200 Output parse the response object to replace the masked tokens with user sent values
1201 """
1202 verbose_proxy_logger.debug(
1203 "PII Masking Args: self.output_parse_pii=%s; type of response=%s", self.output_parse_pii, type(response)
1204 )
1206 if self.apply_to_output is True:
1207 if self._is_anthropic_message_response(response):
1208 return await self._process_anthropic_response_for_pii(
1209 response=cast(dict, response), request_data=data, mode="mask"
1210 )
1211 return await self._mask_output_response(response=response, request_data=data)
1213 if self.output_parse_pii is False and litellm.output_parse_pii is False:
1214 return response
1216 if isinstance(response, ModelResponse) and not isinstance(
1217 response.choices[0], StreamingChoices
1218 ): # /chat/completions requests
1219 await self._process_response_for_pii(
1220 response=response,
1221 request_data=data,
1222 mode="unmask",
1223 )
1224 elif self._is_anthropic_message_response(response):
1225 await self._process_anthropic_response_for_pii(
1226 response=cast(dict, response), request_data=data, mode="unmask"
1227 )
1228 return response
1230 @staticmethod
1231 def _unmask_pii_text(text: str, pii_tokens: dict[str, str]) -> str:
1232 """
1233 Replace PII tokens in *text* with their original values.
1235 Includes a fallback for tokens that were truncated by ``max_tokens``:
1236 if the *end* of ``text`` matches the *beginning* of a token and the
1237 overlap is long enough, the truncated suffix is replaced with the
1238 original value. The minimum overlap length is
1239 ``min(20, len(token) // 2)`` to reduce the risk of false positives
1240 when multiple tokens share a common prefix.
1241 """
1242 for token, original_text in pii_tokens.items():
1243 if token in text:
1244 text = text.replace(token, original_text)
1245 else:
1246 # FALLBACK: Handle truncated tokens (token cut off by max_tokens)
1247 # Only check at the very end of the text.
1248 min_overlap = min(20, len(token) // 2)
1249 for i in range(max(0, len(text) - len(token)), len(text)):
1250 sub = text[i:]
1251 if token.startswith(sub) and len(sub) >= min_overlap:
1252 text = text[:i] + original_text
1253 break
1254 return text
1256 @staticmethod
1257 def _is_anthropic_message_response(
1258 response: ModelResponse | EmbeddingResponse | ImageResponse | dict[str, object],
1259 ) -> bool:
1260 """Check if the response is an Anthropic native message dict."""
1261 return (
1262 isinstance(response, dict)
1263 and response.get("type") == "message"
1264 and isinstance(response.get("content"), list)
1265 )
1267 async def _process_anthropic_response_for_pii(
1268 self,
1269 response: dict,
1270 request_data: dict,
1271 mode: Literal["mask", "unmask"],
1272 ) -> dict:
1273 """
1274 Process an Anthropic native message dict for PII masking/unmasking.
1275 Handles content blocks with type == "text".
1276 """
1277 metadata: Final = (request_data.get("metadata") or {}) if request_data else {}
1278 pii_tokens: Final = metadata.get("pii_tokens", {})
1279 if not pii_tokens and mode == "unmask":
1280 verbose_proxy_logger.debug("No pii_tokens in metadata for Anthropic response unmask")
1281 presidio_config: Final = self.get_presidio_settings_from_request_data(request_data or {})
1283 content: Final = response.get("content")
1284 if not isinstance(content, list):
1285 return response
1287 for block in content:
1288 if not isinstance(block, dict) or block.get("type") != "text":
1289 continue
1290 text_value = block.get("text")
1291 if text_value is None:
1292 continue
1293 if mode == "unmask":
1294 block["text"] = self._unmask_pii_text(text_value, pii_tokens)
1295 elif mode == "mask":
1296 block["text"] = await self.check_pii(
1297 text=text_value,
1298 output_parse_pii=False,
1299 presidio_config=presidio_config,
1300 request_data=request_data,
1301 )
1303 return response
1305 async def _process_response_for_pii(
1306 self,
1307 response: ModelResponse,
1308 request_data: dict,
1309 mode: Literal["mask", "unmask"],
1310 ) -> ModelResponse:
1311 """
1312 Helper to recursively process a ModelResponse for PII.
1313 Handles all choices and tool calls.
1314 """
1315 metadata: Final = (request_data.get("metadata") or {}) if request_data else {}
1316 pii_tokens: Final = metadata.get("pii_tokens", {})
1317 if not pii_tokens and mode == "unmask":
1318 verbose_proxy_logger.debug("No pii_tokens found in request_data['metadata'] — nothing to unmask")
1319 presidio_config: Final = self.get_presidio_settings_from_request_data(request_data or {})
1321 for choice in response.choices:
1322 message = getattr(choice, "message", None)
1323 if message is None:
1324 continue
1326 # 1. Process content
1327 content = getattr(message, "content", None)
1328 if isinstance(content, str):
1329 if mode == "unmask":
1330 message.content = self._unmask_pii_text(content, pii_tokens)
1331 elif mode == "mask":
1332 message.content = await self.check_pii(
1333 text=content,
1334 output_parse_pii=False,
1335 presidio_config=presidio_config,
1336 request_data=request_data,
1337 )
1338 elif isinstance(content, list):
1339 for item in content:
1340 if not isinstance(item, dict):
1341 continue
1342 text_value = item.get("text")
1343 if text_value is None:
1344 continue
1345 if mode == "unmask":
1346 item["text"] = self._unmask_pii_text(text_value, pii_tokens)
1347 elif mode == "mask":
1348 item["text"] = await self.check_pii(
1349 text=text_value,
1350 output_parse_pii=False,
1351 presidio_config=presidio_config,
1352 request_data=request_data,
1353 )
1355 # 2. Process tool calls
1356 tool_calls = getattr(message, "tool_calls", None)
1357 if tool_calls:
1358 for tool_call in tool_calls:
1359 function = getattr(tool_call, "function", None)
1360 if function and hasattr(function, "arguments"):
1361 args = function.arguments
1362 if isinstance(args, str):
1363 if mode == "unmask":
1364 function.arguments = self._unmask_pii_text(args, pii_tokens)
1365 elif mode == "mask":
1366 function.arguments = await self.check_pii(
1367 text=args,
1368 output_parse_pii=False,
1369 presidio_config=presidio_config,
1370 request_data=request_data,
1371 )
1373 # 3. Process legacy function calls
1374 function_call = getattr(message, "function_call", None)
1375 if function_call and hasattr(function_call, "arguments"):
1376 args = function_call.arguments
1377 if isinstance(args, str):
1378 if mode == "unmask":
1379 function_call.arguments = self._unmask_pii_text(args, pii_tokens)
1380 elif mode == "mask":
1381 function_call.arguments = await self.check_pii(
1382 text=args,
1383 output_parse_pii=False,
1384 presidio_config=presidio_config,
1385 request_data=request_data,
1386 )
1387 return response
1389 async def _mask_output_response(
1390 self,
1391 response: ModelResponse | EmbeddingResponse | ImageResponse,
1392 request_data: dict,
1393 ):
1394 """
1395 Apply Presidio masking on model responses (non-streaming).
1396 """
1397 if not isinstance(response, ModelResponse):
1398 return response
1400 # skip streaming here; handled in async_post_call_streaming_iterator_hook
1401 if isinstance(response, ModelResponseStream):
1402 return response
1404 await self._process_response_for_pii(
1405 response=response,
1406 request_data=request_data,
1407 mode="mask",
1408 )
1409 return response
1411 async def _mask_buffered_model_response_stream(
1412 self, all_chunks: Sequence[ModelResponseStream], request_data: dict
1413 ) -> tuple[object, ...]:
1414 from litellm.llms.base_llm.base_model_iterator import (
1415 convert_model_response_to_streaming,
1416 )
1417 from litellm.main import stream_chunk_builder
1418 from litellm.types.utils import ModelResponse
1420 assembled: Final = stream_chunk_builder(chunks=list(all_chunks), messages=request_data.get("messages"))
1421 if not isinstance(assembled, ModelResponse):
1422 return tuple(all_chunks)
1423 await self._process_response_for_pii(response=assembled, request_data=request_data, mode="mask")
1424 return (convert_model_response_to_streaming(assembled),)
1426 async def _stream_apply_output_masking(
1427 self,
1428 response: AsyncIterable[object],
1429 request_data: dict,
1430 ) -> AsyncGenerator[object, None]:
1431 """Apply Presidio masking to streaming output (apply_to_output=True path)."""
1432 all_chunks: list[ModelResponseStream] = []
1433 passthrough_due_to_unknown_stream_shape = False
1434 try:
1435 stream: Final = _coalesce_first_sse_frame(response.__aiter__())
1436 async for chunk in stream:
1437 if isinstance(chunk, ModelResponseStream):
1438 if passthrough_due_to_unknown_stream_shape:
1439 yield chunk
1440 else:
1441 all_chunks.append(chunk)
1442 elif isinstance(chunk, _SsePreface):
1443 yield chunk.raw
1444 elif isinstance(chunk, bytes):
1445 first_frame_is_anthropic = (
1446 not passthrough_due_to_unknown_stream_shape
1447 and not all_chunks
1448 and is_anthropic_sse_stream((chunk,))
1449 )
1450 if not first_frame_is_anthropic:
1451 passthrough_due_to_unknown_stream_shape = (
1452 passthrough_due_to_unknown_stream_shape or not all_chunks
1453 )
1454 yield chunk
1455 continue
1456 for masked_chunk in await self._mask_anthropic_sse_stream(chunk, stream, request_data):
1457 yield masked_chunk
1458 return
1459 else:
1460 for buffered_chunk in _flush_unmaskable_buffer(all_chunks):
1461 yield buffered_chunk
1462 all_chunks = []
1463 passthrough_due_to_unknown_stream_shape = True
1464 yield chunk
1465 if passthrough_due_to_unknown_stream_shape:
1466 verbose_proxy_logger.warning(
1467 "Presidio apply_to_output: streaming response was not a parsed chat completion stream "
1468 "(raw non-Anthropic SSE passthrough or /v1/responses events). "
1469 "Output PII masking was skipped for this response."
1470 )
1471 return
1472 if not all_chunks:
1473 verbose_proxy_logger.warning(
1474 "Presidio apply_to_output: streaming response contained no "
1475 "ModelResponseStream chunks (an empty upstream stream). "
1476 "Output PII masking was skipped for this response."
1477 )
1478 return
1480 for masked_chunk in await self._mask_buffered_model_response_stream(all_chunks, request_data):
1481 yield masked_chunk
1483 except Exception as e:
1484 if not all_chunks or isinstance(e, BlockedPiiEntityError):
1485 raise
1486 verbose_proxy_logger.error("Error masking streaming PII output: %s", e)
1487 for chunk in all_chunks:
1488 yield chunk
1490 async def _mask_anthropic_sse_stream(
1491 self, first_chunk: bytes, rest: AsyncIterator[object], request_data: dict
1492 ) -> tuple[object, ...]:
1493 rest_chunks: Final = [chunk async for chunk in rest] # mutable-ok: tuple() cannot consume an async iterator
1494 chunks: Final = (first_chunk, *rest_chunks)
1495 assembled: Final = assemble_anthropic_sse_stream(chunks, restore_identity=True)
1496 if assembled is None:
1497 verbose_proxy_logger.warning(
1498 "Presidio apply_to_output: raw SSE stream could not be assembled into a response. "
1499 "Output PII masking was skipped for this response."
1500 )
1501 return chunks
1502 original_text: Final = model_response_text(assembled)
1503 await self._process_response_for_pii(response=assembled, request_data=request_data, mode="mask")
1504 if model_response_text(assembled) == original_text:
1505 return chunks
1506 return anthropic_sse_chunks_from_response(assembled)
1508 @staticmethod
1509 def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: dict[str, str]) -> bytes:
1510 try:
1511 text: Final = chunk.decode("utf-8")
1512 except UnicodeDecodeError:
1513 return chunk
1515 result_lines: Final[list[str]] = []
1516 for line in text.split("\n"):
1517 line = line.rstrip("\r")
1518 if line.startswith("data: ") and line != "data: [DONE]":
1519 raw_json = line[6:]
1520 try:
1521 event = json.loads(raw_json)
1522 delta = event.get("delta") if isinstance(event, dict) else None
1523 if (
1524 isinstance(delta, dict)
1525 and event.get("type") == "content_block_delta"
1526 and delta.get("type") == "text_delta"
1527 and isinstance(delta.get("text"), str)
1528 ):
1529 unmasked = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(delta["text"], pii_tokens)
1530 if unmasked != delta["text"]:
1531 event["delta"]["text"] = unmasked
1532 line = "data: " + json.dumps(event, ensure_ascii=False)
1533 except (json.JSONDecodeError, KeyError, TypeError):
1534 pass
1535 result_lines.append(line)
1537 return "\n".join(result_lines).encode("utf-8")
1539 def _unmask_responses_api_completed_chunk(self, chunk: object, pii_tokens: dict[str, str]) -> None:
1540 """
1541 Unmask PII tokens in-place for a ``response.completed`` Responses API event.
1543 The chunk carries a ``response`` attribute (ResponsesAPIResponse) whose
1544 ``output`` list holds message items. Each item has a ``content`` list of
1545 blocks; text blocks expose a ``.text`` string attribute. We walk the tree
1546 and replace every PII token with its original value.
1547 """
1548 response_obj: Final[object] = getattr(chunk, "response", None)
1549 if response_obj is None:
1550 return
1552 output: Final = getattr(response_obj, "output", None) or []
1553 for output_item in output:
1554 content = getattr(output_item, "content", None) or []
1555 for content_block in content:
1556 if isinstance(content_block, dict):
1557 if isinstance(content_block.get("text"), str):
1558 content_block["text"] = self._unmask_pii_text(content_block["text"], pii_tokens)
1559 elif hasattr(content_block, "text") and isinstance(content_block.text, str):
1560 content_block.text = self._unmask_pii_text(content_block.text, pii_tokens)
1562 async def _stream_pii_unmasking(
1563 self,
1564 response: AsyncIterable[object],
1565 request_data: dict,
1566 ) -> AsyncGenerator[object, None]:
1567 """Apply PII unmasking to streaming output (output_parse_pii=True path)."""
1568 from litellm.llms.base_llm.base_model_iterator import (
1569 convert_model_response_to_streaming,
1570 )
1571 from litellm.main import stream_chunk_builder
1572 from litellm.types.utils import ModelResponse
1574 metadata: Final = (request_data.get("metadata") or {}) if request_data else {}
1575 pii_tokens: Final[dict[str, str]] = metadata.get("pii_tokens", {})
1577 remaining_chunks: list[ModelResponseStream] = []
1578 saw_non_chat_chunk = False
1579 try:
1580 async for chunk in response:
1581 if isinstance(chunk, ModelResponseStream):
1582 if saw_non_chat_chunk:
1583 yield chunk
1584 else:
1585 remaining_chunks.append(chunk)
1586 elif isinstance(chunk, bytes):
1587 if pii_tokens:
1588 yield self._unmask_sse_bytes_chunk(chunk, pii_tokens)
1589 else:
1590 yield chunk
1591 continue
1592 else:
1593 # /v1/responses events: unmask response.completed text in-place.
1594 # A mixed stream can't be reassembled, so flush buffered chat
1595 # chunks in order before passthrough instead of dropping them.
1596 if remaining_chunks and not saw_non_chat_chunk:
1597 for buffered_chunk in remaining_chunks:
1598 yield buffered_chunk
1599 remaining_chunks = []
1600 chunk_type = getattr(chunk, "type", None)
1601 if chunk_type == "response.completed" and pii_tokens:
1602 self._unmask_responses_api_completed_chunk(chunk, pii_tokens)
1603 saw_non_chat_chunk = True
1604 yield chunk
1606 if saw_non_chat_chunk:
1607 return
1609 if not remaining_chunks:
1610 return
1612 assembled_model_response: Final = stream_chunk_builder(
1613 chunks=remaining_chunks, messages=request_data.get("messages")
1614 )
1616 if not isinstance(assembled_model_response, ModelResponse):
1617 for chunk in remaining_chunks:
1618 yield chunk
1619 return
1621 self._preserve_usage_from_last_chunk(assembled_model_response, remaining_chunks)
1623 await self._process_response_for_pii(
1624 response=assembled_model_response,
1625 request_data=request_data,
1626 mode="unmask",
1627 )
1629 mock_response_stream: Final = convert_model_response_to_streaming(assembled_model_response)
1630 yield mock_response_stream
1632 except Exception as e:
1633 verbose_proxy_logger.error("Error in PII streaming processing: %s", e)
1634 for chunk in remaining_chunks:
1635 yield chunk
1637 async def async_post_call_streaming_iterator_hook(
1638 self,
1639 user_api_key_dict: UserAPIKeyAuth,
1640 response: AsyncIterable[object],
1641 request_data: dict,
1642 ) -> AsyncGenerator[object, None]:
1643 """
1644 Process streaming response chunks to unmask PII tokens when needed.
1646 Note: the return type includes `bytes` because Anthropic native SSE
1647 streaming sends raw bytes chunks that pass through untransformed.
1648 The base class declares ModelResponseStream only.
1649 """
1650 if self.apply_to_output:
1651 async for chunk in self._stream_apply_output_masking(response, request_data):
1652 yield chunk
1653 return
1655 metadata: Final = (request_data.get("metadata") or {}) if request_data else {}
1656 pii_tokens: Final = metadata.get("pii_tokens", {})
1657 if not pii_tokens and request_data:
1658 verbose_proxy_logger.debug("No pii_tokens in request_data['metadata'] for streaming unmask path")
1659 if not (self.output_parse_pii and pii_tokens):
1660 async for chunk in response:
1661 yield chunk
1662 return
1664 async for chunk in self._stream_pii_unmasking(response, request_data):
1665 yield chunk
1667 @staticmethod
1668 def _preserve_usage_from_last_chunk(
1669 assembled_model_response: ModelResponse,
1670 chunks: list[ModelResponseStream],
1671 ) -> None:
1672 """Copy usage metadata from the last chunk when stream_chunk_builder misses it."""
1673 if not getattr(assembled_model_response, "usage", None) and chunks:
1674 last_chunk_usage: Final = getattr(chunks[-1], "usage", None)
1675 if last_chunk_usage:
1676 setattr(assembled_model_response, "usage", last_chunk_usage)
1678 def get_presidio_settings_from_request_data(self, data: dict) -> PresidioPerRequestConfig | None:
1679 if "metadata" in data:
1680 _metadata: Final = data.get("metadata", None)
1681 if _metadata is None:
1682 return None
1683 _guardrail_config: Final = _metadata.get("guardrail_config")
1684 if _guardrail_config:
1685 _presidio_config: Final = PresidioPerRequestConfig(**_guardrail_config)
1686 return _presidio_config
1688 return None
1690 def print_verbose(self, print_statement):
1691 try:
1692 verbose_proxy_logger.debug(print_statement)
1693 if litellm.set_verbose:
1694 print(print_statement) # noqa: T201
1695 except Exception:
1696 pass
1698 @log_guardrail_information
1699 async def apply_guardrail(
1700 self,
1701 inputs: "GenericGuardrailAPIInputs",
1702 request_data: dict,
1703 input_type: Literal["request", "response"],
1704 logging_obj: Optional["LiteLLMLoggingObj"] = None,
1705 ) -> "GenericGuardrailAPIInputs":
1706 """
1707 UI will call this function to check:
1708 1. If the connection to the guardrail is working
1709 2. When Testing the guardrail with some text, this function will be called with the input text and returns a text after applying the guardrail
1710 """
1711 texts: Final = inputs.get("texts", [])
1713 # When input_type is "response" and pii_tokens are available,
1714 # unmask the text instead of masking it.
1715 metadata: Final = (request_data.get("metadata") or {}) if request_data else {}
1716 pii_tokens: Final = metadata.get("pii_tokens", {})
1718 new_texts: Final = []
1719 if input_type == "response" and pii_tokens:
1720 for text in texts:
1721 new_texts.append(self._unmask_pii_text(text, pii_tokens))
1722 else:
1723 for text in texts:
1724 modified_text = await self.check_pii(
1725 text=text,
1726 output_parse_pii=self.output_parse_pii,
1727 presidio_config=None,
1728 request_data=request_data or {},
1729 )
1730 new_texts.append(modified_text)
1731 inputs["texts"] = new_texts
1732 return inputs
1734 def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None:
1735 """
1736 Update the guardrails litellm params in memory
1737 """
1738 super().update_in_memory_litellm_params(litellm_params)
1739 if self.apply_to_output:
1740 self.output_parse_pii = False
1741 if litellm_params.pii_entities_config:
1742 self.pii_entities_config = litellm_params.pii_entities_config
1743 if litellm_params.presidio_score_thresholds:
1744 self.presidio_score_thresholds = litellm_params.presidio_score_thresholds
1745 if litellm_params.presidio_entities_deny_list:
1746 self.presidio_entities_deny_list = litellm_params.presidio_entities_deny_list
1747 if litellm_params.presidio_analyze_chunk_size_bytes is not None:
1748 # Same validation as __init__: a non-positive value from a guardrail
1749 # update must not silently disable detection via degenerate chunking.
1750 self.presidio_analyze_chunk_size_bytes = self._coerce_analyze_chunk_size(
1751 litellm_params.presidio_analyze_chunk_size_bytes
1752 )