mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(ztds): resolve cache plaintext leak with deepcopy, eliminate quadratic tokenization, and expand sk- key regex
This commit is contained in:
parent
e23994a36c
commit
9dcd3a4708
2 changed files with 84 additions and 23 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue