173 lines
5.2 KiB
Python
173 lines
5.2 KiB
Python
"""知识库 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,
|
||
) |