From 429b2c3546a062fff6778220e05effc3c76b5a7f Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 25 Aug 2026 14:56:00 +0000 Subject: [PATCH] fix(model_armor): assemble raw Anthropic SSE streams in post-call streaming hook Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../model_armor/model_armor.py | 40 ++++- .../guardrail_hooks/test_model_armor.py | 156 ++++++++++++++++++ 2 files changed, 192 insertions(+), 4 deletions(-) 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 d187b5b12e9..419eb93f506 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -841,14 +841,25 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): 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_chunks_from_response, + anthropic_sse_error_frames, + assemble_anthropic_sse_stream, + is_raw_sse_stream, + ) # Collect all chunks all_chunks: Final[list[ModelResponseStream]] = [] async for chunk in response: all_chunks.append(chunk) - # Build complete response - assembled_response: Final = stream_chunk_builder(chunks=all_chunks) + # /v1/messages arrives as raw SSE frames, which stream_chunk_builder cannot assemble + raw_sse: Final = is_raw_sse_stream(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 isinstance(assembled_response, ModelResponse): # Extract content @@ -868,7 +879,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 +895,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 +915,10 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): choice.message.content = sanitized_content # Return sanitized stream + if raw_sse: + for sse_chunk in anthropic_sse_chunks_from_response(assembled_response): + yield sse_chunk + return mock_response: Final = MockResponseIterator(model_response=assembled_response) async for chunk in mock_response: yield chunk @@ -909,6 +926,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 +944,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 +956,13 @@ 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 # Return original chunks if no sanitization needed for chunk in all_chunks: 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 da66c36328e..a6ffc9503e2 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 @@ -623,6 +623,162 @@ 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}}}\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[bytes, ...] = _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 + + +@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."""