我们现在开始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 🙌

        github地址:https://github.com/wbshbrsjlr/zy-docmind-2026

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