Files
ai-video-fullstack/backend/services/brains/workflow_brain.py

778 lines
31 KiB
Python

"""Pipecat Flows-backed Workflow v3 brain."""
from __future__ import annotations
import asyncio
from copy import deepcopy
from typing import Any
from loguru import logger
from models import AssistantConfig, RuntimeTool
from db.session import SessionLocal
from pipecat.flows import (
ContextStrategy,
ContextStrategyConfig,
FlowManager,
FlowsFunctionSchema,
NodeConfig,
)
from pipecat.frames.frames import (
LLMRunFrame,
LLMUpdateSettingsFrame,
OutputTransportMessageUrgentFrame,
)
from pipecat.processors.aggregators.llm_context import LLMContext
from pipecat.processors.frame_processor import FrameProcessor
from services.brains.base import BaseBrain, BrainRuntime, BrainSpec
from services.knowledge import search as search_knowledge
from services.runtime_variables import DynamicVariableStore
from services.tool_executor import ToolExecutionError, ToolExecutor
from services.workflow.agent import WorkflowAgentStage
from services.workflow.models import (
RouteStatus,
WorkflowRuntimeState,
WorkflowStatus,
)
from services.workflow.output import WorkflowOutput
from services.workflow.routing import WorkflowEdgeEvaluator
from services.workflow_engine import WorkflowEngine
from services.workflow_router import WorkflowLLMRouter
MAX_AUTOMATIC_HOPS = 50
class WorkflowBrain(BaseBrain):
spec = BrainSpec(
type="workflow",
supported_runtime_modes=frozenset({"pipeline"}),
owns_context=True,
)
def __init__(self, cfg_or_graph: AssistantConfig | dict[str, Any]):
cfg = cfg_or_graph if isinstance(cfg_or_graph, AssistantConfig) else None
graph = deepcopy(cfg.graph if cfg is not None else cfg_or_graph)
if cfg is not None:
# Graph v3 owns Workflow defaults. Keep older saved graphs compatible
# by filling the new interaction settings from the assistant row.
settings = graph.setdefault("settings", {})
settings.setdefault("enableInterrupt", cfg.enableInterrupt)
settings.setdefault("turnConfig", deepcopy(cfg.turnConfig))
self._engine = WorkflowEngine(graph or {})
if not self._engine.has_graph() or not self._engine.start_id:
raise ValueError("WorkflowBrain 缺少有效的 Start 节点")
self._cfg = cfg
self._store = DynamicVariableStore.from_config(cfg or AssistantConfig(type="workflow"))
self._tools = ToolExecutor(self._store)
self._tool_by_id: dict[str, RuntimeTool] = {
tool.id: tool for tool in (cfg.tools if cfg else [])
}
self._runtime: BrainRuntime | None = None
self._manager: FlowManager | None = None
self._router = WorkflowLLMRouter(cfg or AssistantConfig(type="workflow"))
self._edge_evaluator = WorkflowEdgeEvaluator(
self._engine,
self._store,
self._router_for_node,
)
self._state = WorkflowRuntimeState(current_node_id=self._engine.start_id)
self._turn_lock = asyncio.Lock()
self._output: WorkflowOutput | None = None
self._agent_stage: WorkflowAgentStage | None = None
self._ended = False
self._greeting_context_message: dict[str, str] | None = None
self._startup_waiting_for_greeting = False
async def greeting(self, cfg: AssistantConfig) -> str:
return self._engine.greeting(self._store) or cfg.greeting
def system_prompt(self, cfg: AssistantConfig) -> str:
return self._store.render(self._engine.global_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:
if runtime.worker is None or runtime.context_aggregator is None:
raise RuntimeError("WorkflowBrain 需要 PipelineWorker 和 context aggregator pair")
self._cfg = cfg
self._runtime = runtime
self._store = DynamicVariableStore.from_config(cfg)
self._tools = ToolExecutor(self._store)
self._tool_by_id = {tool.id: tool for tool in cfg.tools}
self._router = WorkflowLLMRouter(cfg)
self._edge_evaluator = WorkflowEdgeEvaluator(
self._engine,
self._store,
self._router_for_node,
)
self._state = WorkflowRuntimeState(current_node_id=self._engine.start_id)
self._turn_lock = asyncio.Lock()
self._output = WorkflowOutput(self._store, runtime)
self._agent_stage = WorkflowAgentStage(
cfg=cfg,
engine=self._engine,
store=self._store,
runtime=runtime,
)
self._ended = False
self._greeting_context_message = None
self._startup_waiting_for_greeting = False
self._manager = FlowManager(
worker=runtime.worker,
llm=runtime.llm,
context_aggregator=runtime.context_aggregator,
transport=runtime.transport,
global_functions=runtime.flow_global_functions,
)
self._manager.state["variables"] = self._store.values
def prepare_greeting_context(
self,
greeting: str,
context: LLMContext,
) -> dict[str, str] | None:
message = super().prepare_greeting_context(greeting, context)
self._greeting_context_message = deepcopy(message) if message else None
return message
async def on_connected(self, *, greeting_pending: bool = False) -> None:
self._state.enter(self._engine.start_id, WorkflowStatus.STARTING)
await self._emit_node_active(self._engine.start_id)
await self._emit_variables(
reason="initialized",
node_id=self._engine.start_id,
)
if self._manager is None:
raise RuntimeError("Workflow FlowManager 尚未初始化")
self._startup_waiting_for_greeting = greeting_pending
if greeting_pending:
# Keep the Workflow on Start until the transport confirms that the
# shared greeting has finished. This prevents an initial Agent's
# fixed speech (or generated reply) from racing the greeting.
await self._manager.initialize(
self._passive_node_config(self._engine.start_id)
)
logger.info("工作流等待 Start 开场白播放完毕")
return
node_config = await self._initial_node_config()
await self._manager.initialize(node_config)
await self._after_node_activated(node_config)
logger.info(f"工作流模式启用: 当前节点={self._manager.current_node}")
async def on_greeting_finished(self) -> None:
"""Enter the first node only after Start's greeting reaches playback end."""
if not self._startup_waiting_for_greeting or self._ended:
return
self._startup_waiting_for_greeting = False
manager = self._require_manager()
if manager.current_node != self._engine.start_id:
return
node_config = await self._initial_node_config()
if node_config.get("name") == self._engine.start_id:
self._state.enter(self._engine.start_id, WorkflowStatus.WAITING_USER)
return
await manager.set_node_from_config(node_config)
await self._after_node_activated(node_config)
logger.info(f"Start 开场白结束,进入节点: {manager.current_node}")
async def _initial_node_config(self) -> NodeConfig:
"""Only a default-only Start advances before the first user turn."""
outgoing = self._engine.outgoing(self._engine.start_id)
has_condition = any(
self._engine.edge_mode(edge) != "always" for edge in outgoing
)
if has_condition:
self._state.enter(self._engine.start_id, WorkflowStatus.WAITING_USER)
return self._passive_node_config(self._engine.start_id)
edge = next(
(
candidate
for candidate in outgoing
if self._engine.edge_mode(candidate) == "always"
),
None,
)
return (
await self._follow_edge(edge)
if edge
else self._passive_node_config(self._engine.start_id)
)
async def on_client_ready(self) -> None:
"""Replay state that may have been emitted before WebRTC data was ready."""
await self._require_output().mark_client_ready()
current_node = (
str(self._manager.current_node)
if self._manager and self._manager.current_node
else self._state.current_node_id
)
if current_node != self._state.current_node_id:
self._state.current_node_id = current_node
await self._emit_node_active(current_node)
await self._emit_variables(
reason="client_ready",
node_id=current_node,
)
def record_user_message(self, content: str) -> None:
if content and not self._ended:
self._store.record("user", content)
async def on_user_turn_end(self, content: str) -> bool:
"""Route a complete user turn before the active stage may reply."""
if not content or self._ended:
return True
async with self._turn_lock:
return await self._handle_user_turn_end(content)
async def _handle_user_turn_end(self, content: str) -> bool:
"""Serialized implementation so one user turn cannot transition twice."""
self.record_user_message(content)
self._state.begin_user_turn(content)
manager = self._require_manager()
current = self._state.current_node_id
if not current:
return True
self._state.status = WorkflowStatus.ROUTING
decision = await self._edge_evaluator.evaluate(current)
if decision.status == RouteStatus.ERROR:
await self._require_output().emit_error(
decision.error or "工作流路由失败",
node_id=current,
code="workflow_routing_error",
)
return await self._continue_current_node_after_no_transition(current)
if decision.edge and manager.current_node == current:
next_config = await self._follow_edge(
decision.edge,
triggering_user_text=content,
)
await manager.set_node_from_config(next_config)
await self._after_node_activated(
next_config,
triggering_user_text=content,
)
return True
return await self._continue_current_node_after_no_transition(current)
async def _continue_current_node_after_no_transition(
self,
node_id: str,
) -> bool:
if self._engine.node_type(node_id) != "agent":
self._state.enter(node_id, WorkflowStatus.WAITING_USER)
return True
await self._refresh_agent_prompt(node_id)
self._state.enter(node_id, WorkflowStatus.RUNNING_AGENT)
await self._require_runtime().queue_frame(LLMRunFrame())
return True
async def _select_edge(
self,
node_id: str,
) -> dict | None:
"""Compatibility helper used by automatic-node traversal and tests."""
decision = await self._edge_evaluator.evaluate(node_id)
if decision.status == RouteStatus.ERROR:
await self._require_output().emit_error(
decision.error or "工作流路由失败",
node_id=node_id,
code="workflow_routing_error",
)
return None
return decision.edge
async def on_assistant_text_end(
self,
_turn_id: str,
content: str,
interrupted: bool,
) -> None:
if not content or interrupted or self._ended:
return
self._store.record("agent", content, completed_agent_turn=True)
self._state.consume_user_turn()
if self._engine.node_type(self._state.current_node_id) == "agent":
self._state.status = WorkflowStatus.WAITING_USER
async def _refresh_agent_prompt(self, node_id: str) -> None:
await self._require_agent_stage().refresh_prompt(node_id)
def _agent_role_message(self, node_id: str) -> str:
return self._require_agent_stage().role_message(node_id)
def _router_for_node(self, node_id: str) -> WorkflowLLMRouter:
if self._agent_stage is None:
return self._router
return self._agent_stage.router_for_node(node_id, self._router)
async def _apply_agent_stage(self, node_id: str) -> None:
await self._require_agent_stage().apply(node_id)
def _agent_config(
self,
node_id: str,
leading_messages: list[dict[str, str]] | None = None,
) -> NodeConfig:
stage = self._engine.agent_stage_config(node_id)
functions: list[FlowsFunctionSchema] = []
for tool_id in stage.tool_ids:
tool = self._tool_by_id.get(str(tool_id))
if tool and tool.type == "http":
functions.append(self._flow_tool(tool, node_id))
knowledge_function = self._knowledge_function(node_id)
if knowledge_function:
functions.append(knowledge_function)
if stage.vision_enabled and self._require_runtime().vision_function:
functions.append(self._require_runtime().vision_function)
return self._require_agent_stage().node_config(
node_id,
functions=functions,
greeting_context_message=self._greeting_context_message,
leading_messages=leading_messages,
)
async def _after_node_activated(
self,
node_config: NodeConfig,
*,
triggering_user_text: str = "",
) -> None:
"""Publish activation and perform exactly one Agent entry behavior."""
node_id = str(node_config.get("name") or "")
node_type = self._engine.node_type(node_id)
if node_type != "agent":
if node_type == "start":
self._state.enter(node_id, WorkflowStatus.WAITING_USER)
return
await self._emit_node_active(node_id)
data = self._engine.data(node_id)
entry_mode = str(data.get("entryMode") or "wait_user")
if entry_mode == "fixed_speech":
entry_speech = self._store.render(str(data.get("entrySpeech") or ""))
await self._queue_visible_speech(
entry_speech,
source="workflow-fixed-reply",
node_id=node_id,
)
self._state.enter(node_id, WorkflowStatus.WAITING_USER)
self._state.consume_user_turn()
return
should_run = entry_mode == "generate" or bool(triggering_user_text)
if should_run:
self._state.enter(node_id, WorkflowStatus.RUNNING_AGENT)
await self._require_runtime().queue_frame(LLMRunFrame())
return
self._state.enter(node_id, WorkflowStatus.WAITING_USER)
async def _queue_visible_speech(
self,
text: str,
*,
source: str = "workflow-speech",
node_id: str | None = None,
) -> None:
await self._require_output().speak(
text,
source=source,
node_id=node_id,
)
def _passive_node_config(
self,
node_id: str,
task_messages: list[dict[str, str]] | None = None,
) -> NodeConfig:
"""Keep a non-conversational terminal node active without ending the call."""
return {
"name": node_id,
"role_message": self._store.render(self._engine.global_prompt()),
"task_messages": list(task_messages or []),
"functions": [],
"context_strategy": ContextStrategyConfig(strategy=ContextStrategy.APPEND),
"respond_immediately": False,
}
def _flow_tool(self, tool: RuntimeTool, node_id: str) -> FlowsFunctionSchema:
properties, required = self._tools.schema_parts(tool)
self._tools.register_secrets(tool)
async def handler(args, _flow_manager):
transition_id = self._state.transition_id
try:
result = await self._tools.execute(tool, dict(args or {}))
except ToolExecutionError as exc:
return {"status": "error", "message": str(exc)}
if (
self._state.current_node_id != node_id
or self._state.transition_id != transition_id
):
return {
"status": "stale",
"message": "工具完成时当前 Agent 已经切换,结果不再触发路由。",
}
updated_variables = list(result.get("updated_variables") or [])
if updated_variables:
await self._emit_variables(
reason="tool",
node_id=node_id,
changed=updated_variables,
)
await self._refresh_agent_prompt(node_id)
edge = self._engine.deterministic_edge(
node_id,
self._store,
include_default=False,
)
if edge:
next_config = await self._follow_edge(
edge,
triggering_user_text=(
self._state.pending_user_turn.text
if self._state.pending_user_turn
else ""
),
)
return result, self._flow_managed_transition_config(
next_config,
triggering_user_text=(
self._state.pending_user_turn.text
if self._state.pending_user_turn
else ""
),
)
return result
return FlowsFunctionSchema(
name=tool.function_name,
description=tool.description or f"调用 {tool.name}",
properties=properties,
required=required,
handler=handler,
cancel_on_interruption=True,
)
def _flow_managed_transition_config(
self,
node_config: NodeConfig,
*,
triggering_user_text: str,
) -> NodeConfig:
"""Let FlowManager finish a tool result before activating its target.
Function-returned transitions are the one place where FlowManager must
schedule the LLM run. Its pre-action only updates our explicit runtime
state; it never performs a second LLM run.
"""
node_id = str(node_config.get("name") or "")
if self._engine.node_type(node_id) != "agent":
return node_config
entry_mode = str(
self._engine.data(node_id).get("entryMode") or "wait_user"
)
should_run = entry_mode == "generate" or bool(triggering_user_text)
configured = dict(node_config)
configured["respond_immediately"] = (
should_run and entry_mode != "fixed_speech"
)
configured["pre_actions"] = [
{
"type": "workflow_function_transition_entry",
"node_id": node_id,
"entry_mode": entry_mode,
"should_run": should_run,
"handler": self._activate_from_flow_transition,
}
]
return configured
async def _activate_from_flow_transition(
self,
action: dict,
_flow_manager: FlowManager,
) -> None:
"""Apply visible entry state without manually queueing an LLM run."""
node_id = str(action.get("node_id") or "")
await self._emit_node_active(node_id)
entry_mode = str(action.get("entry_mode") or "wait_user")
if entry_mode == "fixed_speech":
entry_speech = self._store.render(
str(self._engine.data(node_id).get("entrySpeech") or "")
)
await self._queue_visible_speech(
entry_speech,
source="workflow-fixed-reply",
node_id=node_id,
)
self._state.enter(node_id, WorkflowStatus.WAITING_USER)
self._state.consume_user_turn()
elif action.get("should_run"):
self._state.enter(node_id, WorkflowStatus.RUNNING_AGENT)
else:
self._state.enter(node_id, WorkflowStatus.WAITING_USER)
def _knowledge_function(self, node_id: str) -> FlowsFunctionSchema | None:
stage = self._engine.agent_stage_config(node_id)
knowledge_id = str(stage.knowledge_base_id or "")
if not knowledge_id or stage.knowledge_mode != "on_demand":
return None
cfg = self._cfg or AssistantConfig(type="workflow")
knowledge = cfg.workflow_knowledge_bases.get(knowledge_id)
description = "在当前 Agent 绑定的知识库中检索资料。"
if knowledge:
description += f"知识库:{knowledge.name}{knowledge.description}"
async def handler(args, _flow_manager):
query = str((args or {}).get("query") or "").strip()
if not query:
return {"status": "error", "message": "检索问题为空"}
try:
async with SessionLocal() as session:
results = await search_knowledge(
session,
knowledge_id,
query,
top_k=stage.knowledge_top_n,
score_threshold=stage.knowledge_score_threshold,
)
return {"status": "ok", "results": results}
except Exception as exc: # noqa: BLE001 - tool errors are returned to the LLM
logger.warning(f"Workflow 知识库检索失败:{exc}")
return {"status": "error", "message": "知识库检索暂时不可用"}
return FlowsFunctionSchema(
name="search_knowledge_base",
description=description,
properties={
"query": {"type": "string", "description": "完整问题或检索关键词"}
},
required=["query"],
handler=handler,
cancel_on_interruption=True,
)
async def _follow_edge(
self,
edge: dict,
*,
triggering_user_text: str = "",
) -> NodeConfig:
self._state.begin_transition()
leading_messages: list[dict[str, str]] = []
speech = self._engine.edge_transition_speech(edge)
if speech:
content = self._store.render(speech).strip()
if content:
await self._queue_visible_speech(
content,
source="workflow-edge-transition",
node_id=str(edge.get("target") or "") or None,
)
leading_messages.append(
{"role": "assistant", "content": content}
)
return await self._resolve_path(
str(edge.get("target") or ""),
leading_messages=leading_messages,
triggering_user_text=triggering_user_text,
)
async def _resolve_path(
self,
node_id: str,
*,
leading_messages: list[dict[str, str]] | None = None,
triggering_user_text: str = "",
) -> NodeConfig:
context_messages = list(leading_messages or [])
for hop in range(MAX_AUTOMATIC_HOPS):
self._state.automatic_hops = hop
node_type = self._engine.node_type(node_id)
if node_type == "agent":
await self._apply_agent_stage(node_id)
agent_messages = context_messages
if (
triggering_user_text
and self._engine.data(node_id).get("contextPolicy") == "fresh"
):
agent_messages = [
{"role": "user", "content": triggering_user_text},
*context_messages,
]
return self._agent_config(node_id, agent_messages)
if node_type == "end":
await self._enter_end(node_id)
return self._passive_node_config(node_id, context_messages)
if node_type == "action":
await self._enter_action(node_id)
elif node_type == "handoff":
await self._enter_handoff(node_id)
elif node_type == "start":
self._state.enter(node_id, WorkflowStatus.WAITING_USER)
await self._emit_node_active(node_id)
else:
raise RuntimeError(f"工作流指向未知节点:{node_id}")
if not self._engine.has_outgoing(node_id):
return self._passive_node_config(node_id, context_messages)
edge = await self._select_edge(node_id)
if not edge:
return self._passive_node_config(node_id, context_messages)
self._state.begin_transition()
speech = self._engine.edge_transition_speech(edge)
if speech:
content = self._store.render(speech).strip()
if content:
target_id = str(edge.get("target") or "")
await self._queue_visible_speech(
content,
source="workflow-edge-transition",
node_id=target_id or None,
)
context_messages.append(
{"role": "assistant", "content": content}
)
node_id = str(edge.get("target") or "")
raise RuntimeError("工作流连续自动跳转超过安全上限")
async def _enter_action(self, node_id: str) -> None:
self._state.enter(node_id, WorkflowStatus.RUNNING_ACTION)
await self._emit_node_active(node_id)
data = self._engine.data(node_id)
tool_id = str(data.get("toolId") or "")
tool = self._tool_by_id.get(tool_id)
if not tool:
self._store.values["system__last_action_status"] = "error"
self._store.values["system__last_action_error"] = f"工具不存在:{tool_id}"
return
try:
arguments = self._store.render_data(data.get("arguments") or {})
result = await self._tools.execute(
tool,
arguments,
result_assignments=data.get("resultAssignments") or {},
)
updated_variables = list(result.get("updated_variables") or [])
if updated_variables:
await self._emit_variables(
reason="action",
node_id=node_id,
changed=updated_variables,
)
self._store.values["system__last_action_status"] = "ok"
self._store.values["system__last_action_error"] = ""
except (ToolExecutionError, ValueError) as exc:
self._store.values["system__last_action_status"] = "error"
self._store.values["system__last_action_error"] = str(exc)[:2048]
async def _enter_handoff(self, node_id: str) -> None:
self._state.enter(node_id, WorkflowStatus.HANDOFF)
await self._emit_node_active(node_id)
data = self._engine.data(node_id)
message = self._store.render(str(data.get("message") or ""))
await self._require_runtime().queue_frame(
OutputTransportMessageUrgentFrame(
message={
"type": "handoff-requested",
"nodeId": node_id,
"targetType": data.get("targetType", "human"),
"target": data.get("target", ""),
"message": message,
}
)
)
if message:
await self._queue_visible_speech(message)
self._store.values["system__handoff_status"] = "requested"
async def _enter_end(self, node_id: str) -> None:
self._ended = True
self._state.enter(node_id, WorkflowStatus.ENDED)
self._state.finish()
await self._emit_node_active(node_id)
runtime = self._require_runtime()
if runtime.set_knowledge_scope:
runtime.set_knowledge_scope({"mode": "disabled"})
if runtime.set_vision_scope:
runtime.set_vision_scope({"enabled": False})
if runtime.set_input_enabled:
runtime.set_input_enabled(False)
data = self._engine.data(node_id)
message = self._store.render(str(data.get("message") or ""))
scope = str(data.get("scope") or "session")
if scope == "flow":
await runtime.queue_frame(
OutputTransportMessageUrgentFrame(
message={"type": "flow-ended", "nodeId": node_id}
)
)
if message:
await self._queue_visible_speech(message)
return
runtime.call_end.begin("workflow_completed")
if message:
await self._queue_visible_speech(message)
arm_tracked = getattr(runtime.call_end, "arm_after_tracked_speech", None)
if callable(arm_tracked):
await arm_tracked()
elif message:
runtime.call_end.arm_after_speech()
else:
await runtime.call_end.finish()
async def _emit_node_active(self, node_id: str | None) -> None:
await self._require_output().emit_node_active(node_id)
async def _emit_variables(
self,
*,
reason: str,
node_id: str | None,
changed: list[str] | None = None,
) -> None:
"""Publish a safe snapshot so Workflow debug mirrors runtime state."""
await self._require_output().emit_variables(
reason=reason,
node_id=node_id,
changed=changed,
)
def _require_runtime(self) -> BrainRuntime:
if self._runtime is None:
raise RuntimeError("WorkflowBrain 尚未绑定 pipeline runtime")
return self._runtime
def _require_manager(self) -> FlowManager:
if self._manager is None:
raise RuntimeError("Workflow FlowManager 尚未初始化")
return self._manager
def _require_output(self) -> WorkflowOutput:
if self._output is None:
# A few focused unit tests bind BrainRuntime directly. Lazily create
# the output adapter so those tests exercise the same production path.
self._output = WorkflowOutput(self._store, self._require_runtime())
return self._output
def _require_agent_stage(self) -> WorkflowAgentStage:
if self._agent_stage is None:
self._agent_stage = WorkflowAgentStage(
cfg=self._cfg or AssistantConfig(type="workflow"),
engine=self._engine,
store=self._store,
runtime=self._require_runtime(),
)
return self._agent_stage