Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈 - L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合) - 文档三级总结与 2.5 级回退 - query 解析路由与分类 - /admin 管理页面
This commit is contained in:
@@ -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
|
||||
@@ -0,0 +1,351 @@
|
||||
"""Qdrant 向量数据库服务封装
|
||||
|
||||
分层摘要索引 RAG 的存储层:
|
||||
- doc_l1:文档级总结(dense + sparse)
|
||||
- doc_l2 / doc_l3:大纲节点(dense)
|
||||
- chunks:原文分块(dense + sparse)
|
||||
|
||||
提供集合初始化(幂等)、按层 upsert、dense / hybrid(RRF 融合)检索接口。
|
||||
"""
|
||||
|
||||
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_paths:must 条件 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)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user