feat: 使用redis来缓存用户信息
This commit is contained in:
@@ -10,9 +10,11 @@ from starlette.responses import JSONResponse
|
||||
from app.config.env import env
|
||||
from app.controller.add_user_route import unauthorized_exception, get_current_user
|
||||
from app.model.ApiSecretModel import ApiSecretService
|
||||
from app.model.UserModel import PublicUser
|
||||
from app.utils.CrpyUtils import TokenInfo, CryptUtils
|
||||
from app.utils.api_secret_utils import api_secret_utils, ApiSecretStatus
|
||||
from app.utils.db_utils import async_session
|
||||
from app.utils.redis_utils import get_redis_cache
|
||||
|
||||
|
||||
def add_app_middlewares(app: FastAPI):
|
||||
@@ -63,14 +65,23 @@ def add_app_middlewares(app: FastAPI):
|
||||
except InvalidTokenError:
|
||||
raise unauthorized_exception
|
||||
|
||||
# 根据token信息获取用户信息
|
||||
async with async_session() as session:
|
||||
try:
|
||||
public_user = await get_current_user(session, token)
|
||||
request.state.user = public_user
|
||||
request.state.token = token
|
||||
except InvalidTokenError:
|
||||
raise unauthorized_exception
|
||||
async def get_default_cache_value():
|
||||
async with async_session() as session:
|
||||
try:
|
||||
public_user = await get_current_user(session, token)
|
||||
return public_user.model_dump(mode="json")
|
||||
except InvalidTokenError:
|
||||
return None
|
||||
|
||||
user_dict = await get_redis_cache(
|
||||
f"access_token_username_{token_info.get('username')}",
|
||||
get_default_cache_value
|
||||
)
|
||||
if not user_dict:
|
||||
raise unauthorized_exception
|
||||
|
||||
request.state.user = PublicUser(**user_dict)
|
||||
request.state.token = token
|
||||
|
||||
# 如果请求的是api接口
|
||||
# api开头的接口需要额外验证秘钥
|
||||
|
||||
+19
-5
@@ -8,6 +8,7 @@ from sqlmodel import Field, Relationship, select
|
||||
from app.model.BasicModel import BasicModel
|
||||
from app.model.PosModel import PosModel
|
||||
from app.utils.create_module_service import create_model_service
|
||||
from app.utils.redis_utils import remove_redis_cache
|
||||
|
||||
|
||||
class UserValidate(str, Enum):
|
||||
@@ -93,11 +94,24 @@ class UserServiceModel(PublicUser, table=True):
|
||||
pass
|
||||
|
||||
|
||||
async def before_update(row_dict, session):
|
||||
user_cls: UserServiceModel = await UserService.query_item(session, {"id": row_dict["id"]})
|
||||
await remove_redis_cache(f"access_token_username_{user_cls.username}")
|
||||
|
||||
|
||||
async def before_delete(row_dict, session):
|
||||
user_cls: UserServiceModel = await UserService.query_item(session, {"id": row_dict["id"]})
|
||||
await remove_redis_cache(f"access_token_username_{user_cls.username}")
|
||||
|
||||
|
||||
UserService = create_model_service(
|
||||
Cls=UserServiceModel,
|
||||
custom_query=(lambda: select(UserServiceModel)
|
||||
.options(
|
||||
selectinload(UserServiceModel.position).
|
||||
selectinload(PosModel.organization)
|
||||
))
|
||||
custom_query=(
|
||||
lambda: select(UserServiceModel)
|
||||
.options(
|
||||
selectinload(UserServiceModel.position).
|
||||
selectinload(PosModel.organization)
|
||||
)),
|
||||
before_update=before_update,
|
||||
before_delete=before_delete
|
||||
)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Annotated
|
||||
from typing import Annotated, AsyncContextManager
|
||||
|
||||
import redis.asyncio as redis
|
||||
from fastapi import Depends
|
||||
@@ -38,7 +38,7 @@ class RedisUtils():
|
||||
|
||||
# 获取一个Redis连接客户端
|
||||
@asynccontextmanager
|
||||
async def get_redis_connection(self):
|
||||
async def get_redis_connection(self) -> AsyncContextManager[redis.Redis]:
|
||||
if not self.redis_pool:
|
||||
raise Exception("Redis 连接池未初始化")
|
||||
redis_client = redis.Redis(connection_pool=self.redis_pool)
|
||||
@@ -63,3 +63,26 @@ async def get_redis_client() -> redis.Redis:
|
||||
|
||||
|
||||
RedisClientDep = Annotated[redis.Redis, Depends(get_redis_client)]
|
||||
|
||||
|
||||
# 从redis中获取key的缓存,如果没有值则执行默认值获取函数,并保存到redis中
|
||||
async def get_redis_cache(key: str, default_value_getter):
|
||||
async with redis_utils.get_redis_connection() as redis_client:
|
||||
mapping = await redis_client.hgetall(key)
|
||||
exists = bool(mapping)
|
||||
if exists:
|
||||
return mapping
|
||||
else:
|
||||
value = await default_value_getter()
|
||||
print("value", value)
|
||||
if value is not None:
|
||||
new_value = {k: v for k, v in value.items() if v is not None}
|
||||
await redis_client.hset(key, mapping=new_value)
|
||||
return value
|
||||
|
||||
|
||||
# 删除redis缓存
|
||||
async def remove_redis_cache(key: str):
|
||||
async with redis_utils.get_redis_connection() as redis_client:
|
||||
result = await redis_client.delete(key)
|
||||
return result > 0 # 返回是否删除成功
|
||||
|
||||
Reference in New Issue
Block a user