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