store.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327
  1. """SQLite 数据层:query 生成结果 + 真实搜索结果,统一进 data/app.db。
  2. 为什么不再用散落 json:量大了「一次性 fetch 整个 json」撑不住,要能分页/筛选/多 run 累积。
  3. 设计:
  4. queries —— 一条生成的 query 一行;各打法字段不同,统一塞进 axes(json) 列
  5. search_results —— 一条「query × 平台」一行(把 douyin/weixin 嵌套拍平),方便按平台/成功筛
  6. 媒体文件仍躺 data/search/,库里只存相对路径(cover/video)。
  7. 纯 stdlib sqlite3,无新依赖。写入幂等:同 run 重导先清该 run。
  8. """
  9. from __future__ import annotations
  10. import json
  11. import sqlite3
  12. from pathlib import Path
  13. from typing import Any, Iterable, Optional
  14. ROOT = Path(__file__).resolve().parent.parent
  15. DB_PATH = ROOT / "data" / "app.db"
  16. _SCHEMA = """
  17. CREATE TABLE IF NOT EXISTS queries (
  18. id INTEGER PRIMARY KEY AUTOINCREMENT,
  19. run_id TEXT NOT NULL,
  20. method TEXT NOT NULL, -- 打法名(实质 × 创作阶段 × 需求 等)
  21. query TEXT NOT NULL,
  22. axes TEXT NOT NULL DEFAULT '{}', -- 该打法的正交轴 json(实质/形式/阶段…)
  23. created_at INTEGER NOT NULL
  24. );
  25. CREATE INDEX IF NOT EXISTS ix_queries_run ON queries(run_id);
  26. CREATE INDEX IF NOT EXISTS ix_queries_method ON queries(method);
  27. CREATE TABLE IF NOT EXISTS search_results (
  28. id INTEGER PRIMARY KEY AUTOINCREMENT,
  29. run_id TEXT NOT NULL,
  30. method TEXT NOT NULL,
  31. query TEXT NOT NULL,
  32. platform TEXT NOT NULL, -- douyin / weixin
  33. ok INTEGER NOT NULL, -- 1=搜到且有媒体, 0=失败/空
  34. title TEXT,
  35. url TEXT,
  36. cover TEXT, -- /data/search/... 相对路径
  37. video TEXT, -- 抖音才有
  38. extra TEXT NOT NULL DEFAULT '{}', -- nick/error 等附加 json
  39. created_at INTEGER NOT NULL
  40. );
  41. CREATE INDEX IF NOT EXISTS ix_sr_run ON search_results(run_id);
  42. CREATE INDEX IF NOT EXISTS ix_sr_platform ON search_results(platform);
  43. CREATE INDEX IF NOT EXISTS ix_sr_method ON search_results(method);
  44. CREATE TABLE IF NOT EXISTS runs (
  45. run_id TEXT PRIMARY KEY,
  46. kind TEXT NOT NULL, -- queries / search
  47. note TEXT,
  48. created_at INTEGER NOT NULL
  49. );
  50. -- 帖子级「创作知识 / 非创作知识」分类(按 url,一帖判一次)。acquisition/classify.py 写入。
  51. CREATE TABLE IF NOT EXISTS post_class (
  52. url TEXT PRIMARY KEY,
  53. is_creation INTEGER, -- 1=创作知识 / 0=非
  54. reason TEXT,
  55. knowledge TEXT, -- 提取出的完整创作知识(创作帖才有)
  56. points TEXT, -- 结构化知识点 json(按卡片/分段)
  57. created_at INTEGER NOT NULL
  58. );
  59. """
  60. # 旧库迁移:post_class 早期没有 knowledge/points 列,补上(已存在则忽略)
  61. _MIGRATIONS = [
  62. "ALTER TABLE post_class ADD COLUMN knowledge TEXT",
  63. "ALTER TABLE post_class ADD COLUMN points TEXT",
  64. ]
  65. def connect(db_path: Path | str = DB_PATH) -> sqlite3.Connection:
  66. """打开(必要时建)库,建表,行按 dict 取。busy_timeout 让多进程并发写不报锁。"""
  67. p = Path(db_path)
  68. p.parent.mkdir(parents=True, exist_ok=True)
  69. conn = sqlite3.connect(str(p), timeout=30)
  70. conn.row_factory = sqlite3.Row
  71. conn.execute("PRAGMA busy_timeout=30000") # 撞锁等 30s 而非报错(与微信补下载并发安全)
  72. conn.executescript(_SCHEMA)
  73. for sql in _MIGRATIONS:
  74. try:
  75. conn.execute(sql)
  76. except sqlite3.OperationalError:
  77. pass # 列已存在
  78. return conn
  79. def _now(ts: Optional[int]) -> int:
  80. # 调用方传入时间戳(脚本侧用 int(time.time()));缺省 0,避免本模块碰 Date/random。
  81. return int(ts) if ts is not None else 0
  82. def _record_run(conn: sqlite3.Connection, run_id: str, kind: str, note: str, ts: int) -> None:
  83. conn.execute(
  84. "INSERT OR REPLACE INTO runs(run_id, kind, note, created_at) VALUES(?,?,?,?)",
  85. (run_id, kind, note, ts),
  86. )
  87. # ---- 写入 ----------------------------------------------------------------
  88. # 哪些键是「轴」(除 query 外的结构化字段),存进 axes blob
  89. _NON_AXIS = {"query"}
  90. def insert_queries(conn: sqlite3.Connection, run_id: str, method: str,
  91. items: Iterable[dict], *, ts: Optional[int] = None) -> int:
  92. """写一批生成的 query。items=[{query, <各轴>...}]。同 (run_id,method) 先清后写(幂等)。"""
  93. t = _now(ts)
  94. conn.execute("DELETE FROM queries WHERE run_id=? AND method=?", (run_id, method))
  95. rows = []
  96. for it in items:
  97. q = (it or {}).get("query")
  98. if not q:
  99. continue
  100. axes = {k: v for k, v in it.items() if k not in _NON_AXIS}
  101. rows.append((run_id, method, q, json.dumps(axes, ensure_ascii=False), t))
  102. conn.executemany(
  103. "INSERT INTO queries(run_id, method, query, axes, created_at) VALUES(?,?,?,?,?)", rows)
  104. _record_run(conn, run_id, "queries", method, t)
  105. conn.commit()
  106. return len(rows)
  107. def _flatten_search(rec: dict) -> list[dict]:
  108. """把一条 {method, query, douyin/weixin/xiaohongshu:{ok:[...],error}} 拍成「每个结果一行」。
  109. 某渠道有结果 → 每个结果一行 ok=1;空/失败 → 一行 ok=0 记 error(保留「搜过但无结果」状态)。
  110. 只处理记录里实际出现的渠道(向后兼容只有抖音/微信的旧记录)。"""
  111. out = []
  112. method, query = rec.get("method", ""), rec.get("query", "")
  113. for platform in ("douyin", "weixin", "xiaohongshu"):
  114. if platform not in rec:
  115. continue
  116. v = rec.get(platform) or {}
  117. hits = v.get("ok") or []
  118. if hits:
  119. for h in hits:
  120. extra = {k: h[k] for k in ("nick", "images", "body_text") if k in h}
  121. out.append({"platform": platform, "ok": 1, "title": h.get("title"),
  122. "url": h.get("url"), "cover": h.get("cover"),
  123. "video": h.get("video"), "extra": extra})
  124. else:
  125. out.append({"platform": platform, "ok": 0, "extra": {"error": v.get("error") or "未搜到"}})
  126. return [{**r, "method": method, "query": query} for r in out]
  127. def insert_search_results(conn: sqlite3.Connection, run_id: str,
  128. records: Iterable[dict], *, ts: Optional[int] = None) -> int:
  129. """写一批真实搜索结果。records 为 run_search.py 产出的嵌套结构,按平台拍平入库。"""
  130. t = _now(ts)
  131. conn.execute("DELETE FROM search_results WHERE run_id=?", (run_id,))
  132. rows = []
  133. for rec in records:
  134. for r in _flatten_search(rec):
  135. rows.append((run_id, r["method"], r["query"], r["platform"], r["ok"],
  136. r.get("title"), r.get("url"), r.get("cover"), r.get("video"),
  137. json.dumps(r.get("extra", {}), ensure_ascii=False), t))
  138. conn.executemany(
  139. "INSERT INTO search_results(run_id, method, query, platform, ok, title, url, "
  140. "cover, video, extra, created_at) VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
  141. _record_run(conn, run_id, "search", f"{len(rows)} rows", t)
  142. conn.commit()
  143. return len(rows)
  144. # ---- 读取(分页 + 筛选)---------------------------------------------------
  145. def _page(conn: sqlite3.Connection, base: str, where: list[str], params: list[Any],
  146. page: int, size: int) -> dict:
  147. """通用分页:返回 {total, page, size, items}。"""
  148. size = max(1, min(int(size), 200))
  149. page = max(1, int(page))
  150. clause = (" WHERE " + " AND ".join(where)) if where else ""
  151. total = conn.execute(f"SELECT COUNT(*) FROM ({base}{clause})", params).fetchone()[0]
  152. rows = conn.execute(f"{base}{clause} ORDER BY id LIMIT ? OFFSET ?",
  153. params + [size, (page - 1) * size]).fetchall()
  154. return {"total": total, "page": page, "size": size, "items": [dict(r) for r in rows]}
  155. def list_queries(conn: sqlite3.Connection, *, method: Optional[str] = None,
  156. run_id: Optional[str] = None, page: int = 1, size: int = 30) -> dict:
  157. where, params = [], []
  158. if method:
  159. where.append("method=?"); params.append(method)
  160. if run_id:
  161. where.append("run_id=?"); params.append(run_id)
  162. res = _page(conn, "SELECT * FROM queries", where, params, page, size)
  163. for it in res["items"]:
  164. it["axes"] = json.loads(it.get("axes") or "{}")
  165. return res
  166. def list_search(conn: sqlite3.Connection, *, platform: Optional[str] = None,
  167. method: Optional[str] = None, ok: Optional[bool] = None,
  168. run_id: Optional[str] = None, query: Optional[str] = None,
  169. page: int = 1, size: int = 30) -> dict:
  170. where, params = [], []
  171. if platform:
  172. where.append("platform=?"); params.append(platform)
  173. if method:
  174. where.append("method=?"); params.append(method)
  175. if ok is not None:
  176. where.append("ok=?"); params.append(1 if ok else 0)
  177. if run_id:
  178. where.append("run_id=?"); params.append(run_id)
  179. if query:
  180. where.append("query=?"); params.append(query) # 按 query 文本取该条的全部搜索结果
  181. res = _page(conn, "SELECT * FROM search_results", where, params, page, size)
  182. cm = class_map(conn, {it.get("url") for it in res["items"] if it.get("url")})
  183. for it in res["items"]:
  184. it["extra"] = json.loads(it.get("extra") or "{}")
  185. it["cls"] = cm.get(it.get("url")) # 创作知识分类(无则 None=未分类)
  186. return res
  187. # ---- 帖子级创作知识分类 ----------------------------------------------------
  188. def get_unclassified_posts(conn: sqlite3.Connection) -> list[dict]:
  189. """库里有结果、但还没分类的去重帖子(按 url)。带 title/body_text/本地图片,供 classify 用。"""
  190. rows = conn.execute(
  191. "SELECT url, platform, title, cover, extra FROM search_results "
  192. "WHERE ok=1 AND url IS NOT NULL AND url NOT IN (SELECT url FROM post_class) "
  193. "GROUP BY url"
  194. ).fetchall()
  195. out = []
  196. for r in rows:
  197. extra = json.loads(r["extra"] or "{}")
  198. imgs = extra.get("images") or ([r["cover"]] if r["cover"] else [])
  199. out.append({"url": r["url"], "platform": r["platform"], "title": r["title"] or "",
  200. "body_text": extra.get("body_text", ""), "images": imgs})
  201. return out
  202. def upsert_class(conn: sqlite3.Connection, url: str, is_creation: int, reason: str, ts: int,
  203. knowledge: str = "", points: str = "") -> None:
  204. conn.execute(
  205. "INSERT OR REPLACE INTO post_class(url, is_creation, reason, knowledge, points, created_at) "
  206. "VALUES(?,?,?,?,?,?)", (url, is_creation, reason, knowledge, points, ts))
  207. conn.commit()
  208. def posts_to_classify(conn: sqlite3.Connection, platforms: list[str]) -> list[dict]:
  209. """指定平台的全部去重帖(ok=1),带 title/body_text/本地图/本地视频,供(重)分类。
  210. 不管是否已分类——上层 upsert 覆盖,便于换更忠实的判法重判。"""
  211. qs = ",".join("?" * len(platforms))
  212. rows = conn.execute(
  213. f"SELECT url, platform, title, cover, video, extra FROM search_results "
  214. f"WHERE ok=1 AND url IS NOT NULL AND platform IN ({qs}) GROUP BY url", platforms
  215. ).fetchall()
  216. out = []
  217. for r in rows:
  218. e = json.loads(r["extra"] or "{}")
  219. imgs = e.get("images") or ([r["cover"]] if r["cover"] else [])
  220. out.append({"url": r["url"], "platform": r["platform"], "title": r["title"] or "",
  221. "body_text": e.get("body_text", ""), "images": imgs, "video": r["video"]})
  222. return out
  223. def class_map(conn: sqlite3.Connection, urls) -> dict:
  224. """{url: {is_creation, reason, knowledge, points}},给取数时挂分类 + 提取的知识点。"""
  225. urls = [u for u in urls if u]
  226. if not urls:
  227. return {}
  228. qs = ",".join("?" * len(urls))
  229. rows = conn.execute(
  230. f"SELECT url, is_creation, reason, knowledge, points FROM post_class WHERE url IN ({qs})", urls
  231. ).fetchall()
  232. out = {}
  233. for r in rows:
  234. out[r["url"]] = {"is_creation": r["is_creation"], "reason": r["reason"],
  235. "knowledge": r["knowledge"] or "",
  236. "points": json.loads(r["points"]) if r["points"] else []}
  237. return out
  238. def class_counts(conn: sqlite3.Connection) -> dict:
  239. y = conn.execute("SELECT COUNT(*) FROM post_class WHERE is_creation=1").fetchone()[0]
  240. n = conn.execute("SELECT COUNT(*) FROM post_class WHERE is_creation=0").fetchone()[0]
  241. return {"creation": y, "non_creation": n}
  242. # ---- 微信正文补下载 --------------------------------------------------------
  243. def weixin_urls_missing_body(conn: sqlite3.Connection) -> list[str]:
  244. """所有微信帖(distinct url, ok=1)中 extra 还没补正文(body_text)的,供补下载。"""
  245. rows = conn.execute(
  246. "SELECT url, extra FROM search_results WHERE platform='weixin' AND ok=1 "
  247. "AND url IS NOT NULL GROUP BY url"
  248. ).fetchall()
  249. return [r["url"] for r in rows if not json.loads(r["extra"] or "{}").get("body_text")]
  250. def update_post_content(conn: sqlite3.Connection, url: str, body_text: str,
  251. images: list[str]) -> None:
  252. """把正文 + 本地图片写回该 url 的所有行的 extra(其余字段保留)。"""
  253. for r in conn.execute("SELECT id, extra FROM search_results WHERE url=?", (url,)).fetchall():
  254. e = json.loads(r["extra"] or "{}")
  255. e["body_text"] = body_text
  256. if images:
  257. e["images"] = images
  258. conn.execute("UPDATE search_results SET extra=? WHERE id=?",
  259. (json.dumps(e, ensure_ascii=False), r["id"]))
  260. conn.commit()
  261. def search_summary(conn: sqlite3.Connection) -> dict:
  262. """每条 query 搜到几个结果:{query文本: {total, ok}}。给列表页判断按钮显不显示、显几个。"""
  263. rows = conn.execute(
  264. "SELECT query, COUNT(*) AS total, SUM(ok) AS ok_n FROM search_results GROUP BY query"
  265. ).fetchall()
  266. return {r["query"]: {"total": r["total"], "ok": r["ok_n"] or 0} for r in rows}
  267. def list_runs(conn: sqlite3.Connection) -> list[dict]:
  268. rows = conn.execute("SELECT * FROM runs ORDER BY created_at DESC, run_id").fetchall()
  269. return [dict(r) for r in rows]
  270. def list_methods(conn: sqlite3.Connection, table: str = "queries") -> list[str]:
  271. """某表里有哪些 method(给前端筛选下拉)。"""
  272. if table not in ("queries", "search_results"):
  273. raise ValueError(table)
  274. rows = conn.execute(f"SELECT DISTINCT method FROM {table} ORDER BY method").fetchall()
  275. return [r[0] for r in rows]