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 <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-14 16:51:59 +05:30
parent d58b2f1b44
commit 4b759c0019
No known key found for this signature in database
3 changed files with 180 additions and 44 deletions

View file

@ -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.

View file

@ -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)

View file

@ -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
# ---------------------------------------------------------------