mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
b55622e8ba
commit
464516a188
4 changed files with 1063 additions and 983 deletions
File diff suppressed because it is too large
Load diff
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)))
|
||||
Loading…
Add table
Reference in a new issue