mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #39036 from BerriAI/litellm_fix_stream_modify_response_chunks
fix(guardrails): deliver modify_response block as valid SSE on streaming chat and Responses
This commit is contained in:
commit
92d453373a
13 changed files with 1189 additions and 17 deletions
|
|
@ -117,7 +117,7 @@
|
|||
"limit": 111
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 695
|
||||
"limit": 692
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 5
|
||||
|
|
|
|||
|
|
@ -155,8 +155,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.
|
||||
|
|
|
|||
|
|
@ -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 stream_item_field(item, "type") == "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:
|
||||
|
|
|
|||
|
|
@ -14,9 +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 (
|
||||
|
|
@ -24,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,
|
||||
|
|
@ -32,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
|
||||
|
|
@ -49,7 +56,10 @@ from litellm.types.utils import (
|
|||
if TYPE_CHECKING:
|
||||
from fastapi import HTTPException
|
||||
|
||||
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
|
||||
|
||||
|
|
@ -1005,3 +1015,129 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
else:
|
||||
# Subsequent chunks - clear the text
|
||||
content_item["text"] = ""
|
||||
|
||||
def _check_streaming_has_ended(self, responses_so_far: Sequence[object]) -> bool:
|
||||
"""
|
||||
True once any relayed chunk carries a non-null ``finish_reason``.
|
||||
|
||||
The unified guardrail's ``end_of_stream_only`` streaming path probes
|
||||
this via ``hasattr`` to withhold the terminal chunks until
|
||||
end-of-stream moderation runs, so a block can replace the finish
|
||||
instead of trailing after a ``finish_reason`` the client already saw.
|
||||
"""
|
||||
return any(
|
||||
stream_item_field(choice, "finish_reason") is not None
|
||||
for item in responses_so_far
|
||||
for choice in _stream_chunk_choices(item)
|
||||
)
|
||||
|
||||
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 _stream_chunk_choices(item: object) -> Sequence[object]:
|
||||
choices: Final = stream_item_field(item, "choices")
|
||||
if isinstance(choices, Sequence) and not isinstance(choices, (str, bytes)):
|
||||
return choices
|
||||
return ()
|
||||
|
||||
|
||||
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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -28,12 +28,16 @@ Output: response.output is List[GenericResponseOutputItem] where each has:
|
|||
- text: str
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
||||
|
||||
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
|
||||
from openai.types.responses.tool_param import FunctionToolParam
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -41,17 +45,33 @@ 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,
|
||||
ErrorEvent,
|
||||
ErrorEventError,
|
||||
OpenAIMcpServerTool,
|
||||
OutputItemAddedEvent,
|
||||
OutputItemDoneEvent,
|
||||
OutputTextDeltaEvent,
|
||||
OutputTextDoneEvent,
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
|
|
@ -63,11 +83,13 @@ from litellm.types.utils import GenericGuardrailAPIInputs
|
|||
if TYPE_CHECKING:
|
||||
from fastapi import HTTPException
|
||||
|
||||
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):
|
||||
|
|
@ -865,3 +887,331 @@ 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: close the
|
||||
output item still open on the wire, 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, serialize_as_any=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
|
||||
part: Final[_BlockedContentPart] = {"type": "output_text", "text": exc.message, "annotations": ()}
|
||||
done_part: Final[_BlockedDoneContentPart] = {
|
||||
"type": "output_text",
|
||||
"text": exc.message,
|
||||
"annotations": (),
|
||||
"logprobs": None,
|
||||
}
|
||||
return (
|
||||
*_open_item_closing_events(responses_so_far),
|
||||
OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=output_index,
|
||||
item=item,
|
||||
),
|
||||
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,
|
||||
),
|
||||
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 _BlockedItemPayload(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[GenericResponseOutputItem, ...]]
|
||||
status: ReadOnly[str]
|
||||
usage: ReadOnly[ResponseAPIUsage]
|
||||
|
||||
|
||||
def _blocked_output_item(exc: "ModifyResponseException") -> GenericResponseOutputItem:
|
||||
payload: Final[_BlockedItemPayload] = {
|
||||
"type": "message",
|
||||
"id": f"msg_{uuid.uuid4()}",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": ({"type": "output_text", "text": exc.message, "annotations": ()},),
|
||||
}
|
||||
return GenericResponseOutputItem.model_validate(payload)
|
||||
|
||||
|
||||
def _blocked_response(
|
||||
exc: "ModifyResponseException",
|
||||
response_id: str,
|
||||
model: str,
|
||||
output_item: GenericResponseOutputItem | 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
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _OpenItemState:
|
||||
item_id: str
|
||||
item_type: str
|
||||
role: str
|
||||
output_index: int
|
||||
content_index: int
|
||||
text: str
|
||||
part_open: bool
|
||||
payload: object
|
||||
|
||||
|
||||
def _open_item_state(responses_so_far: Sequence[object]) -> _OpenItemState | None:
|
||||
typed: Final = tuple((stream_item_field(event, "type"), event) for event in responses_so_far)
|
||||
added: Final = tuple(
|
||||
(added_index, stream_item_field(event, "item"))
|
||||
for event_type, event in typed
|
||||
if event_type == "response.output_item.added"
|
||||
and isinstance(added_index := stream_item_field(event, "output_index"), int)
|
||||
)
|
||||
done_indices: Final = frozenset(
|
||||
done_index
|
||||
for event_type, event in typed
|
||||
if event_type == "response.output_item.done"
|
||||
and isinstance(done_index := stream_item_field(event, "output_index"), int)
|
||||
)
|
||||
open_added: Final = tuple((index, payload) for index, payload in added if index not in done_indices)
|
||||
if not open_added:
|
||||
return None
|
||||
output_index, item_payload = open_added[-1]
|
||||
if item_payload is None:
|
||||
return None
|
||||
item_id: Final = stream_item_field(item_payload, "id")
|
||||
if not isinstance(item_id, str) or not item_id:
|
||||
return None
|
||||
raw_type: Final = stream_item_field(item_payload, "type")
|
||||
raw_role: Final = stream_item_field(item_payload, "role")
|
||||
part_added: Final = tuple(
|
||||
part_index
|
||||
for event_type, event in typed
|
||||
if event_type == "response.content_part.added"
|
||||
and stream_item_field(event, "item_id") == item_id
|
||||
and isinstance(part_index := stream_item_field(event, "content_index"), int)
|
||||
)
|
||||
part_done: Final = frozenset(
|
||||
part_done_index
|
||||
for event_type, event in typed
|
||||
if event_type == "response.content_part.done"
|
||||
and stream_item_field(event, "item_id") == item_id
|
||||
and isinstance(part_done_index := stream_item_field(event, "content_index"), int)
|
||||
)
|
||||
open_parts: Final = tuple(index for index in part_added if index not in part_done)
|
||||
text: Final = "".join(
|
||||
delta
|
||||
for event_type, event in typed
|
||||
if event_type == "response.output_text.delta"
|
||||
and stream_item_field(event, "item_id") == item_id
|
||||
and isinstance(delta := stream_item_field(event, "delta"), str)
|
||||
)
|
||||
return _OpenItemState(
|
||||
item_id=item_id,
|
||||
item_type=raw_type if isinstance(raw_type, str) and raw_type else "message",
|
||||
role=raw_role if isinstance(raw_role, str) and raw_role else "assistant",
|
||||
output_index=output_index,
|
||||
content_index=open_parts[-1] if open_parts else 0,
|
||||
text=text,
|
||||
part_open=bool(open_parts),
|
||||
payload=item_payload,
|
||||
)
|
||||
|
||||
|
||||
_item_fields_adapter: Final = TypeAdapter(Mapping[str, object])
|
||||
_no_item_fields: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _incomplete_item_fields(payload: object) -> Mapping[str, object]:
|
||||
raw: Final = payload.model_dump() if isinstance(payload, BaseModel) else payload
|
||||
if not isinstance(raw, dict):
|
||||
return _no_item_fields
|
||||
return _item_fields_adapter.validate_python(raw)
|
||||
|
||||
|
||||
def _open_item_closing_events(responses_so_far: Sequence[object]) -> Sequence[ResponsesAPIStreamingResponse]:
|
||||
"""Close the output item still in progress on the relayed stream before the
|
||||
block item is appended: strict Responses clients reject a
|
||||
``response.completed`` that arrives while an earlier ``output_item.added``
|
||||
was never closed. A message item closes ``completed`` with exactly the text
|
||||
the client has received so far; any other item type (a function call the
|
||||
guardrail rejected, for instance) closes ``incomplete`` so the synthetic
|
||||
done event can never authorize acting on it."""
|
||||
open_item: Final = _open_item_state(responses_so_far)
|
||||
if open_item is None:
|
||||
return ()
|
||||
if open_item.item_type != "message":
|
||||
return (
|
||||
OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=open_item.output_index,
|
||||
item=BaseLiteLLMOpenAIResponseObject.model_validate(
|
||||
MappingProxyType({**_incomplete_item_fields(open_item.payload), "status": "incomplete"})
|
||||
),
|
||||
),
|
||||
)
|
||||
partial_part: Final[_BlockedContentPart] = {
|
||||
"type": "output_text",
|
||||
"text": open_item.text,
|
||||
"annotations": (),
|
||||
}
|
||||
closed_payload: Final[_BlockedItemPayload] = {
|
||||
"type": open_item.item_type,
|
||||
"id": open_item.item_id,
|
||||
"status": "completed",
|
||||
"role": open_item.role,
|
||||
"content": (partial_part,),
|
||||
}
|
||||
item_done: Final = OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=open_item.output_index,
|
||||
item=GenericResponseOutputItem.model_validate(closed_payload),
|
||||
)
|
||||
if not open_item.part_open:
|
||||
return (item_done,)
|
||||
partial_done_part: Final[_BlockedDoneContentPart] = {
|
||||
"type": "output_text",
|
||||
"text": open_item.text,
|
||||
"annotations": (),
|
||||
"logprobs": None,
|
||||
}
|
||||
return (
|
||||
OutputTextDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
|
||||
item_id=open_item.item_id,
|
||||
output_index=open_item.output_index,
|
||||
content_index=open_item.content_index,
|
||||
text=open_item.text,
|
||||
),
|
||||
ContentPartDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_DONE,
|
||||
item_id=open_item.item_id,
|
||||
output_index=open_item.output_index,
|
||||
content_index=open_item.content_index,
|
||||
part=ContentPartDonePartOutputText.model_validate(partial_done_part),
|
||||
),
|
||||
item_done,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1023,7 +1023,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,
|
||||
|
|
@ -1090,7 +1090,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,
|
||||
|
|
@ -1274,10 +1274,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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1559,3 +1559,87 @@ 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}
|
||||
|
||||
|
||||
class TestCheckStreamingHasEnded:
|
||||
"""_check_streaming_has_ended lets end_of_stream_only withhold the finish chunk until moderation"""
|
||||
|
||||
def test_empty_and_content_only_chunks_are_not_ended(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
assert handler._check_streaming_has_ended([]) is False
|
||||
content_only = [
|
||||
{"id": "chatcmpl-live", "choices": [{"index": 0, "delta": {"content": "hi"}, "finish_reason": None}]},
|
||||
{"id": "chatcmpl-live", "choices": []},
|
||||
{"id": "chatcmpl-live", "usage": {"prompt_tokens": 1, "completion_tokens": 1}},
|
||||
]
|
||||
assert handler._check_streaming_has_ended(content_only) is False
|
||||
|
||||
def test_dict_finish_chunk_marks_stream_ended(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = [
|
||||
{"id": "chatcmpl-live", "choices": [{"index": 0, "delta": {"content": "hi"}, "finish_reason": None}]},
|
||||
{"id": "chatcmpl-live", "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]},
|
||||
]
|
||||
assert handler._check_streaming_has_ended(chunks) is True
|
||||
|
||||
def test_object_finish_chunk_marks_stream_ended(self):
|
||||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = [
|
||||
ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content=None), finish_reason="stop")]
|
||||
)
|
||||
]
|
||||
assert handler._check_streaming_has_ended(chunks) is True
|
||||
|
|
|
|||
|
|
@ -1321,3 +1321,219 @@ 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.done"
|
||||
assert payloads[0]["output_index"] == 2
|
||||
assert payloads[0]["item"]["id"] == "msg_orig"
|
||||
assert payloads[0]["item"]["status"] == "completed"
|
||||
assert types[1] == "response.output_item.added"
|
||||
assert payloads[1]["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}
|
||||
|
||||
def test_continuation_reads_usage_from_typed_completed_event(self):
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
original = [
|
||||
ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=ResponsesAPIResponse.model_validate(
|
||||
{
|
||||
"id": "resp_live",
|
||||
"created_at": 1,
|
||||
"model": "gpt-5.4-mini",
|
||||
"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=[]
|
||||
)
|
||||
)
|
||||
completed = payloads[-1]["response"]
|
||||
assert completed["usage"]["input_tokens"] == 7
|
||||
assert completed["usage"]["output_tokens"] == 21
|
||||
assert completed["usage"]["total_tokens"] == 28
|
||||
|
||||
def test_continuation_closes_open_item_given_pydantic_events_with_enum_types(self):
|
||||
from litellm.types.llms.openai import (
|
||||
BaseLiteLLMOpenAIResponseObject,
|
||||
ContentPartAddedEvent,
|
||||
OutputItemAddedEvent,
|
||||
OutputTextDeltaEvent,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
open_item = GenericResponseOutputItem.model_validate(
|
||||
{"type": "message", "id": "msg_live", "status": "in_progress", "role": "assistant", "content": []}
|
||||
)
|
||||
yielded = [
|
||||
OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, output_index=0, item=open_item
|
||||
),
|
||||
ContentPartAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED,
|
||||
item_id="msg_live",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
part=BaseLiteLLMOpenAIResponseObject.model_validate(
|
||||
{"type": "output_text", "text": "", "annotations": []}
|
||||
),
|
||||
),
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_live",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta="partial ",
|
||||
),
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_live",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta="text",
|
||||
),
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=yielded), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert types[:3] == [
|
||||
"response.output_text.done",
|
||||
"response.content_part.done",
|
||||
"response.output_item.done",
|
||||
]
|
||||
assert payloads[0]["text"] == "partial text"
|
||||
assert payloads[2]["item"]["id"] == "msg_live"
|
||||
assert payloads[2]["item"]["status"] == "completed"
|
||||
assert payloads[2]["item"]["content"][0]["text"] == "partial text"
|
||||
assert types[3] == "response.output_item.added"
|
||||
assert payloads[3]["output_index"] == 1
|
||||
|
||||
def test_continuation_closes_open_function_call_as_incomplete(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
yielded = [
|
||||
{"type": "response.created", "response": {"id": "resp_live", "model": "gpt-5.4-mini"}},
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"output_index": 0,
|
||||
"item": {
|
||||
"id": "fc_live",
|
||||
"type": "function_call",
|
||||
"status": "in_progress",
|
||||
"call_id": "call_1",
|
||||
"name": "run_payment",
|
||||
"arguments": "",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.function_call_arguments.delta",
|
||||
"item_id": "fc_live",
|
||||
"output_index": 0,
|
||||
"delta": '{"amount": 100}',
|
||||
},
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=yielded), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert types[0] == "response.output_item.done"
|
||||
closed = payloads[0]["item"]
|
||||
assert closed["id"] == "fc_live"
|
||||
assert closed["type"] == "function_call"
|
||||
assert closed["status"] == "incomplete"
|
||||
assert closed["name"] == "run_payment"
|
||||
assert "content" not in closed
|
||||
assert types[1] == "response.output_item.added"
|
||||
assert payloads[1]["output_index"] == 1
|
||||
assert types[-1] == "response.completed"
|
||||
|
||||
def test_continuation_without_open_item_emits_no_closing_events(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
yielded = [
|
||||
{"type": "response.created", "response": {"id": "resp_live", "model": "gpt-5.4-mini"}},
|
||||
{"type": "response.in_progress", "response": {"id": "resp_live"}},
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=yielded), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert types[0] == "response.output_item.added"
|
||||
assert types[-1] == "response.completed"
|
||||
dones = [payload for payload in payloads if payload["type"] == "response.output_item.done"]
|
||||
assert len(dones) == 1
|
||||
assert dones[0]["item"]["content"][0]["text"] == "Blocked by policy."
|
||||
|
|
|
|||
|
|
@ -5524,7 +5524,9 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
|
|||
"""Regression for PR #38722: a topicPolicy DENY caught by the end-of-stream
|
||||
scan used to raise after SSE headers were flushed, so the client saw a
|
||||
silently truncated stream. The unified hook must emit the chat in-stream
|
||||
error frame instead."""
|
||||
error frame instead. The finish chunk is withheld while the end-of-stream
|
||||
scan runs, so on a block it is dropped rather than relayed before the
|
||||
frame."""
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
|
||||
unified_guardrail as unified_module,
|
||||
|
|
@ -5582,8 +5584,9 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
|
|||
finally:
|
||||
unified_module.endpoint_guardrail_translation_mappings = None
|
||||
|
||||
assert len(out) == 3
|
||||
assert len(out) == 2
|
||||
assert isinstance(out[0], ModelResponseStream)
|
||||
assert out[0].choices[0].finish_reason is None
|
||||
frame = out[-1]
|
||||
assert isinstance(frame, bytes)
|
||||
payload = json.loads(frame.decode()[len("data: ") :])
|
||||
|
|
|
|||
|
|
@ -0,0 +1,327 @@
|
|||
"""
|
||||
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, Dict, Literal, Optional, Tuple, Union
|
||||
|
||||
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."
|
||||
|
||||
JsonPayload = Dict[str, object]
|
||||
StreamChunk = Union[ModelResponseStream, JsonPayload, bytes]
|
||||
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
class _PassingGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that always lets response scans through unchanged."""
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
return inputs
|
||||
|
||||
|
||||
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[JsonPayload, 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[Union[ModelResponseStream, JsonPayload], None],
|
||||
sampling_rate: int = 1,
|
||||
end_of_stream_only: bool = False,
|
||||
buffer_until_moderated: bool = False,
|
||||
blocks: bool = True,
|
||||
) -> Tuple[StreamChunk, ...]:
|
||||
guardrail = (
|
||||
_BlockingGuardrail(guardrail_name="test-blocking-guardrail", event_hook="post_call")
|
||||
if blocks
|
||||
else _PassingGuardrail(guardrail_name="test-passing-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": [guardrail.guardrail_name]},
|
||||
}
|
||||
|
||||
return tuple(
|
||||
[
|
||||
chunk
|
||||
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,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _sse_payloads(collected: Tuple[StreamChunk, ...]) -> Tuple[JsonPayload, ...]:
|
||||
return tuple(
|
||||
json.loads(line[len("data:") :].strip())
|
||||
for chunk in collected
|
||||
if isinstance(chunk, bytes)
|
||||
for block in chunk.decode().split("\n\n")
|
||||
for line in block.strip().split("\n")
|
||||
if line.startswith("data:")
|
||||
)
|
||||
|
||||
|
||||
def _assert_no_error_frame(collected: Tuple[StreamChunk, ...]) -> 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():
|
||||
"""Regression for bugbot's finish-ordering finding: in end_of_stream_only
|
||||
mode the original finish chunk must be withheld until moderation decides,
|
||||
so a block's content_filter finish is the only stream terminator a client
|
||||
ever sees - never policy text trailing after finish_reason stop."""
|
||||
collected = await _run_hook("/v1/chat/completions", _chat_stream(end=True), end_of_stream_only=True)
|
||||
_assert_no_error_frame(collected)
|
||||
forwarded = [chunk for chunk in collected if isinstance(chunk, ModelResponseStream)]
|
||||
assert forwarded, "content chunks still stream to the client before end-of-stream moderation"
|
||||
assert all(choice.finish_reason is None for chunk in forwarded for choice in chunk.choices), (
|
||||
"the original finish chunk must be withheld until moderation decides"
|
||||
)
|
||||
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_chat_end_of_stream_pass_releases_withheld_finish_chunk():
|
||||
"""When end-of-stream moderation passes, the withheld finish chunk is
|
||||
released so a clean stream still terminates normally."""
|
||||
collected = await _run_hook(
|
||||
"/v1/chat/completions", _chat_stream(end=True), end_of_stream_only=True, blocks=False
|
||||
)
|
||||
assert not [chunk for chunk in collected if isinstance(chunk, bytes)], (
|
||||
"a clean stream must carry no synthetic block frames"
|
||||
)
|
||||
forwarded = [chunk for chunk in collected if isinstance(chunk, ModelResponseStream)]
|
||||
finish_reasons = [choice.finish_reason for chunk in forwarded for choice in chunk.choices]
|
||||
assert finish_reasons[-1] == "stop", "the withheld finish chunk must be released after moderation passes"
|
||||
assert all(reason is None for reason in finish_reasons[:-1])
|
||||
|
||||
|
||||
@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 and bugbot's unclosed-item
|
||||
finding: after events were already forwarded, the block first closes the
|
||||
output item still open on the wire, then appends the replacement item under
|
||||
the same response id, and closes with response.completed - never a second
|
||||
response.created and never a completed response with an item left open."""
|
||||
collected = await _run_hook("/v1/responses", _responses_stream(end=False))
|
||||
_assert_no_error_frame(collected)
|
||||
forwarded = [chunk for chunk in collected if isinstance(chunk, dict)]
|
||||
forwarded_types = [chunk["type"] for chunk in forwarded]
|
||||
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[-1] == "response.completed"
|
||||
|
||||
all_events = forwarded + list(payloads)
|
||||
opened = sorted(event["output_index"] for event in all_events if event["type"] == "response.output_item.added")
|
||||
closed = sorted(event["output_index"] for event in all_events if event["type"] == "response.output_item.done")
|
||||
assert opened == closed, "every output item opened on the stream must be closed before response.completed"
|
||||
original_done_position = block_types.index("response.output_item.done")
|
||||
block_item_position = block_types.index("response.output_item.added")
|
||||
assert original_done_position < block_item_position, (
|
||||
"the in-progress original item must be closed before the block item is appended"
|
||||
)
|
||||
assert payloads[original_done_position]["item"]["id"] == "msg_orig"
|
||||
assert payloads[block_item_position]["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}
|
||||
|
|
@ -1844,7 +1844,8 @@ class TestStreamingHttpErrorFrames:
|
|||
|
||||
out = await _drive_stream(UnifiedLLMGuardrails(), guardrail, chunks)
|
||||
|
||||
assert out[:2] == chunks
|
||||
assert out[0] == chunks[0]
|
||||
assert chunks[1] not in out
|
||||
frame = out[-1]
|
||||
assert isinstance(frame, bytes)
|
||||
text = frame.decode()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22367
|
||||
"limit": 22364
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26777
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue