mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): fail closed on deep llm_shield_proxy requests, widen coverage
- Request walks no longer skip what lies past their depth bound. Content nested past it, and tool inputs or schemas past the new JSON bound, now block the request instead of reaching the provider unredacted. The old depth test asserted the skip; it now asserts the block. - Tool and output schemas are walked by their JSON Schema structure, and give up `title`, `examples` and `default` as well as `description`. `enum` and `const` still go out as sent. - Responses events are matched by shape: any `*.delta` with a string delta is a token stream, and any `*.done` restores every non-identifier text field plus the `part` or `item` it repeats. This covers reasoning_summary_part.done and MCP arguments, and future families. Audio deltas are left alone. - An SSE stream whose first chunk ends partway through a field name (`b"eve"`) is no longer taken for a non-SSE stream.
This commit is contained in:
parent
40c26178fa
commit
d7d608fc52
2 changed files with 317 additions and 118 deletions
|
|
@ -82,13 +82,15 @@ JsonBody: TypeAlias = dict
|
|||
# One redactable span: the text as it stands, and the write that puts the
|
||||
# replacement back where it came from.
|
||||
# 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.
|
||||
# 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
|
||||
|
||||
# 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
|
||||
# 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, for the same reason as above.
|
||||
_MAX_JSON_DEPTH: Final = 64
|
||||
|
||||
_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list.
|
||||
|
||||
|
|
@ -108,31 +110,50 @@ _ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_
|
|||
# `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",
|
||||
)
|
||||
)
|
||||
# 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":")
|
||||
|
||||
# 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",
|
||||
}
|
||||
# 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"))
|
||||
|
||||
# JSON Schema keywords that hold free text an application writes, and so can hold PII.
|
||||
# `enum` and `const` are deliberately absent: the model has to reproduce those values
|
||||
# exactly, and one redacted into the non-restorable vault would come back as a stand-in.
|
||||
_SCHEMA_TEXT_KEYWORDS: Final = frozenset(("description", "title"))
|
||||
_SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default"))
|
||||
|
||||
# JSON Schema keywords whose value is a map of name -> subschema, a single subschema, or
|
||||
# a list of subschemas. Knowing which is which is what lets the walk tell a property
|
||||
# *named* "description" apart from the `description` keyword.
|
||||
_SCHEMA_MAP_KEYWORDS: Final = frozenset(("properties", "patternProperties", "$defs", "definitions", "dependentSchemas"))
|
||||
_SCHEMA_KEYWORDS: Final = frozenset(
|
||||
(
|
||||
"items",
|
||||
"additionalProperties",
|
||||
"additionalItems",
|
||||
"unevaluatedProperties",
|
||||
"unevaluatedItems",
|
||||
"propertyNames",
|
||||
"contains",
|
||||
"not",
|
||||
"if",
|
||||
"then",
|
||||
"else",
|
||||
)
|
||||
)
|
||||
_SCHEMA_LIST_KEYWORDS: Final = frozenset(("allOf", "anyOf", "oneOf", "prefixItems"))
|
||||
|
||||
# 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.
|
||||
|
|
@ -184,12 +205,21 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None:
|
|||
_collect_entry(prompt, index, slots)
|
||||
|
||||
|
||||
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_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.
|
||||
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.
|
||||
|
|
@ -202,8 +232,8 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None:
|
|||
if isinstance(content, str):
|
||||
_collect(node, "content", slots)
|
||||
continue
|
||||
if depth >= _MAX_CONTENT_DEPTH:
|
||||
continue
|
||||
if depth >= _MAX_CONTENT_DEPTH and content:
|
||||
raise _RequestTooDeep("content")
|
||||
for part in content if isinstance(content, list) else ():
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
|
|
@ -212,8 +242,9 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None:
|
|||
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)
|
||||
# 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)
|
||||
if "content" in part:
|
||||
pending.append((part, depth + 1))
|
||||
|
||||
|
|
@ -292,11 +323,11 @@ def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) -
|
|||
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.
|
||||
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. Names, types, `enum` and `const` 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`.
|
||||
|
|
@ -309,27 +340,36 @@ def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> No
|
|||
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)
|
||||
_collect_schema_text(holder.get("parameters"), privileged)
|
||||
_collect_schema_text(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.
|
||||
def _collect_schema_text(schema: object, privileged: _SlotSink) -> None:
|
||||
"""Collects the free text 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.
|
||||
That is every `description` and `title` string, and every string inside `examples`
|
||||
and `default`. The walk follows the schema's own structure -- `properties` and the
|
||||
other subschema keywords -- rather than every nested dict, which is what tells a
|
||||
property *named* "description" (a subschema, descended into) from the `description`
|
||||
keyword (text, collected). Nested past `_MAX_JSON_DEPTH`, the request is refused.
|
||||
"""
|
||||
pending: Final[list] = [(schema, 0)] # mutable-ok: local walk stack.
|
||||
while pending:
|
||||
node, depth = pending.pop()
|
||||
if depth > _MAX_SCHEMA_DEPTH:
|
||||
if not isinstance(node, dict):
|
||||
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)))
|
||||
if depth > _MAX_JSON_DEPTH:
|
||||
raise _RequestTooDeep("schema")
|
||||
for keyword, value in tuple(node.items()):
|
||||
if keyword in _SCHEMA_TEXT_KEYWORDS:
|
||||
_collect(node, keyword, privileged)
|
||||
elif keyword in _SCHEMA_VALUE_KEYWORDS:
|
||||
_collect(node, keyword, privileged)
|
||||
_collect_json_leaves(value, privileged, strict=True)
|
||||
elif keyword in _SCHEMA_MAP_KEYWORDS and isinstance(value, dict):
|
||||
pending.extend((child, depth + 1) for child in value.values())
|
||||
elif keyword in _SCHEMA_KEYWORDS or keyword in _SCHEMA_LIST_KEYWORDS:
|
||||
pending.extend((child, depth + 1) for child in (value if isinstance(value, list) else (value,)))
|
||||
|
||||
|
||||
def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None:
|
||||
|
|
@ -338,7 +378,7 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged
|
|||
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 descriptions go
|
||||
`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 = data.get("prediction")
|
||||
|
|
@ -346,11 +386,14 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged
|
|||
_collect(prediction, "content", slots)
|
||||
_collect_text_parts(prediction, "content", slots)
|
||||
response_format: Final = data.get("response_format")
|
||||
if isinstance(response_format, dict):
|
||||
_collect_schema_descriptions(response_format.get("json_schema"), privileged)
|
||||
text_options: Final = data.get("text")
|
||||
if isinstance(text_options, dict):
|
||||
_collect_schema_descriptions(text_options.get("format"), privileged)
|
||||
for wrapper in (
|
||||
response_format.get("json_schema") if isinstance(response_format, dict) else None,
|
||||
text_options.get("format") if isinstance(text_options, dict) else None,
|
||||
):
|
||||
if isinstance(wrapper, dict):
|
||||
_collect(wrapper, "description", privileged)
|
||||
_collect_schema_text(wrapper.get("schema"), privileged)
|
||||
|
||||
|
||||
def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None:
|
||||
|
|
@ -400,20 +443,26 @@ def _write_field(holder: object, name: str, value: str) -> None:
|
|||
setattr(holder, name, value)
|
||||
|
||||
|
||||
def _collect_json_leaves(node: object, slots: _SlotSink, depth: int = 0) -> None:
|
||||
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_CONTENT_DEPTH` for the same reason the request walk is: the shape is 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
|
||||
tool input cannot spend stack frames proportional to attacker-chosen depth.
|
||||
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] = [(node, depth)] # mutable-ok: local walk stack.
|
||||
pending: Final[list] = [(node, 0)] # mutable-ok: local walk stack.
|
||||
while pending:
|
||||
current, current_depth = pending.pop()
|
||||
if current_depth > _MAX_CONTENT_DEPTH:
|
||||
if current_depth > _MAX_JSON_DEPTH:
|
||||
if strict and isinstance(current, (dict, list)) and current:
|
||||
raise _RequestTooDeep("json")
|
||||
continue
|
||||
if isinstance(current, dict):
|
||||
for key in tuple(current):
|
||||
|
|
@ -453,9 +502,10 @@ def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]:
|
|||
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.
|
||||
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"):
|
||||
|
|
@ -466,7 +516,7 @@ def _collect_response_item(item: object, slots: _SlotSink) -> None:
|
|||
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"):
|
||||
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)))
|
||||
|
|
@ -481,6 +531,23 @@ async def _rehydrate_slots(slots: Sequence[_Slot], rehydrate: _Rehydrate) -> Non
|
|||
write(replacement)
|
||||
|
||||
|
||||
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.
|
||||
|
||||
|
|
@ -517,20 +584,25 @@ class _AnthropicSSERestorer:
|
|||
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.
|
||||
self._is_sse: bool | None = None # rebind-ok: undecided until the opening bytes settle it.
|
||||
|
||||
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:
|
||||
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
|
||||
|
|
@ -549,8 +621,12 @@ class _AnthropicSSERestorer:
|
|||
|
||||
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
|
||||
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""
|
||||
|
|
@ -637,16 +713,19 @@ class _AnthropicSSERestorer:
|
|||
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 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:
|
||||
|
||||
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.
|
||||
- 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:
|
||||
|
|
@ -660,25 +739,22 @@ class _ResponsesStreamRestorer:
|
|||
kind: Final = _responses_event_type(event)
|
||||
if kind is None:
|
||||
return (event,)
|
||||
if kind in _RESPONSES_DELTA_EVENTS:
|
||||
if kind.endswith(".delta") and kind not in _RESPONSES_BINARY_DELTAS:
|
||||
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":
|
||||
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")
|
||||
_collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view.
|
||||
elif kind == "response.output_item.done":
|
||||
if part is not None:
|
||||
_collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view.
|
||||
_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,)
|
||||
return (*flushed, event)
|
||||
|
||||
async def finish(self) -> tuple[object, ...]:
|
||||
"""Flushes every stream the provider never closed, e.g. a truncated reply."""
|
||||
|
|
@ -725,12 +801,21 @@ def _responses_stream_key(event: object, kind: str) -> tuple:
|
|||
)
|
||||
|
||||
|
||||
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)),)
|
||||
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.
|
||||
"""
|
||||
fields: Final = event if isinstance(event, dict) else getattr(event, "__dict__", None)
|
||||
if not isinstance(fields, dict):
|
||||
return
|
||||
for name, value in tuple(fields.items()):
|
||||
if not isinstance(name, str) or name in _RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"):
|
||||
continue
|
||||
if isinstance(value, str) and value:
|
||||
slots.append((value, lambda new, n=name: _write_field(event, n, new)))
|
||||
|
||||
|
||||
class LLMShieldProxyGuardrail(CustomGuardrail):
|
||||
|
|
@ -919,7 +1004,13 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
|||
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True:
|
||||
return data
|
||||
|
||||
slots, privileged = self._locate_request_texts(data)
|
||||
try:
|
||||
slots, privileged = self._locate_request_texts(data)
|
||||
except _RequestTooDeep as exc:
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=f"Request {exc} nests deeper than LLM Shield Proxy inspects; blocking the request.",
|
||||
) from exc
|
||||
if not slots and not privileged:
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -480,33 +480,59 @@ class TestRequestCoverage:
|
|||
assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deeply_nested_tool_results_are_bounded(self):
|
||||
"""Nesting is caller controlled, so the descent has to stop somewhere.
|
||||
async def test_nesting_past_the_bound_blocks_the_request(self):
|
||||
"""Nesting is caller controlled, so the descent has to stop somewhere -- and where
|
||||
it stops, the request must not go out.
|
||||
|
||||
The walk must terminate on a payload built to be pathological, rather than
|
||||
following it as far as it goes.
|
||||
This test used to assert the opposite: that text past the bound was skipped. That
|
||||
sent `past-the-bound@example.com` to the provider unredacted while the guardrail
|
||||
reported as enabled.
|
||||
"""
|
||||
guardrail = _guardrail()
|
||||
|
||||
captured: list = []
|
||||
|
||||
async def echo(url, headers, json, timeout): # noqa: ARG001
|
||||
captured.append(json["texts"])
|
||||
return _response({"texts": list(json["texts"])})
|
||||
|
||||
guardrail.async_handler.post = AsyncMock(side_effect=echo) # type: ignore[method-assign]
|
||||
mock = _mock_post(guardrail)
|
||||
|
||||
deep: dict = {"type": "tool_result", "content": "past-the-bound@example.com"}
|
||||
for _ in range(200):
|
||||
deep = {"type": "tool_result", "content": [deep]}
|
||||
data = {"messages": [{"role": "user", "content": [{"type": "text", "text": "shallow"}, deep]}]}
|
||||
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion")
|
||||
mock.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deep_tool_input_blocks_the_request(self):
|
||||
"""A tool_use input past the JSON bound must not be forwarded half-redacted."""
|
||||
guardrail = _guardrail()
|
||||
mock = _mock_post(guardrail)
|
||||
|
||||
deep: dict = {"email": "past-the-bound@example.com"}
|
||||
for _ in range(100):
|
||||
deep = {"next": deep}
|
||||
block = {"type": "tool_use", "id": "t1", "name": "f", "input": deep}
|
||||
data = {"messages": [{"role": "assistant", "content": [block]}]}
|
||||
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion")
|
||||
mock.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realistic_nesting_is_redacted_in_full(self):
|
||||
"""The bounds are far past real payloads: a tool input nested inside a tool result,
|
||||
several JSON levels deep, is redacted whole rather than refused."""
|
||||
guardrail = _guardrail()
|
||||
_mock_post(guardrail, {"texts": ["[EMAIL_1]"]})
|
||||
|
||||
tool_use = {
|
||||
"type": "tool_use",
|
||||
"id": "t1",
|
||||
"name": "f",
|
||||
"input": {"a": {"b": {"c": {"d": {"to": "x@example.com"}}}}},
|
||||
}
|
||||
data = {"messages": [{"role": "user", "content": [{"type": "tool_result", "content": [tool_use]}]}]}
|
||||
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion")
|
||||
|
||||
sent = captured[0]
|
||||
assert "shallow" in sent
|
||||
assert "past-the-bound@example.com" not in sent, "the walk followed the chain past its bound"
|
||||
assert len(sent) < 200
|
||||
assert tool_use["input"]["a"]["b"]["c"]["d"]["to"] == "[EMAIL_1]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_prompt_object_variables_are_redacted(self):
|
||||
|
|
@ -607,8 +633,9 @@ class TestRequestCoverage:
|
|||
|
||||
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."""
|
||||
def test_tool_schemas_give_up_their_free_text_and_nothing_else(self):
|
||||
"""Descriptions, titles, examples and defaults are collected. Names, types, enum
|
||||
and const values must reach the model exactly as sent."""
|
||||
data = {
|
||||
"tools": [
|
||||
{
|
||||
|
|
@ -618,12 +645,16 @@ class TestRequestCoverage:
|
|||
"description": "top",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"title": "title",
|
||||
"properties": {
|
||||
# A property that is itself named "description".
|
||||
"description": {"type": "string", "description": "named"},
|
||||
"kind": {"type": "string", "enum": ["a", "b"], "description": "enum"},
|
||||
"kind": {"type": "string", "enum": ["a", "b"], "const": "a", "description": "enum"},
|
||||
"deep": {"type": "array", "items": {"type": "object", "description": "nested"}},
|
||||
"to": {"type": "string", "examples": ["example"], "default": "default"},
|
||||
"choice": {"anyOf": [{"type": "object", "default": {"who": "object-default"}}]},
|
||||
},
|
||||
"$defs": {"shared": {"description": "defined"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
|
@ -631,7 +662,26 @@ class TestRequestCoverage:
|
|||
}
|
||||
_, privileged = LLMShieldProxyGuardrail._locate_request_texts(data)
|
||||
|
||||
assert sorted(text for text, _ in privileged) == ["enum", "named", "nested", "top"]
|
||||
assert sorted(text for text, _ in privileged) == [
|
||||
"default",
|
||||
"defined",
|
||||
"enum",
|
||||
"example",
|
||||
"named",
|
||||
"nested",
|
||||
"object-default",
|
||||
"title",
|
||||
"top",
|
||||
]
|
||||
|
||||
def test_schema_nesting_past_the_bound_is_refused(self):
|
||||
schema: dict = {"type": "object", "description": "past-the-bound@example.com"}
|
||||
for _ in range(100):
|
||||
schema = {"type": "object", "properties": {"next": schema}}
|
||||
data = {"tools": [{"type": "function", "function": {"name": "f", "parameters": schema}}]}
|
||||
|
||||
with pytest.raises(Exception, match="schema"):
|
||||
LLMShieldProxyGuardrail._locate_request_texts(data)
|
||||
|
||||
|
||||
class TestRestoration:
|
||||
|
|
@ -1320,6 +1370,19 @@ class TestAnthropicStreamRestoration:
|
|||
|
||||
assert out == chunks
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("cut", [1, 3, 5, 6])
|
||||
async def test_a_field_name_split_by_the_first_chunk_still_reads_as_sse(self, cut: int):
|
||||
"""`b"eve"` then `b"nt: ..."` is still SSE; deciding on the first chunk alone
|
||||
would pass the whole stream through with its placeholders."""
|
||||
guardrail, _ = _shielded(self.VAULT)
|
||||
raw = b"".join(_text_block_stream("Mail [EMAIL_1]"))
|
||||
|
||||
out = await _restore_stream(guardrail, [raw[:cut], raw[cut:]])
|
||||
|
||||
deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"]
|
||||
assert "".join(deltas) == "Mail a@example.com"
|
||||
|
||||
|
||||
class TestResponsesStreamRestoration:
|
||||
"""/v1/responses streams are typed events, with no `choices` to walk."""
|
||||
|
|
@ -1415,3 +1478,48 @@ class TestResponsesStreamRestoration:
|
|||
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"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reasoning_summary_part_done_is_restored(self):
|
||||
"""The summary part repeats the whole summary text after its deltas."""
|
||||
guardrail, _ = _shielded(self.VAULT)
|
||||
part = SimpleNamespace(type="summary_text", text="asked about [EMAIL_1]")
|
||||
event = SimpleNamespace(
|
||||
type="response.reasoning_summary_part.done", item_id="rs_1", output_index=0, summary_index=0, part=part
|
||||
)
|
||||
|
||||
await _restore_stream(guardrail, [event])
|
||||
|
||||
assert part.text == "asked about a@example.com"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_call_arguments_are_restored(self):
|
||||
"""A stream family outside the chat-era set: matched by shape, not by name."""
|
||||
guardrail, _ = _shielded(self.VAULT)
|
||||
deltas = [
|
||||
{"type": "response.mcp_call_arguments.delta", "item_id": "mcp_1", "output_index": 0, "delta": d}
|
||||
for d in ('{"to": "[EMAI', 'L_1]"}')
|
||||
]
|
||||
done = {
|
||||
"type": "response.mcp_call_arguments.done",
|
||||
"item_id": "mcp_1",
|
||||
"output_index": 0,
|
||||
"arguments": '{"to": "[EMAIL_1]"}',
|
||||
}
|
||||
|
||||
out = await _restore_stream(guardrail, [*deltas, done])
|
||||
|
||||
assert json.loads("".join(e["delta"] for e in out[:-1])) == {"to": "a@example.com"}
|
||||
assert json.loads(out[-1]["arguments"]) == {"to": "a@example.com"}
|
||||
assert out[-1]["item_id"] == "mcp_1", "identifiers are not text and stay as sent"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_deltas_are_not_sent_to_the_shield(self):
|
||||
"""Audio arrives base64-encoded; restoring it would cost a round trip for nothing."""
|
||||
guardrail, shield = _shielded(self.VAULT)
|
||||
audio = {"type": "response.audio.delta", "item_id": "a_1", "output_index": 0, "delta": "UklGRiQAAABXQVZF"}
|
||||
|
||||
out = await _restore_stream(guardrail, [audio])
|
||||
|
||||
assert out == [audio]
|
||||
assert shield.urls == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue