feat(workflow): enhance message stages and image routing
This commit is contained in:
@@ -1021,7 +1021,7 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
"type": "message",
|
||||
"data": {
|
||||
"speech": "请问您怎么称呼?",
|
||||
"showMessage": False,
|
||||
"completionPolicy": "playback",
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -1457,11 +1457,10 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
"type": "message",
|
||||
"data": {
|
||||
"speech": "请先确认 {{customer}} 的重要信息。",
|
||||
"showMessage": True,
|
||||
"title": "重要提示",
|
||||
"message": "请核对客户信息。",
|
||||
"confirmLabel": "确认",
|
||||
"requireConfirmation": True,
|
||||
"completionPolicy": "confirmation",
|
||||
},
|
||||
},
|
||||
],
|
||||
@@ -1585,6 +1584,131 @@ class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertTrue(result.succeeded)
|
||||
self.assertEqual(input_states, [False, True])
|
||||
|
||||
async def test_interruptible_message_forwards_multimodal_input_once(self):
|
||||
graph = {
|
||||
"specVersion": 3,
|
||||
"settings": {},
|
||||
"nodes": [
|
||||
{"id": "start", "type": "start", "data": {}},
|
||||
{
|
||||
"id": "message",
|
||||
"type": "message",
|
||||
"data": {
|
||||
"speech": "请按提示操作,也可以直接告诉我需求。",
|
||||
"completionPolicy": "interruptible",
|
||||
},
|
||||
},
|
||||
{
|
||||
"id": "agent",
|
||||
"type": "agent",
|
||||
"data": {
|
||||
"prompt": "处理用户输入",
|
||||
"contextPolicy": "fresh",
|
||||
},
|
||||
},
|
||||
],
|
||||
"edges": [
|
||||
{
|
||||
"id": "start-message",
|
||||
"source": "start",
|
||||
"target": "message",
|
||||
"data": {"mode": "always"},
|
||||
},
|
||||
{
|
||||
"id": "message-agent",
|
||||
"source": "message",
|
||||
"target": "agent",
|
||||
"data": {"mode": "always"},
|
||||
},
|
||||
],
|
||||
}
|
||||
brain = WorkflowBrain(graph)
|
||||
queued = []
|
||||
input_states = []
|
||||
|
||||
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
|
||||
|
||||
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)
|
||||
|
||||
async def queue_frame(frame):
|
||||
queued.append(frame)
|
||||
|
||||
call_end = PlaybackCallEnd()
|
||||
manager = FakeManager()
|
||||
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,
|
||||
set_input_enabled=input_states.append,
|
||||
)
|
||||
brain._manager = manager
|
||||
|
||||
await brain.on_connected()
|
||||
await asyncio.sleep(0)
|
||||
self.assertEqual(manager.current_node, "message")
|
||||
self.assertIsNotNone(call_end.completion)
|
||||
self.assertTrue(all(input_states))
|
||||
|
||||
image_message = {
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "已发送一张图片"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,AA=="},
|
||||
},
|
||||
],
|
||||
}
|
||||
await brain.on_user_turn_end(
|
||||
"已发送一张图片",
|
||||
user_message=image_message,
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
self.assertEqual(manager.current_node, "agent")
|
||||
self.assertEqual(
|
||||
[config["name"] for config in manager.configs].count("agent"),
|
||||
1,
|
||||
)
|
||||
self.assertTrue(call_end.completion.cancelled())
|
||||
self.assertTrue(all(input_states))
|
||||
self.assertEqual(
|
||||
manager.configs[-1]["task_messages"],
|
||||
[image_message],
|
||||
)
|
||||
self.assertEqual(
|
||||
sum(isinstance(frame, LLMRunFrame) for frame in queued),
|
||||
1,
|
||||
)
|
||||
self.assertTrue(
|
||||
any(
|
||||
isinstance(frame, OutputTransportMessageUrgentFrame)
|
||||
and frame.message.get("event") == "message_interrupted"
|
||||
for frame in queued
|
||||
)
|
||||
)
|
||||
|
||||
async def test_message_between_agents_resumes_after_playback(self):
|
||||
graph = {
|
||||
"specVersion": 3,
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import unittest
|
||||
|
||||
from pipecat.frames.frames import LLMMessagesAppendFrame, UserImageRawFrame
|
||||
from services.pipecat.pipeline import _multimodal_user_input_frame
|
||||
from services.pipecat.processors import UserInputError, parse_user_input
|
||||
|
||||
|
||||
@@ -52,6 +54,33 @@ class UserInputParserTests(unittest.TestCase):
|
||||
}
|
||||
)
|
||||
|
||||
def test_native_image_uses_the_standard_multimodal_user_turn_path(self):
|
||||
image = UserImageRawFrame(
|
||||
image=bytes([220, 40, 40] * 16 * 16),
|
||||
size=(16, 16),
|
||||
format="RGB",
|
||||
)
|
||||
|
||||
frame = _multimodal_user_input_frame(
|
||||
image,
|
||||
"请根据用户刚提交的图片进行回复。",
|
||||
)
|
||||
|
||||
self.assertIsInstance(frame, LLMMessagesAppendFrame)
|
||||
self.assertTrue(frame.run_llm)
|
||||
self.assertEqual(frame.messages[0]["role"], "user")
|
||||
content = frame.messages[0]["content"]
|
||||
self.assertEqual(
|
||||
content[0],
|
||||
{"type": "text", "text": "请根据用户刚提交的图片进行回复。"},
|
||||
)
|
||||
self.assertEqual(content[1]["type"], "image_url")
|
||||
self.assertTrue(
|
||||
content[1]["image_url"]["url"].startswith(
|
||||
"data:image/jpeg;base64,"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -200,32 +200,53 @@ class WorkflowGraphTests(unittest.TestCase):
|
||||
message = next(
|
||||
node for node in normalized["nodes"] if node["type"] == "message"
|
||||
)
|
||||
self.assertFalse(message["data"]["showMessage"])
|
||||
self.assertFalse(message["data"]["requireConfirmation"])
|
||||
self.assertEqual(message["data"]["completionPolicy"], "playback")
|
||||
self.assertNotIn("requireConfirmation", message["data"])
|
||||
self.assertNotIn("showMessage", message["data"])
|
||||
self.assertEqual(message["data"]["confirmLabel"], "确认")
|
||||
|
||||
message["data"].update(
|
||||
{
|
||||
"speech": "",
|
||||
"showMessage": True,
|
||||
"title": "重要提示",
|
||||
"message": "",
|
||||
"confirmLabel": "确认",
|
||||
"requireConfirmation": True,
|
||||
"completionPolicy": "confirmation",
|
||||
}
|
||||
)
|
||||
errors = validate_graph(normalized)
|
||||
self.assertTrue(any("弹窗消息必须为" in error for error in errors))
|
||||
self.assertTrue(any("必须配置播报内容" in error for error in errors))
|
||||
|
||||
message["data"].update(
|
||||
{
|
||||
"speech": "请确认",
|
||||
"showMessage": False,
|
||||
"requireConfirmation": True,
|
||||
}
|
||||
{"speech": "", "completionPolicy": "playback"}
|
||||
)
|
||||
errors = validate_graph(normalized)
|
||||
self.assertTrue(any("等待确认时必须显示弹窗" in error for error in errors))
|
||||
self.assertTrue(any("必须配置播报内容" in error for error in errors))
|
||||
|
||||
message["data"]["completionPolicy"] = "unknown"
|
||||
errors = validate_graph(normalized)
|
||||
self.assertTrue(any("完成策略无效" in error for error in errors))
|
||||
|
||||
legacy = valid_graph()
|
||||
legacy["nodes"].append(
|
||||
{
|
||||
"id": "legacy-message",
|
||||
"type": "message",
|
||||
"data": {
|
||||
"speech": "请确认",
|
||||
"showMessage": True,
|
||||
"title": "提示",
|
||||
"message": "请确认",
|
||||
"confirmLabel": "确认",
|
||||
"requireConfirmation": True,
|
||||
},
|
||||
}
|
||||
)
|
||||
legacy_message = normalize_graph(legacy)["nodes"][-1]["data"]
|
||||
self.assertEqual(legacy_message["completionPolicy"], "confirmation")
|
||||
self.assertNotIn("requireConfirmation", legacy_message)
|
||||
self.assertNotIn("showMessage", legacy_message)
|
||||
|
||||
def test_voice_resource_creates_isolated_runtime_config(self):
|
||||
base = AssistantConfig(type="workflow", asr="default", voice="default")
|
||||
|
||||
Reference in New Issue
Block a user