fix(guardrails): deliver modify_response block as valid SSE on streaming chat and Responses

A guardrail modify_response verdict on a streaming request only produced a
proper replacement on /v1/messages: the chat completions and Responses API
translations had no build_block_sse_chunks, so the ModifyResponseException
re-raised and surfaced as an in-stream 500 error frame (or a whole-request
500 in buffered mode) instead of the documented 200 replacement.

Implement build_block_sse_chunks for both OpenAI translations: chat emits a
content delta plus a finish_reason content_filter chunk with real usage;
Responses emits the typed event sequence (standalone via
build_synthetic_response_events pre-stream, or an output-item continuation
under the in-progress response id mid-stream) ending in response.completed.
This commit is contained in:
mateo-berri 2026-08-31 16:01:39 -07:00
parent d44d281d1d
commit 7edf5b36cf
12 changed files with 769 additions and 17 deletions

View file

@ -117,7 +117,7 @@
"limit": 117
},
"reportUnnecessaryComparison": {
"limit": 697
"limit": 696
},
"reportUnnecessaryContains": {
"limit": 5

View file

@ -144,7 +144,7 @@ class AnthropicMessagesHandler(BaseTranslation):
self,
exc: "ModifyResponseException",
stream_started: bool = False,
responses_so_far: list[object] | None = None,
responses_so_far: Sequence[object] | None = None,
) -> list[bytes]:
"""
Build an Anthropic SSE sequence delivering the guardrail block message
@ -162,7 +162,7 @@ class AnthropicMessagesHandler(BaseTranslation):
would make Anthropic clients reject the stream.
"""
if stream_started:
return self._block_continuation_chunks(exc, responses_so_far or [])
return self._block_continuation_chunks(exc, responses_so_far or ())
return self._standalone_block_chunks(exc)
def _standalone_block_chunks(self, exc: "ModifyResponseException") -> list[bytes]:
@ -187,7 +187,9 @@ class AnthropicMessagesHandler(BaseTranslation):
)
return list(FakeAnthropicMessagesStreamIterator(response=block_response))
def _block_continuation_chunks(self, exc: "ModifyResponseException", responses_so_far: list[object]) -> list[bytes]:
def _block_continuation_chunks(
self, exc: "ModifyResponseException", responses_so_far: Sequence[object]
) -> list[bytes]:
"""Continue an already-started message: close the open content block,
append the block message as a new text block, then end the message --
without a second message_start."""
@ -237,7 +239,7 @@ class AnthropicMessagesHandler(BaseTranslation):
@staticmethod
def _content_block_state(
responses_so_far: list[object],
responses_so_far: Sequence[object],
) -> tuple[int | None, int | None]:
"""From the SSE chunks already sent to the client, return (open
content-block index or None, highest content-block index seen or None).

View file

@ -1,4 +1,5 @@
from abc import ABC, abstractmethod
from collections.abc import Sequence
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Final, Optional
@ -127,8 +128,8 @@ class BaseTranslation(ABC):
self,
exc: "ModifyResponseException",
stream_started: bool = False,
responses_so_far: list[Any] | None = None,
) -> list[bytes] | None:
responses_so_far: Sequence[Any] | None = None,
) -> Sequence[bytes] | None:
"""
Build the streaming chunks that deliver a guardrail block message and
cleanly terminate the stream in this provider's wire format.

View file

@ -124,6 +124,61 @@ def blocked_responses_api_usage(original_response: object) -> ResponseAPIUsage:
)
def stream_item_field(item: object, field: str) -> object | None:
if isinstance(item, dict):
return item.get(field)
return getattr(item, field, None)
def blocked_chat_stream_usage(original_response: object) -> tuple[int, int]:
"""
``(prompt_tokens, completion_tokens)`` for a synthetic guardrail-blocked
chat completions stream.
A mid-stream block carries the chunks received so far as a list; real usage
rides on the final chunk when the upstream sent one
(``stream_options.include_usage``). Non-list originals defer to
``blocked_response_usage``.
"""
if not isinstance(original_response, list):
usage: Final = blocked_response_usage(original_response)
return usage.get("input_tokens", 0), usage.get("output_tokens", 0)
usage_obj: Final = next(
(
chunk_usage
for item in reversed(original_response)
if (chunk_usage := stream_item_field(item, "usage")) is not None
),
None,
)
return (
_usage_tokens(usage_obj, "prompt_tokens", "input_tokens"),
_usage_tokens(usage_obj, "completion_tokens", "output_tokens"),
)
def blocked_responses_stream_usage(original_response: object) -> ResponseAPIUsage:
"""
``ResponseAPIUsage`` for a synthetic guardrail-blocked /v1/responses stream.
A mid-stream block carries the events received so far as a list; real usage
rides on the ``response.completed`` event's response when the upstream sent
one. Non-list originals defer to ``blocked_responses_api_usage``.
"""
if not isinstance(original_response, list):
return blocked_responses_api_usage(original_response)
completed: Final = next(
(
response
for item in reversed(original_response)
if str(stream_item_field(item, "type") or "") == "response.completed"
and (response := stream_item_field(item, "response")) is not None
),
None,
)
return blocked_responses_api_usage(completed)
def effective_skip_system_message_for_guardrail(guardrail_to_apply: Any) -> bool:
per: Final = getattr(guardrail_to_apply, "skip_system_message_in_guardrail", None)
if per is not None:

View file

