编程 pydantic-graph:用类型注解搭状态机,但先确认你真的需要图

2026-10-07 00:04:14

pydantic-graph:用类型注解搭状态机,但先确认你真的需要图

pydantic-graph 是 Pydantic 团队做的 async 图 / 状态机库,节点(node)和边(edge)都由 Python 类型注解定义。作为 Pydantic AI 的一部分开发,但它不依赖 pydantic-ai,本身是与 GenAI 无关的纯图状态机库。

  • 文档:
  • 项目:

依赖关系上,pydantic-graph 是 pydantic-ai 的必需依赖、pydantic-ai-slim 的可选依赖,也可以单独装:

pip install pydantic-graph
# 或
uv add pydantic-graph

官方对它的态度:钉枪

如果 Pydantic AI 的 agent 是锤子,multi-agent workflow 是大锤,那 graph 就是钉枪。钉枪看起来酷,但设置成本高得多,而且不会让你成为更好的建造者。

graph 是强力工具,但不是所有任务都合适。如果你不确定基于 graph 的方案是不是好主意,那它可能就没必要。只有在控制流本身是难点时——需要显式状态、分支、可恢复的转移、要把流程画成图——才考虑它。大多数 agent 用普通 Python 加 multi-agent 模式是更短的路。

关键类型

GraphRunContext:图运行时的上下文,类似 Pydantic AI 的 RunContext,持有图的状态(state)和依赖(deps),运行时传给节点。对状态类型 StateT 泛型。

End:返回值,表示图运行应结束。对图的返回类型 RunEndT 泛型。

Nodes:BaseNode 的子类定义图中的执行节点。节点通常是 dataclass,一般包含调用节点时所需的参数字段、run 方法里的业务逻辑,以及 run 方法的返回注解——pydantic-graph 读它来决定节点的出边。

节点对三类东西泛型:state(必须与所处图的状态类型一致,无状态图用 None)、deps(图的依赖类型,默认 object,不用依赖可省略)、图的返回类型(仅当节点返回 End 时相关,RunEndT 默认 Never,不返回 End 可省略,返回则必须写)。

一个中间节点示例

from dataclasses import dataclass
from pydantic_graph import BaseNode, GraphRunContext

@dataclass
class MyNode(BaseNode[MyState]):  # 状态是 MyState,不能结束运行,故省略 RunEndT
foo: int

async def run(
self,
ctx: GraphRunContext[MyState],
) -> AnotherNode:  # 返回类型决定出边
...
return AnotherNode()

让它可选地结束(foo 能被 5 整除就结束):

from dataclasses import dataclass
from pydantic_graph import BaseNode, End, GraphRunContext

@dataclass
class MyNode(BaseNode[MyState, object, int]):  # 泛型参数按位置:state, deps, 返回类型
foo: int

async def run(
self,
ctx: GraphRunContext[MyState],
) -> AnotherNode | End[int]:  # 联合返回类型 = 多条出边
if self.foo % 5 == 0:
return End(self.foo)
else:
return AnotherNode()

泛型参数只能按位置传,所以要写返回类型时必须把 deps 的默认值 object 也写上。

GraphBuilder 与完整示例

Graph 是由 GraphBuilder 产生的可执行图。builder 是把 step 函数、BaseNode 类和它们之间的边组装成图的入口,对 state(StateT)、deps(DepsT)、input(InputT,初始输入类型)、output(OutputT,最终输出类型)泛型。

从两个 BaseNode 子类构建的简单图:

from __future__ import annotations
from dataclasses import dataclass
from pydantic_graph import BaseNode, End, GraphBuilder, GraphRunContext, StepContext

@dataclass
class DivisibleBy5(BaseNode[None, object, int]):
foo: int

async def run(self, ctx: GraphRunContext[None]) -> Increment | End[int]:
if self.foo % 5 == 0:
return End(self.foo)
else:
return Increment(self.foo)

@dataclass
class Increment(BaseNode[None]):
foo: int

async def run(self, ctx: GraphRunContext[None]) -> DivisibleBy5:
return DivisibleBy5(self.foo + 1)

g = GraphBuilder(input_type=int, output_type=int)

@g.step
async def start(ctx: StepContext[None, None, int]) -> DivisibleBy5:
return DivisibleBy5(ctx.inputs)

g.add(
g.node(DivisibleBy5),
g.node(Increment),
g.edge_from(g.start_node).to(start),
)

fives_graph = g.build()

async def main():
result = await fives_graph.run(inputs=4)
print(result)
#> 5

要点:

  • 每个 BaseNode 子类用 g.node() 注册,出边由每个节点 run 的返回类型推断;
  • g.edge_from(g.start_node).to(start) 把起始节点接到入口 step;
  • g.build() 返回可执行的 Graph;
  • graph.run() 是 async 的,返回 End 节点的原始输出值。

可以用 print(fives_graph) 或 fives_graph.render() 生成 mermaid 图。

带状态的图

图里的 state 概念提供了一个可选方式:在节点运行时访问和修改一个对象(通常是 dataclass 或 Pydantic 模型)。把图想成生产线,state 就是被沿线传递、被每个节点逐步构建的引擎。

自动售货机例子:

from __future__ import annotations
from dataclasses import dataclass
from rich.prompt import Prompt
from pydantic_graph import BaseNode, End, GraphBuilder, GraphRunContext, StepContext

@dataclass
class MachineState:
user_balance: float = 0.0
product: str | None = None

@dataclass
class InsertCoin(BaseNode[MachineState]):
async def run(self, ctx: GraphRunContext[MachineState]) -> CoinsInserted:
return CoinsInserted(float(Prompt.ask('Insert coins')))

@dataclass
class CoinsInserted(BaseNode[MachineState]):
amount: float

async def run(self, ctx: GraphRunContext[MachineState]) -> SelectProduct | Purchase:
ctx.state.user_balance += self.amount
if ctx.state.product is not None:
return Purchase(ctx.state.product)
else:
return SelectProduct()

@dataclass
class SelectProduct(BaseNode[MachineState]):
async def run(self, ctx: GraphRunContext[MachineState]) -> Purchase:
return Purchase(Prompt.ask('Select product'))

PRODUCT_PRICES = {
'water': 1.25,
'soda': 1.50,
'crisps': 1.75,
'chocolate': 2.00,
}

@dataclass
class Purchase(BaseNode[MachineState, object, None]):
product: str

async def run(
self, ctx: GraphRunContext[MachineState]
) -> End[None] | InsertCoin | SelectProduct:
if price := PRODUCT_PRICES.get(self.product):
ctx.state.product = self.product
if ctx.state.user_balance >= price:
ctx.state.user_balance -= price
return End(None)
else:
diff = price - ctx.state.user_balance
print(f'Not enough money for {self.product}, need {diff:0.2f} more')
return InsertCoin()
else:
print(f'No such product: {self.product}, try again')
return SelectProduct()

g = GraphBuilder(state_type=MachineState)

@g.step
async def start(ctx: StepContext[MachineState, None, None]) -> InsertCoin:
return InsertCoin()

g.add(
g.node(InsertCoin),
g.node(CoinsInserted),
g.node(SelectProduct),
g.node(Purchase),
g.edge_from(g.start_node).to(start),
)

vending_machine_graph = g.build()

async def main():
state = MachineState()
await vending_machine_graph.run(state=state)
print(f'purchase successful item={state.product} change={state.user_balance:0.2f}')
#> purchase successful item=crisps change=0.25

要点:run 方法的返回类型很关键,用来确定出边,也是渲染 mermaid 图和运行时检测异常行为的依据(运行时会强制校验)。CoinsInserted 的 run 返回联合类型,表示可能有多条出边。Purchase 能结束运行,所以必须设置 RunEndT 泛型参数(这里图返回类型是 None,所以写 End[None])。

其他

  • 图可以迭代:Agent 在 Pydantic AI 底层就是用 pydantic-graph 编排 model 请求与响应的处理;需要更细粒度控制时用 Agent.iter。
  • 依赖注入:图和节点都支持 deps 泛型,用法与 Pydantic AI 的依赖注入一致。
  • mermaid 图:print(graph) 或 graph.render()。

如果你的分支和状态用几个 if/else 加 dataclass 就能说清,这套图带来的类型校验和可渲染性,未必抵得过多出来的那层抽象。

推荐文章

程序员茄子在线接单