fix(guardrails/purview): harden user-id resolution and broaden DLP text

Prefer API key and proxy-injected metadata over client metadata for Entra
identity. Scan full message transcript pre-call and all completion choices
post-call. Align logging-only hook with the same user-id rules.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-14 16:27:35 +05:30
parent 2525883796
commit d58b2f1b44
No known key found for this signature in database
3 changed files with 223 additions and 50 deletions

View file

@ -1,11 +1,12 @@
import time
import uuid
from collections import OrderedDict
from types import SimpleNamespace
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_last_user_message,
get_str_from_messages,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -252,21 +253,49 @@ class PurviewGuardrailBase:
) -> Optional[str]:
"""Resolve the Entra user object ID from request data or auth context.
Resolution order:
1. ``metadata[user_id_field]`` (explicit per-request mapping)
2. ``user_api_key_dict.user_id``
3. ``user_api_key_dict.end_user_id``
Trust order (strongest first) so client ``metadata[user_id_field]`` cannot
impersonate another Entra user for Purview ``protectionScopes`` / ``processContent``:
1. ``user_api_key_dict.user_id`` — LiteLLM key / internal user
2. ``user_api_key_dict.end_user_id`` — end-user on the API key
3. ``metadata["user_api_key_user_id"]`` — proxy-injected from the key (when present)
4. ``metadata[user_id_field]`` — caller-supplied; used only when none of the above apply
"""
metadata = data.get("metadata") or data.get("litellm_metadata") or {}
uid = metadata.get(self.user_id_field)
if uid:
return str(uid)
if hasattr(user_api_key_dict, "user_id") and user_api_key_dict.user_id:
return str(user_api_key_dict.user_id)
if hasattr(user_api_key_dict, "end_user_id") and user_api_key_dict.end_user_id:
return str(user_api_key_dict.end_user_id)
uid = metadata.get("user_api_key_user_id")
if uid:
return str(uid)
uid = metadata.get(self.user_id_field)
if uid:
return str(uid)
return None
@staticmethod
def _logging_kwargs_metadata(kwargs: Dict[str, Any]) -> Dict[str, Any]:
"""Metadata dict from ``model_call_details`` / logging kwargs."""
litellm_params = kwargs.get("litellm_params") or {}
if not isinstance(litellm_params, dict):
return {}
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]:
"""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"),
end_user_id=md.get("user_api_key_end_user_id"),
)
return self._resolve_user_id({"metadata": md}, shim)
# ------------------------------------------------------------------
# Policy action evaluation
# ------------------------------------------------------------------
@ -285,9 +314,15 @@ class PurviewGuardrailBase:
return False
# ------------------------------------------------------------------
# User prompt extraction
# Prompt text for DLP
# ------------------------------------------------------------------
def get_user_prompt(self, messages: List["AllMessageValues"]) -> Optional[str]:
"""Get the last consecutive block of user messages as a single string."""
return get_last_user_message(messages)
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.
"""
if not messages:
return None
text = get_str_from_messages(messages).strip()
return text or None

View file

