mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
d58b2f1b44
commit
4b759c0019
3 changed files with 180 additions and 44 deletions
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue