426 lines
14 KiB
Python
426 lines
14 KiB
Python
"""
|
|
文档服务
|
|
|
|
文档上传、管理、状态控制
|
|
"""
|
|
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
|