mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
85f2d04a83
commit
c92bb00a23
3 changed files with 167 additions and 75 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue