Files
2026-06-08 18:14:59 +08:00

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