mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
test(ztds): add proxy lifecycle, non-message sanitization, failure and streaming tests
This commit is contained in:
parent
577b54b1ec
commit
7dd8632471
1 changed files with 83 additions and 3 deletions
|
|
@ -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__":
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue