diff --git a/litellm/proxy/guardrails/guardrail_hooks/ztds.py b/litellm/proxy/guardrails/guardrail_hooks/ztds.py index 9a1b0b698c9..85efbca26ad 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ztds.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ztds.py @@ -14,7 +14,7 @@ Invariants Enforced: import re import uuid -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple try: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -48,9 +48,10 @@ class ZTDSGuardrail(CustomGuardrail): enabled_entities: Optional[List[str]] = None, reverse_on_output: bool = True, enforce_zero_egress: bool = True, + guardrail_name: Optional[str] = "ztds", **kwargs: Any, ): - super().__init__(**kwargs) + super().__init__(guardrail_name=guardrail_name, **kwargs) self.enabled_entities = enabled_entities or list(self.PATTERNS.keys()) self.reverse_on_output = reverse_on_output self.enforce_zero_egress = enforce_zero_egress @@ -61,27 +62,32 @@ class ZTDSGuardrail(CustomGuardrail): def sanitize_text(self, text: str, session_id: str) -> Tuple[str, Dict[str, str]]: """ - Deterministically sanitizes a text string in volatile memory. - Returns: (sanitized_text, token_map) + In-memory single-pass deterministic tokenization. + Guarantees zero network calls and deterministic surrogate assignment within session scope. """ + if not text or not isinstance(text, str): + return text, {} + if session_id not in self._session_maps: self._session_maps[session_id] = {} + if session_id not in self._entity_maps: self._entity_maps[session_id] = {} token_map = self._session_maps[session_id] entity_map = self._entity_maps[session_id] - sanitized = text + sanitized = text for entity_type in self.enabled_entities: pattern = self.PATTERNS.get(entity_type) if not pattern: continue + # Process matches in reverse string order to preserve exact substring indices matches = list(pattern.finditer(sanitized)) - # Sort in reverse order of start position to safely replace in string - for match in sorted(matches, key=lambda m: m.start(), reverse=True): + for match in reversed(matches): original = match.group(0) - # Reuse deterministic surrogate if same entity seen in session + + # Deterministic Reversible Tokenization (Invariant 2) if original in entity_map: token = entity_map[original] else: @@ -224,7 +230,45 @@ class ZTDSGuardrail(CustomGuardrail): """ LiteLLM post-call failure hook: ensures volatile RAM zeroization when upstream provider calls fail. """ - session_id = data.get("_ztds_session_id") + session_id = data.get("_ztds_session_id") if isinstance(data, dict) else None if session_id: self.zeroize_session(session_id) + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: Any, + response: Any, + request_data: dict, + ) -> AsyncGenerator[Any, None]: + """ + LiteLLM streaming iterator hook: restores tokens across streaming response chunks in volatile RAM + and guarantees Theorem 2 zeroization upon stream completion or error. + """ + session_id = request_data.get("_ztds_session_id") if isinstance(request_data, dict) else None + try: + async for chunk in response: + if session_id and self.reverse_on_output: + if hasattr(chunk, "choices") and chunk.choices: + for choice in chunk.choices: + delta = getattr(choice, "delta", None) + if delta and hasattr(delta, "content") and isinstance(delta.content, str): + delta.content = self.restore_text(delta.content, session_id) + elif isinstance(chunk, dict) and "choices" in chunk: + for choice in chunk["choices"]: + delta = choice.get("delta") if isinstance(choice, dict) else None + if delta and isinstance(delta, dict) and isinstance(delta.get("content"), str): + delta["content"] = self.restore_text(delta["content"], session_id) + yield chunk + finally: + if session_id: + self.zeroize_session(session_id) + + async def async_post_call_streaming_hook( + self, + user_api_key_dict: Any, + response: str, + ) -> Any: + """ + LiteLLM post-call streaming hook fallback. + """ + return response