fix(realtime): inject turn_detection into first session.update for deferred mode

- Instead of sending turn_detection as separate message (which gets dropped), inject it into the first client session.update
- This ensures guardrails work correctly in deferred mode
- Add test for turn_detection injection in deferred mode

Made-with: Cursor
This commit is contained in:
Sameer Kankute 2026-04-27 19:27:23 +05:30
parent 86f5c21025
commit 6274bc0d7b
No known key found for this signature in database
2 changed files with 74 additions and 0 deletions

View file

@ -636,6 +636,27 @@ class RealTimeStreaming:
## LOGGING
self.store_input(message=message)
## GUARDRAIL: Inject turn_detection into first session.update if needed
try:
msg_obj = json.loads(message)
if (
msg_obj.get("type") == "session.update"
and self.session_configuration_request is None
and not self._guardrail_turn_detection_update_sent
and self._has_audio_transcription_guardrails()
):
# Inject turn_detection into the first session.update
session = msg_obj.setdefault("session", {})
session.setdefault("turn_detection", {})["create_response"] = False
message = json.dumps(msg_obj)
self._guardrail_turn_detection_update_sent = True
verbose_logger.debug(
"Injected turn_detection into first session.update for audio transcription guardrails"
)
except (json.JSONDecodeError, AttributeError):
pass
## FORWARD TO BACKEND
if self.provider_config:
message = self.provider_config.transform_realtime_request(

View file

@ -1117,3 +1117,56 @@ async def test_guardrail_update_respects_idempotency_flag():
# Second call should be a no-op (idempotent)
await streaming._maybe_send_guardrail_turn_detection_update()
assert backend_ws.send.await_count == 1 # Still 1, not 2
@pytest.mark.asyncio
async def test_guardrail_turn_detection_injected_into_first_session_update_deferred_mode():
"""Verify turn_detection is injected into first session.update in deferred mode."""
client_ws = AsyncMock()
client_ws.receive_text = AsyncMock(
side_effect=[
json.dumps({
"type": "session.update",
"session": {
"modalities": ["text", "audio"],
"tools": [{"type": "function", "name": "get_weather"}],
}
}),
ConnectionClosed(None, None),
]
)
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
provider_config = MagicMock()
transformed_messages = []
def mock_transform(msg, model, session_config):
transformed_messages.append((msg, session_config))
return [msg] # Pass through for simplicity
provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform)
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
# Simulate first session.update in deferred mode
await streaming.client_ack_messages()
# Verify turn_detection was injected into the session.update
assert len(transformed_messages) == 1
transformed_msg, session_config = transformed_messages[0]
msg_obj = json.loads(transformed_msg)
assert msg_obj["type"] == "session.update"
assert "turn_detection" in msg_obj["session"]
assert msg_obj["session"]["turn_detection"]["create_response"] is False
assert streaming._guardrail_turn_detection_update_sent is True