497 lines
19 KiB
Python
497 lines
19 KiB
Python
"""Local prompt assistant, including prompt-only reusable tools."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from collections.abc import Awaitable
|
|
from typing import Any
|
|
from uuid import uuid4
|
|
|
|
from loguru import logger
|
|
from models import AssistantConfig
|
|
from pipecat.adapters.schemas.function_schema import FunctionSchema
|
|
from pipecat.frames.frames import OutputTransportMessageUrgentFrame, TTSSpeakFrame
|
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
|
from pipecat.processors.frame_processor import FrameProcessor
|
|
from pipecat.services.llm_service import (
|
|
FunctionCallParams,
|
|
FunctionCallResultProperties,
|
|
)
|
|
from pipecat.utils.time import time_now_iso8601
|
|
|
|
from services.brains.base import (
|
|
BaseBrain,
|
|
BrainRuntime,
|
|
BrainSpec,
|
|
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
|
|
|
|
|
|
PREFLIGHT_TIMEOUT_SECONDS = 30
|
|
|
|
|
|
class PromptBrain(BaseBrain):
|
|
spec = BrainSpec(
|
|
type="prompt",
|
|
supported_runtime_modes=frozenset({"pipeline", "realtime"}),
|
|
owns_context=True,
|
|
)
|
|
|
|
def __init__(self, cfg: AssistantConfig):
|
|
self._cfg = cfg
|
|
self._dynamic_enabled = True
|
|
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._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:
|
|
# 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
|
|
|
|
def build_llm(self, cfg: AssistantConfig, context: LLMContext) -> FrameProcessor:
|
|
from services.pipecat.service_factory import create_llm
|
|
|
|
return create_llm(cfg)
|
|
|
|
async def setup(self, cfg: AssistantConfig, runtime: BrainRuntime) -> None:
|
|
self._runtime = runtime
|
|
self._tools.set_client_tools(runtime.client_tools)
|
|
self._actions = ActionRunner(
|
|
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._preflight_finished = False
|
|
self._opening_started = False
|
|
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
|
|
)
|
|
schemas: list[FunctionSchema] = []
|
|
for tool in cfg.tools:
|
|
if llm_tool_ids is not None and tool.id not in llm_tool_ids:
|
|
continue
|
|
if tool.type == "end_call":
|
|
schema, handler = self._make_end_call_tool(tool, runtime)
|
|
elif tool.type in {"http", "mcp", "client"}:
|
|
schema, handler = self._make_remote_tool(tool, runtime)
|
|
else:
|
|
continue
|
|
schemas.append(schema)
|
|
policy = policy_for_tool(tool)
|
|
runtime.llm.register_function(
|
|
tool.function_name,
|
|
handler,
|
|
cancel_on_interruption=policy.cancel_on_interruption,
|
|
)
|
|
runtime.set_tools(schemas)
|
|
|
|
async def run_preflight(self) -> None:
|
|
if self._preflight_finished:
|
|
return
|
|
try:
|
|
async with asyncio.timeout(PREFLIGHT_TIMEOUT_SECONDS):
|
|
succeeded = await self._run_startup_actions("preflight")
|
|
except TimeoutError as exc:
|
|
raise RuntimeError("Prompt preflight 超过 30 秒安全上限") from exc
|
|
if not succeeded:
|
|
raise RuntimeError("必需的 Prompt preflight Action 执行失败")
|
|
self._preflight_finished = True
|
|
|
|
async def on_connected(self, *, greeting_pending: bool = False) -> None:
|
|
if (
|
|
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:
|
|
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
|
|
self._opening_finished = True
|
|
self._opening_input_blocked = False
|
|
|
|
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 {}
|
|
return [
|
|
action
|
|
for action in startup.get("actions") or []
|
|
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")
|
|
tool_id = str(action.get("tool_id") or action.get("toolId") or "")
|
|
tool = self._tool_by_id.get(tool_id)
|
|
invocation_id = self._actions.new_invocation_id()
|
|
logger.info(
|
|
f"执行 Prompt {phase} Action: action={action_id} tool={tool_id}"
|
|
)
|
|
outcome = await self._actions.execute(
|
|
tool,
|
|
action.get("arguments") or {},
|
|
invocation_id=invocation_id,
|
|
)
|
|
if outcome.updated_variables:
|
|
self._refresh_prompt()
|
|
if outcome.status == ActionStatus.SUCCESS:
|
|
continue
|
|
if outcome.status == ActionStatus.CANCELLED:
|
|
return False
|
|
if bool(action.get("required", True)):
|
|
logger.warning(
|
|
f"必需的 Prompt {phase} Action 失败: "
|
|
f"action={action_id} error={outcome.error}"
|
|
)
|
|
return False
|
|
logger.warning(
|
|
f"忽略可选 Prompt {phase} Action 失败: "
|
|
f"action={action_id} error={outcome.error}"
|
|
)
|
|
return True
|
|
|
|
async def _fail_opening(self, message: str) -> None:
|
|
self._startup_failed = True
|
|
runtime = self._runtime
|
|
if runtime is None or runtime.call_end.ending:
|
|
return
|
|
await runtime.queue_frame(
|
|
OutputTransportMessageUrgentFrame(
|
|
message={"type": "startup-action-error", "message": message}
|
|
)
|
|
)
|
|
runtime.call_end.begin("startup_action_failed")
|
|
await runtime.call_end.finish()
|
|
|
|
def record_user_message(self, content: str) -> None:
|
|
if not self._dynamic_enabled:
|
|
return
|
|
self._store.record("user", content)
|
|
self._refresh_prompt()
|
|
|
|
async def on_session_update(
|
|
self,
|
|
dynamic_variables: dict[str, Any],
|
|
) -> SessionVariableUpdate:
|
|
changed = self._store.assign_declared_many(dynamic_variables)
|
|
if changed:
|
|
self._refresh_prompt()
|
|
return SessionVariableUpdate(
|
|
changed=changed,
|
|
dynamic_variables=self._store.public_values(),
|
|
)
|
|
|
|
async def on_assistant_text_start(self, _turn_id: str) -> None:
|
|
if self._runtime is not None:
|
|
self._runtime.call_end.begin_response()
|
|
|
|
async def on_assistant_text_end(
|
|
self,
|
|
_turn_id: str,
|
|
content: str,
|
|
interrupted: bool,
|
|
) -> None:
|
|
if content and not interrupted:
|
|
self._store.record("agent", content, completed_agent_turn=True)
|
|
self._refresh_prompt()
|
|
if (
|
|
self._waiting_for_generated_end_speech
|
|
and self._runtime is not None
|
|
and self._runtime.call_end.ending
|
|
):
|
|
self._waiting_for_generated_end_speech = False
|
|
await self._runtime.call_end.finish_after_current_speech(
|
|
has_text=bool(content.strip()) and not interrupted
|
|
)
|
|
|
|
def _refresh_prompt(self) -> None:
|
|
if self._dynamic_enabled and self._runtime is not None:
|
|
self._runtime.set_system_prompt(self._store.render(self._cfg.prompt))
|
|
|
|
def _make_remote_tool(self, tool, runtime: BrainRuntime):
|
|
properties, required = self._tools.schema_parts(tool)
|
|
self._tools.register_secrets(tool)
|
|
policy = policy_for_tool(tool)
|
|
|
|
async def return_result(params: FunctionCallParams, result: dict) -> None:
|
|
if not policy.runs_llm_after_result:
|
|
await params.result_callback(
|
|
result,
|
|
properties=FunctionCallResultProperties(run_llm=False),
|
|
)
|
|
else:
|
|
await params.result_callback(result)
|
|
|
|
async def call_tool(params: FunctionCallParams) -> None:
|
|
try:
|
|
result = await self._tools.execute(tool, dict(params.arguments or {}))
|
|
if result["updated_variables"]:
|
|
self._refresh_prompt()
|
|
await return_result(params, result)
|
|
except (ToolExecutionError, ValueError) as exc:
|
|
await return_result(
|
|
params,
|
|
{"status": "error", "message": f"工具调用失败: {exc}"},
|
|
)
|
|
|
|
schema = FunctionSchema(
|
|
name=tool.function_name,
|
|
description=tool.description or f"调用 {tool.name}",
|
|
properties=properties,
|
|
required=required,
|
|
)
|
|
return schema, call_tool
|
|
|
|
def _make_end_call_tool(self, tool, runtime: BrainRuntime):
|
|
config = (tool.definition or {}).get("config") or {}
|
|
message_type = str(config.get("message_type") or "none")
|
|
custom_message = str(config.get("custom_message") or "").strip()
|
|
capture_reason = bool(config.get("capture_reason", True))
|
|
|
|
async def end_call(params: FunctionCallParams) -> None:
|
|
reason = str(params.arguments.get("reason") or "end_call_tool").strip()
|
|
uses_custom_message = message_type == "custom" and bool(custom_message)
|
|
self._waiting_for_generated_end_speech = not uses_custom_message
|
|
runtime.call_end.begin(reason)
|
|
await params.result_callback(
|
|
{"status": "success", "action": "ending_call"},
|
|
properties=FunctionCallResultProperties(run_llm=False),
|
|
)
|
|
|
|
if not uses_custom_message:
|
|
# The model may have already streamed a spoken goodbye before
|
|
# invoking this tool. Decide at assistant-text-end whether to
|
|
# wait for that TTS audio or finish immediately for tool-only calls.
|
|
return
|
|
|
|
turn_id = uuid4().hex
|
|
timestamp = time_now_iso8601()
|
|
for message in (
|
|
{
|
|
"type": "assistant-text-start",
|
|
"turn_id": turn_id,
|
|
"timestamp": timestamp,
|
|
},
|
|
{
|
|
"type": "assistant-text-delta",
|
|
"turn_id": turn_id,
|
|
"delta": custom_message,
|
|
},
|
|
{
|
|
"type": "assistant-text-end",
|
|
"turn_id": turn_id,
|
|
"content": custom_message,
|
|
"interrupted": False,
|
|
},
|
|
):
|
|
await runtime.queue_frame(
|
|
OutputTransportMessageUrgentFrame(message=message)
|
|
)
|
|
runtime.call_end.arm_after_speech()
|
|
await runtime.queue_frame(
|
|
TTSSpeakFrame(custom_message, append_to_context=False)
|
|
)
|
|
|
|
properties = (
|
|
{
|
|
"reason": {
|
|
"type": "string",
|
|
"description": "结束本次通话的简短原因。",
|
|
}
|
|
}
|
|
if capture_reason
|
|
else {}
|
|
)
|
|
schema = FunctionSchema(
|
|
name=tool.function_name,
|
|
description=tool.description or "结束当前通话。",
|
|
properties=properties,
|
|
required=["reason"] if capture_reason else [],
|
|
)
|
|
return schema, end_call
|