mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): scan streaming tool calls with guardrails
This commit is contained in:
parent
b5d1fb7298
commit
9bccb5bdc0
2 changed files with 86 additions and 1 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue