# -*- 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); """) conn.commit() def create_session(session_id: Optional[str] = None, title: str = "") -> str: """创建新会话,返回 session_id""" conn = _get_connection() sid = session_id or str(uuid.uuid4()) conn.execute( "INSERT OR IGNORE INTO sessions (session_id, title) VALUES (?, ?)", (sid, title), ) 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) VALUES (?)", (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) -> list[dict]: """查询会话列表(按更新时间倒序)""" conn = _get_connection() rows = conn.execute( """SELECT session_id, title, created_at, updated_at, (SELECT COUNT(*) FROM messages WHERE messages.session_id = sessions.session_id) AS message_count FROM sessions ORDER BY 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() # 应用启动时自动初始化 init_db()