From 4311f521a74a7cac36a5ee3b29523724bd8de824 Mon Sep 17 00:00:00 2001 From: martsforever Date: Sat, 27 Sep 2025 21:06:15 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E4=BD=BF=E7=94=A8redis=E6=9D=A5?= =?UTF-8?q?=E7=BC=93=E5=AD=98=E7=94=A8=E6=88=B7=E4=BF=A1=E6=81=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/middlewares/app_middlewares.py | 27 +++++++++++++++++++-------- app/model/UserModel.py | 24 +++++++++++++++++++----- app/utils/redis_utils.py | 27 +++++++++++++++++++++++++-- 3 files changed, 63 insertions(+), 15 deletions(-) diff --git a/app/middlewares/app_middlewares.py b/app/middlewares/app_middlewares.py index c9fc467..a0624fe 100644 --- a/app/middlewares/app_middlewares.py +++ b/app/middlewares/app_middlewares.py @@ -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开头的接口需要额外验证秘钥 diff --git a/app/model/UserModel.py b/app/model/UserModel.py index 811fafb..7748550 100644 --- a/app/model/UserModel.py +++ b/app/model/UserModel.py @@ -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 ) diff --git a/app/utils/redis_utils.py b/app/utils/redis_utils.py index a795189..282b308 100644 --- a/app/utils/redis_utils.py +++ b/app/utils/redis_utils.py @@ -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 # 返回是否删除成功