test(ztds): add proxy lifecycle, non-message sanitization, failure and streaming tests

This commit is contained in:
Ilya Sibiryakov 2026-09-27 09:24:27 +03:00
parent 577b54b1ec
commit 7dd8632471

View file

@ -24,10 +24,36 @@ class MockModelResponse:
self.choices = [MockChoice(content)]
class MockDelta:
def __init__(self, content):
self.content = content
class MockStreamChoice:
def __init__(self, content):
self.delta = MockDelta(content)
class MockStreamChunk:
def __init__(self, content):
self.choices = [MockStreamChoice(content)]
class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase):
def setUp(self):
self.guardrail = ZTDSGuardrail()
def test_constructor_accepts_proxy_kwargs(self):
"""Verify proxy instantiation with standard guardrail configuration kwargs."""
g = ZTDSGuardrail(
guardrail_name="ztds",
event_hook=["pre_call", "post_call"],
default_on=True,
reverse_on_output=True,
)
self.assertEqual(g.guardrail_name, "ztds")
self.assertTrue(g.reverse_on_output)
def test_deterministic_surrogate_tokenization(self):
"""Invariant 2: Identical cleartext entities must receive identical tokens in session."""
session_id = "test-session-1"
@ -80,8 +106,9 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase):
self.assertIn("[IBAN_TOKEN_1]", user_content)
self.assertTrue(modified_data["metadata"]["ztds_sanitized"])
session_id = modified_data["_ztds_session_id"]
# Check session table exists in volatile RAM before post-call
self.assertIn("call-101", self.guardrail._session_maps)
self.assertIn(session_id, self.guardrail._session_maps)
# 2. Simulate model response that mentions the token
model_reply = "Verified account for [EMAIL_TOKEN_1] linked to [IBAN_TOKEN_1]."
@ -100,8 +127,61 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase):
self.assertNotIn("[EMAIL_TOKEN_1]", final_text)
# Invariant 3 / Theorem 2: Session tables MUST be completely zeroized from RAM
self.assertNotIn("call-101", self.guardrail._session_maps)
self.assertNotIn("call-101", self.guardrail._entity_maps)
self.assertNotIn(session_id, self.guardrail._session_maps)
self.assertNotIn(session_id, self.guardrail._entity_maps)
async def test_non_message_content_sanitization(self):
"""Sanitization of prompt (legacy completions) and input (embeddings/moderation)."""
data = {
"litellm_call_id": "call-202",
"prompt": "Prompt with secret sk-live12345678901234567890 and email test@corp.com",
"input": ["Batch item with email user@corp.com", "Plain string"],
}
modified = await self.guardrail.async_pre_call_hook({}, {}, data, "completion")
self.assertNotIn("sk-live12345678901234567890", modified["prompt"])
self.assertIn("[API_SECRET_TOKEN_1]", modified["prompt"])
self.assertNotIn("user@corp.com", modified["input"][0])
self.assertIn("[EMAIL_TOKEN_2]", modified["input"][0])
async def test_failure_hook_zeroizes_ram(self):
"""Theorem 2: When upstream provider fails, RAM tables must be completely wiped."""
data = {
"litellm_call_id": "call-303",
"messages": [{"role": "user", "content": "Sensitive secret sk-live12345678901234567890"}]
}
modified = await self.guardrail.async_pre_call_hook({}, {}, data, "chat_completion")
session_id = modified["_ztds_session_id"]
self.assertIn(session_id, self.guardrail._session_maps)
# Trigger failure hook
await self.guardrail.async_post_call_failure_hook(modified, {}, Exception("Upstream 500"))
self.assertNotIn(session_id, self.guardrail._session_maps)
async def test_streaming_hook_restores_and_zeroizes(self):
"""Streaming response chunks are unmasked and RAM is zeroized upon completion."""
data = {
"litellm_call_id": "call-404",
"messages": [{"role": "user", "content": "Hello user@corp.com"}]
}
modified = await self.guardrail.async_pre_call_hook({}, {}, data, "chat_completion")
session_id = modified["_ztds_session_id"]
self.assertIn(session_id, self.guardrail._session_maps)
async def fake_stream():
yield MockStreamChunk("Result for ")
yield MockStreamChunk("[EMAIL_TOKEN_1]")
yield MockStreamChunk(" confirmed.")
chunks = []
async for chunk in self.guardrail.async_post_call_streaming_iterator_hook({}, fake_stream(), modified):
chunks.append(chunk.choices[0].delta.content)
full_output = "".join(chunks)
self.assertIn("user@corp.com", full_output)
self.assertNotIn("[EMAIL_TOKEN_1]", full_output)
# Theorem 2 verification
self.assertNotIn(session_id, self.guardrail._session_maps)
if __name__ == "__main__":