"""M6 Qdrant 批量入库脚本。 输入: data/embeddings/{date}/*.json (M5 产物,含 vector + payload) 目标: Qdrant Collection a_share_news 策略: 幂等 upsert(url_hash 做 point ID,同 ID 覆盖不重复计数) 用法: uv run python -m scripts.run_qdrant_ingest # 入库今日 uv run python -m scripts.run_qdrant_ingest --date 20260616 uv run python -m scripts.run_qdrant_ingest --recreate # 重建 collection + 全量入库 uv run python -m scripts.run_qdrant_ingest --memory # 内存模式(测试) uv run python -m scripts.run_qdrant_ingest --limit 10 # 试跑 N 条 """ from __future__ import annotations import argparse import json import sys from datetime import date from pathlib import Path from typing import Any from loguru import logger from vectorstore import SearchResult, VectorStore, make_qdrant_client 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}", ) log_path = Path("logs") / "qdrant.log" log_path.parent.mkdir(parents=True, exist_ok=True) logger.add(log_path, level="DEBUG", rotation="10 MB", retention=5, encoding="utf-8") def _load_embedding_result(path: Path) -> dict[str, Any] | None: try: return json.loads(path.read_text(encoding="utf-8")) except (json.JSONDecodeError, OSError) as e: logger.warning("跳过损坏文件 {}: {}", path, e) return None def _result_to_point( obj: dict[str, Any], events_dir: Path | None = None, ) -> dict[str, Any] | None: """把 EmbeddingResult JSON dict 转换为 Qdrant Point 格式。 Payload 包含 title/url/source_id/publish_time/event/计数 等。 事件字段优先从 M4 ExtractedEvent 补充(EmbeddingResult 本身不含 event)。 """ vector = obj.get("vector") if not vector: return None url_hash = obj["url_hash"] payload = { "url_hash": url_hash, "title": obj.get("title") or "", "url": obj.get("url") or "", "source_id": obj.get("source_id") or "", "publish_time": obj.get("publish_time"), "char_count": obj.get("char_count"), "word_count": obj.get("word_count"), "event": { "stock_codes": [], "company_names": [], "industries": [], "sentiment": "neutral", "importance": 1, "event_type": "其他", "summary": "", }, } # 优先从 M4 ExtractedEvent JSON 补事件字段 ev = _load_event_from_m4(events_dir, url_hash) if events_dir else None if ev is not None: payload["event"] = ev elif "event" in obj and isinstance(obj["event"], dict): ev_src = obj["event"] payload["event"]["stock_codes"] = list(ev_src.get("stock_codes") or []) payload["event"]["company_names"] = list(ev_src.get("company_names") or []) payload["event"]["industries"] = list(ev_src.get("industries") or []) payload["event"]["sentiment"] = ev_src.get("sentiment") or "neutral" payload["event"]["importance"] = ev_src.get("importance") or 1 payload["event"]["event_type"] = ev_src.get("event_type") or "其他" payload["event"]["summary"] = ev_src.get("summary") or "" return {"id": url_hash, "vector": vector, "payload": payload} 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" 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 ev_inner = obj.get("event") if not isinstance(ev_inner, dict): return None return { "stock_codes": list(ev_inner.get("stock_codes") or []), "company_names": list(ev_inner.get("company_names") or []), "industries": list(ev_inner.get("industries") or []), "sentiment": ev_inner.get("sentiment") or "neutral", "importance": ev_inner.get("importance") or 1, "event_type": ev_inner.get("event_type") or "其他", "summary": ev_inner.get("summary") or "", } def main() -> int: parser = argparse.ArgumentParser(description="A 股新闻 Qdrant 入库 (M6)") parser.add_argument("--embeddings-root", default="data/embeddings") parser.add_argument( "--date", default=date.today().strftime("%Y%m%d"), help="日期 YYYYMMDD,默认今日", ) parser.add_argument("--limit", type=int, default=0, help="最多处理 N 条,0=不限") parser.add_argument("--recreate", action="store_true", help="入库前删除并重建 collection") parser.add_argument("--host", default=None, help="Qdrant host(HTTP 模式)") parser.add_argument("--port", type=int, default=None, help="Qdrant HTTP port") parser.add_argument("--path", default="data/qdrant_storage", help="本地文件模式路径(默认,无需 Docker)") parser.add_argument("--collection", default=None, help="Collection 名") parser.add_argument("--memory", action="store_true", help="内存模式(仅测试)") parser.add_argument("--log-level", default="INFO") args = parser.parse_args() _setup_logger(args.log_level) # 收集文件 emb_dir = Path(args.embeddings_root) / args.date events_dir = Path("data/events") / args.date # M4 产物(补 event 字段) if not emb_dir.is_dir(): logger.error("嵌入目录不存在: {}", emb_dir) return 2 files = sorted(emb_dir.glob("*.json")) if args.limit: files = files[: args.limit] if not files: logger.error("{} 下无嵌入文件", emb_dir) return 2 logger.info("待入库文章数: {} (date={}), events补: {}", len(files), args.date, "yes" if events_dir.is_dir() else "no") # 转换 points: list[dict[str, Any]] = [] for fp in files: obj = _load_embedding_result(fp) if obj is None: continue pt = _result_to_point(obj, events_dir=events_dir if events_dir.is_dir() else None) if pt is not None: points.append(pt) if not points: logger.error("所有嵌入文件均无法解析") return 2 client = make_qdrant_client( host=args.host, port=args.port, memory=args.memory, path=args.path if not args.host and not args.memory else None, ) store = VectorStore(client, collection_name=args.collection) # 初始化 collection try: store.init_collection(recreate=args.recreate) except Exception: # noqa: BLE001 - Qdrant 未启动时会爆连接错误 logger.exception("无法连接 Qdrant,请先 docker compose --profile m6 up -d: {}") return 3 # 入库 before = store.count() store.upsert(points) after = store.count() logger.info("入库完成: 前 {} -> 后 {} (净增 {})", before, after, after - before) # 简单自检:用第一条向量做检索,验证可查回 if points and not args.memory: probe = points[0] try: results: list[SearchResult] = store.query( query_vector=probe["vector"], top_k=1, ) if results: r = results[0] logger.info("自检 OK: top-1 title={!r} score={:.4f}", r.title[:30], r.score) else: logger.warning("自检:检索返回空") except Exception as e: # noqa: BLE001 logger.warning("自检失败(不阻塞): {}", e) store.close() return 0 if __name__ == "__main__": raise SystemExit(main())