mirror of
https://github.com/dograh-hq/dograh.git
synced 2026-07-25 12:01:04 +02:00
fix: fix transition logic for realtime providers
This commit is contained in:
parent
348cd8427b
commit
960d183c82
7 changed files with 415 additions and 33 deletions
|
|
@ -1415,6 +1415,7 @@ class TestCustomToolManagerUnit:
|
|||
# Verify handler was registered
|
||||
assert "api_call" in registered_handlers
|
||||
assert registered_kwargs["api_call"]["timeout_secs"] == pytest.approx(5)
|
||||
assert registered_kwargs["api_call"]["is_node_transition"] is False
|
||||
|
||||
# Now test that the handler works
|
||||
handler = registered_handlers["api_call"]
|
||||
|
|
@ -1445,6 +1446,41 @@ class TestCustomToolManagerUnit:
|
|||
# Verify result was returned
|
||||
assert result_received["status"] == "success"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("category", ["end_call", "transfer_call"])
|
||||
async def test_register_handlers_marks_call_control_tools_as_node_transitions(
|
||||
self, category
|
||||
):
|
||||
from api.services.workflow.pipecat_engine_custom_tools import CustomToolManager
|
||||
|
||||
mock_engine = Mock()
|
||||
mock_engine._get_organization_id = AsyncMock(return_value=1)
|
||||
mock_engine.llm.register_function = Mock()
|
||||
manager = CustomToolManager(mock_engine)
|
||||
tool = MockToolModel(
|
||||
tool_uuid=f"{category}-uuid",
|
||||
name=category.replace("_", " "),
|
||||
description=f"Perform {category}",
|
||||
category=category,
|
||||
definition={
|
||||
"schema_version": 1,
|
||||
"type": category,
|
||||
"config": {},
|
||||
},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"api.services.workflow.pipecat_engine_custom_tools.db_client.get_tools_by_uuids",
|
||||
new=AsyncMock(return_value=[tool]),
|
||||
):
|
||||
await manager.register_handlers([tool.tool_uuid])
|
||||
|
||||
mock_engine.llm.register_function.assert_called_once()
|
||||
assert (
|
||||
mock_engine.llm.register_function.call_args.kwargs["is_node_transition"]
|
||||
is True
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transfer_call_renders_destination_from_initial_context(self):
|
||||
"""Transfer call tools resolve destination templates before provider calls."""
|
||||
|
|
|
|||
|
|
@ -1,11 +1,20 @@
|
|||
import asyncio
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from pipecat.frames.frames import TranscriptionFrame, TTSSpeakFrame
|
||||
from pipecat.frames.frames import (
|
||||
NodeTransitionStartedFrame,
|
||||
TranscriptionFrame,
|
||||
TTSSpeakFrame,
|
||||
)
|
||||
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||
from pipecat.processors.aggregators.llm_response_universal import (
|
||||
LLMContextAggregatorPair,
|
||||
)
|
||||
from pipecat.processors.frame_processor import FrameDirection
|
||||
from pipecat.services.llm_service import FunctionCallFromLLM
|
||||
|
||||
from api.services.pipecat.realtime.gemini_live import DograhGeminiLiveLLMService
|
||||
|
||||
|
|
@ -163,3 +172,194 @@ async def test_tts_greeting_waits_for_gemini_session_before_sending_prompt():
|
|||
assert prompt.endswith('"Hello from Dograh."')
|
||||
assert service._run_llm_when_session_ready is False
|
||||
assert service._pending_initial_greeting_text is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transition_call_flushes_pending_transcription_before_execution():
|
||||
service = _make_service()
|
||||
service._NODE_TRANSITION_TRANSCRIPTION_GRACE_SECONDS = 0
|
||||
service.register_function(
|
||||
"transition_to_next_node",
|
||||
AsyncMock(),
|
||||
is_node_transition=True,
|
||||
)
|
||||
service._user_transcription_buffer = "My last answer"
|
||||
|
||||
events = []
|
||||
|
||||
async def _push_transcription(text, result=None):
|
||||
events.append(("transcription", text))
|
||||
|
||||
async def _run_function_calls(function_calls):
|
||||
events.append(("function", function_calls[0].function_name))
|
||||
|
||||
service._push_user_transcription = AsyncMock(side_effect=_push_transcription)
|
||||
service.run_function_calls = AsyncMock(side_effect=_run_function_calls)
|
||||
service.create_task = lambda coro, name=None: asyncio.create_task(coro, name=name)
|
||||
|
||||
function_call = FunctionCallFromLLM(
|
||||
context=LLMContext(),
|
||||
tool_call_id="call-transition",
|
||||
function_name="transition_to_next_node",
|
||||
arguments={},
|
||||
)
|
||||
|
||||
await service._run_or_defer_function_calls([function_call])
|
||||
transition_task = service._transition_function_call_task
|
||||
assert transition_task is not None
|
||||
await transition_task
|
||||
|
||||
assert events == [
|
||||
("transcription", "My last answer"),
|
||||
("function", "transition_to_next_node"),
|
||||
]
|
||||
assert service._user_transcription_buffer == ""
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_transition_call_is_not_deferred_while_bot_is_responding():
|
||||
service = _make_service()
|
||||
service._bot_is_responding = True
|
||||
service.run_function_calls = AsyncMock()
|
||||
|
||||
function_call = FunctionCallFromLLM(
|
||||
context=LLMContext(),
|
||||
tool_call_id="call-ordinary",
|
||||
function_name="look_up_account",
|
||||
arguments={},
|
||||
)
|
||||
|
||||
await service._run_or_defer_function_calls([function_call])
|
||||
|
||||
service.run_function_calls.assert_awaited_once_with([function_call])
|
||||
assert service._pending_node_transition_function_calls == []
|
||||
assert service._transition_function_call_task is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_node_transition_call_is_deferred_while_bot_is_responding():
|
||||
service = _make_service()
|
||||
service._bot_is_responding = True
|
||||
service.register_function(
|
||||
"transition_to_next_node",
|
||||
AsyncMock(),
|
||||
is_node_transition=True,
|
||||
)
|
||||
service.run_function_calls = AsyncMock()
|
||||
|
||||
function_call = FunctionCallFromLLM(
|
||||
context=LLMContext(),
|
||||
tool_call_id="call-transition",
|
||||
function_name="transition_to_next_node",
|
||||
arguments={},
|
||||
)
|
||||
|
||||
await service._run_or_defer_function_calls([function_call])
|
||||
|
||||
service.run_function_calls.assert_not_awaited()
|
||||
assert service._pending_node_transition_function_calls == [function_call]
|
||||
assert service._transition_function_call_task is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initial_system_instruction_update_is_not_a_node_handoff():
|
||||
service = _make_service()
|
||||
service._connect = AsyncMock()
|
||||
service._disconnect = AsyncMock()
|
||||
|
||||
handled = await service._handle_changed_settings(
|
||||
{"system_instruction": "old prompt"}
|
||||
)
|
||||
|
||||
assert handled == {"system_instruction"}
|
||||
service._connect.assert_awaited_once_with()
|
||||
service._disconnect.assert_not_awaited()
|
||||
assert service._awaiting_node_transition_context is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_node_transition_uses_fresh_connection_instead_of_stale_handle():
|
||||
service = _make_service()
|
||||
service._session = _FakeSession()
|
||||
service._session_resumption_handle = "stale-handle"
|
||||
service._disconnect = AsyncMock()
|
||||
service._connect = AsyncMock()
|
||||
|
||||
handled = await service._handle_changed_settings(
|
||||
{"system_instruction": "old prompt"}
|
||||
)
|
||||
|
||||
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._connect.assert_awaited_once_with(session_resumption_handle=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fresh_transition_session_waits_for_updated_context_before_ready():
|
||||
service = _make_service()
|
||||
service._handled_initial_context = True
|
||||
service._awaiting_node_transition_context = True
|
||||
service._node_transition_context_received = False
|
||||
service._process_completed_function_calls = AsyncMock()
|
||||
service._drain_pending_tool_results = AsyncMock()
|
||||
|
||||
async def _seed_context():
|
||||
service._ready_for_realtime_input = True
|
||||
|
||||
service._create_initial_response = AsyncMock(side_effect=_seed_context)
|
||||
|
||||
session = _FakeSession()
|
||||
await service._handle_session_ready(session)
|
||||
|
||||
assert service._ready_for_realtime_input is False
|
||||
service._create_initial_response.assert_not_awaited()
|
||||
|
||||
context = _make_tool_result_context("call-transition")
|
||||
await service._handle_context(context)
|
||||
|
||||
service._process_completed_function_calls.assert_awaited_once_with(
|
||||
send_new_results=False
|
||||
)
|
||||
service._create_initial_response.assert_awaited_once()
|
||||
service._drain_pending_tool_results.assert_awaited_once()
|
||||
assert service._ready_for_realtime_input is True
|
||||
assert service._awaiting_node_transition_context is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_node_transition_frame_commits_user_transcript_to_context():
|
||||
context = LLMContext()
|
||||
context_aggregator = LLMContextAggregatorPair(context, realtime_service_mode=True)
|
||||
user_aggregator = context_aggregator.user()
|
||||
user_aggregator.push_context_frame = AsyncMock()
|
||||
|
||||
await user_aggregator._handle_transcription(
|
||||
TranscriptionFrame(
|
||||
text="The answer that selected this node",
|
||||
user_id="",
|
||||
timestamp="now",
|
||||
)
|
||||
)
|
||||
context_aggregation_event = asyncio.Event()
|
||||
await user_aggregator._handle_node_transition_started(
|
||||
NodeTransitionStartedFrame(
|
||||
function_calls=[
|
||||
FunctionCallFromLLM(
|
||||
context=context,
|
||||
tool_call_id="call-transition",
|
||||
function_name="transition_to_next_node",
|
||||
arguments={},
|
||||
)
|
||||
],
|
||||
context_aggregation_event=context_aggregation_event,
|
||||
)
|
||||
)
|
||||
|
||||
assert context.messages[-1] == {
|
||||
"role": "user",
|
||||
"content": "The answer that selected this node",
|
||||
}
|
||||
user_aggregator.push_context_frame.assert_awaited_once()
|
||||
assert context_aggregation_event.is_set()
|
||||
|
|
|
|||
|
|
@ -182,6 +182,7 @@ class TestPipecatEngineToolCalls:
|
|||
|
||||
# Assert that the context was updated with END_CALL_SYSTEM_PROMPT
|
||||
assert llm._settings.system_instruction == END_CALL_SYSTEM_PROMPT
|
||||
assert llm._functions["end_call"].is_node_transition is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_builtin_and_transition_calls_through_engine_1(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue