Files
ai-agent-admin/backend-fastapi/ai_platform/knowledge/services/indexing_service.py
T
2026-06-08 18:14:59 +08:00

541 lines
21 KiB
Python

"""
文档索引服务
负责文档处理管道:文本提取 → 分块 → 向量化 → 入库
分段数据存入业务数据库,向量数据存入 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)