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 <Sameerlite@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-05-14 13:30:35 +00:00
parent f2113fd565
commit 514f99fbf3
No known key found for this signature in database
3 changed files with 555 additions and 17 deletions

View file

@ -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

View file

@ -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

View file

@ -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
# ---------------------------------------------------------------