diff --git a/tests/guardrails_tests/test_ztds_guardrail.py b/tests/guardrails_tests/test_ztds_guardrail.py index 2d715ab40d8..be4b6a85da1 100644 --- a/tests/guardrails_tests/test_ztds_guardrail.py +++ b/tests/guardrails_tests/test_ztds_guardrail.py @@ -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__":