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