From 4b759c00192d2147ab2e88ce70d8ecf81201e528 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 14 May 2026 16:51:59 +0530 Subject: [PATCH] fix(guardrails/purview): scan /v1/completions prompt and TextChoices Normalize text-completion prompts (string or list of strings); skip token-id-only prompts. Run post-call DLP on TextCompletionResponse choices. Extend logging_only hook for text_completion. Add tests and completion_prompt_to_str helper. Co-authored-by: Cursor --- .../guardrail_hooks/microsoft_purview/base.py | 27 ++++ .../microsoft_purview/purview_dlp.py | 123 +++++++++++------- .../guardrail_hooks/test_microsoft_purview.py | 74 +++++++++++ 3 files changed, 180 insertions(+), 44 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index ba74e72f578..9986982d554 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -317,6 +317,33 @@ class PurviewGuardrailBase: # Prompt text for DLP # ------------------------------------------------------------------ + @staticmethod + def completion_prompt_to_str(prompt: Any) -> Optional[str]: + """Normalize OpenAI ``/v1/completions`` ``prompt`` for text DLP. + + Supports string prompts and list-of-string prompts. List-of-token-id prompts + are skipped (no plaintext for Purview to evaluate). + """ + if prompt is None: + return None + if isinstance(prompt, str): + stripped = prompt.strip() + return stripped or None + if isinstance(prompt, list) and prompt: + if all(isinstance(x, str) for x in prompt): + joined = "\n".join(s.strip() for s in prompt if isinstance(s, str)) + return joined.strip() or None + if all(isinstance(x, int) for x in prompt): + verbose_proxy_logger.debug( + "Purview DLP: completions prompt is token ids only; skipping text scan" + ) + return None + str_parts = [x for x in prompt if isinstance(x, str)] + if str_parts: + joined = "\n".join(s.strip() for s in str_parts) + return joined.strip() or None + return None + def get_prompt_text_for_dlp(self, messages: List["AllMessageValues"]) -> Optional[str]: """Concatenate text from every chat message (all roles) for pre-call DLP. 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 6a7cc7548c2..97f360baeaf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -11,7 +11,7 @@ import asyncio import uuid from concurrent.futures import ThreadPoolExecutor from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Type, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Type, Union, cast from fastapi import HTTPException @@ -176,19 +176,23 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) return data + prompt_text: Optional[str] = None messages: Optional[List] = data.get("messages") - if not messages: + if messages: + prompt_text = self.get_prompt_text_for_dlp(cast(List[Any], messages)) + elif call_type in ("text_completion", "atext_completion"): + prompt_text = self.completion_prompt_to_str(data.get("prompt")) + + if not prompt_text: return data - prompt_text = self.get_prompt_text_for_dlp(messages) - if prompt_text: - await self._check_content( - user_id=user_id, - text=prompt_text, - activity="uploadText", - request_data=data, - block_on_violation=True, - ) + await self._check_content( + user_id=user_id, + text=prompt_text, + activity="uploadText", + request_data=data, + block_on_violation=True, + ) return None # ------------------------------------------------------------------ @@ -202,7 +206,12 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): response: Union[Any, "ModelResponse", "EmbeddingResponse", "ImageResponse"], ) -> Any: """Check LLM response against Purview DLP policies.""" - from litellm.types.utils import Choices, ModelResponse + from litellm.types.utils import ( + Choices, + ModelResponse, + TextChoices, + TextCompletionResponse, + ) user_id = self._resolve_user_id(data, user_api_key_dict) if not user_id: @@ -211,8 +220,16 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) return response - if isinstance(response, ModelResponse) and response.choices: - parts: List[str] = [] + 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 @@ -224,15 +241,16 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) if isinstance(raw, str) and raw.strip(): parts.append(raw) - if parts: - combined = "\n\n---\n\n".join(parts) - await self._check_content( - user_id=user_id, - text=combined, - activity="downloadText", - request_data=data, - block_on_violation=True, - ) + + if parts: + combined = "\n\n---\n\n".join(parts) + await self._check_content( + user_id=user_id, + text=combined, + activity="downloadText", + request_data=data, + block_on_violation=True, + ) return response # ------------------------------------------------------------------ @@ -279,23 +297,39 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): return kwargs, result # Log prompt (uploadText) + prompt_text: Optional[str] = None messages = kwargs.get("messages") if messages: - prompt_text = self.get_prompt_text_for_dlp(messages) - if prompt_text: - await self._check_content( - user_id=user_id, - text=prompt_text, - activity="uploadText", - request_data=kwargs, - block_on_violation=False, - ) + prompt_text = self.get_prompt_text_for_dlp(cast(List[Any], messages)) + elif call_type in ("text_completion", "atext_completion"): + prompt_text = self.completion_prompt_to_str(kwargs.get("prompt")) + + if prompt_text: + await self._check_content( + user_id=user_id, + text=prompt_text, + activity="uploadText", + request_data=kwargs, + block_on_violation=False, + ) # Log response (downloadText) - from litellm.types.utils import Choices, ModelResponse + from litellm.types.utils import ( + Choices, + ModelResponse, + TextChoices, + TextCompletionResponse, + ) - if isinstance(result, ModelResponse) and result.choices: - parts: List[str] = [] + 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 @@ -307,15 +341,16 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) if isinstance(raw, str) and raw.strip(): parts.append(raw) - if parts: - combined = "\n\n---\n\n".join(parts) - await self._check_content( - user_id=user_id, - text=combined, - activity="downloadText", - request_data=kwargs, - block_on_violation=False, - ) + + if parts: + combined = "\n\n---\n\n".join(parts) + await self._check_content( + user_id=user_id, + text=combined, + activity="downloadText", + request_data=kwargs, + block_on_violation=False, + ) except Exception as e: verbose_proxy_logger.error("Purview audit logging error: %s", e) 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 138c8608b37..4841682cd2e 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 @@ -127,6 +127,29 @@ class TestShouldBlock: assert PurviewGuardrailBase._should_block(response) is True +# --------------------------------------------------------------- +# completion prompt normalization (text completions API) +# --------------------------------------------------------------- + + +class TestCompletionPromptToStr: + def test_string_prompt(self): + assert PurviewGuardrailBase.completion_prompt_to_str(" hi ") == "hi" + + def test_list_of_strings(self): + assert ( + PurviewGuardrailBase.completion_prompt_to_str(["a", "b"]) + == "a\nb" + ) + + def test_token_ids_returns_none(self): + assert PurviewGuardrailBase.completion_prompt_to_str([1, 2, 3]) is None + + def test_empty(self): + assert PurviewGuardrailBase.completion_prompt_to_str("") is None + assert PurviewGuardrailBase.completion_prompt_to_str([]) is None + + # --------------------------------------------------------------- # User ID resolution # --------------------------------------------------------------- @@ -437,6 +460,57 @@ class TestPostCallHook: assert "Second completion body" in combined +class TestTextCompletionHooks: + @pytest.mark.asyncio + async def test_pre_call_text_completion_uses_prompt(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={"prompt": "Completions API prompt body"}, + call_type="text_completion", + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["text"] == "Completions API prompt body" + assert mock_check.call_args.kwargs["activity"] == "uploadText" + + @pytest.mark.asyncio + async def test_post_call_text_completion_all_choices(self): + from litellm.types.utils import TextChoices, TextCompletionResponse + + guardrail = _make_guardrail() + response = TextCompletionResponse( + model="gpt-3.5-turbo-instruct", + choices=[ + TextChoices(text="alpha", index=0), + TextChoices(text="beta", index=1), + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_post_call_success_hook( + data={"metadata": {"user_id": "user-123"}}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + + mock_check.assert_called_once() + combined = mock_check.call_args.kwargs["text"] + assert "alpha" in combined + assert "beta" in combined + + # --------------------------------------------------------------- # Logging hook user resolution # ---------------------------------------------------------------