Files
2026-07-18 16:13:52 +08:00

217 lines
5.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Qdrant 入库管道 + 语义搜索。
输入: data/embeddings/{YYYYMMDD}/{url_hash}.jsonM5 向量) + data/events/M4 元数据)
动作: upsert 到 Qdrant collection
搜索: embed query → Qdrant query → 返回 SearchResult 列表
"""
import json
import logging
from datetime import datetime
from pathlib import Path
from crawler.utils import get_news_day
from embedding.client import (
embed_batch,
load_embedding_config,
make_embedding_client,
)
from vectorstore.client import VectorStore, make_qdrant_client
from vectorstore.models import SearchFilter, SearchResult
logger = logging.getLogger(__name__)
def _load_embedding_files(date_str: str) -> list[dict]:
"""加载指定日期的嵌入向量文件(含对应的 M4 元数据)。
Args:
date_str: 日期 YYYYMMDD
Returns:
dict 列表,含 url_hash / vector / article 信息
"""
embedding_dir = Path(f"data/embeddings/{date_str}")
event_dir = Path(f"data/events/{date_str}")
if not embedding_dir.exists():
return []
items: list[dict] = []
for emb_file in sorted(embedding_dir.glob("*.json")):
if emb_file.name == "index.json":
continue
try:
emb_data = json.loads(emb_file.read_text(encoding="utf-8"))
# 加载对应的 M4 事件文章获取元数据
event_file = event_dir / emb_file.name
article_data = {}
if event_file.exists():
article_data = json.loads(event_file.read_text(encoding="utf-8"))
items.append({
"url_hash": emb_data["url_hash"],
"source_id": emb_data.get("source_id", ""),
"vector": emb_data["vector"],
"article": article_data,
})
except (json.JSONDecodeError, Exception) as e:
logger.warning("加载嵌入文件失败 %s: %s", emb_file, e)
return items
def _build_payload(article_data: dict) -> dict:
"""从 M4 文章数据构建 Qdrant payload。
Args:
article_data: EnTranslatedArticle 的 dict
Returns:
payload dict
"""
content_zh = article_data.get("content_zh", "")
return {
"title": article_data.get("title", ""),
"title_zh": article_data.get("title_zh", ""),
"url": article_data.get("url", ""),
"source_id": article_data.get("source_id", ""),
"source_name": article_data.get("source_name", ""),
"publish_time": article_data.get("publish_time", ""),
"events": article_data.get("events", []),
"word_count": article_data.get("word_count", 0),
"word_count_zh": article_data.get("word_count_zh", 0),
"content_zh_preview": content_zh[:300] if content_zh else "",
}
def ingest_all_embeddings(
date_str: str | None = None,
*,
recreate: bool = False,
) -> dict:
"""将所有 M5 向量入库 Qdrant。
Args:
date_str: 日期 YYYYMMDD,默认当前新闻日
recreate: 是否重建 collection
Returns:
统计 dict
"""
if date_str is None:
date_str = get_news_day()
logger.info("══════ 开始 Qdrant 入库,日期: %s ══════", date_str)
items = _load_embedding_files(date_str)
if not items:
logger.warning("嵌入目录无数据: data/embeddings/%s/", date_str)
return {"date": date_str, "total": 0, "ingested": 0, "failed": 0, "elapsed_sec": 0}
start_time = datetime.now()
# 构造 Qdrant 客户端
client = make_qdrant_client()
store = VectorStore(client)
try:
# 初始化 collection
store.init_collection(recreate=recreate)
# 构建 points
points: list[dict] = []
for item in items:
payload = _build_payload(item["article"])
points.append({
"id": item["url_hash"],
"vector": item["vector"],
"payload": payload,
})
# 批量写入
ingested = store.upsert(points)
failed = len(points) - ingested
finally:
store.close()
elapsed = (datetime.now() - start_time).total_seconds()
logger.info(
"══════ Qdrant 入库完成: %d 条,耗时 %.1f 秒 ══════",
ingested, elapsed,
)
return {
"date": date_str,
"total": len(items),
"ingested": ingested,
"failed": failed,
"elapsed_sec": elapsed,
}
def search_news(
query: str,
*,
top_k: int = 10,
search_filter: SearchFilter | None = None,
score_threshold: float | None = None,
) -> list[SearchResult]:
"""语义搜索新闻。
流程:
1. 将查询文本向量化(使用 M5 Embedding 服务)
2. Qdrant 语义检索
Args:
query: 中文搜索查询
top_k: 返回条数
search_filter: 可选过滤条件
score_threshold: 最低相似度阈值
Returns:
SearchResult 列表
"""
# 1. 向量化查询
emb_config = load_embedding_config()
emb_client = make_embedding_client(emb_config)
try:
vectors = embed_batch(emb_client, emb_config, [query])
if not vectors:
logger.error("查询向量化失败")
return []
query_vector = vectors[0]
finally:
emb_client.close()
# 2. Qdrant 检索
qdrant = make_qdrant_client()
store = VectorStore(qdrant)
try:
results = store.query(
query_vector=query_vector,
top_k=top_k,
search_filter=search_filter,
score_threshold=score_threshold,
)
finally:
store.close()
return results
def get_collection_info() -> dict:
"""获取 Qdrant collection 信息。"""
client = make_qdrant_client()
store = VectorStore(client)
try:
info = store.info()
return info.model_dump()
finally:
store.close()