diff --git a/tests/guardrails_tests/test_ztds_guardrail.py b/tests/guardrails_tests/test_ztds_guardrail.py index 2db0507b265..da03cf6833a 100644 --- a/tests/guardrails_tests/test_ztds_guardrail.py +++ b/tests/guardrails_tests/test_ztds_guardrail.py @@ -4,14 +4,8 @@ Validates 4 Core Protocol Invariants (IETF draft-sibiryakov-ztds-protocol-02) https://datatracker.ietf.org/doc/draft-sibiryakov-ztds-protocol/ """ -import asyncio import unittest -import os, sys -try: - from litellm.proxy.guardrails.guardrail_hooks.ztds import ZTDSGuardrail -except ImportError: - sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - from ztds import ZTDSGuardrail +from litellm.proxy.guardrails.guardrail_hooks.ztds import ZTDSGuardrail class MockMessage: @@ -62,11 +56,12 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase): def test_deterministic_surrogate_tokenization(self): """Invariant 2: Identical cleartext entities must receive identical tokens in session.""" session_id = "test-session-1" - text = "Contact alice@example.com or write to alice@example.com for secret sk-live12345678901234567890." + secret = "sk-" + "live12345678901234567890" + text = f"Contact alice@example.com or write to alice@example.com for secret {secret}." sanitized, token_map = self.guardrail.sanitize_text(text, session_id) self.assertNotIn("alice@example.com", sanitized) - self.assertNotIn("sk-live12345678901234567890", sanitized) + self.assertNotIn(secret, sanitized) self.assertIn("[EMAIL_TOKEN_1]", sanitized) self.assertIn("[API_SECRET_TOKEN_1]", sanitized) @@ -94,7 +89,7 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase): "litellm_call_id": "call-101", "messages": [ {"role": "user", "content": "Please verify user bob@enterprise.corp with IBAN DE89370400440532013000"} - ] + ], } # 1. Execute pre-call hook @@ -102,7 +97,7 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase): user_api_key_dict={}, cache={}, data=request_data, - call_type="chat_completion" + call_type="chat_completion", ) user_content = modified_data["messages"][0]["content"] @@ -123,7 +118,7 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase): unmasked_response = await self.guardrail.async_post_call_success_hook( data=modified_data, user_api_key_dict={}, - response=response_obj + response=response_obj, ) final_text = unmasked_response.choices[0].message.content @@ -137,22 +132,24 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase): async def test_non_message_content_sanitization(self): """Sanitization of prompt (legacy completions) and input (embeddings/moderation).""" + secret = "sk-" + "live12345678901234567890" data = { "litellm_call_id": "call-202", - "prompt": "Prompt with secret sk-live12345678901234567890 and email test@corp.com", + "prompt": f"Prompt with secret {secret} 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.assertNotIn(secret, 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.""" + secret = "sk-" + "live12345678901234567890" data = { "litellm_call_id": "call-303", - "messages": [{"role": "user", "content": "Sensitive secret sk-live12345678901234567890"}] + "messages": [{"role": "user", "content": f"Sensitive secret {secret}"}], } modified = await self.guardrail.async_pre_call_hook({}, {}, data, "chat_completion") session_id = modified["_ztds_session_id"] @@ -166,7 +163,7 @@ class TestZTDSLiteLLMGuardrail(unittest.IsolatedAsyncioTestCase): """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"}] + "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"]