diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index 79af171ae2d..40a36a6e502 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -1,3 +1,4 @@ +import asyncio import time import uuid from collections import OrderedDict @@ -5,6 +6,7 @@ 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.url_utils import encode_url_path_segment from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_content_list_to_str, ) @@ -66,6 +68,12 @@ class PurviewGuardrailBase: OrderedDict() ) self._scope_cache_maxsize = 1000 + self._cache_lock = asyncio.Lock() + + @staticmethod + def _encode_graph_user_id(user_id: str) -> str: + """Percent-encode Entra user id for Graph ``/users/{id}/...`` path segments.""" + return encode_url_path_segment(user_id, field_name="user_id") # ------------------------------------------------------------------ # OAuth2 token management @@ -74,8 +82,9 @@ class PurviewGuardrailBase: async def _get_access_token(self) -> str: """Acquire or return cached OAuth2 token via client_credentials grant.""" now = time.time() - if self._token_cache and self._token_cache[1] > now + 60: - return self._token_cache[0] + async with self._cache_lock: + if self._token_cache and self._token_cache[1] > now + 60: + return self._token_cache[0] url = TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id) data = { @@ -93,7 +102,8 @@ class PurviewGuardrailBase: token_data = response.json() access_token = token_data["access_token"] expires_in = int(token_data.get("expires_in", 3599)) - self._token_cache = (access_token, now + expires_in) + async with self._cache_lock: + self._token_cache = (access_token, now + expires_in) verbose_proxy_logger.debug( "Purview: acquired new OAuth2 token (expires_in=%ds)", expires_in ) @@ -144,15 +154,17 @@ class PurviewGuardrailBase: Returns: Tuple of (etag, scope_response). """ - cached = self._scope_cache.get(user_id) + encoded_user_id = self._encode_graph_user_id(user_id) now = time.time() - if cached and (now - cached[2]) < SCOPE_CACHE_TTL_SECONDS: - self._scope_cache.move_to_end(user_id) - return cached[0], cached[1] + async with self._cache_lock: + cached = self._scope_cache.get(user_id) + if cached and (now - cached[2]) < SCOPE_CACHE_TTL_SECONDS: + self._scope_cache.move_to_end(user_id) + return cached[0], cached[1] url = ( - f"{GRAPH_API_BASE}/users/{user_id}" + f"{GRAPH_API_BASE}/users/{encoded_user_id}" "/dataSecurityAndGovernance/protectionScopes/compute" ) body: Dict[str, Any] = { @@ -168,14 +180,15 @@ class PurviewGuardrailBase: response_json, response_headers = await self._graph_post(url, body) 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) + async with self._cache_lock: + 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) return etag, response_json # ------------------------------------------------------------------ @@ -199,8 +212,9 @@ class PurviewGuardrailBase: etag: Cached ETag from protectionScopes/compute. correlation_id: Optional conversation/thread ID. """ + encoded_user_id = self._encode_graph_user_id(user_id) url = ( - f"{GRAPH_API_BASE}/users/{user_id}" + f"{GRAPH_API_BASE}/users/{encoded_user_id}" "/dataSecurityAndGovernance/processContent" ) body: Dict[str, Any] = { @@ -244,7 +258,8 @@ class PurviewGuardrailBase: # If policies changed, invalidate scope cache so next call re-fetches. if response_json.get("protectionScopeState") == "modified": - self._scope_cache.pop(user_id, None) + async with self._cache_lock: + self._scope_cache.pop(user_id, None) return response_json 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 059e14a5e35..366541955cb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -227,13 +227,13 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) input_data = data.get("input") - if input_data is None: + if input_data is None and not data.get("instructions"): return None - if isinstance(input_data, str): - return input_data.strip() or None try: + # Always transform via messages so ``instructions`` become a system message + # (string ``input`` alone would skip instructions and bypass DLP). messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( - input=input_data, + input=input_data if input_data is not None else "", responses_api_request=data, ) return self.get_prompt_text_for_dlp(cast(List[Any], messages)) @@ -255,31 +255,30 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): ) -> Optional[str]: """Resolve user ID for blocking (pre_call / post_call) DLP hooks. - Tries trusted sources first (``_resolve_trusted_user_id``). When none - are available the method falls back to the full resolution (which - includes caller-supplied ``metadata[user_id_field]``) and logs a - **SECURITY WARNING** so operators know the DLP evaluation is tied to - an identity the caller could have forged. + Uses only trusted proxy-authenticated sources (``_resolve_trusted_user_id``). + Caller-supplied ``metadata[user_id_field]`` is rejected (fail closed) because + it can impersonate another Entra user's Purview policy. - Returns ``None`` (and logs a warning) when no identity at all can be - resolved — the calling hook should skip the DLP check in that case. + Returns ``None`` when no trusted identity exists — hooks skip the DLP check. """ trusted_id = self._resolve_trusted_user_id(data, user_api_key_dict) if trusted_id: return trusted_id - caller_id = self._resolve_user_id(data, user_api_key_dict) - if caller_id: - verbose_proxy_logger.warning( - "Purview DLP [SECURITY]: no proxy-authenticated user identity found; " - "using caller-supplied user_id=%r for blocking DLP evaluation. " - "Bind a user_id to the API key to prevent identity spoofing.", - caller_id, + if self._resolve_user_id(data, user_api_key_dict): + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: No proxy-authenticated user identity; " + "bind user_id to the API key (caller-supplied metadata cannot " + "be used for blocking DLP)" + ), + }, ) - return caller_id verbose_proxy_logger.warning( - "Purview DLP: no user_id resolved; skipping DLP check" + "Purview DLP: no trusted user_id resolved; skipping DLP check" ) return None @@ -309,15 +308,15 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): raw_prompt = data.get("prompt") prompt_text = self.completion_prompt_to_str(raw_prompt) if raw_prompt is not None and prompt_text is None: - # Prompt was provided but resolved to None — the only case where - # completion_prompt_to_str returns None for non-empty input is a - # pure token-id array. Token IDs cannot be scanned by Purview; - # pass through without DLP (content opaque at the text layer). - verbose_proxy_logger.warning( - "Purview DLP: prompt is token-id array and cannot be scanned; " - "request will proceed without DLP evaluation" + raise HTTPException( + status_code=400, + detail={ + "error": ( + "Microsoft Purview DLP: Token-id completion prompts " + "cannot be scanned for DLP in blocking mode" + ), + }, ) - return data elif call_type in ("responses", "aresponses"): prompt_text = self._responses_api_input_to_str(data) 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 3719d8a762d..3bd91d85d2a 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 @@ -381,8 +381,8 @@ class TestPostCallHook: 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"), + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), response=response, ) @@ -417,8 +417,8 @@ class TestPostCallHook: with pytest.raises(HTTPException) as exc_info: await guardrail.async_post_call_success_hook( - data={"metadata": {"user_id": "user-123"}}, - user_api_key_dict=UserAPIKeyAuth(api_key="test"), + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), response=response, ) @@ -471,8 +471,8 @@ class TestPostCallHook: 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"), + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), response=response, ) @@ -522,8 +522,8 @@ class TestTextCompletionHooks: 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"), + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), response=response, ) @@ -619,6 +619,53 @@ class TestResponsesAPIHooks: mock_check.assert_not_called() + @pytest.mark.asyncio + async def test_pre_call_responses_string_input_includes_instructions(self): + """Benign string ``input`` must still scan ``instructions`` (system message).""" + 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": "benign user text", + "instructions": "SYSTEM_SENSITIVE in instructions", + }, + call_type="responses", + ) + + mock_check.assert_called_once() + sent = mock_check.call_args.kwargs["text"] + assert "benign user text" in sent + assert "SYSTEM_SENSITIVE in instructions" in sent + + @pytest.mark.asyncio + async def test_pre_call_responses_instructions_only(self): + """Requests with only ``instructions`` (no ``input``) must still be scanned.""" + 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={"instructions": "policy text in instructions only"}, + call_type="responses", + ) + + mock_check.assert_called_once() + assert "policy text in instructions only" in mock_check.call_args.kwargs[ + "text" + ] + @pytest.mark.asyncio async def test_post_call_responses_api_output_text(self): """Post-call hook must scan text from ``ResponsesAPIResponse.output``.""" @@ -647,8 +694,8 @@ class TestResponsesAPIHooks: 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"), + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), response=response, ) @@ -673,8 +720,8 @@ class TestResponsesAPIHooks: 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"), + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), response=response, ) @@ -1658,10 +1705,10 @@ class TestCompletionResponseTextPartsToolCalls: mock_check.return_value = {"policyActions": []} await guardrail.async_post_call_success_hook( - data={"metadata": {"user_id": "user-123"}}, + data={}, user_api_key_dict=__import__( "litellm.proxy._types", fromlist=["UserAPIKeyAuth"] - ).UserAPIKeyAuth(api_key="test"), + ).UserAPIKeyAuth(api_key="test", user_id="user-123"), response=response, ) @@ -1670,6 +1717,41 @@ class TestCompletionResponseTextPartsToolCalls: assert '{"credit_card": "4111-1111-1111-1111"}' in sent_text +# --------------------------------------------------------------- +# Graph user id path encoding +# --------------------------------------------------------------- + + +class TestGraphUserIdEncoding: + def test_encode_graph_user_id_percent_encodes_special_chars(self): + from urllib.parse import quote + + raw = "user/with%special" + encoded = PurviewGuardrailBase._encode_graph_user_id(raw) + assert encoded == quote(raw, safe="") + assert "/" not in encoded + + @pytest.mark.asyncio + async def test_compute_protection_scopes_uses_encoded_path(self): + guardrail = _make_guardrail() + guardrail._scope_cache.clear() + + mock_resp = _mock_scope_response() + + async def _capture_post(url, **kwargs): + assert "/users/" in url + assert "user%2Fwith%25special" in url + return mock_resp + + guardrail.async_handler.post = AsyncMock(side_effect=_capture_post) + + with patch.object(guardrail, "_get_access_token", new_callable=AsyncMock) as mock_token: + mock_token.return_value = "tok" + await guardrail._compute_protection_scopes("user/with%special") + + guardrail.async_handler.post.assert_called_once() + + # --------------------------------------------------------------- # _resolve_trusted_user_id # --------------------------------------------------------------- @@ -1766,16 +1848,14 @@ class TestResolveUserIdForBlocking: assert result == "trusted-111" assert "SECURITY" not in caplog.text - def test_caller_supplied_id_returns_with_security_warning(self, caplog): - import logging - + def test_caller_supplied_id_raises_http_exception(self): guardrail = _make_guardrail() auth = UserAPIKeyAuth(api_key="test") data = {"metadata": {"user_id": "caller-supplied-999"}} - with caplog.at_level(logging.WARNING): - result = guardrail._resolve_user_id_for_blocking(data, auth) - assert result == "caller-supplied-999" - assert "SECURITY" in caplog.text + with pytest.raises(HTTPException) as exc_info: + guardrail._resolve_user_id_for_blocking(data, auth) + assert exc_info.value.status_code == 400 + assert "proxy-authenticated" in str(exc_info.value.detail) def test_no_id_returns_none(self, caplog): import logging @@ -1794,17 +1874,15 @@ class TestResolveUserIdForBlocking: class TestTokenIdPromptHandling: @pytest.mark.asyncio - async def test_token_id_prompt_passes_through_with_warning(self, caplog): - """Pure token-id prompts cannot be scanned; they should pass through.""" - import logging - + async def test_token_id_prompt_raises_in_blocking_mode(self): + """Pure token-id prompts must be rejected in blocking pre_call mode.""" guardrail = _make_guardrail() with patch.object( guardrail, "_check_content", new_callable=AsyncMock ) as mock_check: - with caplog.at_level(logging.WARNING): - result = await guardrail.async_pre_call_hook( + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"), cache=None, data={"prompt": [1, 2, 3, 100, 200]}, @@ -1812,8 +1890,8 @@ class TestTokenIdPromptHandling: ) mock_check.assert_not_called() - assert result is not None - assert "token-id" in caplog.text.lower() + assert exc_info.value.status_code == 400 + assert "Token-id" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_missing_prompt_skips_without_warning(self, caplog):