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

1125 lines
41 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
知识库 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])