Files
amb_rag/backend/app/api/knowledge_bases.py

173 lines
5.2 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 路由。"""
from fastapi import APIRouter, Depends, Query
from sqlalchemy.orm import Session
from app.api.deps import get_current_user, get_db
from app.models.user import User
from app.schemas.knowledge_base import (
KbCreateRequest,
KbListResponse,
KbResponse,
KbSetExpiryRequest,
KbTokenResponse,
KbUpdateRequest,
)
from app.services.kb_service import KnowledgeBaseService
router = APIRouter(prefix="/knowledge-bases", tags=["knowledge-bases"])
@router.get("", response_model=KbListResponse)
def list_knowledge_bases(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbListResponse:
svc = KnowledgeBaseService(db)
items, total = svc.list_by_user(user, page=page, page_size=page_size)
doc_count_map = db and svc._kb_repo.get_document_count_map(user.id)
return KbListResponse(
items=[_to_response(kb, doc_count_map.get(kb.id, 0)) for kb in items],
total=total,
)
@router.post("", response_model=KbResponse, status_code=201)
def create_knowledge_base(
body: KbCreateRequest,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbResponse:
svc = KnowledgeBaseService(db)
kb, token = svc.create(user, body.name, body.description)
return _to_response(kb, ai_url=f"/k/{token}")
@router.get("/{kb_id}", response_model=KbResponse)
def get_knowledge_base(
kb_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbResponse:
svc = KnowledgeBaseService(db)
kb = svc.get_or_404(kb_id, user)
doc_count = svc._kb_repo.count_documents(kb_id)
return _to_response(kb, doc_count)
@router.put("/{kb_id}", response_model=KbResponse)
def update_knowledge_base(
kb_id: str,
body: KbUpdateRequest,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbResponse:
svc = KnowledgeBaseService(db)
kb = svc.update(kb_id, user, body.name, body.description)
return _to_response(kb)
@router.delete("/{kb_id}", status_code=204)
def delete_knowledge_base(
kb_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> None:
svc = KnowledgeBaseService(db)
svc.delete(kb_id, user)
@router.post("/{kb_id}/regenerate-token", response_model=KbTokenResponse)
def regenerate_token(
kb_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbTokenResponse:
svc = KnowledgeBaseService(db)
kb, token = svc.regenerate_token(kb_id, user)
return KbTokenResponse(
token=token, ai_url=f"/k/{token}", token_hint=kb.token_hint,
expires_at=kb.token_expires_at, is_expired=False,
)
@router.post("/{kb_id}/enable", response_model=KbResponse)
def enable_knowledge_base(
kb_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbResponse:
svc = KnowledgeBaseService(db)
kb = svc.set_enabled(kb_id, user, True)
return _to_response(kb)
@router.post("/{kb_id}/disable", response_model=KbResponse)
def disable_knowledge_base(
kb_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbResponse:
svc = KnowledgeBaseService(db)
kb = svc.set_enabled(kb_id, user, False)
return _to_response(kb)
@router.get("/{kb_id}/link", response_model=KbTokenResponse)
def get_knowledge_base_link(
kb_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbTokenResponse:
"""获取完整 AI 链接(解密 token+ 有效期信息。"""
from app.services.kb_public_service import KbPublicService
svc = KnowledgeBaseService(db)
kb = svc.get_or_404(kb_id, user)
token = svc.get_full_token(kb)
if token is None:
return KbTokenResponse(
token="", ai_url="", token_hint=kb.token_hint or "",
expires_at=None, is_expired=False,
)
return KbTokenResponse(
token=token,
ai_url=f"/k/{token}",
token_hint=kb.token_hint or "",
expires_at=kb.token_expires_at,
is_expired=KbPublicService.is_expired(kb),
)
@router.post("/{kb_id}/set-expiry", response_model=KbResponse)
def set_link_expiry(
kb_id: str,
body: KbSetExpiryRequest,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> KbResponse:
"""设置 AI 链接有效期。expires_in_minutes 为 null 表示长期有效。"""
svc = KnowledgeBaseService(db)
kb = svc.set_expiry(kb_id, user, body.expires_in_minutes)
doc_count = svc._kb_repo.count_documents(kb_id)
return KbResponse(
id=kb.id, name=kb.name, description=kb.description, enabled=kb.enabled,
token_hint=kb.token_hint, document_count=doc_count,
created_at=kb.created_at, updated_at=kb.updated_at,
)
def _to_response(kb, doc_count: int = 0, ai_url: str | None = None) -> KbResponse:
return KbResponse(
id=kb.id,
name=kb.name,
description=kb.description,
enabled=kb.enabled,
token_hint=kb.token_hint,
ai_url=ai_url,
document_count=doc_count,
created_at=kb.created_at,
updated_at=kb.updated_at,
)