diff --git a/backend/services/brains/prompt_brain.py b/backend/services/brains/prompt_brain.py index 01724ca..e1d9a1a 100644 --- a/backend/services/brains/prompt_brain.py +++ b/backend/services/brains/prompt_brain.py @@ -43,6 +43,7 @@ from services.message_stage import ( MessageDisplaySpec, MessageStageRunner, MessageStageSpec, + confirmation_context_message, ) from services.runtime_variables import DynamicVariableError, DynamicVariableStore from services.system_tools import state_update_properties, system_tool_kind @@ -226,8 +227,12 @@ class PromptBrain(BaseBrain): self.prepare_greeting_context(speech, runtime.context) try: if opening_message is not None: + message_spec = self._opening_message_stage_spec( + speech, + opening_message, + ) message_result = await self._message_stages.run( - self._opening_message_stage_spec(speech, opening_message), + message_spec, speak=self._speak_opening, set_input_enabled=runtime.set_input_enabled, input_already_blocked=self._opening_input_blocked, @@ -241,6 +246,12 @@ class PromptBrain(BaseBrain): message_result.error or "开场消息显示失败" ) return + context_message = confirmation_context_message( + message_result, + source="prompt-opening", + ) + if context_message is not None: + runtime.context.add_message(context_message) if opening_actions: result = await self._action_stages.run( diff --git a/backend/services/brains/workflow_brain.py b/backend/services/brains/workflow_brain.py index d4c7ee7..07b4eb9 100644 --- a/backend/services/brains/workflow_brain.py +++ b/backend/services/brains/workflow_brain.py @@ -58,6 +58,7 @@ from services.message_stage import ( MessageStageResult, MessageStageRunner, MessageStageSpec, + confirmation_context_message, ) from services.runtime_variables import DynamicVariableError, DynamicVariableStore from services.system_tools import state_update_properties, system_tool_kind @@ -1421,6 +1422,12 @@ class WorkflowBrain(BaseBrain): context_message = fixed_speech_context_message(result.speech) if context_message is not None: context_messages.append(context_message) + confirmation_message = confirmation_context_message( + result, + source=f"workflow-message:{continuation.node_id}", + ) + if confirmation_message is not None: + context_messages.append(confirmation_message) if not self._engine.has_outgoing(continuation.node_id): self._state.enter( diff --git a/backend/services/message_stage.py b/backend/services/message_stage.py index 923c302..320193b 100644 --- a/backend/services/message_stage.py +++ b/backend/services/message_stage.py @@ -17,6 +17,7 @@ from services.message_policy import ( BUILTIN_SHOW_MESSAGE = "show_message" +MESSAGE_CONFIRMATION_CONTEXT_MARKER = "[客户端交互事件,不是语音转写]" SpeechCompletion = Awaitable[None] | None Speak = Callable[[str], Awaitable[SpeechCompletion]] StartedHook = Callable[[], Awaitable[None]] @@ -40,15 +41,54 @@ class MessageStageSpec: completion_policy: MessageCompletionPolicy = MESSAGE_PLAYBACK +@dataclass(frozen=True) +class MessageConfirmation: + """The exact dialog and user action that completed a confirmation stage.""" + + title: str + message: str + confirm_label: str + action: str + + @dataclass(frozen=True) class MessageStageResult: """Result used by Workflow routing and Prompt opening failure handling.""" succeeded: bool speech: str = "" - action: str | None = None + confirmation: MessageConfirmation | None = None error: str | None = None + @property + def action(self) -> str | None: + """Keep existing trace and routing callers independent of result shape.""" + return self.confirmation.action if self.confirmation is not None else None + + +def confirmation_context_message( + result: MessageStageResult, + *, + source: str, +) -> dict[str, str] | None: + """Represent a confirmed client dialog as one provider-neutral user event.""" + confirmation = result.confirmation + if not result.succeeded or confirmation is None: + return None + return { + "role": "user", + "content": "\n".join( + [ + MESSAGE_CONFIRMATION_CONTEXT_MARKER, + f"用户已阅读消息弹窗并点击“{confirmation.confirm_label}”。", + f"弹窗标题:{confirmation.title}", + f"弹窗内容:{confirmation.message}", + f"操作结果:{confirmation.action}", + f"事件来源:{source}", + ] + ), + } + class MessageStageRunner: """Run one atomic user-visible message stage. @@ -110,12 +150,12 @@ class MessageStageRunner: if speech and speak is not None: playback_completion = await speak(speech) - action: str | None = None + confirmation: MessageConfirmation | None = None if spec.display is not None: result = await self._show_message(spec, speech=speech) if not result.succeeded: return result - action = result.action + confirmation = result.confirmation # Confirmation is the gate. It deliberately does not wait for the # audio completion future, so the user can continue immediately. @@ -128,7 +168,7 @@ class MessageStageRunner: result = MessageStageResult( succeeded=True, speech=speech, - action=action, + confirmation=confirmation, ) return result except asyncio.CancelledError: @@ -207,8 +247,23 @@ class MessageStageRunner: if isinstance(data, dict) else None ) + if require_confirmation and action is None: + return MessageStageResult( + succeeded=False, + speech=speech, + error="客户端未返回确认操作", + ) return MessageStageResult( succeeded=True, speech=speech, - action=action, + confirmation=( + MessageConfirmation( + title=display.title, + message=display.message, + confirm_label=display.confirm_label, + action=action, + ) + if require_confirmation and action is not None + else None + ), ) diff --git a/backend/services/workflow/realtime.py b/backend/services/workflow/realtime.py index 0f90e02..dccc39c 100644 --- a/backend/services/workflow/realtime.py +++ b/backend/services/workflow/realtime.py @@ -22,6 +22,7 @@ from services.message_stage import ( MessageDisplaySpec, MessageStageRunner, MessageStageSpec, + confirmation_context_message, ) from services.pipecat.call_lifecycle import playback_marker_for from services.pipecat.realtime_tools import RealtimeTool, RealtimeToolResult @@ -893,7 +894,17 @@ class WorkflowRealtimeController: node_id=node_id, code="workflow_message_error", ) - return result.succeeded + return False + context_message = confirmation_context_message( + result, + source=f"workflow-message:{node_id}", + ) + if context_message is not None: + await self._runtime.realtime.send_text( + context_message["content"], + run_immediately=False, + ) + return True async def _enter_action(self, node_id: str) -> bool: self._state.enter(node_id, WorkflowStatus.RUNNING_ACTION) diff --git a/backend/tests/test_brains.py b/backend/tests/test_brains.py index 7618531..5dd970a 100644 --- a/backend/tests/test_brains.py +++ b/backend/tests/test_brains.py @@ -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( { diff --git a/backend/tests/test_fixed_speech_playback.py b/backend/tests/test_fixed_speech_playback.py index f3e741e..0a7d411 100644 --- a/backend/tests/test_fixed_speech_playback.py +++ b/backend/tests/test_fixed_speech_playback.py @@ -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()