From 6ed6eb24f1061943398e609d37793ba4669ef323 Mon Sep 17 00:00:00 2001 From: Ilya Sibiryakov Date: Sun, 27 Sep 2026 09:17:44 +0300 Subject: [PATCH] fix(guardrails): harden session isolation, payload coverage, and failure cleanup in ZTDS --- .../proxy/guardrails/guardrail_hooks/ztds.py | 81 ++++++++++++++----- 1 file changed, 62 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/ztds.py b/litellm/proxy/guardrails/guardrail_hooks/ztds.py index b979df558ee..9a1b0b698c9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ztds.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ztds.py @@ -14,14 +14,15 @@ Invariants Enforced: import re import uuid -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any, Dict, List, Optional, Tuple try: from litellm.integrations.custom_guardrail import CustomGuardrail except ImportError: # Standalone fallback when running outside full LiteLLM package class CustomGuardrail: - pass + def __init__(self, **kwargs: Any) -> None: + pass class ZTDSGuardrail(CustomGuardrail): @@ -33,7 +34,7 @@ class ZTDSGuardrail(CustomGuardrail): # Comprehensive zero-egress regex patterns for sensitive identifiers PATTERNS: Dict[str, re.Pattern] = { - "EMAIL": re.compile(r"\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,7}\b"), + "EMAIL": re.compile(r"\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,24}\b"), "IPV4": re.compile(r"\b(?:\d{1,3}\.){3}\d{1,3}\b"), "IBAN": re.compile(r"\b[A-Z]{2}[0-9]{2}[A-Z0-9]{4}[0-9]{7}([A-Z0-9]?){0,16}\b"), "CREDIT_CARD": re.compile(r"\b(?:\d{4}[-\s]?){3}\d{4}\b"), @@ -47,8 +48,9 @@ class ZTDSGuardrail(CustomGuardrail): enabled_entities: Optional[List[str]] = None, reverse_on_output: bool = True, enforce_zero_egress: bool = True, + **kwargs: Any, ): - super().__init__() + super().__init__(**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 @@ -126,12 +128,15 @@ class ZTDSGuardrail(CustomGuardrail): call_type: str, ) -> Dict[str, Any]: """ - LiteLLM pre-call hook: intercepts outgoing messages and sanitizes all content. + LiteLLM pre-call hook: intercepts outgoing messages, prompts, and inputs and sanitizes all content. + Generates an internal random nonce to prevent cross-tenant ID collisions. Zero network sockets are opened during this operation. """ - session_id = data.get("litellm_call_id") or str(uuid.uuid4()) + raw_call_id = data.get("litellm_call_id") or "call" + session_id = f"{raw_call_id}_{uuid.uuid4().hex}" data["_ztds_session_id"] = session_id + # 1. Sanitize messages array (chat completions) messages = data.get("messages") if isinstance(messages, list): for message in messages: @@ -146,6 +151,28 @@ class ZTDSGuardrail(CustomGuardrail): if isinstance(chunk, dict) and chunk.get("type") == "text": chunk["text"], _ = self.sanitize_text(chunk.get("text", ""), session_id) + # 2. Sanitize prompt field (legacy completions) + if "prompt" in data: + prompt = data["prompt"] + if isinstance(prompt, str): + data["prompt"], _ = self.sanitize_text(prompt, session_id) + elif isinstance(prompt, list): + data["prompt"] = [ + self.sanitize_text(p, session_id)[0] if isinstance(p, str) else p + for p in prompt + ] + + # 3. Sanitize input field (moderations, embeddings, responses) + if "input" in data: + raw_input = data["input"] + if isinstance(raw_input, str): + data["input"], _ = self.sanitize_text(raw_input, session_id) + elif isinstance(raw_input, list): + data["input"] = [ + self.sanitize_text(item, session_id)[0] if isinstance(item, str) else item + for item in raw_input + ] + # Attach ZTDS audit receipt to metadata if "metadata" not in data or data["metadata"] is None: data["metadata"] = {} @@ -163,25 +190,41 @@ class ZTDSGuardrail(CustomGuardrail): ) -> Any: """ LiteLLM post-call success hook: restores cleartext entities in volatile RAM and zeroizes session map. + Guarantees Theorem 2 cleanup in finally block regardless of reverse_on_output configuration. """ session_id = data.get("_ztds_session_id") - if not session_id or not self.reverse_on_output: + if not session_id: return response try: - # Process standard ModelResponse object - if hasattr(response, "choices") and response.choices: - for choice in response.choices: - if hasattr(choice, "message") and hasattr(choice.message, "content"): - if isinstance(choice.message.content, str): - choice.message.content = self.restore_text(choice.message.content, session_id) - # Process dictionary response fallback - elif isinstance(response, dict) and "choices" in response: - for choice in response["choices"]: - if "message" in choice and "content" in choice["message"]: - choice["message"]["content"] = self.restore_text(choice["message"]["content"], session_id) + if self.reverse_on_output: + # Process standard ModelResponse object + if hasattr(response, "choices") and response.choices: + for choice in response.choices: + if hasattr(choice, "message") and hasattr(choice.message, "content"): + if isinstance(choice.message.content, str): + choice.message.content = self.restore_text(choice.message.content, session_id) + # Process dictionary response fallback + elif isinstance(response, dict) and "choices" in response: + for choice in response["choices"]: + if "message" in choice and "content" in choice["message"]: + choice["message"]["content"] = self.restore_text(choice["message"]["content"], session_id) finally: - # Theorem 2: Guarantee RAM zeroization even if response handling fails + # Theorem 2: Guarantee RAM zeroization even if reverse_on_output is False or response handling fails self.zeroize_session(session_id) return response + + async def async_post_call_failure_hook( + self, + data: Dict[str, Any], + user_api_key_dict: Any, + error: Exception, + ) -> None: + """ + LiteLLM post-call failure hook: ensures volatile RAM zeroization when upstream provider calls fail. + """ + session_id = data.get("_ztds_session_id") + if session_id: + self.zeroize_session(session_id) +