| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419 |
- # -*- 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
|