217 lines
5.9 KiB
Python
217 lines
5.9 KiB
Python
"""Qdrant 入库管道 + 语义搜索。
|
||
|
||
输入: data/embeddings/{YYYYMMDD}/{url_hash}.json(M5 向量) + 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()
|