@ -180,11 +180,11 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
if not messages:
return data
user_prompt = self.get_user_prompt(messages)
if user_prompt:
prompt_text = self.get_prompt_text_for_dlp(messages)
if prompt_text:
await self._check_content(
user_id=user_id,
text=user_prompt,
text=prompt_text,
activity="uploadText",
request_data=data,
block_on_violation=True,
@ -211,16 +211,24 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
)
return response
if (
isinstance(response, ModelResponse)
and response.choices
and isinstance(response.choices[0], Choices)
):
content = response.choices[0].message.content or ""
if content:
if isinstance(response, ModelResponse) and response.choices:
parts: List[str] = []
for choice in response.choices:
if not isinstance(choice, Choices):
continue
msg = choice.message
if msg is None:
continue
raw = msg.get("content") if isinstance(msg, dict) else getattr(
msg, "content", None
)
if isinstance(raw, str) and raw.strip():
parts.append(raw)
if parts:
combined = "\n\n---\n\n".join(parts)
await self._check_content(
user_id=user_id,
text=content,
text=combined,
activity="downloadText",
request_data=data,
block_on_violation=True,
@ -264,10 +272,7 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
Errors are logged but never raised — this mode is non-blocking.
"""
try:
metadata = kwargs.get("metadata") or kwargs.get("litellm_metadata") or {}
user_id = metadata.get(self.user_id_field) or kwargs.get(
"user_api_key_user_id"
)
user_id = self._resolve_user_id_from_logging_kwargs(kwargs)
if not user_id:
verbose_proxy_logger.debug("Purview audit: no user_id, skipping")
@ -276,11 +281,11 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
# Log prompt (uploadText)
messages = kwargs.get("messages")
if messages:
user_prompt = self.get_user_prompt(messages)
if user_prompt:
prompt_text = self.get_prompt_text_for_dlp(messages)
if prompt_text:
await self._check_content(
user_id=user_id,
text=user_prompt,
text=prompt_text,
activity="uploadText",
request_data=kwargs,
block_on_violation=False,
@ -290,16 +295,27 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
from litellm.types.utils import Choices, ModelResponse
if isinstance(result, ModelResponse) and result.choices:
if isinstance(result.choices[0], Choices):
content = result.choices[0].message.content or ""
if content:
await self._check_content(
user_id=user_id,
text=content,
activity="downloadText",
request_data=kwargs,
block_on_violation=False,
)
parts: List[str] = []
for choice in result.choices:
if not isinstance(choice, Choices):
continue
msg = choice.message
if msg is None:
continue
raw = msg.get("content") if isinstance(msg, dict) else getattr(
msg, "content", None
)
if isinstance(raw, str) and raw.strip():
parts.append(raw)
if parts:
combined = "\n\n---\n\n".join(parts)
await self._check_content(
user_id=user_id,
text=combined,
activity="downloadText",
request_data=kwargs,
block_on_violation=False,
)
except Exception as e:
verbose_proxy_logger.error("Purview audit logging error: %s", e)

View file

@ -133,15 +133,36 @@ class TestShouldBlock:
class TestResolveUserId:
def test_from_metadata(self):
def test_from_metadata_when_no_auth_identity(self):
guardrail = _make_guardrail()
data = {"metadata": {"user_id": "entra-user-123"}}
assert guardrail._resolve_user_id(data, Mock()) == "entra-user-123"
auth = UserAPIKeyAuth(api_key="test-key-no-user")
assert guardrail._resolve_user_id(data, auth) == "entra-user-123"
def test_custom_field(self):
def test_authenticated_user_id_overrides_metadata(self):
"""Key user_id must win over spoofed metadata[user_id_field]."""
guardrail = _make_guardrail()
data = {"metadata": {"user_id": "spoofed-entra-id"}}
auth = UserAPIKeyAuth(api_key="test", user_id="real-entra-id")
assert guardrail._resolve_user_id(data, auth) == "real-entra-id"
def test_user_api_key_metadata_before_custom_field(self):
"""Proxy-injected user_api_key_user_id wins over arbitrary metadata field."""
guardrail = _make_guardrail(user_id_field="entra_id")
data = {
"metadata": {
"user_api_key_user_id": "from-proxy-111",
"entra_id": "metadata-222",
}
}
auth = UserAPIKeyAuth(api_key="test")
assert guardrail._resolve_user_id(data, auth) == "from-proxy-111"
def test_custom_field_when_no_stronger_source(self):
guardrail = _make_guardrail(user_id_field="entra_id")
data = {"metadata": {"entra_id": "custom-user-456"}}
assert guardrail._resolve_user_id(data, Mock()) == "custom-user-456"
auth = UserAPIKeyAuth(api_key="test")
assert guardrail._resolve_user_id(data, auth) == "custom-user-456"
def test_from_user_api_key_dict_user_id(self):
guardrail = _make_guardrail()
@ -150,16 +171,20 @@ class TestResolveUserId:
def test_from_end_user_id(self):
guardrail = _make_guardrail()
auth = Mock()
auth.user_id = None
auth.end_user_id = "end-user-101"
auth = UserAPIKeyAuth(api_key="test", end_user_id="end-user-101")
assert guardrail._resolve_user_id({}, auth) == "end-user-101"
def test_end_user_id_after_key_user_id(self):
"""When both key user_id and end_user_id exist, key user_id is used first."""
guardrail = _make_guardrail()
auth = UserAPIKeyAuth(
api_key="test", user_id="key-owner", end_user_id="end-user-101"
)
assert guardrail._resolve_user_id({}, auth) == "key-owner"
def test_none_when_missing(self):
guardrail = _make_guardrail()
auth = Mock()
auth.user_id = None
auth.end_user_id = None
auth = UserAPIKeyAuth(api_key="test")
assert guardrail._resolve_user_id({}, auth) is None
@ -255,6 +280,38 @@ class TestPreCallHook:
mock_check.assert_not_called()
class TestPreCallFullTranscript:
@pytest.mark.asyncio
async def test_pre_call_sends_all_message_roles_to_dlp(self):
"""DLP text must include system / prior turns, not only the last user block."""
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={
"messages": [
{"role": "system", "content": "SYSTEM_SENSITIVE"},
{"role": "user", "content": "EARLIER_USER"},
{"role": "assistant", "content": "reply"},
{"role": "user", "content": "final benign"},
]
},
call_type="completion",
)
mock_check.assert_called_once()
sent = mock_check.call_args.kwargs["text"]
assert "SYSTEM_SENSITIVE" in sent
assert "EARLIER_USER" in sent
assert "final benign" in sent
# ---------------------------------------------------------------
# Post-call hook
# ---------------------------------------------------------------
@ -346,6 +403,71 @@ class TestPostCallHook:
mock_check.assert_not_called()
assert result is response
@pytest.mark.asyncio
async def test_post_call_scans_all_choices(self):
from litellm.types.utils import Choices, Message, ModelResponse
guardrail = _make_guardrail()
response = ModelResponse(
choices=[
Choices(
index=0, message=Message(content="First completion", role="assistant")
),
Choices(
index=1,
message=Message(content="Second completion body", role="assistant"),
),
],
)
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=UserAPIKeyAuth(api_key="test"),
response=response,
)
mock_check.assert_called_once()
combined = mock_check.call_args.kwargs["text"]
assert "First completion" in combined
assert "Second completion body" in combined
# ---------------------------------------------------------------
# Logging hook user resolution
# ---------------------------------------------------------------
class TestLoggingResolveUserId:
def test_logging_prefers_user_api_key_user_id_in_metadata(self):
guardrail = _make_guardrail()
kwargs = {
"litellm_params": {
"metadata": {
"user_api_key_user_id": "trusted-from-proxy",
"user_id": "metadata-spoof",
}
}
}
assert (
guardrail._resolve_user_id_from_logging_kwargs(kwargs)
== "trusted-from-proxy"
)
def test_logging_falls_back_to_user_id_field(self):
guardrail = _make_guardrail()
kwargs = {
"litellm_params": {"metadata": {"user_id": "only-metadata-user"}}
}
assert (
guardrail._resolve_user_id_from_logging_kwargs(kwargs)
== "only-metadata-user"
)
# ---------------------------------------------------------------
# _check_content — integration-level