From 4b21583005bbda288caac14960bae67076a98b2d Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 1 Sep 2026 12:56:09 -0700 Subject: [PATCH] fix(model_armor): classify the stream surface and fail closed when it cannot be assembled Decide the wire format explicitly instead of inferring it from a boolean pair, so an opaque raw SSE stream (the Google :streamGenerateContent route) is never refused in Anthropic framing, and a stream that cannot be assembled is blocked rather than released unscanned unless fail_on_error is disabled. Also scan Responses tool-call arguments, read the body only off a terminal Responses event, and record the applied guardrail on the fail-closed path. --- litellm/proxy/guardrails/anthropic_sse.py | 59 ++- .../model_armor/model_armor.py | 375 ++++++++++-------- .../guardrail_hooks/test_model_armor.py | 198 +++++++-- 3 files changed, 430 insertions(+), 202 deletions(-) diff --git a/litellm/proxy/guardrails/anthropic_sse.py b/litellm/proxy/guardrails/anthropic_sse.py index 5c635ecfb49..28220f09f00 100644 --- a/litellm/proxy/guardrails/anthropic_sse.py +++ b/litellm/proxy/guardrails/anthropic_sse.py @@ -13,6 +13,19 @@ from typing import Final from litellm.types.utils import Choices, ModelResponse +_ANTHROPIC_EVENT_TYPES: Final = frozenset( + { + "message_start", + "message_delta", + "message_stop", + "content_block_start", + "content_block_delta", + "content_block_stop", + "ping", + "error", + } +) + def is_raw_sse_stream(all_chunks: Sequence[object]) -> bool: return any(isinstance(chunk, (str, bytes)) for chunk in all_chunks) @@ -30,23 +43,43 @@ def _joined_sse_stream(all_chunks: Sequence[object]) -> str | None: return None -def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None: +def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]: from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( AnthropicPassthroughLoggingHandler, ) + return tuple( + event_data + for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses + if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing + ) + + +def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None: return next( ( message - for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses - if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing - and event_data.get("type") == "message_start" - and isinstance(message := event_data.get("message"), dict) + for event_data in _parsed_sse_events(sse_stream) + if event_data.get("type") == "message_start" and isinstance(message := event_data.get("message"), dict) ), None, ) +def is_anthropic_sse_stream(all_chunks: Sequence[object]) -> bool: + """Whether raw SSE frames are Anthropic Messages events. + + ``is_raw_sse_stream`` only says the chunks are unparsed bytes, and ``/v1/messages`` is not the + only endpoint that streams those: the Google ``:streamGenerateContent`` route marks its own + stream raw too. Reading its frames as Anthropic ones would refuse the response in a wire format + its client cannot parse, so the surface is decided on the event types actually present. + """ + sse_stream: Final = _joined_sse_stream(all_chunks) + if sse_stream is None: + return False + return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in _parsed_sse_events(sse_stream)) + + def assemble_anthropic_sse_stream( all_chunks: Sequence[object], *, restore_identity: bool = False ) -> ModelResponse | None: @@ -119,20 +152,18 @@ def is_sse_error_stream(all_chunks: Sequence[object]) -> bool: them would hide the refusal the client is owed. Covers both wire forms a guardrail emits: the Anthropic ``error`` event and the chat-completions ``{"error": ...}`` payload. """ + if not all(isinstance(chunk, (str, bytes)) for chunk in all_chunks): + # A stream mixing typed chunks with an error frame still carries content to scan, and the + # frames-only join below would drop exactly the part that has to be scanned + return False sse_stream: Final = _joined_sse_stream(all_chunks) if sse_stream is None: return False - from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( - AnthropicPassthroughLoggingHandler, + events: Final = _parsed_sse_events(sse_stream) + return len(events) > 0 and all( + event.get("type") == "error" or isinstance(event.get("error"), Mapping) for event in events ) - events: Final = tuple( - event_data - for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses - if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses - ) - return len(events) > 0 and all(event.get("type") == "error" or "error" in event for event in events) - def anthropic_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]: from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( 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 14df0b71e1d..baa87055507 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -1,4 +1,5 @@ from collections.abc import AsyncGenerator, Mapping, Sequence +from enum import Enum, auto from typing import TYPE_CHECKING, Any, Final, Literal import httpx @@ -25,15 +26,13 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.llms.openai.responses.guardrail_translation.handler import ( - OpenAIResponsesHandler, -) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.anthropic_sse import ( anthropic_sse_chunks_from_response, anthropic_sse_error_frames, assemble_anthropic_sse_stream, + is_anthropic_sse_stream, is_raw_sse_stream, is_sse_error_stream, ) @@ -42,7 +41,11 @@ from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import ( plan_file_scans, ) from litellm.types.guardrails import GuardrailEventHooks, LitellmParams -from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionToolCallChunk, + ResponsesAPIResponse, +) from litellm.types.utils import ( CallTypes, CallTypesLiteral, @@ -56,6 +59,18 @@ from litellm.types.utils import ( GUARDRAIL_NAME: Final = "model_armor" +# Only these carry the finished output; response.created carries an empty body +_RESPONSES_TERMINAL_EVENT_TYPES: Final = frozenset({"response.completed", "response.incomplete", "response.failed"}) + + +class _StreamSurface(Enum): + """Wire format of a buffered streaming response, which decides how it is read and how it is refused.""" + + CHAT_COMPLETIONS = auto() + ANTHROPIC_MESSAGES = auto() + RESPONSES = auto() + OPAQUE_SSE = auto() + class ModelArmorAPIError(Exception): """Model Armor API failure (non-2xx), distinct from a content-block decision so @@ -843,42 +858,73 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): return response @staticmethod - def _final_responses_api_response( - all_chunks: Sequence[object], - ) -> tuple[bool, ResponsesAPIResponse | None]: - """Detect a ``/v1/responses`` event stream and return its final response body. + def _is_terminal_error_stream(all_chunks: Sequence[object]) -> bool: + """Whether the buffered stream is only the refusal an earlier guardrail in the chain emitted. - Returns ``(is_responses_stream, final_response)`` so the caller can tell a stream - that is not a Responses stream apart from one whose terminal ``response.completed`` - event never arrived. + post_call guardrails are composed, so this hook can be handed the terminal error items a + preceding one produced. They carry no message to scan, and replacing them would hide the + refusal the client is owed. """ - events: Final = tuple( - (event_type, getattr(chunk, "response", None)) + if all(getattr(chunk, "type", None) == "error" for chunk in all_chunks): + return True + return is_sse_error_stream(all_chunks) + + @staticmethod + def _classify_stream(all_chunks: Sequence[object]) -> _StreamSurface: + """Wire format the buffered chunks belong to.""" + if is_raw_sse_stream(all_chunks): + return ( + _StreamSurface.ANTHROPIC_MESSAGES if is_anthropic_sse_stream(all_chunks) else _StreamSurface.OPAQUE_SSE + ) + if any( + isinstance(event_type := getattr(chunk, "type", None), str) and event_type.startswith("response.") for chunk in all_chunks - if isinstance(event_type := getattr(chunk, "type", None), str) - ) - return ( - any(event_type.startswith("response.") for event_type, _ in events), - next( - (body for _, body in reversed(events) if isinstance(body, ResponsesAPIResponse)), - None, + ): + return _StreamSurface.RESPONSES + return _StreamSurface.CHAT_COMPLETIONS + + @staticmethod + def _final_responses_api_response(all_chunks: Sequence[object]) -> ResponsesAPIResponse | None: + """Response body carried by a terminal ``/v1/responses`` event. + + A stream cut short before it completes has to read as unassembled rather than as a clean + empty response: ``response.created`` also carries a body, but an empty one, and scanning + that would release every buffered delta unscanned. + """ + return next( + ( + body + for chunk in reversed(all_chunks) + if getattr(chunk, "type", None) in _RESPONSES_TERMINAL_EVENT_TYPES + and isinstance(body := getattr(chunk, "response", None), ResponsesAPIResponse) ), + None, ) @staticmethod def _responses_api_response_text(response: ResponsesAPIResponse) -> str: - """Concatenate the output text carried by a Responses API response.""" + """Text to scan in a Responses API response, tool-call arguments included. + + Tool calls are folded in because ``get_content_from_model_response`` folds them into what + the chat surface scans, and a Responses turn can carry its whole payload in them. + """ + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + texts: Final[list[str]] = [] # mutable-ok: the shared extractor below appends into caller-owned lists + tool_calls: Final[list[ChatCompletionToolCallChunk]] = [] # mutable-ok: the same extractor's tool-call sink handler: Final = OpenAIResponsesHandler() for output_idx, output_item in enumerate(response.output or ()): handler._extract_output_text_and_images( # pyright: ignore[reportPrivateUsage] # the shared Responses output extractor; forking it would duplicate per-item parsing - output_item, - output_idx, - texts, - [], # mutable-ok: the extractor's images sink, unused here - [], # mutable-ok: the extractor's task-mapping sink, unused here + output_item=output_item, + output_idx=output_idx, + texts_to_check=texts, + images_to_check=[], # mutable-ok: the extractor's images sink, unused here + task_mappings=[], # mutable-ok: the extractor's task-mapping sink, unused here + tool_calls_to_check=tool_calls, ) - return "".join(texts) + return "".join((*texts, *(json.dumps(tool_call) for tool_call in tool_calls))) def _extract_streaming_content(self, assembled_response: object) -> str: """Text to scan from an assembled stream, for every endpoint shape this hook serves.""" @@ -897,35 +943,56 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): def _assemble_chat_completion_stream( all_chunks: list[Any], # mutable-ok: stream_chunk_builder only accepts a mutable list ) -> ModelResponse | TextCompletionResponse | None: - """Assemble chat-completion chunks, returning ``None`` when they are not chat deltas.""" + """Assemble chat-completion chunks, returning ``None`` when they cannot be assembled.""" from litellm.main import stream_chunk_builder try: return stream_chunk_builder(chunks=all_chunks) except Exception as exc: - verbose_proxy_logger.warning( - "Model Armor: could not assemble the streamed response for scanning (%s), forwarding it unscanned", - exc, - ) + verbose_proxy_logger.warning("Model Armor: chat-completion stream assembly failed (%s)", exc) return None + def _assemble_stream( + self, all_chunks: Sequence[object], surface: _StreamSurface + ) -> ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None: + """Assemble the buffered stream into the scannable response its surface produces.""" + if surface is _StreamSurface.ANTHROPIC_MESSAGES: + return assemble_anthropic_sse_stream(all_chunks, restore_identity=True) + if surface is _StreamSurface.RESPONSES: + return self._final_responses_api_response(all_chunks) + if surface is _StreamSurface.OPAQUE_SSE: + return None + return self._assemble_chat_completion_stream(list(all_chunks)) + @staticmethod - def _stream_error_items( - exc: HTTPException, - error_obj: Mapping[str, object], - *, - raw_sse: bool, - responses_stream: bool, - ) -> Sequence[object]: + def _error_payload(exc: HTTPException) -> Mapping[str, object]: + """Error object for a terminal stream item, carrying the status the frame would otherwise lose.""" + detail: Final = exc.detail if isinstance(exc.detail, Mapping) else {"message": str(exc.detail)} + error_value: Final = detail.get("error", detail) + return { + **(dict(error_value) if isinstance(error_value, Mapping) else {"message": str(error_value)}), + "code": str(exc.status_code), + } + + @staticmethod + def _build_responses_error_items(exc: HTTPException) -> Sequence[object] | None: + """Responses API error events for a failure discovered after the stream started.""" + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + return OpenAIResponsesHandler().build_stream_error_items(exc, responses_so_far=None) + + def _stream_error_items(self, exc: HTTPException, *, surface: _StreamSurface) -> Sequence[object]: """Frame a guardrail failure as terminal stream items in this endpoint's wire format.""" - chat_completions_form: Final = (f"data: {json.dumps({'error': error_obj})}\n\n",) - if raw_sse: - return anthropic_sse_error_frames(str(error_obj.get("message", ""))) - if responses_stream: - return ( - OpenAIResponsesHandler().build_stream_error_items(exc, responses_so_far=None) or chat_completions_form - ) - return chat_completions_form + payload: Final = self._error_payload(exc) + if surface is _StreamSurface.ANTHROPIC_MESSAGES: + return anthropic_sse_error_frames(str(payload.get("message", ""))) + if surface is _StreamSurface.RESPONSES and (responses_items := self._build_responses_error_items(exc)): + return responses_items + # Also the fallback when a surface cannot frame its own error: create_response() reads the + # status back out of this form, so the refusal keeps its code instead of arriving as a 200 + return (f"data: {json.dumps({'error': payload})}\n\n",) async def async_post_call_streaming_iterator_hook( self, @@ -936,78 +1003,93 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): """Process streaming response chunks.""" from litellm.llms.base_llm.base_model_iterator import MockResponseIterator + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) # Collect all chunks all_chunks: Final[list[Any]] = [] async for chunk in response: all_chunks.append(chunk) - raw_sse: Final = is_raw_sse_stream(all_chunks) - responses_stream, final_responses_api_response = ( - (False, None) if raw_sse else self._final_responses_api_response(all_chunks) - ) + if not all_chunks or self._is_terminal_error_stream(all_chunks): + for chunk in all_chunks: + yield chunk + return + + surface: Final = self._classify_stream(all_chunks) # Build complete response - assembled_response: Final = ( - assemble_anthropic_sse_stream(all_chunks, restore_identity=True) - if raw_sse - else final_responses_api_response - if responses_stream - else self._assemble_chat_completion_stream(all_chunks) - ) + assembled_response: Final = self._assemble_stream(all_chunks, surface) + + if assembled_response is None: + if not self.optional_params.get("fail_on_error", True): + verbose_proxy_logger.warning( + "Model Armor: streamed response could not be assembled for scanning, " + "forwarding it unscanned because fail_on_error is disabled" + ) + for chunk in all_chunks: + yield chunk + return - if ( - assembled_response is None - and (raw_sse or responses_stream) - and not (raw_sse and is_sse_error_stream(all_chunks)) - ): # Forwarding an unscannable stream would silently disable the guardrail, so fail closed - unscannable: Final = HTTPException( - status_code=500, - detail=f"{self.guardrail_name}: streamed response could not be assembled for scanning, blocking it", - ) + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) for error_item in self._stream_error_items( - unscannable, - {"message": str(unscannable.detail), "code": "500"}, - raw_sse=raw_sse, - responses_stream=responses_stream, + HTTPException( + status_code=500, + detail=f"{self.guardrail_name}: streamed response could not be assembled for scanning, blocking it", + ), + surface=surface, ): yield error_item return - if assembled_response is not None: - # Extract content - content: Final = self._extract_streaming_content(assembled_response) + # Extract content + content: Final = self._extract_streaming_content(assembled_response) - if content: - try: - # Check with Model Armor - armor_response: Final = await self.make_model_armor_request( - content=content, - source="model_response", - request_data=request_data, - ) + if not content: + verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail") + for chunk in all_chunks: + yield chunk + return - # Attach Model Armor response & status to this request's metadata to avoid race conditions - 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" + try: + # Check with Model Armor + armor_response: Final = await self.make_model_armor_request( + content=content, + source="model_response", + request_data=request_data, + ) + + # Attach Model Armor response & status to this request's metadata to avoid race conditions + 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" + + # 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): + raise HTTPException( + status_code=400, + detail=self._build_block_error_detail( + "Streaming response blocked by Model Armor", + armor_response, + ), + ) + + # Apply sanitization if enabled + if self.mask_response_content: + sanitized_content: Final = self._get_sanitized_content(armor_response) + if sanitized_content and sanitized_content != content: + if not isinstance(assembled_response, ModelResponse): + verbose_proxy_logger.warning( + "Model Armor: sanitized content cannot be re-emitted on this " + "streaming endpoint, blocking the response instead" ) - - # Add guardrail to applied_guardrails BEFORE potential blocking - # This ensures guardrail is recorded even when it blocks the request - from litellm.proxy.common_utils.callback_utils import ( - add_guardrail_to_applied_guardrails_header, - ) - - 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): raise HTTPException( status_code=400, detail=self._build_block_error_detail( @@ -1016,72 +1098,37 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): ), ) - # Apply sanitization if enabled - if self.mask_response_content: - sanitized_content: Final = self._get_sanitized_content(armor_response) - if sanitized_content and sanitized_content != content: - if not isinstance(assembled_response, ModelResponse): - verbose_proxy_logger.warning( - "Model Armor: sanitized content cannot be re-emitted on this " - "streaming endpoint, blocking the response instead" - ) - raise HTTPException( - status_code=400, - detail=self._build_block_error_detail( - "Streaming response blocked by Model Armor", - armor_response, - ), - ) + # Update assembled response + self._apply_sanitized_content(assembled_response, sanitized_content) - # Update assembled response - self._apply_sanitized_content(assembled_response, 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 - return - - except ModelArmorAPIError as e: - if self.optional_params.get("fail_on_error", True): - error_obj = {"message": e.detail, "code": "500"} - for error_item in self._stream_error_items( - HTTPException(status_code=500, detail=e.detail), - error_obj, - raw_sse=raw_sse, - responses_stream=responses_stream, - ): - yield error_item + # Return sanitized stream + if surface is _StreamSurface.ANTHROPIC_MESSAGES: + for sse_chunk in anthropic_sse_chunks_from_response(assembled_response): + yield sse_chunk return - except HTTPException as e: - # Yield the error as a terminal stream item so create_response() detects - # it and returns a proper JSON error response with the correct status code. - # (Raising from a generator hits create_response's generic except → 500.) - detail: Final = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)} - error_value: Final = detail.get("error", detail) - if isinstance(error_value, dict): - error_obj = dict(error_value) - else: - error_obj = {"message": str(error_value)} - error_obj["code"] = str(e.status_code) - for error_item in self._stream_error_items( - e, - error_obj, - raw_sse=raw_sse, - responses_stream=responses_stream, - ): - yield error_item + mock_response: Final = MockResponseIterator(model_response=assembled_response) + async for chunk in mock_response: + yield chunk return - except Exception as e: - verbose_proxy_logger.error("Model Armor streaming error: %s", str(e), exc_info=True) - if self.optional_params.get("fail_on_error", True): - raise - else: - verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail") + + except ModelArmorAPIError as e: + if self.optional_params.get("fail_on_error", True): + for error_item in self._stream_error_items( + HTTPException(status_code=500, detail=e.detail), surface=surface + ): + yield error_item + return + except HTTPException as e: + # Yield the error as a terminal stream item so create_response() detects it and returns + # a proper JSON error response with the correct status code. Raising from a generator + # instead hits create_response's generic except and becomes a 500. + for error_item in self._stream_error_items(e, surface=surface): + yield error_item + return + except Exception as e: + verbose_proxy_logger.error("Model Armor streaming error: %s", str(e), exc_info=True) + if self.optional_params.get("fail_on_error", True): + raise # 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 2dcd594bb18..eb76deb935c 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 @@ -15,9 +15,7 @@ import litellm.types.utils from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError -from litellm.llms.openai.responses.guardrail_translation.handler import ( - OpenAIResponsesHandler, -) +from litellm.proxy.guardrails.anthropic_sse import anthropic_sse_error_frames from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.model_armor import ModelArmorGuardrail from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import ( @@ -3892,7 +3890,7 @@ def _responses_api_events(): ) -async def _drain_surface_hook(guardrail, chunks): +async def _drain_surface_hook(guardrail, chunks, request_data=None): async def _stream(): for chunk in chunks: yield chunk @@ -3902,7 +3900,9 @@ async def _drain_surface_hook(guardrail, chunks): async for item in guardrail.async_post_call_streaming_iterator_hook( user_api_key_dict=UserAPIKeyAuth(), response=_stream(), - request_data={ + request_data=request_data + if request_data is not None + else { "model": "claude-haiku", "messages": [{"role": "user", "content": "show me a card"}], "metadata": {"guardrails": ["model-armor-test"]}, @@ -4059,6 +4059,7 @@ async def test_streaming_api_failure_frames_error_per_surface(surface): id="anthropic-sse-without-message-start", ), pytest.param(None, id="responses-stream-without-completed-event"), + pytest.param("created", id="responses-stream-cut-off-after-response-created"), ], ) async def test_streaming_hook_fails_closed_when_a_surface_stream_cannot_be_assembled(chunks): @@ -4070,16 +4071,17 @@ async def test_streaming_hook_fails_closed_when_a_surface_stream_cannot_be_assem ResponsesAPIStreamEvents, ) - if chunks is None: - chunks = ( - OutputTextDeltaEvent( - type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, - item_id="msg_1", - output_index=0, - content_index=0, - delta="hi", - ), + if chunks is None or chunks == "created": + delta = OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta="my card is 4111-1111-1111-1111", ) + # response.created carries a ResponsesAPIResponse too, but an empty one: reading the body + # off it would scan "" and release every buffered delta unscanned + chunks = (delta,) if chunks is None else (_responses_created_event(), delta) guardrail = _surface_guardrail() post = _armor_post_mock(_MODEL_ARMOR_CLEAN) @@ -4145,8 +4147,6 @@ async def test_streaming_hook_forwards_a_preceding_guardrails_error_item(): async def test_streaming_hook_forwards_a_preceding_guardrails_error_frame(chunks): """Chained post_call guardrails hand each other their output. An earlier guardrail's error frame carries no message to assemble, and replacing it would hide the real refusal.""" - from litellm.proxy.guardrails.anthropic_sse import anthropic_sse_error_frames - if chunks is None: chunks = anthropic_sse_error_frames("Streaming response blocked by the first guardrail") guardrail = _surface_guardrail() @@ -4163,17 +4163,167 @@ async def test_streaming_hook_forwards_a_preceding_guardrails_error_frame(chunks async def test_streaming_responses_error_falls_back_to_sse_when_the_handler_declines(): """build_stream_error_items may return None, which must not swallow the block into a clean 200: the refusal falls back to the chat-completions SSE form that still carries the status.""" - guardrail = _surface_guardrail() + from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import ( + _StreamSurface, + ) + + class _DecliningGuardrail(ModelArmorGuardrail): + @staticmethod + def _build_responses_error_items(exc): + return None + + guardrail = _DecliningGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) exc = HTTPException(status_code=400, detail={"message": "blocked"}) - with patch.object(OpenAIResponsesHandler, "build_stream_error_items", return_value=None): - items = guardrail._stream_error_items( - exc, - {"message": "blocked", "code": "400"}, - raw_sse=False, - responses_stream=True, - ) + items = guardrail._stream_error_items(exc, surface=_StreamSurface.RESPONSES) assert len(items) == 1 assert '"code": "400"' in items[0] assert "blocked" in items[0] + + +def _responses_created_event(): + from litellm.types.llms.openai import ( + ResponseCreatedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + return ResponseCreatedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_CREATED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o-mini", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + + +@pytest.mark.asyncio +async def test_streaming_hook_refuses_an_opaque_sse_stream_without_anthropic_framing(): + """/v1/messages is not the only endpoint that streams raw SSE: the Google generateContent + route marks its own stream raw too. Its frames carry no Anthropic event types, so refusing + them in Anthropic's format would hand a Google client a body it cannot parse.""" + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + chunks = (b'data: {"candidates":[{"content":{"parts":[{"text":"my card is 4111"}]}}]}\n\n',) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + assert tuple(delivered) != chunks + body = "".join(item.decode() if isinstance(item, bytes) else item for item in delivered) + assert "could not be assembled for scanning" in body + assert "event: error" not in body + assert '"code": "500"' in body + + +@pytest.mark.asyncio +async def test_streaming_unassemblable_stream_is_forwarded_when_fail_on_error_is_disabled(): + """fail_on_error: false is a deliberate choice to degrade open, and it governs every other + path in this hook. The fail-closed refusal has to honour it too.""" + guardrail = _surface_guardrail(fail_on_error=False) + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + chunks = ( + b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,' + b'"delta":{"type":"text_delta","text":"hi"}}\n\n', + ) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + assert tuple(delivered) == chunks + + +@pytest.mark.asyncio +async def test_streaming_fail_closed_records_the_applied_guardrail(): + """A refusal that no header or log attributes to the guardrail leaves on-call unable to tell + a guardrail block apart from a provider failure.""" + guardrail = _surface_guardrail() + request_data = { + "model": "claude-haiku", + "messages": [{"role": "user", "content": "show me a card"}], + "metadata": {"guardrails": ["model-armor-test"]}, + } + chunks = ( + b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,' + b'"delta":{"type":"text_delta","text":"hi"}}\n\n', + ) + + with patch.object(guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_CLEAN)): + await _drain_surface_hook(guardrail, chunks, request_data=request_data) + + assert request_data["metadata"]["applied_guardrails"] == ["model-armor-test"] + + +@pytest.mark.asyncio +async def test_streaming_responses_tool_call_output_is_scanned(): + """An agentic /v1/responses turn can carry its whole payload in tool-call arguments, which + is what the chat surface already folds into the scanned text.""" + from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + completed = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o-mini", + object="response", + output=[ + { + "type": "function_call", + "id": "fc_1", + "call_id": "call_1", + "name": "send_email", + "arguments": '{"body": "my card is 4111-1111-1111-1111"}', + } + ], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, (completed,)) + + post.assert_called_once() + scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"] + assert "4111-1111-1111-1111" in scanned + assert tuple(delivered) == (completed,) + + +@pytest.mark.asyncio +async def test_streaming_hook_refuses_a_content_stream_that_ends_with_an_error_frame(): + """The chain-aware passthrough must stay narrow. A stream carrying real content plus a + trailing error frame is not a bare refusal to forward: the assembler cannot read it, and + releasing it would ship the buffered content unscanned.""" + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + chunks = (*_ANTHROPIC_SSE_CHUNKS, *anthropic_sse_error_frames("upstream gave up")) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + body = b"".join(delivered) + assert b"4111-1111-1111-1111" not in body + assert b"could not be assembled for scanning" in body