fix(purview): fix LRU cache refresh position and add Responses API scanning

Two fixes to the Microsoft Purview DLP guardrail:

1. LRU cache bug (base.py): When a stale scope cache entry was re-fetched,
   the assignment  updated the value but
   Python's OrderedDict.__setitem__ preserves the original insertion order for
   existing keys. This left the refreshed entry near the front of the dict,
   making it the first candidate for LRU eviction via popitem(last=False).
   Fix: call move_to_end(user_id) after every write to an existing key.

2. Responses API coverage gap (purview_dlp.py): Requests to /v1/responses use
   an 'input' field instead of 'messages' or 'prompt', so the pre-call hook
   returned without scanning the content. Similarly, post-call hook did not
   handle ResponsesAPIResponse.output. Fix: add _responses_api_input_to_str()
   helper and handle 'responses'/'aresponses' call types in async_pre_call_hook,
   async_post_call_success_hook (via _completion_response_text_parts), and
   async_logging_hook.

Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-05-14 11:49:54 +00:00
parent f20cf90db6
commit 9b69e66dc1
No known key found for this signature in database
3 changed files with 286 additions and 6 deletions

View file

@ -169,6 +169,10 @@ class PurviewGuardrailBase:
etag = response_headers.get("etag", response_headers.get("ETag", ""))
self._scope_cache[user_id] = (etag, response_json, now)
# Move refreshed entry to the end so it is treated as most-recently-used.
# OrderedDict.__setitem__ preserves existing insertion order for known
# keys, so an explicit move_to_end() call is required.
self._scope_cache.move_to_end(user_id)
# Evict least-recently-used entry when cache exceeds max size.
while len(self._scope_cache) > self._scope_cache_maxsize:
self._scope_cache.popitem(last=False)
@ -287,11 +291,14 @@ class PurviewGuardrailBase:
md = litellm_params.get("metadata")
return md if isinstance(md, dict) else {}
def _resolve_user_id_from_logging_kwargs(self, kwargs: Dict[str, Any]) -> Optional[str]:
def _resolve_user_id_from_logging_kwargs(
self, kwargs: Dict[str, Any]
) -> Optional[str]:
"""Same trust order as ``_resolve_user_id`` for logging-only hooks (no ``UserAPIKeyAuth``)."""
md = self._logging_kwargs_metadata(kwargs)
shim = SimpleNamespace(
user_id=md.get("user_api_key_user_id") or kwargs.get("user_api_key_user_id"),
user_id=md.get("user_api_key_user_id")
or kwargs.get("user_api_key_user_id"),
end_user_id=md.get("user_api_key_end_user_id"),
)
return self._resolve_user_id({"metadata": md}, shim)
@ -344,7 +351,9 @@ class PurviewGuardrailBase:
return joined.strip() or None
return None
def get_prompt_text_for_dlp(self, messages: List["AllMessageValues"]) -> Optional[str]:
def get_prompt_text_for_dlp(
self, messages: List["AllMessageValues"]
) -> 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.

View file

