mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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.
This commit is contained in:
parent
12e3b14a03
commit
f1a689059f
2 changed files with 822 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue