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