fix(purview): message separator, non-blocking logging_hook, TextChoices type error

Three bugs fixed in the Microsoft Purview DLP guardrail:

1. get_prompt_text_for_dlp message separator (base.py)
   - Previously called get_str_from_messages() which concatenated all message
     texts with NO separator, so 'end of msg1' + 'start of msg2' became
     'end of msg1start of msg2'.
   - Now joins per-message text with '\n\n' via convert_content_list_to_str(),
     preserving DLP pattern detection accuracy across message boundaries.

2. logging_hook blocking the event loop thread (purview_dlp.py)
   - Previously called future.result() which blocked the calling thread
     (often the event loop thread) for the entire round-trip of two sequential
     Microsoft Graph API calls (_compute_protection_scopes + _process_content).
   - Now fires and forgets: when called inside a running loop, schedules the
     coroutine with loop.create_task(); otherwise spawns a daemon thread.
     Returns (kwargs, result) immediately in both cases.
   - Removes unused concurrent.futures.ThreadPoolExecutor import; adds threading.

3. Incompatible assignment type error (purview_dlp.py:180)
   - mypy inferred 'choice' as TextChoices from the first loop body, then
     flagged the assignment in the second loop as incompatible with Choices.
   - Fixed by using distinct loop variable names: text_choice (TextChoices) and
     chat_choice (Choices).

Tests: 7 new tests added covering the separator fix (TestGetPromptTextForDlp)
and the non-blocking logging_hook (TestLoggingHookNonBlocking).

Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-05-14 12:08:03 +00:00
parent 9b69e66dc1
commit f2113fd565
No known key found for this signature in database
3 changed files with 179 additions and 39 deletions

View file

@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
convert_content_list_to_str,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -356,9 +356,14 @@ class PurviewGuardrailBase:
) -> Optional[str]:
"""Concatenate text from every chat message (all roles) for pre-call DLP.
Evaluates the same payload the model receives, not only the trailing user turn.
Evaluates the same payload the model receives, not only the trailing user
turn. Each message is separated by ``\\n\\n`` so that tokens at message
boundaries are not merged (e.g., ``"end of msg1\\n\\nstart of msg2"``
rather than ``"end of msg1start of msg2"``), which preserves DLP pattern
detection accuracy across message boundaries.
"""
if not messages:
return None
text = get_str_from_messages(messages).strip()
parts = [convert_content_list_to_str(message=msg).strip() for msg in messages]
text = "\n\n".join(p for p in parts if p)
return text or None

View file

@ -8,8 +8,8 @@ Supports three modes:
"""
import asyncio
import threading
import uuid
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Type, Union, cast
@ -166,10 +166,10 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
"""Collect non-empty assistant text segments from chat, text completions, or responses API."""
parts: List[str] = []
if isinstance(result, TextCompletionResponse) and result.choices:
for choice in result.choices:
if not isinstance(choice, TextChoices):
for text_choice in result.choices:
if not isinstance(text_choice, TextChoices):
continue
raw = choice.get("text")
raw = text_choice.get("text")
if isinstance(raw, str) and raw.strip():
parts.append(raw)
elif isinstance(result, ResponsesAPIResponse):
@ -177,10 +177,10 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
if text and text.strip():
parts.append(text)
elif isinstance(result, ModelResponse) and result.choices:
for choice in result.choices:
if not isinstance(choice, Choices):
for chat_choice in result.choices:
if not isinstance(chat_choice, Choices):
continue
msg = choice.message
msg = chat_choice.message
if msg is None:
continue
raw = (
@ -300,28 +300,46 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
def logging_hook(
self, kwargs: dict, result: Any, call_type: str
) -> Tuple[dict, Any]:
"""Sync wrapper for async_logging_hook (follows Presidio pattern)."""
"""Fire-and-forget async audit logging; returns original (kwargs, result) immediately.
def run_in_new_loop() -> Tuple[dict, Any]:
new_loop = asyncio.new_event_loop()
Unlike the Presidio pattern (which does local text manipulation),
``async_logging_hook`` makes two sequential network calls to the
Microsoft Graph API. Blocking the calling thread — or worse, the
event loop thread — until those HTTP round-trips complete would
significantly degrade throughput. Since the hook is audit-only and
always returns ``(kwargs, result)`` unchanged, we can schedule the
work without waiting and return immediately.
"""
async def _log_safe() -> None:
try:
asyncio.set_event_loop(new_loop)
return new_loop.run_until_complete(
self.async_logging_hook(
kwargs=kwargs, result=result, call_type=call_type
)
await self.async_logging_hook(
kwargs=kwargs, result=result, call_type=call_type
)
except Exception as exc:
verbose_proxy_logger.error(
"Purview audit background logging error: %s", exc
)
finally:
new_loop.close()
asyncio.set_event_loop(None)
try:
_ = asyncio.get_running_loop()
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(run_in_new_loop)
return future.result()
loop = asyncio.get_running_loop()
loop.create_task(_log_safe())
except RuntimeError:
return run_in_new_loop()
# No running event loop — run in a background daemon thread so
# the caller still isn't blocked.
def _run_in_new_loop() -> None:
new_loop = asyncio.new_event_loop()
try:
asyncio.set_event_loop(new_loop)
new_loop.run_until_complete(_log_safe())
finally:
new_loop.close()
asyncio.set_event_loop(None)
thread = threading.Thread(target=_run_in_new_loop, daemon=True)
thread.start()
return kwargs, result
async def async_logging_hook(
self, kwargs: dict, result: Any, call_type: str

View file

@ -1,5 +1,6 @@
"""Unit tests for the Microsoft Purview DLP guardrail."""
import asyncio
import time
from unittest.mock import AsyncMock, Mock, patch
@ -137,10 +138,7 @@ class TestCompletionPromptToStr:
assert PurviewGuardrailBase.completion_prompt_to_str(" hi ") == "hi"
def test_list_of_strings(self):
assert (
PurviewGuardrailBase.completion_prompt_to_str(["a", "b"])
== "a\nb"
)
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
@ -457,7 +455,8 @@ class TestPostCallHook:
response = ModelResponse(
choices=[
Choices(
index=0, message=Message(content="First completion", role="assistant")
index=0,
message=Message(content="First completion", role="assistant"),
),
Choices(
index=1,
@ -595,9 +594,7 @@ class TestResponsesAPIHooks:
user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"),
cache=None,
data={
"input": [
{"role": "user", "content": "Secret phrase: alpha bravo"}
]
"input": [{"role": "user", "content": "Secret phrase: alpha bravo"}]
},
call_type="responses",
)
@ -637,7 +634,9 @@ class TestResponsesAPIHooks:
"id": "msg-1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "card 4111-1111-1111-1111"}],
"content": [
{"type": "output_text", "text": "card 4111-1111-1111-1111"}
],
}
],
)
@ -746,9 +745,7 @@ class TestLoggingResolveUserId:
def test_logging_falls_back_to_user_id_field(self):
guardrail = _make_guardrail()
kwargs = {
"litellm_params": {"metadata": {"user_id": "only-metadata-user"}}
}
kwargs = {"litellm_params": {"metadata": {"user_id": "only-metadata-user"}}}
assert (
guardrail._resolve_user_id_from_logging_kwargs(kwargs)
== "only-metadata-user"
@ -1014,7 +1011,9 @@ class TestScopeCaching:
assert mock_post.call_count == 4
assert "user-a" in guardrail._scope_cache, "user-a was wrongly evicted"
assert "user-b" not in guardrail._scope_cache, "user-b should have been evicted"
assert (
"user-b" not in guardrail._scope_cache
), "user-b should have been evicted"
assert "user-c" in guardrail._scope_cache
@pytest.mark.asyncio
@ -1042,6 +1041,124 @@ class TestScopeCaching:
assert "user-1" not in guardrail._scope_cache
# ---------------------------------------------------------------
# get_prompt_text_for_dlp — message separator
# ---------------------------------------------------------------
class TestGetPromptTextForDlp:
def test_single_message_no_extra_separator(self):
"""A single message is returned as-is (no leading/trailing separator)."""
guardrail = _make_guardrail()
result = guardrail.get_prompt_text_for_dlp(
[{"role": "user", "content": "Hello"}]
)
assert result == "Hello"
def test_messages_separated_by_double_newline(self):
"""Adjacent messages must NOT be concatenated without a separator.
Before the fix, "end of msg1" + "start of msg2" became
"end of msg1start of msg2", mangling DLP pattern detection.
"""
guardrail = _make_guardrail()
result = guardrail.get_prompt_text_for_dlp(
[
{"role": "system", "content": "end of msg1"},
{"role": "user", "content": "start of msg2"},
]
)
assert result is not None
assert "end of msg1" in result
assert "start of msg2" in result
# Separator must be present between messages
assert "end of msg1start of msg2" not in result
assert "end of msg1\n\nstart of msg2" in result
def test_empty_messages_returns_none(self):
guardrail = _make_guardrail()
assert guardrail.get_prompt_text_for_dlp([]) is None
def test_whitespace_only_messages_skipped(self):
guardrail = _make_guardrail()
result = guardrail.get_prompt_text_for_dlp(
[
{"role": "system", "content": " "},
{"role": "user", "content": "real content"},
]
)
assert result == "real content"
def test_multi_role_conversation_preserves_all_content(self):
guardrail = _make_guardrail()
result = guardrail.get_prompt_text_for_dlp(
[
{"role": "system", "content": "SYSTEM"},
{"role": "user", "content": "USER1"},
{"role": "assistant", "content": "ASSISTANT"},
{"role": "user", "content": "USER2"},
]
)
assert result is not None
for token in ("SYSTEM", "USER1", "ASSISTANT", "USER2"):
assert token in result
# ---------------------------------------------------------------
# logging_hook — non-blocking fire-and-forget
# ---------------------------------------------------------------
class TestLoggingHookNonBlocking:
@pytest.mark.asyncio
async def test_logging_hook_does_not_block_running_loop(self):
"""logging_hook must return immediately without blocking the event loop.
Before the fix, logging_hook called future.result() which blocked the
event loop thread for the full round-trip of the two Graph API calls.
"""
guardrail = _make_guardrail(logging_only=True)
call_count = 0
async def slow_async_hook(**_kwargs):
nonlocal call_count
await asyncio.sleep(0.05)
call_count += 1
return _kwargs.get("kwargs", {}), _kwargs.get("result")
with patch.object(guardrail, "async_logging_hook", side_effect=slow_async_hook):
# Call logging_hook from within a running event loop
result = guardrail.logging_hook(
kwargs={"messages": [{"role": "user", "content": "test"}]},
result=None,
call_type="completion",
)
# Must return (kwargs, result) unchanged without waiting for async work
assert result[0]["messages"][0]["content"] == "test"
assert result[1] is None
def test_logging_hook_returns_original_kwargs_and_result(self):
"""Return value must be the original (kwargs, result) tuple unchanged."""
guardrail = _make_guardrail(logging_only=True)
kwargs = {"messages": [{"role": "user", "content": "hello"}]}
result_obj = {"some": "result"}
with patch.object(
guardrail,
"async_logging_hook",
new_callable=AsyncMock,
return_value=(kwargs, result_obj),
):
out = guardrail.logging_hook(
kwargs=kwargs,
result=result_obj,
call_type="completion",
)
assert out == (kwargs, result_obj)
# ---------------------------------------------------------------
# Initializer validation
# ---------------------------------------------------------------