留言 1:"MCP 做的是把这套流程自动化:定时触发 → 执行 → 格式化 → 推送。省的不是写 SQL 的活,是每天重复执行的活。这个定时触发怎么做的?"
留言 2:"应该让 AI 自己读表 DDL,根据自然语言生成 SQL。"
好,安排。本文在上篇代码基础上,一次性解决这两个问题,再额外送你 5 个生产环境必用的扩展功能。全文 5000 字,附完整可运行代码,建议收藏。
上篇我们搭了一个基础版 MCP 服务器,连了 SQLite,暴露了两个工具:
query_sales:按产品名查销售数据get_sales_summary:销售汇总统计它能跑通,但放到真实业务里,有几个明显短板:
本文就是来解决这些问题的。
升级后的 MCP 服务器,从"两个工具的小白版"变成了"八个工具的生产版":
这是被问最多的功能。核心思路:用 apscheduler 做定时任务,到点自动调用 MCP 查询,生成 Markdown 报告,通过 Webhook 推送到群聊。
先抽象一个通知层,兼容企业微信、钉钉、飞书、邮件。实际接入时填你自己的 Webhook 地址。
# notifier.py - 推送通知模块,兼容多平台
"""推送通知模块,支持企业微信、钉钉、飞书、邮件。
实际使用时,在配置文件中填入对应的 Webhook 地址或 SMTP 信息。
"""
from __future__ import annotations
import json
import smtplib
from abc import ABC, abstractmethod
from email.mime.text import MIMEText
from typing import Any
import requests
class Notifier(ABC):
"""通知器抽象基类。"""
@abstractmethod
def send(self, title: str, content: str) -> dict[str, Any]:
"""发送通知。
Args:
title: 通知标题
content: 通知内容(Markdown 格式)
Returns:
包含发送结果的字典
"""
...
class WeComNotifier(Notifier):
"""企业微信机器人通知器。
Args:
webhook_url: 企业微信机器人 Webhook 地址
"""
def __init__(self, webhook_url: str) -> None:
self.webhook_url = webhook_url
def send(self, title: str, content: str) -> dict[str, Any]:
"""通过企业微信机器人发送 Markdown 消息。"""
payload = {
"msgtype": "markdown",
"markdown": {"content": f"**{title}**\n\n{content}"},
}
try:
resp = requests.post(
self.webhook_url,
json=payload,
timeout=30,
)
resp.raise_for_status()
return {"success": True, "platform": "wecom", "response": resp.json()}
except Exception as exc:
return {"success": False, "platform": "wecom", "error": str(exc)}
class DingTalkNotifier(Notifier):
"""钉钉机器人通知器。
Args:
webhook_url: 钉钉机器人 Webhook 地址
secret: 加签密钥(可选)
"""
def __init__(self, webhook_url: str, secret: str | None = None) -> None:
self.webhook_url = webhook_url
self.secret = secret
def _sign(self) -> str | None:
"""生成加签签名,若未配置 secret 则返回 None。"""
if not self.secret:
return None
import time
import hmac
import hashlib
import base64
timestamp = str(int(time.time() * 1000))
string_to_sign = f"{timestamp}\n{self.secret}"
hmac_code = hmac.new(
self.secret.encode("utf-8"),
string_to_sign.encode("utf-8"),
digestmod=hashlib.sha256,
).digest()
sign = base64.b64encode(hmac_code).decode("utf-8")
return f"×tamp={timestamp}&sign={sign}"
def send(self, title: str, content: str) -> dict[str, Any]:
"""通过钉钉机器人发送 Markdown 消息。"""
sign = self._sign()
url = f"{self.webhook_url}{sign}" if sign else self.webhook_url
payload = {
"msgtype": "markdown",
"markdown": {
"title": title,
"text": f"### {title}\n\n{content}",
},
}
try:
resp = requests.post(url, json=payload, timeout=30)
resp.raise_for_status()
return {"success": True, "platform": "dingtalk", "response": resp.json()}
except Exception as exc:
return {"success": False, "platform": "dingtalk", "error": str(exc)}
class LarkNotifier(Notifier):
"""飞书机器人通知器。
Args:
webhook_url: 飞书机器人 Webhook 地址
"""
def __init__(self, webhook_url: str) -> None:
self.webhook_url = webhook_url
def send(self, title: str, content: str) -> dict[str, Any]:
"""通过飞书机器人发送富文本消息。"""
payload = {
"msg_type": "post",
"content": {
"post": {
"zh_cn": {
"title": title,
"content": [
[{"tag": "text", "text": content}],
],
}
}
},
}
try:
resp = requests.post(self.webhook_url, json=payload, timeout=30)
resp.raise_for_status()
return {"success": True, "platform": "lark", "response": resp.json()}
except Exception as exc:
return {"success": False, "platform": "lark", "error": str(exc)}
class EmailNotifier(Notifier):
"""邮件通知器。
Args:
smtp_host: SMTP 服务器地址
smtp_port: SMTP 端口
username: 邮箱账号
password: 邮箱密码或授权码
to_addrs: 收件人列表
"""
def __init__(
self,
smtp_host: str,
smtp_port: int,
username: str,
password: str,
to_addrs: list[str],
) -> None:
self.smtp_host = smtp_host
self.smtp_port = smtp_port
self.username = username
self.password = password
self.to_addrs = to_addrs
def send(self, title: str, content: str) -> dict[str, Any]:
"""发送 HTML 邮件。"""
msg = MIMEText(content, "html", "utf-8")
msg["Subject"] = title
msg["From"] = self.username
msg["To"] = ", ".join(self.to_addrs)
try:
with smtplib.SMTP_SSL(self.smtp_host, self.smtp_port) as server:
server.login(self.username, self.password)
server.sendmail(self.username, self.to_addrs, msg.as_string())
return {"success": True, "platform": "email", "to": self.to_addrs}
except Exception as exc:
return {"success": False, "platform": "email", "error": str(exc)}
def create_notifier(config: dict[str, Any]) -> Notifier:
"""根据配置工厂化创建通知器。
Args:
config: 包含 platform 和其他平台特定参数的字典
Returns:
对应平台的 Notifier 实例
Raises:
ValueError: 不支持的 platform 类型
"""
platform = config.get("platform")
if platform == "wecom":
return WeComNotifier(config["webhook_url"])
if platform == "dingtalk":
return DingTalkNotifier(config["webhook_url"], config.get("secret"))
if platform == "lark":
return LarkNotifier(config["webhook_url"])
if platform == "email":
return EmailNotifier(
config["smtp_host"],
config["smtp_port"],
config["username"],
config["password"],
config["to_addrs"],
)
raise ValueError(f"不支持的通知平台: {platform}")
这是定时任务的核心。每天 9:00 自动查询数据库,生成 Markdown 报告,推送到群聊。
# scheduler.py - 定时调度器:自动查询 → 格式化 → 推送
"""定时调度器,每天自动生成销售日报并推送。
用法:
python scheduler.py
配置:修改本文件底部的 CONFIG 字典填入你的推送信息。
"""
from __future__ import annotations
import json
import sqlite3
from datetime import datetime, timedelta
from typing import Any
from apscheduler.schedulers.background import BackgroundScheduler
from apscheduler.triggers.cron import CronTrigger
from notifier import create_notifier, Notifier
# --------------------------- 配置区 ---------------------------
# 改成你的实际数据库路径
DB_PATH = "test.db"
# 推送配置示例(支持同时推多个平台)
NOTIFIER_CONFIGS: list[dict[str, Any]] = [
# 示例 1:企业微信(取消注释并填入真实 URL 即可使用)
# {"platform": "wecom", "webhook_url": "https://qyapi.weixin.qq.com/cgi-bin/webhook/..."},
# 示例 2:钉钉
# {"platform": "dingtalk", "webhook_url": "https://oapi.dingtalk.com/robot/...", "secret": "SECxxx"},
# 示例 3:飞书
# {"platform": "lark", "webhook_url": "https://open.feishu.cn/open-apis/bot/..."},
]
# 定时规则:每天 09:00 执行
CRON_HOUR = 9
CRON_MINUTE = 0
# --------------------------------------------------------------
class ReportGenerator:
"""销售日报生成器。
从 SQLite 拉取昨日销售数据,格式化为 Markdown 报告。
"""
def __init__(self, db_path: str) -> None:
self.db_path = db_path
def _query_yesterday_sales(self) -> list[dict[str, Any]]:
"""查询昨日销售明细。"""
yesterday = (datetime.now() - timedelta(days=1)).strftime("%Y-%m-%d")
conn = sqlite3.connect(self.db_path)
conn.row_factory = sqlite3.Row
cursor = conn.cursor()
cursor.execute(
"SELECT * FROM sales WHERE sale_date = ? ORDER BY amount DESC",
(yesterday,),
)
rows = [dict(row) for row in cursor.fetchall()]
conn.close()
return rows
def _query_summary(self) -> dict[str, Any]:
"""查询昨日汇总统计。"""
yesterday = (datetime.now() - timedelta(days=1)).strftime("%Y-%m-%d")
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
# 昨日总销售额
cursor.execute(
"SELECT SUM(amount) FROM sales WHERE sale_date = ?",
(yesterday,),
)
total = cursor.fetchone()[0] or 0
# 各产品销量
cursor.execute(
"""
SELECT product, COUNT(*) as count, SUM(amount) as total
FROM sales WHERE sale_date = ?
GROUP BY product ORDER BY total DESC
""",
(yesterday,),
)
by_product = [
{"product": r[0], "count": r[1], "total": r[2]}
for r in cursor.fetchall()
]
conn.close()
return {"total_amount": total, "by_product": by_product, "date": yesterday}
def generate(self) -> tuple[str, str]:
"""生成 Markdown 报告。
Returns:
(title, content) 元组
"""
summary = self._query_summary()
details = self._query_yesterday_sales()
date_str = summary["date"]
title = f"📊 销售日报 ({date_str})"
lines: list[str] = [
f"## 📊 销售日报 ({date_str})",
"",
f"**昨日总销售额:¥{summary['total_amount']:,.2f}**",
"",
"### 各产品销售排名",
"",
"| 产品 | 销量 | 销售额 |",
"|------|------|--------|",
]
for item in summary["by_product"]:
lines.append(
f"| {item['product']} | {item['count']} | ¥{item['total']:,.2f} |"
)
lines += [
"",
"### 销售明细",
"",
"| ID | 产品 | 金额 | 日期 |",
"|----|------|------|------|",
]
for row in details:
lines.append(
f"| {row['id']} | {row['product']} | ¥{row['amount']:,.2f} | {row['sale_date']} |"
)
lines += [
"",
"---",
f"*报告生成时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}*",
]
return title, "\n".join(lines)
class ScheduledJob:
"""封装定时任务逻辑。"""
def __init__(self, db_path: str, notifier_configs: list[dict[str, Any]]) -> None:
self.generator = ReportGenerator(db_path)
self.notifiers: list[Notifier] = [
create_notifier(cfg) for cfg in notifier_configs
]
def run(self) -> None:
"""执行一次日报生成与推送。"""
print(f"[{datetime.now()}] 开始生成日报...")
title, content = self.generator.generate()
for notifier in self.notifiers:
result = notifier.send(title, content)
status = "✅ 成功" if result["success"] else "❌ 失败"
print(f" [{status}] {result['platform']}: {result.get('response', result.get('error'))}")
print(f"[{datetime.now()}] 日报推送完成。\n")
def main() -> None:
"""启动后台调度器。"""
job = ScheduledJob(DB_PATH, NOTIFIER_CONFIGS)
# 立即执行一次(方便测试)
job.run()
scheduler = BackgroundScheduler()
scheduler.add_job(
job.run,
trigger=CronTrigger(hour=CRON_HOUR, minute=CRON_MINUTE),
id="daily_sales_report",
replace_existing=True,
)
scheduler.start()
print(f"定时任务已启动:每天 {CRON_HOUR:02d}:{CRON_MINUTE:02d} 自动生成并推送日报")
print("按 Ctrl+C 停止")
try:
# 保持主线程存活
import time
while True:
time.sleep(1)
except (KeyboardInterrupt, SystemExit):
scheduler.shutdown()
print("调度器已停止")
if __name__ == "__main__":
main()
安装依赖:pip install apscheduler requests
关键点解读:
BackgroundScheduler:后台运行,不阻塞主线程CronTrigger(hour=9, minute=0):每天 9:00 精确触发replace_existing=True:重复启动时自动覆盖旧任务,防重复notifier 模块抽象层:换一个平台只需要改配置,不用改业务代码这是第二条留言的核心诉求。我们新增两个工具:
explain_table:返回指定表的完整结构(字段名、类型、约束)natural_language_query:接收自然语言,AI 先读表结构,再生成 SQL 执行先多建几张表、多塞点数据,方便演示自然语言查询的通用性。
# init_db.py - 初始化数据库(增强版,新增用户表和订单表)
"""初始化 SQLite 测试数据库,包含 sales、users、orders 三张表。"""
from __future__ import annotations
import sqlite3
from datetime import datetime, timedelta
import random
DB_PATH = "test.db"
def init_sales(cursor: sqlite3.Cursor) -> None:
"""初始化销售表。"""
cursor.execute("""
CREATE TABLE IF NOT EXISTS sales (
id INTEGER PRIMARY KEY AUTOINCREMENT,
product TEXT NOT NULL,
amount REAL NOT NULL,
sale_date TEXT NOT NULL,
region TEXT DEFAULT '华东'
)
""")
products = ["iPhone 15", "MacBook Pro", "AirPods Pro", "iPad Air", "Apple Watch"]
regions = ["华东", "华南", "华北", "西南"]
base_data = [
("iPhone 15", 5999.00, "2026-08-01"),
("MacBook Pro", 14999.00, "2026-08-02"),
("AirPods Pro", 1899.00, "2026-08-02"),
("iPad Air", 4399.00, "2026-08-03"),
("iPhone 15", 5999.00, "2026-08-03"),
]
cursor.executemany(
"INSERT INTO sales (product, amount, sale_date) VALUES (?, ?, ?)",
base_data,
)
# 再生成 30 天的随机数据,方便演示自然语言查询
for i in range(50):
product = random.choice(products)
amount = random.choice([5999, 14999, 1899, 4399, 2999])
date = (datetime(2026, 8, 1) + timedelta(days=random.randint(0, 30))).strftime("%Y-%m-%d")
region = random.choice(regions)
cursor.execute(
"INSERT INTO sales (product, amount, sale_date, region) VALUES (?, ?, ?, ?)",
(product, float(amount), date, region),
)
def init_users(cursor: sqlite3.Cursor) -> None:
"""初始化用户表。"""
cursor.execute("""
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT NOT NULL UNIQUE,
email TEXT NOT NULL,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
status INTEGER DEFAULT 1 CHECK (status IN (0, 1))
)
""")
users = [
("alice", "alice@example.com", "2026-07-01"),
("bob", "bob@example.com", "2026-07-15"),
("charlie", "charlie@example.com", "2026-08-01"),
]
cursor.executemany(
"INSERT INTO users (username, email, created_at) VALUES (?, ?, ?)",
users,
)
def init_orders(cursor: sqlite3.Cursor) -> None:
"""初始化订单表。"""
cursor.execute("""
CREATE TABLE IF NOT EXISTS orders (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
product TEXT NOT NULL,
amount REAL NOT NULL,
status TEXT DEFAULT 'pending' CHECK (status IN ('pending', 'paid', 'shipped', 'cancelled')),
created_at TEXT NOT NULL DEFAULT (datetime('now')),
FOREIGN KEY (user_id) REFERENCES users(id)
)
""")
orders = [
(1, "iPhone 15", 5999.00, "paid"),
(2, "MacBook Pro", 14999.00, "shipped"),
(3, "AirPods Pro", 1899.00, "pending"),
(1, "iPad Air", 4399.00, "paid"),
]
cursor.executemany(
"INSERT INTO orders (user_id, product, amount, status) VALUES (?, ?, ?, ?)",
orders,
)
def main() -> None:
"""主入口:创建数据库并插入测试数据。"""
conn = sqlite3.connect(DB_PATH)
cursor = conn.cursor()
init_sales(cursor)
init_users(cursor)
init_orders(cursor)
conn.commit()
conn.close()
print("数据库初始化完成,已创建 sales、users、orders 三张表。")
if __name__ == "__main__":
main()
自然语言生成 SQL 的核心逻辑。为了控制篇幅,这里用规则模板 + LLM 调用两种模式。你实际接入时,可以把 generate_sql 换成调用 OpenAI/Claude API。
# sql_generator.py - 自然语言生成 SQL + 安全校验
"""自然语言转 SQL 生成器,内置 SQL 注入防护和全表扫描检测。
设计原则:
1. 先读表结构(DDL),让 AI 知道字段名和类型
2. 再基于结构生成 SQL
3. 最后过一遍安全校验,拦截危险语句
"""
from __future__ import annotations
import re
import sqlite3
from typing import Any
class SQLGenerator:
"""基于自然语言描述生成安全 SQL。
Args:
db_path: SQLite 数据库路径
"""
# 危险操作黑名单(正则)
DANGEROUS_PATTERNS = [
r"\bDROP\b",
r"\bDELETE\b(?!.*WHERE)", # 不带 WHERE 的 DELETE
r"\bUPDATE\b(?!.*WHERE)", # 不带 WHERE 的 UPDATE
r"\bINSERT\b.*\bINTO\b.*\bSELECT\b", # INSERT INTO ... SELECT 注入
r";\s*--", # 注释注入
r";\s*/\*", # 块注释注入
]
# 常见自然语言到 SQL 的映射模板(兜底规则)
TEMPLATES: list[tuple[list[str], str]] = [
(
["总销售额", "全部销售额", "销售额总计"],
"SELECT SUM(amount) as total_sales FROM {table}",
),
(
["销量", "多少条", "多少笔", "总数"],
"SELECT COUNT(*) as total_count FROM {table}",
),
(
["平均", "均值"],
"SELECT AVG(amount) as avg_amount FROM {table}",
),
(
["最高", "最大", "最贵"],
"SELECT MAX(amount) as max_amount FROM {table}",
),
(
["最低", "最小", "最便宜"],
"SELECT MIN(amount) as min_amount FROM {table}",
),
]
def __init__(self, db_path: str) -> None:
self.db_path = db_path
def get_table_schema(self, table_name: str) -> dict[str, Any]:
"""获取指定表的完整结构(DDL + 字段信息)。
Args:
table_name: 表名
Returns:
包含 columns、indexes、foreign_keys、ddl 的字典
"""
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
# 1. 表结构(字段信息)
cursor.execute(f"PRAGMA table_info({table_name})")
columns = [
{
"name": row[1],
"type": row[2],
"notnull": bool(row[3]),
"default": row[4],
"pk": bool(row[5]),
}
for row in cursor.fetchall()
]
# 2. 索引
cursor.execute(f"PRAGMA index_list({table_name})")
indexes = [{"name": row[1], "unique": bool(row[2])} for row in cursor.fetchall()]
# 3. 外键
cursor.execute(f"PRAGMA foreign_key_list({table_name})")
foreign_keys = [
{
"id": row[0],
"seq": row[1],
"table": row[2],
"from": row[3],
"to": row[4],
}
for row in cursor.fetchall()
]
# 4. 建表语句(近似 DDL)
cursor.execute(
"SELECT sql FROM sqlite_master WHERE type='table' AND name=?",
(table_name,),
)
ddl_row = cursor.fetchone()
ddl = ddl_row[0] if ddl_row else None
conn.close()
return {
"table_name": table_name,
"columns": columns,
"indexes": indexes,
"foreign_keys": foreign_keys,
"ddl": ddl,
}
def _match_template(self, nl_query: str, table_name: str) -> str | None:
"""用规则模板匹配自然语言查询,快速生成 SQL。
Args:
nl_query: 自然语言描述
table_name: 目标表名
Returns:
匹配到的 SQL 模板,或 None
"""
query_lower = nl_query.lower()
for keywords, template in self.TEMPLATES:
if any(kw in query_lower for kw in keywords):
return template.format(table=table_name)
return None
def generate_sql(self, nl_query: str, table_name: str) -> str:
"""将自然语言描述转换为 SQL。
当前实现先用规则模板匹配,未匹配时回退到基于 DDL 的提示组装。
生产环境建议接入 LLM API(OpenAI/Claude)替换此实现。
Args:
nl_query: 自然语言描述,如"上个月总销售额"
table_name: 目标表名
Returns:
生成的 SQL 语句
"""
# 第一步:尝试规则模板匹配
sql = self._match_template(nl_query, table_name)
if sql:
return sql
# 第二步:读取表结构,组装提示(这里模拟 LLM 的行为)
schema = self.get_table_schema(table_name)
columns_desc = ", ".join(
f"{c['name']} ({c['type']})" for c in schema["columns"]
)
# 简单启发式规则(兜底)
query_lower = nl_query.lower()
if "按" in query_lower and "分组" in query_lower:
# 按 XX 分组
match = re.search(r"按(\w+)分组", query_lower)
group_col = match.group(1) if match else schema["columns"][0]["name"]
return f"SELECT {group_col}, COUNT(*) as cnt, SUM(amount) as total FROM {table_name} GROUP BY {group_col}"
if "最近" in query_lower or "最新" in query_lower:
date_col = next(
(c["name"] for c in schema["columns"] if "date" in c["name"].lower() or "time" in c["name"].lower()),
schema["columns"][0]["name"],
)
return f"SELECT * FROM {table_name} ORDER BY {date_col} DESC LIMIT 10"
# 默认:查全部,限制 100 条
return f"SELECT * FROM {table_name} LIMIT 100"
def validate_sql(self, sql: str) -> dict[str, Any]:
"""校验 SQL 安全性。
检查项:
1. 是否包含危险操作(DROP、无 WHERE 的 DELETE/UPDATE)
2. 是否包含 SQL 注入特征(注释、堆叠查询)
3. 是否会导致全表扫描(无 WHERE 的 SELECT)
Args:
sql: 待校验的 SQL 语句
Returns:
包含 valid(是否通过)、reason(失败原因)、warnings(警告)的字典
"""
upper = sql.upper()
result = {"valid": True, "reason": None, "warnings": []}
# 1. 危险操作黑名单
for pattern in self.DANGEROUS_PATTERNS:
if re.search(pattern, upper, re.IGNORECASE):
result["valid"] = False
result["reason"] = f"检测到危险操作,匹配正则:{pattern}"
return result
# 2. SQL 注入特征(多语句)
if ";" in sql and not sql.strip().endswith(";"):
result["valid"] = False
result["reason"] = "检测到多语句堆叠,疑似 SQL 注入"
return result
# 3. 全表扫描警告(无 WHERE 的 SELECT)
if re.search(r"^\s*SELECT\s", upper, re.IGNORECASE) and "WHERE" not in upper:
result["warnings"].append("警告:该查询无 WHERE 条件,可能导致全表扫描")
return result
关键设计说明:
get_table_schema:用 PRAGMA table_info() 等 SQLite 原生命令获取完整表结构,比手动解析 CREATE TABLE 更可靠validate_sql:三层防御——危险操作拦截、注入特征检测、全表扫描警告TEMPLATES:兜底规则,确保常见查询不走 LLM 也能快速响应这是全文最长的代码块,但也是最值钱的。它把前面所有功能串在一起,加上权限控制、审计日志、缓存、数据导出。
# server.py - 生产级 MCP 服务器:定时推送 + 自然语言查询 + 权限 + 审计 + 缓存
"""生产级 MCP 服务器,集成 SQLite 数据库连接、自然语言 SQL 生成、
权限控制、审计日志、查询缓存、数据导出等功能。
用法:
1. 先运行 python init_db.py 初始化数据库
2. 再运行 python server.py 启动 MCP 服务器
3. 在 Claude Desktop / Cursor 中配置 MCP 连接
环境变量:
MCP_TOKEN: 访问令牌(生产环境必须设置,否则只读模式)
"""
from __future__ import annotations
import hashlib
import json
import os
import sqlite3
import time
from datetime import datetime, timedelta
from typing import Any
from mcp.server import Server
from mcp.server.stdio import stdio_server
from mcp.types import Tool, TextContent
from sql_generator import SQLGenerator
# ========================= 配置区 =========================
DB_PATH = "test.db"
AUDIT_DB_PATH = "audit.db" # 审计日志独立存储,防篡改
CACHE_TTL_SECONDS = 60 # 查询缓存 60 秒
TOKEN = os.getenv("MCP_TOKEN", "") # 生产环境必须设置
# ==========================================================
# ------------------------- 连接池 -------------------------
class ConnectionPool:
"""SQLite 连接池(简化版,生产环境建议用 sqlalchemy.pool)。
SQLite 是文件级锁,多线程同时写会冲突。连接池通过复用连接 + WAL 模式缓解。
"""
_instance: ConnectionPool | None = None
def __new__(cls, *args: Any, **kwargs: Any) -> ConnectionPool:
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._initialized = False
return cls._instance
def __init__(self, db_path: str, max_size: int = 5) -> None:
if self._initialized:
return
self.db_path = db_path
self.max_size = max_size
self._pool: list[sqlite3.Connection] = []
self._lock = False
self._initialized = True
# 初始化 WAL 模式(提升并发性能)
conn = sqlite3.connect(db_path, timeout=10)
conn.execute("PRAGMA journal_mode=WAL")
conn.close()
def get(self) -> sqlite3.Connection:
"""获取一个连接。简化实现:直接新建,但有超时控制。"""
return sqlite3.connect(self.db_path, timeout=10)
# ------------------------- 审计日志 -------------------------
class AuditLogger:
"""审计日志记录器,记录每次查询的时间、调用方、工具名、参数。"""
def __init__(self, db_path: str) -> None:
self.db_path = db_path
self._init_db()
def _init_db(self) -> None:
"""初始化审计日志表。"""
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
cursor.execute("""
CREATE TABLE IF NOT EXISTS audit_logs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
timestamp TEXT NOT NULL DEFAULT (datetime('now')),
tool_name TEXT NOT NULL,
arguments TEXT,
result_summary TEXT,
client_info TEXT
)
""")
conn.commit()
conn.close()
def log(self, tool_name: str, arguments: dict[str, Any], result_summary: str, client_info: str = "") -> None:
"""记录一次工具调用。"""
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
cursor.execute(
"INSERT INTO audit_logs (tool_name, arguments, result_summary, client_info) VALUES (?, ?, ?, ?)",
(tool_name, json.dumps(arguments, ensure_ascii=False), result_summary[:500], client_info),
)
conn.commit()
conn.close()
# ------------------------- 查询缓存 -------------------------
class QueryCache:
"""简单内存缓存,基于查询参数哈希,TTL 过期。"""
def __init__(self, ttl_seconds: int = 60) -> None:
self.ttl_seconds = ttl_seconds
self._cache: dict[str, dict[str, Any]] = {}
def _key(self, tool_name: str, arguments: dict[str, Any]) -> str:
"""生成缓存键。"""
raw = f"{tool_name}:{json.dumps(arguments, sort_keys=True, ensure_ascii=False)}"
return hashlib.md5(raw.encode()).hexdigest()
def get(self, tool_name: str, arguments: dict[str, Any]) -> list[TextContent] | None:
"""获取缓存结果。"""
key = self._key(tool_name, arguments)
entry = self._cache.get(key)
if not entry:
return None
if time.time() - entry["timestamp"] > self.ttl_seconds:
del self._cache[key]
return None
return entry["data"]
def set(self, tool_name: str, arguments: dict[str, Any], data: list[TextContent]) -> None:
"""写入缓存。"""
key = self._key(tool_name, arguments)
self._cache[key] = {"timestamp": time.time(), "data": data}
def clear(self) -> None:
"""清空缓存。"""
self._cache.clear()
# ------------------------- 权限控制 -------------------------
class AuthGuard:
"""简单的 Token 权限校验。"""
def __init__(self, expected_token: str) -> None:
self.expected_token = expected_token
self.enabled = bool(expected_token)
def check(self, token: str | None) -> bool:
"""校验 Token。未启用时直接通过。"""
if not self.enabled:
return True
return token == self.expected_token
# ========================= 初始化组件 =========================
pool = ConnectionPool(DB_PATH)
audit = AuditLogger(AUDIT_DB_PATH)
cache = QueryCache(CACHE_TTL_SECONDS)
auth = AuthGuard(TOKEN)
sql_generator = SQLGenerator(DB_PATH)
app = Server("sqlite-assistant-pro")
# =============================================================
@app.list_tools()
async def list_tools() -> list[Tool]:
"""注册所有工具——AI 客户端根据这个列表"看见"能力。"""
return [
# --- 基础查询(保留上篇功能) ---
Tool(
name="query_sales",
description="查询销售数据。可以按产品名筛选,也可以查全部。",
inputSchema={
"type": "object",
"properties": {
"product": {
"type": "string",
"description": "要查询的产品名称,如 'iPhone 15'。留空则查全部。"
}
},
"required": []
}
),
Tool(
name="get_sales_summary",
description="获取销售汇总统计:总销售额、各产品销量、最近日期。",
inputSchema={
"type": "object",
"properties": {},
"required": []
}
),
# --- 新功能:表结构解释 ---
Tool(
name="explain_table",
description="获取指定表的完整结构(字段、类型、约束、索引、外键)。",
inputSchema={
"type": "object",
"properties": {
"table_name": {
"type": "string",
"description": "表名,如 'sales'、'users'、'orders'"
}
},
"required": ["table_name"]
}
),
# --- 新功能:自然语言查询 ---
Tool(
name="natural_language_query",
description="用自然语言描述你的查询需求,AI 会自动读取表结构、生成 SQL 并执行。",
inputSchema={
"type": "object",
"properties": {
"table_name": {
"type": "string",
"description": "要查询的表名"
},
"query": {
"type": "string",
"description": "自然语言描述,如 '上个月总销售额'、'按产品分组统计销量'"
}
},
"required": ["table_name", "query"]
}
),
# --- 新功能:带条件通用查询 ---
Tool(
name="query_with_filter",
description="通用条件查询:支持多字段、多条件组合,不再硬编码 WHERE。",
inputSchema={
"type": "object",
"properties": {
"table_name": {
"type": "string",
"description": "表名"
},
"columns": {
"type": "array",
"items": {"type": "string"},
"description": "要查询的字段列表,留空则查全部"
},
"filters": {
"type": "object",
"description": "条件字典,如 {'product': 'iPhone 15', 'region': '华东'}"
},
"order_by": {
"type": "string",
"description": "排序字段,如 'sale_date DESC'"
},
"limit": {
"type": "integer",
"description": "限制返回条数,默认 100",
"default": 100
}
},
"required": ["table_name"]
}
),
# --- 新功能:数据导出 ---
Tool(
name="export_data",
description="将查询结果导出为 CSV 或 Excel 格式,返回下载链接。",
inputSchema={
"type": "object",
"properties": {
"table_name": {
"type": "string",
"description": "表名"
},
"format": {
"type": "string",
"enum": ["csv", "excel"],
"description": "导出格式"
},
"filters": {
"type": "object",
"description": "筛选条件(同 query_with_filter)"
}
},
"required": ["table_name", "format"]
}
),
# --- 新功能:缓存管理 ---
Tool(
name="clear_cache",
description="清空查询缓存,强制下次查询走数据库。",
inputSchema={
"type": "object",
"properties": {},
"required": []
}
),
]
@app.call_tool()
async def call_tool(name: str, arguments: dict[str, Any]) -> list[TextContent]:
"""处理所有工具调用——这是 MCP 服务器的核心处理逻辑。"""
# 1. 权限校验(如果配置了 TOKEN)
token = arguments.pop("_token", None)
if not auth.check(token):
return [TextContent(type="text", text='{"error": "Unauthorized: invalid token"}')]
# 2. 缓存命中检查
cached = cache.get(name, arguments)
if cached is not None:
audit.log(name, arguments, "CACHE_HIT", "")
return cached
# 3. 执行具体逻辑
result = await _execute_tool(name, arguments)
# 4. 写入缓存 + 审计日志
cache.set(name, arguments, result)
summary = result[0].text[:200] if result else ""
audit.log(name, arguments, summary, "")
return result
async def _execute_tool(name: str, arguments: dict[str, Any]) -> list[TextContent]:
"""实际执行工具逻辑的分发器。"""
conn = pool.get()
cursor = conn.cursor()
try:
# ---------- 1. query_sales(上篇保留) ----------
if name == "query_sales":
product = arguments.get("product", "")
if product:
cursor.execute(
"SELECT * FROM sales WHERE product LIKE ? ORDER BY sale_date DESC",
(f"%{product}%",)
)
else:
cursor.execute("SELECT * FROM sales ORDER BY sale_date DESC")
rows = cursor.fetchall()
columns = [desc[0] for desc in cursor.description]
results = [dict(zip(columns, row)) for row in rows]
return [TextContent(type="text", text=json.dumps(results, ensure_ascii=False, indent=2))]
# ---------- 2. get_sales_summary(上篇保留) ----------
elif name == "get_sales_summary":
cursor.execute("SELECT SUM(amount) FROM sales")
total = cursor.fetchone()[0] or 0
cursor.execute("""
SELECT product, COUNT(*) as count, SUM(amount) as total
FROM sales GROUP BY product ORDER BY total DESC
""")
by_product = [{"product": r[0], "count": r[1], "total": r[2]} for r in cursor.fetchall()]
cursor.execute("SELECT MAX(sale_date) FROM sales")
latest = cursor.fetchone()[0]
summary = {
"total_amount": total,
"by_product": by_product,
"latest_date": latest
}
return [TextContent(type="text", text=json.dumps(summary, ensure_ascii=False, indent=2))]
# ---------- 3. explain_table(新:读 DDL) ----------
elif name == "explain_table":
table_name = arguments.get("table_name", "")
if not table_name:
return [TextContent(type="text", text='{"error": "table_name 不能为空"}')]
schema = sql_generator.get_table_schema(table_name)
return [TextContent(type="text", text=json.dumps(schema, ensure_ascii=False, indent=2))]
# ---------- 4. natural_language_query(新:自然语言 → SQL) ----------
elif name == "natural_language_query":
table_name = arguments.get("table_name", "")
query = arguments.get("query", "")
if not table_name or not query:
return [TextContent(type="text", text='{"error": "table_name 和 query 都不能为空"}')]
# 生成 SQL
sql = sql_generator.generate_sql(query, table_name)
# 安全校验
validation = sql_generator.validate_sql(sql)
if not validation["valid"]:
return [TextContent(type="text", text=json.dumps({
"error": "SQL 校验失败",
"reason": validation["reason"],
"generated_sql": sql
}, ensure_ascii=False, indent=2))]
# 执行
cursor.execute(sql)
rows = cursor.fetchall()
columns = [desc[0] for desc in cursor.description] if cursor.description else []
results = [dict(zip(columns, row)) for row in rows] if columns else []
response = {
"generated_sql": sql,
"warnings": validation.get("warnings", []),
"results": results,
"count": len(results)
}
return [TextContent(type="text", text=json.dumps(response, ensure_ascii=False, indent=2))]
# ---------- 5. query_with_filter(新:通用条件查询) ----------
elif name == "query_with_filter":
table_name = arguments.get("table_name", "")
columns = arguments.get("columns", [])
filters = arguments.get("filters", {})
order_by = arguments.get("order_by", "")
limit = arguments.get("limit", 100)
if not table_name:
return [TextContent(type="text", text='{"error": "table_name 不能为空"}')]
# 构建 SQL
col_str = ", ".join(columns) if columns else "*"
sql = f"SELECT {col_str} FROM {table_name}"
params: list[Any] = []
if filters:
conditions = []
for key, value in filters.items():
conditions.append(f"{key} = ?")
params.append(value)
sql += " WHERE " + " AND ".join(conditions)
if order_by:
sql += f" ORDER BY {order_by}"
sql += f" LIMIT {limit}"
# 安全校验
validation = sql_generator.validate_sql(sql)
if not validation["valid"]:
return [TextContent(type="text", text=json.dumps({"error": validation["reason"]}, ensure_ascii=False))]
cursor.execute(sql, params)
rows = cursor.fetchall()
cols = [desc[0] for desc in cursor.description] if cursor.description else []
results = [dict(zip(cols, row)) for row in rows] if cols else []
return [TextContent(type="text", text=json.dumps({"sql": sql, "results": results}, ensure_ascii=False, indent=2))]
# ---------- 6. export_data(新:数据导出) ----------
elif name == "export_data":
table_name = arguments.get("table_name", "")
fmt = arguments.get("format", "csv")
filters = arguments.get("filters", {})
if not table_name:
return [TextContent(type="text", text='{"error": "table_name 不能为空"}')]
# 先查询数据
sql = f"SELECT * FROM {table_name}"
params: list[Any] = []
if filters:
conditions = []
for key, value in filters.items():
conditions.append(f"{key} = ?")
params.append(value)
sql += " WHERE " + " AND ".join(conditions)
sql += " LIMIT 10000" # 导出限制 1 万条
cursor.execute(sql, params)
rows = cursor.fetchall()
cols = [desc[0] for desc in cursor.description] if cursor.description else []
if not cols:
return [TextContent(type="text", text='{"error": "表为空或不存在"}')]
# 生成导出文件
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
filename = f"{table_name}_{timestamp}"
if fmt == "csv":
import csv
filepath = f"{filename}.csv"
with open(filepath, "w", newline="", encoding="utf-8-sig") as f:
writer = csv.writer(f)
writer.writerow(cols)
writer.writerows(rows)
return [TextContent(type="text", text=f"CSV 导出成功:{filepath}({len(rows)} 条记录)")]
elif fmt == "excel":
try:
import openpyxl
except ImportError:
return [TextContent(type="text", text='{"error": "请先安装 openpyxl:pip install openpyxl"}')]
filepath = f"{filename}.xlsx"
wb = openpyxl.Workbook()
ws = wb.active
ws.title = table_name
ws.append(cols)
for row in rows:
ws.append(row)
wb.save(filepath)
return [TextContent(type="text", text=f"Excel 导出成功:{filepath}({len(rows)} 条记录)")]
else:
return [TextContent(type="text", text='{"error": "不支持的格式,请用 csv 或 excel"}')]
# ---------- 7. clear_cache(新:清空缓存) ----------
elif name == "clear_cache":
cache.clear()
return [TextContent(type="text", text='{"status": "缓存已清空"}')]
# ---------- 未知工具 ----------
else:
return [TextContent(type="text", text=f'{{"error": "未知工具:{name}"}}')]
except sqlite3.Error as exc:
return [TextContent(type="text", text=f'{{"error": "数据库错误:{str(exc)}"}}')]
finally:
conn.close()
async def main() -> None:
"""启动 MCP 服务器。"""
async with stdio_server() as (read_stream, write_stream):
await app.run(read_stream, write_stream)
if __name__ == "__main__":
import asyncio
asyncio.run(main())
代码结构总览:
| 模块 | 职责 | 关键类/函数 |
|---|---|---|
| 连接池 | 复用连接、WAL 模式 | ConnectionPool |
| 审计日志 | 记录谁在查什么 | AuditLogger |
| 查询缓存 | 60 秒 TTL 缓存 | QueryCache |
| 权限控制 | Token 校验 | AuthGuard |
| SQL 生成 | 自然语言 → SQL | SQLGenerator |
| 工具注册 | 暴露 7 个工具 | list_tools |
| 执行分发 | 处理所有调用 | call_tool → _execute_tool |
上面 server.py 的 export_data 工具已经实现了这个功能。补充两点说明:
csv 模块,带 UTF-8-BOM(兼容 Excel 中文不乱码)openpyxl,没有会自动提示安装使用示例(在 Claude 中输入):
把最近一周的销售数据导出成 Excel
AI 会自动调用 query_with_filter 查数据,再调用 export_data 生成文件。
现象:定时任务和 MCP 查询同时执行,报 database is locked
原因:SQLite 默认是 journal_mode=DELETE,写操作会锁整个文件
解决:conn.execute("PRAGMA journal_mode=WAL") # 写前日志,读写不阻塞,加上 timeout=10,连接时等 10 秒再报错。
现象:AI 说"帮我删了测试数据",结果生成 DELETE FROM sales
原因:没有 WHERE 的 DELETE 被允许了
解决:validate_sql 中加入正则 r"\bDELETE\b(?!.*WHERE)",拦截无 WHERE 的 DELETE。
现象:AI 不调用 natural_language_query,自己编 SQL
原因:description 里没写清楚"什么时候用它"
解决:description 要包含场景关键词:"用自然语言描述你的查询需求,AI 会自动读取表结构、生成 SQL 并执行。"这样 AI 就知道"人话 → 用它,SQL → 不用它"。
现象:改了数据库,AI 返回的还是旧数据
原因:QueryCache 60 秒 TTL,期间不查数据库
解决:暴露 clear_cache 工具,让 AI 可以主动清缓存。或者在定时任务推送后自动清缓存。
现象:AI 说"导出成功",但文件在哪?
原因:MCP 服务器在子进程运行,工作目录和 AI 客户端不一样
解决:导出时用绝对路径,或在 description 里写清楚"文件保存在服务器运行目录"。
对比上篇,这次升级带来了质的飞跃:
| 能力 | 上篇 | 本文 |
|---|---|---|
| 查询方式 | 手动触发 | 自然语言 + 定时自动 |
| 表结构理解 | 无 | AI 自动读 DDL |
| 安全性 | 裸奔 | Token + 审计 + SQL 校验 |
| 性能 | 每次查库 | 60 秒缓存 |
| 分享 | 复制 JSON | 自动推送到群聊 |
| 导出 | 无 | CSV / Excel |
| 通用性 | 硬编码 WHERE | 任意条件组合 |
init_db.py 初始化测试数据python server.py,在 Claude/Cursor 里测试每个工具scheduler.py 的 Webhook,接入你的企业微信/钉钉SQLGenerator.generate_sql 换成调用 OpenAI API,实现真正的 LLM 驱动MCP 的本质是给 AI 装"手"。上篇装了一只手,这篇装了一整套工具箱 + 自动流水线。2026 年,不会用 MCP 就像 2023 年不会写提示词——但比提示词更狠,因为它直接改变 AI 的能力边界。