From 9bccb5bdc02d1e7555e63a01e361c98fdef8473a Mon Sep 17 00:00:00 2001 From: Ultronen <82553854+Ultronen@users.noreply.github.com> Date: Sun, 6 Sep 2026 18:16:42 +0800 Subject: [PATCH] fix(proxy): scan streaming tool calls with guardrails --- litellm/proxy/utils.py | 32 ++++++++++- .../proxy_logging/test_streaming_hooks.py | 55 +++++++++++++++++++ 2 files changed, 86 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b53b6bf04f2..e1ec9444a6e 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -903,13 +903,38 @@ def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, return complete_response +def _streaming_guardrail_response_text(*, complete_response: str, response: object) -> str: + if not isinstance(response, (ModelResponse, ModelResponseStream)): + return complete_response + response_dict: Final = response.model_dump(mode="json", exclude_none=True) + structured_choices: Final = tuple( + {key: delta[key] for key in ("tool_calls", "function_call") if delta.get(key) is not None} + for choice in response_dict.get("choices", []) + if isinstance(choice, dict) + for delta in (choice.get("delta"),) + if isinstance(delta, dict) and any(delta.get(key) is not None for key in ("tool_calls", "function_call")) + ) + if not structured_choices: + return complete_response + return _StreamingHookResponseText( + json.dumps( + {"content": str(complete_response), "choices": structured_choices}, + ensure_ascii=False, + separators=(",", ":"), + sort_keys=True, + ) + ) + + def _is_unchanged_structured_streaming_hook_response( *, callback_response: object, complete_response: str, response_str: str, response: object ) -> bool: - if response_str != "" or not isinstance(response, (ModelResponse, ModelResponseStream)): + if not isinstance(response, (ModelResponse, ModelResponseStream)): return False if isinstance(complete_response, _StreamingHookResponseText): return callback_response is complete_response + if response_str != "": + return False return callback_response == complete_response @@ -3539,6 +3564,11 @@ class ProxyLogging: str_so_far=str_so_far, response=response, ) + if isinstance(_callback, CustomGuardrail): + complete_response = _streaming_guardrail_response_text( + complete_response=complete_response, + response=response, + ) callback_response: ( str | ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index 5478bac75bf..97573e62b7e 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -18,6 +18,7 @@ import pytest from fastapi import HTTPException import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( @@ -364,6 +365,60 @@ async def test_async_post_call_streaming_hook_preserves_tool_calls_when_callback assert suppressed == "" +@pytest.mark.asyncio +async def test_async_post_call_streaming_hook_exposes_tool_calls_to_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch +): + class _BlockingGuardrail(CustomGuardrail): + def should_run_guardrail(self, data, event_type): + return True + + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + if "blocked-command" in response: + return "data: blocked\n\n" + return response + + monkeypatch.setattr( + litellm, + "callbacks", + [_BlockingGuardrail(guardrail_name="tool-call-scanner")], + ) + + response = litellm.ModelResponseStream( + id="chatcmpl-guardrail-tools", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "id": "call-shell", + "type": "function", + "function": { + "name": "run_command", + "arguments": '{"command":"blocked-command"}', + }, + } + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + + out = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=response, + user_api_key_dict=make_user_api_key_auth(), + ) + + assert out == "data: blocked\n\n" + + @pytest.mark.asyncio async def test_async_post_call_streaming_hook_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch): class _Per(CustomLogger):