mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix(realtime): run guardrails on function_call_output content
Tool result outputs are client-controlled and fed to the model, so they must pass the same content checks as user text messages. Otherwise an attacker can smuggle blocked content into a function_call_output and have the model process it.
This commit is contained in:
parent
295a9e6e13
commit
16e44cbee8
2 changed files with 174 additions and 0 deletions
|
|
@ -879,6 +879,24 @@ class RealTimeStreaming:
|
|||
# include it as response instructions.
|
||||
self._pending_guardrail_message = combined_text
|
||||
continue # don't forward the original blocked message
|
||||
elif item.get("type") == "function_call_output":
|
||||
# Tool results are client-controlled and fed to the
|
||||
# model; check them with the same guardrail used for
|
||||
# user text so an attacker cannot smuggle blocked
|
||||
# content into a function_call_output.
|
||||
output = item.get("output", "")
|
||||
output_text = (
|
||||
output
|
||||
if isinstance(output, str)
|
||||
else json.dumps(output)
|
||||
)
|
||||
if output_text:
|
||||
blocked = await self.run_realtime_guardrails(
|
||||
output_text
|
||||
)
|
||||
if blocked:
|
||||
self._pending_guardrail_message = output_text
|
||||
continue
|
||||
|
||||
if (
|
||||
msg_type == "response.create"
|
||||
|
|
|
|||
|
|
@ -873,6 +873,162 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error():
|
|||
litellm.callbacks = [] # cleanup
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realtime_function_call_output_guardrail_blocks_and_returns_error():
|
||||
"""
|
||||
Test that a client-supplied function_call_output whose content triggers a
|
||||
guardrail is blocked: it is not forwarded to the backend, and an error
|
||||
event is sent to the client.
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
class BlockingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self, inputs, request_data, input_type, logging_obj=None
|
||||
):
|
||||
texts = inputs.get("texts", [])
|
||||
for text in texts:
|
||||
if "@" in text:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "email address detected"},
|
||||
)
|
||||
return inputs
|
||||
|
||||
guardrail = BlockingGuardrail(
|
||||
guardrail_name="email-blocker",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
default_on=True,
|
||||
)
|
||||
litellm.callbacks = [guardrail]
|
||||
|
||||
client_ws = MagicMock()
|
||||
client_ws.send_text = AsyncMock()
|
||||
|
||||
backend_ws = MagicMock()
|
||||
backend_ws.send = AsyncMock()
|
||||
backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None))
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.pre_call = MagicMock()
|
||||
|
||||
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
|
||||
|
||||
item_create_msg = json.dumps(
|
||||
{
|
||||
"type": "conversation.item.create",
|
||||
"item": {
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_123",
|
||||
"output": "Tool says: my email is test@example.com",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
client_ws.receive_text = AsyncMock(
|
||||
side_effect=[
|
||||
item_create_msg,
|
||||
Exception("connection closed"),
|
||||
]
|
||||
)
|
||||
|
||||
await streaming.client_ack_messages()
|
||||
|
||||
sent_texts = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list]
|
||||
error_events = [e for e in sent_texts if e.get("type") == "error"]
|
||||
assert len(error_events) == 1, f"Expected one error event, got: {sent_texts}"
|
||||
assert error_events[0]["error"]["type"] == "guardrail_violation"
|
||||
|
||||
sent_to_backend = [c.args[0] for c in backend_ws.send.call_args_list if c.args]
|
||||
forwarded_tool_outputs = [
|
||||
json.loads(m)
|
||||
for m in sent_to_backend
|
||||
if isinstance(m, str)
|
||||
and json.loads(m).get("type") == "conversation.item.create"
|
||||
and json.loads(m).get("item", {}).get("type") == "function_call_output"
|
||||
]
|
||||
assert len(forwarded_tool_outputs) == 0, (
|
||||
f"Blocked function_call_output should not be forwarded, got: "
|
||||
f"{forwarded_tool_outputs}"
|
||||
)
|
||||
|
||||
litellm.callbacks = [] # cleanup
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realtime_function_call_output_guardrail_allows_clean_output():
|
||||
"""
|
||||
Test that a clean function_call_output passes through and reaches the backend
|
||||
when guardrails are configured.
|
||||
"""
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
class BlockingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self, inputs, request_data, input_type, logging_obj=None
|
||||
):
|
||||
return inputs
|
||||
|
||||
guardrail = BlockingGuardrail(
|
||||
guardrail_name="noop",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
default_on=True,
|
||||
)
|
||||
litellm.callbacks = [guardrail]
|
||||
|
||||
client_ws = MagicMock()
|
||||
client_ws.send_text = AsyncMock()
|
||||
|
||||
backend_ws = MagicMock()
|
||||
backend_ws.send = AsyncMock()
|
||||
backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None))
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.pre_call = MagicMock()
|
||||
|
||||
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
|
||||
|
||||
item_create_msg = json.dumps(
|
||||
{
|
||||
"type": "conversation.item.create",
|
||||
"item": {
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_456",
|
||||
"output": '{"temperature": 72, "unit": "F"}',
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
client_ws.receive_text = AsyncMock(
|
||||
side_effect=[
|
||||
item_create_msg,
|
||||
Exception("connection closed"),
|
||||
]
|
||||
)
|
||||
|
||||
await streaming.client_ack_messages()
|
||||
|
||||
sent_to_backend = [c.args[0] for c in backend_ws.send.call_args_list if c.args]
|
||||
forwarded = [
|
||||
json.loads(m)
|
||||
for m in sent_to_backend
|
||||
if isinstance(m, str)
|
||||
and json.loads(m).get("type") == "conversation.item.create"
|
||||
and json.loads(m).get("item", {}).get("type") == "function_call_output"
|
||||
]
|
||||
assert (
|
||||
len(forwarded) == 1
|
||||
), f"Clean function_call_output should be forwarded, got: {forwarded}"
|
||||
|
||||
litellm.callbacks = [] # cleanup
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realtime_text_input_guardrail_uses_pre_call_mode():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue