Initial commit: QMDSearch 分层信息检索服务

- FastAPI + Qdrant + Redis + Ollama 技术栈
- L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合)
- 文档三级总结与 2.5 级回退
- query 解析路由与分类
- /admin 管理页面
This commit is contained in:
2026-07-29 21:24:40 +08:00
commit 51dc8dc4f6
83 changed files with 10794 additions and 0 deletions
View File
+62
View File
@@ -0,0 +1,62 @@
"""Ollama 客户端
通过 Ollama HTTP API 调用本地模型进行文本生成。
"""
import httpx
import structlog
from app.config import settings
logger = structlog.get_logger()
class OllamaClient:
"""Ollama HTTP API 客户端"""
def __init__(
self,
base_url: str | None = None,
model: str | None = None,
timeout: float = 120.0,
) -> None:
self.base_url = (base_url or settings.ollama_base_url).rstrip("/")
self.model = model or settings.ollama_model
self.timeout = timeout
async def generate(self, prompt: str, json_mode: bool = False) -> str:
"""调用 Ollama 生成文本
Args:
prompt: 输入提示词
json_mode: 为 True 时启用 Ollama 原生 JSON 约束输出(payload 加 format: json
Returns:
生成的文本内容
"""
url = f"{self.base_url}/api/generate"
payload = {
"model": self.model,
"prompt": prompt,
"stream": False,
}
if json_mode:
payload["format"] = "json"
async with httpx.AsyncClient(timeout=self.timeout) as client:
resp = await client.post(url, json=payload)
resp.raise_for_status()
data = resp.json()
result = data.get("response", "")
logger.debug("Ollama 生成完成", model=self.model, output_length=len(result))
return result
async def is_available(self) -> bool:
"""检查 Ollama 服务是否可用"""
try:
async with httpx.AsyncClient(timeout=5.0) as client:
resp = await client.get(f"{self.base_url}/api/tags")
return resp.status_code == 200
except httpx.HTTPError:
return False
+351
View File
@@ -0,0 +1,351 @@
"""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)
+75
View File
@@ -0,0 +1,75 @@
"""Redis 缓存服务
为检索结果与 query 解析结果提供 JSON 缓存。所有操作容错:
Redis 不可用或缓存数据异常时降级为未命中 / 写入失败,绝不影响主流程。
"""
import json
from typing import Any
import structlog
from redis import asyncio as redis_async
from app.config import settings
logger = structlog.get_logger()
class RedisCache:
"""Redis JSON 缓存客户端(懒连接,全操作容错)"""
def __init__(self) -> None:
self._client: redis_async.Redis | None = None
def _get_client(self) -> redis_async.Redis:
"""懒创建 Redis 客户端(TCP 连接在首次执行命令时才建立)"""
if self._client is None:
self._client = redis_async.from_url(settings.redis_url, decode_responses=True)
return self._client
async def get_json(self, key: str) -> dict[str, Any] | None:
"""读取缓存并反序列化为 dict
任何异常(连接失败 / JSON 解析失败)或值非 dict 时都降级为未命中,返回 None。
"""
try:
raw = await self._get_client().get(key)
if raw is None:
return None
data = json.loads(raw)
return data if isinstance(data, dict) else None
except Exception:
logger.warning("Redis 读取缓存失败,降级为未命中", key=key, exc_info=True)
return None
async def set_json(self, key: str, value: dict[str, Any], ttl: int | None = None) -> bool:
"""写入缓存(SETEX),ttl 缺省取 settings.cache_ttl
异常时记录日志并返回 False,不影响主流程。
"""
try:
await self._get_client().setex(
key, ttl if ttl is not None else settings.cache_ttl, json.dumps(value, ensure_ascii=False)
)
return True
except Exception:
logger.warning("Redis 写入缓存失败", key=key, exc_info=True)
return False
async def close(self) -> None:
"""关闭底层连接"""
if self._client is not None:
await self._client.aclose()
self._client = None
# 模块级懒加载单例,避免每请求重建客户端
_cache: RedisCache | None = None
def get_cache() -> RedisCache:
"""获取全局 RedisCache 单例"""
global _cache
if _cache is None:
_cache = RedisCache()
return _cache