211 lines
7.8 KiB
Python
211 lines
7.8 KiB
Python
"""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())
|