From f20cf90db614ef94859d862a3d104865bccd3319 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 14 May 2026 11:34:00 +0000 Subject: [PATCH] fix(purview-dlp): return data after DLP pass; per-call executor; dedupe text extraction async_pre_call_hook now returns the request dict after a successful check so callers match skip-path behavior. logging_hook uses a fresh ThreadPoolExecutor per invocation like Presidio to avoid single-worker starvation. Response text extraction is centralized in _completion_response_text_parts. Co-authored-by: Sameer Kankute --- .../microsoft_purview/purview_dlp.py | 101 +++++++----------- .../guardrail_hooks/test_microsoft_purview.py | 23 ++++ 2 files changed, 62 insertions(+), 62 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py index 97f360baeaf..893fa39eac9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -21,7 +21,13 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import GuardrailStatus +from litellm.types.utils import ( + Choices, + GuardrailStatus, + ModelResponse, + TextChoices, + TextCompletionResponse, +) from .base import PurviewGuardrailBase @@ -34,7 +40,6 @@ if TYPE_CHECKING: CallTypesLiteral, EmbeddingResponse, ImageResponse, - ModelResponse, ) @@ -78,7 +83,6 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) self._logging_only = logging_only self.guardrail_provider = "microsoft_purview" - self._executor = ThreadPoolExecutor(max_workers=1) verbose_proxy_logger.info( "Initialized Microsoft Purview DLP Guardrail: %s (logging_only=%s)", guardrail_name, @@ -156,6 +160,31 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): return response + @staticmethod + def _completion_response_text_parts(result: Any) -> List[str]: + """Collect non-empty assistant text segments from chat or text completions.""" + parts: List[str] = [] + if isinstance(result, TextCompletionResponse) and result.choices: + for choice in result.choices: + if not isinstance(choice, TextChoices): + continue + raw = choice.get("text") + if isinstance(raw, str) and raw.strip(): + parts.append(raw) + elif isinstance(result, ModelResponse) and result.choices: + for choice in result.choices: + if not isinstance(choice, Choices): + continue + msg = choice.message + if msg is None: + continue + raw = msg.get("content") if isinstance(msg, dict) else getattr( + msg, "content", None + ) + if isinstance(raw, str) and raw.strip(): + parts.append(raw) + return parts + # ------------------------------------------------------------------ # Pre-call hook — DLP on prompts # ------------------------------------------------------------------ @@ -193,7 +222,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): request_data=data, block_on_violation=True, ) - return None + return data # ------------------------------------------------------------------ # Post-call hook — DLP on responses @@ -203,16 +232,9 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): self, data: dict, user_api_key_dict: "UserAPIKeyAuth", - response: Union[Any, "ModelResponse", "EmbeddingResponse", "ImageResponse"], + response: Union[Any, ModelResponse, "EmbeddingResponse", "ImageResponse"], ) -> Any: """Check LLM response against Purview DLP policies.""" - from litellm.types.utils import ( - Choices, - ModelResponse, - TextChoices, - TextCompletionResponse, - ) - user_id = self._resolve_user_id(data, user_api_key_dict) if not user_id: verbose_proxy_logger.warning( @@ -220,27 +242,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) return response - parts: List[str] = [] - - if isinstance(response, TextCompletionResponse) and response.choices: - for choice in response.choices: - if not isinstance(choice, TextChoices): - continue - raw = choice.get("text") - if isinstance(raw, str) and raw.strip(): - parts.append(raw) - elif isinstance(response, ModelResponse) and response.choices: - for choice in response.choices: - if not isinstance(choice, Choices): - continue - msg = choice.message - if msg is None: - continue - raw = msg.get("content") if isinstance(msg, dict) else getattr( - msg, "content", None - ) - if isinstance(raw, str) and raw.strip(): - parts.append(raw) + parts = self._completion_response_text_parts(response) if parts: combined = "\n\n---\n\n".join(parts) @@ -277,8 +279,9 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): try: _ = asyncio.get_running_loop() - future = self._executor.submit(run_in_new_loop) - return future.result() + with ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit(run_in_new_loop) + return future.result() except RuntimeError: return run_in_new_loop() @@ -314,33 +317,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) # Log response (downloadText) - from litellm.types.utils import ( - Choices, - ModelResponse, - TextChoices, - TextCompletionResponse, - ) - - parts: List[str] = [] - if isinstance(result, TextCompletionResponse) and result.choices: - for choice in result.choices: - if not isinstance(choice, TextChoices): - continue - raw = choice.get("text") - if isinstance(raw, str) and raw.strip(): - parts.append(raw) - elif isinstance(result, ModelResponse) and result.choices: - for choice in result.choices: - if not isinstance(choice, Choices): - continue - msg = choice.message - if msg is None: - continue - raw = msg.get("content") if isinstance(msg, dict) else getattr( - msg, "content", None - ) - if isinstance(raw, str) and raw.strip(): - parts.append(raw) + parts = self._completion_response_text_parts(result) if parts: combined = "\n\n---\n\n".join(parts) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py index 4841682cd2e..0080ca28eb1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py @@ -237,6 +237,29 @@ class TestPreCallHook: assert mock_check.call_args.kwargs["activity"] == "uploadText" assert mock_check.call_args.kwargs["block_on_violation"] is True + @pytest.mark.asyncio + async def test_pre_call_success_returns_request_data(self): + """After a successful DLP pass, the hook must return the same data dict (not None).""" + guardrail = _make_guardrail() + payload = { + "messages": [{"role": "user", "content": "Hello, how are you?"}], + "litellm_call_id": "call-abc", + } + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + out = await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data=payload, + call_type="completion", + ) + + assert out is payload + @pytest.mark.asyncio async def test_pre_call_block(self): guardrail = _make_guardrail()