mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(purview): fail-closed blocking DLP; revert directory-based UI HTML
Blocking hooks now require UserAPIKeyAuth user_id/end_user_id only (no spoofable metadata), re-raise Responses API transform errors, scan streamed text completions, and reject requests with no bound identity. Reverts the accidental directory-based Next.js output fromcc47081(c70c4303b7). Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
cc47081cf5
commit
212815dd5a
38 changed files with 110 additions and 73 deletions
|
|
@ -293,6 +293,10 @@ class PurviewGuardrailBase:
|
|||
return trusted
|
||||
|
||||
metadata = data.get("metadata") or data.get("litellm_metadata") or {}
|
||||
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)
|
||||
|
|
@ -313,27 +317,19 @@ class PurviewGuardrailBase:
|
|||
) -> Optional[str]:
|
||||
"""Resolve user ID from trusted (proxy-authenticated) sources only.
|
||||
|
||||
Identical trust order to ``_resolve_user_id`` levels 1-3, but intentionally
|
||||
omits level 4 (``metadata[user_id_field]``) because that value is supplied
|
||||
by the caller and can be forged. Use this in blocking modes so that
|
||||
content is always evaluated against the actual authenticated user's Purview
|
||||
policy, not one chosen by the caller.
|
||||
Uses only ``UserAPIKeyAuth.user_id`` and ``UserAPIKeyAuth.end_user_id``.
|
||||
Intentionally omits ``metadata[user_id_field]`` and
|
||||
``metadata["user_api_key_user_id"]`` because those can be supplied or
|
||||
spoofed by the caller when the API key has no bound user.
|
||||
|
||||
Returns ``None`` when no authenticated identity is available, even if a
|
||||
caller-supplied value exists. Callers should then decide whether to fail
|
||||
open (skip) or fail closed (block).
|
||||
Returns ``None`` when no authenticated identity is available. Blocking
|
||||
hooks must fail closed rather than skip the DLP check.
|
||||
"""
|
||||
metadata = data.get("metadata") or data.get("litellm_metadata") or {}
|
||||
|
||||
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)
|
||||
|
||||
return None
|
||||
|
||||
def _resolve_user_id_from_logging_kwargs(
|
||||
|
|
|
|||
|
|
@ -241,20 +241,12 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
)
|
||||
return self.get_prompt_text_for_dlp(cast(List[Any], messages))
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: failed to transform responses API input",
|
||||
exc_info=True,
|
||||
)
|
||||
if raise_on_failure:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: Responses API input could "
|
||||
"not be transformed for DLP scanning in blocking mode"
|
||||
),
|
||||
},
|
||||
)
|
||||
raise
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
|
@ -272,7 +264,8 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
Caller-supplied ``metadata[user_id_field]`` is rejected (fail closed) because
|
||||
it can impersonate another Entra user's Purview policy.
|
||||
|
||||
Returns ``None`` when no trusted identity exists — hooks skip the DLP check.
|
||||
Raises ``HTTPException`` when no trusted identity exists or when only
|
||||
caller-supplied metadata is available (fail closed).
|
||||
"""
|
||||
trusted_id = self._resolve_trusted_user_id(data, user_api_key_dict)
|
||||
if trusted_id:
|
||||
|
|
@ -290,10 +283,15 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
},
|
||||
)
|
||||
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: no trusted user_id resolved; skipping DLP check"
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: No proxy-authenticated user identity; "
|
||||
"bind user_id to the API key for blocking DLP"
|
||||
),
|
||||
},
|
||||
)
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Pre-call hook — DLP on prompts
|
||||
|
|
@ -404,18 +402,12 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
|
||||
assembled_response = stream_chunk_builder(chunks=all_chunks)
|
||||
|
||||
if not isinstance(assembled_response, ModelResponse):
|
||||
# Non-chat response (e.g. embeddings) — pass through unchanged.
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
user_id = self._resolve_user_id_for_blocking(request_data, user_api_key_dict)
|
||||
if user_id:
|
||||
|
||||
if isinstance(assembled_response, TextCompletionResponse):
|
||||
parts = self._completion_response_text_parts(assembled_response)
|
||||
if parts:
|
||||
combined = "\n\n---\n\n".join(parts)
|
||||
# Raises HTTPException(400) on violation — no chunks are yielded.
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=combined,
|
||||
|
|
@ -423,8 +415,29 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
request_data=request_data,
|
||||
block_on_violation=True,
|
||||
)
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
# DLP passed (or skipped) — re-yield chunks from the assembled response.
|
||||
if not isinstance(assembled_response, ModelResponse):
|
||||
# Non-chat response (e.g. embeddings) — pass through unchanged.
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
parts = self._completion_response_text_parts(assembled_response)
|
||||
if parts:
|
||||
combined = "\n\n---\n\n".join(parts)
|
||||
# Raises HTTPException(400) on violation — no chunks are yielded.
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=combined,
|
||||
activity="downloadText",
|
||||
request_data=request_data,
|
||||
block_on_violation=True,
|
||||
)
|
||||
|
||||
# DLP passed — re-yield chunks from the assembled chat response.
|
||||
mock_response = MockResponseIterator(model_response=assembled_response)
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
|
|
|
|||
|
|
@ -290,22 +290,22 @@ class TestPreCallHook:
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_no_user_id_skips(self):
|
||||
async def test_pre_call_no_user_id_raises(self):
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail, "_check_content", new_callable=AsyncMock
|
||||
) as mock_check:
|
||||
result = await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data={"messages": [{"role": "user", "content": "Hello"}]},
|
||||
call_type="completion",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
cache=None,
|
||||
data={"messages": [{"role": "user", "content": "Hello"}]},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
mock_check.assert_not_called()
|
||||
# Returns data when skipping
|
||||
assert result is not None
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_no_messages_skips(self):
|
||||
|
|
@ -427,7 +427,7 @@ class TestPostCallHook:
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_no_user_id_skips(self):
|
||||
async def test_post_call_no_user_id_raises(self):
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
|
|
@ -440,14 +440,15 @@ class TestPostCallHook:
|
|||
with patch.object(
|
||||
guardrail, "_check_content", new_callable=AsyncMock
|
||||
) as mock_check:
|
||||
result = await guardrail.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
response=response,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
response=response,
|
||||
)
|
||||
|
||||
mock_check.assert_not_called()
|
||||
assert result is response
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_scans_all_choices(self):
|
||||
|
|
@ -1807,11 +1808,12 @@ class TestResolveTrustedUserId:
|
|||
auth = UserAPIKeyAuth(api_key="test", end_user_id="end-user-222")
|
||||
assert guardrail._resolve_trusted_user_id({}, auth) == "end-user-222"
|
||||
|
||||
def test_trusted_user_id_from_proxy_injected_metadata(self):
|
||||
def test_metadata_user_api_key_user_id_not_trusted_without_auth(self):
|
||||
"""Metadata user_api_key_user_id is not trusted when the key has no user_id."""
|
||||
guardrail = _make_guardrail()
|
||||
auth = UserAPIKeyAuth(api_key="test")
|
||||
data = {"metadata": {"user_api_key_user_id": "proxy-user-333"}}
|
||||
assert guardrail._resolve_trusted_user_id(data, auth) == "proxy-user-333"
|
||||
assert guardrail._resolve_trusted_user_id(data, auth) is None
|
||||
|
||||
def test_trusted_user_id_returns_none_for_caller_supplied_only(self):
|
||||
"""Caller-supplied metadata must NOT be returned by _resolve_trusted_user_id."""
|
||||
|
|
@ -1896,14 +1898,13 @@ class TestResolveUserIdForBlocking:
|
|||
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
|
||||
|
||||
def test_no_id_raises_http_exception(self):
|
||||
guardrail = _make_guardrail()
|
||||
auth = UserAPIKeyAuth(api_key="test")
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = guardrail._resolve_user_id_for_blocking({}, auth)
|
||||
assert result is None
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
guardrail._resolve_user_id_for_blocking({}, auth)
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "bind user_id" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
|
|
@ -2077,8 +2078,8 @@ class TestStreamingIteratorHook:
|
|||
assert len(chunks) == 0 # No chunks yielded before the block
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_no_user_id_passes_through(self):
|
||||
"""No resolvable user_id → stream passes through without DLP scan."""
|
||||
async def test_streaming_no_user_id_raises_before_yield(self):
|
||||
"""No resolvable user_id → fail closed before any chunk is yielded."""
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
|
|
@ -2092,33 +2093,60 @@ class TestStreamingIteratorHook:
|
|||
]
|
||||
)
|
||||
|
||||
async def fake_response_stream():
|
||||
yield assembled_response
|
||||
|
||||
with patch(
|
||||
"litellm.main.stream_chunk_builder", return_value=assembled_response
|
||||
):
|
||||
chunks = []
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"), # no user_id
|
||||
response=fake_response_stream(),
|
||||
request_data={},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert len(chunks) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_text_completion_scans_before_yield(self):
|
||||
"""Streamed /v1/completions must be DLP-scanned via TextCompletionResponse."""
|
||||
from litellm.types.utils import TextChoices, TextCompletionResponse
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
assembled_response = TextCompletionResponse(
|
||||
model="gpt-3.5-turbo-instruct",
|
||||
choices=[TextChoices(text="completion body", index=0)],
|
||||
)
|
||||
|
||||
async def fake_response_stream():
|
||||
yield assembled_response
|
||||
|
||||
with (
|
||||
patch("litellm.main.stream_chunk_builder", return_value=assembled_response),
|
||||
patch(
|
||||
"litellm.llms.base_llm.base_model_iterator.MockResponseIterator"
|
||||
) as mock_iterator_cls,
|
||||
"litellm.main.stream_chunk_builder", return_value=assembled_response
|
||||
),
|
||||
patch.object(
|
||||
guardrail, "_check_content", new_callable=AsyncMock
|
||||
) as mock_check,
|
||||
):
|
||||
|
||||
async def _iter_chunks():
|
||||
yield assembled_response
|
||||
|
||||
mock_iterator_cls.return_value.__aiter__ = lambda s: _iter_chunks()
|
||||
mock_check.return_value = {"policyActions": []}
|
||||
|
||||
chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test"), # no user_id
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"),
|
||||
response=fake_response_stream(),
|
||||
request_data={},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
mock_check.assert_not_called()
|
||||
mock_check.assert_called_once()
|
||||
assert mock_check.call_args.kwargs["text"] == "completion body"
|
||||
assert len(chunks) > 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue