mirror of
https://github.com/dograh-hq/dograh.git
synced 2026-07-25 12:01:04 +02:00
fix: allow transient disconnect for reconnect for gemini
This commit is contained in:
parent
96cde1b767
commit
eae30b3b21
3 changed files with 55 additions and 4 deletions
|
|
@ -214,6 +214,23 @@ class DograhGeminiLiveLLMService(GeminiLiveLLMService):
|
|||
)
|
||||
self._schedule_node_transition_function_calls(fcs)
|
||||
|
||||
async def _disconnect_for_reconnect(self) -> bool:
|
||||
"""Disconnect without discarding a pending graceful shutdown.
|
||||
|
||||
Returns:
|
||||
``True`` when the caller should open a new Gemini session.
|
||||
``False`` when a deferred :class:`EndFrame` was released instead.
|
||||
"""
|
||||
await self._disconnect(preserve_pending_end_frame=True)
|
||||
if not self._end_frame_pending_bot_turn_finished:
|
||||
return True
|
||||
|
||||
logger.info(
|
||||
"Releasing deferred EndFrame instead of reconnecting Gemini service"
|
||||
)
|
||||
await self._release_deferred_end_frame()
|
||||
return False
|
||||
|
||||
async def _reconnect_for_node_transition(self) -> None:
|
||||
"""Start a fresh connection and wait to seed the completed context.
|
||||
|
||||
|
|
@ -227,7 +244,14 @@ class DograhGeminiLiveLLMService(GeminiLiveLLMService):
|
|||
self._node_transition_context_received = False
|
||||
self._node_transition_context_seed_started = False
|
||||
self._session_resumption_handle = None
|
||||
await self._disconnect()
|
||||
should_open_new_session = await self._disconnect_for_reconnect()
|
||||
if not should_open_new_session:
|
||||
# The helper released a deferred EndFrame, so graceful shutdown now
|
||||
# owns the lifecycle and this node transition must not reconnect.
|
||||
self._awaiting_node_transition_context = False
|
||||
self._node_transition_context_received = False
|
||||
self._node_transition_context_seed_started = False
|
||||
return
|
||||
await self._connect(session_resumption_handle=None)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import asyncio
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from pipecat.frames.frames import (
|
||||
EndFrame,
|
||||
NodeTransitionStartedFrame,
|
||||
TranscriptionFrame,
|
||||
TTSSpeakFrame,
|
||||
|
|
@ -292,10 +293,36 @@ async def test_node_transition_uses_fresh_connection_instead_of_stale_handle():
|
|||
assert handled == {"system_instruction"}
|
||||
assert service._session_resumption_handle is None
|
||||
assert service._awaiting_node_transition_context is True
|
||||
service._disconnect.assert_awaited_once()
|
||||
service._disconnect.assert_awaited_once_with(preserve_pending_end_frame=True)
|
||||
service._connect.assert_awaited_once_with(session_resumption_handle=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_node_transition_releases_deferred_end_frame_instead_of_reconnecting():
|
||||
service = _make_service()
|
||||
service._session = _FakeSession()
|
||||
service._connect = AsyncMock()
|
||||
service.queue_frame = AsyncMock()
|
||||
service._bot_is_responding = True
|
||||
|
||||
end_frame = EndFrame(reason="user_idle_max_duration_exceeded")
|
||||
timeout_task = MagicMock()
|
||||
timeout_task.done.return_value = False
|
||||
service._end_frame_pending_bot_turn_finished = end_frame
|
||||
service._end_frame_deferral_timeout_task = timeout_task
|
||||
|
||||
await service._reconnect_for_node_transition()
|
||||
|
||||
service.queue_frame.assert_awaited_once_with(end_frame)
|
||||
service._connect.assert_not_awaited()
|
||||
timeout_task.cancel.assert_called_once_with()
|
||||
assert service._end_frame_pending_bot_turn_finished is None
|
||||
assert service._end_frame_deferral_timeout_task is None
|
||||
assert service._awaiting_node_transition_context is False
|
||||
assert service._node_transition_context_received is False
|
||||
assert service._node_transition_context_seed_started is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fresh_transition_session_waits_for_updated_context_before_ready():
|
||||
service = _make_service()
|
||||
|
|
|
|||
2
pipecat
2
pipecat
|
|
@ -1 +1 @@
|
|||
Subproject commit aadd1d5dd606d2871b082e6f2ca1ad1eee53785b
|
||||
Subproject commit 0de21c0bc9fa4b42f07df5a57b44a759dadc1560
|
||||
Loading…
Add table
Add a link
Reference in a new issue