feat(guardrails): deliver tool-call rewrites into buffered streams

A post_call pipeline guardrail that rewrites a streamed tool call (its
arguments or its name) now has that rewrite written back across the buffered
chunks on chat, Responses, and Messages streams, so the client receives the
rewritten tool call instead of the original. The chat handler rewrites the
first fragment of each tool-call index and blanks the rest, the Responses
handler syncs the function_call output items and their argument events, and
the Messages handler rewrites the tool_use content_block_start and
input_json_delta events in both dict and SSE-bytes chunks.

The delivers_ended_stream_text_rewrites flag becomes
delivers_ended_stream_rewrites, since the write-back now covers both text and
tool calls, and the executor only discards a tool-call rewrite on translations
without write-back or on a shape the translation refuses.
This commit is contained in:
mateo-berri 2026-09-08 11:38:18 -07:00
parent 08b60c409a
commit cbe340a31c
11 changed files with 618 additions and 72 deletions

View file

@ -13,10 +13,11 @@ Pattern Overview:
"""
import json
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Callable, Mapping, Sequence
from copy import deepcopy
from dataclasses import dataclass
from itertools import chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable
from typing_extensions import ReadOnly, TypedDict, assert_never
@ -41,6 +42,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
merge_guardrailed_scoped_messages,
merge_returned_tools_into_request_tools,
scoped_structured_message_indices,
stream_item_field,
stream_item_fingerprint,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
@ -153,6 +155,28 @@ class ExtractedInput:
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
@dataclass(frozen=True, slots=True)
class _ToolCallShape:
name: str | None
arguments: str
_SSEEventRewriter = Callable[[Mapping[str, object]], Mapping[str, object] | None]
def _tool_call_shapes(tool_calls: Sequence[object]) -> tuple[_ToolCallShape, ...]:
"""The guardrail-visible shape of each tool call, whether the guardrail handed
back the ``ChatCompletionMessageToolCall`` objects it was given or plain dicts."""
functions: Final = tuple(stream_item_field(tool_call, "function") for tool_call in tool_calls)
return tuple(
_ToolCallShape(
name=name if isinstance(name := stream_item_field(function, "name"), str) else None,
arguments=arguments if isinstance(arguments := stream_item_field(function, "arguments"), str) else "",
)
for function in functions
)
class _AnthropicSSEDelta(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[str]
@ -170,7 +194,7 @@ class AnthropicMessagesHandler(BaseTranslation):
them through guardrail rewrites; downstream provider handling is out of scope.
"""
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
def __init__(self):
super().__init__()
@ -1050,6 +1074,7 @@ class AnthropicMessagesHandler(BaseTranslation):
first_choice.message.tool_calls,
)
string_so_far = first_choice.message.content
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_list or ())
guardrail_inputs: Final = GenericGuardrailAPIInputs()
if string_so_far:
guardrail_inputs["texts"] = [string_so_far]
@ -1084,6 +1109,19 @@ class AnthropicMessagesHandler(BaseTranslation):
and guardrailed_texts[0] != string_so_far
):
self._write_ended_stream_text_rewrite(responses_so_far, guardrailed_texts[0])
if deliver_ended_stream_rewrites:
returned_tool_calls: Final = _guardrailed_inputs.get("tool_calls")
self._write_ended_stream_tool_call_rewrites(
responses_so_far,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=_tool_call_shapes(
returned_tool_calls
if isinstance(returned_tool_calls, list)
and len(returned_tool_calls) == len(pre_guardrail_tool_calls)
else tool_calls_list or ()
),
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
else:
verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
return responses_so_far
@ -1212,38 +1250,120 @@ class AnthropicMessagesHandler(BaseTranslation):
"""Deliver an ended-stream guardrail text rewrite by rewriting the
buffered chunks in place: the first ``text_delta`` carries the full
rewritten text and every later one is blanked, leaving the surrounding
message and content-block framing untouched. Handles both chunk formats
this stream carries (parsed event dicts and raw SSE bytes)."""
message and content-block framing untouched."""
replacements: Final = chain((rewritten_text,), repeat(""))
for idx, item in enumerate(responses_so_far):
if isinstance(item, dict):
delta = item.get("delta")
if item.get("type") == "content_block_delta" and isinstance(delta, dict):
if delta.get("type") == "text_delta":
delta["text"] = next(replacements)
elif isinstance(item, (bytes, bytearray)):
responses_so_far[idx] = ( # rebind-ok: delivers the rewrite into the caller's buffer
AnthropicMessagesHandler._rewrite_sse_text_deltas(bytes(item), replacements)
)
def rewrite_text_delta(event: Mapping[str, object]) -> Mapping[str, object] | None:
delta: Final = event.get("delta")
if event.get("type") != "content_block_delta" or not isinstance(delta, Mapping):
return None
if delta.get("type") != "text_delta":
return None
return {**event, "delta": {**delta, "text": next(replacements)}}
AnthropicMessagesHandler._rewrite_ended_stream_events(responses_so_far, rewrite_text_delta)
@classmethod
def _write_ended_stream_tool_call_rewrites(
cls,
responses_so_far: list[Any], # mutable-ok: rewrites the caller's buffered chunks in place
*,
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
guardrail_name: str,
) -> None:
"""Deliver ended-stream guardrail tool-call rewrites by rewriting the
buffered chunks in place: the rebuilt response lists tool calls in the
order of the stream's ``tool_use`` blocks, so the nth rewritten call lands
on the nth block, its first ``input_json_delta`` carrying the full rewritten
arguments, every later one blanked, and ``content_block_start`` carrying the
rewritten name. Blocks that do not line up with the rebuilt tool calls make
the rewrite undeliverable, so the pipeline executor discards it and releases
the original chunks."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
block_indices: Final = tuple(
index
for item in responses_so_far
for event in cls._iter_sse_events(item)
if event.get("type") == "content_block_start"
and isinstance(block := event.get("content_block"), Mapping)
and block.get("type") == "tool_use"
and isinstance(index := event.get("index"), int)
)
if len(block_indices) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
rewrites_by_block: Final = MappingProxyType(
{
index: after
for index, before, after in zip(block_indices, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
}
)
argument_replacements: Final = MappingProxyType(
{index: chain((rewrite.arguments,), repeat("")) for index, rewrite in rewrites_by_block.items()}
)
def rewrite_tool_use(event: Mapping[str, object]) -> Mapping[str, object] | None:
index: Final = event.get("index")
if not isinstance(index, int) or index not in rewrites_by_block:
return None
match event.get("type"):
case "content_block_start":
block: Final = event.get("content_block")
name: Final = rewrites_by_block[index].name
if not isinstance(block, Mapping) or name is None:
return None
return {**event, "content_block": {**block, "name": name}}
case "content_block_delta":
delta: Final = event.get("delta")
if not isinstance(delta, Mapping) or delta.get("type") != "input_json_delta":
return None
return {**event, "delta": {**delta, "partial_json": next(argument_replacements[index])}}
case _:
return None
cls._rewrite_ended_stream_events(responses_so_far, rewrite_tool_use)
@staticmethod
def _rewrite_sse_text_deltas(sse_bytes: bytes, replacements: "Iterator[str]") -> bytes:
"""Rewrite every ``text_delta`` data line in one SSE chunk with the next
replacement text, leaving all other events and framing byte-identical."""
def _rewrite_ended_stream_events(
responses_so_far: list[Any], # mutable-ok: rewrites the caller's buffered chunks in place
rewrite_event: _SSEEventRewriter,
) -> None:
"""Replace every buffered event ``rewrite_event`` returns a rewrite for, in
both chunk formats this stream carries (parsed event dicts and raw SSE
bytes), leaving every other event and the framing untouched."""
rewritten_items: Final = tuple(
AnthropicMessagesHandler._rewrite_buffered_item(item, rewrite_event) for item in responses_so_far
)
responses_so_far[:] = rewritten_items # rebind-ok: delivers the rewrites into the caller's buffer
@staticmethod
def _rewrite_buffered_item(item: object, rewrite_event: _SSEEventRewriter) -> object:
if isinstance(item, dict):
rewritten: Final = rewrite_event(_as_str_mapping(item))
return item if rewritten is None else dict(rewritten)
if isinstance(item, (bytes, bytearray)):
return AnthropicMessagesHandler._rewrite_sse_events(bytes(item), rewrite_event)
return item
@staticmethod
def _rewrite_sse_events(sse_bytes: bytes, rewrite_event: _SSEEventRewriter) -> bytes:
"""Rewrite the data lines of one SSE chunk that ``rewrite_event`` rewrites,
leaving all other events and framing byte-identical."""
try:
decoded: Final = sse_bytes.decode("utf-8")
except UnicodeDecodeError:
return sse_bytes
return "\n\n".join(
AnthropicMessagesHandler._rewrite_sse_block(block, replacements) for block in decoded.split("\n\n")
"\n".join(AnthropicMessagesHandler._rewrite_sse_line(line, rewrite_event) for line in block.split("\n"))
for block in decoded.split("\n\n")
).encode("utf-8")
@staticmethod
def _rewrite_sse_block(block: str, replacements: "Iterator[str]") -> str:
return "\n".join(AnthropicMessagesHandler._rewrite_sse_line(line, replacements) for line in block.split("\n"))
@staticmethod
def _rewrite_sse_line(line: str, replacements: "Iterator[str]") -> str:
def _rewrite_sse_line(line: str, rewrite_event: _SSEEventRewriter) -> str:
if not line.startswith("data:"):
return line
try:
@ -1252,14 +1372,10 @@ class AnthropicMessagesHandler(BaseTranslation):
)
except json.JSONDecodeError:
return line
if not isinstance(data, dict) or data.get("type") != "content_block_delta":
if not isinstance(data, dict):
return line
delta: Final = data.get("delta")
if not isinstance(delta, dict) or delta.get("type") != "text_delta":
return line
return "data: " + json.dumps(
{**data, "delta": {**delta, "text": next(replacements)}} # mutable-ok: json.dumps needs plain dicts
)
rewritten: Final = rewrite_event(_as_str_mapping(data))
return line if rewritten is None else "data: " + json.dumps(rewritten)
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
stream_ended: Final = self._check_streaming_has_ended(responses_so_far)

View file

@ -52,13 +52,14 @@ class StreamingScanKey:
class BaseTranslation(ABC):
delivers_ended_stream_text_rewrites: ClassVar[bool] = False
delivers_ended_stream_rewrites: ClassVar[bool] = False
"""Whether ``process_output_streaming_response`` accepts
``deliver_ended_stream_rewrites=True`` and, on an ended (fully buffered)
stream, writes guardrail text rewrites back across ``responses_so_far`` so
a buffered pipeline can release rewritten chunks. Tool-call rewrites, and
text rewrites on every other translation, are undeliverable: the pipeline
executor discards them and releases the original chunks."""
stream, writes guardrail text and tool-call rewrites back across
``responses_so_far`` so a buffered pipeline can release rewritten chunks,
raising ``UndeliverableStreamRewrite`` for a shape it cannot place. Rewrites
on every other translation are undeliverable: the pipeline executor
discards them and releases the original chunks."""
@staticmethod
def transform_user_api_key_dict_to_metadata(
@ -175,9 +176,9 @@ class BaseTranslation(ABC):
transformations (see ``StreamTransformSink``); base handlers ignore it.
``deliver_ended_stream_rewrites`` is passed True only when the caller
holds the whole buffered stream and the subclass declares
``delivers_ended_stream_text_rewrites``: the handler then writes
guardrail text rewrites back across ``responses_so_far`` instead of
discarding them.
``delivers_ended_stream_rewrites``: the handler then writes
guardrail text and tool-call rewrites back across ``responses_so_far``
instead of discarding them.
"""
return responses_so_far

View file

@ -49,6 +49,8 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
coerce_stream_holdback_value,
)
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Choices,
GenericGuardrailAPIInputs,
ModelResponse,
@ -78,7 +80,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
"""
@ -610,13 +612,14 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
deliver_ended_stream_rewrites: bool,
) -> None:
"""Ended-stream path: rebuild the full response, run the non-streaming
output guardrail against it, and (when opted in) write any text rewrite
back across the buffered chunks."""
output guardrail against it, and (when opted in) write any text or
tool-call rewrite back across the buffered chunks."""
model_response: Final = cast(
ModelResponse,
stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj),
)
pre_guardrail_texts: Final = self._string_choice_contents(model_response)
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
await self.process_output_response(
response=model_response,
guardrail_to_apply=guardrail_to_apply,
@ -624,13 +627,21 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
if deliver_ended_stream_rewrites:
await self._write_ended_stream_text_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_texts=pre_guardrail_texts,
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
if not deliver_ended_stream_rewrites:
return
guardrail_name: Final = guardrail_to_apply.guardrail_name or "unknown"
await self._write_ended_stream_text_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_texts=pre_guardrail_texts,
guardrail_name=guardrail_name,
)
self._write_ended_stream_tool_call_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
guardrail_name=guardrail_name,
)
def build_stream_error_items(
self,
@ -1043,6 +1054,71 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
task_mappings=[(target_choice_index, None) for _ in changed], # mutable-ok: callee takes lists
)
@staticmethod
def _function_tool_call_shapes(response: "ModelResponse") -> tuple[tuple[str | None, str], ...]:
return tuple(
(tool_call.function.name, tool_call.function.arguments)
for choice in response.choices
for tool_call in choice.message.tool_calls or ()
if isinstance(tool_call, ChatCompletionMessageToolCall)
)
@staticmethod
def _function_tool_call_fragments(
responses_so_far: Sequence["ModelResponseStream"],
) -> tuple[tuple[ChatCompletionDeltaToolCall, ...], ...]:
"""Group the stream's function tool-call fragments by their tool-call index, in
the index order ``stream_chunk_builder`` lists the rebuilt tool calls, keeping
only the indices the builder keeps (an id and a name somewhere in the stream)."""
fragments: Final = tuple(
tool_call
for response in responses_so_far
for choice in response.choices
for tool_call in choice.delta.tool_calls or ()
if isinstance(tool_call, ChatCompletionDeltaToolCall)
)
identified: Final = frozenset(fragment.index for fragment in fragments if fragment.id)
named: Final = frozenset(fragment.index for fragment in fragments if fragment.function.name)
return tuple(
tuple(fragment for fragment in fragments if fragment.index == index) for index in sorted(identified & named)
)
def _write_ended_stream_tool_call_rewrites(
self,
responses_so_far: list["ModelResponseStream"], # mutable-ok: rewrites the caller's buffered chunks in place
guardrailed_response: "ModelResponse",
pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...],
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail tool-call rewrites back across the buffered
chunks: the rewritten name and full arguments land in the tool call's first
fragment and the arguments of its later fragments are blanked, mirroring the
text write-back. A rewrite on a stream carrying more than one distinct choice
index, or whose fragments do not line up with the rebuilt tool calls, is
reported as undeliverable, so the pipeline executor discards it and releases
the original chunks."""
post_guardrail_tool_calls: Final = self._function_tool_call_shapes(guardrailed_response)
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
stream_choice_indices: Final = frozenset(
choice.index for response in responses_so_far for choice in response.choices
)
fragments_by_tool_call: Final = self._function_tool_call_fragments(responses_so_far)
if len(stream_choice_indices) != 1 or len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
for before, (name, arguments), fragments in zip(
pre_guardrail_tool_calls, post_guardrail_tool_calls, fragments_by_tool_call
):
if (name, arguments) == before:
continue
head, *tail = fragments
head.function.name = name
head.function.arguments = arguments
for fragment in tail:
fragment.function.arguments = ""
async def _apply_guardrail_responses_to_output_streaming(
self,
responses: list["ModelResponseStream"],

View file

@ -101,6 +101,18 @@ if TYPE_CHECKING:
from litellm.types.llms.openai import ResponseInputParam
class _ToolCallShape(NamedTuple):
name: str | None
arguments: str
def _tool_call_shapes(tool_calls: Sequence[ChatCompletionToolCallChunk]) -> tuple[_ToolCallShape, ...]:
return tuple(
_ToolCallShape(name=tool_call["function"].get("name"), arguments=tool_call["function"].get("arguments", ""))
for tool_call in tool_calls
)
class ResponseOutputEnvelope(TypedDict, total=False):
"""Dict form of a Responses API response, as far as guardrail write-back reads it."""
@ -340,7 +352,7 @@ class OpenAIResponsesHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
"""
@ -754,6 +766,7 @@ class OpenAIResponsesHandler(BaseTranslation):
if response_model:
inputs["model"] = response_model
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,
@ -762,6 +775,12 @@ class OpenAIResponsesHandler(BaseTranslation):
)
guardrailed_texts: Final = guardrailed_inputs.get("texts", [])
returned_tool_calls: Final = guardrailed_inputs.get("tool_calls")
post_guardrail_tool_calls: Final = _tool_call_shapes(
returned_tool_calls
if isinstance(returned_tool_calls, list) and len(returned_tool_calls) == len(tool_calls_to_check)
else tool_calls_to_check
)
# Write guardrailed texts back into the output items in-place.
# final_chunk is a reference into responses_so_far so this
@ -784,6 +803,13 @@ class OpenAIResponsesHandler(BaseTranslation):
stream_events=responses_so_far[:-1],
rewrites_by_position=rewrites_by_position,
)
self._deliver_ended_stream_tool_call_rewrites(
responses_so_far=responses_so_far,
outputs=outputs,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=post_guardrail_tool_calls,
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
return responses_so_far
# ------------------------------------------------------------------ #
@ -894,6 +920,89 @@ class OpenAIResponsesHandler(BaseTranslation):
continue
OpenAIResponsesHandler._write_event_field(content[content_idx], "text", rewritten)
def _deliver_ended_stream_tool_call_rewrites(
self,
responses_so_far: Sequence[object],
outputs: Sequence[object],
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail tool-call rewrites into the completed
envelope's ``function_call`` items and sync the earlier stream events,
keyed by ``output_index``. The guardrail sees the envelope's function
calls in output order, which is how a rewritten call finds its item; a
rewrite whose calls do not line up with the envelope is reported as
undeliverable, so the pipeline executor discards it and releases the
original events."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
function_call_indices: Final = tuple(
output_idx
for output_idx, output_item in enumerate(outputs)
if stream_item_field(output_item, "type") == "function_call"
)
if len(function_call_indices) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
rewrites_by_output_index: Final = MappingProxyType(
{
output_idx: after
for output_idx, before, after in zip(
function_call_indices, pre_guardrail_tool_calls, post_guardrail_tool_calls
)
if after != before
}
)
for output_idx, rewrite in rewrites_by_output_index.items():
self._write_function_call_item(outputs[output_idx], rewrite.name, rewrite.arguments)
self._sync_stream_events_with_tool_call_rewrites(
stream_events=responses_so_far[:-1],
rewrites_by_output_index=rewrites_by_output_index,
)
def _sync_stream_events_with_tool_call_rewrites(
self,
stream_events: Sequence[object],
rewrites_by_output_index: Mapping[int, _ToolCallShape],
) -> None:
"""Sync pre-completion function-call events with the rewritten completed
response: the first ``function_call_arguments.delta`` for a rewritten call
carries the full rewritten arguments and the rest are blanked, while
``function_call_arguments.done`` and ``output_item.done`` carry the full
rewritten arguments and ``output_item.added`` / ``output_item.done`` the
rewritten name, so every event a client may read agrees with the
rewritten ``response.completed`` payload."""
delta_replacements: Final = MappingProxyType(
{index: chain((rewrite.arguments,), repeat("")) for index, rewrite in rewrites_by_output_index.items()}
)
for event in stream_events:
output_index = stream_item_field(event, "output_index")
if not isinstance(output_index, int) or output_index not in rewrites_by_output_index:
continue
rewrite = rewrites_by_output_index[output_index]
match stream_item_field(event, "type"):
case "response.function_call_arguments.delta":
self._write_event_field(event, "delta", next(delta_replacements[output_index]))
case "response.function_call_arguments.done":
self._write_event_field(event, "arguments", rewrite.arguments)
case "response.output_item.added":
self._write_function_call_item(stream_item_field(event, "item"), rewrite.name, None)
case "response.output_item.done":
self._write_function_call_item(stream_item_field(event, "item"), rewrite.name, rewrite.arguments)
case _:
pass
@staticmethod
def _write_function_call_item(item: object, name: str | None, arguments: str | None) -> None:
if not (isinstance(item, dict) or hasattr(item, "get")):
return
if name is not None:
OpenAIResponsesHandler._write_event_field(item, "name", name)
if arguments is not None:
OpenAIResponsesHandler._write_event_field(item, "arguments", arguments)
def _check_streaming_has_ended(self, responses_so_far: Sequence[object]) -> bool:
"""
Check if the streaming has ended.

View file

@ -85,10 +85,10 @@ def _logged_by_inner_guardrail(method: _GuardrailMethodT) -> _GuardrailMethodT:
class _StreamRewriteObserver(CustomGuardrail):
"""Stand-in handed to the endpoint translation in place of a streaming pipeline step's
guardrail. It records whether the guardrail returned different output than it was given,
which for guardrails like Bedrock's ANONYMIZED action is only known at runtime. Text
rewrites are deliverable on translations that write them back across the buffered chunks
(``delivers_ended_stream_text_rewrites``); tool-call rewrites and text rewrites on any
other translation are discarded by the executor, which releases the original chunks.
which for guardrails like Bedrock's ANONYMIZED action is only known at runtime. Text and
tool-call rewrites are deliverable on translations that write them back across the
buffered chunks (``delivers_ended_stream_rewrites``); rewrites on any other translation
are discarded by the executor, which releases the original chunks.
The inner guardrail's ``apply_guardrail`` already records the guardrail information
and span, so the observer's stays out of ``log_guardrail_information``."""
@ -290,13 +290,13 @@ class PipelineExecutor:
litellm_logging_obj: "LiteLLMLoggingObj | None",
) -> None:
"""Run one streaming post_call step through the endpoint translation, delivering
text rewrites on translations that support ended-stream write-back. A rewrite that
cannot reach the client yet (a tool-call rewrite, a text rewrite on a translation
without write-back, or one the translation refused with
``UndeliverableStreamRewrite``) is discarded: the buffered chunks go back to the
originals and the step passes, so the client gets the stream the merge base sent."""
text and tool-call rewrites on translations that support ended-stream write-back. A
rewrite that cannot reach the client yet (one on a translation without write-back, or
one the translation refused with ``UndeliverableStreamRewrite``) is discarded: the
buffered chunks go back to the originals and the step passes, so the client gets the
stream the merge base sent."""
observer: Final = _StreamRewriteObserver(callback)
deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_text_rewrites
deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_rewrites
originals: Final = copy.deepcopy(streaming_chunks)
try:
if deliver_rewrites:
@ -319,7 +319,7 @@ class PipelineExecutor:
except UndeliverableStreamRewrite:
_release_original_chunks(step.guardrail, streaming_chunks, originals)
else:
if observer.rewrote_tool_calls or (observer.rewrote_texts and not deliver_rewrites):
if not deliver_rewrites and (observer.rewrote_texts or observer.rewrote_tool_calls):
_release_original_chunks(step.guardrail, streaming_chunks, originals)
if not callback.records_own_guardrail_information:
add_guardrail_to_applied_guardrails_header(request_data=hook_input, guardrail_name=step.guardrail)

View file

@ -3445,11 +3445,11 @@ class ProxyLogging:
assembled output through the endpoint guardrail translation, the same
machinery flat post_call guardrails use at end of stream. An allow
releases the buffered chunks: verbatim when no guardrail rewrote the
output, rewritten in place when one rewrote text and the translation
delivers ended-stream rewrites (later steps then re-scan the rewritten
chunks, so rewrites chain). A rewrite the translation cannot deliver
yet (a tool-call rewrite, or a text rewrite on a route without
write-back) is discarded by the executor and the original chunks are
output, rewritten in place when one rewrote text or a tool call and the
translation delivers ended-stream rewrites (later steps then re-scan the
rewritten chunks, so rewrites chain). A rewrite the translation cannot
deliver yet (one on a route without write-back, or a shape the route
refuses) is discarded by the executor and the original chunks are
released, as is a buffered shape no translation resolves; a block or
modify_response terminates with the translation's block chunks or the
raised error.

View file

@ -315,6 +315,72 @@ class TestAnthropicMessagesHandlerStreamingOutputProcessing:
assert "event: message_start" in raw and "event: message_stop" in raw
assert '"stop_reason": "end_turn"' in raw
@staticmethod
def _ended_tool_use_sse_chunks() -> list:
events = [
("message_start", {"type": "message_start", "message": {"id": "msg_1", "type": "message", "role": "assistant", "model": "claude-sonnet-4-5", "content": [], "stop_reason": None, "usage": {"input_tokens": 1, "output_tokens": 0}}}),
("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "tool_use", "id": "toolu_1", "name": "lookup_fruit", "input": {}}}),
("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "input_json_delta", "partial_json": ""}}),
("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "input_json_delta", "partial_json": '{"fruit": "persim'}}),
("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "input_json_delta", "partial_json": 'mon"}'}}),
("content_block_stop", {"type": "content_block_stop", "index": 0}),
("message_delta", {"type": "message_delta", "delta": {"stop_reason": "tool_use", "stop_sequence": None}, "usage": {"output_tokens": 2}}),
("message_stop", {"type": "message_stop"}),
]
return [f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode() for name, payload in events]
@staticmethod
def _argument_masking_guardrail() -> CustomGuardrail:
class MaskArguments(CustomGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
for tool_call in inputs.get("tool_calls", []):
tool_call.function.arguments = '{"fruit": "[MASKED]"}'
return inputs
return MaskArguments(guardrail_name="test")
@staticmethod
def _partial_jsons(chunks: list) -> list:
return [
json.loads(line[len("data:") :].strip())["delta"]["partial_json"]
for chunk in chunks
for line in chunk.decode().split("\n")
if line.startswith("data:") and json.loads(line[len("data:") :].strip()).get("type") == "content_block_delta"
]
@pytest.mark.asyncio
async def test_deliver_ended_stream_rewrites_writes_tool_use_input_back_into_sse_chunks(self):
handler = AnthropicMessagesHandler()
chunks = self._ended_tool_use_sse_chunks()
result = await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=self._argument_masking_guardrail(),
litellm_logging_obj=MagicMock(),
deliver_ended_stream_rewrites=True,
)
assert result is chunks
assert self._partial_jsons(chunks) == ['{"fruit": "[MASKED]"}', "", ""]
raw = b"".join(chunks).decode()
assert '"name": "lookup_fruit"' in raw and '"id": "toolu_1"' in raw
assert '"stop_reason": "tool_use"' in raw
assert "persim" not in raw
@pytest.mark.asyncio
async def test_ended_stream_tool_use_rewrite_leaves_chunks_untouched_by_default(self):
handler = AnthropicMessagesHandler()
chunks = self._ended_tool_use_sse_chunks()
original = [bytes(chunk) for chunk in chunks]
await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=self._argument_masking_guardrail(),
litellm_logging_obj=MagicMock(),
)
assert chunks == original
@pytest.mark.asyncio
async def test_ended_stream_rewrite_leaves_chunks_untouched_by_default(self):
handler = AnthropicMessagesHandler()

View file

@ -1113,6 +1113,79 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
assert chunks[1].choices[0].delta.content in (None, "")
assert chunks[1].choices[0].finish_reason == "stop"
@staticmethod
def _ended_tool_call_stream_chunks() -> list:
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
Delta,
Function,
ModelResponseStream,
StreamingChoices,
)
def chunk(tool_call: ChatCompletionDeltaToolCall | None, finish_reason: Optional[str] = None):
return ModelResponseStream(
id="chatcmpl-123",
created=1234567890,
model="gpt-4",
object="chat.completion.chunk",
choices=[
StreamingChoices(
index=0,
delta=Delta(tool_calls=[tool_call] if tool_call else None),
finish_reason=finish_reason,
)
],
)
def fragment(arguments: str, name: Optional[str] = None, call_id: Optional[str] = None):
return ChatCompletionDeltaToolCall(
id=call_id, index=0, type="function", function=Function(name=name, arguments=arguments)
)
return [
chunk(fragment("", name="lookup_fruit", call_id="call_1")),
chunk(fragment('{"fruit":')),
chunk(fragment(' "persimmon"}')),
chunk(None, finish_reason="tool_calls"),
]
@pytest.mark.asyncio
async def test_deliver_ended_stream_rewrites_writes_tool_call_arguments_back_into_chunks(self):
handler = OpenAIChatCompletionsHandler()
guardrail = MockGuardrail(guardrail_name="test")
chunks = self._ended_tool_call_stream_chunks()
result = await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=guardrail,
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
assert result is chunks
fragments = [chunk.choices[0].delta.tool_calls for chunk in chunks[:3]]
assert [fragment[0].function.arguments for fragment in fragments] == ['{"fruit": "PERSIMMON"}', "", ""]
assert fragments[0][0].function.name == "lookup_fruit"
assert fragments[0][0].id == "call_1"
assert chunks[3].choices[0].delta.tool_calls is None
assert chunks[3].choices[0].finish_reason == "tool_calls"
@pytest.mark.asyncio
async def test_ended_stream_tool_call_rewrite_leaves_chunks_untouched_by_default(self):
handler = OpenAIChatCompletionsHandler()
guardrail = MockGuardrail(guardrail_name="test")
chunks = self._ended_tool_call_stream_chunks()
await handler.process_output_streaming_response(
responses_so_far=chunks,
guardrail_to_apply=guardrail,
litellm_logging_obj=None,
)
fragments = [chunk.choices[0].delta.tool_calls for chunk in chunks[:3]]
assert [fragment[0].function.arguments for fragment in fragments] == ["", '{"fruit":', ' "persimmon"}']
@pytest.mark.asyncio
async def test_ended_stream_rewrite_leaves_chunks_untouched_by_default(self):
handler = OpenAIChatCompletionsHandler()

View file

@ -1195,6 +1195,94 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
assert events[4]["item"]["content"][0]["text"] == "hello [MASKED]"
assert events[5]["response"]["output"][0]["content"][0]["text"] == "hello [MASKED]"
@staticmethod
def _ended_function_call_stream_events() -> List[dict]:
def item(arguments: str, status: str) -> dict:
return {
"type": "function_call",
"id": "fc_123",
"call_id": "call_123",
"name": "lookup_fruit",
"arguments": arguments,
"status": status,
}
return [
{"type": "response.output_item.added", "output_index": 0, "item": item("", "in_progress")},
{"type": "response.function_call_arguments.delta", "item_id": "fc_123", "output_index": 0, "delta": '{"fruit":'},
{"type": "response.function_call_arguments.delta", "item_id": "fc_123", "output_index": 0, "delta": ' "persimmon"}'},
{
"type": "response.function_call_arguments.done",
"item_id": "fc_123",
"output_index": 0,
"arguments": '{"fruit": "persimmon"}',
},
{"type": "response.output_item.done", "output_index": 0, "item": item('{"fruit": "persimmon"}', "completed")},
{
"type": "response.completed",
"response": {
"id": "resp_123",
"model": "gpt-4o",
"output": [item('{"fruit": "persimmon"}', "completed")],
"status": "completed",
},
},
]
@staticmethod
def _argument_masking_guardrail() -> CustomGuardrail:
class MaskArguments(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
tool_calls = [
{**tool_call, "function": {**tool_call["function"], "arguments": '{"fruit": "[MASKED]"}'}}
for tool_call in inputs.get("tool_calls", [])
]
return {**inputs, "tool_calls": tool_calls}
return MaskArguments(guardrail_name="test-mask-arguments")
@pytest.mark.asyncio
async def test_deliver_ended_stream_rewrites_syncs_function_call_events(self):
handler = OpenAIResponsesHandler()
events = self._ended_function_call_stream_events()
result = await handler.process_output_streaming_response(
responses_so_far=events,
guardrail_to_apply=self._argument_masking_guardrail(),
litellm_logging_obj=None,
deliver_ended_stream_rewrites=True,
)
assert result is events
assert events[0]["item"]["arguments"] == ""
assert events[1]["delta"] == '{"fruit": "[MASKED]"}'
assert events[2]["delta"] == ""
assert events[3]["arguments"] == '{"fruit": "[MASKED]"}'
assert events[4]["item"]["arguments"] == '{"fruit": "[MASKED]"}'
assert events[5]["response"]["output"][0]["arguments"] == '{"fruit": "[MASKED]"}'
assert events[5]["response"]["output"][0]["name"] == "lookup_fruit"
@pytest.mark.asyncio
async def test_ended_stream_function_call_rewrite_leaves_events_untouched_by_default(self):
handler = OpenAIResponsesHandler()
events = self._ended_function_call_stream_events()
await handler.process_output_streaming_response(
responses_so_far=events,
guardrail_to_apply=self._argument_masking_guardrail(),
litellm_logging_obj=None,
)
assert events[1]["delta"] == '{"fruit":'
assert events[3]["arguments"] == '{"fruit": "persimmon"}'
assert events[5]["response"]["output"][0]["arguments"] == '{"fruit": "persimmon"}'
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_type", ["response.incomplete", "response.failed"])
async def test_deliver_ended_stream_rewrites_syncs_non_completed_terminals(self, terminal_type):

View file

@ -924,7 +924,7 @@ class _TextReturningGuardrail(CustomGuardrail):
class _TextTranslation:
delivers_ended_stream_text_rewrites = False
delivers_ended_stream_rewrites = False
def __init__(self):
self.seen_guardrail_names = []
@ -946,7 +946,7 @@ class _WritingTranslation:
"""Writes the guardrail's text (and tool-call) outputs back into the buffered chunks the way the
chat/Responses/Messages handlers do on an ended stream."""
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
async def process_output_streaming_response(
self,
@ -970,7 +970,7 @@ class _WritingTranslation:
class _RefusingTranslation:
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
async def process_output_streaming_response(
self,
@ -1088,13 +1088,27 @@ async def test_streaming_step_delivers_text_rewrite_through_writing_translation(
@pytest.mark.asyncio
async def test_streaming_step_discards_tool_call_rewrite_and_restores_written_text(monkeypatch, caplog):
async def test_streaming_step_delivers_tool_call_rewrite_through_writing_translation(monkeypatch, caplog):
monkeypatch.setattr(litellm, "callbacks", [_TextAndToolCallRewritingGuardrail(rewrite_tool_call=True)])
chunks = [_chunk()]
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
result = await _run_streaming_step(_WritingTranslation(), chunks)
assert result.terminal_action == "allow"
assert chunks[0]["text"] == "hello [MASKED]"
assert chunks[0]["tool_call"]["function"]["arguments"] == '{"ssn": "[MASKED]"}'
assert not any("discarded" in record.getMessage() for record in caplog.records)
@pytest.mark.asyncio
async def test_streaming_step_discards_tool_call_rewrite_when_translation_lacks_write_back(monkeypatch, caplog):
monkeypatch.setattr(litellm, "callbacks", [_TextAndToolCallRewritingGuardrail(rewrite_tool_call=True)])
chunks = [_chunk()]
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
result = await _run_streaming_step(_TextTranslation(), chunks)
_assert_passed_with_discard_warning(result, caplog)
assert chunks == [_chunk()]

View file

@ -1836,7 +1836,7 @@ def _echoed_tool_call_dicts(arguments: str) -> List[Dict[str, Any]]:
@pytest.mark.asyncio
@pytest.mark.parametrize("on_fail, on_error", [("block", None), ("next", "next")])
async def test_streaming_iterator_hook_pipeline_releases_originals_on_runtime_tool_call_rewrite(
async def test_streaming_iterator_hook_pipeline_delivers_runtime_tool_call_rewrite(
proxy_logging, make_user_api_key_auth, monkeypatch, on_fail, on_error, caplog
):
transform = lambda inputs: {"tool_calls": _echoed_tool_call_dicts('{"ssn": "[MASKED]"}')} # noqa: E731
@ -1855,9 +1855,12 @@ async def test_streaming_iterator_hook_pipeline_releases_originals_on_runtime_to
delivered.append(item)
assert len(delivered) == 2
assert delivered[0].choices[0].delta.tool_calls[0].function.arguments == '{"ssn": "123"}'
delivered_tool_call = delivered[0].choices[0].delta.tool_calls[0]
assert delivered_tool_call.function.arguments == '{"ssn": "[MASKED]"}'
assert delivered_tool_call.function.name == "lookup"
assert delivered_tool_call.id == "call_1"
assert delivered[1].choices[0].finish_reason == "tool_calls"
assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog))
assert not any("discarded" in message for message in _warnings(caplog))
@pytest.mark.asyncio