feat: 打通多源新闻记录链路,检索/日报/知识库可见多源 (方案A+B)
- 模型层: 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
This commit is contained in:
@@ -0,0 +1,165 @@
|
||||
"""一次性回填脚本:为 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())
|
||||
|
||||
+19
-10
@@ -89,8 +89,11 @@ def _build_text_from_event(
|
||||
event_path: Path,
|
||||
processed_root: Path,
|
||||
day: str,
|
||||
) -> tuple[str, Article, str | None] | None:
|
||||
"""从 ExtractedEvent JSON 构造嵌入文本与文章元数据。"""
|
||||
) -> tuple[str, Article, str | None, list[str] | None] | None:
|
||||
"""从 ExtractedEvent JSON 构造嵌入文本与文章元数据。
|
||||
|
||||
返回 (text, article, summary, sources);sources 为 M3 多源记录。
|
||||
"""
|
||||
try:
|
||||
obj: dict[str, Any] = json.loads(event_path.read_text(encoding="utf-8"))
|
||||
except json.JSONDecodeError as e:
|
||||
@@ -98,6 +101,7 @@ def _build_text_from_event(
|
||||
return None
|
||||
|
||||
article, head, summary = _from_event_dict(obj)
|
||||
sources = obj.get("sources") or None
|
||||
# 真正的正文要去 processed/ 找
|
||||
real_article = _load_article_by_hash(processed_root, day, article.url_hash)
|
||||
if real_article is None:
|
||||
@@ -110,17 +114,21 @@ def _build_text_from_event(
|
||||
update={"publish_time": article.publish_time or real_article.publish_time}
|
||||
)
|
||||
text = compose_text(real_article, head=head, summary=summary)
|
||||
return text, real_article, summary
|
||||
return text, real_article, summary, sources
|
||||
|
||||
|
||||
def _build_text_from_article(article_path: Path) -> tuple[str, Article, str | None] | None:
|
||||
def _build_text_from_article(
|
||||
article_path: Path,
|
||||
) -> tuple[str, Article, str | None, list[str] | None] | None:
|
||||
try:
|
||||
obj = json.loads(article_path.read_text(encoding="utf-8"))
|
||||
article = Article.model_validate(obj)
|
||||
except (json.JSONDecodeError, ValidationError) as e:
|
||||
logger.warning("跳过损坏 article 文件 {}: {}", article_path, e)
|
||||
return None
|
||||
return compose_text(article), article, None
|
||||
# deduped uniques JSON 含 sources 多源字段;processed 产物无此字段 → None
|
||||
sources = obj.get("sources") or None
|
||||
return compose_text(article), article, None, sources
|
||||
|
||||
|
||||
def _collect_inputs(args: argparse.Namespace) -> list[tuple[Path, str]]:
|
||||
@@ -192,8 +200,8 @@ async def _run(args: argparse.Namespace) -> int:
|
||||
return 0
|
||||
logger.info("待嵌入文章数: {} (跳过已处理 {}; input={})", len(files), skipped, args.input)
|
||||
|
||||
# 准备每篇文本
|
||||
prepared: list[tuple[str, Article, str | None]] = []
|
||||
# 准备每篇文本: (文本, 文章, 摘要, 多源列表)
|
||||
prepared: list[tuple[str, Article, str | None, list[str] | None]] = []
|
||||
for fp, kind in files:
|
||||
if kind == "event":
|
||||
built = _build_text_from_event(fp, Path(args.processed_root), args.date)
|
||||
@@ -234,12 +242,12 @@ async def _run(args: argparse.Namespace) -> int:
|
||||
batch_size = args.batch_size
|
||||
for i in range(0, len(prepared), batch_size):
|
||||
batch = prepared[i : i + batch_size]
|
||||
texts = [t for t, _, _ in batch]
|
||||
texts = [t for t, _, _, _ in batch]
|
||||
try:
|
||||
vectors = await provider.embed_batch(texts)
|
||||
except EmbeddingError as e:
|
||||
logger.warning("批 {} 嵌入失败: {}", i // batch_size, e)
|
||||
for _, art, _ in batch:
|
||||
for _, art, _, _ in batch:
|
||||
fail_cnt += 1
|
||||
with failed_path.open("a", encoding="utf-8") as f:
|
||||
f.write(
|
||||
@@ -252,7 +260,7 @@ async def _run(args: argparse.Namespace) -> int:
|
||||
)
|
||||
continue
|
||||
|
||||
for (text, article, _summary), vec in zip(batch, vectors, strict=True):
|
||||
for (text, article, _summary, sources), vec in zip(batch, vectors, strict=True):
|
||||
if len(vec) != provider.dim:
|
||||
logger.warning(
|
||||
"维度不一致 url_hash={} 实际={} 预期={}",
|
||||
@@ -261,6 +269,7 @@ async def _run(args: argparse.Namespace) -> int:
|
||||
result = EmbeddingResult(
|
||||
url_hash=article.url_hash,
|
||||
source_id=article.source_id,
|
||||
sources=sources or [],
|
||||
title=article.title,
|
||||
text=text,
|
||||
vector=vec,
|
||||
|
||||
@@ -52,18 +52,26 @@ def _result_to_point(
|
||||
) -> dict[str, Any] | None:
|
||||
"""把 EmbeddingResult JSON dict 转换为 Qdrant Point 格式。
|
||||
|
||||
Payload 包含 title/url/source_id/publish_time/event/计数 等。
|
||||
Payload 包含 title/url/source_id/sources/publish_time/event/计数 等。
|
||||
事件字段优先从 M4 ExtractedEvent 补充(EmbeddingResult 本身不含 event)。
|
||||
sources 多源字段优先取 M5 产物,缺失时回查 M4,再兜底 [主源]。
|
||||
"""
|
||||
vector = obj.get("vector")
|
||||
if not vector:
|
||||
return None
|
||||
url_hash = obj["url_hash"]
|
||||
source_id = obj.get("source_id") or ""
|
||||
sources = obj.get("sources") or None
|
||||
if not sources and events_dir is not None:
|
||||
sources = _load_sources_from_m4(events_dir, url_hash)
|
||||
if not sources:
|
||||
sources = [source_id] if source_id else []
|
||||
payload = {
|
||||
"url_hash": url_hash,
|
||||
"title": obj.get("title") or "",
|
||||
"url": obj.get("url") or "",
|
||||
"source_id": obj.get("source_id") or "",
|
||||
"source_id": source_id,
|
||||
"sources": sources,
|
||||
"publish_time": obj.get("publish_time"),
|
||||
"char_count": obj.get("char_count"),
|
||||
"word_count": obj.get("word_count"),
|
||||
@@ -93,6 +101,21 @@ def _result_to_point(
|
||||
return {"id": url_hash, "vector": vector, "payload": payload}
|
||||
|
||||
|
||||
def _load_sources_from_m4(events_dir: Path, url_hash: str) -> list[str] | None:
|
||||
"""从 M4 ExtractedEvent JSON 读取 sources 多源字段(旧 M5 产物兜底用)。"""
|
||||
event_file = events_dir / f"{url_hash}.json"
|
||||
if not event_file.is_file():
|
||||
return None
|
||||
try:
|
||||
obj = json.loads(event_file.read_text(encoding="utf-8"))
|
||||
except (json.JSONDecodeError, OSError):
|
||||
return None
|
||||
sources = obj.get("sources")
|
||||
if isinstance(sources, list) and sources:
|
||||
return [s for s in sources if s]
|
||||
return None
|
||||
|
||||
|
||||
def _load_event_from_m4(events_dir: Path, url_hash: str) -> dict[str, Any] | None:
|
||||
"""从 M4 ExtractedEvent JSON 中提取事件 payload 子集。"""
|
||||
event_file = events_dir / f"{url_hash}.json"
|
||||
|
||||
Reference in New Issue
Block a user