SQL生成信息过滤

之前将三路召回的信息整理成两份核心上下文:

table_infos # 按表组织好的表结构上下文

metric_infos # 整理后的指标上下文

但这里的“整理好”,还不等于“足够精确”。

本章要实现的三个节点,就是在 SQL 生成前做最后一轮准备:

  1. filter_table # 过滤表和字段  把业务上下文筛干净
  2. filter_metric # 过滤指标   把业务上下文筛干净
  3. add_extra_context # 添加日期、数据库等额外上下文  把运行环境信息补完整

把它们放回完整链路中看,位置如下:

用户问题
  -> extract_keywords
  -> recall_column / recall_metric / recall_value
  -> merge_retrieved_info
  -> filter_table / filter_metric
  -> add_extra_context
  -> generate_sql

image

过滤不是重新检索,而是在召回和合并之后,把“可能相关”的上下文压缩成“本次 SQL 真正需要”的上下文。

节点处理对象输出结果作用
filter_table table_infos 过滤后的表和字段 保留本次查询真正需要的 schema
filter_metric metric_infos 过滤后的指标 保留本次查询真正需要的业务指标
add_extra_context 当前日期、数据库连接 date_infodb_info 补齐相对时间和 SQL 方言信息

为什么要把上下文转为YAML

table_infos 和 metric_infos 在 Python 程序里是列表、字典等对象。它们不能直接“作为对象”交给大模型,只能先转成文本.

常见做法:

  1. Python 对象 -> JSON 字符串
  2. Python 对象 -> YAML 字符串

本项目里选择 YAML:  yaml.dump(table_infos, allow_unicode=True, sort_keys=False)

这样做有几个好处:

  • 层级结构清楚,表、字段、字段属性更容易看;
  • allow_unicode=True 可以让中文正常显示,而不是转成 Unicode 编码;
  • sort_keys=False 可以保留原有字段顺序,不会按字母顺序打乱;
  • 放进提示词后,比 Python 对象的默认打印形式更适合模型阅读。

表过滤节点- filter_table

filter_table 读取的是上一章生成的:state["table_infos"]

输入给大模型的大致内容包括:

  • 用户问题 query
  • 候选表及字段信息 table_infos

模型需要返回一个简单的 JSON 对象,表示“保留哪些表,以及每张表里保留哪些字段”。这里特意让模型返回“选择结果”,而不是完整的过滤后表结构。原因很实际:完整的 table_infos 层级比较深,有表名、

示例值、别名等信息。让模型原样重写这整套结构,出错概率会更高。

更稳的做法是:让模型只做选择题,返回表名和字段名;真正的裁剪由程序完成。

1. filter_table提示词注意什么

提示词的定位是:让模型扮演查询规划专家,在候选表和字段中选出回答当前问题必须使用的部分。

核心规则可以概括成几条:

  • 只能从候选表和候选字段中选择;
  • 不能新增表、不能新增字段、不能修改字段名;
  • 每张保留下来的表,至少要有一个字段被选中;
  • 字段是否保留,以“本次查询是否实际使用”为标准;
  • 如果选择多张表,要保留 join 所需的主外键字段;
  • 只输出 JSON 对象,不输出解释文字。

2. filter_table核心代码

import yaml
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, TableInfoState
from app.core.log import logger
from app.prompt.prompt_loader import load_prompt


async def filter_table(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    """根据用户问题裁剪候选表结构上下文"""

    writer = runtime.stream_writer
    writer("过滤表信息")

    query = state["query"]
    table_infos: list[TableInfoState] = state["table_infos"]

    # table_infos 是嵌套结构,转成 YAML 后更适合放进提示词,也保留中文字段说明
    prompt = PromptTemplate(
        template=load_prompt("filter_table_info"),
        input_variables=["query", "table_infos"],
    )
    # filter_table_info prompt 要求模型只输出 JSON 对象:表名 -> 字段名列表
    output_parser = JsonOutputParser()
    # LCEL 管道:填充提示词 -> 调用模型 -> 解析 JSON
    chain = prompt | llm | output_parser

    result = await chain.ainvoke(
        {
            "query": query,
            "table_infos": yaml.dump(table_infos, allow_unicode=True, sort_keys=False),   # 把 table_infos 序列化成 YAML 后传给模型
        }
    )
    # 模型只负责选择,程序根据选择结果从原始 TableInfoState 中裁剪,避免模型重写复杂结构出错
    filtered_table_infos: list[TableInfoState] = []
    # 根据模型返回的字典裁剪原始表结构。
    # 这里没有在遍历 table_infos 时直接 remove 元素。因为遍历列表时删除元素,容易造成索引错位,导致漏处理。更稳的方式是新建一个列表
    # 需要保留的表,再追加进去。
    for table_info in table_infos:
        if table_info["name"] in result:
            table_info["columns"] = [
                column_info
                # 字段过滤则使用列表推导式
                for column_info in table_info["columns"]
                if column_info["name"] in result[table_info["name"]]
            ]
            filtered_table_infos.append(table_info)

    logger.info(
        f"过滤后的表信息:{[filtered_table_info['name'] for filtered_table_info in filtered_table_infos]}"
    )
    return {"table_infos": filtered_table_infos}

 3. filter_metric指标过滤节点

 metric_infos 里可能有多个候选指标,这些指标可能都和用户问题有一点关系,但不一定都需要进入最终 SQL 生成

所以 filter_metric 会让大模型根据用户问题,从候选指标中选择真正需要的指标

4. filter_metric提示词注意什么

  • 只能从候选指标中选择;
  • 不能新增指标,不能修改指标含义;
  • 只有本次问题确实用于度量、统计、对比、计算结果的指标才保留;
  • 仅用于筛选、分组或限定范围的字段,不视为指标;
  • 如果候选指标都不适合,返回空数组;
  • 只输出 JSON 数组,不输出解释文字。

5. filter_metric核心代码

import yaml
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, MetricInfoState
from app.core.log import logger
from app.prompt.prompt_loader import load_prompt


async def filter_metric(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    """根据用户问题裁剪候选指标上下文"""

    writer = runtime.stream_writer
    writer("过滤指标信息")

    query = state["query"]
    metric_infos: list[MetricInfoState] = state["metric_infos"]

    # metric_infos 转成 YAML 后作为候选项交给模型,模型只需要返回被选中的指标名称
    prompt = PromptTemplate(
        template=load_prompt("filter_metric_info"),
        input_variables=["query", "metric_infos"],
    )
    # filter_metric_info prompt 要求模型只输出 JSON 数组
    output_parser = JsonOutputParser()
    # LCEL 管道:填充提示词 -> 调用模型 -> 解析 JSON
    chain = prompt | llm | output_parser

    result = await chain.ainvoke(
        {
            "query": query,
            "metric_infos": yaml.dump(
                metric_infos, allow_unicode=True, sort_keys=False
            ),
        }
    )
    # 用模型返回的指标名称过滤原始结构,保留描述 依赖字段 别名等完整上下文
    filtered_metric_infos = [
        metric_info
        for metric_info in metric_infos
        if metric_info["name"] in result
    ]

    logger.info(
        f"过滤后的指标信息:{[filtered_metric_info['name'] for filtered_metric_info in filtered_metric_infos]}"
    )
    return {"metric_infos": filtered_metric_infos}

 6. 过滤节点在图里的并行关系

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")

 7. 额外上下文节点add_extra_context

表和指标过滤完成后,业务上下文已经比较干净了。但 SQL 生成还需要一些不属于表结构的外部信息。

例如用户问:本季度,上个月,模型要把它们转换成 SQL 条件,就必须知道当前日期。所以 add_extra_context 会补两类信息:

date_info # 当前日期、星期、季                                    db_info # 数据库方言、数据库版本

8. 补齐State和Context

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


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


class DataAgentState(TypedDict):
    # 前面章节已有字段省略...

    date_info: DateInfoState  # 当前日期 星期和季度信息
    db_info: DBInfoState  # 数据库方言和版本信息

DateInfoState 和 DBInfoState 定义了两份上下文的结构。DataAgentState 里加入 date_infodb_info 后,后续的 generate_sql 节点就可以稳定读取它们

add_extra_context 还需要访问数仓数据库,查询数据库版本和方言,所以 DataAgentContext 中要有 dw_mysql_repository

这里使用的是 DWMySQLRepository,不是 MetaMySQLRepository

原因是:

  • Meta MySQL 保存表、字段、指标等元数据;
  • DW MySQL 模拟真实数仓,SQL 最终也会在这里校验和执行;
  • 数据库方言和版本应该以真实执行 SQL 的数据库为准。

9. add_extra_context核心代码

from datetime import date

from langgraph.runtime import Runtime

from app.agent.context import DataAgentContext
from app.agent.state import DataAgentState, DateInfoState, DBInfoState
from app.core.log import logger


async def add_extra_context(state: DataAgentState, runtime: Runtime[DataAgentContext]):
    """补齐 SQL 生成所需的日期和数据库环境信息"""

    dw_mysql_repository = runtime.context["dw_mysql_repository"]

    # 当前日期信息会帮助模型处理“今天 本月 本季度 最近 N 天”等相对时间表达
    today = date.today()
    date_str = today.strftime("%Y-%m-%d")
    weekday = today.strftime("%A")
    quarter = f"Q{(today.month - 1) // 3 + 1}"
    date_info = DateInfoState(date=date_str, weekday=weekday, quarter=quarter)

    # 数据库方言和版本会影响函数名 日期运算 limit 语法等 SQL 细节
    db = await dw_mysql_repository.get_db_info()
    db_info = DBInfoState(**db)
    logger.info(f"数据库信息:{db_info}")
    logger.info(f"日期信息:{date_info}")
    return {"date_info": date_info, "db_info": db_info}

这段代码做两件事:

  1. 生成当前日期信息
    1. date_str 是日期字符串
    2. weekday 是星期几
    3. quarter 是季度
  2. 查询数据库信息 
    1. get_db_info() 返回的是一个字典:{"dialect": "mysql", "version": "8.0.x"}
    2. 字段名刚好和 DBInfoState 对应,所以可以用 **db 解包创建

10. DWMySQLRepository如何获取数据库信息

from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession


class DWMySQLRepository:
    """负责查询数仓真实表结构和字段样例值"""

    def __init__(self, session: AsyncSession):
        self.session = session

    async def get_db_info(self):
        """读取当前数仓数据库的方言和版本,供 SQL 生成提示词使用"""

        sql = "select version()"
        result = await self.session.execute(text(sql))
        version = result.scalar()

        # dialect 来自 SQLAlchemy 当前绑定的数据库方言,例如 mysql
        dialect = self.session.bind.dialect.name
        return {"dialect": dialect, "version": version}

这里获取了两个信息。

  1. 数据库版本
    1. select version()这条 SQL 在 MySQL 中会返回当前数据库版本。因为结果只有一行一列,所以代码里直接使用version = result.scalar()
  2. 数据库方言
    1. 当前项目使用 SQLAlchemy 管理数据库连接,dialect.name 可以拿到当前连接对应的数据库类型

11. 在graph.py中传入DW repository

因为 add_extra_context 会读取runtime.context["dw_mysql_repository"].

所以本地测试工作流时,需要初始化 DW 数据库连接,并把 DWMySQLRepository 放进 context

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,
):
    meta_mysql_repository = MetaMySQLRepository(meta_session)
    dw_mysql_repository = DWMySQLRepository(dw_session)

    context = DataAgentContext(
        column_qdrant_repository=column_qdrant_repository,
        embedding_client=embedding_client_manager.client,
        metric_qdrant_repository=metric_qdrant_repository,
        value_es_repository=value_es_repository,
        meta_mysql_repository=meta_mysql_repository,
        dw_mysql_repository=dw_mysql_repository,
    )

这里同时创建了 meta_mysql_repository 和 dw_mysql_repository

  • 合并召回信息需要查元数据库,所以用 meta_mysql_repository
  • 添加额外上下文、后续校验和执行 SQL 需要面向数仓,所以用 dw_mysql_repository

几个注意点

1. 不要边遍历边删除元素

过滤表时,不建议一边遍历 table_infos,一边调用 remove 删除元素。

更稳的做法是新建一个 filtered_table_infos,把需要保留的表追加进去。字段列表也可以通过列表推导式生成一个新的列表

2. 指标过滤允许返回空值

不是所有问题都需要业务指标,比如查询华北有哪些分店。这个问题更像明细查询或对象查询,不一定要用 GMV订单数 这类指标。允许 filter_metric 返回 []

 

posted @ 2026-06-01 16:49  幻影之舞  阅读(13)  评论(0)    收藏  举报