| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263 |
- """SQLite 数据层:query 生成结果 + 真实搜索结果,统一进 data/app.db。
- 为什么不再用散落 json:量大了「一次性 fetch 整个 json」撑不住,要能分页/筛选/多 run 累积。
- 设计:
- queries —— 一条生成的 query 一行;各打法字段不同,统一塞进 axes(json) 列
- search_results —— 一条「query × 平台」一行(把 douyin/weixin 嵌套拍平),方便按平台/成功筛
- 媒体文件仍躺 data/search/,库里只存相对路径(cover/video)。
- 纯 stdlib sqlite3,无新依赖。写入幂等:同 run 重导先清该 run。
- """
- from __future__ import annotations
- import json
- import sqlite3
- from pathlib import Path
- from typing import Any, Iterable, Optional
- ROOT = Path(__file__).resolve().parent.parent
- DB_PATH = ROOT / "data" / "app.db"
- _SCHEMA = """
- CREATE TABLE IF NOT EXISTS queries (
- id INTEGER PRIMARY KEY AUTOINCREMENT,
- run_id TEXT NOT NULL,
- method TEXT NOT NULL, -- 打法名(实质 × 创作阶段 × 需求 等)
- query TEXT NOT NULL,
- axes TEXT NOT NULL DEFAULT '{}', -- 该打法的正交轴 json(实质/形式/阶段…)
- created_at INTEGER NOT NULL
- );
- CREATE INDEX IF NOT EXISTS ix_queries_run ON queries(run_id);
- CREATE INDEX IF NOT EXISTS ix_queries_method ON queries(method);
- CREATE TABLE IF NOT EXISTS search_results (
- id INTEGER PRIMARY KEY AUTOINCREMENT,
- run_id TEXT NOT NULL,
- method TEXT NOT NULL,
- query TEXT NOT NULL,
- platform TEXT NOT NULL, -- douyin / weixin
- ok INTEGER NOT NULL, -- 1=搜到且有媒体, 0=失败/空
- title TEXT,
- url TEXT,
- cover TEXT, -- /data/search/... 相对路径
- video TEXT, -- 抖音才有
- extra TEXT NOT NULL DEFAULT '{}', -- nick/error 等附加 json
- created_at INTEGER NOT NULL
- );
- CREATE INDEX IF NOT EXISTS ix_sr_run ON search_results(run_id);
- CREATE INDEX IF NOT EXISTS ix_sr_platform ON search_results(platform);
- CREATE INDEX IF NOT EXISTS ix_sr_method ON search_results(method);
- CREATE TABLE IF NOT EXISTS runs (
- run_id TEXT PRIMARY KEY,
- kind TEXT NOT NULL, -- queries / search
- note TEXT,
- created_at INTEGER NOT NULL
- );
- -- 帖子级「创作知识 / 非创作知识」分类(按 url,一帖判一次)。acquisition/classify.py 写入。
- CREATE TABLE IF NOT EXISTS post_class (
- url TEXT PRIMARY KEY,
- is_creation INTEGER, -- 1=创作知识 / 0=非
- reason TEXT,
- created_at INTEGER NOT NULL
- );
- """
- def connect(db_path: Path | str = DB_PATH) -> sqlite3.Connection:
- """打开(必要时建)库,建表,行按 dict 取。"""
- p = Path(db_path)
- p.parent.mkdir(parents=True, exist_ok=True)
- conn = sqlite3.connect(str(p))
- conn.row_factory = sqlite3.Row
- conn.executescript(_SCHEMA)
- return conn
- def _now(ts: Optional[int]) -> int:
- # 调用方传入时间戳(脚本侧用 int(time.time()));缺省 0,避免本模块碰 Date/random。
- return int(ts) if ts is not None else 0
- def _record_run(conn: sqlite3.Connection, run_id: str, kind: str, note: str, ts: int) -> None:
- conn.execute(
- "INSERT OR REPLACE INTO runs(run_id, kind, note, created_at) VALUES(?,?,?,?)",
- (run_id, kind, note, ts),
- )
- # ---- 写入 ----------------------------------------------------------------
- # 哪些键是「轴」(除 query 外的结构化字段),存进 axes blob
- _NON_AXIS = {"query"}
- def insert_queries(conn: sqlite3.Connection, run_id: str, method: str,
- items: Iterable[dict], *, ts: Optional[int] = None) -> int:
- """写一批生成的 query。items=[{query, <各轴>...}]。同 (run_id,method) 先清后写(幂等)。"""
- t = _now(ts)
- conn.execute("DELETE FROM queries WHERE run_id=? AND method=?", (run_id, method))
- rows = []
- for it in items:
- q = (it or {}).get("query")
- if not q:
- continue
- axes = {k: v for k, v in it.items() if k not in _NON_AXIS}
- rows.append((run_id, method, q, json.dumps(axes, ensure_ascii=False), t))
- conn.executemany(
- "INSERT INTO queries(run_id, method, query, axes, created_at) VALUES(?,?,?,?,?)", rows)
- _record_run(conn, run_id, "queries", method, t)
- conn.commit()
- return len(rows)
- def _flatten_search(rec: dict) -> list[dict]:
- """把一条 {method, query, douyin/weixin/xiaohongshu:{ok:[...],error}} 拍成「每个结果一行」。
- 某渠道有结果 → 每个结果一行 ok=1;空/失败 → 一行 ok=0 记 error(保留「搜过但无结果」状态)。
- 只处理记录里实际出现的渠道(向后兼容只有抖音/微信的旧记录)。"""
- out = []
- method, query = rec.get("method", ""), rec.get("query", "")
- for platform in ("douyin", "weixin", "xiaohongshu"):
- if platform not in rec:
- continue
- v = rec.get(platform) or {}
- hits = v.get("ok") or []
- if hits:
- for h in hits:
- extra = {k: h[k] for k in ("nick", "images", "body_text") if k in h}
- out.append({"platform": platform, "ok": 1, "title": h.get("title"),
- "url": h.get("url"), "cover": h.get("cover"),
- "video": h.get("video"), "extra": extra})
- else:
- out.append({"platform": platform, "ok": 0, "extra": {"error": v.get("error") or "未搜到"}})
- return [{**r, "method": method, "query": query} for r in out]
- def insert_search_results(conn: sqlite3.Connection, run_id: str,
- records: Iterable[dict], *, ts: Optional[int] = None) -> int:
- """写一批真实搜索结果。records 为 run_search.py 产出的嵌套结构,按平台拍平入库。"""
- t = _now(ts)
- conn.execute("DELETE FROM search_results WHERE run_id=?", (run_id,))
- rows = []
- for rec in records:
- for r in _flatten_search(rec):
- rows.append((run_id, r["method"], r["query"], r["platform"], r["ok"],
- r.get("title"), r.get("url"), r.get("cover"), r.get("video"),
- json.dumps(r.get("extra", {}), ensure_ascii=False), t))
- conn.executemany(
- "INSERT INTO search_results(run_id, method, query, platform, ok, title, url, "
- "cover, video, extra, created_at) VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
- _record_run(conn, run_id, "search", f"{len(rows)} rows", t)
- conn.commit()
- return len(rows)
- # ---- 读取(分页 + 筛选)---------------------------------------------------
- def _page(conn: sqlite3.Connection, base: str, where: list[str], params: list[Any],
- page: int, size: int) -> dict:
- """通用分页:返回 {total, page, size, items}。"""
- size = max(1, min(int(size), 200))
- page = max(1, int(page))
- clause = (" WHERE " + " AND ".join(where)) if where else ""
- total = conn.execute(f"SELECT COUNT(*) FROM ({base}{clause})", params).fetchone()[0]
- rows = conn.execute(f"{base}{clause} ORDER BY id LIMIT ? OFFSET ?",
- params + [size, (page - 1) * size]).fetchall()
- return {"total": total, "page": page, "size": size, "items": [dict(r) for r in rows]}
- def list_queries(conn: sqlite3.Connection, *, method: Optional[str] = None,
- run_id: Optional[str] = None, page: int = 1, size: int = 30) -> dict:
- where, params = [], []
- if method:
- where.append("method=?"); params.append(method)
- if run_id:
- where.append("run_id=?"); params.append(run_id)
- res = _page(conn, "SELECT * FROM queries", where, params, page, size)
- for it in res["items"]:
- it["axes"] = json.loads(it.get("axes") or "{}")
- return res
- def list_search(conn: sqlite3.Connection, *, platform: Optional[str] = None,
- method: Optional[str] = None, ok: Optional[bool] = None,
- run_id: Optional[str] = None, query: Optional[str] = None,
- page: int = 1, size: int = 30) -> dict:
- where, params = [], []
- if platform:
- where.append("platform=?"); params.append(platform)
- if method:
- where.append("method=?"); params.append(method)
- if ok is not None:
- where.append("ok=?"); params.append(1 if ok else 0)
- if run_id:
- where.append("run_id=?"); params.append(run_id)
- if query:
- where.append("query=?"); params.append(query) # 按 query 文本取该条的全部搜索结果
- res = _page(conn, "SELECT * FROM search_results", where, params, page, size)
- cm = class_map(conn, {it.get("url") for it in res["items"] if it.get("url")})
- for it in res["items"]:
- it["extra"] = json.loads(it.get("extra") or "{}")
- it["cls"] = cm.get(it.get("url")) # 创作知识分类(无则 None=未分类)
- return res
- # ---- 帖子级创作知识分类 ----------------------------------------------------
- def get_unclassified_posts(conn: sqlite3.Connection) -> list[dict]:
- """库里有结果、但还没分类的去重帖子(按 url)。带 title/body_text/本地图片,供 classify 用。"""
- rows = conn.execute(
- "SELECT url, platform, title, cover, extra FROM search_results "
- "WHERE ok=1 AND url IS NOT NULL AND url NOT IN (SELECT url FROM post_class) "
- "GROUP BY url"
- ).fetchall()
- out = []
- for r in rows:
- extra = json.loads(r["extra"] or "{}")
- imgs = extra.get("images") or ([r["cover"]] if r["cover"] else [])
- out.append({"url": r["url"], "platform": r["platform"], "title": r["title"] or "",
- "body_text": extra.get("body_text", ""), "images": imgs})
- return out
- def upsert_class(conn: sqlite3.Connection, url: str, is_creation: int, reason: str, ts: int) -> None:
- conn.execute("INSERT OR REPLACE INTO post_class(url, is_creation, reason, created_at) VALUES(?,?,?,?)",
- (url, is_creation, reason, ts))
- conn.commit()
- def class_map(conn: sqlite3.Connection, urls) -> dict:
- """{url: {is_creation, reason}},给取数时挂分类。"""
- urls = [u for u in urls if u]
- if not urls:
- return {}
- qs = ",".join("?" * len(urls))
- rows = conn.execute(f"SELECT url, is_creation, reason FROM post_class WHERE url IN ({qs})", urls).fetchall()
- return {r["url"]: {"is_creation": r["is_creation"], "reason": r["reason"]} for r in rows}
- def class_counts(conn: sqlite3.Connection) -> dict:
- y = conn.execute("SELECT COUNT(*) FROM post_class WHERE is_creation=1").fetchone()[0]
- n = conn.execute("SELECT COUNT(*) FROM post_class WHERE is_creation=0").fetchone()[0]
- return {"creation": y, "non_creation": n}
- def search_summary(conn: sqlite3.Connection) -> dict:
- """每条 query 搜到几个结果:{query文本: {total, ok}}。给列表页判断按钮显不显示、显几个。"""
- rows = conn.execute(
- "SELECT query, COUNT(*) AS total, SUM(ok) AS ok_n FROM search_results GROUP BY query"
- ).fetchall()
- return {r["query"]: {"total": r["total"], "ok": r["ok_n"] or 0} for r in rows}
- def list_runs(conn: sqlite3.Connection) -> list[dict]:
- rows = conn.execute("SELECT * FROM runs ORDER BY created_at DESC, run_id").fetchall()
- return [dict(r) for r in rows]
- def list_methods(conn: sqlite3.Connection, table: str = "queries") -> list[str]:
- """某表里有哪些 method(给前端筛选下拉)。"""
- if table not in ("queries", "search_results"):
- raise ValueError(table)
- rows = conn.execute(f"SELECT DISTINCT method FROM {table} ORDER BY method").fetchall()
- return [r[0] for r in rows]
|