From d598d1bf96a05cd24fc3fbb6bfa7deb2ef31d596 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sun, 13 Sep 2026 11:52:14 -0500 Subject: [PATCH] fix(guardrails): restore tool calls in the LLM Shield guardrail The request walk redacted a tool call's `arguments` -- plus the legacy `function_call`, Anthropic `tool_use.input` leaves and the Responses API's `function_call` / `function_call_output` fields -- while the response walk restored only `message.content`. A placeholder therefore reached the caller inside a tool call, and nothing raised. This is the same change as the out-of-tree example adapter this file is copied from, kept body-identical on purpose: the response side now collects every restorable span in one positional rehydrate batch, streaming keeps a window per (choice index, tool-call index) and flushes each into the chunk carrying the finish_reason, and `apply_guardrail` restores `inputs["tool_calls"]` on the response side. The declared limit on restoring values inside a JSON string is documented in the module. --- .../llm_shield_proxy/llm_shield_proxy.py | 369 ++++++++++++++---- 1 file changed, 298 insertions(+), 71 deletions(-) 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 3cab7b20cef..e093199578a 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 @@ -85,8 +85,11 @@ _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's p # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. -# Sliding windows keyed by streaming choice index, threaded through one stream. -_CarryWindows: TypeAlias = dict # mutable-ok: per-choice windows advanced in place. +# 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. +_CarryWindows: TypeAlias = dict # mutable-ok: per-stream windows advanced in place. # A caller-owned list whose entries are rewritten in place, such as a Completions # `prompt` sent as an array of strings. @@ -224,6 +227,64 @@ def _choice_index(choice: object) -> int: 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. + """ + if isinstance(holder, dict): + return holder.get(name) + return getattr(holder, name, None) + + +def _write_field(holder: object, name: str, value: str) -> None: + """Writes one string field back into a dict or an object. Pairs with _read_field.""" + if isinstance(holder, dict): + holder[name] = value + else: + setattr(holder, name, value) + + +def _collect_json_leaves(node: object, slots: _SlotSink, depth: int = 0) -> 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. + """ + if depth > _MAX_CONTENT_DEPTH: + return + if isinstance(node, dict): + for key in tuple(node): + value = node[key] + if isinstance(value, str) and value: + slots.append((value, lambda new, d=node, k=key: d.__setitem__(k, new))) + else: + _collect_json_leaves(value, slots, depth + 1) + return + if isinstance(node, list): + for index, value in enumerate(node): + if isinstance(value, str) and value: + slots.append((value, lambda new, entries=node, i=index: entries.__setitem__(i, new))) + else: + _collect_json_leaves(value, slots, depth + 1) + + +def _carry_sort_key(key: tuple) -> tuple: + """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) + + class LLMShieldProxyGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -414,6 +475,15 @@ class LLMShieldProxyGuardrail(CustomGuardrail): for (_, write), replacement in zip(slots, redacted): write(replacement) + # KNOWN LIMIT: a tool call's `arguments` is a JSON *string*, and a restored value is + # spliced into it as raw text. If the original value contained a double quote, a + # backslash or a newline, the reassembled document is no longer valid JSON for a + # strict parser. The proxy's own tool-argument rehydration has the same property + # (`_rehydrate_json_response` in api/main.py), so this is a pre-existing limit of the + # product rather than one introduced here. Escaping is deliberately NOT applied as a + # fix: a fragment is an arbitrary slice of a JSON document, so the code cannot tell + # whether the position it writes is inside a string literal, and escaping + # unconditionally would corrupt the values that are not. @log_guardrail_information async def async_post_call_success_hook( self, @@ -428,27 +498,46 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if self._is_anthropic_message_response(response): return await self._restore_anthropic_response(response, data) - text_blocks: Final = self._responses_api_text_blocks(response) - if text_blocks: - return await self._restore_responses_api_response(response, text_blocks, data) + response_slots: Final = self._responses_api_slots(response) + if response_slots: + return await self._restore_responses_api_response(response, response_slots, data) choices: Final = getattr(response, "choices", None) if not choices: return response - pending: Final = tuple( - (choice.message, choice.message.content) - for choice in choices - if getattr(choice, "message", None) is not None - and isinstance(getattr(choice.message, "content", None), str) - and choice.message.content - ) + # One batch for every restorable span in the reply, collected in document order: + # the shield maps its answers back by position. A second round trip is not an + # option here -- /v1/guard/rehydrate caps a batch at 256 texts and 1,000,000 + # characters, and `_same_length_or_raise` is what guarantees the positional + # mapping -- so a reply carrying more spans than that fails closed, which is this + # guardrail's posture everywhere else. + pending: Final[list] = [] # mutable-ok: local accumulator, frozen before use. + for choice in choices: + message: Final = getattr(choice, "message", None) + if message is None: + continue + content: Final = getattr(message, "content", None) + if isinstance(content, str) and content: + pending.append((content, lambda new, m=message: setattr(m, "content", new))) + # A tool call's `arguments` is model-generated text and the request path + # redacts it, so leaving it unrestored hands the application a placeholder to + # invoke a tool with. These are Pydantic objects on this path, not dicts. + for tool_call in getattr(message, "tool_calls", None) or (): + function: Final = getattr(tool_call, "function", None) + arguments: Final = getattr(function, "arguments", None) if function is not None else None + if isinstance(arguments, str) and arguments: + pending.append((arguments, lambda new, f=function: setattr(f, "arguments", new))) + legacy: Final = getattr(message, "function_call", None) + legacy_arguments: Final = getattr(legacy, "arguments", None) if legacy is not None else None + if isinstance(legacy_arguments, str) and legacy_arguments: + pending.append((legacy_arguments, lambda new, fn=legacy: setattr(fn, "arguments", new))) if not pending: return response - restored: Final = await self._rehydrate(tuple(text for _, text in pending), self._session_id(data)) - for (message, _), replacement in zip(pending, restored): - message.content = replacement + restored: Final = await self._rehydrate(tuple(text for text, _ in pending), self._session_id(data)) + for (_, write), replacement in zip(pending, restored): + write(replacement) return response @staticmethod @@ -461,60 +550,65 @@ class LLMShieldProxyGuardrail(CustomGuardrail): ) async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: - """Restores text blocks in an Anthropic native message reply. + """Restores text blocks and tool inputs in an Anthropic native message reply. This shape has no `choices`, so without its own branch the reply would go back to the caller still carrying placeholders. + + A `tool_use` block's payload is `input`, an arbitrary JSON object rather than a + string, and the request path redacts its string leaves -- so the reply's leaves + have to come back or the application invokes the tool with placeholders. """ - blocks: Final = tuple( - block - for block in response["content"] - if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) - ) - if not blocks: + slots: Final[list] = [] # mutable-ok: accumulator, frozen before use. + for block in response["content"]: + if not isinstance(block, dict): + continue + kind: Final = block.get("type") + if kind == "text" and isinstance(block.get("text"), str) and block["text"]: + slots.append((block["text"], lambda new, b=block: b.__setitem__("text", new))) + elif kind == "tool_use" and isinstance(block.get("input"), dict): + _collect_json_leaves(block["input"], slots) + if not slots: return response - restored: Final = await self._rehydrate(tuple(block["text"] for block in blocks), self._session_id(data)) - for block, replacement in zip(blocks, restored): - block["text"] = replacement + restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, restored): + write(replacement) return response @staticmethod - def _responses_api_text_blocks(response: object) -> Sequence[object]: - """Text blocks in a Responses API reply. + def _responses_api_slots(response: object) -> Sequence[_Slot]: + """Restorable spans in a Responses API reply. That shape carries `output` items rather than `choices`, so it needs its own walk; without one the reply goes back to the caller still holding - placeholders even though the request was redacted correctly. Blocks come - through as dicts or as objects depending on how far the reply has been + placeholders even though the request was redacted correctly. Items and blocks + come through as dicts or as objects depending on how far the reply has been deserialised, so both are handled. + + The item-level fields mirror `_collect_responses_fields`, which walks the same + fields on the request side -- a function_call item holds `arguments`, a + function_call_output holds `output` -- so the two directions stay symmetric. """ - blocks: Final[list[object]] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[list] = [] # mutable-ok: accumulator, frozen on return. for item in getattr(response, "output", None) or (): for block in getattr(item, "content", None) or (): - if isinstance(block, dict): - if isinstance(block.get("text"), str) and block["text"]: - blocks.append(block) - elif isinstance(getattr(block, "text", None), str) and block.text: - blocks.append(block) - return tuple(blocks) - - @staticmethod - def _block_text(block: object) -> str: - return block["text"] if isinstance(block, dict) else block.text + text: Final = _read_field(block, "text") + if isinstance(text, str) and text: + slots.append((text, lambda new, b=block: _write_field(b, "text", new))) + for field in ("arguments", "output"): + value: Final = _read_field(item, field) + if isinstance(value, str) and value: + slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) + return tuple(slots) async def _restore_responses_api_response( - self, response: Any, blocks: Sequence[object], data: MutableRequest + self, response: Any, slots: Sequence[_Slot], data: MutableRequest ) -> Any: """Puts the original values back into a Responses API reply.""" - restored: Final = await self._rehydrate( - tuple(self._block_text(block) for block in blocks), self._session_id(data) - ) - for block, replacement in zip(blocks, restored): - if isinstance(block, dict): - block["text"] = replacement - else: - block.text = replacement + restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, restored): + write(replacement) return response async def async_post_call_streaming_iterator_hook( @@ -525,10 +619,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail): ) -> AsyncGenerator[Any, None]: """Restores original values incrementally, without buffering the stream. - Each choice is its own token stream, so the sliding window is tracked per - choice index. One shared window would splice the characters held back for - one choice onto the next. The windows are locals of this generator, so they - are scoped to a single stream and cannot leak between concurrent requests. + Each choice -- and each tool call within a choice -- is its own token stream, so + the sliding window is tracked per (choice index, tool call) pair. One shared + window would splice the characters held back for one stream onto another. The + windows are locals of this generator, so they are scoped to a single stream and + cannot leak between concurrent requests. """ if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: async for chunk in response: @@ -536,7 +631,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): return session_id: Final = self._session_id(request_data) - carries: Final[dict] = {} # mutable-ok: per-choice windows, local to this stream. + carries: Final[dict] = {} # mutable-ok: per-stream windows, local to this generator. last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: @@ -551,49 +646,151 @@ class LLMShieldProxyGuardrail(CustomGuardrail): yield trailing async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None: - """Restores one choice's delta, advancing that choice's own window.""" + """Restores one choice's delta, advancing that choice's own windows. + + Content and each tool call are separate token streams, so each gets its own + window: `(choice_index, None)` for content, `(choice_index, tool_call_index)` for + one tool call's accumulating `arguments`. A shared window would splice the text + held back for one stream onto another. + """ delta: Final = getattr(choice, "delta", None) if delta is None: return index: Final = _choice_index(choice) - carry: Final = carries.get(index, "") - text: Final = getattr(delta, "content", None) is_final: Final = bool(getattr(choice, "finish_reason", None)) + await self._restore_content_window(delta, (index, None), carries, session_id, is_final) + + for tool_call in getattr(delta, "tool_calls", None) or (): + await self._restore_tool_call_window(tool_call, index, carries, session_id) + + if is_final: + # A client parses a tool call's arguments when it sees the finish_reason, so + # every window this choice still holds has to land in *this* chunk. Flushing + # after it produces argument JSON the client has already stopped waiting for. + await self._flush_finished_choice(delta, index, carries, session_id) + + async def _restore_content_window( + self, + delta: Any, + key: tuple, + carries: _CarryWindows, + session_id: str, + is_final: bool, + ) -> None: + """Restores one delta's content through its own window.""" + carry: Final = carries.get(key, "") + text: Final = getattr(delta, "content", None) + if not isinstance(text, str) or not text: # Nothing to restore here, but a final chunk still has to flush the window. if is_final and carry: - flushed, flushed_carry = await self._stream_step("", carry, True, session_id) - carries[index] = flushed_carry # rebind-ok: this choice's window advances. + flushed, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. if flushed: delta.content = flushed return emitted, remaining = await self._stream_step(text, carry, is_final, session_id) - carries[index] = remaining # rebind-ok: this choice's window advances. + carries[key] = remaining # rebind-ok: this stream's window advances. delta.content = emitted + async def _restore_tool_call_window( + self, + tool_call: Any, + choice_index: int, + carries: _CarryWindows, + session_id: str, + ) -> None: + """Restores one streamed tool call's argument fragment. + + A tool call's `arguments` is a JSON document delivered as fragments that clients + concatenate per tool-call index, so each index gets a window of its own rather + than sharing the content stream's. + """ + tool_index: Final = _read_field(tool_call, "index") + if not isinstance(tool_index, int): + return + function: Final = _read_field(tool_call, "function") + if function is None: + return + arguments: Final = _read_field(function, "arguments") + if not isinstance(arguments, str) or not arguments: + return + + key: Final = (choice_index, tool_index) + emitted, remaining = await self._stream_step(arguments, carries.get(key, ""), False, session_id) + carries[key] = remaining # rebind-ok: this tool call's window advances. + _write_field(function, "arguments", emitted) + + async def _flush_finished_choice( + self, + delta: Any, + choice_index: int, + carries: _CarryWindows, + session_id: str, + ) -> None: + """Emits everything this finishing choice still holds, into this chunk. + + A client parses a tool call's `arguments` when the chunk carrying the + finish_reason arrives, so a flush delivered afterwards is too late -- the client + has already tried to parse truncated JSON. Content lands back on `content`; held + tool-call text is appended as an index-only continuation entry, which is the shape + clients concatenate by index, so no id or name is needed. Appending is correct + even when this chunk already carried a fragment for that tool call. + """ + continuations: Final[list] = [] # mutable-ok: built into this chunk's delta. + for key in sorted((held for held in carries if held[0] == choice_index), key=_carry_sort_key): + carry = carries[key] + if not carry: + continue + _, tool_index = key + text, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if not text: + continue + if tool_index is None: + delta.content = text + else: + continuations.append({"index": tool_index, "function": {"arguments": text}}) + if continuations: + existing: Final[list] = list(getattr(delta, "tool_calls", None) or []) + delta.tool_calls = existing + continuations + async def _flush_trailing( self, last_chunk: Any, carries: _CarryWindows, session_id: str ) -> AsyncGenerator[Any, None]: - """Empties every window still holding text, one chunk per choice. + """Empties every window still holding text, one chunk per window. + + This is the net for a stream that ended with no finish_reason at all; a stream + that ended with one is flushed into its own terminal chunk by + `_flush_finished_choice`, because that is the moment a client parses tool + arguments. Driven by the windows rather than by the last chunk's choices. A choice that finished earlier is not present in the terminal chunk, and flushing only what that chunk carries would drop its held text and truncate its answer. """ - for index in sorted(carries): - carry = carries[index] + for key in sorted(carries, key=_carry_sort_key): + carry = carries[key] if not carry: continue + choice_index, tool_index = key text, remaining = await self._stream_step("", carry, True, session_id) - carries[index] = remaining # rebind-ok: this choice's window advances. + carries[key] = remaining # rebind-ok: this stream's window advances. if not text: continue - chunk = self._chunk_for_choice(last_chunk, index) + chunk = self._chunk_for_choice(last_chunk, choice_index) if chunk is None: continue - chunk.choices[0].delta.content = text + if tool_index is None: + chunk.choices[0].delta.content = text + else: + # The copy carried this chunk's own content and tool calls, both already + # delivered. Replace rather than append, and drop the content, or the + # client sees them twice. + chunk.choices[0].delta.content = None + chunk.choices[0].delta.tool_calls = [{"index": tool_index, "function": {"arguments": text}}] yield chunk @staticmethod @@ -645,16 +842,46 @@ class LLMShieldProxyGuardrail(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: - texts: Final = inputs.get("texts") - if not texts: + """Unified entry point: what the UI's Test guardrail button and the translation + handlers call. + + `tool_calls` is handled on the response side only. LiteLLM populates the field here, + and on a reply it holds the model's tool arguments -- the same text the native hook + restores, and restoring one but not the other would leave the placeholder on + whichever path ran. The request side is left to the native pre-call hook, because + redacting it here as well would redact it twice. + """ + text_list: Final[list] = list(inputs.get("texts") or ()) + tool_calls: Final[list] = list(inputs.get("tool_calls") or ()) if input_type == "response" else [] + if not text_list and not tool_calls: return inputs + # Copied rather than mutated: the caller's tool calls are theirs to own, and this + # method's contract is to hand back a new mapping. + restored_calls: Final[list] = [copy.deepcopy(call) for call in tool_calls] # mutable-ok: a new list. + spans: Final[list] = list(text_list) # mutable-ok: ordered batch, frozen before the call. + writers: Final[list] = [] # mutable-ok: one per span appended below. + for call in restored_calls: + function: Final = _read_field(call, "function") + arguments: Final = _read_field(function, "arguments") if function is not None else None + if isinstance(arguments, str) and arguments: + spans.append(arguments) + writers.append(lambda new, f=function: _write_field(f, "arguments", new)) + replaced: Final = ( - await self._redact(tuple(texts), self._mint_session_id(request_data)) + await self._redact(tuple(spans), self._mint_session_id(request_data)) if input_type == "request" - else await self._rehydrate(tuple(texts), self._session_id(request_data)) + else await self._rehydrate(tuple(spans), self._session_id(request_data)) ) + restored_values: Final[list] = list(replaced) + + for write, replacement in zip(writers, restored_values[len(text_list) :]): + write(replacement) # Return a new mapping rather than rewriting the caller's, so this stays a # pure transform of the inputs it was handed. - merged: Final[JsonBody] = {**inputs, "texts": list(replaced)} # mutable-ok: TypedDict. + merged: Final[JsonBody] = {**inputs} # mutable-ok: TypedDict. + if text_list: + merged["texts"] = restored_values[: len(text_list)] + if restored_calls: + merged["tool_calls"] = restored_calls return merged