mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
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:
parent
d44d281d1d
commit
7edf5b36cf
12 changed files with 769 additions and 17 deletions
|
|
@ -117,7 +117,7 @@
|
|||
"limit": 117
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 697
|
||||
"limit": 696
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 5
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,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}
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22704
|
||||
"limit": 22698
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26854
|
||||
"limit": 26853
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue