- 模型层: EmbeddingResult/SearchResult 新增 sources 字段(主源居首,旧产物兜底) - M5 run_embedding: events/deduped 产物透传 sources 进 EmbeddingResult - M6 run_qdrant_ingest: payload 写入 sources(M5 → 回查 M4 → 兜底 [主源]) - vectorstore: 检索读取 payload.sources - 日报 HTML / CLI search / MCP: 多源显示「财联社 / 新浪 [多源]」 - 新增 scripts/backfill_qdrant_sources.py: 指纹库为权威源,scroll+upsert 回填存量 (本地模式 set_payload 逐点 0.65s 不可行,改走 ingest 同款快速路径) - 新增 tests/test_multisource.py 10 个;全量 276 passed
166 lines
5.4 KiB
Python
166 lines
5.4 KiB
Python
"""一次性回填脚本:为 Qdrant 存量 point 的 payload 补充 sources 多源字段。
|
|
|
|
背景:
|
|
M5/M6 旧链路未把 M3 多源记录写入 EmbeddingResult / Qdrant payload,
|
|
导致知识库检索只能看到单源。新链路修复后,此脚本为存量数据补齐。
|
|
|
|
数据来源(按优先级):
|
|
1. 指纹库 data/dedup/fingerprints.sqlite3 的 source_ids 列(权威多源记录);
|
|
2. source_ids 为 NULL 的旧行回退 [source_id];
|
|
3. 指纹库查不到的 point 回退 payload.source_id。
|
|
|
|
实现:
|
|
本地文件模式 Qdrant 的 set_payload 逐点极慢(≈0.65s/点,33k 条约 6 小时),
|
|
故采用与正常 ingest 相同的快速路径:
|
|
scroll 分页取全部 point(含向量)→ 合并 sources → 批量 upsert 回写。
|
|
|
|
注意:
|
|
Qdrant 为本地文件模式,同一时刻仅允许一个进程打开;
|
|
执行前请确认无 pipeline 子进程正在运行(避开调度窗口)。
|
|
|
|
用法:
|
|
uv run python -m scripts.backfill_qdrant_sources # 实际执行
|
|
uv run python -m scripts.backfill_qdrant_sources --dry-run # 只统计不写入
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import sqlite3
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
from loguru import logger
|
|
from qdrant_client.http.models import PointStruct
|
|
|
|
from vectorstore import DEFAULT_COLLECTION, make_qdrant_client
|
|
|
|
_SCROLL_LIMIT = 1000
|
|
_UPSERT_BATCH = 100
|
|
|
|
|
|
def _setup_logger(level: str) -> None:
|
|
logger.remove()
|
|
logger.add(
|
|
sys.stderr,
|
|
level=level,
|
|
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {name} | {message}",
|
|
)
|
|
|
|
|
|
def _load_fingerprint_sources(db_path: Path) -> dict[str, list[str]]:
|
|
"""从指纹库读 url_hash -> sources(source_ids 为 NULL 时回退 [source_id])。"""
|
|
conn = sqlite3.connect(db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
rows = conn.execute(
|
|
"SELECT url_hash, source_id, source_ids FROM fingerprints"
|
|
).fetchall()
|
|
conn.close()
|
|
|
|
result: dict[str, list[str]] = {}
|
|
for r in rows:
|
|
sources: list[str] = []
|
|
if r["source_ids"]:
|
|
try:
|
|
val = json.loads(r["source_ids"])
|
|
if isinstance(val, list):
|
|
sources = [s for s in val if s]
|
|
except (TypeError, ValueError):
|
|
sources = []
|
|
if not sources:
|
|
sources = [r["source_id"]] if r["source_id"] else []
|
|
# 保证主源居首
|
|
if r["source_id"] and r["source_id"] in sources:
|
|
sources = [r["source_id"], *[s for s in sources if s != r["source_id"]]]
|
|
if sources:
|
|
result[r["url_hash"]] = sources
|
|
return result
|
|
|
|
|
|
def _build_points(
|
|
records: list,
|
|
fp_sources: dict[str, list[str]],
|
|
) -> list[PointStruct]:
|
|
"""把 scroll 记录转回 PointStruct,payload 补入 sources 字段。"""
|
|
points: list[PointStruct] = []
|
|
for rec in records:
|
|
payload = dict(rec.payload or {})
|
|
url_hash = payload.get("url_hash") or ""
|
|
source_id = payload.get("source_id") or ""
|
|
sources = fp_sources.get(url_hash) or ([source_id] if source_id else [])
|
|
if sources:
|
|
payload["sources"] = sources
|
|
points.append(PointStruct(
|
|
id=rec.id,
|
|
vector=rec.vector,
|
|
payload=payload,
|
|
))
|
|
return points
|
|
|
|
|
|
def main() -> int:
|
|
parser = argparse.ArgumentParser(description="Qdrant 存量 payload 回填 sources 多源字段")
|
|
parser.add_argument("--db", default="data/dedup/fingerprints.sqlite3",
|
|
help="指纹库路径")
|
|
parser.add_argument("--collection", default=DEFAULT_COLLECTION)
|
|
parser.add_argument("--dry-run", action="store_true", help="只统计不写入")
|
|
parser.add_argument("--log-level", default="INFO")
|
|
args = parser.parse_args()
|
|
|
|
_setup_logger(args.log_level)
|
|
|
|
fp_sources = _load_fingerprint_sources(Path(args.db))
|
|
logger.info("指纹库加载完成: {} 条,其中多源 {} 条",
|
|
len(fp_sources), sum(1 for v in fp_sources.values() if len(v) > 1))
|
|
|
|
client = make_qdrant_client()
|
|
|
|
# scroll 分页 + (非 dry-run) 批量 upsert 回写
|
|
total = 0
|
|
multi_cnt = 0
|
|
updated = 0
|
|
offset = None
|
|
while True:
|
|
records, offset = client.scroll(
|
|
collection_name=args.collection,
|
|
limit=_SCROLL_LIMIT,
|
|
offset=offset,
|
|
with_payload=True,
|
|
with_vectors=True,
|
|
)
|
|
if not records:
|
|
if offset is None:
|
|
break
|
|
continue
|
|
total += len(records)
|
|
multi_cnt += sum(
|
|
1 for r in records
|
|
if len((r.payload or {}).get("sources") or []) > 1
|
|
)
|
|
if not args.dry_run:
|
|
points = _build_points(records, fp_sources)
|
|
for i in range(0, len(points), _UPSERT_BATCH):
|
|
client.upsert(
|
|
collection_name=args.collection,
|
|
points=points[i : i + _UPSERT_BATCH],
|
|
)
|
|
updated += len(points)
|
|
if offset is None:
|
|
break
|
|
|
|
logger.info("Qdrant 现存 point: {} 条(其中 payload 已含多源 {} 条)", total, multi_cnt)
|
|
if args.dry_run:
|
|
logger.info("--dry-run: 不执行写入")
|
|
client.close()
|
|
return 0
|
|
|
|
logger.info("回填完成: 重新 upsert {} 条 point(payload 已补 sources)", updated)
|
|
client.close()
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|
|
|