fix: record message confirmations in context

This commit is contained in:
Xin Wang
2026-08-11 14:15:10 +08:00
parent 86639692ba
commit fc026e765e
6 changed files with 273 additions and 8 deletions

View File

@@ -31,6 +31,7 @@ from services.brains.dify_llm import (
)
from services.brains.workflow_brain import ConfiguredFlowManager, WorkflowBrain
from services.fixed_speech import FIXED_SPEECH_CONTEXT_MARKER
from services.message_stage import MESSAGE_CONFIRMATION_CONTEXT_MARKER
from services.runtime_variables import prepare_dynamic_config
from services.action_runtime import ActionOutcome, ActionStatus
from services.workflow.models import (
@@ -418,6 +419,7 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
confirmation_started = asyncio.Event()
user_confirmed = asyncio.Event()
client_calls = []
context = LLMContext(messages=[])
class FakeClientTools:
async def call(self, function_name, arguments, **options):
@@ -429,7 +431,7 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
context=context,
llm=FakeLLM(),
queue_frame=queue_frame,
set_system_prompt=lambda _prompt: None,
@@ -458,6 +460,13 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(input_states, [False])
self.assertEqual(called_tool_ids, [])
self.assertFalse(
any(
MESSAGE_CONFIRMATION_CONTEXT_MARKER
in str(message.get("content") or "")
for message in context.get_messages()
)
)
self.assertEqual(client_calls[0][0], "show_message")
self.assertFalse(client_calls[0][1]["dismissible"])
self.assertTrue(client_calls[0][2]["interrupt_on_result"])
@@ -481,6 +490,17 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
await opening_task
self.assertEqual(input_states, [False, True])
self.assertEqual(called_tool_ids, ["opening_data"])
confirmation_messages = [
message
for message in context.get_messages()
if MESSAGE_CONFIRMATION_CONTEXT_MARKER
in str(message.get("content") or "")
]
self.assertEqual(len(confirmation_messages), 1)
self.assertEqual(confirmation_messages[0]["role"], "user")
self.assertIn("点击“确认”", confirmation_messages[0]["content"])
self.assertIn("弹窗标题:重要提示", confirmation_messages[0]["content"])
self.assertIn("弹窗内容:请确认已阅读。", confirmation_messages[0]["content"])
self.assertEqual(
sum(isinstance(frame, LLMRunFrame) for frame in queued),
1,
@@ -500,6 +520,14 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
# Replayed client-ready must not execute startup actions twice.
await brain.on_client_ready()
self.assertEqual(brain._actions.execute.await_count, 1)
self.assertEqual(
sum(
MESSAGE_CONFIRMATION_CONTEXT_MARKER
in str(message.get("content") or "")
for message in context.get_messages()
),
1,
)
self.assertEqual(
sum(isinstance(frame, LLMRunFrame) for frame in queued),
1,
@@ -2277,10 +2305,106 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
result = await message_task
self.assertTrue(result.succeeded)
self.assertEqual(result.action, "confirmed")
self.assertIsNotNone(result.confirmation)
self.assertEqual(result.confirmation.title, "重要提示")
self.assertEqual(result.confirmation.message, "请核对客户信息。")
self.assertEqual(result.confirmation.confirm_label, "确认")
self.assertFalse(call_end.playback_completion.done())
self.assertEqual(input_states, [False, True])
self.assertEqual(events[-1], "message_completed")
async def test_confirmation_message_is_forwarded_to_next_agent_context(self):
graph = {
"specVersion": 3,
"settings": {},
"nodes": [
{"id": "start", "type": "start", "data": {}},
{
"id": "message",
"type": "message",
"data": {
"title": "办理提示",
"message": "确认后开始办理。",
"confirmLabel": "继续办理",
"completionPolicy": "confirmation",
},
},
{
"id": "agent",
"type": "agent",
"data": {"prompt": "收集事故信息"},
},
],
"edges": [
{
"id": "start-message",
"source": "start",
"target": "message",
"data": {"mode": "always"},
},
{
"id": "message-agent",
"source": "message",
"target": "agent",
"data": {"mode": "always"},
},
],
}
brain = WorkflowBrain(graph)
class FakeManager:
def __init__(self):
self.current_node = None
self.configs = []
async def initialize(self, config):
self.current_node = config["name"]
self.configs.append(config)
async def set_node_from_config(self, config):
self.current_node = config["name"]
self.configs.append(config)
class FakeClientTools:
async def call(self, *_args, **_kwargs):
return {"status": "ok", "data": {"action": "confirmed"}}
manager = FakeManager()
client_tools = FakeClientTools()
brain._runtime = BrainRuntime(
context=LLMContext(messages=[]),
llm=FakeLLM(),
queue_frame=noop_queue_frame,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=FakeCallEnd(),
client_tools=client_tools,
set_input_enabled=lambda _enabled: None,
)
brain._manager = manager
brain._message_stages.set_client_tools(client_tools)
await brain.on_connected()
for _ in range(10):
await asyncio.sleep(0)
if manager.current_node == "agent":
break
self.assertEqual(manager.current_node, "agent")
confirmation_messages = [
message
for message in manager.configs[-1]["task_messages"]
if MESSAGE_CONFIRMATION_CONTEXT_MARKER
in str(message.get("content") or "")
]
self.assertEqual(len(confirmation_messages), 1)
self.assertEqual(confirmation_messages[0]["role"], "user")
self.assertIn("点击“继续办理”", confirmation_messages[0]["content"])
self.assertIn(
"事件来源workflow-message:message",
confirmation_messages[0]["content"],
)
async def test_speech_only_message_waits_for_transport_playback(self):
brain = WorkflowBrain(
{

View File

@@ -12,6 +12,7 @@ from pipecat.frames.frames import (
)
import services.brains # Initialize the brain registry before realtime imports.
from services.fixed_speech import FixedSpeechOutput
from services.message_stage import MESSAGE_CONFIRMATION_CONTEXT_MARKER
from services.pipecat.call_lifecycle import (
CallEndCoordinator,
FixedSpeechPlaybackMarkerFrame,
@@ -141,6 +142,62 @@ class FixedSpeechPlaybackTest(unittest.IsolatedAsyncioTestCase):
self.assertEqual(reasons, ["workflow_completed"])
async def test_realtime_confirmation_is_appended_without_running_model(self):
graph = {
"specVersion": 3,
"settings": {},
"nodes": [
{
"id": "message",
"type": "message",
"data": {
"title": "办理提示",
"message": "确认后开始办理。",
"confirmLabel": "继续办理",
"completionPolicy": "confirmation",
},
}
],
"edges": [],
}
appended = []
input_states = []
class FakeRealtime:
async def send_text(self, text, *, run_immediately=True):
appended.append((text, run_immediately))
class FakeClientTools:
async def call(self, *_args, **_kwargs):
return {"status": "ok", "data": {"action": "confirmed"}}
async def queue_frame(_frame):
pass
controller = WorkflowRealtimeController(
cfg=AssistantConfig(type="workflow", graph=graph),
engine=WorkflowEngine(graph),
store=DynamicVariableStore({}),
runtime=SimpleNamespace(
realtime=FakeRealtime(),
queue_frame=queue_frame,
call_end=SimpleNamespace(ending=False),
session_id="test-session",
client_tools=FakeClientTools(),
set_input_enabled=input_states.append,
capture_image=None,
),
)
succeeded = await controller._enter_message("message")
self.assertTrue(succeeded)
self.assertEqual(input_states, [False, True])
self.assertEqual(len(appended), 1)
self.assertFalse(appended[0][1])
self.assertIn(MESSAGE_CONFIRMATION_CONTEXT_MARKER, appended[0][0])
self.assertIn("点击“继续办理”", appended[0][0])
if __name__ == "__main__":
unittest.main()