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"}
查询结果:
{"type": "result", "data": [{"region_name": "华北", "order_amount": 1000}, ...]}错误信息:
{"type": "error", "message": "数据查询失败,数据库连接超时"}二、代码组织规划
问数智能体的核心代码主要位于 data-agent/app/agent 目录下。
在智能体执行过程中,部分节点需要访问元数据或向量/检索系统以获取辅助信息,相关的查询与访问逻辑统一封装在 data-agent/app/repositories 目录中,供各个 Agent 节点按需调用。
除代码逻辑外,智能体运行还高度依赖提示词(Prompt)。所有提示词均集中管理在 data-agent/prompts 目录下。
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中编写如下代码,用于加载提示词:
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中编写如下代码:
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中编写如下代码:
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 字段写入节点对照表:
| 字段 | 写入节点 |
|---|---|
| keywords | extract_keywords |
| retrieved_columns | recall_column |
| retrieved_values | recall_value |
| retrieved_metrics | recall_metric |
| table_infos / metric_infos | merge_retrieved_info |
| date_info / db_info | add_extra_context |
| sql | generate_sql / correct_sql |
| error | validate_sql |
4.4 Context 定义
🟡 【P1 看注释就行】
DataAgentContext定义了节点运行时所需的依赖(Repository + Embedding 客户端),通过 LangGraph 的context参数注入,避免在节点函数内直接创建客户端。
在data-agent/app/agent/context.py中编写如下代码:
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 分词提取中文关键词,按词性过滤(名词、动词、专有名词等)。亮点:将完整查询语句也作为关键词,保证原始信息不丢失。
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,去重合并结果。
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)}")
raise5.3 recall_metric 节点
结构与 recall_column 类似,使用 metric_qdrant_repository 查询指标信息:
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 语义检索)。因为字段取值是确定性的维度值(如"华北"、"黄金会员"),全文匹配比语义相似度更精确。
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)}")
raise5.5 merge_retrieved_info 节点
🔥 【P0 必须理解】 这是整个工作流中最复杂的节点,完成 5 件事:
- 将指标信息的相关字段补充到字段列表中
- 将 ES 召回的值合并到对应字段的 examples 中
- 按字段所属表进行分组
- 为每个表显式补充主键和外键字段(LLM 生成 SQL 时 JOIN 需要)
- 构造
TableInfoState[]和MetricInfoState[]供后续节点使用
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 |
七、本阶段文件索引
| 优先级 | 文件 | 路径 |
|---|---|---|
| 🔥 P0 | state.py | app/agent/state.py |
| 🔥 P0 | context.py | app/agent/context.py |
| 🔥 P0 | extract_keywords.py | app/agent/nodes/extract_keywords.py |
| 🔥 P0 | recall_column.py | app/agent/nodes/recall_column.py |
| 🔥 P0 | recall_value.py | app/agent/nodes/recall_value.py |
| 🔥 P0 | recall_metric.py | app/agent/nodes/recall_metric.py |
| 🔥 P0 | merge_retrieved_info.py | app/agent/nodes/merge_retrieved_info.py |
| 🟡 P1 | llm.py | app/agent/llm.py |
| 🟡 P1 | prompt_loader.py | app/prompt/prompt_loader.py |
| 🟢 P2 | extend_keywords_for_*.prompt (3 个) | prompts/ |