SQL的生成与执行闭环

SQL生成前有那些信息

query         # 用户原始问题
table_infos   # 可使用的表、字段、字段类型、字段描述、示例值
metric_infos  # 可参考的业务指标、指标口径、依赖字段
date_info     # 当前日期、星期、季度
db_info       # 当前数据库方言和版本

如果缺少这些上下文,模型也能写出一条“看起来像 SQL”的语句,但那条 SQL 很可能用错表、用错字段、写错指标口径,甚至不符合当前数据库版本。所以本章的重点不是让模型自由发挥,而是让模型在前面整理好的约束范围内生成 SQL

生成SQL-generate_sql

llm结合上述信息,生成sql.

其中最关键的约束是:输出必须只包含一条完整 SQL 语句的纯文本.需要注意的是,很多大模型在写代码时,会习惯性输出 Markdown 代码块.

但项目拿到 SQL 后要直接交给数据库校验和执行。如果 SQL 字符串里混入 Markdown 代码块,数据库并不认识这些符号,就会报语法错误

所以提示词必须明确要求:

  1. 不要解释
  2. 不要 Markdown 代码块
  3. 不要 ```sql
  4. 只输出 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}

两个点值得注意:

  1. 结构化上下文会先转成 YAML。
    1. yaml.dump(table_infos, allow_unicode=True, sort_keys=False) 
  2. 输出解析器使用的是 StrOutputParser,不是 JsonOutputParser
    1. 因为这一节点只需要一条 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,由图结构来决定

用条件分支决定校正还是执行

image

这就是本章的核心闭环:不是一次生成就盲目执行,而是先校验;校验失败就把错误信息带回模型,让它基于真实错误修正 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_sqlcorrect_sql validate_sqlcorrect_sqlrun_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

 

posted @ 2026-06-04 14:22  幻影之舞  阅读(20)  评论(0)    收藏  举报