content_db.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369
  1. """SQLite 内容数据库操作类"""
  2. import json
  3. import sqlite3
  4. from pathlib import Path
  5. from typing import Optional
  6. class ContentDB:
  7. """SQLite 内容数据库操作类,每个文档对应一个独立的 SQLite 数据库文件"""
  8. def __init__(self, db_path: str):
  9. self.db_path = Path(db_path)
  10. self.conn: Optional[sqlite3.Connection] = None
  11. def connect(self):
  12. """连接数据库"""
  13. self.conn = sqlite3.connect(str(self.db_path))
  14. self.conn.row_factory = sqlite3.Row
  15. return self
  16. def close(self):
  17. """关闭连接"""
  18. if self.conn:
  19. self.conn.close()
  20. self.conn = None
  21. def __enter__(self):
  22. """上下文管理器入口"""
  23. return self.connect()
  24. def __exit__(self, exc_type, exc_val, exc_tb):
  25. """上下文管理器退出"""
  26. self.close()
  27. def create_tables(self):
  28. """创建 document_blocks 表"""
  29. self.conn.execute("""
  30. CREATE TABLE IF NOT EXISTS document_blocks (
  31. id TEXT PRIMARY KEY,
  32. block_order INTEGER NOT NULL,
  33. type TEXT NOT NULL,
  34. level INTEGER DEFAULT 0,
  35. "index" INTEGER DEFAULT 0,
  36. content TEXT NOT NULL,
  37. word_style TEXT DEFAULT '',
  38. style TEXT DEFAULT '{}',
  39. metadata TEXT DEFAULT '{}'
  40. )
  41. """)
  42. self.conn.execute('CREATE INDEX IF NOT EXISTS idx_block_order ON document_blocks(block_order)')
  43. self.conn.execute('CREATE INDEX IF NOT EXISTS idx_type ON document_blocks(type)')
  44. self.conn.execute('CREATE INDEX IF NOT EXISTS idx_level ON document_blocks(level)')
  45. self.conn.commit()
  46. def insert_blocks(self, blocks: list[dict]):
  47. """批量插入 blocks"""
  48. for block in blocks:
  49. content = block['content']
  50. if isinstance(content, (dict, list)):
  51. content = json.dumps(content, ensure_ascii=False)
  52. self.conn.execute("""
  53. INSERT INTO document_blocks
  54. (id, block_order, type, level, "index", content, word_style, style, metadata)
  55. VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
  56. """, (
  57. block['id'],
  58. block['block_order'],
  59. block['type'],
  60. block.get('level', 0),
  61. block.get('index', 0),
  62. content,
  63. block.get('word_style', ''),
  64. json.dumps(block.get('style', {}), ensure_ascii=False),
  65. json.dumps(block.get('metadata', {}), ensure_ascii=False)
  66. ))
  67. self.conn.commit()
  68. def get_blocks(self, order_by: str = 'block_order') -> list[dict]:
  69. """查询所有 blocks"""
  70. cursor = self.conn.execute(f"""
  71. SELECT * FROM document_blocks
  72. ORDER BY {order_by}
  73. """)
  74. rows = cursor.fetchall()
  75. return [self._row_to_dict(row) for row in rows]
  76. def get_block_by_id(self, block_id: str) -> Optional[dict]:
  77. """按 ID 查询单个 block"""
  78. cursor = self.conn.execute("""
  79. SELECT * FROM document_blocks WHERE id = ?
  80. """, (block_id,))
  81. row = cursor.fetchone()
  82. return self._row_to_dict(row) if row else None
  83. def update_block(self, block_id: str, updates: dict):
  84. """更新单个 block"""
  85. set_clauses = []
  86. params = []
  87. if 'content' in updates:
  88. content = updates['content']
  89. if isinstance(content, (dict, list)):
  90. content = json.dumps(content, ensure_ascii=False)
  91. set_clauses.append('content = ?')
  92. params.append(content)
  93. if 'style' in updates:
  94. set_clauses.append('style = ?')
  95. params.append(json.dumps(updates['style'], ensure_ascii=False))
  96. if 'word_style' in updates:
  97. set_clauses.append('word_style = ?')
  98. params.append(updates['word_style'])
  99. if 'metadata' in updates:
  100. set_clauses.append('metadata = ?')
  101. params.append(json.dumps(updates['metadata'], ensure_ascii=False))
  102. if not set_clauses:
  103. return
  104. params.append(block_id)
  105. sql = f"UPDATE document_blocks SET {', '.join(set_clauses)} WHERE id = ?"
  106. self.conn.execute(sql, params)
  107. self.conn.commit()
  108. def delete_block(self, block_id: str):
  109. """删除单个 block"""
  110. self.conn.execute('DELETE FROM document_blocks WHERE id = ?', (block_id,))
  111. self.conn.commit()
  112. def search_blocks(self, query: str, block_type: Optional[str] = None) -> list[dict]:
  113. """搜索 blocks"""
  114. sql = "SELECT * FROM document_blocks WHERE content LIKE ?"
  115. params = [f'%{query}%']
  116. if block_type:
  117. sql += " AND type = ?"
  118. params.append(block_type)
  119. sql += " ORDER BY block_order"
  120. cursor = self.conn.execute(sql, params)
  121. rows = cursor.fetchall()
  122. return [self._row_to_dict(row) for row in rows]
  123. def get_headings(self) -> list[dict]:
  124. """获取所有标题块"""
  125. cursor = self.conn.execute("""
  126. SELECT * FROM document_blocks
  127. WHERE type = 'heading'
  128. ORDER BY block_order
  129. """)
  130. rows = cursor.fetchall()
  131. return [self._row_to_dict(row) for row in rows]
  132. def get_stats(self) -> dict:
  133. """获取统计信息"""
  134. cursor = self.conn.execute("""
  135. SELECT
  136. type,
  137. COUNT(*) as count
  138. FROM document_blocks
  139. GROUP BY type
  140. """)
  141. stats = {row['type']: row['count'] for row in cursor.fetchall()}
  142. cursor = self.conn.execute("SELECT COUNT(*) as total FROM document_blocks")
  143. total = cursor.fetchone()['total']
  144. return {
  145. 'total': total,
  146. 'by_type': stats
  147. }
  148. def _build_type_filter(self, block_type: str, level: Optional[int] = None) -> tuple[str, list]:
  149. """构建类型过滤条件(用于 index 查询)"""
  150. if block_type == 'heading' and level is not None:
  151. return "type = ? AND level = ?", [block_type, level]
  152. return "type = ?", [block_type]
  153. def _query_next_value(self, field: str, after_order: int, type_filter: str, params: list) -> Optional[int]:
  154. """查询下一个值(index 或 block_order)"""
  155. sql = f'SELECT "{field}" FROM document_blocks WHERE {type_filter} AND block_order > ? ORDER BY block_order LIMIT 1'
  156. cursor = self.conn.execute(sql, params + [after_order])
  157. row = cursor.fetchone()
  158. return row[field] if row else None
  159. def _query_prev_value(self, field: str, after_order: int, type_filter: str, params: list) -> Optional[int]:
  160. """查询前一个值(index 或 block_order)"""
  161. sql = f'SELECT "{field}" FROM document_blocks WHERE {type_filter} AND block_order <= ? ORDER BY block_order DESC LIMIT 1'
  162. cursor = self.conn.execute(sql, params + [after_order])
  163. row = cursor.fetchone()
  164. return row[field] if row else None
  165. def _query_max_value(self, field: str, type_filter: str, params: list) -> Optional[int]:
  166. """查询最大值(index 或 block_order)"""
  167. sql = f'SELECT MAX("{field}") as max_val FROM document_blocks'
  168. if type_filter:
  169. sql += f' WHERE {type_filter}'
  170. cursor = self.conn.execute(sql, params)
  171. row = cursor.fetchone()
  172. return row['max_val']
  173. def _calculate_sparse_value(
  174. self,
  175. field: str,
  176. after_block_id: Optional[str],
  177. type_filter: str,
  178. params: list,
  179. default_min: int,
  180. rebalance_func: callable
  181. ) -> int:
  182. """通用稀疏值计算逻辑"""
  183. if after_block_id:
  184. after_block = self.get_block_by_id(after_block_id)
  185. if not after_block:
  186. raise ValueError(f"Block not found: {after_block_id}")
  187. after_order = after_block['block_order']
  188. next_val = self._query_next_value(field, after_order, type_filter, params)
  189. if next_val is not None:
  190. prev_val = self._query_prev_value(field, after_order, type_filter, params)
  191. if prev_val is None:
  192. prev_val = default_min
  193. gap = next_val - prev_val
  194. if gap <= 1:
  195. rebalance_func(prev_val, next_val)
  196. next_val = self._query_next_value(field, after_order, type_filter, params)
  197. if next_val is None:
  198. next_val = prev_val + 200
  199. return (prev_val + next_val) // 2
  200. else:
  201. max_val = self._query_max_value(field, type_filter, params)
  202. return (max_val if max_val is not None else default_min) + 100
  203. else:
  204. max_val = self._query_max_value(field, type_filter, params)
  205. return (max_val if max_val is not None else default_min) + 100
  206. def calculate_index_for_insert(self, after_block_id: str, new_type: str, new_level: int = 0) -> int:
  207. """计算插入 block 的 index(稀疏排序)"""
  208. type_filter, params = self._build_type_filter(new_type, new_level)
  209. return self._calculate_sparse_value(
  210. field='index',
  211. after_block_id=after_block_id,
  212. type_filter=type_filter,
  213. params=params,
  214. default_min=-100,
  215. rebalance_func=lambda prev, next: self._rebalance_indexes_between(new_type, new_level, prev, next)
  216. )
  217. def calculate_block_order_for_insert(self, after_block_id: str = None) -> int:
  218. """计算插入 block 的 block_order(稀疏排序)"""
  219. return self._calculate_sparse_value(
  220. field='block_order',
  221. after_block_id=after_block_id,
  222. type_filter='1=1',
  223. params=[],
  224. default_min=0,
  225. rebalance_func=self._rebalance_block_orders_between
  226. )
  227. def _rebalance_indexes_between(self, block_type: str, level: int, start_index: int, end_index: int):
  228. """局部重排:重新分配区间内同类型同级别 blocks 的 index"""
  229. if block_type == 'heading':
  230. cursor = self.conn.execute("""
  231. SELECT id, "index"
  232. FROM document_blocks
  233. WHERE type = ? AND level = ? AND "index" > ? AND "index" < ?
  234. ORDER BY "index"
  235. """, (block_type, level, start_index, end_index))
  236. else:
  237. cursor = self.conn.execute("""
  238. SELECT id, "index"
  239. FROM document_blocks
  240. WHERE type = ? AND "index" > ? AND "index" < ?
  241. ORDER BY "index"
  242. """, (block_type, start_index, end_index))
  243. blocks = cursor.fetchall()
  244. if not blocks:
  245. return
  246. count = len(blocks)
  247. gap = end_index - start_index
  248. step = gap // (count + 1)
  249. new_index = start_index
  250. for block in blocks:
  251. new_index += step
  252. self.conn.execute(
  253. 'UPDATE document_blocks SET "index" = ? WHERE id = ?',
  254. (new_index, block['id'])
  255. )
  256. self.conn.commit()
  257. def _rebalance_block_orders_between(self, start_order: int, end_order: int):
  258. """局部重排:重新分配区间内所有 blocks 的 block_order"""
  259. cursor = self.conn.execute("""
  260. SELECT id, block_order
  261. FROM document_blocks
  262. WHERE block_order > ? AND block_order < ?
  263. ORDER BY block_order
  264. """, (start_order, end_order))
  265. blocks = cursor.fetchall()
  266. if not blocks:
  267. return
  268. count = len(blocks)
  269. gap = end_order - start_order
  270. step = gap // (count + 1)
  271. new_order = start_order
  272. for block in blocks:
  273. new_order += step
  274. self.conn.execute(
  275. 'UPDATE document_blocks SET block_order = ? WHERE id = ?',
  276. (new_order, block['id'])
  277. )
  278. self.conn.commit()
  279. def _row_to_dict(self, row) -> dict:
  280. """将 sqlite3.Row 转换为字典"""
  281. d = dict(row)
  282. if d.get('style'):
  283. try:
  284. d['style'] = json.loads(d['style'])
  285. except (json.JSONDecodeError, TypeError):
  286. d['style'] = {}
  287. else:
  288. d['style'] = {}
  289. if d.get('metadata'):
  290. try:
  291. d['metadata'] = json.loads(d['metadata'])
  292. except (json.JSONDecodeError, TypeError):
  293. d['metadata'] = {}
  294. else:
  295. d['metadata'] = {}
  296. if d.get('content'):
  297. try:
  298. d['content'] = json.loads(d['content'])
  299. except (json.JSONDecodeError, TypeError):
  300. pass
  301. return d
  302. @staticmethod
  303. def generate_block_id(block_type: str, level: int, index: int) -> str:
  304. """生成 Block ID"""
  305. if block_type == 'heading':
  306. return f'block-h{level}-{index}'
  307. elif block_type == 'paragraph':
  308. return f'block-p-{index}'
  309. elif block_type == 'table':
  310. return f'block-table-{index}'
  311. elif block_type == 'image':
  312. return f'block-img-{index}'
  313. elif block_type == 'toc':
  314. return f'block-toc-{index}'
  315. else:
  316. return f'block-{block_type}-{index}'