From f1a689059f4ce5c22a46857e5206a8bd138a0a0e Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 09:12:40 -0500 Subject: [PATCH] feat(guardrails): restore llm_shield_proxy placeholders on native streams Anthropic /v1/messages and /v1/responses streams have no `choices`, so the streaming hook passed them through with placeholders still in them. Both are now restored incrementally, with the same per-stream windows as chat: - /v1/messages arrives as raw SSE. Frames are cut at event boundaries, text_delta and input_json_delta are restored per block index, and held text is emitted as one more delta ahead of content_block_stop. Signed thinking deltas, frames from other endpoints and non-SSE raw streams pass through unchanged. - /v1/responses events are restored per item and part. Held text goes out as a copy of the stream's last delta before its .done event, and the events that repeat the reply (.done, content_part.done, output_item.done, response.completed) are restored in full. The request side now also redacts Anthropic tool_use inputs and Responses reasoning summaries, and sends tool and function descriptions (including parameter schema descriptions) and the user / safety_identifier fields to the non-restorable vault, like system prompts. Tool results stay restorable: the model reads them to answer, so restoring them returns what the caller would have seen without the guardrail. --- .../llm_shield_proxy/llm_shield_proxy.py | 468 +++++++++++++++++- .../guardrail_hooks/test_llm_shield_proxy.py | 371 ++++++++++++++ 2 files changed, 822 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 8e3e6d81dbc..ebafb10831a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -6,9 +6,14 @@ # +-------------------------------------------------------------+ import copy +import functools +import json import os +import re import uuid -from collections.abc import AsyncGenerator, Callable, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence +from enum import Enum +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ @@ -80,8 +85,54 @@ JsonBody: TypeAlias = dict # bound is what stops a crafted one from becoming an unbounded walk. _MAX_CONTENT_DEPTH: Final = 8 +# How far a tool's parameter schema is followed. Deeper than content: every nested +# object costs two levels (`properties`, then the property), and a description missed +# here goes to the provider in the clear. +_MAX_SCHEMA_DEPTH: Final = 32 + _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. +# One incremental rehydration step for a stream the caller has already bound to its +# vault: (new text, carried window, final) -> (text safe to emit, window still held). +_StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] # mutable-ok: Callable's param list. + +# A batch rehydration already bound to the request's vault. +_Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] # mutable-ok: Callable's param list. + +# Anthropic /v1/messages delta types that carry restorable text, and the field holding +# it. `thinking_delta` is left out on purpose: a thinking block is signed, and one +# rewritten here fails verification when the client sends it back on the next turn. +_ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_json_delta": "partial_json"}) + +# A blank line ends an SSE event. Frames are cut there, never inside an event, so a +# `data:` line split across two network chunks is parsed only once it is whole. +_SSE_EVENT_BOUNDARY: Final = re.compile(rb"(\r?\n\r?\n)") + +# Responses API events whose `delta` is model text. Each belongs to the stream that the +# matching `.done` event in `_RESPONSES_DONE_FIELDS` closes. +_RESPONSES_DELTA_EVENTS: Final = frozenset( + ( + "response.output_text.delta", + "response.refusal.delta", + "response.function_call_arguments.delta", + "response.reasoning_summary_text.delta", + ) +) + +# The `.done` event that closes each delta stream, and the field that repeats the +# stream's full text on it. +_RESPONSES_DONE_FIELDS: Final = MappingProxyType( + { + "response.output_text.done": "text", + "response.refusal.done": "refusal", + "response.function_call_arguments.done": "arguments", + "response.reasoning_summary_text.done": "text", + } +) + +# Terminal Responses API events that repeat the whole reply under `response`. +_RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) + # The accumulator the collectors below append into. It never escapes # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. @@ -158,6 +209,11 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None: continue # Image and audio parts have no text and fall through untouched. _collect(part, "text", slots) + if part.get("type") == "tool_use": + # A replayed Anthropic tool call. Its `input` is a JSON object rather than + # a string, so a value can sit at any depth -- the reply side walks the + # same leaves when it restores one. + _collect_json_leaves(part.get("input"), slots, depth + 1) if "content" in part: pending.append((part, depth + 1)) @@ -220,6 +276,72 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged # A function_call item holds `arguments`; a function_call_output holds `output`. _collect(item, "arguments", slots) _collect(item, "output", slots) + # A replayed reasoning item carries the model's summary of its own reasoning, + # which quotes whatever the conversation contained. + _collect_text_parts(item, "summary", slots) + + +def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) -> None: + """Collects the `text` of every part in the list held at `key`.""" + parts: Final = container.get(key) + for part in parts if isinstance(parts, list) else (): + if isinstance(part, dict): + _collect(part, "text", slots) + + +def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> None: + """Tool definitions are application-authored free text bound for the provider. + + A description -- on the tool, or on any property of its parameter schema -- is where + callers put examples and customer context, so it carries PII as often as a prompt + does. It is collected into the privileged sink, like a system prompt: redacted + outbound, and never restorable from the reply. Names, types and enum values are left + as sent, because the model has to reproduce them exactly for a call to route. + + Covers Chat `tools[].function`, the legacy `functions[]`, and the flat tool shape the + Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. + """ + for key in ("tools", "functions"): + declared = data.get(key) + for tool in declared if isinstance(declared, list) else (): + if not isinstance(tool, dict): + continue + function = tool.get("function") + for holder in (tool, function) if isinstance(function, dict) else (tool,): + _collect(holder, "description", privileged) + _collect_schema_descriptions(holder.get("parameters"), privileged) + _collect_schema_descriptions(holder.get("input_schema"), privileged) + + +def _collect_schema_descriptions(schema: object, privileged: _SlotSink) -> None: + """Collects every string `description` in a JSON schema, at any depth. + + Only `description` is free text. A property that is itself *named* "description" + holds a schema object rather than a string, so it is descended into, not collected. + Walked with an explicit stack and a depth bound, like the other request walks. + """ + pending: Final[list] = [(schema, 0)] # mutable-ok: local walk stack. + while pending: + node, depth = pending.pop() + if depth > _MAX_SCHEMA_DEPTH: + continue + if isinstance(node, dict): + _collect(node, "description", privileged) + pending.extend((value, depth + 1) for value in node.values() if isinstance(value, (dict, list))) + elif isinstance(node, list): + pending.extend((value, depth + 1) for value in node if isinstance(value, (dict, list))) + + +def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None: + """`user` and `safety_identifier` are forwarded to the provider and often hold an email. + + Only detected PII is replaced, so an opaque id reaches the provider unchanged. LiteLLM's + own end-user spend tracking reads the id resolved at authentication, before this hook + runs, so rewriting the field here does not move spend. Nothing restores these from a + reply, hence the privileged sink. + """ + _collect(data, "user", privileged) + _collect(data, "safety_identifier", privileged) def _choice_index(choice: object) -> int: @@ -240,6 +362,15 @@ def _read_field(holder: object, name: str) -> object: return getattr(holder, name, None) +def _read_list(holder: object, name: str) -> Sequence[object]: + """Reads a list field from a dict or an object; anything else reads as empty. + + The entries are the reply's own objects, so writing through them edits the reply. + """ + value: Final = _read_field(holder, name) + return tuple(value) if isinstance(value, (list, tuple)) else () + + def _write_field(holder: object, name: str, value: str) -> None: """Writes one string field back into a dict or an object. Pairs with _read_field.""" if isinstance(holder, dict): @@ -298,6 +429,289 @@ def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: return [{"index": tool_index, "function": {"arguments": text}}] # mutable-ok: delta.tool_calls is a list. +def _collect_response_item(item: object, slots: _SlotSink) -> None: + """Restorable spans in one Responses API output item, dict or object. + + Mirrors `_collect_responses_fields` on the request side -- a function_call item holds + `arguments`, a function_call_output holds `output`, a reasoning item holds `summary` + parts -- so the two directions stay symmetric. + """ + for block in _read_list(item, "content"): + for field in ("text", "refusal"): + text = _read_field(block, field) + if isinstance(text, str) and text: + slots.append((text, lambda new, b=block, f=field: _write_field(b, f, new))) + for part in _read_list(item, "summary"): + text = _read_field(part, "text") + if isinstance(text, str) and text: + slots.append((text, lambda new, p=part: _write_field(p, "text", new))) + for field in ("arguments", "output"): + value = _read_field(item, field) + if isinstance(value, str) and value: + slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) + + +async def _rehydrate_slots(slots: Sequence[_Slot], rehydrate: _Rehydrate) -> None: + """Restores every span in `slots` in one batch and writes each result back.""" + if not slots: + return + restored: Final = await rehydrate(tuple(text for text, _ in slots)) + for (_, write), replacement in zip(slots, restored): + write(replacement) + + +def _responses_event_type(chunk: object) -> str | None: + """The event type of a Responses API stream event, or None for any other chunk. + + The type arrives as a plain string on dicts and as a str-valued Enum on LiteLLM's + event models. The Enum is unwrapped because it does not hash like its value, so it + would miss every lookup in the event tables above. + """ + if isinstance(chunk, (bytes, str)): + return None + kind: Final = _read_field(chunk, "type") + value: Final = kind.value if isinstance(kind, Enum) else kind + return value if isinstance(value, str) and value.startswith("response.") else None + + +class _AnthropicSSERestorer: + """Restores an Anthropic `/v1/messages` stream, which reaches the hook as raw SSE. + + Each content block is its own token stream with its own window, keyed by the block's + `index`: `text_delta` carries prose and `input_json_delta` a tool call's arguments. + When a block stops, whatever its window still holds is emitted as one more delta for + that block, just ahead of the `content_block_stop` frame, so the client has the whole + block before it is told the block is complete. + + Frames are processed whole. A network chunk can end in the middle of an event, so the + unfinished tail is kept until the rest arrives; that delays one partial event, never + a completed one. A frame that is not an Anthropic event -- another endpoint's SSE, or + anything that fails to parse -- is passed through byte for byte, and a raw stream that + does not open like SSE at all is passed through chunk by chunk, never buffered. + """ + + def __init__(self, step: _StreamStep) -> None: + self._step: Final = step + self._carries: Final[dict[int, str]] = {} # mutable-ok: per-block windows advanced in place. + self._delta_types: Final[dict[int, str]] = {} # mutable-ok: each block's delta type, for its flush. + self._pending = b"" # rebind-ok: the unfinished tail of the stream. + self._as_text = False # rebind-ok: set once if the stream arrives as str rather than bytes. + self._is_sse: bool | None = None # rebind-ok: decided once, by the stream's first chunk. + + async def feed(self, chunk: bytes | str) -> tuple[bytes | str, ...]: + """Restores every event this chunk completes; holds back an unfinished tail.""" + if isinstance(chunk, str): + self._as_text = True + raw: Final = chunk.encode("utf-8") if isinstance(chunk, str) else chunk + if self._is_sse is None and raw.strip(): + # An SSE stream opens with a field or a comment. Anything else (a JSON array + # streamed in pieces, say) has no event boundaries to wait for. + self._is_sse = raw.lstrip().startswith((b"event:", b"data:", b":")) + if not self._is_sse: + return (chunk,) + buffered: Final = self._pending + raw + boundaries: Final = tuple(_SSE_EVENT_BOUNDARY.finditer(buffered)) + if not boundaries: + self._pending = buffered + return () + cut: Final = boundaries[-1].end() + self._pending = buffered[cut:] + # With a capturing group, split alternates event, separator, ..., and ends in the + # empty remainder after the last separator. + parts: Final = _SSE_EVENT_BOUNDARY.split(buffered[:cut]) + restored: Final = tuple( + [ # mutable-ok: an await needs a list comprehension; frozen at once. + await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2) + ] + ) + return self._emit(b"".join(restored)) + + async def finish(self) -> tuple[bytes | str, ...]: + """Emits an unterminated final event and any window a block never closed.""" + tail: Final = await self._restore_event(self._pending) if self._pending.strip() else self._pending + self._pending = b"" + flushed: Final = await self._flush_all() + # The tail had no blank line after it; one is needed before another frame follows. + separator: Final = b"\n\n" if tail.strip() and flushed else b"" + return self._emit(tail + separator + flushed) + + def _emit(self, frames: bytes) -> tuple[bytes | str, ...]: + if not frames: + return () + return (frames.decode("utf-8") if self._as_text else frames,) + + async def _restore_event(self, block: bytes) -> bytes: + """Rewrites one SSE event, or returns it untouched if it carries nothing to restore.""" + try: + lines: Final = block.decode("utf-8").split("\n") + except UnicodeDecodeError: + return block + data_lines: Final = tuple(index for index, line in enumerate(lines) if line.startswith("data:")) + if len(data_lines) != 1: + return block + line: Final = lines[data_lines[0]] + try: + event: Final = json.loads(line[len("data:") :]) + except ValueError: + return block + if not isinstance(event, dict): + return block + kind: Final = event.get("type") + index: Final = event.get("index") + if kind == "content_block_stop" and isinstance(index, int): + return await self._flush(index) + block + if kind == "message_stop": + return await self._flush_all() + block + if kind != "content_block_delta" or not await self._restore_delta(event): + return block + ending: Final = "\r" if line.endswith("\r") else "" + rewritten: Final = ( + *lines[: data_lines[0]], + f"data: {json.dumps(event, ensure_ascii=False)}{ending}", + *lines[data_lines[0] + 1 :], + ) + return "\n".join(rewritten).encode("utf-8") + + async def _restore_delta(self, event: MutableRequest) -> bool: + """Advances one block's window through this delta. False if it holds no text.""" + index: Final = event.get("index") + delta: Final = event.get("delta") + if not isinstance(index, int) or not isinstance(delta, dict): + return False + delta_type: Final = delta.get("type") + if not isinstance(delta_type, str): + return False + field: Final = _ANTHROPIC_DELTA_FIELDS.get(delta_type) + text: Final = delta.get(field) if field is not None else None + if field is None or not isinstance(text, str) or not text: + return False + emitted, remaining = await self._step(text, self._carries.get(index, ""), False) + self._carries[index] = remaining + self._delta_types[index] = delta_type + delta[field] = emitted + return True + + async def _flush(self, index: int) -> bytes: + """One synthetic delta frame carrying whatever `index`'s window still holds.""" + carry: Final = self._carries.pop(index, "") + delta_type: Final = self._delta_types.pop(index, None) + field: Final = _ANTHROPIC_DELTA_FIELDS.get(delta_type) if isinstance(delta_type, str) else None + if not carry or field is None: + return b"" + text, _ = await self._step("", carry, True) + if not text: + return b"" + event: Final[JsonBody] = { # mutable-ok: serialised on the next line. + "type": "content_block_delta", + "index": index, + "delta": {"type": delta_type, field: text}, + } + return f"event: content_block_delta\ndata: {json.dumps(event, ensure_ascii=False)}\n\n".encode() + + async def _flush_all(self) -> bytes: + flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) # mutable-ok: frozen. + return b"".join(flushed) + + +class _ResponsesStreamRestorer: + """Restores a Responses API event stream. + + Every delta stream -- one output_text content part, one refusal, one function call's + arguments, one reasoning summary part -- gets its own window, keyed by the event + family, the item id and the part index. When its `.done` event arrives, whatever the + window still holds goes out first, as a copy of that stream's last delta event -- so + it carries the stream's own ids, and repeats that event's `sequence_number` -- and + the `.done` event's full text is then restored in one call. + + The events that repeat the reply wholesale -- `content_part.done`, + `output_item.done`, and `response.completed` / `response.incomplete` -- are restored + the same way the non-streaming reply is. + """ + + def __init__(self, step: _StreamStep, rehydrate: _Rehydrate) -> None: + self._step: Final = step + self._rehydrate: Final = rehydrate + self._carries: Final[dict[tuple, str]] = {} # mutable-ok: per-stream windows advanced in place. + self._last_deltas: Final[dict[tuple, object]] = {} # mutable-ok: newest delta per stream. + + async def restore(self, event: object) -> tuple[object, ...]: + """The events to emit in place of `event`: any flush, then the event itself.""" + kind: Final = _responses_event_type(event) + if kind is None: + return (event,) + if kind in _RESPONSES_DELTA_EVENTS: + await self._restore_delta(event, kind) + return (event,) + done_field: Final = _RESPONSES_DONE_FIELDS.get(kind) + if done_field is not None: + flushed: Final = await self._flush(_responses_stream_key(event, kind)) + await _rehydrate_slots(_field_slot(event, done_field), self._rehydrate) + return (*flushed, event) + slots: Final[_SlotSink] = [] # mutable-ok: accumulator, restored in one batch. + if kind == "response.content_part.done": + part: Final = _read_field(event, "part") + _collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view. + elif kind == "response.output_item.done": + _collect_response_item(_read_field(event, "item"), slots) + elif kind in _RESPONSES_TERMINAL_EVENTS: + for item in _read_list(_read_field(event, "response"), "output"): + _collect_response_item(item, slots) + await _rehydrate_slots(slots, self._rehydrate) + return (event,) + + async def finish(self) -> tuple[object, ...]: + """Flushes every stream the provider never closed, e.g. a truncated reply.""" + flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) # mutable-ok: frozen. + return tuple(event for events in flushed for event in events) + + async def _restore_delta(self, event: object, kind: str) -> None: + text: Final = _read_field(event, "delta") + if not isinstance(text, str) or not text: + return + key: Final = _responses_stream_key(event, kind) + emitted, remaining = await self._step(text, self._carries.get(key, ""), False) + self._carries[key] = remaining + self._last_deltas[key] = event + _write_field(event, "delta", emitted) + + async def _flush(self, key: tuple) -> tuple[object, ...]: + carry: Final = self._carries.pop(key, "") + template: Final = self._last_deltas.pop(key, None) + if not carry or template is None: + return () + text, _ = await self._step("", carry, True) + if not text: + return () + flush: Final = copy.deepcopy(template) + _write_field(flush, "delta", text) + return (flush,) + + +def _responses_stream_key(event: object, kind: str) -> tuple: + """Identifies the delta stream an event belongs to, the same for its delta and done. + + The family is the event type without its `.delta` / `.done` suffix, so an output_text + stream and a refusal stream on the same part never share a window. + """ + family: Final = kind.rsplit(".", 1)[0] + part_index: Final = _read_field(event, "content_index") + summary_index: Final = _read_field(event, "summary_index") + return ( + family, + _read_field(event, "item_id"), + _read_field(event, "output_index"), + part_index if part_index is not None else summary_index, + ) + + +def _field_slot(holder: object, field: str) -> Sequence[_Slot]: + """The one restorable span at `field` on `holder`, if it holds text.""" + text: Final = _read_field(holder, field) + if not isinstance(text, str) or not text: + return () + return ((text, lambda new: _write_field(holder, field, new)),) + + class LLMShieldProxyGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -443,9 +857,16 @@ class LLMShieldProxyGuardrail(CustomGuardrail): The split exists because the response is restored against one vault only. Server-authored spans -- system and developer turns, Anthropic's top-level - `system`, the Responses API `instructions` -- go into a vault nothing is - ever restored against, so a caller who gets the model to echo one of their - placeholders back receives the placeholder, not the value behind it. + `system`, the Responses API `instructions`, tool definitions -- go into a + vault nothing is ever restored against, so a caller who gets the model to + echo one of their placeholders back receives the placeholder, not the value + behind it. End-user identifiers go there too: nothing in a reply needs them. + + Tool *results* stay on the caller's side deliberately. The model reads them in + order to answer, so it can already repeat anything in them; restoring the + placeholder gives the caller the answer they would have had without this + guardrail, and an agent that reads a file and quotes an address from it needs + that address back. """ slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. privileged: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. @@ -458,6 +879,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): _collect_responses_fields(data, slots, privileged) _collect_prompt(data, slots) _collect_system(data, privileged) + _collect_tool_definitions(data, privileged) + _collect_end_user_ids(data, privileged) return tuple(slots), tuple(privileged) # --- hooks -------------------------------------------------------------------- @@ -609,23 +1032,14 @@ class LLMShieldProxyGuardrail(CustomGuardrail): fields on the request side -- a function_call item holds `arguments`, a function_call_output holds `output` -- so the two directions stay symmetric. """ - slots: Final[list] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. for item in getattr(response, "output", None) or (): - for block in getattr(item, "content", None) or (): - text = _read_field(block, "text") - if isinstance(text, str) and text: - slots.append((text, lambda new, b=block: _write_field(b, "text", new))) - for field in ("arguments", "output"): - value = _read_field(item, field) - if isinstance(value, str) and value: - slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) + _collect_response_item(item, slots) return tuple(slots) async def _restore_responses_api_response(self, response: Any, slots: Sequence[_Slot], data: MutableRequest) -> Any: """Puts the original values back into a Responses API reply.""" - restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) - for (_, write), replacement in zip(slots, restored): - write(replacement) + await _rehydrate_slots(slots, functools.partial(self._rehydrate, session_id=self._session_id(data))) return response async def async_post_call_streaming_iterator_hook( @@ -641,6 +1055,10 @@ class LLMShieldProxyGuardrail(CustomGuardrail): window would splice the characters held back for one stream onto another. The windows are locals of this generator, so they are scoped to a single stream and cannot leak between concurrent requests. + + The two native stream shapes have no `choices` and are restored by their own + walkers, with the same per-stream windows: Anthropic `/v1/messages` arrives as raw + SSE frames, and the Responses API as typed events. """ if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: async for chunk in response: @@ -648,16 +1066,32 @@ class LLMShieldProxyGuardrail(CustomGuardrail): return session_id: Final = self._session_id(request_data) + step: Final = functools.partial(self._stream_step, session_id=session_id) + rehydrate: Final = functools.partial(self._rehydrate, session_id=session_id) + sse: Final = _AnthropicSSERestorer(step) + events: Final = _ResponsesStreamRestorer(step, rehydrate) carries: Final[dict] = {} # mutable-ok: per-stream windows, local to this generator. last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: + if isinstance(chunk, (bytes, str)): + for frames in await sse.feed(chunk): + yield frames + continue + if _responses_event_type(chunk) is not None: + for event in await events.restore(chunk): + yield event + continue last_chunk = chunk for choice in getattr(chunk, "choices", None) or (): await self._restore_choice(choice, carries, session_id) yield chunk - # A stream that ended without a finish_reason can still leave text held back. + # A stream that ended early can still leave text held back, in any shape. + for frames in await sse.finish(): + yield frames + for event in await events.finish(): + yield event if last_chunk is not None and any(carries.values()): async for trailing in self._flush_trailing(last_chunk, carries, session_id): yield trailing diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 0197b83c9f0..a24b3845f9b 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -13,6 +13,12 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_shield_proxy.llm_shield_proxy ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import ( + FunctionCallArgumentsDeltaEvent, + OutputTextDeltaEvent, + OutputTextDoneEvent, + ResponsesAPIStreamEvents, +) from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices @@ -66,6 +72,81 @@ def _field(holder: object, name: str) -> object: return holder.get(name) if isinstance(holder, dict) else getattr(holder, name) +class _FakeShield: + """The three guard endpoints over one fixed vault, placeholder -> original. + + The stream endpoint holds back a trailing `[` that has not closed yet, which is the + behaviour that makes a placeholder split across two chunks come out whole. + """ + + def __init__(self, vault: dict[str, str]) -> None: + self.vault = vault + self.urls: list[str] = [] + + def _restore(self, text: str) -> str: + for placeholder, original in self.vault.items(): + text = text.replace(placeholder, original) + return text + + async def post(self, url: str, headers: dict, json: dict, timeout: float) -> Response: + self.urls.append(url) + if url.endswith("/rehydrate/stream"): + text = self._restore(json["carry"] + json["text"]) + opening = text.rfind("[") + if json["final"] or opening == -1 or "]" in text[opening:]: + return _response({"text": text, "carry": ""}) + return _response({"text": text[:opening], "carry": text[opening:]}) + return _response({"texts": [self._restore(text) for text in json["texts"]]}) + + +def _shielded(vault: dict[str, str]) -> tuple[LLMShieldProxyGuardrail, _FakeShield]: + guardrail = _guardrail(event_hook="post_call") + shield = _FakeShield(vault) + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + return guardrail, shield + + +def _sse(event: dict) -> bytes: + return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() + + +def _sse_events(frames: list) -> list[dict]: + """Parses emitted SSE output, whatever its chunking, back into event payloads.""" + raw = b"".join(frame.encode() if isinstance(frame, str) else frame for frame in frames).decode() + return [ + json.loads(line[len("data:") :]) + for event in raw.split("\n\n") + for line in event.split("\n") + if line.startswith("data:") + ] + + +def _text_block_stream(*deltas: str) -> list[bytes]: + """An Anthropic /v1/messages stream with one text block made of `deltas`.""" + return [ + _sse({"type": "message_start", "message": {"id": "msg_1", "role": "assistant", "content": []}}), + _sse({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + *( + _sse({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": d}}) + for d in deltas + ), + _sse({"type": "content_block_stop", "index": 0}), + _sse({"type": "message_stop"}), + ] + + +async def _restore_stream(guardrail: LLMShieldProxyGuardrail, chunks: list) -> list: + async def stream(): + for chunk in chunks: + yield chunk + + return await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + def test_llm_shield_guardrail_config(monkeypatch: pytest.MonkeyPatch): """Should register through init_guardrails_v2 like any other provider.""" monkeypatch.setattr(litellm, "guardrail_name_config_map", {}) @@ -479,6 +560,79 @@ class TestRequestCoverage: assert data["messages"][2]["tool_calls"][0]["function"]["arguments"] == "c" assert data["input"] == "d" + @pytest.mark.asyncio + async def test_anthropic_tool_use_input_is_redacted(self): + """A replayed tool_use block carries its arguments as a JSON object, not a string.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "messages": [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "t1", + "name": "send", + "input": {"to": "jane.doe@example.com", "meta": {"phone": "555-0100"}}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["jane.doe@example.com", "555-0100"] + block = data["messages"][0]["content"][0] + assert block["input"] == {"to": "[EMAIL_1]", "meta": {"phone": "[PHONE_1]"}} + assert block["name"] == "send", "the tool name has to arrive unchanged for the call to route" + + @pytest.mark.asyncio + async def test_responses_reasoning_summary_is_redacted(self): + """A replayed reasoning item quotes the conversation in its summary parts.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["user asked about [EMAIL_1]"]}) + + data = { + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "user asked about jane.doe@example.com"}], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["summary"][0]["text"] == "user asked about [EMAIL_1]" + + def test_tool_schemas_give_up_descriptions_and_nothing_else(self): + """Only free text is collected; names, types and enum values must reach the model.""" + data = { + "tools": [ + { + "type": "function", + "function": { + "name": "lookup", + "description": "top", + "parameters": { + "type": "object", + "properties": { + # A property that is itself named "description". + "description": {"type": "string", "description": "named"}, + "kind": {"type": "string", "enum": ["a", "b"], "description": "enum"}, + "deep": {"type": "array", "items": {"type": "object", "description": "nested"}}, + }, + }, + }, + } + ] + } + _, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert sorted(text for text, _ in privileged) == ["enum", "named", "nested", "top"] + class TestRestoration: @pytest.mark.asyncio @@ -653,6 +807,30 @@ class TestVaultIsolation: id="anthropic-top-level-system", ), pytest.param({"instructions": "S", "input": "U"}, id="responses-instructions"), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"type": "function", "function": {"name": "f", "description": "S"}}], + }, + id="chat-tool-description", + ), + pytest.param( + {"input": "U", "tools": [{"type": "function", "name": "f", "description": "S"}]}, + id="responses-tool-description", + ), + pytest.param( + {"messages": [{"role": "user", "content": "U"}], "functions": [{"name": "f", "description": "S"}]}, + id="legacy-function-description", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"name": "f", "input_schema": {"properties": {"to": {"description": "S"}}}}], + }, + id="anthropic-schema-description", + ), + pytest.param({"messages": [{"role": "user", "content": "U"}], "user": "S"}, id="end-user-id"), + pytest.param({"input": "U", "safety_identifier": "S"}, id="safety-identifier"), ], ) def test_server_authored_text_is_split_from_the_callers(self, data: dict) -> None: @@ -1016,3 +1194,196 @@ class TestApplyGuardrailToolCalls: assert merged["tool_calls"][0]["function"]["arguments"] == '{"email": "a@b.com"}' assert inputs["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + +class TestAnthropicStreamRestoration: + """/v1/messages streams reach the hook as raw SSE frames, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @pytest.mark.asyncio + async def test_split_placeholder_is_restored_and_never_fragmented(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMA", "IL_1] now")) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + + @pytest.mark.asyncio + async def test_held_text_lands_before_its_block_stops(self): + """A trailing `[` that never became a placeholder is still part of the answer.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMAIL_1], x = a[")) + + types = [e["type"] for e in _sse_events(out)] + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com, x = a[" + assert types.index("content_block_stop") > max(i for i, t in enumerate(types) if t == "content_block_delta") + + @pytest.mark.asyncio + async def test_events_split_across_network_chunks_are_restored(self): + """A chunk can end mid-event; the frame is parsed once it is whole.""" + guardrail, _ = _shielded(self.VAULT) + raw = b"".join(_text_block_stream("Mail [EMA", "IL_1] now")) + + out = await _restore_stream(guardrail, [raw[i : i + 7] for i in range(0, len(raw), 7)]) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + + @pytest.mark.asyncio + async def test_str_frames_stay_str(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [frame.decode() for frame in _text_block_stream("[EMAIL_1]")]) + + assert all(isinstance(frame, str) for frame in out) + assert "a@example.com" in "".join(out) + + @pytest.mark.asyncio + async def test_tool_input_json_is_restored(self): + guardrail, _ = _shielded(self.VAULT) + frames = [ + _sse({"type": "content_block_start", "index": 1, "content_block": {"type": "tool_use", "input": {}}}), + *( + _sse( + { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "input_json_delta", "partial_json": p}, + } + ) + for p in ('{"to": "[EMAI', 'L_1]"}') + ), + _sse({"type": "content_block_stop", "index": 1}), + ] + + out = await _restore_stream(guardrail, frames) + + partial = "".join(e["delta"]["partial_json"] for e in _sse_events(out) if e["type"] == "content_block_delta") + assert json.loads(partial) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_signed_thinking_and_foreign_frames_pass_through_byte_for_byte(self): + """Rewriting a signed thinking block breaks it; other frames are not ours to touch.""" + guardrail, shield = _shielded(self.VAULT) + thinking = {"type": "thinking_delta", "thinking": "[EMAIL_1]"} + frames = [ + _sse({"type": "content_block_delta", "index": 0, "delta": thinking}), + b'data: {"candidates": [{"content": {"parts": [{"text": "[EMAIL_1]"}]}}]}\n\n', + b"data: not json\n\n", + ] + + out = await _restore_stream(guardrail, frames) + + assert b"".join(out) == b"".join(frames) + assert shield.urls == [] + + @pytest.mark.asyncio + async def test_a_raw_stream_that_is_not_sse_is_never_buffered(self): + """Without event boundaries to wait for, buffering would hold the whole reply.""" + guardrail, _ = _shielded(self.VAULT) + chunks = [b'[{"candidates": []}', b', {"candidates": []}]'] + + out = await _restore_stream(guardrail, chunks) + + assert out == chunks + + +class TestResponsesStreamRestoration: + """/v1/responses streams are typed events, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @staticmethod + def _text_delta(delta: str, sequence_number: int, content_index: int = 0) -> OutputTextDeltaEvent: + return OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=content_index, + delta=delta, + sequence_number=sequence_number, + ) + + @pytest.mark.asyncio + async def test_deltas_and_done_text_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + done = OutputTextDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + item_id="msg_1", + output_index=0, + content_index=0, + text="Mail [EMAIL_1] x[", + ) + + out = await _restore_stream( + guardrail, [self._text_delta("Mail [EMA", 1), self._text_delta("IL_1] x[", 2), done] + ) + + deltas = [e.delta for e in out if isinstance(e, OutputTextDeltaEvent)] + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + assert "".join(deltas) == "Mail a@example.com x[" + assert out[-1].text == "Mail a@example.com x[", "the done event repeats the full, restored text" + assert isinstance(out[-2], OutputTextDeltaEvent), "held text must land before the done event" + + @pytest.mark.asyncio + async def test_function_call_arguments_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + FunctionCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id="fc_1", + output_index=1, + delta=part, + ) + for part in ('{"to": "[EMAI', 'L_1]"}') + ] + + out = await _restore_stream(guardrail, events) + + assert json.loads("".join(e.delta for e in out)) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_a_truncated_stream_still_flushes(self): + """No done event at all: whatever the window holds goes out at the end.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [self._text_delta("see [EMAIL_1] a[", 1)]) + + assert "".join(e.delta for e in out) == "see a@example.com a[" + + @pytest.mark.asyncio + async def test_completed_response_is_restored(self): + """The terminal event repeats the whole reply, and clients read it as the answer.""" + guardrail, _ = _shielded(self.VAULT) + block = {"type": "output_text", "text": "Mail [EMAIL_1]"} + call = SimpleNamespace(type="function_call", arguments='{"to": "[EMAIL_1]"}') + completed = SimpleNamespace( + type="response.completed", + response=SimpleNamespace(output=[SimpleNamespace(content=[block]), call]), + ) + + await _restore_stream(guardrail, [completed]) + + assert block["text"] == "Mail a@example.com" + assert call.arguments == '{"to": "a@example.com"}' + + @pytest.mark.asyncio + async def test_streams_on_different_parts_do_not_share_a_window(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + self._text_delta("one [EMA", 1), + self._text_delta("two", 2, content_index=1), + self._text_delta("IL_1]", 3), + ] + + out = await _restore_stream(guardrail, events) + + by_part: dict[int, str] = {} + for event in out: + by_part[event.content_index] = by_part.get(event.content_index, "") + event.delta + assert by_part == {0: "one a@example.com", 1: "two"}