667 lines
24 KiB
Python
667 lines
24 KiB
Python
"""
|
||
检索服务
|
||
|
||
支持向量检索、全文检索、混合检索(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 返回 answer,segment_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 []
|