# -*- coding: utf-8 -*- """ 数据库模块 - SQLite 聊天历史存储 功能: - 存储用户会话和聊天消息 - 支持按会话ID查询历史记录 - 支持分页查询 表结构: - sessions: 会话表(session_id, created_at, updated_at) - messages: 消息表(id, session_id, role, content, created_at) """ import sqlite3 import threading import uuid from datetime import datetime from pathlib import Path from typing import Optional # 数据库文件路径 DB_PATH = Path(__file__).parent.parent / "data" / "chat_history.db" # 线程本地存储(每个线程独立连接) _local = threading.local() def _get_connection() -> sqlite3.Connection: """获取当前线程的数据库连接""" if not hasattr(_local, "conn") or _local.conn is None: DB_PATH.parent.mkdir(parents=True, exist_ok=True) _local.conn = sqlite3.connect(str(DB_PATH), check_same_thread=False) _local.conn.row_factory = sqlite3.Row _local.conn.execute("PRAGMA journal_mode=WAL") # 提高并发性能 _local.conn.execute("PRAGMA foreign_keys=ON") return _local.conn def init_db(): """初始化数据库表结构""" conn = _get_connection() conn.executescript(""" -- 会话表 CREATE TABLE IF NOT EXISTS sessions ( session_id TEXT PRIMARY KEY, title TEXT DEFAULT '', created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')), updated_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')) ); -- 消息表 CREATE TABLE IF NOT EXISTS messages ( id INTEGER PRIMARY KEY AUTOINCREMENT, session_id TEXT NOT NULL, role TEXT NOT NULL CHECK(role IN ('user', 'assistant', 'system', 'tool')), content TEXT NOT NULL DEFAULT '', tool_calls TEXT DEFAULT NULL, created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')), FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE ); -- 索引 CREATE INDEX IF NOT EXISTS idx_messages_session ON messages(session_id, created_at); CREATE INDEX IF NOT EXISTS idx_sessions_updated ON sessions(updated_at DESC); """) # 迁移:为已有数据库添加 mode 列(如果不存在) _migrate_add_mode(conn) # 迁移:添加 user_name 列(用户隔离),旧数据统一归为 admin _migrate_add_user_name(conn) # 迁移:添加上下文用量表 _migrate_context_usage(conn) # 迁移:添加 sort_order 列(会话置顶) _migrate_add_sort_order(conn) # 迁移:添加 user_preferences 表(用户偏好记忆) _migrate_user_preferences(conn) conn.commit() def _migrate_add_mode(conn: sqlite3.Connection): """迁移:为 sessions 表添加 mode 列(如果尚不存在)。""" try: conn.execute("SELECT mode FROM sessions LIMIT 0") except sqlite3.OperationalError: conn.execute("ALTER TABLE sessions ADD COLUMN mode TEXT DEFAULT 'plan'") def _migrate_add_user_name(conn: sqlite3.Connection): """迁移:为 sessions 表添加 user_name 列(用户隔离),旧数据统一归 admin。""" try: conn.execute("SELECT user_name FROM sessions LIMIT 0") except sqlite3.OperationalError: conn.execute("ALTER TABLE sessions ADD COLUMN user_name TEXT DEFAULT 'admin'") # 填充历史数据中的 NULL 值 conn.execute("UPDATE sessions SET user_name = 'admin' WHERE user_name IS NULL") def _migrate_context_usage(conn: sqlite3.Connection): """迁移:创建会话上下文用量表。""" conn.execute(""" CREATE TABLE IF NOT EXISTS session_context_usage ( id INTEGER PRIMARY KEY AUTOINCREMENT, session_id TEXT NOT NULL, total_limit INTEGER NOT NULL DEFAULT 1000000, current_usage INTEGER NOT NULL DEFAULT 0, messages_tokens INTEGER DEFAULT 0, tools_tokens INTEGER DEFAULT 0, skills_tokens INTEGER DEFAULT 0, mcp_tokens INTEGER DEFAULT 0, system_prompt_tokens INTEGER DEFAULT 0, other_tokens INTEGER DEFAULT 0, updated_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')), FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE ) """) conn.execute(""" CREATE INDEX IF NOT EXISTS idx_context_usage_session ON session_context_usage(session_id, updated_at DESC) """) def _migrate_add_sort_order(conn: sqlite3.Connection): """迁移:为 sessions 表添加 sort_order 列(置顶排序)。""" try: conn.execute("SELECT sort_order FROM sessions LIMIT 0") except sqlite3.OperationalError: conn.execute("ALTER TABLE sessions ADD COLUMN sort_order INTEGER DEFAULT 0") def _migrate_user_preferences(conn: sqlite3.Connection): """迁移:创建用户偏好记忆表。""" conn.execute(""" CREATE TABLE IF NOT EXISTS user_preferences ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_name TEXT NOT NULL, category TEXT DEFAULT '通用', content TEXT NOT NULL, keywords TEXT DEFAULT '', created_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')), updated_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime')) ) """) conn.execute(""" CREATE INDEX IF NOT EXISTS idx_prefs_user ON user_preferences(user_name) """) def create_session(session_id: Optional[str] = None, title: str = "", user_name: str = "admin") -> str: """创建新会话,返回 session_id""" conn = _get_connection() sid = session_id or str(uuid.uuid4()) conn.execute( "INSERT OR IGNORE INTO sessions (session_id, title, user_name) VALUES (?, ?, ?)", (sid, title, user_name), ) conn.commit() return sid def save_message(session_id: str, role: str, content: str, tool_calls: Optional[str] = None): """保存一条消息""" conn = _get_connection() # 确保会话存在 conn.execute( "INSERT OR IGNORE INTO sessions (session_id, user_name) VALUES (?, 'admin')", (session_id,), ) conn.execute( "INSERT INTO messages (session_id, role, content, tool_calls) VALUES (?, ?, ?, ?)", (session_id, role, content, tool_calls), ) # 更新会话时间戳 conn.execute( "UPDATE sessions SET updated_at = datetime('now', 'localtime') WHERE session_id = ?", (session_id,), ) conn.commit() def get_messages(session_id: str, limit: int = 50, offset: int = 0) -> list[dict]: """查询会话消息(分页)""" conn = _get_connection() rows = conn.execute( """SELECT id, session_id, role, content, tool_calls, created_at FROM messages WHERE session_id = ? ORDER BY created_at ASC LIMIT ? OFFSET ?""", (session_id, limit, offset), ).fetchall() return [dict(row) for row in rows] def get_sessions(limit: int = 20, offset: int = 0, user_name: str | None = None) -> list[dict]: """查询会话列表(按更新时间倒序),可按用户名过滤。""" conn = _get_connection() if user_name: rows = conn.execute( """SELECT session_id, title, user_name, sort_order, created_at, updated_at, (SELECT COUNT(*) FROM messages WHERE messages.session_id = sessions.session_id) AS message_count FROM sessions WHERE user_name = ? ORDER BY sort_order DESC, updated_at DESC LIMIT ? OFFSET ?""", (user_name, limit, offset), ).fetchall() else: rows = conn.execute( """SELECT session_id, title, user_name, sort_order, created_at, updated_at, (SELECT COUNT(*) FROM messages WHERE messages.session_id = sessions.session_id) AS message_count FROM sessions ORDER BY sort_order DESC, updated_at DESC LIMIT ? OFFSET ?""", (limit, offset), ).fetchall() return [dict(row) for row in rows] def delete_session(session_id: str): """删除会话及其所有消息""" conn = _get_connection() conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,)) conn.execute("DELETE FROM sessions WHERE session_id = ?", (session_id,)) conn.commit() def update_session_title(session_id: str, title: str): """更新会话标题""" conn = _get_connection() conn.execute( "UPDATE sessions SET title = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?", (title, session_id), ) conn.commit() # ── 会话权限模式 ── VALID_MODES = {"plan", "full"} def set_session_mode(session_id: str, mode: str): """设置会话的权限模式。 Args: session_id: 会话 ID mode: "plan" | "full" """ if mode not in VALID_MODES: raise ValueError(f"无效的模式: {mode},可选: {', '.join(sorted(VALID_MODES))}") conn = _get_connection() conn.execute( "INSERT OR IGNORE INTO sessions (session_id, user_name) VALUES (?, 'admin')", (session_id,), ) conn.execute( "UPDATE sessions SET mode = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?", (mode, session_id), ) conn.commit() def get_session_mode(session_id: str) -> str: """获取会话的权限模式,默认返回 'plan'。""" conn = _get_connection() row = conn.execute( "SELECT mode FROM sessions WHERE session_id = ?", (session_id,), ).fetchone() if row and row["mode"]: return row["mode"] return "plan" # 应用启动时自动初始化 init_db() # ── 会话置顶 ── def pin_session(session_id: str, pinned: bool = True): """置顶/取消置顶会话。 Args: session_id: 会话 ID pinned: True=置顶(sort_order=1),False=取消(sort_order=0) """ conn = _get_connection() conn.execute( "UPDATE sessions SET sort_order = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?", (1 if pinned else 0, session_id), ) conn.commit() def is_session_pinned(session_id: str) -> bool: """查询会话是否已置顶。""" conn = _get_connection() row = conn.execute( "SELECT sort_order FROM sessions WHERE session_id = ?", (session_id,), ).fetchone() return bool(row and row["sort_order"] and row["sort_order"] > 0) # ── 会话上下文用量 ── def save_context_usage( session_id: str, total_limit: int, current_usage: int, messages_tokens: int = 0, mcp_tokens: int = 0, skills_tokens: int = 0, system_prompt_tokens: int = 0, other_tokens: int = 0, ): """保存或更新会话的上下文用量快照。每次任务完成后插入一条新记录。 mcp_tokens 已合并本地工具 + MCP 远程工具;tools_tokens 列保留但写 0(兼容旧表结构)。 """ conn = _get_connection() conn.execute( """INSERT INTO session_context_usage (session_id, total_limit, current_usage, messages_tokens, tools_tokens, skills_tokens, mcp_tokens, system_prompt_tokens, other_tokens) VALUES (?, ?, ?, ?, 0, ?, ?, ?, ?)""", (session_id, total_limit, current_usage, messages_tokens, skills_tokens, mcp_tokens, system_prompt_tokens, other_tokens), ) conn.commit() def get_latest_context_usage(session_id: str) -> dict | None: """查询会话最近一次的上下文用量。""" conn = _get_connection() row = conn.execute( """SELECT * FROM session_context_usage WHERE session_id = ? ORDER BY updated_at DESC LIMIT 1""", (session_id,), ).fetchone() return dict(row) if row else None # ── 用户偏好记忆 ── def save_user_preference(user_name: str, content: str, keywords: str = "", category: str = "通用") -> int: """保存一条用户偏好/习惯。 Args: user_name: 用户名 content: 偏好内容(如"习惯使用 m³/s 作为风速单位") keywords: 触发关键词,逗号分隔 category: 偏好分类 Returns: 新插入记录的 id """ conn = _get_connection() cur = conn.execute( """INSERT INTO user_preferences (user_name, category, content, keywords) VALUES (?, ?, ?, ?)""", (user_name, category, content, keywords), ) conn.commit() return cur.lastrowid def get_user_preferences(user_name: str) -> list[dict]: """查询某用户的所有偏好记录。 Args: user_name: 用户名 Returns: 偏好记录列表,按更新时间倒序 """ conn = _get_connection() rows = conn.execute( """SELECT id, user_name, category, content, keywords, created_at, updated_at FROM user_preferences WHERE user_name = ? ORDER BY updated_at DESC""", (user_name,), ).fetchall() return [dict(row) for row in rows] def delete_user_preference(pref_id: int, user_name: str) -> bool: """删除一条用户偏好(带用户名校验,防止越权)。 Args: pref_id: 偏好记录 ID user_name: 当前用户名 Returns: True 表示删除成功,False 表示记录不存在或不属于该用户 """ conn = _get_connection() cur = conn.execute( "DELETE FROM user_preferences WHERE id = ? AND user_name = ?", (pref_id, user_name), ) conn.commit() return cur.rowcount > 0