store.py 11 KB

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