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.
This commit is contained in:
Ninad Phalak 2026-09-13 11:52:14 -05:00
parent 579b4c30be
commit d598d1bf96
No known key found for this signature in database
GPG key ID: 59119ED515433744

View file

@ -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. # _locate_request_texts, which freezes it into a tuple before returning.
_SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors.
# Sliding windows keyed by streaming choice index, threaded through one stream. # Sliding windows keyed by (choice index, tool-call index | None), threaded through one
_CarryWindows: TypeAlias = dict # mutable-ok: per-choice windows advanced in place. # 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 # A caller-owned list whose entries are rewritten in place, such as a Completions
# `prompt` sent as an array of strings. # `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 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): class LLMShieldProxyGuardrail(CustomGuardrail):
"""Redacts PII before it leaves the proxy and restores it in the response. """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): for (_, write), replacement in zip(slots, redacted):
write(replacement) 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 @log_guardrail_information
async def async_post_call_success_hook( async def async_post_call_success_hook(
self, self,
@ -428,27 +498,46 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
if self._is_anthropic_message_response(response): if self._is_anthropic_message_response(response):
return await self._restore_anthropic_response(response, data) return await self._restore_anthropic_response(response, data)
text_blocks: Final = self._responses_api_text_blocks(response) response_slots: Final = self._responses_api_slots(response)
if text_blocks: if response_slots:
return await self._restore_responses_api_response(response, text_blocks, data) return await self._restore_responses_api_response(response, response_slots, data)
choices: Final = getattr(response, "choices", None) choices: Final = getattr(response, "choices", None)
if not choices: if not choices:
return response return response
pending: Final = tuple( # One batch for every restorable span in the reply, collected in document order:
(choice.message, choice.message.content) # the shield maps its answers back by position. A second round trip is not an
for choice in choices # option here -- /v1/guard/rehydrate caps a batch at 256 texts and 1,000,000
if getattr(choice, "message", None) is not None # characters, and `_same_length_or_raise` is what guarantees the positional
and isinstance(getattr(choice.message, "content", None), str) # mapping -- so a reply carrying more spans than that fails closed, which is this
and choice.message.content # 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: if not pending:
return response return response
restored: Final = await self._rehydrate(tuple(text for _, text in pending), self._session_id(data)) restored: Final = await self._rehydrate(tuple(text for text, _ in pending), self._session_id(data))
for (message, _), replacement in zip(pending, restored): for (_, write), replacement in zip(pending, restored):
message.content = replacement write(replacement)
return response return response
@staticmethod @staticmethod
@ -461,60 +550,65 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
) )
async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: 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 This shape has no `choices`, so without its own branch the reply would go
back to the caller still carrying placeholders. 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( slots: Final[list] = [] # mutable-ok: accumulator, frozen before use.
block for block in response["content"]:
for block in response["content"] if not isinstance(block, dict):
if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) continue
) kind: Final = block.get("type")
if not blocks: 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 return response
restored: Final = await self._rehydrate(tuple(block["text"] for block in blocks), self._session_id(data)) restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data))
for block, replacement in zip(blocks, restored): for (_, write), replacement in zip(slots, restored):
block["text"] = replacement write(replacement)
return response return response
@staticmethod @staticmethod
def _responses_api_text_blocks(response: object) -> Sequence[object]: def _responses_api_slots(response: object) -> Sequence[_Slot]:
"""Text blocks in a Responses API reply. """Restorable spans in a Responses API reply.
That shape carries `output` items rather than `choices`, so it needs its own 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 walk; without one the reply goes back to the caller still holding
placeholders even though the request was redacted correctly. Blocks come placeholders even though the request was redacted correctly. Items and blocks
through as dicts or as objects depending on how far the reply has been come through as dicts or as objects depending on how far the reply has been
deserialised, so both are handled. 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 item in getattr(response, "output", None) or ():
for block in getattr(item, "content", None) or (): for block in getattr(item, "content", None) or ():
if isinstance(block, dict): text: Final = _read_field(block, "text")
if isinstance(block.get("text"), str) and block["text"]: if isinstance(text, str) and text:
blocks.append(block) slots.append((text, lambda new, b=block: _write_field(b, "text", new)))
elif isinstance(getattr(block, "text", None), str) and block.text: for field in ("arguments", "output"):
blocks.append(block) value: Final = _read_field(item, field)
return tuple(blocks) if isinstance(value, str) and value:
slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new)))
@staticmethod return tuple(slots)
def _block_text(block: object) -> str:
return block["text"] if isinstance(block, dict) else block.text
async def _restore_responses_api_response( async def _restore_responses_api_response(
self, response: Any, blocks: Sequence[object], data: MutableRequest self, response: Any, slots: Sequence[_Slot], data: MutableRequest
) -> Any: ) -> Any:
"""Puts the original values back into a Responses API reply.""" """Puts the original values back into a Responses API reply."""
restored: Final = await self._rehydrate( restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data))
tuple(self._block_text(block) for block in blocks), self._session_id(data) for (_, write), replacement in zip(slots, restored):
) write(replacement)
for block, replacement in zip(blocks, restored):
if isinstance(block, dict):
block["text"] = replacement
else:
block.text = replacement
return response return response
async def async_post_call_streaming_iterator_hook( async def async_post_call_streaming_iterator_hook(
@ -525,10 +619,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
) -> AsyncGenerator[Any, None]: ) -> AsyncGenerator[Any, None]:
"""Restores original values incrementally, without buffering the stream. """Restores original values incrementally, without buffering the stream.
Each choice is its own token stream, so the sliding window is tracked per Each choice -- and each tool call within a choice -- is its own token stream, so
choice index. One shared window would splice the characters held back for the sliding window is tracked per (choice index, tool call) pair. One shared
one choice onto the next. The windows are locals of this generator, so they window would splice the characters held back for one stream onto another. The
are scoped to a single stream and cannot leak between concurrent requests. 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: if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True:
async for chunk in response: async for chunk in response:
@ -536,7 +631,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
return return
session_id: Final = self._session_id(request_data) 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. last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush.
async for chunk in response: async for chunk in response:
@ -551,49 +646,151 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
yield trailing yield trailing
async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None: 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) delta: Final = getattr(choice, "delta", None)
if delta is None: if delta is None:
return return
index: Final = _choice_index(choice) 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)) 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: if not isinstance(text, str) or not text:
# Nothing to restore here, but a final chunk still has to flush the window. # Nothing to restore here, but a final chunk still has to flush the window.
if is_final and carry: if is_final and carry:
flushed, flushed_carry = await self._stream_step("", carry, True, session_id) flushed, remaining = await self._stream_step("", carry, True, session_id)
carries[index] = flushed_carry # rebind-ok: this choice's window advances. carries[key] = remaining # rebind-ok: this stream's window advances.
if flushed: if flushed:
delta.content = flushed delta.content = flushed
return return
emitted, remaining = await self._stream_step(text, carry, is_final, session_id) 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 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( async def _flush_trailing(
self, last_chunk: Any, carries: _CarryWindows, session_id: str self, last_chunk: Any, carries: _CarryWindows, session_id: str
) -> AsyncGenerator[Any, None]: ) -> 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 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 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. that chunk carries would drop its held text and truncate its answer.
""" """
for index in sorted(carries): for key in sorted(carries, key=_carry_sort_key):
carry = carries[index] carry = carries[key]
if not carry: if not carry:
continue continue
choice_index, tool_index = key
text, remaining = await self._stream_step("", carry, True, session_id) 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: if not text:
continue continue
chunk = self._chunk_for_choice(last_chunk, index) chunk = self._chunk_for_choice(last_chunk, choice_index)
if chunk is None: if chunk is None:
continue 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 yield chunk
@staticmethod @staticmethod
@ -645,16 +842,46 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
input_type: Literal["request", "response"], input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None, logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs: ) -> GenericGuardrailAPIInputs:
texts: Final = inputs.get("texts") """Unified entry point: what the UI's Test guardrail button and the translation
if not texts: 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 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 = ( 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" 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 # Return a new mapping rather than rewriting the caller's, so this stays a
# pure transform of the inputs it was handed. # 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 return merged