fix(ztds): resolve cache plaintext leak with deepcopy, eliminate quadratic tokenization, and expand sk- key regex

This commit is contained in:
Ilya Sibiryakov 2026-10-02 15:45:47 +03:00
parent e23994a36c
commit 9dcd3a4708
2 changed files with 84 additions and 23 deletions

View file

@ -14,6 +14,7 @@ Invariants Enforced:
from __future__ import annotations
import copy
import re
import uuid
from collections.abc import AsyncGenerator, AsyncIterable, Mapping
@ -51,7 +52,7 @@ class ZTDSGuardrail(CustomGuardrail):
"SSN": re.compile(r"\b\d{3}-\d{2}-\d{4}\b"),
"PHONE": re.compile(r"\b(?:\+?\d{1,3}[-.\s]?)?\(?\d{3}\)?[-.\s]?\d{3}[-.\s]?\d{4}\b"),
"API_SECRET": re.compile(
r"\b(?:sk-(?:proj-)?[A-Za-z0-9_-]{20,}|ghp_[A-Za-z0-9]{20,}|eyJ[A-Za-z0-9_-]{20,}\.[A-Za-z0-9_-]{20,}\.[A-Za-z0-9_-]{20,})\b"
r"\b(?:sk-[a-zA-Z0-9_-]{20,}|ghp_[a-zA-Z0-9]{20,}|eyJ[a-zA-Z0-9_-]{20,}\.[a-zA-Z0-9_-]{20,}\.[a-zA-Z0-9_-]{20,})\b"
),
}
)
@ -97,6 +98,9 @@ class ZTDSGuardrail(CustomGuardrail):
entity_map = self._entity_maps[session_id]
caller_set = self._caller_tokens[session_id]
# Pre-index existing bracketed tokens in text to prevent collisions in O(1)
existing_tokens = set(self.TOKEN_PATTERN.findall(text))
sanitized = text
for entity_type in self.enabled_entities:
if entity_type == "EMAIL" and "@" not in sanitized:
@ -106,30 +110,29 @@ class ZTDSGuardrail(CustomGuardrail):
if not pattern:
continue
# Process matches in reverse string order to preserve exact substring indices
matches = tuple(pattern.finditer(sanitized))
for match in reversed(matches):
original = match.group(0)
# Per-entity surrogate counter to eliminate quadratic scans over token_map
entity_counter = sum(1 for k in token_map if k.startswith(f"[{entity_type}_TOKEN_"))
# Deterministic Reversible Tokenization (Invariant 2) with Collision Avoidance
def _replace_match(match: re.Match, et: str = entity_type) -> str:
nonlocal entity_counter
original = match.group(0)
if original in entity_map:
token = entity_map[original]
else:
count = sum(1 for k in token_map if k.startswith(f"[{entity_type}_TOKEN_")) + 1
while True:
candidate = f"[{entity_type}_TOKEN_{count}]"
if candidate not in text and candidate not in token_map:
entity_counter += 1
candidate = f"[{et}_TOKEN_{entity_counter}]"
if candidate not in existing_tokens and candidate not in token_map:
token = candidate
break
count += 1
token_map[token] = original
entity_map[original] = token
if is_caller_visible:
caller_set.add(token)
return token
start, end = match.span()
sanitized = sanitized[:start] + token + sanitized[end:]
sanitized = pattern.sub(_replace_match, sanitized)
return sanitized, token_map
@ -245,6 +248,7 @@ class ZTDSGuardrail(CustomGuardrail):
"""
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.
Deep-copies response before unmasking so that upstream shared caches retain sanitized surrogates.
"""
session_id = data.get("_ztds_session_id")
if not session_id or not isinstance(session_id, str):
@ -252,9 +256,10 @@ class ZTDSGuardrail(CustomGuardrail):
try:
if self.reverse_on_output:
caller_response = copy.deepcopy(response)
# Process standard ModelResponse object
if hasattr(response, "choices") and response.choices:
for choice in response.choices:
if hasattr(caller_response, "choices") and caller_response.choices:
for choice in caller_response.choices:
if (
hasattr(choice, "message")
and hasattr(choice.message, "content")
@ -262,10 +267,13 @@ class ZTDSGuardrail(CustomGuardrail):
):
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)
elif isinstance(caller_response, dict) and "choices" in caller_response:
for choice in caller_response["choices"]:
if isinstance(choice, dict) and "message" in choice and isinstance(choice["message"], dict):
content = choice["message"].get("content")
if isinstance(content, str):
choice["message"]["content"] = self.restore_text(content, session_id)
return caller_response
finally:
# Theorem 2: Guarantee RAM zeroization even if reverse_on_output is False or response handling fails
self.zeroize_session(session_id)
@ -294,22 +302,26 @@ class ZTDSGuardrail(CustomGuardrail):
"""
LiteLLM streaming iterator hook: restores tokens across streaming response chunks in volatile RAM
and guarantees Theorem 2 zeroization upon stream completion or error.
Deep-copies chunk before unmasking so upstream completion cache retains sanitized surrogates.
"""
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 isinstance(session_id, str) and self.reverse_on_output:
if hasattr(chunk, "choices") and chunk.choices:
for choice in chunk.choices:
caller_chunk = copy.deepcopy(chunk)
if hasattr(caller_chunk, "choices") and caller_chunk.choices:
for choice in caller_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"]:
elif isinstance(caller_chunk, dict) and "choices" in caller_chunk:
for choice in caller_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
yield caller_chunk
else:
yield chunk
finally:
if session_id and isinstance(session_id, str):
self.zeroize_session(session_id)

View file

@ -261,6 +261,55 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase):
restored = self.guardrail.restore_text(sanitized, session_id)
self.assertEqual(restored, raw)
def test_anthropic_and_hyphenated_api_keys(self):
"""Verify detection of Anthropic sk-ant- keys and hyphenated API tokens."""
session_id = "test-session-anthropic"
secret = "sk-ant-api03-abcdefghijklmnopqrstuvwxyz123456"
raw = f"Anthropic token: {secret}"
sanitized, _ = self.guardrail.sanitize_text(raw, session_id)
self.assertNotIn(secret, sanitized)
self.assertIn("[API_SECRET_TOKEN_1]", sanitized)
def test_high_volume_linear_tokenization_performance(self):
"""Verify that 2,000 distinct email tokens execute in linear time (< 0.5s)."""
import time
session_id = "test-session-scale"
payload = " ".join(f"user_{i}@enterprise-corp.com" for i in range(2000))
t0 = time.perf_counter()
sanitized, token_map = self.guardrail.sanitize_text(payload, session_id)
elapsed = time.perf_counter() - t0
self.assertLess(elapsed, 0.5, f"Quadratic tokenization regression: took {elapsed:.4f}s")
self.assertEqual(len(token_map), 2000)
self.assertNotIn("user_0@enterprise-corp.com", sanitized)
async def test_streaming_chunk_deepcopy_preserves_cache_immutability(self):
"""Veria AI security fix: stream chunk deepcopy ensures upstream cache retains surrogates."""
session_id = "test-session-stream-cache"
self.guardrail.sanitize_text("user@corp.com", session_id)
original_chunk = MockStreamChunk("Here is [EMAIL_TOKEN_1]")
async def _generator():
yield original_chunk
request_data = {"_ztds_session_id": session_id}
stream_iter = self.guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict={},
response=_generator(),
request_data=request_data,
)
chunks_received = []
async for c in stream_iter:
chunks_received.append(c)
# Caller receives restored cleartext
self.assertEqual(chunks_received[0].choices[0].delta.content, "Here is user@corp.com")
# Original chunk object retains sanitized surrogate for completion cache
self.assertEqual(original_chunk.choices[0].delta.content, "Here is [EMAIL_TOKEN_1]")
if __name__ == "__main__":
unittest.main()