@ -25,6 +25,7 @@ from litellm.types.utils import (
Choices,
GuardrailStatus,
ModelResponse,
ResponsesAPIResponse,
TextChoices,
TextCompletionResponse,
)
@ -162,7 +163,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
@staticmethod
def _completion_response_text_parts(result: Any) -> List[str]:
"""Collect non-empty assistant text segments from chat or text completions."""
"""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:
@ -171,6 +172,10 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
raw = choice.get("text")
if isinstance(raw, str) and raw.strip():
parts.append(raw)
elif isinstance(result, ResponsesAPIResponse):
text = result.output_text
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):
@ -178,13 +183,44 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
msg = choice.message
if msg is None:
continue
raw = msg.get("content") if isinstance(msg, dict) else getattr(
msg, "content", None
raw = (
msg.get("content")
if isinstance(msg, dict)
else getattr(msg, "content", None)
)
if isinstance(raw, str) and raw.strip():
parts.append(raw)
return parts
def _responses_api_input_to_str(self, data: Dict[str, Any]) -> Optional[str]:
"""Extract DLP-scannable text from a Responses API request ``input`` field.
``input`` may be a plain string or a list of input items (messages). In
the latter case the items are converted to chat messages via the standard
LiteLLM transformation and then concatenated by ``get_prompt_text_for_dlp``.
"""
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
input_data = data.get("input")
if input_data is None:
return None
if isinstance(input_data, str):
return input_data.strip() or None
try:
messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=input_data,
responses_api_request=data,
)
return self.get_prompt_text_for_dlp(cast(List[Any], messages))
except Exception:
verbose_proxy_logger.debug(
"Purview DLP: failed to transform responses API input; skipping scan",
exc_info=True,
)
return None
# ------------------------------------------------------------------
# Pre-call hook — DLP on prompts
# ------------------------------------------------------------------
@ -211,6 +247,8 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
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"))
elif call_type in ("responses", "aresponses"):
prompt_text = self._responses_api_input_to_str(data)
if not prompt_text:
return data
@ -306,6 +344,8 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
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"))
elif call_type in ("responses", "aresponses"):
prompt_text = self._responses_api_input_to_str(kwargs)
if prompt_text:
await self._check_content(

View file

@ -534,6 +534,195 @@ class TestTextCompletionHooks:
assert "beta" in combined
# ---------------------------------------------------------------
# Responses API hooks
# ---------------------------------------------------------------
class TestResponsesAPIHooks:
@pytest.mark.asyncio
async def test_pre_call_responses_api_string_input(self):
"""Pre-call hook must scan plain-string ``input`` on responses call type."""
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={"input": "SSN: 123-45-6789"},
call_type="responses",
)
mock_check.assert_called_once()
assert mock_check.call_args.kwargs["activity"] == "uploadText"
assert "SSN: 123-45-6789" in mock_check.call_args.kwargs["text"]
@pytest.mark.asyncio
async def test_pre_call_aresponses_string_input(self):
"""Pre-call hook must scan ``input`` on ``aresponses`` call type too."""
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={"input": "sensitive content"},
call_type="aresponses",
)
mock_check.assert_called_once()
assert "sensitive content" in mock_check.call_args.kwargs["text"]
@pytest.mark.asyncio
async def test_pre_call_responses_api_list_input(self):
"""Pre-call hook must extract text from structured list ``input``."""
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={
"input": [
{"role": "user", "content": "Secret phrase: alpha bravo"}
]
},
call_type="responses",
)
mock_check.assert_called_once()
assert "Secret phrase: alpha bravo" in mock_check.call_args.kwargs["text"]
@pytest.mark.asyncio
async def test_pre_call_responses_api_no_input_skips(self):
"""Pre-call hook must not call _check_content when ``input`` is absent."""
guardrail = _make_guardrail()
with patch.object(
guardrail, "_check_content", new_callable=AsyncMock
) as mock_check:
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"),
cache=None,
data={},
call_type="responses",
)
mock_check.assert_not_called()
@pytest.mark.asyncio
async def test_post_call_responses_api_output_text(self):
"""Post-call hook must scan text from ``ResponsesAPIResponse.output``."""
from litellm.types.llms.openai import ResponsesAPIResponse
guardrail = _make_guardrail()
response = ResponsesAPIResponse(
id="resp-1",
created_at=0,
output=[
{
"type": "message",
"id": "msg-1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "card 4111-1111-1111-1111"}],
}
],
)
with patch.object(
guardrail, "_check_content", new_callable=AsyncMock
) as mock_check:
mock_check.return_value = {"policyActions": []}
result = 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()
assert mock_check.call_args.kwargs["activity"] == "downloadText"
assert "card 4111-1111-1111-1111" in mock_check.call_args.kwargs["text"]
assert result is response
@pytest.mark.asyncio
async def test_post_call_responses_api_empty_output_skips(self):
"""Post-call hook must not call _check_content when output has no text."""
from litellm.types.llms.openai import ResponsesAPIResponse
guardrail = _make_guardrail()
response = ResponsesAPIResponse(
id="resp-2",
created_at=0,
output=[],
)
with patch.object(
guardrail, "_check_content", new_callable=AsyncMock
) as mock_check:
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_not_called()
@pytest.mark.asyncio
async def test_logging_hook_responses_api_input_and_output(self):
"""Logging hook must scan both ``input`` and ``ResponsesAPIResponse.output``."""
from litellm.types.llms.openai import ResponsesAPIResponse
guardrail = _make_guardrail(logging_only=True)
result_response = ResponsesAPIResponse(
id="resp-3",
created_at=0,
output=[
{
"type": "message",
"id": "msg-2",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "response body"}],
}
],
)
with patch.object(
guardrail, "_check_content", new_callable=AsyncMock
) as mock_check:
mock_check.return_value = {"policyActions": []}
await guardrail.async_logging_hook(
kwargs={
"input": "prompt body",
"litellm_params": {"metadata": {"user_id": "user-123"}},
},
result=result_response,
call_type="responses",
)
assert mock_check.call_count == 2
activities = {c.kwargs["activity"] for c in mock_check.call_args_list}
assert activities == {"uploadText", "downloadText"}
texts = {c.kwargs["text"] for c in mock_check.call_args_list}
assert any("prompt body" in t for t in texts)
assert any("response body" in t for t in texts)
# ---------------------------------------------------------------
# Logging hook user resolution
# ---------------------------------------------------------------
@ -786,6 +975,48 @@ class TestScopeCaching:
assert "user-a" in guardrail._scope_cache
assert "user-b" not in guardrail._scope_cache
@pytest.mark.asyncio
async def test_scope_cache_refresh_moves_to_end_of_lru(self):
"""Refreshing a stale entry must move it to the MRU end of the OrderedDict.
Before the fix, OrderedDict.__setitem__ preserved the original insertion
position for existing keys, causing the just-refreshed entry to be the
next candidate for LRU eviction.
"""
guardrail = _make_guardrail()
guardrail._scope_cache_maxsize = 2
scope_payload = (
{"value": []},
{"ETag": "scope-etag"},
)
with patch.object(
guardrail, "_graph_post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = scope_payload
# Populate cache: user-a (older), user-b (newer)
await guardrail._compute_protection_scopes("user-a")
await guardrail._compute_protection_scopes("user-b")
assert mock_post.call_count == 2
# Expire user-a's entry so it is re-fetched on the next access.
old_etag, old_scope, _ = guardrail._scope_cache["user-a"]
guardrail._scope_cache["user-a"] = (old_etag, old_scope, 0.0)
# Re-fetch user-a — should move it to the MRU end.
await guardrail._compute_protection_scopes("user-a")
assert mock_post.call_count == 3
# Adding a third user must evict user-b (the true LRU), not user-a.
await guardrail._compute_protection_scopes("user-c")
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-c" in guardrail._scope_cache
@pytest.mark.asyncio
async def test_scope_invalidated_on_modified(self):
guardrail = _make_guardrail()