mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(model_armor): handle Anthropic Messages and Responses streams in post_call
The post_call streaming hook buffered every chunk and fed it to stream_chunk_builder, which only understands chat-completion deltas. /v1/messages streams raw Anthropic SSE bytes and /v1/responses streams typed Responses events, so both raised litellm.APIError and surfaced to the client as a 500 on every streamed request. Assemble each surface with its own reader, frame guardrail failures as terminal items in that surface's wire format, and pass the stream through unscanned when it cannot be assembled instead of raising.
This commit is contained in:
parent
ec3f8183c3
commit
67a9932f7d
3 changed files with 588 additions and 15 deletions
|
|
@ -111,6 +111,29 @@ def anthropic_sse_error_frames(message: str) -> tuple[bytes, ...]:
|
|||
)
|
||||
|
||||
|
||||
def is_sse_error_stream(all_chunks: Sequence[object]) -> bool:
|
||||
"""Whether the buffered stream carries nothing but error frames.
|
||||
|
||||
post_call guardrails run in a chain, so a hook can be handed the terminal error frames an
|
||||
earlier guardrail emitted when it blocked. Those carry no message to assemble, and replacing
|
||||
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.
|
||||
"""
|
||||
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 = 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 (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
|
|
|
|||
|
|
@ -25,14 +25,24 @@ 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_raw_sse_stream,
|
||||
is_sse_error_stream,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
|
||||
MODEL_ARMOR_MAX_FILE_SIZE_BYTES,
|
||||
plan_file_scans,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
CallTypesLiteral,
|
||||
|
|
@ -41,6 +51,7 @@ from litellm.types.utils import (
|
|||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StandardLoggingGuardrailInformation,
|
||||
TextCompletionResponse,
|
||||
)
|
||||
|
||||
GUARDRAIL_NAME: Final = "model_armor"
|
||||
|
|
@ -831,6 +842,91 @@ 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.
|
||||
|
||||
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.
|
||||
"""
|
||||
events: Final = tuple(
|
||||
(event_type, getattr(chunk, "response", None))
|
||||
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,
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _responses_api_response_text(response: ResponsesAPIResponse) -> str:
|
||||
"""Concatenate the output text carried by a Responses API response."""
|
||||
texts: Final[list[str]] = [] # mutable-ok: the shared extractor below appends into caller-owned lists
|
||||
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
|
||||
)
|
||||
return "".join(texts)
|
||||
|
||||
def _extract_streaming_content(self, assembled_response: object) -> str:
|
||||
"""Text to scan from an assembled stream, for every endpoint shape this hook serves."""
|
||||
if isinstance(assembled_response, ResponsesAPIResponse):
|
||||
return self._responses_api_response_text(assembled_response)
|
||||
return self._extract_content_from_response(assembled_response)
|
||||
|
||||
@staticmethod
|
||||
def _apply_sanitized_content(assembled_response: ModelResponse, sanitized_content: str) -> None:
|
||||
"""Replace every non-empty choice message with the Model Armor sanitized text."""
|
||||
for choice in assembled_response.choices:
|
||||
if isinstance(choice, Choices) and choice.message.content:
|
||||
choice.message.content = sanitized_content
|
||||
|
||||
@staticmethod
|
||||
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."""
|
||||
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,
|
||||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _stream_error_items(
|
||||
exc: HTTPException,
|
||||
error_obj: Mapping[str, object],
|
||||
*,
|
||||
raw_sse: bool,
|
||||
responses_stream: bool,
|
||||
) -> 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
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -840,19 +936,48 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
"""Process streaming response chunks."""
|
||||
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
|
||||
# Collect all chunks
|
||||
all_chunks: Final[list[ModelResponseStream]] = []
|
||||
all_chunks: Final[list[Any]] = []
|
||||
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)
|
||||
responses_stream, final_responses_api_response = (
|
||||
(False, None) if raw_sse else self._final_responses_api_response(all_chunks)
|
||||
)
|
||||
|
||||
if isinstance(assembled_response, ModelResponse):
|
||||
# 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)
|
||||
)
|
||||
|
||||
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",
|
||||
)
|
||||
for error_item in self._stream_error_items(
|
||||
unscannable,
|
||||
{"message": str(unscannable.detail), "code": "500"},
|
||||
raw_sse=raw_sse,
|
||||
responses_stream=responses_stream,
|
||||
):
|
||||
yield error_item
|
||||
return
|
||||
|
||||
if assembled_response is not None:
|
||||
# Extract content
|
||||
content: Final = self._extract_content_from_response(assembled_response)
|
||||
content: Final = self._extract_streaming_content(assembled_response)
|
||||
|
||||
if content:
|
||||
try:
|
||||
|
|
@ -895,13 +1020,27 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
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
|
||||
for choice in assembled_response.choices:
|
||||
if isinstance(choice, Choices):
|
||||
if choice.message.content:
|
||||
choice.message.content = sanitized_content
|
||||
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
|
||||
|
|
@ -910,11 +1049,17 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
except ModelArmorAPIError as e:
|
||||
if self.optional_params.get("fail_on_error", True):
|
||||
error_obj = {"message": e.detail, "code": "500"}
|
||||
yield f"data: {json.dumps({'error': error_obj})}\n\n"
|
||||
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
|
||||
except HTTPException as e:
|
||||
# Yield error as SSE event so create_response() detects it and
|
||||
# returns a proper JSON error response with the correct status code.
|
||||
# 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)
|
||||
|
|
@ -923,7 +1068,13 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
else:
|
||||
error_obj = {"message": str(error_value)}
|
||||
error_obj["code"] = str(e.status_code)
|
||||
yield f"data: {json.dumps({'error': error_obj})}\n\n"
|
||||
for error_item in self._stream_error_items(
|
||||
e,
|
||||
error_obj,
|
||||
raw_sse=raw_sse,
|
||||
responses_stream=responses_stream,
|
||||
):
|
||||
yield error_item
|
||||
return
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Model Armor streaming error: %s", str(e), exc_info=True)
|
||||
|
|
|
|||
|
|
@ -15,6 +15,9 @@ 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._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.model_armor import ModelArmorGuardrail
|
||||
from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import (
|
||||
|
|
@ -3778,3 +3781,399 @@ async def test_moderation_hook_skips_chat_traffic_when_configured_for_during_mcp
|
|||
|
||||
assert result == data
|
||||
mock_post.assert_not_called()
|
||||
|
||||
|
||||
_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 card is 4111-1111-1111-1111"}}\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',
|
||||
)
|
||||
|
||||
_MODEL_ARMOR_CLEAN = {"sanitizationResult": {"filterMatchState": "NO_MATCH_FOUND"}}
|
||||
|
||||
_MODEL_ARMOR_BLOCK = {
|
||||
"sanitizationResult": {
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"sdp": {
|
||||
"sdpFilterResult": {
|
||||
"inspectResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"findings": [
|
||||
{"infoType": "CREDIT_CARD_NUMBER", "likelihood": "VERY_LIKELY"}
|
||||
],
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
# 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.
|
||||
_MODEL_ARMOR_SANITIZED = {
|
||||
"sanitizedText": "my card is [REDACTED]",
|
||||
"sanitizationResult": {"filterMatchState": "NO_MATCH_FOUND"},
|
||||
}
|
||||
|
||||
|
||||
def _surface_guardrail(**kwargs):
|
||||
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
|
||||
|
||||
|
||||
def _armor_post_mock(payload):
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json = AsyncMock(return_value=payload)
|
||||
return AsyncMock(return_value=mock_response)
|
||||
|
||||
|
||||
async def _anthropic_sse_stream():
|
||||
for chunk in _ANTHROPIC_SSE_CHUNKS:
|
||||
yield chunk
|
||||
|
||||
|
||||
def _responses_api_events():
|
||||
from litellm.types.llms.openai import (
|
||||
OutputTextDeltaEvent,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
completed = ResponsesAPIResponse(
|
||||
id="resp_1",
|
||||
created_at=0,
|
||||
model="gpt-4o-mini",
|
||||
object="response",
|
||||
output=[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_1",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "my card is 4111-1111-1111-1111"}],
|
||||
}
|
||||
],
|
||||
parallel_tool_calls=False,
|
||||
tool_choice="auto",
|
||||
tools=[],
|
||||
)
|
||||
return (
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_1",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta="my card is 4111-1111-1111-1111",
|
||||
),
|
||||
ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=completed,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def _drain_surface_hook(guardrail, chunks):
|
||||
async def _stream():
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
return [
|
||||
item
|
||||
async for item in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
response=_stream(),
|
||||
request_data={
|
||||
"model": "claude-haiku",
|
||||
"messages": [{"role": "user", "content": "show me a card"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]},
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_scans_raw_anthropic_sse_instead_of_crashing():
|
||||
"""A /v1/messages stream arrives as raw SSE bytes and must be assembled, then scanned.
|
||||
|
||||
Regression for the 500 `Error building chunks for logging/streaming usage calculation`:
|
||||
stream_chunk_builder calls .get() on each chunk, which raises on bytes.
|
||||
"""
|
||||
guardrail = _surface_guardrail()
|
||||
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", post):
|
||||
delivered = await _drain_surface_hook(guardrail, _ANTHROPIC_SSE_CHUNKS)
|
||||
|
||||
post.assert_called_once()
|
||||
scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"]
|
||||
assert "my card is 4111-1111-1111-1111" in scanned
|
||||
assert tuple(delivered) == _ANTHROPIC_SSE_CHUNKS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_scans_responses_api_events_instead_of_crashing():
|
||||
"""A /v1/responses stream arrives as typed Responses events, which stream_chunk_builder
|
||||
cannot subscript. The final response.completed event carries the text to scan."""
|
||||
guardrail = _surface_guardrail()
|
||||
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
|
||||
events = _responses_api_events()
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", post):
|
||||
delivered = await _drain_surface_hook(guardrail, events)
|
||||
|
||||
post.assert_called_once()
|
||||
scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"]
|
||||
assert scanned == "my card is 4111-1111-1111-1111"
|
||||
assert tuple(delivered) == events
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_block_emits_anthropic_error_frame():
|
||||
"""A block on /v1/messages must terminate the stream in Anthropic's error format.
|
||||
|
||||
The OpenAI-shaped `data: {"error": ...}` frame the chat surface uses is rejected by
|
||||
Anthropic clients.
|
||||
"""
|
||||
guardrail = _surface_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_BLOCK)
|
||||
):
|
||||
delivered = await _drain_surface_hook(guardrail, _ANTHROPIC_SSE_CHUNKS)
|
||||
|
||||
body = b"".join(delivered)
|
||||
assert b"event: error" in body
|
||||
assert b'"type": "error"' in body
|
||||
assert b"guardrail_error" in body
|
||||
assert b"Streaming response blocked by Model Armor" in body
|
||||
assert b"4111-1111-1111-1111" not in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_block_emits_responses_api_error_event():
|
||||
"""A block on /v1/responses must terminate the stream with a Responses ErrorEvent."""
|
||||
from litellm.types.llms.openai import ErrorEvent
|
||||
|
||||
guardrail = _surface_guardrail()
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_BLOCK)
|
||||
):
|
||||
delivered = await _drain_surface_hook(guardrail, _responses_api_events())
|
||||
|
||||
assert len(delivered) == 1
|
||||
error_event = delivered[0]
|
||||
assert isinstance(error_event, ErrorEvent)
|
||||
assert error_event.error.type == "guardrail_error"
|
||||
assert error_event.error.code == "400"
|
||||
assert error_event.error.message == "Streaming response blocked by Model Armor"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_masking_re_emits_anthropic_sse_with_sanitized_text():
|
||||
"""mask_response_content on /v1/messages must ship the sanitized text, not the original."""
|
||||
guardrail = _surface_guardrail(mask_response_content=True)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_SANITIZED)
|
||||
):
|
||||
delivered = await _drain_surface_hook(guardrail, _ANTHROPIC_SSE_CHUNKS)
|
||||
|
||||
body = b"".join(delivered)
|
||||
assert b"[REDACTED]" in body
|
||||
assert b"4111-1111-1111-1111" not in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_masking_blocks_responses_api_stream():
|
||||
"""A Responses event stream cannot be rebuilt from sanitized text, so releasing it would
|
||||
ship the content the guardrail just rewrote. It is blocked instead."""
|
||||
from litellm.types.llms.openai import ErrorEvent
|
||||
|
||||
guardrail = _surface_guardrail(mask_response_content=True)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_SANITIZED)
|
||||
):
|
||||
delivered = await _drain_surface_hook(guardrail, _responses_api_events())
|
||||
|
||||
assert len(delivered) == 1
|
||||
assert isinstance(delivered[0], ErrorEvent)
|
||||
assert delivered[0].error.code == "400"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("surface", ["anthropic_sse", "responses"])
|
||||
async def test_streaming_api_failure_frames_error_per_surface(surface):
|
||||
"""A Model Armor outage with fail_on_error must terminate the stream in the endpoint's
|
||||
own error format rather than leaking an OpenAI SSE frame onto it."""
|
||||
from litellm.types.llms.openai import ErrorEvent
|
||||
|
||||
guardrail = _surface_guardrail(fail_on_error=True)
|
||||
chunks = _ANTHROPIC_SSE_CHUNKS if surface == "anthropic_sse" else _responses_api_events()
|
||||
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 500
|
||||
mock_response.text = "Internal Server Error"
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler, "post", AsyncMock(return_value=mock_response)
|
||||
):
|
||||
delivered = await _drain_surface_hook(guardrail, chunks)
|
||||
|
||||
assert len(delivered) >= 1
|
||||
if surface == "anthropic_sse":
|
||||
assert b"event: error" in b"".join(delivered)
|
||||
else:
|
||||
assert isinstance(delivered[0], ErrorEvent)
|
||||
assert delivered[0].error.code == "500"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"chunks",
|
||||
[
|
||||
pytest.param(
|
||||
(b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,'
|
||||
b'"delta":{"type":"text_delta","text":"hi"}}\n\n',),
|
||||
id="anthropic-sse-without-message-start",
|
||||
),
|
||||
pytest.param(None, id="responses-stream-without-completed-event"),
|
||||
],
|
||||
)
|
||||
async def test_streaming_hook_fails_closed_when_a_surface_stream_cannot_be_assembled(chunks):
|
||||
"""Forwarding an unscannable /v1/messages or /v1/responses stream would silently disable the
|
||||
guardrail, so the stream is refused in its own wire format instead of released unscanned."""
|
||||
from litellm.types.llms.openai import (
|
||||
ErrorEvent,
|
||||
OutputTextDeltaEvent,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
if chunks is None:
|
||||
chunks = (
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_1",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta="hi",
|
||||
),
|
||||
)
|
||||
guardrail = _surface_guardrail()
|
||||
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", post):
|
||||
delivered = await _drain_surface_hook(guardrail, chunks)
|
||||
|
||||
post.assert_not_called()
|
||||
assert tuple(delivered) != tuple(chunks)
|
||||
if isinstance(chunks[0], bytes):
|
||||
joined = b"".join(item.encode() if isinstance(item, str) else item for item in delivered).decode()
|
||||
assert "event: error" in joined
|
||||
assert "could not be assembled for scanning" in joined
|
||||
return
|
||||
assert len(delivered) == 1
|
||||
assert isinstance(delivered[0], ErrorEvent)
|
||||
assert "could not be assembled for scanning" in delivered[0].error.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_forwards_a_preceding_guardrails_error_item():
|
||||
"""A guardrail earlier in the post_call chain replaces the stream with its own terminal
|
||||
error item. That item is not a chat delta, and feeding it to stream_chunk_builder is what
|
||||
surfaced the ticket's 500, so it has to be forwarded untouched instead."""
|
||||
from litellm.types.llms.openai import (
|
||||
ErrorEvent,
|
||||
ErrorEventError,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
chunks = (
|
||||
ErrorEvent(
|
||||
type=ResponsesAPIStreamEvents.ERROR,
|
||||
sequence_number=1,
|
||||
error=ErrorEventError(
|
||||
type="guardrail_error",
|
||||
code="400",
|
||||
message="Streaming response blocked by Model Armor",
|
||||
param=None,
|
||||
),
|
||||
),
|
||||
)
|
||||
guardrail = _surface_guardrail()
|
||||
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
|
||||
|
||||
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
|
||||
@pytest.mark.parametrize(
|
||||
"chunks",
|
||||
[
|
||||
pytest.param(None, id="anthropic-error-event"),
|
||||
pytest.param(
|
||||
('data: {"error": {"message": "Streaming response blocked by the first guardrail", "code": "400"}}\n\n',),
|
||||
id="chat-completions-error-payload",
|
||||
),
|
||||
],
|
||||
)
|
||||
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()
|
||||
post = _armor_post_mock(_MODEL_ARMOR_CLEAN)
|
||||
|
||||
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_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()
|
||||
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,
|
||||
)
|
||||
|
||||
assert len(items) == 1
|
||||
assert '"code": "400"' in items[0]
|
||||
assert "blocked" in items[0]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue