6ae91679e2
- .env.example 嵌入默认改为本地 BGE-M3(EMBEDDING_PROVIDER=local, dim=1024),NAS 全新部署即用 - 新增 app/services/reranker.py:基于 Ollama /api/rerank 的 cross-encoder 重排服务 - retriever 在 chunk 候选阶段接入语义精排,调用失败优雅降级为 RRF 顺序 - docker-compose 拉取 qwen3-reranker:0.6b,OLLAMA_MAX_LOADED_MODELS 提至 3 - .gitignore 排除 .workbuddy/ 与 .trae/(防止误提交项目记忆与 IDE spec) - 新增 reranker 单元与集成测试,全量测试 529 passed
218 lines
9.9 KiB
Python
218 lines
9.9 KiB
Python
"""分层检索引擎
|
||
|
||
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 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,
|
||
)
|
||
from app.services.reranker import RerankerService, create_reranker_service
|
||
|
||
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,
|
||
reranker: RerankerService | 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()
|
||
self.reranker = reranker if reranker is not None else create_reranker_service()
|
||
|
||
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(await self._rerank(query_text, 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 = await self._rerank(query_text, 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
|
||
|
||
async def _rerank(
|
||
self, query_text: str, points: list[models.ScoredPoint], final_k: int
|
||
) -> list[models.ScoredPoint]:
|
||
"""对 chunk 候选池做语义精排(cross-encoder)
|
||
|
||
重排未启用、候选为空或调用异常时,退化为 RRF 融合的原始顺序(取前 final_k),
|
||
保证检索链路在任何情况下都不因重排失败而中断。
|
||
"""
|
||
if self.reranker is None or not points:
|
||
return points[:final_k]
|
||
documents = [(p.payload or {}).get("text", "") for p in points]
|
||
try:
|
||
scores = await self.reranker.rerank(query_text, documents)
|
||
except Exception:
|
||
logger.warning("重排调用失败,退化为 RRF 融合顺序", exc_info=True)
|
||
return points[:final_k]
|
||
# 按相关性分数降序重排,写回 score 字段为相关性分数
|
||
ordered = sorted(range(len(points)), key=lambda i: scores[i], reverse=True)
|
||
return [points[i].model_copy(update={"score": scores[i]}) for i in ordered[: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,
|
||
)
|