Skip to content

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, ...]} 指导哪些字段保留。

python
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)}")
        raise

2.2 filter_metric 节点

与 filter_table 类似,用 LLM 过滤无关的指标:

python
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)}")
        raise

2.3 add_extra_context 节点

🟢 【P2 后面可以查】 提供当前日期信息和数据仓库版本信息,帮助 LLM 处理"去年"、"Q1"等相对时间表达,并了解数据库方言。

python
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"})
        raise

2.4 generate_sql 节点

🔥 【P0 必须理解】 核心节点。将经过过滤的表格信息、指标信息、日期信息、数据库信息以 YAML 格式序列化,与用户问题一起输入 LLM,使用 generate_sql.prompt 模板生成 SQL。yaml.dump() 是关键的序列化步骤——它让结构化数据以清晰格式呈现给 LLM。

python
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)}")
        raise

2.5 validate_sql 节点

🔥 【P0 必须理解】 使用 MySQL 的 EXPLAIN 命令验证 SQL 语法正确性。如果抛出异常,将错误信息写入 state["error"];如果成功,将 error 置为 None。后续条件边根据 error 是否为 None 决定走 execute_sql 还是 correct_sql。

python
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 工作流中只有失败才会触发。

python
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"})
        raise

2.7 execute_sql 节点

🟢 【P2 后面可以查】 最终节点。执行 SQL 并通过 runtime.stream_writer 发送结果(SSE 格式的 result 事件)。

python
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 必须理解】 这是整个项目的核心编排文件。关键设计点:

  1. 并行边extract_keywords → 同时指向 recall_column/recall_value/recall_metric
  2. 合并点:三个召回节点 → 同时指向 merge_retrieved_info
  3. 并行过滤merge_retrieved_info → 同时指向 filter_table/filter_metric
  4. 条件边validate_sql → 根据 state["error"] 选择 execute_sqlcorrect_sql
python
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 中添加测试代码,运行观察输出:

python
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 可能不准确人工 reviewEXPLAIN 自动验证 + LLM 校正回退准确率提升 60%
多表 JOIN 场景下字段选择困难手动匹配filter_table + filter_metric 双重过滤字段选择准确率 95%+
时间表达不统一手动转换add_extra_context 注入当前日期减少时间相关错误 80%

六、本阶段文件索引

优先级文件路径
🔥 P0graph.pyapp/agent/graph.py
🔥 P0generate_sql.pyapp/agent/nodes/generate_sql.py
🔥 P0validate_sql.pyapp/agent/nodes/validate_sql.py
🔥 P0filter_table.pyapp/agent/nodes/filter_table.py
🟡 P1filter_metric.pyapp/agent/nodes/filter_metric.py
🟡 P1add_extra_context.pyapp/agent/nodes/add_extra_context.py
🟡 P1correct_sql.pyapp/agent/nodes/correct_sql.py
🟢 P2execute_sql.pyapp/agent/nodes/execute_sql.py
🟢 P2generate_sql.promptprompts/generate_sql.prompt
🟢 P2correct_sql.promptprompts/correct_sql.prompt
🟢 P2filter_table_info.promptprompts/filter_table_info.prompt
🟢 P2filter_metric_info.promptprompts/filter_metric_info.prompt

OPC 超级个体实战指南