feat: add deterministic message interaction stages
This commit is contained in:
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
@@ -25,10 +26,18 @@ from services.brains.base import (
|
||||
SessionVariableUpdate,
|
||||
)
|
||||
from services.action_runtime import (
|
||||
ActionOutcome,
|
||||
ActionInvocationCancelled,
|
||||
ActionRunner,
|
||||
ActionStatus,
|
||||
)
|
||||
from services.action_stage import ActionStageRunner, ActionStageSpec, StageAction
|
||||
from services.fixed_speech import FixedSpeechOutput
|
||||
from services.message_stage import (
|
||||
MessageDisplaySpec,
|
||||
MessageStageRunner,
|
||||
MessageStageSpec,
|
||||
)
|
||||
from services.runtime_variables import DynamicVariableStore
|
||||
from services.tool_executor import ToolExecutionError, ToolExecutor
|
||||
from services.tool_policy import policy_for_tool
|
||||
@@ -50,17 +59,31 @@ class PromptBrain(BaseBrain):
|
||||
self._store = DynamicVariableStore.from_config(cfg)
|
||||
self._tools = ToolExecutor(self._store)
|
||||
self._actions = ActionRunner(self._tools)
|
||||
self._action_stages = ActionStageRunner(self._actions)
|
||||
self._message_stages = MessageStageRunner()
|
||||
self._tool_by_id = {tool.id: tool for tool in cfg.tools}
|
||||
self._runtime: BrainRuntime | None = None
|
||||
self._output: FixedSpeechOutput | None = None
|
||||
self._waiting_for_generated_end_speech = False
|
||||
self._greeting_finished = True
|
||||
self._preflight_finished = False
|
||||
self._opening_started = False
|
||||
self._opening_finished = False
|
||||
self._opening_input_blocked = False
|
||||
self._startup_failed = False
|
||||
|
||||
async def greeting(self, cfg: AssistantConfig) -> str:
|
||||
return self._store.render(cfg.greeting) if self._dynamic_enabled else cfg.greeting
|
||||
# The built-in opening Message owns the greeting so speech and the
|
||||
# client dialog can start as one atomic stage.
|
||||
if self._opening_message() is not None:
|
||||
return ""
|
||||
return self._render_greeting(cfg)
|
||||
|
||||
def _render_greeting(self, cfg: AssistantConfig) -> str:
|
||||
return (
|
||||
self._store.render(cfg.greeting)
|
||||
if self._dynamic_enabled
|
||||
else cfg.greeting
|
||||
)
|
||||
|
||||
def system_prompt(self, cfg: AssistantConfig) -> str:
|
||||
return self._store.render(cfg.prompt) if self._dynamic_enabled else cfg.prompt
|
||||
@@ -77,12 +100,15 @@ class PromptBrain(BaseBrain):
|
||||
self._tools,
|
||||
is_session_ending=lambda: runtime.call_end.ending,
|
||||
)
|
||||
self._action_stages = ActionStageRunner(self._actions)
|
||||
self._message_stages = MessageStageRunner(runtime.client_tools)
|
||||
self._output = FixedSpeechOutput(self._store, runtime)
|
||||
self._tool_by_id = {tool.id: tool for tool in cfg.tools}
|
||||
self._waiting_for_generated_end_speech = False
|
||||
self._greeting_finished = True
|
||||
self._preflight_finished = False
|
||||
self._opening_started = False
|
||||
self._opening_finished = not bool(self._startup_actions("opening"))
|
||||
self._opening_finished = not self._has_opening_stage()
|
||||
self._opening_input_blocked = False
|
||||
self._startup_failed = False
|
||||
llm_tool_ids = (
|
||||
set(cfg.llm_tool_ids) if cfg.llm_tool_ids is not None else None
|
||||
@@ -119,32 +145,83 @@ class PromptBrain(BaseBrain):
|
||||
self._preflight_finished = True
|
||||
|
||||
async def on_connected(self, *, greeting_pending: bool = False) -> None:
|
||||
self._greeting_finished = not greeting_pending
|
||||
if (
|
||||
self._startup_actions("opening")
|
||||
self._has_opening_stage()
|
||||
and self._runtime is not None
|
||||
and self._runtime.set_input_enabled is not None
|
||||
):
|
||||
self._runtime.set_input_enabled(False)
|
||||
self._opening_input_blocked = True
|
||||
|
||||
async def on_client_ready(self) -> None:
|
||||
if self._output is not None:
|
||||
await self._output.mark_client_ready()
|
||||
if self._opening_started or self._opening_finished or self._startup_failed:
|
||||
return
|
||||
self._opening_started = True
|
||||
runtime = self._runtime
|
||||
if runtime is None:
|
||||
raise RuntimeError("PromptBrain 尚未初始化")
|
||||
opening_message = self._opening_message()
|
||||
opening_actions = self._startup_actions("opening")
|
||||
speech = (
|
||||
self._render_greeting(self._cfg).strip()
|
||||
if opening_message is not None
|
||||
else ""
|
||||
)
|
||||
if speech:
|
||||
self.prepare_greeting_context(speech, runtime.context)
|
||||
try:
|
||||
succeeded = await self._run_startup_actions("opening")
|
||||
if opening_message is not None:
|
||||
message_result = await self._message_stages.run(
|
||||
self._opening_message_stage_spec(speech, opening_message),
|
||||
speak=self._speak_opening,
|
||||
set_input_enabled=runtime.set_input_enabled,
|
||||
input_already_blocked=self._opening_input_blocked,
|
||||
release_input_on_success=not bool(opening_actions),
|
||||
release_input_on_failure=False,
|
||||
)
|
||||
if not message_result.succeeded:
|
||||
await self._fail_opening(
|
||||
message_result.error or "开场消息显示失败"
|
||||
)
|
||||
return
|
||||
|
||||
if opening_actions:
|
||||
result = await self._action_stages.run(
|
||||
self._opening_actions_stage_spec(),
|
||||
set_input_enabled=runtime.set_input_enabled,
|
||||
input_already_blocked=self._opening_input_blocked,
|
||||
release_input_on_failure=False,
|
||||
on_outcome=self._publish_opening_outcome,
|
||||
)
|
||||
if not result.succeeded:
|
||||
await self._fail_opening("必需的开场 Action 执行失败")
|
||||
return
|
||||
except ActionInvocationCancelled:
|
||||
self._startup_failed = True
|
||||
raise
|
||||
if not succeeded:
|
||||
await self._fail_opening("必需的开场 Action 执行失败")
|
||||
return
|
||||
self._opening_finished = True
|
||||
self._release_startup_gate_if_ready()
|
||||
self._opening_input_blocked = False
|
||||
|
||||
async def on_greeting_finished(self) -> None:
|
||||
self._greeting_finished = True
|
||||
self._release_startup_gate_if_ready()
|
||||
async def _speak_opening(self, content: str) -> Awaitable[None] | None:
|
||||
if self._output is None:
|
||||
raise RuntimeError("Prompt 固定播报输出尚未初始化")
|
||||
return await self._output.speak(
|
||||
content,
|
||||
source="prompt-opening-speech",
|
||||
record_history=False,
|
||||
)
|
||||
|
||||
def _opening_message(self) -> dict[str, Any] | None:
|
||||
startup = self._cfg.startup if isinstance(self._cfg.startup, dict) else {}
|
||||
value = startup.get("opening_message", startup.get("openingMessage"))
|
||||
return value if isinstance(value, dict) else None
|
||||
|
||||
def _has_opening_stage(self) -> bool:
|
||||
return self._opening_message() is not None or bool(
|
||||
self._startup_actions("opening")
|
||||
)
|
||||
|
||||
def _startup_actions(self, phase: str) -> list[dict[str, Any]]:
|
||||
startup = self._cfg.startup if isinstance(self._cfg.startup, dict) else {}
|
||||
@@ -154,6 +231,78 @@ class PromptBrain(BaseBrain):
|
||||
if isinstance(action, dict) and action.get("phase", "opening") == phase
|
||||
]
|
||||
|
||||
def _opening_actions_stage_spec(self) -> ActionStageSpec:
|
||||
actions = tuple(
|
||||
StageAction(
|
||||
id=str(action.get("id") or "startup_action"),
|
||||
tool=self._tool_by_id.get(
|
||||
str(action.get("tool_id") or action.get("toolId") or "")
|
||||
),
|
||||
arguments=action.get("arguments") or {},
|
||||
required=bool(action.get("required", True)),
|
||||
invocation_id=self._actions.new_invocation_id(),
|
||||
)
|
||||
for action in self._startup_actions("opening")
|
||||
)
|
||||
return ActionStageSpec(
|
||||
actions=actions,
|
||||
input_policy="block",
|
||||
)
|
||||
|
||||
def _opening_message_stage_spec(
|
||||
self,
|
||||
speech: str,
|
||||
config: dict[str, Any],
|
||||
) -> MessageStageSpec:
|
||||
return MessageStageSpec(
|
||||
speech=speech,
|
||||
display=MessageDisplaySpec(
|
||||
title=self._store.render(
|
||||
str(config.get("title") or "重要提示")
|
||||
).strip(),
|
||||
message=self._store.render(
|
||||
str(config.get("message") or "")
|
||||
).strip(),
|
||||
confirm_label=self._store.render(
|
||||
str(
|
||||
config.get("confirm_label")
|
||||
or config.get("confirmLabel")
|
||||
or "确认"
|
||||
)
|
||||
).strip(),
|
||||
),
|
||||
require_confirmation=True,
|
||||
)
|
||||
|
||||
async def _publish_opening_outcome(
|
||||
self,
|
||||
action: StageAction,
|
||||
outcome: ActionOutcome,
|
||||
) -> None:
|
||||
if outcome.updated_variables:
|
||||
self._refresh_prompt()
|
||||
if self._runtime is not None:
|
||||
await self._runtime.queue_frame(
|
||||
OutputTransportMessageUrgentFrame(
|
||||
message={
|
||||
"type": "startup-action-result",
|
||||
"actionId": action.id,
|
||||
"phase": "opening",
|
||||
"outcome": outcome.trace_payload(),
|
||||
}
|
||||
)
|
||||
)
|
||||
if outcome.status == ActionStatus.FAILURE and action.required:
|
||||
logger.warning(
|
||||
f"必需的 Prompt opening Action 失败: "
|
||||
f"action={action.id} error={outcome.error}"
|
||||
)
|
||||
elif outcome.status == ActionStatus.FAILURE:
|
||||
logger.warning(
|
||||
f"忽略可选 Prompt opening Action 失败: "
|
||||
f"action={action.id} error={outcome.error}"
|
||||
)
|
||||
|
||||
async def _run_startup_actions(self, phase: str) -> bool:
|
||||
for action in self._startup_actions(phase):
|
||||
action_id = str(action.get("id") or "startup_action")
|
||||
@@ -170,17 +319,6 @@ class PromptBrain(BaseBrain):
|
||||
)
|
||||
if outcome.updated_variables:
|
||||
self._refresh_prompt()
|
||||
if phase == "opening" and self._runtime is not None:
|
||||
await self._runtime.queue_frame(
|
||||
OutputTransportMessageUrgentFrame(
|
||||
message={
|
||||
"type": "startup-action-result",
|
||||
"actionId": action_id,
|
||||
"phase": phase,
|
||||
"outcome": outcome.trace_payload(),
|
||||
}
|
||||
)
|
||||
)
|
||||
if outcome.status == ActionStatus.SUCCESS:
|
||||
continue
|
||||
if outcome.status == ActionStatus.CANCELLED:
|
||||
@@ -197,18 +335,6 @@ class PromptBrain(BaseBrain):
|
||||
)
|
||||
return True
|
||||
|
||||
def _release_startup_gate_if_ready(self) -> None:
|
||||
runtime = self._runtime
|
||||
if (
|
||||
runtime is not None
|
||||
and runtime.set_input_enabled is not None
|
||||
and self._greeting_finished
|
||||
and self._opening_finished
|
||||
and not self._startup_failed
|
||||
and not runtime.call_end.ending
|
||||
):
|
||||
runtime.set_input_enabled(True)
|
||||
|
||||
async def _fail_opening(self, message: str) -> None:
|
||||
self._startup_failed = True
|
||||
runtime = self._runtime
|
||||
|
||||
Reference in New Issue
Block a user