@ -14,8 +14,14 @@ Pattern Overview:
This pattern can be replicated for other message formats (e.g., Anthropic).
"""
import json
import time
import uuid
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final, Union, cast
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import (
@ -23,6 +29,7 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
StreamTransformSink,
)
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_chat_stream_usage,
effective_scan_only_tool_results_for_guardrail,
effective_skip_system_message_for_guardrail,
effective_skip_tool_message_for_guardrail,
@ -31,6 +38,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
openai_tool_name,
role_out_of_guardrail_scope,
scoped_structured_message_indices,
stream_item_field,
)
from litellm.main import stream_chunk_builder
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
@ -46,7 +54,10 @@ from litellm.types.utils import (
)
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -1000,3 +1011,107 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
else:
# Subsequent chunks - clear the text
content_item["text"] = ""
def build_block_sse_chunks(
self,
exc: "ModifyResponseException",
stream_started: bool = False,
responses_so_far: Sequence[object] | None = None,
) -> Sequence[bytes]:
"""
Build OpenAI chat-completions SSE chunks that deliver the guardrail
block message and terminate the stream cleanly, mirroring the
non-streaming block response: ``finish_reason`` ``content_filter`` plus
the real usage the upstream call consumed.
- ``stream_started`` False (buffered / pre-stream): nothing has been
sent, so open a standalone completion with a ``role`` delta.
- ``stream_started`` True (sampling / mid-stream): chunks already
reached the client, so continue the in-progress completion (reuse its
id/created/model, content-only delta).
The proxy's data generator appends ``data: [DONE]`` itself.
"""
chunk_id, created, model = _blocked_stream_identity(exc, responses_so_far or ())
prompt_tokens, completion_tokens = blocked_chat_stream_usage(exc.original_response)
continuation_delta: Final[_BlockedChunkDelta] = {"content": exc.message}
standalone_delta: Final[_BlockedChunkDelta] = {"role": "assistant", "content": exc.message}
message_chunk: Final[_BlockedChunk] = {
"id": chunk_id,
"object": "chat.completion.chunk",
"created": created,
"model": model,
"choices": (
{
"index": 0,
"delta": continuation_delta if stream_started else standalone_delta,
"finish_reason": None,
},
),
}
final_chunk: Final[_BlockedChunk] = {
"id": chunk_id,
"object": "chat.completion.chunk",
"created": created,
"model": model,
"choices": ({"index": 0, "delta": {}, "finish_reason": "content_filter"},),
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
}
return _chat_sse_chunk(message_chunk), _chat_sse_chunk(final_chunk)
class _BlockedChunkDelta(TypedDict, total=False):
role: ReadOnly[str]
content: ReadOnly[str]
class _BlockedChunkChoice(TypedDict):
index: ReadOnly[int]
delta: ReadOnly[_BlockedChunkDelta]
finish_reason: ReadOnly[str | None]
class _BlockedChunkUsage(TypedDict):
prompt_tokens: ReadOnly[int]
completion_tokens: ReadOnly[int]
total_tokens: ReadOnly[int]
class _BlockedChunk(TypedDict):
id: ReadOnly[str]
object: ReadOnly[str]
created: ReadOnly[int]
model: ReadOnly[str]
choices: ReadOnly[tuple[_BlockedChunkChoice, ...]]
usage: NotRequired[ReadOnly[_BlockedChunkUsage]]
def _chat_sse_chunk(payload: _BlockedChunk) -> bytes:
return f"data: {json.dumps(payload)}\n\n".encode()
def _blocked_stream_identity(
exc: "ModifyResponseException", responses_so_far: Sequence[object]
) -> tuple[str, int, str]:
identified: Final = next(
(
(chunk_id, item)
for item in responses_so_far
if isinstance(chunk_id := stream_item_field(item, "id"), str) and chunk_id
),
None,
)
if identified is None:
return f"chatcmpl-{uuid.uuid4()}", int(time.time()), exc.model
chunk_id, source = identified
created: Final = stream_item_field(source, "created")
model: Final = stream_item_field(source, "model")
return (
chunk_id,
created if isinstance(created, int) else int(time.time()),
model if isinstance(model, str) and model else exc.model,
)

View file

@ -28,6 +28,8 @@ Output: response.output is List[GenericResponseOutputItem] where each has:
- text: str
"""
import time
import uuid
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final, Union, cast
@ -41,15 +43,31 @@ from litellm.completion_extras.litellm_responses_transformation.transformation i
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_responses_stream_usage,
stream_item_field,
)
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from litellm.types.llms.openai import (
AllMessageValues,
BaseLiteLLMOpenAIResponseObject,
ChatCompletionToolCallChunk,
ChatCompletionToolParam,
ContentPartAddedEvent,
ContentPartDoneEvent,
ContentPartDonePartOutputText,
OpenAIMcpServerTool,
OutputItemAddedEvent,
OutputItemDoneEvent,
OutputTextDeltaEvent,
OutputTextDoneEvent,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
ResponsesAPIStreamingResponse,
)
from litellm.types.responses.main import (
GenericResponseOutputItem,
@ -59,11 +77,13 @@ from litellm.types.responses.main import (
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import ResponseInputParam
from litellm.types.utils import ResponsesAPIResponse
class ResponseOutputEnvelope(TypedDict, total=False):
@ -802,3 +822,183 @@ class OpenAIResponsesHandler(BaseTranslation):
content[content_idx]["text"] = guardrail_response
elif hasattr(content[content_idx], "text"):
content[content_idx].text = guardrail_response
def build_block_sse_chunks(
self,
exc: "ModifyResponseException",
stream_started: bool = False,
responses_so_far: Sequence[object] | None = None,
) -> Sequence[bytes]:
"""
Build Responses API SSE events that deliver the guardrail block message
and terminate the stream cleanly, mirroring the non-streaming block
response: a completed response whose only output is the violation text,
with the real usage the upstream call consumed.
- ``stream_started`` False (buffered / pre-stream): nothing has been
sent, so emit the full synthetic sequence (``response.created``
through ``response.completed``).
- ``stream_started`` True (sampling / mid-stream): events already
reached the client, so continue the in-progress response: deliver the
block message as a new output item under the same response id and
close with a ``response.completed`` carrying only the replacement
item.
The proxy's data generator appends ``data: [DONE]`` itself.
"""
events: Final = (
self._block_continuation_events(exc, responses_so_far or ())
if stream_started
else self._standalone_block_events(exc)
)
return tuple(
f"data: {event.model_dump_json(exclude_none=True, exclude_unset=True)}\n\n".encode() for event in events
)
@staticmethod
def _standalone_block_events(exc: "ModifyResponseException") -> Sequence[ResponsesAPIStreamingResponse]:
from litellm.responses.streaming_iterator import build_synthetic_response_events
return build_synthetic_response_events(
transformed=_blocked_response(exc, response_id=f"resp_{uuid.uuid4()}", model=exc.model),
logging_obj=None,
chunk_size=max(len(exc.message), 1),
)
@staticmethod
def _block_continuation_events(
exc: "ModifyResponseException", responses_so_far: Sequence[object]
) -> Sequence[ResponsesAPIStreamingResponse]:
response_id, model, output_index = _continuation_identity(exc, responses_so_far)
item: Final = _blocked_output_item(exc)
item_id: Final = item["id"]
item_model: Final = BaseLiteLLMOpenAIResponseObject.model_validate(item)
part: Final[_BlockedContentPart] = {"type": "output_text", "text": exc.message, "annotations": ()}
done_part: Final[_BlockedDoneContentPart] = {
"type": "output_text",
"text": exc.message,
"annotations": (),
"logprobs": None,
}
return (
OutputItemAddedEvent(
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
output_index=output_index,
item=item_model,
),
ContentPartAddedEvent(
type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED,
item_id=item_id,
output_index=output_index,
content_index=0,
part=BaseLiteLLMOpenAIResponseObject.model_validate(part),
),
OutputTextDeltaEvent(
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
item_id=item_id,
output_index=output_index,
content_index=0,
delta=exc.message,
),
OutputTextDoneEvent(
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
item_id=item_id,
output_index=output_index,
content_index=0,
text=exc.message,
),
ContentPartDoneEvent(
type=ResponsesAPIStreamEvents.CONTENT_PART_DONE,
item_id=item_id,
output_index=output_index,
content_index=0,
part=ContentPartDonePartOutputText.model_validate(done_part),
),
OutputItemDoneEvent(
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
output_index=output_index,
item=item_model,
),
ResponseCompletedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=_blocked_response(exc, response_id=response_id, model=model, output_item=item),
),
)
class _BlockedContentPart(TypedDict):
type: ReadOnly[str]
text: ReadOnly[str]
annotations: ReadOnly[tuple[object, ...]]
class _BlockedDoneContentPart(TypedDict):
type: ReadOnly[str]
text: ReadOnly[str]
annotations: ReadOnly[tuple[object, ...]]
logprobs: ReadOnly[None]
class _BlockedOutputItem(TypedDict):
type: ReadOnly[str]
id: ReadOnly[str]
status: ReadOnly[str]
role: ReadOnly[str]
content: ReadOnly[tuple[_BlockedContentPart, ...]]
class _BlockedResponsePayload(TypedDict):
id: ReadOnly[str]
object: ReadOnly[str]
created_at: ReadOnly[int]
model: ReadOnly[str]
output: ReadOnly[tuple[_BlockedOutputItem, ...]]
status: ReadOnly[str]
usage: ReadOnly[ResponseAPIUsage]
def _blocked_output_item(exc: "ModifyResponseException") -> _BlockedOutputItem:
item: Final[_BlockedOutputItem] = {
"type": "message",
"id": f"msg_{uuid.uuid4()}",
"status": "completed",
"role": "assistant",
"content": ({"type": "output_text", "text": exc.message, "annotations": ()},),
}
return item
def _blocked_response(
exc: "ModifyResponseException",
response_id: str,
model: str,
output_item: _BlockedOutputItem | None = None,
) -> ResponsesAPIResponse:
payload: Final[_BlockedResponsePayload] = {
"id": response_id,
"object": "response",
"created_at": int(time.time()),
"model": model,
"output": (output_item if output_item is not None else _blocked_output_item(exc),),
"status": "completed",
"usage": blocked_responses_stream_usage(exc.original_response),
}
return ResponsesAPIResponse.model_validate(payload)
def _continuation_identity(exc: "ModifyResponseException", responses_so_far: Sequence[object]) -> tuple[str, str, int]:
responses: Final = tuple(
response for item in responses_so_far if (response := stream_item_field(item, "response")) is not None
)
response_id: Final = next(
(rid for response in responses if isinstance(rid := stream_item_field(response, "id"), str) and rid),
f"resp_{uuid.uuid4()}",
)
model: Final = next(
(m for response in responses if isinstance(m := stream_item_field(response, "model"), str) and m),
exc.model,
)
indices: Final = tuple(
index for item in responses_so_far if isinstance(index := stream_item_field(item, "output_index"), int)
)
return response_id, model, max(indices) + 1 if indices else 0

View file

@ -969,7 +969,7 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
transformed: ResponsesAPIResponse,
logging_obj: LiteLLMLoggingObj,
) -> None:
self._events: list[ResponsesAPIStreamingResponse] = _build_synthetic_response_events(
self._events: Sequence[ResponsesAPIStreamingResponse] = build_synthetic_response_events(
transformed=transformed,
logging_obj=logging_obj,
chunk_size=self.CHUNK_SIZE,
@ -1036,7 +1036,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
transformed: ResponsesAPIResponse,
logging_obj: LiteLLMLoggingObj,
) -> None:
self._events = _build_synthetic_response_events(
self._events = build_synthetic_response_events(
transformed=transformed,
logging_obj=logging_obj,
chunk_size=MockResponsesAPIStreamingIterator.CHUNK_SIZE,
@ -1218,10 +1218,10 @@ def _add_text_like_part_events(
)
def _build_synthetic_response_events(
def build_synthetic_response_events(
*,
transformed: ResponsesAPIResponse,
logging_obj: LiteLLMLoggingObj,
logging_obj: LiteLLMLoggingObj | None,
chunk_size: int,
) -> list[ResponsesAPIStreamingResponse]:
openai_types: Final = _get_openai_response_types()

View file

@ -841,7 +841,7 @@ def test_build_synthetic_response_events_covers_annotations_function_calls_and_r
)
try:
events = streaming_module._build_synthetic_response_events(
events = streaming_module.build_synthetic_response_events(
transformed=transformed,
logging_obj=logging_obj,
chunk_size=5,

View file

@ -1559,3 +1559,54 @@ class TestScanOnlyToolResults:
assert data["messages"][3]["content"] == "page says [BLOCKED] here"
assert data["messages"][3]["tool_call_id"] == "call_1"
assert data["messages"][4]["content"] == "and then?"
class TestBuildBlockSseChunks:
"""build_block_sse_chunks turns a streaming ModifyResponseException into 200 SSE chunks"""
def _exc(self, original_response=None):
from litellm.exceptions import ModifyResponseException
return ModifyResponseException(
message="Blocked by policy.",
model="gpt-5.4-mini",
request_data={},
guardrail_name="test",
original_response=original_response,
)
def _payloads(self, chunks):
return [json.loads(chunk.decode().removeprefix("data: ").strip()) for chunk in chunks]
def test_standalone_block_uses_fresh_identity_and_zero_usage(self):
handler = OpenAIChatCompletionsHandler()
first, final = self._payloads(handler.build_block_sse_chunks(self._exc(), stream_started=False))
assert first["id"].startswith("chatcmpl-")
assert first["model"] == "gpt-5.4-mini"
assert first["choices"][0]["delta"] == {"role": "assistant", "content": "Blocked by policy."}
assert first["choices"][0]["finish_reason"] is None
assert final["choices"][0]["delta"] == {}
assert final["choices"][0]["finish_reason"] == "content_filter"
assert final["usage"] == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
def test_continuation_reuses_stream_identity_and_real_usage(self):
handler = OpenAIChatCompletionsHandler()
yielded = [
{"id": "chatcmpl-live", "created": 1724900000, "model": "gpt-5.4-mini-2026-01-01"},
]
original = yielded + [
{"id": "chatcmpl-live", "usage": {"prompt_tokens": 11, "completion_tokens": 5}},
]
first, final = self._payloads(
handler.build_block_sse_chunks(
self._exc(original_response=original), stream_started=True, responses_so_far=yielded
)
)
assert (first["id"], first["created"], first["model"]) == (
"chatcmpl-live",
1724900000,
"gpt-5.4-mini-2026-01-01",
)
assert first["choices"][0]["delta"] == {"content": "Blocked by policy."}
assert final["id"] == "chatcmpl-live"
assert final["usage"] == {"prompt_tokens": 11, "completion_tokens": 5, "total_tokens": 16}

View file

@ -1229,3 +1229,66 @@ class TestOpenAIResponsesHandlerToolInjection:
names = [t.get("name") for t in result["tools"]]
assert "get_weather" in names
assert "injected_tool" in names
class TestBuildBlockSseChunks:
"""build_block_sse_chunks turns a streaming ModifyResponseException into 200 SSE events"""
def _exc(self, original_response=None):
from litellm.exceptions import ModifyResponseException
return ModifyResponseException(
message="Blocked by policy.",
model="gpt-5.4-mini",
request_data={},
guardrail_name="test",
original_response=original_response,
)
def _payloads(self, chunks):
import json
return [json.loads(chunk.decode().removeprefix("data: ").strip()) for chunk in chunks]
def test_standalone_block_emits_complete_synthetic_stream(self):
handler = OpenAIResponsesHandler()
payloads = self._payloads(handler.build_block_sse_chunks(self._exc(), stream_started=False))
types = [payload["type"] for payload in payloads]
assert types[0] == "response.created"
assert types[-1] == "response.completed"
completed = payloads[-1]["response"]
assert completed["id"].startswith("resp_")
assert completed["model"] == "gpt-5.4-mini"
assert completed["output"][0]["content"][0]["text"] == "Blocked by policy."
def test_continuation_appends_item_at_next_output_index_with_real_usage(self):
handler = OpenAIResponsesHandler()
yielded = [
{"type": "response.created", "response": {"id": "resp_live", "model": "gpt-5.4-mini-2026-01-01"}},
{"type": "response.output_item.added", "output_index": 2, "item": {"id": "msg_orig"}},
]
original = yielded + [
{
"type": "response.completed",
"response": {
"id": "resp_live",
"model": "gpt-5.4-mini-2026-01-01",
"output": [],
"usage": {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28},
},
}
]
payloads = self._payloads(
handler.build_block_sse_chunks(
self._exc(original_response=original), stream_started=True, responses_so_far=yielded
)
)
types = [payload["type"] for payload in payloads]
assert "response.created" not in types
assert types[0] == "response.output_item.added"
assert payloads[0]["output_index"] == 3
completed = payloads[-1]["response"]
assert completed["id"] == "resp_live"
assert completed["model"] == "gpt-5.4-mini-2026-01-01"
assert completed["output"][0]["content"][0]["text"] == "Blocked by policy."
assert completed["usage"] == {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28}

View file

@ -0,0 +1,265 @@
"""
Regression tests for blocking an OpenAI-format streaming response from the
unified guardrail post-call streaming iterator hook.
When a guardrail's ``apply_guardrail`` raises ``ModifyResponseException``
while (or at the end of) a chat completions or Responses API stream is being
relayed, the hook must emit a well-formed SSE termination sequence carrying
the block message - NOT a bare ``data: {"error": ...}`` blob that surfaces as
an HTTP 500 error frame and truncates the stream.
"""
import json
from typing import Any, AsyncGenerator, List, Literal, Optional
import pytest
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.types.utils import (
Delta,
GenericGuardrailAPIInputs,
ModelResponseStream,
StreamingChoices,
)
BLOCK_MESSAGE = "This response was replaced by policy."
class _BlockingGuardrail(CustomGuardrail):
"""Mock guardrail that always blocks response scans by raising ModifyResponseException."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
raise ModifyResponseException(
message=BLOCK_MESSAGE,
model="gpt-5.4-mini",
request_data=request_data,
guardrail_name=self.guardrail_name,
)
def _chat_chunk(delta: Delta, finish_reason: Optional[str] = None) -> ModelResponseStream:
return ModelResponseStream(
id="chatcmpl-live",
created=1724900000,
model="gpt-5.4-mini",
choices=[StreamingChoices(index=0, delta=delta, finish_reason=finish_reason)],
)
async def _chat_stream(end: bool) -> AsyncGenerator[ModelResponseStream, None]:
yield _chat_chunk(Delta(role="assistant", content="This "))
for text in ["is ", "the ", "original ", "answer."]:
yield _chat_chunk(Delta(content=text))
if end:
yield _chat_chunk(Delta(), finish_reason="stop")
async def _responses_stream(end: bool) -> AsyncGenerator[dict, None]:
original_text = "This is the original answer."
response_envelope = {"id": "resp_live", "model": "gpt-5.4-mini", "status": "in_progress", "output": []}
yield {"type": "response.created", "response": response_envelope}
yield {"type": "response.in_progress", "response": response_envelope}
yield {
"type": "response.output_item.added",
"output_index": 0,
"item": {"id": "msg_orig", "type": "message", "role": "assistant", "content": []},
}
yield {
"type": "response.content_part.added",
"item_id": "msg_orig",
"output_index": 0,
"content_index": 0,
"part": {"type": "output_text", "text": "", "annotations": []},
}
for delta in ["This ", "is ", "the ", "original ", "answer."]:
yield {
"type": "response.output_text.delta",
"item_id": "msg_orig",
"output_index": 0,
"content_index": 0,
"delta": delta,
}
yield {
"type": "response.output_text.done",
"item_id": "msg_orig",
"output_index": 0,
"content_index": 0,
"text": original_text,
}
if end:
yield {
"type": "response.completed",
"response": {
"id": "resp_live",
"model": "gpt-5.4-mini",
"status": "completed",
"output": [
{
"id": "msg_orig",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": original_text, "annotations": []}],
}
],
"usage": {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28},
},
}
async def _run_hook(
route: str,
stream: AsyncGenerator[Any, None],
sampling_rate: int = 1,
end_of_stream_only: bool = False,
buffer_until_moderated: bool = False,
) -> List[Any]:
guardrail = _BlockingGuardrail(guardrail_name="test-blocking-guardrail", event_hook="post_call")
guardrail.streaming_sampling_rate = sampling_rate
guardrail.streaming_end_of_stream_only = end_of_stream_only
guardrail.streaming_buffer_until_moderated = buffer_until_moderated
unified_guardrail = UnifiedLLMGuardrails()
user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route=route)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-blocking-guardrail"]},
}
collected: List[Any] = []
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=stream,
request_data=request_data,
):
collected.append(chunk)
return collected
def _sse_payloads(collected: List[Any]) -> List[dict]:
payloads = []
for chunk in collected:
if not isinstance(chunk, bytes):
continue
for block in chunk.decode().split("\n\n"):
for line in block.strip().split("\n"):
if line.startswith("data:"):
payloads.append(json.loads(line[len("data:") :].strip()))
return payloads
def _assert_no_error_frame(collected: List[Any]) -> None:
raw = "".join(chunk.decode() for chunk in collected if isinstance(chunk, bytes))
assert '"error"' not in raw, f"unexpected error blob in stream: {raw!r}"
@pytest.mark.asyncio
async def test_chat_pre_stream_block_emits_standalone_completion():
"""Block on the first chunk: a standalone completion opens with a role delta
and ends with finish_reason content_filter."""
collected = await _run_hook("/v1/chat/completions", _chat_stream(end=False))
_assert_no_error_frame(collected)
payloads = _sse_payloads(collected)
assert payloads, "no block SSE chunks were emitted"
assert payloads[0]["choices"][0]["delta"] == {"role": "assistant", "content": BLOCK_MESSAGE}
assert payloads[-1]["choices"][0]["finish_reason"] == "content_filter"
@pytest.mark.asyncio
async def test_chat_mid_stream_block_continues_the_completion():
"""Regression for the LIT-6496 500 error frame: after chunks were already
forwarded, the block continues the same completion id and terminates with
finish_reason content_filter instead of raising into an error blob."""
collected = await _run_hook("/v1/chat/completions", _chat_stream(end=False), sampling_rate=5)
_assert_no_error_frame(collected)
forwarded = [chunk for chunk in collected if isinstance(chunk, ModelResponseStream)]
assert forwarded, "original chunks should have streamed before the block"
payloads = _sse_payloads(collected)
assert payloads, "no block SSE chunks were emitted"
assert all(payload["id"] == "chatcmpl-live" for payload in payloads), (
"block chunks must continue the in-progress completion, not start a new one"
)
assert payloads[0]["choices"][0]["delta"] == {"content": BLOCK_MESSAGE}
assert payloads[-1]["choices"][0]["finish_reason"] == "content_filter"
@pytest.mark.asyncio
async def test_chat_end_of_stream_block_terminates_cleanly():
collected = await _run_hook("/v1/chat/completions", _chat_stream(end=True), end_of_stream_only=True)
_assert_no_error_frame(collected)
payloads = _sse_payloads(collected)
assert BLOCK_MESSAGE in json.dumps(payloads)
assert payloads[-1]["choices"][0]["finish_reason"] == "content_filter"
@pytest.mark.asyncio
async def test_responses_buffered_block_emits_full_event_sequence():
"""Buffered moderation blocks before anything streams: a complete synthetic
Responses stream from response.created through response.completed carrying
the block message, with the original content never released."""
collected = await _run_hook("/v1/responses", _responses_stream(end=True), buffer_until_moderated=True)
_assert_no_error_frame(collected)
assert not [chunk for chunk in collected if isinstance(chunk, dict)], (
"buffered original chunks must never be released after a block"
)
payloads = _sse_payloads(collected)
event_types = [payload["type"] for payload in payloads]
assert event_types[0] == "response.created"
assert "response.output_text.delta" in event_types
assert event_types[-1] == "response.completed"
completed = payloads[-1]["response"]
assert completed["status"] == "completed"
assert completed["output"][0]["content"][0]["text"] == BLOCK_MESSAGE
assert completed["usage"] == {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28}
assert "original answer" not in json.dumps(payloads)
@pytest.mark.asyncio
async def test_responses_mid_stream_block_continues_the_response():
"""Regression for the LIT-6496 500 error frame: after events were already
forwarded, the block appends a new output item under the same response id
and closes with response.completed - never a second response.created."""
collected = await _run_hook("/v1/responses", _responses_stream(end=False))
_assert_no_error_frame(collected)
forwarded_types = [chunk["type"] for chunk in collected if isinstance(chunk, dict)]
assert "response.created" in forwarded_types, "original events should have streamed before the block"
payloads = _sse_payloads(collected)
assert payloads, "no block SSE chunks were emitted"
block_types = [payload["type"] for payload in payloads]
assert "response.created" not in block_types, "a mid-stream block must not restart the response"
assert block_types[0] == "response.output_item.added"
assert block_types[-1] == "response.completed"
assert payloads[0]["output_index"] == 1, "the block item must continue after the original output item"
completed = payloads[-1]["response"]
assert completed["id"] == "resp_live"
assert completed["output"][0]["content"][0]["text"] == BLOCK_MESSAGE
@pytest.mark.asyncio
async def test_responses_end_of_stream_block_reports_original_usage():
collected = await _run_hook("/v1/responses", _responses_stream(end=True), end_of_stream_only=True)
_assert_no_error_frame(collected)
forwarded_types = [chunk["type"] for chunk in collected if isinstance(chunk, dict)]
assert "response.completed" not in forwarded_types, (
"the original terminal event must be withheld and replaced by the block sequence"
)
payloads = _sse_payloads(collected)
completed = payloads[-1]["response"]
assert payloads[-1]["type"] == "response.completed"
assert completed["id"] == "resp_live"
assert completed["output"][0]["content"][0]["text"] == BLOCK_MESSAGE
assert completed["usage"] == {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28}

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22704
"limit": 22698
},
"LIT002": {
"limit": 26854
"limit": 26853
},
"LIT003": {
"limit": 269