mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(purview): comprehensive security hardening — identity spoofing, streaming bypass, token-id gap
Some checks failed
Unit Tests: Caching (Redis) / caching-redis (push) Has been cancelled
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-utils (push) Has been cancelled
Unit Tests: Proxy DB Operations / auth-checks (push) Has been cancelled
Unit Tests: Proxy DB Operations / budgets (push) Has been cancelled
Unit Tests: Proxy DB Operations / custom-logging (push) Has been cancelled
Unit Tests: Proxy DB Operations / db-and-spend (push) Has been cancelled
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Has been cancelled
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Has been cancelled
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Has been cancelled
Unit Tests: Proxy DB Operations / key-generation (push) Has been cancelled
Unit Tests: Proxy DB Operations / logging-misc (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-runtime (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-server-core (push) Has been cancelled
Unit Tests: Proxy DB Operations / schema-migration (push) Has been cancelled
Some checks failed
Unit Tests: Caching (Redis) / caching-redis (push) Has been cancelled
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-utils (push) Has been cancelled
Unit Tests: Proxy DB Operations / auth-checks (push) Has been cancelled
Unit Tests: Proxy DB Operations / budgets (push) Has been cancelled
Unit Tests: Proxy DB Operations / custom-logging (push) Has been cancelled
Unit Tests: Proxy DB Operations / db-and-spend (push) Has been cancelled
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Has been cancelled
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Has been cancelled
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Has been cancelled
Unit Tests: Proxy DB Operations / key-generation (push) Has been cancelled
Unit Tests: Proxy DB Operations / logging-misc (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-runtime (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-server-core (push) Has been cancelled
Unit Tests: Proxy DB Operations / schema-migration (push) Has been cancelled
Four security issues addressed:
1. end_user_id kwargs fallback missing in _resolve_user_id_from_logging_kwargs
user_id already fell back to kwargs.get("user_api_key_user_id") when absent
from metadata, but end_user_id only checked md.get("user_api_key_end_user_id")
with no kwargs-level fallback. Added or kwargs.get("user_api_key_end_user_id").
2. Streaming responses bypassed post_call blocking
async_post_call_success_hook only runs on assembled non-streaming responses.
For streaming requests the proxy already delivered all content before the
hook ran, so raising HTTPException there had no effect. Added
async_post_call_streaming_iterator_hook which buffers the entire stream,
assembles it via stream_chunk_builder, runs the Purview DLP check, and only
then re-yields chunks via MockResponseIterator. If a violation is detected the
exception is raised before any bytes reach the client. The proxy automatically
skips async_post_call_success_hook for guardrails that define this method,
preventing duplicate scans.
3. Caller-controlled Purview user identity in blocking modes
When a LiteLLM API key has no bound user_id the guardrail fell back to
metadata[user_id_field], which is supplied by the caller. A caller could set
this to any Entra object ID whose Purview policies are more permissive and
bypass DLP. Added _resolve_trusted_user_id() that only returns identities
from the proxy auth system (user_api_key_dict.user_id, end_user_id, or
proxy-injected metadata["user_api_key_user_id"]). Added
_resolve_user_id_for_blocking() used by all blocking-mode hooks: tries
trusted sources first; if only caller-supplied is available, logs a
SECURITY WARNING and still proceeds (backward compat); if nothing resolves,
skips with a warning.
4. Token-id prompt DLP bypass
When /v1/completions received a pure token-id array prompt,
completion_prompt_to_str() returned None and the pre_call hook silently
skipped the Purview scan. An authenticated caller could tokenize blocked
text and send it without DLP evaluation. The hook now detects this case
(raw_prompt present but prompt_text None) and logs a WARNING while letting
the request pass through — token-id payloads are opaque at the text layer
and cannot be scanned. This makes the gap explicit rather than silent.
Tests: 94 total, all passing.
Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com>
This commit is contained in:
parent
c70c4303b7
commit
2346ae21e9
3 changed files with 489 additions and 13 deletions
|
|
@ -291,6 +291,34 @@ class PurviewGuardrailBase:
|
|||
md = litellm_params.get("metadata")
|
||||
return md if isinstance(md, dict) else {}
|
||||
|
||||
def _resolve_trusted_user_id(
|
||||
self, data: Dict[str, Any], user_api_key_dict: Any
|
||||
) -> 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.
|
||||
|
||||
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).
|
||||
"""
|
||||
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(
|
||||
self, kwargs: Dict[str, Any]
|
||||
) -> Optional[str]:
|
||||
|
|
@ -299,7 +327,8 @@ class PurviewGuardrailBase:
|
|||
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"),
|
||||
end_user_id=md.get("user_api_key_end_user_id")
|
||||
or kwargs.get("user_api_key_end_user_id"),
|
||||
)
|
||||
return self._resolve_user_id({"metadata": md}, shim)
|
||||
|
||||
|
|
|
|||
|
|
@ -11,7 +11,18 @@ import asyncio
|
|||
import threading
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Type, Union, cast
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -25,6 +36,7 @@ from litellm.types.utils import (
|
|||
Choices,
|
||||
GuardrailStatus,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
ResponsesAPIResponse,
|
||||
TextChoices,
|
||||
TextCompletionResponse,
|
||||
|
|
@ -232,6 +244,45 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
)
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Identity resolution for blocking modes
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _resolve_user_id_for_blocking(
|
||||
self,
|
||||
data: Dict[str, Any],
|
||||
user_api_key_dict: Any,
|
||||
) -> 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.
|
||||
|
||||
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.
|
||||
"""
|
||||
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,
|
||||
)
|
||||
return caller_id
|
||||
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: no user_id resolved; skipping DLP check"
|
||||
)
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Pre-call hook — DLP on prompts
|
||||
# ------------------------------------------------------------------
|
||||
|
|
@ -245,19 +296,28 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
call_type: "CallTypesLiteral",
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Check user prompt against Purview DLP policies before LLM call."""
|
||||
user_id = self._resolve_user_id(data, user_api_key_dict)
|
||||
user_id = self._resolve_user_id_for_blocking(data, user_api_key_dict)
|
||||
if not user_id:
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: No user_id found, skipping pre-call check"
|
||||
)
|
||||
return data
|
||||
|
||||
prompt_text: Optional[str] = None
|
||||
is_text_completion = call_type in ("text_completion", "atext_completion")
|
||||
messages: Optional[List] = data.get("messages")
|
||||
if messages:
|
||||
prompt_text = self.get_prompt_text_for_dlp(cast(List[Any], messages))
|
||||
elif call_type in ("text_completion", "atext_completion"):
|
||||
prompt_text = self.completion_prompt_to_str(data.get("prompt"))
|
||||
elif is_text_completion:
|
||||
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"
|
||||
)
|
||||
return data
|
||||
elif call_type in ("responses", "aresponses"):
|
||||
prompt_text = self._responses_api_input_to_str(data)
|
||||
|
||||
|
|
@ -283,12 +343,14 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
response: Union[Any, ModelResponse, "EmbeddingResponse", "ImageResponse"],
|
||||
) -> Any:
|
||||
"""Check LLM response against Purview DLP policies."""
|
||||
user_id = self._resolve_user_id(data, user_api_key_dict)
|
||||
"""Check LLM response against Purview DLP policies (non-streaming only).
|
||||
|
||||
Streaming responses are handled by ``async_post_call_streaming_iterator_hook``
|
||||
which buffers all chunks before scanning. The proxy automatically skips
|
||||
this hook for requests that have a streaming iterator hook defined.
|
||||
"""
|
||||
user_id = self._resolve_user_id_for_blocking(data, user_api_key_dict)
|
||||
if not user_id:
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: No user_id found, skipping post-call check"
|
||||
)
|
||||
return response
|
||||
|
||||
parts = self._completion_response_text_parts(response)
|
||||
|
|
@ -304,6 +366,57 @@ class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail):
|
|||
)
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
"""Check streaming LLM responses against Purview DLP policies.
|
||||
|
||||
All chunks are buffered before the DLP scan so that no content is
|
||||
delivered to the client if a policy violation is detected. After a
|
||||
clean scan the assembled response is re-yielded chunk-by-chunk via a
|
||||
``MockResponseIterator`` so the caller receives normal streaming output.
|
||||
|
||||
The proxy automatically skips ``async_post_call_success_hook`` for
|
||||
guardrails that define this method, preventing duplicate scans.
|
||||
"""
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
|
||||
# Buffer the entire stream before any DLP scan.
|
||||
all_chunks: List[ModelResponseStream] = []
|
||||
async for chunk in response:
|
||||
all_chunks.append(chunk)
|
||||
|
||||
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:
|
||||
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 (or skipped) — re-yield chunks from the assembled response.
|
||||
mock_response = MockResponseIterator(model_response=assembled_response)
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Logging-only hook — audit without blocking
|
||||
# ------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -1670,6 +1670,340 @@ class TestCompletionResponseTextPartsToolCalls:
|
|||
assert '{"credit_card": "4111-1111-1111-1111"}' in sent_text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# _resolve_trusted_user_id
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestResolveTrustedUserId:
|
||||
def test_trusted_user_id_from_api_key_dict(self):
|
||||
guardrail = _make_guardrail()
|
||||
auth = UserAPIKeyAuth(api_key="test", user_id="auth-user-111")
|
||||
assert guardrail._resolve_trusted_user_id({}, auth) == "auth-user-111"
|
||||
|
||||
def test_trusted_user_id_from_end_user_id(self):
|
||||
guardrail = _make_guardrail()
|
||||
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):
|
||||
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"
|
||||
|
||||
def test_trusted_user_id_returns_none_for_caller_supplied_only(self):
|
||||
"""Caller-supplied metadata must NOT be returned by _resolve_trusted_user_id."""
|
||||
guardrail = _make_guardrail()
|
||||
auth = UserAPIKeyAuth(api_key="test")
|
||||
data = {"metadata": {"user_id": "caller-supplied-444"}}
|
||||
assert guardrail._resolve_trusted_user_id(data, auth) is None
|
||||
|
||||
def test_trusted_prefers_key_user_id_over_end_user_id(self):
|
||||
guardrail = _make_guardrail()
|
||||
auth = UserAPIKeyAuth(
|
||||
api_key="test", user_id="key-owner", end_user_id="end-user"
|
||||
)
|
||||
assert guardrail._resolve_trusted_user_id({}, auth) == "key-owner"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# _resolve_user_id_from_logging_kwargs — end_user_id fallback
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLoggingEndUserIdFallback:
|
||||
def test_end_user_id_from_metadata(self):
|
||||
guardrail = _make_guardrail()
|
||||
kwargs = {
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"user_api_key_end_user_id": "end-user-from-metadata",
|
||||
}
|
||||
}
|
||||
}
|
||||
result = guardrail._resolve_user_id_from_logging_kwargs(kwargs)
|
||||
assert result == "end-user-from-metadata"
|
||||
|
||||
def test_end_user_id_fallback_from_top_level_kwargs(self):
|
||||
"""end_user_id must fall back to kwargs-level key when absent from metadata."""
|
||||
guardrail = _make_guardrail()
|
||||
kwargs = {
|
||||
"user_api_key_end_user_id": "end-user-from-kwargs",
|
||||
"litellm_params": {"metadata": {}},
|
||||
}
|
||||
result = guardrail._resolve_user_id_from_logging_kwargs(kwargs)
|
||||
assert result == "end-user-from-kwargs"
|
||||
|
||||
def test_metadata_end_user_id_wins_over_kwargs_level(self):
|
||||
"""Metadata value must win over top-level kwargs when both exist."""
|
||||
guardrail = _make_guardrail()
|
||||
kwargs = {
|
||||
"user_api_key_end_user_id": "kwargs-end-user",
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"user_api_key_end_user_id": "metadata-end-user",
|
||||
}
|
||||
},
|
||||
}
|
||||
result = guardrail._resolve_user_id_from_logging_kwargs(kwargs)
|
||||
assert result == "metadata-end-user"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# _resolve_user_id_for_blocking — security warning path
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestResolveUserIdForBlocking:
|
||||
def test_trusted_id_returned_without_warning(self, caplog):
|
||||
import logging
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
auth = UserAPIKeyAuth(api_key="test", user_id="trusted-111")
|
||||
with caplog.at_level(logging.WARNING):
|
||||
result = guardrail._resolve_user_id_for_blocking({}, auth)
|
||||
assert result == "trusted-111"
|
||||
assert "SECURITY" not in caplog.text
|
||||
|
||||
def test_caller_supplied_id_returns_with_security_warning(self, caplog):
|
||||
import logging
|
||||
|
||||
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
|
||||
|
||||
def test_no_id_returns_none(self, caplog):
|
||||
import logging
|
||||
|
||||
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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Token-id prompt handling in pre_call blocking mode
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
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
|
||||
|
||||
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(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"),
|
||||
cache=None,
|
||||
data={"prompt": [1, 2, 3, 100, 200]},
|
||||
call_type="text_completion",
|
||||
)
|
||||
|
||||
mock_check.assert_not_called()
|
||||
assert result is not None
|
||||
assert "token-id" in caplog.text.lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_prompt_skips_without_warning(self, caplog):
|
||||
"""No prompt at all → silently skip (not a token-id bypass case)."""
|
||||
import logging
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail, "_check_content", new_callable=AsyncMock
|
||||
) as mock_check:
|
||||
with caplog.at_level(logging.WARNING):
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="u1"),
|
||||
cache=None,
|
||||
data={},
|
||||
call_type="text_completion",
|
||||
)
|
||||
|
||||
mock_check.assert_not_called()
|
||||
assert "token-id" not in caplog.text.lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_string_prompt_still_scanned(self):
|
||||
"""Normal string prompts must still be sent to Purview."""
|
||||
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="u1"),
|
||||
cache=None,
|
||||
data={"prompt": "sensitive text"},
|
||||
call_type="text_completion",
|
||||
)
|
||||
|
||||
mock_check.assert_called_once()
|
||||
assert mock_check.call_args.kwargs["text"] == "sensitive text"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Streaming iterator hook
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStreamingIteratorHook:
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_clean_response_yields_all_chunks(self):
|
||||
"""Clean stream: all chunks must be re-yielded after DLP passes."""
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
assembled_response = ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
message=Message(content="safe response", role="assistant"),
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
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,
|
||||
patch.object(
|
||||
guardrail, "_check_content", new_callable=AsyncMock
|
||||
) as mock_check,
|
||||
):
|
||||
mock_check.return_value = {"policyActions": []}
|
||||
|
||||
async def _iter_chunks():
|
||||
yield assembled_response
|
||||
|
||||
mock_iterator_cls.return_value.__aiter__ = lambda s: _iter_chunks()
|
||||
|
||||
chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"),
|
||||
response=fake_response_stream(),
|
||||
request_data={"metadata": {"user_id": "user-123"}},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
mock_check.assert_called_once()
|
||||
assert mock_check.call_args.kwargs["activity"] == "downloadText"
|
||||
assert len(chunks) > 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_violation_raises_before_any_chunk(self):
|
||||
"""A policy violation must raise HTTPException before yielding any chunk."""
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
assembled_response = ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
message=Message(
|
||||
content="SSN: 123-45-6789",
|
||||
role="assistant",
|
||||
),
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
async def fake_response_stream():
|
||||
yield assembled_response
|
||||
|
||||
with (
|
||||
patch("litellm.main.stream_chunk_builder", return_value=assembled_response),
|
||||
patch.object(
|
||||
guardrail,
|
||||
"_check_content",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Microsoft Purview DLP: Content blocked by policy"
|
||||
},
|
||||
),
|
||||
),
|
||||
):
|
||||
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", user_id="user-123"
|
||||
),
|
||||
response=fake_response_stream(),
|
||||
request_data={"metadata": {"user_id": "user-123"}},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
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."""
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
assembled_response = ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
message=Message(content="some content", role="assistant"),
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
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,
|
||||
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()
|
||||
|
||||
chunks = []
|
||||
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)
|
||||
|
||||
mock_check.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Auto-discovery registration
|
||||
# ---------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue