SQL的生成与执行闭环
SQL生成前有那些信息
query # 用户原始问题
table_infos # 可使用的表、字段、字段类型、字段描述、示例值
metric_infos # 可参考的业务指标、指标口径、依赖字段
date_info # 当前日期、星期、季度
db_info # 当前数据库方言和版本
如果缺少这些上下文,模型也能写出一条“看起来像 SQL”的语句,但那条 SQL 很可能用错表、用错字段、写错指标口径,甚至不符合当前数据库版本。所以本章的重点不是让模型自由发挥,而是让模型在前面整理好的约束范围内生成 SQL
生成SQL-generate_sql
llm结合上述信息,生成sql.
其中最关键的约束是:输出必须只包含一条完整 SQL 语句的纯文本.需要注意的是,很多大模型在写代码时,会习惯性输出 Markdown 代码块.
但项目拿到 SQL 后要直接交给数据库校验和执行。如果 SQL 字符串里混入 Markdown 代码块,数据库并不认识这些符号,就会报语法错误
所以提示词必须明确要求:
- 不要解释
- 不要 Markdown 代码块
- 不要 ```sql
- 只输出 SQL 本身
核心代码
async def generate_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer("生成SQL")
# 这些上下文都由前置节点准备完成,模型只在给定表、字段、指标口径范围内生成 SQL
table_infos = state["table_infos"]
metric_infos = state["metric_infos"]
date_info = state["date_info"]
db_info = state["db_info"]
query = state["query"]
prompt = PromptTemplate(
template=load_prompt("generate_sql"),
input_variables=["table_infos", "metric_infos", "date_info", "db_info", "query"],
)
# 生成 SQL 只需要一段纯文本,所以这里使用 StrOutputParser
output_parser = StrOutputParser()
chain = prompt | llm | output_parser
result = await chain.ainvoke(
{
# YAML 更适合放进提示词:保留嵌套结构、顺序和中文说明,方便模型理解表字段关系
"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),
"query": query,
}
)
logger.info(f"生成的SQL:{result}")
return {"sql": result}
两个点值得注意:
- 结构化上下文会先转成 YAML。
- yaml.dump(table_infos, allow_unicode=True, sort_keys=False)
- 输出解析器使用的是
StrOutputParser,不是JsonOutputParser- 因为这一节点只需要一条 SQL 字符串,不需要 JSON 对象
校验SQL: validate_sql
生成 SQL 之后,不能直接执行,大模型也仍然可能出错。常见错误包括:
- 表名写错;
- 字段名写错;
- 字段别名写错;
join条件不完整;- 聚合函数和
group by不匹配; - SQL 方言不符合当前数据库;
- 使用了提示词里没有提供的表或字段。
explain <generated_sql>
EXPLAIN 原本是用来看 SQL 执行计划的,但这里主要利用它的副作用:让数据库提前解析 SQL,若 SQL 里有字段不存在,数据库会直接报错。
这条错误信息非常有用。后面 correct_sql 可以把它连同原 SQL 一起交给大模型,让模型做有依据的修正
validate_sql核心代码
async def validate_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer("校验SQL")
# 读取 generate_sql 写入状态的 SQL。
sql = state["sql"]
# SQL 可用性必须交给真实数仓判断,这里从运行时 context 中取 DW Repository。
dw_mysql_repository: DWMySQLRepository = runtime.context["dw_mysql_repository"]
try:
# validate 内部使用 explain <sql>,只关心数据库能否成功解析这条 SQL。
await dw_mysql_repository.validate(sql)
logger.info("SQL语法正确")
return {"error": None}
except Exception as e:
# 不直接抛异常,而是把错误信息写入 state,交给条件边判断。
logger.info(f"SQL语法错误:{str(e)}")
return {"error": str(e)}
校验失败时,节点没有直接抛异常,而是把错误转成字符串,写入状态,return {"error": str(e)},这是为了交给 LangGraph 的条件分支处理
也就是说,validate_sql 只负责判断 SQL 是否有问题,不负责决定下一步去哪。下一步走 run_sql 还是 correct_sql,由图结构来决定
用条件分支决定校正还是执行

这就是本章的核心闭环:不是一次生成就盲目执行,而是先校验;校验失败就把错误信息带回模型,让它基于真实错误修正 SQL
graph_builder.add_edge("generate_sql", "validate_sql")
graph_builder.add_conditional_edges(
source="validate_sql",
path=lambda state: "run_sql" if state["error"] is None else "correct_sql",
path_map={"run_sql": "run_sql", "correct_sql": "correct_sql"},
)
graph_builder.add_edge("correct_sql", "run_sql")
graph_builder.add_edge("run_sql", END)
当前是用的简化版;生成 SQL -> 校验一次 -> 如果失败则校正一次 -> 执行
而真实项目中应该多次循环: 生成 SQL -> 校验 -> 校正 -> 再校验 -> 最多重试 N 次 -> 执行或返回失败
校正SQL: correct_sql
在 SQL 校验失败时,根据错误信息修正 SQL。它不是重新生成一条完全不同的 SQL,而是在尽量保持原业务语义不变的前提下,修复导致 SQL 无法执行的问题
校正提示词关注什么
- 必须严格基于 SQL 执行错误信息进行修正;
- 不得改变原始业务语义;
- 不得新增、删除或替换统计指标、统计维度、过滤条件或时间范围;
- 只能使用上下文中真实存在的表和字段;
- 如果涉及指标计算,要继续遵循指标信息中的业务口径;
- 修正后的 SQL 仍然只能是一条查询 SQL;
- 输出仍然必须是纯文本 SQL,不能带 Markdown 代码块。
核心代码
async def correct_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer("校正SQL")
# 校正 SQL 仍然需要完整上下文,避免模型只根据报错修语法却改丢业务语义
table_infos = state["table_infos"]
metric_infos = state["metric_infos"]
date_info = state["date_info"]
db_info = state["db_info"]
query = state["query"]
# sql 是待修正的候选 SQL,error 是数据库 explain 返回的具体错误信息
sql = state["sql"]
error = state["error"]
prompt = PromptTemplate(
template=load_prompt("correct_sql"),
input_variables=[
"table_infos",
"metric_infos",
"date_info",
"db_info",
"query",
"sql",
"error",
],
)
output_parser = StrOutputParser()
chain = prompt | llm | output_parser
result = await chain.ainvoke(
{
# 与生成节点保持一致,用 YAML 向模型提供稳定、可读的结构化上下文
"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),
"query": query,
"sql": sql,
"error": error,
}
)
logger.info(f"校正后的SQL:{result}")
return {"sql": result}
执行SQL: run_sql
async def run_sql(state: DataAgentState, runtime: Runtime[DataAgentContext]):
writer = runtime.stream_writer
writer("执行SQL")
# 这里拿到的是 generate_sql 生成的 SQL,
# 或 correct_sql 修正后覆盖进去的 SQL。
sql = state["sql"]
dw_mysql_repository = runtime.context["dw_mysql_repository"]
result = await dw_mysql_repository.run(sql)
logger.info(f"SQL执行结果:{result}")
这段代码很短,因为真正的数据库访问被封装到了 DWMySQLRepository。
当前实现里 run_sql 会执行 SQL 并把结果写入日志,但还没有把 result 返回到 DataAgentState,所以本章本地测试时主要通过日志确认查询结果。后面进入 API 封装时,如果要把最终结果返回给前端,就需要在接口层或执行节点里继续补充结果输出
DWMySQLRepository.run
async def run(self, sql: str) -> list[dict]:
# 执行 SQL
result = await self.session.execute(text(sql))
# 把查询结果转成字典列表
return [dict(row) for row in result.mappings().fetchall()]
补齐字段状态
新增了两个关键状态字段:
| 字段 | 谁写入 | 谁读取 | 作用 |
|---|---|---|---|
sql |
generate_sql、correct_sql |
validate_sql、correct_sql、run_sql |
保存候选或修正后的 SQL |
error |
validate_sql |
graph 条件分支、correct_sql |
保存 SQL 校验错误 |
可以按下面的流向理解:
generate_sql 写入 sql
-> validate_sql 读取 sql,并写入 error
-> graph 根据 error 判断分支
-> correct_sql 读取 sql 和 error,并覆盖 sql
-> run_sql 读取最终 sql

浙公网安备 33010602011771号