DocMind-python-agent-V2:Text-to-SQL
我们现在开始V2版本功能的开发,基于github上的
https://github.com/wbshbrsjlr/zy-docmind-2026
DocMind 目前只能“读”(文档),还不能“算”(数据库),而真实企业场景里,这两者缺一不可。
Text-to-SQL 是考验 LLM 推理能力的经典场景:模型需要把自然语言“翻译”成精确的 SQL 语法、处理表结构、应对报错并重试。这个过程中,Agent 会展示完整的 “观察(Query)→ 思考(Generate SQL)→ 行动(Execute)→ 观察(Result)” 循环。没有它,你的 Agent 只是一个“检索路由”,有了它,Agent 才真正开始“动脑子”。
我们开始开发 text-to-sql的相关功能,使用postgreSql数据库版本16:
需要新增/修改的文件清单
| 文件路径 | 操作 | 说明 |
|---|---|---|
core/config.py | 修改 | 增加 PostgreSQL 配置项 |
tools/sql_tool.py | 新增 | Text-to-SQL 工具核心代码 |
tools/__init__.py | 修改 | 导出新工具 |
main.py | 修改 | 注册 SQL 工具 |
.env | 修改 | 添加数据库连接信息(不提交到 Git) |
1. 添加依赖
uv add psycopg2-binary
2. 修改 core/config.py
# @file:core/config.py
import os
from dotenv import load_dotenv
load_dotenv()
class Config:
"""集中管理所有配置项"""
DEEPSEEK_API_KEY: str = os.getenv("DEEPSEEK_API_KEY", "")
DEEPSEEK_BASE_URL: str = os.getenv("DEEPSEEK_BASE_URL", "https://api.deepseek.com")
LLM_MODEL: str = os.getenv("LLM_MODEL", "deepseek-v4-flash")
OLLAMA_BASE_URL: str = os.getenv("OLLAMA_BASE_URL", "http://127.0.0.1:11434")
EMBEDDING_MODEL: str = os.getenv("EMBEDDING_MODEL", "qwen3-embedding:4b")
CHUNK_SIZE: int = int(os.getenv("CHUNK_SIZE", "500"))
CHUNK_OVERLAP: int = int(os.getenv("CHUNK_OVERLAP", "50"))
# --- 🆕 PostgreSQL 配置() ---
PG_HOST: str = os.getenv("PG_HOST", "127.0.0.1")
PG_PORT: int = int(os.getenv("PG_PORT", "5432"))
PG_DB: str = os.getenv("PG_DB", "postgres")
PG_USER: str = os.getenv("PG_USER", "postgres")
PG_PASSWORD: str = os.getenv("PG_PASSWORD", "")
@classmethod
def validate(cls):
if not cls.DEEPSEEK_API_KEY:
raise ValueError("请在 .env 文件中设置 DEEPSEEK_API_KEY")
print(f"✅ 配置验证通过,Ollama 地址: {cls.OLLAMA_BASE_URL}")
print(f"✅ PostgreSQL 目标: {cls.PG_HOST}:{cls.PG_PORT}/{cls.PG_DB}")
3.新增 tools/sql_tool.py
# tools/sql_tool.py
# 导入 psycopg2,用于连接和操作 PostgreSQL 数据库
import psycopg2
# 从 psycopg2 中导入 OperationalError,用于捕获数据库操作异常
from psycopg2 import OperationalError
# 从 pydantic 导入 BaseModel 和 Field,用于定义工具输入参数的结构和描述
from pydantic import BaseModel, Field
# 从 langchain.tools 导入 tool 装饰器,用于将函数注册为 LangChain 工具
from langchain.tools import tool
# 从 core.config 导入 Config,获取数据库连接配置(主机、端口、库名、用户名、密码等)
from core.config import Config
# 定义工具输入参数的 Pydantic 模型,用于校验和描述
class SQLInput(BaseModel):
# query 字段为字符串类型,必须提供,描述为只读 SQL 查询,强制要求以 SELECT 开头
query: str = Field(
description="合法的只读 SQL 查询语句,必须以 SELECT 开头,不要包含 INSERT/UPDATE/DELETE。"
)
# 使用 @tool 装饰器将此函数声明为 LangChain 工具,指定参数模型和工具描述
@tool(args_schema=SQLInput, description="查询 PostgreSQL 数据库中的结构化数据(如财务表、销售记录)。仅支持 SELECT。")
def query_sql_database(query: str) -> str:
"""
连接远程 PostgreSQL(192.168.0.31)执行查询,返回 Markdown 表格。
"""
# 去除查询字符串首尾的空白字符
clean_query = query.strip()
# 检查是否以 SELECT 开头(不区分大小写),否则拒绝执行并返回错误信息
if not clean_query.upper().startswith("SELECT"):
return "❌ 错误:出于安全考虑,仅支持 SELECT 只读查询。"
# 构造数据库连接参数字典,从 Config 中读取配置
conn_params = {
"host": Config.PG_HOST, # 数据库主机地址
"port": Config.PG_PORT, # 数据库端口
"dbname": Config.PG_DB, # 数据库名称
"user": Config.PG_USER, # 数据库用户名
"password": Config.PG_PASSWORD, # 数据库密码
"connect_timeout": 5, # 连接超时秒数
"client_encoding": "utf8", # 客户端字符编码,防止中文乱码
}
try:
# 使用参数建立数据库连接
conn = psycopg2.connect(**conn_params)
# 设置会话为只读模式,并禁用自动提交(确保只读)
conn.set_session(readonly=True, autocommit=False)
# 创建游标对象用于执行 SQL
cur = conn.cursor()
# 执行清理后的 SQL 查询
cur.execute(clean_query)
# 如果游标没有描述信息(即无结果集),则关闭游标和连接,返回执行成功无数据
if cur.description is None:
cur.close()
conn.close()
return "✅ 查询执行成功,无返回数据。"
# 从游标描述中提取列名作为表头
headers = [desc[0] for desc in cur.description]
# 获取所有查询结果行
rows = cur.fetchall()
# 定义最大返回行数,防止结果过大
MAX_ROWS = 30
truncated = False
# 如果结果行数超过最大行数,则截断并标记为已截断
if len(rows) > MAX_ROWS:
rows = rows[:MAX_ROWS]
truncated = True
# 如果截断后没有行,返回执行成功但结果为空
if not rows:
return "✅ 查询执行成功,但结果为空(0 行)。"
# 构造 Markdown 表格输出
md_lines = []
# 添加表头行,用管道符分隔各列
md_lines.append("| " + " | ".join(headers) + " |")
# 添加分隔行,每列使用三个短横线
md_lines.append("| " + " | ".join(["---"] * len(headers)) + " |")
# 遍历每一行数据,将每个单元格转换为字符串,若为 None 则置为空字符串
for row in rows:
formatted = [str(cell) if cell is not None else "" for cell in row]
md_lines.append("| " + " | ".join(formatted) + " |")
# 将列表合并为多行字符串
result_md = "\n".join(md_lines)
# 如果结果被截断,添加警告信息
if truncated:
result_md += f"\n\n⚠️ 结果超过 {MAX_ROWS} 行,仅展示前 {MAX_ROWS} 行。"
# 关闭游标和连接,释放资源
cur.close()
conn.close()
# 返回 Markdown 表格结果
return result_md
# 捕获数据库连接或操作错误(OperationalError)
except OperationalError as e:
# 返回友好的错误提示,包含主机、端口以及排查建议
return (
f"❌ 数据库连接失败({Config.PG_HOST}:{Config.PG_PORT})。\n"
f"请检查:\n"
f"1. Docker 容器是否运行\n"
f"2. 端口映射是否正确 (-p 5432:5432)\n"
f"3. 用户名/密码是否正确\n"
f"错误详情: {e}"
)
# 捕获其他所有异常,返回通用错误信息
except Exception as e:
return f"❌ SQL 执行异常: {e}"
代码中的双星号 ** 是 Python 中的字典解包(Dictionary Unpacking)语法,也叫关键字参数解包。简单来说:** 会把一个字典拆解成 key=value 形式的关键字参数,传递给函数。
加了 ** 后,psycopg2.connect(**conn_params) 就等同于写成这样:
psycopg2.connect( host="localhost", port=5432, dbname="mydb", user="admin", password="123456" )
4.tools/__init__.py(修改,新增导出)
# tools/__init__.py
from tools.rag_tool import query_knowledge_base
from tools.sql_tool import query_sql_database # 新增
__all__ = ["query_knowledge_base", "query_sql_database"]
5.main.py(修改两处)
# @file : main.py
# 从 langchain_core.messages 导入 AIMessage 类,用于处理 AI 回复消息
from langchain_core.messages import AIMessage
# 从 langchain_openai 导入 ChatOpenAI 类,用于初始化大语言模型客户端
from langchain_openai import ChatOpenAI
# 从 langchain.agents 导入 create_agent 函数,用于创建智能体
from langchain.agents import create_agent
# 从 langgraph.checkpoint.memory 导入 MemorySaver,用于对话状态记忆
from langgraph.checkpoint.memory import MemorySaver
# 从核心配置模块导入 Config,获取应用配置
from core.config import Config
# 从核心 HTTP 客户端模块导入 LoggingHttpClient,用于带日志的 HTTP 请求
from core.http_client import LoggingHttpClient
# 从工具模块导入知识库查询工具函数
from tools.rag_tool import query_knowledge_base
from tools.sql_tool import query_sql_database # 🆕 新增导入
# 从向量存储索引管道模块导入索引函数
from stores.index_pipeline import index_pipeline
# ---------- 初始化 ----------
# 验证配置项是否完整有效(如 API 密钥等)
Config.validate()
# 创建带有日志功能的 HTTP 客户端,设置超时时间为 60 秒
http_client = LoggingHttpClient(timeout=60.0)
# 初始化 ChatOpenAI 大语言模型实例
llm = ChatOpenAI(
# 模型名称(从配置中读取)
model=Config.LLM_MODEL,
# API 密钥(从配置中读取 DeepSeek API 密钥)
api_key=Config.DEEPSEEK_API_KEY,
# API 基础 URL(从配置中读取 DeepSeek 端点)
base_url=Config.DEEPSEEK_BASE_URL,
# 温度参数设为 0.7,控制回复的随机性
temperature=0.7,
# 传入自定义 HTTP 客户端
http_client=http_client
)
# ---------- 启动时索引 ----------
# 打印分隔线和启动信息
print("=" * 60)
print("DocMind 启动中...")
print("=" * 60)
# 调用索引管道,指定文档目录和 LLM 实例
index_pipeline(docs_dir="./docs", llm=llm)
# ---------- 创建 Agent ----------
# 创建内存检查点保存器,用于存储对话历史
memory = MemorySaver()
# 使用 create_agent 创建智能体
agent = create_agent(
# 指定使用的模型
model=llm,
# 提供工具列表(知识库查询工具)
tools=[query_knowledge_base, query_sql_database], # 🆕 加入第二个工具
# 设置系统提示词,指导智能体行为
system_prompt=(
"你是企业文档智能助手 DocMind 。\n"
"1. 当用户询问 PDF/Word/Excel 里的内容时,使用 query_knowledge_base。\n"
"2. 当用户询问数据报表、财务数字、统计指标时,请根据问题生成 SQL,使用 query_sql_database 查询。\n"
"3. 如果问题需要同时用到两者,你可以先查数据库,再查文档,最后汇总。\n"
"4. 回答要注明数据来源(数据库或文档)。"
),
# 传入检查点保存器,实现对话记忆
checkpointer=memory
)
# ---------- 对话循环 ----------
# 打印就绪信息和提示
print("\n" + "=" * 60)
print("DocMind 已就绪")
print("输入 'exit' 退出")
print("=" * 60)
# 设置线程 ID,用于区分不同用户的对话
thread_id = "user_001"
# 进入无限循环,持续接收用户输入
while True:
# 获取用户输入(去除首尾空白)
user_input = input("\n你: ")
# 如果用户输入 'exit'(不区分大小写),则退出循环
if user_input.lower() == "exit":
break
# 如果输入为空或仅有空白,则跳过本次循环
if not user_input.strip():
continue
# 调用智能体的 invoke 方法,传入用户消息和配置(包含线程 ID)
result = agent.invoke(
{"messages": [{"role": "user", "content": user_input}]},
config={"configurable": {"thread_id": thread_id}}
)
# 初始化变量,用于存储最后一条 AI 消息
last_ai_msg = None
# 反向遍历智能体返回的消息列表(从后向前)
for msg in reversed(result['messages']):
# 如果消息是 AIMessage 类型且包含内容
if isinstance(msg, AIMessage) and msg.content:
# 赋值并跳出循环
last_ai_msg = msg
break
# 打印分隔线
print("-" * 60)
# 如果找到了 AI 消息,则打印其内容
if last_ai_msg:
print(f"🤖 DocMind: {last_ai_msg.content}")
else:
# 否则打印未获取到有效回复的提示
print("🤖 DocMind: (未获取到有效回复)")
# 打印分隔线
print("-" * 60)
6.在env中添加postgresql内容:
# PostgreSQL 连接 PG_HOST=yourIp PG_PORT=5432 PG_DB=yourDbName PG_USER=postgres PG_PASSWORD=yourPwd
7.准备数据,在postgresql中执行:
-- 创建销售数据表
CREATE TABLE IF NOT EXISTS sales_data (
id SERIAL PRIMARY KEY,
product_name VARCHAR(100) NOT NULL,
region VARCHAR(50) NOT NULL,
amount DECIMAL(10, 2) NOT NULL,
sale_date DATE NOT NULL
);-- 插入 10 条测试数据
INSERT INTO sales_data (product_name, region, amount, sale_date) VALUES
('智能助手', '华东', 1500.00, '2026-01-15'),
('智能助手', '华南', 1200.00, '2026-02-20'),
('智能助手', '华东', 1800.00, '2026-03-10'),
('文档解析', '华东', 800.00, '2026-01-22'),
('文档解析', '华北', 950.00, '2026-02-14'),
('文档解析', '华南', 700.00, '2026-03-05'),
('数据中台', '华南', 1100.00, '2026-01-30'),
('数据中台', '华东', 1300.00, '2026-02-28'),
('数据中台', '华北', 900.00, '2026-03-15'),
('数据中台', '华东', 1600.00, '2026-03-28');
8.启动main,输入 "查询 sales_data 表的所有数据" ,便可以看到如下信息了:

大家可能比较疑惑,我们声明tool的时候并没有将表结构传入给大模型,那他是如何知道的并成功写出对应sql的呢?下面的问答告诉大家:
你: 你是怎么知道表结构的
------------------------------------------------------------
🤖 DocMind: 问得好!我来给你还原一下我是怎么知道表结构的 👇### 步骤一:查看数据库中有哪些表
我执行了以下 SQL 查询:
```sql
SELECT table_name
FROM information_schema.tables
WHERE table_schema = 'public'
```
得到了一个表:**`sales_data`**### 步骤二:查看该表有哪些列
我执行了以下 SQL 查询:
```sql
SELECT column_name, data_type
FROM information_schema.columns
WHERE table_schema = 'public'
AND table_name = 'sales_data'
```
得到了表中的 **5 个字段**:| 字段名 | 数据类型 |
|--------|---------|
| `id` | 整数 |
| `product_name` | 字符串 |
| `region` | 字符串 |
| `amount` | 数值 |
| `sale_date` | 日期 |---
### 总结
我并不是"事先知道"表结构的,而是**实时查询了 PostgreSQL 自带的系统表**(`information_schema`),它就像是数据库的"户口本",记录着所有表、字段、数据类型等信息。然后根据这些信息,我才写出了正确的查询语句来找到最近的销售记录。这种方式的好处是——不管你的数据库结构怎么变,我都能动态适应 😊
------------------------------------------------------------你: 也就是说你调用了至少两次query_sql_database方法?一次查表结构,一次查数据?
------------------------------------------------------------
🤖 DocMind: 你说得没错,不过我实际上调用了 **3 次** `query_sql_database`,不是 2 次 😄来看看完整的调用链:
---
### 第 1 次调用 — 查有哪些表
```sql
SELECT table_name FROM information_schema.tables WHERE table_schema = 'public'
```
> 结果:发现了一个表叫 `sales_data`### 第 2 次调用 — 查有哪些字段
```sql
SELECT column_name, data_type FROM information_schema.columns
WHERE table_schema = 'public' AND table_name = 'sales_data'
```
> 结果:知道了 `id`、`product_name`、`region`、`amount`、`sale_date` 这 5 个字段### 第 3 次调用 — 正式查数据
```sql
SELECT ... FROM sales_data ORDER BY sale_date DESC LIMIT 1
```
> 结果:拿到了最近一笔销售记录---
所以我的工作流程是:**先摸清数据库的"底细"(表结构),再精准地查询你要的数据**,这样即使表名、字段名我不认识,也能灵活适配,不会写错 SQL 🙌
更多推荐


所有评论(0)