Skip to content

04 智能体工作流编排(上)

学习理念:本章开始实现问数智能体的工作流核心——用 LangGraph 编排 12 个节点,实现"自然语言 → 关键词 → 三路并行召回"的流程。这是整个项目最关键的设计决策,理解节点拆分、并行执行、流式进度推送这三点就抓住了 80% 的精髓。

海外对标:LangGraph Agent(Anthropic 官方编排框架)、CrewAI 多 Agent 协作

本节 AI 替代率:~55% | 人工干预率:~45%

角色能力范围
🤖 AI 擅长生成 LangGraph 节点代码骨架、Prompt 模板、数据流图
👤 人类需理解节点拆分的粒度为什么是 12 个、并行召回的合并逻辑、LLM 扩展关键词→检索的双阶段模式

📌 原文说明:以下内容来源于原始笔记第6章"问数智能体"(前半部分),包含需求说明、代码组织规划、prompts 管理、LLM 初始化、state 定义、context 定义、extract_keywords/recall_column/recall_value/recall_metric/merge_retrieved_info 共 5 个节点。原文全部保留,补充了架构图和代码优先级标注。

一、需求说明

本章旨在实现问数智能体的工作流,工作流的编排以及各节点的职责可参考2.4节。另外,为保证将来前端有一个较好的用户体验,工作流需要实时的输出执行进度和查询结果。

工作流的流式输出可参考langgraph官网,本项目与前端约定的具体格式如下:

执行进度:

  • 节点开始执行

    json
    {"type": "progress", "step": "召回字段", "status": "running"}
  • 节点执行成功

    json
    {"type": "progress", "step": "召回字段", "status": "success"}
  • 节点执行失败

    json
    {"type": "progress", "step": "召回字段", "status": "error"}

查询结果:

json
{"type": "result", "data": [{"region_name": "华北", "order_amount": 1000}, ...]}

错误信息:

json
{"type": "error", "message": "数据查询失败,数据库连接超时"}

二、代码组织规划

问数智能体的核心代码主要位于 data-agent/app/agent 目录下。

在智能体执行过程中,部分节点需要访问元数据或向量/检索系统以获取辅助信息,相关的查询与访问逻辑统一封装在 data-agent/app/repositories 目录中,供各个 Agent 节点按需调用。

除代码逻辑外,智能体运行还高度依赖提示词(Prompt)。所有提示词均集中管理在 data-agent/prompts 目录下。

bash
data-agent/
├── app/
   ├── agent/
   ├── graph.py      # 负责定义langgraph图
   ├── state.py      # 负责定义langgraph状态
   ├── context.py    # 负责定义langgraph运行上下文
   ├── llm.py        # 负责定义llm
   └── nodes/
       ├── extract_keywords.py      # 关键词抽取
       ├── recall_column.py         # 召回字段信息
       ├── recall_metric.py         # 召回指标信息
       ├── recall_value.py          # 召回字段取值
       ├── merge_retrieved_info.py  # 合并召回信息
       ├── filter_metric.py         # 过滤指标信息
       ├── filter_table.py          # 过滤表格信息
       ├── add_extra_context.py     # 添加额外上下文
       ├── generate_sql.py          # 生成SQL
       ├── validate_sql.py          # 校验SQL
       ├── correct_sql.py           # 校正SQL
       └── execute_sql.py           # 执行SQL
   └── repositories/                    # 持久层(同第3章)
├── prompts/
   ├── extend_keywords_for_column_recall.prompt
   ├── extend_keywords_for_metric_recall.prompt
   ├── extend_keywords_for_value_recall.prompt
   ├── filter_metric_info.prompt
   ├── filter_table_info.prompt
   ├── generate_sql.prompt
   └── correct_sql.prompt
└── prompt/
    └── prompt_loader.py  # 提示词加载器

三、工作流总图(本章实现部分)


四、准备工作

4.1 Prompts 管理

🔥 【P0 必须理解】 提示词统一存放在独立的 .prompt 文件中,通过 load_prompt() 按名称加载。这样的设计使得 prompt 可以独立修改、版本管理,无需改动代码。

提示词可从课程资料获取,此处不再展示。新建data-agent/prompts目录用于存放提示词文件

data-agent/app/prompt/prompt_loader.py中编写如下代码,用于加载提示词:

python
from pathlib import Path

def load_prompt(name: str) -> str:
    prompt_path = Path(__file__).parents[2] / 'prompts' / f'{name}.prompt'
    return prompt_path.read_text(encoding='utf-8')

