diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 0ce7e63488c..595ce96e534 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index c13625296f8..98269add1ad 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -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 () )