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.
|
# _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
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue