chat_store.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462
  1. # -*- coding: utf-8 -*-
  2. """
  3. 数据库模块 - SQLite 聊天历史存储
  4. 功能:
  5. - 存储用户会话和聊天消息
  6. - 支持按会话ID查询历史记录
  7. - 支持分页查询
  8. 表结构:
  9. - sessions: 会话表(session_id, created_at, updated_at)
  10. - messages: 消息表(id, session_id, role, content, created_at)
  11. """
  12. import sqlite3
  13. import threading
  14. import uuid
  15. from datetime import datetime
  16. from pathlib import Path
  17. from typing import Optional
  18. # 数据库文件路径
  19. DB_PATH = Path(__file__).parent.parent / "data" / "chat_history.db"
  20. # 线程本地存储(每个线程独立连接)
  21. _local = threading.local()
  22. def _get_connection() -> sqlite3.Connection:
  23. """获取当前线程的数据库连接"""
  24. if not hasattr(_local, "conn") or _local.conn is None:
  25. DB_PATH.parent.mkdir(parents=True, exist_ok=True)
  26. _local.conn = sqlite3.connect(str(DB_PATH), check_same_thread=False)
  27. _local.conn.row_factory = sqlite3.Row
  28. _local.conn.execute("PRAGMA journal_mode=WAL") # 提高并发性能
  29. _local.conn.execute("PRAGMA foreign_keys=ON")
  30. return _local.conn
  31. def init_db():
  32. """初始化数据库表结构"""
  33. conn = _get_connection()
  34. conn.executescript("""
  35. -- 会话表
  36. CREATE TABLE IF NOT EXISTS sessions (
  37. session_id TEXT PRIMARY KEY,
  38. title TEXT DEFAULT '',
  39. created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')),
  40. updated_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime'))
  41. );
  42. -- 消息表
  43. CREATE TABLE IF NOT EXISTS messages (
  44. id INTEGER PRIMARY KEY AUTOINCREMENT,
  45. session_id TEXT NOT NULL,
  46. role TEXT NOT NULL CHECK(role IN ('user', 'assistant', 'system', 'tool')),
  47. content TEXT NOT NULL DEFAULT '',
  48. tool_calls TEXT DEFAULT NULL,
  49. created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')),
  50. FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE
  51. );
  52. -- 索引
  53. CREATE INDEX IF NOT EXISTS idx_messages_session
  54. ON messages(session_id, created_at);
  55. CREATE INDEX IF NOT EXISTS idx_sessions_updated
  56. ON sessions(updated_at DESC);
  57. """)
  58. # 迁移:为已有数据库添加 mode 列(如果不存在)
  59. _migrate_add_mode(conn)
  60. # 迁移:添加 user_name 列(用户隔离),旧数据统一归为 admin
  61. _migrate_add_user_name(conn)
  62. # 迁移:添加上下文用量表
  63. _migrate_context_usage(conn)
  64. # 迁移:添加 sort_order 列(会话置顶)
  65. _migrate_add_sort_order(conn)
  66. # 迁移:添加 user_preferences 表(用户偏好记忆)
  67. _migrate_user_preferences(conn)
  68. conn.commit()
  69. def _migrate_add_mode(conn: sqlite3.Connection):
  70. """迁移:为 sessions 表添加 mode 列(如果尚不存在)。"""
  71. try:
  72. conn.execute("SELECT mode FROM sessions LIMIT 0")
  73. except sqlite3.OperationalError:
  74. conn.execute("ALTER TABLE sessions ADD COLUMN mode TEXT DEFAULT 'plan'")
  75. def _migrate_add_user_name(conn: sqlite3.Connection):
  76. """迁移:为 sessions 表添加 user_name 列(用户隔离),旧数据统一归 admin。"""
  77. try:
  78. conn.execute("SELECT user_name FROM sessions LIMIT 0")
  79. except sqlite3.OperationalError:
  80. conn.execute("ALTER TABLE sessions ADD COLUMN user_name TEXT DEFAULT 'admin'")
  81. # 填充历史数据中的 NULL 值
  82. conn.execute("UPDATE sessions SET user_name = 'admin' WHERE user_name IS NULL")
  83. def _migrate_context_usage(conn: sqlite3.Connection):
  84. """迁移:创建会话上下文用量表。"""
  85. conn.execute("""
  86. CREATE TABLE IF NOT EXISTS session_context_usage (
  87. id INTEGER PRIMARY KEY AUTOINCREMENT,
  88. session_id TEXT NOT NULL,
  89. total_limit INTEGER NOT NULL DEFAULT 1000000,
  90. current_usage INTEGER NOT NULL DEFAULT 0,
  91. messages_tokens INTEGER DEFAULT 0,
  92. tools_tokens INTEGER DEFAULT 0,
  93. skills_tokens INTEGER DEFAULT 0,
  94. mcp_tokens INTEGER DEFAULT 0,
  95. system_prompt_tokens INTEGER DEFAULT 0,
  96. other_tokens INTEGER DEFAULT 0,
  97. updated_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')),
  98. FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE
  99. )
  100. """)
  101. conn.execute("""
  102. CREATE INDEX IF NOT EXISTS idx_context_usage_session
  103. ON session_context_usage(session_id, updated_at DESC)
  104. """)
  105. def _migrate_add_sort_order(conn: sqlite3.Connection):
  106. """迁移:为 sessions 表添加 sort_order 列(置顶排序)。"""
  107. try:
  108. conn.execute("SELECT sort_order FROM sessions LIMIT 0")
  109. except sqlite3.OperationalError:
  110. conn.execute("ALTER TABLE sessions ADD COLUMN sort_order INTEGER DEFAULT 0")
  111. def _migrate_user_preferences(conn: sqlite3.Connection):
  112. """迁移:创建用户偏好记忆表。"""
  113. conn.execute("""
  114. CREATE TABLE IF NOT EXISTS user_preferences (
  115. id INTEGER PRIMARY KEY AUTOINCREMENT,
  116. user_name TEXT NOT NULL,
  117. category TEXT DEFAULT '通用',
  118. content TEXT NOT NULL,
  119. keywords TEXT DEFAULT '',
  120. created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')),
  121. updated_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime'))
  122. )
  123. """)
  124. conn.execute("""
  125. CREATE INDEX IF NOT EXISTS idx_prefs_user
  126. ON user_preferences(user_name)
  127. """)
  128. def create_session(session_id: Optional[str] = None, title: str = "",
  129. user_name: str = "admin") -> str:
  130. """创建新会话,返回 session_id"""
  131. conn = _get_connection()
  132. sid = session_id or str(uuid.uuid4())
  133. conn.execute(
  134. "INSERT OR IGNORE INTO sessions (session_id, title, user_name) VALUES (?, ?, ?)",
  135. (sid, title, user_name),
  136. )
  137. conn.commit()
  138. return sid
  139. def save_message(session_id: str, role: str, content: str,
  140. tool_calls: Optional[str] = None):
  141. """保存一条消息"""
  142. conn = _get_connection()
  143. # 确保会话存在
  144. conn.execute(
  145. "INSERT OR IGNORE INTO sessions (session_id, user_name) VALUES (?, 'admin')",
  146. (session_id,),
  147. )
  148. conn.execute(
  149. "INSERT INTO messages (session_id, role, content, tool_calls) VALUES (?, ?, ?, ?)",
  150. (session_id, role, content, tool_calls),
  151. )
  152. # 更新会话时间戳
  153. conn.execute(
  154. "UPDATE sessions SET updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  155. (session_id,),
  156. )
  157. conn.commit()
  158. def get_messages(session_id: str, limit: int = 50,
  159. offset: int = 0) -> list[dict]:
  160. """查询会话消息(分页)"""
  161. conn = _get_connection()
  162. rows = conn.execute(
  163. """SELECT id, session_id, role, content, tool_calls, created_at
  164. FROM messages
  165. WHERE session_id = ?
  166. ORDER BY created_at ASC
  167. LIMIT ? OFFSET ?""",
  168. (session_id, limit, offset),
  169. ).fetchall()
  170. return [dict(row) for row in rows]
  171. def get_sessions(limit: int = 20, offset: int = 0,
  172. user_name: str | None = None) -> list[dict]:
  173. """查询会话列表(按更新时间倒序),可按用户名过滤。"""
  174. conn = _get_connection()
  175. if user_name:
  176. rows = conn.execute(
  177. """SELECT session_id, title, user_name, sort_order, created_at, updated_at,
  178. (SELECT COUNT(*) FROM messages WHERE messages.session_id = sessions.session_id) AS message_count
  179. FROM sessions
  180. WHERE user_name = ?
  181. ORDER BY sort_order DESC, updated_at DESC
  182. LIMIT ? OFFSET ?""",
  183. (user_name, limit, offset),
  184. ).fetchall()
  185. else:
  186. rows = conn.execute(
  187. """SELECT session_id, title, user_name, sort_order, created_at, updated_at,
  188. (SELECT COUNT(*) FROM messages WHERE messages.session_id = sessions.session_id) AS message_count
  189. FROM sessions
  190. ORDER BY sort_order DESC, updated_at DESC
  191. LIMIT ? OFFSET ?""",
  192. (limit, offset),
  193. ).fetchall()
  194. return [dict(row) for row in rows]
  195. def get_session_info(session_id: str) -> dict | None:
  196. """查询单个会话的元数据(标题、用户名、模式等)。
  197. Args:
  198. session_id: 会话 ID
  199. Returns:
  200. dict 或 None(会话不存在时)
  201. """
  202. conn = _get_connection()
  203. row = conn.execute(
  204. """SELECT session_id, title, user_name, mode, sort_order, created_at, updated_at
  205. FROM sessions WHERE session_id = ?""",
  206. (session_id,),
  207. ).fetchone()
  208. return dict(row) if row else None
  209. def search_sessions(user_name: str, keyword: str, limit: int = 10) -> list[dict]:
  210. """按标题模糊搜索用户的会话。
  211. Args:
  212. user_name: 用户名
  213. keyword: 搜索关键词(支持 SQL LIKE 通配符自动包裹)
  214. limit: 最大返回条数
  215. Returns:
  216. 匹配的会话列表,按 sort_order DESC, updated_at DESC 排序
  217. """
  218. conn = _get_connection()
  219. like = f"%{keyword}%"
  220. rows = conn.execute(
  221. """SELECT session_id, title, user_name, sort_order, updated_at,
  222. (SELECT COUNT(*) FROM messages WHERE messages.session_id = sessions.session_id) AS message_count
  223. FROM sessions
  224. WHERE user_name = ? AND title LIKE ? AND title != ''
  225. ORDER BY sort_order DESC, updated_at DESC
  226. LIMIT ?""",
  227. (user_name, like, limit),
  228. ).fetchall()
  229. return [dict(row) for row in rows]
  230. def delete_session(session_id: str):
  231. """删除会话及其所有消息"""
  232. conn = _get_connection()
  233. conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,))
  234. conn.execute("DELETE FROM sessions WHERE session_id = ?", (session_id,))
  235. conn.commit()
  236. def update_session_title(session_id: str, title: str):
  237. """更新会话标题"""
  238. conn = _get_connection()
  239. conn.execute(
  240. "UPDATE sessions SET title = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  241. (title, session_id),
  242. )
  243. conn.commit()
  244. # ── 会话权限模式 ──
  245. VALID_MODES = {"plan", "full"}
  246. def set_session_mode(session_id: str, mode: str):
  247. """设置会话的权限模式。
  248. Args:
  249. session_id: 会话 ID
  250. mode: "plan" | "full"
  251. """
  252. if mode not in VALID_MODES:
  253. raise ValueError(f"无效的模式: {mode},可选: {', '.join(sorted(VALID_MODES))}")
  254. conn = _get_connection()
  255. conn.execute(
  256. "INSERT OR IGNORE INTO sessions (session_id, user_name) VALUES (?, 'admin')",
  257. (session_id,),
  258. )
  259. conn.execute(
  260. "UPDATE sessions SET mode = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  261. (mode, session_id),
  262. )
  263. conn.commit()
  264. def get_session_mode(session_id: str) -> str:
  265. """获取会话的权限模式,默认返回 'plan'。"""
  266. conn = _get_connection()
  267. row = conn.execute(
  268. "SELECT mode FROM sessions WHERE session_id = ?",
  269. (session_id,),
  270. ).fetchone()
  271. if row and row["mode"]:
  272. return row["mode"]
  273. return "plan"
  274. # 应用启动时自动初始化
  275. init_db()
  276. # ── 会话置顶 ──
  277. def pin_session(session_id: str, pinned: bool = True):
  278. """置顶/取消置顶会话。
  279. Args:
  280. session_id: 会话 ID
  281. pinned: True=置顶(sort_order=1),False=取消(sort_order=0)
  282. """
  283. conn = _get_connection()
  284. conn.execute(
  285. "UPDATE sessions SET sort_order = ? WHERE session_id = ?",
  286. (1 if pinned else 0, session_id),
  287. )
  288. conn.commit()
  289. def is_session_pinned(session_id: str) -> bool:
  290. """查询会话是否已置顶。"""
  291. conn = _get_connection()
  292. row = conn.execute(
  293. "SELECT sort_order FROM sessions WHERE session_id = ?",
  294. (session_id,),
  295. ).fetchone()
  296. return bool(row and row["sort_order"] and row["sort_order"] > 0)
  297. # ── 会话上下文用量 ──
  298. def save_context_usage(
  299. session_id: str,
  300. total_limit: int,
  301. current_usage: int,
  302. messages_tokens: int = 0,
  303. mcp_tokens: int = 0,
  304. skills_tokens: int = 0,
  305. system_prompt_tokens: int = 0,
  306. other_tokens: int = 0,
  307. ):
  308. """保存或更新会话的上下文用量快照。每次任务完成后插入一条新记录。
  309. mcp_tokens 已合并本地工具 + MCP 远程工具;tools_tokens 列保留但写 0(兼容旧表结构)。
  310. """
  311. conn = _get_connection()
  312. conn.execute(
  313. """INSERT INTO session_context_usage
  314. (session_id, total_limit, current_usage,
  315. messages_tokens, tools_tokens, skills_tokens,
  316. mcp_tokens, system_prompt_tokens, other_tokens)
  317. VALUES (?, ?, ?, ?, 0, ?, ?, ?, ?)""",
  318. (session_id, total_limit, current_usage,
  319. messages_tokens, skills_tokens,
  320. mcp_tokens, system_prompt_tokens, other_tokens),
  321. )
  322. conn.commit()
  323. def get_latest_context_usage(session_id: str) -> dict | None:
  324. """查询会话最近一次的上下文用量。"""
  325. conn = _get_connection()
  326. row = conn.execute(
  327. """SELECT * FROM session_context_usage
  328. WHERE session_id = ?
  329. ORDER BY updated_at DESC
  330. LIMIT 1""",
  331. (session_id,),
  332. ).fetchone()
  333. return dict(row) if row else None
  334. # ── 用户偏好记忆 ──
  335. def save_user_preference(user_name: str, content: str,
  336. keywords: str = "", category: str = "通用") -> int:
  337. """保存一条用户偏好/习惯。
  338. Args:
  339. user_name: 用户名
  340. content: 偏好内容(如"习惯使用 m³/s 作为风速单位")
  341. keywords: 触发关键词,逗号分隔
  342. category: 偏好分类
  343. Returns:
  344. 新插入记录的 id
  345. """
  346. conn = _get_connection()
  347. cur = conn.execute(
  348. """INSERT INTO user_preferences (user_name, category, content, keywords)
  349. VALUES (?, ?, ?, ?)""",
  350. (user_name, category, content, keywords),
  351. )
  352. conn.commit()
  353. return cur.lastrowid
  354. def get_user_preferences(user_name: str) -> list[dict]:
  355. """查询某用户的所有偏好记录。
  356. Args:
  357. user_name: 用户名
  358. Returns:
  359. 偏好记录列表,按更新时间倒序
  360. """
  361. conn = _get_connection()
  362. rows = conn.execute(
  363. """SELECT id, user_name, category, content, keywords, created_at, updated_at
  364. FROM user_preferences
  365. WHERE user_name = ?
  366. ORDER BY updated_at DESC""",
  367. (user_name,),
  368. ).fetchall()
  369. return [dict(row) for row in rows]
  370. def delete_user_preference(pref_id: int, user_name: str) -> bool:
  371. """删除一条用户偏好(带用户名校验,防止越权)。
  372. Args:
  373. pref_id: 偏好记录 ID
  374. user_name: 当前用户名
  375. Returns:
  376. True 表示删除成功,False 表示记录不存在或不属于该用户
  377. """
  378. conn = _get_connection()
  379. cur = conn.execute(
  380. "DELETE FROM user_preferences WHERE id = ? AND user_name = ?",
  381. (pref_id, user_name),
  382. )
  383. conn.commit()
  384. return cur.rowcount > 0