store.py 9.2 KB

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