From 9b69e66dc1a9cf3c58bfe2ae540d29b8ba25bf92 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 14 May 2026 11:49:54 +0000 Subject: [PATCH] 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 --- .../guardrail_hooks/microsoft_purview/base.py | 15 +- .../microsoft_purview/purview_dlp.py | 46 +++- .../guardrail_hooks/test_microsoft_purview.py | 231 ++++++++++++++++++ 3 files changed, 286 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index 9986982d554..91579a2bf04 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -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. diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py index 893fa39eac9..96106032a62 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py index 0080ca28eb1..dc0fe05e2cf 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py @@ -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()