diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 01770f83141..80d4446339a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 64768b39070..0929341028e 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -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 == []