fix: record message confirmations in context
This commit is contained in:
@@ -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