From 514f99fbf31a6f4deec862d6b027112d71f83687 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 14 May 2026 13:30:35 +0000 Subject: [PATCH] fix(purview): suppress API errors in logging-only mode and scan tool-call arguments Three issues fixed: 1. _check_content except block re-raised unconditionally even when block_on_violation=False. The docstring promised 'log only - do not raise' but network/API errors always propagated. Fixed by checking block_on_violation before re-raising; when False, log a warning and continue. 2. async_logging_hook used a single try/except wrapping both the prompt and response audit calls. When the first _check_content (uploadText) raised due to an API error the second call (downloadText) was silently skipped. Fixed by giving each audit call its own try/except so both always run independently. 3. convert_content_list_to_str() only reads message.content, so tool_calls[].function.arguments and function_call.arguments were invisible to the Purview pre-call and post-call scans. An authenticated caller could embed sensitive text in tool-call arguments and bypass DLP. Fixed by: - Adding PurviewGuardrailBase._extract_tool_call_args_from_message() which handles both dict and object-style messages, covering both tool_calls[] arrays and the legacy function_call field. - Updating get_prompt_text_for_dlp() to include those arguments alongside message content (request/prompt path). - Changing _completion_response_text_parts() from @staticmethod to an instance method and adding tool-call argument extraction for ModelResponse choices (response path). Co-authored-by: Sameer Kankute --- .../guardrail_hooks/microsoft_purview/base.py | 68 ++- .../microsoft_purview/purview_dlp.py | 44 +- .../guardrail_hooks/test_microsoft_purview.py | 460 ++++++++++++++++++ 3 files changed, 555 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index a3fe2c9ed69..ae58978d419 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -351,6 +351,57 @@ class PurviewGuardrailBase: return joined.strip() or None return None + @staticmethod + def _extract_tool_call_args_from_message(message: Any) -> List[str]: + """Return plaintext arguments strings from tool_calls and function_call fields. + + Covers both the request path (assistant messages in chat histories that + carry tool_calls / function_call) and the response path (model-generated + tool calls returned in a ModelResponse). Both dict-style and object-style + representations are handled. + """ + args: List[str] = [] + + # tool_calls: [{"function": {"arguments": "..."}}] + tool_calls = ( + message.get("tool_calls") + if isinstance(message, dict) + else getattr(message, "tool_calls", None) + ) + if tool_calls: + for tc in tool_calls: + fn = ( + tc.get("function") + if isinstance(tc, dict) + else getattr(tc, "function", None) + ) + if fn is None: + continue + arguments = ( + fn.get("arguments") + if isinstance(fn, dict) + else getattr(fn, "arguments", None) + ) + if isinstance(arguments, str) and arguments.strip(): + args.append(arguments) + + # Legacy function_call: {"arguments": "..."} + function_call = ( + message.get("function_call") + if isinstance(message, dict) + else getattr(message, "function_call", None) + ) + if function_call is not None: + arguments = ( + function_call.get("arguments") + if isinstance(function_call, dict) + else getattr(function_call, "arguments", None) + ) + if isinstance(arguments, str) and arguments.strip(): + args.append(arguments) + + return args + def get_prompt_text_for_dlp( self, messages: List["AllMessageValues"] ) -> Optional[str]: @@ -361,9 +412,22 @@ class PurviewGuardrailBase: 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. + + Tool-call arguments (``tool_calls[].function.arguments`` and + ``function_call.arguments``) are included alongside message content so + that sensitive data hidden in function arguments is not bypassed. """ if not messages: return None - parts = [convert_content_list_to_str(message=msg).strip() for msg in messages] - text = "\n\n".join(p for p in parts if p) + parts: List[str] = [] + for msg in messages: + segments: List[str] = [] + content = convert_content_list_to_str(message=msg).strip() + if content: + segments.append(content) + segments.extend(self._extract_tool_call_args_from_message(msg)) + combined = "\n".join(segments) + if combined.strip(): + parts.append(combined.strip()) + text = "\n\n".join(parts) return text or None 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 f5dd5b05f88..e14eab6de24 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -135,9 +135,14 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): if self._should_block(response): status = "guardrail_intervened" - except Exception: + except Exception as exc: status = "guardrail_failed_to_respond" - raise + if block_on_violation: + raise + verbose_proxy_logger.warning( + "Purview DLP: API/network error in logging-only mode (not re-raised): %s", + exc, + ) finally: end_time = datetime.now() self.add_standard_logging_guardrail_information_to_request_data( @@ -161,9 +166,13 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): return response - @staticmethod - def _completion_response_text_parts(result: Any) -> List[str]: - """Collect non-empty assistant text segments from chat, text completions, or responses API.""" + def _completion_response_text_parts(self, result: Any) -> List[str]: + """Collect non-empty text segments from chat, text completions, or responses API. + + Includes assistant message content *and* model-generated tool-call + arguments so that sensitive data returned inside function calls is not + missed by the DLP scan. + """ parts: List[str] = [] if isinstance(result, TextCompletionResponse) and result.choices: for text_choice in result.choices: @@ -190,6 +199,8 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) if isinstance(raw, str) and raw.strip(): parts.append(raw) + # Include tool-call arguments returned by the model + parts.extend(self._extract_tool_call_args_from_message(msg)) return parts def _responses_api_input_to_str(self, data: Dict[str, Any]) -> Optional[str]: @@ -347,15 +358,16 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): """Send both prompt and response to Purview for audit logging. Errors are logged but never raised — this mode is non-blocking. + Each audit call (prompt and response) is wrapped in its own try/except + so a failure on the first does not prevent the second from running. """ + user_id = self._resolve_user_id_from_logging_kwargs(kwargs) + if not user_id: + verbose_proxy_logger.debug("Purview audit: no user_id, skipping") + return kwargs, result + + # Log prompt (uploadText) try: - user_id = self._resolve_user_id_from_logging_kwargs(kwargs) - - if not user_id: - verbose_proxy_logger.debug("Purview audit: no user_id, skipping") - return kwargs, result - - # Log prompt (uploadText) prompt_text: Optional[str] = None messages = kwargs.get("messages") if messages: @@ -373,10 +385,12 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): request_data=kwargs, block_on_violation=False, ) + except Exception as e: + verbose_proxy_logger.error("Purview audit logging error (prompt): %s", e) - # Log response (downloadText) + # Log response (downloadText) — runs regardless of prompt audit outcome + try: parts = self._completion_response_text_parts(result) - if parts: combined = "\n\n---\n\n".join(parts) await self._check_content( @@ -387,6 +401,6 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): block_on_violation=False, ) except Exception as e: - verbose_proxy_logger.error("Purview audit logging error: %s", e) + verbose_proxy_logger.error("Purview audit logging error (response): %s", e) return kwargs, result 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 7c29e96f3ef..22e7d0c0c5c 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 @@ -1210,6 +1210,466 @@ class TestInitializerValidation: initialize_guardrail(litellm_params, {"guardrail_name": "test"}) +# --------------------------------------------------------------- +# _check_content — API error handling with block_on_violation=False +# --------------------------------------------------------------- + + +class TestCheckContentApiErrorHandling: + @pytest.mark.asyncio + async def test_api_error_reraises_when_block_on_violation_true(self): + """API/network errors must propagate when block_on_violation=True.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + side_effect=RuntimeError("network failure"), + ): + with pytest.raises(RuntimeError, match="network failure"): + await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + @pytest.mark.asyncio + async def test_api_error_not_reraised_when_block_on_violation_false(self): + """API/network errors must be swallowed (logged only) when block_on_violation=False.""" + guardrail = _make_guardrail() + + with patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + side_effect=RuntimeError("network failure"), + ): + # Must NOT raise — should return empty dict + result = await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=False, + ) + + assert isinstance(result, dict) + + @pytest.mark.asyncio + async def test_process_content_error_not_reraised_when_block_on_violation_false( + self, + ): + """Errors from _process_content itself must also be suppressed in logging-only mode.""" + guardrail = _make_guardrail() + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + side_effect=ConnectionError("timeout"), + ), + ): + result = await guardrail._check_content( + user_id="user-1", + text="some content", + activity="uploadText", + request_data={}, + block_on_violation=False, + ) + + assert isinstance(result, dict) + + +# --------------------------------------------------------------- +# async_logging_hook — independent prompt/response audit calls +# --------------------------------------------------------------- + + +class TestAsyncLoggingHookIndependence: + @pytest.mark.asyncio + async def test_response_audit_runs_even_if_prompt_audit_fails(self): + """A failure in the prompt audit must not prevent the response audit from running.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail(logging_only=True) + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message(content="response text", role="assistant"), + ) + ], + ) + + call_activities: list = [] + + async def fake_check_content(**kwargs): + activity = kwargs.get("activity") + if activity == "uploadText": + raise RuntimeError("simulated prompt API failure") + call_activities.append(activity) + return {"policyActions": []} + + with patch.object(guardrail, "_check_content", side_effect=fake_check_content): + await guardrail.async_logging_hook( + kwargs={ + "messages": [{"role": "user", "content": "prompt"}], + "litellm_params": {"metadata": {"user_id": "user-123"}}, + }, + result=response, + call_type="completion", + ) + + # The response audit must still have been attempted + assert "downloadText" in call_activities + + @pytest.mark.asyncio + async def test_prompt_audit_runs_even_if_response_audit_fails(self): + """A failure in the response audit must not affect the prompt audit result.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail(logging_only=True) + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message(content="response text", role="assistant"), + ) + ], + ) + + call_activities: list = [] + + async def fake_check_content(**kwargs): + activity = kwargs.get("activity") + if activity == "downloadText": + raise RuntimeError("simulated response API failure") + call_activities.append(activity) + return {"policyActions": []} + + with patch.object(guardrail, "_check_content", side_effect=fake_check_content): + await guardrail.async_logging_hook( + kwargs={ + "messages": [{"role": "user", "content": "prompt"}], + "litellm_params": {"metadata": {"user_id": "user-123"}}, + }, + result=response, + call_type="completion", + ) + + assert "uploadText" in call_activities + + @pytest.mark.asyncio + async def test_logging_hook_returns_original_when_both_audits_fail(self): + """async_logging_hook must always return (kwargs, result) even if both audits fail.""" + guardrail = _make_guardrail(logging_only=True) + + with patch.object( + guardrail, + "_check_content", + new_callable=AsyncMock, + side_effect=RuntimeError("total failure"), + ): + kwargs = { + "messages": [{"role": "user", "content": "prompt"}], + "litellm_params": {"metadata": {"user_id": "user-123"}}, + } + result_obj = {"some": "result"} + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, + result=result_obj, + call_type="completion", + ) + + assert out_kwargs is kwargs + assert out_result is result_obj + + +# --------------------------------------------------------------- +# Tool-call argument extraction +# --------------------------------------------------------------- + + +class TestExtractToolCallArgs: + def test_dict_message_with_tool_calls(self): + msg = { + "role": "assistant", + "content": None, + "tool_calls": [ + {"function": {"arguments": '{"ssn": "123-45-6789"}'}}, + {"function": {"arguments": '{"card": "4111-1111-1111-1111"}'}}, + ], + } + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert '{"ssn": "123-45-6789"}' in args + assert '{"card": "4111-1111-1111-1111"}' in args + + def test_dict_message_with_function_call(self): + msg = { + "role": "assistant", + "content": None, + "function_call": {"name": "lookup", "arguments": '{"query": "secret"}'}, + } + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert '{"query": "secret"}' in args + + def test_object_message_with_tool_calls(self): + from litellm.types.utils import Message + + msg = Message( + role="assistant", + content=None, + tool_calls=[ + { + "id": "tc1", + "type": "function", + "function": {"name": "fn", "arguments": '{"x": 1}'}, + }, + ], + ) + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert '{"x": 1}' in args + + def test_message_with_no_tool_calls(self): + msg = {"role": "user", "content": "hello"} + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert args == [] + + def test_empty_arguments_skipped(self): + msg = { + "role": "assistant", + "content": None, + "tool_calls": [{"function": {"arguments": " "}}], + } + args = MicrosoftPurviewDLPGuardrail._extract_tool_call_args_from_message(msg) + assert args == [] + + +# --------------------------------------------------------------- +# Tool-call arguments included in DLP text extraction (prompt) +# --------------------------------------------------------------- + + +class TestGetPromptTextToolCalls: + def test_tool_call_args_included_in_prompt_scan(self): + """Sensitive data in tool_calls[].function.arguments must appear in DLP text.""" + guardrail = _make_guardrail() + messages = [ + {"role": "user", "content": "benign query"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "tc1", + "type": "function", + "function": { + "name": "lookup", + "arguments": '{"ssn": "123-45-6789"}', + }, + } + ], + }, + ] + text = guardrail.get_prompt_text_for_dlp(messages) + assert text is not None + assert "benign query" in text + assert '{"ssn": "123-45-6789"}' in text + + def test_function_call_args_included_in_prompt_scan(self): + """Legacy function_call.arguments must also appear in DLP text.""" + guardrail = _make_guardrail() + messages = [ + { + "role": "assistant", + "content": "Calling function", + "function_call": { + "name": "search", + "arguments": '{"credit_card": "4111-1111-1111-1111"}', + }, + } + ] + text = guardrail.get_prompt_text_for_dlp(messages) + assert text is not None + assert "Calling function" in text + assert '{"credit_card": "4111-1111-1111-1111"}' in text + + def test_content_only_message_unchanged(self): + """Messages without tool calls must still produce the same output.""" + guardrail = _make_guardrail() + messages = [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "Tell me a joke."}, + ] + text = guardrail.get_prompt_text_for_dlp(messages) + assert text is not None + assert "You are helpful." in text + assert "Tell me a joke." in text + + @pytest.mark.asyncio + async def test_pre_call_hook_scans_tool_call_args(self): + """async_pre_call_hook must include tool_call arguments in the text sent to Purview.""" + 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=__import__( + "litellm.proxy._types", fromlist=["UserAPIKeyAuth"] + ).UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={ + "messages": [ + {"role": "user", "content": "benign"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "tc1", + "type": "function", + "function": { + "name": "do_thing", + "arguments": '{"password": "hunter2"}', + }, + } + ], + }, + ] + }, + call_type="completion", + ) + + mock_check.assert_called_once() + sent_text = mock_check.call_args.kwargs["text"] + assert '{"password": "hunter2"}' in sent_text + + +# --------------------------------------------------------------- +# Tool-call arguments included in DLP text extraction (response) +# --------------------------------------------------------------- + + +class TestCompletionResponseTextPartsToolCalls: + def test_response_tool_call_args_included(self): + """Model-generated tool_call arguments must appear in the DLP scan text.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + role="assistant", + content=None, + tool_calls=[ + { + "id": "tc1", + "type": "function", + "function": { + "name": "exfil", + "arguments": '{"data": "secret-value"}', + }, + } + ], + ), + ) + ], + ) + parts = guardrail._completion_response_text_parts(response) + assert any("secret-value" in p for p in parts) + + def test_response_with_content_and_tool_calls(self): + """Both message content and tool_call arguments must be included.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + role="assistant", + content="Here is the result", + tool_calls=[ + { + "id": "tc2", + "type": "function", + "function": { + "name": "fn", + "arguments": '{"ssn": "123-45-6789"}', + }, + } + ], + ), + ) + ], + ) + parts = guardrail._completion_response_text_parts(response) + combined = " ".join(parts) + assert "Here is the result" in combined + assert '{"ssn": "123-45-6789"}' in combined + + @pytest.mark.asyncio + async def test_post_call_hook_scans_response_tool_call_args(self): + """async_post_call_success_hook must send tool_call arguments to Purview.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + role="assistant", + content=None, + tool_calls=[ + { + "id": "tc3", + "type": "function", + "function": { + "name": "retrieve", + "arguments": '{"credit_card": "4111-1111-1111-1111"}', + }, + } + ], + ), + ) + ], + ) + + 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=__import__( + "litellm.proxy._types", fromlist=["UserAPIKeyAuth"] + ).UserAPIKeyAuth(api_key="test"), + response=response, + ) + + mock_check.assert_called_once() + sent_text = mock_check.call_args.kwargs["text"] + assert '{"credit_card": "4111-1111-1111-1111"}' in sent_text + + # --------------------------------------------------------------- # Auto-discovery registration # ---------------------------------------------------------------