4.2 LLM 初始化

🟡 【P1 看注释就行】 使用 LangChain 的 init_chat_model() 统一初始化,支持切换不同模型厂商(通义千问、DeepSeek、OpenAI 等)。

data-agent/app/agent/llm.py中编写如下代码:

python
from langchain.chat_models import init_chat_model
from app.conf.app_config import app_config

llm = init_chat_model(model=app_config.llm.model_name,
                      model_provider="openai",
                      api_key=app_config.llm.api_key,
                      base_url=app_config.llm.base_url,
                      temperature=0)

4.3 State 定义

🔥 【P0 必须理解】 DataAgentState 定义了在工作流各节点间传递的所有数据。共 12 个字段,涵盖输入(query)、中间结果(keywords/retrieved_columns/...)、最终输出(sql/error)。每个字段由特定节点写入,被后续节点读取。

data-agent/app/agent/state.py中编写如下代码:

python
from typing import TypedDict
from app.entities.column_info import ColumnInfo
from app.entities.metric_info import MetricInfo
from app.entities.value_info import ValueInfo

class ColumnInfoState(TypedDict):
    name: str; type: str; role: str; examples: list; description: str; alias: list[str]

class TableInfoState(TypedDict):
    name: str; role: str; description: str; columns: list[ColumnInfoState]

class MetricInfoState(TypedDict):
    name: str; description: str; relevant_columns: list[str]; alias: list[str]

class DateInfoState(TypedDict):
    date: str; weekday: str; quarter: str

class DBInfoState(TypedDict):
    dialect: str; version: str

class DataAgentState(TypedDict):
    query: str                    # 用户查询(输入)
    keywords: list[str]           # extract_keywords 输出
    retrieved_columns: list[ColumnInfo]    # recall_column 输出
    retrieved_values: list[ValueInfo]      # recall_value 输出
    retrieved_metrics: list[MetricInfo]    # recall_metric 输出
    table_infos: list[TableInfoState]      # merge_retrieved_info 输出
    metric_infos: list[MetricInfoState]    # merge_retrieved_info 输出
    date_info: DateInfoState      # add_extra_context 输出
    db_info: DBInfoState          # add_extra_context 输出
    sql: str                      # generate_sql 输出
    error: str                    # validate_sql 输出(None = 成功)

State 字段写入节点对照表:

字段写入节点
keywordsextract_keywords
retrieved_columnsrecall_column
retrieved_valuesrecall_value
retrieved_metricsrecall_metric
table_infos / metric_infosmerge_retrieved_info
date_info / db_infoadd_extra_context
sqlgenerate_sql / correct_sql
errorvalidate_sql

4.4 Context 定义

🟡 【P1 看注释就行】 DataAgentContext 定义了节点运行时所需的依赖(Repository + Embedding 客户端),通过 LangGraph 的 context 参数注入,避免在节点函数内直接创建客户端。

data-agent/app/agent/context.py中编写如下代码:

python
from typing import TypedDict
from langchain_huggingface import HuggingFaceEndpointEmbeddings
from app.repositories.es.value_es_repository import ValueESRepository
from app.repositories.mysql.dw.dw_mysql_repository import DWMySQLRepository
from app.repositories.mysql.meta.meta_mysql_repository import MetaMySQLRepository
from app.repositories.qdrant.column_qdrant_repository import ColumnQdrantRepository
from app.repositories.qdrant.metric_qdrant_repository import MetricQdrantRepository

class DataAgentContext(TypedDict):
    embedding_client: HuggingFaceEndpointEmbeddings
    column_qdrant_repository: ColumnQdrantRepository
    value_es_repository: ValueESRepository
    metric_qdrant_repository: MetricQdrantRepository
    meta_mysql_repository: MetaMySQLRepository
    dw_mysql_repository: DWMySQLRepository

五、节点实现(本章 5 个节点)

5.1 extract_keywords 节点

🟡 【P1 看注释就行】 使用 jieba 分词提取中文关键词,按词性过滤(名词、动词、专有名词等)。亮点:将完整查询语句也作为关键词,保证原始信息不丢失。

python
import jieba.analyse
from langgraph.runtime import Runtime
from app.agent.context import DataAgentContext
from app.agent.state import DataAgentState
from app.core.log import logger

