feat: 用户登录注册相关接口调整,登录成功之后返回access_token以及refresh_token
This commit is contained in:
@@ -84,8 +84,12 @@ def add_user_route(app: FastAPI):
|
|||||||
|
|
||||||
public_user = PublicUser(**user.model_dump())
|
public_user = PublicUser(**user.model_dump())
|
||||||
|
|
||||||
active_user_token = CryptUtils.create_access_token(public_user.username, expires_delta=timedelta(days=365 * 3))
|
verify_user_token = CryptUtils.create_token(
|
||||||
active_url = f"{env.server_domain}:{env.server_port}/verify?token={active_user_token}"
|
username=public_user.username,
|
||||||
|
type="verify",
|
||||||
|
expires_delta=timedelta(days=365 * 3),
|
||||||
|
)
|
||||||
|
active_url = f"{env.server_domain}:{env.server_port}/verify?token={verify_user_token}"
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"result": public_user,
|
"result": public_user,
|
||||||
@@ -95,9 +99,10 @@ def add_user_route(app: FastAPI):
|
|||||||
# 验证用户账号接口
|
# 验证用户账号接口
|
||||||
@app.get("/verify")
|
@app.get("/verify")
|
||||||
async def _verify(token: str, session: AsyncSessionDep):
|
async def _verify(token: str, session: AsyncSessionDep):
|
||||||
username = CryptUtils.get_username_from_token(token)
|
token_info = CryptUtils.get_token_info(token)
|
||||||
|
username = token_info.get('username')
|
||||||
|
|
||||||
if not username:
|
if not username or token_info.get('type') != 'verify':
|
||||||
return {"result": None, "error": "token无效或者已经过期"}
|
return {"result": None, "error": "token无效或者已经过期"}
|
||||||
|
|
||||||
query = select(UserModel).where(UserModel.username == username)
|
query = select(UserModel).where(UserModel.username == username)
|
||||||
@@ -133,14 +138,21 @@ def add_user_route(app: FastAPI):
|
|||||||
headers={"WWW-Authenticate": "Bearer"},
|
headers={"WWW-Authenticate": "Bearer"},
|
||||||
)
|
)
|
||||||
|
|
||||||
token = Token(
|
access_token = CryptUtils.create_token(
|
||||||
token=CryptUtils.create_access_token(user.username),
|
username=user.username,
|
||||||
token_type="Bearer",
|
type="access",
|
||||||
|
expires_delta=timedelta(minutes=env.jwt_access_token_expire_minutes),
|
||||||
|
)
|
||||||
|
refresh_token = CryptUtils.create_token(
|
||||||
|
username=user.username,
|
||||||
|
type="refresh",
|
||||||
|
expires_delta=timedelta(minutes=env.jwt_refresh_token_expire_minutes),
|
||||||
)
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"result": user,
|
"result": user,
|
||||||
"token": token,
|
"access_token": access_token,
|
||||||
|
"refresh_token": refresh_token,
|
||||||
}
|
}
|
||||||
|
|
||||||
# 获取用户信息接口
|
# 获取用户信息接口
|
||||||
@@ -187,7 +199,7 @@ unauthorized_exception = HTTPException(
|
|||||||
# 获取当前用户信息,通过注入的token来获取当前用户信息,如果token有效则返回用户信息,无效则抛出异常
|
# 获取当前用户信息,通过注入的token来获取当前用户信息,如果token有效则返回用户信息,无效则抛出异常
|
||||||
async def get_current_user(session: AsyncSessionDep, token: str = Depends(oauth2_scheme)):
|
async def get_current_user(session: AsyncSessionDep, token: str = Depends(oauth2_scheme)):
|
||||||
try:
|
try:
|
||||||
username = CryptUtils.get_username_from_token(token)
|
username = CryptUtils.get_token_info(token).get('username')
|
||||||
if not username:
|
if not username:
|
||||||
raise unauthorized_exception
|
raise unauthorized_exception
|
||||||
except InvalidTokenError:
|
except InvalidTokenError:
|
||||||
|
|||||||
+24
-10
@@ -1,4 +1,5 @@
|
|||||||
from datetime import timedelta, datetime, timezone
|
from datetime import timedelta, datetime, timezone
|
||||||
|
from typing import TypedDict, Literal, TypeAlias
|
||||||
|
|
||||||
import jwt
|
import jwt
|
||||||
from passlib.context import CryptContext
|
from passlib.context import CryptContext
|
||||||
@@ -7,6 +8,17 @@ from app.config.env import env
|
|||||||
|
|
||||||
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
||||||
|
|
||||||
|
# token的类型,access用于接口认证,refresh用于刷新access token,verify用于激活用户账号
|
||||||
|
AccessTokenType: TypeAlias = Literal["access", "refresh", "verify"]
|
||||||
|
|
||||||
|
|
||||||
|
class TokenInfo(TypedDict):
|
||||||
|
# 用户名信息
|
||||||
|
username: str
|
||||||
|
# token过期时间
|
||||||
|
exp: datetime
|
||||||
|
type: AccessTokenType
|
||||||
|
|
||||||
|
|
||||||
class CryptUtils:
|
class CryptUtils:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -18,17 +30,19 @@ class CryptUtils:
|
|||||||
return pwd_context.verify(plain_password, hashed_password)
|
return pwd_context.verify(plain_password, hashed_password)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create_access_token(username: str, expires_delta: timedelta | None = None):
|
def create_token(
|
||||||
data: dict = {"sub": username}
|
username: str,
|
||||||
if expires_delta:
|
type: AccessTokenType,
|
||||||
expire = datetime.now(timezone.utc) + expires_delta
|
expires_delta: timedelta
|
||||||
else:
|
):
|
||||||
expire = datetime.now(timezone.utc) + timedelta(minutes=env.jwt_access_token_expire_minutes)
|
data: TokenInfo = {
|
||||||
data.update({'exp': expire})
|
"username": username,
|
||||||
|
"type": type,
|
||||||
|
"exp": datetime.now(timezone.utc) + expires_delta
|
||||||
|
}
|
||||||
return jwt.encode(data, env.jwt_secret_key, env.jwt_algorithm)
|
return jwt.encode(data, env.jwt_secret_key, env.jwt_algorithm)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_username_from_token(token: str):
|
def get_token_info(token: str) -> TokenInfo:
|
||||||
data = jwt.decode(token, env.jwt_secret_key, algorithms=[env.jwt_algorithm])
|
data = jwt.decode(token, env.jwt_secret_key, algorithms=[env.jwt_algorithm])
|
||||||
username = data.get("sub")
|
return data
|
||||||
return username
|
|
||||||
|
|||||||
Reference in New Issue
Block a user