"""SQLite 指纹存储。 注意:SimHash 是 64 位无符号整数,SQLite INTEGER 是 64 位有符号 (范围 [-2^63, 2^63-1])。直接存可能溢出/转负数,虽然 XOR 仍然 正确但语义混乱。这里统一存为 16 位 hex TEXT,避免符号问题。 source_ids 列存 JSON 数组文本(同一内容组全部来源);旧库无此列时 自动 ALTER TABLE 迁移,旧数据读取时回退为 [source_id]。 """ from __future__ import annotations import json import sqlite3 from datetime import datetime, timedelta from pathlib import Path from typing import Any from loguru import logger from .models import Fingerprint DEFAULT_DB_PATH = Path("data/dedup/fingerprints.sqlite3") _SCHEMA_SQL = """ CREATE TABLE IF NOT EXISTS fingerprints ( url_hash TEXT PRIMARY KEY, content_hash TEXT NOT NULL, simhash_hex TEXT NOT NULL, source_id TEXT NOT NULL, url TEXT NOT NULL, title TEXT NOT NULL, publish_date TEXT, ingested_at TEXT NOT NULL, source_ids TEXT ); CREATE INDEX IF NOT EXISTS idx_content_hash ON fingerprints(content_hash); CREATE INDEX IF NOT EXISTS idx_publish_date ON fingerprints(publish_date); CREATE INDEX IF NOT EXISTS idx_source_id ON fingerprints(source_id); """ # 兼容旧库:为已存在但缺少 source_ids 列的表补列 _ALTER_SQL = "ALTER TABLE fingerprints ADD COLUMN source_ids TEXT" def _to_hex(simhash: int) -> str: return f"{simhash:016x}" def _from_hex(hex_str: str) -> int: return int(hex_str, 16) def _to_sources_json(source_ids: list[str]) -> str: return json.dumps(source_ids, ensure_ascii=False) def _from_sources_json(raw: str | None, fallback: str) -> list[str]: """解析 source_ids 列;NULL/损坏时回退 [主源]。""" if not raw: return [fallback] try: val = json.loads(raw) except (TypeError, ValueError): return [fallback] if isinstance(val, list) and val: # 保证主源在首位(兼容手改/旧数据) cleaned = [s for s in val if s and s != fallback] return [fallback, *cleaned] return [fallback] def _row_to_fp(row: sqlite3.Row) -> Fingerprint: return Fingerprint( url_hash=row["url_hash"], content_hash=row["content_hash"], simhash=_from_hex(row["simhash_hex"]), source_id=row["source_id"], url=row["url"], title=row["title"], publish_date=row["publish_date"], ingested_at=datetime.fromisoformat(row["ingested_at"]), source_ids=_from_sources_json(row["source_ids"], row["source_id"]), ) class FingerprintStore: """SQLite 包装。线程不安全(每个线程请新建实例)。""" def __init__(self, db_path: str | Path = DEFAULT_DB_PATH) -> None: self.db_path = Path(db_path) self.db_path.parent.mkdir(parents=True, exist_ok=True) self._conn: sqlite3.Connection = sqlite3.connect( self.db_path, isolation_level=None ) self._conn.row_factory = sqlite3.Row self._conn.executescript(_SCHEMA_SQL) self._migrate_source_ids() logger.debug("打开指纹库: {}", self.db_path) def _migrate_source_ids(self) -> None: """旧库兼容:为缺少 source_ids 列的表补列(新库无需执行)。""" try: self._conn.execute(_ALTER_SQL) logger.info("指纹库迁移:为 fingerprints 表新增 source_ids 列") except sqlite3.OperationalError: logger.debug("source_ids 列已存在,跳过迁移") def close(self) -> None: self._conn.close() def __enter__(self) -> FingerprintStore: return self def __exit__(self, *_: Any) -> None: self.close() # ------------------------------------------------------------------ # # 查询 # ------------------------------------------------------------------ # def get_by_url_hash(self, url_hash: str) -> Fingerprint | None: row = self._conn.execute( "SELECT * FROM fingerprints WHERE url_hash = ?", (url_hash,) ).fetchone() return _row_to_fp(row) if row else None def find_by_content_hash(self, content_hash: str) -> Fingerprint | None: """返回任一匹配项。""" row = self._conn.execute( "SELECT * FROM fingerprints WHERE content_hash = ? LIMIT 1", (content_hash,), ).fetchone() return _row_to_fp(row) if row else None def candidates_for_simhash( self, publish_date: str | None, window_days: int, ) -> list[Fingerprint]: """返回 publish_date ± window_days 内的指纹候选。 publish_date 为 None 时,不限定窗口(返回全部,慎用)。 """ if publish_date is None or window_days < 0: rows = self._conn.execute("SELECT * FROM fingerprints").fetchall() return [_row_to_fp(r) for r in rows] try: center = datetime.strptime(publish_date, "%Y-%m-%d") except ValueError: logger.debug("publish_date 不可解析: {!r},退化为全表扫描", publish_date) rows = self._conn.execute("SELECT * FROM fingerprints").fetchall() return [_row_to_fp(r) for r in rows] lo = (center - timedelta(days=window_days)).strftime("%Y-%m-%d") hi = (center + timedelta(days=window_days)).strftime("%Y-%m-%d") rows = self._conn.execute( "SELECT * FROM fingerprints " "WHERE publish_date IS NULL OR (publish_date >= ? AND publish_date <= ?)", (lo, hi), ).fetchall() return [_row_to_fp(r) for r in rows] def count(self) -> int: return self._conn.execute("SELECT COUNT(*) FROM fingerprints").fetchone()[0] def count_by_source(self) -> dict[str, int]: rows = self._conn.execute( "SELECT source_id, COUNT(*) AS n FROM fingerprints GROUP BY source_id" ).fetchall() return {r["source_id"]: r["n"] for r in rows} def date_range(self) -> tuple[str | None, str | None]: row = self._conn.execute( "SELECT MIN(publish_date) AS lo, MAX(publish_date) AS hi FROM fingerprints" ).fetchone() return (row["lo"], row["hi"]) if row else (None, None) # ------------------------------------------------------------------ # # 写入 # ------------------------------------------------------------------ # def upsert(self, fp: Fingerprint) -> None: self._conn.execute( "INSERT OR REPLACE INTO fingerprints " "(url_hash, content_hash, simhash_hex, source_id, url, title, " " publish_date, ingested_at, source_ids) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)", ( fp.url_hash, fp.content_hash, _to_hex(fp.simhash), fp.source_id, fp.url, fp.title, fp.publish_date, fp.ingested_at.isoformat(), _to_sources_json(fp.source_ids), ), ) def delete(self, url_hash: str) -> None: self._conn.execute("DELETE FROM fingerprints WHERE url_hash = ?", (url_hash,)) def clear(self) -> None: """清空指纹库,主要用于测试。""" self._conn.execute("DELETE FROM fingerprints")