This commit is contained in:
devin-ai-integration[bot] 2026-08-27 16:33:33 -04:00 • committed by GitHub
commit 7992398d96
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 662 additions and 7 deletions

View file

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

View file

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

View file

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

View file

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

View file

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