diff --git a/backend/app/admin/api/v1/auth/auth.py b/backend/app/admin/api/v1/auth/auth.py index e9c2a48d..3b9f808a 100644 --- a/backend/app/admin/api/v1/auth/auth.py +++ b/backend/app/admin/api/v1/auth/auth.py @@ -2,7 +2,7 @@ # -*- coding: utf-8 -*- from typing import Annotated -from fastapi import APIRouter, Depends, Query, Request +from fastapi import APIRouter, Depends, Request, Response from fastapi.security import HTTPBasicCredentials from fastapi_limiter.depends import RateLimiter from starlette.background import BackgroundTasks @@ -28,18 +28,20 @@ async def swagger_login(obj: Annotated[HTTPBasicCredentials, Depends()]) -> GetS description='json 格式登录, 仅支持在第三方api工具调试, 例如: postman', dependencies=[Depends(RateLimiter(times=5, minutes=1))], ) -async def user_login(request: Request, obj: AuthLoginParam, background_tasks: BackgroundTasks) -> ResponseModel: - data = await auth_service.login(request=request, obj=obj, background_tasks=background_tasks) +async def user_login( + request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks +) -> ResponseModel: + data = await auth_service.login(request=request, response=response, obj=obj, background_tasks=background_tasks) return response_base.success(data=data) @router.post('/token/new', summary='创建新 token', dependencies=[DependsJwtAuth]) -async def create_new_token(request: Request, refresh_token: Annotated[str, Query(...)]) -> ResponseModel: - data = await auth_service.new_token(request=request, refresh_token=refresh_token) +async def create_new_token(request: Request, response: Response) -> ResponseModel: + data = await auth_service.new_token(request=request, response=response) return response_base.success(data=data) @router.post('/logout', summary='用户登出', dependencies=[DependsJwtAuth]) -async def user_logout(request: Request) -> ResponseModel: - await auth_service.logout(request=request) +async def user_logout(request: Request, response: Response) -> ResponseModel: + await auth_service.logout(request=request, response=response) return response_base.success() diff --git a/backend/app/admin/schema/token.py b/backend/app/admin/schema/token.py index 61244e59..73276ea7 100644 --- a/backend/app/admin/schema/token.py +++ b/backend/app/admin/schema/token.py @@ -19,10 +19,8 @@ class AccessTokenBase(SchemaBase): class GetNewToken(AccessTokenBase): - refresh_token: str - refresh_token_type: str = 'Bearer' - refresh_token_expire_time: datetime + pass -class GetLoginToken(GetNewToken): +class GetLoginToken(AccessTokenBase): user: GetUserInfoNoRelationDetail diff --git a/backend/app/admin/service/auth_service.py b/backend/app/admin/service/auth_service.py index 32b2fd31..a8e38d52 100644 --- a/backend/app/admin/service/auth_service.py +++ b/backend/app/admin/service/auth_service.py @@ -1,6 +1,6 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- -from fastapi import Request +from fastapi import Request, Response from fastapi.security import HTTPBasicCredentials from starlette.background import BackgroundTask, BackgroundTasks @@ -38,12 +38,14 @@ class AuthService: raise errors.AuthorizationError(msg='用户名或密码有误') elif not current_user.status: raise errors.AuthorizationError(msg='用户已被锁定, 请联系统管理员') - access_token, _ = await create_access_token(str(current_user.id), multi_login=current_user.is_multi_login) + access_token = await create_access_token(str(current_user.id), current_user.is_multi_login) await user_dao.update_login_time(db, obj.username) - return access_token, current_user + return access_token.access_token, current_user @staticmethod - async def login(*, request: Request, obj: AuthLoginParam, background_tasks: BackgroundTasks) -> GetLoginToken: + async def login( + *, request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks + ) -> GetLoginToken: async with async_db_session.begin() as db: try: current_user = await user_dao.get_by_username(db, obj.username) @@ -61,14 +63,8 @@ class AuthService: if captcha_code.lower() != obj.captcha.lower(): raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR) current_user_id = current_user.id - access_token, access_token_expire_time = await create_access_token( - str(current_user_id), multi_login=current_user.is_multi_login - ) - refresh_token, refresh_token_expire_time = await create_refresh_token( - sub=str(current_user_id), - expire_time=access_token_expire_time, - multi_login=current_user.is_multi_login, - ) + access_token = await create_access_token(str(current_user_id), current_user.is_multi_login) + refresh_token = await create_refresh_token(str(current_user_id), current_user.is_multi_login) except errors.NotFoundError as e: raise errors.NotFoundError(msg=e.msg) except (errors.AuthorizationError, errors.CustomError) as e: @@ -102,19 +98,29 @@ class AuthService: ) await redis_client.delete(f'{admin_settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{request.state.ip}') await user_dao.update_login_time(db, obj.username) + response.set_cookie( + settings.COOKIE_REFRESH_TOKEN_KEY, + refresh_token.refresh_token, + settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, + refresh_token.refresh_token_expire_time, + ) await db.refresh(current_user) data = GetLoginToken( - access_token=access_token, - refresh_token=refresh_token, - access_token_expire_time=access_token_expire_time, - refresh_token_expire_time=refresh_token_expire_time, + access_token=access_token.access_token, + access_token_expire_time=access_token.access_token_expire_time, user=current_user, # type: ignore ) return data @staticmethod - async def new_token(*, request: Request, refresh_token: str) -> GetNewToken: - user_id = jwt_decode(refresh_token) + async def new_token(*, request: Request, response: Response) -> GetNewToken: + refresh_token = request.cookies.get(settings.COOKIE_REFRESH_TOKEN_KEY) + if not refresh_token: + raise errors.TokenError(msg='Refresh Token 丢失,请重新登录') + try: + user_id = jwt_decode(refresh_token) + except Exception: + raise errors.TokenError(msg='Refresh Token 无效') if request.user.id != user_id: raise errors.TokenError(msg='Refresh Token 无效') async with async_db_session() as db: @@ -130,23 +136,34 @@ class AuthService: refresh_token=refresh_token, multi_login=current_user.is_multi_login, ) + response.set_cookie( + settings.COOKIE_REFRESH_TOKEN_KEY, + new_token.new_refresh_token, + settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, + new_token.new_refresh_token_expire_time, + ) data = GetNewToken( access_token=new_token.new_access_token, access_token_expire_time=new_token.new_access_token_expire_time, - refresh_token=new_token.new_refresh_token, - refresh_token_expire_time=new_token.new_refresh_token_expire_time, ) return data @staticmethod - async def logout(*, request: Request) -> None: + async def logout(*, request: Request, response: Response) -> None: token = await get_token(request) + refresh_token = request.cookies.get(settings.COOKIE_REFRESH_TOKEN_KEY) + response.delete_cookie(settings.COOKIE_REFRESH_TOKEN_KEY) if request.user.is_multi_login: key = f'{settings.TOKEN_REDIS_PREFIX}:{request.user.id}:{token}' await redis_client.delete(key) + if refresh_token: + key = f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{request.user.id}:{refresh_token}' + await redis_client.delete(key) else: key_prefix = f'{settings.TOKEN_REDIS_PREFIX}:{request.user.id}:' await redis_client.delete_prefix(key_prefix) + key_prefix = f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{request.user.id}:' + await redis_client.delete_prefix(key_prefix) auth_service = AuthService() diff --git a/backend/app/admin/service/oauth2_service.py b/backend/app/admin/service/oauth2_service.py index 320a000d..e2a6724d 100644 --- a/backend/app/admin/service/oauth2_service.py +++ b/backend/app/admin/service/oauth2_service.py @@ -1,7 +1,7 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- from fast_captcha import text_captcha -from fastapi import BackgroundTasks, Request +from fastapi import BackgroundTasks, Request, Response from backend.app.admin.conf import admin_settings from backend.app.admin.crud.crud_user import user_dao @@ -13,6 +13,7 @@ from backend.app.admin.service.login_log_service import LoginLogService from backend.common.enums import LoginLogStatusType, UserSocialType from backend.common.exception.errors import AuthorizationError from backend.common.security import jwt +from backend.core.conf import settings from backend.database.db_mysql import async_db_session from backend.database.db_redis import redis_client from backend.utils.timezone import timezone @@ -21,7 +22,7 @@ from backend.utils.timezone import timezone class OAuth2Service: @staticmethod async def create_with_login( - *, request: Request, background_tasks: BackgroundTasks, user: dict, social: UserSocialType + *, request: Request, response: Response, background_tasks: BackgroundTasks, user: dict, social: UserSocialType ) -> GetLoginToken | None: async with async_db_session.begin() as db: # 获取 OAuth2 平台用户信息 @@ -54,12 +55,8 @@ class OAuth2Service: new_user_social = CreateUserSocialParam(source=social.value, uid=str(_id), user_id=sys_user.id) await user_social_dao.create(db, new_user_social) # 创建 token - access_token, access_token_expire_time = await jwt.create_access_token( - str(sys_user.id), multi_login=sys_user.is_multi_login - ) - refresh_token, refresh_token_expire_time = await jwt.create_refresh_token( - str(sys_user.id), access_token_expire_time, multi_login=sys_user.is_multi_login - ) + access_token = await jwt.create_access_token(str(sys_user.id), sys_user.is_multi_login) + refresh_token = await jwt.create_refresh_token(str(sys_user.id), multi_login=sys_user.is_multi_login) await user_dao.update_login_time(db, sys_user.username) await db.refresh(sys_user) login_log = dict( @@ -72,11 +69,15 @@ class OAuth2Service: ) background_tasks.add_task(LoginLogService.create, **login_log) await redis_client.delete(f'{admin_settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{request.state.ip}') + response.set_cookie( + settings.COOKIE_REFRESH_TOKEN_KEY, + refresh_token.refresh_token, + settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, + refresh_token.refresh_token_expire_time, + ) data = GetLoginToken( - access_token=access_token, - refresh_token=refresh_token, - access_token_expire_time=access_token_expire_time, - refresh_token_expire_time=refresh_token_expire_time, + access_token=access_token.access_token, + access_token_expire_time=access_token.access_token_expire_time, user=sys_user, # type: ignore ) return data diff --git a/backend/common/dataclasses.py b/backend/common/dataclasses.py index aeaca534..124d1d37 100644 --- a/backend/common/dataclasses.py +++ b/backend/common/dataclasses.py @@ -26,7 +26,7 @@ class UserAgentInfo: @dataclasses.dataclass -class RequestCallNextReturn: +class RequestCallNext: code: str msg: str status: StatusType @@ -35,8 +35,20 @@ class RequestCallNextReturn: @dataclasses.dataclass -class NewTokenReturn: +class NewToken: new_access_token: str - new_refresh_token: str new_access_token_expire_time: datetime + new_refresh_token: str new_refresh_token_expire_time: datetime + + +@dataclasses.dataclass +class AccessToken: + access_token: str + access_token_expire_time: datetime + + +@dataclasses.dataclass +class RefreshToken: + refresh_token: str + refresh_token_expire_time: datetime diff --git a/backend/common/response/response_schema.py b/backend/common/response/response_schema.py index 8346be2a..cf3804da 100644 --- a/backend/common/response/response_schema.py +++ b/backend/common/response/response_schema.py @@ -57,7 +57,7 @@ class ResponseBase: @router.get('/test') def test() -> ResponseModel: - return await response_base.success(data={'test': 'test'}) + return response_base.success(data={'test': 'test'}) """ @staticmethod diff --git a/backend/common/security/jwt.py b/backend/common/security/jwt.py index 9c2075b3..61ec623b 100644 --- a/backend/common/security/jwt.py +++ b/backend/common/security/jwt.py @@ -1,6 +1,6 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- -from datetime import datetime, timedelta +from datetime import timedelta from asgiref.sync import sync_to_async from fastapi import Depends, Request @@ -11,7 +11,7 @@ from passlib.context import CryptContext from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.model import User -from backend.common.dataclasses import NewTokenReturn +from backend.common.dataclasses import AccessToken, NewToken, RefreshToken from backend.common.exception.errors import AuthorizationError, TokenError from backend.core.conf import settings from backend.database.db_redis import redis_client @@ -45,83 +45,78 @@ def password_verify(plain_password: str, hashed_password: str) -> bool: return pwd_context.verify(plain_password, hashed_password) -async def create_access_token(sub: str, expires_delta: timedelta | None = None, **kwargs) -> tuple[str, datetime]: +async def create_access_token(sub: str, multi_login: bool) -> AccessToken: """ Generate encryption token :param sub: The subject/userid of the JWT - :param expires_delta: Increased expiry time + :param multi_login: multipoint login for user :return: """ - if expires_delta: - expire = timezone.now() + expires_delta - expire_seconds = int(expires_delta.total_seconds()) - else: - expire = timezone.now() + timedelta(seconds=settings.TOKEN_EXPIRE_SECONDS) - expire_seconds = settings.TOKEN_EXPIRE_SECONDS - multi_login = kwargs.pop('multi_login', None) - to_encode = {'exp': expire, 'sub': sub, **kwargs} - token = jwt.encode(to_encode, settings.TOKEN_SECRET_KEY, settings.TOKEN_ALGORITHM) + expire = timezone.now() + timedelta(seconds=settings.TOKEN_EXPIRE_SECONDS) + expire_seconds = settings.TOKEN_EXPIRE_SECONDS + + to_encode = {'exp': expire, 'sub': sub} + access_token = jwt.encode(to_encode, settings.TOKEN_SECRET_KEY, settings.TOKEN_ALGORITHM) + if multi_login is False: key_prefix = f'{settings.TOKEN_REDIS_PREFIX}:{sub}' await redis_client.delete_prefix(key_prefix) - key = f'{settings.TOKEN_REDIS_PREFIX}:{sub}:{token}' - await redis_client.setex(key, expire_seconds, token) - return token, expire + + key = f'{settings.TOKEN_REDIS_PREFIX}:{sub}:{access_token}' + await redis_client.setex(key, expire_seconds, access_token) + return AccessToken(access_token=access_token, access_token_expire_time=expire) -async def create_refresh_token(sub: str, expire_time: datetime | None = None, **kwargs) -> tuple[str, datetime]: +async def create_refresh_token(sub: str, multi_login: bool) -> RefreshToken: """ Generate encryption refresh token, only used to create a new token :param sub: The subject/userid of the JWT - :param expire_time: expiry time + :param multi_login: multipoint login for user :return: """ - if expire_time: - expire = expire_time + timedelta(seconds=settings.TOKEN_REFRESH_EXPIRE_SECONDS) - expire_datetime = timezone.f_datetime(expire_time) - current_datetime = timezone.now() - if expire_datetime < current_datetime: - raise TokenError(msg='Refresh Token 已过期') - expire_seconds = int((expire_datetime - current_datetime).total_seconds()) - else: - expire = timezone.now() + timedelta(seconds=settings.TOKEN_EXPIRE_SECONDS) - expire_seconds = settings.TOKEN_REFRESH_EXPIRE_SECONDS - multi_login = kwargs.pop('multi_login', None) - to_encode = {'exp': expire, 'sub': sub, **kwargs} + expire = timezone.now() + timedelta(seconds=settings.TOKEN_REFRESH_EXPIRE_SECONDS) + expire_seconds = settings.TOKEN_REFRESH_EXPIRE_SECONDS + + to_encode = {'exp': expire, 'sub': sub} refresh_token = jwt.encode(to_encode, settings.TOKEN_SECRET_KEY, settings.TOKEN_ALGORITHM) + if multi_login is False: key_prefix = f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{sub}' await redis_client.delete_prefix(key_prefix) + key = f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{sub}:{refresh_token}' await redis_client.setex(key, expire_seconds, refresh_token) - return refresh_token, expire + return RefreshToken(refresh_token=refresh_token, refresh_token_expire_time=expire) -async def create_new_token(sub: str, token: str, refresh_token: str, **kwargs) -> NewTokenReturn: +async def create_new_token(sub: str, token: str, refresh_token: str, multi_login: bool) -> NewToken: """ Generate new token :param sub: :param token :param refresh_token: + :param multi_login: :return: """ redis_refresh_token = await redis_client.get(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{sub}:{refresh_token}') if not redis_refresh_token or redis_refresh_token != refresh_token: raise TokenError(msg='Refresh Token 已过期') - new_access_token, new_access_token_expire_time = await create_access_token(sub, **kwargs) - new_refresh_token, new_refresh_token_expire_time = await create_refresh_token(sub, **kwargs) + + new_access_token = await create_access_token(sub, multi_login) + new_refresh_token = await create_refresh_token(sub, multi_login) + token_key = f'{settings.TOKEN_REDIS_PREFIX}:{sub}:{token}' refresh_token_key = f'{settings.TOKEN_REDIS_PREFIX}:{sub}:{refresh_token}' await redis_client.delete(token_key) await redis_client.delete(refresh_token_key) - return NewTokenReturn( - new_access_token=new_access_token, - new_refresh_token=new_refresh_token, - new_access_token_expire_time=new_access_token_expire_time, - new_refresh_token_expire_time=new_refresh_token_expire_time, + return NewToken( + new_access_token=new_access_token.access_token, + new_access_token_expire_time=new_access_token.access_token_expire_time, + new_refresh_token=new_refresh_token.refresh_token, + new_refresh_token_expire_time=new_refresh_token.refresh_token_expire_time, ) diff --git a/backend/core/conf.py b/backend/core/conf.py index 2ac86c2d..de4b6059 100644 --- a/backend/core/conf.py +++ b/backend/core/conf.py @@ -84,13 +84,17 @@ class Settings(BaseSettings): # Token TOKEN_ALGORITHM: str = 'HS256' # 算法 TOKEN_EXPIRE_SECONDS: int = 60 * 60 * 24 * 1 # 过期时间,单位:秒 - TOKEN_REFRESH_EXPIRE_SECONDS: int = 60 * 60 * 24 * 7 # 刷新过期时间,单位:秒 + TOKEN_REFRESH_EXPIRE_SECONDS: int = 60 * 60 * 24 * 7 # refresh token 过期时间,单位:秒 TOKEN_REDIS_PREFIX: str = 'fba:token' TOKEN_REFRESH_REDIS_PREFIX: str = 'fba:token:refresh' TOKEN_EXCLUDE: list[str] = [ # JWT / RBAC 白名单 f'{API_V1_STR}/auth/login', ] + # Cookies + COOKIE_REFRESH_TOKEN_KEY: str = 'fba_refresh_token' + COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS: int = TOKEN_REFRESH_EXPIRE_SECONDS + # Sys User USER_REDIS_PREFIX: str = 'fba:user' USER_REDIS_EXPIRE_SECONDS: int = 60 * 60 * 24 * 7 diff --git a/backend/middleware/opera_log_middleware.py b/backend/middleware/opera_log_middleware.py index f6539dc9..22f9cad8 100644 --- a/backend/middleware/opera_log_middleware.py +++ b/backend/middleware/opera_log_middleware.py @@ -10,7 +10,7 @@ from starlette.requests import Request from backend.app.admin.schema.opera_log import CreateOperaLogParam from backend.app.admin.service.opera_log_service import OperaLogService -from backend.common.dataclasses import RequestCallNextReturn +from backend.common.dataclasses import RequestCallNext from backend.common.enums import OperaLogCipherType, StatusType from backend.common.log import log from backend.core.conf import settings @@ -90,7 +90,7 @@ class OperaLogMiddleware(BaseHTTPMiddleware): return res.response - async def execute_request(self, request: Request, call_next) -> RequestCallNextReturn: + async def execute_request(self, request: Request, call_next) -> RequestCallNext: """执行请求""" code = 200 msg = 'Success' @@ -108,7 +108,7 @@ class OperaLogMiddleware(BaseHTTPMiddleware): status = StatusType.disable err = e - return RequestCallNextReturn(code=str(code), msg=msg, status=status, err=err, response=response) + return RequestCallNext(code=str(code), msg=msg, status=status, err=err, response=response) @staticmethod @sync_to_async