mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
86f5c21025
commit
6274bc0d7b
2 changed files with 74 additions and 0 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue