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:
+3
-1
@@ -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"):
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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])}")
|
||||||
|
|||||||
@@ -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 ""
|
||||||
|
|||||||
@@ -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,
|
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,
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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"]
|
||||||
@@ -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"),
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user