chat_store.py 4.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154
  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. conn.commit()
  59. def create_session(session_id: Optional[str] = None, title: str = "") -> str:
  60. """创建新会话,返回 session_id"""
  61. conn = _get_connection()
  62. sid = session_id or str(uuid.uuid4())
  63. conn.execute(
  64. "INSERT OR IGNORE INTO sessions (session_id, title) VALUES (?, ?)",
  65. (sid, title),
  66. )
  67. conn.commit()
  68. return sid
  69. def save_message(session_id: str, role: str, content: str,
  70. tool_calls: Optional[str] = None):
  71. """保存一条消息"""
  72. conn = _get_connection()
  73. # 确保会话存在
  74. conn.execute(
  75. "INSERT OR IGNORE INTO sessions (session_id) VALUES (?)",
  76. (session_id,),
  77. )
  78. conn.execute(
  79. "INSERT INTO messages (session_id, role, content, tool_calls) VALUES (?, ?, ?, ?)",
  80. (session_id, role, content, tool_calls),
  81. )
  82. # 更新会话时间戳
  83. conn.execute(
  84. "UPDATE sessions SET updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  85. (session_id,),
  86. )
  87. conn.commit()
  88. def get_messages(session_id: str, limit: int = 50,
  89. offset: int = 0) -> list[dict]:
  90. """查询会话消息(分页)"""
  91. conn = _get_connection()
  92. rows = conn.execute(
  93. """SELECT id, session_id, role, content, tool_calls, created_at
  94. FROM messages
  95. WHERE session_id = ?
  96. ORDER BY created_at ASC
  97. LIMIT ? OFFSET ?""",
  98. (session_id, limit, offset),
  99. ).fetchall()
  100. return [dict(row) for row in rows]
  101. def get_sessions(limit: int = 20, offset: int = 0) -> list[dict]:
  102. """查询会话列表(按更新时间倒序)"""
  103. conn = _get_connection()
  104. rows = conn.execute(
  105. """SELECT session_id, title, created_at, updated_at,
  106. (SELECT COUNT(*) FROM messages WHERE messages.session_id = sessions.session_id) AS message_count
  107. FROM sessions
  108. ORDER BY updated_at DESC
  109. LIMIT ? OFFSET ?""",
  110. (limit, offset),
  111. ).fetchall()
  112. return [dict(row) for row in rows]
  113. def delete_session(session_id: str):
  114. """删除会话及其所有消息"""
  115. conn = _get_connection()
  116. conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,))
  117. conn.execute("DELETE FROM sessions WHERE session_id = ?", (session_id,))
  118. conn.commit()
  119. def update_session_title(session_id: str, title: str):
  120. """更新会话标题"""
  121. conn = _get_connection()
  122. conn.execute(
  123. "UPDATE sessions SET title = ?, updated_at = datetime('now', 'localtime') WHERE session_id = ?",
  124. (title, session_id),
  125. )
  126. conn.commit()
  127. # 应用启动时自动初始化
  128. init_db()