This commit is contained in:
Sameer Kankute 2026-04-27 19:05:59 +05:30
parent b719e35ed9
commit 817792c6cc
No known key found for this signature in database
3 changed files with 80 additions and 8 deletions

View file

@ -78,6 +78,9 @@ class RealTimeStreaming:
# Track whether session.created has already been sent to the client
# (e.g. synthetic event in deferred setup mode).
self._session_created_sent_to_client: bool = False
# Track whether we have already sent the guardrail turn-detection update
# that disables provider auto-response for transcription guardrails.
self._guardrail_turn_detection_update_sent: bool = False
def _should_store_message(
self,
@ -245,6 +248,22 @@ class RealTimeStreaming:
except (json.JSONDecodeError, TypeError):
return
async def _maybe_send_guardrail_turn_detection_update(self) -> None:
"""Disable provider auto-response once when transcription guardrails are enabled."""
if self._guardrail_turn_detection_update_sent:
return
if not self._has_audio_transcription_guardrails():
return
await self._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {"turn_detection": {"create_response": False}},
}
)
)
self._guardrail_turn_detection_update_sent = True
def _has_realtime_guardrails(self) -> bool:
"""Return True if any callback is registered for realtime guardrail event types."""
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -446,6 +465,7 @@ class RealTimeStreaming:
event_str = json.dumps(event)
if isinstance(event, dict) and event.get("type") == "session.created":
if self._session_created_sent_to_client:
await self._maybe_send_guardrail_turn_detection_update()
verbose_logger.debug(
"Skipping duplicate session.created from provider stream"
)
@ -459,14 +479,7 @@ class RealTimeStreaming:
):
self.store_message(event_str)
await self.websocket.send_text(event_str)
await self._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {"turn_detection": {"create_response": False}},
}
)
)
await self._maybe_send_guardrail_turn_detection_update()
continue
## GUARDRAIL: run on transcription events in provider_config path too
if (

View file

@ -5167,6 +5167,7 @@ class BaseLLMHTTPHandler:
if synthetic_session is not None:
await websocket.send_text(json.dumps(synthetic_session))
realtime_streaming._session_created_sent_to_client = True
await realtime_streaming._maybe_send_guardrail_turn_detection_update()
verbose_logger.debug(
"Sent synthetic session.created to client to unblock connection"
)

View file

@ -1022,3 +1022,61 @@ async def test_provider_path_suppresses_duplicate_session_created_after_syntheti
assert not any(
payload.get("type") == "session.created" for payload in sent_payloads
), f"Expected duplicate session.created to be suppressed, got: {sent_payloads}"
@pytest.mark.asyncio
async def test_duplicate_session_created_still_triggers_guardrail_turn_detection_update():
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)]
)
backend_ws.send = AsyncMock()
provider_config = MagicMock()
provider_config.transform_realtime_response = MagicMock(
return_value={
"response": [
{
"type": "session.created",
"event_id": "event_1",
"session": {"id": "sess_1", "modalities": ["audio"]},
}
],
"current_output_item_id": None,
"current_response_id": None,
"current_delta_chunks": [],
"current_conversation_id": None,
"current_item_chunks": [],
"current_delta_type": None,
"session_configuration_request": None,
}
)
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
# Synthetic session.created already sent by llm_http_handler.
streaming._session_created_sent_to_client = True
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
streaming._send_to_backend = AsyncMock() # type: ignore[method-assign]
await streaming.backend_to_client_send_messages()
# Duplicate session.created should still cause the one-time guardrail
# turn_detection update to be sent to backend.
assert streaming._send_to_backend.await_count == 1
sent_update = json.loads(streaming._send_to_backend.await_args_list[0].args[0])
assert sent_update["type"] == "session.update"
assert sent_update["session"]["turn_detection"]["create_response"] is False