Files
QMDSearch/app/core/retriever.py
T
kplam 2ab8b56a01 feat: 完成全量功能开发,包括前端管理后台与后端服务优化
此提交实现了完整的知识库管理系统:
1. 新增Vue3 + Antd Vue前端管理后台,包含登录、文档管理、检索、类目设置等完整页面
2. 重构后端LLM调用抽象层,支持Ollama与OpenAI兼容服务动态切换
3. 调整默认嵌入模型配置为本地bge-m3模式
4. 优化入库任务去重逻辑与缓存清理机制
5. 完善Docker镜像构建与docker-compose部署配置
6. 修复多项测试用例与兼容性问题
7. 新增运行时配置API,支持动态调整系统参数
2026-07-31 12:05:25 +08:00

195 lines
8.6 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.
"""分层检索引擎
L1(文档总结)→ L2(章节大纲)→ L3(小节定位)→ chunks(原文)逐层收窄:
- L1 无候选文档:直接全库 chunk 兜底,fallback=True
- L2 无命中:全部候选文档回退到 L3 的 doc 级查询(b 路)
- 2.5 级文档无 L2 节点,天然落入 L3 b 路(仅按 doc 过滤)
- L3 两路(aL2 命中文档按 section 过滤;b:其余文档仅按 doc 过滤)RRF 融合;
L3 无命中时 chunk 层回退为 L1 候选文档级检索
"""
from collections.abc import Iterable
import structlog
from qdrant_client import models
from app.config import settings
from app.core.embeddings import EmbeddingService, create_embedding_service
from app.core.query_parser import QueryParser, RouteDecision
from app.core.ranker import finalize, rrf_fuse
from app.core.result_summarizer import ResultSummarizer
from app.core.sparse import SparseEncoder
from app.models.knowledge import load_taxonomy
from app.models.search import ExtractedInfo, SearchHit, SearchRequest, SearchResponse
from app.services.llm import create_llm_client
from app.services.qdrant import (
COLLECTION_CHUNKS,
COLLECTION_L1,
COLLECTION_L2,
COLLECTION_L3,
SPARSE_COLLECTIONS,
QdrantService,
)
logger = structlog.get_logger()
def _unique(values: Iterable[str | None]) -> list[str]:
"""去重保序并丢弃空值"""
return list(dict.fromkeys(v for v in values if v))
class Retriever:
"""分层检索引擎,依赖均可注入(默认自建,便于测试替换)"""
def __init__(
self,
qdrant: QdrantService | None = None,
query_parser: QueryParser | None = None,
embedding: EmbeddingService | None = None,
sparse_encoder: SparseEncoder | None = None,
result_summarizer: ResultSummarizer | None = None,
) -> None:
self.qdrant = qdrant or QdrantService()
self.query_parser = query_parser or QueryParser(
ollama=create_llm_client("query"),
taxonomy=load_taxonomy(settings.taxonomy_path),
)
self.embedding = embedding or create_embedding_service()
self.sparse_encoder = sparse_encoder or SparseEncoder()
self.result_summarizer = result_summarizer or ResultSummarizer()
async def search(self, request: SearchRequest) -> SearchResponse:
"""分层检索主流程"""
route = await self.query_parser.parse_and_route(request.query)
query_text = route.parsed.rewrite or request.query
dense = (await self.embedding.embed([query_text]))[0]
sparse = self.sparse_encoder.encode(query_text) if settings.sparse_enabled else None
# 路由兜底时不做类目过滤,避免丢召回
categories = None if route.fallback else route.filter_categories
# L1:文档级检索,产出候选文档
l1_filter = QdrantService.build_filter(categories=categories)
l1_hits = await self._search_collection(COLLECTION_L1, dense, sparse, settings.l1_doc_top_n, l1_filter)
logger.info("L1 检索完成", hits=len(l1_hits), categories=categories)
if not l1_hits:
# L1 无候选文档 → 全库 chunk 兜底
chunk_hits = await self._search_collection(COLLECTION_CHUNKS, dense, sparse, settings.retrieval_top_k, None)
logger.info("L1 无命中,全库 chunk 兜底", hits=len(chunk_hits))
hits = self._to_hits(finalize(chunk_hits, self._final_k(request)))
return await self._build_response(request, route, hits, fallback=True)
doc_ids = _unique((p.payload or {}).get("doc_id") for p in l1_hits)
# L2:候选文档内检索章节大纲
l2_filter = QdrantService.build_filter(categories=categories, doc_ids=doc_ids)
l2_hits = await self._search_collection(
COLLECTION_L2, dense, sparse, settings.l2_section_top_n * len(doc_ids), l2_filter
)
logger.info("L2 检索完成", hits=len(l2_hits))
l2_doc_ids = _unique((p.payload or {}).get("doc_id") for p in l2_hits)
l2_section_paths = _unique((p.payload or {}).get("section_path") for p in l2_hits)
# 无 L2 命中的文档(2.5 级文档或该 doc 的 L2 未命中)走 L3 b 路
remaining_doc_ids = [d for d in doc_ids if d not in l2_doc_ids]
# L3:两路查询后 RRF 融合
l3_lists: list[list[models.ScoredPoint]] = []
if l2_doc_ids:
l3_filter_a = QdrantService.build_filter(
categories=categories, doc_ids=l2_doc_ids, section_paths=l2_section_paths
)
l3_lists.append(await self._search_collection(COLLECTION_L3, dense, sparse, settings.l3_top_n, l3_filter_a))
if remaining_doc_ids:
l3_filter_b = QdrantService.build_filter(categories=categories, doc_ids=remaining_doc_ids)
l3_lists.append(await self._search_collection(COLLECTION_L3, dense, sparse, settings.l3_top_n, l3_filter_b))
l3_hits = rrf_fuse(l3_lists)
logger.info("L3 检索完成", hits=len(l3_hits))
# chunk 层:L3 有命中按 section 收窄;无命中回退为 L1 候选文档级检索
if l3_hits:
l3_doc_ids = _unique((p.payload or {}).get("doc_id") for p in l3_hits)
# L3 命中 section_path 全为空(如 2.5 级文档)时仅按 doc 过滤
l3_section_paths = _unique((p.payload or {}).get("section_path") for p in l3_hits)
chunk_filter = QdrantService.build_filter(
categories=categories, doc_ids=l3_doc_ids, section_paths=l3_section_paths
)
else:
logger.info("L3 无命中,回退到 L1 候选文档级 chunk 检索")
chunk_filter = QdrantService.build_filter(categories=categories, doc_ids=doc_ids)
chunk_hits = await self._search_collection(
COLLECTION_CHUNKS, dense, sparse, settings.retrieval_top_k, chunk_filter
)
logger.info("chunk 检索完成", hits=len(chunk_hits))
final_points = finalize(rrf_fuse([chunk_hits]), self._final_k(request))
hits = self._to_hits(final_points)
return await self._build_response(request, route, hits, fallback=route.fallback)
async def _search_collection(
self,
collection: str,
dense: list[float],
sparse: tuple[list[int], list[float]] | None,
limit: int,
query_filter: models.Filter | None,
) -> list[models.ScoredPoint]:
"""按集合是否支持 sparse 选择 hybrid 或 dense 检索"""
if sparse is not None and collection in SPARSE_COLLECTIONS:
return await self.qdrant.search_hybrid(collection, dense, sparse, limit, query_filter=query_filter)
return await self.qdrant.search_dense(collection, dense, limit, query_filter=query_filter)
@staticmethod
def _final_k(request: SearchRequest) -> int:
"""最终返回数:请求指定优先,否则用配置默认值"""
return request.top_k or settings.retrieval_final_k
@staticmethod
def _to_hits(points: list[models.ScoredPoint]) -> list[SearchHit]:
"""chunk 点组装为 SearchHit,字段取自 chunk payloaddoc_summary 仅上下文标注)"""
hits: list[SearchHit] = []
for p in points:
payload = p.payload or {}
hits.append(
SearchHit(
text=payload.get("text", ""),
doc_id=payload.get("doc_id", ""),
title=payload.get("title", ""),
section_path=payload.get("section_path") or "",
score=p.score,
doc_summary=payload.get("doc_summary") or "",
)
)
return hits
async def _build_response(
self,
request: SearchRequest,
route: RouteDecision,
hits: list[SearchHit],
*,
fallback: bool,
) -> SearchResponse:
"""构造最终响应:组装 AI 提取信息,按需生成结果总结"""
parsed = route.parsed
extracted = ExtractedInfo(
rewrite=parsed.rewrite,
keywords=parsed.keywords,
entities=parsed.entities,
intent=parsed.intent,
time_range=parsed.time_range,
categories=[c.name for c in parsed.categories],
)
summary: str | None = None
if request.summarize:
summary = await self.result_summarizer.summarize(request.query, hits)
return SearchResponse(
query=request.query,
hits=hits,
routed_categories=route.filter_categories or [],
fallback=fallback,
extracted_info=extracted,
summary=summary,
)