fix(ztds): add streaming iterator hook and proxy constructor kwargs

This commit is contained in:
Ilya Sibiryakov 2026-09-27 09:24:21 +03:00
parent 6076be7d78
commit 577b54b1ec

View file

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