Files
ai-safety-platform/app/api/v1/endpoints/auth.py
T
baozaotumao2025 14722be770 first commit
2026-07-18 21:00:26 +08:00

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)