mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): harden session isolation, payload coverage, and failure cleanup in ZTDS
This commit is contained in:
parent
ea7d0ba0c4
commit
6ed6eb24f1
1 changed files with 62 additions and 19 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue