From 80828310d6ab175238b1cf4017ded00f99dfdfdd Mon Sep 17 00:00:00 2001 From: simon Date: Sat, 22 Aug 2026 22:26:10 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=89=93=E9=80=9A=E5=A4=9A=E6=BA=90?= =?UTF-8?q?=E6=96=B0=E9=97=BB=E8=AE=B0=E5=BD=95=E9=93=BE=E8=B7=AF,?= =?UTF-8?q?=E6=A3=80=E7=B4=A2/=E6=97=A5=E6=8A=A5/=E7=9F=A5=E8=AF=86?= =?UTF-8?q?=E5=BA=93=E5=8F=AF=E8=A7=81=E5=A4=9A=E6=BA=90=20(=E6=96=B9?= =?UTF-8?q?=E6=A1=88A+B)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 模型层: 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 --- a_share_cli/main.py | 4 +- continuation.md | 30 ++++ embedding/models.py | 19 ++- mcp_server/tools.py | 6 +- scheduler/reporter.py | 5 +- scripts/backfill_qdrant_sources.py | 165 ++++++++++++++++++++++ scripts/run_embedding.py | 29 ++-- scripts/run_qdrant_ingest.py | 27 +++- tests/test_multisource.py | 217 +++++++++++++++++++++++++++++ vectorstore/client.py | 1 + vectorstore/models.py | 4 + 11 files changed, 490 insertions(+), 17 deletions(-) create mode 100644 scripts/backfill_qdrant_sources.py create mode 100644 tests/test_multisource.py diff --git a/a_share_cli/main.py b/a_share_cli/main.py index c00c98b..c54d171 100644 --- a/a_share_cli/main.py +++ b/a_share_cli/main.py @@ -255,7 +255,9 @@ def cmd_search(args: argparse.Namespace) -> int: ev.get("sentiment"), "" ) 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"): print(f" 代码: {','.join(ev['stock_codes'])}") if ev.get("event_type"): diff --git a/continuation.md b/continuation.md index 266cb05..d87cb9e 100644 --- a/continuation.md +++ b/continuation.md @@ -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 高重复率返回码语义 **问题**:`run_dedup` 重复率 > 5% 时返回 1;`scheduler/pipeline.py` 曾把 `dedup` 的 rc=1 无条件视为成功并打日志"无新数据场景"。结果:高重复率(可能是正常无新数据,也可能是抓取源/指纹库异常)被一刀切掩盖,真实异常无法上报。 diff --git a/embedding/models.py b/embedding/models.py index cb60b3b..6de9e11 100644 --- a/embedding/models.py +++ b/embedding/models.py @@ -4,8 +4,9 @@ from __future__ import annotations from datetime import datetime from enum import StrEnum +from typing import Self -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator class EmbeddingProviderType(StrEnum): @@ -19,7 +20,11 @@ class EmbeddingResult(BaseModel): """单篇文章的嵌入结果(落盘格式)。""" 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="原文标题(便于人工检索)") text: str = Field( ..., description="实际送入 embedder 的文本(已截断/拼接)" @@ -33,6 +38,16 @@ class EmbeddingResult(BaseModel): char_count: int = Field(default=0, ge=0, description="text 字符数,便于排查") 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: return ( f"[{self.source_id}] {self.title[:30]} " diff --git a/mcp_server/tools.py b/mcp_server/tools.py index b61579e..cc8b7a7 100644 --- a/mcp_server/tools.py +++ b/mcp_server/tools.py @@ -63,6 +63,7 @@ def _search( "title": r.title, "url": r.url, "source": r.source_id, + "sources": r.sources if r.sources else ([r.source_id] if r.source_id else []), "score": round(r.score, 4), "publish_time": r.publish_time.isoformat() if r.publish_time else None, "event": { @@ -100,7 +101,10 @@ def _fmt_results(hits: list[dict[str, Any]], query: str) -> str: ev.get("sentiment"), "" ) 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 '未知'}") if ev.get("company_names"): lines.append(f"- 公司: {', '.join(ev['company_names'][:5])}") diff --git a/scheduler/reporter.py b/scheduler/reporter.py index e9f9498..abe9cd7 100644 --- a/scheduler/reporter.py +++ b/scheduler/reporter.py @@ -807,7 +807,10 @@ def _render_event_table(events: list[dict], show_source: bool = True, imp = ev.get("importance", 0) imp_cls = f"imp-{imp}" if imp >= 4 else "" 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'{src} 📰' codes_in_event = {c.strip().split(".")[0] for c in (ev.get("stock_codes") or [])} star = "⭐ " if codes_in_event & wl_codes else "" code_str = f" ({','.join(list(codes_in_event)[:3])})" if codes_in_event else "" diff --git a/scripts/backfill_qdrant_sources.py b/scripts/backfill_qdrant_sources.py new file mode 100644 index 0000000..0dcd3ae --- /dev/null +++ b/scripts/backfill_qdrant_sources.py @@ -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()) + diff --git a/scripts/run_embedding.py b/scripts/run_embedding.py index 5e5e088..496b6e9 100644 --- a/scripts/run_embedding.py +++ b/scripts/run_embedding.py @@ -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, diff --git a/scripts/run_qdrant_ingest.py b/scripts/run_qdrant_ingest.py index fabe77e..8752c82 100644 --- a/scripts/run_qdrant_ingest.py +++ b/scripts/run_qdrant_ingest.py @@ -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" diff --git a/tests/test_multisource.py b/tests/test_multisource.py new file mode 100644 index 0000000..a3ffed8 --- /dev/null +++ b/tests/test_multisource.py @@ -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"] diff --git a/vectorstore/client.py b/vectorstore/client.py index 596afbe..1cf69a4 100644 --- a/vectorstore/client.py +++ b/vectorstore/client.py @@ -257,6 +257,7 @@ class VectorStore: title=payload.get("title") or "", url=payload.get("url") or "", source_id=payload.get("source_id") or "", + sources=payload.get("sources") or [], publish_time=publish_time, event=payload.get("event"), char_count=payload.get("char_count"), diff --git a/vectorstore/models.py b/vectorstore/models.py index 66cbd27..227360e 100644 --- a/vectorstore/models.py +++ b/vectorstore/models.py @@ -31,6 +31,10 @@ class SearchResult(BaseModel): title: str url: str source_id: str + sources: list[str] = Field( + default_factory=list, + description="全部来源(主源居首),来自 M3 去重多源记录;旧数据可能为空", + ) publish_time: datetime | None = None event: dict[str, Any] | None = None # EventExtraction 展开的 dict char_count: int | None = None