refactor(guardrails): split llm_shield_proxy into payload, request walk and stream modules

The module had grown past 1,700 lines. Shared payload types and helpers move to
payload.py, the request walk to request_walk.py and the stream restorers to
stream_restorers.py; llm_shield_proxy.py keeps the guardrail class. No behaviour
change.
This commit is contained in:
Ninad Phalak 2026-10-05 09:58:48 +00:00
parent b55622e8ba
commit 464516a188
No known key found for this signature in database
4 changed files with 1063 additions and 983 deletions

View file

@ -0,0 +1,206 @@
import copy
import functools
from collections.abc import Awaitable, Callable, Sequence
from typing import (
Final,
TypeAlias,
)
# The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites
# the caller's payload in place, which is the entire point of the hook.
MutableRequest: TypeAlias = dict[str, object]
# A JSON body on its way to httpx, which requires a real dict rather than a view.
JsonBody: TypeAlias = dict[str, object]
# How far a JSON value -- a tool input, a parameter schema -- is followed on the request
# side. Legitimate JSON nests far deeper than content blocks do, so the bound is
# generous; past it the request is refused, since text past the bound would reach the
# provider unredacted.
MAX_JSON_DEPTH: Final = 64
# One redactable span: the text as it stands, and the write that puts the
# replacement back where it came from.
Slot: TypeAlias = tuple[str, Callable[[str], None]]
# 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]]]
# A batch rehydration already bound to the request's vault.
Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]]
# 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]
# A caller-owned list whose entries are rewritten in place, such as a Completions
# `prompt` sent as an array of strings.
MutableSeq: TypeAlias = list[object]
def as_object(value: object) -> MutableRequest | None:
"""`value` as a JSON object, or None.
`isinstance(value, dict)` alone leaves the keys and values unknown to the type
checker. A JSON object's keys are strings, so the type is stated once, here.
"""
return value if isinstance(value, dict) else None
def as_array(value: object) -> MutableSeq | None:
"""`value` as a JSON array, or None. See `as_object`."""
return value if isinstance(value, list) else None
def detached(value: object) -> object:
"""A deep copy of a reply or chunk, for restoring without touching LiteLLM's own object.
LiteLLM keeps the object it handed the hooks to fill its response cache and its
logs, so writing restored plaintext into that object would put it there too.
"""
return copy.deepcopy(value)
def is_container(value: object) -> bool:
"""Whether `value` is a JSON object or array, without narrowing it to unknown types."""
return isinstance(value, (dict, list))
def collect(container: MutableRequest, key: str, slots: SlotSink) -> None:
"""Records the string at `key`, along with the write that replaces it."""
value: Final = container.get(key)
if isinstance(value, str) and value:
slots.append((value, functools.partial(container.__setitem__, key)))
def collect_entry(entries: MutableSeq, index: int, slots: SlotSink) -> None:
"""Records a string held directly in a list, rather than under a key."""
value: Final = entries[index]
if isinstance(value, str) and value:
slots.append((value, functools.partial(entries.__setitem__, index)))
class RequestTooDeep(Exception):
"""A request nests text past a walk's bound.
Skipping the rest would forward it unredacted while the guardrail reports as
enabled, so the pre-call hook refuses the request instead.
"""
def collect_text_parts(container: MutableRequest, key: str, slots: SlotSink) -> None:
"""Collects the `text` of every part in the list held at `key`."""
for entry in as_array(container.get(key)) or ():
part = as_object(entry)
if part is not None:
collect(part, "text", slots)
def choice_index(choice: object) -> int:
"""Streaming choices are matched across chunks by their index."""
index: Final = getattr(choice, "index", 0)
return index if isinstance(index, int) else 0
def read_field(holder: object, name: str) -> object:
"""Reads one field from a dict or from an object.
LiteLLM's replies arrive as Pydantic models on some paths and as plain dicts on
others, depending how far they have been deserialised, so every response walk here
has to handle both shapes.
"""
fields: Final = as_object(holder)
if fields is not None:
return fields.get(name)
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)
if isinstance(value, tuple):
return value
return tuple(as_array(value) or ())
def write_field(holder: object, name: str, value: object) -> None:
"""Writes one string field back into a dict or an object. Pairs with read_field."""
if isinstance(holder, dict):
holder[name] = value
else:
setattr(holder, name, value)
def collect_json_leaves(node: object, slots: SlotSink, *, strict: bool = False) -> None:
"""Collects every string leaf of a JSON-ish structure, with a write-back per leaf.
An Anthropic `tool_use` block carries `input`, an arbitrary JSON object rather than a
string, so a value worth restoring can sit at any depth. Bounded by `MAX_JSON_DEPTH`:
the shape is caller or model controlled, and the bound is what stops a crafted one from
becoming an unbounded descent. Walked with an explicit stack rather than recursively,
so a deeply nested value cannot spend stack frames proportional to attacker-chosen
depth.
`strict` is for the request side, where a leaf left behind would reach the provider
unredacted: past the bound it raises `RequestTooDeep`. On the reply side a leaf past
the bound just keeps its placeholder, which leaks nothing, so it is skipped.
"""
pending: Final[list[tuple[object, int]]] = [(node, 0)] # mutable-ok: local walk stack.
while pending:
current, current_depth = pending.pop()
if current_depth > MAX_JSON_DEPTH:
if strict and is_container(current) and current:
raise RequestTooDeep("json")
continue
current_object = as_object(current)
if current_object is not None:
for key in tuple(current_object):
value = current_object[key]
if isinstance(value, str) and value:
slots.append((value, functools.partial(current_object.__setitem__, key)))
else:
pending.append((value, current_depth + 1))
continue
entries = as_array(current)
if entries is not None:
for index, value in enumerate(entries):
if isinstance(value, str) and value:
slots.append((value, functools.partial(entries.__setitem__, index)))
else:
pending.append((value, current_depth + 1))
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 or
mcp_call item holds `arguments`, their outputs `output`, a reasoning item `summary`
parts -- so the two directions stay symmetric. A custom tool call carries `input` and
a code interpreter call `code`, both model-written.
"""
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", "input", "code"):
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)

View file

@ -0,0 +1,402 @@
from collections.abc import Sequence
from typing import (
Final,
)
from .payload import (
MAX_JSON_DEPTH,
MutableRequest,
RequestTooDeep,
Slot,
SlotSink,
as_array,
as_object,
collect,
collect_entry,
collect_json_leaves,
collect_text_parts,
is_container,
read_list,
)
# Roles whose text the application author wrote and the caller never sees. Their
# PII is still redacted outbound, but it is not restorable from the reply.
PRIVILEGED_ROLES: Final = frozenset({"system", "developer"})
# How far a tool_result chain is followed. Real payloads nest one or two deep; the
# bound is what stops a crafted one from becoming an unbounded walk. A request that
# nests deeper is refused rather than forwarded, because text past the bound would
# otherwise reach the provider unredacted.
MAX_CONTENT_DEPTH: Final = 8
# JSON Schema keywords whose value has to reach the model or a validator verbatim, so the
# schema walk leaves them alone: types, formats, patterns, references and
# required-property lists. Everything else is scanned.
SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset(
(
"type",
"format",
"pattern",
"required",
"dependentRequired",
"propertyOrdering",
"discriminator",
"contentEncoding",
"contentMediaType",
"$ref",
"$id",
"$schema",
"$anchor",
"$dynamicRef",
"$dynamicAnchor",
"$recursiveRef",
"$recursiveAnchor",
"$vocabulary",
)
)
# Keywords holding JSON values rather than schemas: every string in them is collected,
# whatever the keys around it are called.
SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default"))
# Keywords holding the literal values the model must reproduce. These go to the CALLER's
# vault, not the privileged one: the model emits the stand-in in its tool arguments or
# structured output, and restoring the reply turns it back into the value the schema
# allows, so the call still routes. In the non-restorable vault it would come back as a
# stand-in no validator accepts.
SCHEMA_LITERAL_KEYWORDS: Final = frozenset(("enum", "const"))
# Keywords whose value maps names to subschemas. Their keys are property names, not
# keywords, so a property called `type` or `enum` is walked like any other subschema.
SCHEMA_MAP_KEYWORDS: Final = frozenset(
("properties", "patternProperties", "$defs", "definitions", "dependentSchemas", "dependencies")
)
def collect_prompt(data: MutableRequest, slots: SlotSink) -> None:
"""The Completions API sends its text in `prompt`, and its tail in `suffix`."""
collect(data, "suffix", slots)
prompt: Final = data.get("prompt")
if isinstance(prompt, str):
collect(data, "prompt", slots)
return
prompt_object: Final = as_object(prompt)
if prompt_object is not None:
# A Responses API PromptObject. `variables` are substituted into the stored
# prompt on the provider side, so they are caller text. `id` and `version`
# identify which prompt to use and must arrive unchanged. A variable is a string
# or a typed input such as `{"type": "input_text", "text": ...}`.
variables: Final = as_object(prompt_object.get("variables"))
if variables is not None:
for name in tuple(variables):
collect(variables, name, slots)
typed = as_object(variables[name])
if typed is not None:
collect(typed, "text", slots)
return
entries: Final = as_array(prompt)
if entries is None:
return
for index in range(len(entries)):
collect_entry(entries, index, slots)
def collect_content(container: MutableRequest, slots: SlotSink) -> None:
"""Collects `content`, a string or a list of typed parts.
An Anthropic tool_result nests its own content, so this has to descend. It walks
with an explicit stack and a depth bound rather than by recursion: the nesting is
caller controlled, and an unbounded descent is a JSON bomb. Content nested past the
bound raises `RequestTooDeep` rather than being skipped.
"""
# Walked in document order: the shield maps its replies back by position, so the
# order spans are collected in is part of the contract.
pending: Final[list[tuple[MutableRequest, int]]] = [(container, 0)] # mutable-ok: local queue, never escapes.
cursor = 0 # rebind-ok: advances through the queue.
while cursor < len(pending):
node, depth = pending[cursor]
cursor += 1
content = node.get("content")
if isinstance(content, str):
collect(node, "content", slots)
continue
if depth >= MAX_CONTENT_DEPTH and content:
raise RequestTooDeep("content")
for item in as_array(content) or ():
part = as_object(item)
if part is 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. Its own JSON bound applies, not the
# content one, and past it the request is refused.
collect_json_leaves(part.get("input"), slots, strict=True)
source = as_object(part.get("source")) if part.get("type") == "document" else None
if source is not None:
# An Anthropic document carries text inline: a `text` source holds it in
# `data`, a `content` source as a string or blocks, walked like any other
# content. Base64, URL and file sources are binary or remote, and pass
# untouched. Its `title` and `context` are caller text too.
collect(part, "title", slots)
collect(part, "context", slots)
if source.get("type") == "text":
collect(source, "data", slots)
elif source.get("type") == "content":
pending.append((source, depth + 1))
if "content" in part:
pending.append((part, depth + 1))
def collect_participant_name(message: MutableRequest, slots: SlotSink) -> None:
"""Redacts `name` where it identifies a person, never where it names a function.
On a user or assistant turn `name` is the participant, which is personal data.
On a tool or function turn the same field carries the function's name and has
to reach the provider unchanged, or the call no longer routes.
"""
if message.get("role") in ("tool", "function"):
return
collect(message, "name", slots)
def collect_tool_arguments(message: MutableRequest, slots: SlotSink) -> None:
"""Tool arguments carry the values a user asked the model to act on."""
for tool_call in read_list(message, "tool_calls"):
tool_call_object = as_object(tool_call)
function = as_object(tool_call_object.get("function")) if tool_call_object is not None else None
if function is not None:
collect(function, "arguments", slots)
legacy: Final = as_object(message.get("function_call"))
if legacy is not None:
collect(legacy, "arguments", slots)
def collect_system(data: MutableRequest, slots: SlotSink) -> None:
"""Anthropic's /v1/messages carries its system prompt at the top level."""
system: Final = data.get("system")
if isinstance(system, str):
collect(data, "system", slots)
return
collect_text_parts(data, "system", slots)
def collect_responses_fields(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None:
"""The Responses API sends text outside `messages`, in `instructions` and `input`.
`instructions` is written by the application, not by the caller, so it is
collected into the privileged sink; `input` is the caller's own text, except for
system and developer items in it, which go to the privileged sink like their Chat
counterparts.
"""
collect(data, "instructions", privileged)
request_input: Final = data.get("input")
if isinstance(request_input, str):
collect(data, "input", slots)
return
entries: Final = as_array(request_input)
if entries is None:
return
for index, entry in enumerate(entries):
if isinstance(entry, str):
# The embeddings and moderations shape: `input` as an array of strings.
collect_entry(entries, index, slots)
continue
item = as_object(entry)
if item is None:
continue
collect_content(item, privileged if item.get("role") in PRIVILEGED_ROLES else slots)
# A function_call item holds `arguments`; a function_call_output holds `output`,
# as a string or as a list of input_text parts. A custom_tool_call holds `input`
# and a code_interpreter_call `code` -- the fields the reply side restores.
collect(item, "arguments", slots)
collect(item, "output", slots)
collect_text_parts(item, "output", slots)
collect(item, "input", slots)
collect(item, "code", 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_tool_definitions(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None:
"""Tool definitions are application-authored free text bound for the provider.
A tool's description and the free text in its parameter schema are where callers put
examples and customer context, so they carry PII as often as a prompt does. They are
collected into the privileged sink, like a system prompt: redacted outbound, and never
restorable from the reply. `enum` and `const` values are the exception, and go to the
caller's vault -- see `SCHEMA_LITERAL_KEYWORDS`. Names and types are left as sent.
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"):
for entry in as_array(data.get(key)) or ():
tool = as_object(entry)
if tool is None:
continue
function = as_object(tool.get("function"))
for holder in (tool, function) if function is not None else (tool,):
collect(holder, "description", privileged)
collect_schema_text(holder.get("parameters"), slots, privileged)
collect_schema_text(holder.get("input_schema"), slots, privileged)
def collect_schema_text(schema: object, slots: SlotSink, privileged: SlotSink) -> None:
"""Collects the text in a JSON Schema, at any depth.
Scan by default: every string is collected except under the keywords in
`SCHEMA_STRUCTURAL_KEYWORDS`, whose values must go out verbatim. A list of keywords
*to* collect would leak every one it forgot -- draft-07 `dependencies`, a `$comment`,
a vendor `x-` extension -- which is how this walk started out. Free text goes to the
privileged sink; `enum` / `const` literals go to the caller's, so the model's use of
them is restored.
Structure matters in two places. Under `properties` and the other name -> subschema
maps, keys are property names rather than keywords, so a property called `type` is a
subschema to walk, not a keyword to skip. And `examples` / `default` hold JSON values,
so all their strings are collected whatever the keys around them are called. Nested
past `MAX_JSON_DEPTH`, the request is refused.
"""
pending: Final[list[tuple[object, int]]] = [(schema, 0)] # mutable-ok: local walk stack.
while pending:
node, depth = pending.pop()
if depth > MAX_JSON_DEPTH:
if is_container(node) and node:
raise RequestTooDeep("schema")
continue
entries = as_array(node)
if entries is not None:
for index, item in enumerate(entries):
collect_entry(entries, index, privileged)
if is_container(item):
pending.append((item, depth + 1))
continue
schema_object = as_object(node)
if schema_object is None:
continue
for keyword, value in tuple(schema_object.items()):
if keyword in SCHEMA_STRUCTURAL_KEYWORDS:
continue
subschemas = as_object(value) if keyword in SCHEMA_MAP_KEYWORDS else None
if keyword in SCHEMA_LITERAL_KEYWORDS:
collect(schema_object, keyword, slots)
collect_json_leaves(value, slots, strict=True)
elif keyword in SCHEMA_VALUE_KEYWORDS:
collect(schema_object, keyword, privileged)
collect_json_leaves(value, privileged, strict=True)
elif subschemas is not None:
pending.extend((child, depth + 1) for child in subschemas.values())
elif isinstance(value, str):
collect(schema_object, keyword, privileged)
elif is_container(value):
pending.append((value, depth + 1))
def collect_output_contracts(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None:
"""Text the caller sends to shape the reply rather than to prompt it.
A predicted output (`prediction.content`) is the caller's own draft of the answer, so
it goes with their text: the model largely repeats it, and it has to come back. A
structured-output schema -- Chat `response_format.json_schema`, Responses
`text.format` -- is application-authored like a tool schema, so its free text goes
to the privileged sink, and its names and types stay as sent.
"""
prediction: Final = as_object(data.get("prediction"))
if prediction is not None:
collect(prediction, "content", slots)
collect_text_parts(prediction, "content", slots)
response_format: Final = as_object(data.get("response_format"))
text_options: Final = as_object(data.get("text"))
for declared in (
response_format.get("json_schema") if response_format is not None else None,
text_options.get("format") if text_options is not None else None,
):
wrapper = as_object(declared)
if wrapper is not None:
collect(wrapper, "description", privileged)
collect_schema_text(wrapper.get("schema"), slots, privileged)
def collect_user_locations(data: MutableRequest, privileged: SlotSink) -> None:
"""Web search forwards the user's approximate location, whose `city` and `region`
are free text and can hold a street address.
Chat carries it in `web_search_options.user_location.approximate`; the Responses
and Anthropic web-search tools carry it flat on the tool's `user_location`. Nothing
restores it from a reply, hence the privileged sink.
"""
options: Final = data.get("web_search_options")
tools: Final = as_array(data.get("tools")) or ()
for declared in (options, *tools):
holder = as_object(declared)
location = as_object(holder.get("user_location")) if holder is not None else None
if location is None:
continue
approximate = as_object(location.get("approximate"))
for container in (location, approximate) if approximate is not None else (location,):
collect(container, "city", privileged)
collect(container, "region", privileged)
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 locate_request_texts(
data: MutableRequest,
) -> tuple[Sequence[Slot], Sequence[Slot]]:
"""Finds every redactable span, split by whether the caller can see it.
Anything missed here reaches the provider in the clear while the guardrail
still reports as enabled, so the walk covers every request shape that
carries text.
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`, tool and output schemas -- 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.
`extra_body` is walked the same way as the request itself. LiteLLM merges it over
the transformed request just before sending, so a field there -- `input`,
`messages`, `system` -- replaces the redacted one on the wire.
"""
slots: Final[SlotSink] = []
privileged: Final[SlotSink] = []
for payload in (data, as_object(data.get("extra_body"))):
if payload is None:
continue
for entry in read_list(payload, "messages"):
message = as_object(entry)
if message is not None:
sink = privileged if message.get("role") in PRIVILEGED_ROLES else slots
collect_content(message, sink)
collect_participant_name(message, sink)
collect_tool_arguments(message, sink)
collect_responses_fields(payload, slots, privileged)
collect_prompt(payload, slots)
collect_system(payload, privileged)
collect_tool_definitions(payload, slots, privileged)
collect_output_contracts(payload, slots, privileged)
collect_user_locations(payload, privileged)
collect_end_user_ids(payload, privileged)
return tuple(slots), tuple(privileged)

View file

@ -0,0 +1,368 @@
import copy
import functools
import itertools
import json
import re
from enum import Enum
from types import MappingProxyType
from typing import (
Final,
TypeAlias,
)
from .payload import (
JsonBody,
MutableRequest,
Rehydrate,
SlotSink,
StreamStep,
as_object,
collect_response_item,
read_field,
read_list,
rehydrate_slots,
write_field,
)
# 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)")
# What an SSE stream can open with: one of its fields, or a `:` comment.
SSE_OPENINGS: Final = (b"event:", b"data:", b"id:", b"retry:", b":")
# Responses API delta events whose `delta` is not text. Audio arrives base64-encoded;
# sending it through the shield would cost a round trip per chunk to restore nothing.
RESPONSES_BINARY_DELTAS: Final = frozenset(("response.audio.delta",))
# Fields on a Responses API event that identify something rather than say something.
# Every other string field on a `.done` event is model text and is restored, so an event
# type added upstream is covered by default instead of leaking a placeholder.
RESPONSES_STRUCTURAL_FIELDS: Final = frozenset(
("type", "id", "item_id", "call_id", "name", "server_label", "status", "obfuscation")
)
# Terminal Responses API events that repeat the whole reply under `response`.
RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete"))
# Sliding windows keyed by (choice index, tool-call index | None), threaded through one
# stream. `None` is the content channel; an int is one tool call's accumulating
# `arguments`. Content and each tool call are separate token streams, so each needs its
# own window -- one shared window would splice one stream's held-back tail onto another.
CarryKey: TypeAlias = tuple[int, int | None]
CarryWindows: TypeAlias = dict[CarryKey, str]
# A Responses API delta stream: (event family, item id, output index, part index).
ResponsesStreamKey: TypeAlias = tuple[str, object, object, object]
def carry_sort_key(key: CarryKey) -> tuple[int, int]:
"""Orders streaming windows without ever comparing None to an int.
`sorted()` over the raw keys raises as soon as one choice holds both a content window
and a tool-call window, because `None < 0` is not orderable. Content sorts first, then
tool calls by their index.
"""
choice_index, tool_index = key
return (choice_index, -1 if tool_index is None else tool_index)
def continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]:
"""A `tool_calls` delta carrying `text` as an index-only continuation.
Clients concatenate tool-call fragments by index, so no id or name is needed.
"""
return [{"index": tool_index, "function": {"arguments": text}}]
def opens_like_sse(head: bytes) -> bool | None:
"""Whether a raw stream is SSE, judged by its opening bytes; None while undecidable.
An SSE stream opens with a field name or a `:` comment. Anything else -- a JSON array
streamed in pieces, say -- has no event boundaries to wait for. A chunk that ends
partway through a field name decides nothing yet, so that case waits for more.
"""
opening: Final = head.lstrip()
if not opening:
return None
if opening.startswith(SSE_OPENINGS):
return True
if any(field.startswith(opening) for field in SSE_OPENINGS):
return None
return False
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""
self._as_text = False
self._is_sse: bool | None = None
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
if self._is_sse is False:
return (chunk,)
raw: Final = chunk.encode("utf-8") if isinstance(chunk, str) else chunk
buffered: Final = self._pending + raw
if self._is_sse is None:
self._is_sse = opens_like_sse(buffered)
if self._is_sse is None:
# Too little has arrived to tell -- `b"eve"` could still become `event:`.
self._pending = buffered
return ()
if not self._is_sse:
self._pending = b""
return self._emit(buffered)
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(
[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."""
held: Final = self._pending
self._pending = b""
if not self._is_sse:
# The stream ended before it could be told apart from SSE: hand it back as is.
return self._emit(held)
tail: Final = await self._restore_event(held) if held.strip() else held
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:
parsed: Final[object] = json.loads(line[len("data:") :])
except ValueError:
return block
event: Final = as_object(parsed)
if event is None:
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] = {
"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)])
return b"".join(flushed)
class ResponsesStreamRestorer:
"""Restores a Responses API event stream.
The event families are matched by shape rather than listed, so a text stream the
API adds later is restored by default instead of leaking a placeholder:
- Any `*.delta` event whose `delta` is a string is a token stream (output_text,
refusal, function-call and MCP arguments, reasoning summaries, ...). Each gets its
own window, keyed by the family, the item id and the part index.
- Any `*.done` event closes the stream of the same family. Whatever its 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`. Then every text
field on the done event is restored in full: its string fields other than
identifiers, plus any `part` or `item` it repeats.
- `response.completed` / `response.incomplete` repeat the whole reply, and 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[ResponsesStreamKey, str]] = {} # mutable-ok: per-stream windows advanced in place.
self._last_deltas: Final[dict[ResponsesStreamKey, 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.endswith(".delta") and kind not in RESPONSES_BINARY_DELTAS:
await self._restore_delta(event, kind)
return (event,)
slots: Final[SlotSink] = []
flushed: Final = await self._flush(responses_stream_key(event, kind)) if kind.endswith(".done") else ()
if kind.endswith(".done"):
collect_event_text(event, slots)
part: Final = read_field(event, "part")
if part is not None:
collect_response_item({"content": [part]}, slots)
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 (*flushed, 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)])
return tuple(itertools.chain.from_iterable(flushed))
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: ResponsesStreamKey) -> 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) -> ResponsesStreamKey:
"""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 collect_event_text(event: object, slots: SlotSink) -> None:
"""Collects every top-level text field of a Responses API event, dict or model.
Scan by default, with identifiers excluded, rather than a list of known fields: the
`.done` event of each stream family names its text differently (`text`, `refusal`,
`arguments`, ...), and a family added upstream would otherwise leak a placeholder.
"""
# A model's fields live in its `__dict__`; an empty dict has none either way.
attributes: Final[object] = getattr(event, "__dict__", None)
fields: Final = as_object(event) or as_object(attributes)
if fields is None:
return
for name, value in tuple(fields.items()):
if name in RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"):
continue
if isinstance(value, str) and value:
slots.append((value, functools.partial(write_field, event, name)))