Build lightweight AI agent admin

This commit is contained in:
Codex
2026-06-08 18:14:59 +08:00
commit e164840f43
2530 changed files with 435693 additions and 0 deletions
@@ -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('&nbsp;', ' ')
text = text.replace('&lt;', '<')
text = text.replace('&gt;', '>')
text = text.replace('&amp;', '&')
text = text.replace('&quot;', '"')
text = text.replace('&#39;', "'")
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 APIJina/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 返回 answersegment_id 用 annotation id
results.append(RetrievalResult(
segment_id=str(ann.id),
document_id='',
document_name='[Q&A]',
knowledge_base_id=str(ann.knowledge_base_id),
content=ann.answer,
score=round(score_map.get(hit.id, 0.0), 4),
token_count=0,
metadata={'type': 'annotation', 'question': ann.question},
))
# 更新标注命中次数
if annotation_ids:
try:
from sqlalchemy import update
await self._db.execute(
update(KnowledgeAnnotation)
.where(KnowledgeAnnotation.id.in_(annotation_ids))
.values(hit_count=func.coalesce(KnowledgeAnnotation.hit_count, 0) + 1)
)
await self._db.commit()
except Exception:
pass
return results
except Exception as e:
logger.warning(f'标注匹配失败: {e}')
return []