Initial commit
This commit is contained in:
@@ -0,0 +1,210 @@
|
||||
"""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())
|
||||
Reference in New Issue
Block a user