Files
2026-06-08 18:14:59 +08:00

667 lines
24 KiB
Python
Raw Permalink 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.
"""
检索服务
支持向量检索、全文检索、混合检索(RRF 融合)
向量检索通过 Qdrant 向量数据库实现,全文检索通过业务数据库 SQL 实现
"""
import logging
import time
from typing import List, Optional, Dict, Any
from sqlalchemy import select, func
from sqlalchemy.ext.asyncio import AsyncSession
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
from ai_platform.knowledge.services.embedding_service import EmbeddingService
from ai_platform.knowledge.schemas.segment_schema import RetrievalResult
from ai_platform.knowledge.vector_store import get_vector_store
logger = logging.getLogger(__name__)
# RRF 融合常数
RRF_K = 60
class RetrievalService:
"""
检索服务
支持三种检索模式:
- vector: 纯向量检索(通过 Qdrant
- fulltext: 纯全文检索(通过业务数据库 LIKE + 关键词匹配)
- hybrid: 混合检索(向量 + 全文,RRF 融合排序)
"""
def __init__(self, db: AsyncSession):
self._db = db
self._embedding_service = EmbeddingService(db)
self._vector_store = get_vector_store()
async def retrieve(
self,
query: str,
knowledge_base_ids: List[str],
top_k: int = 5,
score_threshold: float = 0.5,
retrieval_mode: Optional[str] = None,
rerank_enabled: Optional[bool] = None,
rerank_model_id: Optional[str] = None,
metadata_filter: Optional[Dict[str, Any]] = None,
) -> List[RetrievalResult]:
"""
检索知识库
Args:
query: 查询文本
knowledge_base_ids: 知识库 ID 列表
top_k: 返回数量
score_threshold: 相似度阈值
retrieval_mode: 检索模式(不传则使用第一个知识库的配置)
rerank_enabled: 是否启用重排序(不传则使用知识库配置)
rerank_model_id: 重排序模型 ID(不传则使用知识库配置)
metadata_filter: 元数据过滤条件
Returns:
检索结果列表
"""
if not query or not knowledge_base_ids:
return []
start_time = time.time()
# 获取知识库配置
kb_map = await self._get_knowledge_bases(knowledge_base_ids)
if not kb_map:
return []
first_kb = list(kb_map.values())[0]
# 经济模式强制使用全文检索
is_economy = getattr(first_kb, 'indexing_technique', 'high_quality') == 'economy'
# 确定检索模式
if is_economy:
retrieval_mode = 'fulltext'
elif not retrieval_mode:
retrieval_mode = first_kb.retrieval_mode or 'hybrid'
# 确定 rerank 配置(参数优先,否则使用知识库配置;经济模式禁用 rerank)
if is_economy:
rerank_enabled = False
else:
if rerank_enabled is None:
rerank_enabled = first_kb.rerank_enabled or False
if not rerank_model_id:
rerank_model_id = first_kb.rerank_model_id
# 获取 embedding 模型(使用第一个知识库的配置)
embedding_model_id = first_kb.embedding_model_id
# 如果启用了 rerank,初始检索多取一些候选结果
candidate_multiplier = 3 if rerank_enabled and rerank_model_id else 2
candidate_top_k = top_k * candidate_multiplier
results = []
if retrieval_mode == 'vector':
results = await self._vector_search(
query, knowledge_base_ids, embedding_model_id,
top_k=candidate_top_k, score_threshold=score_threshold,
dimensions=first_kb.embedding_dimensions,
metadata_filter=metadata_filter,
)
for r in results:
r.match_source = 'vector'
elif retrieval_mode == 'fulltext':
results = await self._fulltext_search(
query, knowledge_base_ids, top_k=candidate_top_k,
metadata_filter=metadata_filter,
)
for r in results:
r.match_source = 'fulltext'
elif retrieval_mode == 'hybrid':
# 混合检索:向量 + 全文,RRF 融合
vector_results = await self._vector_search(
query, knowledge_base_ids, embedding_model_id,
top_k=candidate_top_k, score_threshold=score_threshold,
dimensions=first_kb.embedding_dimensions,
metadata_filter=metadata_filter,
)
for r in vector_results:
r.match_source = 'vector'
fulltext_results = await self._fulltext_search(
query, knowledge_base_ids, top_k=candidate_top_k,
metadata_filter=metadata_filter,
)
for r in fulltext_results:
r.match_source = 'fulltext'
results = self._rrf_merge(vector_results, fulltext_results)
# 多知识库权重加权
if len(kb_map) > 1:
for r in results:
kb = kb_map.get(r.knowledge_base_id)
weight = getattr(kb, 'retrieval_weight', 1.0) or 1.0 if kb else 1.0
if weight != 1.0:
r.score = round(r.score * weight, 4)
results.sort(key=lambda x: x.score, reverse=True)
# 过滤低分结果(rerank 前先粗筛)
if not rerank_enabled:
results = [r for r in results if r.score >= score_threshold]
# Rerank 重排序
if rerank_enabled and rerank_model_id and results:
results = await self._rerank_results(query, results, rerank_model_id, top_k)
# rerank 后再按阈值过滤
results = [r for r in results if r.score >= score_threshold]
# 截断到 top_k
results = results[:top_k]
# 内容级去重(多知识库检索时可能有重复内容)
results = self._deduplicate_results(results)
# 标注优先匹配:将匹配到的标注结果插入到最前面
annotation_results = await self._match_annotations(
query, knowledge_base_ids, embedding_model_id,
score_threshold=score_threshold,
dimensions=first_kb.embedding_dimensions,
)
if annotation_results:
for r in annotation_results:
r.match_source = 'annotation'
# 标注结果置顶,去重后合并
existing_ids = {r.segment_id for r in annotation_results}
results = annotation_results + [r for r in results if r.segment_id not in existing_ids]
results = results[:top_k]
# 填充知识库名称和文档名称
await self._fill_names(results, kb_map)
# 填充父分段内容(Small-to-Big 模式)
await self._fill_parent_content(results)
# 更新命中次数
segment_ids = [r.segment_id for r in results]
if segment_ids:
await self._update_hit_counts(segment_ids)
elapsed = int((time.time() - start_time) * 1000)
rerank_info = ', rerank=ON' if rerank_enabled else ''
logger.info(f'检索完成: {len(results)} 条结果, 耗时 {elapsed}ms, 模式={retrieval_mode}{rerank_info}')
return results
async def _rerank_results(
self,
query: str,
results: List[RetrievalResult],
rerank_model_id: str,
top_n: int,
) -> List[RetrievalResult]:
"""使用 Rerank 模型对检索结果重排序"""
from ai_platform.knowledge.services.rerank_service import RerankService
try:
rerank_service = RerankService(self._db)
documents = [r.content for r in results]
rerank_results = await rerank_service.rerank(
model_id=rerank_model_id,
query=query,
documents=documents,
top_n=top_n,
)
# 按 rerank 分数重新排列结果
reranked = []
for rr in rerank_results:
if 0 <= rr.index < len(results):
result = results[rr.index]
result.score = round(rr.relevance_score, 4)
reranked.append(result)
logger.info(f'Rerank 完成: {len(results)} -> {len(reranked)} 条结果')
return reranked
except Exception as e:
logger.warning(f'Rerank 失败,使用原始排序: {e}')
return results
async def _vector_search(
self,
query: str,
knowledge_base_ids: List[str],
embedding_model_id: str,
top_k: int = 10,
score_threshold: float = 0.0,
dimensions: Optional[int] = None,
metadata_filter: Optional[Dict[str, Any]] = None,
) -> List[RetrievalResult]:
"""向量检索(通过 Qdrant 余弦相似度)"""
if not embedding_model_id:
logger.warning('未配置 Embedding 模型,跳过向量检索')
return []
try:
query_embedding = await self._embedding_service.embed_text(
model_id=embedding_model_id,
text=query,
dimensions=dimensions,
)
except Exception as e:
logger.error(f'查询向量化失败: {e}')
return []
# 对每个知识库分别搜索(每个知识库对应一个 Qdrant collection
all_hits = []
for kb_id in knowledge_base_ids:
hits = await self._vector_store.search(
knowledge_base_id=kb_id,
query_vector=query_embedding,
top_k=top_k,
score_threshold=score_threshold,
)
all_hits.extend(hits)
if not all_hits:
return []
# 按分数排序
all_hits.sort(key=lambda h: h.score, reverse=True)
all_hits = all_hits[:top_k]
# 从业务数据库获取分段详情
segment_ids = [h.id for h in all_hits]
score_map = {h.id: h.score for h in all_hits}
seg_result = await self._db.execute(
select(KnowledgeSegment).where(
KnowledgeSegment.id.in_(segment_ids),
KnowledgeSegment.is_deleted == False,
KnowledgeSegment.enabled == True,
)
)
segments = {str(s.id): s for s in seg_result.scalars().all()}
results = []
for hit in all_hits:
seg = segments.get(hit.id)
if not seg:
continue
# Q&A 模式:question 用于匹配,返回 answer 作为 content
content = seg.answer if seg.answer else seg.content
meta = dict(seg.extra_metadata) if seg.extra_metadata else {}
if seg.answer:
meta['question'] = seg.content
meta['chunk_mode'] = 'qa'
results.append(RetrievalResult(
segment_id=str(seg.id),
document_id=str(seg.document_id),
knowledge_base_id=str(seg.knowledge_base_id),
content=content,
score=round(score_map.get(hit.id, 0.0), 4),
token_count=seg.token_count or 0,
page_number=seg.page_number,
metadata=meta,
keywords=seg.keywords,
))
return results
async def _fulltext_search(
self,
query: str,
knowledge_base_ids: List[str],
top_k: int = 10,
metadata_filter: Optional[Dict[str, Any]] = None,
) -> List[RetrievalResult]:
"""全文检索(基于 ORM LIKE 和关键词匹配,兼容所有数据库)"""
import re
keywords = re.split(r'[\s,,。.!?;;、]+', query)
keywords = [k.strip() for k in keywords if k.strip() and len(k.strip()) >= 2]
if not keywords:
keywords = [query.strip()]
# 使用 SQLAlchemy ORM 构建查询(兼容 PG / MySQL 等)
from sqlalchemy import or_
keyword_conditions = [
func.lower(KnowledgeSegment.content).contains(kw.lower())
for kw in keywords[:5]
]
conditions = [
KnowledgeSegment.knowledge_base_id.in_(knowledge_base_ids),
KnowledgeSegment.is_deleted == False,
KnowledgeSegment.enabled == True,
or_(*keyword_conditions),
]
# 元数据过滤
if metadata_filter:
conditions.extend(self._build_metadata_conditions(metadata_filter))
stmt = (
select(KnowledgeSegment)
.where(*conditions)
.order_by(KnowledgeSegment.char_count.asc())
.limit(top_k)
)
result = await self._db.execute(stmt)
segments = result.scalars().all()
# 计算简单的关键词匹配分数
results = []
for seg in segments:
content_lower = seg.content.lower()
match_count = sum(1 for kw in keywords if kw.lower() in content_lower)
score = match_count / len(keywords) if keywords else 0
# Q&A 模式:返回 answer 作为 content
content = seg.answer if seg.answer else seg.content
meta = dict(seg.extra_metadata) if seg.extra_metadata else {}
if seg.answer:
meta['question'] = seg.content
meta['chunk_mode'] = 'qa'
results.append(RetrievalResult(
segment_id=str(seg.id),
document_id=str(seg.document_id),
knowledge_base_id=str(seg.knowledge_base_id),
content=content,
score=round(score, 4),
token_count=seg.token_count or 0,
page_number=seg.page_number,
metadata=meta,
keywords=seg.keywords,
))
results.sort(key=lambda x: x.score, reverse=True)
return results
def _rrf_merge(
self,
vector_results: List[RetrievalResult],
fulltext_results: List[RetrievalResult],
) -> List[RetrievalResult]:
"""
RRF (Reciprocal Rank Fusion) 融合排序
RRF_score = sum(1 / (k + rank_i)) for each result list
"""
scores = {} # segment_id -> (rrf_score, result)
# 向量检索结果排名
for rank, result in enumerate(vector_results):
rrf_score = 1.0 / (RRF_K + rank + 1)
if result.segment_id in scores:
old_score, old_result = scores[result.segment_id]
scores[result.segment_id] = (old_score + rrf_score, old_result)
else:
scores[result.segment_id] = (rrf_score, result)
# 全文检索结果排名
for rank, result in enumerate(fulltext_results):
rrf_score = 1.0 / (RRF_K + rank + 1)
if result.segment_id in scores:
old_score, old_result = scores[result.segment_id]
scores[result.segment_id] = (old_score + rrf_score, old_result)
else:
scores[result.segment_id] = (rrf_score, result)
# 按 RRF 分数排序
sorted_items = sorted(scores.values(), key=lambda x: x[0], reverse=True)
if not sorted_items:
return []
# 归一化分数到 0-1
# RRF 单条结果的理论最大分数为 2/(k+1)(同时出现在两个列表的第一名)
# 使用理论最大值归一化,避免单条结果被归一化为 100%
theoretical_max = 2.0 / (RRF_K + 1)
results = []
for rrf_score, result in sorted_items:
normalized_score = min(rrf_score / theoretical_max, 1.0)
result.score = round(normalized_score, 4)
results.append(result)
return results
@staticmethod
def _deduplicate_results(results: List[RetrievalResult], similarity_threshold: float = 0.95) -> List[RetrievalResult]:
"""
内容级去重(多知识库检索时可能有重复内容)
使用内容前 200 字符的相似度判断是否重复,保留分数最高的。
"""
if len(results) <= 1:
return results
deduplicated = []
seen_contents = []
for r in results:
content_key = r.content[:200].strip().lower()
is_dup = False
for seen in seen_contents:
# 简单的字符重叠率判断
if content_key == seen:
is_dup = True
break
# 如果前 200 字符有 95% 以上重叠,视为重复
shorter = min(len(content_key), len(seen))
if shorter > 0:
common = sum(1 for a, b in zip(content_key, seen) if a == b)
if common / shorter >= similarity_threshold:
is_dup = True
break
if not is_dup:
deduplicated.append(r)
seen_contents.append(content_key)
return deduplicated
async def _fill_parent_content(self, results: List[RetrievalResult]):
"""填充父分段内容(Small-to-Big 模式)"""
if not results:
return
# 获取所有 segment_id,查询是否有 parent_segment_id
segment_ids = [r.segment_id for r in results if r.segment_id]
if not segment_ids:
return
seg_result = await self._db.execute(
select(KnowledgeSegment.id, KnowledgeSegment.parent_segment_id).where(
KnowledgeSegment.id.in_(segment_ids),
KnowledgeSegment.is_deleted == False,
)
)
parent_map = {}
for row in seg_result:
if row.parent_segment_id:
parent_map[str(row.id)] = row.parent_segment_id
if not parent_map:
return
# 批量获取父分段内容
parent_ids = list(set(parent_map.values()))
parent_result = await self._db.execute(
select(KnowledgeSegment.id, KnowledgeSegment.content).where(
KnowledgeSegment.id.in_(parent_ids),
KnowledgeSegment.is_deleted == False,
)
)
parent_content_map = {str(row.id): row.content for row in parent_result}
# 填充到结果中
for r in results:
parent_id = parent_map.get(r.segment_id)
if parent_id:
r.parent_content = parent_content_map.get(str(parent_id))
@staticmethod
def _build_metadata_conditions(metadata_filter: Dict[str, Any]) -> list:
"""构建元数据过滤条件(基于 JSON 字段,跨数据库兼容)"""
from app.db_compat import json_extract
conditions = []
for key, value in metadata_filter.items():
if value is not None:
# 使用跨数据库兼容的 json_extract 函数
try:
conditions.append(
json_extract(KnowledgeSegment.extra_metadata, key) == str(value)
)
except Exception:
pass
return conditions
async def _get_knowledge_bases(self, kb_ids: List[str]) -> Dict[str, KnowledgeBase]:
"""批量获取知识库"""
result = await self._db.execute(
select(KnowledgeBase).where(
KnowledgeBase.id.in_(kb_ids),
KnowledgeBase.is_deleted == False,
)
)
kbs = result.scalars().all()
return {str(kb.id): kb for kb in kbs}
async def _fill_names(self, results: List[RetrievalResult], kb_map: Dict[str, KnowledgeBase]):
"""填充知识库名称和文档名称"""
if not results:
return
# 获取文档名称
doc_ids = list({r.document_id for r in results})
doc_result = await self._db.execute(
select(KnowledgeDocument.id, KnowledgeDocument.name).where(
KnowledgeDocument.id.in_(doc_ids)
)
)
doc_name_map = {row.id: row.name for row in doc_result}
for result in results:
result.document_name = doc_name_map.get(result.document_id, '')
kb = kb_map.get(result.knowledge_base_id)
result.knowledge_base_name = kb.name if kb else ''
async def _update_hit_counts(self, segment_ids: List[str]):
"""更新分段命中次数(兼容所有数据库)"""
if not segment_ids:
return
try:
from sqlalchemy import update
await self._db.execute(
update(KnowledgeSegment)
.where(KnowledgeSegment.id.in_(segment_ids))
.values(hit_count=func.coalesce(KnowledgeSegment.hit_count, 0) + 1)
)
await self._db.commit()
except Exception as e:
logger.warning(f'更新命中次数失败: {e}')
async def _match_annotations(
self,
query: str,
knowledge_base_ids: List[str],
embedding_model_id: Optional[str],
score_threshold: float = 0.5,
dimensions: Optional[int] = None,
max_results: int = 3,
) -> List[RetrievalResult]:
"""
匹配标注(Q&A 对)
通过向量相似度匹配标注的 question,返回对应的 answer。
标注结果优先级高于普通分段。
"""
from ai_platform.knowledge.models import KnowledgeAnnotation
if not embedding_model_id:
return []
try:
# 向量化查询
query_embedding = await self._embedding_service.embed_text(
model_id=embedding_model_id,
text=query,
dimensions=dimensions if dimensions else None,
)
# 在 Qdrant 中搜索标注向量(payload.type == 'annotation'
all_hits = []
for kb_id in knowledge_base_ids:
try:
hits = await self._vector_store.search(
knowledge_base_id=kb_id,
query_vector=query_embedding,
top_k=max_results,
score_threshold=score_threshold,
filter_conditions={'type': 'annotation'},
)
all_hits.extend(hits)
except Exception as e:
logger.warning(f'标注向量搜索失败 (kb={kb_id}): {e}')
continue
if not all_hits:
return []
# 按分数排序取 top
all_hits.sort(key=lambda h: h.score, reverse=True)
all_hits = all_hits[:max_results]
# 从数据库获取标注详情
annotation_ids = [h.id for h in all_hits]
score_map = {h.id: h.score for h in all_hits}
ann_result = await self._db.execute(
select(KnowledgeAnnotation).where(
KnowledgeAnnotation.id.in_(annotation_ids),
KnowledgeAnnotation.is_deleted == False,
KnowledgeAnnotation.enabled == True,
)
)
annotations = {str(a.id): a for a in ann_result.scalars().all()}
results = []
for hit in all_hits:
ann = annotations.get(hit.id)
if not ann:
continue
# 标注结果:content 返回 answersegment_id 用 annotation id
results.append(RetrievalResult(
segment_id=str(ann.id),
document_id='',
document_name='[Q&A]',
knowledge_base_id=str(ann.knowledge_base_id),
content=ann.answer,
score=round(score_map.get(hit.id, 0.0), 4),
token_count=0,
metadata={'type': 'annotation', 'question': ann.question},
))
# 更新标注命中次数
if annotation_ids:
try:
from sqlalchemy import update
await self._db.execute(
update(KnowledgeAnnotation)
.where(KnowledgeAnnotation.id.in_(annotation_ids))
.values(hit_count=func.coalesce(KnowledgeAnnotation.hit_count, 0) + 1)
)
await self._db.commit()
except Exception:
pass
return results
except Exception as e:
logger.warning(f'标注匹配失败: {e}')
return []