first commit
This commit is contained in:
@@ -0,0 +1,317 @@
|
||||
import logging
|
||||
import time
|
||||
from typing import Annotated
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, Request, Response, Security, status
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.deps import get_database
|
||||
from app.api.security import bearer_scheme
|
||||
from app.core.config import Settings, get_settings
|
||||
from app.core.exceptions import (
|
||||
ConfigurationError,
|
||||
ConflictError,
|
||||
ForbiddenError,
|
||||
NotFoundError,
|
||||
RefreshNotDueError,
|
||||
UnauthorizedError,
|
||||
)
|
||||
from app.db.repositories.repositories import UserRepository
|
||||
from app.schemas.api import (
|
||||
ChangePasswordRequest,
|
||||
CreateUserRequest,
|
||||
ErrorResponse,
|
||||
LoginRequest,
|
||||
ResetPasswordRequest,
|
||||
RefreshTokenRequest,
|
||||
TokenResponse,
|
||||
UpdateUserRequest,
|
||||
UserResponse,
|
||||
)
|
||||
from app.services.auth_service import (
|
||||
derive_refresh_token,
|
||||
issue_refresh_token,
|
||||
issue_token,
|
||||
token_digest,
|
||||
verify_password,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.post("/auth/login", response_model=TokenResponse, summary="登录并获取访问令牌")
|
||||
def login(
|
||||
payload: LoginRequest,
|
||||
db: Session = Depends(get_database),
|
||||
settings: Settings = Depends(get_settings),
|
||||
) -> TokenResponse:
|
||||
user = UserRepository(db).get_active_by_username(payload.username)
|
||||
if not user or not verify_password(payload.password, user.password_hash):
|
||||
logger.warning("Login failed username=%s", payload.username)
|
||||
raise UnauthorizedError("用户名或密码错误")
|
||||
try:
|
||||
family_id = str(uuid4())
|
||||
refresh_token = issue_refresh_token()
|
||||
refresh_expires_at = int(time.time()) + settings.auth_refresh_token_ttl_seconds
|
||||
token, expires_at = issue_token(user.id, settings, user.token_version, family_id)
|
||||
UserRepository(db).create_refresh_token(
|
||||
user_id=user.id,
|
||||
family_id=family_id,
|
||||
token_digest=token_digest(refresh_token),
|
||||
token_version=user.token_version,
|
||||
expires_at=refresh_expires_at,
|
||||
access_expires_at=expires_at,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise ConfigurationError("认证令牌配置无效") from exc
|
||||
logger.info("Login succeeded user_id=%s", user.id)
|
||||
return TokenResponse(
|
||||
access_token=token,
|
||||
expires_at=expires_at,
|
||||
refresh_token=refresh_token,
|
||||
refresh_expires_at=refresh_expires_at,
|
||||
)
|
||||
|
||||
|
||||
def _token_response(row, refresh_token: str, settings: Settings) -> TokenResponse:
|
||||
access_token, _ = issue_token(
|
||||
row.user_id,
|
||||
settings,
|
||||
row.token_version,
|
||||
row.family_id,
|
||||
row.access_expires_at,
|
||||
)
|
||||
return TokenResponse(
|
||||
access_token=access_token,
|
||||
expires_at=row.access_expires_at,
|
||||
refresh_token=refresh_token,
|
||||
refresh_expires_at=row.expires_at,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/auth/refresh",
|
||||
response_model=TokenResponse,
|
||||
summary="轮换刷新令牌并签发新的访问令牌",
|
||||
responses={
|
||||
401: {"model": ErrorResponse, "description": "刷新令牌或会话已失效。"},
|
||||
409: {"model": ErrorResponse, "description": "访问令牌尚未进入续约窗口。"},
|
||||
},
|
||||
)
|
||||
def refresh(
|
||||
payload: RefreshTokenRequest,
|
||||
idempotency_key: Annotated[str, Header(alias="Idempotency-Key", min_length=1, max_length=128)],
|
||||
db: Session = Depends(get_database),
|
||||
settings: Settings = Depends(get_settings),
|
||||
) -> TokenResponse:
|
||||
repo = UserRepository(db)
|
||||
digest = token_digest(payload.refresh_token)
|
||||
row = repo.get_refresh_token(digest)
|
||||
now = int(time.time())
|
||||
user = repo.get_active(row.user_id) if row else None
|
||||
if (
|
||||
not row
|
||||
or not user
|
||||
or row.revoked_at is not None
|
||||
or row.expires_at <= now
|
||||
or row.access_expires_at is None
|
||||
or user.token_version != row.token_version
|
||||
):
|
||||
raise UnauthorizedError("需要有效的刷新令牌")
|
||||
|
||||
replacement = derive_refresh_token(payload.refresh_token, idempotency_key, settings)
|
||||
replacement_digest = token_digest(replacement)
|
||||
if row.consumed_at is not None:
|
||||
if row.idempotency_key == idempotency_key and row.replacement_digest == replacement_digest:
|
||||
return _token_response(row, replacement, settings)
|
||||
repo.revoke_refresh_family(row.family_id, now)
|
||||
raise UnauthorizedError("刷新令牌已被重放,会话已撤销")
|
||||
|
||||
refresh_after = row.access_expires_at - settings.auth_token_refresh_window_seconds
|
||||
if now < refresh_after:
|
||||
raise RefreshNotDueError(
|
||||
"访问令牌尚未进入续约窗口", details={"refresh_after": refresh_after}
|
||||
)
|
||||
|
||||
access_expires_at = now + settings.auth_token_ttl_seconds
|
||||
refresh_expires_at = now + settings.auth_refresh_token_ttl_seconds
|
||||
if repo.rotate_refresh_token(
|
||||
row,
|
||||
idempotency_key=idempotency_key,
|
||||
replacement_digest=replacement_digest,
|
||||
access_expires_at=access_expires_at,
|
||||
refresh_expires_at=refresh_expires_at,
|
||||
consumed_at=now,
|
||||
):
|
||||
row.access_expires_at = access_expires_at
|
||||
row.expires_at = refresh_expires_at
|
||||
return _token_response(row, replacement, settings)
|
||||
|
||||
row = repo.get_refresh_token(digest)
|
||||
if (
|
||||
row
|
||||
and row.idempotency_key == idempotency_key
|
||||
and row.replacement_digest == replacement_digest
|
||||
):
|
||||
return _token_response(row, replacement, settings)
|
||||
if row:
|
||||
repo.revoke_refresh_family(row.family_id, now)
|
||||
raise UnauthorizedError("刷新令牌已被重放,会话已撤销")
|
||||
|
||||
|
||||
@router.get(
|
||||
"/auth/me",
|
||||
response_model=UserResponse,
|
||||
summary="读取当前登录用户",
|
||||
description="返回当前 Bearer 令牌对应的已启用用户及其管理员标志。",
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def get_current_user(request: Request, db: Session = Depends(get_database)) -> UserResponse:
|
||||
user = UserRepository(db).get_active(request.state.user_id)
|
||||
if not user:
|
||||
raise UnauthorizedError("需要有效的访问令牌")
|
||||
return UserResponse.model_validate(user, from_attributes=True)
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/auth/password",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="修改当前登录用户密码",
|
||||
description=(
|
||||
"校验当前密码后修改本人密码。成功后该用户已签发的所有 Bearer 令牌"
|
||||
"立即失效,客户端应使用新密码重新登录。"
|
||||
),
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def change_password(
|
||||
payload: ChangePasswordRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_database),
|
||||
) -> Response:
|
||||
repo = UserRepository(db)
|
||||
user = repo.get_active(request.state.user_id)
|
||||
if not user or not verify_password(payload.current_password, user.password_hash):
|
||||
raise UnauthorizedError("当前密码错误")
|
||||
repo.reset_password(user, payload.new_password)
|
||||
logger.info("User changed own password user_id=%s", user.id)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
def _admin(request: Request, db: Session) -> UserRepository:
|
||||
repo = UserRepository(db)
|
||||
user = repo.get_active(request.state.user_id)
|
||||
if not user or not user.is_admin:
|
||||
raise ForbiddenError("需要管理员权限")
|
||||
return repo
|
||||
|
||||
|
||||
@router.get(
|
||||
"/auth/users",
|
||||
response_model=list[UserResponse],
|
||||
summary="列出本地用户",
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def list_users(request: Request, db: Session = Depends(get_database)) -> list[UserResponse]:
|
||||
return [
|
||||
UserResponse.model_validate(user, from_attributes=True)
|
||||
for user in _admin(request, db).list()
|
||||
]
|
||||
|
||||
|
||||
@router.post(
|
||||
"/auth/users",
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
response_model=UserResponse,
|
||||
summary="创建本地用户",
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def create_user(
|
||||
payload: CreateUserRequest, request: Request, db: Session = Depends(get_database)
|
||||
) -> UserResponse:
|
||||
repo = _admin(request, db)
|
||||
try:
|
||||
user = repo.create(**payload.model_dump())
|
||||
except IntegrityError as exc:
|
||||
db.rollback()
|
||||
raise ConflictError("用户名已存在") from exc
|
||||
logger.info("User created user_id=%s by_user_id=%s", user.id, request.state.user_id)
|
||||
return UserResponse.model_validate(user, from_attributes=True)
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/auth/users/{user_id}",
|
||||
response_model=UserResponse,
|
||||
summary="启用或禁用本地用户",
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def update_user(
|
||||
user_id: int, payload: UpdateUserRequest, request: Request, db: Session = Depends(get_database)
|
||||
) -> UserResponse:
|
||||
repo = _admin(request, db)
|
||||
if user_id == request.state.user_id and not payload.is_active:
|
||||
raise ConflictError("不能禁用当前登录的管理员")
|
||||
user = repo.get(user_id)
|
||||
if not user:
|
||||
raise NotFoundError("用户不存在")
|
||||
updated = repo.set_active(user, payload.is_active)
|
||||
logger.info(
|
||||
"User active state changed user_id=%s by_user_id=%s", user_id, request.state.user_id
|
||||
)
|
||||
return UserResponse.model_validate(updated, from_attributes=True)
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/auth/users/{user_id}/password",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="重置本地用户密码",
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def reset_password(
|
||||
user_id: int,
|
||||
payload: ResetPasswordRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_database),
|
||||
) -> Response:
|
||||
repo = _admin(request, db)
|
||||
user = repo.get(user_id)
|
||||
if not user:
|
||||
raise NotFoundError("用户不存在")
|
||||
repo.reset_password(user, payload.password)
|
||||
logger.info("User password reset user_id=%s by_user_id=%s", user_id, request.state.user_id)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/auth/users/{user_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="删除本地用户",
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def delete_user(user_id: int, request: Request, db: Session = Depends(get_database)) -> Response:
|
||||
repo = _admin(request, db)
|
||||
if user_id == request.state.user_id:
|
||||
raise ConflictError("不能删除当前登录的管理员")
|
||||
user = repo.get(user_id)
|
||||
if not user:
|
||||
raise NotFoundError("用户不存在")
|
||||
repo.delete(user)
|
||||
logger.info("User deleted user_id=%s by_user_id=%s", user_id, request.state.user_id)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/auth/logout",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="撤销当前访问令牌",
|
||||
dependencies=[Security(bearer_scheme)],
|
||||
)
|
||||
def logout(request: Request, db: Session = Depends(get_database)) -> Response:
|
||||
repo = UserRepository(db)
|
||||
repo.revoke_token(token_digest(request.headers["Authorization"][7:]))
|
||||
if request.state.session_id:
|
||||
repo.revoke_refresh_family(request.state.session_id, int(time.time()))
|
||||
logger.info("Token revoked user_id=%s", request.state.user_id)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
Reference in New Issue
Block a user