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:
2026-08-22 22:26:10 +08:00
parent 8fa27ad65b
commit 80828310d6
11 changed files with 490 additions and 17 deletions
+3 -1
View File
@@ -255,7 +255,9 @@ def cmd_search(args: argparse.Namespace) -> int:
ev.get("sentiment"), "" ev.get("sentiment"), ""
) )
print(f"{i}. {sentiment_icon} {h.title}") print(f"{i}. {sentiment_icon} {h.title}")
print(f" 来源: {h.source_id} | 相似度: {h.score:.4f} | 时间: {h.publish_time}") src_display = " / ".join(h.sources) if len(h.sources) > 1 else (h.sources[0] if h.sources else h.source_id)
multi_tag = " [多源]" if len(h.sources) > 1 else ""
print(f" 来源: {src_display}{multi_tag} | 相似度: {h.score:.4f} | 时间: {h.publish_time}")
if ev.get("stock_codes"): if ev.get("stock_codes"):
print(f" 代码: {','.join(ev['stock_codes'])}") print(f" 代码: {','.join(ev['stock_codes'])}")
if ev.get("event_type"): if ev.get("event_type"):
+30
View File
@@ -4,6 +4,36 @@
--- ---
## 本次完成 (2026-08-22) — 多源新闻记录链路修复(方案 A + B)
**用户需求**:一条新闻有多个来源时,全部来源都要记录并可见;此前"找不到多源"。
**诊断结论**:
- M3→M4 多源记录逻辑**本身正常**:指纹库 227 条多源、当日 sources.json 11 条、events 313 条全含 sources、MySQL news_event.sources 填充率 100% 且有 7 条真实多源(如许家印案 `["cls","sina"]`)
- 用户"看不到"的原因:① 展示层(日报 HTML/CLI/MCP)只渲染单源 `source_id`;② 知识库链路 M5→M6 **真断点**——`EmbeddingResult` 无 `sources` 字段、Qdrant payload 只写 `source_id`、`SearchResult` 无 `sources`
**修复内容**:
- B1 模型层: `embedding/models.py` `EmbeddingResult.sources` + validator(主源居首/去重/旧产物兜底);`vectorstore/models.py` `SearchResult.sources`
- B2 `scripts/run_embedding.py`: `_build_text_from_event/_build_text_from_article` 透传 sources 进 EmbeddingResult
- B3 `scripts/run_qdrant_ingest.py`: payload 写 `sources`(M5 产物 → 回查 M4 → 兜底 [主源])
- B4 `vectorstore/client.py`: 检索读取 payload.sources
- A1 `scheduler/reporter.py`: 日报 HTML 多源时显示「财联社 / 新浪 📰」
- A2 `a_share_cli/main.py` + `mcp_server/tools.py`: 检索展示「cecn / cscn [多源]」,MCP 返回 sources 数组
- 存量回填 `scripts/backfill_qdrant_sources.py`(新):以指纹库 source_ids 为权威源,scroll+批量 upsert 回写 3.3 万条 payload
- ⚠️ 坑:本地文件模式 Qdrant 的 `set_payload` 逐点极慢(0.65s/点,33k 条约 6h),改为与 ingest 相同的 scroll+upsert 快速路径;且本地模式单进程锁,必须避开调度窗口
**验证**:
- 新增 `tests/test_multisource.py` 10 个(validator/`_result_to_point` 三优先级/文本构造透传/回填加载)
- 全量 **276 passed**(3 个 crawler 基线失败与本次无关);ruff 干净(reporter 5 个 N806/SIM115 为既有问题未动)
- 今日 M5/M6 重跑:313 条全含 sources(11 条多源);CLI 检索显示「来源: cecn / cscn [多源]」✓,MCP 返回 `sources: ['cecn','cscn']` ✓
- 存量回填:待执行(等 22:00 pipeline 结束后运行,避开 Qdrant 单进程锁)
**待办/遗留**:
- 回填脚本执行 + 回填后检索验证
- P1-2(补跑 steps 含 cninfo)、P1-3(时区一致性),待用户决策
---
## 本次完成 (2026-08-22) — P1-1 修复 dedup 高重复率返回码语义 ## 本次完成 (2026-08-22) — P1-1 修复 dedup 高重复率返回码语义
**问题**:`run_dedup` 重复率 > 5% 时返回 1;`scheduler/pipeline.py` 曾把 `dedup` 的 rc=1 无条件视为成功并打日志"无新数据场景"。结果:高重复率(可能是正常无新数据,也可能是抓取源/指纹库异常)被一刀切掩盖,真实异常无法上报。 **问题**:`run_dedup` 重复率 > 5% 时返回 1;`scheduler/pipeline.py` 曾把 `dedup` 的 rc=1 无条件视为成功并打日志"无新数据场景"。结果:高重复率(可能是正常无新数据,也可能是抓取源/指纹库异常)被一刀切掩盖,真实异常无法上报。
+17 -2
View File
@@ -4,8 +4,9 @@ from __future__ import annotations
from datetime import datetime from datetime import datetime
from enum import StrEnum from enum import StrEnum
from typing import Self
from pydantic import BaseModel, Field from pydantic import BaseModel, Field, model_validator
class EmbeddingProviderType(StrEnum): class EmbeddingProviderType(StrEnum):
@@ -19,7 +20,11 @@ class EmbeddingResult(BaseModel):
"""单篇文章的嵌入结果(落盘格式)。""" """单篇文章的嵌入结果(落盘格式)。"""
url_hash: str = Field(..., description="主键,与 Article.url_hash 一致") url_hash: str = Field(..., description="主键,与 Article.url_hash 一致")
source_id: str = Field(..., description="来源源 id") source_id: str = Field(..., description="主来源源 id")
sources: list[str] = Field(
default_factory=list,
description="该唯一新闻的全部来源(主源 source_id 居首),来自 M3 去重多源记录",
)
title: str = Field(..., description="原文标题(便于人工检索)") title: str = Field(..., description="原文标题(便于人工检索)")
text: str = Field( text: str = Field(
..., description="实际送入 embedder 的文本(已截断/拼接)" ..., description="实际送入 embedder 的文本(已截断/拼接)"
@@ -33,6 +38,16 @@ class EmbeddingResult(BaseModel):
char_count: int = Field(default=0, ge=0, description="text 字符数,便于排查") char_count: int = Field(default=0, ge=0, description="text 字符数,便于排查")
publish_time: datetime | None = None publish_time: datetime | None = None
@model_validator(mode="after")
def _ensure_sources(self) -> Self:
"""保证 sources 非空、去重且以主源 source_id 开头(旧产物无字段时兜底)。"""
seen: list[str] = []
for s in [self.source_id, *self.sources]:
if s and s not in seen:
seen.append(s)
self.sources = seen
return self
def short_summary(self) -> str: def short_summary(self) -> str:
return ( return (
f"[{self.source_id}] {self.title[:30]} " f"[{self.source_id}] {self.title[:30]} "
+5 -1
View File
@@ -63,6 +63,7 @@ def _search(
"title": r.title, "title": r.title,
"url": r.url, "url": r.url,
"source": r.source_id, "source": r.source_id,
"sources": r.sources if r.sources else ([r.source_id] if r.source_id else []),
"score": round(r.score, 4), "score": round(r.score, 4),
"publish_time": r.publish_time.isoformat() if r.publish_time else None, "publish_time": r.publish_time.isoformat() if r.publish_time else None,
"event": { "event": {
@@ -100,7 +101,10 @@ def _fmt_results(hits: list[dict[str, Any]], query: str) -> str:
ev.get("sentiment"), "" ev.get("sentiment"), ""
) )
lines.append(f"### {i}. {h['title']}") lines.append(f"### {i}. {h['title']}")
lines.append(f"- 来源: {h['source']} | 相似度: {h['score']} | {sentiment}") src_list = h.get("sources") or ([h["source"]] if h.get("source") else [])
src_display = " / ".join(src_list)
multi_tag = " [多源]" if len(src_list) > 1 else ""
lines.append(f"- 来源: {src_display}{multi_tag} | 相似度: {h['score']} | {sentiment}")
lines.append(f"- 时间: {h['publish_time'] or '未知'}") lines.append(f"- 时间: {h['publish_time'] or '未知'}")
if ev.get("company_names"): if ev.get("company_names"):
lines.append(f"- 公司: {', '.join(ev['company_names'][:5])}") lines.append(f"- 公司: {', '.join(ev['company_names'][:5])}")
+4 -1
View File
@@ -807,7 +807,10 @@ def _render_event_table(events: list[dict], show_source: bool = True,
imp = ev.get("importance", 0) imp = ev.get("importance", 0)
imp_cls = f"imp-{imp}" if imp >= 4 else "" imp_cls = f"imp-{imp}" if imp >= 4 else ""
title = e["title"][:70] title = e["title"][:70]
src = _source_name(e.get("source_id", "")) src_ids = e.get("sources") or [e.get("source_id", "")]
src = " / ".join(_source_name(s) for s in src_ids if s)
if len(src_ids) > 1:
src = f'<span title="多源新闻">{src}</span> 📰'
codes_in_event = {c.strip().split(".")[0] for c in (ev.get("stock_codes") or [])} codes_in_event = {c.strip().split(".")[0] for c in (ev.get("stock_codes") or [])}
star = "⭐ " if codes_in_event & wl_codes else "" star = "⭐ " if codes_in_event & wl_codes else ""
code_str = f" <small>({','.join(list(codes_in_event)[:3])})</small>" if codes_in_event else "" code_str = f" <small>({','.join(list(codes_in_event)[:3])})</small>" if codes_in_event else ""
+165
View File
@@ -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
View File
@@ -89,8 +89,11 @@ def _build_text_from_event(
event_path: Path, event_path: Path,
processed_root: Path, processed_root: Path,
day: str, day: str,
) -> tuple[str, Article, str | None] | None: ) -> tuple[str, Article, str | None, list[str] | None] | None:
"""从 ExtractedEvent JSON 构造嵌入文本与文章元数据。""" """从 ExtractedEvent JSON 构造嵌入文本与文章元数据。
返回 (text, article, summary, sources);sources 为 M3 多源记录。
"""
try: try:
obj: dict[str, Any] = json.loads(event_path.read_text(encoding="utf-8")) obj: dict[str, Any] = json.loads(event_path.read_text(encoding="utf-8"))
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
@@ -98,6 +101,7 @@ def _build_text_from_event(
return None return None
article, head, summary = _from_event_dict(obj) article, head, summary = _from_event_dict(obj)
sources = obj.get("sources") or None
# 真正的正文要去 processed/ 找 # 真正的正文要去 processed/ 找
real_article = _load_article_by_hash(processed_root, day, article.url_hash) real_article = _load_article_by_hash(processed_root, day, article.url_hash)
if real_article is None: 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} update={"publish_time": article.publish_time or real_article.publish_time}
) )
text = compose_text(real_article, head=head, summary=summary) 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: try:
obj = json.loads(article_path.read_text(encoding="utf-8")) obj = json.loads(article_path.read_text(encoding="utf-8"))
article = Article.model_validate(obj) article = Article.model_validate(obj)
except (json.JSONDecodeError, ValidationError) as e: except (json.JSONDecodeError, ValidationError) as e:
logger.warning("跳过损坏 article 文件 {}: {}", article_path, e) logger.warning("跳过损坏 article 文件 {}: {}", article_path, e)
return None 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]]: def _collect_inputs(args: argparse.Namespace) -> list[tuple[Path, str]]:
@@ -192,8 +200,8 @@ async def _run(args: argparse.Namespace) -> int:
return 0 return 0
logger.info("待嵌入文章数: {} (跳过已处理 {}; input={})", len(files), skipped, args.input) 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: for fp, kind in files:
if kind == "event": if kind == "event":
built = _build_text_from_event(fp, Path(args.processed_root), args.date) 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 batch_size = args.batch_size
for i in range(0, len(prepared), batch_size): for i in range(0, len(prepared), batch_size):
batch = prepared[i : i + batch_size] batch = prepared[i : i + batch_size]
texts = [t for t, _, _ in batch] texts = [t for t, _, _, _ in batch]
try: try:
vectors = await provider.embed_batch(texts) vectors = await provider.embed_batch(texts)
except EmbeddingError as e: except EmbeddingError as e:
logger.warning("批 {} 嵌入失败: {}", i // batch_size, e) logger.warning("批 {} 嵌入失败: {}", i // batch_size, e)
for _, art, _ in batch: for _, art, _, _ in batch:
fail_cnt += 1 fail_cnt += 1
with failed_path.open("a", encoding="utf-8") as f: with failed_path.open("a", encoding="utf-8") as f:
f.write( f.write(
@@ -252,7 +260,7 @@ async def _run(args: argparse.Namespace) -> int:
) )
continue 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: if len(vec) != provider.dim:
logger.warning( logger.warning(
"维度不一致 url_hash={} 实际={} 预期={}", "维度不一致 url_hash={} 实际={} 预期={}",
@@ -261,6 +269,7 @@ async def _run(args: argparse.Namespace) -> int:
result = EmbeddingResult( result = EmbeddingResult(
url_hash=article.url_hash, url_hash=article.url_hash,
source_id=article.source_id, source_id=article.source_id,
sources=sources or [],
title=article.title, title=article.title,
text=text, text=text,
vector=vec, vector=vec,
+25 -2
View File
@@ -52,18 +52,26 @@ def _result_to_point(
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
"""把 EmbeddingResult JSON dict 转换为 Qdrant Point 格式。 """把 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)。 事件字段优先从 M4 ExtractedEvent 补充(EmbeddingResult 本身不含 event)。
sources 多源字段优先取 M5 产物,缺失时回查 M4,再兜底 [主源]。
""" """
vector = obj.get("vector") vector = obj.get("vector")
if not vector: if not vector:
return None return None
url_hash = obj["url_hash"] 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 = { payload = {
"url_hash": url_hash, "url_hash": url_hash,
"title": obj.get("title") or "", "title": obj.get("title") or "",
"url": obj.get("url") 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"), "publish_time": obj.get("publish_time"),
"char_count": obj.get("char_count"), "char_count": obj.get("char_count"),
"word_count": obj.get("word_count"), "word_count": obj.get("word_count"),
@@ -93,6 +101,21 @@ def _result_to_point(
return {"id": url_hash, "vector": vector, "payload": payload} 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: def _load_event_from_m4(events_dir: Path, url_hash: str) -> dict[str, Any] | None:
"""从 M4 ExtractedEvent JSON 中提取事件 payload 子集。""" """从 M4 ExtractedEvent JSON 中提取事件 payload 子集。"""
event_file = events_dir / f"{url_hash}.json" event_file = events_dir / f"{url_hash}.json"
+217
View File
@@ -0,0 +1,217 @@
"""多源记录端到端链路测试 (方案 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"]
+1
View File
@@ -257,6 +257,7 @@ class VectorStore:
title=payload.get("title") or "", title=payload.get("title") or "",
url=payload.get("url") or "", url=payload.get("url") or "",
source_id=payload.get("source_id") or "", source_id=payload.get("source_id") or "",
sources=payload.get("sources") or [],
publish_time=publish_time, publish_time=publish_time,
event=payload.get("event"), event=payload.get("event"),
char_count=payload.get("char_count"), char_count=payload.get("char_count"),
+4
View File
@@ -31,6 +31,10 @@ class SearchResult(BaseModel):
title: str title: str
url: str url: str
source_id: str source_id: str
sources: list[str] = Field(
default_factory=list,
description="全部来源(主源居首),来自 M3 去重多源记录;旧数据可能为空",
)
publish_time: datetime | None = None publish_time: datetime | None = None
event: dict[str, Any] | None = None # EventExtraction 展开的 dict event: dict[str, Any] | None = None # EventExtraction 展开的 dict
char_count: int | None = None char_count: int | None = None