"""分层检索引擎 L1(文档总结)→ L2(章节大纲)→ L3(小节定位)→ chunks(原文)逐层收窄: - L1 无候选文档:直接全库 chunk 兜底,fallback=True - L2 无命中:全部候选文档回退到 L3 的 doc 级查询(b 路) - 2.5 级文档无 L2 节点,天然落入 L3 b 路(仅按 doc 过滤) - L3 两路(a:L2 命中文档按 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 payload(doc_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, )