fix(guardrails): harden session isolation, payload coverage, and failure cleanup in ZTDS

This commit is contained in:
Ilya Sibiryakov 2026-09-27 09:17:44 +03:00
parent ea7d0ba0c4
commit 6ed6eb24f1

View file

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