fix: fix transition logic for realtime providers

This commit is contained in:
Abhishek Kumar 2026-07-15 16:45:16 +05:30
parent 348cd8427b
commit 960d183c82
7 changed files with 415 additions and 33 deletions

View file

@ -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."""

View file

@ -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()

View file

@ -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(