mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(ztds): add streaming iterator hook and proxy constructor kwargs
This commit is contained in:
parent
6076be7d78
commit
577b54b1ec
1 changed files with 53 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue