mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
579b4c30be
commit
d598d1bf96
1 changed files with 298 additions and 71 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue