玄同 765

大语言模型 (LLM) 开发工程师 | 中国传媒大学 · 数字媒体技术(智能交互与游戏设计)

CSDN · 个人主页 | GitHub · Follow


关于作者

  • 深耕领域:大语言模型开发 / RAG 知识库 / AI Agent 落地 / 模型微调
  • 技术栈:Python | RAG (LangChain / Dify + Milvus) | FastAPI + Docker
  • 工程能力:专注模型工程化部署、知识库构建与优化,擅长全流程解决方案

「让 AI 交互更智能,让技术落地更高效」
欢迎技术探讨与项目合作,解锁大模型与智能交互的无限可能!


从零构建企业级 Agent 编排框架:基于 FastAPI 的 LangGraph 风格框架设计与实现

零、前置知识:从零理解 Agent 编排

在深入探讨如何构建 Agent 编排框架之前,让我们先建立一些基础概念。如果你已经熟悉 有向无环图(DAG)状态管理FastAPI 这些概念,可以直接跳过这一节进入正文。但对于刚接触这些概念的朋友,这部分内容会帮你建立直观的理解,让后续的代码阅读更加顺畅。

0.1 什么是有向无环图(DAG)?

生活中的例子:

想象一下你在规划一次旅行。假设你要从北京出发,经过西安、成都,最后到达拉萨。这个行程有一个关键特点:你可以往西走,但不能往回走——你不可能从成都突然回到北京再重新出发。这种"只能前进,不能回头"的路线结构,就是有向无环图在现实世界中的映射。

北京
起点

西安

成都

拉萨
终点

在计算机科学中,有向无环图(Directed Acyclic Graph,简称 DAG)是一种数据结构,它由节点(Vertex)和有向边(Edge)组成,其中:

  • 节点代表计算单元或数据点,就像旅行中的每个城市
  • 有向边代表从一个节点到另一个节点的路径,就像从北京到西安的航班
  • 无环意味着没有循环路径,你不能从成都飞回北京

关键要点

  1. 拓扑排序:DAG 的重要特性是可以进行拓扑排序——把所有节点排列成一个线性序列,使得每条边的起点都在终点之前。这对于安排任务的执行顺序至关重要
  2. 依赖管理:DAG 天然适合表达"任务B依赖于任务A完成"这种依赖关系
  3. 并行执行:如果两个节点没有依赖关系,它们可以并行执行,充分利用计算资源

0.2 什么是 Reducer 模式?

生活中的例子:

想象你在整理一张表格。每周你都会收到新的数据,需要添加到表格里。传统做法是覆盖——直接用新数据替换旧数据。但有时候你不想丢失历史记录,而是想把新数据追加到末尾。这就是 Reducer 模式的核心思想。

输入

旧状态
[消息1, 消息2]

新增量
[消息3]

Reducer 合并
merge(old, new)

新状态
[消息1, 消息2, 消息3]

自动将新数据
追加到列表末尾

在编程中,Reducer 是一个函数,它接收"旧状态"和"新数据",返回"合并后的新状态"。Reducer 模式的核心优势在于:

  1. 不可变性:不直接修改旧状态,而是返回新状态,避免意外的副作用
  2. 可组合性:多个 Reducer 可以组合在一起,分别处理状态的不同部分
  3. 可预测性:给定相同的输入,总是产生相同的输出

在 Agent 框架中,Reducer 模式用于解决一个关键问题:当多个节点并行执行时,它们各自返回的状态增量应该如何合并?例如,节点A说"添加消息X",节点B说"添加消息Y",我们需要一种机制把X和Y都保留下来,而不是让后执行的覆盖先执行的。

0.3 FastAPI 异步编程基础

为什么要了解异步?

传统的 Web 框架(如 Django)采用同步模式:当一个用户请求数据库查询时,整个服务器都在等待数据库返回结果,期间无法处理其他用户的请求。这就像餐厅只有一个服务员,所有客人都得排队等一个人点完菜。

FastAPI 基于 asyncio 构建,采用异步模式。当一个请求在等待数据库时,服务器可以同时处理其他请求。这就像餐厅有多个服务员,每个客人都有专人服务,效率大幅提升。

关键概念

  • async def:声明一个异步函数,表示函数内部可以使用 await 等待异步操作
  • await:暂停当前协程,等待另一个协程完成后再继续
  • asyncio.gather:并发执行多个协程,等待所有协程完成后汇总结果

0.4 本章小结

现在你已经掌握了理解 Agent 编排框架所需的三个核心概念:

  1. DAG(有向无环图):用于表达任务之间的依赖关系和执行流程
  2. Reducer 模式:用于智能合并多个节点返回的状态增量
  3. FastAPI 异步编程:用于构建高并发的 Web 服务

这些概念将贯穿整个框架的设计。理解它们有助于你在阅读后续代码时理解"为什么要这样设计",而不仅仅是"代码做了什么"。


一、痛点场景:为什么需要自研 Agent 框架?

你是否有过这样的经历?你信心满满地基于 LangChain 构建了一个看似完美的 Agent 工作流,结果上线后问题层出不穷——工作流执行到一半卡死了,你不知道卡在哪里;用户刷新页面后整个对话状态丢失了,你不得不让用户重新开始;想要在工作流中加一个条件分支,发现 LangChain 的 Chain 根本不支持;更糟糕的是,某天 LangChain 升级了一个版本,你的代码全部报错了。

这些问题都指向了一个根本性的问题:LangChain 及其背后的 LangGraph 虽然好用,但它们是通用的轮子,不一定适合你的具体业务场景。

作为一个有多年 LLM 应用开发经验的工程师,我深知选择开源框架的利弊。LangChain/LangGraph 确实降低了 AI 应用的开发门槛,但它们的设计初衷是“通用”,而不是“企业级定制”。当你需要在生产环境中构建一个稳定、可控、可扩展的 Agent 系统时,你需要的往往不是更多的包装,而是对底层逻辑的完全掌控。

在这篇文章中,我将手把手教你如何基于 FastAPI 从零构建一个企业级的 Agent 编排框架——我们称之为 AgentFlow。这个框架会借鉴 LangGraph 的核心设计理念(节点、边、状态机),但在此基础上进行深度定制,以满足真实业务场景的需求。

更重要的是,我们会详细讲解如何集成 Milvus(向量存储)Neo4j(图数据库)PostgreSQL(关系型存储) 来构建一个完整的记忆系统。这不是简单地把三个数据库拼在一起,而是一个有设计层次、有业务逻辑的记忆架构。


一、核心架构设计理念

1.1 为什么选择图结构?

在深入代码之前,我们需要先理解一个根本性的设计决策:为什么 Agent 编排要使用有向无环图(DAG)而不是线性链?

问题背景:线性链的局限性

想象你正在开发一个智能客服系统。用户提交一个问题后,你希望系统能够:

  1. 首先理解用户的问题(意图识别)
  2. 然后查询相关知识库(知识检索)
  3. 如果知识库中没有满意答案,转人工(条件判断)
  4. 最后生成回答(回答生成)

这个流程用传统的 LangChain Chain 来实现会非常困难,因为 Chain 是线性的——每个步骤只能依赖前一步的结果,无法实现"如果…那么…"的条件分支。更重要的是,当知识检索返回多个结果时,我们可能希望并行处理这些结果来提高效率。

核心思想:DAG 赋能动态路由

**有向无环图(DAG)**正是为解决这些问题而设计的。在 DAG 结构中:

  • 节点(Node)代表计算单元,如"意图识别"、"知识检索"等
  • (Edge)代表节点之间的流向
  • 条件边允许根据状态动态决定下一步走哪条路径
  • 并行边允许同时执行多个分支
形象比喻:地铁线路图

如果你坐过北京或上海的地铁,可能会注意到地铁线路图并不是一条直线,而是一个复杂的网络。从A站到B站可能有多种路线,取决于你在换乘站的选择。DAG 之于 Agent,就像地铁线路图之于出行——它提供了灵活性选择权,而不是把乘客禁锢在单一路线上。

Yes|复杂

No|简单

用户提交问题

解析问题

识别知识

问题复杂度?

深度分析

简单解答

推荐练习

生成反馈

返回结果

与 LangChain Chain 的对比
特性 LangChain Chain AgentFlow DAG
执行顺序 严格线性 支持条件分支、并行
状态传递 全局状态 节点间可传递部分状态
可视化 困难 天然支持图可视化
条件逻辑 需要额外封装 原生支持
扩展性 链式组合 节点/边自由组合

这是一个典型的条件分支场景。传统的 LangChain Chain 是线性的——你无法轻松实现"如果…那么…"的逻辑。这就是 DAG 结构的优势:节点代表计算单元,边代表控制流,条件边可以在运行时动态决定下一步走哪条路径。

1.2 框架整体架构

AgentFlow 框架的整体架构可以分为五层:

Memory Layer - 数据持久化

短期记忆
PostgreSQL

长期记忆
Milvus

图记忆
Neo4j

Node Layer - 计算单元

LLM节点

工具节点

条件分支

并行执行

Core Engine - 编排核心

图编译器

执行引擎

状态管理

检查点

Gateway Layer - 流量控制

认证授权

限流熔断

请求路由

API Layer - 请求入口

HTTP
Request-Response

WebSocket
双向通信

SSE
服务器推送

让我逐一解释每一层的职责:

API 层负责接收外部请求。在这个层面,我们需要支持多种通讯协议——HTTP 用于简单的请求-响应,WebSocket 用于需要双向通讯的场景,SSE 用于需要流式输出的场景。

网关层处理横切关注点。认证确保请求来自合法用户,限流保护系统不被流量冲垮,路由将请求分发到正确的处理流程。

核心编排层是框架的心脏。图编译器负责将用户定义的工作流编译成可执行的 DAG。执行引擎按照拓扑顺序执行节点。状态管理维护工作流执行过程中的状态。检查点用于实现中断恢复。

节点层定义了工作流中的各种计算单元。LLM 节点调用大语言模型,工具节点执行外部工具,条件节点根据状态决定分支,并行节点同时执行多个任务。

记忆层是数据持久化的基础。短期记忆用 PostgreSQL 存储对话历史,长期记忆用 Milvus 实现语义检索,图记忆用 Neo4j 维护实体关系。


二、核心数据模型设计

2.1 状态(State)的定义

问题背景:为什么需要特殊的状态管理?

在传统的 Web 应用中,请求和响应通常是独立的——每个请求携带所有必要信息,处理完成后返回结果。但在 Agent 工作流中情况不同:一个工作流可能包含多个节点,每个节点都可能产生新的数据,这些数据需要传递给后续节点。更复杂的是,我们还可能需要支持中断恢复——如果工作流执行到一半失败了,能否从中断点继续?

这就需要一个特殊的状态管理机制:Reducer 模式

核心思想:部分状态更新

传统方式下,节点返回完整状态:

# 传统方式:节点返回完整状态
def node_a(state):
    return {"messages": [...], "context": {...}, "result": "..."}

Reducer 方式下,节点只返回增量:

# Reducer 方式:节点只返回增量
def node_a(state):
    return {"messages": [new_message]}  # 只返回新增的消息

框架会自动调用 Reducer 函数将增量合并到现有状态中。

状态模型详解

状态是工作流执行的核心。我们使用 Pydantic 来定义类型安全的状态模型:

from typing import TypedDict, List, Dict, Any, Optional, Annotated
from pydantic import BaseModel, Field
from datetime import datetime
from enum import Enum


class NodeStatus(str, Enum):
    """
    节点执行状态枚举

    每个节点在执行过程中会经历多种状态。
    跟踪这些状态对于调试和监控至关重要。
    
    状态流转:
    PENDING → RUNNING → COMPLETED/FAILED/SKIPPED
    """
    PENDING = "pending"           # 待执行:节点已加入图,但尚未开始执行
    RUNNING = "running"           # 执行中:节点正在执行,可能耗时较长
    COMPLETED = "completed"        # 执行完成:节点成功完成,正常进入下一节点
    FAILED = "failed"             # 执行失败:节点执行出错,可能需要重试或人工介入
    SKIPPED = "skipped"           # 跳过(条件不满足):因条件判断未通过而跳过


class MessageRole(str, Enum):
    """消息角色枚举
    
    对应 OpenAI 的消息角色定义:
    - system: 系统提示词,定义 Agent 的行为准则
    - user: 用户消息,来自最终用户
    - assistant: Assistant 消息,来自 LLM
    - tool: 工具结果消息,包含工具调用的返回值
    """
    SYSTEM = "system"
    USER = "user"
    ASSISTANT = "assistant"
    TOOL = "tool"


class Message(BaseModel):
    """
    对话消息模型

    记录对话中的每一条消息。
    这是短期记忆的核心数据结构。

    Attributes:
        role: 消息角色(system/user/assistant/tool)
        content: 消息内容
        tool_calls: 工具调用列表(如果有)
        tool_call_id: 工具调用 ID(用于 tool 消息)
        metadata: 附加元数据
    """
    role: MessageRole
    content: str
    tool_calls: Optional[List[Dict[str, Any]]] = None
    tool_call_id: Optional[str] = None
    metadata: Dict[str, Any] = Field(default_factory=dict)
    created_at: datetime = Field(default_factory=datetime.now)


"""
    真正的 Reducer 模式状态定义
    
    使用 Annotated 声明状态更新逻辑,框架层会自动调用 reducer 进行状态合并。
    每次节点返回的增量状态会与现有状态智能合并,而不是粗暴覆盖。
    
    Reducer 模式的核心优势:
    1. 并行安全:多个节点返回的增量会自动合并,不会相互覆盖
    2. 不可变性:状态不可变更新,避免竞态条件
    3. 可追溯:每次更新都有明确的对账关系
"""
import operator
from typing import Annotated


def merge_messages(old: List[Message], new: List[Message]) -> List[Message]:
    """
    消息追加 Reducer
    
    自动将新的消息列表追加到现有列表末尾。
    这是最常用的 Reducer,适用于多轮对话场景。
    
    工作原理:
    假设 old = [msg1, msg2],new = [msg3]
    则结果 = [msg1, msg2, msg3]
    
    Args:
        old: 现有的消息列表
        new: 新增的消息列表
        
    Returns:
        合并后的完整消息列表
    """
    return old + new


def merge_dict(old: Dict[str, Any], new: Dict[str, Any]) -> Dict[str, Any]:
    """
    字典浅合并 Reducer
    
    将新的键值对合并到现有字典中。
    如果键相同,新值覆盖旧值。
    
    工作原理:
    假设 old = {"a": 1, "b": 2},new = {"b": 3, "c": 4}
    则结果 = {"a": 1, "b": 3, "c": 4}
    
    Args:
        old: 现有字典
        new: 新增字典
        
    Returns:
        合并后的字典
    """
    merged = old.copy()
    merged.update(new)
    return merged


def merge_node_status(old: Dict[str, "NodeStatus"], new: Dict[str, "NodeStatus"]) -> Dict[str, "NodeStatus"]:
    """
    节点状态合并 Reducer
    
    合并节点执行状态,保留已完成的状态。
    后执行的状态会覆盖先前的状态(如 FAILED 会覆盖 RUNNING)。
    
    Args:
        old: 现有节点状态字典
        new: 新增节点状态字典
        
    Returns:
        合并后的节点状态字典
    """
    merged = old.copy()
    merged.update(new)
    return merged


class AgentState(TypedDict, total=False):
    """
    Agent 工作流状态 - Reducer 模式
    
    这是整个框架最核心的数据结构。
    使用 Annotated + Reducer 函数实现真正的状态自动合并。
    
    与普通 TypedDict 的区别:
    - 普通模式:节点返回 {"messages": [...]} 会直接覆盖原有 messages
    - Reducer 模式:节点返回 {"messages": [new_msg]} 会自动追加到列表末尾
    
    使用示例:
        # 节点只需返回增量
        return {"messages": [new_message]}
        
        # 执行引擎会自动执行 merge_messages(old, [new_message])
        # 最终 messages = old_messages + [new_message]
    """
    session_id: str
    user_id: str
    # 核心:每次返回的 messages 会自动 append,而不是覆盖
    messages: Annotated[List[Message], merge_messages]
    current_node: str
    # 节点状态也会智能合并
    node_status: Annotated[Dict[str, NodeStatus], merge_node_status]
    # 上下文字典智能合并
    context: Annotated[Dict[str, Any], merge_dict]
    tools: List[Dict[str, Any]]
    final_response: Optional[str]
    metadata: Annotated[Dict[str, Any], merge_dict]
    checkpoint_id: Optional[str]
接口文档:状态字段说明
字段名 类型 说明 Reducer
session_id str 会话唯一标识,用于关联同一用户的多次对话 无(不更新)
user_id str 用户标识,用于多用户隔离 无(不更新)
messages List[Message] 对话历史,LLM 生成的回复会追加到此列表 merge_messages
current_node str 当前正在执行的节点名称 无(由引擎设置)
node_status Dict[str, NodeStatus] 各节点的执行状态 merge_node_status
context Dict[str, Any] 节点间传递的临时数据 merge_dict
tools List[Dict] 可用工具列表
final_response Optional[str] 最终响应文本
metadata Dict[str, Any] 附加元数据 merge_dict
checkpoint_id Optional[str] 检查点ID,用于中断恢复

2.2 节点(Node)的抽象

问题背景:为什么需要统一的节点接口?

在 Agent 工作流中,我们需要执行各种不同类型的操作:调用 LLM、执行工具、做条件判断、并行处理任务…如果每种操作都用不同的代码来实现,工作流将变得难以维护和扩展。

核心思想:统一的节点接口

我们定义一个抽象基类 BaseNode,所有节点都必须继承这个类并实现 execute 方法。这样,无论节点内部做什么复杂的操作,对外都表现为统一的接口:“接收状态,返回状态”。

形象比喻:流水线上的工位

想象汽车组装流水线。每个工位(节点)负责不同的任务:安装发动机、安装轮胎、喷涂车漆。每个工位只需要做好自己的工作,然后把半成品传递给下一个工位。我们的节点也是类似:每个节点接收当前状态,做自己的处理,然后返回更新后的状态。

节点是工作流中的基本计算单元。我们定义一个基类来统一所有节点的行为:

from typing import Callable, Dict, Any, TypeVar, Generic
from abc import ABC, abstractmethod
import asyncio
from datetime import datetime


T = TypeVar('T', bound=AgentState)


class BaseNode(ABC, Generic[T]):
    """
    节点基类

    所有节点都需要继承这个基类。
    它定义了节点的基本接口和通用功能。

    设计原则:
    1. 节点应该是纯函数(尽量避免副作用)
    2. 节点只返回状态的部分更新(使用 Reducer 模式)
    3. 节点应该是可序列化的(支持检查点)

    Attributes:
        name: 节点名称,必须唯一
        description: 节点描述,用于文档和调试
        metadata: 节点元数据
    """

    def __init__(self, name: str, description: str = ""):
        """
        初始化节点

        Args:
            name: 节点唯一标识
            description: 节点功能描述
        """
        self.name = name
        self.description = description
        self.metadata: Dict[str, Any] = {}
        self._execution_count = 0

    @abstractmethod
    async def execute(self, state: T) -> T:
        """
        执行节点逻辑

        这是节点的核心方法。每个子类必须实现这个方法。

        Args:
            state: 当前工作流状态

        Returns:
            更新后的状态(部分更新)
        """
        pass

    async def before_execute(self, state: T) -> T:
        """
        执行前钩子

        可以在节点执行前做一些准备工作。
        默认实现不做任何事情。

        Args:
            state: 当前状态

        Returns:
            状态(可能被修改)
        """
        return state

    async def after_execute(self, state: T) -> T:
        """
        执行后钩子

        可以在节点执行后做一些清理工作。
        默认实现不做任何事情。

        Args:
            state: 执行后的状态

        Returns:
            状态(可能被修改)
        """
        return state

    def __repr__(self):
        return f"<{self.__class__.__name__}(name='{self.name}')>"


class FunctionNode(BaseNode[T]):
    """
    函数节点

    使用简单的 Python 函数来定义节点逻辑。
    适合简单的数据转换或函数调用。

    使用示例:

    def extract_keywords(state: AgentState) -> AgentState:
        keywords = extract_keywords_from_text(state['messages'][-1].content)
        return {'context': {'keywords': keywords}}

    node = FunctionNode(
        name='extract_keywords',
        func=extract_keywords
    )
    """

    def __init__(
        self,
        name: str,
        func: Callable[[T], T],
        description: str = ""
    ):
        """
        初始化函数节点

        Args:
            name: 节点名称
            func: 执行函数,接收状态,返回部分更新
            description: 节点描述
        """
        super().__init__(name, description)
        self.func = func

    async def execute(self, state: T) -> T:
        """执行函数节点"""
        # 支持同步和异步函数
        if asyncio.iscoroutinefunction(self.func):
            return await self.func(state)
        else:
            return self.func(state)


class LLMNode(BaseNode[T]):
    """
    LLM 节点

    专门用于调用大语言模型的节点。
    支持流式输出、工具调用等功能。

    Attributes:
        model: 模型名称
        temperature: 温度参数
        max_tokens: 最大 token 数
        system_prompt: 系统提示词
    """

    def __init__(
        self,
        name: str,
        model: str = "gpt-3.5-turbo",
        temperature: float = 0.7,
        max_tokens: int = 2000,
        system_prompt: str = "You are a helpful assistant.",
        description: str = ""
    ):
        """
        初始化 LLM 节点

        Args:
            name: 节点名称
            model: 模型名称
            temperature: 温度参数 (0-2)
                - 0.0-0.3: 确定性输出,更适合需要准确答案的场景
                - 0.7: 平衡,创造性和准确性兼顾
                - 1.0-2.0: 更多样性,适合创意写作
            max_tokens: 最大生成 token 数
            system_prompt: 系统提示词
            description: 节点描述
        """
        super().__init__(name, description)
        self.model = model
        self.temperature = temperature
        self.max_tokens = max_tokens
        self.system_prompt = system_prompt

    async def execute(self, state: T) -> T:
        """
        执行 LLM 调用

        核心逻辑:
        1. 构建消息列表
        2. 调用 LLM API
        3. 处理响应(包括工具调用)
        4. 返回更新后的状态
        """
        # 步骤 1:构建消息列表
        # 从状态中获取历史消息
        messages = self._build_messages(state)

        # 步骤 2:调用 LLM(这里用模拟)
        # 实际项目中应该调用真实的 LLM API
        response_content = await self._call_llm(messages)

        # 步骤 3:创建新的消息
        new_message = Message(
            role=MessageRole.ASSISTANT,
            content=response_content
        )

        # 步骤 4:只返回增量!真正的 Reducer 模式!
        # 执行引擎会自动调用 merge_messages 合并到现有状态
        return {
            "messages": [new_message],
            "context": {"last_response": response_content}
        }

    def _build_messages(self, state: T) -> List[Dict[str, Any]]:
        """构建消息列表"""
        messages = []

        # 添加系统消息
        messages.append({
            "role": "system",
            "content": self.system_prompt
        })

        # 添加历史消息
        for msg in state.get('messages', []):
            messages.append({
                "role": msg.role.value,
                "content": msg.content
            })

        return messages

    async def _call_llm(self, messages: List[Dict[str, Any]]) -> str:
        """
        调用 LLM API

        这是一个抽象方法,实际项目中需要实现具体的 API 调用。
        可以对接 OpenAI、Claude、Azure OpenAI 等。
        
        接口说明:
        - 参数:messages - 消息列表
        - 返回:LLM 生成的文本内容
        """
        # 模拟 LLM 调用
        await asyncio.sleep(0.1)
        return "这是一个模拟的 LLM 响应。在实际项目中,这里会调用真实的 LLM API。"


class ToolNode(BaseNode[T]):
    """
    工具节点

    执行外部工具或函数的节点。
    支持同步和异步工具。

    Attributes:
        tools: 工具定义列表
    """

    def __init__(
        self,
        name: str,
        tools: List[Dict[str, Any]],
        description: str = ""
    ):
        """
        初始化工具节点

        Args:
            name: 节点名称
            tools: 工具定义列表,每个工具包含 name、description、function 等
            description: 节点描述
        """
        super().__init__(name, description)
        self.tools = tools
        self._tool_registry = {tool['name']: tool for tool in tools}

    async def execute(self, state: T) -> T:
        """
        执行工具调用

        从 LLM 的响应中解析工具调用,
        执行工具,然后返回结果。
        """
        # 步骤 1:获取最后一次 LLM 响应
        messages = state.get('messages', [])
        if not messages or messages[-1].role != MessageRole.ASSISTANT:
            return state

        last_message = messages[-1]

        # 步骤 2:检查是否有工具调用
        tool_calls = last_message.tool_calls
        if not tool_calls:
            return state

        # 步骤 3:执行工具
        tool_results = []
        for tool_call in tool_calls:
            result = await self._execute_tool(tool_call)
            tool_results.append(result)

        # 步骤 4:添加工具结果到消息历史
        new_messages = list(state.get('messages', []))
        for result in tool_results:
            new_messages.append(Message(
                role=MessageRole.TOOL,
                content=result['output'],
                tool_call_id=result['tool_call_id']
            ))

        # 步骤 5:更新状态
        new_state = dict(state)
        new_state['messages'] = new_messages

        return new_state

    async def _execute_tool(self, tool_call: Dict[str, Any]) -> Dict[str, Any]:
        """
        执行单个工具调用

        Args:
            tool_call: 工具调用对象,包含 id、name、arguments

        Returns:
            工具执行结果
        """
        tool_name = tool_call.get('name')
        arguments = tool_call.get('arguments', {})

        if tool_name not in self._tool_registry:
            return {
                'tool_call_id': tool_call.get('id'),
                'tool_name': tool_name,
                'output': f"Error: Tool '{tool_name}' not found"
            }

        tool = self._tool_registry[tool_name]
        func = tool.get('function')

        # 执行工具
        try:
            if asyncio.iscoroutinefunction(func):
                output = await func(**arguments)
            else:
                output = func(**arguments)
        except Exception as e:
            output = f"Error: {str(e)}"

        return {
            'tool_call_id': tool_call.get('id'),
            'tool_name': tool_name,
            'output': str(output)
        }


class ConditionalNode(BaseNode[T]):
    """
    条件节点

    根据状态决定下一步走哪个分支。
    类似于编程中的 switch/case 或 if/else。

    Attributes:
        conditions: 条件到节点名的映射
        default: 默认节点名
    """

    def __init__(
        self,
        name: str,
        conditions: Dict[str, str],
        default: str = "__end__",
        description: str = ""
    ):
        """
        初始化条件节点

        Args:
            name: 节点名称
            conditions: 条件映射,{条件值: 目标节点名}
            default: 默认目标节点
            description: 节点描述
        """
        super().__init__(name, description)
        self.conditions = conditions
        self.default = default

    async def execute(self, state: T) -> T:
        """
        执行条件判断

        这个方法需要被重写以实现具体的判断逻辑。
        默认实现返回 default 节点。
        """
        # 实际项目中,这里应该根据状态判断
        # 这里用简化逻辑作为示例
        result = self.evaluate_condition(state)
        next_node = self.conditions.get(result, self.default)

        new_state = dict(state)
        new_state['context'] = new_state.get('context', {})
        new_state['context']['next_node'] = next_node

        return new_state

    def evaluate_condition(self, state: T) -> str:
        """
        评估条件

        子类应该重写这个方法来实现具体的判断逻辑。

        Args:
            state: 当前状态

        Returns:
            条件值,用于查找目标节点
        """
        return "default"
节点类型快速参考
节点类型 用途 适用场景
FunctionNode 执行简单的 Python 函数 数据转换、计算任务
LLMNode 调用大语言模型 生成文本、理解意图
ToolNode 执行外部工具 API 调用、数据库查询
ConditionalNode 条件分支 业务逻辑判断、路由选择

2.3 边(Edge)的抽象

边定义了节点之间的流向:

from typing import Callable, Dict, Any, List, Optional
from enum import Enum


class EdgeType(str, Enum):
    """边类型枚举"""
    NORMAL = "normal"           # 普通边:固定流向
    CONDITIONAL = "conditional" # 条件边:动态路由
    START = "start"             # 起始边
    END = "end"                 # 终止边


class BaseEdge:
    """
    边基类

    边定义了节点之间的流向。
    """

    def __init__(
        self,
        source: str,
        target: str,
        edge_type: EdgeType = EdgeType.NORMAL
    ):
        """
        初始化边

        Args:
            source: 源节点名称
            target: 目标节点名称
            edge_type: 边类型
        """
        self.source = source
        self.target = target
        self.edge_type = edge_type


class ConditionalEdge(BaseEdge):
    """
    条件边

    根据条件函数的返回值动态决定目标节点。

    Attributes:
        condition_func: 条件函数,接收状态,返回目标节点名
        condition_mapping: 条件值到节点名的映射
    """

    def __init__(
        self,
        source: str,
        condition_func: Callable[[Any], str],
        condition_mapping: Dict[str, str]
    ):
        """
        初始化条件边

        Args:
            source: 源节点名称
            condition_func: 条件函数
            condition_mapping: 条件值到节点名的映射
        """
        super().__init__(source, "", EdgeType.CONDITIONAL)
        self.condition_func = condition_func
        self.condition_mapping = condition_mapping

    def get_target(self, state: Any) -> str:
        """
        根据状态获取目标节点

        Args:
            state: 当前状态

        Returns:
            目标节点名称
        """
        # 执行条件函数
        result = self.condition_func(state)

        # 查找对应的目标节点
        return self.condition_mapping.get(result, "__end__")

三、图编译器与执行引擎

3.1 图的定义与编译

问题背景:为什么要编译而不是直接执行?

你可能会问:为什么要有"编译"这个步骤?直接在定义图的时候就执行不行吗?

答案是:编译过程有几个重要作用:

  1. 验证图的合法性:检查是否有环、节点是否都存在
  2. 优化执行顺序:通过拓扑排序确定最佳执行顺序
  3. 构建索引:加速运行时查找节点和边
  4. 序列化支持:编译后的图可以保存到磁盘,支持重启恢复
核心思想:先定义后执行

用户使用流畅的链式 API 定义工作流,然后调用 compile() 方法将图编译成可执行的形式。这个过程类似于 SQL 的"准备语句"——先解析和验证,再执行。

现在我们有了节点和边的抽象,接下来是图编译器——它负责将用户定义的工作流编译成可执行的 DAG:

from typing import Dict, Set, List, Optional, Any
import networkx as nx
from collections import defaultdict


class Graph:
    """
    有向无环图 (DAG)

    这是 AgentFlow 框架的核心数据结构。
    它使用 NetworkX 来维护图的拓扑结构。

    Attributes:
        nodes: 节点字典
        edges: 边列表
        graph: NetworkX 有向图
        entry_point: 入口节点
    """

    def __init__(self):
        """初始化空图"""
        self.nodes: Dict[str, BaseNode] = {}
        self.edges: List[BaseEdge] = []
        self.graph = nx.DiGraph()
        self.entry_point: Optional[str] = None

    def add_node(self, node: BaseNode):
        """
        添加节点到图

        Args:
            node: 要添加的节点

        Raises:
            ValueError: 如果节点名称已存在
        """
        if node.name in self.nodes:
            raise ValueError(f"Node '{node.name}' already exists")

        self.nodes[node.name] = self.graph.add_node(node.name)
        self.nodes[node.name] = node

    def add_edge(
        self,
        source: str,
        target: str,
        edge_type: EdgeType = EdgeType.NORMAL
    ):
        """
        添加边到图

        Args:
            source: 源节点名称
            target: 目标节点名称
            edge_type: 边类型

        Raises:
            ValueError: 如果节点不存在
        """
        # 验证节点存在
        if source not in self.nodes:
            raise ValueError(f"Source node '{source}' does not exist")
        if target not in self.nodes:
            raise ValueError(f"Target node '{target}' does not exist")

        # 创建边对象
        edge = BaseEdge(source, target, edge_type)
        self.edges.append(edge)

        # 添加到 NetworkX 图
        self.graph.add_edge(source, target)

    def add_conditional_edges(
        self,
        source: str,
        condition_edge: ConditionalEdge
    ):
        """
        添加条件边

        条件边需要特殊处理,因为目标节点在运行时才确定。

        Args:
            source: 源节点名称
            condition_edge: 条件边对象
        """
        if source not in self.nodes:
            raise ValueError(f"Source node '{source}' does not exist")

        self.edges.append(condition_edge)
        # 条件边暂时不添加到 NetworkX(目标不固定)

    def set_entry_point(self, node_name: str):
        """
        设置入口节点

        Args:
            node_name: 入口节点名称
        """
        if node_name not in self.nodes:
            raise ValueError(f"Entry point '{node_name}' does not exist")

        self.entry_point = node_name

    def validate(self) -> bool:
        """
        验证图的合法性

        检查:
        1. 是否有入口节点
        2. 图是否有环(DAG 必要条件)
        3. 所有边是否连接到存在的节点

        Returns:
            是否通过验证

        Raises:
            ValueError: 如果图不合法
        """
        # 检查入口节点
        if not self.entry_point:
            raise ValueError("Entry point not set")

        # 检查是否有环
        if not nx.is_directed_acyclic_graph(self.graph):
            raise ValueError("Graph contains cycles")

        # 检查边连接
        for edge in self.edges:
            if edge.source not in self.nodes or edge.target not in self.nodes:
                raise ValueError(f"Invalid edge: {edge.source} -> {edge.target}")

        return True

    def get_topological_order(self) -> List[str]:
        """
        获取拓扑排序结果

        这是执行节点的顺序。
        注意:这只是静态分析,条件边的目标在运行时才确定。

        Returns:
            节点名称列表
        """
        try:
            return list(nx.topological_sort(self.graph))
        except nx.NetworkXError as e:
            raise ValueError(f"Failed to get topological order: {e}")

    def get_node_dependencies(self, node_name: str) -> Set[str]:
        """
        获取节点的依赖节点

        Args:
            node_name: 节点名称

        Returns:
            依赖节点名称集合
        """
        if node_name not in self.graph:
            return set()

        # 获取所有前驱节点
        return set(self.graph.predecessors(node_name))

    def get_next_nodes(self, current_node: str, state: Any) -> List[str]:
        """
        获取当前节点的下一个节点

        这个方法会考虑条件边,根据状态动态决定目标。

        Args:
            current_node: 当前节点名称
            state: 当前状态

        Returns:
            下一个节点名称列表
        """
        next_nodes = []

        # 遍历所有边
        for edge in self.edges:
            if edge.source != current_node:
                continue

            if edge.edge_type == EdgeType.NORMAL:
                # 普通边,直接添加目标
                next_nodes.append(edge.target)

            elif edge.edge_type == EdgeType.CONDITIONAL:
                # 条件边,根据状态计算目标
                if isinstance(edge, ConditionalEdge):
                    target = edge.get_target(state)
                    if target and target != "__end__":
                        next_nodes.append(target)

        return next_nodes


class StateGraph:
    """
    状态图

    这是用户使用的主要接口。
    它封装了图构建的细节,提供链式 API。

    使用示例:

    from agentflow import StateGraph, START, END
    from agentflow.nodes import LLMNode, ToolNode

    graph = StateGraph(AgentState)

    # 添加节点
    graph.add_node(LLMNode("llm_node"))
    graph.add_node(ToolNode("tool_node"))

    # 添加边
    graph.add_edge(START, "llm_node")
    graph.add_edge("llm_node", "tool_node")
    graph.add_edge("tool_node", END)

    # 编译
    compiled = graph.compile()
    """

    START = "__start__"
    END = "__end__"

    def __init__(self, state_class: type):
        """
        初始化状态图

        Args:
            state_class: 状态类(TypedDict 或 Pydantic BaseModel)
        """
        self.state_class = state_class
        self.graph = Graph()
        self._conditional_edges: Dict[str, ConditionalEdge] = {}

    def add_node(self, node: BaseNode) -> 'StateGraph':
        """
        添加节点

        Args:
            node: 节点对象

        Returns:
            self(支持链式调用)
        """
        self.graph.add_node(node)
        return self

    def add_edge(self, source: str, target: str) -> 'StateGraph':
        """
        添加普通边

        Args:
            source: 源节点名称(可以使用 START/END 常量)
            target: 目标节点名称

        Returns:
            self
        """
        # 转换 START/END
        source = self._resolve_node_name(source)
        target = self._resolve_node_name(target)

        # 入口点特殊处理
        if source == self.START:
            self.graph.set_entry_point(target)
        else:
            self.graph.add_edge(source, target)

        return self

    def add_conditional_edges(
        self,
        source: str,
        condition_func: Callable[[Any], str],
        path_mapping: Dict[str, str]
    ) -> 'StateGraph':
        """
        添加条件边

        Args:
            source: 源节点名称
            condition_func: 条件函数
            path_mapping: 条件值到目标节点的映射

        Returns:
            self
        """
        source = self._resolve_node_name(source)

        # 创建条件边
        condition_edge = ConditionalEdge(
            source=source,
            condition_func=condition_func,
            condition_mapping=path_mapping
        )

        self._conditional_edges[source] = condition_edge
        self.graph.add_conditional_edges(source, condition_edge)

        return self

    def compile(self, checkpointer: Any = None) -> 'CompiledGraph':
        """
        编译图

        Args:
            checkpointer: 检查点存储器

        Returns:
            可执行的编译图对象
        """
        # 验证图
        self.graph.validate()

        # 返回编译后的图
        return CompiledGraph(
            graph=self.graph,
            state_class=self.state_class,
            checkpointer=checkpointer
        )

    def _resolve_node_name(self, name: str) -> str:
        """
        解析节点名称

        将 START/END 常量转换为实际的节点名称。

        Args:
            name: 节点名称

        Returns:
            解析后的名称
        """
        if name == self.START:
            if not self.graph.entry_point:
                raise ValueError("Entry point not set")
            return self.graph.entry_point
        return name
拓扑排序示例:理解执行顺序

假设我们定义了以下工作流:

  • 节点A(入口)→ 节点B → 节点C
  • 节点A → 节点D

拓扑排序的结果可能是:A → B → D → CA → D → B → C

无论哪种顺序,A 永远在最前面(因为它是入口),C 永远在最后(因为它没有后继)。NetworkX 的 topological_sort 函数会返回一个有效的顺序。

3.2 执行引擎

问题背景:执行引擎的核心职责

如果说图编译器是"规划师",那执行引擎就是"建筑师"。它负责按照规划执行每一个步骤,处理各种边界情况,并确保整个过程稳定可靠。

执行引擎需要处理的核心问题包括:

  1. 状态流转:如何从一个节点传递到下一个节点
  2. 条件分支:如何根据状态动态选择分支
  3. 并行执行:如何并发执行多个节点并合并结果
  4. 错误处理:节点失败时如何处理
  5. 中断恢复:如何支持从检查点恢复
分步执行流程图解

让我们用一个具体例子来理解执行引擎的工作流程。假设有如下工作流:

复杂

简单

开始

执行
节点A

条件
判断

执行
节点C

执行
节点D

执行
节点E

结束

执行引擎的工作步骤

第1步:初始化
  - current_state = 初始状态
  - current_node = "节点A"(入口节点)
  - 队列:[A]

第2步:执行节点A
  - 从队列取出 A,执行
  - A 执行完成,更新状态
  - 获取 A 的下一个节点:B
  - 队列:[B]

第3步:执行条件判断
  - 从队列取出 B,执行
  - B 根据状态判断,返回 "复杂"
  - 获取 B 的下一个节点:C
  - 队列:[C]

第4步:执行节点C
  - 从队列取出 C,执行
  - C 执行完成
  - 获取 C 的下一个节点:E
  - 队列:[E]

第5步:执行节点E
  - 从队列取出 E,执行
  - E 执行完成
  - 获取 E 的下一个节点:无
  - 队列:[]

执行结束

现在是最核心的部分——执行引擎。它负责按照正确的顺序执行节点:


由于博客内容超出上限,后续内容请点击此处链接阅读!

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