diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index c1b46ea4f7f..ba74e72f578 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -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 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 8d7aa9d3d74..6a7cc7548c2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -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) 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 f10b93fa85a..138c8608b37 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 @@ -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