diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index baa87055507..26d04d89c2b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 7ac8918a23e..4b274fba116 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -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"