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

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

View File

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

View File

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

View File

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

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