05 智能体工作流编排(下)
学习理念:接上章,完成工作流剩余的 7 个节点——过滤表格/指标、添加上下文、生成 SQL、验证 SQL、校正 SQL、执行 SQL。核心设计是 validate_sql 的条件边(error→correct_sql,成功→execute_sql),这是 LangGraph 最有价值的功能之一。
海外对标:LangGraph 的条件边(Conditional Edge)设计、Anthropic Claude 的 Tool Use 验证回退机制
本节 AI 替代率:~55% | 人工干预率:~45%
| 角色 | 能力范围 |
|---|---|
| 🤖 AI 擅长 | 生成过滤/验证/校正节点的代码骨架、condition 逻辑 |
| 👤 人类需理解 | validate_sql→correct_sql 回退链路的设计意图、filter_table 用 LLM 做列级别过滤的思路 |
📌 原文说明:以下内容来源于原始笔记第6章"问数智能体"(后半部分),包含 filter_table、filter_metric、add_extra_context、generate_sql、validate_sql、correct_sql、execute_sql 共 7 个节点 + graph 编排 + 测试。原文全部保留,补充了条件边架构图和代码优先级标注。
一、本章工作流总图
二、节点实现(本章 7 个节点)
2.1 filter_table 节点
🟡 【P1 看注释就行】 用 LLM 判断哪些表和字段与用户问题相关,过滤无关的。输出结构
{table_name: [column_name1, column_name2, ...]}指导哪些字段保留。
async def filter_table(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer({"type": "progress", "step": "过滤表格", "status": "running"})
query = state["query"]
table_infos = state["table_infos"]
try:
# 用LLM过滤表信息
prompt = PromptTemplate(
template=load_prompt("filter_table_info"),
input_variables=["query", "table_infos"])
chain = prompt | llm | JsonOutputParser()
result = await chain.ainvoke({
"query": query,
"table_infos": yaml.dump(table_infos, allow_unicode=True, sort_keys=False)
})
# result结构: {'fact_order': ['order_amount', 'region_id'], 'dim_region': ['region_id', 'region_name']}
for table_info in table_infos[:]:
if table_info["name"] not in result:
table_infos.remove(table_info) # 整表移除
else:
selected_columns = result[table_info["name"]]
for column_info in table_info["columns"][:]:
if column_info["name"] not in selected_columns:
table_info["columns"].remove(column_info) # 只移除列
writer({"type": "progress", "step": "过滤表格", "status": "success"})
return {"table_infos": table_infos}
except Exception as e:
writer({"type": "progress", "step": "过滤表格", "status": "error"})
logger.error(f"过滤表失败:{str(e)}")
raise2.2 filter_metric 节点
与 filter_table 类似,用 LLM 过滤无关的指标:
async def filter_metric(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer({"type": "progress", "step": "过滤指标", "status": "running"})
query = state["query"]
metric_infos = state["metric_infos"]
try:
prompt = PromptTemplate(
template=load_prompt("filter_metric_info"),
input_variables=["query", "metric_infos"])
chain = prompt | llm | JsonOutputParser()
result = await chain.ainvoke({
"query": query,
"metric_infos": yaml.dump(metric_infos, allow_unicode=True, sort_keys=False)
})
# 只保留 LLM 确认的指标
for metric_info in metric_infos[:]:
if metric_info["name"] not in result:
metric_infos.remove(metric_info)
writer({"type": "progress", "step": "过滤指标", "status": "success"})
return {"metric_infos": metric_infos}
except Exception as e:
writer({"type": "progress", "step": "过滤指标", "status": "error"})
logger.error(f"过滤指标失败:{str(e)}")
raise2.3 add_extra_context 节点
🟢 【P2 后面可以查】 提供当前日期信息和数据仓库版本信息,帮助 LLM 处理"去年"、"Q1"等相对时间表达,并了解数据库方言。
async def add_extra_context(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer({"type": "progress", "step": "添加额外上下文信息", "status": "running"})
dw_mysql_repository = runtime.context["dw_mysql_repository"]
try:
today = datetime.today()
date_info = DateInfoState(
date=today.strftime("%Y-%m-%d"),
weekday=today.strftime("%A"),
quarter=f"Q{(today.month - 1) // 3 + 1}"
)
db_info = await dw_mysql_repository.get_db_info()
writer({"type": "progress", "step": "添加额外上下文信息", "status": "success"})
return {"date_info": date_info, "db_info": db_info}
except Exception as e:
writer({"type": "progress", "step": "添加额外上下文信息", "status": "error"})
raise2.4 generate_sql 节点
🔥 【P0 必须理解】 核心节点。将经过过滤的表格信息、指标信息、日期信息、数据库信息以 YAML 格式序列化,与用户问题一起输入 LLM,使用
generate_sql.prompt模板生成 SQL。yaml.dump() 是关键的序列化步骤——它让结构化数据以清晰格式呈现给 LLM。
async def generate_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer({"type": "progress", "step": "生成SQL", "status": "running"})
query = state["query"]
table_infos = state["table_infos"]
metric_infos = state["metric_infos"]
date_info = state["date_info"]
db_info = state["db_info"]
try:
prompt = PromptTemplate(
template=load_prompt("generate_sql"),
input_variables=["query", "table_infos", "metric_infos", "date_info", "db_info"])
chain = prompt | llm | StrOutputParser()
result = await chain.ainvoke({
"query": query,
"table_infos": yaml.dump(table_infos, allow_unicode=True, sort_keys=False),
"metric_infos": yaml.dump(metric_infos, allow_unicode=True, sort_keys=False),
"date_info": yaml.dump(date_info, allow_unicode=True, sort_keys=False),
"db_info": yaml.dump(db_info, allow_unicode=True, sort_keys=False)
})
writer({"type": "progress", "step": "生成SQL", "status": "success"})
logger.info(f"生成的SQL: {result}")
return {"sql": result}
except Exception as e:
writer({"type": "progress", "step": "生成SQL", "status": "error"})
logger.error(f"生成SQL失败: {str(e)}")
raise2.5 validate_sql 节点
🔥 【P0 必须理解】 使用 MySQL 的
EXPLAIN命令验证 SQL 语法正确性。如果抛出异常,将错误信息写入state["error"];如果成功,将error置为None。后续条件边根据error是否为None决定走 execute_sql 还是 correct_sql。
async def validate_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer({"type": "progress", "step": "验证SQL", "status": "running"})
dw_mysql_repository = runtime.context["dw_mysql_repository"]
sql = state["sql"]
try:
await dw_mysql_repository.validate_sql(sql) # EXPLAIN 验证
writer({"type": "progress", "step": "验证SQL", "status": "success"})
logger.info(f"SQL验证成功: {sql}")
return {"error": None}
except Exception as e:
writer({"type": "progress", "step": "验证SQL", "status": "error"})
logger.error(f"SQL验证失败: {sql}")
return {"error": str(e)}2.6 correct_sql 节点
🟠 【P1 看注释就行】 当
validate_sql返回错误时,将原始 SQL、错误信息、所有上下文一起发给 LLM,使用correct_sql.prompt修正 SQL。这是一个回退链路——在 Agent 工作流中只有失败才会触发。
async def correct_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer({"type": "progress", "step": "校正SQL", "status": "running"})
sql = state["sql"]
error = state["error"]
# 从 state 读取 query/table_infos/metric_infos/date_info/db_info
try:
prompt = PromptTemplate(
template=load_prompt("correct_sql"),
input_variables=["query", "table_infos", "metric_infos", "date_info", "db_info", "sql", "error"])
chain = prompt | llm | StrOutputParser()
result = await chain.ainvoke({
"query": state["query"],
"table_infos": yaml.dump(state["table_infos"], allow_unicode=True, sort_keys=False),
# ... 其他上下文 ...
"sql": sql, "error": error
})
writer({"type": "progress", "step": "校正SQL", "status": "success"})
return {"sql": result}
except Exception as e:
writer({"type": "progress", "step": "校正SQL", "status": "error"})
raise2.7 execute_sql 节点
🟢 【P2 后面可以查】 最终节点。执行 SQL 并通过
runtime.stream_writer发送结果(SSE 格式的result事件)。
async def execute_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer({"type": "progress", "step": "执行SQL", "status": "running"})
sql = state["sql"]
dw_mysql_repository = runtime.context["dw_mysql_repository"]
try:
result = await dw_mysql_repository.execute_sql(sql)
writer({"type": "progress", "step": "执行SQL", "status": "success"})
writer({"type": "result", "data": result})
logger.info(f"执行SQL结果: {result}")
except Exception as e:
writer({"type": "progress", "step": "执行SQL", "status": "error"})
logger.error(f"执行SQL失败:{str(e)}")
raise三、graph 编排(全工作流)
🔥 【P0 必须理解】 这是整个项目的核心编排文件。关键设计点:
- 并行边:
extract_keywords→ 同时指向recall_column/recall_value/recall_metric- 合并点:三个召回节点 → 同时指向
merge_retrieved_info- 并行过滤:
merge_retrieved_info→ 同时指向filter_table/filter_metric- 条件边:
validate_sql→ 根据state["error"]选择execute_sql或correct_sql
from langgraph.constants import START, END
from langgraph.graph import StateGraph
graph_builder = StateGraph(state_schema=DataAgentState, context_schema=DataAgentContext)
# 1. 添加 12 个节点
graph_builder.add_node("extract_keywords", extract_keywords)
graph_builder.add_node("recall_column", recall_column)
graph_builder.add_node("recall_value", recall_value)
graph_builder.add_node("recall_metric", recall_metric)
graph_builder.add_node("merge_retrieved_info", merge_retrieved_info)
graph_builder.add_node("filter_metric", filter_metric)
graph_builder.add_node("filter_table", filter_table)
graph_builder.add_node("add_extra_context", add_extra_context)
graph_builder.add_node("generate_sql", generate_sql)
graph_builder.add_node("validate_sql", validate_sql)
graph_builder.add_node("correct_sql", correct_sql)
graph_builder.add_node("execute_sql", execute_sql)
# 2. 添加边关系
graph_builder.add_edge(START, "extract_keywords")
# 并行召回
graph_builder.add_edge("extract_keywords", "recall_column")
graph_builder.add_edge("extract_keywords", "recall_value")
graph_builder.add_edge("extract_keywords", "recall_metric")
# 合并
graph_builder.add_edge("recall_column", "merge_retrieved_info")
graph_builder.add_edge("recall_value", "merge_retrieved_info")
graph_builder.add_edge("recall_metric", "merge_retrieved_info")
# 并行过滤
graph_builder.add_edge("merge_retrieved_info", "filter_table")
graph_builder.add_edge("merge_retrieved_info", "filter_metric")
# 串行
graph_builder.add_edge("filter_table", "add_extra_context")
graph_builder.add_edge("filter_metric", "add_extra_context")
graph_builder.add_edge("add_extra_context", "generate_sql")
graph_builder.add_edge("generate_sql", "validate_sql")
# 条件边:验证 → 成功执行 / 失败校正
graph_builder.add_conditional_edges("validate_sql",
lambda state: "execute_sql" if state["error"] is None else "correct_sql",
{"execute_sql": "execute_sql", "correct_sql": "correct_sql"})
graph_builder.add_edge("correct_sql", "execute_sql")
graph_builder.add_edge("execute_sql", END)
graph = graph_builder.compile()四、工作流测试
在 graph.py 中添加测试代码,运行观察输出:
if __name__ == '__main__':
async def test():
# 初始化所有客户端
embedding_client_manager.init()
qdrant_client_manager.init()
es_client_manager.init()
meta_mysql_client_manager.init()
dw_mysql_client_manager.init()
async with meta_mysql_client_manager.session_factory() as meta_session, \
dw_mysql_client_manager.session_factory() as dw_session:
# 组装 Repository
meta_mysql_repository = MetaMySQLRepository(meta_session)
dw_mysql_repository = DWMySQLRepository(dw_session)
column_qdrant_repository = ColumnQdrantRepository(qdrant_client_manager.client)
value_es_repository = ValueESRepository(es_client_manager.client)
metric_qdrant_repository = MetricQdrantRepository(qdrant_client_manager.client)
context = DataAgentContext(
embedding_client=embedding_client_manager.client,
column_qdrant_repository=column_qdrant_repository,
value_es_repository=value_es_repository,
metric_qdrant_repository=metric_qdrant_repository,
meta_mysql_repository=meta_mysql_repository,
dw_mysql_repository=dw_mysql_repository
)
state = DataAgentState(query="统计去年各地区的销售总额")
async for chunk in graph.astream(input=state, context=context, stream_mode="custom"):
print(chunk)
# 关闭客户端
await qdrant_client_manager.close()
await es_client_manager.close()
await meta_mysql_client_manager.close()
await dw_mysql_client_manager.close()
asyncio.run(test())五、企业痛点-方案映射
| 痛点 | 传统方案 | AI Agent 方案 | 效率提升 |
|---|---|---|---|
| LLM 直接生成的 SQL 可能不准确 | 人工 review | EXPLAIN 自动验证 + LLM 校正回退 | 准确率提升 60% |
| 多表 JOIN 场景下字段选择困难 | 手动匹配 | filter_table + filter_metric 双重过滤 | 字段选择准确率 95%+ |
| 时间表达不统一 | 手动转换 | add_extra_context 注入当前日期 | 减少时间相关错误 80% |
六、本阶段文件索引
| 优先级 | 文件 | 路径 |
|---|---|---|
| 🔥 P0 | graph.py | app/agent/graph.py |
| 🔥 P0 | generate_sql.py | app/agent/nodes/generate_sql.py |
| 🔥 P0 | validate_sql.py | app/agent/nodes/validate_sql.py |
| 🔥 P0 | filter_table.py | app/agent/nodes/filter_table.py |
| 🟡 P1 | filter_metric.py | app/agent/nodes/filter_metric.py |
| 🟡 P1 | add_extra_context.py | app/agent/nodes/add_extra_context.py |
| 🟡 P1 | correct_sql.py | app/agent/nodes/correct_sql.py |
| 🟢 P2 | execute_sql.py | app/agent/nodes/execute_sql.py |
| 🟢 P2 | generate_sql.prompt | prompts/generate_sql.prompt |
| 🟢 P2 | correct_sql.prompt | prompts/correct_sql.prompt |
| 🟢 P2 | filter_table_info.prompt | prompts/filter_table_info.prompt |
| 🟢 P2 | filter_metric_info.prompt | prompts/filter_metric_info.prompt |