mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
refactor(realtime-guardrails): extract _send_output_text_done helper, fix silent except, clean up test imports
- Extract duplicated response.text.done guardrail handling into _send_output_text_done() — was copy-pasted across both _handle_provider_config_message and _handle_raw_backend_message - Add verbose_logger.debug on malformed backend message instead of silent pass - Move fastapi HTTPException import to top of test file - Remove double import litellm in _make_streaming helper - Drop unused test imports (asyncio, unittest, patch, Any, Optional)
This commit is contained in:
parent
9578efef61
commit
a7f1b840b1
2 changed files with 46 additions and 73 deletions
|
|
@ -368,6 +368,43 @@ class RealTimeStreaming:
|
|||
return True, safe_msg
|
||||
return False, ""
|
||||
|
||||
async def _send_output_text_done(
|
||||
self, event: Any, event_str: str
|
||||
) -> None:
|
||||
"""
|
||||
Handle a response.text.done event through the output guardrail.
|
||||
|
||||
If blocked: discards buffered deltas, sends replacement error delta+done.
|
||||
If clean: flushes buffered deltas then forwards the done event.
|
||||
"""
|
||||
full_text = event.get("text", "")
|
||||
blocked_out, error_msg = await self.run_realtime_output_guardrails(full_text)
|
||||
if blocked_out:
|
||||
self._pending_output_text_events.clear()
|
||||
error_delta_str = json.dumps(
|
||||
{
|
||||
"type": "response.text.delta",
|
||||
"delta": error_msg,
|
||||
"content_index": event.get("content_index", 0),
|
||||
"item_id": event.get("item_id", ""),
|
||||
"output_index": event.get("output_index", 0),
|
||||
"response_id": event.get("response_id", ""),
|
||||
}
|
||||
)
|
||||
error_done = dict(event)
|
||||
error_done["text"] = error_msg
|
||||
error_done_str = json.dumps(error_done)
|
||||
self.store_message(error_done_str)
|
||||
await self.websocket.send_text(error_delta_str)
|
||||
await self.websocket.send_text(error_done_str)
|
||||
else:
|
||||
for pending in self._pending_output_text_events:
|
||||
self.store_message(pending)
|
||||
await self.websocket.send_text(pending)
|
||||
self._pending_output_text_events.clear()
|
||||
self.store_message(event_str)
|
||||
await self.websocket.send_text(event_str)
|
||||
|
||||
async def _handle_provider_config_message(self, raw_response) -> None:
|
||||
"""Process a backend message when a provider_config is set (transformed path)."""
|
||||
returned_object = self.provider_config.transform_realtime_response( # type: ignore[union-attr]
|
||||
|
|
@ -444,35 +481,7 @@ class RealTimeStreaming:
|
|||
self._pending_output_text_events.append(event_str)
|
||||
continue
|
||||
if isinstance(event, dict) and event.get("type") == "response.text.done":
|
||||
full_text = event.get("text", "")
|
||||
blocked_out, error_msg = await self.run_realtime_output_guardrails(
|
||||
full_text
|
||||
)
|
||||
if blocked_out:
|
||||
self._pending_output_text_events.clear()
|
||||
error_delta_str = json.dumps(
|
||||
{
|
||||
"type": "response.text.delta",
|
||||
"delta": error_msg,
|
||||
"content_index": event.get("content_index", 0),
|
||||
"item_id": event.get("item_id", ""),
|
||||
"output_index": event.get("output_index", 0),
|
||||
"response_id": event.get("response_id", ""),
|
||||
}
|
||||
)
|
||||
error_done = dict(event)
|
||||
error_done["text"] = error_msg
|
||||
error_done_str = json.dumps(error_done)
|
||||
self.store_message(error_done_str)
|
||||
await self.websocket.send_text(error_delta_str)
|
||||
await self.websocket.send_text(error_done_str)
|
||||
else:
|
||||
for pending in self._pending_output_text_events:
|
||||
self.store_message(pending)
|
||||
await self.websocket.send_text(pending)
|
||||
self._pending_output_text_events.clear()
|
||||
self.store_message(event_str)
|
||||
await self.websocket.send_text(event_str)
|
||||
await self._send_output_text_done(event, event_str)
|
||||
continue
|
||||
## LOGGING
|
||||
self.store_message(event_str)
|
||||
|
|
@ -535,39 +544,11 @@ class RealTimeStreaming:
|
|||
return True
|
||||
|
||||
if event_obj.get("type") == "response.text.done":
|
||||
full_text = event_obj.get("text", "")
|
||||
blocked_out, error_msg = await self.run_realtime_output_guardrails(
|
||||
full_text
|
||||
)
|
||||
if blocked_out:
|
||||
self._pending_output_text_events.clear()
|
||||
error_delta_str = json.dumps(
|
||||
{
|
||||
"type": "response.text.delta",
|
||||
"delta": error_msg,
|
||||
"content_index": event_obj.get("content_index", 0),
|
||||
"item_id": event_obj.get("item_id", ""),
|
||||
"output_index": event_obj.get("output_index", 0),
|
||||
"response_id": event_obj.get("response_id", ""),
|
||||
}
|
||||
)
|
||||
error_done = dict(event_obj)
|
||||
error_done["text"] = error_msg
|
||||
error_done_str = json.dumps(error_done)
|
||||
self.store_message(error_done_str)
|
||||
await self.websocket.send_text(error_delta_str)
|
||||
await self.websocket.send_text(error_done_str)
|
||||
else:
|
||||
for pending in self._pending_output_text_events:
|
||||
self.store_message(pending)
|
||||
await self.websocket.send_text(pending)
|
||||
self._pending_output_text_events.clear()
|
||||
self.store_message(raw_response)
|
||||
await self.websocket.send_text(raw_response)
|
||||
await self._send_output_text_done(event_obj, raw_response)
|
||||
return True
|
||||
|
||||
except (json.JSONDecodeError, AttributeError):
|
||||
pass
|
||||
verbose_logger.debug("[realtime] skipped malformed backend message")
|
||||
return False
|
||||
|
||||
async def backend_to_client_send_messages(self):
|
||||
|
|
|
|||
|
|
@ -5,13 +5,14 @@ Tests that response.text.delta events are buffered and checked on
|
|||
response.text.done, blocking bad output and forwarding clean output.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import unittest
|
||||
from typing import Any, List, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from typing import List
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class _FakeWebSocket:
|
||||
|
|
@ -56,8 +57,6 @@ class _BlockingGuardrail(CustomGuardrail):
|
|||
texts = inputs.get("texts", [])
|
||||
for text in texts:
|
||||
if "bomb" in text.lower():
|
||||
from fastapi import HTTPException
|
||||
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "Response blocked by content filter."},
|
||||
|
|
@ -96,14 +95,7 @@ def _make_streaming(guardrail=None):
|
|||
# Patch store_message to be a no-op (avoids JSON parsing complexity)
|
||||
streaming.store_message = MagicMock()
|
||||
|
||||
if guardrail is not None:
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = [guardrail]
|
||||
else:
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
litellm.callbacks = [guardrail] if guardrail is not None else []
|
||||
|
||||
return streaming, client_ws, backend_ws
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue