LangGraph高级状态设计-Reducer与状态分层
*LangGraph 三层状态设计——全局状态、节点私有状态、子图状态,以及三种 Reducer 合并策略*
LangGraph 高级状态设计:Reducer 与状态分层
1.1 默认 TypedDict 的覆盖问题
LangGraph 三层状态设计——全局状态、节点私有状态、子图状态,以及三种 Reducer 合并策略
LangGraph 的状态用 TypedDict(Python 内置的类型工具,用于定义字典的结构和字段类型,让代码编辑器能检查字段名称和类型是否正确)来定义。
LangGraph 状态的默认更新策略是覆盖:节点返回的字典中,哪个字段有值,就把该字段的当前值完全替换。
from typing import TypedDict
from langgraph.graph import StateGraph, START, END
class SimpleState(TypedDict):
messages: list[str]
count: int
def node_a(state: SimpleState) -> dict:
return {"messages": ["来自节点A的消息"], "count": 1}
def node_b(state: SimpleState) -> dict:
return {"messages": ["来自节点B的消息"], "count": 2}
如果 node_a 和 node_b 并行执行,最终 messages 的值只会是其中一个节点的返回值,另一个被覆盖丢失了。即使串行执行,节点 B 的返回也会完全替换节点 A 写入的内容,而不是追加。
这在对话历史、搜索结果列表、并行任务结果汇总等场景中是致命的缺陷。解决方案是 Reducer。
1.2 Annotated + Reducer:自定义合并逻辑
Python 的 Annotated 类型允许在类型注解中附加额外信息。LangGraph 利用这个特性,让开发者为每个字段指定一个 reducer 函数(归约函数:决定"新值如何与旧值合并"的逻辑,例如追加、去重、取最新),在状态更新时控制如何合并新旧值。
from typing import Annotated, TypedDict
class State(TypedDict):
# operator.add 就是 reducer:新值追加到旧列表末尾
messages: Annotated[list[str], operator.add]
# 普通字段:新值覆盖旧值(默认行为)
current_step: str
1.2.1 operator.add:列表追加
import operator
from typing import Annotated, TypedDict
class ConversationState(TypedDict):
messages: Annotated[list[str], operator.add]
summary: str
def add_message_a(state: ConversationState) -> dict:
return {"messages": ["用户:你好"]}
def add_message_b(state: ConversationState) -> dict:
return {"messages": ["助手:你好,有什么可以帮你?"]}
# 执行后 messages = ["用户:你好", "助手:你好,有什么可以帮你?"]
# 而不是只保留最后一个节点的返回值
operator.add 对列表来说等同于 list_a + list_b,即追加操作。
1.2.2 自定义 Reducer:去重
from typing import Annotated, TypedDict
def deduplicate(existing: list[str], new: list[str]) -> list[str]:
"""合并列表并去重,保持原有顺序"""
seen = set(existing)
result = list(existing)
for item in new:
if item not in seen:
result.append(item)
seen.add(item)
return result
class SearchState(TypedDict):
# 搜索结果去重合并,避免重复条目
search_results: Annotated[list[str], deduplicate]
query: str
1.2.3 自定义 Reducer:取最新值(带时间戳)
from typing import Annotated, TypedDict
from datetime import datetime
def keep_latest(
existing: dict | None,
new: dict | None
) -> dict | None:
"""保留时间戳更新的值"""
if existing is None:
return new
if new is None:
return existing
# 比较时间戳,保留较新的
existing_ts = existing.get("timestamp", "")
new_ts = new.get("timestamp", "")
return new if new_ts >= existing_ts else existing
class MonitorState(TypedDict):
# 多个监控节点并行写入,只保留最新的
latest_status: Annotated[dict | None, keep_latest]
alerts: Annotated[list[str], operator.add]
1.2.4 自定义 Reducer:限制列表长度
def make_bounded_list(max_size: int):
"""工厂函数:生成有长度限制的列表 reducer"""
def bounded_add(existing: list, new: list) -> list:
combined = existing + new
# 保留最新的 max_size 个元素
return combined[-max_size:]
return bounded_add
class AgentState(TypedDict):
# 最多保留最近 10 条消息,防止 context 过长
messages: Annotated[list[str], make_bounded_list(10)]
tool_calls: Annotated[list[dict], make_bounded_list(50)]
1.3 MessagesState:内置消息列表
处理对话历史是最常见的场景。LangGraph 内置了 MessagesState,省去手动定义 Annotated[list[BaseMessage], ...] 的样板代码:
from langgraph.graph import MessagesState
from langchain_core.messages import HumanMessage, AIMessage, SystemMessage
# MessagesState 等价于:
# class MessagesState(TypedDict):
# messages: Annotated[list[AnyMessage], add_messages]
class MyAgentState(MessagesState):
# 继承 MessagesState,添加自定义字段
user_id: str
session_context: dict
add_messages 是 LangGraph 内置的 reducer,比简单的 operator.add 更智能:
- 追加新消息
- 如果新消息的
id与已有消息相同,则更新(而非追加),支持消息编辑场景
from langgraph.graph import MessagesState
from langchain_openai import ChatOpenAI
from langgraph.graph import StateGraph, START, END
llm = ChatOpenAI(model="gpt-4o-mini")
def chat_node(state: MessagesState) -> dict:
response = llm.invoke(state["messages"])
return {"messages": [response]} # 直接追加 AIMessage
builder = StateGraph(MessagesState)
builder.add_node("chat", chat_node)
builder.add_edge(START, "chat")
builder.add_edge("chat", END)
graph = builder.compile()
result = graph.invoke({"messages": [HumanMessage(content="什么是 LangGraph?")]})
1.4 状态分层:InputState / OutputState / PrivateState
随着图的复杂度增加,状态结构也会膨胀。一个大型 Agent 可能有数十个字段,其中很多是内部中间计算结果,不应该暴露给 API 调用方。
LangGraph 支持为同一个图定义不同的"状态视图":
- InputState:API 入口只接受这些字段
- OutputState:API 出口只返回这些字段
- PrivateState:图内部节点之间传递的完整状态(包含中间字段)
1.4.1 代码示例:只接收 query,只返回 answer
from typing import Annotated, TypedDict
import operator
from langgraph.graph import StateGraph, START, END
from langchain_openai import ChatOpenAI
from langchain_core.messages import HumanMessage
# 1. 定义三层状态
class InputState(TypedDict):
"""外部 API 只传入这些字段"""
query: str
class OutputState(TypedDict):
"""外部 API 只返回这些字段"""
answer: str
sources: list[str]
class PrivateState(TypedDict):
"""图内部完整状态,包含中间处理字段"""
query: str # 来自 InputState
expanded_queries: list[str] # 查询扩展结果(内部用)
retrieved_chunks: Annotated[list[str], operator.add] # 检索到的文档片段
reranked_chunks: list[str] # 重排序后的片段(内部用)
answer: str # 最终答案
sources: list[str] # 来源列表
intermediate_scores: list[float] # 相关性评分(内部用,不暴露)
# 2. 定义各节点
llm = ChatOpenAI(model="gpt-4o-mini")
def expand_query(state: PrivateState) -> dict:
"""将原始查询扩展为多个变体"""
response = llm.invoke([
HumanMessage(content=f"""将以下问题改写为3个不同表述,用于检索:
问题:{state['query']}
只输出改写后的问题,每行一个,不加序号。""")
])
expanded = [state["query"]] + response.content.strip().split("\n")
return {"expanded_queries": expanded[:4]} # 最多4个
def retrieve_documents(state: PrivateState) -> dict:
"""模拟文档检索(实际应对接向量数据库)"""
# 实际使用时替换为真实的检索逻辑
mock_chunks = [
f"关于'{q}'的相关文档片段..." for q in state["expanded_queries"]
]
scores = [0.9, 0.85, 0.7, 0.6][:len(mock_chunks)]
return {
"retrieved_chunks": mock_chunks,
"intermediate_scores": scores,
}
def rerank_and_filter(state: PrivateState) -> dict:
"""重排序并过滤低质量文档"""
# 按分数过滤,只保留相关性 > 0.7 的片段
filtered = [
chunk for chunk, score
in zip(state["retrieved_chunks"], state.get("intermediate_scores", []))
if score > 0.7
]
return {"reranked_chunks": filtered or state["retrieved_chunks"][:2]}
def generate_answer(state: PrivateState) -> dict:
"""基于检索结果生成最终答案"""
context = "\n\n".join(state["reranked_chunks"])
response = llm.invoke([
HumanMessage(content=f"""基于以下文档回答问题:
文档内容:
{context}
问题:{state['query']}
请给出准确、简洁的回答,并标注信息来源。""")
])
return {
"answer": response.content,
"sources": [f"文档片段 {i+1}" for i in range(len(state["reranked_chunks"]))]
}
# 3. 构建图,指定输入/输出状态类型
builder = StateGraph(
PrivateState,
input=InputState, # 指定输入状态类型
output=OutputState, # 指定输出状态类型
)
builder.add_node("expand_query", expand_query)
builder.add_node("retrieve_documents", retrieve_documents)
builder.add_node("rerank_and_filter", rerank_and_filter)
builder.add_node("generate_answer", generate_answer)
builder.add_edge(START, "expand_query")
builder.add_edge("expand_query", "retrieve_documents")
builder.add_edge("retrieve_documents", "rerank_and_filter")
builder.add_edge("rerank_and_filter", "generate_answer")
builder.add_edge("generate_answer", END)
graph = builder.compile()
# 4. 调用:只传入 query,只得到 answer 和 sources
result = graph.invoke({"query": "LangGraph 的 Send API 怎么使用?"})
print(result)
# 输出:{"answer": "...", "sources": [...]}
# intermediate_scores, retrieved_chunks 等内部字段不会出现在结果中
1.4.2 状态分层架构图
1.5 用 Pydantic 替代 TypedDict
TypedDict 只做类型注解,运行时不做任何验证。Pydantic(一个 Python 数据验证库,声明字段后会在运行时自动检查类型是否符合要求)提供字段验证、默认值、类型强制转换,适合对数据质量有要求的场景。
from pydantic import BaseModel, Field, field_validator, model_config
from typing import Annotated
import operator
class AgentState(BaseModel):
"""用 Pydantic 定义状态,支持字段验证和默认值(Pydantic v2)"""
model_config = {"validate_assignment": True} # 允许字段赋值时校验
query: str = Field(..., min_length=1, max_length=1000, description="用户查询")
max_retries: int = Field(default=3, ge=1, le=10, description="最大重试次数")
messages: list[dict] = Field(default_factory=list)
retrieved_docs: list[str] = Field(default_factory=list)
answer: str = Field(default="")
metadata: dict = Field(default_factory=dict)
@field_validator("query")
@classmethod
def strip_query(cls, v: str) -> str:
"""自动去除首尾空格"""
return v.strip()
TypedDict 与 Pydantic 的选型建议:
| 场景 | 推荐方案 | 原因 |
|---|---|---|
| 快速原型、内部工具 | TypedDict | 零开销,语法简单 |
| 对外 API、生产系统 | Pydantic | 自动验证,减少运行时错误 |
| 需要默认值 | Pydantic | TypedDict 不支持默认值 |
| 需要字段文档 | Pydantic Field | 更清晰的 schema |
| 性能敏感路径 | TypedDict | Pydantic 有额外验证开销 |
1.6 大型项目的状态设计原则
1.6.1 原则一:最小暴露
每个节点只应该访问它真正需要的字段。通过 InputState/OutputState/PrivateState 分层,确保中间计算结果不泄漏到外部接口。
1.6.2 原则二:字段语义清晰
状态字段命名应该能直接表达其生命周期:
class WellDesignedState(TypedDict):
# 输入字段:来自外部,不会被内部节点修改
user_query: str
user_id: str
# 中间字段:节点计算结果,命名包含阶段信息
parsed_intent: str # parse 阶段产出
retrieval_results: list # retrieval 阶段产出
ranked_results: list # ranking 阶段产出
# 输出字段:最终结果
final_answer: str
confidence_score: float
1.6.3 原则三:Reducer 必须是纯函数
Reducer 函数会被 LangGraph 多次调用(包括状态重放、检查点恢复),必须是无副作用的纯函数(纯函数:相同输入永远得到相同输出,且不会修改函数外部的任何变量):
# 错误示范:有副作用
external_log = []
def bad_reducer(existing: list, new: list) -> list:
external_log.append(f"合并了 {len(new)} 条") # 副作用!
return existing + new
# 正确示范:纯函数
def good_reducer(existing: list, new: list) -> list:
return existing + new
1.6.4 原则四:为并行写入设计 Reducer
只要存在并行节点向同一字段写入的可能,就必须为该字段定义 Reducer。遗漏 Reducer 是并行场景中最常见的 bug 来源:
class ParallelSafeState(TypedDict):
# 所有可能被并行写入的字段都要定义 Reducer
search_results: Annotated[list[str], operator.add]
extracted_entities: Annotated[list[str], deduplicate]
error_logs: Annotated[list[str], operator.add]
# 只会被单个节点写入的字段,无需 Reducer
final_summary: str # 只有最后的汇总节点写入
status: str # 只有状态控制节点写入