This commit is contained in:
amb
2026-09-01 13:00:36 +08:00
parent 1d8621717a
commit dfd38c99a0
35 changed files with 3141 additions and 21 deletions
+129
View File
@@ -0,0 +1,129 @@
"""认证路由:注册、登录、登出、用户信息。"""
from fastapi import APIRouter, Depends, Response
from sqlalchemy.orm import Session
from app.api.deps import get_current_user, get_db
from app.core.session import get_cookie_params
from app.models.user import User
from app.schemas.user import (
LoginRequest,
MeResponse,
RegisterRequest,
StorageInfoResponse,
UpdateMeRequest,
UserResponse,
)
from app.services.auth_service import AuthService
router = APIRouter(prefix="/auth", tags=["auth"])
@router.post("/register", response_model=UserResponse, status_code=201)
def register(body: RegisterRequest, response: Response, db: Session = Depends(get_db)) -> UserResponse:
"""注册新用户并自动登录。"""
auth_service = AuthService(db)
user, token = auth_service.register(body.username, body.email, body.password)
response.set_cookie(value=token, **get_cookie_params())
return _user_response(user)
@router.post("/login", response_model=UserResponse)
def login(body: LoginRequest, response: Response, db: Session = Depends(get_db)) -> UserResponse:
"""登录。"""
auth_service = AuthService(db)
user, token = auth_service.login(body.username_or_email, body.password)
response.set_cookie(value=token, **get_cookie_params())
return _user_response(user)
@router.post("/logout", status_code=204)
def logout(
response: Response,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> None:
"""登出:删除 session 并清除 Cookie。"""
from app.core.session import SESSION_COOKIE_NAME
token = "" # 不需要实际 tokenlogout 内部通过 user_id 找 session
auth_service = AuthService(db)
# 直接删除所有该用户的 session(MVP 简化:单设备)
from app.core.session import _store
to_delete = [t for t, e in _store.items() if e.user_id == user.id]
for t in to_delete:
auth_service.logout(t)
response.delete_cookie(SESSION_COOKIE_NAME, path="/")
return None
@router.get("/me", response_model=MeResponse)
def get_me(user: User = Depends(get_current_user)) -> MeResponse:
"""获取当前用户信息。"""
from app.core.config import get_settings
settings = get_settings()
return MeResponse(
id=user.id,
username=user.username,
email=user.email,
status=user.status,
storage_used=user.storage_used,
storage_quota=settings.default_storage_quota,
created_at=user.created_at,
)
@router.patch("/me", response_model=MeResponse)
def update_me(
body: UpdateMeRequest,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> MeResponse:
"""修改密码。"""
auth_service = AuthService(db)
if body.password is not None:
auth_service.change_password(user, body.password)
from app.core.config import get_settings
settings = get_settings()
return MeResponse(
id=user.id,
username=user.username,
email=user.email,
status=user.status,
storage_used=user.storage_used,
storage_quota=settings.default_storage_quota,
created_at=user.created_at,
)
@router.get("/storage", response_model=StorageInfoResponse)
def get_storage(user: User = Depends(get_current_user)) -> StorageInfoResponse:
"""获取存储用量信息。"""
from app.core.config import get_settings
settings = get_settings()
return StorageInfoResponse(
storage_used=user.storage_used,
storage_quota=settings.default_storage_quota,
storage_used_mb=round(user.storage_used / (1024 * 1024), 2),
storage_quota_mb=round(settings.default_storage_quota / (1024 * 1024), 2),
)
def _user_response(user: User) -> UserResponse:
from app.core.config import get_settings
settings = get_settings()
return UserResponse(
id=user.id,
username=user.username,
email=user.email,
status=user.status,
storage_used=user.storage_used,
storage_quota=settings.default_storage_quota,
created_at=user.created_at,
)
+31
View File
@@ -0,0 +1,31 @@
"""API 依赖注入。
- get_current_user: 从 Cookie 读 session → 查库 → 返回 User
- 所有需登录的路由用 Depends(get_current_user)
"""
from fastapi import Depends, Request
from sqlalchemy.orm import Session
from app.core.db import get_session as get_db_session
from app.core.session import SESSION_COOKIE_NAME
from app.models.user import User
from app.services.auth_service import AuthService
def get_db() -> Session:
"""数据库会话依赖(同步 SQLAlchemy)。"""
yield from get_db_session()
def get_current_user(
request: Request,
db: Session = Depends(get_db),
) -> User:
"""从 Cookie 中读取 session token,返回当前登录用户。
未登录或 session 过期抛出 AuthRequiredError401)。
"""
token = request.cookies.get(SESSION_COOKIE_NAME)
auth_service = AuthService(db)
return auth_service.get_current_user(token)
+140
View File
@@ -0,0 +1,140 @@
"""文档 API 路由。"""
from fastapi import APIRouter, Depends, Query, UploadFile, File, Form
from sqlalchemy.orm import Session
from app.api.deps import get_current_user, get_db
from app.models.user import User
from app.schemas.document import (
DocumentListResponse,
DocumentResponse,
DocumentUpdateRequest,
DocumentUploadResponse,
)
from app.services.doc_service import DocumentService
router = APIRouter(prefix="/documents", tags=["documents"])
@router.post("/upload", response_model=DocumentUploadResponse, status_code=201)
async def upload_document(
kb_id: str = Form(...),
file: UploadFile = File(...),
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> DocumentUploadResponse:
"""上传文档到知识库。"""
content = await file.read()
svc = DocumentService(db)
doc = svc.upload(
user=user,
kb_id=kb_id,
filename=file.filename or "unknown",
content=content,
)
return DocumentUploadResponse(
id=doc.id,
original_filename=doc.original_filename,
file_size=doc.file_size,
status=doc.status,
message="文档上传成功,等待解析。",
)
@router.get("", response_model=DocumentListResponse)
def list_documents(
kb_id: str = Query(...),
page: int = Query(1, ge=1),
page_size: int = Query(50, ge=1, le=200),
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> DocumentListResponse:
svc = DocumentService(db)
items, total = svc.list_by_knowledge_base(kb_id, user, page=page, page_size=page_size)
return DocumentListResponse(
items=[_to_response(doc) for doc in items],
total=total,
)
@router.get("/{doc_id}", response_model=DocumentResponse)
def get_document(
doc_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> DocumentResponse:
svc = DocumentService(db)
doc = svc.get_or_404(doc_id, user)
return _to_response(doc)
@router.put("/{doc_id}", response_model=DocumentResponse)
def update_document(
doc_id: str,
body: DocumentUpdateRequest,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> DocumentResponse:
svc = DocumentService(db)
doc = svc.get_or_404(doc_id, user)
update_fields = {}
if body.title is not None:
update_fields["title"] = body.title
if body.description is not None:
update_fields["description"] = body.description
if body.keywords is not None:
update_fields["keywords"] = body.keywords
if body.category_id is not None:
update_fields["category_id"] = body.category_id
if update_fields:
from app.repositories.doc_repo import DocumentRepository
DocumentRepository(db).update(doc, **update_fields)
db.commit()
return _to_response(doc)
@router.delete("/{doc_id}", status_code=204)
def delete_document(
doc_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> None:
svc = DocumentService(db)
svc.delete(doc_id, user)
@router.post("/{doc_id}/reprocess", response_model=DocumentResponse)
def reprocess_document(
doc_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> DocumentResponse:
"""重新解析文档。"""
svc = DocumentService(db)
doc = svc.get_or_404(doc_id, user)
svc._process_document(doc)
db.refresh(doc)
return _to_response(doc)
def _to_response(doc) -> DocumentResponse:
return DocumentResponse(
id=doc.id,
knowledge_base_id=doc.knowledge_base_id,
original_filename=doc.original_filename,
file_size=doc.file_size,
mime_type=doc.mime_type,
file_ext=doc.file_ext,
sha256=doc.sha256,
title=doc.title,
description=doc.description,
keywords=doc.keywords,
content_summary=doc.content_summary,
status=doc.status,
error_code=doc.error_code,
doc_token_hint=doc.doc_token_hint,
category_id=doc.category_id,
created_at=doc.created_at,
updated_at=doc.updated_at,
)
+140
View File
@@ -0,0 +1,140 @@
"""知识库 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,
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)
@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)。"""
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 "")
return KbTokenResponse(token=token, ai_url=f"/k/{token}", token_hint=kb.token_hint or "")
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,
)