fix(guardrails): preserve tool context during streamed text scans

This commit is contained in:
albertbausili 2026-09-28 12:36:42 +02:00
parent cb01259970
commit 400f845b5f
2 changed files with 19 additions and 4 deletions

View file

@ -794,6 +794,14 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
self.merge_user_api_key_metadata_into_request(request_data, user_api_key_dict)
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
if self._streamed_tool_call_fingerprints(responses_so_far):
assembled: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
inputs["tool_calls"] = [
converted
for choice in assembled.choices
for tool_call in choice.message.tool_calls or ()
if (converted := self._convert_tool_call_to_dict(tool_call)) is not None
]
if responses_so_far and getattr(responses_so_far[0], "model", None):
inputs["model"] = responses_so_far[0].model
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(

View file

@ -946,10 +946,11 @@ def _delta_text(item):
class _ToolRedactingGuardrail(CustomGuardrail):
def __init__(self) -> None:
def __init__(self, require_tool_context: bool = False) -> 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
self.require_tool_context = require_tool_context
async def apply_guardrail(
self,
@ -961,7 +962,10 @@ class _ToolRedactingGuardrail(CustomGuardrail):
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", ()))
texts: Final = tuple(
"checked:" + text.replace("SECRET", "MASKED") if calls or not self.require_tool_context else text
for text in inputs.get("texts", ())
)
return {
**inputs,
"texts": list(texts),
@ -984,10 +988,11 @@ class TestStreamingTransform:
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler})
@pytest.mark.asyncio
@pytest.mark.parametrize("require_tool_context", [False, True])
@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
self, include_text: bool, tool_count: int, require_tool_context: bool
) -> None:
chunks: Final = (
*([_stream_chunk("hello SECRET")] if include_text else []),
@ -1001,7 +1006,9 @@ class TestStreamingTransform:
]), finish_reason="tool_calls")]),
ModelResponseStream(choices=[], usage={"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}),
)
out: Final = await _drive_stream(UnifiedLLMGuardrails(), _ToolRedactingGuardrail(), chunks)
out: Final = await _drive_stream(
UnifiedLLMGuardrails(), _ToolRedactingGuardrail(require_tool_context=require_tool_context), chunks
)
calls: Final = tuple(
call for chunk in out for choice in chunk.choices for call in choice.delta.tool_calls or ()
)