paper_reader_history.py 3.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107
  1. """阅读历史记录 —— 对话会话持久化、恢复与上下文延续."""
  2. from __future__ import annotations
  3. import sqlite3
  4. import time
  5. from contextlib import contextmanager
  6. from ...utils.common import exec_sql
  7. def ensure_tables(db_path: str) -> None:
  8. exec_sql(db_path,
  9. "CREATE TABLE IF NOT EXISTS paper_reader_turns(id INTEGER PRIMARY KEY AUTOINCREMENT,paper_id INTEGER NOT NULL,role TEXT NOT NULL,content TEXT NOT NULL,created_at INTEGER NOT NULL)",
  10. "CREATE INDEX IF NOT EXISTS idx_paper_reader_turns_paper ON paper_reader_turns(paper_id,created_at)",
  11. )
  12. @contextmanager
  13. def _conn(db_path: str, *, row_factory=None):
  14. conn = sqlite3.connect(db_path)
  15. if row_factory:
  16. conn.row_factory = row_factory
  17. try:
  18. yield conn
  19. conn.commit()
  20. finally:
  21. conn.close()
  22. def append_turn(db_path: str, *, paper_id: int, role: str, content: str) -> None:
  23. ensure_tables(db_path)
  24. role2 = (role or "").strip().lower()
  25. if role2 not in ("user", "assistant"):
  26. role2 = "user"
  27. text = (content or "").strip()
  28. if not text:
  29. return
  30. now = int(time.time())
  31. with _conn(db_path) as conn:
  32. conn.execute(
  33. "INSERT INTO paper_reader_turns(paper_id,role,content,created_at) VALUES(?,?,?,?)",
  34. (int(paper_id), role2, text, now),
  35. )
  36. def prepend_turn(
  37. db_path: str,
  38. *,
  39. paper_id: int,
  40. role: str,
  41. content: str,
  42. before_created_at: int,
  43. ) -> None:
  44. ensure_tables(db_path)
  45. role2 = (role or "").strip().lower()
  46. if role2 not in ("user", "assistant"):
  47. role2 = "user"
  48. text = (content or "").strip()
  49. if not text:
  50. return
  51. ts = int(before_created_at) - 1
  52. if ts < 0:
  53. ts = 0
  54. now = int(time.time())
  55. if ts >= now:
  56. ts = now - 1
  57. with _conn(db_path) as conn:
  58. conn.execute(
  59. "INSERT INTO paper_reader_turns(paper_id,role,content,created_at) VALUES(?,?,?,?)",
  60. (int(paper_id), role2, text, ts),
  61. )
  62. def ensure_opening_turn(db_path: str, *, paper_id: int, opening_text: str) -> None:
  63. op = (opening_text or "").strip()
  64. if not op:
  65. return
  66. turns = list_turns(db_path, paper_id=int(paper_id), limit=5)
  67. if not turns:
  68. append_turn(db_path, paper_id=int(paper_id), role="assistant", content=op)
  69. return
  70. first = turns[0]
  71. r0 = (first.get("role") or "").strip().lower()
  72. c0 = (first.get("content") or "").strip()
  73. if r0 == "assistant" and c0 == op:
  74. return
  75. if r0 == "assistant":
  76. return
  77. if r0 == "user":
  78. try:
  79. ts0 = int(first.get("created_at") or 0)
  80. except Exception:
  81. ts0 = int(time.time())
  82. prepend_turn(db_path, paper_id=int(paper_id), role="assistant", content=op, before_created_at=ts0)
  83. return
  84. def list_turns(db_path: str, *, paper_id: int, limit: int = 200) -> list[dict[str, str | None]]:
  85. ensure_tables(db_path)
  86. with _conn(db_path, row_factory=sqlite3.Row) as conn:
  87. rows = conn.execute(
  88. "SELECT role,content,created_at FROM paper_reader_turns WHERE paper_id=? ORDER BY created_at ASC,id ASC LIMIT ?",
  89. (int(paper_id), int(limit)),
  90. ).fetchall()
  91. out: list[dict[str, str | None]] = []
  92. for r in rows:
  93. out.append({
  94. "role": (r["role"] or "").strip(),
  95. "content": (r["content"] or "").strip(),
  96. "created_at": int(r["created_at"] or 0),
  97. })
  98. return out