318 lines
11 KiB
Python
318 lines
11 KiB
Python
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)
|