Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
"""
|
||||
知识库服务
|
||||
"""
|
||||
@@ -0,0 +1,159 @@
|
||||
"""
|
||||
文档预处理/清洗服务
|
||||
|
||||
参考 Dify 的 DatasetProcessRule,支持可配置的文本清洗规则。
|
||||
在文本提取之后、分块之前执行。
|
||||
"""
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 默认预处理规则
|
||||
DEFAULT_PROCESS_RULES: Dict[str, Any] = {
|
||||
"pre_processing_rules": [
|
||||
{"id": "remove_extra_spaces", "enabled": True},
|
||||
{"id": "remove_urls_emails", "enabled": False},
|
||||
{"id": "remove_html_tags", "enabled": False},
|
||||
{"id": "remove_consecutive_newlines", "enabled": True},
|
||||
{"id": "remove_trailing_whitespace", "enabled": True},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class CleaningService:
|
||||
"""
|
||||
文本清洗服务
|
||||
|
||||
支持的清洗规则:
|
||||
- remove_extra_spaces: 合并连续空格为单个空格
|
||||
- remove_urls_emails: 移除 URL 和邮箱地址
|
||||
- remove_html_tags: 移除 HTML 标签
|
||||
- remove_consecutive_newlines: 合并连续空行(3+)为双空行
|
||||
- remove_trailing_whitespace: 去除行尾空白
|
||||
"""
|
||||
|
||||
# 规则处理器映射
|
||||
RULE_PROCESSORS = {
|
||||
"remove_extra_spaces": "_remove_extra_spaces",
|
||||
"remove_urls_emails": "_remove_urls_emails",
|
||||
"remove_html_tags": "_remove_html_tags",
|
||||
"remove_consecutive_newlines": "_remove_consecutive_newlines",
|
||||
"remove_trailing_whitespace": "_remove_trailing_whitespace",
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def clean(cls, text: str, process_rules: Optional[Dict[str, Any]] = None) -> str:
|
||||
"""
|
||||
根据预处理规则清洗文本
|
||||
|
||||
Args:
|
||||
text: 原始文本
|
||||
process_rules: 预处理规则配置,为 None 则使用默认规则
|
||||
|
||||
Returns:
|
||||
清洗后的文本
|
||||
"""
|
||||
if not text:
|
||||
return text
|
||||
|
||||
rules = process_rules or DEFAULT_PROCESS_RULES
|
||||
pre_rules = rules.get("pre_processing_rules", [])
|
||||
|
||||
original_length = len(text)
|
||||
|
||||
for rule in pre_rules:
|
||||
rule_id = rule.get("id", "")
|
||||
enabled = rule.get("enabled", False)
|
||||
|
||||
if not enabled:
|
||||
continue
|
||||
|
||||
processor_name = cls.RULE_PROCESSORS.get(rule_id)
|
||||
if not processor_name:
|
||||
logger.warning(f"未知的预处理规则: {rule_id}")
|
||||
continue
|
||||
|
||||
processor = getattr(cls, processor_name, None)
|
||||
if processor:
|
||||
text = processor(text)
|
||||
|
||||
cleaned_length = len(text)
|
||||
if original_length != cleaned_length:
|
||||
logger.info(
|
||||
f"文本清洗完成: {original_length} -> {cleaned_length} 字符 "
|
||||
f"(减少 {original_length - cleaned_length})"
|
||||
)
|
||||
|
||||
return text.strip()
|
||||
|
||||
@staticmethod
|
||||
def _remove_extra_spaces(text: str) -> str:
|
||||
"""合并连续空格为单个空格(保留换行符)"""
|
||||
# 只处理同一行内的连续空格,不影响换行
|
||||
lines = text.split('\n')
|
||||
cleaned_lines = []
|
||||
for line in lines:
|
||||
cleaned_lines.append(re.sub(r'[ \t]+', ' ', line))
|
||||
return '\n'.join(cleaned_lines)
|
||||
|
||||
@staticmethod
|
||||
def _remove_urls_emails(text: str) -> str:
|
||||
"""移除 URL 和邮箱地址"""
|
||||
# 移除 URL
|
||||
text = re.sub(
|
||||
r'https?://[^\s<>"{}|\\^`\[\]]+',
|
||||
'',
|
||||
text,
|
||||
)
|
||||
# 移除邮箱
|
||||
text = re.sub(
|
||||
r'[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}',
|
||||
'',
|
||||
text,
|
||||
)
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def _remove_html_tags(text: str) -> str:
|
||||
"""移除 HTML 标签,保留文本内容"""
|
||||
# 移除 script 和 style 标签及其内容
|
||||
text = re.sub(r'<script[^>]*>.*?</script>', '', text, flags=re.DOTALL | re.IGNORECASE)
|
||||
text = re.sub(r'<style[^>]*>.*?</style>', '', text, flags=re.DOTALL | re.IGNORECASE)
|
||||
# 移除所有 HTML 标签
|
||||
text = re.sub(r'<[^>]+>', '', text)
|
||||
# 解码常见 HTML 实体
|
||||
text = text.replace(' ', ' ')
|
||||
text = text.replace('<', '<')
|
||||
text = text.replace('>', '>')
|
||||
text = text.replace('&', '&')
|
||||
text = text.replace('"', '"')
|
||||
text = text.replace(''', "'")
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def _remove_consecutive_newlines(text: str) -> str:
|
||||
"""合并连续空行(3个以上换行)为双换行"""
|
||||
return re.sub(r'\n{3,}', '\n\n', text)
|
||||
|
||||
@staticmethod
|
||||
def _remove_trailing_whitespace(text: str) -> str:
|
||||
"""去除每行行尾空白"""
|
||||
return '\n'.join(line.rstrip() for line in text.split('\n'))
|
||||
|
||||
@classmethod
|
||||
def get_default_rules(cls) -> Dict[str, Any]:
|
||||
"""获取默认预处理规则"""
|
||||
return DEFAULT_PROCESS_RULES.copy()
|
||||
|
||||
@classmethod
|
||||
def get_available_rules(cls) -> List[Dict[str, str]]:
|
||||
"""获取所有可用的预处理规则"""
|
||||
return [
|
||||
{"id": "remove_extra_spaces", "label": "合并连续空格"},
|
||||
{"id": "remove_urls_emails", "label": "移除 URL 和邮箱"},
|
||||
{"id": "remove_html_tags", "label": "移除 HTML 标签"},
|
||||
{"id": "remove_consecutive_newlines", "label": "合并连续空行"},
|
||||
{"id": "remove_trailing_whitespace", "label": "去除行尾空白"},
|
||||
]
|
||||
@@ -0,0 +1,425 @@
|
||||
"""
|
||||
文档服务
|
||||
|
||||
文档上传、管理、状态控制
|
||||
"""
|
||||
import logging
|
||||
from typing import Optional, List, Tuple
|
||||
|
||||
from sqlalchemy import select, func, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DocumentService:
|
||||
"""文档服务"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
|
||||
async def get_list(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
name: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeDocument], int]:
|
||||
"""获取文档列表"""
|
||||
query = select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
)
|
||||
|
||||
if name:
|
||||
query = query.where(KnowledgeDocument.name.ilike(f"%{name}%"))
|
||||
if status:
|
||||
query = query.where(KnowledgeDocument.status == status)
|
||||
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeDocument.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_by_id(self, doc_id: str) -> Optional[KnowledgeDocument]:
|
||||
"""获取文档详情"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.id == doc_id,
|
||||
KnowledgeDocument.is_deleted == False
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def add_document(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
file_id: str,
|
||||
name: Optional[str] = None,
|
||||
) -> KnowledgeDocument:
|
||||
"""
|
||||
添加文档到知识库
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
file_id: 文件管理系统中的文件 ID
|
||||
name: 文档名称(不传则从文件信息获取)
|
||||
"""
|
||||
# 验证知识库存在
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == knowledge_base_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if not kb:
|
||||
raise ValueError('知识库不存在')
|
||||
|
||||
# 获取文件信息
|
||||
from core.file_manager.model import FileManager
|
||||
file_result = await self._db.execute(
|
||||
select(FileManager).where(
|
||||
FileManager.id == file_id,
|
||||
FileManager.is_deleted == False
|
||||
)
|
||||
)
|
||||
file_info = file_result.scalar_one_or_none()
|
||||
if not file_info:
|
||||
raise ValueError('文件不存在')
|
||||
|
||||
# 检查是否已添加(通过 file_id 去重)
|
||||
existing = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.file_id == file_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
)
|
||||
)
|
||||
if existing.scalars().first():
|
||||
raise ValueError('该文件已添加到知识库')
|
||||
|
||||
# 通过文件 MD5 检测内容重复(跨知识库)
|
||||
duplicate_warning = None
|
||||
if file_info.md5:
|
||||
dup_result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.content_hash == file_info.md5,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
KnowledgeDocument.knowledge_base_id != knowledge_base_id,
|
||||
).limit(1)
|
||||
)
|
||||
dup_doc = dup_result.scalar_one_or_none()
|
||||
if dup_doc:
|
||||
duplicate_warning = f'该文件内容与其他知识库中的文档 "{dup_doc.name}" 重复'
|
||||
logger.info(f'文档内容重复检测: file_id={file_id}, 重复文档={dup_doc.id}')
|
||||
|
||||
# 同知识库内容去重(严格阻止)
|
||||
same_kb_dup = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.content_hash == file_info.md5,
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
).limit(1)
|
||||
)
|
||||
if same_kb_dup.scalar_one_or_none():
|
||||
raise ValueError('该知识库中已存在相同内容的文档')
|
||||
|
||||
doc = KnowledgeDocument(
|
||||
knowledge_base_id=knowledge_base_id,
|
||||
file_id=file_id,
|
||||
name=name or file_info.name,
|
||||
file_type=file_info.file_ext or '',
|
||||
file_size=file_info.size or 0,
|
||||
content_hash=file_info.md5 or '',
|
||||
status='pending',
|
||||
duplicate_warning=duplicate_warning,
|
||||
)
|
||||
self._db.add(doc)
|
||||
await self._db.commit()
|
||||
await self._db.refresh(doc)
|
||||
|
||||
return doc
|
||||
|
||||
async def batch_add_documents(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
file_ids: List[str],
|
||||
) -> List[KnowledgeDocument]:
|
||||
"""批量添加文档"""
|
||||
docs = []
|
||||
for file_id in file_ids:
|
||||
try:
|
||||
doc = await self.add_document(knowledge_base_id, file_id)
|
||||
docs.append(doc)
|
||||
except ValueError as e:
|
||||
logger.warning(f'添加文档失败 (file_id={file_id}): {e}')
|
||||
continue
|
||||
return docs
|
||||
|
||||
async def delete_document(self, doc_id: str) -> bool:
|
||||
"""删除文档(软删除,同时删除分段 + 清理 Qdrant 向量)"""
|
||||
doc = await self.get_by_id(doc_id)
|
||||
if not doc:
|
||||
return False
|
||||
|
||||
doc.is_deleted = True
|
||||
|
||||
# 软删除关联分段
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == doc_id
|
||||
).values(is_deleted=True)
|
||||
)
|
||||
|
||||
# 更新知识库统计
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service._update_kb_stats(doc.knowledge_base_id)
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
# 从 Qdrant 删除该文档的所有向量
|
||||
try:
|
||||
from ai_platform.knowledge.vector_store import get_vector_store
|
||||
vector_store = get_vector_store()
|
||||
await vector_store.delete_by_filter(
|
||||
doc.knowledge_base_id,
|
||||
filter_conditions={'document_id': doc_id},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f'从 Qdrant 删除文档向量失败: {e}')
|
||||
|
||||
return True
|
||||
|
||||
async def toggle_document(self, doc_id: str, enabled: bool) -> Optional[KnowledgeDocument]:
|
||||
"""启用/禁用文档"""
|
||||
doc = await self.get_by_id(doc_id)
|
||||
if not doc:
|
||||
return None
|
||||
|
||||
doc.enabled = enabled
|
||||
|
||||
# 同时启用/禁用关联分段
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == doc_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
).values(enabled=enabled)
|
||||
)
|
||||
|
||||
await self._db.commit()
|
||||
await self._db.refresh(doc)
|
||||
return doc
|
||||
|
||||
async def get_segments(
|
||||
self,
|
||||
document_id: str,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
keyword: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeSegment], int]:
|
||||
"""获取文档的分段列表"""
|
||||
query = select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == document_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
|
||||
if keyword:
|
||||
query = query.where(KnowledgeSegment.content.ilike(f"%{keyword}%"))
|
||||
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeSegment.position.asc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_kb_segments(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
keyword: Optional[str] = None,
|
||||
enabled: Optional[bool] = None,
|
||||
embedding_status: Optional[str] = None,
|
||||
metadata_key: Optional[str] = None,
|
||||
metadata_value: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeSegment], int]:
|
||||
"""获取知识库的所有分段"""
|
||||
from app.db_compat import json_extract, json_has_key
|
||||
|
||||
query = select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
|
||||
if keyword:
|
||||
query = query.where(KnowledgeSegment.content.ilike(f"%{keyword}%"))
|
||||
if enabled is not None:
|
||||
query = query.where(KnowledgeSegment.enabled == enabled)
|
||||
if embedding_status:
|
||||
query = query.where(KnowledgeSegment.embedding_status == embedding_status)
|
||||
if metadata_key:
|
||||
if metadata_value:
|
||||
query = query.where(
|
||||
json_extract(KnowledgeSegment.extra_metadata, metadata_key) == metadata_value
|
||||
)
|
||||
else:
|
||||
query = query.where(
|
||||
json_has_key(KnowledgeSegment.extra_metadata, metadata_key)
|
||||
)
|
||||
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeSegment.document_id, KnowledgeSegment.position.asc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def update_segment(
|
||||
self,
|
||||
segment_id: str,
|
||||
content: Optional[str] = None,
|
||||
keywords: Optional[List[str]] = None,
|
||||
enabled: Optional[bool] = None,
|
||||
extra_metadata: Optional[dict] = None,
|
||||
) -> Optional[KnowledgeSegment]:
|
||||
"""更新分段"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.id == segment_id,
|
||||
KnowledgeSegment.is_deleted == False
|
||||
)
|
||||
)
|
||||
segment = result.scalar_one_or_none()
|
||||
if not segment:
|
||||
return None
|
||||
|
||||
need_reindex = False
|
||||
if content is not None and content != segment.content:
|
||||
segment.content = content
|
||||
segment.char_count = len(content)
|
||||
segment.embedding_status = 'pending'
|
||||
need_reindex = True
|
||||
if keywords is not None:
|
||||
segment.keywords = keywords
|
||||
if enabled is not None:
|
||||
segment.enabled = enabled
|
||||
if extra_metadata is not None:
|
||||
segment.extra_metadata = extra_metadata
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
# 如果内容变更,重新向量化
|
||||
if need_reindex:
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service.index_segment(segment.knowledge_base_id, segment_id)
|
||||
|
||||
await self._db.refresh(segment)
|
||||
return segment
|
||||
|
||||
async def add_segment(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
document_id: str,
|
||||
content: str,
|
||||
keywords: Optional[List[str]] = None,
|
||||
) -> KnowledgeSegment:
|
||||
"""手动添加分段"""
|
||||
# 获取当前最大 position
|
||||
max_pos_result = await self._db.execute(
|
||||
select(func.max(KnowledgeSegment.position)).where(
|
||||
KnowledgeSegment.document_id == document_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
)
|
||||
max_pos = max_pos_result.scalar() or 0
|
||||
|
||||
segment = KnowledgeSegment(
|
||||
knowledge_base_id=knowledge_base_id,
|
||||
document_id=document_id,
|
||||
position=max_pos + 1,
|
||||
content=content,
|
||||
char_count=len(content),
|
||||
keywords=keywords,
|
||||
embedding_status='pending',
|
||||
enabled=True,
|
||||
)
|
||||
self._db.add(segment)
|
||||
await self._db.commit()
|
||||
await self._db.refresh(segment)
|
||||
|
||||
# 向量化
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service.index_segment(knowledge_base_id, segment.id)
|
||||
|
||||
# 更新统计
|
||||
await indexing_service._update_kb_stats(knowledge_base_id)
|
||||
await self._db.commit()
|
||||
|
||||
await self._db.refresh(segment)
|
||||
return segment
|
||||
|
||||
async def delete_segment(self, segment_id: str) -> bool:
|
||||
"""删除分段(软删除 + 清理 Qdrant 向量)"""
|
||||
from sqlalchemy import update
|
||||
|
||||
# 先查询获取 kb_id
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeSegment.knowledge_base_id).where(
|
||||
KnowledgeSegment.id == segment_id,
|
||||
KnowledgeSegment.is_deleted == False
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
if not row:
|
||||
return False
|
||||
|
||||
kb_id = str(row[0])
|
||||
|
||||
# 直接 SQL UPDATE 避免并发场景下的 StaleDataError
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment)
|
||||
.where(KnowledgeSegment.id == segment_id)
|
||||
.values(is_deleted=True)
|
||||
)
|
||||
await self._db.commit()
|
||||
|
||||
# 从 Qdrant 删除该分段的向量
|
||||
try:
|
||||
from ai_platform.knowledge.vector_store import get_vector_store
|
||||
vector_store = get_vector_store()
|
||||
await vector_store.delete(kb_id, [str(segment_id)])
|
||||
except Exception as e:
|
||||
logger.warning(f'从 Qdrant 删除分段向量失败: {e}')
|
||||
|
||||
# 更新统计
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service._update_kb_stats(kb_id)
|
||||
await self._db.commit()
|
||||
|
||||
return True
|
||||
@@ -0,0 +1,206 @@
|
||||
"""
|
||||
Embedding 服务
|
||||
|
||||
通过 OpenAI 兼容的 Embeddings API 将文本转换为向量
|
||||
支持所有兼容 OpenAI 接口的提供商(OpenAI、Qwen、Ollama 等)
|
||||
"""
|
||||
import logging
|
||||
import math
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 单次批量请求的默认最大文本数
|
||||
DEFAULT_BATCH_SIZE = 50
|
||||
|
||||
# 不同提供商的批次大小限制
|
||||
PROVIDER_BATCH_SIZE = {
|
||||
'qwen': 10, # 阿里云 DashScope 限制 10
|
||||
'dashscope': 10,
|
||||
'siliconflow': 10, # 硅基流动限制较小
|
||||
'ollama': 1, # Ollama 通常逐条处理
|
||||
}
|
||||
|
||||
|
||||
|
||||
class EmbeddingService:
|
||||
"""
|
||||
Embedding 服务
|
||||
|
||||
通过模型 ID 获取对应的提供商,调用 OpenAI 兼容的 Embeddings API
|
||||
"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
self._client_cache = {}
|
||||
|
||||
async def _get_client_and_model(self, model_id: str):
|
||||
"""
|
||||
根据模型 ID 获取异步客户端和模型名称
|
||||
|
||||
Returns:
|
||||
(async_client, model_name, max_tokens, provider_type)
|
||||
"""
|
||||
from ai_platform.models import LLMModel, LLMProvider
|
||||
|
||||
result = await self._db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == model_id,
|
||||
LLMModel.is_active == True,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
model = result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise ValueError(f'Embedding 模型不存在或已禁用: {model_id}')
|
||||
if model.model_type != 'embedding':
|
||||
raise ValueError(f'模型 {model.display_name} 不是 Embedding 类型')
|
||||
|
||||
provider_result = await self._db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == model.provider_id,
|
||||
LLMProvider.is_active == True,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise ValueError('Embedding 模型对应的提供商不存在或已禁用')
|
||||
|
||||
cache_key = str(provider.id)
|
||||
if cache_key not in self._client_cache:
|
||||
import httpx
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
# 根据提供商类型确定 base_url
|
||||
if provider.provider_type == 'ollama':
|
||||
base_url = (provider.ollama_host or 'http://localhost:11434').rstrip('/') + '/v1'
|
||||
else:
|
||||
from ai_platform.providers.registry import ProviderRegistry
|
||||
provider_class = ProviderRegistry.get(provider.provider_type)
|
||||
default_base = getattr(provider_class, 'DEFAULT_API_BASE', 'https://api.openai.com/v1') if provider_class else 'https://api.openai.com/v1'
|
||||
base_url = provider.api_base or default_base
|
||||
|
||||
self._client_cache[cache_key] = AsyncOpenAI(
|
||||
api_key=provider.api_key or 'ollama',
|
||||
base_url=base_url,
|
||||
timeout=httpx.Timeout(120.0, connect=30.0),
|
||||
max_retries=5,
|
||||
)
|
||||
|
||||
return self._client_cache[cache_key], model.model_name, model.context_window or 8191, provider.provider_type
|
||||
|
||||
async def embed_text(
|
||||
self,
|
||||
model_id: str,
|
||||
text: str,
|
||||
dimensions: Optional[int] = None,
|
||||
) -> List[float]:
|
||||
"""
|
||||
将单个文本转换为向量
|
||||
|
||||
Args:
|
||||
model_id: Embedding 模型 ID
|
||||
text: 文本内容
|
||||
dimensions: 向量维度(可选,部分模型支持)
|
||||
|
||||
Returns:
|
||||
向量列表 List[float]
|
||||
"""
|
||||
results = await self.embed_texts(model_id, [text], dimensions)
|
||||
return results[0]
|
||||
|
||||
async def embed_texts(
|
||||
self,
|
||||
model_id: str,
|
||||
texts: List[str],
|
||||
dimensions: Optional[int] = None,
|
||||
) -> List[List[float]]:
|
||||
"""
|
||||
批量将文本转换为向量
|
||||
|
||||
Args:
|
||||
model_id: Embedding 模型 ID
|
||||
texts: 文本列表
|
||||
dimensions: 向量维度(可选)
|
||||
|
||||
Returns:
|
||||
向量列表 List[List[float]]
|
||||
"""
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
client, model_name, max_tokens, provider_type = await self._get_client_and_model(model_id)
|
||||
|
||||
# 根据提供商类型确定批次大小
|
||||
batch_size = PROVIDER_BATCH_SIZE.get(provider_type, DEFAULT_BATCH_SIZE)
|
||||
|
||||
# 预处理:截断超长文本
|
||||
processed_texts = []
|
||||
for text in texts:
|
||||
if not text or not text.strip():
|
||||
processed_texts.append(" ")
|
||||
else:
|
||||
# 粗略估算 token 数(中文约 1 字 = 1.5 token,英文约 4 字符 = 1 token)
|
||||
# 保守截断到 max_tokens * 2 个字符
|
||||
max_chars = max_tokens * 2
|
||||
if len(text) > max_chars:
|
||||
processed_texts.append(text[:max_chars])
|
||||
else:
|
||||
processed_texts.append(text)
|
||||
|
||||
# 分批处理
|
||||
all_embeddings = [None] * len(processed_texts)
|
||||
total_batches = math.ceil(len(processed_texts) / batch_size)
|
||||
|
||||
for batch_idx in range(total_batches):
|
||||
start = batch_idx * batch_size
|
||||
end = min(start + batch_size, len(processed_texts))
|
||||
batch_texts = processed_texts[start:end]
|
||||
|
||||
try:
|
||||
kwargs = {
|
||||
'model': model_name,
|
||||
'input': batch_texts,
|
||||
}
|
||||
# 仅对明确支持 dimensions 参数的模型传递该参数
|
||||
if dimensions:
|
||||
model_lower = model_name.lower()
|
||||
# OpenAI text-embedding-3 系列原生支持任意 dimensions
|
||||
if 'text-embedding-3' in model_lower:
|
||||
kwargs['dimensions'] = dimensions
|
||||
# DashScope text-embedding-v3 只接受 [64,128,256,512,768,1024]
|
||||
elif 'text-embedding-v3' in model_lower and dimensions in (64, 128, 256, 512, 768, 1024):
|
||||
kwargs['dimensions'] = dimensions
|
||||
|
||||
response = await client.embeddings.create(**kwargs)
|
||||
|
||||
for item in response.data:
|
||||
all_embeddings[start + item.index] = item.embedding
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f'Embedding 批次 {batch_idx + 1}/{total_batches} 失败: {e}')
|
||||
raise ValueError(f'Embedding 调用失败: {str(e)}')
|
||||
|
||||
# 检查是否所有向量都已生成
|
||||
for i, emb in enumerate(all_embeddings):
|
||||
if emb is None:
|
||||
raise ValueError(f'第 {i} 个文本的向量未生成')
|
||||
|
||||
return all_embeddings
|
||||
|
||||
async def get_embedding_dimensions(self, model_id: str) -> int:
|
||||
"""
|
||||
获取模型的向量维度(通过嵌入一个测试文本来检测)
|
||||
|
||||
Args:
|
||||
model_id: Embedding 模型 ID
|
||||
|
||||
Returns:
|
||||
向量维度
|
||||
"""
|
||||
test_embedding = await self.embed_text(model_id, "test")
|
||||
return len(test_embedding)
|
||||
@@ -0,0 +1,91 @@
|
||||
"""
|
||||
索引进度推送服务
|
||||
|
||||
通过 Redis Pub/Sub 推送索引进度,前端通过 SSE 订阅。
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Redis 频道前缀
|
||||
CHANNEL_PREFIX = "knowledge:indexing:progress:"
|
||||
|
||||
|
||||
class IndexingProgressService:
|
||||
"""索引进度推送服务"""
|
||||
|
||||
@staticmethod
|
||||
def _channel(knowledge_base_id: str) -> str:
|
||||
return f"{CHANNEL_PREFIX}{knowledge_base_id}"
|
||||
|
||||
@classmethod
|
||||
async def publish(
|
||||
cls,
|
||||
knowledge_base_id: str,
|
||||
document_id: str,
|
||||
step: str,
|
||||
progress: float,
|
||||
message: str = "",
|
||||
document_name: str = "",
|
||||
error: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
发布索引进度事件
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
document_id: 文档 ID
|
||||
step: 当前步骤 (extracting/cleaning/chunking/vectorizing/completed/failed)
|
||||
progress: 进度 0.0 ~ 1.0
|
||||
message: 进度描述
|
||||
document_name: 文档名称
|
||||
error: 错误信息(仅 failed 步骤)
|
||||
"""
|
||||
try:
|
||||
from utils.redis import RedisClient
|
||||
client = await RedisClient.get_client()
|
||||
|
||||
event = {
|
||||
"document_id": document_id,
|
||||
"document_name": document_name,
|
||||
"step": step,
|
||||
"progress": round(progress, 2),
|
||||
"message": message,
|
||||
}
|
||||
if error:
|
||||
event["error"] = error
|
||||
|
||||
await client.publish(
|
||||
cls._channel(knowledge_base_id),
|
||||
json.dumps(event, ensure_ascii=False),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"发布索引进度失败: {e}")
|
||||
|
||||
@classmethod
|
||||
async def subscribe(cls, knowledge_base_id: str):
|
||||
"""
|
||||
订阅索引进度事件(异步生成器,用于 SSE)
|
||||
|
||||
Yields:
|
||||
dict: 进度事件
|
||||
"""
|
||||
from utils.redis import RedisClient
|
||||
client = await RedisClient.get_client()
|
||||
pubsub = client.pubsub()
|
||||
channel = cls._channel(knowledge_base_id)
|
||||
|
||||
await pubsub.subscribe(channel)
|
||||
try:
|
||||
async for message in pubsub.listen():
|
||||
if message["type"] == "message":
|
||||
try:
|
||||
data = json.loads(message["data"])
|
||||
yield data
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
continue
|
||||
finally:
|
||||
await pubsub.unsubscribe(channel)
|
||||
await pubsub.close()
|
||||
@@ -0,0 +1,540 @@
|
||||
"""
|
||||
文档索引服务
|
||||
|
||||
负责文档处理管道:文本提取 → 分块 → 向量化 → 入库
|
||||
分段数据存入业务数据库,向量数据存入 Qdrant 向量数据库
|
||||
"""
|
||||
import hashlib
|
||||
import logging
|
||||
import math
|
||||
from datetime import datetime
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from sqlalchemy import select, func, delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
|
||||
from ai_platform.knowledge.chunking import get_chunker
|
||||
from ai_platform.knowledge.services.embedding_service import EmbeddingService
|
||||
from ai_platform.knowledge.vector_store import get_vector_store, VectorPoint
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 向量化批次大小
|
||||
EMBEDDING_BATCH_SIZE = 50
|
||||
|
||||
|
||||
class IndexingService:
|
||||
"""
|
||||
文档索引服务
|
||||
|
||||
处理管道:
|
||||
1. 从文件管理系统提取文本内容
|
||||
2. 按知识库配置的策略分块
|
||||
3. 调用 Embedding 模型向量化
|
||||
4. 将分段写入业务数据库,向量写入 Qdrant
|
||||
"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
self._embedding_service = EmbeddingService(db)
|
||||
self._vector_store = get_vector_store()
|
||||
|
||||
async def index_document(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
document_id: str,
|
||||
) -> Tuple[int, int]:
|
||||
"""
|
||||
索引单个文档
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
document_id: 文档 ID
|
||||
|
||||
Returns:
|
||||
(segment_count, token_count) 分段数和 Token 数
|
||||
"""
|
||||
# 1. 获取知识库配置
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == knowledge_base_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if not kb:
|
||||
raise ValueError(f'知识库不存在: {knowledge_base_id}')
|
||||
|
||||
if not kb.embedding_model_id:
|
||||
raise ValueError('知识库未配置 Embedding 模型')
|
||||
|
||||
# 2. 获取文档
|
||||
doc_result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.id == document_id,
|
||||
KnowledgeDocument.is_deleted == False
|
||||
)
|
||||
)
|
||||
doc = doc_result.scalar_one_or_none()
|
||||
if not doc:
|
||||
raise ValueError(f'文档不存在: {document_id}')
|
||||
|
||||
# 更新状态为 indexing
|
||||
doc.status = 'indexing'
|
||||
doc.indexing_started_at = datetime.now()
|
||||
doc.error_message = None
|
||||
await self._db.commit()
|
||||
|
||||
try:
|
||||
from ai_platform.knowledge.services.indexing_progress_service import IndexingProgressService
|
||||
|
||||
# 3. 提取文本
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='extracting', progress=0.1,
|
||||
message='正在提取文本内容...', document_name=doc.name,
|
||||
)
|
||||
text_content = await self._extract_text(doc.file_id)
|
||||
if not text_content or not text_content.strip():
|
||||
raise ValueError('文档内容为空,无法索引')
|
||||
|
||||
# 3.5 预处理/清洗
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='cleaning', progress=0.2,
|
||||
message='正在预处理/清洗文本...', document_name=doc.name,
|
||||
)
|
||||
from ai_platform.knowledge.services.cleaning_service import CleaningService
|
||||
text_content = CleaningService.clean(text_content, kb.process_rules)
|
||||
|
||||
# 计算内容哈希(用于去重)
|
||||
content_hash = hashlib.md5(text_content.encode('utf-8')).hexdigest()
|
||||
doc.content_hash = content_hash
|
||||
|
||||
# 4. 分块
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='chunking', progress=0.3,
|
||||
message='正在分块...', document_name=doc.name,
|
||||
)
|
||||
chunk_strategy = kb.chunk_strategy or 'recursive'
|
||||
chunk_kwargs = {}
|
||||
|
||||
# Q&A 模式需要 LLM 调用函数
|
||||
if chunk_strategy == 'qa':
|
||||
chunk_kwargs['llm_caller'] = self._create_llm_caller(kb)
|
||||
|
||||
chunker = get_chunker(
|
||||
strategy=chunk_strategy,
|
||||
chunk_size=kb.chunk_size or 500,
|
||||
chunk_overlap=kb.chunk_overlap or 50,
|
||||
separator=kb.separator,
|
||||
**chunk_kwargs,
|
||||
)
|
||||
|
||||
doc_metadata = {
|
||||
'document_id': document_id,
|
||||
'document_name': doc.name,
|
||||
'file_type': doc.file_type,
|
||||
}
|
||||
|
||||
# Q&A 模式使用异步分块
|
||||
if chunk_strategy == 'qa' and hasattr(chunker, 'chunk_async'):
|
||||
chunks = await chunker.chunk_async(text_content, metadata=doc_metadata)
|
||||
else:
|
||||
chunks = chunker.chunk(text_content, metadata=doc_metadata)
|
||||
|
||||
if not chunks:
|
||||
raise ValueError('文档分块结果为空')
|
||||
|
||||
# 5. 删除旧的分段(重新索引场景)
|
||||
await self._delete_document_segments(document_id, knowledge_base_id)
|
||||
|
||||
# 判断索引模式
|
||||
is_economy = (kb.indexing_technique == 'economy')
|
||||
|
||||
# 6. 确保 Qdrant collection 存在(经济模式跳过)
|
||||
if not is_economy:
|
||||
# 自动检测并修正 embedding 维度
|
||||
try:
|
||||
real_dim = await self._embedding_service.get_embedding_dimensions(kb.embedding_model_id)
|
||||
if real_dim != kb.embedding_dimensions:
|
||||
logger.info(f'修正 embedding 维度: {kb.embedding_dimensions} → {real_dim}')
|
||||
kb.embedding_dimensions = real_dim
|
||||
await self._db.commit()
|
||||
except Exception as e:
|
||||
logger.warning(f'自动检测 embedding 维度失败: {e}')
|
||||
|
||||
vector_size = kb.embedding_dimensions or 1536
|
||||
await self._vector_store.ensure_collection(knowledge_base_id, vector_size)
|
||||
|
||||
# 7. 入库(分批处理)
|
||||
segment_count = 0
|
||||
total_token_count = 0
|
||||
total_char_count = 0
|
||||
failed_embedding_count = 0
|
||||
|
||||
total_batches = math.ceil(len(chunks) / EMBEDDING_BATCH_SIZE)
|
||||
|
||||
for batch_idx in range(total_batches):
|
||||
start = batch_idx * EMBEDDING_BATCH_SIZE
|
||||
end = min(start + EMBEDDING_BATCH_SIZE, len(chunks))
|
||||
batch_chunks = chunks[start:end]
|
||||
|
||||
# 高质量模式:批量向量化;经济模式:跳过
|
||||
if is_economy:
|
||||
embeddings = [None] * len(batch_chunks)
|
||||
else:
|
||||
batch_texts = [c.content for c in batch_chunks]
|
||||
try:
|
||||
embeddings = await self._embedding_service.embed_texts(
|
||||
model_id=kb.embedding_model_id,
|
||||
texts=batch_texts,
|
||||
dimensions=kb.embedding_dimensions if kb.embedding_dimensions else None,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f'向量化批次 {batch_idx + 1}/{total_batches} 失败: {e}')
|
||||
embeddings = [None] * len(batch_texts)
|
||||
failed_embedding_count += len(batch_texts)
|
||||
|
||||
# 创建分段记录(业务数据库)+ 收集向量点(Qdrant)
|
||||
vector_points = []
|
||||
for i, chunk in enumerate(batch_chunks):
|
||||
embedding = embeddings[i] if i < len(embeddings) else None
|
||||
char_count = len(chunk.content)
|
||||
token_count = self._estimate_tokens(chunk.content)
|
||||
word_count = self._count_words(chunk.content)
|
||||
|
||||
# 自动提取关键词(高质量和经济模式均提取,增强全文检索)
|
||||
keywords = chunk.metadata.get('keywords')
|
||||
if not keywords:
|
||||
keywords = self._extract_keywords(chunk.content)
|
||||
|
||||
# 经济模式下 embedding_status 标记为 'skipped'
|
||||
if is_economy:
|
||||
emb_status = 'skipped'
|
||||
else:
|
||||
emb_status = 'completed' if embedding else 'failed'
|
||||
|
||||
segment = KnowledgeSegment(
|
||||
knowledge_base_id=knowledge_base_id,
|
||||
document_id=document_id,
|
||||
position=start + i,
|
||||
content=chunk.content,
|
||||
answer=chunk.metadata.get('answer'),
|
||||
token_count=token_count,
|
||||
char_count=char_count,
|
||||
word_count=word_count,
|
||||
page_number=chunk.metadata.get('page_number'),
|
||||
keywords=keywords,
|
||||
extra_metadata=chunk.metadata,
|
||||
embedding_status=emb_status,
|
||||
enabled=True,
|
||||
)
|
||||
self._db.add(segment)
|
||||
await self._db.flush()
|
||||
|
||||
# 收集向量点,稍后批量写入 Qdrant(经济模式跳过)
|
||||
if embedding and not is_economy:
|
||||
vector_points.append(VectorPoint(
|
||||
id=str(segment.id),
|
||||
vector=embedding,
|
||||
payload={
|
||||
'document_id': document_id,
|
||||
'knowledge_base_id': knowledge_base_id,
|
||||
'position': start + i,
|
||||
},
|
||||
))
|
||||
|
||||
segment_count += 1
|
||||
total_token_count += token_count
|
||||
total_char_count += char_count
|
||||
|
||||
# 提交业务数据库
|
||||
await self._db.commit()
|
||||
|
||||
# 批量写入 Qdrant(经济模式跳过)
|
||||
if vector_points:
|
||||
await self._vector_store.upsert(knowledge_base_id, vector_points)
|
||||
|
||||
step_label = '关键词提取中' if is_economy else '向量化中'
|
||||
batch_progress = 0.3 + 0.6 * (end / len(chunks))
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='vectorizing',
|
||||
progress=batch_progress,
|
||||
message=f'{step_label} {end}/{len(chunks)}',
|
||||
document_name=doc.name,
|
||||
)
|
||||
logger.info(f'文档 {doc.name} 索引进度: {end}/{len(chunks)}')
|
||||
|
||||
# 7. 更新文档状态
|
||||
if not is_economy and failed_embedding_count > 0:
|
||||
if failed_embedding_count >= segment_count:
|
||||
doc.status = 'failed'
|
||||
doc.error_message = f'所有 {segment_count} 个分段向量化失败'
|
||||
else:
|
||||
doc.status = 'completed'
|
||||
doc.error_message = f'{failed_embedding_count}/{segment_count} 个分段向量化失败'
|
||||
else:
|
||||
doc.status = 'completed'
|
||||
doc.segment_count = segment_count
|
||||
doc.token_count = total_token_count
|
||||
doc.char_count = total_char_count
|
||||
doc.indexing_completed_at = datetime.now()
|
||||
|
||||
# 8. 更新知识库统计
|
||||
await self._update_kb_stats(knowledge_base_id)
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='completed', progress=1.0,
|
||||
message=f'索引完成: {segment_count} 个分段',
|
||||
document_name=doc.name,
|
||||
)
|
||||
logger.info(f'文档 {doc.name} 索引完成: {segment_count} 个分段, {total_token_count} tokens')
|
||||
return segment_count, total_token_count
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f'文档索引失败: {e}')
|
||||
doc.status = 'failed'
|
||||
doc.error_message = str(e)[:500]
|
||||
await self._db.commit()
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='failed', progress=0.0,
|
||||
message='索引失败', document_name=doc.name,
|
||||
error=str(e)[:200],
|
||||
)
|
||||
raise
|
||||
|
||||
async def reindex_document(self, knowledge_base_id: str, document_id: str) -> Tuple[int, int]:
|
||||
"""重新索引文档(删除旧分段后重新处理)"""
|
||||
return await self.index_document(knowledge_base_id, document_id)
|
||||
|
||||
async def index_segment(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
segment_id: str,
|
||||
) -> bool:
|
||||
"""
|
||||
为单个分段生成向量(用于手动添加或更新分段后)
|
||||
"""
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == knowledge_base_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if not kb or not kb.embedding_model_id:
|
||||
return False
|
||||
|
||||
seg_result = await self._db.execute(
|
||||
select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.id == segment_id,
|
||||
KnowledgeSegment.is_deleted == False
|
||||
)
|
||||
)
|
||||
segment = seg_result.scalar_one_or_none()
|
||||
if not segment:
|
||||
return False
|
||||
|
||||
try:
|
||||
embedding = await self._embedding_service.embed_text(
|
||||
model_id=kb.embedding_model_id,
|
||||
text=segment.content,
|
||||
dimensions=kb.embedding_dimensions if kb.embedding_dimensions else None,
|
||||
)
|
||||
|
||||
# 确保 collection 存在
|
||||
vector_size = kb.embedding_dimensions or 1536
|
||||
await self._vector_store.ensure_collection(knowledge_base_id, vector_size)
|
||||
|
||||
# 写入 Qdrant
|
||||
await self._vector_store.upsert(knowledge_base_id, [VectorPoint(
|
||||
id=str(segment_id),
|
||||
vector=embedding,
|
||||
payload={
|
||||
'document_id': str(segment.document_id),
|
||||
'knowledge_base_id': knowledge_base_id,
|
||||
'position': segment.position or 0,
|
||||
},
|
||||
)])
|
||||
|
||||
segment.embedding_status = 'completed'
|
||||
await self._db.commit()
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f'分段向量化失败: {e}')
|
||||
segment.embedding_status = 'failed'
|
||||
await self._db.commit()
|
||||
return False
|
||||
|
||||
async def _extract_text(self, file_id: str) -> str:
|
||||
"""从文件管理系统提取文本内容(启用 OCR 支持图片和扫描版 PDF)"""
|
||||
from core.file_manager.service import FileManagerService
|
||||
|
||||
text_content = await FileManagerService.get_file_text_content(
|
||||
self._db, file_id, enable_ocr=True
|
||||
)
|
||||
if not text_content:
|
||||
raise ValueError('无法提取文件文本内容')
|
||||
return text_content
|
||||
|
||||
async def _delete_document_segments(self, document_id: str, knowledge_base_id: str):
|
||||
"""删除文档的所有分段(业务数据库 + Qdrant)"""
|
||||
# 先从 Qdrant 删除该文档的所有向量
|
||||
try:
|
||||
await self._vector_store.delete_by_filter(
|
||||
knowledge_base_id,
|
||||
filter_conditions={'document_id': document_id},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f'从 Qdrant 删除文档向量失败: {e}')
|
||||
|
||||
# 再从业务数据库删除分段记录
|
||||
await self._db.execute(
|
||||
delete(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == document_id
|
||||
)
|
||||
)
|
||||
|
||||
async def _update_kb_stats(self, knowledge_base_id: str):
|
||||
"""更新知识库统计信息"""
|
||||
# 文档数
|
||||
doc_count_result = await self._db.execute(
|
||||
select(func.count()).select_from(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
)
|
||||
)
|
||||
doc_count = doc_count_result.scalar() or 0
|
||||
|
||||
# 分段数和 Token 数
|
||||
seg_stats = await self._db.execute(
|
||||
select(
|
||||
func.count(),
|
||||
func.coalesce(func.sum(KnowledgeSegment.token_count), 0),
|
||||
func.coalesce(func.sum(KnowledgeSegment.char_count), 0),
|
||||
).where(
|
||||
KnowledgeSegment.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
)
|
||||
row = seg_stats.one()
|
||||
seg_count = row[0] or 0
|
||||
total_tokens = row[1] or 0
|
||||
total_chars = row[2] or 0
|
||||
|
||||
# 更新知识库
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(KnowledgeBase.id == knowledge_base_id)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if kb:
|
||||
kb.document_count = doc_count
|
||||
kb.segment_count = seg_count
|
||||
kb.total_token_count = total_tokens
|
||||
kb.total_char_count = total_chars
|
||||
|
||||
def _create_llm_caller(self, kb: KnowledgeBase):
|
||||
"""
|
||||
创建 Q&A 分块所需的 LLM 调用函数
|
||||
|
||||
使用知识库所属应用中配置的第一个 chat 类型模型。
|
||||
"""
|
||||
db = self._db
|
||||
|
||||
async def llm_caller(system_prompt: str, user_prompt: str) -> str:
|
||||
from ai_platform.services.llm_service import LLMService
|
||||
|
||||
# 查找可用的 chat 模型
|
||||
from ai_platform.models import LLMModel
|
||||
model_result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.model_type == 'chat',
|
||||
LLMModel.is_active == True,
|
||||
LLMModel.is_deleted == False,
|
||||
).limit(1)
|
||||
)
|
||||
model = model_result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise ValueError('未找到可用的 chat 模型,无法进行 Q&A 拆分')
|
||||
|
||||
llm_service = LLMService(db)
|
||||
messages = []
|
||||
if system_prompt:
|
||||
messages.append({'role': 'system', 'content': system_prompt})
|
||||
messages.append({'role': 'user', 'content': user_prompt})
|
||||
|
||||
response = await llm_service.chat_async(
|
||||
model_id=str(model.id),
|
||||
messages=messages,
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
)
|
||||
return response.content or ''
|
||||
|
||||
return llm_caller
|
||||
|
||||
@staticmethod
|
||||
def _extract_keywords(text: str, max_keywords: int = 10) -> List[str]:
|
||||
"""
|
||||
从文本中提取关键词(经济模式使用)
|
||||
|
||||
使用简单的 TF 统计提取高频词,无需外部依赖。
|
||||
"""
|
||||
import re
|
||||
if not text:
|
||||
return []
|
||||
|
||||
# 中文分词(简单按标点和空格分割)
|
||||
# 提取中文词组(2-4字)和英文单词
|
||||
chinese_words = re.findall(r'[\u4e00-\u9fff]{2,4}', text)
|
||||
english_words = [w.lower() for w in re.findall(r'[a-zA-Z]{3,}', text)]
|
||||
|
||||
all_words = chinese_words + english_words
|
||||
|
||||
# 停用词(简单列表)
|
||||
stop_words = {
|
||||
'的', '了', '在', '是', '我', '有', '和', '就', '不', '人', '都', '一',
|
||||
'一个', '上', '也', '很', '到', '说', '要', '去', '你', '会', '着',
|
||||
'没有', '看', '好', '自己', '这', '他', '她', '它', '我们', '他们',
|
||||
'可以', '这个', '那个', '什么', '如果', '因为', '所以', '但是', '而且',
|
||||
'the', 'and', 'for', 'are', 'but', 'not', 'you', 'all', 'can',
|
||||
'had', 'her', 'was', 'one', 'our', 'out', 'has', 'have', 'been',
|
||||
'this', 'that', 'with', 'from', 'they', 'will', 'would', 'there',
|
||||
}
|
||||
|
||||
# 词频统计
|
||||
word_freq = {}
|
||||
for word in all_words:
|
||||
if word in stop_words or len(word) < 2:
|
||||
continue
|
||||
word_freq[word] = word_freq.get(word, 0) + 1
|
||||
|
||||
# 按频率排序取 top
|
||||
sorted_words = sorted(word_freq.items(), key=lambda x: x[1], reverse=True)
|
||||
return [w for w, _ in sorted_words[:max_keywords]]
|
||||
|
||||
@staticmethod
|
||||
def _count_words(text: str) -> int:
|
||||
"""计算词数(中文按字计算,英文按空格分词)"""
|
||||
if not text:
|
||||
return 0
|
||||
import re
|
||||
# 中文字符数
|
||||
chinese_chars = len(re.findall(r'[\u4e00-\u9fff]', text))
|
||||
# 英文单词数
|
||||
english_words = len(re.findall(r'[a-zA-Z]+', text))
|
||||
return chinese_chars + english_words
|
||||
|
||||
@staticmethod
|
||||
def _estimate_tokens(text: str) -> int:
|
||||
"""粗略估算文本的 Token 数"""
|
||||
if not text:
|
||||
return 0
|
||||
# 中文约 1 字 = 1.5 token,英文约 4 字符 = 1 token
|
||||
# 简单混合估算
|
||||
chinese_chars = sum(1 for c in text if '\u4e00' <= c <= '\u9fff')
|
||||
other_chars = len(text) - chinese_chars
|
||||
return int(chinese_chars * 1.5 + other_chars / 4)
|
||||
@@ -0,0 +1,267 @@
|
||||
"""
|
||||
知识库服务
|
||||
|
||||
知识库 CRUD 操作
|
||||
|
||||
数据权限:
|
||||
- 使用 get_list_with_data_scope() 自动应用数据权限
|
||||
- 支持本人、本部门、本部门及下级、全部等数据范围
|
||||
"""
|
||||
import logging
|
||||
from typing import Optional, List, Tuple
|
||||
|
||||
from sqlalchemy import select, func, or_, and_
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
|
||||
from ai_platform.knowledge.schemas.knowledge_base_schema import KnowledgeBaseCreate, KnowledgeBaseUpdate
|
||||
from app.data_scope_utils import get_data_scope_filter, apply_data_scope_to_conditions
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 资源类型(用于数据权限配置)
|
||||
RESOURCE_TYPE = "knowledge_base"
|
||||
RESOURCE_DISPLAY_NAME = "知识库管理"
|
||||
|
||||
|
||||
class KnowledgeService:
|
||||
"""知识库服务"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
|
||||
async def get_list(
|
||||
self,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
name: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
application_id: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeBase], int]:
|
||||
"""获取知识库列表"""
|
||||
query = select(KnowledgeBase).where(KnowledgeBase.is_deleted == False)
|
||||
|
||||
if application_id:
|
||||
query = query.where(or_(
|
||||
KnowledgeBase.application_id == application_id,
|
||||
and_(KnowledgeBase.application_id.is_(None), KnowledgeBase.is_global == True)
|
||||
))
|
||||
if name:
|
||||
query = query.where(KnowledgeBase.name.ilike(f"%{name}%"))
|
||||
if status:
|
||||
query = query.where(KnowledgeBase.status == status)
|
||||
|
||||
# 总数
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeBase.sort.desc(), KnowledgeBase.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_list_with_data_scope(
|
||||
self,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
name: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
application_id: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeBase], int]:
|
||||
"""
|
||||
获取知识库列表(带数据权限过滤)
|
||||
|
||||
自动从上下文获取当前用户信息,应用数据权限过滤
|
||||
"""
|
||||
conditions = [KnowledgeBase.is_deleted == False]
|
||||
|
||||
if application_id:
|
||||
conditions.append(or_(
|
||||
KnowledgeBase.application_id == application_id,
|
||||
and_(KnowledgeBase.application_id.is_(None), KnowledgeBase.is_global == True)
|
||||
))
|
||||
if name:
|
||||
conditions.append(KnowledgeBase.name.ilike(f"%{name}%"))
|
||||
if status:
|
||||
conditions.append(KnowledgeBase.status == status)
|
||||
|
||||
# 获取数据权限过滤条件并应用
|
||||
data_scope_filter = await get_data_scope_filter(self._db, RESOURCE_TYPE)
|
||||
scope_conditions = apply_data_scope_to_conditions(KnowledgeBase, data_scope_filter)
|
||||
conditions.extend(scope_conditions)
|
||||
|
||||
# 总数
|
||||
query = select(KnowledgeBase).where(and_(*conditions))
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeBase.sort.desc(), KnowledgeBase.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_by_id(self, kb_id: str) -> Optional[KnowledgeBase]:
|
||||
"""获取知识库详情"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == kb_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_by_code(self, code: str) -> Optional[KnowledgeBase]:
|
||||
"""根据编码获取知识库"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.code == code,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def create(self, data: KnowledgeBaseCreate) -> KnowledgeBase:
|
||||
"""创建知识库"""
|
||||
# 检查编码唯一性
|
||||
existing = await self.get_by_code(data.code)
|
||||
if existing:
|
||||
raise ValueError(f'知识库编码 {data.code} 已存在')
|
||||
|
||||
kb_data = data.model_dump()
|
||||
|
||||
# 自动检测 embedding 模型的真实维度
|
||||
if data.embedding_model_id:
|
||||
try:
|
||||
from ai_platform.knowledge.services.embedding_service import EmbeddingService
|
||||
embedding_service = EmbeddingService(self._db)
|
||||
real_dim = await embedding_service.get_embedding_dimensions(data.embedding_model_id)
|
||||
kb_data['embedding_dimensions'] = real_dim
|
||||
logger.info(f'自动检测 embedding 维度: {real_dim}')
|
||||
except Exception as e:
|
||||
logger.warning(f'自动检测 embedding 维度失败,使用默认值: {e}')
|
||||
|
||||
kb = KnowledgeBase(**kb_data)
|
||||
|
||||
# 自动填充创建人和部门
|
||||
from utils.context import get_current_user_info_from_context
|
||||
user_info = get_current_user_info_from_context()
|
||||
if user_info:
|
||||
if not kb.sys_creator_id:
|
||||
kb.sys_creator_id = user_info.get('user_id')
|
||||
if not kb.sys_dept_id and user_info.get('dept_id'):
|
||||
kb.sys_dept_id = user_info.get('dept_id')
|
||||
|
||||
self._db.add(kb)
|
||||
await self._db.commit()
|
||||
await self._db.refresh(kb)
|
||||
return kb
|
||||
|
||||
async def update(self, kb_id: str, data: KnowledgeBaseUpdate) -> Optional[KnowledgeBase]:
|
||||
"""更新知识库"""
|
||||
kb = await self.get_by_id(kb_id)
|
||||
if not kb:
|
||||
return None
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
|
||||
# 如果更换了 embedding 模型,自动重新检测维度
|
||||
new_model_id = update_data.get('embedding_model_id')
|
||||
if new_model_id and new_model_id != kb.embedding_model_id:
|
||||
try:
|
||||
from ai_platform.knowledge.services.embedding_service import EmbeddingService
|
||||
embedding_service = EmbeddingService(self._db)
|
||||
real_dim = await embedding_service.get_embedding_dimensions(new_model_id)
|
||||
update_data['embedding_dimensions'] = real_dim
|
||||
logger.info(f'更换模型后自动检测 embedding 维度: {real_dim}')
|
||||
except Exception as e:
|
||||
logger.warning(f'自动检测 embedding 维度失败: {e}')
|
||||
|
||||
for key, value in update_data.items():
|
||||
setattr(kb, key, value)
|
||||
|
||||
await self._db.commit()
|
||||
await self._db.refresh(kb)
|
||||
return kb
|
||||
|
||||
async def delete(self, kb_id: str) -> bool:
|
||||
"""删除知识库(软删除 + 清理 Qdrant collection)"""
|
||||
kb = await self.get_by_id(kb_id)
|
||||
if not kb:
|
||||
return False
|
||||
|
||||
kb.is_deleted = True
|
||||
|
||||
# 同时软删除所有文档和分段
|
||||
doc_result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == kb_id,
|
||||
KnowledgeDocument.is_deleted == False
|
||||
)
|
||||
)
|
||||
docs = doc_result.scalars().all()
|
||||
for doc in docs:
|
||||
doc.is_deleted = True
|
||||
|
||||
# 软删除分段
|
||||
from sqlalchemy import update
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment).where(
|
||||
KnowledgeSegment.knowledge_base_id == kb_id
|
||||
).values(is_deleted=True)
|
||||
)
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
# 删除 Qdrant 中对应的 collection
|
||||
try:
|
||||
from ai_platform.knowledge.vector_store import get_vector_store
|
||||
vector_store = get_vector_store()
|
||||
await vector_store.delete_collection(kb_id)
|
||||
except Exception as e:
|
||||
logger.warning(f'删除 Qdrant collection 失败: {e}')
|
||||
|
||||
return True
|
||||
|
||||
async def get_simple_list(self, application_id: Optional[str] = None) -> List[dict]:
|
||||
"""获取知识库简单列表(用于下拉选择)"""
|
||||
query = select(
|
||||
KnowledgeBase.id,
|
||||
KnowledgeBase.name,
|
||||
KnowledgeBase.code,
|
||||
KnowledgeBase.document_count,
|
||||
KnowledgeBase.segment_count,
|
||||
).where(
|
||||
KnowledgeBase.is_deleted == False,
|
||||
KnowledgeBase.status == 'active',
|
||||
)
|
||||
|
||||
if application_id:
|
||||
query = query.where(or_(
|
||||
KnowledgeBase.application_id == application_id,
|
||||
and_(KnowledgeBase.application_id.is_(None), KnowledgeBase.is_global == True)
|
||||
))
|
||||
|
||||
query = query.order_by(KnowledgeBase.sort.desc(), KnowledgeBase.sys_create_datetime.desc())
|
||||
result = await self._db.execute(query)
|
||||
rows = result.all()
|
||||
|
||||
return [
|
||||
{
|
||||
'id': row.id,
|
||||
'name': row.name,
|
||||
'code': row.code,
|
||||
'document_count': row.document_count or 0,
|
||||
'segment_count': row.segment_count or 0,
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
@@ -0,0 +1,182 @@
|
||||
"""
|
||||
Rerank 重排序服务
|
||||
|
||||
通过 Rerank 模型对检索结果进行重新排序,提升检索质量。
|
||||
支持 Jina/Cohere 风格的 Rerank API(大多数提供商兼容此接口)。
|
||||
"""
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RerankResult:
|
||||
"""重排序结果"""
|
||||
index: int
|
||||
relevance_score: float
|
||||
|
||||
|
||||
class RerankService:
|
||||
"""
|
||||
Rerank 重排序服务
|
||||
|
||||
通过模型 ID 获取对应的提供商,调用 Rerank API 对文档进行重排序。
|
||||
支持两种 API 风格:
|
||||
- Jina/Cohere 风格:POST /v1/rerank
|
||||
- OpenAI 兼容风格(部分提供商)
|
||||
"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
self._client_cache = {}
|
||||
|
||||
async def _get_client_config(self, model_id: str):
|
||||
"""
|
||||
根据模型 ID 获取 API 配置
|
||||
|
||||
Returns:
|
||||
(base_url, api_key, model_name)
|
||||
"""
|
||||
from ai_platform.models import LLMModel, LLMProvider
|
||||
|
||||
result = await self._db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == model_id,
|
||||
LLMModel.is_active == True,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
model = result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise ValueError(f'Rerank 模型不存在或已禁用: {model_id}')
|
||||
if model.model_type != 'rerank':
|
||||
raise ValueError(f'模型 {model.display_name} 不是 Rerank 类型')
|
||||
|
||||
provider_result = await self._db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == model.provider_id,
|
||||
LLMProvider.is_active == True,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise ValueError('Rerank 模型对应的提供商不存在或已禁用')
|
||||
|
||||
if provider.provider_type == 'ollama':
|
||||
base_url = (provider.ollama_host or 'http://localhost:11434').rstrip('/') + '/v1'
|
||||
else:
|
||||
base_url = provider.api_base or 'https://api.openai.com/v1'
|
||||
|
||||
api_key = provider.api_key or 'ollama'
|
||||
|
||||
return base_url, api_key, model.model_name
|
||||
|
||||
async def rerank(
|
||||
self,
|
||||
model_id: str,
|
||||
query: str,
|
||||
documents: List[str],
|
||||
top_n: Optional[int] = None,
|
||||
) -> List[RerankResult]:
|
||||
"""
|
||||
对文档列表进行重排序
|
||||
|
||||
Args:
|
||||
model_id: Rerank 模型 ID
|
||||
query: 查询文本
|
||||
documents: 待排序的文档列表
|
||||
top_n: 返回前 N 个结果(默认返回全部)
|
||||
|
||||
Returns:
|
||||
按相关性降序排列的 RerankResult 列表
|
||||
"""
|
||||
if not documents:
|
||||
return []
|
||||
|
||||
if top_n is None:
|
||||
top_n = len(documents)
|
||||
|
||||
base_url, api_key, model_name = await self._get_client_config(model_id)
|
||||
|
||||
try:
|
||||
return await self._call_rerank_api(
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
model_name=model_name,
|
||||
query=query,
|
||||
documents=documents,
|
||||
top_n=top_n,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f'Rerank 调用失败: {e}')
|
||||
raise ValueError(f'Rerank 调用失败: {str(e)}')
|
||||
|
||||
async def _call_rerank_api(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
model_name: str,
|
||||
query: str,
|
||||
documents: List[str],
|
||||
top_n: int,
|
||||
) -> List[RerankResult]:
|
||||
"""
|
||||
调用 Rerank API(Jina/Cohere 兼容风格)
|
||||
|
||||
POST {base_url}/rerank
|
||||
{
|
||||
"model": "...",
|
||||
"query": "...",
|
||||
"documents": ["...", "..."],
|
||||
"top_n": 5
|
||||
}
|
||||
|
||||
Response:
|
||||
{
|
||||
"results": [
|
||||
{"index": 0, "relevance_score": 0.95},
|
||||
{"index": 2, "relevance_score": 0.87},
|
||||
...
|
||||
]
|
||||
}
|
||||
"""
|
||||
import httpx
|
||||
|
||||
url = base_url.rstrip('/') + '/rerank'
|
||||
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
}
|
||||
|
||||
payload = {
|
||||
'model': model_name,
|
||||
'query': query,
|
||||
'documents': documents,
|
||||
'top_n': top_n,
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=60) as client:
|
||||
response = await client.post(url, json=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
# 解析结果(兼容 Jina/Cohere/通义千问 等格式)
|
||||
raw_results = data.get('results', [])
|
||||
results = []
|
||||
for item in raw_results:
|
||||
results.append(RerankResult(
|
||||
index=item.get('index', 0),
|
||||
relevance_score=item.get('relevance_score', 0.0),
|
||||
))
|
||||
|
||||
# 按相关性降序排序
|
||||
results.sort(key=lambda r: r.relevance_score, reverse=True)
|
||||
|
||||
return results
|
||||
@@ -0,0 +1,666 @@
|
||||
"""
|
||||
检索服务
|
||||
|
||||
支持向量检索、全文检索、混合检索(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 []
|
||||
Reference in New Issue
Block a user