store.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305
  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. created_at INTEGER NOT NULL
  56. );
  57. """
  58. def connect(db_path: Path | str = DB_PATH) -> sqlite3.Connection:
  59. """打开(必要时建)库,建表,行按 dict 取。busy_timeout 让多进程并发写不报锁。"""
  60. p = Path(db_path)
  61. p.parent.mkdir(parents=True, exist_ok=True)
  62. conn = sqlite3.connect(str(p), timeout=30)
  63. conn.row_factory = sqlite3.Row
  64. conn.execute("PRAGMA busy_timeout=30000") # 撞锁等 30s 而非报错(与微信补下载并发安全)
  65. conn.executescript(_SCHEMA)
  66. return conn
  67. def _now(ts: Optional[int]) -> int:
  68. # 调用方传入时间戳(脚本侧用 int(time.time()));缺省 0,避免本模块碰 Date/random。
  69. return int(ts) if ts is not None else 0
  70. def _record_run(conn: sqlite3.Connection, run_id: str, kind: str, note: str, ts: int) -> None:
  71. conn.execute(
  72. "INSERT OR REPLACE INTO runs(run_id, kind, note, created_at) VALUES(?,?,?,?)",
  73. (run_id, kind, note, ts),
  74. )
  75. # ---- 写入 ----------------------------------------------------------------
  76. # 哪些键是「轴」(除 query 外的结构化字段),存进 axes blob
  77. _NON_AXIS = {"query"}
  78. def insert_queries(conn: sqlite3.Connection, run_id: str, method: str,
  79. items: Iterable[dict], *, ts: Optional[int] = None) -> int:
  80. """写一批生成的 query。items=[{query, <各轴>...}]。同 (run_id,method) 先清后写(幂等)。"""
  81. t = _now(ts)
  82. conn.execute("DELETE FROM queries WHERE run_id=? AND method=?", (run_id, method))
  83. rows = []
  84. for it in items:
  85. q = (it or {}).get("query")
  86. if not q:
  87. continue
  88. axes = {k: v for k, v in it.items() if k not in _NON_AXIS}
  89. rows.append((run_id, method, q, json.dumps(axes, ensure_ascii=False), t))
  90. conn.executemany(
  91. "INSERT INTO queries(run_id, method, query, axes, created_at) VALUES(?,?,?,?,?)", rows)
  92. _record_run(conn, run_id, "queries", method, t)
  93. conn.commit()
  94. return len(rows)
  95. def _flatten_search(rec: dict) -> list[dict]:
  96. """把一条 {method, query, douyin/weixin/xiaohongshu:{ok:[...],error}} 拍成「每个结果一行」。
  97. 某渠道有结果 → 每个结果一行 ok=1;空/失败 → 一行 ok=0 记 error(保留「搜过但无结果」状态)。
  98. 只处理记录里实际出现的渠道(向后兼容只有抖音/微信的旧记录)。"""
  99. out = []
  100. method, query = rec.get("method", ""), rec.get("query", "")
  101. for platform in ("douyin", "weixin", "xiaohongshu"):
  102. if platform not in rec:
  103. continue
  104. v = rec.get(platform) or {}
  105. hits = v.get("ok") or []
  106. if hits:
  107. for h in hits:
  108. extra = {k: h[k] for k in ("nick", "images", "body_text") if k in h}
  109. out.append({"platform": platform, "ok": 1, "title": h.get("title"),
  110. "url": h.get("url"), "cover": h.get("cover"),
  111. "video": h.get("video"), "extra": extra})
  112. else:
  113. out.append({"platform": platform, "ok": 0, "extra": {"error": v.get("error") or "未搜到"}})
  114. return [{**r, "method": method, "query": query} for r in out]
  115. def insert_search_results(conn: sqlite3.Connection, run_id: str,
  116. records: Iterable[dict], *, ts: Optional[int] = None) -> int:
  117. """写一批真实搜索结果。records 为 run_search.py 产出的嵌套结构,按平台拍平入库。"""
  118. t = _now(ts)
  119. conn.execute("DELETE FROM search_results WHERE run_id=?", (run_id,))
  120. rows = []
  121. for rec in records:
  122. for r in _flatten_search(rec):
  123. rows.append((run_id, r["method"], r["query"], r["platform"], r["ok"],
  124. r.get("title"), r.get("url"), r.get("cover"), r.get("video"),
  125. json.dumps(r.get("extra", {}), ensure_ascii=False), t))
  126. conn.executemany(
  127. "INSERT INTO search_results(run_id, method, query, platform, ok, title, url, "
  128. "cover, video, extra, created_at) VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
  129. _record_run(conn, run_id, "search", f"{len(rows)} rows", t)
  130. conn.commit()
  131. return len(rows)
  132. # ---- 读取(分页 + 筛选)---------------------------------------------------
  133. def _page(conn: sqlite3.Connection, base: str, where: list[str], params: list[Any],
  134. page: int, size: int) -> dict:
  135. """通用分页:返回 {total, page, size, items}。"""
  136. size = max(1, min(int(size), 200))
  137. page = max(1, int(page))
  138. clause = (" WHERE " + " AND ".join(where)) if where else ""
  139. total = conn.execute(f"SELECT COUNT(*) FROM ({base}{clause})", params).fetchone()[0]
  140. rows = conn.execute(f"{base}{clause} ORDER BY id LIMIT ? OFFSET ?",
  141. params + [size, (page - 1) * size]).fetchall()
  142. return {"total": total, "page": page, "size": size, "items": [dict(r) for r in rows]}
  143. def list_queries(conn: sqlite3.Connection, *, method: Optional[str] = None,
  144. run_id: Optional[str] = None, page: int = 1, size: int = 30) -> dict:
  145. where, params = [], []
  146. if method:
  147. where.append("method=?"); params.append(method)
  148. if run_id:
  149. where.append("run_id=?"); params.append(run_id)
  150. res = _page(conn, "SELECT * FROM queries", where, params, page, size)
  151. for it in res["items"]:
  152. it["axes"] = json.loads(it.get("axes") or "{}")
  153. return res
  154. def list_search(conn: sqlite3.Connection, *, platform: Optional[str] = None,
  155. method: Optional[str] = None, ok: Optional[bool] = None,
  156. run_id: Optional[str] = None, query: Optional[str] = None,
  157. page: int = 1, size: int = 30) -> dict:
  158. where, params = [], []
  159. if platform:
  160. where.append("platform=?"); params.append(platform)
  161. if method:
  162. where.append("method=?"); params.append(method)
  163. if ok is not None:
  164. where.append("ok=?"); params.append(1 if ok else 0)
  165. if run_id:
  166. where.append("run_id=?"); params.append(run_id)
  167. if query:
  168. where.append("query=?"); params.append(query) # 按 query 文本取该条的全部搜索结果
  169. res = _page(conn, "SELECT * FROM search_results", where, params, page, size)
  170. cm = class_map(conn, {it.get("url") for it in res["items"] if it.get("url")})
  171. for it in res["items"]:
  172. it["extra"] = json.loads(it.get("extra") or "{}")
  173. it["cls"] = cm.get(it.get("url")) # 创作知识分类(无则 None=未分类)
  174. return res
  175. # ---- 帖子级创作知识分类 ----------------------------------------------------
  176. def get_unclassified_posts(conn: sqlite3.Connection) -> list[dict]:
  177. """库里有结果、但还没分类的去重帖子(按 url)。带 title/body_text/本地图片,供 classify 用。"""
  178. rows = conn.execute(
  179. "SELECT url, platform, title, cover, extra FROM search_results "
  180. "WHERE ok=1 AND url IS NOT NULL AND url NOT IN (SELECT url FROM post_class) "
  181. "GROUP BY url"
  182. ).fetchall()
  183. out = []
  184. for r in rows:
  185. extra = json.loads(r["extra"] or "{}")
  186. imgs = extra.get("images") or ([r["cover"]] if r["cover"] else [])
  187. out.append({"url": r["url"], "platform": r["platform"], "title": r["title"] or "",
  188. "body_text": extra.get("body_text", ""), "images": imgs})
  189. return out
  190. def upsert_class(conn: sqlite3.Connection, url: str, is_creation: int, reason: str, ts: int) -> None:
  191. conn.execute("INSERT OR REPLACE INTO post_class(url, is_creation, reason, created_at) VALUES(?,?,?,?)",
  192. (url, is_creation, reason, ts))
  193. conn.commit()
  194. def posts_to_classify(conn: sqlite3.Connection, platforms: list[str]) -> list[dict]:
  195. """指定平台的全部去重帖(ok=1),带 title/body_text/本地图/本地视频,供(重)分类。
  196. 不管是否已分类——上层 upsert 覆盖,便于换更忠实的判法重判。"""
  197. qs = ",".join("?" * len(platforms))
  198. rows = conn.execute(
  199. f"SELECT url, platform, title, cover, video, extra FROM search_results "
  200. f"WHERE ok=1 AND url IS NOT NULL AND platform IN ({qs}) GROUP BY url", platforms
  201. ).fetchall()
  202. out = []
  203. for r in rows:
  204. e = json.loads(r["extra"] or "{}")
  205. imgs = e.get("images") or ([r["cover"]] if r["cover"] else [])
  206. out.append({"url": r["url"], "platform": r["platform"], "title": r["title"] or "",
  207. "body_text": e.get("body_text", ""), "images": imgs, "video": r["video"]})
  208. return out
  209. def class_map(conn: sqlite3.Connection, urls) -> dict:
  210. """{url: {is_creation, reason}},给取数时挂分类。"""
  211. urls = [u for u in urls if u]
  212. if not urls:
  213. return {}
  214. qs = ",".join("?" * len(urls))
  215. rows = conn.execute(f"SELECT url, is_creation, reason FROM post_class WHERE url IN ({qs})", urls).fetchall()
  216. return {r["url"]: {"is_creation": r["is_creation"], "reason": r["reason"]} for r in rows}
  217. def class_counts(conn: sqlite3.Connection) -> dict:
  218. y = conn.execute("SELECT COUNT(*) FROM post_class WHERE is_creation=1").fetchone()[0]
  219. n = conn.execute("SELECT COUNT(*) FROM post_class WHERE is_creation=0").fetchone()[0]
  220. return {"creation": y, "non_creation": n}
  221. # ---- 微信正文补下载 --------------------------------------------------------
  222. def weixin_urls_missing_body(conn: sqlite3.Connection) -> list[str]:
  223. """所有微信帖(distinct url, ok=1)中 extra 还没补正文(body_text)的,供补下载。"""
  224. rows = conn.execute(
  225. "SELECT url, extra FROM search_results WHERE platform='weixin' AND ok=1 "
  226. "AND url IS NOT NULL GROUP BY url"
  227. ).fetchall()
  228. return [r["url"] for r in rows if not json.loads(r["extra"] or "{}").get("body_text")]
  229. def update_post_content(conn: sqlite3.Connection, url: str, body_text: str,
  230. images: list[str]) -> None:
  231. """把正文 + 本地图片写回该 url 的所有行的 extra(其余字段保留)。"""
  232. for r in conn.execute("SELECT id, extra FROM search_results WHERE url=?", (url,)).fetchall():
  233. e = json.loads(r["extra"] or "{}")
  234. e["body_text"] = body_text
  235. if images:
  236. e["images"] = images
  237. conn.execute("UPDATE search_results SET extra=? WHERE id=?",
  238. (json.dumps(e, ensure_ascii=False), r["id"]))
  239. conn.commit()
  240. def search_summary(conn: sqlite3.Connection) -> dict:
  241. """每条 query 搜到几个结果:{query文本: {total, ok}}。给列表页判断按钮显不显示、显几个。"""
  242. rows = conn.execute(
  243. "SELECT query, COUNT(*) AS total, SUM(ok) AS ok_n FROM search_results GROUP BY query"
  244. ).fetchall()
  245. return {r["query"]: {"total": r["total"], "ok": r["ok_n"] or 0} for r in rows}
  246. def list_runs(conn: sqlite3.Connection) -> list[dict]:
  247. rows = conn.execute("SELECT * FROM runs ORDER BY created_at DESC, run_id").fetchall()
  248. return [dict(r) for r in rows]
  249. def list_methods(conn: sqlite3.Connection, table: str = "queries") -> list[str]:
  250. """某表里有哪些 method(给前端筛选下拉)。"""
  251. if table not in ("queries", "search_results"):
  252. raise ValueError(table)
  253. rows = conn.execute(f"SELECT DISTINCT method FROM {table} ORDER BY method").fetchall()
  254. return [r[0] for r in rows]