diff --git a/litellm/proxy/guardrails/guardrail_hooks/ztds.py b/litellm/proxy/guardrails/guardrail_hooks/ztds.py index 3522d070dfb..69b1f29d486 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ztds.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ztds.py @@ -17,14 +17,14 @@ from __future__ import annotations import re import uuid from collections.abc import AsyncGenerator -from typing import Any, ClassVar +from typing import ClassVar try: from litellm.integrations.custom_guardrail import CustomGuardrail except ImportError: # Standalone fallback when running outside full LiteLLM package class CustomGuardrail: - def __init__(self, **kwargs: Any) -> None: + def __init__(self, **kwargs: object) -> None: for k, v in kwargs.items(): setattr(self, k, v) @@ -55,8 +55,8 @@ class ZTDSGuardrail(CustomGuardrail): reverse_on_output: bool = True, enforce_zero_egress: bool = True, guardrail_name: str | None = "ztds", - **kwargs: Any, - ): + **kwargs: object, + ) -> None: super().__init__(guardrail_name=guardrail_name, **kwargs) self.enabled_entities = enabled_entities or list(self.PATTERNS.keys()) self.reverse_on_output = reverse_on_output @@ -156,11 +156,11 @@ class ZTDSGuardrail(CustomGuardrail): async def async_pre_call_hook( self, - user_api_key_dict: Any, - cache: Any, - data: dict[str, Any], + user_api_key_dict: object, + cache: object, + data: dict[str, object], call_type: str, - ) -> dict[str, Any]: + ) -> dict[str, object]: """ LiteLLM pre-call hook: intercepts outgoing messages, prompts, and inputs and sanitizes all content. Generates an internal random nonce to prevent cross-tenant ID collisions. @@ -223,16 +223,16 @@ class ZTDSGuardrail(CustomGuardrail): async def async_post_call_success_hook( self, - data: dict[str, Any], - user_api_key_dict: Any, - response: Any, - ) -> Any: + data: dict[str, object], + user_api_key_dict: object, + response: object, + ) -> object: """ 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. """ session_id = data.get("_ztds_session_id") - if not session_id: + if not session_id or not isinstance(session_id, str): return response try: @@ -259,51 +259,52 @@ class ZTDSGuardrail(CustomGuardrail): async def async_post_call_failure_hook( self, - data: dict[str, Any], - user_api_key_dict: Any, + data: dict[str, object], + user_api_key_dict: object, error: Exception, ) -> None: """ LiteLLM post-call failure hook: ensures volatile RAM zeroization when upstream provider calls fail. """ session_id = data.get("_ztds_session_id") if isinstance(data, dict) else None - if session_id: + if session_id and isinstance(session_id, str): self.zeroize_session(session_id) async def async_post_call_streaming_iterator_hook( self, - user_api_key_dict: Any, - response: Any, - request_data: dict, - ) -> AsyncGenerator[Any, None]: + user_api_key_dict: object, + response: object, + request_data: dict[str, object], + ) -> AsyncGenerator[object, None]: """ LiteLLM streaming iterator hook: restores tokens across streaming response chunks in volatile RAM and guarantees Theorem 2 zeroization upon stream completion or error. """ 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 self.reverse_on_output: - if hasattr(chunk, "choices") and chunk.choices: - for choice in 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"]: - 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 + if hasattr(response, "__aiter__"): + async for chunk in response: # type: ignore[union-attr] + 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: + 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"]: + 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 finally: - if session_id: + if session_id and isinstance(session_id, str): self.zeroize_session(session_id) async def async_post_call_streaming_hook( self, - user_api_key_dict: Any, + user_api_key_dict: object, response: str, - ) -> Any: + ) -> object: """ LiteLLM post-call streaming hook fallback. """