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