Files
news/tests/test_multisource.py
T
simon 80828310d6 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
2026-08-22 22:26:10 +08:00

218 lines
6.6 KiB
Python

"""多源记录端到端链路测试 (方案 A + B)。
覆盖:
- EmbeddingResult.sources validator:主源居首 / 去重 / 旧产物兜底
- run_qdrant_ingest._result_to_point:payload 写入 sources(
M5 产物优先 → 回查 M4 → 兜底 [主源])
- run_embedding 文本构造透传 sources
- 回填脚本 _load_fingerprint_sources:source_ids 解析与兜底
"""
from __future__ import annotations
import json
import sqlite3
from pathlib import Path
from embedding import EmbeddingResult
# --------------------------------------------------------------------------- #
# EmbeddingResult.sources validator
# --------------------------------------------------------------------------- #
def _emb(**overrides) -> EmbeddingResult:
base = {
"url_hash": "a" * 16,
"source_id": "cls",
"title": "t",
"text": "x",
"vector": [0.1] * 4,
"dim": 4,
"provider": "dashscope",
"model": "m",
}
base.update(overrides)
return EmbeddingResult(**base)
def test_embedding_sources_default_fallback_to_main() -> None:
"""旧产物无 sources → 兜底为 [主源]。"""
r = _emb()
assert r.sources == ["cls"]
def test_embedding_sources_main_source_first_and_dedup() -> None:
"""主源居首且去重保序。"""
r = _emb(sources=["sina", "cls", "eastmoney"])
assert r.sources == ["cls", "sina", "eastmoney"]
def test_embedding_sources_keeps_multi() -> None:
r = _emb(sources=["cls", "eastmoney"])
assert r.sources == ["cls", "eastmoney"]
# --------------------------------------------------------------------------- #
# run_qdrant_ingest._result_to_point
# --------------------------------------------------------------------------- #
def test_result_to_point_sources_from_m5() -> None:
"""M5 产物自带 sources → 直接写入 payload。"""
from scripts.run_qdrant_ingest import _result_to_point
obj = {
"url_hash": "b" * 16,
"vector": [0.1] * 4,
"title": "t",
"url": "https://x",
"source_id": "cls",
"sources": ["cls", "eastmoney"],
}
pt = _result_to_point(obj)
assert pt is not None
assert pt["payload"]["sources"] == ["cls", "eastmoney"]
def test_result_to_point_sources_fallback_m4(tmp_path: Path) -> None:
"""M5 产物无 sources → 回查 M4 events 的 sources。"""
from scripts.run_qdrant_ingest import _result_to_point
h = "c" * 16
ev_dir = tmp_path / "events"
ev_dir.mkdir()
(ev_dir / f"{h}.json").write_text(
json.dumps({"sources": ["sina", "yicai"], "event": {}}),
encoding="utf-8",
)
obj = {
"url_hash": h,
"vector": [0.1] * 4,
"title": "t",
"url": "https://x",
"source_id": "sina",
}
pt = _result_to_point(obj, events_dir=ev_dir)
assert pt is not None
assert pt["payload"]["sources"] == ["sina", "yicai"]
def test_result_to_point_sources_fallback_main_only() -> None:
"""M5/M4 均无 sources → 兜底 [主源]。"""
from scripts.run_qdrant_ingest import _result_to_point
obj = {
"url_hash": "d" * 16,
"vector": [0.1] * 4,
"title": "t",
"url": "https://x",
"source_id": "zqrb",
}
pt = _result_to_point(obj)
assert pt is not None
assert pt["payload"]["sources"] == ["zqrb"]
# --------------------------------------------------------------------------- #
# run_embedding 文本构造透传 sources
# --------------------------------------------------------------------------- #
def test_build_text_from_event_returns_sources(tmp_path: Path) -> None:
from scripts.run_embedding import _build_text_from_event
h = "e" * 16
ev_dir = tmp_path / "events"
ev_dir.mkdir()
proc_root = tmp_path / "processed"
proc_root.mkdir()
ev = {
"url_hash": h,
"url": "https://x",
"source_id": "cls",
"title": "标题",
"publish_time": None,
"sources": ["cls", "sina"],
"event": {"summary": "摘要", "sentiment": "neutral", "importance": 2,
"event_type": "其他"},
}
fp = ev_dir / f"{h}.json"
fp.write_text(json.dumps(ev, ensure_ascii=False), encoding="utf-8")
built = _build_text_from_event(fp, proc_root, "20260616")
assert built is not None
_text, _article, _summary, sources = built
assert sources == ["cls", "sina"]
def test_build_text_from_article_reads_sources(tmp_path: Path) -> None:
from scripts.run_embedding import _build_text_from_article
h = "f" * 16
art = {
"source_id": "cls",
"url": "https://x",
"url_hash": h,
"title": "标题",
"content": "正文内容",
"word_count": 4,
"sources": ["cls", "eastmoney"],
}
fp = tmp_path / f"{h}.json"
fp.write_text(json.dumps(art, ensure_ascii=False), encoding="utf-8")
built = _build_text_from_article(fp)
assert built is not None
_text, _article, _summary, sources = built
assert sources == ["cls", "eastmoney"]
# --------------------------------------------------------------------------- #
# 回填脚本 _load_fingerprint_sources
# --------------------------------------------------------------------------- #
def _make_fp_db(tmp_path: Path) -> Path:
db = tmp_path / "fp.sqlite3"
conn = sqlite3.connect(db)
conn.execute(
"CREATE TABLE fingerprints (url_hash TEXT PRIMARY KEY, source_id TEXT, "
"source_ids TEXT)"
)
conn.execute(
"INSERT INTO fingerprints VALUES (?, ?, ?)",
("h1", "cls", json.dumps(["cls", "sina"])),
)
conn.execute(
"INSERT INTO fingerprints VALUES (?, ?, ?)", ("h2", "zqrb", None)
)
conn.commit()
conn.close()
return db
def test_load_fingerprint_sources_parses_and_falls_back(tmp_path: Path) -> None:
from scripts.backfill_qdrant_sources import _load_fingerprint_sources
db = _make_fp_db(tmp_path)
res = _load_fingerprint_sources(db)
assert res["h1"] == ["cls", "sina"]
assert res["h2"] == ["zqrb"] # source_ids NULL → 回退 [主源]
def test_load_fingerprint_sources_main_first(tmp_path: Path) -> None:
"""source_ids 顺序异常时保证主源居首。"""
from scripts.backfill_qdrant_sources import _load_fingerprint_sources
db = tmp_path / "fp.sqlite3"
conn = sqlite3.connect(db)
conn.execute(
"CREATE TABLE fingerprints (url_hash TEXT PRIMARY KEY, source_id TEXT, "
"source_ids TEXT)"
)
conn.execute(
"INSERT INTO fingerprints VALUES (?, ?, ?)",
("h1", "cls", json.dumps(["sina", "cls"])),
)
conn.commit()
conn.close()
res = _load_fingerprint_sources(db)
assert res["h1"] == ["cls", "sina"]