1125 lines
41 KiB
Python
1125 lines
41 KiB
Python
"""
|
||
知识库 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])
|