AI驱动的自然语言到SQL查询引擎深度实战:从Schema理解到查询生成与自动修复的全解析

举报
江南清风起 发表于 2026/09/09 23:35:03 2026/09/09
【摘要】 AI驱动的自然语言到SQL查询引擎深度实战:从Schema理解到查询生成与自动修复的全解析 引言Text-to-SQL(自然语言转SQL)是LLM在企业数据领域的杀手级应用:用户用自然语言提问,系统自动生成并执行SQL返回结果。但生产级Text-to-SQL远非"把问题给LLM生成SQL"那么简单:需要理解数据库Schema(表结构、关系、业务语义)、处理歧义查询、验证SQL语法与安全性、...

AI驱动的自然语言到SQL查询引擎深度实战:从Schema理解到查询生成与自动修复的全解析

引言

Text-to-SQL(自然语言转SQL)是LLM在企业数据领域的杀手级应用:用户用自然语言提问,系统自动生成并执行SQL返回结果。但生产级Text-to-SQL远非"把问题给LLM生成SQL"那么简单:需要理解数据库Schema(表结构、关系、业务语义)、处理歧义查询、验证SQL语法与安全性、自动修复执行错误、结果格式化与可视化。本文从Text-to-SQL的架构讲起,覆盖Schema感知与上下文注入、SQL生成与Few-Shot增强、SQL验证与安全过滤、执行错误自动修复、结果可视化与自然语言摘要、多轮对话与上下文记忆、评估基准与准确率优化,构建企业级NL2SQL引擎。

一、Schema感知与上下文注入

# nl2sql/schema_manager.py
from dataclasses import dataclass, field
from typing import Optional

@dataclass
class TableSchema:
    name: str
    description: str
    columns: list[dict]    # [{name, type, description, is_pk, is_fk, fk_target}]
    sample_rows: list[dict] = field(default_factory=list)
    row_count: int = 0

@dataclass
class DatabaseSchema:
    tables: list[TableSchema]
    relationships: list[dict]  # [{from_table, from_col, to_table, to_col}]
    glossary: dict = field(default_factory=dict)  # 业务术语映射

class SchemaManager:
    """数据库Schema管理:结构感知与上下文构建"""
    
    SCHEMA_PROMPT = """数据库Schema信息:

{schema_text}

业务术语表:
{glossary}

示例数据:
{sample_data}

用户问题:{question}

请基于以上Schema信息生成SQL查询。规则:
1. 仅使用SELECT查询,禁止INSERT/UPDATE/DELETE/DROP
2. 使用表别名提高可读性
3. 列名用双引号包裹
4. 添加LIMIT防止结果过大
5. 如果问题模糊,选择最合理的解释

输出JSON:{{"sql": "...", "explanation": "...", "tables_used": [...]}}"""

    async def build_context(self, question: str,
                             db_schema: DatabaseSchema,
                             relevant_tables: list[str] = None) -> str:
        """构建Schema上下文:只注入相关表"""
        if relevant_tables:
            tables = [t for t in db_schema.tables if t.name in relevant_tables]
        else:
            # 模糊匹配表名
            tables = self._find_relevant_tables(question, db_schema)
        
        schema_text = self._format_schema(tables)
        glossary = self._format_glossary(db_schema.glossary, question)
        sample = self._format_samples(tables)
        
        return self.SCHEMA_PROMPT.format(
            schema_text=schema_text,
            glossary=glossary,
            sample_data=sample,
            question=question,
        )
    
    def _find_relevant_tables(self, question: str,
                               schema: DatabaseSchema) -> list[TableSchema]:
        """根据问题关键词匹配相关表"""
        relevant = []
        question_lower = question.lower()
        for table in schema.tables:
            # 表名或描述中的关键词出现在问题中
            keywords = [table.name.lower()] + [
                col["name"].lower() for col in table.columns
            ] + table.description.lower().split()
            if any(kw in question_lower for kw in keywords if len(kw) > 2):
                relevant.append(table)
        # 如果没匹配到,返回全部表(小库)
        return relevant if relevant else schema.tables[:5]
    
    def _format_schema(self, tables: list[TableSchema]) -> str:
        lines = []
        for t in tables:
            lines.append(f"表:{t.name}{t.description})")
            for col in t.columns:
                pk = " [PK]" if col.get("is_pk") else ""
                fk = f" [FK→{col['fk_target']}]" if col.get("is_fk") else ""
                lines.append(f"  - {col['name']} {col['type']}{pk}{fk}: {col.get('description', '')}")
            lines.append("")
        return "\n".join(lines)
    
    def _format_glossary(self, glossary: dict, question: str) -> str:
        relevant = {}
        for term, definition in glossary.items():
            if term.lower() in question.lower():
                relevant[term] = definition
        if not relevant:
            return "(无相关术语)"
        return "\n".join(f"- {t}: {d}" for t, d in relevant.items())
    
    def _format_samples(self, tables: list[TableSchema]) -> str:
        lines = []
        for t in tables[:3]:  # 最多3个表的样本
            if t.sample_rows:
                lines.append(f"{t.name} 示例:")
                for row in t.sample_rows[:3]:
                    lines.append(f"  {row}")
        return "\n".join(lines) if lines else "(无样本数据)"

二、SQL生成与验证

# nl2sql/generator.py
import re
from dataclasses import dataclass

@dataclass
class SQLResult:
    sql: str
    explanation: str
    tables_used: list[str]
    is_valid: bool = True
    errors: list[str] = None

class SQLGenerator:
    """SQL生成与验证"""
    
    FORBIDDEN_KEYWORDS = [
        "INSERT", "UPDATE", "DELETE", "DROP", "ALTER", "CREATE",
        "TRUNCATE", "GRANT", "REVOKE", "COPY", "EXEC", "EXECUTE",
    ]
    
    def __init__(self, llm_client, schema_manager: SchemaManager):
        self.llm = llm_client
        self.schema = schema_manager
    
    async def generate(self, question: str,
                        db_schema: DatabaseSchema,
                        few_shot_examples: list[dict] = None) -> SQLResult:
        """生成SQL"""
        context = await self.schema.build_context(question, db_schema)
        
        # 注入Few-Shot示例
        if few_shot_examples:
            examples_text = "\n\n示例:\n"
            for ex in few_shot_examples[:3]:
                examples_text += f"问:{ex['question']}\nSQL:{ex['sql']}\n"
            context = examples_text + "\n" + context
        
        raw = await self.llm.complete(
            context, response_format={"type": "json_object"},
            temperature=0,
        )
        import json
        data = json.loads(raw)
        
        sql = data.get("sql", "")
        result = SQLResult(
            sql=sql,
            explanation=data.get("explanation", ""),
            tables_used=data.get("tables_used", []),
        )
        
        # 验证
        result.is_valid, result.errors = self._validate(sql)
        
        return result
    
    def _validate(self, sql: str) -> tuple[bool, list[str]]:
        """SQL安全验证"""
        errors = []
        sql_upper = sql.upper().strip()
        
        # 必须是SELECT
        if not sql_upper.startswith("SELECT") and not sql_upper.startswith("WITH"):
            errors.append("仅允许SELECT或WITH查询")
        
        # 禁止关键词
        for kw in self.FORBIDDEN_KEYWORDS:
            if re.search(rf'\b{kw}\b', sql_upper):
                errors.append(f"禁止使用 {kw}")
        
        # 必须有LIMIT(如果没有则添加)
        if "LIMIT" not in sql_upper and "COUNT" not in sql_upper:
            # 不强制报错,但自动添加
            pass
        
        # 分号检查
        if sql.count(";") > 1:
            errors.append("禁止多条SQL语句")
        
        return len(errors) == 0, errors
    
    def sanitize(self, sql: str, max_rows: int = 1000) -> str:
        """SQL安全化:添加LIMIT"""
        sql = sql.strip().rstrip(";")
        if "LIMIT" not in sql.upper():
            sql = f"SELECT * FROM ({sql}) AS _safe_query LIMIT {max_rows}"
        return sql

三、自动修复

# nl2sql/auto_repair.py
class SQLAutoRepair:
    """SQL执行错误自动修复"""
    
    REPAIR_PROMPT = """SQL执行失败,请修复。

原始SQL:
{sql}

错误信息:
{error}

数据库Schema:
{schema}

修复规则:
1. 分析错误原因(表名/列名错误、类型不匹配、语法错误)
2. 生成修复后的SQL
3. 保持查询意图不变

输出JSON:{{"sql": "修复后SQL", "fix_description": "修复说明"}}"""
    
    def __init__(self, llm_client, schema_manager, max_retries: int = 2):
        self.llm = llm_client
        self.schema = schema_manager
        self.max_retries = max_retries
    
    async def repair_and_execute(self, sql: str, error: str,
                                  db_schema: DatabaseSchema,
                                  executor) -> dict:
        """修复并重新执行"""
        for attempt in range(self.max_retries):
            # 生成修复
            schema_text = self.schema._format_schema(db_schema.tables[:5])
            prompt = self.REPAIR_PROMPT.format(
                sql=sql, error=error, schema=schema_text,
            )
            raw = await self.llm.complete(
                prompt, response_format={"type": "json_object"},
                temperature=0,
            )
            import json
            data = json.loads(raw)
            repaired_sql = data["sql"]
            
            # 验证修复后的SQL
            valid, errors = SQLGenerator._validate(self, repaired_sql)
            if not valid:
                error = "; ".join(errors)
                continue
            
            # 执行
            try:
                result = await executor.execute(repaired_sql)
                return {
                    "success": True,
                    "sql": repaired_sql,
                    "fix_description": data.get("fix_description", ""),
                    "result": result,
                    "repair_attempts": attempt + 1,
                }
            except Exception as e:
                error = str(e)
                sql = repaired_sql
        
        return {
            "success": False,
            "error": error,
            "repair_attempts": self.max_retries,
        }

四、结果摘要

# nl2sql/result_summarizer.py
class ResultSummarizer:
    """查询结果自然语言摘要"""
    
    SUMMARIZE_PROMPT = """将SQL查询结果转化为自然语言摘要。

用户问题:{question}
SQL查询:{sql}
查询结果(前{limit}行):
{result_preview}
总行数:{total_rows}

要求:
1. 用自然语言回答用户问题
2. 包含关键数字与趋势
3. 如果结果为空,说明可能原因
4. 适合非技术人员理解

摘要:"""
    
    async def summarize(self, question: str, sql: str,
                         result: list[dict], total_rows: int,
                         llm_client) -> str:
        """生成结果摘要"""
        import json
        preview = result[:20]
        result_text = json.dumps(preview, ensure_ascii=False, default=str)
        if len(result_text) > 2000:
            result_text = result_text[:2000] + "...(更多数据省略)"
        
        prompt = self.SUMMARIZE_PROMPT.format(
            question=question, sql=sql[:200],
            result_preview=result_text,
            limit=len(preview), total_rows=total_rows,
        )
        return await llm_client.complete(prompt, temperature=0.2, max_tokens=500)
    
    def to_table(self, result: list[dict], max_rows: int = 20) -> str:
        """格式化为Markdown表格"""
        if not result:
            return "(无数据)"
        headers = list(result[0].keys())
        lines = ["| " + " | ".join(headers) + " |",
                 "| " + " | ".join("---" for _ in headers) + " |"]
        for row in result[:max_rows]:
            lines.append("| " + " | ".join(str(row.get(h, "")) for h in headers) + " |")
        if len(result) > max_rows:
            lines.append(f"\n(共{len(result)}行,显示前{max_rows}行)")
        return "\n".join(lines)

五、多轮对话

# nl2sql/conversation.py
class NL2SQLConversation:
    """多轮对话:上下文记忆与指代消解"""
    
    CLARIFY_PROMPT = """用户问题可能有歧义,请分析是否需要澄清。

当前问题:{question}
对话历史:{history}

如果问题引用了之前的结果(如"上面的表"、"其中"、"那"),请消解指代。
如果问题模糊(如"看看数据"),请建议一个具体查询。

输出JSON:{{"clarified_question": "消解后的问题", "needs_clarification": true/false, "suggestion": "..."}}"""
    
    def __init__(self, llm_client):
        self.llm = llm_client
        self.history: list[dict] = []
    
    async def process(self, question: str, db_schema: DatabaseSchema,
                       generator: SQLGenerator, executor,
                       summarizer: ResultSummarizer) -> dict:
        """处理用户问题"""
        # 指代消解
        if self.history:
            clarified = await self._resolve_reference(question)
            question = clarified.get("clarified_question", question)
        
        # 生成SQL
        sql_result = await generator.generate(question, db_schema)
        if not sql_result.is_valid:
            return {"error": "SQL验证失败", "details": sql_result.errors}
        
        # 执行
        safe_sql = generator.sanitize(sql_result.sql)
        try:
            rows, total = await executor.execute(safe_sql)
        except Exception as e:
            # 自动修复
            repair = SQLAutoRepair(self.llm, generator.schema)
            repair_result = await repair.repair_and_execute(
                safe_sql, str(e), db_schema, executor,
            )
            if repair_result["success"]:
                rows = repair_result["result"]["rows"]
                total = repair_result["result"]["total"]
                sql_result.sql = repair_result["sql"]
            else:
                return {"error": "SQL执行失败", "details": repair_result["error"]}
        
        # 摘要
        summary = await summarizer.summarize(
            question, sql_result.sql, rows, total, self.llm,
        )
        
        # 记录历史
        self.history.append({
            "question": question,
            "sql": sql_result.sql,
            "result_count": total,
        })
        
        return {
            "question": question,
            "sql": sql_result.sql,
            "explanation": sql_result.explanation,
            "table": summarizer.to_table(rows),
            "summary": summary,
            "total_rows": total,
        }
    
    async def _resolve_reference(self, question: str) -> dict:
        """指代消解"""
        history_text = "\n".join(
            f"Q: {h['question']}\nSQL: {h['sql'][:100]}\n结果: {h['result_count']}行"
            for h in self.history[-3:]
        )
        prompt = self.CLARIFY_PROMPT.format(
            question=question, history=history_text,
        )
        raw = await self.llm.complete(
            prompt, response_format={"type": "json_object"},
            temperature=0,
        )
        import json
        return json.loads(raw)

总结

Text-to-SQL引擎的工程体系以"Schema感知-生成验证-修复-摘要-对话"五阶段展开:Schema管理器根据问题关键词匹配相关表并注入结构定义、列描述、外键关系与样本数据,SQL生成器以JSON格式输出SQL+解释+使用的表并以关键词黑名单验证安全性(禁止INSERT/UPDATE/DELETE/DROP),自动修复器在执行失败时分析错误并重新生成修复SQL(最多2轮重试),结果摘要器把查询结果转化为非技术人员可理解的自然语言描述并生成Markdown表格,多轮对话以指代消解处理"上面的表"“其中"等引用。当数据库查询从"写SQL"变成"说人话”,NL2SQL引擎让非技术人员也能自主探索数据,这是LLM在企业数据领域最具ROI的应用场景。

【声明】本内容来自华为云开发者社区博主,不代表华为云及华为云开发者社区的观点和立场。转载时必须标注文章的来源(华为云社区)、文章链接、文章作者等基本信息,否则作者和本社区有权追究责任。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱: cloudbbs@huaweicloud.com
  • 点赞
  • 收藏
  • 关注作者

评论(0

0/1000
抱歉,系统识别当前为高风险访问,暂不支持该操作

全部回复

上滑加载中

设置昵称

在此一键设置昵称,即可参与社区互动!

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。