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.
_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