chat_store.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419
  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 delete_session(session_id: str):
  196. """删除会话及其所有消息"""
  197. conn = _get_connection()
  198. conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,))
  199. conn.execute("DELETE FROM sessions WHERE session_id = ?", (session_id,))
  200. conn.commit()
  201. def update_session_title(session_id: str, title: str):
  202. """更新会话标题"""
  203. conn = _get_connection()
  204. conn.execute(
  205. "UPDATE sessions SET title = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  206. (title, session_id),
  207. )
  208. conn.commit()
  209. # ── 会话权限模式 ──
  210. VALID_MODES = {"plan", "full"}
  211. def set_session_mode(session_id: str, mode: str):
  212. """设置会话的权限模式。
  213. Args:
  214. session_id: 会话 ID
  215. mode: "plan" | "full"
  216. """
  217. if mode not in VALID_MODES:
  218. raise ValueError(f"无效的模式: {mode},可选: {', '.join(sorted(VALID_MODES))}")
  219. conn = _get_connection()
  220. conn.execute(
  221. "INSERT OR IGNORE INTO sessions (session_id, user_name) VALUES (?, 'admin')",
  222. (session_id,),
  223. )
  224. conn.execute(
  225. "UPDATE sessions SET mode = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  226. (mode, session_id),
  227. )
  228. conn.commit()
  229. def get_session_mode(session_id: str) -> str:
  230. """获取会话的权限模式,默认返回 'plan'。"""
  231. conn = _get_connection()
  232. row = conn.execute(
  233. "SELECT mode FROM sessions WHERE session_id = ?",
  234. (session_id,),
  235. ).fetchone()
  236. if row and row["mode"]:
  237. return row["mode"]
  238. return "plan"
  239. # 应用启动时自动初始化
  240. init_db()
  241. # ── 会话置顶 ──
  242. def pin_session(session_id: str, pinned: bool = True):
  243. """置顶/取消置顶会话。
  244. Args:
  245. session_id: 会话 ID
  246. pinned: True=置顶(sort_order=1),False=取消(sort_order=0)
  247. """
  248. conn = _get_connection()
  249. conn.execute(
  250. "UPDATE sessions SET sort_order = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  251. (1 if pinned else 0, session_id),
  252. )
  253. conn.commit()
  254. def is_session_pinned(session_id: str) -> bool:
  255. """查询会话是否已置顶。"""
  256. conn = _get_connection()
  257. row = conn.execute(
  258. "SELECT sort_order FROM sessions WHERE session_id = ?",
  259. (session_id,),
  260. ).fetchone()
  261. return bool(row and row["sort_order"] and row["sort_order"] > 0)
  262. # ── 会话上下文用量 ──
  263. def save_context_usage(
  264. session_id: str,
  265. total_limit: int,
  266. current_usage: int,
  267. messages_tokens: int = 0,
  268. mcp_tokens: int = 0,
  269. skills_tokens: int = 0,
  270. system_prompt_tokens: int = 0,
  271. other_tokens: int = 0,
  272. ):
  273. """保存或更新会话的上下文用量快照。每次任务完成后插入一条新记录。
  274. mcp_tokens 已合并本地工具 + MCP 远程工具;tools_tokens 列保留但写 0(兼容旧表结构)。
  275. """
  276. conn = _get_connection()
  277. conn.execute(
  278. """INSERT INTO session_context_usage
  279. (session_id, total_limit, current_usage,
  280. messages_tokens, tools_tokens, skills_tokens,
  281. mcp_tokens, system_prompt_tokens, other_tokens)
  282. VALUES (?, ?, ?, ?, 0, ?, ?, ?, ?)""",
  283. (session_id, total_limit, current_usage,
  284. messages_tokens, skills_tokens,
  285. mcp_tokens, system_prompt_tokens, other_tokens),
  286. )
  287. conn.commit()
  288. def get_latest_context_usage(session_id: str) -> dict | None:
  289. """查询会话最近一次的上下文用量。"""
  290. conn = _get_connection()
  291. row = conn.execute(
  292. """SELECT * FROM session_context_usage
  293. WHERE session_id = ?
  294. ORDER BY updated_at DESC
  295. LIMIT 1""",
  296. (session_id,),
  297. ).fetchone()
  298. return dict(row) if row else None
  299. # ── 用户偏好记忆 ──
  300. def save_user_preference(user_name: str, content: str,
  301. keywords: str = "", category: str = "通用") -> int:
  302. """保存一条用户偏好/习惯。
  303. Args:
  304. user_name: 用户名
  305. content: 偏好内容(如"习惯使用 m³/s 作为风速单位")
  306. keywords: 触发关键词,逗号分隔
  307. category: 偏好分类
  308. Returns:
  309. 新插入记录的 id
  310. """
  311. conn = _get_connection()
  312. cur = conn.execute(
  313. """INSERT INTO user_preferences (user_name, category, content, keywords)
  314. VALUES (?, ?, ?, ?)""",
  315. (user_name, category, content, keywords),
  316. )
  317. conn.commit()
  318. return cur.lastrowid
  319. def get_user_preferences(user_name: str) -> list[dict]:
  320. """查询某用户的所有偏好记录。
  321. Args:
  322. user_name: 用户名
  323. Returns:
  324. 偏好记录列表,按更新时间倒序
  325. """
  326. conn = _get_connection()
  327. rows = conn.execute(
  328. """SELECT id, user_name, category, content, keywords, created_at, updated_at
  329. FROM user_preferences
  330. WHERE user_name = ?
  331. ORDER BY updated_at DESC""",
  332. (user_name,),
  333. ).fetchall()
  334. return [dict(row) for row in rows]
  335. def delete_user_preference(pref_id: int, user_name: str) -> bool:
  336. """删除一条用户偏好(带用户名校验,防止越权)。
  337. Args:
  338. pref_id: 偏好记录 ID
  339. user_name: 当前用户名
  340. Returns:
  341. True 表示删除成功,False 表示记录不存在或不属于该用户
  342. """
  343. conn = _get_connection()
  344. cur = conn.execute(
  345. "DELETE FROM user_preferences WHERE id = ? AND user_name = ?",
  346. (pref_id, user_name),
  347. )
  348. conn.commit()
  349. return cur.rowcount > 0