mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge c9f907e2f4 into 3746ba58d7
This commit is contained in:
commit
7992398d96
5 changed files with 662 additions and 7 deletions
|
|
@ -111,6 +111,70 @@ def anthropic_sse_error_frames(message: str) -> tuple[bytes, ...]:
|
|||
)
|
||||
|
||||
|
||||
def _is_text_delta_event(event: str) -> bool:
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
|
||||
event_data: Final = AnthropicPassthroughLoggingHandler._extract_sse_data(event) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
|
||||
if event_data is None or event_data.get("type") != "content_block_delta":
|
||||
return False
|
||||
delta: Final = event_data.get("delta")
|
||||
return isinstance(delta, dict) and delta.get("type") == "text_delta"
|
||||
|
||||
|
||||
def _text_delta_frame(text: str, index: object) -> bytes:
|
||||
payload: Final = { # mutable-ok: json.dumps serializes a dict
|
||||
"type": "content_block_delta",
|
||||
"index": index if isinstance(index, int) else 0,
|
||||
"delta": {"type": "text_delta", "text": text}, # mutable-ok: json.dumps serializes a dict
|
||||
}
|
||||
return f"event: content_block_delta\ndata: {json.dumps(payload)}\n\n".encode()
|
||||
|
||||
|
||||
def _content_block_index(event: str) -> object:
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
|
||||
event_data: Final = AnthropicPassthroughLoggingHandler._extract_sse_data(event) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
|
||||
return None if event_data is None else event_data.get("index")
|
||||
|
||||
|
||||
def rewrite_anthropic_sse_text(all_chunks: Sequence[object], replacement: str) -> tuple[bytes, ...] | None:
|
||||
"""Re-emit the upstream frames with the assistant text replaced by ``replacement``.
|
||||
|
||||
Rebuilding the stream from an assembled ``ModelResponse`` loses everything the assembler does
|
||||
not model, such as the cache, server-tool-use and service-tier usage fields Anthropic sends in
|
||||
``message_start``. Rewriting the frames in place keeps them. Returns None when the stream holds
|
||||
no text delta to replace.
|
||||
|
||||
Model Armor sanitizes the whole assistant text as one string, so a response split over several
|
||||
text blocks gets all of its sanitized text in the first block and the later text blocks come out
|
||||
empty. Block order and every start/stop pair still hold, and concatenating the text yields
|
||||
exactly what Model Armor returned.
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
|
||||
sse_stream: Final = _joined_sse_stream(all_chunks)
|
||||
if sse_stream is None:
|
||||
return None
|
||||
events: Final = AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
|
||||
text_positions: Final = tuple(position for position, event in enumerate(events) if _is_text_delta_event(event))
|
||||
if not text_positions:
|
||||
return None
|
||||
dropped: Final = frozenset(text_positions[1:])
|
||||
return tuple(
|
||||
_text_delta_frame(replacement, _content_block_index(event))
|
||||
if position == text_positions[0]
|
||||
else f"{event}\n\n".encode()
|
||||
for position, event in enumerate(events)
|
||||
if position not in dropped
|
||||
)
|
||||
|
||||
|
||||
def anthropic_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]:
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
|
@ -7,6 +9,8 @@ from .model_armor import ModelArmorGuardrail
|
|||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
_NO_EXTRAS: Final[Mapping[str, object]] = MappingProxyType({}) # mutable-ok: MappingProxyType needs a dict
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
|
||||
import litellm
|
||||
|
|
@ -14,6 +18,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
ModelArmorGuardrail,
|
||||
)
|
||||
|
||||
streaming_params: Final[Mapping[str, object]] = litellm_params.model_extra or _NO_EXTRAS
|
||||
_model_armor_callback: Final = ModelArmorGuardrail(
|
||||
guardrail_name=guardrail.get("guardrail_name", ""),
|
||||
event_hook=litellm_params.mode,
|
||||
|
|
@ -28,6 +33,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
fail_on_error=litellm_params.fail_on_error,
|
||||
skip_unscannable_attachments=litellm_params.skip_unscannable_attachments,
|
||||
sanitize_error_detail=litellm_params.sanitize_error_detail,
|
||||
streaming_buffer_until_moderated=streaming_params.get("streaming_buffer_until_moderated"),
|
||||
streaming_sampling_rate=streaming_params.get("streaming_sampling_rate"),
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
from collections.abc import AsyncGenerator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
import json
|
||||
|
|
@ -15,6 +16,7 @@ from litellm.caching import DualCache
|
|||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
|
|
@ -37,6 +39,7 @@ from litellm.types.utils import (
|
|||
CallTypes,
|
||||
CallTypesLiteral,
|
||||
Choices,
|
||||
GenericGuardrailAPIInputs,
|
||||
GuardrailStatus,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
|
|
@ -82,6 +85,10 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
- Post-call sanitization (sanitizeModelResponse)
|
||||
"""
|
||||
|
||||
# apply_guardrail only exists to scan streamed chunks: file scanning, masking and the MCP
|
||||
# hooks need the native lifecycle hooks, so they must not be routed to the shared runner
|
||||
use_native_lifecycle_hooks: ClassVar[bool] = True
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
return [
|
||||
|
|
@ -123,6 +130,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
self.credentials = credentials
|
||||
self.api_endpoint = api_endpoint
|
||||
self.sanitize_error_detail = sanitize_error_detail is not False
|
||||
self.streaming_buffer_until_moderated = kwargs.get("streaming_buffer_until_moderated") is not False
|
||||
self.streaming_sampling_rate = kwargs.get("streaming_sampling_rate") or 5
|
||||
|
||||
# Store optional params
|
||||
self.optional_params = kwargs
|
||||
|
|
@ -188,6 +197,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None:
|
||||
super().update_in_memory_litellm_params(litellm_params)
|
||||
self.sanitize_error_detail = self.sanitize_error_detail is not False
|
||||
self.streaming_buffer_until_moderated = self.streaming_buffer_until_moderated is not False
|
||||
self.streaming_sampling_rate = self.streaming_sampling_rate or 5
|
||||
|
||||
def _log_request_debug(
|
||||
self,
|
||||
|
|
@ -831,6 +842,61 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
|
||||
return response
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict, # mutable-ok: CustomGuardrail.apply_guardrail signature
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Scan bare texts, which lets the shared runner drive Model Armor over a live stream."""
|
||||
texts: Final = inputs.get("texts") or ()
|
||||
scanned: Final = [ # mutable-ok: GenericGuardrailAPIInputs.texts is a list
|
||||
await self._scan_text(text, input_type=input_type, request_data=request_data) if text else text
|
||||
for text in texts
|
||||
]
|
||||
return {**inputs, "texts": scanned} # mutable-ok: GenericGuardrailAPIInputs is a plain TypedDict
|
||||
|
||||
async def _scan_text(
|
||||
self,
|
||||
text: str,
|
||||
input_type: Literal["request", "response"],
|
||||
request_data: dict, # mutable-ok: forwarded to make_model_armor_request
|
||||
) -> str:
|
||||
masking_enabled: Final = self.mask_request_content if input_type == "request" else self.mask_response_content
|
||||
try:
|
||||
armor_response: Final = await self.make_model_armor_request(
|
||||
content=text,
|
||||
source="user_prompt" if input_type == "request" else "model_response",
|
||||
request_data=request_data,
|
||||
)
|
||||
except ModelArmorAPIError as e:
|
||||
self._raise_if_fail_closed(e)
|
||||
return text
|
||||
|
||||
if self._should_block_content(armor_response, allow_sanitization=masking_enabled):
|
||||
message: Final = f"{input_type.capitalize()} blocked by Model Armor"
|
||||
if input_type == "request":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=self._build_block_error_detail(message, armor_response),
|
||||
)
|
||||
raise ModifyResponseException(
|
||||
message=message,
|
||||
model=request_data.get("model") or "unknown",
|
||||
request_data=request_data,
|
||||
guardrail_name=self.guardrail_name or GUARDRAIL_NAME,
|
||||
original_response=request_data.get("response"),
|
||||
)
|
||||
|
||||
sanitized: Final = self._get_sanitized_content(armor_response) if masking_enabled else None
|
||||
return sanitized or text
|
||||
|
||||
def _streams_incrementally(self) -> bool:
|
||||
"""Sanitization needs the whole response, so only allow/block configs can scan mid-stream."""
|
||||
return not self.streaming_buffer_until_moderated and not self.mask_response_content
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -839,16 +905,43 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
"""Process streaming response chunks."""
|
||||
|
||||
if self._streams_incrementally():
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
async for streamed_chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=response,
|
||||
request_data=request_data,
|
||||
guardrail_to_apply=self,
|
||||
buffer_until_moderated_default=False,
|
||||
):
|
||||
yield streamed_chunk
|
||||
return
|
||||
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.proxy.guardrails.anthropic_sse import (
|
||||
anthropic_sse_error_frames,
|
||||
assemble_anthropic_sse_stream,
|
||||
is_raw_sse_stream,
|
||||
rewrite_anthropic_sse_text,
|
||||
)
|
||||
|
||||
# Collect all chunks
|
||||
all_chunks: Final[list[ModelResponseStream]] = []
|
||||
all_chunks: Final[list[object]] = [] # mutable-ok: the whole stream has to be held to assemble it
|
||||
async for chunk in response:
|
||||
all_chunks.append(chunk)
|
||||
|
||||
# Build complete response
|
||||
assembled_response: Final = stream_chunk_builder(chunks=all_chunks)
|
||||
raw_sse: Final = is_raw_sse_stream(all_chunks)
|
||||
chat_stream: Final = bool(all_chunks) and all(isinstance(chunk, ModelResponseStream) for chunk in all_chunks)
|
||||
assembled_response: Final = (
|
||||
assemble_anthropic_sse_stream(all_chunks, restore_identity=True)
|
||||
if raw_sse
|
||||
else stream_chunk_builder(chunks=all_chunks)
|
||||
if chat_stream
|
||||
else None
|
||||
)
|
||||
|
||||
if isinstance(assembled_response, ModelResponse):
|
||||
# Extract content
|
||||
|
|
@ -868,7 +961,9 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
_, metadata = get_or_create_metadata_bucket(request_data)
|
||||
metadata["_model_armor_response"] = self._build_logging_response(armor_response)
|
||||
metadata["_model_armor_status"] = (
|
||||
"blocked" if self._should_block_content(armor_response) else "success"
|
||||
"blocked"
|
||||
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content)
|
||||
else "success"
|
||||
)
|
||||
|
||||
# Add guardrail to applied_guardrails BEFORE potential blocking
|
||||
|
|
@ -882,7 +977,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
)
|
||||
|
||||
# Check if blocked
|
||||
if self._should_block_content(armor_response):
|
||||
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=self._build_block_error_detail(
|
||||
|
|
@ -902,6 +997,13 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
choice.message.content = sanitized_content
|
||||
|
||||
# Return sanitized stream
|
||||
if raw_sse:
|
||||
rewritten: Final = rewrite_anthropic_sse_text(all_chunks, sanitized_content)
|
||||
for sse_chunk in rewritten or anthropic_sse_error_frames(
|
||||
f"{self.guardrail_name}: sanitized response could not be re-emitted, blocking it"
|
||||
):
|
||||
yield sse_chunk
|
||||
return
|
||||
mock_response: Final = MockResponseIterator(model_response=assembled_response)
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
|
|
@ -909,6 +1011,10 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
|
||||
except ModelArmorAPIError as e:
|
||||
if self.optional_params.get("fail_on_error", True):
|
||||
if raw_sse:
|
||||
for error_frame in anthropic_sse_error_frames(e.detail):
|
||||
yield error_frame
|
||||
return
|
||||
error_obj = {"message": e.detail, "code": "500"}
|
||||
yield f"data: {json.dumps({'error': error_obj})}\n\n"
|
||||
return
|
||||
|
|
@ -923,6 +1029,10 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
else:
|
||||
error_obj = {"message": str(error_value)}
|
||||
error_obj["code"] = str(e.status_code)
|
||||
if raw_sse:
|
||||
for error_frame in anthropic_sse_error_frames(str(error_obj.get("message", error_obj))):
|
||||
yield error_frame
|
||||
return
|
||||
yield f"data: {json.dumps({'error': error_obj})}\n\n"
|
||||
return
|
||||
except Exception as e:
|
||||
|
|
@ -931,6 +1041,18 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
raise
|
||||
else:
|
||||
verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail")
|
||||
elif raw_sse:
|
||||
# Forwarding an unscannable stream would silently disable the guardrail, so fail closed
|
||||
for error_frame in anthropic_sse_error_frames(
|
||||
f"{self.guardrail_name}: streamed response could not be assembled for scanning, blocking it"
|
||||
):
|
||||
yield error_frame
|
||||
return
|
||||
elif not chat_stream:
|
||||
verbose_proxy_logger.warning(
|
||||
"Model Armor: streaming response contained unsupported event objects "
|
||||
"(e.g. /v1/responses events); output scanning was skipped for this response."
|
||||
)
|
||||
|
||||
# Return original chunks if no sanitization needed
|
||||
for chunk in all_chunks:
|
||||
|
|
|
|||
|
|
@ -26,6 +26,26 @@ class ModelArmorGuardrailConfigModel(GuardrailConfigModel):
|
|||
),
|
||||
)
|
||||
|
||||
streaming_buffer_until_moderated: bool | None = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"True (default) withholds every streamed chunk until Model Armor has scanned the "
|
||||
"assembled response, so nothing unscanned reaches the client but time to first token "
|
||||
"grows to the full generation time. False streams chunks as they arrive and scans them "
|
||||
"as the response grows, at the cadence set by streaming_sampling_rate, terminating the "
|
||||
"stream when a scan blocks. Ignored when mask_response_content is True, since "
|
||||
"sanitization needs the whole response."
|
||||
),
|
||||
)
|
||||
streaming_sampling_rate: int | None = Field(
|
||||
default=5,
|
||||
ge=1,
|
||||
description=(
|
||||
"When streaming_buffer_until_moderated is False, scan every Nth streamed chunk. Lower "
|
||||
"values catch violations sooner, at the cost of more Model Armor calls."
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
"""Return the UI-friendly name for Model Armor guardrail"""
|
||||
|
|
|
|||
|
|
@ -623,6 +623,448 @@ async def test_model_armor_streaming_block_yields_sse_error():
|
|||
assert int(error_data["error"]["code"]) == 400
|
||||
|
||||
|
||||
_ANTHROPIC_SSE_CHUNKS = (
|
||||
b'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_1","type":"message",'
|
||||
b'"role":"assistant","model":"claude","content":[],"usage":{"input_tokens":5,"output_tokens":0,'
|
||||
b'"cache_read_input_tokens":4,"service_tier":"standard"}}}\n\n',
|
||||
b'event: content_block_start\ndata: {"type":"content_block_start","index":0,'
|
||||
b'"content_block":{"type":"text","text":""}}\n\n',
|
||||
b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,'
|
||||
b'"delta":{"type":"text_delta","text":"my ssn is 123-45-6789"}}\n\n',
|
||||
b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n',
|
||||
b'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},'
|
||||
b'"usage":{"output_tokens":9}}\n\n',
|
||||
b'event: message_stop\ndata: {"type":"message_stop"}\n\n',
|
||||
)
|
||||
|
||||
|
||||
def _sse_armor_guardrail(**kwargs: object) -> ModelArmorGuardrail:
|
||||
guardrail = ModelArmorGuardrail(
|
||||
template_id="test-template",
|
||||
project_id="test-project",
|
||||
location="us-central1",
|
||||
guardrail_name="model-armor-test",
|
||||
**kwargs,
|
||||
)
|
||||
guardrail._ensure_access_token_async = AsyncMock(
|
||||
return_value=("test-token", "test-project")
|
||||
)
|
||||
return guardrail
|
||||
|
||||
|
||||
async def _drain_armor_streaming_hook(
|
||||
guardrail: ModelArmorGuardrail, chunks: tuple[object, ...] = _ANTHROPIC_SSE_CHUNKS
|
||||
) -> list[object]:
|
||||
async def _stream():
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
return [
|
||||
chunk
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
response=_stream(),
|
||||
request_data={
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "what is my ssn"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]},
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def _armor_api_response(sanitization_result: dict) -> AsyncMock:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json = AsyncMock(return_value={"sanitizationResult": sanitization_result})
|
||||
return mock_response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_scans_raw_anthropic_sse_instead_of_crashing():
|
||||
"""A /v1/messages stream arrives as raw SSE frames and must be assembled, then scanned.
|
||||
|
||||
Regression for `500 Error building chunks for logging/streaming usage calculation`:
|
||||
stream_chunk_builder subscripts each chunk, which raises TypeError on bytes.
|
||||
"""
|
||||
guardrail = _sse_armor_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(return_value=_armor_api_response({"filterMatchState": "NO_MATCH_FOUND"})),
|
||||
) as mock_post:
|
||||
delivered = await _drain_armor_streaming_hook(guardrail)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
assert "my ssn is 123-45-6789" in json.dumps(mock_post.call_args.kwargs.get("json"))
|
||||
assert tuple(delivered) == _ANTHROPIC_SSE_CHUNKS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_blocks_raw_anthropic_sse_with_error_frame():
|
||||
guardrail = _sse_armor_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(
|
||||
return_value=_armor_api_response(
|
||||
{
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"sdp": {
|
||||
"sdpFilterResult": {
|
||||
"inspectResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"findings": [{"infoType": "US_SOCIAL_SECURITY_NUMBER"}],
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
),
|
||||
):
|
||||
delivered = await _drain_armor_streaming_hook(guardrail)
|
||||
|
||||
body = b"".join(chunk for chunk in delivered if isinstance(chunk, bytes))
|
||||
assert b"event: error" in body
|
||||
assert b"123-45-6789" not in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_masks_raw_anthropic_sse():
|
||||
guardrail = _sse_armor_guardrail(mask_response_content=True)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(
|
||||
return_value=_armor_api_response(
|
||||
{
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"sdp": {
|
||||
"sdpFilterResult": {
|
||||
"deidentifyResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"data": {"text": "my ssn is [REDACTED]"},
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
),
|
||||
):
|
||||
delivered = await _drain_armor_streaming_hook(guardrail)
|
||||
|
||||
body = b"".join(chunk for chunk in delivered if isinstance(chunk, bytes))
|
||||
assert b"[REDACTED]" in body
|
||||
assert b"123-45-6789" not in body
|
||||
# the masked stream is a rewrite of the upstream frames, so usage the assembler does not
|
||||
# model (cache counts, service tier) still reaches the client
|
||||
assert b'"cache_read_input_tokens":4' in body
|
||||
assert b'"service_tier":"standard"' in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_masked_multi_block_anthropic_stream_keeps_every_block():
|
||||
"""Model Armor sanitizes the whole text at once, so block 1 carries it and block 2 empties out.
|
||||
|
||||
The blocks and their start/stop pairs still have to survive in order, otherwise a client
|
||||
tracking content block indexes breaks.
|
||||
"""
|
||||
multi_block_chunks = (
|
||||
_ANTHROPIC_SSE_CHUNKS[0],
|
||||
_ANTHROPIC_SSE_CHUNKS[1],
|
||||
_ANTHROPIC_SSE_CHUNKS[2],
|
||||
_ANTHROPIC_SSE_CHUNKS[3],
|
||||
b'event: content_block_start\ndata: {"type":"content_block_start","index":1,'
|
||||
b'"content_block":{"type":"text","text":""}}\n\n',
|
||||
b'event: content_block_delta\ndata: {"type":"content_block_delta","index":1,'
|
||||
b'"delta":{"type":"text_delta","text":" and my card is 4111111111111111"}}\n\n',
|
||||
b'event: content_block_stop\ndata: {"type":"content_block_stop","index":1}\n\n',
|
||||
_ANTHROPIC_SSE_CHUNKS[4],
|
||||
_ANTHROPIC_SSE_CHUNKS[5],
|
||||
)
|
||||
guardrail = _sse_armor_guardrail(mask_response_content=True)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(
|
||||
return_value=_armor_api_response(
|
||||
{
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"sdp": {
|
||||
"sdpFilterResult": {
|
||||
"deidentifyResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"data": {"text": "my ssn is [REDACTED] and my card is [REDACTED]"},
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
),
|
||||
):
|
||||
delivered = await _drain_armor_streaming_hook(guardrail, multi_block_chunks)
|
||||
|
||||
body = b"".join(chunk for chunk in delivered if isinstance(chunk, bytes))
|
||||
assert b"123-45-6789" not in body
|
||||
assert b"4111111111111111" not in body
|
||||
assert b"my ssn is [REDACTED] and my card is [REDACTED]" in body
|
||||
assert body.count(b'"type":"content_block_start"') == 2
|
||||
assert body.count(b'"type": "content_block_delta"') == 1
|
||||
assert body.count(b'"type":"content_block_stop"') == 2
|
||||
assert body.index(b'"index":1,"content_block"') > body.index(b'"index":0,"content_block"')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_passes_through_responses_api_events():
|
||||
"""/v1/responses streams deliver event objects stream_chunk_builder cannot assemble.
|
||||
|
||||
Regression for `500 Error building chunks for logging/streaming usage calculation`.
|
||||
"""
|
||||
from litellm.types.llms.openai import GenericEvent, ResponsesAPIStreamEvents
|
||||
|
||||
guardrail = _sse_armor_guardrail()
|
||||
events = (
|
||||
GenericEvent(type=ResponsesAPIStreamEvents.RESPONSE_CREATED),
|
||||
GenericEvent(type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED),
|
||||
)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock()) as mock_post:
|
||||
delivered = await _drain_armor_streaming_hook(guardrail, chunks=events)
|
||||
|
||||
mock_post.assert_not_called()
|
||||
assert tuple(delivered) == events
|
||||
|
||||
|
||||
async def _iter_chunks(chunks: tuple[object, ...]):
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
|
||||
async def _chunks_produced_before_first_delivery(guardrail: ModelArmorGuardrail) -> int:
|
||||
"""How much of the upstream stream is consumed before the client sees anything.
|
||||
|
||||
1 means the guardrail streams as it scans; the full chunk count means it buffers the
|
||||
whole response first, which is what makes time to first token equal generation time.
|
||||
"""
|
||||
produced: list[object] = []
|
||||
|
||||
async def _stream():
|
||||
for chunk in _ANTHROPIC_SSE_CHUNKS:
|
||||
produced.append(chunk)
|
||||
yield chunk
|
||||
|
||||
delivered = guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(request_route="/v1/messages"),
|
||||
response=_stream(),
|
||||
request_data={
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "what is my ssn"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]},
|
||||
},
|
||||
)
|
||||
try:
|
||||
await delivered.__anext__()
|
||||
finally:
|
||||
await delivered.aclose()
|
||||
return len(produced)
|
||||
|
||||
|
||||
def test_config_disables_streaming_buffer():
|
||||
from litellm.proxy.guardrails.guardrail_hooks.model_armor import initialize_guardrail
|
||||
from litellm.types.guardrails import LitellmParams
|
||||
|
||||
guardrail = initialize_guardrail(
|
||||
LitellmParams(
|
||||
guardrail="model_armor",
|
||||
mode="post_call",
|
||||
template_id="test-template",
|
||||
project_id="test-project",
|
||||
streaming_buffer_until_moderated=False,
|
||||
streaming_sampling_rate=2,
|
||||
),
|
||||
{"guardrail_name": "model-armor-test"},
|
||||
)
|
||||
|
||||
assert guardrail._streams_incrementally() is True
|
||||
assert guardrail.streaming_sampling_rate == 2
|
||||
|
||||
|
||||
def test_default_config_buffers_streams():
|
||||
from litellm.proxy.guardrails.guardrail_hooks.model_armor import initialize_guardrail
|
||||
from litellm.types.guardrails import LitellmParams
|
||||
|
||||
guardrail = initialize_guardrail(
|
||||
LitellmParams(
|
||||
guardrail="model_armor",
|
||||
mode="post_call",
|
||||
template_id="test-template",
|
||||
project_id="test-project",
|
||||
),
|
||||
{"guardrail_name": "model-armor-default"},
|
||||
)
|
||||
|
||||
assert guardrail._streams_incrementally() is False
|
||||
assert guardrail.streaming_sampling_rate == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_buffers_whole_response_by_default():
|
||||
guardrail = _sse_armor_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(return_value=_armor_api_response({"filterMatchState": "NO_MATCH_FOUND"})),
|
||||
):
|
||||
produced = await _chunks_produced_before_first_delivery(guardrail)
|
||||
|
||||
assert produced == len(_ANTHROPIC_SSE_CHUNKS)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_streams_while_scanning_when_buffering_disabled():
|
||||
guardrail = _sse_armor_guardrail(
|
||||
streaming_buffer_until_moderated=False,
|
||||
streaming_sampling_rate=1,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(return_value=_armor_api_response({"filterMatchState": "NO_MATCH_FOUND"})),
|
||||
):
|
||||
produced = await _chunks_produced_before_first_delivery(guardrail)
|
||||
|
||||
assert produced == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_keeps_buffering_when_masking_responses():
|
||||
"""Sanitization needs the assembled response, so masking configs cannot stream as they scan."""
|
||||
guardrail = _sse_armor_guardrail(
|
||||
mask_response_content=True,
|
||||
streaming_buffer_until_moderated=False,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(return_value=_armor_api_response({"filterMatchState": "NO_MATCH_FOUND"})),
|
||||
):
|
||||
produced = await _chunks_produced_before_first_delivery(guardrail)
|
||||
|
||||
assert produced == len(_ANTHROPIC_SSE_CHUNKS)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completions_streams_while_scanning_when_buffering_disabled():
|
||||
guardrail = _sse_armor_guardrail(
|
||||
streaming_buffer_until_moderated=False,
|
||||
streaming_sampling_rate=1,
|
||||
)
|
||||
produced: list[object] = []
|
||||
|
||||
async def _stream():
|
||||
for text in ("hello ", "world"):
|
||||
chunk = litellm.ModelResponseStream(
|
||||
choices=[
|
||||
litellm.types.utils.StreamingChoices(delta=litellm.types.utils.Delta(content=text))
|
||||
]
|
||||
)
|
||||
produced.append(chunk)
|
||||
yield chunk
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(return_value=_armor_api_response({"filterMatchState": "NO_MATCH_FOUND"})),
|
||||
):
|
||||
delivered = guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(request_route="/v1/chat/completions"),
|
||||
response=_stream(),
|
||||
request_data={
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]},
|
||||
},
|
||||
)
|
||||
first = await delivered.__anext__()
|
||||
assert len(produced) == 1
|
||||
assert first.choices[0].delta.content == "hello "
|
||||
await delivered.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_incremental_streaming_blocks_before_flagged_text_is_delivered():
|
||||
guardrail = _sse_armor_guardrail(
|
||||
streaming_buffer_until_moderated=False,
|
||||
streaming_sampling_rate=1,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
AsyncMock(
|
||||
return_value=_armor_api_response(
|
||||
{
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"sdp": {
|
||||
"sdpFilterResult": {
|
||||
"inspectResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"findings": [{"infoType": "US_SOCIAL_SECURITY_NUMBER"}],
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
),
|
||||
):
|
||||
delivered = [
|
||||
chunk
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(request_route="/v1/messages"),
|
||||
response=_iter_chunks(_ANTHROPIC_SSE_CHUNKS),
|
||||
request_data={
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "what is my ssn"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]},
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
body = b"".join(chunk if isinstance(chunk, bytes) else str(chunk).encode() for chunk in delivered)
|
||||
assert b"123-45-6789" not in body
|
||||
assert b"Model Armor" in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_fails_closed_on_unparseable_raw_sse():
|
||||
guardrail = _sse_armor_guardrail()
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", AsyncMock()) as mock_post:
|
||||
delivered = await _drain_armor_streaming_hook(
|
||||
guardrail, chunks=(b"data: not anthropic\n\n",)
|
||||
)
|
||||
|
||||
mock_post.assert_not_called()
|
||||
body = b"".join(chunk for chunk in delivered if isinstance(chunk, bytes))
|
||||
assert b"event: error" in body
|
||||
assert b"not anthropic" not in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_armor_api_failure_raises_sanitized_error():
|
||||
"""Test that Model Armor API failures raise HTTP 400, not the upstream status code."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue