Add token related interfaces (#495)

This commit is contained in:
Wu Clan
2025-01-23 13:12:46 +08:00
committed by GitHub
parent 7553ccf14d
commit 3b65679fb0
14 changed files with 306 additions and 116 deletions
+9
View File
@@ -40,15 +40,24 @@ class NewToken:
new_access_token_expire_time: datetime
new_refresh_token: str
new_refresh_token_expire_time: datetime
session_uuid: str
@dataclasses.dataclass
class AccessToken:
access_token: str
access_token_expire_time: datetime
session_uuid: str
@dataclasses.dataclass
class RefreshToken:
refresh_token: str
refresh_token_expire_time: datetime
@dataclasses.dataclass
class TokenPayload:
id: int
session_uuid: str
expire_time: datetime
+65 -37
View File
@@ -1,6 +1,9 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import json
from datetime import timedelta
from uuid import uuid4
from fastapi import Depends, Request
from fastapi.security import HTTPBearer
@@ -13,7 +16,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from backend.app.admin.model import User
from backend.app.admin.schema.user import CurrentUserIns
from backend.common.dataclasses import AccessToken, NewToken, RefreshToken
from backend.common.dataclasses import AccessToken, NewToken, RefreshToken, TokenPayload
from backend.common.exception.errors import AuthorizationError, TokenError
from backend.core.conf import settings
from backend.database.db import async_db_session
@@ -49,78 +52,101 @@ def password_verify(plain_password: str, hashed_password: str) -> bool:
return password_hash.verify(plain_password, hashed_password)
async def create_access_token(sub: str, multi_login: bool) -> AccessToken:
async def create_access_token(user_id: str, multi_login: bool, **kwargs) -> AccessToken:
"""
Generate encryption token
:param sub: The subject/userid of the JWT
:param multi_login: multipoint login for user
:param user_id: The user id of the JWT
:param multi_login: Multipoint login for user
:param kwargs: Token extra information
:return:
"""
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)
session_uuid = str(uuid4())
access_token = jwt.encode(
{'session_uuid': session_uuid, 'exp': expire, 'sub': user_id},
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)
await redis_client.delete_prefix(f'{settings.TOKEN_REDIS_PREFIX}:{user_id}')
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)
await redis_client.setex(
f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:{session_uuid}',
settings.TOKEN_EXPIRE_SECONDS,
access_token,
)
# Token 附加信息单独存储
if kwargs:
await redis_client.setex(
f'{settings.TOKEN_EXTRA_INFO_REDIS_PREFIX}:{session_uuid}',
settings.TOKEN_EXPIRE_SECONDS,
json.dumps(kwargs, ensure_ascii=False),
)
return AccessToken(access_token=access_token, access_token_expire_time=expire, session_uuid=session_uuid)
async def create_refresh_token(sub: str, multi_login: bool) -> RefreshToken:
async def create_refresh_token(user_id: 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 user_id: The user id of the JWT
:param multi_login: multipoint login for user
:return:
"""
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)
refresh_token = jwt.encode(
{'exp': expire, 'sub': user_id},
settings.TOKEN_SECRET_KEY,
settings.TOKEN_ALGORITHM,
)
if multi_login is False:
key_prefix = f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{sub}'
key_prefix = f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}'
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)
await redis_client.setex(
f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}:{refresh_token}',
settings.TOKEN_REFRESH_EXPIRE_SECONDS,
refresh_token,
)
return RefreshToken(refresh_token=refresh_token, refresh_token_expire_time=expire)
async def create_new_token(sub: str, token: str, refresh_token: str, multi_login: bool) -> NewToken:
async def create_new_token(user_id: str, token: str, refresh_token: str, multi_login: bool, **kwargs) -> NewToken:
"""
Generate new token
:param sub:
:param user_id:
:param token
:param refresh_token:
:param multi_login:
:param kwargs: Access token extra information
:return:
"""
redis_refresh_token = await redis_client.get(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{sub}:{refresh_token}')
redis_refresh_token = await redis_client.get(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}:{refresh_token}')
if not redis_refresh_token or redis_refresh_token != refresh_token:
raise TokenError(msg='Refresh Token 已过期')
new_access_token = await create_access_token(sub, multi_login)
new_refresh_token = await create_refresh_token(sub, multi_login)
token_payload = jwt_decode(token)
new_access_token = await create_access_token(user_id, multi_login, **kwargs)
new_refresh_token = await create_refresh_token(user_id, multi_login)
keys = [
f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:{token_payload.session_uuid}',
f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}:{refresh_token}',
]
for key in keys:
await redis_client.delete(key)
token_key = f'{settings.TOKEN_REDIS_PREFIX}:{sub}:{token}'
refresh_token_key = f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{sub}:{refresh_token}'
await redis_client.delete(token_key)
await redis_client.delete(refresh_token_key)
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,
session_uuid=new_access_token.session_uuid,
)
@@ -137,7 +163,7 @@ def get_token(request: Request) -> str:
return token
def jwt_decode(token: str) -> int:
def jwt_decode(token: str) -> TokenPayload:
"""
Decode token
@@ -146,14 +172,16 @@ def jwt_decode(token: str) -> int:
"""
try:
payload = jwt.decode(token, settings.TOKEN_SECRET_KEY, algorithms=[settings.TOKEN_ALGORITHM])
user_id = int(payload.get('sub'))
session_uuid = payload.get('session_uuid') or 'debug'
user_id = payload.get('sub')
expire_time = payload.get('exp')
if not user_id:
raise TokenError(msg='Token 无效')
except ExpiredSignatureError:
raise TokenError(msg='Token 已过期')
except (JWTError, Exception):
raise TokenError(msg='Token 无效')
return user_id
return TokenPayload(id=int(user_id), session_uuid=session_uuid, expire_time=expire_time)
async def get_current_user(db: AsyncSession, pk: int) -> User:
@@ -203,9 +231,9 @@ async def jwt_authentication(token: str) -> CurrentUserIns:
:param token:
:return:
"""
user_id = jwt_decode(token)
key = f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:{token}'
token_verify = await redis_client.get(key)
token_payload = jwt_decode(token)
user_id = token_payload.id
token_verify = await redis_client.get(f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:{token_payload.session_uuid}')
if not token_verify:
raise TokenError(msg='Token 已过期')
cache_user = await redis_client.get(f'{settings.JWT_USER_REDIS_PREFIX}:{user_id}')
+12 -5
View File
@@ -6,6 +6,7 @@ from backend.app.task.conf import task_settings
from backend.common.log import log
from backend.common.security.jwt import jwt_authentication
from backend.core.conf import settings
from backend.database.redis import redis_client
sio = socketio.AsyncServer(
# 此配置是为了集成 celery 实现消息订阅,如果你不使用 celery,可以直接删除此配置,不会造成任何影响
@@ -29,16 +30,20 @@ sio = socketio.AsyncServer(
@sio.event
async def connect(sid, environ, auth):
"""当客户端连接时触发"""
if not auth:
print('ws 连接失败:无授权')
log.error('ws 连接失败:无授权')
return False
session_uuid = auth.get('session_uuid')
token = auth.get('token')
if not token:
print('ws 连接失败:无 token 授权')
if not token or not session_uuid:
log.error('ws 连接失败:授权失败,请检查')
return False
if token == 'internal':
# 免授权直连
if token == settings.WS_NO_AUTH_MARKER:
await redis_client.sadd(settings.TOKEN_ONLINE_REDIS_PREFIX, session_uuid)
return True
try:
@@ -47,9 +52,11 @@ async def connect(sid, environ, auth):
log.info(f'ws 连接失败:{e}')
return False
await redis_client.sadd(settings.TOKEN_ONLINE_REDIS_PREFIX, session_uuid)
return True
@sio.event
async def disconnect(sid):
pass
"""当客户端断开连接时触发"""
await redis_client.spop(settings.TOKEN_ONLINE_REDIS_PREFIX)