mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): apply buffered tool argument rewrites
This commit is contained in:
parent
080b5a486a
commit
5c437f7efb
6 changed files with 143 additions and 78 deletions
|
|
@ -556,7 +556,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
terminate the stream. Text rewrites are not propagated to the client here
|
||||
(see ``_process_streaming_transform`` for the incremental_diff path) unless
|
||||
``deliver_ended_stream_rewrites`` opts the ended-stream branch in."""
|
||||
has_stream_ended: Final = self._first_choice_has_finished(responses_so_far)
|
||||
has_stream_ended: Final = deliver_ended_stream_rewrites or self._first_choice_has_finished(responses_so_far)
|
||||
|
||||
if has_stream_ended:
|
||||
await self._process_ended_stream(
|
||||
|
|
@ -1132,20 +1132,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
def _function_tool_call_fragments(
|
||||
responses_so_far: Sequence["ModelResponseStream"],
|
||||
) -> tuple[tuple[ChatCompletionDeltaToolCall, ...], ...]:
|
||||
"""Group the stream's function tool-call fragments by their tool-call index, in
|
||||
the index order ``stream_chunk_builder`` lists the rebuilt tool calls, keeping
|
||||
only the indices the builder keeps (an id and a name somewhere in the stream)."""
|
||||
fragments: Final = tuple(
|
||||
tool_call
|
||||
(choice.index, tool_call)
|
||||
for response in responses_so_far
|
||||
for choice in response.choices
|
||||
for tool_call in choice.delta.tool_calls or ()
|
||||
if isinstance(tool_call, ChatCompletionDeltaToolCall)
|
||||
)
|
||||
identified: Final = frozenset(fragment.index for fragment in fragments if fragment.id)
|
||||
named: Final = frozenset(fragment.index for fragment in fragments if fragment.function.name)
|
||||
identified: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.id)
|
||||
named: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.function.name)
|
||||
return tuple(
|
||||
tuple(fragment for fragment in fragments if fragment.index == index) for index in sorted(identified & named)
|
||||
tuple(fragment for choice, fragment in fragments if (choice, fragment.index) == key)
|
||||
for key in sorted(identified & named)
|
||||
)
|
||||
|
||||
def _write_ended_stream_tool_call_rewrites(
|
||||
|
|
@ -1155,28 +1153,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...],
|
||||
guardrail_name: str,
|
||||
) -> None:
|
||||
"""Write ended-stream guardrail tool-call rewrites back across the buffered
|
||||
chunks: the rewritten name and full arguments land in the tool call's first
|
||||
fragment and the arguments of its later fragments are blanked, mirroring the
|
||||
text write-back. A rewrite on a stream carrying more than one distinct choice
|
||||
index, or whose fragments do not line up with the rebuilt tool calls, is
|
||||
reported as undeliverable, so the pipeline executor discards it and releases
|
||||
the original chunks."""
|
||||
post_guardrail_tool_calls: Final = self._function_tool_call_shapes(guardrailed_response)
|
||||
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
|
||||
return
|
||||
stream_choice_indices: Final = frozenset(
|
||||
choice.index for response in responses_so_far for choice in response.choices
|
||||
)
|
||||
fragments_by_tool_call: Final = self._function_tool_call_fragments(responses_so_far)
|
||||
if len(stream_choice_indices) != 1:
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
raise UndeliverableStreamRewrite(
|
||||
guardrail_name,
|
||||
f"the stream carries {len(stream_choice_indices)} choices and tool-call rewrites are only written "
|
||||
"back on single-choice streams",
|
||||
)
|
||||
if len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ guardrails:
|
|||
collector_key: os.environ/TRUSTGUARD_COLLECTOR_KEY # tgcol_… ; optional if the API key is bound
|
||||
unreachable_fallback: fail_closed
|
||||
timeout: 5
|
||||
streaming_transform_mode: block_only # incremental_diff to stream redacted output
|
||||
streaming_transform_mode: incremental_diff
|
||||
default_on: true
|
||||
```
|
||||
|
||||
|
|
@ -62,9 +62,9 @@ Under `incremental_diff` the reply is held until the end-of-stream evaluate retu
|
|||
|
||||
`incremental_diff` covers OpenAI chat completions streaming; other surfaces fall back to `block_only`.
|
||||
|
||||
Under `incremental_diff`, LiteLLM buffers tool-call deltas until TrustGuard inspects the assembled response. A blocking verdict releases no tool-call arguments or finish signal. Allowed tool calls retain their original deltas and order
|
||||
Under `incremental_diff`, LiteLLM buffers tool-call deltas until TrustGuard inspects the assembled response. A blocking verdict releases no tool-call arguments or finish signal. Tool calls retain their IDs and order, with transformed arguments written into the buffered deltas before delivery
|
||||
|
||||
The default `block_only` mode still forwards tool-call deltas before inspection. Use `incremental_diff` or non-streaming requests when tool calls must be checked before delivery. Streamed tool-call rewrites are not supported; use non-streaming requests for transformed tool arguments
|
||||
The default `block_only` mode still forwards tool-call deltas before inspection. Use `incremental_diff` or non-streaming requests when tool calls must be checked before delivery. The example selects `incremental_diff` for inspected text and tool arguments
|
||||
|
||||
## References
|
||||
|
||||
|
|
|
|||
|
|
@ -736,28 +736,14 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
async for out in _round(item, is_final=False):
|
||||
yield out
|
||||
|
||||
# v1 does not transform streamed tool calls, but they must still go
|
||||
# through the guardrail's block decision. Run the block_only inspection
|
||||
# over the full assembled response so tool calls cannot bypass it.
|
||||
#
|
||||
# Pass a deep copy of responses_so_far — the block path routes through
|
||||
# ``_process_streaming_block_only`` which mutates ``delta.content``
|
||||
# in-place on the chunk objects it receives. For an n>1 chunk carrying
|
||||
# text on one choice and tool_calls (with finish_reason) on another,
|
||||
# ``has_stream_ended`` reads ``choices[0]`` alone and can miss the
|
||||
# terminal signal, letting the block path rewrite the raw accumulator.
|
||||
# The subsequent final ``_round`` would then re-read the already-mutated
|
||||
# text, producing double-application for a non-idempotent guardrail or a
|
||||
# ``stream_transform_underflow`` 400 from mismatched prefixes. A shallow
|
||||
# list copy wouldn't help — the mutation is on the chunk objects
|
||||
# themselves — so we deepcopy.
|
||||
if saw_tool_calls:
|
||||
async for out in self._inspect_full_response_for_block(
|
||||
inspected_responses: Final = copy.deepcopy(responses_so_far)
|
||||
async for out in self._inspect_full_response(
|
||||
endpoint_translation=endpoint_translation,
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
request_data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
responses_so_far=copy.deepcopy(responses_so_far),
|
||||
responses_so_far=inspected_responses,
|
||||
responses_yielded=responses_yielded,
|
||||
):
|
||||
yield out
|
||||
|
|
@ -771,7 +757,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
finish_reason_per_choice=finish_reason_per_choice,
|
||||
held_choices=_held_choices(held_chars_per_choice),
|
||||
)
|
||||
for buffered_item in responses_so_far
|
||||
for buffered_item in inspected_responses
|
||||
if self._chunk_has_tool_calls(buffered_item)
|
||||
)
|
||||
|
||||
|
|
@ -826,7 +812,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
responses_yielded.append(trailing)
|
||||
yield trailing
|
||||
|
||||
async def _inspect_full_response_for_block(
|
||||
async def _inspect_full_response(
|
||||
self,
|
||||
*,
|
||||
endpoint_translation: _EndpointTranslation,
|
||||
|
|
@ -836,16 +822,8 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
responses_so_far: Sequence[object],
|
||||
responses_yielded: Sequence[object],
|
||||
) -> AsyncGenerator[object, None]:
|
||||
"""Run the block-only guardrail inspection over the full assembled
|
||||
response (text + tool calls) so nothing bypasses the block decision.
|
||||
|
||||
The guardrail's returned transforms are discarded here (v1 does not
|
||||
transform tool calls); only its block decision matters. A block is
|
||||
surfaced the same way as elsewhere: ModifyResponseException terminates the
|
||||
stream via the shared block handler; a GenericGuardrailAPI block raises and
|
||||
propagates, matching block_only.
|
||||
"""
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
try:
|
||||
await endpoint_translation.process_output_streaming_response(
|
||||
|
|
@ -855,7 +833,10 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
stream_transform_sink=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
except UndeliverableStreamRewrite as exc:
|
||||
raise HTTPException(status_code=400, detail="Guardrail stream rewrite could not be applied") from exc
|
||||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = responses_so_far
|
||||
|
|
|
|||
|
|
@ -54,7 +54,7 @@ class NeuralTrustGuardrailConfigModel(GuardrailConfigModel):
|
|||
description=(
|
||||
"How a `transform` verdict reaches a streaming client. `block_only` (default) streams the raw "
|
||||
"model chunks, so `block` and `ask` still end the stream but the redacted text is dropped. "
|
||||
"`incremental_diff` withholds the model chunks and streams TrustGuard's rewritten text instead: "
|
||||
"`incremental_diff` withholds the model chunks and streams TrustGuard's rewritten text and tool arguments: "
|
||||
"the reply arrives once the end-of-stream evaluate returns, and a blocking verdict ends the "
|
||||
"stream with nothing already sent. OpenAI chat completions streaming only."
|
||||
),
|
||||
|
|
|
|||
|
|
@ -2,7 +2,10 @@
|
|||
|
||||
import logging
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from typing import Final, Literal
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -41,7 +44,10 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai
|
|||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.utils import CallTypes, Delta, ModelResponseStream, StreamingChoices
|
||||
from litellm.types.utils import (
|
||||
CallTypes, ChatCompletionMessageToolCall, Delta, GenericGuardrailAPIInputs,
|
||||
ModelResponseStream, StreamingChoices,
|
||||
)
|
||||
|
||||
|
||||
class RecordingGuardrail(CustomGuardrail):
|
||||
|
|
@ -947,6 +953,36 @@ def _delta_text(item):
|
|||
return item.choices[0].delta.content or ""
|
||||
|
||||
|
||||
class _ToolRedactingGuardrail(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(guardrail_name="tool-redactor", event_hook=GuardrailEventHooks.post_call, default_on=True)
|
||||
self.streaming_transform_mode = "incremental_diff"
|
||||
self.streaming_sampling_rate = 1
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: object | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
calls: Final = TypeAdapter(tuple[ChatCompletionMessageToolCall, ...]).validate_python(
|
||||
inputs.get("tool_calls", ())
|
||||
)
|
||||
texts: Final = tuple("checked:" + text.replace("SECRET", "MASKED") for text in inputs.get("texts", ()))
|
||||
return {
|
||||
**inputs,
|
||||
"texts": list(texts),
|
||||
"stream_holdback_chars": [len(text) for text in texts],
|
||||
"tool_calls": [
|
||||
call.model_copy(update={"function": call.function.model_copy(update={
|
||||
"arguments": call.function.arguments.replace("SECRET", "MASKED"),
|
||||
})})
|
||||
for call in calls
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class TestStreamingTransform:
|
||||
"""Streaming text-transformation (incremental_diff) path on the OpenAI chat
|
||||
completions streaming surface."""
|
||||
|
|
@ -955,6 +991,77 @@ class TestStreamingTransform:
|
|||
def _use_openai_handler_mapping(self, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler})
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("include_text", [False, True])
|
||||
@pytest.mark.parametrize("tool_count", [1, 2])
|
||||
async def test_buffered_tool_arguments_are_rewritten_before_delivery(
|
||||
self, include_text: bool, tool_count: int
|
||||
) -> None:
|
||||
chunks: Final = (
|
||||
*([_stream_chunk("hello SECRET")] if include_text else []),
|
||||
ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[
|
||||
{"index": index, "id": f"call_{index}", "type": "function",
|
||||
"function": {"name": "contact", "arguments": '{"contact":"SEC'}}
|
||||
for index in range(tool_count)
|
||||
]))]),
|
||||
ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[
|
||||
{"index": index, "function": {"arguments": 'RET"}'}} for index in range(tool_count)
|
||||
]), finish_reason="tool_calls")]),
|
||||
ModelResponseStream(choices=[], usage={"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}),
|
||||
)
|
||||
out: Final = await _drive_stream(UnifiedLLMGuardrails(), _ToolRedactingGuardrail(), chunks)
|
||||
calls: Final = tuple(
|
||||
call for chunk in out for choice in chunk.choices for call in choice.delta.tool_calls or ()
|
||||
)
|
||||
for index in range(tool_count):
|
||||
arguments: Final = "".join(call.function.arguments or "" for call in calls if call.index == index)
|
||||
assert arguments == '{"contact":"MASKED"}'
|
||||
assert next(call.id for call in calls if call.index == index and call.id) == f"call_{index}"
|
||||
assert "".join(_delta_text(chunk) for chunk in out) == ("checked:hello MASKED" if include_text else "")
|
||||
assert any(choice.finish_reason == "tool_calls" for chunk in out for choice in chunk.choices)
|
||||
assert out[-1].usage.total_tokens == 18
|
||||
assert all("SECRET" not in chunk.model_dump_json() for chunk in out)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_rewrites_keep_completion_choices_separate(self) -> None:
|
||||
async def response() -> AsyncIterator[ModelResponseStream]:
|
||||
yield ModelResponseStream(choices=[StreamingChoices(
|
||||
index=index,
|
||||
delta=Delta(tool_calls=[{"index": 0, "id": f"call_{index}", "type": "function",
|
||||
"function": {"name": "contact", "arguments": '{"contact":"SECRET"}'}}]),
|
||||
finish_reason="tool_calls",
|
||||
) for index in range(2)])
|
||||
|
||||
iterator: Final = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(request_route="/v1/chat/completions"),
|
||||
response=response(), request_data={"guardrail_to_apply": _ToolRedactingGuardrail()},
|
||||
)
|
||||
chunks: Final = tuple([chunk async for chunk in iterator])
|
||||
calls: Final = tuple(
|
||||
(choice.index, call.id, call.function.arguments)
|
||||
for chunk in chunks for choice in chunk.choices for call in choice.delta.tool_calls or ()
|
||||
)
|
||||
assert calls == ((0, "call_0", '{"contact":"MASKED"}'), (1, "call_1", '{"contact":"MASKED"}'))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_undeliverable_rewrite_is_a_closed_failure(self) -> None:
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
class UndeliverableTranslation(OpenAIChatCompletionsHandler):
|
||||
async def process_output_streaming_response(
|
||||
self, *args: object, **kwargs: object
|
||||
) -> list[ModelResponseStream]:
|
||||
raise UndeliverableStreamRewrite("tool-redactor", "unmappable tool fragments")
|
||||
|
||||
iterator: Final = UnifiedLLMGuardrails()._inspect_full_response(
|
||||
endpoint_translation=UndeliverableTranslation(), guardrail_to_apply=_ToolRedactingGuardrail(),
|
||||
request_data={}, user_api_key_dict=UserAPIKeyAuth(), responses_so_far=(), responses_yielded=(),
|
||||
)
|
||||
with pytest.raises(unified_module.HTTPException) as error:
|
||||
await anext(iterator)
|
||||
assert error.value.status_code == 400
|
||||
assert error.value.detail == "Guardrail stream rewrite could not be applied"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_only_drops_text_rewrites(self):
|
||||
"""Default block_only: the guardrail's uppercasing never reaches the
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ with guardrail transformations, including tool calls.
|
|||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Literal, Optional
|
||||
from typing import Any, Final, Literal, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -1400,24 +1400,21 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
] == [("call_1", '{"fruit": "persimmon"}'), ("call_2", '{"fruit": "durian"}')]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_ended_stream_tool_call_rewrite_on_multi_choice_stream_fails_closed(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = self._two_choice_tool_call_stream_chunks()
|
||||
|
||||
with pytest.raises(UndeliverableStreamRewrite, match="the stream carries 2 choices") as raised:
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=MockGuardrail(guardrail_name="test"),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
|
||||
assert raised.value.guardrail_name == "test"
|
||||
assert raised.value.reason == (
|
||||
"the stream carries 2 choices and tool-call rewrites are only written back on single-choice streams"
|
||||
async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None:
|
||||
handler: Final = OpenAIChatCompletionsHandler()
|
||||
chunks: Final = self._two_choice_tool_call_stream_chunks()
|
||||
await handler.process_output_streaming_response(
|
||||
responses_so_far=chunks,
|
||||
guardrail_to_apply=MockGuardrail(guardrail_name="test"),
|
||||
litellm_logging_obj=None,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
arguments: Final = tuple(
|
||||
"".join(call.function.arguments for chunk in chunks for choice in chunk.choices
|
||||
if choice.index == index for call in choice.delta.tool_calls or ())
|
||||
for index in range(2)
|
||||
)
|
||||
assert arguments == ('{"fruit": "PERSIMMON"}', '{"fruit": "DURIAN"}')
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deliver_ended_stream_clean_multi_choice_stream_released_untouched(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue