LangChain-Text-to-SQL-让LLM直接查数据库
企业中大量有价值的数据存在关系型数据库里,但能写 SQL 的人远比有业务问题的人少。产品经理想知道"上周各地区的退款率是多少",需要找数据分析师,等排期,写查询,出报告——整个链路可能要几天。Text-to-SQL 技术让 LLM 充当 SQL 翻译器,用户用自然语言提问,系统自动生成并执行 SQL,返回自然语言结果。
LangChain Text-to-SQL:让 LLM 直接查数据库
企业中大量有价值的数据存在关系型数据库里,但能写 SQL 的人远比有业务问题的人少。产品经理想知道"上周各地区的退款率是多少",需要找数据分析师,等排期,写查询,出报告——整个链路可能要几天。Text-to-SQL 技术让 LLM 充当 SQL 翻译器,用户用自然语言提问,系统自动生成并执行 SQL,返回自然语言结果。
1.1 企业场景与技术挑战
Text-to-SQL完整流程——从自然语言问题到LLM理解Schema、生成SQL、执行、格式化结果输出
Text-to-SQL(文字转 SQL:用自然语言提问,系统自动将问题翻译成 SQL 查询语句并执行)的典型应用场景:
- 业务数据自助查询:运营人员直接提问,无需等待数据团队
- 报表系统自然语言入口:在 BI 工具(Business Intelligence,商业智能工具:用于数据分析和可视化的软件,如 Tableau、Power BI)上叠加自然语言查询层
- 客服数据查询:客服系统中直接查询用户订单、账户信息
- 内部数据探查工具:开发人员快速探索不熟悉的数据库
主要技术挑战:
| 挑战 | 具体问题 | 解决思路 |
|---|---|---|
| Schema 理解(Schema:数据库的结构定义,描述有哪些表、每张表有哪些字段) | 表名/列名不直观,如 ord_dtl(订单明细) |
提供 Schema 描述和注释 |
| 多表联查 | 需要正确推断表关系和 JOIN 条件 | 在提示词中说明外键关系 |
| 方言差异 | MySQL/PostgreSQL/SQLite 语法不同 | 明确告知目标数据库类型 |
| SQL 注入 | 用户输入可能被拼接成恶意 SQL(攻击者在输入中夹带 DROP TABLE 等破坏性指令) |
只读权限 + 参数化查询 |
| 模糊语义 | "最近"是最近一天还是一周? | 业务规则文档化 |
1.2 LangChain SQL 工具链核心组件
from langchain_community.utilities import SQLDatabase
from langchain.chains import create_sql_query_chain
from langchain_community.tools.sql_database.tool import QuerySQLDataBaseTool
from langchain_openai import ChatOpenAI
# 1. SQLDatabase:数据库连接封装
# 支持 SQLAlchemy 连接字符串(SQLite/PostgreSQL/MySQL/MSSQL)
db = SQLDatabase.from_uri(
"sqlite:///./business.db",
# 可选:只暴露部分表(安全控制)
include_tables=["orders", "products", "customers", "order_items"],
# 可选:每表采样的行数,用于 schema 推断
sample_rows_in_table_info=3
)
# 查看 LangChain 提取到的 Schema 信息
print(db.get_table_info())
# 输出示例:
# CREATE TABLE orders (
# id INTEGER PRIMARY KEY,
# customer_id INTEGER,
# total_amount DECIMAL(10,2),
# status VARCHAR(20),
# created_at TIMESTAMP
# )
# /* 3 sample rows:
# id | customer_id | total_amount | status | created_at
# 1 | 42 | 299.00 | paid | 2024-01-15
# ...*/
# 2. create_sql_query_chain:生成 SQL 的链
llm = ChatOpenAI(model="gpt-4o", temperature=0)
sql_query_chain = create_sql_query_chain(llm, db)
# 测试 SQL 生成
sql = sql_query_chain.invoke({"question": "最近 7 天销售额最高的 5 个产品是哪些?"})
print(sql)
# SELECT p.name, SUM(oi.quantity * oi.price) as revenue
# FROM order_items oi
# JOIN orders o ON oi.order_id = o.id
# JOIN products p ON oi.product_id = p.id
# WHERE o.created_at >= DATE('now', '-7 days')
# GROUP BY p.id, p.name
# ORDER BY revenue DESC
# LIMIT 5
# 3. QuerySQLDataBaseTool:执行 SQL 的工具
execute_tool = QuerySQLDataBaseTool(db=db)
result = execute_tool.invoke(sql)
print(result)
1.3 完整实战:自然语言→SQL→执行→自然语言结果
from langchain_openai import ChatOpenAI
from langchain_community.utilities import SQLDatabase
from langchain.chains import create_sql_query_chain
from langchain_community.tools.sql_database.tool import QuerySQLDataBaseTool
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnablePassthrough
from operator import itemgetter
import re
llm = ChatOpenAI(model="gpt-4o", temperature=0)
# 连接数据库(示例使用 SQLite)
db = SQLDatabase.from_uri(
"sqlite:///./ecommerce.db",
include_tables=["orders", "products", "customers", "order_items"]
)
# -------- Step 1:SQL 生成链 --------
sql_chain = create_sql_query_chain(llm, db)
# -------- Step 2:SQL 清理(处理模型有时输出 markdown 代码块)--------
def clean_sql(sql_with_possible_markdown: str) -> str:
"""清理模型输出中可能的 markdown 代码块"""
# 去掉 ```sql ... ``` 包裹
sql = re.sub(r'```(?:sql)?\n?(.*?)\n?```', r'\1', sql_with_possible_markdown, flags=re.DOTALL)
# 去掉 SQLQuery: 前缀(某些模型会加)
sql = re.sub(r'^SQLQuery:\s*', '', sql.strip(), flags=re.IGNORECASE)
return sql.strip()
# -------- Step 3:SQL 执行工具 --------
execute_sql = QuerySQLDataBaseTool(db=db)
# -------- Step 4:结果自然语言化 --------
answer_prompt = ChatPromptTemplate.from_messages([
("system", """你是业务数据分析师,擅长将数据库查询结果转化为清晰的业务洞察。
规则:
- 以清晰易懂的语言回答用户问题
- 对数字进行适当的格式化(金额加货币符号,百分比保留两位小数)
- 如果结果为空,说明"暂无数据"
- 如果 SQL 执行出错,解释可能的原因"""),
("human", """用户问题:{question}
执行的 SQL:
{sql}
查询结果:
{result}
请用自然语言回答用户的问题。""")
])
answer_chain = answer_prompt | llm | StrOutputParser()
# -------- 完整流水线 --------
full_pipeline = (
RunnablePassthrough.assign(
sql=lambda x: clean_sql(sql_chain.invoke({"question": x["question"]}))
)
| RunnablePassthrough.assign(
result=lambda x: execute_sql.invoke(x["sql"])
)
| answer_chain
)
# 测试
questions = [
"本月销售额是多少?",
"哪些客户的订单从未完成?",
"最畅销的 3 个产品类别是什么?",
]
for q in questions:
print(f"\n问题:{q}")
answer = full_pipeline.invoke({"question": q})
print(f"回答:{answer}")
1.4 安全性:只读权限与 SQL 注入防护
Text-to-SQL 的最大安全风险是生成 DELETE/DROP/UPDATE 等破坏性 SQL。
1.4.1 数据库只读权限配置
# PostgreSQL:创建只读用户
# CREATE USER readonly_user WITH PASSWORD 'secure_password';
# GRANT CONNECT ON DATABASE mydb TO readonly_user;
# GRANT USAGE ON SCHEMA public TO readonly_user;
# GRANT SELECT ON ALL TABLES IN SCHEMA public TO readonly_user;
# SQLite:使用 URI 参数开启只读模式
db_readonly = SQLDatabase.from_uri("file:./ecommerce.db?mode=ro&uri=true")
# SQLAlchemy 级别的只读连接
from sqlalchemy import create_engine, event
engine = create_engine("postgresql://readonly_user:password@localhost/mydb")
# 拦截所有写操作
@event.listens_for(engine, "before_execute")
def prevent_write(conn, clauseelement, multiparams, params, execution_options):
sql_upper = str(clauseelement).upper().strip()
forbidden_keywords = ["INSERT", "UPDATE", "DELETE", "DROP", "CREATE", "ALTER", "TRUNCATE"]
for kw in forbidden_keywords:
if sql_upper.startswith(kw):
raise ValueError(f"只允许 SELECT 查询,拒绝执行:{kw} 操作")
1.4.2 SQL 验证层
def validate_sql(sql: str) -> tuple[bool, str]:
"""在执行前验证 SQL 的安全性"""
sql_upper = sql.upper().strip()
# 禁止写操作
forbidden = ["INSERT", "UPDATE", "DELETE", "DROP", "CREATE",
"ALTER", "TRUNCATE", "GRANT", "REVOKE", "EXEC", "EXECUTE"]
for kw in forbidden:
if re.search(rf'\b{kw}\b', sql_upper):
return False, f"禁止执行 {kw} 操作"
# 禁止注释(可能绕过过滤)
if "--" in sql or "/*" in sql:
return False, "SQL 中不允许包含注释"
# 必须是 SELECT 语句
if not sql_upper.lstrip().startswith("SELECT"):
return False, "只允许 SELECT 查询"
return True, "ok"
def safe_execute(sql: str, db: SQLDatabase) -> str:
"""带验证的 SQL 执行"""
is_valid, reason = validate_sql(sql)
if not is_valid:
return f"SQL 安全验证失败:{reason}"
try:
return db.run(sql)
except Exception as e:
return f"SQL 执行错误:{str(e)}"
1.5 难点处理
1.5.1 中文列名与业务术语映射
很多业务数据库的表名、列名使用缩写或英文,与用户的自然语言描述有差距:
# 在 System Prompt 中添加业务术语说明
SCHEMA_CONTEXT = """
数据库业务背景说明:
【表说明】
- orders: 订单主表,记录每笔交易
- order_items: 订单明细,每行对应订单中的一件商品
- customers: 客户信息
- products: 商品信息
【关键字段说明】
- orders.status: 订单状态,可选值:'pending'(待付款), 'paid'(已付款), 'shipped'(已发货), 'completed'(已完成), 'cancelled'(已取消)
- orders.total_amount: 订单总金额,单位元
- products.category: 商品类别,如 '电子产品', '服装', '食品'
【业务规则】
- "最近 7 天" = created_at >= DATE('now', '-7 days')(SQLite 语法)
- "本月" = strftime('%Y-%m', created_at) = strftime('%Y-%m', 'now')
- "活跃用户" = 最近 30 天内有 completed 订单的用户
"""
# 自定义 create_sql_query_chain 的 Prompt
from langchain_core.prompts import PromptTemplate
custom_prompt = PromptTemplate.from_template(
"""你是一名精通 SQLite 的 SQL 专家。
{schema_context}
数据库 Schema:
{table_info}
用户问题:{input}
写一条 SQLite SELECT 语句回答这个问题。
只输出 SQL,不要加任何解释或 markdown 代码块。
SQL 中的字符串使用单引号。
""".replace("{schema_context}", SCHEMA_CONTEXT)
)
1.5.2 模糊查询处理
# 在生成 SQL 之前,先澄清模糊意图
clarify_prompt = ChatPromptTemplate.from_template("""
用户问题:{question}
判断这个问题是否包含模糊时间范围(如"最近"、"近期"、"上个季度"等)。
如果有,输出 JSON:{{"ambiguous": true, "clarification_needed": "需要澄清的具体问题"}}
如果没有,输出 JSON:{{"ambiguous": false}}
只输出 JSON,不要其他内容。
""")
import json
from langchain_core.runnables import RunnableBranch, RunnableLambda
clarify_chain = clarify_prompt | llm | StrOutputParser()
def parse_clarify_result(x: dict) -> dict:
try:
result = json.loads(clarify_chain.invoke({"question": x["question"]}))
return {**x, "needs_clarification": result.get("ambiguous", False),
"clarification": result.get("clarification_needed", "")}
except:
return {**x, "needs_clarification": False}
clarification_branch = RunnableBranch(
(
lambda x: x.get("needs_clarification"),
RunnableLambda(lambda x: f"请问您说的'{x['clarification']}'具体是指什么时间范围?")
),
full_pipeline # 不需要澄清,直接执行
)
smart_pipeline = RunnableLambda(parse_clarify_result) | clarification_branch
1.6 进阶:结合 Schema 描述提升准确率
提供更丰富的 Schema 上下文,是提升 SQL 生成准确率最有效的手段:
# 使用 SQLDatabase 的 table_info_schema 方法获取详细 Schema
print(db.get_table_info(table_names=["orders", "order_items"]))
# 手动编写业务级 Schema 描述(推荐)
DETAILED_SCHEMA = """
Table: orders
- id: 订单 ID,主键
- customer_id: 客户 ID,外键关联 customers.id
- total_amount: 订单总金额(含税,单位:元)
- discount_amount: 优惠金额
- status: 状态(pending/paid/shipped/completed/cancelled)
- created_at: 下单时间
- paid_at: 付款时间(未付款则为 NULL)
Table: order_items
- id: 明细 ID
- order_id: 关联 orders.id
- product_id: 关联 products.id
- quantity: 购买数量
- unit_price: 下单时的单价(快照,不随商品价格变化)
关联关系:
- 一个订单(orders)包含多个明细(order_items)
- 订单金额 = SUM(order_items.quantity * order_items.unit_price) - discount_amount
"""
1.7 Text-to-SQL 与 RAG 的结合:混合查询
真实企业场景中,用户的问题往往需要同时查结构化数据和非结构化文档:
# 混合查询:结构化 + 非结构化
hybrid_prompt = ChatPromptTemplate.from_messages([
("system", """你是业务数据分析助手。根据需要从两个来源获取信息:
1. SQL 查询:用于精确的数字、统计数据
2. 文档知识库(RAG):用于业务规则、产品说明、操作手册
现有 SQL 查询结果:
{sql_result}
知识库相关内容:
{rag_context}
综合以上信息,回答用户问题。"""),
("human", "{question}")
])
hybrid_chain = (
RunnablePassthrough.assign(
sql_result=lambda x: safe_execute(
clean_sql(sql_chain.invoke({"question": x["question"]})), db
),
rag_context=lambda x: format_docs(retriever.invoke(x["question"]))
)
| hybrid_prompt
| llm
| StrOutputParser()
)
1.8 Text-to-SQL 完整流程
Text-to-SQL 的准确率很大程度上取决于 Schema 描述的质量。模型生成 SQL 的过程是在已知 Schema 上做语义推断——Schema 描述越清晰,业务规则越明确,生成就越准确。先把表名、字段名、业务含义写清楚,比调 Prompt 技巧更有效。