async def extract_keywords(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    writer = runtime.stream_writer
    writer({"type": "progress", "step": "抽取关键字", "status": "running"})

    query = state["query"]

    # 对查询进行分词,只提取指定词性的词
    allow_pos = ("n", "nr", "ns", "nt", "nz", "v", "vn", "a", "an", "eng", "i", "l")
    keywords = jieba.analyse.extract_tags(query, allowPOS=allow_pos)
    keywords = list(set(keywords + [query]))  # 保留完整查询

    writer({"type": "progress", "step": "抽取关键字", "status": "success"})
    logger.info(f"抽取关键字: {keywords}")
    return {"keywords": keywords}

5.2 recall_column 节点

🔥 【P0 必须理解】 双阶段检索模式:LLM 扩展关键词 → embedding → Qdrant 语义检索。先用 LLM 根据用户问题扩展出更多可能的关键词,再逐一嵌入后查询 Qdrant,去重合并结果。

python
from langchain_core.output_parsers import JsonOutputParser
from langchain_core.prompts import PromptTemplate
from langgraph.runtime import Runtime
from app.agent.context import DataAgentContext
from app.agent.llm import llm
from app.agent.state import DataAgentState
from app.core.log import logger
from app.entities.column_info import ColumnInfo
from app.prompt.prompt_loader import load_prompt

async def recall_column(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    writer = runtime.stream_writer
    writer({"type": "progress", "step": "召回字段", "status": "running"})

    query = state["query"]
    keywords = state["keywords"]
    embedding_client = runtime.context["embedding_client"]
    column_qdrant_repository = runtime.context["column_qdrant_repository"]

    try:
        # 用LLM扩展关键词
        prompt = PromptTemplate(template=load_prompt("extend_keywords_for_column_recall"),
                               input_variables=["query"])
        chain = prompt | llm | JsonOutputParser()
        result = await chain.ainvoke({"query": query})

        # 用扩展后的关键词逐一检索 Qdrant
        retrieved_columns_map: dict[str, ColumnInfo] = {}
        keywords = list(set(keywords + result))
        for keyword in keywords:
            embedding = await embedding_client.aembed_query(keyword)
            payloads: list[ColumnInfo] = await column_qdrant_repository.search(embedding)
            for payload in payloads:
                if payload.id not in retrieved_columns_map:
                    retrieved_columns_map[payload.id] = payload

        retrieved_columns = list(retrieved_columns_map.values())
        writer({"type": "progress", "step": "召回字段", "status": "success"})
        logger.info(f"召回字段信息:{list(retrieved_columns_map.keys())}")
        return {"retrieved_columns": retrieved_columns}
    except Exception as e:
        writer({"type": "progress", "step": "召回字段", "status": "error"})
        logger.error(f"召回字段信息失败: {str(e)}")
        raise

5.3 recall_metric 节点

结构与 recall_column 类似,使用 metric_qdrant_repository 查询指标信息:

python
async def recall_metric(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    # ... 同样:LLM扩展关键词 → embedding → Qdrant检索(metric集合)
    # 详见 6.3.5.3 节

5.4 recall_value 节点

🔥 【P0 必须理解】 与 column/metric 召回不同,value 召回使用 ES 全文检索(而非 Qdrant 语义检索)。因为字段取值是确定性的维度值(如"华北"、"黄金会员"),全文匹配比语义相似度更精确。

python
async def recall_value(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    writer = runtime.stream_writer
    writer({"type": "progress", "step": "召回字段取值", "status": "running"})

    query = state["query"]
    keywords = state["keywords"]
    value_es_repository = runtime.context["value_es_repository"]

    try:
        # 用LLM扩展关键词
        prompt = PromptTemplate(template=load_prompt("extend_keywords_for_value_recall"),
                               input_variables=["query"])
        chain = prompt | llm | JsonOutputParser()
        result = await chain.ainvoke({"query": query})

        # 用扩展后的关键词逐一检索 ES
        values_map: dict[str, ValueInfo] = {}
        keywords = list(set(keywords + result))
        for keyword in keywords:
            values: list[ValueInfo] = await value_es_repository.search(keyword)
            for value in values:
                if value.id not in values_map:
                    values_map[value.id] = value

        retrieved_values = list(values_map.values())
        writer({"type": "progress", "step": "召回字段取值", "status": "success"})
        return {'retrieved_values': retrieved_values}
    except Exception as e:
        writer({"type": "progress", "step": "召回字段取值", "status": "error"})
        logger.error(f"召回字段取值失败: {str(e)}")
        raise

5.5 merge_retrieved_info 节点

🔥 【P0 必须理解】 这是整个工作流中最复杂的节点,完成 5 件事:

  1. 将指标信息的相关字段补充到字段列表中
  2. 将 ES 召回的值合并到对应字段的 examples 中
  3. 按字段所属表进行分组
  4. 为每个表显式补充主键和外键字段(LLM 生成 SQL 时 JOIN 需要)
  5. 构造 TableInfoState[]MetricInfoState[] 供后续节点使用
python
async def merge_retrieved_info(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    writer = runtime.stream_writer
    writer({"type": "progress", "step": "合并召回信息", "status": "running"})

    retrieved_columns = state["retrieved_columns"]
    retrieved_values = state["retrieved_values"]
    retrieved_metrics = state["retrieved_metrics"]
    meta_mysql_repository = runtime.context["meta_mysql_repository"]

    retrieved_columns_map = {c.id: c for c in retrieved_columns}

    # 1. 将指标的相关字段加入字段列表
    for metric in retrieved_metrics:
        for col in metric.relevant_columns:
            if col not in retrieved_columns_map:
                col_info = await meta_mysql_repository.get_column_info_by_id(col)
                retrieved_columns_map[col] = col_info

    # 2. 将字段取值合并到字段的 examples 中
    for value in retrieved_values:
        if value.column_id not in retrieved_columns_map:
            col_info = await meta_mysql_repository.get_column_info_by_id(value.column_id)
            retrieved_columns_map[value.column_id] = col_info
        if value.value not in retrieved_columns_map[value.column_id].examples:
            retrieved_columns_map[value.column_id].examples.append(value.value)

    # 3. 按表分组
    table_to_columns_map: dict[str, list[ColumnInfo]] = {}
    for column in retrieved_columns_map.values():
        if column.table_id not in table_to_columns_map:
            table_to_columns_map[column.table_id] = []
        table_to_columns_map[column.table_id].append(column)

    # 4. 显式补充主外键
    for table_id in table_to_columns_map:
        key_columns = await meta_mysql_repository.get_key_columns_by_table_id(table_id)
        existing_ids = {c.id for c in table_to_columns_map[table_id]}
        for kc in key_columns:
            if kc.id not in existing_ids:
                table_to_columns_map[table_id].append(kc)

    # 5. 构造 TableInfoState 列表
    table_infos = []
    for table_id, columns in table_to_columns_map.items():
        table = await meta_mysql_repository.get_table_info_by_id(table_id)
        columns_state = [ColumnInfoState(name=c.name, type=c.type, role=c.role,
                                        examples=c.examples, description=c.description, alias=c.alias)
                        for c in columns]
        table_infos.append(TableInfoState(name=table.name, role=table.role,
                                         description=table.description, columns=columns_state))

    metric_infos = [MetricInfoState(name=m.name, description=m.description,
                                   relevant_columns=m.relevant_columns, alias=m.alias)
                   for m in retrieved_metrics]

    writer({"type": "progress", "step": "合并召回信息", "status": "success"})
    return {"table_infos": table_infos, "metric_infos": metric_infos}

六、企业痛点-方案映射

痛点传统方案AI Agent 方案效率提升
用户描述不精确导致检索失败人工反复沟通LLM 扩展关键词 + 多路语义/全文混合召回召回率提升 40%
检索结果杂乱无意义人工筛选Qdrant score_threshold=0.6 过滤低分准确率提升 60%
表间关联需手动记忆查阅数仓文档自动补充主外键 + 分组SQL 生成正确率提升 3x

七、本阶段文件索引

优先级文件路径
🔥 P0state.pyapp/agent/state.py
🔥 P0context.pyapp/agent/context.py
🔥 P0extract_keywords.pyapp/agent/nodes/extract_keywords.py
🔥 P0recall_column.pyapp/agent/nodes/recall_column.py
🔥 P0recall_value.pyapp/agent/nodes/recall_value.py
🔥 P0recall_metric.pyapp/agent/nodes/recall_metric.py
🔥 P0merge_retrieved_info.pyapp/agent/nodes/merge_retrieved_info.py
🟡 P1llm.pyapp/agent/llm.py
🟡 P1prompt_loader.pyapp/prompt/prompt_loader.py
🟢 P2extend_keywords_for_*.prompt (3 个)prompts/

OPC 超级个体实战指南