Initial commit

This commit is contained in:
2026-07-18 15:51:01 +08:00
commit f2c80c5a9c
799 changed files with 133475 additions and 0 deletions
+210
View File
@@ -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())