fix(purview): fail closed on identity spoofing, token prompts, and path encoding

Encode Entra user IDs in Graph paths, guard caches with asyncio.Lock, scan
Responses API instructions with string input, reject caller-only metadata and
token-id completion prompts in blocking mode, and revert unrelated UI HTML
restructure from the PR branch.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-22 18:29:47 +05:30
parent 85f2d04a83
commit c92bb00a23
No known key found for this signature in database
3 changed files with 167 additions and 75 deletions

View file

@ -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

View file

@ -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)

View file

@ -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):