feat: add prompt opening behavior modes

This commit is contained in:
Xin Wang
2026-08-04 13:29:21 +08:00
parent 78bba090a2
commit 5bf5987fe4
11 changed files with 235 additions and 97 deletions

View File

@@ -366,6 +366,7 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
tools=[tool],
startup={
"execution_mode": "sequential",
"opening_mode": "confirmation",
"entry_mode": "generate",
"opening_message": {
"title": "重要提示",
@@ -485,6 +486,7 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
type="prompt",
greeting="请先确认",
startup={
"opening_mode": "confirmation",
"entry_mode": "wait_user",
"opening_message": {
"title": "重要提示",
@@ -527,7 +529,10 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
cfg = AssistantConfig(
type="prompt",
greeting="固定开场白",
startup={"entry_mode": "generate"},
startup={
"opening_mode": "playback",
"entry_mode": "generate",
},
)
brain = build_brain(cfg)
queued = []
@@ -561,45 +566,22 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
1,
)
async def test_opening_confirmation_can_keep_speech_playing_before_entry(self):
async def test_interruptible_opening_user_input_suppresses_entry_reply(self):
cfg = AssistantConfig(
type="prompt",
greeting="必须完整播放的开场白",
greeting="可以打断的开场白",
startup={
"opening_mode": "interruptible",
"entry_mode": "generate",
"opening_message": {
"title": "重要提示",
"message": "请确认已阅读。",
"skip_speech_on_confirm": False,
},
},
)
brain = build_brain(cfg)
queued = []
input_states = []
message_displayed = asyncio.Event()
class TrackedCallEnd(FakeCallEnd):
def __init__(self):
super().__init__()
self.playback_completion = None
def track_speech(self):
self.tracked_speeches += 1
self.playback_completion = asyncio.get_running_loop().create_future()
return self.playback_completion
class FakeClientTools:
async def call(self, *_args, **options):
self.options = options
message_displayed.set()
return {"status": "ok", "data": {"action": "confirmed"}}
async def queue_frame(frame):
queued.append(frame)
call_end = TrackedCallEnd()
client_tools = FakeClientTools()
await brain.setup(
cfg,
BrainRuntime(
@@ -608,24 +590,52 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
queue_frame=queue_frame,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=call_end,
client_tools=client_tools,
call_end=FakeCallEnd(),
set_input_enabled=input_states.append,
),
)
await brain.on_connected(greeting_pending=False)
opening_task = asyncio.create_task(brain.on_client_ready())
await message_displayed.wait()
await asyncio.sleep(0)
await brain.on_connected(greeting_pending=True)
self.assertEqual(input_states, [])
await brain.on_interruption_processed()
await brain.on_greeting_finished()
self.assertFalse(client_tools.options["interrupt_on_result"])
self.assertFalse(opening_task.done())
self.assertFalse(any(isinstance(frame, LLMRunFrame) for frame in queued))
call_end.playback_completion.set_result(None)
await opening_task
self.assertEqual(input_states, [False, True])
async def test_interruptible_opening_natural_finish_runs_entry_behavior(self):
cfg = AssistantConfig(
type="prompt",
greeting="可以打断的开场白",
startup={
"opening_mode": "interruptible",
"entry_mode": "generate",
},
)
brain = build_brain(cfg)
queued = []
input_states = []
async def queue_frame(frame):
queued.append(frame)
await brain.setup(
cfg,
BrainRuntime(
context=LLMContext(messages=[]),
llm=FakeLLM(),
queue_frame=queue_frame,
set_system_prompt=lambda _prompt: None,
set_tools=lambda _tools: None,
call_end=FakeCallEnd(),
set_input_enabled=input_states.append,
),
)
await brain.on_connected(greeting_pending=True)
await brain.on_greeting_finished()
self.assertEqual(input_states, [])
self.assertEqual(
sum(isinstance(frame, LLMRunFrame) for frame in queued),
1,

View File

@@ -52,6 +52,7 @@ class _Brain:
self.prepared_greeting = ""
self.greeting_pending = False
self.greeting_finished = 0
self.interruptions_processed = 0
def prepare_greeting_context(self, greeting, _context):
self.prepared_greeting = greeting
@@ -62,6 +63,9 @@ class _Brain:
async def on_greeting_finished(self):
self.greeting_finished += 1
async def on_interruption_processed(self):
self.interruptions_processed += 1
async def on_client_ready(self):
for content, timestamp in (
("Message 节点播报", "2026-07-14T10:00:00.200+00:00"),
@@ -111,6 +115,7 @@ class PipelineEventTest(unittest.IsolatedAsyncioTestCase):
)
self.assertEqual(acknowledgements, [True])
self.assertEqual(brain.interruptions_processed, 1)
async def test_greeting_keeps_playback_timestamp_until_client_ready(self):
transport = _EventSource()

View File

@@ -109,7 +109,7 @@ class StartupActionValidationTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(body.startup.actions, [])
self.assertEqual(body.startup.opening_message.confirm_label, "确认")
self.assertTrue(body.startup.opening_message.skip_speech_on_confirm)
self.assertEqual(body.startup.opening_mode, "confirmation")
self.assertEqual(body.startup.entry_mode, "generate")
async def test_startup_tool_does_not_need_conversation_binding(self):