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