Files
QMDSearch/app/services/qdrant.py
T
kplam 51dc8dc4f6 Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈
- L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合)
- 文档三级总结与 2.5 级回退
- query 解析路由与分类
- /admin 管理页面
2026-07-29 21:24:40 +08:00

352 lines
14 KiB
Python
Raw 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 向量数据库服务封装
分层摘要索引 RAG 的存储层:
- doc_l1:文档级总结(dense + sparse
- doc_l2 / doc_l3:大纲节点(dense
- chunks:原文分块(dense + sparse
提供集合初始化(幂等)、按层 upsert、dense / hybridRRF 融合)检索接口。
"""
import uuid
from typing import Any
import structlog
from qdrant_client import AsyncQdrantClient, models
from app.config import settings
logger = structlog.get_logger()
# 集合名
COLLECTION_L1 = "doc_l1"
COLLECTION_L2 = "doc_l2"
COLLECTION_L3 = "doc_l3"
COLLECTION_CHUNKS = "chunks"
ALL_COLLECTIONS = (COLLECTION_L1, COLLECTION_L2, COLLECTION_L3, COLLECTION_CHUNKS)
# 命名向量
VECTOR_DENSE = "dense"
VECTOR_SPARSE = "sparse"
# 需要 sparse 向量的集合
SPARSE_COLLECTIONS = (COLLECTION_L1, COLLECTION_CHUNKS)
# 需要建立 KEYWORD payload 索引的字段
PAYLOAD_INDEX_FIELDS = ("doc_id", "category", "tags", "section_path")
# upsert_nodes 集合名 -> point id 层级前缀
_NODE_ID_PREFIX = {COLLECTION_L2: "l2", COLLECTION_L3: "l3"}
# 稀疏向量统一表示:(indices, values)
SparseVectorTuple = tuple[list[int], list[float]]
def _point_id(key: str) -> str:
"""由确定性 key 生成 UUID5 point id,保证重复入库幂等覆盖"""
return str(uuid.uuid5(uuid.NAMESPACE_URL, key))
class QdrantService:
"""Qdrant 异步客户端封装"""
def __init__(self, client: AsyncQdrantClient | None = None) -> None:
# 允许注入自定义 client(测试可用 location=":memory:" 的本地模式)
self._client = client or AsyncQdrantClient(host=settings.qdrant_host, port=settings.qdrant_port)
@property
def client(self) -> AsyncQdrantClient:
return self._client
async def ensure_collections(self) -> None:
"""初始化 4 个集合与 payload 索引(幂等,已存在则跳过)"""
existing = await self._client.get_collections()
existing_names = {c.name for c in existing.collections}
for collection in ALL_COLLECTIONS:
if collection in existing_names:
continue
sparse_config = None
if settings.sparse_enabled and collection in SPARSE_COLLECTIONS:
sparse_config = {VECTOR_SPARSE: models.SparseVectorParams(modifier=models.Modifier.IDF)}
dense_params = models.VectorParams(size=settings.embedding_dimension, distance=models.Distance.COSINE)
await self._client.create_collection(
collection_name=collection,
vectors_config={VECTOR_DENSE: dense_params},
sparse_vectors_config=sparse_config,
)
logger.info("创建 Qdrant 集合", collection=collection, sparse=sparse_config is not None)
# payload 索引:先查已有索引,缺失才创建,保证幂等
for collection in ALL_COLLECTIONS:
info = await self._client.get_collection(collection)
indexed = set(info.payload_schema.keys()) if info.payload_schema else set()
for field in PAYLOAD_INDEX_FIELDS:
if field in indexed:
continue
await self._client.create_payload_index(
collection_name=collection,
field_name=field,
field_schema=models.PayloadSchemaType.KEYWORD,
)
logger.debug("payload 索引就绪", collection=collection)
# ---------- upsert ----------
async def upsert_l1(
self,
doc_id: str,
title: str,
summary: str,
category: str,
tags: list[str],
dense_vector: list[float],
sparse_vector: SparseVectorTuple | None = None,
) -> None:
"""写入 L1 文档总结,payload 含 doc_id/title/category/tags/text(=summary)"""
vector: dict[str, Any] = {VECTOR_DENSE: dense_vector}
if sparse_vector is not None:
vector[VECTOR_SPARSE] = models.SparseVector(indices=sparse_vector[0], values=sparse_vector[1])
point = models.PointStruct(
id=_point_id(f"{doc_id}:l1"),
vector=vector,
payload={"doc_id": doc_id, "title": title, "category": category, "tags": tags, "text": summary},
)
await self._client.upsert(collection_name=COLLECTION_L1, points=[point])
async def upsert_nodes(self, collection: str, nodes: list[dict[str, Any]]) -> None:
"""批量写入 L2/L3 大纲节点
每个 node 含 doc_id/section_path/text/category/tags/dense_vector。
"""
if collection not in _NODE_ID_PREFIX:
raise ValueError(f"upsert_nodes 仅支持 {COLLECTION_L2}/{COLLECTION_L3},收到: {collection}")
prefix = _NODE_ID_PREFIX[collection]
points = [
models.PointStruct(
id=_point_id(f"{node['doc_id']}:{prefix}:{i}"),
vector={VECTOR_DENSE: node["dense_vector"]},
payload={
"doc_id": node["doc_id"],
"section_path": node["section_path"],
"text": node["text"],
"category": node["category"],
"tags": node["tags"],
},
)
for i, node in enumerate(nodes)
]
if points:
await self._client.upsert(collection_name=collection, points=points)
async def upsert_chunks(self, chunks: list[dict[str, Any]]) -> None:
"""批量写入原文 chunk
每个 chunk 含 doc_id/chunk_index/text/section_path/title/category/tags/dense_vector
sparse_vector 可选((indices, values) 形式)。
"""
points = []
for chunk in chunks:
vector: dict[str, Any] = {VECTOR_DENSE: chunk["dense_vector"]}
sparse = chunk.get("sparse_vector")
if sparse is not None:
vector[VECTOR_SPARSE] = models.SparseVector(indices=sparse[0], values=sparse[1])
points.append(
models.PointStruct(
id=_point_id(f"{chunk['doc_id']}:chunk:{chunk['chunk_index']}"),
vector=vector,
payload={
"doc_id": chunk["doc_id"],
"chunk_index": chunk["chunk_index"],
"text": chunk["text"],
"section_path": chunk["section_path"],
"title": chunk["title"],
"category": chunk["category"],
"tags": chunk["tags"],
# 所属文档 L1 总结,仅作检索结果上下文标注
"doc_summary": chunk.get("doc_summary", ""),
},
)
)
if points:
await self._client.upsert(collection_name=COLLECTION_CHUNKS, points=points)
# ---------- 查询 ----------
async def search_dense(
self,
collection: str,
vector: list[float],
limit: int,
query_filter: models.Filter | None = None,
) -> list[models.ScoredPoint]:
"""dense 命名向量检索"""
resp = await self._client.query_points(
collection_name=collection,
query=vector,
using=VECTOR_DENSE,
limit=limit,
query_filter=query_filter,
)
return resp.points
async def search_hybrid(
self,
collection: str,
dense_vector: list[float],
sparse: SparseVectorTuple,
limit: int,
query_filter: models.Filter | None = None,
) -> list[models.ScoredPoint]:
"""dense + sparse 两路 prefetch,服务端 RRF 融合"""
resp = await self._client.query_points(
collection_name=collection,
prefetch=[
models.Prefetch(query=dense_vector, using=VECTOR_DENSE, limit=limit, filter=query_filter),
models.Prefetch(
query=models.SparseVector(indices=sparse[0], values=sparse[1]),
using=VECTOR_SPARSE,
limit=limit,
filter=query_filter,
),
],
query=models.FusionQuery(fusion=models.Fusion.RRF),
limit=limit,
query_filter=query_filter,
)
return resp.points
# ---------- 管理操作 ----------
async def count(self, collection: str) -> int:
"""精确统计集合中点的总数"""
result = await self._client.count(collection_name=collection, exact=True)
return result.count
async def scroll_l1(self, limit: int = 20, offset: str | None = None) -> tuple[list[dict[str, Any]], str | None]:
"""分页浏览 L1 文档列表(不取向量)
返回 (items, next_offset)item 含 doc_id/title/category/tags/summary(=payload text)
next_offset 为下一页游标,无更多数据时为 None。
"""
records, next_offset = await self._client.scroll(
collection_name=COLLECTION_L1,
limit=limit,
offset=offset,
with_payload=["doc_id", "title", "category", "tags", "text"],
with_vectors=False,
)
items = []
for record in records:
payload = record.payload or {}
items.append(
{
"doc_id": payload.get("doc_id", ""),
"title": payload.get("title", ""),
"category": payload.get("category", ""),
"tags": payload.get("tags", []),
"summary": payload.get("text", ""),
}
)
return items, str(next_offset) if next_offset is not None else None
async def get_doc_detail(self, doc_id: str) -> dict[str, Any] | None:
"""获取文档详情:L1 记录 + L2/L3 全部节点 + chunks 数量,文档不存在返回 None"""
doc_filter = self.build_filter(doc_ids=[doc_id])
l1_records, _ = await self._client.scroll(
collection_name=COLLECTION_L1,
scroll_filter=doc_filter,
limit=1,
with_payload=True,
with_vectors=False,
)
if not l1_records:
return None
l2_nodes = await self._scroll_payloads(COLLECTION_L2, doc_filter)
l3_nodes = await self._scroll_payloads(COLLECTION_L3, doc_filter)
chunks_count = await self._client.count(
collection_name=COLLECTION_CHUNKS,
count_filter=doc_filter,
exact=True,
)
return {
"l1": l1_records[0].payload or {},
"l2_nodes": l2_nodes,
"l3_nodes": l3_nodes,
"chunks_count": chunks_count.count,
}
async def delete_by_doc_id(self, doc_id: str) -> dict[str, int]:
"""删除四层集合中该 doc_id 的所有点,返回各集合删除数量(不存在的 doc_id 全 0,幂等)"""
doc_filter = self.build_filter(doc_ids=[doc_id])
deleted: dict[str, int] = {}
for collection in ALL_COLLECTIONS:
# 先记录删除前数量,再按过滤器删除
before = await self._client.count(collection_name=collection, count_filter=doc_filter, exact=True)
await self._client.delete(
collection_name=collection,
points_selector=models.FilterSelector(filter=doc_filter),
)
deleted[collection] = before.count
logger.info("按 doc_id 删除文档数据", doc_id=doc_id, deleted=deleted)
return deleted
async def _scroll_payloads(
self,
collection: str,
scroll_filter: models.Filter | None,
page_size: int = 256,
) -> list[dict[str, Any]]:
"""循环 scroll 取出所有匹配点的 payload(含 section_path/text 等全字段)"""
payloads: list[dict[str, Any]] = []
offset = None
while True:
records, next_offset = await self._client.scroll(
collection_name=collection,
scroll_filter=scroll_filter,
limit=page_size,
offset=offset,
with_payload=True,
with_vectors=False,
)
payloads.extend(record.payload or {} for record in records)
if next_offset is None:
return payloads
offset = next_offset
# ---------- 过滤构造 ----------
@staticmethod
def build_filter(
categories: list[str] | None = None,
doc_ids: list[str] | None = None,
section_paths: list[str] | None = None,
) -> models.Filter | None:
"""构造 payload 过滤器
- categories:主类硬过滤 + 多标签软召回(category 或 tags 命中其一即召回,min_should=1
- doc_ids / section_pathsmust 条件 MatchAny
- 全空返回 None
"""
must: list[models.FieldCondition] = []
if doc_ids:
must.append(models.FieldCondition(key="doc_id", match=models.MatchAny(any=doc_ids)))
if section_paths:
must.append(models.FieldCondition(key="section_path", match=models.MatchAny(any=section_paths)))
min_should = None
if categories:
min_should = models.MinShould(
conditions=[
models.FieldCondition(key="category", match=models.MatchAny(any=categories)),
models.FieldCondition(key="tags", match=models.MatchAny(any=categories)),
],
min_count=1,
)
if not must and min_should is None:
return None
return models.Filter(must=must or None, min_should=min_should)