"""Local prompt assistant, including prompt-only reusable tools.""" from __future__ import annotations import asyncio 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 ( ActionInvocationCancelled, ActionRunner, ActionStatus, ) 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._tool_by_id = {tool.id: tool for tool in cfg.tools} self._runtime: BrainRuntime | 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._startup_failed = False async def 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._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._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: self._greeting_finished = not greeting_pending if ( self._startup_actions("opening") and self._runtime is not None and self._runtime.set_input_enabled is not None ): self._runtime.set_input_enabled(False) async def on_client_ready(self) -> None: if self._opening_started or self._opening_finished or self._startup_failed: return self._opening_started = True try: succeeded = await self._run_startup_actions("opening") 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() async def on_greeting_finished(self) -> None: self._greeting_finished = True self._release_startup_gate_if_ready() 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 ] 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 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: 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 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 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