feat: add deterministic message interaction stages
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, patch
|
||||
@@ -29,7 +30,7 @@ from services.brains.dify_llm import (
|
||||
)
|
||||
from services.brains.workflow_brain import WorkflowBrain
|
||||
from services.runtime_variables import prepare_dynamic_config
|
||||
from services.action_runtime import ActionError, ActionOutcome, ActionStatus
|
||||
from services.action_runtime import ActionOutcome, ActionStatus
|
||||
from services.workflow.models import (
|
||||
LLMRouteResult,
|
||||
RouteStatus,
|
||||
@@ -350,30 +351,32 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
["preflight_1", "preflight_2"],
|
||||
)
|
||||
|
||||
async def test_opening_actions_wait_for_confirmation_and_greeting(self):
|
||||
tools = [
|
||||
RuntimeTool(
|
||||
id=f"opening_{index}",
|
||||
name=f"开场动作 {index}",
|
||||
function_name="show_message" if index == 1 else "load_opening_data",
|
||||
type="client" if index == 1 else "http",
|
||||
)
|
||||
for index in (1, 2)
|
||||
]
|
||||
async def test_opening_stage_starts_speech_and_releases_on_confirmation(self):
|
||||
tool = RuntimeTool(
|
||||
id="opening_data",
|
||||
name="加载开场数据",
|
||||
function_name="load_opening_data",
|
||||
type="http",
|
||||
)
|
||||
cfg = AssistantConfig(
|
||||
type="prompt",
|
||||
tools=tools,
|
||||
greeting="请阅读并确认重要信息",
|
||||
tools=[tool],
|
||||
startup={
|
||||
"execution_mode": "sequential",
|
||||
"opening_message": {
|
||||
"title": "重要提示",
|
||||
"message": "请确认已阅读。",
|
||||
"confirm_label": "确认",
|
||||
},
|
||||
"actions": [
|
||||
{
|
||||
"id": f"opening_{index}",
|
||||
"id": "opening_data",
|
||||
"phase": "opening",
|
||||
"tool_id": f"opening_{index}",
|
||||
"tool_id": "opening_data",
|
||||
"arguments": {},
|
||||
"required": True,
|
||||
}
|
||||
for index in (1, 2)
|
||||
],
|
||||
},
|
||||
)
|
||||
@@ -384,6 +387,17 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def queue_frame(frame):
|
||||
queued.append(frame)
|
||||
|
||||
confirmation_started = asyncio.Event()
|
||||
user_confirmed = asyncio.Event()
|
||||
client_calls = []
|
||||
|
||||
class FakeClientTools:
|
||||
async def call(self, function_name, arguments, **options):
|
||||
client_calls.append((function_name, arguments, options))
|
||||
confirmation_started.set()
|
||||
await user_confirmed.wait()
|
||||
return {"status": "ok", "data": {"action": "confirmed"}}
|
||||
|
||||
await brain.setup(
|
||||
cfg,
|
||||
BrainRuntime(
|
||||
@@ -393,28 +407,51 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
set_system_prompt=lambda _prompt: None,
|
||||
set_tools=lambda _tools: None,
|
||||
call_end=FakeCallEnd(),
|
||||
client_tools=FakeClientTools(),
|
||||
set_input_enabled=input_states.append,
|
||||
),
|
||||
)
|
||||
brain._actions.execute = AsyncMock(
|
||||
side_effect=[
|
||||
ActionOutcome(
|
||||
invocation_id=f"act_{index}",
|
||||
status=ActionStatus.SUCCESS,
|
||||
duration_ms=index,
|
||||
)
|
||||
for index in (1, 2)
|
||||
]
|
||||
)
|
||||
called_tool_ids = []
|
||||
|
||||
await brain.on_connected(greeting_pending=True)
|
||||
await brain.on_client_ready()
|
||||
async def execute(tool, *_args, **_kwargs):
|
||||
called_tool_ids.append(tool.id)
|
||||
return ActionOutcome(
|
||||
invocation_id=f"act_{len(called_tool_ids)}",
|
||||
status=ActionStatus.SUCCESS,
|
||||
duration_ms=len(called_tool_ids),
|
||||
)
|
||||
|
||||
brain._actions.execute = AsyncMock(side_effect=execute)
|
||||
|
||||
self.assertEqual(await brain.greeting(cfg), "")
|
||||
await brain.on_connected(greeting_pending=False)
|
||||
opening_task = asyncio.create_task(brain.on_client_ready())
|
||||
await confirmation_started.wait()
|
||||
|
||||
self.assertEqual(input_states, [False])
|
||||
called_tool_ids = [
|
||||
call.args[0].id for call in brain._actions.execute.await_args_list
|
||||
]
|
||||
self.assertEqual(called_tool_ids, ["opening_1", "opening_2"])
|
||||
self.assertEqual(called_tool_ids, [])
|
||||
self.assertEqual(client_calls[0][0], "show_message")
|
||||
self.assertFalse(client_calls[0][1]["dismissible"])
|
||||
self.assertTrue(
|
||||
any(
|
||||
isinstance(frame, TTSSpeakFrame)
|
||||
and frame.text == "请阅读并确认重要信息"
|
||||
for frame in queued
|
||||
)
|
||||
)
|
||||
self.assertTrue(
|
||||
any(
|
||||
isinstance(frame, OutputTransportMessageUrgentFrame)
|
||||
and frame.message.get("type") == "transcript"
|
||||
and frame.message.get("content") == "请阅读并确认重要信息"
|
||||
for frame in queued
|
||||
)
|
||||
)
|
||||
|
||||
user_confirmed.set()
|
||||
await opening_task
|
||||
self.assertEqual(input_states, [False, True])
|
||||
self.assertEqual(called_tool_ids, ["opening_data"])
|
||||
self.assertEqual(
|
||||
len(
|
||||
[
|
||||
@@ -424,36 +461,22 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
and frame.message.get("type") == "startup-action-result"
|
||||
]
|
||||
),
|
||||
2,
|
||||
1,
|
||||
)
|
||||
|
||||
await brain.on_greeting_finished()
|
||||
self.assertEqual(input_states, [False, True])
|
||||
|
||||
# Replayed client-ready must not execute startup actions twice.
|
||||
await brain.on_client_ready()
|
||||
self.assertEqual(brain._actions.execute.await_count, 2)
|
||||
self.assertEqual(brain._actions.execute.await_count, 1)
|
||||
|
||||
async def test_required_opening_failure_keeps_input_blocked_and_ends_call(self):
|
||||
tool = RuntimeTool(
|
||||
id="opening_message",
|
||||
name="开场确认",
|
||||
function_name="show_message",
|
||||
type="client",
|
||||
)
|
||||
cfg = AssistantConfig(
|
||||
type="prompt",
|
||||
tools=[tool],
|
||||
startup={
|
||||
"actions": [
|
||||
{
|
||||
"id": "opening_message",
|
||||
"phase": "opening",
|
||||
"tool_id": tool.id,
|
||||
"arguments": {},
|
||||
"required": True,
|
||||
}
|
||||
]
|
||||
"opening_message": {
|
||||
"title": "重要提示",
|
||||
"message": "请确认已阅读。",
|
||||
"confirm_label": "确认",
|
||||
}
|
||||
},
|
||||
)
|
||||
brain = build_brain(cfg)
|
||||
@@ -463,6 +486,10 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def queue_frame(_frame):
|
||||
pass
|
||||
|
||||
class FailingClientTools:
|
||||
async def call(self, *_args, **_kwargs):
|
||||
return {"status": "error", "message": "客户端未显示消息"}
|
||||
|
||||
await brain.setup(
|
||||
cfg,
|
||||
BrainRuntime(
|
||||
@@ -472,20 +499,10 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
set_system_prompt=lambda _prompt: None,
|
||||
set_tools=lambda _tools: None,
|
||||
call_end=call_end,
|
||||
client_tools=FailingClientTools(),
|
||||
set_input_enabled=input_states.append,
|
||||
),
|
||||
)
|
||||
brain._actions.execute = AsyncMock(
|
||||
return_value=ActionOutcome(
|
||||
invocation_id="act_failed",
|
||||
status=ActionStatus.FAILURE,
|
||||
duration_ms=1,
|
||||
error=ActionError(
|
||||
code="tool_error",
|
||||
message="用户未确认",
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
await brain.on_connected(greeting_pending=False)
|
||||
await brain.on_client_ready()
|
||||
@@ -948,7 +965,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertNotIn("fetch_user_image", custom_config["role_message"])
|
||||
self.assertFalse(scopes[-1]["enabled"])
|
||||
|
||||
async def test_initial_fixed_speech_waits_for_start_greeting_to_finish(self):
|
||||
async def test_initial_fixed_speech_starts_without_workflow_greeting(self):
|
||||
brain = WorkflowBrain(
|
||||
{
|
||||
"specVersion": 3,
|
||||
@@ -1006,12 +1023,14 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
)
|
||||
brain._manager = FakeManager()
|
||||
|
||||
await brain.on_connected(greeting_pending=True)
|
||||
|
||||
self.assertEqual(brain._manager.current_node, "start")
|
||||
self.assertFalse(any(isinstance(frame, TTSSpeakFrame) for frame in queued))
|
||||
|
||||
await brain.on_greeting_finished()
|
||||
self.assertNotIn("greeting", brain._engine.data("start"))
|
||||
self.assertEqual(
|
||||
await brain.greeting(
|
||||
AssistantConfig(type="workflow", greeting="旧助手级开场白")
|
||||
),
|
||||
"",
|
||||
)
|
||||
await brain.on_connected()
|
||||
|
||||
self.assertEqual(brain._manager.current_node, "agent")
|
||||
fixed_speech_frames = [
|
||||
@@ -1020,8 +1039,8 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertEqual(len(fixed_speech_frames), 1)
|
||||
self.assertEqual(fixed_speech_frames[0].text, "请问您怎么称呼?")
|
||||
|
||||
# Playback notifications may be duplicated by a transport reconnect;
|
||||
# the initial entry behavior must still run only once.
|
||||
# Workflow no longer owns a greeting playback lifecycle. Stray generic
|
||||
# transport notifications must not repeat Agent entry behavior.
|
||||
await brain.on_greeting_finished()
|
||||
self.assertEqual(
|
||||
len([frame for frame in queued if isinstance(frame, TTSSpeakFrame)]),
|
||||
@@ -1372,6 +1391,148 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
await brain._enter_action("queue_action")
|
||||
self.assertEqual(input_states, ["executing"])
|
||||
|
||||
async def test_message_starts_speech_and_releases_on_confirmation(self):
|
||||
brain = WorkflowBrain(
|
||||
AssistantConfig(
|
||||
type="workflow",
|
||||
graph={
|
||||
"specVersion": 3,
|
||||
"settings": {},
|
||||
"nodes": [
|
||||
{"id": "start", "type": "start", "data": {}},
|
||||
{
|
||||
"id": "message",
|
||||
"type": "message",
|
||||
"data": {
|
||||
"speech": "请先确认 {{customer}} 的重要信息。",
|
||||
"showMessage": True,
|
||||
"title": "重要提示",
|
||||
"message": "请核对客户信息。",
|
||||
"confirmLabel": "确认",
|
||||
"requireConfirmation": True,
|
||||
},
|
||||
},
|
||||
],
|
||||
"edges": [],
|
||||
},
|
||||
)
|
||||
)
|
||||
brain._store.values["customer"] = "王先生"
|
||||
events = []
|
||||
|
||||
async def queue_frame(frame):
|
||||
if isinstance(frame, TTSSpeakFrame):
|
||||
events.append(("speech", frame.text))
|
||||
|
||||
class OrderedCallEnd(FakeCallEnd):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.playback_completion = None
|
||||
|
||||
def track_speech(self):
|
||||
self.tracked_speeches += 1
|
||||
events.append("tracked")
|
||||
self.playback_completion = asyncio.get_running_loop().create_future()
|
||||
return self.playback_completion
|
||||
|
||||
message_started = asyncio.Event()
|
||||
user_confirmed = asyncio.Event()
|
||||
|
||||
class FakeClientTools:
|
||||
async def call(self, function_name, arguments, **options):
|
||||
self.function_name = function_name
|
||||
self.arguments = arguments
|
||||
self.options = options
|
||||
events.append("message_displayed")
|
||||
message_started.set()
|
||||
await user_confirmed.wait()
|
||||
return {"status": "ok", "data": {"action": "confirmed"}}
|
||||
|
||||
input_states = []
|
||||
call_end = OrderedCallEnd()
|
||||
client_tools = FakeClientTools()
|
||||
brain._runtime = BrainRuntime(
|
||||
context=LLMContext(messages=[]),
|
||||
llm=FakeLLM(),
|
||||
queue_frame=queue_frame,
|
||||
set_system_prompt=lambda _prompt: None,
|
||||
set_tools=lambda _tools: None,
|
||||
call_end=call_end,
|
||||
client_tools=client_tools,
|
||||
set_input_enabled=input_states.append,
|
||||
)
|
||||
brain._message_stages.set_client_tools(client_tools)
|
||||
|
||||
message_task = asyncio.create_task(brain._enter_message("message"))
|
||||
await message_started.wait()
|
||||
|
||||
self.assertEqual(
|
||||
events,
|
||||
[
|
||||
"tracked",
|
||||
("speech", "请先确认 王先生 的重要信息。"),
|
||||
"message_displayed",
|
||||
],
|
||||
)
|
||||
self.assertEqual(input_states, [False])
|
||||
self.assertEqual(client_tools.function_name, "show_message")
|
||||
self.assertEqual(client_tools.options["response_wait_mode"], "session")
|
||||
|
||||
user_confirmed.set()
|
||||
result = await message_task
|
||||
self.assertTrue(result.succeeded)
|
||||
self.assertEqual(result.action, "confirmed")
|
||||
self.assertFalse(call_end.playback_completion.done())
|
||||
self.assertEqual(input_states, [False, True])
|
||||
|
||||
async def test_speech_only_message_waits_for_transport_playback(self):
|
||||
brain = WorkflowBrain(
|
||||
{
|
||||
"specVersion": 3,
|
||||
"settings": {},
|
||||
"nodes": [
|
||||
{"id": "start", "type": "start", "data": {}},
|
||||
{
|
||||
"id": "message",
|
||||
"type": "message",
|
||||
"data": {"speech": "正在为您准备服务。"},
|
||||
},
|
||||
],
|
||||
"edges": [],
|
||||
}
|
||||
)
|
||||
|
||||
class PlaybackCallEnd(FakeCallEnd):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.completion = None
|
||||
|
||||
def track_speech(self):
|
||||
self.completion = asyncio.get_running_loop().create_future()
|
||||
return self.completion
|
||||
|
||||
call_end = PlaybackCallEnd()
|
||||
input_states = []
|
||||
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=call_end,
|
||||
set_input_enabled=input_states.append,
|
||||
)
|
||||
|
||||
message_task = asyncio.create_task(brain._enter_message("message"))
|
||||
await asyncio.sleep(0)
|
||||
self.assertFalse(message_task.done())
|
||||
self.assertEqual(input_states, [False])
|
||||
|
||||
call_end.completion.set_result(None)
|
||||
result = await message_task
|
||||
self.assertTrue(result.succeeded)
|
||||
self.assertEqual(input_states, [False, True])
|
||||
|
||||
async def test_nodes_without_outgoing_edges_remain_active(self):
|
||||
queued = []
|
||||
|
||||
@@ -1501,7 +1662,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
{
|
||||
"id": "start",
|
||||
"type": "start",
|
||||
"data": {"name": "Start", "greeting": "你想做什么?"},
|
||||
"data": {"name": "Start"},
|
||||
},
|
||||
{
|
||||
"id": "eat",
|
||||
@@ -1815,10 +1976,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
{
|
||||
"id": "start",
|
||||
"type": "start",
|
||||
"data": {
|
||||
"name": "Start",
|
||||
"greeting": "欢迎,{{user_name}}",
|
||||
},
|
||||
"data": {"name": "Start"},
|
||||
},
|
||||
{
|
||||
"id": "agent",
|
||||
@@ -1925,13 +2083,8 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
)
|
||||
await brain.setup(cfg, runtime)
|
||||
greeting = await brain.greeting(cfg)
|
||||
self.assertEqual(greeting, "欢迎,王先生")
|
||||
greeting_message = {
|
||||
"role": "system",
|
||||
"content": f"{GREETING_CONTEXT_MARKER}\n欢迎,王先生",
|
||||
}
|
||||
brain.prepare_greeting_context(greeting, context)
|
||||
self.assertEqual(context.get_messages(), [greeting_message])
|
||||
self.assertEqual(greeting, "")
|
||||
self.assertEqual(context.get_messages(), [])
|
||||
await brain.on_connected()
|
||||
self.assertEqual(brain._manager.current_node, "agent")
|
||||
variable_events = [
|
||||
@@ -1980,7 +2133,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
agent_config = brain._agent_config("agent")
|
||||
self.assertIn("王先生", agent_config["role_message"])
|
||||
self.assertIn("工作流路由已在用户一轮输入结束时完成", agent_config["role_message"])
|
||||
self.assertEqual(agent_config["task_messages"], [greeting_message])
|
||||
self.assertEqual(agent_config["task_messages"], [])
|
||||
self.assertFalse(agent_config["respond_immediately"])
|
||||
self.assertFalse(any(isinstance(frame, LLMRunFrame) for frame in worker.frames))
|
||||
self.assertEqual(
|
||||
@@ -2005,10 +2158,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertNotIn("pre_actions", fixed_config)
|
||||
self.assertEqual(
|
||||
fixed_config["task_messages"],
|
||||
[
|
||||
greeting_message,
|
||||
{"role": "assistant", "content": "您好,王先生"},
|
||||
],
|
||||
[{"role": "assistant", "content": "您好,王先生"}],
|
||||
)
|
||||
self.assertEqual(
|
||||
brain._agent_config(
|
||||
@@ -2016,7 +2166,6 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
[{"role": "assistant", "content": "正在进入下一阶段"}],
|
||||
)["task_messages"],
|
||||
[
|
||||
greeting_message,
|
||||
{"role": "assistant", "content": "正在进入下一阶段"},
|
||||
{"role": "assistant", "content": "您好,王先生"},
|
||||
],
|
||||
@@ -2034,10 +2183,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
]
|
||||
self.assertEqual(
|
||||
context_updates[-1].messages,
|
||||
[
|
||||
greeting_message,
|
||||
{"role": "assistant", "content": "您好,王先生"},
|
||||
],
|
||||
[{"role": "assistant", "content": "您好,王先生"}],
|
||||
)
|
||||
self.assertFalse(
|
||||
any(
|
||||
|
||||
Reference in New Issue
Block a user