fix(model_armor): let a streamed de-identify match mask instead of blocking

A de-identify template reports MATCH_FOUND for every redaction it makes. The
streaming block check omitted allow_sanitization, so with mask_response_content
enabled that match read as a refusal and the client got a 400 where the
non-streaming sibling returned the redacted text. Pass the flag through, as the
non-streaming hook already does, and stamp the logged status from the same
decision so the spend row agrees with what the client received.

Also drop Any from the chat-completion assembler's parameter; stream_chunk_builder
takes a bare list, so list[object] carries the mutability requirement without
erasing the element type.
This commit is contained in:
Yucheng Zhu 2026-09-01 13:43:05 -07:00
parent 8bd62b4366
commit c3e7224ee2
2 changed files with 119 additions and 7 deletions

View file

@ -941,7 +941,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
@staticmethod
def _assemble_chat_completion_stream(
all_chunks: list[Any], # mutable-ok: stream_chunk_builder only accepts a mutable list
all_chunks: list[object], # mutable-ok: stream_chunk_builder only accepts a mutable list
) -> ModelResponse | TextCompletionResponse | None:
"""Assemble chat-completion chunks, returning ``None`` when they cannot be assembled."""
from litellm.main import stream_chunk_builder
@ -1065,14 +1065,19 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
if isinstance(request_data, dict):
_, 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"
metadata["_model_armor_status"] = (
"blocked"
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content)
else "success"
)
# Add guardrail to applied_guardrails BEFORE potential blocking
# This ensures guardrail is recorded even when it blocks the request
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
# Check if blocked
if self._should_block_content(armor_response):
# Check if blocked. Mirrors the non-streaming sibling: with masking on, a de-identify
# match is a redaction to apply below, not a refusal
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content):
raise HTTPException(
status_code=400,
detail=self._build_block_error_detail(

View file

@ -3814,14 +3814,54 @@ _MODEL_ARMOR_BLOCK = {
}
}
# The streaming hook checks _should_block_content without allow_sanitization, so a
# deidentifyResult MATCH_FOUND blocks rather than masks. The root-level sanitizedText
# fallback in _get_sanitized_content is the shape that reaches the masking branch.
# The root-level sanitizedText fallback in _get_sanitized_content, i.e. a rewrite that trips no
# named filter
_MODEL_ARMOR_SANITIZED = {
"sanitizedText": "my card is [REDACTED]",
"sanitizationResult": {"filterMatchState": "NO_MATCH_FOUND"},
}
# The shape a real de-identify template returns: the SDP filter both matches and hands back the
# rewritten text, so whether it blocks or masks is decided by allow_sanitization alone
_MODEL_ARMOR_DEIDENTIFIED = {
"sanitizationResult": {
"filterMatchState": "MATCH_FOUND",
"filterResults": {
"sdp": {
"sdpFilterResult": {
"deidentifyResult": {
"matchState": "MATCH_FOUND",
"data": {"text": "my card is [REDACTED]"},
}
}
}
},
}
}
def _chat_completion_chunks():
"""The chat-completions surface: typed ModelResponseStream chunks."""
return (
litellm.types.utils.ModelResponseStream(
choices=[
litellm.types.utils.StreamingChoices(
index=0,
delta=litellm.types.utils.Delta(content="my card is 4111-1111-1111-1111"),
)
]
),
litellm.types.utils.ModelResponseStream(
choices=[
litellm.types.utils.StreamingChoices(
index=0,
delta=litellm.types.utils.Delta(content=""),
finish_reason="stop",
)
]
),
)
def _surface_guardrail(**kwargs):
guardrail = ModelArmorGuardrail(
@ -4391,3 +4431,70 @@ async def test_streaming_hook_does_not_forward_typed_chunks_that_end_with_an_err
body = b"".join(item if isinstance(item, bytes) else str(item).encode() for item in delivered)
assert b"4111-1111-1111-1111" not in body
assert b"could not be assembled for scanning" in body
def _delivered_bytes(delivered):
return b"".join(
item
if isinstance(item, bytes)
else item.encode()
if isinstance(item, str)
else str(item.model_dump() if hasattr(item, "model_dump") else item).encode()
for item in delivered
)
@pytest.mark.asyncio
@pytest.mark.parametrize("chunks, case", [(None, "chat_completions"), (_ANTHROPIC_SSE_CHUNKS, "anthropic_sse")])
async def test_streaming_deidentify_match_masks_when_masking_is_enabled(chunks, case):
"""A de-identify template reports MATCH_FOUND for every redaction it makes, so reading that
match as a refusal makes mask_response_content unusable on a stream: the client gets an error
where its non-streaming sibling gets redacted text. The block check has to allow sanitization
exactly as the non-streaming hook does."""
guardrail = _surface_guardrail(mask_response_content=True)
post = _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)
with patch.object(guardrail.async_handler, "post", post):
delivered = await _drain_surface_hook(
guardrail, _chat_completion_chunks() if chunks is None else chunks
)
body = _delivered_bytes(delivered)
assert b"[REDACTED]" in body, case
assert b"4111-1111-1111-1111" not in body, case
assert b"blocked by Model Armor" not in body, case
@pytest.mark.asyncio
@pytest.mark.parametrize("chunks, case", [(None, "chat_completions"), (_ANTHROPIC_SSE_CHUNKS, "anthropic_sse")])
async def test_streaming_deidentify_match_still_blocks_when_masking_is_disabled(chunks, case):
"""Without mask_response_content there is nowhere to put the rewritten text, so the same
de-identify match must still end the stream rather than release the original."""
guardrail = _surface_guardrail()
post = _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)
with patch.object(guardrail.async_handler, "post", post):
delivered = await _drain_surface_hook(
guardrail, _chat_completion_chunks() if chunks is None else chunks
)
body = _delivered_bytes(delivered)
assert b"Streaming response blocked by Model Armor" in body, case
assert b"4111-1111-1111-1111" not in body, case
@pytest.mark.asyncio
async def test_streaming_deidentify_match_logs_masked_run_as_success_not_blocked():
"""The status stamped on request metadata feeds the spend log, so it has to agree with what
the client actually received: a masked stream is a success, not a block."""
guardrail = _surface_guardrail(mask_response_content=True)
request_data = {
"model": "claude-haiku",
"messages": [{"role": "user", "content": "show me a card"}],
"metadata": {"guardrails": ["model-armor-test"]},
}
with patch.object(guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)):
await _drain_surface_hook(guardrail, _chat_completion_chunks(), request_data=request_data)
assert request_data["metadata"]["_model_armor_status"] == "success"