fix(proxy): scan streaming tool calls with guardrails

This commit is contained in:
Ultronen 2026-09-06 18:16:42 +08:00
parent b5d1fb7298
commit 9bccb5bdc0
2 changed files with 86 additions and 1 deletions

View file

@ -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
)

View file

@ -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):