Files
ai-video-fullstack/backend/services/brains/prompt_brain.py
2026-08-02 00:07:35 +08:00

371 lines
14 KiB
Python

"""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