""" 知识库 API """ import logging from typing import Optional, List from fastapi import APIRouter, Depends, HTTPException, Query, BackgroundTasks from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.database import get_db from app.base_schema import PaginatedResponse, ResponseModel from ai_platform.models import LLMModel from core.application.model import Application from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument from ai_platform.knowledge.schemas.knowledge_base_schema import ( KnowledgeBaseCreate, KnowledgeBaseUpdate, KnowledgeBaseResponse, KnowledgeBaseListResponse, ) from ai_platform.knowledge.schemas.document_schema import ( DocumentUploadInput, DocumentBatchUploadInput, DocumentResponse, DocumentListResponse, ) from ai_platform.knowledge.schemas.segment_schema import ( SegmentResponse, SegmentListResponse, SegmentUpdateInput, SegmentCreateInput, RetrievalInput, RetrievalResponse, ChunkPreviewInput, ChunkPreviewResponse, ) from ai_platform.knowledge.schemas.annotation_schema import ( AnnotationCreateInput, AnnotationUpdateInput, AnnotationResponse, ) from ai_platform.knowledge.models import KnowledgeAnnotation from ai_platform.knowledge.schemas.retrieval_log_schema import RetrievalLogResponse logger = logging.getLogger(__name__) router = APIRouter(prefix="/knowledge", tags=["AI-知识库"]) # ==================== 知识库 CRUD ==================== @router.get("/list", response_model=PaginatedResponse[KnowledgeBaseListResponse], summary="知识库列表") async def list_knowledge_bases( name: Optional[str] = Query(None, description="名称"), status: Optional[str] = Query(None, description="状态"), application_id: Optional[str] = Query(None, alias="applicationId", description="所属应用ID"), page: int = Query(1, ge=1, description="页码"), page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"), db: AsyncSession = Depends(get_db), ): """获取知识库列表(自动应用数据权限)""" from ai_platform.knowledge.services.knowledge_service import KnowledgeService service = KnowledgeService(db) items, total = await service.get_list_with_data_scope( page=page, page_size=page_size, name=name, status=status, application_id=application_id, ) # 批量查询应用名称和模型名称 app_ids = list({kb.application_id for kb in items if kb.application_id}) app_name_map = {} if app_ids: app_result = await db.execute( select(Application.id, Application.name).where(Application.id.in_(app_ids)) ) app_name_map = {row.id: row.name for row in app_result} model_ids = list({kb.embedding_model_id for kb in items if kb.embedding_model_id}) model_name_map = {} if model_ids: model_result = await db.execute( select(LLMModel.id, LLMModel.display_name).where(LLMModel.id.in_(model_ids)) ) model_name_map = {row.id: row.display_name for row in model_result} response_items = [] for kb in items: response_items.append({ "id": kb.id, "application_id": kb.application_id, "application_name": app_name_map.get(kb.application_id, ""), "is_global": kb.is_global or False, "name": kb.name, "code": kb.code, "description": kb.description or "", "icon": kb.icon or "", "embedding_model_name": model_name_map.get(kb.embedding_model_id, ""), "document_count": kb.document_count or 0, "segment_count": kb.segment_count or 0, "status": kb.status or "active", "sys_create_datetime": kb.sys_create_datetime, }) return PaginatedResponse(items=response_items, total=total) @router.get("/simple", summary="知识库简单列表(下拉选择用)") async def list_knowledge_bases_simple( application_id: Optional[str] = Query(None, alias="applicationId", description="所属应用ID"), db: AsyncSession = Depends(get_db), ): """获取知识库简单列表""" from ai_platform.knowledge.services.knowledge_service import KnowledgeService service = KnowledgeService(db) return await service.get_simple_list(application_id) # ==================== 检索日志 ==================== @router.get("/retrieval-logs", response_model=PaginatedResponse[RetrievalLogResponse], summary="检索日志列表") async def list_retrieval_logs( knowledge_base_id: Optional[str] = Query(None, alias="knowledgeBaseId", description="知识库ID筛选"), keyword: Optional[str] = Query(None, description="查询关键词搜索"), page: int = Query(1, ge=1), page_size: int = Query(20, ge=1, le=100, alias="pageSize"), db: AsyncSession = Depends(get_db), ): """获取检索日志列表""" from sqlalchemy import func as sa_func from ai_platform.knowledge.models import KnowledgeRetrievalLog query = select(KnowledgeRetrievalLog).where( KnowledgeRetrievalLog.is_deleted == False, ) if knowledge_base_id: from app.db_compat import json_contains query = query.where( json_contains(KnowledgeRetrievalLog.knowledge_base_ids, knowledge_base_id) ) if keyword: query = query.where(KnowledgeRetrievalLog.query.ilike(f"%{keyword}%")) count_query = select(sa_func.count()).select_from(query.subquery()) total = (await db.execute(count_query)).scalar() or 0 query = query.order_by(KnowledgeRetrievalLog.sys_create_datetime.desc()) query = query.offset((page - 1) * page_size).limit(page_size) result = await db.execute(query) items = result.scalars().all() return PaginatedResponse(items=items, total=total) @router.get("/{kb_id}", response_model=KnowledgeBaseResponse, summary="知识库详情") async def get_knowledge_base(kb_id: str, db: AsyncSession = Depends(get_db)): """获取知识库详情""" from ai_platform.knowledge.services.knowledge_service import KnowledgeService service = KnowledgeService(db) kb = await service.get_by_id(kb_id) if not kb: raise HTTPException(status_code=404, detail="知识库不存在") # 获取模型名称 model_name = "" if kb.embedding_model_id: model_result = await db.execute( select(LLMModel.display_name).where(LLMModel.id == kb.embedding_model_id) ) row = model_result.scalar_one_or_none() if row: model_name = row return _build_kb_response(kb, model_name) @router.post("", response_model=KnowledgeBaseResponse, summary="创建知识库") async def create_knowledge_base(data: KnowledgeBaseCreate, db: AsyncSession = Depends(get_db)): """创建知识库""" from ai_platform.knowledge.services.knowledge_service import KnowledgeService service = KnowledgeService(db) try: kb = await service.create(data) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) return _build_kb_response(kb) @router.put("/{kb_id}", response_model=KnowledgeBaseResponse, summary="更新知识库") async def update_knowledge_base( kb_id: str, data: KnowledgeBaseUpdate, db: AsyncSession = Depends(get_db), ): """更新知识库""" from ai_platform.knowledge.services.knowledge_service import KnowledgeService service = KnowledgeService(db) kb = await service.update(kb_id, data) if not kb: raise HTTPException(status_code=404, detail="知识库不存在") return _build_kb_response(kb) @router.delete("/{kb_id}", response_model=ResponseModel, summary="删除知识库") async def delete_knowledge_base(kb_id: str, db: AsyncSession = Depends(get_db)): """删除知识库""" from ai_platform.knowledge.services.knowledge_service import KnowledgeService service = KnowledgeService(db) success = await service.delete(kb_id) if not success: raise HTTPException(status_code=404, detail="知识库不存在") return ResponseModel(message="删除成功") # ==================== 文档管理 ==================== @router.get("/{kb_id}/documents", response_model=PaginatedResponse[DocumentListResponse], summary="文档列表") async def list_documents( kb_id: str, name: Optional[str] = Query(None, description="名称"), status: Optional[str] = Query(None, description="状态"), page: int = Query(1, ge=1, description="页码"), page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"), db: AsyncSession = Depends(get_db), ): """获取知识库文档列表""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) items, total = await service.get_list( knowledge_base_id=kb_id, page=page, page_size=page_size, name=name, status=status, ) response_items = [_build_doc_list_response(doc) for doc in items] return PaginatedResponse(items=response_items, total=total) @router.post("/{kb_id}/documents", response_model=DocumentResponse, summary="添加文档") async def add_document( kb_id: str, data: DocumentUploadInput, background_tasks: BackgroundTasks, db: AsyncSession = Depends(get_db), ): """添加文档到知识库(添加后自动开始索引)""" from ai_platform.knowledge.services.document_service import DocumentService from ai_platform.knowledge.services.indexing_service import IndexingService doc_service = DocumentService(db) try: doc = await doc_service.add_document(kb_id, data.file_id, data.name) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) # 后台异步索引 background_tasks.add_task(_index_document_task, kb_id, doc.id) return _build_doc_response(doc) @router.post("/{kb_id}/documents/batch", response_model=ResponseModel, summary="批量添加文档") async def batch_add_documents( kb_id: str, data: DocumentBatchUploadInput, background_tasks: BackgroundTasks, db: AsyncSession = Depends(get_db), ): """批量添加文档""" from ai_platform.knowledge.services.document_service import DocumentService doc_service = DocumentService(db) docs = await doc_service.batch_add_documents(kb_id, data.file_ids) # 后台并发索引所有文档 doc_ids = [doc.id for doc in docs] background_tasks.add_task(_batch_index_documents_task, kb_id, doc_ids) return ResponseModel(message=f"成功添加 {len(docs)} 个文档,正在后台索引") @router.get("/{kb_id}/documents/{doc_id}", response_model=DocumentResponse, summary="文档详情") async def get_document(kb_id: str, doc_id: str, db: AsyncSession = Depends(get_db)): """获取文档详情""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) doc = await service.get_by_id(doc_id) if not doc or doc.knowledge_base_id != kb_id: raise HTTPException(status_code=404, detail="文档不存在") return _build_doc_response(doc) @router.delete("/{kb_id}/documents/{doc_id}", response_model=ResponseModel, summary="删除文档") async def delete_document(kb_id: str, doc_id: str, db: AsyncSession = Depends(get_db)): """删除文档""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) success = await service.delete_document(doc_id) if not success: raise HTTPException(status_code=404, detail="文档不存在") return ResponseModel(message="删除成功") @router.post("/{kb_id}/documents/{doc_id}/reindex", response_model=ResponseModel, summary="重新索引文档") async def reindex_document( kb_id: str, doc_id: str, background_tasks: BackgroundTasks, db: AsyncSession = Depends(get_db), ): """重新索引文档""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) doc = await service.get_by_id(doc_id) if not doc or doc.knowledge_base_id != kb_id: raise HTTPException(status_code=404, detail="文档不存在") background_tasks.add_task(_index_document_task, kb_id, doc_id) return ResponseModel(message="已开始重新索引") @router.post("/{kb_id}/documents/reindex-all", response_model=ResponseModel, summary="批量重新索引所有文档") async def reindex_all_documents( kb_id: str, background_tasks: BackgroundTasks, db: AsyncSession = Depends(get_db), ): """修改分块策略后一键重新索引知识库所有文档""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) docs, total = await service.get_list(kb_id, page=1, page_size=9999, status='completed') # 也包含 failed 的文档 failed_docs, _ = await service.get_list(kb_id, page=1, page_size=9999, status='failed') all_docs = docs + failed_docs if not all_docs: return ResponseModel(message="没有可重新索引的文档") doc_ids = [doc.id for doc in all_docs] background_tasks.add_task(_batch_index_documents_task, kb_id, doc_ids) return ResponseModel(message=f"已开始重新索引 {len(doc_ids)} 个文档") @router.put("/{kb_id}/documents/{doc_id}/toggle", response_model=DocumentResponse, summary="启用/禁用文档") async def toggle_document( kb_id: str, doc_id: str, enabled: bool = Query(..., description="是否启用"), db: AsyncSession = Depends(get_db), ): """启用/禁用文档""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) doc = await service.toggle_document(doc_id, enabled) if not doc: raise HTTPException(status_code=404, detail="文档不存在") return _build_doc_response(doc) # ==================== 分块预览 ==================== @router.post("/{kb_id}/chunk-preview", response_model=ChunkPreviewResponse, summary="分块预览") async def chunk_preview( kb_id: str, data: ChunkPreviewInput, db: AsyncSession = Depends(get_db), ): """预览文档分块结果(不实际入库)""" from ai_platform.knowledge.services.knowledge_service import KnowledgeService from ai_platform.knowledge.chunking import get_chunker from ai_platform.knowledge.services.cleaning_service import CleaningService from core.file_manager.service import FileManagerService # 验证知识库存在 service = KnowledgeService(db) kb = await service.get_by_id(kb_id) if not kb: raise HTTPException(status_code=404, detail="知识库不存在") # 提取文本 try: text_content = await FileManagerService.get_file_text_content(db, data.file_id) if not text_content or not text_content.strip(): raise HTTPException(status_code=400, detail="文件内容为空") except HTTPException: raise except Exception as e: raise HTTPException(status_code=400, detail=f"无法提取文件内容: {str(e)}") # 预处理/清洗 process_rules = data.process_rules or kb.process_rules if process_rules: text_content = CleaningService.clean(text_content, process_rules) # 分块 strategy = data.chunk_strategy or kb.chunk_strategy or 'recursive' chunk_size = data.chunk_size or kb.chunk_size or 500 chunk_overlap = data.chunk_overlap if data.chunk_overlap is not None else (kb.chunk_overlap or 50) # Q&A 模式不支持预览(需要 LLM 调用) if strategy == 'qa': raise HTTPException(status_code=400, detail="Q&A 分块策略不支持预览") chunker = get_chunker( strategy=strategy, chunk_size=chunk_size, chunk_overlap=chunk_overlap, separator=data.separator or kb.separator, ) chunks = chunker.chunk(text_content, metadata={}) # 构建预览结果 preview_items = [] for i, chunk in enumerate(chunks): char_count = len(chunk.content) # 简单 token 估算 token_count = max(1, int(char_count / 1.5)) word_count = len(chunk.content.split()) preview_items.append({ "position": i, "content": chunk.content, "char_count": char_count, "token_count": token_count, "word_count": word_count, "answer": chunk.metadata.get('answer'), "metadata": chunk.metadata if chunk.metadata else None, }) return ChunkPreviewResponse( chunks=preview_items, total=len(preview_items), strategy=strategy, chunk_size=chunk_size, chunk_overlap=chunk_overlap, ) # ==================== 分段管理 ==================== @router.get("/{kb_id}/segments", response_model=PaginatedResponse[SegmentListResponse], summary="知识库分段列表") async def list_kb_segments( kb_id: str, keyword: Optional[str] = Query(None, description="关键词搜索"), enabled: Optional[bool] = Query(None, description="是否启用"), embedding_status: Optional[str] = Query(None, alias="embeddingStatus", description="向量化状态: completed/failed/pending/skipped"), metadata_key: Optional[str] = Query(None, alias="metadataKey", description="元数据key"), metadata_value: Optional[str] = Query(None, alias="metadataValue", description="元数据value"), page: int = Query(1, ge=1, description="页码"), page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"), db: AsyncSession = Depends(get_db), ): """获取知识库所有分段""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) items, total = await service.get_kb_segments( knowledge_base_id=kb_id, page=page, page_size=page_size, keyword=keyword, enabled=enabled, embedding_status=embedding_status, metadata_key=metadata_key, metadata_value=metadata_value, ) # 获取文档名称映射 doc_ids = list({seg.document_id for seg in items}) doc_name_map = {} if doc_ids: doc_result = await db.execute( select(KnowledgeDocument.id, KnowledgeDocument.name).where( KnowledgeDocument.id.in_(doc_ids) ) ) doc_name_map = {row.id: row.name for row in doc_result} response_items = [_build_segment_list_response(seg, doc_name_map) for seg in items] return PaginatedResponse(items=response_items, total=total) @router.get("/{kb_id}/documents/{doc_id}/segments", response_model=PaginatedResponse[SegmentListResponse], summary="文档分段列表") async def list_document_segments( kb_id: str, doc_id: str, keyword: Optional[str] = Query(None, description="关键词搜索"), page: int = Query(1, ge=1, description="页码"), page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"), db: AsyncSession = Depends(get_db), ): """获取文档的分段列表""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) items, total = await service.get_segments( document_id=doc_id, page=page, page_size=page_size, keyword=keyword, ) # 获取文档名称 doc_result = await db.execute( select(KnowledgeDocument.name).where(KnowledgeDocument.id == doc_id) ) doc_name = doc_result.scalar_one_or_none() or "" doc_name_map = {doc_id: doc_name} response_items = [_build_segment_list_response(seg, doc_name_map) for seg in items] return PaginatedResponse(items=response_items, total=total) @router.put("/{kb_id}/segments/{segment_id}", response_model=SegmentListResponse, summary="更新分段") async def update_segment( kb_id: str, segment_id: str, data: SegmentUpdateInput, db: AsyncSession = Depends(get_db), ): """更新分段内容""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) segment = await service.update_segment( segment_id=segment_id, content=data.content, keywords=data.keywords, enabled=data.enabled, extra_metadata=data.extra_metadata, ) if not segment: raise HTTPException(status_code=404, detail="分段不存在") return _build_segment_list_response(segment, {}) @router.post("/{kb_id}/documents/{doc_id}/segments", response_model=SegmentListResponse, summary="手动添加分段") async def add_segment( kb_id: str, doc_id: str, data: SegmentCreateInput, db: AsyncSession = Depends(get_db), ): """手动添加分段""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) segment = await service.add_segment( knowledge_base_id=kb_id, document_id=doc_id, content=data.content, keywords=data.keywords, ) return _build_segment_list_response(segment, {}) @router.delete("/{kb_id}/segments/{segment_id}", response_model=ResponseModel, summary="删除分段") async def delete_segment(kb_id: str, segment_id: str, db: AsyncSession = Depends(get_db)): """删除分段""" from ai_platform.knowledge.services.document_service import DocumentService service = DocumentService(db) success = await service.delete_segment(segment_id) if not success: raise HTTPException(status_code=404, detail="分段不存在") return ResponseModel(message="删除成功") # ==================== 标注(Q&A)==================== @router.get("/{kb_id}/annotations", response_model=PaginatedResponse[AnnotationResponse], summary="标注列表") async def list_annotations( kb_id: str, page: int = Query(1, ge=1), page_size: int = Query(20, ge=1, le=100, alias="pageSize"), keyword: Optional[str] = Query(None, description="搜索关键词"), db: AsyncSession = Depends(get_db), ): """获取知识库标注列表""" from sqlalchemy import or_, func as sa_func query = select(KnowledgeAnnotation).where( KnowledgeAnnotation.knowledge_base_id == kb_id, KnowledgeAnnotation.is_deleted == False, ) if keyword: query = query.where(or_( KnowledgeAnnotation.question.contains(keyword), KnowledgeAnnotation.answer.contains(keyword), )) # 总数 count_query = select(sa_func.count()).select_from(query.subquery()) total = (await db.execute(count_query)).scalar() or 0 # 分页 query = query.order_by(KnowledgeAnnotation.sys_create_datetime.desc()) query = query.offset((page - 1) * page_size).limit(page_size) result = await db.execute(query) items = result.scalars().all() return PaginatedResponse(items=items, total=total) @router.post("/{kb_id}/annotations", response_model=AnnotationResponse, summary="创建标注") async def create_annotation( kb_id: str, data: AnnotationCreateInput, background_tasks: BackgroundTasks, db: AsyncSession = Depends(get_db), ): """创建知识库标注(Q&A 对)""" annotation = KnowledgeAnnotation( knowledge_base_id=kb_id, question=data.question, answer=data.answer, ) db.add(annotation) await db.commit() await db.refresh(annotation) # 后台向量化 question background_tasks.add_task(_vectorize_annotation, kb_id, str(annotation.id)) return annotation @router.put("/{kb_id}/annotations/{annotation_id}", response_model=AnnotationResponse, summary="更新标注") async def update_annotation( kb_id: str, annotation_id: str, data: AnnotationUpdateInput, background_tasks: BackgroundTasks, db: AsyncSession = Depends(get_db), ): """更新标注""" result = await db.execute( select(KnowledgeAnnotation).where( KnowledgeAnnotation.id == annotation_id, KnowledgeAnnotation.knowledge_base_id == kb_id, KnowledgeAnnotation.is_deleted == False, ) ) annotation = result.scalar_one_or_none() if not annotation: raise HTTPException(status_code=404, detail="标注不存在") question_changed = False update_data = data.model_dump(exclude_unset=True) for key, value in update_data.items(): if key == 'question' and value != annotation.question: question_changed = True setattr(annotation, key, value) await db.commit() await db.refresh(annotation) # 如果 question 变了,重新向量化 if question_changed: background_tasks.add_task(_vectorize_annotation, kb_id, annotation_id) return annotation @router.delete("/{kb_id}/annotations/{annotation_id}", response_model=ResponseModel, summary="删除标注") async def delete_annotation( kb_id: str, annotation_id: str, db: AsyncSession = Depends(get_db), ): """删除标注""" result = await db.execute( select(KnowledgeAnnotation).where( KnowledgeAnnotation.id == annotation_id, KnowledgeAnnotation.knowledge_base_id == kb_id, KnowledgeAnnotation.is_deleted == False, ) ) annotation = result.scalar_one_or_none() if not annotation: raise HTTPException(status_code=404, detail="标注不存在") annotation.is_deleted = True await 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, [annotation_id]) except Exception as e: logger.warning(f"删除标注向量失败: {e}") return ResponseModel(message="删除成功") @router.post("/{kb_id}/annotations/reindex", response_model=ResponseModel, summary="重新向量化所有标注") async def reindex_annotations( kb_id: str, background_tasks: BackgroundTasks, db: AsyncSession = Depends(get_db), ): """重新向量化知识库中所有 pending/failed 状态的标注""" result = await db.execute( select(KnowledgeAnnotation).where( KnowledgeAnnotation.knowledge_base_id == kb_id, KnowledgeAnnotation.is_deleted == False, KnowledgeAnnotation.embedding_status.in_(['pending', 'failed']), ) ) annotations = result.scalars().all() for ann in annotations: background_tasks.add_task(_vectorize_annotation, kb_id, str(ann.id)) return ResponseModel(message=f"已提交 {len(annotations)} 条标注的向量化任务") async def _vectorize_annotation(kb_id: str, annotation_id: str): """后台任务:向量化标注的 question""" from app.database import AsyncSessionLocal from ai_platform.knowledge.services.embedding_service import EmbeddingService from ai_platform.knowledge.vector_store import get_vector_store, VectorPoint try: async with AsyncSessionLocal() as db: # 获取知识库配置 kb_result = await db.execute( select(KnowledgeBase).where(KnowledgeBase.id == kb_id, KnowledgeBase.is_deleted == False) ) kb = kb_result.scalar_one_or_none() if not kb or not kb.embedding_model_id: return # 获取标注 ann_result = await db.execute( select(KnowledgeAnnotation).where( KnowledgeAnnotation.id == annotation_id, KnowledgeAnnotation.is_deleted == False, ) ) annotation = ann_result.scalar_one_or_none() if not annotation: return # 向量化 question embedding_service = EmbeddingService(db) embedding = await embedding_service.embed_text( model_id=kb.embedding_model_id, text=annotation.question, dimensions=kb.embedding_dimensions if kb.embedding_dimensions else None, ) # 写入 Qdrant(使用 annotation_id 作为向量 ID) vector_store = get_vector_store() vector_size = kb.embedding_dimensions or 1536 await vector_store.ensure_collection(kb_id, vector_size) await vector_store.upsert(kb_id, [VectorPoint( id=annotation_id, vector=embedding, payload={ 'knowledge_base_id': kb_id, 'type': 'annotation', }, )]) annotation.embedding_status = 'completed' await db.commit() logger.info(f"标注 {annotation_id} 向量化完成") except Exception as e: logger.error(f"标注向量化失败: {e}") try: async with AsyncSessionLocal() as db: ann_result = await db.execute( select(KnowledgeAnnotation).where(KnowledgeAnnotation.id == annotation_id) ) annotation = ann_result.scalar_one_or_none() if annotation: annotation.embedding_status = 'failed' await db.commit() except Exception: pass # ==================== 索引进度 SSE ==================== @router.get("/{kb_id}/indexing/progress", summary="订阅索引进度(SSE)") async def indexing_progress_sse(kb_id: str): """通过 Server-Sent Events 推送索引进度""" from fastapi.responses import StreamingResponse from ai_platform.knowledge.services.indexing_progress_service import IndexingProgressService import json import asyncio async def event_generator(): try: async for event in IndexingProgressService.subscribe(kb_id): data = json.dumps(event, ensure_ascii=False) yield f"data: {data}\n\n" # 如果收到 completed 或 failed,延迟后结束 if event.get("step") in ("completed", "failed"): await asyncio.sleep(0.5) except asyncio.CancelledError: pass except Exception as e: logger.warning(f"SSE 连接异常: {e}") return StreamingResponse( event_generator(), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "Connection": "keep-alive", "X-Accel-Buffering": "no", }, ) # ==================== 检索 ==================== @router.post("/retrieve", response_model=RetrievalResponse, summary="知识库检索") async def retrieve(data: RetrievalInput, db: AsyncSession = Depends(get_db)): """检索知识库""" import time from ai_platform.knowledge.services.retrieval_service import RetrievalService start_time = time.time() service = RetrievalService(db) # 解析实际使用的检索模式和 rerank 配置 actual_retrieval_mode = data.retrieval_mode or '' actual_rerank = data.rerank_enabled if not actual_retrieval_mode or actual_rerank is None: kb_result = await db.execute( select(KnowledgeBase).where( KnowledgeBase.id == data.knowledge_base_ids[0], KnowledgeBase.is_deleted == False ) ) first_kb = kb_result.scalar_one_or_none() if first_kb: if not actual_retrieval_mode: actual_retrieval_mode = first_kb.retrieval_mode or 'hybrid' if actual_rerank is None: actual_rerank = first_kb.rerank_enabled or False try: results = await service.retrieve( query=data.query, knowledge_base_ids=data.knowledge_base_ids, top_k=data.top_k, score_threshold=data.score_threshold, retrieval_mode=data.retrieval_mode, rerank_enabled=data.rerank_enabled, rerank_model_id=data.rerank_model_id, metadata_filter=data.metadata_filter, ) except Exception as e: logger.exception(f'检索失败: {e}') raise HTTPException(status_code=500, detail=f"检索失败: {str(e)}") elapsed = int((time.time() - start_time) * 1000) # 记录检索日志 try: from ai_platform.knowledge.models import KnowledgeRetrievalLog log = KnowledgeRetrievalLog( query=data.query, knowledge_base_ids=data.knowledge_base_ids, retrieval_mode=actual_retrieval_mode or 'hybrid', top_k=data.top_k, score_threshold=data.score_threshold, result_count=len(results), results=[ {'segment_id': r.segment_id, 'score': r.score, 'kb_id': r.knowledge_base_id} for r in results ], rerank_applied='true' if actual_rerank else 'false', elapsed_time=elapsed, source='api', ) db.add(log) await db.commit() except Exception as e: logger.warning(f'记录检索日志失败: {e}') return RetrievalResponse( results=results, total=len(results), query=data.query, elapsed_time=elapsed, retrieval_mode=actual_retrieval_mode or 'hybrid', rerank_applied=actual_rerank or False, ) # ==================== 命中统计 ==================== @router.get("/{kb_id}/hit-stats", summary="分段命中统计(热力图数据)") async def get_hit_stats( kb_id: str, top_n: int = Query(50, ge=1, le=200, alias="topN", description="返回 Top N 命中分段"), include_zero: bool = Query(False, alias="includeZero", description="是否包含未命中分段"), db: AsyncSession = Depends(get_db), ): """ 获取知识库分段命中统计,用于热力图可视化。 返回按命中次数降序排列的分段列表。 """ from ai_platform.knowledge.models import KnowledgeSegment, KnowledgeDocument from sqlalchemy import func as sa_func query = select( KnowledgeSegment.id, KnowledgeSegment.document_id, KnowledgeSegment.position, KnowledgeSegment.content, KnowledgeSegment.hit_count, KnowledgeSegment.enabled, ).where( KnowledgeSegment.knowledge_base_id == kb_id, KnowledgeSegment.is_deleted == False, ) if not include_zero: query = query.where(KnowledgeSegment.hit_count > 0) query = query.order_by(KnowledgeSegment.hit_count.desc()).limit(top_n) result = await db.execute(query) rows = result.all() # 获取文档名称 doc_ids = list({r.document_id for r in rows}) doc_names = {} if doc_ids: doc_result = await db.execute( select(KnowledgeDocument.id, KnowledgeDocument.name).where( KnowledgeDocument.id.in_(doc_ids) ) ) doc_names = {row.id: row.name for row in doc_result} # 总命中次数(用于计算百分比) total_hits_result = await db.execute( select(sa_func.coalesce(sa_func.sum(KnowledgeSegment.hit_count), 0)).where( KnowledgeSegment.knowledge_base_id == kb_id, KnowledgeSegment.is_deleted == False, ) ) total_hits = total_hits_result.scalar() or 0 items = [] for r in rows: items.append({ "segment_id": r.id, "document_id": r.document_id, "document_name": doc_names.get(r.document_id, ""), "position": r.position, "content_preview": r.content[:100] if r.content else "", "hit_count": r.hit_count or 0, "hit_percentage": round((r.hit_count or 0) / total_hits * 100, 2) if total_hits > 0 else 0, "enabled": r.enabled, }) # 未命中分段统计 zero_hit_result = await db.execute( select(sa_func.count()).where( KnowledgeSegment.knowledge_base_id == kb_id, KnowledgeSegment.is_deleted == False, KnowledgeSegment.hit_count == 0, ) ) zero_hit_count = zero_hit_result.scalar() or 0 return { "items": items, "total_hits": total_hits, "total_segments": len(items) + zero_hit_count if not include_zero else len(items), "zero_hit_count": zero_hit_count, } # ==================== 辅助函数 ==================== def _build_kb_response(kb: KnowledgeBase, embedding_model_name: str = "") -> dict: return { "id": kb.id, "application_id": kb.application_id, "is_global": kb.is_global or False, "name": kb.name, "code": kb.code, "description": kb.description or "", "icon": kb.icon or "", "embedding_model_id": kb.embedding_model_id, "embedding_model_name": embedding_model_name, "embedding_dimensions": kb.embedding_dimensions or 1536, "chunk_strategy": kb.chunk_strategy or "recursive", "chunk_size": kb.chunk_size or 500, "chunk_overlap": kb.chunk_overlap or 50, "separator": kb.separator, "retrieval_mode": kb.retrieval_mode or "hybrid", "top_k": kb.top_k or 5, "score_threshold": kb.score_threshold or 0.5, "rerank_enabled": kb.rerank_enabled or False, "rerank_model_id": kb.rerank_model_id, "retrieval_weight": kb.retrieval_weight or 1.0, "process_rules": kb.process_rules, "indexing_technique": kb.indexing_technique or "high_quality", "document_count": kb.document_count or 0, "segment_count": kb.segment_count or 0, "total_token_count": kb.total_token_count or 0, "total_char_count": kb.total_char_count or 0, "status": kb.status or "active", "sort": kb.sort or 0, "sys_create_datetime": kb.sys_create_datetime, "sys_update_datetime": kb.sys_update_datetime, } def _build_doc_response(doc: KnowledgeDocument) -> dict: return { "id": doc.id, "knowledge_base_id": doc.knowledge_base_id, "file_id": doc.file_id, "name": doc.name, "file_type": doc.file_type or "", "file_size": doc.file_size or 0, "content_hash": doc.content_hash or "", "segment_count": doc.segment_count or 0, "token_count": doc.token_count or 0, "char_count": doc.char_count or 0, "status": doc.status or "pending", "error_message": doc.error_message or "", "enabled": doc.enabled if doc.enabled is not None else True, "indexing_started_at": doc.indexing_started_at, "indexing_completed_at": doc.indexing_completed_at, "sys_create_datetime": doc.sys_create_datetime, "sys_update_datetime": doc.sys_update_datetime, } def _build_doc_list_response(doc: KnowledgeDocument) -> dict: return { "id": doc.id, "knowledge_base_id": doc.knowledge_base_id, "file_id": doc.file_id, "name": doc.name, "file_type": doc.file_type or "", "file_size": doc.file_size or 0, "segment_count": doc.segment_count or 0, "token_count": doc.token_count or 0, "status": doc.status or "pending", "duplicate_warning": doc.duplicate_warning or "", "enabled": doc.enabled if doc.enabled is not None else True, "sys_create_datetime": doc.sys_create_datetime, } def _build_segment_list_response(seg, doc_name_map: dict) -> dict: return { "id": seg.id, "document_id": seg.document_id, "document_name": doc_name_map.get(seg.document_id, ""), "position": seg.position or 0, "content": seg.content or "", "answer": getattr(seg, 'answer', None) or "", "token_count": seg.token_count or 0, "char_count": seg.char_count or 0, "word_count": getattr(seg, 'word_count', None) or 0, "page_number": getattr(seg, 'page_number', None), "keywords": (seg.keywords or []) if hasattr(seg, 'keywords') else [], "extra_metadata": getattr(seg, 'extra_metadata', None), "enabled": seg.enabled if seg.enabled is not None else True, "hit_count": seg.hit_count or 0, "embedding_status": seg.embedding_status or "pending", "sys_create_datetime": seg.sys_create_datetime, } async def _index_document_task(knowledge_base_id: str, document_id: str): """后台索引文档任务""" from app.database import AsyncSessionLocal from ai_platform.knowledge.services.indexing_service import IndexingService async with AsyncSessionLocal() as db: try: service = IndexingService(db) await service.index_document(knowledge_base_id, document_id) except Exception as e: logger.exception(f'后台索引文档失败: {e}') async def _batch_index_documents_task(knowledge_base_id: str, document_ids: list, max_concurrency: int = 3): """后台并发索引多个文档""" import asyncio from app.database import AsyncSessionLocal from ai_platform.knowledge.services.indexing_service import IndexingService semaphore = asyncio.Semaphore(max_concurrency) async def _index_one(doc_id: str): async with semaphore: async with AsyncSessionLocal() as db: try: service = IndexingService(db) await service.index_document(knowledge_base_id, doc_id) except Exception as e: logger.exception(f'并发索引文档失败 (doc_id={doc_id}): {e}') await asyncio.gather(*[_index_one(doc_id) for doc_id in document_ids])