541 lines
21 KiB
Python
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)
|