fix: record message confirmations in context
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user