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:
Ishaan Jaffer 2026-02-25 20:57:56 -08:00
parent 9578efef61
commit a7f1b840b1
2 changed files with 46 additions and 73 deletions

View file

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

View file

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