diff --git a/litellm/proxy/guardrails/anthropic_sse.py b/litellm/proxy/guardrails/anthropic_sse.py index 50c05daee11..b505583cf3f 100644 --- a/litellm/proxy/guardrails/anthropic_sse.py +++ b/litellm/proxy/guardrails/anthropic_sse.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py index eda505e2453..7b9cf308ba6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py @@ -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) 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..fc67afcabcc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -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: diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py b/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py index ec7e1595215..5f0b59a6924 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py @@ -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""" 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..43ef96822c2 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,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."""